diff --git a/codex-rs/memories/mcp/src/schema.rs b/codex-rs/memories/mcp/src/schema.rs index 5adad642a..3e212cf29 100644 --- a/codex-rs/memories/mcp/src/schema.rs +++ b/codex-rs/memories/mcp/src/schema.rs @@ -71,40 +71,22 @@ pub(crate) fn read_output_schema() -> JsonObject { pub(crate) fn search_input_schema() -> JsonObject { json_schema(json!({ - "anyOf": [ - { - "type": "object", - "properties": { - "query": { "type": "string" }, - "match_mode": { "type": "string", "enum": ["any", "all"] }, - "path": { "type": "string" }, - "cursor": { "type": "string" }, - "context_lines": { "type": "integer", "minimum": 0 }, - "case_sensitive": { "type": "boolean" }, - "max_results": { "type": "integer", "minimum": 1 } - }, - "required": ["query"], - "additionalProperties": false + "type": "object", + "properties": { + "queries": { + "type": "array", + "items": { "type": "string" }, + "minItems": 1 }, - { - "type": "object", - "properties": { - "queries": { - "type": "array", - "items": { "type": "string" }, - "minItems": 1 - }, - "match_mode": { "type": "string", "enum": ["any", "all"] }, - "path": { "type": "string" }, - "cursor": { "type": "string" }, - "context_lines": { "type": "integer", "minimum": 0 }, - "case_sensitive": { "type": "boolean" }, - "max_results": { "type": "integer", "minimum": 1 } - }, - "required": ["queries"], - "additionalProperties": false - } - ] + "match_mode": { "type": "string", "enum": ["any", "all"] }, + "path": { "type": "string" }, + "cursor": { "type": "string" }, + "context_lines": { "type": "integer", "minimum": 0 }, + "case_sensitive": { "type": "boolean" }, + "max_results": { "type": "integer", "minimum": 1 } + }, + "required": ["queries"], + "additionalProperties": false })) } diff --git a/codex-rs/memories/mcp/src/server.rs b/codex-rs/memories/mcp/src/server.rs index b9e6d1f76..2633d9448 100644 --- a/codex-rs/memories/mcp/src/server.rs +++ b/codex-rs/memories/mcp/src/server.rs @@ -54,10 +54,10 @@ struct ReadArgs { max_lines: Option, } -#[derive(Deserialize)] +#[derive(Debug, Deserialize)] +#[serde(deny_unknown_fields)] struct SearchArgs { - query: Option, - queries: Option>, + queries: Vec, match_mode: Option, path: Option, cursor: Option, @@ -147,7 +147,7 @@ impl ServerHandler for MemoriesMcpServer { } SEARCH_TOOL_NAME => { let args: SearchArgs = parse_args(value)?; - let request = args.into_request()?; + let request = args.into_request(); json!( self.backend .search(request) @@ -229,21 +229,9 @@ fn parse_args Deserialize<'de>>(value: serde_json::Value) -> Result< } impl SearchArgs { - fn into_request(self) -> Result { - let queries = match (self.query, self.queries) { - (Some(query), None) => Ok(vec![query]), - (None, Some(queries)) => Ok(queries), - (Some(_), Some(_)) => Err(McpError::invalid_params( - "provide either 'query' or 'queries', but not both".to_string(), - None, - )), - (None, None) => Err(McpError::invalid_params( - "missing required field: 'query' or 'queries'".to_string(), - None, - )), - }?; - Ok(SearchMemoriesRequest { - queries, + fn into_request(self) -> SearchMemoriesRequest { + SearchMemoriesRequest { + queries: self.queries, match_mode: self.match_mode.unwrap_or(SearchMatchMode::Any), path: self.path, cursor: self.cursor, @@ -254,7 +242,7 @@ impl SearchArgs { DEFAULT_SEARCH_MAX_RESULTS, MAX_SEARCH_RESULTS, ), - }) + } } } @@ -281,30 +269,6 @@ mod tests { use pretty_assertions::assert_eq; use serde_json::json; - #[test] - fn search_args_accept_legacy_single_query() { - let args: SearchArgs = parse_args(json!({ - "query": "needle", - "match_mode": "all" - })) - .expect("legacy query args should parse"); - - let request = args.into_request().expect("query should convert"); - - assert_eq!( - request, - SearchMemoriesRequest { - queries: vec!["needle".to_string()], - match_mode: SearchMatchMode::All, - path: None, - cursor: None, - context_lines: 0, - case_sensitive: true, - max_results: DEFAULT_SEARCH_MAX_RESULTS, - } - ); - } - #[test] fn search_args_accept_multiple_queries() { let args: SearchArgs = parse_args(json!({ @@ -313,7 +277,7 @@ mod tests { })) .expect("multi-query args should parse"); - let request = args.into_request().expect("queries should convert"); + let request = args.into_request(); assert_eq!( request, @@ -330,17 +294,25 @@ mod tests { } #[test] - fn search_args_reject_both_query_forms() { - let args: SearchArgs = parse_args(json!({ + fn search_args_reject_legacy_single_query() { + let err = parse_args::(json!({ "query": "needle", - "queries": ["needle"] })) - .expect("args should parse before conversion"); + .expect_err("legacy query field should be rejected"); - let err = args - .into_request() - .expect_err("query and queries should be mutually exclusive"); + assert!(err.message.contains("unknown field")); + assert!(err.message.contains("query")); + } - assert!(err.message.contains("either 'query' or 'queries'")); + #[test] + fn search_args_reject_unknown_fields() { + let err = parse_args::(json!({ + "queries": ["needle"], + "query": "needle" + })) + .expect_err("unknown fields should be rejected"); + + assert!(err.message.contains("unknown field")); + assert!(err.message.contains("query")); } }