use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use bytes::{Buf, Bytes};
use http::{Request, Response, StatusCode};
use http_body::Body;
use serde::de::{DeserializeSeed, MapAccess, SeqAccess, Visitor};
use serde::{Deserialize, Deserializer};
use tracing::warn;
use crate::http_layer::{InvalidScope, ScopeVerdict, checked_scopes, judge_scopes, scope_refusal};
use crate::token::TokenRejection;
pub const DEFAULT_BODY_LIMIT: usize = 1024 * 1024;
pub const MIN_BODY_LIMIT: usize = 4 * 1024;
pub const MAX_BODY_LIMIT: usize = 64 * 1024 * 1024;
#[derive(Clone, Debug)]
pub struct McpToolScopes {
rules: Arc<Rules>,
}
#[derive(Clone, Debug)]
struct Rules {
default: Vec<String>,
tools: Vec<(String, Vec<String>)>,
strictest: Vec<String>,
always_scoped: bool,
body_limit: usize,
static_bypasses: bool,
}
impl Rules {
fn with_strictest(mut self) -> Self {
let mut all: Vec<String> = Vec::new();
for scope in self
.default
.iter()
.chain(self.tools.iter().flat_map(|(_, scopes)| scopes))
{
if !all.contains(scope) {
all.push(scope.clone());
}
}
self.strictest = all;
self.always_scoped =
!self.default.is_empty() && self.tools.iter().all(|(_, scopes)| !scopes.is_empty());
self
}
fn for_tool(&self, name: &str) -> &[String] {
self.tools
.iter()
.find(|(tool, _)| tool == name)
.map_or(&self.default, |(_, scopes)| scopes)
}
fn for_body(&self, body: &[u8]) -> Vec<String> {
if body.is_empty() {
return self.default.clone();
}
let mut needs_default = false;
let mut needs_strictest = false;
let mut needs_tool = vec![false; self.tools.len()];
let mut seen_any = false;
let tools = &self.tools;
let readable = for_each_message(
body,
&mut |name| {
tools
.iter()
.position(|(tool, _)| tool == name)
.unwrap_or(UNKNOWN_TOOL)
},
&mut |message| {
seen_any = true;
match message {
Message::NotToolCall => needs_default = true,
Message::ToolCall(UNKNOWN_TOOL) => needs_default = true,
Message::ToolCall(i) => needs_tool[i] = true,
Message::Ambiguous => needs_strictest = true,
}
},
);
if !readable || needs_strictest {
return self.strictest.clone();
}
if !seen_any {
return self.default.clone();
}
let mut all: Vec<String> = Vec::new();
let needed = needs_default
.then_some(&self.default[..])
.into_iter()
.chain(
self.tools
.iter()
.zip(&needs_tool)
.filter(|(_, needed)| **needed)
.map(|((_, scopes), _)| &scopes[..]),
);
for scopes in needed {
for scope in scopes {
if !all.contains(scope) {
all.push(scope.clone());
}
}
}
all
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum McpScopesError {
#[error(transparent)]
InvalidScope(#[from] InvalidScope),
#[error("body limit {0} is outside {MIN_BODY_LIMIT}..={MAX_BODY_LIMIT}")]
BodyLimitOutOfRange(usize),
}
#[allow(clippy::new_without_default)]
impl McpToolScopes {
pub fn new() -> Self {
Self {
rules: Arc::new(Rules {
default: Vec::new(),
tools: Vec::new(),
strictest: Vec::new(),
always_scoped: false,
body_limit: DEFAULT_BODY_LIMIT,
static_bypasses: false,
}),
}
}
fn update(self, f: impl FnOnce(&mut Rules)) -> Self {
let mut rules = Arc::unwrap_or_clone(self.rules);
f(&mut rules);
Self {
rules: Arc::new(rules.with_strictest()),
}
}
pub fn default(self, scopes: impl IntoIterator<Item = impl Into<String>>) -> Self {
self.try_default(scopes)
.unwrap_or_else(|e| panic!("McpToolScopes::default: {e}"))
}
pub fn try_default(
self,
scopes: impl IntoIterator<Item = impl Into<String>>,
) -> Result<Self, McpScopesError> {
let scopes = checked_scopes(scopes)?;
Ok(self.update(|rules| rules.default = scopes))
}
pub fn tool(
self,
name: impl Into<String>,
scopes: impl IntoIterator<Item = impl Into<String>>,
) -> Self {
self.try_tool(name, scopes)
.unwrap_or_else(|e| panic!("McpToolScopes::tool: {e}"))
}
pub fn try_tool(
self,
name: impl Into<String>,
scopes: impl IntoIterator<Item = impl Into<String>>,
) -> Result<Self, McpScopesError> {
let name = name.into();
let scopes = checked_scopes(scopes)?;
Ok(self.update(|rules| {
rules.tools.retain(|(tool, _)| *tool != name);
rules.tools.push((name, scopes));
}))
}
pub fn body_limit(self, bytes: usize) -> Self {
self.try_body_limit(bytes)
.unwrap_or_else(|e| panic!("McpToolScopes::body_limit: {e}"))
}
pub fn try_body_limit(self, bytes: usize) -> Result<Self, McpScopesError> {
if !(MIN_BODY_LIMIT..=MAX_BODY_LIMIT).contains(&bytes) {
return Err(McpScopesError::BodyLimitOutOfRange(bytes));
}
Ok(self.update(|rules| rules.body_limit = bytes))
}
pub fn static_token_bypasses_scopes(self) -> Self {
self.update(|rules| rules.static_bypasses = true)
}
pub fn scopes_for_tool(&self, tool: &str) -> &[String] {
self.rules.for_tool(tool)
}
}
impl<S> tower_layer::Layer<S> for McpToolScopes {
type Service = McpToolScopesService<S>;
fn layer(&self, inner: S) -> Self::Service {
McpToolScopesService {
rules: Arc::clone(&self.rules),
inner,
}
}
}
#[derive(Clone, Debug)]
pub struct McpToolScopesService<S> {
rules: Arc<Rules>,
inner: S,
}
enum ReadError {
TooLarge,
Failed,
}
async fn read_capped<B: Body + Unpin>(
mut body: B,
limit: usize,
announced: u64,
) -> Result<Bytes, ReadError> {
if body.size_hint().lower() > limit as u64 {
return Err(ReadError::TooLarge);
}
let reserve = usize::try_from(announced).map_or(limit, |n| n.min(limit));
let mut buf: Vec<u8> = Vec::with_capacity(reserve);
loop {
let frame = std::future::poll_fn(|cx| Pin::new(&mut body).poll_frame(cx)).await;
match frame {
None => return Ok(Bytes::from(buf)),
Some(Err(_)) => return Err(ReadError::Failed),
Some(Ok(frame)) => {
let Ok(mut data) = frame.into_data() else {
continue;
};
if buf.len().saturating_add(data.remaining()) > limit {
return Err(ReadError::TooLarge);
}
while data.has_remaining() {
let chunk = data.chunk();
let n = chunk.len();
buf.extend_from_slice(chunk);
data.advance(n);
}
}
}
}
}
fn plain<B: Default>(status: StatusCode) -> Response<B> {
let mut response = Response::new(B::default());
*response.status_mut() = status;
response
}
impl<S, ReqBody, ResBody> tower_service::Service<Request<ReqBody>> for McpToolScopesService<S>
where
S: tower_service::Service<Request<ReqBody>, Response = Response<ResBody>>
+ Clone
+ Send
+ 'static,
S::Future: Send + 'static,
ReqBody: Body + From<Bytes> + Send + 'static,
ReqBody::Data: Send,
ReqBody::Error: Send,
ResBody: Default + 'static,
{
type Response = Response<ResBody>;
type Error = S::Error;
type Future =
Pin<Box<dyn Future<Output = Result<Response<ResBody>, S::Error>> + Send + 'static>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, request: Request<ReqBody>) -> Self::Future {
let clone = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, clone);
let rules = Arc::clone(&self.rules);
Box::pin(async move {
let (parts, body) = request.into_parts();
if let ScopeVerdict::NoLayer = judge_scopes(&parts, &[], false) {
return Ok(scope_refusal(
&parts,
&TokenRejection::Missing,
&[],
"McpToolScopes",
));
}
let content_length = parts
.headers
.get(http::header::CONTENT_LENGTH)
.and_then(|v| v.to_str().ok())
.and_then(|v| v.trim().parse::<u64>().ok());
let hint = body.size_hint();
if rules.always_scoped
&& parts
.extensions
.get::<crate::authenticate::Credential>()
.is_none()
{
return Ok(scope_refusal(
&parts,
&TokenRejection::Missing,
&rules.default,
"McpToolScopes",
));
}
let (required, body) = {
let announced = content_length.unwrap_or(0).max(hint.lower());
let read = if announced > rules.body_limit as u64 {
Err(ReadError::TooLarge)
} else {
read_capped(Box::pin(body), rules.body_limit, announced).await
};
let bytes = match read {
Ok(bytes) => bytes,
Err(ReadError::TooLarge) => {
warn!(
path = %parts.uri.path(),
limit = rules.body_limit,
"MCP request body exceeds the limit; refusing it unread"
);
return Ok(plain(StatusCode::PAYLOAD_TOO_LARGE));
}
Err(ReadError::Failed) => {
warn!(
path = %parts.uri.path(),
"MCP request body could not be read; refusing the request"
);
return Ok(plain(StatusCode::BAD_REQUEST));
}
};
(rules.for_body(&bytes), ReqBody::from(bytes))
};
match judge_scopes(&parts, &required, rules.static_bypasses) {
ScopeVerdict::Pass => inner.call(Request::from_parts(parts, body)).await,
ScopeVerdict::Refuse(rejection) => Ok(scope_refusal(
&parts,
&rejection,
&required,
"McpToolScopes",
)),
ScopeVerdict::NoLayer => Ok(scope_refusal(
&parts,
&TokenRejection::Missing,
&required,
"McpToolScopes",
)),
}
})
}
}
pub(crate) const UNKNOWN_TOOL: usize = usize::MAX;
#[derive(Debug, PartialEq, Eq)]
pub(crate) enum Message {
NotToolCall,
ToolCall(usize),
Ambiguous,
}
#[cfg(any(test, fuzzing))]
#[derive(Debug, PartialEq, Eq)]
pub(crate) enum NamedMessage {
NotToolCall,
ToolCall(String),
Ambiguous,
}
#[cfg(any(test, fuzzing))]
#[derive(Debug, PartialEq, Eq)]
pub(crate) enum Classified {
Messages(Vec<NamedMessage>),
Unreadable,
}
#[cfg(any(test, fuzzing))]
pub(crate) fn classify(body: &[u8]) -> Classified {
let mut names: Vec<String> = Vec::new();
let mut messages = Vec::new();
let readable = for_each_message(
body,
&mut |name| {
names.push(name.to_string());
names.len() - 1
},
&mut |m| messages.push(m),
);
if !readable {
return Classified::Unreadable;
}
Classified::Messages(
messages
.into_iter()
.map(|m| match m {
Message::NotToolCall => NamedMessage::NotToolCall,
Message::ToolCall(i) => NamedMessage::ToolCall(names[i].clone()),
Message::Ambiguous => NamedMessage::Ambiguous,
})
.collect(),
)
}
pub(crate) fn for_each_message(
body: &[u8],
tool: &mut dyn FnMut(&str) -> usize,
sink: &mut dyn FnMut(Message),
) -> bool {
let mut deserializer = serde_json::Deserializer::from_slice(body);
TopSeed { tool, sink }
.deserialize(&mut deserializer)
.is_ok()
&& deserializer.end().is_ok()
}
macro_rules! accept_scalars {
($value:expr) => {
fn visit_bool<E>(self, _: bool) -> Result<Self::Value, E> {
Ok($value)
}
fn visit_i64<E>(self, _: i64) -> Result<Self::Value, E> {
Ok($value)
}
fn visit_u64<E>(self, _: u64) -> Result<Self::Value, E> {
Ok($value)
}
fn visit_f64<E>(self, _: f64) -> Result<Self::Value, E> {
Ok($value)
}
fn visit_str<E>(self, _: &str) -> Result<Self::Value, E> {
Ok($value)
}
fn visit_unit<E>(self) -> Result<Self::Value, E> {
Ok($value)
}
};
}
macro_rules! accept_non_string_scalars {
($value:expr) => {
fn visit_bool<E>(self, _: bool) -> Result<Self::Value, E> {
Ok($value)
}
fn visit_i64<E>(self, _: i64) -> Result<Self::Value, E> {
Ok($value)
}
fn visit_u64<E>(self, _: u64) -> Result<Self::Value, E> {
Ok($value)
}
fn visit_f64<E>(self, _: f64) -> Result<Self::Value, E> {
Ok($value)
}
fn visit_unit<E>(self) -> Result<Self::Value, E> {
Ok($value)
}
};
}
struct Validate;
impl<'de> Deserialize<'de> for Validate {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
struct V;
impl<'de> Visitor<'de> for V {
type Value = Validate;
fn expecting(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("any JSON value")
}
accept_scalars!(Validate);
fn visit_map<A: MapAccess<'de>>(self, mut map: A) -> Result<Validate, A::Error> {
while map.next_entry::<Validate, Validate>()?.is_some() {}
Ok(Validate)
}
fn visit_seq<A: SeqAccess<'de>>(self, mut seq: A) -> Result<Validate, A::Error> {
while seq.next_element::<Validate>()?.is_some() {}
Ok(Validate)
}
}
deserializer.deserialize_any(V)
}
}
struct TopSeed<'s> {
tool: &'s mut dyn FnMut(&str) -> usize,
sink: &'s mut dyn FnMut(Message),
}
impl<'de> DeserializeSeed<'de> for TopSeed<'_> {
type Value = ();
fn deserialize<D: Deserializer<'de>>(self, deserializer: D) -> Result<(), D::Error> {
struct V<'s>(TopSeed<'s>);
impl<'de> Visitor<'de> for V<'_> {
type Value = ();
fn expecting(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("a JSON-RPC message or batch")
}
fn visit_map<A: MapAccess<'de>>(self, map: A) -> Result<(), A::Error> {
let message = read_message(map, self.0.tool)?;
(self.0.sink)(message);
Ok(())
}
fn visit_seq<A: SeqAccess<'de>>(self, mut seq: A) -> Result<(), A::Error> {
let TopSeed { tool, sink } = self.0;
while seq
.next_element_seed(ElementSeed {
tool: &mut *tool,
sink: &mut *sink,
})?
.is_some()
{}
Ok(())
}
}
deserializer.deserialize_any(V(self))
}
}
struct ElementSeed<'s> {
tool: &'s mut dyn FnMut(&str) -> usize,
sink: &'s mut dyn FnMut(Message),
}
impl<'de> DeserializeSeed<'de> for ElementSeed<'_> {
type Value = ();
fn deserialize<D: Deserializer<'de>>(self, deserializer: D) -> Result<(), D::Error> {
struct V<'s>(&'s mut dyn FnMut(&str) -> usize);
impl<'de> Visitor<'de> for V<'_> {
type Value = Message;
fn expecting(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("any JSON value")
}
accept_scalars!(Message::Ambiguous);
fn visit_map<A: MapAccess<'de>>(self, map: A) -> Result<Message, A::Error> {
read_message(map, self.0)
}
fn visit_seq<A: SeqAccess<'de>>(self, mut seq: A) -> Result<Message, A::Error> {
while seq.next_element::<Validate>()?.is_some() {}
Ok(Message::Ambiguous)
}
}
let message = deserializer.deserialize_any(V(self.tool))?;
(self.sink)(message);
Ok(())
}
}
enum Key {
Method,
Params,
Name,
Other,
}
impl<'de> Deserialize<'de> for Key {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
struct V;
impl Visitor<'_> for V {
type Value = Key;
fn expecting(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("an object key")
}
fn visit_str<E>(self, key: &str) -> Result<Key, E> {
Ok(match key {
"method" => Key::Method,
"params" => Key::Params,
"name" => Key::Name,
_ => Key::Other,
})
}
}
deserializer.deserialize_any(V)
}
}
struct IsToolsCall(Option<bool>);
impl<'de> Deserialize<'de> for IsToolsCall {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
struct V;
impl<'de> Visitor<'de> for V {
type Value = IsToolsCall;
fn expecting(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("any JSON value")
}
accept_non_string_scalars!(IsToolsCall(None));
fn visit_str<E>(self, v: &str) -> Result<IsToolsCall, E> {
Ok(IsToolsCall(Some(v == "tools/call")))
}
fn visit_map<A: MapAccess<'de>>(self, mut map: A) -> Result<IsToolsCall, A::Error> {
while map.next_entry::<Validate, Validate>()?.is_some() {}
Ok(IsToolsCall(None))
}
fn visit_seq<A: SeqAccess<'de>>(self, mut seq: A) -> Result<IsToolsCall, A::Error> {
while seq.next_element::<Validate>()?.is_some() {}
Ok(IsToolsCall(None))
}
}
deserializer.deserialize_any(V)
}
}
struct NameSeed<'s>(&'s mut dyn FnMut(&str) -> usize);
impl<'de> DeserializeSeed<'de> for NameSeed<'_> {
type Value = Option<usize>;
fn deserialize<D: Deserializer<'de>>(self, deserializer: D) -> Result<Option<usize>, D::Error> {
struct V<'s>(&'s mut dyn FnMut(&str) -> usize);
impl<'de> Visitor<'de> for V<'_> {
type Value = Option<usize>;
fn expecting(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("any JSON value")
}
accept_non_string_scalars!(None);
fn visit_str<E>(self, v: &str) -> Result<Option<usize>, E> {
Ok(Some((self.0)(v)))
}
fn visit_map<A: MapAccess<'de>>(self, mut map: A) -> Result<Option<usize>, A::Error> {
while map.next_entry::<Validate, Validate>()?.is_some() {}
Ok(None)
}
fn visit_seq<A: SeqAccess<'de>>(self, mut seq: A) -> Result<Option<usize>, A::Error> {
while seq.next_element::<Validate>()?.is_some() {}
Ok(None)
}
}
deserializer.deserialize_any(V(self.0))
}
}
struct ParamsSeed<'s>(&'s mut dyn FnMut(&str) -> usize);
impl<'de> DeserializeSeed<'de> for ParamsSeed<'_> {
type Value = Option<usize>;
fn deserialize<D: Deserializer<'de>>(self, deserializer: D) -> Result<Option<usize>, D::Error> {
struct V<'s>(&'s mut dyn FnMut(&str) -> usize);
impl<'de> Visitor<'de> for V<'_> {
type Value = Option<usize>;
fn expecting(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("any JSON value")
}
accept_scalars!(None);
fn visit_map<A: MapAccess<'de>>(self, mut map: A) -> Result<Option<usize>, A::Error> {
let mut name: Option<Option<usize>> = None;
let mut repeated = false;
while let Some(key) = map.next_key::<Key>()? {
match key {
Key::Name if name.is_none() => {
name = Some(map.next_value_seed(NameSeed(&mut *self.0))?);
}
Key::Name => {
repeated = true;
map.next_value::<Validate>()?;
}
_ => {
map.next_value::<Validate>()?;
}
}
}
Ok(match (repeated, name) {
(false, Some(name)) => name,
_ => None,
})
}
fn visit_seq<A: SeqAccess<'de>>(self, mut seq: A) -> Result<Option<usize>, A::Error> {
while seq.next_element::<Validate>()?.is_some() {}
Ok(None)
}
}
deserializer.deserialize_any(V(self.0))
}
}
fn read_message<'de, A: MapAccess<'de>>(
mut map: A,
tool: &mut dyn FnMut(&str) -> usize,
) -> Result<Message, A::Error> {
let mut method: Option<Option<bool>> = None;
let mut params: Option<Option<usize>> = None;
let mut repeated = false;
while let Some(key) = map.next_key::<Key>()? {
match key {
Key::Method if method.is_none() => {
method = Some(map.next_value::<IsToolsCall>()?.0);
}
Key::Params if params.is_none() => {
params = Some(map.next_value_seed(ParamsSeed(&mut *tool))?);
}
Key::Method | Key::Params => {
repeated = true;
map.next_value::<Validate>()?;
}
_ => {
map.next_value::<Validate>()?;
}
}
}
Ok(match (repeated, method) {
(true, _) => Message::Ambiguous,
(false, Some(Some(true))) => match params {
Some(Some(index)) => Message::ToolCall(index),
_ => Message::Ambiguous,
},
_ => Message::NotToolCall,
})
}
#[cfg(all(test, feature = "axum"))]
mod service_tests;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn classification() {
let one = |json: &str| classify(json.as_bytes());
let msgs = |m: Vec<NamedMessage>| Classified::Messages(m);
assert_eq!(
one(
r#"{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"w","arguments":{"name":"x"}}}"#
),
msgs(vec![NamedMessage::ToolCall("w".into())])
);
assert_eq!(
one(r#"{"jsonrpc":"2.0","id":1,"method":"tools/list"}"#),
msgs(vec![NamedMessage::NotToolCall])
);
assert_eq!(
one(r#"{"jsonrpc":"2.0","id":1,"result":{}}"#),
msgs(vec![NamedMessage::NotToolCall])
);
assert_eq!(
one(r#"{"jsonrpc":"2.0","method":"notifications/initialized"}"#),
msgs(vec![NamedMessage::NotToolCall])
);
assert_eq!(
one(r#"{"method":7}"#),
msgs(vec![NamedMessage::NotToolCall])
);
assert_eq!(
one(r#"{"method":"tools\/call","params":{"name":"w"}}"#),
msgs(vec![NamedMessage::ToolCall("w".into())])
);
for json in [
r#"{"method":"tools/call"}"#,
r#"{"method":"tools/call","params":{}}"#,
r#"{"method":"tools/call","params":{"name":1}}"#,
r#"{"method":"tools/call","params":["w"]}"#,
r#"{"method":"tools/call","params":{"name":"r","name":"w"}}"#,
r#"{"method":"tools/list","method":"tools/call","params":{"name":"w"}}"#,
r#"{"method":"tools/call","params":{"name":"r"},"params":{"name":"w"}}"#,
] {
assert_eq!(one(json), msgs(vec![NamedMessage::Ambiguous]), "{json}");
}
assert_eq!(
one(r#"[{"method":"tools/list"},{"method":"tools/call","params":{"name":"w"}},3]"#),
msgs(vec![
NamedMessage::NotToolCall,
NamedMessage::ToolCall("w".into()),
NamedMessage::Ambiguous
])
);
assert_eq!(one("[]"), msgs(vec![]));
for body in [
"",
"{",
"nope",
"7",
"\"tools/call\"",
"{} {}",
"\u{feff}{}",
] {
assert_eq!(one(body), Classified::Unreadable, "{body:?}");
}
}
fn rules() -> McpToolScopes {
McpToolScopes::new()
.default(["mcp:read"])
.tool("write_document", ["mcp:write"])
.tool("admin", ["mcp:write", "mcp:admin"])
}
#[test]
fn requirements_per_body() {
let r = rules();
let r = &r.rules;
let req = |json: &str| r.for_body(json.as_bytes());
assert_eq!(
req(r#"{"method":"tools/call","params":{"name":"write_document"}}"#),
["mcp:write"]
);
assert_eq!(
req(r#"{"method":"tools/call","params":{"name":"other"}}"#),
["mcp:read"]
);
assert_eq!(req(r#"{"method":"initialize"}"#), ["mcp:read"]);
let strictest = ["mcp:read", "mcp:write", "mcp:admin"];
assert_eq!(req("not json"), strictest);
assert_eq!(req(r#"{"method":"tools/call"}"#), strictest);
assert_eq!(
req(
r#"[{"method":"tools/call","params":{"name":"write_document"}},{"method":"tools/list"}]"#
),
["mcp:read", "mcp:write"]
);
assert_eq!(req("[]"), ["mcp:read"]);
assert_eq!(req(""), ["mcp:read"]);
assert_eq!(req(" "), strictest);
assert_eq!(req(r#"{"method":"initialize","x":1e999}"#), strictest);
assert_eq!(
r.for_body(b"{\"method\":\"initialize\",\"x\":\"\xff\"}"),
strictest
);
assert_eq!(req(r#"{"method":"initialize"} x"#), strictest);
assert_eq!(rules().scopes_for_tool("admin"), ["mcp:write", "mcp:admin"]);
assert_eq!(rules().scopes_for_tool("nope"), ["mcp:read"]);
}
#[test]
fn a_repeated_tool_entry_replaces_the_earlier_one() {
let r = McpToolScopes::new().tool("t", ["a"]).tool("t", ["b"]);
assert_eq!(r.scopes_for_tool("t"), ["b"]);
assert_eq!(r.rules.strictest, ["b"]);
}
#[test]
#[should_panic(expected = "is not a valid scope")]
fn an_invalid_scope_panics() {
let _ = McpToolScopes::new().tool("t", ["has space"]);
}
#[test]
fn the_try_forms_return_what_the_panicking_forms_panic_on() {
match McpToolScopes::new().try_tool("t", ["ok", "has space"]) {
Err(McpScopesError::InvalidScope(e)) => assert_eq!(e.scope(), "has space"),
other => panic!("{other:?}"),
}
match McpToolScopes::new().try_default([""]) {
Err(McpScopesError::InvalidScope(e)) => assert_eq!(e.scope(), ""),
other => panic!("{other:?}"),
}
assert_eq!(
McpToolScopes::new()
.try_body_limit(MIN_BODY_LIMIT - 1)
.unwrap_err(),
McpScopesError::BodyLimitOutOfRange(MIN_BODY_LIMIT - 1)
);
let ok = McpToolScopes::new()
.try_default(["r"])
.and_then(|m| m.try_tool("t", ["w"]))
.and_then(|m| m.try_body_limit(MAX_BODY_LIMIT))
.unwrap();
assert_eq!(ok.scopes_for_tool("t"), ["w"]);
assert_eq!(ok.rules.body_limit, MAX_BODY_LIMIT);
}
#[test]
#[should_panic(expected = "body_limit")]
fn a_body_limit_out_of_bounds_panics() {
let _ = McpToolScopes::new().body_limit(MAX_BODY_LIMIT + 1);
}
}