1use std::sync::Arc;
8
9use schemars::JsonSchema;
10use serde::{Deserialize, Serialize};
11use tower_mcp::client::ChannelTransport;
12use tower_mcp::proxy::{AddBackendError, McpProxy};
13use tower_mcp::{CallToolResult, McpRouter, NoParams, SessionHandle, ToolBuilder};
14
15use crate::admin::AdminState;
16use crate::config::ProxyConfig;
17
18#[derive(Clone)]
20struct AdminToolState {
21 admin_state: AdminState,
22 session_handle: SessionHandle,
23 config_snapshot: Arc<String>,
24 proxy: McpProxy,
25}
26
27#[derive(Serialize)]
28struct BackendInfo {
29 namespace: String,
30 healthy: bool,
31 #[serde(skip_serializing_if = "Option::is_none")]
32 last_checked_at: Option<String>,
33 consecutive_failures: u32,
34 #[serde(skip_serializing_if = "Option::is_none")]
35 error: Option<String>,
36 #[serde(skip_serializing_if = "Option::is_none")]
37 transport: Option<String>,
38}
39
40#[derive(Serialize)]
41struct BackendsResult {
42 proxy_name: String,
43 proxy_version: String,
44 backend_count: usize,
45 backends: Vec<BackendInfo>,
46}
47
48#[derive(Serialize)]
49struct SessionResult {
50 active_sessions: usize,
51}
52
53pub async fn register_admin_tools(
63 proxy: &McpProxy,
64 admin_state: AdminState,
65 session_handle: SessionHandle,
66 config: &ProxyConfig,
67 discovery_tools: Option<Vec<tower_mcp::Tool>>,
68) -> Result<(), AddBackendError> {
69 let config_toml =
70 toml::to_string_pretty(config).unwrap_or_else(|e| format!("error serializing: {e}"));
71
72 let search_mode = config.proxy.tool_exposure == crate::config::ToolExposure::Search;
73
74 let state = AdminToolState {
75 admin_state,
76 session_handle,
77 config_snapshot: Arc::new(config_toml),
78 proxy: proxy.clone(),
79 };
80
81 #[cfg(feature = "skills")]
83 let skills = crate::skills::build_skills(state.config_snapshot.clone());
84 #[cfg(not(feature = "skills"))]
85 let skills: Vec<tower_mcp::Prompt> = vec![];
86
87 let router = build_admin_router(state, discovery_tools, search_mode, skills);
88 let transport = ChannelTransport::new(router);
89
90 proxy.add_backend("proxy", transport).await
91}
92
93fn build_admin_router(
94 state: AdminToolState,
95 discovery_tools: Option<Vec<tower_mcp::Tool>>,
96 search_mode: bool,
97 skills: Vec<tower_mcp::Prompt>,
98) -> McpRouter {
99 let state_for_backends = state.clone();
100 let list_backends = ToolBuilder::new("list_backends")
101 .description("List all proxy backends with health status")
102 .handler(move |_: NoParams| {
103 let s = state_for_backends.clone();
104 async move {
105 let health = s.admin_state.health().await;
106 let backends: Vec<BackendInfo> = health
107 .iter()
108 .map(|b| BackendInfo {
109 namespace: b.namespace.clone(),
110 healthy: b.healthy,
111 last_checked_at: b.last_checked_at.map(|t| t.to_rfc3339()),
112 consecutive_failures: b.consecutive_failures,
113 error: b.error.clone(),
114 transport: b.transport.clone(),
115 })
116 .collect();
117
118 let result = BackendsResult {
119 proxy_name: s.admin_state.proxy_name().to_string(),
120 proxy_version: s.admin_state.proxy_version().to_string(),
121 backend_count: s.admin_state.backend_count(),
122 backends,
123 };
124
125 Ok(CallToolResult::text(
126 serde_json::to_string_pretty(&result).unwrap(),
127 ))
128 }
129 })
130 .build();
131
132 let state_for_sessions = state.clone();
133 let session_count = ToolBuilder::new("session_count")
134 .description("Get the number of active MCP sessions")
135 .handler(move |_: NoParams| {
136 let s = state_for_sessions.clone();
137 async move {
138 let count = s.session_handle.session_count().await;
139 let result = SessionResult {
140 active_sessions: count,
141 };
142 Ok(CallToolResult::text(
143 serde_json::to_string_pretty(&result).unwrap(),
144 ))
145 }
146 })
147 .build();
148
149 let config_snapshot = Arc::clone(&state.config_snapshot);
150 let config_tool = ToolBuilder::new("config")
151 .description("Dump the current proxy configuration")
152 .handler(move |_: NoParams| {
153 let config = Arc::clone(&config_snapshot);
154 async move { Ok(CallToolResult::text((*config).clone())) }
155 })
156 .build();
157
158 let state_for_health = state.clone();
159 let health_check = ToolBuilder::new("health_check")
160 .description("Get cached health check results for all backends")
161 .handler(move |_: NoParams| {
162 let s = state_for_health.clone();
163 async move {
164 let health = s.admin_state.health().await;
165 let backends: Vec<BackendInfo> = health
166 .iter()
167 .map(|b| BackendInfo {
168 namespace: b.namespace.clone(),
169 healthy: b.healthy,
170 last_checked_at: b.last_checked_at.map(|t| t.to_rfc3339()),
171 consecutive_failures: b.consecutive_failures,
172 error: b.error.clone(),
173 transport: b.transport.clone(),
174 })
175 .collect();
176 let healthy_count = backends.iter().filter(|b| b.healthy).count();
177 let total = backends.len();
178 let result = HealthCheckResult {
179 status: if healthy_count == total {
180 "healthy"
181 } else {
182 "degraded"
183 }
184 .to_string(),
185 healthy_count,
186 total_count: total,
187 backends,
188 };
189 Ok(CallToolResult::text(
190 serde_json::to_string_pretty(&result).unwrap(),
191 ))
192 }
193 })
194 .build();
195
196 let state_for_add = state.clone();
197 let add_backend = ToolBuilder::new("add_backend")
198 .description("Dynamically add an HTTP backend to the proxy")
199 .handler(move |input: AddBackendInput| {
200 let s = state_for_add.clone();
201 async move {
202 let transport = tower_mcp::client::HttpClientTransport::new(&input.url);
203 match s.proxy.add_backend(&input.name, transport).await {
204 Ok(()) => Ok(CallToolResult::text(format!(
205 "Backend '{}' added successfully at {}",
206 input.name, input.url
207 ))),
208 Err(e) => Ok(CallToolResult::text(format!(
209 "Failed to add backend '{}': {e}",
210 input.name
211 ))),
212 }
213 }
214 })
215 .build();
216
217 let mut router = McpRouter::new()
218 .server_info("mcp-proxy-admin", "0.1.0")
219 .tool(list_backends)
220 .tool(health_check)
221 .tool(session_count)
222 .tool(add_backend)
223 .tool(config_tool);
224
225 if search_mode {
226 let state_for_call = state.clone();
227 let call_tool = ToolBuilder::new("call_tool")
228 .description(
229 "Invoke any backend tool by its fully-qualified name. Use proxy/search_tools \
230 to discover available tools, then call them through this tool.",
231 )
232 .handler(move |input: CallToolInput| {
233 let s = state_for_call.clone();
234 async move {
235 use tower::Service;
236 use tower_mcp::protocol::{CallToolParams, McpRequest, McpResponse, RequestId};
237 use tower_mcp::router::{Extensions, RouterRequest};
238
239 let req = RouterRequest {
240 id: RequestId::Number(0),
241 inner: McpRequest::CallTool(CallToolParams {
242 name: input.name.clone(),
243 arguments: input.arguments.unwrap_or_default().into(),
244 input_responses: None,
245 request_state: None,
246 meta: None,
247 task: None,
248 }),
249 extensions: Extensions::new(),
250 };
251
252 let mut proxy = s.proxy.clone();
253 match proxy.call(req).await {
254 Ok(resp) => match resp.inner {
255 Ok(McpResponse::CallTool(result)) => Ok(result),
256 Ok(_) => Ok(CallToolResult::text(format!(
257 "Unexpected response type for tool '{}'",
258 input.name
259 ))),
260 Err(e) => Ok(CallToolResult::text(format!(
261 "Error calling '{}': {}",
262 input.name, e.message
263 ))),
264 },
265 Err(_) => Ok(CallToolResult::text(format!(
266 "Internal error calling '{}'",
267 input.name
268 ))),
269 }
270 }
271 })
272 .build();
273 router = router.tool(call_tool);
274 }
275
276 if let Some(tools) = discovery_tools {
277 for tool in tools {
278 router = router.tool(tool);
279 }
280 }
281
282 for skill in skills {
284 router = router.prompt(skill);
285 }
286
287 router
288}
289
290#[derive(Serialize)]
291struct HealthCheckResult {
292 status: String,
293 healthy_count: usize,
294 total_count: usize,
295 backends: Vec<BackendInfo>,
296}
297
298#[derive(Debug, Deserialize, JsonSchema)]
299struct AddBackendInput {
300 name: String,
302 url: String,
304}
305
306#[derive(Debug, Deserialize, JsonSchema)]
308struct CallToolInput {
309 name: String,
311 arguments: Option<serde_json::Map<String, serde_json::Value>>,
313}
314
315#[cfg(test)]
316mod tests {
317 use tower::Service;
318 use tower_mcp::client::ChannelTransport;
319 use tower_mcp::protocol::{
320 CallToolParams, ListToolsParams, McpRequest, McpResponse, RequestId,
321 };
322 use tower_mcp::proxy::McpProxy;
323 use tower_mcp::router::{Extensions, RouterRequest};
324 use tower_mcp::{CallToolResult, McpRouter, SessionHandle, ToolBuilder};
325
326 use super::*;
327
328 fn make_session_handle() -> SessionHandle {
329 let svc = tower::util::BoxCloneService::new(tower::service_fn(
330 |_req: tower_mcp::RouterRequest| async {
331 Ok::<_, std::convert::Infallible>(tower_mcp::RouterResponse {
332 id: RequestId::Number(1),
333 inner: Ok(McpResponse::Pong(Default::default())),
334 })
335 },
336 ));
337 let (_, handle) =
338 tower_mcp::transport::http::HttpTransport::from_service(svc).into_router_with_handle();
339 handle
340 }
341
342 fn make_admin_state() -> AdminState {
343 crate::admin::test_admin_state("test-proxy", "0.1.0", 0, vec![])
344 }
345
346 async fn make_test_proxy() -> McpProxy {
347 let router = McpRouter::new().server_info("test", "1.0.0").tool(
348 ToolBuilder::new("ping")
349 .description("Ping")
350 .handler(|_: tower_mcp::NoParams| async move { Ok(CallToolResult::text("pong")) })
351 .build(),
352 );
353
354 McpProxy::builder("test-proxy", "1.0.0")
355 .backend("test", ChannelTransport::new(router))
356 .await
357 .build_strict()
358 .await
359 .unwrap()
360 }
361
362 async fn list_tools(proxy: &mut McpProxy) -> Vec<String> {
363 let req = RouterRequest {
364 id: RequestId::Number(1),
365 inner: McpRequest::ListTools(ListToolsParams {
366 cursor: None,
367 meta: None,
368 }),
369 extensions: Extensions::new(),
370 };
371 let resp = proxy.call(req).await.expect("infallible");
372 match resp.inner.unwrap() {
373 McpResponse::ListTools(result) => result.tools.into_iter().map(|t| t.name).collect(),
374 other => panic!("expected ListTools, got: {other:?}"),
375 }
376 }
377
378 #[tokio::test]
379 async fn test_build_admin_router_has_expected_tools() {
380 let proxy = make_test_proxy().await;
381 let state = AdminToolState {
382 admin_state: make_admin_state(),
383 session_handle: make_session_handle(),
384 config_snapshot: Arc::new("# empty config".to_string()),
385 proxy: proxy.clone(),
386 };
387
388 let router = build_admin_router(state, None, false, vec![]);
389 let transport = ChannelTransport::new(router);
390
391 let mut test_proxy = McpProxy::builder("verify", "1.0.0")
392 .backend("admin", transport)
393 .await
394 .build_strict()
395 .await
396 .unwrap();
397
398 let tools = list_tools(&mut test_proxy).await;
399 assert!(tools.contains(&"admin_list_backends".to_string()));
400 assert!(tools.contains(&"admin_health_check".to_string()));
401 assert!(tools.contains(&"admin_session_count".to_string()));
402 assert!(tools.contains(&"admin_add_backend".to_string()));
403 assert!(tools.contains(&"admin_config".to_string()));
404 assert!(!tools.contains(&"admin_call_tool".to_string()));
406 }
407
408 #[tokio::test]
409 async fn test_search_mode_adds_call_tool() {
410 let proxy = make_test_proxy().await;
411 let state = AdminToolState {
412 admin_state: make_admin_state(),
413 session_handle: make_session_handle(),
414 config_snapshot: Arc::new(String::new()),
415 proxy: proxy.clone(),
416 };
417
418 let router = build_admin_router(state, None, true, vec![]);
419 let transport = ChannelTransport::new(router);
420
421 let mut test_proxy = McpProxy::builder("verify", "1.0.0")
422 .backend("admin", transport)
423 .await
424 .build_strict()
425 .await
426 .unwrap();
427
428 let tools = list_tools(&mut test_proxy).await;
429 assert!(
430 tools.contains(&"admin_call_tool".to_string()),
431 "search mode should add call_tool, got: {tools:?}"
432 );
433 }
434
435 #[tokio::test]
436 async fn test_discovery_tools_included() {
437 let proxy = make_test_proxy().await;
438 let state = AdminToolState {
439 admin_state: make_admin_state(),
440 session_handle: make_session_handle(),
441 config_snapshot: Arc::new(String::new()),
442 proxy: proxy.clone(),
443 };
444
445 let extra_tool = ToolBuilder::new("search_tools")
446 .description("Search for tools")
447 .handler(
448 |_: tower_mcp::NoParams| async move { Ok(CallToolResult::text("search results")) },
449 )
450 .build();
451
452 let router = build_admin_router(state, Some(vec![extra_tool]), false, vec![]);
453 let transport = ChannelTransport::new(router);
454
455 let mut test_proxy = McpProxy::builder("verify", "1.0.0")
456 .backend("admin", transport)
457 .await
458 .build_strict()
459 .await
460 .unwrap();
461
462 let tools = list_tools(&mut test_proxy).await;
463 assert!(
464 tools.contains(&"admin_search_tools".to_string()),
465 "discovery tool should be included, got: {tools:?}"
466 );
467 }
468
469 #[tokio::test]
470 async fn test_config_tool_returns_snapshot() {
471 let config_text = "[proxy]\nname = \"test\"\n".to_string();
472 let proxy = make_test_proxy().await;
473 let state = AdminToolState {
474 admin_state: make_admin_state(),
475 session_handle: make_session_handle(),
476 config_snapshot: Arc::new(config_text.clone()),
477 proxy: proxy.clone(),
478 };
479
480 let router = build_admin_router(state, None, false, vec![]);
481 let transport = ChannelTransport::new(router);
482
483 let mut test_proxy = McpProxy::builder("verify", "1.0.0")
484 .backend("admin", transport)
485 .await
486 .build_strict()
487 .await
488 .unwrap();
489
490 let req = RouterRequest {
491 id: RequestId::Number(1),
492 inner: McpRequest::CallTool(CallToolParams {
493 name: "admin_config".to_string(),
494 arguments: serde_json::json!({}),
495 input_responses: None,
496 request_state: None,
497 meta: None,
498 task: None,
499 }),
500 extensions: Extensions::new(),
501 };
502 let resp = test_proxy.call(req).await.expect("infallible");
503 match resp.inner.unwrap() {
504 McpResponse::CallTool(result) => {
505 let text = result.all_text();
506 assert!(
507 text.contains("[proxy]"),
508 "config tool should return the config snapshot, got: {text}"
509 );
510 }
511 other => panic!("expected CallTool, got: {other:?}"),
512 }
513 }
514}