1use std::{borrow::Cow, fmt, future::Future, io, pin::Pin, sync::Arc};
51
52#[allow(
53 deprecated,
54 reason = "transparent ServerHandler delegation must import legacy logging/subscription parameter types until rmcp removes those methods"
55)]
56use rmcp::{
57 ErrorData, RoleServer, ServerHandler,
58 model::{
59 CallToolRequestParams, CallToolResponse, CallToolResult, CancelTaskParams,
60 CancelledNotificationParam, CompleteRequestParams, CompleteResult, ContentBlock,
61 CustomNotification, CustomRequest, CustomResult, DiscoverResult, GetPromptRequestParams,
62 GetPromptResponse, GetTaskParams, GetTaskResult, InitializeRequestParams, InitializeResult,
63 ListPromptsResult, ListResourceTemplatesResult, ListResourcesResult, ListToolsResult,
64 PaginatedRequestParams, ProgressNotificationParam, ProtocolVersion,
65 ReadResourceRequestParams, ReadResourceResponse, ServerConfig, SetLevelRequestParams,
66 SubscribeRequestParams, SubscriptionFilter, Tool, UnsubscribeRequestParams,
67 UpdateTaskParams,
68 },
69 service::{NotificationContext, RequestContext, SubscriptionContext},
70};
71
72#[derive(Clone)]
74#[non_exhaustive]
75pub struct ToolCallContext {
76 pub tool_name: String,
78 pub arguments: Option<serde_json::Value>,
80 pub identity: Option<String>,
82 pub role: Option<String>,
84 pub sub: Option<String>,
86 pub request_id: Option<String>,
88}
89
90impl ToolCallContext {
91 #[must_use]
96 pub fn for_tool(tool_name: impl Into<String>) -> Self {
97 Self {
98 tool_name: tool_name.into(),
99 arguments: None,
100 identity: None,
101 role: None,
102 sub: None,
103 request_id: None,
104 }
105 }
106}
107
108impl fmt::Debug for ToolCallContext {
109 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
110 let Self {
111 tool_name,
112 arguments,
113 identity,
114 role,
115 sub,
116 request_id,
117 } = self;
118 let mut debug = f.debug_struct("ToolCallContext");
119 debug.field("tool_name", tool_name);
120 if crate::diagnostics::tool_call_arguments() {
121 debug
122 .field("arguments", arguments)
123 .field("identity", identity)
124 .field("role", role)
125 .field("sub", sub);
126 } else {
127 debug
128 .field("arguments", &"[REDACTED]")
129 .field("identity", &"[REDACTED]")
130 .field("role", &"[REDACTED]")
131 .field("sub", &"[REDACTED]");
132 }
133 debug.field("request_id", request_id).finish()
134 }
135}
136
137#[derive(Debug)]
146#[non_exhaustive]
147pub enum HookOutcome {
148 Continue,
150 Deny(ErrorData),
152 Replace(Box<CallToolResult>),
154}
155
156#[derive(Debug, Clone, Copy)]
158#[non_exhaustive]
159pub enum HookDisposition {
160 InnerExecuted,
162 InnerErrored,
164 DeniedBefore,
166 ReplacedBefore,
168 ResultTooLarge,
171}
172
173pub type BeforeHook = Arc<
180 dyn for<'a> Fn(&'a ToolCallContext) -> Pin<Box<dyn Future<Output = HookOutcome> + Send + 'a>>
181 + Send
182 + Sync
183 + 'static,
184>;
185
186pub type AfterHook = Arc<
194 dyn for<'a> Fn(
195 &'a ToolCallContext,
196 HookDisposition,
197 usize,
198 ) -> Pin<Box<dyn Future<Output = ()> + Send + 'a>>
199 + Send
200 + Sync
201 + 'static,
202>;
203
204#[allow(clippy::struct_field_names, reason = "before/after read naturally")]
206#[derive(Clone, Default)]
207#[non_exhaustive]
208pub struct ToolHooks {
209 pub max_result_bytes: Option<usize>,
214 pub before: Option<BeforeHook>,
217 pub after: Option<AfterHook>,
227}
228
229impl ToolHooks {
230 #[must_use]
236 pub fn new() -> Self {
237 Self::default()
238 }
239
240 #[must_use]
242 pub fn with_max_result_bytes(mut self, max: usize) -> Self {
243 self.max_result_bytes = Some(max);
244 self
245 }
246
247 #[must_use]
249 pub fn with_before(mut self, before: BeforeHook) -> Self {
250 self.before = Some(before);
251 self
252 }
253
254 #[must_use]
256 pub fn with_after(mut self, after: AfterHook) -> Self {
257 self.after = Some(after);
258 self
259 }
260}
261
262const _HOOKED_HANDLER_DOC_ANCHOR: &str = "HookedHandler";
263
264impl fmt::Debug for ToolHooks {
265 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
266 f.debug_struct("ToolHooks")
267 .field("max_result_bytes", &self.max_result_bytes)
268 .field("before", &self.before.as_ref().map(|_| "<fn>"))
269 .field("after", &self.after.as_ref().map(|_| "<fn>"))
270 .finish()
271 }
272}
273
274#[derive(Clone)]
276pub struct HookedHandler<H: ServerHandler> {
277 inner: Arc<H>,
278 hooks: Arc<ToolHooks>,
279}
280
281impl<H: ServerHandler> fmt::Debug for HookedHandler<H> {
282 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
283 f.debug_struct("HookedHandler")
284 .field("hooks", &self.hooks)
285 .finish_non_exhaustive()
286 }
287}
288
289#[must_use = "HookedHandler must be wired into a ServerHandler (e.g. via \
294 `serve(..., || hooked)`) to take effect; dropping the returned \
295 value silently disables the supplied hooks"]
296pub fn with_hooks<H: ServerHandler>(inner: H, hooks: Arc<ToolHooks>) -> HookedHandler<H> {
297 HookedHandler {
298 inner: Arc::new(inner),
299 hooks,
300 }
301}
302
303impl<H: ServerHandler> HookedHandler<H> {
304 #[must_use]
306 pub fn inner(&self) -> &H {
307 &self.inner
308 }
309
310 fn build_context(request: &CallToolRequestParams, req_id: Option<String>) -> ToolCallContext {
311 ToolCallContext {
312 tool_name: request.name.to_string(),
313 arguments: request.arguments.clone().map(serde_json::Value::Object),
314 identity: crate::rbac::current_identity(),
315 role: crate::rbac::current_role(),
316 sub: crate::rbac::current_sub(),
317 request_id: req_id,
318 }
319 }
320
321 fn spawn_after(
333 after: Option<&Arc<AfterHookHolder>>,
334 ctx: ToolCallContext,
335 disposition: HookDisposition,
336 size: usize,
337 ) {
338 if let Some(after) = after {
339 use tracing::Instrument;
340
341 let after = Arc::clone(after);
342 let span = tracing::Span::current();
345 let role = crate::rbac::current_role().unwrap_or_default();
349 let identity = crate::rbac::current_identity().unwrap_or_default();
350 let token = crate::rbac::current_token()
351 .unwrap_or_else(|| secrecy::SecretString::from(String::new()));
352 let sub = crate::rbac::current_sub().unwrap_or_default();
353 tokio::spawn(
354 async move {
355 crate::rbac::with_rbac_scope(role, identity, token, sub, async move {
356 let fut = (after.f)(&ctx, disposition, size);
357 fut.await;
358 })
359 .await;
360 }
361 .instrument(span),
362 );
363 }
364 }
365}
366
367struct AfterHookHolder {
371 f: AfterHook,
372}
373
374fn too_large_result(limit: usize, actual: Option<usize>, tool: &str) -> CallToolResult {
380 let actual_desc =
381 actual.map_or_else(|| "an unmeasurable number of".to_owned(), |n| n.to_string());
382 let body = serde_json::json!({
383 "error": "result_too_large",
384 "message": format!(
385 "tool '{tool}' result of {actual_desc} bytes exceeds the configured \
386 max_result_bytes={limit}; ask for a narrower query"
387 ),
388 "limit_bytes": limit,
389 "actual_bytes": actual.map_or_else(
390 || serde_json::Value::from("unknown"),
391 serde_json::Value::from,
392 ),
393 });
394 let mut r = CallToolResult::error(vec![ContentBlock::text(body.to_string())]);
395 r.structured_content = None;
396 r
397}
398
399#[derive(Debug, PartialEq, Eq)]
402enum SizeVerdict {
403 Pass { size: usize },
405 Replace { limit: usize, actual: Option<usize> },
407 PassUnmeasured,
409}
410
411const fn decide_size(size: Option<SizeMeasure>, max: Option<usize>) -> SizeVerdict {
413 match size {
414 Some(SizeMeasure::Exact(size)) => match max {
415 Some(limit) if size > limit => SizeVerdict::Replace {
416 limit,
417 actual: Some(size),
418 },
419 Some(_) | None => SizeVerdict::Pass { size },
420 },
421 Some(SizeMeasure::Exceeded { limit }) => SizeVerdict::Replace {
422 limit,
423 actual: None,
424 },
425 None => match max {
426 Some(limit) => SizeVerdict::Replace {
427 limit,
428 actual: None,
429 },
430 None => SizeVerdict::PassUnmeasured,
431 },
432 }
433}
434
435fn apply_size_cap(
439 result: CallToolResult,
440 max: Option<usize>,
441 tool: &str,
442) -> (CallToolResult, usize, bool) {
443 let size = if max.is_some() {
444 Some(serialized_size(&result, max))
445 } else {
446 None
447 };
448 match decide_size(size, max) {
449 SizeVerdict::Pass { size } => (result, size, false),
450 SizeVerdict::PassUnmeasured => (result, 0, false),
451 SizeVerdict::Replace { limit, actual } => {
452 tracing::warn!(
453 tool = %tool,
454 size_bytes = actual.unwrap_or_default(),
455 size_measured = actual.is_some(),
456 limit_bytes = limit,
457 "tool result exceeds max_result_bytes; replacing with structured error"
458 );
459 let accounted = actual.unwrap_or_else(|| limit.saturating_add(1));
460 (too_large_result(limit, actual, tool), accounted, true)
461 }
462 }
463}
464
465#[allow(
466 deprecated,
467 reason = "transparent ServerHandler delegation must include legacy logging/subscription methods until rmcp removes them"
468)]
469impl<H: ServerHandler> ServerHandler for HookedHandler<H> {
470 async fn ping(&self, context: RequestContext<RoleServer>) -> Result<(), ErrorData> {
471 self.inner.ping(context).await
472 }
473
474 fn get_info(&self) -> ServerConfig {
475 self.inner.get_info()
476 }
477
478 async fn initialize(
479 &self,
480 request: InitializeRequestParams,
481 context: RequestContext<RoleServer>,
482 ) -> Result<InitializeResult, ErrorData> {
483 self.inner.initialize(request, context).await
484 }
485
486 fn negotiate_initialize(
490 &self,
491 request: &InitializeRequestParams,
492 ) -> Result<InitializeResult, ErrorData> {
493 self.inner.negotiate_initialize(request)
494 }
495
496 async fn list_tools(
497 &self,
498 request: Option<PaginatedRequestParams>,
499 context: RequestContext<RoleServer>,
500 ) -> Result<ListToolsResult, ErrorData> {
501 self.inner.list_tools(request, context).await
502 }
503
504 async fn complete(
505 &self,
506 request: CompleteRequestParams,
507 context: RequestContext<RoleServer>,
508 ) -> Result<CompleteResult, ErrorData> {
509 self.inner.complete(request, context).await
510 }
511
512 async fn set_level(
513 &self,
514 request: SetLevelRequestParams,
515 context: RequestContext<RoleServer>,
516 ) -> Result<(), ErrorData> {
517 self.inner.set_level(request, context).await
518 }
519
520 fn get_tool(&self, name: &str) -> Option<Tool> {
521 self.inner.get_tool(name)
522 }
523
524 async fn list_prompts(
525 &self,
526 request: Option<PaginatedRequestParams>,
527 context: RequestContext<RoleServer>,
528 ) -> Result<ListPromptsResult, ErrorData> {
529 self.inner.list_prompts(request, context).await
530 }
531
532 async fn get_prompt(
533 &self,
534 request: GetPromptRequestParams,
535 context: RequestContext<RoleServer>,
536 ) -> Result<GetPromptResponse, ErrorData> {
537 self.inner.get_prompt(request, context).await
538 }
539
540 async fn list_resources(
541 &self,
542 request: Option<PaginatedRequestParams>,
543 context: RequestContext<RoleServer>,
544 ) -> Result<ListResourcesResult, ErrorData> {
545 self.inner.list_resources(request, context).await
546 }
547
548 async fn list_resource_templates(
549 &self,
550 request: Option<PaginatedRequestParams>,
551 context: RequestContext<RoleServer>,
552 ) -> Result<ListResourceTemplatesResult, ErrorData> {
553 self.inner.list_resource_templates(request, context).await
554 }
555
556 async fn read_resource(
557 &self,
558 request: ReadResourceRequestParams,
559 context: RequestContext<RoleServer>,
560 ) -> Result<ReadResourceResponse, ErrorData> {
561 self.inner.read_resource(request, context).await
562 }
563
564 #[allow(
572 clippy::wildcard_enum_match_arm,
573 reason = "CallToolResponse is #[non_exhaustive]; the non-Complete MRTR variants (InputRequired/Task) are passed through unchanged"
574 )]
575 async fn call_tool(
576 &self,
577 request: CallToolRequestParams,
578 context: RequestContext<RoleServer>,
579 ) -> Result<CallToolResponse, ErrorData> {
580 let req_id = Some(format!("{:?}", context.id));
581 let ctx = Self::build_context(&request, req_id);
582 let max = self.hooks.max_result_bytes;
583 let after_holder = self
584 .hooks
585 .after
586 .as_ref()
587 .map(|f| Arc::new(AfterHookHolder { f: Arc::clone(f) }));
588
589 if let Some(before) = self.hooks.before.as_ref() {
591 let outcome = before(&ctx).await;
592 match outcome {
593 HookOutcome::Continue => {}
594 HookOutcome::Deny(err) => {
595 Self::spawn_after(after_holder.as_ref(), ctx, HookDisposition::DeniedBefore, 0);
596 return Err(err);
597 }
598 HookOutcome::Replace(boxed) => {
599 let (final_result, size, capped) = apply_size_cap(*boxed, max, &ctx.tool_name);
600 let disposition = if capped {
601 HookDisposition::ResultTooLarge
602 } else {
603 HookDisposition::ReplacedBefore
604 };
605 Self::spawn_after(after_holder.as_ref(), ctx, disposition, size);
606 return Ok(final_result.into());
607 }
608 }
609 }
610
611 match self.inner.call_tool(request, context).await {
613 Ok(CallToolResponse::Complete(result)) => {
615 let (final_result, size, capped) = apply_size_cap(result, max, &ctx.tool_name);
616 let disposition = if capped {
617 HookDisposition::ResultTooLarge
618 } else {
619 HookDisposition::InnerExecuted
620 };
621 Self::spawn_after(after_holder.as_ref(), ctx, disposition, size);
622 Ok(final_result.into())
623 }
624 Ok(other) => {
627 Self::spawn_after(
628 after_holder.as_ref(),
629 ctx,
630 HookDisposition::InnerExecuted,
631 0,
632 );
633 Ok(other)
634 }
635 Err(e) => {
636 Self::spawn_after(after_holder.as_ref(), ctx, HookDisposition::InnerErrored, 0);
637 Err(e)
638 }
639 }
640 }
641
642 fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> {
646 self.inner.supported_protocol_versions()
647 }
648
649 async fn discover(
650 &self,
651 context: RequestContext<RoleServer>,
652 ) -> Result<DiscoverResult, ErrorData> {
653 self.inner.discover(context).await
654 }
655
656 fn accepted_subscription_filter(
657 &self,
658 requested: &SubscriptionFilter,
659 ) -> Option<SubscriptionFilter> {
660 self.inner.accepted_subscription_filter(requested)
661 }
662
663 async fn listen(&self, context: SubscriptionContext) -> Result<(), ErrorData> {
664 self.inner.listen(context).await
665 }
666
667 async fn subscribe(
668 &self,
669 request: SubscribeRequestParams,
670 context: RequestContext<RoleServer>,
671 ) -> Result<(), ErrorData> {
672 self.inner.subscribe(request, context).await
673 }
674
675 async fn unsubscribe(
676 &self,
677 request: UnsubscribeRequestParams,
678 context: RequestContext<RoleServer>,
679 ) -> Result<(), ErrorData> {
680 self.inner.unsubscribe(request, context).await
681 }
682
683 async fn get_task(
684 &self,
685 request: GetTaskParams,
686 context: RequestContext<RoleServer>,
687 ) -> Result<GetTaskResult, ErrorData> {
688 self.inner.get_task(request, context).await
689 }
690
691 async fn update_task(
692 &self,
693 request: UpdateTaskParams,
694 context: RequestContext<RoleServer>,
695 ) -> Result<(), ErrorData> {
696 self.inner.update_task(request, context).await
697 }
698
699 async fn cancel_task(
700 &self,
701 request: CancelTaskParams,
702 context: RequestContext<RoleServer>,
703 ) -> Result<(), ErrorData> {
704 self.inner.cancel_task(request, context).await
705 }
706
707 async fn on_custom_request(
708 &self,
709 request: CustomRequest,
710 context: RequestContext<RoleServer>,
711 ) -> Result<CustomResult, ErrorData> {
712 self.inner.on_custom_request(request, context).await
713 }
714
715 async fn on_cancelled(
716 &self,
717 notification: CancelledNotificationParam,
718 context: NotificationContext<RoleServer>,
719 ) {
720 self.inner.on_cancelled(notification, context).await;
721 }
722
723 async fn on_progress(
724 &self,
725 notification: ProgressNotificationParam,
726 context: NotificationContext<RoleServer>,
727 ) {
728 self.inner.on_progress(notification, context).await;
729 }
730
731 async fn on_initialized(&self, context: NotificationContext<RoleServer>) {
732 self.inner.on_initialized(context).await;
733 }
734
735 async fn on_roots_list_changed(&self, context: NotificationContext<RoleServer>) {
736 self.inner.on_roots_list_changed(context).await;
737 }
738
739 async fn on_custom_notification(
740 &self,
741 notification: CustomNotification,
742 context: NotificationContext<RoleServer>,
743 ) {
744 self.inner
745 .on_custom_notification(notification, context)
746 .await;
747 }
748}
749
750#[derive(Debug, Clone, Copy, PartialEq, Eq)]
751struct SizeLimitExceeded;
752
753impl fmt::Display for SizeLimitExceeded {
754 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
755 f.write_str("serialized result exceeded configured size cap")
756 }
757}
758
759impl std::error::Error for SizeLimitExceeded {}
760
761struct CountingWriter {
762 bytes: usize,
763 limit: Option<usize>,
764}
765
766impl CountingWriter {
767 const fn unbounded() -> Self {
768 Self {
769 bytes: 0,
770 limit: None,
771 }
772 }
773
774 const fn bounded(limit: usize) -> Self {
775 Self {
776 bytes: 0,
777 limit: Some(limit),
778 }
779 }
780}
781
782impl io::Write for CountingWriter {
783 fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
784 let next = self.bytes.saturating_add(buf.len());
785 if self.limit.is_some_and(|limit| next > limit) {
786 Err(io::Error::other(SizeLimitExceeded))
787 } else {
788 self.bytes = next;
789 Ok(buf.len())
790 }
791 }
792
793 fn flush(&mut self) -> io::Result<()> {
794 Ok(())
795 }
796}
797
798#[derive(Debug, Clone, Copy, PartialEq, Eq)]
800enum SizeMeasure {
801 Exact(usize),
803 Exceeded { limit: usize },
805}
806
807fn serialized_size(result: &CallToolResult, max: Option<usize>) -> SizeMeasure {
809 let mut writer = max.map_or_else(CountingWriter::unbounded, CountingWriter::bounded);
810 match serde_json::to_writer(&mut writer, result) {
811 Ok(()) => SizeMeasure::Exact(writer.bytes),
812 Err(error) if error.io_error_kind() == Some(io::ErrorKind::Other) => {
813 SizeMeasure::Exceeded {
814 limit: max.unwrap_or(writer.bytes),
815 }
816 }
817 Err(_error) => {
818 SizeMeasure::Exact(writer.bytes)
822 }
823 }
824}
825
826#[cfg(test)]
827mod tests {
828 use std::sync::{
829 Arc,
830 atomic::{AtomicUsize, Ordering},
831 };
832
833 #[allow(
834 deprecated,
835 reason = "delegation tests cover legacy logging/subscription methods"
836 )]
837 use rmcp::{
838 ErrorData, RoleServer, ServerHandler,
839 model::{
840 CallToolRequestParams, CallToolResponse, CallToolResult, CancelledNotificationParam,
841 CompleteRequestParams, CompleteResult, CompletionInfo, ContentBlock,
842 CustomNotification, CustomRequest, CustomResult, DiscoverResult,
843 GetPromptRequestParams, GetPromptResult, GetTaskParams, GetTaskResult,
844 ListPromptsResult, ListResourceTemplatesResult, ListResourcesResult, ListToolsResult,
845 PaginatedRequestParams, ProgressNotificationParam, Prompt, PromptMessage,
846 ProtocolVersion, ReadResourceRequestParams, ReadResourceResult, Resource,
847 ResourceContents, ResourceTemplate, Role, ServerConfig, SetLevelRequestParams,
848 SubscribeRequestParams, SubscriptionFilter, UnsubscribeRequestParams, UpdateTaskParams,
849 },
850 service::{RequestContext, SubscriptionContext},
851 };
852 use serde_json::json;
853 use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader, DuplexStream};
854
855 use super::*;
856
857 type DelegationTransport = (
858 DelegationProbe,
859 BufReader<tokio::io::ReadHalf<DuplexStream>>,
860 tokio::io::WriteHalf<DuplexStream>,
861 rmcp::service::RunningService<RoleServer, HookedHandler<DelegationProbe>>,
862 );
863
864 #[derive(Clone, Default)]
865 struct CapturedLogs(Arc<std::sync::Mutex<Vec<u8>>>);
866
867 impl CapturedLogs {
868 fn contents(&self) -> String {
869 let bytes = self.0.lock().map(|guard| guard.clone()).unwrap_or_default();
870 String::from_utf8(bytes).unwrap_or_default()
871 }
872 }
873
874 struct CapturedLogsWriter(Arc<std::sync::Mutex<Vec<u8>>>);
875
876 impl io::Write for CapturedLogsWriter {
877 fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
878 if let Ok(mut guard) = self.0.lock() {
879 guard.extend_from_slice(buf);
880 }
881 Ok(buf.len())
882 }
883
884 fn flush(&mut self) -> io::Result<()> {
885 Ok(())
886 }
887 }
888
889 impl<'a> tracing_subscriber::fmt::MakeWriter<'a> for CapturedLogs {
890 type Writer = CapturedLogsWriter;
891
892 fn make_writer(&'a self) -> Self::Writer {
893 CapturedLogsWriter(Arc::clone(&self.0))
894 }
895 }
896
897 #[derive(Clone, Default)]
899 struct TestHandler {
900 body_bytes: Option<usize>,
902 }
903
904 impl ServerHandler for TestHandler {
905 fn get_info(&self) -> ServerConfig {
906 ServerConfig::default()
907 }
908
909 #[allow(
910 clippy::unused_async_trait_impl,
911 reason = "async is mandated by the rmcp ServerHandler trait signature; this test handler does not await"
912 )]
913 async fn call_tool(
914 &self,
915 _request: CallToolRequestParams,
916 _context: RequestContext<RoleServer>,
917 ) -> Result<CallToolResponse, ErrorData> {
918 let body = "x".repeat(self.body_bytes.unwrap_or(4));
919 Ok(CallToolResult::success(vec![ContentBlock::text(body)]).into())
920 }
921 }
922
923 #[derive(Clone, Default)]
924 struct DelegationProbe {
925 seen: Arc<std::sync::Mutex<Vec<&'static str>>>,
926 notify: Arc<tokio::sync::Notify>,
927 }
928
929 impl DelegationProbe {
930 fn record(&self, method: &'static str) {
931 if let Ok(mut seen) = self.seen.lock() {
932 seen.push(method);
933 }
934 self.notify.notify_waiters();
935 }
936
937 fn seen(&self) -> Vec<&'static str> {
938 self.seen
939 .lock()
940 .map(|seen| seen.clone())
941 .unwrap_or_default()
942 }
943
944 async fn wait_for_seen_count(&self, count: usize) {
945 tokio::time::timeout(std::time::Duration::from_secs(1), async {
946 while self.seen().len() < count {
947 self.notify.notified().await;
948 }
949 })
950 .await
951 .expect("delegated handler methods should be observed");
952 }
953 }
954
955 #[allow(
956 clippy::unused_async_trait_impl,
957 deprecated,
958 reason = "delegation tests cover rmcp async trait methods whose probe implementations return immediately"
959 )]
960 impl ServerHandler for DelegationProbe {
961 fn get_info(&self) -> ServerConfig {
962 ServerConfig::default()
963 }
964
965 async fn ping(&self, _context: RequestContext<RoleServer>) -> Result<(), ErrorData> {
966 self.record("ping");
967 Ok(())
968 }
969
970 async fn complete(
971 &self,
972 _request: CompleteRequestParams,
973 _context: RequestContext<RoleServer>,
974 ) -> Result<CompleteResult, ErrorData> {
975 self.record("complete");
976 let completion = CompletionInfo::with_all_values(vec!["delegated".to_owned()])
977 .expect("single completion is within rmcp max");
978 Ok(CompleteResult::new(completion))
979 }
980
981 async fn set_level(
982 &self,
983 _request: SetLevelRequestParams,
984 _context: RequestContext<RoleServer>,
985 ) -> Result<(), ErrorData> {
986 self.record("set_level");
987 Ok(())
988 }
989
990 async fn subscribe(
991 &self,
992 _request: SubscribeRequestParams,
993 _context: RequestContext<RoleServer>,
994 ) -> Result<(), ErrorData> {
995 self.record("subscribe");
996 Ok(())
997 }
998
999 async fn unsubscribe(
1000 &self,
1001 _request: UnsubscribeRequestParams,
1002 _context: RequestContext<RoleServer>,
1003 ) -> Result<(), ErrorData> {
1004 self.record("unsubscribe");
1005 Ok(())
1006 }
1007
1008 async fn call_tool(
1009 &self,
1010 _request: CallToolRequestParams,
1011 _context: RequestContext<RoleServer>,
1012 ) -> Result<CallToolResponse, ErrorData> {
1013 self.record("call_tool");
1014 Ok(CallToolResult::success(vec![ContentBlock::text("inner")]).into())
1015 }
1016
1017 async fn on_custom_request(
1018 &self,
1019 _request: CustomRequest,
1020 _context: RequestContext<RoleServer>,
1021 ) -> Result<CustomResult, ErrorData> {
1022 self.record("on_custom_request");
1023 Ok(CustomResult::new(json!({ "delegated": true })))
1024 }
1025
1026 async fn on_cancelled(
1027 &self,
1028 _notification: CancelledNotificationParam,
1029 _context: NotificationContext<RoleServer>,
1030 ) {
1031 self.record("on_cancelled");
1032 }
1033
1034 async fn on_progress(
1035 &self,
1036 _notification: ProgressNotificationParam,
1037 _context: NotificationContext<RoleServer>,
1038 ) {
1039 self.record("on_progress");
1040 }
1041
1042 async fn on_initialized(&self, _context: NotificationContext<RoleServer>) {
1043 self.record("on_initialized");
1044 }
1045
1046 async fn on_roots_list_changed(&self, _context: NotificationContext<RoleServer>) {
1047 self.record("on_roots_list_changed");
1048 }
1049
1050 async fn on_custom_notification(
1051 &self,
1052 _notification: CustomNotification,
1053 _context: NotificationContext<RoleServer>,
1054 ) {
1055 self.record("on_custom_notification");
1056 }
1057 }
1058
1059 #[derive(Clone, Default)]
1069 struct PassthroughDefaults<H> {
1070 #[allow(
1071 dead_code,
1072 reason = "deliberately never read: this type overrides nothing, so the probe must stay unreached"
1073 )]
1074 inner: H,
1075 }
1076
1077 impl<H: ServerHandler> PassthroughDefaults<H> {
1078 fn new(inner: H) -> Self {
1079 Self { inner }
1080 }
1081 }
1082
1083 impl<H: ServerHandler> ServerHandler for PassthroughDefaults<H> {}
1084
1085 const SENTINEL_VERSIONS: [ProtocolVersion; 2] =
1089 [ProtocolVersion::V_2025_11_25, ProtocolVersion::V_2026_07_28];
1090
1091 fn probe_capabilities() -> rmcp::model::ServerCapabilities {
1095 rmcp::model::ServerCapabilities::builder()
1096 .enable_prompts()
1097 .enable_resources()
1098 .enable_tools()
1099 .enable_tool_list_changed()
1100 .enable_tasks()
1101 .build()
1102 }
1103
1104 #[derive(Clone, Default)]
1114 struct ForwardingProbe {
1115 seen: Arc<std::sync::Mutex<Vec<&'static str>>>,
1116 }
1117
1118 impl ForwardingProbe {
1119 fn record(&self, method: &'static str) {
1120 if let Ok(mut seen) = self.seen.lock() {
1121 seen.push(method);
1122 }
1123 }
1124
1125 fn seen(&self) -> Vec<&'static str> {
1126 self.seen
1127 .lock()
1128 .map(|seen| seen.clone())
1129 .unwrap_or_default()
1130 }
1131 }
1132
1133 #[allow(
1134 clippy::unused_async_trait_impl,
1135 deprecated,
1136 reason = "coverage drives rmcp's async trait methods, whose probe bodies return immediately"
1137 )]
1138 impl ServerHandler for ForwardingProbe {
1139 fn get_info(&self) -> ServerConfig {
1140 let mut info = ServerConfig::new(probe_capabilities());
1141 info.instructions = Some("forwarding-probe".to_owned());
1142 info
1143 }
1144
1145 fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> {
1146 Cow::Borrowed(&SENTINEL_VERSIONS)
1147 }
1148
1149 fn get_tool(&self, name: &str) -> Option<Tool> {
1150 Some(Tool::new(
1151 name.to_owned(),
1152 "forwarding-probe",
1153 Arc::new(rmcp::model::JsonObject::default()),
1154 ))
1155 }
1156
1157 async fn initialize(
1158 &self,
1159 _request: InitializeRequestParams,
1160 _context: RequestContext<RoleServer>,
1161 ) -> Result<InitializeResult, ErrorData> {
1162 self.record("initialize");
1163 let mut info = InitializeResult::new(probe_capabilities());
1164 info.instructions = Some("forwarding-probe:initialize".to_owned());
1165 Ok(info)
1166 }
1167
1168 async fn discover(
1169 &self,
1170 _context: RequestContext<RoleServer>,
1171 ) -> Result<DiscoverResult, ErrorData> {
1172 self.record("discover");
1173 Ok(DiscoverResult::new(
1174 SENTINEL_VERSIONS.to_vec(),
1175 probe_capabilities(),
1176 ))
1177 }
1178
1179 async fn list_tools(
1180 &self,
1181 _request: Option<PaginatedRequestParams>,
1182 _context: RequestContext<RoleServer>,
1183 ) -> Result<ListToolsResult, ErrorData> {
1184 self.record("list_tools");
1185 Ok(ListToolsResult::with_all_items(vec![Tool::new(
1186 "sentinel-tool",
1187 "forwarding-probe",
1188 Arc::new(rmcp::model::JsonObject::default()),
1189 )]))
1190 }
1191
1192 async fn list_prompts(
1193 &self,
1194 _request: Option<PaginatedRequestParams>,
1195 _context: RequestContext<RoleServer>,
1196 ) -> Result<ListPromptsResult, ErrorData> {
1197 self.record("list_prompts");
1198 Ok(ListPromptsResult::with_all_items(vec![Prompt::new(
1199 "sentinel-prompt",
1200 Some("forwarding-probe"),
1201 None,
1202 )]))
1203 }
1204
1205 async fn list_resources(
1206 &self,
1207 _request: Option<PaginatedRequestParams>,
1208 _context: RequestContext<RoleServer>,
1209 ) -> Result<ListResourcesResult, ErrorData> {
1210 self.record("list_resources");
1211 Ok(ListResourcesResult::with_all_items(vec![Resource::new(
1212 "test://sentinel-resource",
1213 "sentinel-resource",
1214 )]))
1215 }
1216
1217 async fn list_resource_templates(
1218 &self,
1219 _request: Option<PaginatedRequestParams>,
1220 _context: RequestContext<RoleServer>,
1221 ) -> Result<ListResourceTemplatesResult, ErrorData> {
1222 self.record("list_resource_templates");
1223 Ok(ListResourceTemplatesResult::with_all_items(vec![
1224 ResourceTemplate::new("test://sentinel/{id}", "sentinel-template"),
1225 ]))
1226 }
1227
1228 async fn get_prompt(
1229 &self,
1230 _request: GetPromptRequestParams,
1231 _context: RequestContext<RoleServer>,
1232 ) -> Result<GetPromptResponse, ErrorData> {
1233 self.record("get_prompt");
1234 Ok(GetPromptResult::new(vec![PromptMessage::new_text(
1235 Role::User,
1236 "forwarding-probe:get_prompt",
1237 )])
1238 .into())
1239 }
1240
1241 async fn read_resource(
1242 &self,
1243 _request: ReadResourceRequestParams,
1244 _context: RequestContext<RoleServer>,
1245 ) -> Result<ReadResourceResponse, ErrorData> {
1246 self.record("read_resource");
1247 Ok(ReadResourceResult::new(vec![ResourceContents::text(
1248 "forwarding-probe:read_resource",
1249 "test://sentinel-resource",
1250 )])
1251 .into())
1252 }
1253
1254 fn accepted_subscription_filter(
1255 &self,
1256 requested: &SubscriptionFilter,
1257 ) -> Option<SubscriptionFilter> {
1258 self.record("accepted_subscription_filter");
1259 Some(requested.clone())
1260 }
1261
1262 async fn listen(&self, _context: SubscriptionContext) -> Result<(), ErrorData> {
1263 self.record("listen");
1264 Ok(())
1265 }
1266
1267 async fn get_task(
1268 &self,
1269 request: GetTaskParams,
1270 _context: RequestContext<RoleServer>,
1271 ) -> Result<GetTaskResult, ErrorData> {
1272 self.record("get_task");
1273 Ok(GetTaskResult::new(rmcp::model::DetailedTask::new(
1274 rmcp::model::Task::new(
1275 request.task_id,
1276 rmcp::model::TaskStatus::Working,
1277 "2026-01-01T00:00:00Z",
1278 "2026-01-01T00:00:00Z",
1279 ),
1280 rmcp::model::TaskPayload::Working,
1281 )))
1282 }
1283
1284 async fn update_task(
1285 &self,
1286 _request: UpdateTaskParams,
1287 _context: RequestContext<RoleServer>,
1288 ) -> Result<(), ErrorData> {
1289 self.record("update_task");
1290 Err(ErrorData::invalid_request(
1291 "forwarding-probe:update_task",
1292 None,
1293 ))
1294 }
1295
1296 async fn cancel_task(
1297 &self,
1298 _request: CancelTaskParams,
1299 _context: RequestContext<RoleServer>,
1300 ) -> Result<(), ErrorData> {
1301 self.record("cancel_task");
1302 Ok(())
1303 }
1304 }
1305
1306 fn delegation_transport(probe: DelegationProbe, hooks: Arc<ToolHooks>) -> DelegationTransport {
1307 let (client, server) = tokio::io::duplex(16 * 1024);
1308 let (client_read, client_write) = tokio::io::split(client);
1309 let service = rmcp::service::serve_directly::<RoleServer, _, _, io::Error, _>(
1310 with_hooks(probe.clone(), hooks),
1311 server,
1312 None,
1313 );
1314 (probe, BufReader::new(client_read), client_write, service)
1315 }
1316
1317 async fn send_json_rpc(
1318 writer: &mut tokio::io::WriteHalf<DuplexStream>,
1319 reader: &mut BufReader<tokio::io::ReadHalf<DuplexStream>>,
1320 request: serde_json::Value,
1321 ) -> serde_json::Value {
1322 writer
1323 .write_all(request.to_string().as_bytes())
1324 .await
1325 .expect("write request");
1326 writer.write_all(b"\n").await.expect("write newline");
1327 writer.flush().await.expect("flush request");
1328
1329 let mut line = String::new();
1330 reader.read_line(&mut line).await.expect("read response");
1331 serde_json::from_str(&line).expect("response is JSON")
1332 }
1333
1334 async fn send_notification(
1335 writer: &mut tokio::io::WriteHalf<DuplexStream>,
1336 notification: serde_json::Value,
1337 ) {
1338 writer
1339 .write_all(notification.to_string().as_bytes())
1340 .await
1341 .expect("write notification");
1342 writer.write_all(b"\n").await.expect("write newline");
1343 writer.flush().await.expect("flush notification");
1344 }
1345
1346 #[derive(Clone, Default)]
1351 struct NegotiateProbe;
1352
1353 impl ServerHandler for NegotiateProbe {
1354 fn get_info(&self) -> ServerConfig {
1355 ServerConfig::default()
1356 }
1357
1358 fn negotiate_initialize(
1359 &self,
1360 _request: &InitializeRequestParams,
1361 ) -> Result<InitializeResult, ErrorData> {
1362 let mut info = ServerConfig::new(rmcp::model::ServerCapabilities::default());
1363 info.instructions = Some("inner negotiate_initialize override".to_owned());
1364 Ok(info)
1365 }
1366 }
1367
1368 #[test]
1369 fn hooked_handler_preserves_inner_negotiate_initialize_override() {
1370 let handler = with_hooks(NegotiateProbe, Arc::new(ToolHooks::new()));
1371 let request = InitializeRequestParams::new(
1372 rmcp::model::ClientCapabilities::default(),
1373 rmcp::model::Implementation::new("delegation-test-client", "0.0.0"),
1374 );
1375
1376 let result = handler
1381 .negotiate_initialize(&request)
1382 .expect("direct negotiation must succeed");
1383
1384 assert_eq!(
1385 result.instructions.as_deref(),
1386 Some("inner negotiate_initialize override"),
1387 "wrapper must delegate to the inner `negotiate_initialize` override"
1388 );
1389 }
1390
1391 #[tokio::test]
1392 async fn hooked_handler_delegates_ping() {
1393 let (probe, mut reader, mut writer, _service) =
1394 delegation_transport(DelegationProbe::default(), Arc::new(ToolHooks::new()));
1395
1396 let response = send_json_rpc(
1397 &mut writer,
1398 &mut reader,
1399 json!({ "jsonrpc": "2.0", "id": 1, "method": "ping" }),
1400 )
1401 .await;
1402
1403 assert_eq!(response["result"], json!({}));
1404 assert_eq!(probe.seen(), vec!["ping"]);
1405 }
1406
1407 #[tokio::test]
1408 async fn hooked_handler_delegates_notifications() {
1409 let (probe, _reader, mut writer, _service) =
1410 delegation_transport(DelegationProbe::default(), Arc::new(ToolHooks::new()));
1411
1412 send_notification(
1413 &mut writer,
1414 json!({
1415 "jsonrpc": "2.0",
1416 "method": "notifications/cancelled",
1417 "params": { "requestId": 1, "reason": "test" }
1418 }),
1419 )
1420 .await;
1421 send_notification(
1422 &mut writer,
1423 json!({
1424 "jsonrpc": "2.0",
1425 "method": "notifications/progress",
1426 "params": { "progressToken": 1, "progress": 0.5 }
1427 }),
1428 )
1429 .await;
1430 send_notification(
1431 &mut writer,
1432 json!({ "jsonrpc": "2.0", "method": "notifications/initialized" }),
1433 )
1434 .await;
1435 send_notification(
1436 &mut writer,
1437 json!({ "jsonrpc": "2.0", "method": "notifications/roots/list_changed" }),
1438 )
1439 .await;
1440 send_notification(
1441 &mut writer,
1442 json!({ "jsonrpc": "2.0", "method": "notifications/custom/probe" }),
1443 )
1444 .await;
1445
1446 probe.wait_for_seen_count(5).await;
1447 assert_eq!(
1448 probe.seen(),
1449 vec![
1450 "on_cancelled",
1451 "on_progress",
1452 "on_initialized",
1453 "on_roots_list_changed",
1454 "on_custom_notification"
1455 ]
1456 );
1457 }
1458
1459 #[tokio::test]
1460 #[allow(
1461 deprecated,
1462 reason = "set_level is deprecated by rmcp but must delegate"
1463 )]
1464 async fn hooked_handler_delegates_completion_and_level() {
1465 let (probe, mut reader, mut writer, _service) =
1466 delegation_transport(DelegationProbe::default(), Arc::new(ToolHooks::new()));
1467
1468 let completion = send_json_rpc(
1469 &mut writer,
1470 &mut reader,
1471 json!({
1472 "jsonrpc": "2.0",
1473 "id": 1,
1474 "method": "completion/complete",
1475 "params": {
1476 "ref": { "type": "ref/prompt", "name": "prompt" },
1477 "argument": { "name": "arg", "value": "de" }
1478 }
1479 }),
1480 )
1481 .await;
1482 let level = send_json_rpc(
1483 &mut writer,
1484 &mut reader,
1485 json!({
1486 "jsonrpc": "2.0",
1487 "id": 2,
1488 "method": "logging/setLevel",
1489 "params": { "level": "debug" }
1490 }),
1491 )
1492 .await;
1493
1494 assert_eq!(
1495 completion["result"]["completion"]["values"],
1496 json!(["delegated"])
1497 );
1498 assert_eq!(level["result"], json!({}));
1499 assert_eq!(probe.seen(), vec!["complete", "set_level"]);
1500 }
1501
1502 #[tokio::test]
1503 #[allow(
1504 deprecated,
1505 reason = "subscribe/unsubscribe are deprecated by rmcp but must delegate"
1506 )]
1507 async fn hooked_handler_delegates_subscriptions() {
1508 let (probe, mut reader, mut writer, _service) =
1509 delegation_transport(DelegationProbe::default(), Arc::new(ToolHooks::new()));
1510
1511 let subscribe = send_json_rpc(
1512 &mut writer,
1513 &mut reader,
1514 json!({
1515 "jsonrpc": "2.0",
1516 "id": 1,
1517 "method": "resources/subscribe",
1518 "params": { "uri": "file:///tmp/a" }
1519 }),
1520 )
1521 .await;
1522 let unsubscribe = send_json_rpc(
1523 &mut writer,
1524 &mut reader,
1525 json!({
1526 "jsonrpc": "2.0",
1527 "id": 2,
1528 "method": "resources/unsubscribe",
1529 "params": { "uri": "file:///tmp/a" }
1530 }),
1531 )
1532 .await;
1533
1534 assert_eq!(subscribe["result"], json!({}));
1535 assert_eq!(unsubscribe["result"], json!({}));
1536 assert_eq!(probe.seen(), vec!["subscribe", "unsubscribe"]);
1537 }
1538
1539 #[tokio::test]
1540 async fn hooked_handler_delegates_custom_request() {
1541 let (probe, mut reader, mut writer, _service) =
1542 delegation_transport(DelegationProbe::default(), Arc::new(ToolHooks::new()));
1543
1544 let response = send_json_rpc(
1545 &mut writer,
1546 &mut reader,
1547 json!({
1548 "jsonrpc": "2.0",
1549 "id": 1,
1550 "method": "requests/custom/probe",
1551 "params": { "x": true }
1552 }),
1553 )
1554 .await;
1555
1556 assert_eq!(response["result"], json!({ "delegated": true }));
1557 assert_eq!(probe.seen(), vec!["on_custom_request"]);
1558 }
1559
1560 #[tokio::test]
1561 async fn hooked_handler_still_applies_hooks_to_call_tool() {
1562 let before_count = Arc::new(AtomicUsize::new(0));
1563 let before_seen = Arc::clone(&before_count);
1564 let before: BeforeHook = Arc::new(move |_ctx| {
1565 let before_seen = Arc::clone(&before_seen);
1566 Box::pin(async move {
1567 before_seen.fetch_add(1, Ordering::Relaxed);
1568 HookOutcome::Continue
1569 })
1570 });
1571 let after_count = Arc::new(AtomicUsize::new(0));
1572 let after_seen = Arc::clone(&after_count);
1573 let after_notify = Arc::new(tokio::sync::Notify::new());
1574 let after_notify_seen = Arc::clone(&after_notify);
1575 let after: AfterHook = Arc::new(move |_ctx, _disp, _size| {
1576 let after_seen = Arc::clone(&after_seen);
1577 let after_notify_seen = Arc::clone(&after_notify_seen);
1578 Box::pin(async move {
1579 after_seen.fetch_add(1, Ordering::Relaxed);
1580 after_notify_seen.notify_waiters();
1581 })
1582 });
1583 let hooks = Arc::new(
1584 ToolHooks::new()
1585 .with_before(before)
1586 .with_after(after)
1587 .with_max_result_bytes(1024),
1588 );
1589 let (probe, mut reader, mut writer, _service) =
1590 delegation_transport(DelegationProbe::default(), hooks);
1591
1592 let response = send_json_rpc(
1593 &mut writer,
1594 &mut reader,
1595 json!({
1596 "jsonrpc": "2.0",
1597 "id": 1,
1598 "method": "tools/call",
1599 "params": { "name": "probe", "arguments": {} }
1600 }),
1601 )
1602 .await;
1603
1604 tokio::time::timeout(std::time::Duration::from_secs(1), async {
1605 while after_count.load(Ordering::Relaxed) == 0 {
1606 after_notify.notified().await;
1607 }
1608 })
1609 .await
1610 .expect("after hook should run");
1611 assert_eq!(response["result"]["content"][0]["text"], "inner");
1612 assert_eq!(probe.seen(), vec!["call_tool"]);
1613 assert_eq!(before_count.load(Ordering::Relaxed), 1);
1614 assert_eq!(after_count.load(Ordering::Relaxed), 1);
1615 }
1616
1617 const SEMANTIC_DRIVERS: &[(&str, &str)] = &[
1640 ("ping", "hooked_handler_delegates_ping"),
1641 (
1642 "initialize",
1643 "hooked_handler_forwards_initialize_and_discover",
1644 ),
1645 (
1646 "negotiate_initialize",
1647 "hooked_handler_preserves_inner_negotiate_initialize_override",
1648 ),
1649 (
1650 "supported_protocol_versions",
1651 "hooked_handler_forwards_direct_sync_methods",
1652 ),
1653 (
1654 "discover",
1655 "hooked_handler_forwards_initialize_and_discover",
1656 ),
1657 ("complete", "hooked_handler_delegates_completion_and_level"),
1658 ("set_level", "hooked_handler_delegates_completion_and_level"),
1659 (
1660 "get_prompt",
1661 "hooked_handler_forwards_prompt_and_resource_reads",
1662 ),
1663 ("list_prompts", "hooked_handler_forwards_listing_methods"),
1664 ("list_resources", "hooked_handler_forwards_listing_methods"),
1665 (
1666 "list_resource_templates",
1667 "hooked_handler_forwards_listing_methods",
1668 ),
1669 (
1670 "read_resource",
1671 "hooked_handler_forwards_prompt_and_resource_reads",
1672 ),
1673 (
1674 "accepted_subscription_filter",
1675 "hooked_handler_forwards_subscription_lifecycle",
1676 ),
1677 ("listen", "hooked_handler_forwards_subscription_lifecycle"),
1678 ("subscribe", "hooked_handler_delegates_subscriptions"),
1679 ("unsubscribe", "hooked_handler_delegates_subscriptions"),
1680 (
1681 "call_tool",
1682 "hooked_handler_still_applies_hooks_to_call_tool",
1683 ),
1684 ("list_tools", "hooked_handler_forwards_listing_methods"),
1685 ("get_tool", "hooked_handler_forwards_direct_sync_methods"),
1686 (
1687 "on_custom_request",
1688 "hooked_handler_delegates_custom_request",
1689 ),
1690 ("on_cancelled", "hooked_handler_delegates_notifications"),
1691 ("on_progress", "hooked_handler_delegates_notifications"),
1692 ("on_initialized", "hooked_handler_delegates_notifications"),
1693 (
1694 "on_roots_list_changed",
1695 "hooked_handler_delegates_notifications",
1696 ),
1697 (
1698 "on_custom_notification",
1699 "hooked_handler_delegates_notifications",
1700 ),
1701 ("get_info", "hooked_handler_forwards_direct_sync_methods"),
1702 ("get_task", "hooked_handler_forwards_task_methods"),
1703 ("update_task", "hooked_handler_forwards_task_methods"),
1704 ("cancel_task", "hooked_handler_forwards_task_methods"),
1705 ];
1706
1707 #[test]
1712 fn semantic_drivers_table_is_well_formed() {
1713 assert!(!SEMANTIC_DRIVERS.is_empty());
1714 let mut names: Vec<&str> = SEMANTIC_DRIVERS.iter().map(|(name, _)| *name).collect();
1715 let total = names.len();
1716 names.sort_unstable();
1717 names.dedup();
1718 assert_eq!(names.len(), total, "duplicate method in SEMANTIC_DRIVERS");
1719 for (name, driver) in SEMANTIC_DRIVERS {
1720 assert!(!name.is_empty());
1721 assert!(!driver.is_empty());
1722 }
1723 }
1724
1725 fn coverage_meta() -> serde_json::Value {
1729 json!({
1730 "io.modelcontextprotocol/protocolVersion": "2026-07-28",
1731 "io.modelcontextprotocol/clientCapabilities": {
1732 "extensions": { "io.modelcontextprotocol/tasks": {} }
1733 }
1734 })
1735 }
1736
1737 fn coverage_request(id: i64, method: &str, mut params: serde_json::Value) -> serde_json::Value {
1739 let object = params
1740 .as_object_mut()
1741 .expect("coverage requests carry an object of params");
1742 object.insert("_meta".to_owned(), coverage_meta());
1743 json!({ "jsonrpc": "2.0", "id": id, "method": method, "params": params })
1744 }
1745
1746 async fn read_json_rpc(
1750 reader: &mut BufReader<tokio::io::ReadHalf<DuplexStream>>,
1751 ) -> serde_json::Value {
1752 let mut line = String::new();
1753 reader.read_line(&mut line).await.expect("read response");
1754 serde_json::from_str(&line).expect("response is JSON")
1755 }
1756
1757 type ForwardingTransport<P> = (
1758 BufReader<tokio::io::ReadHalf<DuplexStream>>,
1759 tokio::io::WriteHalf<DuplexStream>,
1760 rmcp::service::RunningService<RoleServer, HookedHandler<P>>,
1761 );
1762
1763 fn forwarding_transport<P: ServerHandler>(
1767 inner: P,
1768 hooks: Arc<ToolHooks>,
1769 ) -> ForwardingTransport<P> {
1770 let (client, server) = tokio::io::duplex(16 * 1024);
1771 let (client_read, client_write) = tokio::io::split(client);
1772 let service = rmcp::service::serve_directly::<RoleServer, _, _, io::Error, _>(
1773 with_hooks(inner, hooks),
1774 server,
1775 None,
1776 );
1777 (BufReader::new(client_read), client_write, service)
1778 }
1779
1780 fn coverage_hooks() -> Arc<ToolHooks> {
1781 Arc::new(ToolHooks::new())
1782 }
1783
1784 #[test]
1785 fn hooked_handler_forwards_direct_sync_methods() {
1786 let wrapper = with_hooks(ForwardingProbe::default(), coverage_hooks());
1790 let control = PassthroughDefaults::<ForwardingProbe>::default();
1791
1792 assert_eq!(
1794 wrapper.get_info().instructions.as_deref(),
1795 Some("forwarding-probe")
1796 );
1797 assert_eq!(control.get_info().instructions, None);
1798
1799 let tool = wrapper
1801 .get_tool("sentinel-tool")
1802 .expect("inner get_tool must reach the caller");
1803 assert_eq!(tool.name, "sentinel-tool");
1804 assert_eq!(control.get_tool("sentinel-tool"), None);
1805
1806 assert_eq!(
1808 wrapper.supported_protocol_versions().as_ref(),
1809 SENTINEL_VERSIONS.as_slice()
1810 );
1811 assert_ne!(
1812 control.supported_protocol_versions().as_ref(),
1813 SENTINEL_VERSIONS.as_slice()
1814 );
1815 }
1816
1817 #[tokio::test]
1818 async fn hooked_handler_forwards_initialize_and_discover() {
1819 let probe = ForwardingProbe::default();
1820 let (mut reader, mut writer, _service) =
1821 forwarding_transport(probe.clone(), coverage_hooks());
1822
1823 let initialize = send_json_rpc(
1824 &mut writer,
1825 &mut reader,
1826 json!({
1827 "jsonrpc": "2.0",
1828 "id": 1,
1829 "method": "initialize",
1830 "params": {
1831 "protocolVersion": "2025-11-25",
1832 "capabilities": {},
1833 "clientInfo": { "name": "coverage-driver", "version": "0.0.0" }
1834 }
1835 }),
1836 )
1837 .await;
1838 assert_eq!(
1839 initialize["result"]["instructions"],
1840 json!("forwarding-probe:initialize")
1841 );
1842
1843 let discover = send_json_rpc(
1844 &mut writer,
1845 &mut reader,
1846 coverage_request(2, "server/discover", json!({})),
1847 )
1848 .await;
1849 assert_eq!(
1850 discover["result"]["supportedVersions"],
1851 json!(["2025-11-25", "2026-07-28"])
1852 );
1853
1854 assert_eq!(probe.seen(), vec!["initialize", "discover"]);
1860
1861 let control = ForwardingProbe::default();
1864 let (mut reader, mut writer, _service) =
1865 forwarding_transport(PassthroughDefaults::new(control.clone()), coverage_hooks());
1866 let control_initialize = send_json_rpc(
1867 &mut writer,
1868 &mut reader,
1869 json!({
1870 "jsonrpc": "2.0",
1871 "id": 1,
1872 "method": "initialize",
1873 "params": {
1874 "protocolVersion": "2025-11-25",
1875 "capabilities": {},
1876 "clientInfo": { "name": "coverage-driver", "version": "0.0.0" }
1877 }
1878 }),
1879 )
1880 .await;
1881 let control_discover = send_json_rpc(
1882 &mut writer,
1883 &mut reader,
1884 coverage_request(2, "server/discover", json!({})),
1885 )
1886 .await;
1887 assert_ne!(
1888 control_initialize["result"]["instructions"],
1889 json!("forwarding-probe:initialize")
1890 );
1891 assert_ne!(
1892 control_discover["result"]["supportedVersions"],
1893 json!(["2025-11-25", "2026-07-28"])
1894 );
1895 assert_eq!(control.seen(), Vec::<&str>::new());
1896 }
1897
1898 #[tokio::test]
1899 async fn hooked_handler_forwards_listing_methods() {
1900 let probe = ForwardingProbe::default();
1903 let (mut reader, mut writer, _service) =
1904 forwarding_transport(probe.clone(), coverage_hooks());
1905
1906 let tools = send_json_rpc(
1907 &mut writer,
1908 &mut reader,
1909 coverage_request(1, "tools/list", json!({})),
1910 )
1911 .await;
1912 let prompts = send_json_rpc(
1913 &mut writer,
1914 &mut reader,
1915 coverage_request(2, "prompts/list", json!({})),
1916 )
1917 .await;
1918 let resources = send_json_rpc(
1919 &mut writer,
1920 &mut reader,
1921 coverage_request(3, "resources/list", json!({})),
1922 )
1923 .await;
1924 let templates = send_json_rpc(
1925 &mut writer,
1926 &mut reader,
1927 coverage_request(4, "resources/templates/list", json!({})),
1928 )
1929 .await;
1930
1931 assert_eq!(tools["result"]["tools"][0]["name"], json!("sentinel-tool"));
1932 assert_eq!(
1933 prompts["result"]["prompts"][0]["name"],
1934 json!("sentinel-prompt")
1935 );
1936 assert_eq!(
1937 resources["result"]["resources"][0]["uri"],
1938 json!("test://sentinel-resource")
1939 );
1940 assert_eq!(
1941 templates["result"]["resourceTemplates"][0]["uriTemplate"],
1942 json!("test://sentinel/{id}")
1943 );
1944 assert_eq!(
1945 probe.seen(),
1946 vec![
1947 "list_tools",
1948 "list_prompts",
1949 "list_resources",
1950 "list_resource_templates"
1951 ]
1952 );
1953
1954 let control = ForwardingProbe::default();
1957 let (mut reader, mut writer, _service) =
1958 forwarding_transport(PassthroughDefaults::new(control.clone()), coverage_hooks());
1959 let control_tools = send_json_rpc(
1960 &mut writer,
1961 &mut reader,
1962 coverage_request(1, "tools/list", json!({})),
1963 )
1964 .await;
1965 let control_prompts = send_json_rpc(
1966 &mut writer,
1967 &mut reader,
1968 coverage_request(2, "prompts/list", json!({})),
1969 )
1970 .await;
1971 let control_resources = send_json_rpc(
1972 &mut writer,
1973 &mut reader,
1974 coverage_request(3, "resources/list", json!({})),
1975 )
1976 .await;
1977 let control_templates = send_json_rpc(
1978 &mut writer,
1979 &mut reader,
1980 coverage_request(4, "resources/templates/list", json!({})),
1981 )
1982 .await;
1983 assert_eq!(control_tools["result"]["tools"], json!([]));
1984 assert_eq!(control_prompts["result"]["prompts"], json!([]));
1985 assert_eq!(control_resources["result"]["resources"], json!([]));
1986 assert_eq!(control_templates["result"]["resourceTemplates"], json!([]));
1987 assert_eq!(control.seen(), Vec::<&str>::new());
1988 }
1989
1990 #[tokio::test]
1991 async fn hooked_handler_forwards_prompt_and_resource_reads() {
1992 let probe = ForwardingProbe::default();
1995 let (mut reader, mut writer, _service) =
1996 forwarding_transport(probe.clone(), coverage_hooks());
1997
1998 let prompt = send_json_rpc(
1999 &mut writer,
2000 &mut reader,
2001 coverage_request(1, "prompts/get", json!({ "name": "sentinel-prompt" })),
2002 )
2003 .await;
2004 let resource = send_json_rpc(
2005 &mut writer,
2006 &mut reader,
2007 coverage_request(
2008 2,
2009 "resources/read",
2010 json!({ "uri": "test://sentinel-resource" }),
2011 ),
2012 )
2013 .await;
2014
2015 assert_eq!(
2016 prompt["result"]["messages"][0]["content"]["text"],
2017 json!("forwarding-probe:get_prompt")
2018 );
2019 assert_eq!(
2020 resource["result"]["contents"][0]["text"],
2021 json!("forwarding-probe:read_resource")
2022 );
2023 assert_eq!(probe.seen(), vec!["get_prompt", "read_resource"]);
2024
2025 let control = ForwardingProbe::default();
2028 let (mut reader, mut writer, _service) =
2029 forwarding_transport(PassthroughDefaults::new(control.clone()), coverage_hooks());
2030 let control_prompt = send_json_rpc(
2031 &mut writer,
2032 &mut reader,
2033 coverage_request(1, "prompts/get", json!({ "name": "sentinel-prompt" })),
2034 )
2035 .await;
2036 let control_resource = send_json_rpc(
2037 &mut writer,
2038 &mut reader,
2039 coverage_request(
2040 2,
2041 "resources/read",
2042 json!({ "uri": "test://sentinel-resource" }),
2043 ),
2044 )
2045 .await;
2046 assert_eq!(control_prompt["error"]["code"], json!(-32601));
2047 assert_eq!(control_resource["error"]["code"], json!(-32601));
2048 assert_eq!(control.seen(), Vec::<&str>::new());
2049 }
2050
2051 #[tokio::test]
2052 async fn hooked_handler_forwards_task_methods() {
2053 let probe = ForwardingProbe::default();
2057 let (mut reader, mut writer, _service) =
2058 forwarding_transport(probe.clone(), coverage_hooks());
2059
2060 let get = send_json_rpc(
2061 &mut writer,
2062 &mut reader,
2063 coverage_request(1, "tasks/get", json!({ "taskId": "raw-task" })),
2064 )
2065 .await;
2066 let update = send_json_rpc(
2067 &mut writer,
2068 &mut reader,
2069 coverage_request(
2070 2,
2071 "tasks/update",
2072 json!({ "taskId": "raw-task", "inputResponses": {} }),
2073 ),
2074 )
2075 .await;
2076 let cancel = send_json_rpc(
2077 &mut writer,
2078 &mut reader,
2079 coverage_request(3, "tasks/cancel", json!({ "taskId": "raw-task" })),
2080 )
2081 .await;
2082
2083 assert_eq!(get["result"]["taskId"], json!("raw-task"));
2086 assert_eq!(
2087 update["error"]["message"],
2088 json!("forwarding-probe:update_task")
2089 );
2090 assert_eq!(cancel["result"], json!({ "resultType": "complete" }));
2091 assert_eq!(probe.seen(), vec!["get_task", "update_task", "cancel_task"]);
2092
2093 let control = ForwardingProbe::default();
2096 let (mut reader, mut writer, _service) =
2097 forwarding_transport(PassthroughDefaults::new(control.clone()), coverage_hooks());
2098 let control_get = send_json_rpc(
2099 &mut writer,
2100 &mut reader,
2101 coverage_request(1, "tasks/get", json!({ "taskId": "raw-task" })),
2102 )
2103 .await;
2104 assert_eq!(control_get["error"]["code"], json!(-32601));
2105 assert_eq!(control.seen(), Vec::<&str>::new());
2106 }
2107
2108 #[tokio::test]
2109 async fn hooked_handler_forwards_subscription_lifecycle() {
2110 let probe = ForwardingProbe::default();
2111 let (mut reader, mut writer, _service) =
2112 forwarding_transport(probe.clone(), coverage_hooks());
2113
2114 let ack = send_json_rpc(
2117 &mut writer,
2118 &mut reader,
2119 coverage_request(
2120 1,
2121 "subscriptions/listen",
2122 json!({ "notifications": { "toolsListChanged": true } }),
2123 ),
2124 )
2125 .await;
2126 assert_eq!(
2127 ack["method"],
2128 json!("notifications/subscriptions/acknowledged")
2129 );
2130 let response = read_json_rpc(&mut reader).await;
2131 assert_eq!(response["id"], json!(1));
2132 assert_eq!(response["result"]["resultType"], json!("complete"));
2133
2134 assert_eq!(probe.seen(), vec!["accepted_subscription_filter", "listen"]);
2139
2140 let control = ForwardingProbe::default();
2143 let (mut reader, mut writer, _service) =
2144 forwarding_transport(PassthroughDefaults::new(control.clone()), coverage_hooks());
2145 let control_listen = send_json_rpc(
2146 &mut writer,
2147 &mut reader,
2148 coverage_request(
2149 1,
2150 "subscriptions/listen",
2151 json!({ "notifications": { "toolsListChanged": true } }),
2152 ),
2153 )
2154 .await;
2155 assert_eq!(control_listen["error"]["code"], json!(-32601));
2156 assert_eq!(control.seen(), Vec::<&str>::new());
2157 }
2158
2159 fn ctx(name: &str) -> ToolCallContext {
2160 ToolCallContext {
2161 tool_name: name.to_owned(),
2162 arguments: None,
2163 identity: None,
2164 role: None,
2165 sub: None,
2166 request_id: None,
2167 }
2168 }
2169
2170 fn sensitive_ctx() -> ToolCallContext {
2171 ToolCallContext {
2172 tool_name: "safe-tool-name".to_owned(),
2173 arguments: Some(serde_json::json!({ "password": "argument-secret" })),
2174 identity: Some("identity-secret".to_owned()),
2175 role: Some("role-secret".to_owned()),
2176 sub: Some("sub-secret".to_owned()),
2177 request_id: Some("request-id-visible".to_owned()),
2178 }
2179 }
2180
2181 #[test]
2182 fn tool_call_context_debug_redacts_sensitive_fields_by_default() {
2183 let _guard = crate::diagnostics::ExposureTestGuard::acquire();
2184 crate::diagnostics::set_diagnostic_exposure(
2185 &crate::diagnostics::DiagnosticExposure::default(),
2186 );
2187
2188 let rendered = format!("{:?}", sensitive_ctx());
2189
2190 assert!(rendered.contains("safe-tool-name"));
2191 assert!(rendered.contains("request-id-visible"));
2192 assert!(rendered.contains("[REDACTED]"));
2193 for secret in [
2194 "argument-secret",
2195 "identity-secret",
2196 "role-secret",
2197 "sub-secret",
2198 ] {
2199 assert!(
2200 !rendered.contains(secret),
2201 "ToolCallContext Debug must not contain {secret}: {rendered}"
2202 );
2203 }
2204 }
2205
2206 #[test]
2207 fn tool_call_context_debug_can_show_sensitive_fields_when_enabled() {
2208 let _guard = crate::diagnostics::ExposureTestGuard::acquire();
2209 crate::diagnostics::set_diagnostic_exposure(&crate::diagnostics::DiagnosticExposure {
2210 tool_call_arguments: true,
2211 ..crate::diagnostics::DiagnosticExposure::default()
2212 });
2213
2214 let rendered = format!("{:?}", sensitive_ctx());
2215
2216 for secret in [
2217 "argument-secret",
2218 "identity-secret",
2219 "role-secret",
2220 "sub-secret",
2221 ] {
2222 assert!(
2223 rendered.contains(secret),
2224 "ToolCallContext Debug must contain {secret} when enabled: {rendered}"
2225 );
2226 }
2227 }
2228
2229 #[tokio::test]
2230 async fn size_cap_replaces_oversized_result() {
2231 let inner = TestHandler {
2232 body_bytes: Some(8_192),
2233 };
2234 let hooks = Arc::new(ToolHooks {
2235 max_result_bytes: Some(256),
2236 before: None,
2237 after: None,
2238 });
2239 let hooked = with_hooks(inner, hooks);
2240
2241 let small = CallToolResult::success(vec![ContentBlock::text("ok".to_owned())]);
2242 assert!(exact_size(&small) < 256);
2243
2244 let big = CallToolResult::success(vec![ContentBlock::text("x".repeat(8_192))]);
2245 let size = exact_size(&big);
2246 assert!(size > 256);
2247
2248 let (replaced, accounted, capped) = apply_size_cap(big, Some(256), "whatever");
2249 assert!(capped);
2250 assert_eq!(accounted, 257);
2251 assert_eq!(replaced.is_error, Some(true));
2252 assert!(matches!(
2253 replaced.content.first(),
2254 Some(rmcp::model::ContentBlock::Text(t)) if t.text.contains("result_too_large")
2255 ));
2256
2257 let _ = hooked;
2259 }
2260
2261 fn exact_size(result: &CallToolResult) -> usize {
2262 match serialized_size(result, None) {
2263 SizeMeasure::Exact(size) => size,
2264 SizeMeasure::Exceeded { limit } => {
2265 panic!("unbounded measurement exceeded impossible limit {limit}");
2266 }
2267 }
2268 }
2269
2270 #[test]
2271 fn serialized_size_under_cap_is_exact() {
2272 let result = CallToolResult::success(vec![ContentBlock::text("ok".to_owned())]);
2273 let exact = serde_json::to_vec(&result).unwrap().len();
2274
2275 let measured = serialized_size(&result, Some(exact));
2276
2277 assert_eq!(measured, SizeMeasure::Exact(exact));
2278 }
2279
2280 #[test]
2281 fn serialized_size_over_cap_stops_with_exceeded() {
2282 let result = CallToolResult::success(vec![ContentBlock::text("x".repeat(8_192))]);
2283
2284 let measured = serialized_size(&result, Some(256));
2285
2286 assert_eq!(measured, SizeMeasure::Exceeded { limit: 256 });
2287 }
2288
2289 #[test]
2290 fn over_cap_replacement_does_not_log_serialization_failure() {
2291 let logs = CapturedLogs::default();
2292 let subscriber = tracing_subscriber::fmt()
2293 .with_max_level(tracing::Level::TRACE)
2294 .with_writer(logs.clone())
2295 .with_ansi(false)
2296 .without_time()
2297 .finish();
2298 let _guard = tracing::subscriber::set_default(subscriber);
2299 let result = CallToolResult::success(vec![ContentBlock::text("x".repeat(8_192))]);
2300
2301 let (_final_result, accounted, capped) = apply_size_cap(result, Some(256), "big_tool");
2302
2303 assert!(capped);
2304 assert_eq!(accounted, 257);
2305 assert!(
2306 logs.contents()
2307 .contains("tool result exceeds max_result_bytes")
2308 );
2309 assert!(
2310 !logs.contents().contains("failed to serialize"),
2311 "cap-abort must not be logged as serialization failure: {}",
2312 logs.contents()
2313 );
2314 }
2315
2316 #[test]
2317 fn disabled_result_cap_skips_measurement() {
2318 let result = CallToolResult::success(vec![ContentBlock::text("x".repeat(8_192))]);
2319
2320 let (_final_result, accounted, capped) = apply_size_cap(result, None, "uncapped_tool");
2321
2322 assert!(!capped);
2323 assert_eq!(accounted, 0);
2324 }
2325
2326 #[tokio::test]
2327 async fn before_hook_deny_builds_error() {
2328 let counter = Arc::new(AtomicUsize::new(0));
2329 let c = Arc::clone(&counter);
2330 let before: BeforeHook = Arc::new(move |ctx_ref| {
2331 let c = Arc::clone(&c);
2332 let name = ctx_ref.tool_name.clone();
2333 Box::pin(async move {
2334 c.fetch_add(1, Ordering::Relaxed);
2335 if name == "forbidden" {
2336 HookOutcome::Deny(ErrorData::invalid_request("nope", None))
2337 } else {
2338 HookOutcome::Continue
2339 }
2340 })
2341 });
2342
2343 let hooks = Arc::new(ToolHooks {
2344 max_result_bytes: None,
2345 before: Some(before),
2346 after: None,
2347 });
2348 let hooked = with_hooks(TestHandler::default(), hooks);
2349
2350 let bad_ctx = ctx("forbidden");
2351 let before_fn = hooked.hooks.before.as_ref().unwrap();
2352 let outcome = before_fn(&bad_ctx).await;
2353 assert!(matches!(outcome, HookOutcome::Deny(_)));
2354 assert_eq!(counter.load(Ordering::Relaxed), 1);
2355
2356 let ok_ctx = ctx("allowed");
2357 let outcome2 = before_fn(&ok_ctx).await;
2358 assert!(matches!(outcome2, HookOutcome::Continue));
2359 assert_eq!(counter.load(Ordering::Relaxed), 2);
2360 }
2361
2362 #[test]
2363 fn too_large_result_mentions_limit_and_actual() {
2364 let r = too_large_result(100, Some(500), "my_tool");
2365 let body = serde_json::to_string(&r).unwrap();
2366 assert!(body.contains("result_too_large"));
2367 assert!(body.contains("my_tool"));
2368 assert!(body.contains("100"));
2369 assert!(body.contains("500"));
2370 }
2371
2372 #[test]
2373 fn decide_size_truth_table() {
2374 assert_eq!(
2375 decide_size(Some(SizeMeasure::Exact(10)), Some(100)),
2376 SizeVerdict::Pass { size: 10 }
2377 );
2378 assert_eq!(
2379 decide_size(Some(SizeMeasure::Exact(100)), Some(100)),
2380 SizeVerdict::Pass { size: 100 },
2381 "cap is inclusive: size == limit passes"
2382 );
2383 assert_eq!(
2384 decide_size(Some(SizeMeasure::Exact(101)), Some(100)),
2385 SizeVerdict::Replace {
2386 limit: 100,
2387 actual: Some(101)
2388 }
2389 );
2390 assert_eq!(
2391 decide_size(Some(SizeMeasure::Exact(999)), None),
2392 SizeVerdict::Pass { size: 999 }
2393 );
2394 assert_eq!(
2395 decide_size(None, Some(100)),
2396 SizeVerdict::Replace {
2397 limit: 100,
2398 actual: None
2399 },
2400 "unmeasurable result must fail closed when a cap is configured"
2401 );
2402 assert_eq!(decide_size(None, None), SizeVerdict::PassUnmeasured);
2403 assert_eq!(
2404 decide_size(Some(SizeMeasure::Exceeded { limit: 100 }), Some(100)),
2405 SizeVerdict::Replace {
2406 limit: 100,
2407 actual: None
2408 },
2409 "cap-abort is not an exact measurement"
2410 );
2411 }
2412
2413 #[test]
2414 fn too_large_result_does_not_fabricate_a_size_when_unmeasurable() {
2415 let r = too_large_result(100, None, "my_tool");
2416 let body = serde_json::to_string(&r).unwrap();
2417 assert!(body.contains("result_too_large"));
2418 assert!(body.contains("unknown"));
2419 assert!(
2420 !body.contains("101"),
2421 "the over-limit accounting sentinel must not leak into the client payload"
2422 );
2423 }
2424
2425 #[tokio::test]
2426 async fn replace_outcome_skips_inner_and_returns_payload() {
2427 let before: BeforeHook = Arc::new(|_ctx| {
2430 Box::pin(async {
2431 HookOutcome::Replace(Box::new(CallToolResult::success(vec![ContentBlock::text(
2432 "from-replace".to_owned(),
2433 )])))
2434 })
2435 });
2436 let hooks = Arc::new(ToolHooks {
2437 max_result_bytes: None,
2438 before: Some(before),
2439 after: None,
2440 });
2441 let _hooked = with_hooks(TestHandler::default(), Arc::clone(&hooks));
2442
2443 let outcome = (hooks.before.as_ref().unwrap())(&ctx("any")).await;
2446 let HookOutcome::Replace(boxed) = outcome else {
2447 panic!("expected HookOutcome::Replace");
2448 };
2449 let (result, size, capped) = apply_size_cap(*boxed, None, "any");
2450 assert!(!capped);
2451 assert_eq!(size, 0);
2452 assert!(!result.is_error.unwrap_or(false));
2453 assert!(matches!(
2454 result.content.first(),
2455 Some(rmcp::model::ContentBlock::Text(t)) if t.text == "from-replace"
2456 ));
2457 }
2458
2459 #[tokio::test]
2460 async fn replace_outcome_subject_to_size_cap() {
2461 let huge = CallToolResult::success(vec![ContentBlock::text("y".repeat(8_192))]);
2465 let huge_size = serde_json::to_vec(&huge).unwrap().len();
2466 assert!(huge_size > 256);
2467
2468 let (final_result, accounted, capped) = apply_size_cap(huge, Some(256), "replaced_tool");
2469 assert!(capped);
2470 assert_eq!(accounted, 257);
2471 assert_eq!(final_result.is_error, Some(true));
2472 assert!(matches!(
2473 final_result.content.first(),
2474 Some(rmcp::model::ContentBlock::Text(t)) if t.text.contains("result_too_large")
2475 ));
2476 }
2477
2478 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
2479 async fn after_hook_fires_exactly_once_via_spawn() {
2480 let counter = Arc::new(AtomicUsize::new(0));
2484 let c = Arc::clone(&counter);
2485 let after: AfterHook = Arc::new(move |_ctx, _disp, _size| {
2486 let c = Arc::clone(&c);
2487 Box::pin(async move {
2488 c.fetch_add(1, Ordering::Relaxed);
2489 })
2490 });
2491 let holder = Arc::new(AfterHookHolder { f: after });
2492
2493 HookedHandler::<TestHandler>::spawn_after(
2494 Some(&holder),
2495 ctx("t"),
2496 HookDisposition::InnerExecuted,
2497 42,
2498 );
2499
2500 let deadline = std::time::Instant::now() + std::time::Duration::from_secs(1);
2502 while counter.load(Ordering::Relaxed) == 0 && std::time::Instant::now() < deadline {
2503 tokio::task::yield_now().await;
2504 tokio::time::sleep(std::time::Duration::from_millis(5)).await;
2505 }
2506 assert_eq!(counter.load(Ordering::Relaxed), 1);
2507 }
2508
2509 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
2510 async fn after_hook_panic_is_isolated_from_response_path() {
2511 let after: AfterHook = Arc::new(|_ctx, _disp, _size| {
2515 Box::pin(async {
2516 panic!("intentional panic in after-hook");
2517 })
2518 });
2519 let holder = Arc::new(AfterHookHolder { f: after });
2520
2521 HookedHandler::<TestHandler>::spawn_after(
2522 Some(&holder),
2523 ctx("boom"),
2524 HookDisposition::InnerExecuted,
2525 0,
2526 );
2527
2528 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
2531 let still_alive = tokio::spawn(async { 1_u32 + 2 }).await.unwrap();
2532 assert_eq!(still_alive, 3);
2533 }
2534}