1use std::collections::HashMap;
24use std::convert::Infallible;
25use std::future::Future;
26use std::pin::Pin;
27use std::sync::Arc;
28use std::task::{Context, Poll};
29
30use tower::{Layer, Service};
31use tower_mcp::router::{RouterRequest, RouterResponse};
32use tower_mcp_types::protocol::{McpRequest, McpResponse};
33
34#[derive(Clone)]
47pub struct ParamOverrideLayer {
48 overrides: Vec<ToolOverride>,
49}
50
51impl ParamOverrideLayer {
52 pub fn new(overrides: Vec<ToolOverride>) -> Self {
54 Self { overrides }
55 }
56}
57
58impl<S> Layer<S> for ParamOverrideLayer {
59 type Service = ParamOverrideService<S>;
60
61 fn layer(&self, inner: S) -> Self::Service {
62 ParamOverrideService::new(inner, self.overrides.clone())
63 }
64}
65
66#[derive(Debug, Clone)]
68pub struct ToolOverride {
69 namespaced_tool: String,
71 hide: Vec<String>,
73 defaults: serde_json::Map<String, serde_json::Value>,
75 rename_forward: HashMap<String, String>,
77 rename_reverse: HashMap<String, String>,
79}
80
81impl ToolOverride {
82 pub fn new(namespace: &str, config: &crate::config::ParamOverrideConfig) -> Self {
84 let rename_forward: HashMap<String, String> = config.rename.clone();
85 let rename_reverse: HashMap<String, String> = config
86 .rename
87 .iter()
88 .map(|(orig, new)| (new.clone(), orig.clone()))
89 .collect();
90
91 Self {
92 namespaced_tool: format!("{namespace}{}", config.tool),
93 hide: config.hide.clone(),
94 defaults: config.defaults.clone(),
95 rename_forward,
96 rename_reverse,
97 }
98 }
99}
100
101#[derive(Clone)]
107pub struct ParamOverrideService<S> {
108 inner: S,
109 overrides: Arc<Vec<ToolOverride>>,
110}
111
112impl<S> ParamOverrideService<S> {
113 pub fn new(inner: S, overrides: Vec<ToolOverride>) -> Self {
115 Self {
116 inner,
117 overrides: Arc::new(overrides),
118 }
119 }
120}
121
122fn rewrite_schema(
124 schema: &mut serde_json::Value,
125 hide: &[String],
126 rename_forward: &HashMap<String, String>,
127) {
128 let Some(obj) = schema.as_object_mut() else {
129 return;
130 };
131
132 if let Some(props) = obj.get_mut("properties").and_then(|v| v.as_object_mut()) {
134 for param in hide {
135 props.remove(param);
136 }
137 for (original, renamed) in rename_forward {
138 if let Some(prop_schema) = props.remove(original) {
139 props.insert(renamed.clone(), prop_schema);
140 }
141 }
142 }
143
144 if let Some(required) = obj.get_mut("required").and_then(|v| v.as_array_mut()) {
146 required.retain(|v| {
147 v.as_str()
148 .map(|s| !hide.contains(&s.to_string()))
149 .unwrap_or(true)
150 });
151 for entry in required.iter_mut() {
152 if let Some(s) = entry.as_str()
153 && let Some(new_name) = rename_forward.get(s)
154 {
155 *entry = serde_json::Value::String(new_name.clone());
156 }
157 }
158 }
159}
160
161impl<S> Service<RouterRequest> for ParamOverrideService<S>
162where
163 S: Service<RouterRequest, Response = RouterResponse, Error = Infallible>
164 + Clone
165 + Send
166 + 'static,
167 S::Future: Send,
168{
169 type Response = RouterResponse;
170 type Error = Infallible;
171 type Future = Pin<Box<dyn Future<Output = Result<RouterResponse, Infallible>> + Send>>;
172
173 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
174 self.inner.poll_ready(cx)
175 }
176
177 fn call(&mut self, mut req: RouterRequest) -> Self::Future {
178 let overrides = Arc::clone(&self.overrides);
179
180 if let McpRequest::CallTool(ref mut params) = req.inner {
182 for tool_override in overrides.iter() {
183 if params.name != tool_override.namespaced_tool {
184 continue;
185 }
186
187 if let serde_json::Value::Object(ref mut args) = params.arguments {
189 for (key, value) in &tool_override.defaults {
190 if !args.contains_key(key) {
191 args.insert(key.clone(), value.clone());
192 }
193 }
194
195 let keys_to_rename: Vec<(String, String)> = args
197 .keys()
198 .filter_map(|k| {
199 tool_override
200 .rename_reverse
201 .get(k)
202 .map(|orig| (k.clone(), orig.clone()))
203 })
204 .collect();
205
206 for (new_name, original_name) in keys_to_rename {
207 if let Some(value) = args.remove(&new_name) {
208 args.insert(original_name, value);
209 }
210 }
211 }
212
213 break;
214 }
215 }
216
217 let fut = self.inner.call(req);
218
219 Box::pin(async move {
220 let mut resp = fut.await?;
221
222 if let Ok(McpResponse::ListTools(ref mut result)) = resp.inner {
224 for tool in &mut result.tools {
225 for tool_override in overrides.iter() {
226 if tool.name == tool_override.namespaced_tool {
227 rewrite_schema(
228 &mut tool.input_schema,
229 &tool_override.hide,
230 &tool_override.rename_forward,
231 );
232 break;
233 }
234 }
235 }
236 }
237
238 Ok(resp)
239 })
240 }
241}
242
243#[cfg(test)]
244mod tests {
245 use super::*;
246 use crate::config::ParamOverrideConfig;
247 use crate::test_util::{MockService, call_service};
248 use tower_mcp_types::protocol::{CallToolParams, McpRequest, McpResponse};
249
250 fn mock_with_schema(name: &str, schema: serde_json::Value) -> MockService {
252 use tower_mcp_types::protocol::ToolDefinition;
253 MockService {
254 tools: vec![ToolDefinition {
255 name: name.to_string(),
256 title: None,
257 description: Some(format!("{name} tool")),
258 input_schema: schema,
259 output_schema: None,
260 icons: None,
261 annotations: None,
262 execution: None,
263 meta: None,
264 }],
265 }
266 }
267
268 fn list_dir_schema() -> serde_json::Value {
269 serde_json::json!({
270 "type": "object",
271 "properties": {
272 "path": { "type": "string" },
273 "recursive": { "type": "boolean" },
274 "pattern": { "type": "string" }
275 },
276 "required": ["path"]
277 })
278 }
279
280 fn make_overrides(namespace: &str, configs: Vec<ParamOverrideConfig>) -> Vec<ToolOverride> {
281 configs
282 .iter()
283 .map(|c| ToolOverride::new(namespace, c))
284 .collect()
285 }
286
287 #[tokio::test]
288 async fn test_hide_removes_param_from_schema() {
289 let mock = mock_with_schema("fs/list_directory", list_dir_schema());
290 let overrides = make_overrides(
291 "fs/",
292 vec![ParamOverrideConfig {
293 tool: "list_directory".to_string(),
294 hide: vec!["path".to_string()],
295 defaults: {
296 let mut m = serde_json::Map::new();
297 m.insert("path".to_string(), serde_json::json!("/home/docs"));
298 m
299 },
300 rename: HashMap::new(),
301 }],
302 );
303 let mut svc = ParamOverrideService::new(mock, overrides);
304
305 let resp = call_service(&mut svc, McpRequest::ListTools(Default::default())).await;
306 match resp.inner.unwrap() {
307 McpResponse::ListTools(result) => {
308 let tool = &result.tools[0];
309 let props = tool.input_schema["properties"].as_object().unwrap();
310 assert!(
311 !props.contains_key("path"),
312 "path should be hidden from schema"
313 );
314 assert!(props.contains_key("recursive"), "recursive should remain");
315 assert!(props.contains_key("pattern"), "pattern should remain");
316 let required = tool.input_schema["required"].as_array().unwrap();
318 let req_strs: Vec<&str> = required.iter().map(|v| v.as_str().unwrap()).collect();
319 assert!(!req_strs.contains(&"path"), "path should not be required");
320 }
321 other => panic!("expected ListTools, got: {:?}", other),
322 }
323 }
324
325 #[tokio::test]
326 async fn test_hide_injects_defaults_on_call() {
327 let mock = mock_with_schema("fs/list_directory", list_dir_schema());
328 let overrides = make_overrides(
329 "fs/",
330 vec![ParamOverrideConfig {
331 tool: "list_directory".to_string(),
332 hide: vec!["path".to_string()],
333 defaults: {
334 let mut m = serde_json::Map::new();
335 m.insert("path".to_string(), serde_json::json!("/home/docs"));
336 m
337 },
338 rename: HashMap::new(),
339 }],
340 );
341 let mut svc = ParamOverrideService::new(mock, overrides);
342
343 let resp = call_service(
344 &mut svc,
345 McpRequest::CallTool(CallToolParams {
346 name: "fs/list_directory".to_string(),
347 arguments: serde_json::json!({"recursive": true}),
348 input_responses: None,
349 request_state: None,
350 meta: None,
351 task: None,
352 }),
353 )
354 .await;
355
356 assert!(resp.inner.is_ok(), "call should succeed");
357 }
358
359 #[tokio::test]
360 async fn test_rename_rewrites_schema() {
361 let mock = mock_with_schema("fs/list_directory", list_dir_schema());
362 let overrides = make_overrides(
363 "fs/",
364 vec![ParamOverrideConfig {
365 tool: "list_directory".to_string(),
366 hide: vec![],
367 defaults: serde_json::Map::new(),
368 rename: {
369 let mut m = HashMap::new();
370 m.insert("recursive".to_string(), "deep_search".to_string());
371 m
372 },
373 }],
374 );
375 let mut svc = ParamOverrideService::new(mock, overrides);
376
377 let resp = call_service(&mut svc, McpRequest::ListTools(Default::default())).await;
378 match resp.inner.unwrap() {
379 McpResponse::ListTools(result) => {
380 let tool = &result.tools[0];
381 let props = tool.input_schema["properties"].as_object().unwrap();
382 assert!(
383 !props.contains_key("recursive"),
384 "recursive should be renamed"
385 );
386 assert!(
387 props.contains_key("deep_search"),
388 "deep_search should appear"
389 );
390 assert!(props.contains_key("path"), "path should remain");
391 }
392 other => panic!("expected ListTools, got: {:?}", other),
393 }
394 }
395
396 #[tokio::test]
397 async fn test_rename_reverse_maps_on_call() {
398 let mock = mock_with_schema("fs/list_directory", list_dir_schema());
399 let overrides = make_overrides(
400 "fs/",
401 vec![ParamOverrideConfig {
402 tool: "list_directory".to_string(),
403 hide: vec![],
404 defaults: serde_json::Map::new(),
405 rename: {
406 let mut m = HashMap::new();
407 m.insert("recursive".to_string(), "deep_search".to_string());
408 m
409 },
410 }],
411 );
412 let mut svc = ParamOverrideService::new(mock, overrides);
413
414 let resp = call_service(
416 &mut svc,
417 McpRequest::CallTool(CallToolParams {
418 name: "fs/list_directory".to_string(),
419 arguments: serde_json::json!({"path": "/tmp", "deep_search": true}),
420 input_responses: None,
421 request_state: None,
422 meta: None,
423 task: None,
424 }),
425 )
426 .await;
427
428 assert!(resp.inner.is_ok(), "call should succeed");
429 }
430
431 #[tokio::test]
432 async fn test_hide_and_rename_combined() {
433 let mock = mock_with_schema("fs/list_directory", list_dir_schema());
434 let overrides = make_overrides(
435 "fs/",
436 vec![ParamOverrideConfig {
437 tool: "list_directory".to_string(),
438 hide: vec!["path".to_string()],
439 defaults: {
440 let mut m = serde_json::Map::new();
441 m.insert("path".to_string(), serde_json::json!("/home/docs"));
442 m
443 },
444 rename: {
445 let mut m = HashMap::new();
446 m.insert("recursive".to_string(), "deep_search".to_string());
447 m
448 },
449 }],
450 );
451 let mut svc = ParamOverrideService::new(mock, overrides);
452
453 let resp = call_service(&mut svc, McpRequest::ListTools(Default::default())).await;
455 match resp.inner.unwrap() {
456 McpResponse::ListTools(result) => {
457 let props = result.tools[0].input_schema["properties"]
458 .as_object()
459 .unwrap();
460 assert!(!props.contains_key("path"));
461 assert!(!props.contains_key("recursive"));
462 assert!(props.contains_key("deep_search"));
463 assert!(props.contains_key("pattern"));
464 }
465 other => panic!("expected ListTools, got: {:?}", other),
466 }
467 }
468
469 #[tokio::test]
470 async fn test_non_matching_tool_passes_through() {
471 let mock = mock_with_schema("db/query", list_dir_schema());
472 let overrides = make_overrides(
473 "fs/",
474 vec![ParamOverrideConfig {
475 tool: "list_directory".to_string(),
476 hide: vec!["path".to_string()],
477 defaults: serde_json::Map::new(),
478 rename: HashMap::new(),
479 }],
480 );
481 let mut svc = ParamOverrideService::new(mock, overrides);
482
483 let resp = call_service(&mut svc, McpRequest::ListTools(Default::default())).await;
484 match resp.inner.unwrap() {
485 McpResponse::ListTools(result) => {
486 let props = result.tools[0].input_schema["properties"]
488 .as_object()
489 .unwrap();
490 assert!(props.contains_key("path"), "unmatched tool is untouched");
491 }
492 other => panic!("expected ListTools, got: {:?}", other),
493 }
494 }
495
496 #[tokio::test]
497 async fn test_non_call_tool_passes_through() {
498 let mock = MockService::with_tools(&["fs/list_directory"]);
499 let overrides = make_overrides(
500 "fs/",
501 vec![ParamOverrideConfig {
502 tool: "list_directory".to_string(),
503 hide: vec!["path".to_string()],
504 defaults: serde_json::Map::new(),
505 rename: HashMap::new(),
506 }],
507 );
508 let mut svc = ParamOverrideService::new(mock, overrides);
509
510 let resp = call_service(&mut svc, McpRequest::ListTools(Default::default())).await;
512 assert!(resp.inner.is_ok());
513 }
514
515 #[tokio::test]
516 async fn test_rename_updates_required_array() {
517 let schema = serde_json::json!({
518 "type": "object",
519 "properties": {
520 "path": { "type": "string" },
521 "recursive": { "type": "boolean" }
522 },
523 "required": ["path", "recursive"]
524 });
525 let mock = mock_with_schema("fs/list_directory", schema);
526 let overrides = make_overrides(
527 "fs/",
528 vec![ParamOverrideConfig {
529 tool: "list_directory".to_string(),
530 hide: vec![],
531 defaults: serde_json::Map::new(),
532 rename: {
533 let mut m = HashMap::new();
534 m.insert("recursive".to_string(), "deep_search".to_string());
535 m
536 },
537 }],
538 );
539 let mut svc = ParamOverrideService::new(mock, overrides);
540
541 let resp = call_service(&mut svc, McpRequest::ListTools(Default::default())).await;
542 match resp.inner.unwrap() {
543 McpResponse::ListTools(result) => {
544 let required = result.tools[0].input_schema["required"].as_array().unwrap();
545 let req_strs: Vec<&str> = required.iter().map(|v| v.as_str().unwrap()).collect();
546 assert!(req_strs.contains(&"path"));
547 assert!(req_strs.contains(&"deep_search"));
548 assert!(!req_strs.contains(&"recursive"));
549 }
550 other => panic!("expected ListTools, got: {:?}", other),
551 }
552 }
553
554 #[test]
555 fn test_rewrite_schema_no_properties() {
556 let mut schema = serde_json::json!({"type": "object"});
558 rewrite_schema(&mut schema, &["path".to_string()], &HashMap::new());
559 assert_eq!(schema, serde_json::json!({"type": "object"}));
560 }
561
562 #[test]
563 fn test_rewrite_schema_non_object() {
564 let mut schema = serde_json::json!("string");
566 rewrite_schema(&mut schema, &["path".to_string()], &HashMap::new());
567 assert_eq!(schema, serde_json::json!("string"));
568 }
569
570 #[test]
571 fn test_tool_override_construction() {
572 let config = ParamOverrideConfig {
573 tool: "list_directory".to_string(),
574 hide: vec!["path".to_string()],
575 defaults: {
576 let mut m = serde_json::Map::new();
577 m.insert("path".to_string(), serde_json::json!("/home"));
578 m
579 },
580 rename: {
581 let mut m = HashMap::new();
582 m.insert("recursive".to_string(), "deep_search".to_string());
583 m
584 },
585 };
586 let to = ToolOverride::new("fs/", &config);
587 assert_eq!(to.namespaced_tool, "fs/list_directory");
588 assert_eq!(to.hide, vec!["path"]);
589 assert_eq!(to.rename_forward.get("recursive").unwrap(), "deep_search");
590 assert_eq!(to.rename_reverse.get("deep_search").unwrap(), "recursive");
591 }
592
593 #[tokio::test]
594 async fn test_hidden_default_does_not_overwrite_explicit_arg() {
595 let _mock = mock_with_schema("fs/list_directory", list_dir_schema());
596 let overrides = make_overrides(
597 "fs/",
598 vec![ParamOverrideConfig {
599 tool: "list_directory".to_string(),
600 hide: vec!["path".to_string()],
601 defaults: {
602 let mut m = serde_json::Map::new();
603 m.insert("path".to_string(), serde_json::json!("/home/docs"));
604 m
605 },
606 rename: HashMap::new(),
607 }],
608 );
609
610 let mut req = RouterRequest {
613 id: tower_mcp::protocol::RequestId::Number(1),
614 inner: McpRequest::CallTool(CallToolParams {
615 name: "fs/list_directory".to_string(),
616 arguments: serde_json::json!({"path": "/custom"}),
617 input_responses: None,
618 request_state: None,
619 meta: None,
620 task: None,
621 }),
622 extensions: tower_mcp::router::Extensions::new(),
623 };
624
625 if let McpRequest::CallTool(ref mut params) = req.inner
627 && let serde_json::Value::Object(ref mut args) = params.arguments
628 {
629 let defaults = &overrides[0].defaults;
630 for (key, value) in defaults {
631 if !args.contains_key(key) {
632 args.insert(key.clone(), value.clone());
633 }
634 }
635 assert_eq!(args.get("path").unwrap(), "/custom");
637 }
638 }
639}