use std::{borrow::Cow, fmt, future::Future, io, pin::Pin, sync::Arc};
use rmcp::{
ErrorData, RoleServer, ServerHandler,
model::{
CallToolRequestParams, CallToolResponse, CallToolResult, CancelTaskParams, ContentBlock,
DiscoverResult, GetPromptRequestParams, GetPromptResponse, GetTaskParams, GetTaskResult,
InitializeRequestParams, InitializeResult, ListPromptsResult, ListResourceTemplatesResult,
ListResourcesResult, ListToolsResult, PaginatedRequestParams, ProtocolVersion,
ReadResourceRequestParams, ReadResourceResponse, ServerInfo, SubscriptionFilter, Tool,
UpdateTaskParams,
},
service::{RequestContext, SubscriptionContext},
};
#[derive(Clone)]
#[non_exhaustive]
pub struct ToolCallContext {
pub tool_name: String,
pub arguments: Option<serde_json::Value>,
pub identity: Option<String>,
pub role: Option<String>,
pub sub: Option<String>,
pub request_id: Option<String>,
}
impl ToolCallContext {
#[must_use]
pub fn for_tool(tool_name: impl Into<String>) -> Self {
Self {
tool_name: tool_name.into(),
arguments: None,
identity: None,
role: None,
sub: None,
request_id: None,
}
}
}
impl fmt::Debug for ToolCallContext {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let Self {
tool_name,
arguments,
identity,
role,
sub,
request_id,
} = self;
let mut debug = f.debug_struct("ToolCallContext");
debug.field("tool_name", tool_name);
if crate::diagnostics::tool_call_arguments() {
debug
.field("arguments", arguments)
.field("identity", identity)
.field("role", role)
.field("sub", sub);
} else {
debug
.field("arguments", &"[REDACTED]")
.field("identity", &"[REDACTED]")
.field("role", &"[REDACTED]")
.field("sub", &"[REDACTED]");
}
debug.field("request_id", request_id).finish()
}
}
#[derive(Debug)]
#[non_exhaustive]
pub enum HookOutcome {
Continue,
Deny(ErrorData),
Replace(Box<CallToolResult>),
}
#[derive(Debug, Clone, Copy)]
#[non_exhaustive]
pub enum HookDisposition {
InnerExecuted,
InnerErrored,
DeniedBefore,
ReplacedBefore,
ResultTooLarge,
}
pub type BeforeHook = Arc<
dyn for<'a> Fn(&'a ToolCallContext) -> Pin<Box<dyn Future<Output = HookOutcome> + Send + 'a>>
+ Send
+ Sync
+ 'static,
>;
pub type AfterHook = Arc<
dyn for<'a> Fn(
&'a ToolCallContext,
HookDisposition,
usize,
) -> Pin<Box<dyn Future<Output = ()> + Send + 'a>>
+ Send
+ Sync
+ 'static,
>;
#[allow(clippy::struct_field_names, reason = "before/after read naturally")]
#[derive(Clone, Default)]
#[non_exhaustive]
pub struct ToolHooks {
pub max_result_bytes: Option<usize>,
pub before: Option<BeforeHook>,
pub after: Option<AfterHook>,
}
impl ToolHooks {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn with_max_result_bytes(mut self, max: usize) -> Self {
self.max_result_bytes = Some(max);
self
}
#[must_use]
pub fn with_before(mut self, before: BeforeHook) -> Self {
self.before = Some(before);
self
}
#[must_use]
pub fn with_after(mut self, after: AfterHook) -> Self {
self.after = Some(after);
self
}
}
const _HOOKED_HANDLER_DOC_ANCHOR: &str = "HookedHandler";
impl fmt::Debug for ToolHooks {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ToolHooks")
.field("max_result_bytes", &self.max_result_bytes)
.field("before", &self.before.as_ref().map(|_| "<fn>"))
.field("after", &self.after.as_ref().map(|_| "<fn>"))
.finish()
}
}
#[derive(Clone)]
pub struct HookedHandler<H: ServerHandler> {
inner: Arc<H>,
hooks: Arc<ToolHooks>,
}
impl<H: ServerHandler> fmt::Debug for HookedHandler<H> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("HookedHandler")
.field("hooks", &self.hooks)
.finish_non_exhaustive()
}
}
#[must_use = "HookedHandler must be wired into a ServerHandler (e.g. via \
`serve(..., || hooked)`) to take effect; dropping the returned \
value silently disables the supplied hooks"]
pub fn with_hooks<H: ServerHandler>(inner: H, hooks: Arc<ToolHooks>) -> HookedHandler<H> {
HookedHandler {
inner: Arc::new(inner),
hooks,
}
}
impl<H: ServerHandler> HookedHandler<H> {
#[must_use]
pub fn inner(&self) -> &H {
&self.inner
}
fn build_context(request: &CallToolRequestParams, req_id: Option<String>) -> ToolCallContext {
ToolCallContext {
tool_name: request.name.to_string(),
arguments: request.arguments.clone().map(serde_json::Value::Object),
identity: crate::rbac::current_identity(),
role: crate::rbac::current_role(),
sub: crate::rbac::current_sub(),
request_id: req_id,
}
}
fn spawn_after(
after: Option<&Arc<AfterHookHolder>>,
ctx: ToolCallContext,
disposition: HookDisposition,
size: usize,
) {
if let Some(after) = after {
use tracing::Instrument;
let after = Arc::clone(after);
let span = tracing::Span::current();
let role = crate::rbac::current_role().unwrap_or_default();
let identity = crate::rbac::current_identity().unwrap_or_default();
let token = crate::rbac::current_token()
.unwrap_or_else(|| secrecy::SecretString::from(String::new()));
let sub = crate::rbac::current_sub().unwrap_or_default();
tokio::spawn(
async move {
crate::rbac::with_rbac_scope(role, identity, token, sub, async move {
let fut = (after.f)(&ctx, disposition, size);
fut.await;
})
.await;
}
.instrument(span),
);
}
}
}
struct AfterHookHolder {
f: AfterHook,
}
fn too_large_result(limit: usize, actual: Option<usize>, tool: &str) -> CallToolResult {
let actual_desc =
actual.map_or_else(|| "an unmeasurable number of".to_owned(), |n| n.to_string());
let body = serde_json::json!({
"error": "result_too_large",
"message": format!(
"tool '{tool}' result of {actual_desc} bytes exceeds the configured \
max_result_bytes={limit}; ask for a narrower query"
),
"limit_bytes": limit,
"actual_bytes": actual.map_or_else(
|| serde_json::Value::from("unknown"),
serde_json::Value::from,
),
});
let mut r = CallToolResult::error(vec![ContentBlock::text(body.to_string())]);
r.structured_content = None;
r
}
#[derive(Debug, PartialEq, Eq)]
enum SizeVerdict {
Pass { size: usize },
Replace { limit: usize, actual: Option<usize> },
PassUnmeasured,
}
const fn decide_size(size: Option<SizeMeasure>, max: Option<usize>) -> SizeVerdict {
match size {
Some(SizeMeasure::Exact(size)) => match max {
Some(limit) if size > limit => SizeVerdict::Replace {
limit,
actual: Some(size),
},
Some(_) | None => SizeVerdict::Pass { size },
},
Some(SizeMeasure::Exceeded { limit }) => SizeVerdict::Replace {
limit,
actual: None,
},
None => match max {
Some(limit) => SizeVerdict::Replace {
limit,
actual: None,
},
None => SizeVerdict::PassUnmeasured,
},
}
}
fn apply_size_cap(
result: CallToolResult,
max: Option<usize>,
tool: &str,
) -> (CallToolResult, usize, bool) {
let size = if max.is_some() {
Some(serialized_size(&result, max))
} else {
None
};
match decide_size(size, max) {
SizeVerdict::Pass { size } => (result, size, false),
SizeVerdict::PassUnmeasured => (result, 0, false),
SizeVerdict::Replace { limit, actual } => {
tracing::warn!(
tool = %tool,
size_bytes = actual.unwrap_or_default(),
size_measured = actual.is_some(),
limit_bytes = limit,
"tool result exceeds max_result_bytes; replacing with structured error"
);
let accounted = actual.unwrap_or_else(|| limit.saturating_add(1));
(too_large_result(limit, actual, tool), accounted, true)
}
}
}
impl<H: ServerHandler> ServerHandler for HookedHandler<H> {
fn get_info(&self) -> ServerInfo {
self.inner.get_info()
}
async fn initialize(
&self,
request: InitializeRequestParams,
context: RequestContext<RoleServer>,
) -> Result<InitializeResult, ErrorData> {
self.inner.initialize(request, context).await
}
async fn list_tools(
&self,
request: Option<PaginatedRequestParams>,
context: RequestContext<RoleServer>,
) -> Result<ListToolsResult, ErrorData> {
self.inner.list_tools(request, context).await
}
fn get_tool(&self, name: &str) -> Option<Tool> {
self.inner.get_tool(name)
}
async fn list_prompts(
&self,
request: Option<PaginatedRequestParams>,
context: RequestContext<RoleServer>,
) -> Result<ListPromptsResult, ErrorData> {
self.inner.list_prompts(request, context).await
}
async fn get_prompt(
&self,
request: GetPromptRequestParams,
context: RequestContext<RoleServer>,
) -> Result<GetPromptResponse, ErrorData> {
self.inner.get_prompt(request, context).await
}
async fn list_resources(
&self,
request: Option<PaginatedRequestParams>,
context: RequestContext<RoleServer>,
) -> Result<ListResourcesResult, ErrorData> {
self.inner.list_resources(request, context).await
}
async fn list_resource_templates(
&self,
request: Option<PaginatedRequestParams>,
context: RequestContext<RoleServer>,
) -> Result<ListResourceTemplatesResult, ErrorData> {
self.inner.list_resource_templates(request, context).await
}
async fn read_resource(
&self,
request: ReadResourceRequestParams,
context: RequestContext<RoleServer>,
) -> Result<ReadResourceResponse, ErrorData> {
self.inner.read_resource(request, context).await
}
#[allow(
clippy::wildcard_enum_match_arm,
reason = "CallToolResponse is #[non_exhaustive]; the non-Complete MRTR variants (InputRequired/Task) are passed through unchanged"
)]
async fn call_tool(
&self,
request: CallToolRequestParams,
context: RequestContext<RoleServer>,
) -> Result<CallToolResponse, ErrorData> {
let req_id = Some(format!("{:?}", context.id));
let ctx = Self::build_context(&request, req_id);
let max = self.hooks.max_result_bytes;
let after_holder = self
.hooks
.after
.as_ref()
.map(|f| Arc::new(AfterHookHolder { f: Arc::clone(f) }));
if let Some(before) = self.hooks.before.as_ref() {
let outcome = before(&ctx).await;
match outcome {
HookOutcome::Continue => {}
HookOutcome::Deny(err) => {
Self::spawn_after(after_holder.as_ref(), ctx, HookDisposition::DeniedBefore, 0);
return Err(err);
}
HookOutcome::Replace(boxed) => {
let (final_result, size, capped) = apply_size_cap(*boxed, max, &ctx.tool_name);
let disposition = if capped {
HookDisposition::ResultTooLarge
} else {
HookDisposition::ReplacedBefore
};
Self::spawn_after(after_holder.as_ref(), ctx, disposition, size);
return Ok(final_result.into());
}
}
}
match self.inner.call_tool(request, context).await {
Ok(CallToolResponse::Complete(result)) => {
let (final_result, size, capped) = apply_size_cap(result, max, &ctx.tool_name);
let disposition = if capped {
HookDisposition::ResultTooLarge
} else {
HookDisposition::InnerExecuted
};
Self::spawn_after(after_holder.as_ref(), ctx, disposition, size);
Ok(final_result.into())
}
Ok(other) => {
Self::spawn_after(
after_holder.as_ref(),
ctx,
HookDisposition::InnerExecuted,
0,
);
Ok(other)
}
Err(e) => {
Self::spawn_after(after_holder.as_ref(), ctx, HookDisposition::InnerErrored, 0);
Err(e)
}
}
}
fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> {
self.inner.supported_protocol_versions()
}
async fn discover(
&self,
context: RequestContext<RoleServer>,
) -> Result<DiscoverResult, ErrorData> {
self.inner.discover(context).await
}
fn accepted_subscription_filter(
&self,
requested: &SubscriptionFilter,
) -> Option<SubscriptionFilter> {
self.inner.accepted_subscription_filter(requested)
}
async fn listen(&self, context: SubscriptionContext) -> Result<(), ErrorData> {
self.inner.listen(context).await
}
async fn get_task(
&self,
request: GetTaskParams,
context: RequestContext<RoleServer>,
) -> Result<GetTaskResult, ErrorData> {
self.inner.get_task(request, context).await
}
async fn update_task(
&self,
request: UpdateTaskParams,
context: RequestContext<RoleServer>,
) -> Result<(), ErrorData> {
self.inner.update_task(request, context).await
}
async fn cancel_task(
&self,
request: CancelTaskParams,
context: RequestContext<RoleServer>,
) -> Result<(), ErrorData> {
self.inner.cancel_task(request, context).await
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct SizeLimitExceeded;
impl fmt::Display for SizeLimitExceeded {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("serialized result exceeded configured size cap")
}
}
impl std::error::Error for SizeLimitExceeded {}
struct CountingWriter {
bytes: usize,
limit: Option<usize>,
}
impl CountingWriter {
const fn unbounded() -> Self {
Self {
bytes: 0,
limit: None,
}
}
const fn bounded(limit: usize) -> Self {
Self {
bytes: 0,
limit: Some(limit),
}
}
}
impl io::Write for CountingWriter {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
let next = self.bytes.saturating_add(buf.len());
if self.limit.is_some_and(|limit| next > limit) {
Err(io::Error::other(SizeLimitExceeded))
} else {
self.bytes = next;
Ok(buf.len())
}
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum SizeMeasure {
Exact(usize),
Exceeded { limit: usize },
}
fn serialized_size(result: &CallToolResult, max: Option<usize>) -> SizeMeasure {
let mut writer = max.map_or_else(CountingWriter::unbounded, CountingWriter::bounded);
match serde_json::to_writer(&mut writer, result) {
Ok(()) => SizeMeasure::Exact(writer.bytes),
Err(error) if error.io_error_kind() == Some(io::ErrorKind::Other) => {
SizeMeasure::Exceeded {
limit: max.unwrap_or(writer.bytes),
}
}
Err(_error) => {
SizeMeasure::Exact(writer.bytes)
}
}
}
#[cfg(test)]
mod tests {
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};
use rmcp::{
ErrorData, RoleServer, ServerHandler,
model::{
CallToolRequestParams, CallToolResponse, CallToolResult, ContentBlock, ServerInfo,
},
service::RequestContext,
};
use super::*;
#[derive(Clone, Default)]
struct CapturedLogs(Arc<std::sync::Mutex<Vec<u8>>>);
impl CapturedLogs {
fn contents(&self) -> String {
let bytes = self.0.lock().map(|guard| guard.clone()).unwrap_or_default();
String::from_utf8(bytes).unwrap_or_default()
}
}
struct CapturedLogsWriter(Arc<std::sync::Mutex<Vec<u8>>>);
impl io::Write for CapturedLogsWriter {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
if let Ok(mut guard) = self.0.lock() {
guard.extend_from_slice(buf);
}
Ok(buf.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
impl<'a> tracing_subscriber::fmt::MakeWriter<'a> for CapturedLogs {
type Writer = CapturedLogsWriter;
fn make_writer(&'a self) -> Self::Writer {
CapturedLogsWriter(Arc::clone(&self.0))
}
}
#[derive(Clone, Default)]
struct TestHandler {
body_bytes: Option<usize>,
}
impl ServerHandler for TestHandler {
fn get_info(&self) -> ServerInfo {
ServerInfo::default()
}
#[allow(
clippy::unused_async_trait_impl,
reason = "async is mandated by the rmcp ServerHandler trait signature; this test handler does not await"
)]
async fn call_tool(
&self,
_request: CallToolRequestParams,
_context: RequestContext<RoleServer>,
) -> Result<CallToolResponse, ErrorData> {
let body = "x".repeat(self.body_bytes.unwrap_or(4));
Ok(CallToolResult::success(vec![ContentBlock::text(body)]).into())
}
}
fn ctx(name: &str) -> ToolCallContext {
ToolCallContext {
tool_name: name.to_owned(),
arguments: None,
identity: None,
role: None,
sub: None,
request_id: None,
}
}
fn sensitive_ctx() -> ToolCallContext {
ToolCallContext {
tool_name: "safe-tool-name".to_owned(),
arguments: Some(serde_json::json!({ "password": "argument-secret" })),
identity: Some("identity-secret".to_owned()),
role: Some("role-secret".to_owned()),
sub: Some("sub-secret".to_owned()),
request_id: Some("request-id-visible".to_owned()),
}
}
#[test]
fn tool_call_context_debug_redacts_sensitive_fields_by_default() {
let _guard = crate::diagnostics::ExposureTestGuard::acquire();
crate::diagnostics::set_diagnostic_exposure(
&crate::diagnostics::DiagnosticExposure::default(),
);
let rendered = format!("{:?}", sensitive_ctx());
assert!(rendered.contains("safe-tool-name"));
assert!(rendered.contains("request-id-visible"));
assert!(rendered.contains("[REDACTED]"));
for secret in [
"argument-secret",
"identity-secret",
"role-secret",
"sub-secret",
] {
assert!(
!rendered.contains(secret),
"ToolCallContext Debug must not contain {secret}: {rendered}"
);
}
}
#[test]
fn tool_call_context_debug_can_show_sensitive_fields_when_enabled() {
let _guard = crate::diagnostics::ExposureTestGuard::acquire();
crate::diagnostics::set_diagnostic_exposure(&crate::diagnostics::DiagnosticExposure {
tool_call_arguments: true,
..crate::diagnostics::DiagnosticExposure::default()
});
let rendered = format!("{:?}", sensitive_ctx());
for secret in [
"argument-secret",
"identity-secret",
"role-secret",
"sub-secret",
] {
assert!(
rendered.contains(secret),
"ToolCallContext Debug must contain {secret} when enabled: {rendered}"
);
}
}
#[tokio::test]
async fn size_cap_replaces_oversized_result() {
let inner = TestHandler {
body_bytes: Some(8_192),
};
let hooks = Arc::new(ToolHooks {
max_result_bytes: Some(256),
before: None,
after: None,
});
let hooked = with_hooks(inner, hooks);
let small = CallToolResult::success(vec![ContentBlock::text("ok".to_owned())]);
assert!(exact_size(&small) < 256);
let big = CallToolResult::success(vec![ContentBlock::text("x".repeat(8_192))]);
let size = exact_size(&big);
assert!(size > 256);
let (replaced, accounted, capped) = apply_size_cap(big, Some(256), "whatever");
assert!(capped);
assert_eq!(accounted, 257);
assert_eq!(replaced.is_error, Some(true));
assert!(matches!(
replaced.content.first(),
Some(rmcp::model::ContentBlock::Text(t)) if t.text.contains("result_too_large")
));
let _ = hooked;
}
fn exact_size(result: &CallToolResult) -> usize {
match serialized_size(result, None) {
SizeMeasure::Exact(size) => size,
SizeMeasure::Exceeded { limit } => {
panic!("unbounded measurement exceeded impossible limit {limit}");
}
}
}
#[test]
fn serialized_size_under_cap_is_exact() {
let result = CallToolResult::success(vec![ContentBlock::text("ok".to_owned())]);
let exact = serde_json::to_vec(&result).unwrap().len();
let measured = serialized_size(&result, Some(exact));
assert_eq!(measured, SizeMeasure::Exact(exact));
}
#[test]
fn serialized_size_over_cap_stops_with_exceeded() {
let result = CallToolResult::success(vec![ContentBlock::text("x".repeat(8_192))]);
let measured = serialized_size(&result, Some(256));
assert_eq!(measured, SizeMeasure::Exceeded { limit: 256 });
}
#[test]
fn over_cap_replacement_does_not_log_serialization_failure() {
let logs = CapturedLogs::default();
let subscriber = tracing_subscriber::fmt()
.with_max_level(tracing::Level::TRACE)
.with_writer(logs.clone())
.with_ansi(false)
.without_time()
.finish();
let _guard = tracing::subscriber::set_default(subscriber);
let result = CallToolResult::success(vec![ContentBlock::text("x".repeat(8_192))]);
let (_final_result, accounted, capped) = apply_size_cap(result, Some(256), "big_tool");
assert!(capped);
assert_eq!(accounted, 257);
assert!(
logs.contents()
.contains("tool result exceeds max_result_bytes")
);
assert!(
!logs.contents().contains("failed to serialize"),
"cap-abort must not be logged as serialization failure: {}",
logs.contents()
);
}
#[test]
fn disabled_result_cap_skips_measurement() {
let result = CallToolResult::success(vec![ContentBlock::text("x".repeat(8_192))]);
let (_final_result, accounted, capped) = apply_size_cap(result, None, "uncapped_tool");
assert!(!capped);
assert_eq!(accounted, 0);
}
#[tokio::test]
async fn before_hook_deny_builds_error() {
let counter = Arc::new(AtomicUsize::new(0));
let c = Arc::clone(&counter);
let before: BeforeHook = Arc::new(move |ctx_ref| {
let c = Arc::clone(&c);
let name = ctx_ref.tool_name.clone();
Box::pin(async move {
c.fetch_add(1, Ordering::Relaxed);
if name == "forbidden" {
HookOutcome::Deny(ErrorData::invalid_request("nope", None))
} else {
HookOutcome::Continue
}
})
});
let hooks = Arc::new(ToolHooks {
max_result_bytes: None,
before: Some(before),
after: None,
});
let hooked = with_hooks(TestHandler::default(), hooks);
let bad_ctx = ctx("forbidden");
let before_fn = hooked.hooks.before.as_ref().unwrap();
let outcome = before_fn(&bad_ctx).await;
assert!(matches!(outcome, HookOutcome::Deny(_)));
assert_eq!(counter.load(Ordering::Relaxed), 1);
let ok_ctx = ctx("allowed");
let outcome2 = before_fn(&ok_ctx).await;
assert!(matches!(outcome2, HookOutcome::Continue));
assert_eq!(counter.load(Ordering::Relaxed), 2);
}
#[test]
fn too_large_result_mentions_limit_and_actual() {
let r = too_large_result(100, Some(500), "my_tool");
let body = serde_json::to_string(&r).unwrap();
assert!(body.contains("result_too_large"));
assert!(body.contains("my_tool"));
assert!(body.contains("100"));
assert!(body.contains("500"));
}
#[test]
fn decide_size_truth_table() {
assert_eq!(
decide_size(Some(SizeMeasure::Exact(10)), Some(100)),
SizeVerdict::Pass { size: 10 }
);
assert_eq!(
decide_size(Some(SizeMeasure::Exact(100)), Some(100)),
SizeVerdict::Pass { size: 100 },
"cap is inclusive: size == limit passes"
);
assert_eq!(
decide_size(Some(SizeMeasure::Exact(101)), Some(100)),
SizeVerdict::Replace {
limit: 100,
actual: Some(101)
}
);
assert_eq!(
decide_size(Some(SizeMeasure::Exact(999)), None),
SizeVerdict::Pass { size: 999 }
);
assert_eq!(
decide_size(None, Some(100)),
SizeVerdict::Replace {
limit: 100,
actual: None
},
"unmeasurable result must fail closed when a cap is configured"
);
assert_eq!(decide_size(None, None), SizeVerdict::PassUnmeasured);
assert_eq!(
decide_size(Some(SizeMeasure::Exceeded { limit: 100 }), Some(100)),
SizeVerdict::Replace {
limit: 100,
actual: None
},
"cap-abort is not an exact measurement"
);
}
#[test]
fn too_large_result_does_not_fabricate_a_size_when_unmeasurable() {
let r = too_large_result(100, None, "my_tool");
let body = serde_json::to_string(&r).unwrap();
assert!(body.contains("result_too_large"));
assert!(body.contains("unknown"));
assert!(
!body.contains("101"),
"the over-limit accounting sentinel must not leak into the client payload"
);
}
#[tokio::test]
async fn replace_outcome_skips_inner_and_returns_payload() {
let before: BeforeHook = Arc::new(|_ctx| {
Box::pin(async {
HookOutcome::Replace(Box::new(CallToolResult::success(vec![ContentBlock::text(
"from-replace".to_owned(),
)])))
})
});
let hooks = Arc::new(ToolHooks {
max_result_bytes: None,
before: Some(before),
after: None,
});
let _hooked = with_hooks(TestHandler::default(), Arc::clone(&hooks));
let outcome = (hooks.before.as_ref().unwrap())(&ctx("any")).await;
let HookOutcome::Replace(boxed) = outcome else {
panic!("expected HookOutcome::Replace");
};
let (result, size, capped) = apply_size_cap(*boxed, None, "any");
assert!(!capped);
assert_eq!(size, 0);
assert!(!result.is_error.unwrap_or(false));
assert!(matches!(
result.content.first(),
Some(rmcp::model::ContentBlock::Text(t)) if t.text == "from-replace"
));
}
#[tokio::test]
async fn replace_outcome_subject_to_size_cap() {
let huge = CallToolResult::success(vec![ContentBlock::text("y".repeat(8_192))]);
let huge_size = serde_json::to_vec(&huge).unwrap().len();
assert!(huge_size > 256);
let (final_result, accounted, capped) = apply_size_cap(huge, Some(256), "replaced_tool");
assert!(capped);
assert_eq!(accounted, 257);
assert_eq!(final_result.is_error, Some(true));
assert!(matches!(
final_result.content.first(),
Some(rmcp::model::ContentBlock::Text(t)) if t.text.contains("result_too_large")
));
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn after_hook_fires_exactly_once_via_spawn() {
let counter = Arc::new(AtomicUsize::new(0));
let c = Arc::clone(&counter);
let after: AfterHook = Arc::new(move |_ctx, _disp, _size| {
let c = Arc::clone(&c);
Box::pin(async move {
c.fetch_add(1, Ordering::Relaxed);
})
});
let holder = Arc::new(AfterHookHolder { f: after });
HookedHandler::<TestHandler>::spawn_after(
Some(&holder),
ctx("t"),
HookDisposition::InnerExecuted,
42,
);
let deadline = std::time::Instant::now() + std::time::Duration::from_secs(1);
while counter.load(Ordering::Relaxed) == 0 && std::time::Instant::now() < deadline {
tokio::task::yield_now().await;
tokio::time::sleep(std::time::Duration::from_millis(5)).await;
}
assert_eq!(counter.load(Ordering::Relaxed), 1);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn after_hook_panic_is_isolated_from_response_path() {
let after: AfterHook = Arc::new(|_ctx, _disp, _size| {
Box::pin(async {
panic!("intentional panic in after-hook");
})
});
let holder = Arc::new(AfterHookHolder { f: after });
HookedHandler::<TestHandler>::spawn_after(
Some(&holder),
ctx("boom"),
HookDisposition::InnerExecuted,
0,
);
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let still_alive = tokio::spawn(async { 1_u32 + 2 }).await.unwrap();
assert_eq!(still_alive, 3);
}
}