1use std::{borrow::Cow, fmt, future::Future, pin::Pin, sync::Arc};
40
41use rmcp::{
42 ErrorData, RoleServer, ServerHandler,
43 model::{
44 CallToolRequestParams, CallToolResponse, CallToolResult, CancelTaskParams, ContentBlock,
45 DiscoverResult, GetPromptRequestParams, GetPromptResponse, GetTaskParams, GetTaskResult,
46 InitializeRequestParams, InitializeResult, ListPromptsResult, ListResourceTemplatesResult,
47 ListResourcesResult, ListToolsResult, PaginatedRequestParams, ProtocolVersion,
48 ReadResourceRequestParams, ReadResourceResponse, ServerInfo, SubscriptionFilter, Tool,
49 UpdateTaskParams,
50 },
51 service::{RequestContext, SubscriptionContext},
52};
53
54#[derive(Debug, Clone)]
56#[non_exhaustive]
57pub struct ToolCallContext {
58 pub tool_name: String,
60 pub arguments: Option<serde_json::Value>,
62 pub identity: Option<String>,
64 pub role: Option<String>,
66 pub sub: Option<String>,
68 pub request_id: Option<String>,
70}
71
72impl ToolCallContext {
73 #[must_use]
78 pub fn for_tool(tool_name: impl Into<String>) -> Self {
79 Self {
80 tool_name: tool_name.into(),
81 arguments: None,
82 identity: None,
83 role: None,
84 sub: None,
85 request_id: None,
86 }
87 }
88}
89
90#[derive(Debug)]
99#[non_exhaustive]
100pub enum HookOutcome {
101 Continue,
103 Deny(ErrorData),
105 Replace(Box<CallToolResult>),
107}
108
109#[derive(Debug, Clone, Copy)]
111#[non_exhaustive]
112pub enum HookDisposition {
113 InnerExecuted,
115 InnerErrored,
117 DeniedBefore,
119 ReplacedBefore,
121 ResultTooLarge,
124}
125
126pub type BeforeHook = Arc<
133 dyn for<'a> Fn(&'a ToolCallContext) -> Pin<Box<dyn Future<Output = HookOutcome> + Send + 'a>>
134 + Send
135 + Sync
136 + 'static,
137>;
138
139pub type AfterHook = Arc<
147 dyn for<'a> Fn(
148 &'a ToolCallContext,
149 HookDisposition,
150 usize,
151 ) -> Pin<Box<dyn Future<Output = ()> + Send + 'a>>
152 + Send
153 + Sync
154 + 'static,
155>;
156
157#[allow(clippy::struct_field_names, reason = "before/after read naturally")]
159#[derive(Clone, Default)]
160#[non_exhaustive]
161pub struct ToolHooks {
162 pub max_result_bytes: Option<usize>,
167 pub before: Option<BeforeHook>,
170 pub after: Option<AfterHook>,
174}
175
176impl ToolHooks {
177 #[must_use]
183 pub fn new() -> Self {
184 Self::default()
185 }
186
187 #[must_use]
189 pub fn with_max_result_bytes(mut self, max: usize) -> Self {
190 self.max_result_bytes = Some(max);
191 self
192 }
193
194 #[must_use]
196 pub fn with_before(mut self, before: BeforeHook) -> Self {
197 self.before = Some(before);
198 self
199 }
200
201 #[must_use]
203 pub fn with_after(mut self, after: AfterHook) -> Self {
204 self.after = Some(after);
205 self
206 }
207}
208
209impl fmt::Debug for ToolHooks {
210 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
211 f.debug_struct("ToolHooks")
212 .field("max_result_bytes", &self.max_result_bytes)
213 .field("before", &self.before.as_ref().map(|_| "<fn>"))
214 .field("after", &self.after.as_ref().map(|_| "<fn>"))
215 .finish()
216 }
217}
218
219#[derive(Clone)]
221pub struct HookedHandler<H: ServerHandler> {
222 inner: Arc<H>,
223 hooks: Arc<ToolHooks>,
224}
225
226impl<H: ServerHandler> fmt::Debug for HookedHandler<H> {
227 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
228 f.debug_struct("HookedHandler")
229 .field("hooks", &self.hooks)
230 .finish_non_exhaustive()
231 }
232}
233
234pub fn with_hooks<H: ServerHandler>(inner: H, hooks: Arc<ToolHooks>) -> HookedHandler<H> {
248 HookedHandler {
249 inner: Arc::new(inner),
250 hooks,
251 }
252}
253
254impl<H: ServerHandler> HookedHandler<H> {
255 #[must_use]
257 pub fn inner(&self) -> &H {
258 &self.inner
259 }
260
261 fn build_context(request: &CallToolRequestParams, req_id: Option<String>) -> ToolCallContext {
262 ToolCallContext {
263 tool_name: request.name.to_string(),
264 arguments: request.arguments.clone().map(serde_json::Value::Object),
265 identity: crate::rbac::current_identity(),
266 role: crate::rbac::current_role(),
267 sub: crate::rbac::current_sub(),
268 request_id: req_id,
269 }
270 }
271
272 fn spawn_after(
284 after: Option<&Arc<AfterHookHolder>>,
285 ctx: ToolCallContext,
286 disposition: HookDisposition,
287 size: usize,
288 ) {
289 if let Some(after) = after {
290 use tracing::Instrument;
291
292 let after = Arc::clone(after);
293 let span = tracing::Span::current();
296 let role = crate::rbac::current_role().unwrap_or_default();
300 let identity = crate::rbac::current_identity().unwrap_or_default();
301 let token = crate::rbac::current_token()
302 .unwrap_or_else(|| secrecy::SecretString::from(String::new()));
303 let sub = crate::rbac::current_sub().unwrap_or_default();
304 tokio::spawn(
305 async move {
306 crate::rbac::with_rbac_scope(role, identity, token, sub, async move {
307 let fut = (after.f)(&ctx, disposition, size);
308 fut.await;
309 })
310 .await;
311 }
312 .instrument(span),
313 );
314 }
315 }
316}
317
318struct AfterHookHolder {
322 f: AfterHook,
323}
324
325fn too_large_result(limit: usize, actual: usize, tool: &str) -> CallToolResult {
327 let body = serde_json::json!({
328 "error": "result_too_large",
329 "message": format!(
330 "tool '{tool}' result of {actual} bytes exceeds the configured \
331 max_result_bytes={limit}; ask for a narrower query"
332 ),
333 "limit_bytes": limit,
334 "actual_bytes": actual,
335 });
336 let mut r = CallToolResult::error(vec![ContentBlock::text(body.to_string())]);
337 r.structured_content = None;
338 r
339}
340
341fn serialized_size(result: &CallToolResult) -> usize {
342 serde_json::to_vec(result).map_or(0, |v| v.len())
343}
344
345fn apply_size_cap(
349 result: CallToolResult,
350 max: Option<usize>,
351 tool: &str,
352) -> (CallToolResult, usize, bool) {
353 let size = serialized_size(&result);
354 if let Some(limit) = max
355 && size > limit
356 {
357 tracing::warn!(
358 tool = %tool,
359 size_bytes = size,
360 limit_bytes = limit,
361 "tool result exceeds max_result_bytes; replacing with structured error"
362 );
363 let replaced = too_large_result(limit, size, tool);
364 return (replaced, size, true);
365 }
366 (result, size, false)
367}
368
369impl<H: ServerHandler> ServerHandler for HookedHandler<H> {
370 fn get_info(&self) -> ServerInfo {
371 self.inner.get_info()
372 }
373
374 async fn initialize(
375 &self,
376 request: InitializeRequestParams,
377 context: RequestContext<RoleServer>,
378 ) -> Result<InitializeResult, ErrorData> {
379 self.inner.initialize(request, context).await
380 }
381
382 async fn list_tools(
383 &self,
384 request: Option<PaginatedRequestParams>,
385 context: RequestContext<RoleServer>,
386 ) -> Result<ListToolsResult, ErrorData> {
387 self.inner.list_tools(request, context).await
388 }
389
390 fn get_tool(&self, name: &str) -> Option<Tool> {
391 self.inner.get_tool(name)
392 }
393
394 async fn list_prompts(
395 &self,
396 request: Option<PaginatedRequestParams>,
397 context: RequestContext<RoleServer>,
398 ) -> Result<ListPromptsResult, ErrorData> {
399 self.inner.list_prompts(request, context).await
400 }
401
402 async fn get_prompt(
403 &self,
404 request: GetPromptRequestParams,
405 context: RequestContext<RoleServer>,
406 ) -> Result<GetPromptResponse, ErrorData> {
407 self.inner.get_prompt(request, context).await
408 }
409
410 async fn list_resources(
411 &self,
412 request: Option<PaginatedRequestParams>,
413 context: RequestContext<RoleServer>,
414 ) -> Result<ListResourcesResult, ErrorData> {
415 self.inner.list_resources(request, context).await
416 }
417
418 async fn list_resource_templates(
419 &self,
420 request: Option<PaginatedRequestParams>,
421 context: RequestContext<RoleServer>,
422 ) -> Result<ListResourceTemplatesResult, ErrorData> {
423 self.inner.list_resource_templates(request, context).await
424 }
425
426 async fn read_resource(
427 &self,
428 request: ReadResourceRequestParams,
429 context: RequestContext<RoleServer>,
430 ) -> Result<ReadResourceResponse, ErrorData> {
431 self.inner.read_resource(request, context).await
432 }
433
434 #[allow(
435 clippy::wildcard_enum_match_arm,
436 reason = "CallToolResponse is #[non_exhaustive]; the non-Complete MRTR variants (InputRequired/Task) are passed through unchanged"
437 )]
438 async fn call_tool(
439 &self,
440 request: CallToolRequestParams,
441 context: RequestContext<RoleServer>,
442 ) -> Result<CallToolResponse, ErrorData> {
443 let req_id = Some(format!("{:?}", context.id));
444 let ctx = Self::build_context(&request, req_id);
445 let max = self.hooks.max_result_bytes;
446 let after_holder = self
447 .hooks
448 .after
449 .as_ref()
450 .map(|f| Arc::new(AfterHookHolder { f: Arc::clone(f) }));
451
452 if let Some(before) = self.hooks.before.as_ref() {
454 let outcome = before(&ctx).await;
455 match outcome {
456 HookOutcome::Continue => {}
457 HookOutcome::Deny(err) => {
458 Self::spawn_after(after_holder.as_ref(), ctx, HookDisposition::DeniedBefore, 0);
459 return Err(err);
460 }
461 HookOutcome::Replace(boxed) => {
462 let (final_result, size, capped) = apply_size_cap(*boxed, max, &ctx.tool_name);
463 let disposition = if capped {
464 HookDisposition::ResultTooLarge
465 } else {
466 HookDisposition::ReplacedBefore
467 };
468 Self::spawn_after(after_holder.as_ref(), ctx, disposition, size);
469 return Ok(final_result.into());
470 }
471 }
472 }
473
474 match self.inner.call_tool(request, context).await {
476 Ok(CallToolResponse::Complete(result)) => {
478 let (final_result, size, capped) = apply_size_cap(result, max, &ctx.tool_name);
479 let disposition = if capped {
480 HookDisposition::ResultTooLarge
481 } else {
482 HookDisposition::InnerExecuted
483 };
484 Self::spawn_after(after_holder.as_ref(), ctx, disposition, size);
485 Ok(final_result.into())
486 }
487 Ok(other) => {
490 Self::spawn_after(
491 after_holder.as_ref(),
492 ctx,
493 HookDisposition::InnerExecuted,
494 0,
495 );
496 Ok(other)
497 }
498 Err(e) => {
499 Self::spawn_after(after_holder.as_ref(), ctx, HookDisposition::InnerErrored, 0);
500 Err(e)
501 }
502 }
503 }
504
505 fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> {
509 self.inner.supported_protocol_versions()
510 }
511
512 async fn discover(
513 &self,
514 context: RequestContext<RoleServer>,
515 ) -> Result<DiscoverResult, ErrorData> {
516 self.inner.discover(context).await
517 }
518
519 fn accepted_subscription_filter(
520 &self,
521 requested: &SubscriptionFilter,
522 ) -> Option<SubscriptionFilter> {
523 self.inner.accepted_subscription_filter(requested)
524 }
525
526 async fn listen(&self, context: SubscriptionContext) -> Result<(), ErrorData> {
527 self.inner.listen(context).await
528 }
529
530 async fn get_task(
531 &self,
532 request: GetTaskParams,
533 context: RequestContext<RoleServer>,
534 ) -> Result<GetTaskResult, ErrorData> {
535 self.inner.get_task(request, context).await
536 }
537
538 async fn update_task(
539 &self,
540 request: UpdateTaskParams,
541 context: RequestContext<RoleServer>,
542 ) -> Result<(), ErrorData> {
543 self.inner.update_task(request, context).await
544 }
545
546 async fn cancel_task(
547 &self,
548 request: CancelTaskParams,
549 context: RequestContext<RoleServer>,
550 ) -> Result<(), ErrorData> {
551 self.inner.cancel_task(request, context).await
552 }
553}
554
555#[cfg(test)]
556mod tests {
557 use std::sync::{
558 Arc,
559 atomic::{AtomicUsize, Ordering},
560 };
561
562 use rmcp::{
563 ErrorData, RoleServer, ServerHandler,
564 model::{
565 CallToolRequestParams, CallToolResponse, CallToolResult, ContentBlock, ServerInfo,
566 },
567 service::RequestContext,
568 };
569
570 use super::*;
571
572 #[derive(Clone, Default)]
574 struct TestHandler {
575 body_bytes: Option<usize>,
577 }
578
579 impl ServerHandler for TestHandler {
580 fn get_info(&self) -> ServerInfo {
581 ServerInfo::default()
582 }
583
584 async fn call_tool(
585 &self,
586 _request: CallToolRequestParams,
587 _context: RequestContext<RoleServer>,
588 ) -> Result<CallToolResponse, ErrorData> {
589 let body = "x".repeat(self.body_bytes.unwrap_or(4));
590 Ok(CallToolResult::success(vec![ContentBlock::text(body)]).into())
591 }
592 }
593
594 fn ctx(name: &str) -> ToolCallContext {
595 ToolCallContext {
596 tool_name: name.to_owned(),
597 arguments: None,
598 identity: None,
599 role: None,
600 sub: None,
601 request_id: None,
602 }
603 }
604
605 #[tokio::test]
606 async fn size_cap_replaces_oversized_result() {
607 let inner = TestHandler {
608 body_bytes: Some(8_192),
609 };
610 let hooks = Arc::new(ToolHooks {
611 max_result_bytes: Some(256),
612 before: None,
613 after: None,
614 });
615 let hooked = with_hooks(inner, hooks);
616
617 let small = CallToolResult::success(vec![ContentBlock::text("ok".to_owned())]);
618 assert!(serialized_size(&small) < 256);
619
620 let big = CallToolResult::success(vec![ContentBlock::text("x".repeat(8_192))]);
621 let size = serialized_size(&big);
622 assert!(size > 256);
623
624 let (replaced, accounted, capped) = apply_size_cap(big, Some(256), "whatever");
625 assert!(capped);
626 assert_eq!(accounted, size);
627 assert_eq!(replaced.is_error, Some(true));
628 assert!(matches!(
629 replaced.content.first(),
630 Some(rmcp::model::ContentBlock::Text(t)) if t.text.contains("result_too_large")
631 ));
632
633 let _ = hooked;
635 }
636
637 #[tokio::test]
638 async fn before_hook_deny_builds_error() {
639 let counter = Arc::new(AtomicUsize::new(0));
640 let c = Arc::clone(&counter);
641 let before: BeforeHook = Arc::new(move |ctx_ref| {
642 let c = Arc::clone(&c);
643 let name = ctx_ref.tool_name.clone();
644 Box::pin(async move {
645 c.fetch_add(1, Ordering::Relaxed);
646 if name == "forbidden" {
647 HookOutcome::Deny(ErrorData::invalid_request("nope", None))
648 } else {
649 HookOutcome::Continue
650 }
651 })
652 });
653
654 let hooks = Arc::new(ToolHooks {
655 max_result_bytes: None,
656 before: Some(before),
657 after: None,
658 });
659 let hooked = with_hooks(TestHandler::default(), hooks);
660
661 let bad_ctx = ctx("forbidden");
662 let before_fn = hooked.hooks.before.as_ref().unwrap();
663 let outcome = before_fn(&bad_ctx).await;
664 assert!(matches!(outcome, HookOutcome::Deny(_)));
665 assert_eq!(counter.load(Ordering::Relaxed), 1);
666
667 let ok_ctx = ctx("allowed");
668 let outcome2 = before_fn(&ok_ctx).await;
669 assert!(matches!(outcome2, HookOutcome::Continue));
670 assert_eq!(counter.load(Ordering::Relaxed), 2);
671 }
672
673 #[test]
674 fn too_large_result_mentions_limit_and_actual() {
675 let r = too_large_result(100, 500, "my_tool");
676 let body = serde_json::to_string(&r).unwrap();
677 assert!(body.contains("result_too_large"));
678 assert!(body.contains("my_tool"));
679 assert!(body.contains("100"));
680 assert!(body.contains("500"));
681 }
682
683 #[tokio::test]
684 async fn replace_outcome_skips_inner_and_returns_payload() {
685 let before: BeforeHook = Arc::new(|_ctx| {
688 Box::pin(async {
689 HookOutcome::Replace(Box::new(CallToolResult::success(vec![ContentBlock::text(
690 "from-replace".to_owned(),
691 )])))
692 })
693 });
694 let hooks = Arc::new(ToolHooks {
695 max_result_bytes: None,
696 before: Some(before),
697 after: None,
698 });
699 let _hooked = with_hooks(TestHandler::default(), Arc::clone(&hooks));
700
701 let outcome = (hooks.before.as_ref().unwrap())(&ctx("any")).await;
704 let HookOutcome::Replace(boxed) = outcome else {
705 panic!("expected HookOutcome::Replace");
706 };
707 let (result, size, capped) = apply_size_cap(*boxed, None, "any");
708 assert!(!capped);
709 assert!(size > 0);
710 assert!(!result.is_error.unwrap_or(false));
711 assert!(matches!(
712 result.content.first(),
713 Some(rmcp::model::ContentBlock::Text(t)) if t.text == "from-replace"
714 ));
715 }
716
717 #[tokio::test]
718 async fn replace_outcome_subject_to_size_cap() {
719 let huge = CallToolResult::success(vec![ContentBlock::text("y".repeat(8_192))]);
723 let huge_size = serialized_size(&huge);
724 assert!(huge_size > 256);
725
726 let (final_result, accounted, capped) = apply_size_cap(huge, Some(256), "replaced_tool");
727 assert!(capped);
728 assert_eq!(accounted, huge_size);
729 assert_eq!(final_result.is_error, Some(true));
730 assert!(matches!(
731 final_result.content.first(),
732 Some(rmcp::model::ContentBlock::Text(t)) if t.text.contains("result_too_large")
733 ));
734 }
735
736 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
737 async fn after_hook_fires_exactly_once_via_spawn() {
738 let counter = Arc::new(AtomicUsize::new(0));
742 let c = Arc::clone(&counter);
743 let after: AfterHook = Arc::new(move |_ctx, _disp, _size| {
744 let c = Arc::clone(&c);
745 Box::pin(async move {
746 c.fetch_add(1, Ordering::Relaxed);
747 })
748 });
749 let holder = Arc::new(AfterHookHolder { f: after });
750
751 HookedHandler::<TestHandler>::spawn_after(
752 Some(&holder),
753 ctx("t"),
754 HookDisposition::InnerExecuted,
755 42,
756 );
757
758 let deadline = std::time::Instant::now() + std::time::Duration::from_secs(1);
760 while counter.load(Ordering::Relaxed) == 0 && std::time::Instant::now() < deadline {
761 tokio::task::yield_now().await;
762 tokio::time::sleep(std::time::Duration::from_millis(5)).await;
763 }
764 assert_eq!(counter.load(Ordering::Relaxed), 1);
765 }
766
767 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
768 async fn after_hook_panic_is_isolated_from_response_path() {
769 let after: AfterHook = Arc::new(|_ctx, _disp, _size| {
773 Box::pin(async {
774 panic!("intentional panic in after-hook");
775 })
776 });
777 let holder = Arc::new(AfterHookHolder { f: after });
778
779 HookedHandler::<TestHandler>::spawn_after(
780 Some(&holder),
781 ctx("boom"),
782 HookDisposition::InnerExecuted,
783 0,
784 );
785
786 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
789 let still_alive = tokio::spawn(async { 1_u32 + 2 }).await.unwrap();
790 assert_eq!(still_alive, 3);
791 }
792}