From 6bee02a346d0aa8dc4d5dcb312545fa37408b6ca Mon Sep 17 00:00:00 2001 From: Ahmed Ibrahim Date: Tue, 3 Mar 2026 14:07:51 -0800 Subject: [PATCH 1/5] Build delegated realtime handoff text from all messages (#13395) ## Summary - Route delegated realtime handoff turns from all handoff message texts, preserving order - Fallback to input_transcript only when no messages are present - Add regression coverage for multi-message handoff requests --- codex-rs/core/src/realtime_conversation.rs | 44 +++++++++-- .../core/tests/suite/realtime_conversation.rs | 74 +++++++++++++++++++ 2 files changed, 110 insertions(+), 8 deletions(-) diff --git a/codex-rs/core/src/realtime_conversation.rs b/codex-rs/core/src/realtime_conversation.rs index 8842c41c96..e614be1061 100644 --- a/codex-rs/core/src/realtime_conversation.rs +++ b/codex-rs/core/src/realtime_conversation.rs @@ -373,7 +373,15 @@ pub(crate) async fn handle_audio( } fn realtime_text_from_handoff_request(handoff: &RealtimeHandoffRequested) -> Option { - (!handoff.input_transcript.is_empty()).then(|| handoff.input_transcript.clone()) + let messages = handoff + .messages + .iter() + .map(|message| message.text.as_str()) + .collect::>() + .join("\n"); + (!messages.is_empty()).then_some(messages).or_else(|| { + (!handoff.input_transcript.is_empty()).then(|| handoff.input_transcript.clone()) + }) } fn realtime_api_key( @@ -579,19 +587,39 @@ mod tests { use pretty_assertions::assert_eq; #[test] - fn extracts_text_from_handoff_request_input_transcript() { + fn extracts_text_from_handoff_request_messages() { let handoff = RealtimeHandoffRequested { handoff_id: "handoff_1".to_string(), item_id: "item_1".to_string(), - input_transcript: "hello".to_string(), - messages: vec![RealtimeHandoffMessage { - role: "user".to_string(), - text: "hello".to_string(), - }], + input_transcript: "ignored".to_string(), + messages: vec![ + RealtimeHandoffMessage { + role: "user".to_string(), + text: "hello".to_string(), + }, + RealtimeHandoffMessage { + role: "assistant".to_string(), + text: "hi there".to_string(), + }, + ], }; assert_eq!( realtime_text_from_handoff_request(&handoff), - Some("hello".to_string()) + Some("hello\nhi there".to_string()) + ); + } + + #[test] + fn extracts_text_from_handoff_request_input_transcript_if_messages_missing() { + let handoff = RealtimeHandoffRequested { + handoff_id: "handoff_1".to_string(), + item_id: "item_1".to_string(), + input_transcript: "ignored".to_string(), + messages: vec![], + }; + assert_eq!( + realtime_text_from_handoff_request(&handoff), + Some("ignored".to_string()) ); } diff --git a/codex-rs/core/tests/suite/realtime_conversation.rs b/codex-rs/core/tests/suite/realtime_conversation.rs index 32500299cb..ff4e745dd2 100644 --- a/codex-rs/core/tests/suite/realtime_conversation.rs +++ b/codex-rs/core/tests/suite/realtime_conversation.rs @@ -918,6 +918,80 @@ async fn inbound_handoff_request_starts_turn() -> Result<()> { Ok(()) } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn inbound_handoff_request_uses_all_messages() -> Result<()> { + skip_if_no_network!(Ok(())); + + let api_server = start_mock_server().await; + let response_mock = responses::mount_sse_once( + &api_server, + responses::sse(vec![ + responses::ev_response_created("resp-1"), + responses::ev_assistant_message("msg-1", "ok"), + responses::ev_completed("resp-1"), + ]), + ) + .await; + + let realtime_server = start_websocket_server(vec![vec![vec![ + json!({ + "type": "session.updated", + "session": { "id": "sess_inbound_multi", "instructions": "backend prompt" } + }), + json!({ + "type": "conversation.handoff.requested", + "handoff_id": "handoff_inbound_multi", + "item_id": "item_inbound_multi", + "input_transcript": "ignored", + "messages": [ + { "role": "assistant", "text": "assistant context" }, + { "role": "user", "text": "delegated query" }, + { "role": "assistant", "text": "assist confirm" }, + ] + }), + ]]]) + .await; + + let mut builder = test_codex().with_config({ + let realtime_base_url = realtime_server.uri().to_string(); + move |config| { + config.experimental_realtime_ws_base_url = Some(realtime_base_url); + } + }); + let test = builder.build(&api_server).await?; + + test.codex + .submit(Op::RealtimeConversationStart(ConversationStartParams { + prompt: "backend prompt".to_string(), + session_id: None, + })) + .await?; + + let _ = wait_for_event_match(&test.codex, |msg| match msg { + EventMsg::RealtimeConversationRealtime(RealtimeConversationRealtimeEvent { + payload: RealtimeEvent::SessionUpdated { session_id, .. }, + }) => Some(session_id.clone()), + _ => None, + }) + .await; + + wait_for_event(&test.codex, |event| { + matches!(event, EventMsg::TurnComplete(_)) + }) + .await; + + let request = response_mock.single_request(); + let user_texts = request.message_input_texts("user"); + assert!( + user_texts + .iter() + .any(|text| text == "assistant context\ndelegated query\nassist confirm") + ); + + realtime_server.shutdown().await; + Ok(()) +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn inbound_conversation_item_does_not_start_turn_and_still_forwards_audio() -> Result<()> { skip_if_no_network!(Ok(())); From bab32afa93292b181e350c2c0416f77fd82adf2b Mon Sep 17 00:00:00 2001 From: Eric Traut Date: Tue, 3 Mar 2026 15:32:47 -0700 Subject: [PATCH 2/5] Require deduplicator success before commenting (#13399) Fixed recent regression in issue dedup action --- .github/workflows/issue-deduplicator.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/issue-deduplicator.yml b/.github/workflows/issue-deduplicator.yml index 62af691fe6..2c7c13e688 100644 --- a/.github/workflows/issue-deduplicator.yml +++ b/.github/workflows/issue-deduplicator.yml @@ -335,7 +335,7 @@ jobs: comment-on-issue: name: Comment with potential duplicates needs: select-final - if: ${{ needs.select-final.result != 'skipped' }} + if: ${{ always() && needs.select-final.result == 'success' }} runs-on: ubuntu-latest permissions: contents: read From 041c896509796c438e34406e43473e60b8282cfc Mon Sep 17 00:00:00 2001 From: Ahmed Ibrahim Date: Tue, 3 Mar 2026 14:41:26 -0800 Subject: [PATCH 3/5] Revert "Revert "realtime prompt changes"" (#13398) Reverts openai/codex#13385 --- ...e__remote_compact_resume_restates_realtime_end_shapes.snap | 2 +- ..._turn_compaction_does_not_restate_realtime_end_shapes.snap | 4 ++-- ...mote_pre_turn_compaction_restates_realtime_end_shapes.snap | 2 +- codex-rs/protocol/src/prompts/realtime/realtime_end.md | 4 +--- codex-rs/protocol/src/prompts/realtime/realtime_start.md | 3 +-- 5 files changed, 6 insertions(+), 9 deletions(-) diff --git a/codex-rs/core/tests/suite/snapshots/all__suite__compact_remote__remote_compact_resume_restates_realtime_end_shapes.snap b/codex-rs/core/tests/suite/snapshots/all__suite__compact_remote__remote_compact_resume_restates_realtime_end_shapes.snap index 1b95615375..fc12d431e2 100644 --- a/codex-rs/core/tests/suite/snapshots/all__suite__compact_remote__remote_compact_resume_restates_realtime_end_shapes.snap +++ b/codex-rs/core/tests/suite/snapshots/all__suite__compact_remote__remote_compact_resume_restates_realtime_end_shapes.snap @@ -17,7 +17,7 @@ Scenario: After remote manual /compact and resume, the first resumed turn rebuil 00:compaction:encrypted=true 01:message/developer[2]: [01] - [02] \nRealtime conversation ended.\n\nYou are... + [02] \nRealtime conversation ended.\n\nSubsequ... 02:message/user[2]: [01] [02] > diff --git a/codex-rs/core/tests/suite/snapshots/all__suite__compact_remote__remote_mid_turn_compaction_does_not_restate_realtime_end_shapes.snap b/codex-rs/core/tests/suite/snapshots/all__suite__compact_remote__remote_mid_turn_compaction_does_not_restate_realtime_end_shapes.snap index eb7e583874..b1f83ce4d3 100644 --- a/codex-rs/core/tests/suite/snapshots/all__suite__compact_remote__remote_mid_turn_compaction_does_not_restate_realtime_end_shapes.snap +++ b/codex-rs/core/tests/suite/snapshots/all__suite__compact_remote__remote_mid_turn_compaction_does_not_restate_realtime_end_shapes.snap @@ -12,7 +12,7 @@ Scenario: Remote mid-turn continuation compaction after realtime was closed befo 02:message/developer:\nRealtime conversation started.\n\nYou a... 03:message/user:SETUP_USER 04:message/assistant:REMOTE_SETUP_REPLY -05:message/developer:\nRealtime conversation ended.\n\nYou are... +05:message/developer:\nRealtime conversation ended.\n\nSubsequ... 06:message/user:USER_TWO ## Remote Compaction Request @@ -23,7 +23,7 @@ Scenario: Remote mid-turn continuation compaction after realtime was closed befo 02:message/developer:\nRealtime conversation started.\n\nYou a... 03:message/user:SETUP_USER 04:message/assistant:REMOTE_SETUP_REPLY -05:message/developer:\nRealtime conversation ended.\n\nYou are... +05:message/developer:\nRealtime conversation ended.\n\nSubsequ... 06:message/user:USER_TWO 07:function_call/test_tool 08:function_call_output:unsupported call: test_tool diff --git a/codex-rs/core/tests/suite/snapshots/all__suite__compact_remote__remote_pre_turn_compaction_restates_realtime_end_shapes.snap b/codex-rs/core/tests/suite/snapshots/all__suite__compact_remote__remote_pre_turn_compaction_restates_realtime_end_shapes.snap index 7f14509dd2..57af327d16 100644 --- a/codex-rs/core/tests/suite/snapshots/all__suite__compact_remote__remote_pre_turn_compaction_restates_realtime_end_shapes.snap +++ b/codex-rs/core/tests/suite/snapshots/all__suite__compact_remote__remote_pre_turn_compaction_restates_realtime_end_shapes.snap @@ -17,7 +17,7 @@ Scenario: Remote pre-turn auto-compaction after realtime was closed between turn 00:compaction:encrypted=true 01:message/developer[2]: [01] - [02] \nRealtime conversation ended.\n\nYou are... + [02] \nRealtime conversation ended.\n\nSubsequ... 02:message/user[2]: [01] [02] > diff --git a/codex-rs/protocol/src/prompts/realtime/realtime_end.md b/codex-rs/protocol/src/prompts/realtime/realtime_end.md index 07ccae5286..abec16da03 100644 --- a/codex-rs/protocol/src/prompts/realtime/realtime_end.md +++ b/codex-rs/protocol/src/prompts/realtime/realtime_end.md @@ -1,5 +1,3 @@ Realtime conversation ended. -You are still operating behind an intermediary rather than speaking to the user directly. Use the conversation transcript and current context to decide whether backend help is actually needed, and avoid verbose responses that only add latency. - -Subsequent user input may return to typed text rather than transcript-style text. Do not assume recognition errors or missing punctuation once realtime has ended. Resume normal chat behavior. +Subsequent user input will return to typed text rather than transcript-style text. Do not assume recognition errors or missing punctuation once realtime has ended. Resume normal chat behavior. diff --git a/codex-rs/protocol/src/prompts/realtime/realtime_start.md b/codex-rs/protocol/src/prompts/realtime/realtime_start.md index c7d363eff7..3159a2b07f 100644 --- a/codex-rs/protocol/src/prompts/realtime/realtime_start.md +++ b/codex-rs/protocol/src/prompts/realtime/realtime_start.md @@ -6,5 +6,4 @@ When invoked, you receive the latest conversation transcript and any relevant mo When user text is routed from realtime, treat it as a transcript. It may be unpunctuated or contain recognition errors. -- Ask brief clarification questions when needed. -- Keep responses concise and action-oriented. +- Keep responses concise and action-oriented. Your updates should help the intermediary respond to the user. From 9b004e2db126995307bc32735ad3ee20e6fe5ced Mon Sep 17 00:00:00 2001 From: xl-openai Date: Tue, 3 Mar 2026 15:00:18 -0800 Subject: [PATCH 4/5] Refactor plugin config and cache path (#13333) Update config.toml plugin entries to use @ as the key. Plugin now stays in [plugins/cache/marketplace-name/plugin-name/$version/] Clean up the plugin code structure. Add plugin install functionality (not used yet). --- codex-rs/core/config.schema.json | 6 - codex-rs/core/src/config/types.rs | 1 - codex-rs/core/src/mcp/mod.rs | 15 +- .../src/{plugins.rs => plugins/manager.rs} | 253 ++++++++---- codex-rs/core/src/plugins/manifest.rs | 37 ++ codex-rs/core/src/plugins/mod.rs | 14 + codex-rs/core/src/plugins/store.rs | 374 ++++++++++++++++++ codex-rs/core/tests/suite/plugins.rs | 14 +- 8 files changed, 622 insertions(+), 92 deletions(-) rename codex-rs/core/src/{plugins.rs => plugins/manager.rs} (72%) create mode 100644 codex-rs/core/src/plugins/manifest.rs create mode 100644 codex-rs/core/src/plugins/mod.rs create mode 100644 codex-rs/core/src/plugins/store.rs diff --git a/codex-rs/core/config.schema.json b/codex-rs/core/config.schema.json index 3cd7000913..74af6d50c6 100644 --- a/codex-rs/core/config.schema.json +++ b/codex-rs/core/config.schema.json @@ -1107,14 +1107,8 @@ "enabled": { "default": true, "type": "boolean" - }, - "path": { - "$ref": "#/definitions/AbsolutePathBuf" } }, - "required": [ - "path" - ], "type": "object" }, "ProjectConfig": { diff --git a/codex-rs/core/src/config/types.rs b/codex-rs/core/src/config/types.rs index 8fd5b109d0..6f8c349158 100644 --- a/codex-rs/core/src/config/types.rs +++ b/codex-rs/core/src/config/types.rs @@ -778,7 +778,6 @@ pub struct SkillConfig { #[derive(Serialize, Deserialize, Debug, Clone, PartialEq, Eq, JsonSchema)] #[schemars(deny_unknown_fields)] pub struct PluginConfig { - pub path: AbsolutePathBuf, #[serde(default = "default_enabled")] pub enabled: bool, } diff --git a/codex-rs/core/src/mcp/mod.rs b/codex-rs/core/src/mcp/mod.rs index 4edb45c462..91a7f74db9 100644 --- a/codex-rs/core/src/mcp/mod.rs +++ b/codex-rs/core/src/mcp/mod.rs @@ -417,7 +417,7 @@ mod tests { fs::write(path, contents).unwrap(); } - fn plugin_config_toml(plugin_root: &Path) -> String { + fn plugin_config_toml() -> String { let mut root = toml::map::Map::new(); let mut features = toml::map::Map::new(); @@ -425,14 +425,10 @@ mod tests { root.insert("features".to_string(), Value::Table(features)); let mut plugin = toml::map::Map::new(); - plugin.insert( - "path".to_string(), - Value::String(plugin_root.display().to_string()), - ); plugin.insert("enabled".to_string(), Value::Boolean(true)); let mut plugins = toml::map::Map::new(); - plugins.insert("sample".to_string(), Value::Table(plugin)); + plugins.insert("sample@test".to_string(), Value::Table(plugin)); root.insert("plugins".to_string(), Value::Table(plugins)); toml::to_string(&Value::Table(root)).expect("plugin test config should serialize") @@ -616,7 +612,10 @@ mod tests { #[tokio::test] async fn effective_mcp_servers_include_plugins_without_overriding_user_config() { let codex_home = tempfile::tempdir().expect("tempdir"); - let plugin_root = codex_home.path().join("plugin-sample"); + let plugin_root = codex_home + .path() + .join("plugins/cache") + .join("test/sample/local"); write_file( &plugin_root.join(".codex-plugin/plugin.json"), r#"{"name":"sample"}"#, @@ -638,7 +637,7 @@ mod tests { ); write_file( &codex_home.path().join(CONFIG_TOML_FILE), - &plugin_config_toml(&plugin_root), + &plugin_config_toml(), ); let mut config = ConfigBuilder::default() diff --git a/codex-rs/core/src/plugins.rs b/codex-rs/core/src/plugins/manager.rs similarity index 72% rename from codex-rs/core/src/plugins.rs rename to codex-rs/core/src/plugins/manager.rs index 890d15a323..d98c4379ff 100644 --- a/codex-rs/core/src/plugins.rs +++ b/codex-rs/core/src/plugins/manager.rs @@ -1,4 +1,14 @@ +use super::load_plugin_manifest; +use super::plugin_manifest_name; +use super::store::DEFAULT_PLUGIN_VERSION; +use super::store::PluginId; +use super::store::PluginInstallRequest; +use super::store::PluginInstallResult; +use super::store::PluginStore; +use super::store::PluginStoreError; use crate::config::Config; +use crate::config::ConfigService; +use crate::config::ConfigServiceError; use crate::config::ConfigToml; use crate::config::profile::ConfigProfile; use crate::config::types::McpServerConfig; @@ -7,10 +17,13 @@ use crate::config_loader::ConfigLayerStack; use crate::features::Feature; use crate::features::FeatureOverrides; use crate::features::Features; +use codex_app_server_protocol::ConfigValueWriteParams; +use codex_app_server_protocol::MergeStrategy; use codex_utils_absolute_path::AbsolutePathBuf; use serde::Deserialize; use serde_json::Map as JsonMap; use serde_json::Value as JsonValue; +use serde_json::json; use std::collections::HashMap; use std::fs; use std::path::Path; @@ -18,7 +31,6 @@ use std::path::PathBuf; use std::sync::RwLock; use tracing::warn; -const PLUGIN_MANIFEST_PATH: &str = ".codex-plugin/plugin.json"; const DEFAULT_SKILLS_DIR_NAME: &str = "skills"; const DEFAULT_MCP_CONFIG_FILE: &str = ".mcp.json"; @@ -71,12 +83,16 @@ impl PluginLoadOutcome { } pub struct PluginsManager { + codex_home: PathBuf, + store: PluginStore, cache_by_cwd: RwLock>, } impl PluginsManager { - pub fn new(_codex_home: PathBuf) -> Self { + pub fn new(codex_home: PathBuf) -> Self { Self { + codex_home: codex_home.clone(), + store: PluginStore::new(codex_home), cache_by_cwd: RwLock::new(HashMap::new()), } } @@ -104,7 +120,7 @@ impl PluginsManager { return outcome; } - let outcome = load_plugins_from_layer_stack(config_layer_stack); + let outcome = load_plugins_from_layer_stack(config_layer_stack, &self.store); log_plugin_load_errors(&outcome); let mut cache = match self.cache_by_cwd.write() { Ok(cache) => cache, @@ -128,6 +144,50 @@ impl PluginsManager { Err(err) => err.into_inner().get(cwd).cloned(), } } + + pub async fn install_plugin( + &self, + request: PluginInstallRequest, + ) -> Result { + let store = self.store.clone(); + let result = tokio::task::spawn_blocking(move || store.install(request)) + .await + .map_err(PluginInstallError::join)??; + + ConfigService::new_with_defaults(self.codex_home.clone()) + .write_value(ConfigValueWriteParams { + key_path: format!("plugins.{}", result.plugin_id.as_key()), + value: json!({ + "enabled": true, + }), + merge_strategy: MergeStrategy::Replace, + file_path: None, + expected_version: None, + }) + .await + .map(|_| ()) + .map_err(PluginInstallError::from)?; + + Ok(result) + } +} + +#[derive(Debug, thiserror::Error)] +pub enum PluginInstallError { + #[error("{0}")] + Store(#[from] PluginStoreError), + + #[error("{0}")] + Config(#[from] ConfigServiceError), + + #[error("failed to join plugin install task: {0}")] + Join(#[from] tokio::task::JoinError), +} + +impl PluginInstallError { + fn join(source: tokio::task::JoinError) -> Self { + Self::Join(source) + } } fn plugins_feature_enabled_from_stack(config_layer_stack: &ConfigLayerStack) -> bool { @@ -160,11 +220,6 @@ fn log_plugin_load_errors(outcome: &PluginLoadOutcome) { } } -#[derive(Debug, Default, Deserialize)] -struct PluginManifest { - name: String, -} - #[derive(Debug, Default, Deserialize)] #[serde(rename_all = "camelCase")] struct PluginMcpFile { @@ -172,7 +227,10 @@ struct PluginMcpFile { mcp_servers: HashMap, } -pub fn load_plugins_from_layer_stack(config_layer_stack: &ConfigLayerStack) -> PluginLoadOutcome { +pub(crate) fn load_plugins_from_layer_stack( + config_layer_stack: &ConfigLayerStack, + store: &PluginStore, +) -> PluginLoadOutcome { let mut configured_plugins: Vec<_> = configured_plugins_from_stack(config_layer_stack) .into_iter() .collect(); @@ -181,7 +239,7 @@ pub fn load_plugins_from_layer_stack(config_layer_stack: &ConfigLayerStack) -> P let mut plugins = Vec::with_capacity(configured_plugins.len()); let mut seen_mcp_server_names = HashMap::::new(); for (configured_name, plugin) in configured_plugins { - let loaded_plugin = load_plugin(configured_name.clone(), &plugin); + let loaded_plugin = load_plugin(configured_name.clone(), &plugin, store); for name in loaded_plugin.mcp_servers.keys() { if let Some(previous_plugin) = seen_mcp_server_names.insert(name.clone(), configured_name.clone()) @@ -226,12 +284,18 @@ fn configured_plugins_from_stack( } } -fn load_plugin(config_name: String, plugin: &PluginConfig) -> LoadedPlugin { - let plugin_root = plugin.path.clone(); +fn load_plugin(config_name: String, plugin: &PluginConfig, store: &PluginStore) -> LoadedPlugin { + let plugin_version = DEFAULT_PLUGIN_VERSION.to_string(); + let plugin_root = PluginId::parse(&config_name) + .map(|plugin_id| store.plugin_root(&plugin_id, &plugin_version)); + let root = match &plugin_root { + Ok(plugin_root) => plugin_root.clone(), + Err(_) => store.root().clone(), + }; let mut loaded_plugin = LoadedPlugin { config_name, manifest_name: None, - root: plugin_root.clone(), + root, enabled: plugin.enabled, skill_roots: Vec::new(), mcp_servers: HashMap::new(), @@ -242,6 +306,14 @@ fn load_plugin(config_name: String, plugin: &PluginConfig) -> LoadedPlugin { return loaded_plugin; } + let plugin_root = match plugin_root { + Ok(plugin_root) => plugin_root, + Err(err) => { + loaded_plugin.error = Some(err.to_string()); + return loaded_plugin; + } + }; + if !plugin_root.as_path().is_dir() { loaded_plugin.error = Some("path does not exist or is not a directory".to_string()); return loaded_plugin; @@ -272,33 +344,6 @@ fn load_plugin(config_name: String, plugin: &PluginConfig) -> LoadedPlugin { loaded_plugin } -fn load_plugin_manifest(plugin_root: &Path) -> Option { - let manifest_path = plugin_root.join(PLUGIN_MANIFEST_PATH); - if !manifest_path.is_file() { - return None; - } - let contents = fs::read_to_string(&manifest_path).ok()?; - match serde_json::from_str(&contents) { - Ok(manifest) => Some(manifest), - Err(err) => { - warn!( - path = %manifest_path.display(), - "failed to parse plugin manifest: {err}" - ); - None - } - } -} - -fn plugin_manifest_name(manifest: &PluginManifest, plugin_root: &Path) -> String { - plugin_root - .file_name() - .and_then(|name| name.to_str()) - .filter(|_| manifest.name.trim().is_empty()) - .unwrap_or(&manifest.name) - .to_string() -} - fn default_skill_roots(plugin_root: &Path) -> Vec { let skills_dir = plugin_root.join(DEFAULT_SKILLS_DIR_NAME); if skills_dir.is_dir() { @@ -422,6 +467,7 @@ mod tests { use crate::config::ConfigBuilder; use crate::config::types::McpServerTransportConfig; use pretty_assertions::assert_eq; + use std::fs; use tempfile::TempDir; use toml::Value; @@ -430,11 +476,20 @@ mod tests { fs::write(path, contents).unwrap(); } - fn plugin_config_toml( - plugin_root: &Path, - enabled: bool, - plugins_feature_enabled: bool, - ) -> String { + fn write_plugin(root: &Path, dir_name: &str, manifest_name: &str) { + let plugin_root = root.join(dir_name); + fs::create_dir_all(plugin_root.join(".codex-plugin")).unwrap(); + fs::create_dir_all(plugin_root.join("skills")).unwrap(); + fs::write( + plugin_root.join(".codex-plugin/plugin.json"), + format!(r#"{{"name":"{manifest_name}"}}"#), + ) + .unwrap(); + fs::write(plugin_root.join("skills/SKILL.md"), "skill").unwrap(); + fs::write(plugin_root.join(".mcp.json"), r#"{"mcpServers":{}}"#).unwrap(); + } + + fn plugin_config_toml(enabled: bool, plugins_feature_enabled: bool) -> String { let mut root = toml::map::Map::new(); let mut features = toml::map::Map::new(); @@ -445,14 +500,10 @@ mod tests { root.insert("features".to_string(), Value::Table(features)); let mut plugin = toml::map::Map::new(); - plugin.insert( - "path".to_string(), - Value::String(plugin_root.display().to_string()), - ); plugin.insert("enabled".to_string(), Value::Boolean(enabled)); let mut plugins = toml::map::Map::new(); - plugins.insert("sample".to_string(), Value::Table(plugin)); + plugins.insert("sample@test".to_string(), Value::Table(plugin)); root.insert("plugins".to_string(), Value::Table(plugins)); toml::to_string(&Value::Table(root)).expect("plugin test config should serialize") @@ -471,7 +522,10 @@ mod tests { #[tokio::test] async fn load_plugins_loads_default_skills_and_mcp_servers() { let codex_home = TempDir::new().unwrap(); - let plugin_root = codex_home.path().join("plugin-sample"); + let plugin_root = codex_home + .path() + .join("plugins/cache") + .join("test/sample/local"); write_file( &plugin_root.join(".codex-plugin/plugin.json"), @@ -497,16 +551,13 @@ mod tests { }"#, ); - let outcome = load_plugins_from_config( - &plugin_config_toml(&plugin_root, true, true), - codex_home.path(), - ) - .await; + let outcome = + load_plugins_from_config(&plugin_config_toml(true, true), codex_home.path()).await; assert_eq!( outcome.plugins, vec![LoadedPlugin { - config_name: "sample".to_string(), + config_name: "sample@test".to_string(), manifest_name: Some("sample".to_string()), root: AbsolutePathBuf::try_from(plugin_root.clone()).unwrap(), enabled: true, @@ -544,7 +595,10 @@ mod tests { #[tokio::test] async fn load_plugins_preserves_disabled_plugins_without_effective_contributions() { let codex_home = TempDir::new().unwrap(); - let plugin_root = codex_home.path().join("plugin-sample"); + let plugin_root = codex_home + .path() + .join("plugins/cache") + .join("test/sample/local"); write_file( &plugin_root.join(".codex-plugin/plugin.json"), @@ -562,16 +616,13 @@ mod tests { }"#, ); - let outcome = load_plugins_from_config( - &plugin_config_toml(&plugin_root, false, true), - codex_home.path(), - ) - .await; + let outcome = + load_plugins_from_config(&plugin_config_toml(false, true), codex_home.path()).await; assert_eq!( outcome.plugins, vec![LoadedPlugin { - config_name: "sample".to_string(), + config_name: "sample@test".to_string(), manifest_name: None, root: AbsolutePathBuf::try_from(plugin_root).unwrap(), enabled: false, @@ -605,7 +656,10 @@ mod tests { #[tokio::test] async fn load_plugins_returns_empty_when_feature_disabled() { let codex_home = TempDir::new().unwrap(); - let plugin_root = codex_home.path().join("plugin-sample"); + let plugin_root = codex_home + .path() + .join("plugins/cache") + .join("test/sample/local"); write_file( &plugin_root.join(".codex-plugin/plugin.json"), @@ -616,12 +670,77 @@ mod tests { "---\nname: sample-search\ndescription: search sample data\n---\n", ); + let outcome = + load_plugins_from_config(&plugin_config_toml(true, false), codex_home.path()).await; + + assert_eq!(outcome, PluginLoadOutcome::default()); + } + + #[tokio::test] + async fn load_plugins_rejects_invalid_plugin_keys() { + let codex_home = TempDir::new().unwrap(); + let plugin_root = codex_home + .path() + .join("plugins/cache") + .join("test/sample/local"); + + write_file( + &plugin_root.join(".codex-plugin/plugin.json"), + r#"{"name":"sample"}"#, + ); + + let mut root = toml::map::Map::new(); + let mut features = toml::map::Map::new(); + features.insert("plugins".to_string(), Value::Boolean(true)); + root.insert("features".to_string(), Value::Table(features)); + + let mut plugin = toml::map::Map::new(); + plugin.insert("enabled".to_string(), Value::Boolean(true)); + + let mut plugins = toml::map::Map::new(); + plugins.insert("sample".to_string(), Value::Table(plugin)); + root.insert("plugins".to_string(), Value::Table(plugins)); + let outcome = load_plugins_from_config( - &plugin_config_toml(&plugin_root, true, false), + &toml::to_string(&Value::Table(root)).expect("plugin test config should serialize"), codex_home.path(), ) .await; - assert_eq!(outcome, PluginLoadOutcome::default()); + assert_eq!(outcome.plugins.len(), 1); + assert_eq!( + outcome.plugins[0].error.as_deref(), + Some("invalid plugin key `sample`; expected @") + ); + assert!(outcome.effective_skill_roots().is_empty()); + assert!(outcome.effective_mcp_servers().is_empty()); + } + + #[tokio::test] + async fn install_plugin_updates_config_with_relative_path_and_plugin_key() { + let tmp = tempfile::tempdir().unwrap(); + write_plugin(tmp.path(), "sample-plugin", "sample-plugin"); + + let result = PluginsManager::new(tmp.path().to_path_buf()) + .install_plugin(PluginInstallRequest { + source_path: tmp.path().join("sample-plugin"), + marketplace_name: None, + }) + .await + .unwrap(); + + let installed_path = tmp.path().join("plugins/cache/debug/sample-plugin/local"); + assert_eq!( + result, + PluginInstallResult { + plugin_id: PluginId::new("sample-plugin".to_string(), "debug".to_string()).unwrap(), + plugin_version: "local".to_string(), + installed_path, + } + ); + + let config = fs::read_to_string(tmp.path().join("config.toml")).unwrap(); + assert!(config.contains(r#"[plugins."sample-plugin@debug"]"#)); + assert!(config.contains("enabled = true")); } } diff --git a/codex-rs/core/src/plugins/manifest.rs b/codex-rs/core/src/plugins/manifest.rs new file mode 100644 index 0000000000..ad677de814 --- /dev/null +++ b/codex-rs/core/src/plugins/manifest.rs @@ -0,0 +1,37 @@ +use serde::Deserialize; +use std::fs; +use std::path::Path; + +pub(crate) const PLUGIN_MANIFEST_PATH: &str = ".codex-plugin/plugin.json"; + +#[derive(Debug, Default, Deserialize)] +pub(crate) struct PluginManifest { + name: String, +} + +pub(crate) fn load_plugin_manifest(plugin_root: &Path) -> Option { + let manifest_path = plugin_root.join(PLUGIN_MANIFEST_PATH); + if !manifest_path.is_file() { + return None; + } + let contents = fs::read_to_string(&manifest_path).ok()?; + match serde_json::from_str(&contents) { + Ok(manifest) => Some(manifest), + Err(err) => { + tracing::warn!( + path = %manifest_path.display(), + "failed to parse plugin manifest: {err}" + ); + None + } + } +} + +pub(crate) fn plugin_manifest_name(manifest: &PluginManifest, plugin_root: &Path) -> String { + plugin_root + .file_name() + .and_then(|name| name.to_str()) + .filter(|_| manifest.name.trim().is_empty()) + .unwrap_or(&manifest.name) + .to_string() +} diff --git a/codex-rs/core/src/plugins/mod.rs b/codex-rs/core/src/plugins/mod.rs new file mode 100644 index 0000000000..faf0a8dd0c --- /dev/null +++ b/codex-rs/core/src/plugins/mod.rs @@ -0,0 +1,14 @@ +mod manager; +mod manifest; +mod store; + +pub use manager::LoadedPlugin; +pub use manager::PluginInstallError; +pub use manager::PluginLoadOutcome; +pub use manager::PluginsManager; +pub(crate) use manager::plugin_namespace_for_skill_path; +pub(crate) use manifest::load_plugin_manifest; +pub(crate) use manifest::plugin_manifest_name; +pub use store::PluginId; +pub use store::PluginInstallRequest; +pub use store::PluginInstallResult; diff --git a/codex-rs/core/src/plugins/store.rs b/codex-rs/core/src/plugins/store.rs new file mode 100644 index 0000000000..96c355805e --- /dev/null +++ b/codex-rs/core/src/plugins/store.rs @@ -0,0 +1,374 @@ +use super::load_plugin_manifest; +use super::manifest::PLUGIN_MANIFEST_PATH; +use super::plugin_manifest_name; +use codex_utils_absolute_path::AbsolutePathBuf; +use std::fs; +use std::io; +use std::path::Path; +use std::path::PathBuf; + +const DEFAULT_MARKETPLACE_NAME: &str = "debug"; +pub(crate) const DEFAULT_PLUGIN_VERSION: &str = "local"; +pub(crate) const PLUGINS_CACHE_DIR: &str = "plugins/cache"; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct PluginInstallRequest { + pub source_path: PathBuf, + pub marketplace_name: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct PluginId { + pub plugin_name: String, + pub marketplace_name: String, +} + +impl PluginId { + pub fn new(plugin_name: String, marketplace_name: String) -> Result { + validate_plugin_segment(&plugin_name, "plugin name") + .map_err(PluginStoreError::InvalidPluginKey)?; + validate_plugin_segment(&marketplace_name, "marketplace name") + .map_err(PluginStoreError::InvalidPluginKey)?; + Ok(Self { + plugin_name, + marketplace_name, + }) + } + + pub fn parse(plugin_key: &str) -> Result { + let Some((plugin_name, marketplace_name)) = plugin_key.rsplit_once('@') else { + return Err(PluginStoreError::InvalidPluginKey(format!( + "invalid plugin key `{plugin_key}`; expected @" + ))); + }; + if plugin_name.is_empty() || marketplace_name.is_empty() { + return Err(PluginStoreError::InvalidPluginKey(format!( + "invalid plugin key `{plugin_key}`; expected @" + ))); + } + + Self::new(plugin_name.to_string(), marketplace_name.to_string()).map_err(|err| match err { + PluginStoreError::InvalidPluginKey(message) => { + PluginStoreError::InvalidPluginKey(format!("{message} in `{plugin_key}`")) + } + other => other, + }) + } + + pub fn as_key(&self) -> String { + format!("{}@{}", self.plugin_name, self.marketplace_name) + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct PluginInstallResult { + pub plugin_id: PluginId, + pub plugin_version: String, + pub installed_path: PathBuf, +} + +#[derive(Debug, Clone)] +pub struct PluginStore { + root: AbsolutePathBuf, +} + +impl PluginStore { + pub fn new(codex_home: PathBuf) -> Self { + Self { + root: AbsolutePathBuf::try_from(codex_home.join(PLUGINS_CACHE_DIR)) + .unwrap_or_else(|err| panic!("plugin cache root should be absolute: {err}")), + } + } + + pub fn root(&self) -> &AbsolutePathBuf { + &self.root + } + + pub fn plugin_root(&self, plugin_id: &PluginId, plugin_version: &str) -> AbsolutePathBuf { + AbsolutePathBuf::try_from( + self.root + .as_path() + .join(&plugin_id.marketplace_name) + .join(&plugin_id.plugin_name) + .join(plugin_version), + ) + .unwrap_or_else(|err| panic!("plugin cache path should resolve to an absolute path: {err}")) + } + + pub fn install( + &self, + request: PluginInstallRequest, + ) -> Result { + let source_path = request.source_path; + if !source_path.is_dir() { + return Err(PluginStoreError::InvalidPlugin(format!( + "plugin source path is not a directory: {}", + source_path.display() + ))); + } + + let plugin_name = plugin_name_for_source(&source_path)?; + let marketplace_name = request + .marketplace_name + .filter(|name| !name.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_MARKETPLACE_NAME.to_string()); + let plugin_version = DEFAULT_PLUGIN_VERSION.to_string(); + let plugin_id = match PluginId::new(plugin_name, marketplace_name) { + Ok(plugin_id) => plugin_id, + Err(PluginStoreError::InvalidPluginKey(message)) => { + return Err(PluginStoreError::InvalidPlugin(message)); + } + Err(err) => return Err(err), + }; + let installed_path = self + .plugin_root(&plugin_id, &plugin_version) + .into_path_buf(); + + if let Some(parent) = installed_path.parent() { + fs::create_dir_all(parent).map_err(|err| { + PluginStoreError::io("failed to create plugin cache directory", err) + })?; + } + + remove_existing_target(&installed_path)?; + copy_dir_recursive(&source_path, &installed_path)?; + + Ok(PluginInstallResult { + plugin_id, + plugin_version, + installed_path, + }) + } +} + +#[derive(Debug, thiserror::Error)] +pub enum PluginStoreError { + #[error("{context}: {source}")] + Io { + context: &'static str, + #[source] + source: io::Error, + }, + + #[error("{0}")] + InvalidPlugin(String), + + #[error("{0}")] + InvalidPluginKey(String), +} + +impl PluginStoreError { + fn io(context: &'static str, source: io::Error) -> Self { + Self::Io { context, source } + } +} + +fn plugin_name_for_source(source_path: &Path) -> Result { + let manifest_path = source_path.join(PLUGIN_MANIFEST_PATH); + if !manifest_path.is_file() { + return Err(PluginStoreError::InvalidPlugin(format!( + "missing plugin manifest: {}", + manifest_path.display() + ))); + } + + let manifest = load_plugin_manifest(source_path).ok_or_else(|| { + PluginStoreError::InvalidPlugin(format!( + "missing or invalid plugin manifest: {}", + manifest_path.display() + )) + })?; + + let plugin_name = plugin_manifest_name(&manifest, source_path); + validate_plugin_segment(&plugin_name, "plugin name") + .map_err(PluginStoreError::InvalidPlugin) + .map(|_| plugin_name) +} + +fn validate_plugin_segment(segment: &str, kind: &str) -> Result<(), String> { + if segment.is_empty() { + return Err(format!("invalid {kind}: must not be empty")); + } + if !segment + .chars() + .all(|ch| ch.is_ascii_alphanumeric() || ch == '-' || ch == '_') + { + return Err(format!( + "invalid {kind}: only ASCII letters, digits, `_`, and `-` are allowed" + )); + } + Ok(()) +} + +fn remove_existing_target(path: &Path) -> Result<(), PluginStoreError> { + if !path.exists() { + return Ok(()); + } + + if path.is_dir() { + fs::remove_dir_all(path).map_err(|err| { + PluginStoreError::io("failed to remove existing plugin cache entry", err) + }) + } else { + fs::remove_file(path).map_err(|err| { + PluginStoreError::io("failed to remove existing plugin cache entry", err) + }) + } +} + +fn copy_dir_recursive(source: &Path, target: &Path) -> Result<(), PluginStoreError> { + fs::create_dir_all(target) + .map_err(|err| PluginStoreError::io("failed to create plugin target directory", err))?; + + for entry in fs::read_dir(source) + .map_err(|err| PluginStoreError::io("failed to read plugin source directory", err))? + { + let entry = + entry.map_err(|err| PluginStoreError::io("failed to enumerate plugin source", err))?; + let source_path = entry.path(); + let target_path = target.join(entry.file_name()); + let file_type = entry + .file_type() + .map_err(|err| PluginStoreError::io("failed to inspect plugin source entry", err))?; + + if file_type.is_dir() { + copy_dir_recursive(&source_path, &target_path)?; + } else if file_type.is_file() { + fs::copy(&source_path, &target_path) + .map_err(|err| PluginStoreError::io("failed to copy plugin file", err))?; + } + } + + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use pretty_assertions::assert_eq; + use tempfile::tempdir; + + fn write_plugin(root: &Path, dir_name: &str, manifest_name: &str) { + let plugin_root = root.join(dir_name); + fs::create_dir_all(plugin_root.join(".codex-plugin")).unwrap(); + fs::create_dir_all(plugin_root.join("skills")).unwrap(); + fs::write( + plugin_root.join(".codex-plugin/plugin.json"), + format!(r#"{{"name":"{manifest_name}"}}"#), + ) + .unwrap(); + fs::write(plugin_root.join("skills/SKILL.md"), "skill").unwrap(); + fs::write(plugin_root.join(".mcp.json"), r#"{"mcpServers":{}}"#).unwrap(); + } + + #[test] + fn install_copies_plugin_into_default_marketplace() { + let tmp = tempdir().unwrap(); + write_plugin(tmp.path(), "sample-plugin", "sample-plugin"); + + let result = PluginStore::new(tmp.path().to_path_buf()) + .install(PluginInstallRequest { + source_path: tmp.path().join("sample-plugin"), + marketplace_name: None, + }) + .unwrap(); + + let installed_path = tmp.path().join("plugins/cache/debug/sample-plugin/local"); + assert_eq!( + result, + PluginInstallResult { + plugin_id: PluginId::new("sample-plugin".to_string(), "debug".to_string()).unwrap(), + plugin_version: "local".to_string(), + installed_path: installed_path.clone(), + } + ); + assert!(installed_path.join(".codex-plugin/plugin.json").is_file()); + assert!(installed_path.join("skills/SKILL.md").is_file()); + } + + #[test] + fn install_uses_manifest_name_for_destination_and_key() { + let tmp = tempdir().unwrap(); + write_plugin(tmp.path(), "source-dir", "manifest-name"); + + let result = PluginStore::new(tmp.path().to_path_buf()) + .install(PluginInstallRequest { + source_path: tmp.path().join("source-dir"), + marketplace_name: Some("market".to_string()), + }) + .unwrap(); + + assert_eq!( + result, + PluginInstallResult { + plugin_id: PluginId::new("manifest-name".to_string(), "market".to_string()) + .unwrap(), + plugin_version: "local".to_string(), + installed_path: tmp.path().join("plugins/cache/market/manifest-name/local"), + } + ); + } + + #[test] + fn plugin_root_derives_path_from_key_and_version() { + let tmp = tempdir().unwrap(); + let store = PluginStore::new(tmp.path().to_path_buf()); + let plugin_id = PluginId::new("sample".to_string(), "debug".to_string()).unwrap(); + + assert_eq!( + store.plugin_root(&plugin_id, "local").as_path(), + tmp.path().join("plugins/cache/debug/sample/local") + ); + } + + #[test] + fn plugin_root_rejects_path_separators_in_key_segments() { + let err = PluginId::parse("../../etc@debug").unwrap_err(); + assert_eq!( + err.to_string(), + "invalid plugin name: only ASCII letters, digits, `_`, and `-` are allowed in `../../etc@debug`" + ); + + let err = PluginId::parse("sample@../../etc").unwrap_err(); + assert_eq!( + err.to_string(), + "invalid marketplace name: only ASCII letters, digits, `_`, and `-` are allowed in `sample@../../etc`" + ); + } + + #[test] + fn install_rejects_manifest_names_with_path_separators() { + let tmp = tempdir().unwrap(); + write_plugin(tmp.path(), "source-dir", "../../etc"); + + let err = PluginStore::new(tmp.path().to_path_buf()) + .install(PluginInstallRequest { + source_path: tmp.path().join("source-dir"), + marketplace_name: None, + }) + .unwrap_err(); + + assert_eq!( + err.to_string(), + "invalid plugin name: only ASCII letters, digits, `_`, and `-` are allowed" + ); + } + + #[test] + fn install_rejects_marketplace_names_with_path_separators() { + let tmp = tempdir().unwrap(); + write_plugin(tmp.path(), "source-dir", "sample-plugin"); + + let err = PluginStore::new(tmp.path().to_path_buf()) + .install(PluginInstallRequest { + source_path: tmp.path().join("source-dir"), + marketplace_name: Some("../../etc".to_string()), + }) + .unwrap_err(); + + assert_eq!( + err.to_string(), + "invalid marketplace name: only ASCII letters, digits, `_`, and `-` are allowed" + ); + } +} diff --git a/codex-rs/core/tests/suite/plugins.rs b/codex-rs/core/tests/suite/plugins.rs index ed2ba54a3a..3e641d010d 100644 --- a/codex-rs/core/tests/suite/plugins.rs +++ b/codex-rs/core/tests/suite/plugins.rs @@ -24,7 +24,7 @@ use tempfile::TempDir; use wiremock::MockServer; fn write_plugin_skill_plugin(home: &TempDir) -> std::path::PathBuf { - let plugin_root = home.path().join("plugins/sample"); + let plugin_root = home.path().join("plugins/cache/test/sample/local"); let skill_dir = plugin_root.join("skills/sample-search"); std::fs::create_dir_all(skill_dir.as_path()).expect("create plugin skill dir"); std::fs::create_dir_all(plugin_root.join(".codex-plugin")).expect("create plugin manifest dir"); @@ -40,17 +40,14 @@ fn write_plugin_skill_plugin(home: &TempDir) -> std::path::PathBuf { .expect("write plugin skill"); std::fs::write( home.path().join("config.toml"), - format!( - "[features]\nplugins = true\n\n[plugins.sample]\nenabled = true\npath = \"{}\"\n", - plugin_root.display() - ), + "[features]\nplugins = true\n\n[plugins.\"sample@test\"]\nenabled = true\n", ) .expect("write config"); skill_dir.join("SKILL.md") } fn write_plugin_mcp_plugin(home: &TempDir, command: &str) { - let plugin_root = home.path().join("plugins/sample"); + let plugin_root = home.path().join("plugins/cache/test/sample/local"); std::fs::create_dir_all(plugin_root.join(".codex-plugin")).expect("create plugin manifest dir"); std::fs::write( plugin_root.join(".codex-plugin/plugin.json"), @@ -72,10 +69,7 @@ fn write_plugin_mcp_plugin(home: &TempDir, command: &str) { .expect("write plugin mcp config"); std::fs::write( home.path().join("config.toml"), - format!( - "[features]\nplugins = true\n\n[plugins.sample]\nenabled = true\npath = \"{}\"\n", - plugin_root.display() - ), + "[features]\nplugins = true\n\n[plugins.\"sample@test\"]\nenabled = true\n", ) .expect("write config"); } From 24a2d0c696d8b1b3d0c137f8e67126b33b07d189 Mon Sep 17 00:00:00 2001 From: viyatb-oai Date: Tue, 3 Mar 2026 15:12:06 -0800 Subject: [PATCH 5/5] fix(network-proxy): reject mismatched host headers (#13275) ## Summary - reject plain HTTP absolute-form requests whose Host header does not match the request target authority - add host/port-aware Host header validation for non-default ports - add regression coverage for mismatched Host forwarding and validator edge cases --- codex-rs/network-proxy/src/http_proxy.rs | 139 +++++++++++++++-------- 1 file changed, 94 insertions(+), 45 deletions(-) diff --git a/codex-rs/network-proxy/src/http_proxy.rs b/codex-rs/network-proxy/src/http_proxy.rs index cc08cd3b25..4f88d25383 100644 --- a/codex-rs/network-proxy/src/http_proxy.rs +++ b/codex-rs/network-proxy/src/http_proxy.rs @@ -48,6 +48,8 @@ use rama_http::Request; use rama_http::Response; use rama_http::StatusCode; use rama_http::header; +use rama_http::headers::HeaderMapExt; +use rama_http::headers::Host; use rama_http::layer::remove_header::RemoveResponseHeaderLayer; use rama_http::matcher::MethodMatcher; use rama_http_backend::client::proxy::layer::HttpProxyConnector; @@ -55,7 +57,6 @@ use rama_http_backend::server::HttpServer; use rama_http_backend::server::layer::upgrade::UpgradeLayer; use rama_http_backend::server::layer::upgrade::Upgraded; use rama_net::Protocol; -use rama_net::address::HostWithOptPort; use rama_net::address::ProxyAddress; use rama_net::client::ConnectorService; use rama_net::client::EstablishedClientConnection; @@ -562,18 +563,27 @@ async fn http_plain_proxy( }; } - let authority = match RequestContext::try_from(&req).map(|ctx| ctx.host_with_port()) { - Ok(authority) => authority, + let request_ctx = match RequestContext::try_from(&req) { + Ok(request_ctx) => request_ctx, Err(err) => { warn!("missing host: {err}"); return Ok(text_response(StatusCode::BAD_REQUEST, "missing host")); } }; + let authority = request_ctx.host_with_port(); let host = normalize_host(&authority.host.to_string()); let port = authority.port; - if let Err(err) = validate_plain_http_host_header(&req, &authority) { - warn!("HTTP request host mismatch: {err}"); - return Ok(text_response(StatusCode::BAD_REQUEST, "host mismatch")); + if let Err(reason) = validate_absolute_form_host_header(&req, &request_ctx) { + let client = client.as_deref().unwrap_or_default(); + let host_header = req + .headers() + .get(header::HOST) + .and_then(|value| value.to_str().ok()) + .unwrap_or(""); + warn!( + "request rejected due to mismatched Host header (client={client}, target={host}:{port}, host_header={host_header}, reason={reason})" + ); + return Ok(text_response(StatusCode::BAD_REQUEST, reason)); } let enabled = match app_state .enabled() @@ -757,6 +767,39 @@ fn client_addr(input: &T) -> Option { .map(|info| info.peer_addr().to_string()) } +fn validate_absolute_form_host_header( + req: &Request, + request_ctx: &RequestContext, +) -> Result<(), &'static str> { + if req.uri().scheme_str().is_none() { + return Ok(()); + } + + let Some(host_header) = req + .headers() + .typed_try_get::() + .map_err(|_| "invalid Host header")? + else { + return Ok(()); + }; + + if host_header.0.host != request_ctx.authority.host { + return Err("Host header does not match request target"); + } + + if let Some(host_port) = host_header.0.port { + if Some(host_port) != request_ctx.authority.port { + return Err("Host header does not match request target"); + } + return Ok(()); + } + + if !request_ctx.authority_has_default_port() { + return Err("Host header does not match request target"); + } + + Ok(()) +} fn remove_hop_by_hop_request_headers(headers: &mut HeaderMap) { while let Some(raw_connection) = headers.get(header::CONNECTION).cloned() { headers.remove(header::CONNECTION); @@ -792,45 +835,6 @@ fn remove_hop_by_hop_request_headers(headers: &mut HeaderMap) { } } -fn validate_plain_http_host_header( - req: &Request, - target: &rama_net::address::HostWithPort, -) -> std::result::Result<(), &'static str> { - // Only enforce this in absolute-form requests. Origin-form requests use the Host header as the - // routing authority, so there is no separate target authority to compare against. - if req.uri().authority().is_none() { - return Ok(()); - } - - let Some(raw_host) = req.headers().get(header::HOST) else { - return Ok(()); - }; - let raw_host = raw_host.to_str().map_err(|_| "invalid Host header")?; - let parsed = HostWithOptPort::try_from(raw_host).map_err(|_| "invalid Host header")?; - - let target_host = normalize_host(&target.host.to_string()); - let request_host = normalize_host(&parsed.host.to_string()); - if request_host.is_empty() || request_host != target_host { - return Err("request Host header host does not match target authority"); - } - - let expected_port = target.port; - let request_port = match parsed.port { - Some(port) => port, - None => match req.uri().scheme_str() { - Some("http") => 80, - Some("https") => 443, - Some(_) | None => expected_port, - }, - }; - - if request_port != expected_port { - return Err("request Host header port does not match target authority"); - } - - Ok(()) -} - fn json_blocked(host: &str, reason: &str, details: Option<&PolicyDecisionDetails<'_>>) -> Response { let (message, decision, source, protocol, port) = details .map(|details| { @@ -1136,6 +1140,51 @@ mod tests { assert_eq!(response.unwrap().status(), StatusCode::BAD_REQUEST); } + #[test] + fn validate_absolute_form_host_header_allows_matching_default_port() { + let req = Request::builder() + .method(Method::GET) + .uri("http://example.com/") + .header("host", "example.com") + .body(Body::empty()) + .unwrap(); + + assert_eq!( + validate_absolute_form_host_header(&req, &RequestContext::try_from(&req).unwrap(),), + Ok(()) + ); + } + + #[test] + fn validate_absolute_form_host_header_rejects_mismatched_host() { + let req = Request::builder() + .method(Method::GET) + .uri("http://raw.githubusercontent.com/") + .header("host", "api.github.com") + .body(Body::empty()) + .unwrap(); + + assert_eq!( + validate_absolute_form_host_header(&req, &RequestContext::try_from(&req).unwrap(),), + Err("Host header does not match request target") + ); + } + + #[test] + fn validate_absolute_form_host_header_rejects_missing_non_default_port() { + let req = Request::builder() + .method(Method::GET) + .uri("http://example.com:8080/") + .header("host", "example.com") + .body(Body::empty()) + .unwrap(); + + assert_eq!( + validate_absolute_form_host_header(&req, &RequestContext::try_from(&req).unwrap(),), + Err("Host header does not match request target") + ); + } + #[test] fn remove_hop_by_hop_request_headers_keeps_forwarding_headers() { let mut headers = HeaderMap::new();