Skip to main content

rmcp_server_kit/
tool_hooks.rs

1//! Opt-in tool-call instrumentation for `ServerHandler` implementations.
2//!
3//! [`crate::tool_hooks::HookedHandler`] wraps any [`rmcp::ServerHandler`] with:
4//!
5//! - **Before hooks** (async) that observe `(tool_name, arguments, identity,
6//!   role, sub, request_id)` and may [`HookOutcome::Continue`](crate::tool_hooks::HookOutcome::Continue),
7//!   [`HookOutcome::Deny`](crate::tool_hooks::HookOutcome::Deny), or
8//!   [`HookOutcome::Replace`](crate::tool_hooks::HookOutcome::Replace) the call.
9//! - **After hooks** (async) that observe the same context plus a
10//!   [`HookDisposition`](crate::tool_hooks::HookDisposition) describing how the call resolved and the
11//!   approximate result size in bytes.  After-hooks are spawned via
12//!   `tokio::spawn` and never block the response path.
13//! - **Result-size capping**: serialized tool results larger than
14//!   `max_result_bytes` are replaced with a structured error, preventing
15//!   token-expensive or memory-expensive payloads from reaching clients.
16//!   The cap applies both to inner-handler results and to
17//!   [`HookOutcome::Replace`](crate::tool_hooks::HookOutcome::Replace) payloads.
18//!
19//! # Cancel safety
20//!
21//! The transparent `ServerHandler` delegation methods are cancel-safe with
22//! respect to this wrapper: cancellation only drops the delegated inner
23//! future, so the wrapped handler's own cancel-safety contract is inherited
24//! unchanged.  The `call_tool` implementation on
25//! [`crate::tool_hooks::HookedHandler`] is the exception and is documented
26//! as **NOT cancel-safe** at its definition: once a before-hook or the
27//! inner handler has been awaited, cancellation can prevent the paired
28//! after-hook from being spawned.
29//!
30//! This is entirely **opt-in** at the application layer - `rmcp_server_kit::serve()`
31//! does not wrap handlers automatically.  Applications that want hooks do:
32//!
33//! ```no_run
34//! use std::sync::Arc;
35//! use rmcp_server_kit::tool_hooks::{HookedHandler, HookOutcome, ToolHooks, with_hooks};
36//!
37//! # #[derive(Clone, Default)]
38//! # struct MyHandler;
39//! # impl rmcp::ServerHandler for MyHandler {}
40//! let handler = MyHandler::default();
41//! let hooks = Arc::new(
42//!     ToolHooks::new()
43//!         .with_max_result_bytes(256 * 1024)
44//!         .with_before(Arc::new(|_ctx| Box::pin(async { HookOutcome::Continue })))
45//!         .with_after(Arc::new(|_ctx, _disp, _bytes| Box::pin(async {}))),
46//! );
47//! let _wrapped = with_hooks(handler, hooks);
48//! ```
49
50use 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/// Context passed to before/after hooks for a single tool call.
73#[derive(Clone)]
74#[non_exhaustive]
75pub struct ToolCallContext {
76    /// Tool name being invoked.
77    pub tool_name: String,
78    /// JSON arguments as sent by the client (may be `None`).
79    pub arguments: Option<serde_json::Value>,
80    /// Identity name from the authenticated request, if any.
81    pub identity: Option<String>,
82    /// RBAC role associated with the request, if any.
83    pub role: Option<String>,
84    /// OAuth `sub` claim, if present.
85    pub sub: Option<String>,
86    /// Raw JSON-RPC request id rendered as a string, if available.
87    pub request_id: Option<String>,
88}
89
90impl ToolCallContext {
91    /// Construct a [`ToolCallContext`] with the given tool name and all
92    /// optional fields cleared.  Primarily for use in unit tests and
93    /// benchmarks of user-supplied hooks; the runtime path populates
94    /// these fields from the request and task-local RBAC state.
95    #[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/// Outcome returned by a [`BeforeHook`] to control invocation flow.
138///
139/// - [`HookOutcome::Continue`] - proceed with the wrapped handler.
140/// - [`HookOutcome::Deny`] - reject the call with the supplied
141///   [`ErrorData`]; the inner handler is **not** called.
142/// - [`HookOutcome::Replace`] - return the supplied result instead of
143///   invoking the inner handler.  The result is still subject to
144///   `max_result_bytes` capping.
145#[derive(Debug)]
146#[non_exhaustive]
147pub enum HookOutcome {
148    /// Proceed with the wrapped handler.
149    Continue,
150    /// Reject the call.  The error is propagated to the client as-is.
151    Deny(ErrorData),
152    /// Skip the inner handler and return the supplied result instead.
153    Replace(Box<CallToolResult>),
154}
155
156/// How a tool call resolved, passed to the [`AfterHook`].
157#[derive(Debug, Clone, Copy)]
158#[non_exhaustive]
159pub enum HookDisposition {
160    /// The inner handler ran and returned `Ok`.
161    InnerExecuted,
162    /// The inner handler ran and returned `Err`.
163    InnerErrored,
164    /// The before-hook returned [`HookOutcome::Deny`].
165    DeniedBefore,
166    /// The before-hook returned [`HookOutcome::Replace`].
167    ReplacedBefore,
168    /// The result (from inner or replace) exceeded `max_result_bytes`
169    /// and was substituted with a structured error.
170    ResultTooLarge,
171}
172
173/// Async before-hook callback type.
174///
175/// Returns a [`HookOutcome`] controlling whether the inner handler runs.
176/// The borrow of `ToolCallContext` is held for the duration of the
177/// returned future, which avoids forcing implementations to clone the
178/// context for every invocation.
179pub 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
186/// Async after-hook callback type.
187///
188/// Receives the call context, a [`HookDisposition`] describing how the
189/// call resolved, and the approximate serialized result size in bytes
190/// (`0` for `DeniedBefore` and `InnerErrored`).  Spawned via
191/// `tokio::spawn`, so it must not assume it runs before the response is
192/// flushed.
193pub 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/// Opt-in hooks applied by [`crate::tool_hooks::HookedHandler`].
205#[allow(clippy::struct_field_names, reason = "before/after read naturally")]
206#[derive(Clone, Default)]
207#[non_exhaustive]
208pub struct ToolHooks {
209    /// Hard cap on serialized `CallToolResult` size in bytes.  When
210    /// exceeded, the result is replaced with an `is_error=true` result
211    /// carrying a `result_too_large` structured error.  `None` disables
212    /// the cap.
213    pub max_result_bytes: Option<usize>,
214    /// Optional before-hook invoked after arg deserialization, before
215    /// the wrapped handler is called.
216    pub before: Option<BeforeHook>,
217    /// Optional after-hook invoked once per normally-resolved call - that
218    /// is, on the Deny / Replace / Ok / Err paths.  Spawned via
219    /// `tokio::spawn` and never blocks the response path.
220    ///
221    /// **Not guaranteed under cancellation.**  If the `call_tool` future is
222    /// dropped after a before-hook has run but before the call resolves,
223    /// the paired after-hook is never spawned.  Do not use before/after
224    /// pairing as a mandatory resource guard; make the after-hook
225    /// idempotent or tolerant of missing closes (see [`crate::cancel`]).
226    pub after: Option<AfterHook>,
227}
228
229impl ToolHooks {
230    /// Construct an empty [`ToolHooks`] with no cap and no hooks.
231    ///
232    /// Use the `with_*` builder methods to populate fields; this avoids
233    /// the `#[non_exhaustive]` restriction that prevents struct-literal
234    /// construction from outside the crate.
235    #[must_use]
236    pub fn new() -> Self {
237        Self::default()
238    }
239
240    /// Set the serialized result size cap in bytes.
241    #[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    /// Set the before-hook.
248    #[must_use]
249    pub fn with_before(mut self, before: BeforeHook) -> Self {
250        self.before = Some(before);
251        self
252    }
253
254    /// Set the after-hook.
255    #[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/// `ServerHandler` wrapper that applies [`ToolHooks`].
275#[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/// Construct a [`crate::tool_hooks::HookedHandler`] from an inner handler and hooks.
290///
291/// Returning the wrapped handler is the entire point of this function;
292/// dropping it on the floor would silently disable the supplied hooks.
293#[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    /// Access the wrapped handler.
305    #[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    /// Spawn the after-hook on the current Tokio runtime.  The future
322    /// captures clones of `ctx` and the `Arc<AfterHook>` so it can run
323    /// independently of the request task; panics inside the after-hook
324    /// are caught by Tokio and never poison the response path.
325    ///
326    /// The spawned task is **instrumented** with the request span via
327    /// [`tracing::Instrument`] and re-establishes the per-request RBAC
328    /// task-locals (role, identity, token, sub) via
329    /// [`crate::rbac::with_rbac_scope`]. Without this, after-hooks lose
330    /// their parent span (breaking trace correlation) and observe
331    /// `current_role()` / `current_identity()` as `None`.
332    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            // Capture the request span before leaving the request task so
343            // after-hook log lines are correlated with the originating call.
344            let span = tracing::Span::current();
345            // Snapshot RBAC task-locals; defaults are empty strings so the
346            // re-established scope is a no-op when the request had no
347            // authenticated identity (e.g. health checks, anonymous tools).
348            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
367/// Internal newtype that owns the [`AfterHook`] so we can `Arc::clone`
368/// the *holder* and let the spawned task borrow `ctx` for the lifetime
369/// of the future without lifetime acrobatics in `tokio::spawn`.
370struct AfterHookHolder {
371    f: AfterHook,
372}
373
374/// Structured error body returned when a result exceeds `max_result_bytes`.
375///
376/// `actual` is `None` when the result could not be serialized, so its true
377/// size is unknown. It is rendered as `"unknown"` rather than a fabricated
378/// number -- operators read `actual_bytes` as a measurement.
379fn 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/// Outcome of the `max_result_bytes` policy for a measured -- or
400/// unmeasurable -- result.
401#[derive(Debug, PartialEq, Eq)]
402enum SizeVerdict {
403    /// Within the cap, or no cap configured. Carries the measured size.
404    Pass { size: usize },
405    /// Over the cap, or unmeasurable while a cap is configured.
406    Replace { limit: usize, actual: Option<usize> },
407    /// Unmeasurable and no cap configured: nothing to enforce.
408    PassUnmeasured,
409}
410
411/// Decide what the size cap does, given an optional size-measurement outcome.
412const 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
435/// Apply the `max_result_bytes` cap to a result.  Returns the (possibly
436/// replaced) result, the size used for accounting, and whether the cap
437/// fired.
438fn 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    // Synchronous negotiation helper: no context and no hooks apply, so plain
487    // delegation is the only transparent option -- the trait default would
488    // shadow an inner override.
489    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    // NOT cancel-safe: this awaits consumer-supplied before-hooks and the
565    // consumer's inner handler. After-hooks are dispatched only on the normal
566    // Deny/Replace/Ok/Err paths, so a cancellation between the before-hook and
567    // the response drops the paired after-hook -- an audit hook can therefore
568    // record a started call that is never closed out. Consumers needing
569    // guaranteed pairing should make the after-hook idempotent or run the tool
570    // body detached (see `crate::cancel`).
571    #[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        // Before hook: may Continue, Deny, or Replace.
590        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        // Inner handler.
612        match self.inner.call_tool(request, context).await {
613            // Completed tool result: subject to the size cap + after hook.
614            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            // MRTR input-required / task responses (rmcp 3.0): no CallToolResult
625            // to size-cap, so pass them through unchanged.
626            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    // rmcp 3.0 added task/subscription/discovery request handlers with defaults;
643    // delegate them to `inner` so wrapping a handler that implements those stays
644    // transparent (otherwise the default would shadow the inner implementation).
645    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/// Outcome of measuring serialized result size.
799#[derive(Debug, Clone, Copy, PartialEq, Eq)]
800enum SizeMeasure {
801    /// Exact serialized size in bytes.
802    Exact(usize),
803    /// Serialization crossed the configured size cap and stopped early.
804    Exceeded { limit: usize },
805}
806
807/// Serialized byte length, or a deliberate cap-abort outcome.
808fn 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            // `CallToolResult` is made only of infallibly serializable fields
819            // (`String`, `bool`, arrays/maps, and serde_json::Value`). There is
820            // no inhabitable production value that can reach this branch.
821            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    /// Minimal in-process `ServerHandler` for tests.
898    #[derive(Clone, Default)]
899    struct TestHandler {
900        /// When Some, `call_tool` returns a body of this many 'x' bytes.
901        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    /// The wrapper as it would be if every delegation were deleted: it holds the
1060    /// probe but overrides nothing, so rmcp's default `ServerHandler` bodies
1061    /// answer every call and the inner handler is never reached.
1062    ///
1063    /// Driving a method against this type is the per-method mutation check: if
1064    /// the inner probe still observes the call, the observation came from the
1065    /// harness (or a re-entrant default) rather than from the wrapper's
1066    /// delegation. `inner` is deliberately never read, hence the `dead_code`
1067    /// allow.
1068    #[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    /// Two of rmcp's five known versions: distinguishable from the default
1086    /// (all of `KNOWN_VERSIONS`), while still covering an initialize-capable
1087    /// version and the 2026-07-28 version that `discover` requires.
1088    const SENTINEL_VERSIONS: [ProtocolVersion; 2] =
1089        [ProtocolVersion::V_2025_11_25, ProtocolVersion::V_2026_07_28];
1090
1091    /// Capabilities the probe advertises so the dispatcher routes prompts,
1092    /// resources, tools, tool-list-changed subscriptions and tasks to the
1093    /// wrapper at all.
1094    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    /// Records every method it is called with and answers with a sentinel no
1105    /// rmcp default body can construct, so exact-log equality proves the
1106    /// wrapper forwarded the call.
1107    ///
1108    /// Deliberately silent (no record) for `get_info`,
1109    /// `supported_protocol_versions` and `get_tool`: rmcp calls those outside
1110    /// request dispatch (peer configuration, capability validation), so a
1111    /// record could not be attributed to the driver. Those three are proven by
1112    /// value differential against [`PassthroughDefaults`] instead.
1113    #[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    /// Overrides only `negotiate_initialize`, so the sentinel value can only
1347    /// come from an inner override reaching the caller through the wrapper.
1348    /// `get_info` stays at the test default and `initialize` at the upstream
1349    /// default, so no other path can produce the sentinel.
1350    #[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        // Called directly on the wrapper (the HTTP path routes `initialize`,
1377        // which is already delegated); the trait default here would negotiate
1378        // from `get_info()` + `supported_protocol_versions()`, losing the
1379        // override.
1380        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    // ----------------------------------------------------------------------
1618    // Semantic coverage drivers
1619    //
1620    // Each driver proves, by exact equality on the inner probe's record log
1621    // (and on sentinel values no rmcp default body can produce), that the
1622    // wrapper forwarded the call. Every driver runs twice: once against the
1623    // real wrapper, and once against `PassthroughDefaults` -- the wrapper with
1624    // all delegations deleted -- where the probe must stay untouched. That
1625    // second half is the per-method mutation check.
1626    //
1627    // No assertion here uses `contains`, `is_empty()` or a length threshold:
1628    // rmcp's self-re-entrant defaults (`initialize`, `negotiate_initialize`,
1629    // `discover`) and its ambient handler calls can otherwise make a
1630    // non-forwarding body look green.
1631    // ----------------------------------------------------------------------
1632
1633    /// Maps every method the `HookedHandler` `ServerHandler` impl defines to the
1634    /// test that proves the inner handler was reached.
1635    ///
1636    /// Parsed by `tests/delegation_guard.rs`, which asserts this table covers
1637    /// exactly its `DIRECTLY_DELEGATED ∪ BEHAVIORAL_WRAPPED` classification --
1638    /// so a method cannot be added or reclassified without a driver.
1639    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    /// `SEMANTIC_DRIVERS` is parsed by `tests/delegation_guard.rs`; this
1708    /// in-crate check keeps the constant referenced (so it cannot rot as dead
1709    /// code) and rejects duplicate entries, which the source-level parser
1710    /// cannot see.
1711    #[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    /// Per-request metadata satisfying every gate the drivers cross: a protocol
1726    /// version inside the probe's narrowed list, and client capabilities
1727    /// declaring the tasks extension that `tasks/*` requires.
1728    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    /// A request whose `params._meta` carries [`coverage_meta`].
1738    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    /// Read one server frame that is already queued (used where a single
1747    /// request produces two frames: the subscription acknowledgement, then the
1748    /// response).
1749    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    /// Like [`delegation_transport`], but for a caller-chosen inner handler, so
1764    /// one driver can run against both the real wrapper and
1765    /// [`PassthroughDefaults`].
1766    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        // No record log in this driver: `get_info`, `get_tool` and
1787        // `supported_protocol_versions` are proven by value differential, so
1788        // rmcp's ambient calls cannot contaminate the proof.
1789        let wrapper = with_hooks(ForwardingProbe::default(), coverage_hooks());
1790        let control = PassthroughDefaults::<ForwardingProbe>::default();
1791
1792        // `get_info`: sentinel instructions no default body can produce.
1793        assert_eq!(
1794            wrapper.get_info().instructions.as_deref(),
1795            Some("forwarding-probe")
1796        );
1797        assert_eq!(control.get_info().instructions, None);
1798
1799        // `get_tool`: the default returns `None` unconditionally.
1800        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        // `supported_protocol_versions`: the default is all of KNOWN_VERSIONS.
1807        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        // Exact full sequence (mechanism 3): the two driven methods and nothing
1855        // else. Neither can be produced by a re-entrant default -- the probe
1856        // overrides `initialize` outright, and `discover` is only reachable
1857        // through the wrapper's delegation -- so the expected log contains both
1858        // driven methods, per the ambient-calls invariant.
1859        assert_eq!(probe.seen(), vec!["initialize", "discover"]);
1860
1861        // Negative control: with every delegation deleted, both requests are
1862        // answered by rmcp defaults and the probe stays untouched.
1863        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        // Mechanism 1 (ambient-silent probe): the expected log is exactly the
1901        // driven methods.
1902        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        // Negative control: the same four requests, answered by defaults --
1955        // empty result sets, probe untouched.
1956        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        // Mechanism 1 (ambient-silent probe): the expected log is exactly the
1993        // driven methods.
1994        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        // Negative control: both defaults reject with method-not-found, so the
2026        // client sees an error and the probe is never entered.
2027        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        // Mechanism 1 (ambient-silent probe): rmcp calls `get_info` for the
2054        // tasks-capability gate, which this probe does not record, so the
2055        // expected log is exactly the driven methods.
2056        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        // `DetailedTask` inlines the base task fields at the top level of the
2084        // result, so the echoed sentinel id sits at `result.taskId`.
2085        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        // Negative control: the tasks capability gate rejects both requests
2094        // before dispatch (the passthrough advertises no `tasks` extension).
2095        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        // Two frames come back: the acknowledgement is emitted before `listen`
2115        // runs, the response after it returns.
2116        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        // Exact full sequence (mechanism 3): this arm calls
2135        // `accepted_subscription_filter` before `listen`, and both are the
2136        // wrapper's own delegations -- so the expected log is the two-element
2137        // ordered sequence, not the singleton `["listen"]`.
2138        assert_eq!(probe.seen(), vec!["accepted_subscription_filter", "listen"]);
2139
2140        // Negative control: the default filter hook returns `None`, so rmcp
2141        // rejects the request before `listen` and the probe stays untouched.
2142        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        // Compile-check that HookedHandler instantiates with the test inner.
2258        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        // Returning Replace from before-hook must yield the supplied
2428        // CallToolResult directly, with no need for the inner handler.
2429        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        // Exercise the before-hook closure + apply_size_cap helper directly,
2444        // matching the established test pattern in this module.
2445        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        // A Replace payload that exceeds max_result_bytes must be rewritten
2462        // to result_too_large just like an inner-handler result would be,
2463        // and the disposition must reflect ResultTooLarge.
2464        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        // spawn_after must enqueue the after-hook exactly one time per
2481        // invocation and never block the caller; we wait for the spawned
2482        // task to run by polling the counter with a short timeout.
2483        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        // Wait up to 1s for the spawned task to run.
2501        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        // A panicking after-hook must not affect the request task.  We
2512        // spawn a panicking after-hook and then verify the current task
2513        // can still complete an unrelated future to completion.
2514        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        // Give Tokio a chance to run + abort the panicking task, then
2529        // confirm we're still alive and the runtime is healthy.
2530        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}