diff --git a/codex-rs/Cargo.lock b/codex-rs/Cargo.lock index 54038c8488..8a55a2b6b5 100644 --- a/codex-rs/Cargo.lock +++ b/codex-rs/Cargo.lock @@ -2660,6 +2660,7 @@ dependencies = [ "serde_json", "serial_test", "sha1 0.10.6", + "sha2 0.10.9", "shlex", "similar", "tempfile", @@ -11230,6 +11231,7 @@ dependencies = [ "js-sys", "log", "mime", + "mime_guess", "native-tls", "percent-encoding", "pin-project-lite", diff --git a/codex-rs/codex-mcp/src/connection_manager.rs b/codex-rs/codex-mcp/src/connection_manager.rs index 81553fae2e..94c0259559 100644 --- a/codex-rs/codex-mcp/src/connection_manager.rs +++ b/codex-rs/codex-mcp/src/connection_manager.rs @@ -20,15 +20,12 @@ use crate::codex_apps::CodexAppsToolsCacheKey; use crate::codex_apps::write_cached_codex_apps_tools_if_needed; use crate::elicitation::ElicitationRequestManager; use crate::elicitation::ElicitationReviewerHandle; -use crate::file_transfer::CompleteUploadResult; +use crate::file_transfer::AuthorizeDownloadResult; +use crate::file_transfer::AuthorizeUploadParams; +use crate::file_transfer::AuthorizeUploadResult; use crate::file_transfer::FileUriParams; -use crate::file_transfer::GetDownloadResult; -use crate::file_transfer::METHOD_FILES_COMPLETE_UPLOAD; -use crate::file_transfer::METHOD_FILES_GET_DOWNLOAD; -use crate::file_transfer::METHOD_FILES_PREPARE_UPLOAD; -use crate::file_transfer::McpFileCapabilities; -use crate::file_transfer::PrepareUploadParams; -use crate::file_transfer::PrepareUploadResult; +use crate::file_transfer::METHOD_FILES_AUTHORIZE_DOWNLOAD; +use crate::file_transfer::METHOD_FILES_AUTHORIZE_UPLOAD; use crate::mcp::CODEX_APPS_MCP_SERVER_NAME; use crate::mcp::ToolPluginProvenance; use crate::rmcp_client::AsyncManagedClient; @@ -749,60 +746,32 @@ impl McpConnectionManager { }) } - pub async fn prepare_file_upload( + pub async fn authorize_file_upload( &self, server: &str, - params: PrepareUploadParams, - ) -> Result { - self.require_file_capability(server, |capabilities| capabilities.prepare_upload) - .await?; - self.send_file_request(server, METHOD_FILES_PREPARE_UPLOAD, params) + params: AuthorizeUploadParams, + ) -> Result { + self.send_file_request(server, METHOD_FILES_AUTHORIZE_UPLOAD, params) .await } - pub async fn file_capabilities(&self, server: &str) -> Result { - if !self.mcp_file_transfer_enabled.load(Ordering::Relaxed) { - return Ok(McpFileCapabilities::default()); - } - Ok(self.client_by_name(server).await?.file_capabilities) - } - - pub async fn complete_file_upload( + pub async fn authorize_file_download( &self, server: &str, uri: String, - ) -> Result { - self.require_file_capability(server, |capabilities| capabilities.complete_upload) - .await?; - self.send_file_request(server, METHOD_FILES_COMPLETE_UPLOAD, FileUriParams { uri }) - .await - } - - pub async fn get_file_download(&self, server: &str, uri: String) -> Result { + ) -> Result { if !self.mcp_file_transfer_enabled.load(Ordering::Relaxed) { return Err(anyhow!("MCP file transfer is disabled")); } - // rmcp 1.7 drops the draft top-level `capabilities.files` object. A - // structured `mcp-file://` tool result is itself sufficient evidence - // to try the matching draft download method. - self.send_file_request(server, METHOD_FILES_GET_DOWNLOAD, FileUriParams { uri }) - .await - } - - async fn require_file_capability( - &self, - server: &str, - supported: impl FnOnce(McpFileCapabilities) -> bool, - ) -> Result<()> { - if !self.mcp_file_transfer_enabled.load(Ordering::Relaxed) { - return Err(anyhow!("MCP file transfer is disabled")); - } - if !supported(self.file_capabilities(server).await?) { - return Err(anyhow!( - "MCP server `{server}` does not advertise the required file capability" - )); - } - Ok(()) + // SEP-2631 defines `capabilities.files` on the client. A file URI + // returned by the server is sufficient evidence to try the download + // authorization method without a separate server capability gate. + self.send_file_request( + server, + METHOD_FILES_AUTHORIZE_DOWNLOAD, + FileUriParams { uri }, + ) + .await } async fn send_file_request(&self, server: &str, method: &str, params: P) -> Result diff --git a/codex-rs/codex-mcp/src/connection_manager_tests.rs b/codex-rs/codex-mcp/src/connection_manager_tests.rs index 0a7f7bd851..2a2e8f2987 100644 --- a/codex-rs/codex-mcp/src/connection_manager_tests.rs +++ b/codex-rs/codex-mcp/src/connection_manager_tests.rs @@ -850,12 +850,13 @@ async fn file_methods_are_rejected_while_feature_is_disabled() { ); let error = manager - .prepare_file_upload( + .authorize_file_upload( "server", - PrepareUploadParams { + AuthorizeUploadParams { name: "file.txt".to_string(), mime_type: "text/plain".to_string(), size: 4, + digest: None, }, ) .await diff --git a/codex-rs/codex-mcp/src/file_transfer.rs b/codex-rs/codex-mcp/src/file_transfer.rs index f5964fddf7..d4ec8f3826 100644 --- a/codex-rs/codex-mcp/src/file_transfer.rs +++ b/codex-rs/codex-mcp/src/file_transfer.rs @@ -4,7 +4,6 @@ use std::collections::BTreeMap; use std::collections::BTreeSet; use std::sync::Arc; -use rmcp::model::ServerCapabilities; use rmcp::model::Tool; use serde::Deserialize; use serde::Serialize; @@ -16,13 +15,19 @@ const META_OPENAI_FILE_PARAMS: &str = "openai/fileParams"; const MCP_FILE_HANDLE_GUIDANCE: &str = "Pass an absolute local file path. Do not construct data: URIs, mcp-file:// handles, signed URLs, or file-service payloads."; const OPENAI_FILE_HANDLE_GUIDANCE: &str = "This parameter expects an absolute local file path. If you want to upload a file, provide the absolute path to that file here."; -pub(crate) const METHOD_FILES_PREPARE_UPLOAD: &str = "files/prepareUpload"; -pub(crate) const METHOD_FILES_COMPLETE_UPLOAD: &str = "files/completeUpload"; -pub(crate) const METHOD_FILES_GET_DOWNLOAD: &str = "files/getDownload"; +pub(crate) const METHOD_FILES_AUTHORIZE_UPLOAD: &str = "files/authorizeUpload"; +pub(crate) const METHOD_FILES_AUTHORIZE_DOWNLOAD: &str = "files/authorizeDownload"; #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] -pub struct McpFileValue { +pub struct FileDigest { + pub algorithm: String, + pub value: String, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct FileValue { pub uri: String, #[serde(default)] pub name: Option, @@ -30,25 +35,40 @@ pub struct McpFileValue { pub mime_type: Option, #[serde(default)] pub size: Option, + #[serde(default)] + pub digest: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct FileTransferMultipart { + pub file_field: String, + #[serde(default)] + pub fields: BTreeMap, } #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct FileTransferDescriptor { - #[serde(default)] - pub transport: Option, + pub transport: String, pub method: String, pub url: String, #[serde(default)] + pub headers: BTreeMap, + #[serde(default)] + pub multipart: Option, + #[serde(default)] pub expires_at: Option, } #[derive(Debug, Clone, PartialEq, Eq, Serialize)] #[serde(rename_all = "camelCase")] -pub struct PrepareUploadParams { +pub struct AuthorizeUploadParams { pub name: String, pub mime_type: String, pub size: u64, + #[serde(skip_serializing_if = "Option::is_none")] + pub digest: Option, } #[derive(Debug, Clone, PartialEq, Eq, Serialize)] @@ -57,68 +77,21 @@ pub struct FileUriParams { } #[derive(Debug, Clone, PartialEq, Eq, Deserialize)] -pub struct PrepareUploadResult { - pub file: McpFileValue, - pub transfer: FileTransferDescriptor, +pub struct AuthorizeUploadResult { + pub file: FileValue, + pub upload: FileTransferDescriptor, + #[serde(default)] + pub download: Option, } #[derive(Debug, Clone, PartialEq, Eq, Deserialize)] -pub struct CompleteUploadResult { - pub file: McpFileValue, +pub struct AuthorizedFile { + pub file: FileValue, + #[serde(default)] + pub download: Option, } -#[derive(Debug, Clone, PartialEq, Eq, Deserialize)] -pub struct GetDownloadResult { - pub file: McpFileValue, - pub transfer: FileTransferDescriptor, -} - -#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] -pub struct McpFileCapabilities { - pub prepare_upload: bool, - pub complete_upload: bool, - pub get_download: bool, -} - -impl McpFileCapabilities { - pub(crate) fn from_server_and_tools( - capabilities: &ServerCapabilities, - tools: &[crate::ToolInfo], - ) -> Self { - let extension = capabilities.extensions.as_ref().and_then(|extensions| { - extensions - .get("io.modelcontextprotocol/files") - .or_else(|| extensions.get("files")) - }); - if let Some(extension) = extension { - return Self { - prepare_upload: extension - .get("prepareUpload") - .and_then(JsonValue::as_bool) - .unwrap_or(true), - complete_upload: extension - .get("completeUpload") - .and_then(JsonValue::as_bool) - .unwrap_or(true), - get_download: extension - .get("getDownload") - .and_then(JsonValue::as_bool) - .unwrap_or(true), - }; - } - - let exposes_upload = tools.iter().any(|tool| { - file_input_specs(&tool.tool) - .iter() - .any(FileInputSpec::accepts_upload) - }); - Self { - prepare_upload: exposes_upload, - complete_upload: exposes_upload, - get_download: exposes_upload, - } - } -} +pub type AuthorizeDownloadResult = AuthorizedFile; #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)] #[serde(rename_all = "snake_case")] @@ -170,8 +143,11 @@ pub fn file_input_specs(tool: &Tool) -> Vec { .and_then(JsonValue::as_object) { for (path, property) in properties { + let is_array = property.get("type").and_then(JsonValue::as_str) == Some("array") + || property.get("items").is_some(); let Some(extension) = property .get(MCP_FILE_SCHEMA_EXTENSION) + .or_else(|| property.get("items")?.get(MCP_FILE_SCHEMA_EXTENSION)) .and_then(JsonValue::as_object) else { continue; @@ -185,8 +161,7 @@ pub fn file_input_specs(tool: &Tool) -> Vec { max_size: extension.get("maxSize").and_then(JsonValue::as_u64), transfer_modes, sources: BTreeSet::from([FileInputSource::Mcp]), - is_array: property.get("type").and_then(JsonValue::as_str) == Some("array") - || property.get("items").is_some(), + is_array, }, ); } @@ -250,7 +225,9 @@ fn parse_transfer_modes(value: Option<&JsonValue>) -> BTreeSet let values = match value { Some(JsonValue::Array(values)) => values.as_slice(), Some(_) => return BTreeSet::new(), - None => return BTreeSet::from([FileTransferMode::Inline]), + None => { + return BTreeSet::from([FileTransferMode::Inline, FileTransferMode::Upload]); + } }; values .iter() diff --git a/codex-rs/codex-mcp/src/file_transfer_tests.rs b/codex-rs/codex-mcp/src/file_transfer_tests.rs index ccfa274920..0e46bef1fe 100644 --- a/codex-rs/codex-mcp/src/file_transfer_tests.rs +++ b/codex-rs/codex-mcp/src/file_transfer_tests.rs @@ -1,7 +1,6 @@ use super::*; use pretty_assertions::assert_eq; use rmcp::model::Meta; -use rmcp::model::ServerCapabilities; fn tool(input_schema: JsonValue, meta: Option) -> Tool { let mut tool = Tool::new( @@ -59,7 +58,7 @@ fn parses_and_merges_mcp_and_openai_file_inputs() { } #[test] -fn missing_transfer_modes_defaults_to_inline() { +fn missing_transfer_modes_allows_inline_and_upload() { let tool = tool( serde_json::json!({ "type": "object", @@ -69,7 +68,7 @@ fn missing_transfer_modes_defaults_to_inline() { ); assert_eq!( file_input_specs(&tool)[0].transfer_modes, - BTreeSet::from([FileTransferMode::Inline]) + BTreeSet::from([FileTransferMode::Inline, FileTransferMode::Upload]) ); } @@ -81,11 +80,14 @@ fn malformed_extension_values_are_ignored_and_arrays_are_preserved() { "properties": { "files": { "type": "array", - "items": {"type": "string"}, - "x-mcp-file": { - "accept": ["text/plain", 7, ""], - "maxSize": -1, - "transferModes": ["upload", "future-mode", 7] + "items": { + "type": "string", + "format": "uri", + "x-mcp-file": { + "accept": ["text/plain", 7, ""], + "maxSize": -1, + "transferModes": ["upload", "future-mode", 7] + } } } } @@ -142,92 +144,104 @@ fn mcp_schema_masking_is_gated_while_legacy_masking_is_not() { } #[test] -fn file_capabilities_use_sep_extension_when_available() { - let mut capabilities = ServerCapabilities::default(); - capabilities.extensions = Some(std::collections::BTreeMap::from([( - "io.modelcontextprotocol/files".to_string(), - serde_json::json!({ - "prepareUpload": true, - "completeUpload": true, - "getDownload": false - }) - .as_object() - .expect("object") - .clone(), - )])); - - assert_eq!( - McpFileCapabilities::from_server_and_tools(&capabilities, &[]), - McpFileCapabilities { - prepare_upload: true, - complete_upload: true, - get_download: false, - } - ); -} - -#[test] -fn file_capabilities_infer_openai_compatibility_from_upload_schema() { - let tool = tool( - serde_json::json!({ - "type": "object", - "properties": { - "file": {"x-mcp-file": {"transferModes": ["upload"]}} - } - }), - /*meta*/ None, - ); - let tool = crate::ToolInfo { - server_name: "server".to_string(), - supports_parallel_tool_calls: false, - server_origin: None, - callable_name: "upload".to_string(), - callable_namespace: "server".to_string(), - namespace_description: None, - tool, - connector_id: None, - connector_name: None, - plugin_display_names: Vec::new(), - }; - - assert_eq!( - McpFileCapabilities::from_server_and_tools(&ServerCapabilities::default(), &[tool]), - McpFileCapabilities { - prepare_upload: true, - complete_upload: true, - get_download: true, - } - ); -} - -#[test] -fn transfer_descriptors_accept_draft_and_extended_shapes() { - assert_eq!( - serde_json::from_value::(serde_json::json!({ - "method": "PUT", - "url": "https://example.com/upload" - })) - .expect("draft descriptor"), - FileTransferDescriptor { - transport: None, - method: "PUT".to_string(), - url: "https://example.com/upload".to_string(), - expires_at: None, - } - ); +fn transfer_descriptors_accept_sep_shape() { assert_eq!( serde_json::from_value::(serde_json::json!({ "transport": "https", "method": "GET", "url": "https://example.com/download", + "headers": {"Authorization": "Bearer secret"}, + "multipart": { + "fileField": "payload", + "fields": {"token": "abc123"} + }, "expiresAt": "2030-01-01T00:00:00Z" })) .expect("extended descriptor"), FileTransferDescriptor { - transport: Some("https".to_string()), + transport: "https".to_string(), method: "GET".to_string(), url: "https://example.com/download".to_string(), + headers: BTreeMap::from([("Authorization".to_string(), "Bearer secret".to_string(),)]), + multipart: Some(FileTransferMultipart { + file_field: "payload".to_string(), + fields: BTreeMap::from([("token".to_string(), "abc123".to_string())]), + }), expires_at: Some("2030-01-01T00:00:00Z".to_string()), } ); } + +#[test] +fn authorize_upload_params_match_sep_wire_shape() { + assert_eq!( + serde_json::to_value(AuthorizeUploadParams { + name: "report.pdf".to_string(), + mime_type: "application/pdf".to_string(), + size: 248_123, + digest: Some(FileDigest { + algorithm: "sha-256".to_string(), + value: "digest-value".to_string(), + }), + }) + .expect("serialize params"), + serde_json::json!({ + "name": "report.pdf", + "mimeType": "application/pdf", + "size": 248123, + "digest": {"algorithm": "sha-256", "value": "digest-value"} + }) + ); +} + +#[test] +fn authorize_results_match_sep_wire_shapes() { + let file = serde_json::json!({ + "uri": "mcp-file://server/file-1", + "name": "report.pdf", + "mimeType": "application/pdf", + "size": 248123, + "digest": {"algorithm": "sha-256", "value": "digest-value"} + }); + let upload = serde_json::json!({ + "transport": "https", + "method": "POST", + "url": "https://upload.example.com/file-1", + "headers": {"X-Upload": "token"}, + "multipart": {"fileField": "file", "fields": {"token": "abc123"}}, + "expiresAt": "2030-01-01T00:00:00Z" + }); + let download = serde_json::json!({ + "transport": "https", + "method": "GET", + "url": "https://download.example.com/file-1" + }); + + let result: AuthorizeUploadResult = serde_json::from_value(serde_json::json!({ + "file": file, + "upload": upload, + "download": download + })) + .expect("authorize upload result"); + assert_eq!( + result, + AuthorizeUploadResult { + file: serde_json::from_value(file.clone()).expect("file value"), + upload: serde_json::from_value(upload).expect("upload descriptor"), + download: Some(serde_json::from_value(download.clone()).expect("download descriptor")), + } + ); + + let result: AuthorizeDownloadResult = serde_json::from_value(serde_json::json!({ + "file": file, + "download": download + })) + .expect("authorize download result"); + assert_eq!( + result, + AuthorizedFile { + file: serde_json::from_value(file).expect("file value"), + download: Some(serde_json::from_value(download).expect("download descriptor")), + } + ); +} diff --git a/codex-rs/codex-mcp/src/lib.rs b/codex-rs/codex-mcp/src/lib.rs index fd4877da47..dd06c93045 100644 --- a/codex-rs/codex-mcp/src/lib.rs +++ b/codex-rs/codex-mcp/src/lib.rs @@ -3,16 +3,17 @@ pub use connection_manager::tool_is_model_visible; pub use elicitation::ElicitationReviewRequest; pub use elicitation::ElicitationReviewer; pub use elicitation::ElicitationReviewerHandle; -pub use file_transfer::CompleteUploadResult; +pub use file_transfer::AuthorizeDownloadResult; +pub use file_transfer::AuthorizeUploadParams; +pub use file_transfer::AuthorizeUploadResult; +pub use file_transfer::AuthorizedFile; +pub use file_transfer::FileDigest; pub use file_transfer::FileInputSource; pub use file_transfer::FileInputSpec; pub use file_transfer::FileTransferDescriptor; pub use file_transfer::FileTransferMode; -pub use file_transfer::GetDownloadResult; -pub use file_transfer::McpFileCapabilities; -pub use file_transfer::McpFileValue; -pub use file_transfer::PrepareUploadParams; -pub use file_transfer::PrepareUploadResult; +pub use file_transfer::FileTransferMultipart; +pub use file_transfer::FileValue; pub use file_transfer::file_input_specs; pub use resource_client::McpResourceClient; pub use resource_client::McpResourcePage; diff --git a/codex-rs/codex-mcp/src/rmcp_client.rs b/codex-rs/codex-mcp/src/rmcp_client.rs index df43d516b1..b77516085d 100644 --- a/codex-rs/codex-mcp/src/rmcp_client.rs +++ b/codex-rs/codex-mcp/src/rmcp_client.rs @@ -27,7 +27,6 @@ use crate::codex_apps::normalize_codex_apps_callable_namespace; use crate::codex_apps::normalize_codex_apps_tool_title; use crate::codex_apps::write_cached_codex_apps_tools_if_needed; use crate::elicitation::ElicitationRequestManager; -use crate::file_transfer::McpFileCapabilities; use crate::mcp::CODEX_APPS_MCP_SERVER_NAME; use crate::mcp::ToolPluginProvenance; use crate::runtime::McpRuntimeContext; @@ -94,7 +93,6 @@ pub(crate) struct ManagedClient { pub(crate) tool_timeout: Option, pub(crate) server_instructions: Option, pub(crate) server_supports_sandbox_state_meta_capability: bool, - pub(crate) file_capabilities: McpFileCapabilities, pub(crate) codex_apps_tools_cache_context: Option, } @@ -533,8 +531,6 @@ async fn start_server_task( } let tools = filter_tools(tools, &tool_filter); - let file_capabilities = - McpFileCapabilities::from_server_and_tools(&initialize_result.capabilities, &tools); let managed = ManagedClient { client: Arc::clone(&client), server_info, @@ -543,7 +539,6 @@ async fn start_server_task( tool_filter, server_instructions: initialize_result.instructions, server_supports_sandbox_state_meta_capability, - file_capabilities, codex_apps_tools_cache_context, }; diff --git a/codex-rs/core/Cargo.toml b/codex-rs/core/Cargo.toml index c8e24a46e7..02cafdab42 100644 --- a/codex-rs/core/Cargo.toml +++ b/codex-rs/core/Cargo.toml @@ -89,7 +89,7 @@ mime_guess = { workspace = true } once_cell = { workspace = true } rand = { workspace = true } regex-lite = { workspace = true } -reqwest = { workspace = true, features = ["json", "stream"] } +reqwest = { workspace = true, features = ["json", "multipart", "stream"] } rmcp = { workspace = true, default-features = false, features = [ "base64", "macros", @@ -98,6 +98,7 @@ rmcp = { workspace = true, default-features = false, features = [ ] } serde = { workspace = true, features = ["derive"] } serde_json = { workspace = true } +sha2 = { workspace = true } sha1 = { workspace = true } shlex = { workspace = true } similar = { workspace = true } diff --git a/codex-rs/core/src/mcp_file_transfer.rs b/codex-rs/core/src/mcp_file_transfer.rs index b5f645df28..9b944af0e7 100644 --- a/codex-rs/core/src/mcp_file_transfer.rs +++ b/codex-rs/core/src/mcp_file_transfer.rs @@ -5,11 +5,15 @@ use std::path::PathBuf; use base64::Engine; use base64::engine::general_purpose::STANDARD; +use base64::engine::general_purpose::URL_SAFE_NO_PAD; +use codex_mcp::AuthorizeUploadParams; +use codex_mcp::FileDigest; use codex_mcp::FileInputSpec; -use codex_mcp::PrepareUploadParams; use codex_protocol::mcp::CallToolResult; use codex_protocol::permissions::ReadDenyMatcher; use serde_json::Value as JsonValue; +use sha2::Digest; +use sha2::Sha256; use url::Url; use crate::session::session::Session; @@ -73,10 +77,10 @@ pub(crate) async fn materialize_mcp_file_outputs( ) -> Result { let mut files = HashMap::::new(); for content in &result.content { - collect_output_files(content, &mut files); + collect_output_files(content, &mut files)?; } if let Some(structured_content) = result.structured_content.as_ref() { - collect_output_files(structured_content, &mut files); + collect_output_files(structured_content, &mut files)?; } if files.is_empty() { return Ok(result); @@ -96,9 +100,9 @@ pub(crate) async fn materialize_mcp_file_outputs( let mut replacements = HashMap::new(); for (uri, file) in files { let download = manager - .get_file_download(server, uri.clone()) + .authorize_file_download(server, uri.clone()) .await - .map_err(|error| format!("failed to prepare MCP file download: {error:#}"))?; + .map_err(|error| format!("failed to authorize MCP file download: {error:#}"))?; validate_mcp_file_uri(&download.file.uri)?; if download.file.uri != uri { return Err("MCP download response returned a different file URI".to_string()); @@ -112,16 +116,16 @@ pub(crate) async fn materialize_mcp_file_outputs( .unwrap_or("download"), ); let output_path = unique_output_path(&output_dir, &name).await; + let transfer = download.download.as_ref().ok_or_else(|| { + "MCP download authorization omitted a transfer descriptor".to_string() + })?; let size = download_transfer_file( sess, - &download.transfer, + transfer, &output_path, - download - .file - .size - .filter(|size| *size > 0) - .unwrap_or(DEFAULT_MAX_FILE_BYTES) - .min(DEFAULT_MAX_FILE_BYTES), + DEFAULT_MAX_FILE_BYTES, + download.file.size.or(file.size), + download.file.digest.as_ref().or(file.digest.as_ref()), ) .await?; let local_uri = Url::from_file_path(&output_path) @@ -134,6 +138,7 @@ pub(crate) async fn materialize_mcp_file_outputs( "name": name, "mimeType": download.file.mime_type.or(file.mime_type), "size": size, + "digest": download.file.digest.or(file.digest), }), ); } @@ -150,19 +155,41 @@ pub(crate) async fn materialize_mcp_file_outputs( struct McpOutputFile { name: Option, mime_type: Option, + size: Option, + digest: Option, } -fn collect_output_files(value: &JsonValue, files: &mut HashMap) { +fn collect_output_files( + value: &JsonValue, + files: &mut HashMap, +) -> Result<(), String> { match value { JsonValue::Array(values) => { for value in values { - collect_output_files(value, files); + collect_output_files(value, files)?; } } JsonValue::Object(object) => { + if object.get("type").and_then(JsonValue::as_str) == Some("resource_link") { + return Ok(()); + } if let Some(uri) = object.get("uri").and_then(JsonValue::as_str) - && uri.starts_with("mcp-file://") + && is_file_transfer_uri(uri) { + let size = object + .get("size") + .map(|size| { + size.as_u64().ok_or_else(|| { + "MCP file response returned an invalid file size".to_string() + }) + }) + .transpose()?; + let digest = object + .get("digest") + .cloned() + .map(serde_json::from_value) + .transpose() + .map_err(|_| "MCP file response returned an invalid digest".to_string())?; files .entry(uri.to_string()) .or_insert_with(|| McpOutputFile { @@ -175,15 +202,18 @@ fn collect_output_files(value: &JsonValue, files: &mut HashMap {} } + Ok(()) } fn replace_output_files(value: &mut JsonValue, replacements: &HashMap) { @@ -335,10 +365,12 @@ async fn rewrite_single_file( let mime_type = mime_guess::from_path(&name) .first_raw() .unwrap_or("application/octet-stream"); - if !spec.accepts.is_empty() + let has_mime_constraint = spec.accepts.iter().any(|accept| !accept.starts_with('.')); + if has_mime_constraint && !spec .accepts .iter() + .filter(|accept| !accept.starts_with('.')) .any(|accept| mime_matches(accept, mime_type)) { return Err(format!( @@ -347,35 +379,54 @@ async fn rewrite_single_file( )); } let manager = sess.services.mcp_connection_manager.load_full(); - let capabilities = manager - .file_capabilities(server) - .await - .map_err(|error| format!("failed to inspect MCP file capabilities: {error:#}"))?; - if spec.accepts_upload() && capabilities.prepare_upload && capabilities.complete_upload { + if spec.accepts_upload() { tracing::debug!( transfer_mode = "upload", size_bucket = file_size_bucket(size), "adapting MCP file input" ); - let prepared = manager - .prepare_file_upload( + let digest = FileDigest { + algorithm: "sha-256".to_string(), + value: URL_SAFE_NO_PAD.encode(Sha256::digest(&bytes)), + }; + let authorized = manager + .authorize_file_upload( server, - PrepareUploadParams { - name, + AuthorizeUploadParams { + name: name.clone(), mime_type: mime_type.to_string(), size, + digest: Some(digest.clone()), }, ) .await - .map_err(|error| format!("failed to prepare MCP file upload: {error:#}"))?; - validate_mcp_file_uri(&prepared.file.uri)?; - put_transfer_file(sess, &prepared.transfer, bytes, max_size).await?; - let completed = manager - .complete_file_upload(server, prepared.file.uri) - .await - .map_err(|error| format!("failed to complete MCP file upload: {error:#}"))?; - validate_mcp_file_uri(&completed.file.uri)?; - return Ok(JsonValue::String(completed.file.uri)); + .map_err(|error| format!("failed to authorize MCP file upload: {error:#}"))?; + validate_mcp_file_uri(&authorized.file.uri)?; + if authorized + .file + .size + .is_some_and(|authorized_size| authorized_size != size) + { + return Err("MCP upload authorization returned a different file size".to_string()); + } + if authorized + .file + .digest + .as_ref() + .is_some_and(|authorized_digest| authorized_digest != &digest) + { + return Err("MCP upload authorization returned a different file digest".to_string()); + } + put_transfer_file( + sess, + &authorized.upload, + bytes, + max_size, + authorized.file.name.as_deref().unwrap_or(&name), + authorized.file.mime_type.as_deref().unwrap_or(mime_type), + ) + .await?; + return Ok(JsonValue::String(authorized.file.uri)); } if spec.accepts_inline() { tracing::debug!( @@ -404,13 +455,17 @@ fn model_file_ref(value: &JsonValue) -> Option<&str> { } fn validate_mcp_file_uri(uri: &str) -> Result<(), String> { - if uri.starts_with("mcp-file://") { + if is_file_transfer_uri(uri) { Ok(()) } else { Err("MCP file response returned an invalid file URI".to_string()) } } +fn is_file_transfer_uri(uri: &str) -> bool { + Url::parse(uri).is_ok_and(|uri| !matches!(uri.scheme(), "data" | "file" | "http" | "https")) +} + fn mime_matches(accept: &str, mime_type: &str) -> bool { accept == "*/*" || accept == mime_type diff --git a/codex-rs/core/src/mcp_file_transfer/http.rs b/codex-rs/core/src/mcp_file_transfer/http.rs index c255786ac6..4e73f86b70 100644 --- a/codex-rs/core/src/mcp_file_transfer/http.rs +++ b/codex-rs/core/src/mcp_file_transfer/http.rs @@ -1,15 +1,23 @@ use std::net::IpAddr; use std::time::Duration; +use base64::Engine; +use base64::engine::general_purpose::URL_SAFE_NO_PAD; use codex_network_proxy::NetworkProxy; use futures::StreamExt; use reqwest::Method; +use reqwest::RequestBuilder; use reqwest::header::ACCEPT; use reqwest::header::CONTENT_LENGTH; use reqwest::header::HeaderMap; +use reqwest::header::HeaderName; use reqwest::header::HeaderValue; use reqwest::header::USER_AGENT; +use reqwest::multipart::Form; +use reqwest::multipart::Part; use reqwest::redirect::Policy; +use sha2::Digest; +use sha2::Sha256; use tokio::io::AsyncWriteExt; use url::Url; @@ -23,11 +31,18 @@ pub(super) async fn download_transfer_file( transfer: &codex_mcp::FileTransferDescriptor, output_path: &std::path::Path, max_size: u64, + expected_size: Option, + expected_digest: Option<&codex_mcp::FileDigest>, ) -> Result { let url = validated_transfer_descriptor(transfer, "GET")?; - let response = transfer_client(sess, &url) - .await? - .get(url) + if expected_size.is_some_and(|size| size > max_size) { + return Err(format!("MCP download exceeds the {max_size}-byte limit")); + } + if expected_digest.is_some_and(|digest| digest.algorithm != "sha-256") { + return Err("MCP download uses an unsupported digest algorithm".to_string()); + } + let request = transfer_client(sess, &url).await?.get(url); + let response = apply_transfer_headers(request, transfer, /*multipart*/ false)? .send() .await .map_err(|_| "MCP download transfer request failed".to_string())?; @@ -47,6 +62,7 @@ pub(super) async fn download_transfer_file( .await .map_err(|error| format!("failed to create MCP download: {error}"))?; let mut size = 0_u64; + let mut hasher = expected_digest.map(|_| Sha256::new()); let mut stream = response.bytes_stream(); while let Some(chunk) = stream.next().await { let chunk = chunk.map_err(|error| format!("failed to read MCP download: {error}"))?; @@ -54,11 +70,22 @@ pub(super) async fn download_transfer_file( if size > max_size { return Err(format!("MCP download exceeds the {max_size}-byte limit")); } + if let Some(hasher) = hasher.as_mut() { + hasher.update(&chunk); + } output .write_all(&chunk) .await .map_err(|error| format!("failed to write MCP download: {error}"))?; } + if expected_size.is_some_and(|expected_size| size != expected_size) { + return Err("MCP download size did not match its file metadata".to_string()); + } + if let (Some(hasher), Some(expected_digest)) = (hasher, expected_digest) + && URL_SAFE_NO_PAD.encode(hasher.finalize()) != expected_digest.value + { + return Err("MCP download digest did not match its file metadata".to_string()); + } output .flush() .await @@ -81,6 +108,8 @@ pub(super) async fn put_transfer_file( transfer: &codex_mcp::FileTransferDescriptor, bytes: Vec, max_size: u64, + name: &str, + mime_type: &str, ) -> Result<(), String> { let url = validated_upload_transfer_descriptor(transfer)?; let method = Method::from_bytes(transfer.method.as_bytes()) @@ -89,16 +118,34 @@ pub(super) async fn put_transfer_file( if size > max_size { return Err(format!("MCP upload exceeds the {max_size}-byte limit")); } - let stream = futures::stream::once(async move { Ok::<_, std::io::Error>(bytes) }); let azure_blob_upload = url.host_str().is_some_and(|host| { host.ends_with(".blob.core.windows.net") || host.ends_with(".oaiusercontent.com") }); - let mut request = transfer_client(sess, &url) - .await? - .request(method, url) - .header(CONTENT_LENGTH, size) - .body(reqwest::Body::wrap_stream(stream)); - if azure_blob_upload { + let mut request = transfer_client(sess, &url).await?.request(method, url); + if let Some(multipart) = transfer.multipart.as_ref() { + if transfer.method != "POST" { + return Err("MCP multipart upload method must be POST".to_string()); + } + if multipart.file_field.is_empty() { + return Err("MCP multipart upload file field must not be empty".to_string()); + } + let part = Part::bytes(bytes) + .file_name(name.to_string()) + .mime_str(mime_type) + .map_err(|error| format!("invalid MCP upload MIME type: {error}"))?; + let mut form = Form::new(); + for (field, value) in &multipart.fields { + form = form.text(field.clone(), value.clone()); + } + request = request.multipart(form.part(multipart.file_field.clone(), part)); + } else { + let stream = futures::stream::once(async move { Ok::<_, std::io::Error>(bytes) }); + request = request + .header(CONTENT_LENGTH, size) + .body(reqwest::Body::wrap_stream(stream)); + } + request = apply_transfer_headers(request, transfer, transfer.multipart.is_some())?; + if azure_blob_upload && transfer.multipart.is_none() { request = request.header("x-ms-blob-type", "BlockBlob"); } let response = request @@ -112,15 +159,46 @@ pub(super) async fn put_transfer_file( Ok(()) } +fn apply_transfer_headers( + mut request: RequestBuilder, + transfer: &codex_mcp::FileTransferDescriptor, + multipart: bool, +) -> Result { + for (name, value) in &transfer.headers { + let name = HeaderName::from_bytes(name.as_bytes()) + .map_err(|error| format!("invalid MCP transfer header name: {error}"))?; + let lower_name = name.as_str(); + if matches!( + lower_name, + "connection" + | "content-length" + | "host" + | "proxy-authenticate" + | "proxy-authorization" + | "te" + | "trailer" + | "transfer-encoding" + | "upgrade" + ) { + return Err(format!( + "MCP transfer header `{lower_name}` is not permitted" + )); + } + if multipart && lower_name == "content-type" { + continue; + } + let value = HeaderValue::from_str(value) + .map_err(|error| format!("invalid MCP transfer header value: {error}"))?; + request = request.header(name, value); + } + Ok(request) +} + pub(super) fn validated_transfer_descriptor( transfer: &codex_mcp::FileTransferDescriptor, expected_method: &str, ) -> Result { - if transfer - .transport - .as_deref() - .is_some_and(|value| value != "https") - { + if transfer.transport != "https" { return Err("MCP transfer transport must be HTTPS".to_string()); } if transfer.method != expected_method { diff --git a/codex-rs/core/src/mcp_file_transfer_tests.rs b/codex-rs/core/src/mcp_file_transfer_tests.rs index 69394b7818..7ad55e23f9 100644 --- a/codex-rs/core/src/mcp_file_transfer_tests.rs +++ b/codex-rs/core/src/mcp_file_transfer_tests.rs @@ -2,6 +2,7 @@ use super::*; use codex_mcp::FileInputSource; use codex_mcp::FileTransferMode; use pretty_assertions::assert_eq; +use std::collections::BTreeMap; use std::collections::BTreeSet; use tempfile::tempdir; use wiremock::Mock; @@ -35,13 +36,17 @@ async fn transfer_rejects_non_https_remote_urls() { put_transfer_file( &session, &codex_mcp::FileTransferDescriptor { - transport: Some("https".to_string()), + transport: "https".to_string(), method: "PUT".to_string(), url: "http://example.com/upload".to_string(), + headers: Default::default(), + multipart: None, expires_at: None, }, Vec::new(), /*max_size*/ 0, + /*name*/ "file.txt", + /*mime_type*/ "text/plain", ) .await .expect_err("remote HTTP must be rejected"), @@ -58,7 +63,8 @@ fn output_file_detection_requires_a_structured_file_value() { "text": "mcp-file://server/not-a-file-value" }), &mut files, - ); + ) + .expect("collect output files"); assert_eq!(files.len(), 1); assert!(files.contains_key("mcp-file://server/file_1")); } @@ -96,9 +102,11 @@ fn sanitizes_download_filenames() { fn rejects_expired_transfer_descriptors() { let error = validated_transfer_descriptor( &codex_mcp::FileTransferDescriptor { - transport: Some("https".to_string()), + transport: "https".to_string(), method: "GET".to_string(), url: "https://example.com/file".to_string(), + headers: Default::default(), + multipart: None, expires_at: Some("2020-01-01T00:00:00Z".to_string()), }, "GET", @@ -118,12 +126,67 @@ fn matches_exact_and_wildcard_mime_types() { #[test] fn validates_opaque_mcp_file_uris() { assert_eq!(validate_mcp_file_uri("mcp-file://server/file_1"), Ok(())); + assert_eq!(validate_mcp_file_uri("artifact://server/file_1"), Ok(())); assert_eq!( validate_mcp_file_uri("https://example.com/signed?secret=value"), Err("MCP file response returned an invalid file URI".to_string()) ); } +#[tokio::test] +async fn multipart_upload_honors_authorized_fields_headers_and_file_part() { + let (session, _) = crate::session::tests::make_session_and_context().await; + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/upload")) + .and(header("x-upload-token", "authorized")) + .respond_with(ResponseTemplate::new(204)) + .expect(1) + .mount(&server) + .await; + put_transfer_file( + &session, + &codex_mcp::FileTransferDescriptor { + transport: "https".to_string(), + method: "POST".to_string(), + url: format!("{}/upload", server.uri()), + headers: BTreeMap::from([ + ( + "content-type".to_string(), + "multipart/form-data".to_string(), + ), + ("x-upload-token".to_string(), "authorized".to_string()), + ]), + multipart: Some(codex_mcp::FileTransferMultipart { + file_field: "payload".to_string(), + fields: BTreeMap::from([("policy".to_string(), "signed-policy".to_string())]), + }), + expires_at: None, + }, + b"multipart bytes".to_vec(), + /*max_size*/ 32, + /*name*/ "report.txt", + /*mime_type*/ "text/plain", + ) + .await + .expect("multipart upload succeeds"); + + let requests = server.received_requests().await.expect("received requests"); + let request = requests.first().expect("upload request"); + let content_type = request + .headers + .get("content-type") + .expect("content type") + .to_str() + .expect("valid content type"); + assert!(content_type.starts_with("multipart/form-data; boundary=")); + let body = String::from_utf8_lossy(&request.body); + assert!(body.contains("name=\"policy\"")); + assert!(body.contains("signed-policy")); + assert!(body.contains("name=\"payload\"; filename=\"report.txt\"")); + assert!(body.contains("multipart bytes")); +} + #[test] fn rejects_non_public_transfer_addresses() { for address in [ @@ -216,13 +279,17 @@ async fn upload_transfer_streams_the_exact_file_bytes() { put_transfer_file( &session, &codex_mcp::FileTransferDescriptor { - transport: Some("https".to_string()), + transport: "https".to_string(), method: "PUT".to_string(), url: format!("{}/upload", server.uri()), + headers: Default::default(), + multipart: None, expires_at: None, }, b"stream me".to_vec(), /*max_size*/ 32, + /*name*/ "file.txt", + /*mime_type*/ "text/plain", ) .await .expect("upload succeeds"); @@ -243,13 +310,17 @@ async fn upload_transfer_accepts_post_descriptors() { put_transfer_file( &session, &codex_mcp::FileTransferDescriptor { - transport: None, + transport: "https".to_string(), method: "POST".to_string(), url: format!("{}/upload", server.uri()), + headers: Default::default(), + multipart: None, expires_at: None, }, b"stream me".to_vec(), /*max_size*/ 32, + /*name*/ "file.txt", + /*mime_type*/ "text/plain", ) .await .expect("upload succeeds"); @@ -271,13 +342,17 @@ async fn download_transfer_materializes_exact_bytes() { let size = download_transfer_file( &session, &codex_mcp::FileTransferDescriptor { - transport: Some("https".to_string()), + transport: "https".to_string(), method: "GET".to_string(), url: format!("{}/download", server.uri()), + headers: Default::default(), + multipart: None, expires_at: None, }, &output, /*max_size*/ 32, + /*expected_size*/ None, + /*expected_digest*/ None, ) .await .expect("download succeeds"); @@ -289,6 +364,45 @@ async fn download_transfer_materializes_exact_bytes() { ); } +#[tokio::test] +async fn download_rejects_digest_mismatch_and_removes_partial_file() { + let (session, _) = crate::session::tests::make_session_and_context().await; + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/download")) + .and(header("authorization", "Bearer signed")) + .respond_with(ResponseTemplate::new(200).set_body_bytes(b"download me")) + .mount(&server) + .await; + let directory = tempdir().expect("temp dir"); + let output = directory.path().join("download.txt"); + let error = download_transfer_file( + &session, + &codex_mcp::FileTransferDescriptor { + transport: "https".to_string(), + method: "GET".to_string(), + url: format!("{}/download", server.uri()), + headers: BTreeMap::from([("authorization".to_string(), "Bearer signed".to_string())]), + multipart: None, + expires_at: None, + }, + &output, + /*max_size*/ 32, + /*expected_size*/ Some(11), + /*expected_digest*/ + Some(&codex_mcp::FileDigest { + algorithm: "sha-256".to_string(), + value: "invalid".to_string(), + }), + ) + .await + .expect_err("digest mismatch"); + + assert_eq!(error, "MCP download digest did not match its file metadata"); + assert!(!output.exists()); + assert!(!output.with_extension("part").exists()); +} + #[tokio::test] async fn transfer_errors_do_not_expose_signed_urls() { let (session, _) = crate::session::tests::make_session_and_context().await; @@ -305,13 +419,17 @@ async fn transfer_errors_do_not_expose_signed_urls() { let error = download_transfer_file( &session, &codex_mcp::FileTransferDescriptor { - transport: None, + transport: "https".to_string(), method: "GET".to_string(), url: format!("{}/download?sig={secret}", server.uri()), + headers: Default::default(), + multipart: None, expires_at: None, }, &output, /*max_size*/ 32, + /*expected_size*/ None, + /*expected_digest*/ None, ) .await .expect_err("failed transfer"); @@ -338,13 +456,17 @@ async fn failed_download_removes_partial_file() { let error = download_transfer_file( &session, &codex_mcp::FileTransferDescriptor { - transport: None, + transport: "https".to_string(), method: "GET".to_string(), url: format!("{}/download", server.uri()), + headers: Default::default(), + multipart: None, expires_at: None, }, &output, /*max_size*/ 32, + /*expected_size*/ None, + /*expected_digest*/ None, ) .await .expect_err("oversized transfer");