1use std::future::Future;
42use std::pin::Pin;
43use std::task::{Context, Poll};
44use std::time::Instant;
45
46use tower::Layer;
47use tower_service::Service;
48use tracing::Level;
49
50use crate::protocol::McpRequest;
51use crate::router::{RouterRequest, RouterResponse, ToolAnnotationsMap};
52
53#[derive(Debug, Clone, Copy)]
68pub struct AuditLayer {
69 level: Level,
70}
71
72impl Default for AuditLayer {
73 fn default() -> Self {
74 Self::new()
75 }
76}
77
78impl AuditLayer {
79 pub fn new() -> Self {
81 Self { level: Level::INFO }
82 }
83
84 pub fn level(mut self, level: Level) -> Self {
88 self.level = level;
89 self
90 }
91}
92
93impl<S> Layer<S> for AuditLayer {
94 type Service = AuditService<S>;
95
96 fn layer(&self, inner: S) -> Self::Service {
97 AuditService {
98 inner,
99 level: self.level,
100 }
101 }
102}
103
104#[derive(Debug, Clone)]
108pub struct AuditService<S> {
109 inner: S,
110 level: Level,
111}
112
113struct AuditInfo {
115 method: String,
116 request_id: String,
117 tool: Option<String>,
118 resource_uri: Option<String>,
119 prompt: Option<String>,
120 read_only: Option<bool>,
121 destructive: Option<bool>,
122}
123
124impl AuditInfo {
125 fn extract(req: &RouterRequest) -> Self {
126 let method = req.inner.method_name().to_string();
127 let request_id = format!("{:?}", req.id);
128
129 let mut info = Self {
130 method,
131 request_id,
132 tool: None,
133 resource_uri: None,
134 prompt: None,
135 read_only: None,
136 destructive: None,
137 };
138
139 match &req.inner {
140 McpRequest::CallTool(params) => {
141 info.tool = Some(params.name.clone());
142
143 if let Some(annotations) = req.extensions.get::<ToolAnnotationsMap>() {
144 info.read_only = Some(annotations.is_read_only(¶ms.name));
145 info.destructive = Some(annotations.is_destructive(¶ms.name));
146 }
147 }
148 McpRequest::ReadResource(params) => {
149 info.resource_uri = Some(params.uri.clone());
150 }
151 McpRequest::GetPrompt(params) => {
152 info.prompt = Some(params.name.clone());
153 }
154 McpRequest::SubscribeResource(params) => {
155 info.resource_uri = Some(params.uri.clone());
156 }
157 McpRequest::UnsubscribeResource(params) => {
158 info.resource_uri = Some(params.uri.clone());
159 }
160 _ => {}
161 }
162
163 info
164 }
165}
166
167const JSONRPC_INVALID_PARAMS: i32 = -32602;
169
170impl<S> Service<RouterRequest> for AuditService<S>
171where
172 S: Service<RouterRequest, Response = RouterResponse> + Clone + Send + 'static,
173 S::Error: Send,
174 S::Future: Send,
175{
176 type Response = RouterResponse;
177 type Error = S::Error;
178 type Future = Pin<Box<dyn Future<Output = Result<RouterResponse, S::Error>> + Send>>;
179
180 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
181 self.inner.poll_ready(cx)
182 }
183
184 fn call(&mut self, req: RouterRequest) -> Self::Future {
185 let info = AuditInfo::extract(&req);
186 let start = Instant::now();
187 let fut = self.inner.call(req);
188 let level = self.level;
189
190 Box::pin(async move {
191 let result = fut.await;
192 let duration_ms = start.elapsed().as_secs_f64() * 1000.0;
193
194 if let Ok(response) = &result {
195 let (status, error) = match &response.inner {
196 Ok(_) => ("success", None),
197 Err(err) => {
198 let s = if err.code == JSONRPC_INVALID_PARAMS {
199 "denied"
200 } else {
201 "error"
202 };
203 (s, Some((err.code, err.message.as_str())))
204 }
205 };
206
207 emit_audit_event(level, &info, duration_ms, status, error);
208 }
209
210 result
211 })
212 }
213}
214
215fn emit_audit_event(
220 level: Level,
221 info: &AuditInfo,
222 duration_ms: f64,
223 status: &str,
224 error: Option<(i32, &str)>,
225) {
226 let method = info.method.as_str();
227 let request_id = info.request_id.as_str();
228 let tool = info.tool.as_deref();
229 let resource_uri = info.resource_uri.as_deref();
230 let prompt = info.prompt.as_deref();
231 let read_only = info.read_only;
232 let destructive = info.destructive;
233
234 match (level, error) {
235 (Level::TRACE, None) => {
236 tracing::trace!(target: "mcp::audit", method, request_id, ?tool, ?resource_uri, ?prompt, duration_ms, status, ?read_only, ?destructive, "audit")
237 }
238 (Level::TRACE, Some((code, msg))) => {
239 tracing::trace!(target: "mcp::audit", method, request_id, ?tool, ?resource_uri, ?prompt, duration_ms, status, error_code = code, error_message = msg, ?read_only, ?destructive, "audit")
240 }
241 (Level::DEBUG, None) => {
242 tracing::debug!(target: "mcp::audit", method, request_id, ?tool, ?resource_uri, ?prompt, duration_ms, status, ?read_only, ?destructive, "audit")
243 }
244 (Level::DEBUG, Some((code, msg))) => {
245 tracing::debug!(target: "mcp::audit", method, request_id, ?tool, ?resource_uri, ?prompt, duration_ms, status, error_code = code, error_message = msg, ?read_only, ?destructive, "audit")
246 }
247 (Level::INFO, None) => {
248 tracing::info!(target: "mcp::audit", method, request_id, ?tool, ?resource_uri, ?prompt, duration_ms, status, ?read_only, ?destructive, "audit")
249 }
250 (Level::INFO, Some((code, msg))) => {
251 tracing::info!(target: "mcp::audit", method, request_id, ?tool, ?resource_uri, ?prompt, duration_ms, status, error_code = code, error_message = msg, ?read_only, ?destructive, "audit")
252 }
253 (Level::WARN, None) => {
254 tracing::warn!(target: "mcp::audit", method, request_id, ?tool, ?resource_uri, ?prompt, duration_ms, status, ?read_only, ?destructive, "audit")
255 }
256 (Level::WARN, Some((code, msg))) => {
257 tracing::warn!(target: "mcp::audit", method, request_id, ?tool, ?resource_uri, ?prompt, duration_ms, status, error_code = code, error_message = msg, ?read_only, ?destructive, "audit")
258 }
259 (Level::ERROR, None) => {
260 tracing::error!(target: "mcp::audit", method, request_id, ?tool, ?resource_uri, ?prompt, duration_ms, status, ?read_only, ?destructive, "audit")
261 }
262 (Level::ERROR, Some((code, msg))) => {
263 tracing::error!(target: "mcp::audit", method, request_id, ?tool, ?resource_uri, ?prompt, duration_ms, status, error_code = code, error_message = msg, ?read_only, ?destructive, "audit")
264 }
265 }
266}
267
268#[cfg(test)]
269mod tests {
270 use super::*;
271 use crate::protocol::{CallToolParams, GetPromptParams, ReadResourceParams, RequestId};
272 use crate::router::Extensions;
273 use std::collections::HashMap;
274
275 #[test]
276 fn test_layer_creation() {
277 let layer = AuditLayer::new();
278 assert_eq!(layer.level, Level::INFO);
279 }
280
281 #[test]
282 fn test_layer_with_custom_level() {
283 let layer = AuditLayer::new().level(Level::DEBUG);
284 assert_eq!(layer.level, Level::DEBUG);
285 }
286
287 #[test]
288 fn test_layer_default() {
289 let layer = AuditLayer::default();
290 assert_eq!(layer.level, Level::INFO);
291 }
292
293 #[test]
294 fn test_audit_info_tool_call() {
295 let req = RouterRequest {
296 id: RequestId::Number(1),
297 inner: McpRequest::CallTool(CallToolParams {
298 input_responses: None,
299 request_state: None,
300 name: "my_tool".to_string(),
301 arguments: serde_json::json!({}),
302 meta: None,
303 task: None,
304 }),
305 extensions: Extensions::new(),
306 };
307
308 let info = AuditInfo::extract(&req);
309 assert_eq!(info.method, "tools/call");
310 assert_eq!(info.tool, Some("my_tool".to_string()));
311 assert!(info.resource_uri.is_none());
312 assert!(info.prompt.is_none());
313 }
314
315 #[test]
316 fn test_audit_info_resource_read() {
317 let req = RouterRequest {
318 id: RequestId::Number(2),
319 inner: McpRequest::ReadResource(ReadResourceParams {
320 input_responses: None,
321 request_state: None,
322 uri: "file:///test.txt".to_string(),
323 meta: None,
324 }),
325 extensions: Extensions::new(),
326 };
327
328 let info = AuditInfo::extract(&req);
329 assert_eq!(info.method, "resources/read");
330 assert!(info.tool.is_none());
331 assert_eq!(info.resource_uri, Some("file:///test.txt".to_string()));
332 }
333
334 #[test]
335 fn test_audit_info_prompt_get() {
336 let req = RouterRequest {
337 id: RequestId::Number(3),
338 inner: McpRequest::GetPrompt(GetPromptParams {
339 input_responses: None,
340 request_state: None,
341 name: "review".to_string(),
342 arguments: HashMap::new(),
343 meta: None,
344 }),
345 extensions: Extensions::new(),
346 };
347
348 let info = AuditInfo::extract(&req);
349 assert_eq!(info.method, "prompts/get");
350 assert!(info.tool.is_none());
351 assert_eq!(info.prompt, Some("review".to_string()));
352 }
353
354 #[test]
355 fn test_audit_info_ping() {
356 let req = RouterRequest {
357 id: RequestId::Number(4),
358 inner: McpRequest::Ping,
359 extensions: Extensions::new(),
360 };
361
362 let info = AuditInfo::extract(&req);
363 assert_eq!(info.method, "ping");
364 assert!(info.tool.is_none());
365 assert!(info.resource_uri.is_none());
366 assert!(info.prompt.is_none());
367 }
368
369 #[tokio::test]
370 async fn test_passthrough() {
371 let router = crate::McpRouter::new().server_info("test", "1.0.0");
372 let layer = AuditLayer::new();
373 let mut service = layer.layer(router);
374
375 let req = RouterRequest {
376 id: RequestId::Number(1),
377 inner: McpRequest::Ping,
378 extensions: Extensions::new(),
379 };
380
381 let result = Service::call(&mut service, req).await;
382 assert!(result.is_ok());
383 assert!(result.unwrap().inner.is_ok());
384 }
385
386 #[tokio::test]
387 async fn test_tool_call_audit() {
388 let tool = crate::ToolBuilder::new("test_tool")
389 .description("A test tool")
390 .handler(|_: serde_json::Value| async move { Ok(crate::CallToolResult::text("done")) })
391 .build();
392
393 let router = crate::McpRouter::new()
394 .server_info("test", "1.0.0")
395 .tool(tool);
396 let layer = AuditLayer::new();
397 let mut service = layer.layer(router);
398
399 let req = RouterRequest {
400 id: RequestId::Number(1),
401 inner: McpRequest::CallTool(CallToolParams {
402 input_responses: None,
403 request_state: None,
404 name: "test_tool".to_string(),
405 arguments: serde_json::json!({}),
406 meta: None,
407 task: None,
408 }),
409 extensions: Extensions::new(),
410 };
411
412 let result = Service::call(&mut service, req).await;
413 assert!(result.is_ok());
414 }
415
416 #[tokio::test]
417 async fn test_error_audit() {
418 let router = crate::McpRouter::new().server_info("test", "1.0.0");
419 let layer = AuditLayer::new();
420 let mut service = layer.layer(router);
421
422 let req = RouterRequest {
423 id: RequestId::Number(1),
424 inner: McpRequest::CallTool(CallToolParams {
425 input_responses: None,
426 request_state: None,
427 name: "nonexistent".to_string(),
428 arguments: serde_json::json!({}),
429 meta: None,
430 task: None,
431 }),
432 extensions: Extensions::new(),
433 };
434
435 let result = Service::call(&mut service, req).await;
436 assert!(result.is_ok());
437 assert!(result.unwrap().inner.is_err());
438 }
439}