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::{
ClaimClause, InvalidScope, ScopeVerdict, checked_scopes, judge_scopes, scope_refusal,
scope_refusal_with_claims,
};
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, Default, PartialEq, Eq)]
struct Requirement {
scopes: Vec<String>,
claims: Vec<ClaimClause>,
}
impl Requirement {
fn is_empty(&self) -> bool {
self.scopes.is_empty() && self.claims.is_empty()
}
fn absorb(&mut self, other: &Requirement) {
for scope in &other.scopes {
if !self.scopes.contains(scope) {
self.scopes.push(scope.clone());
}
}
for clause in &other.claims {
if !self.claims.contains(clause) {
self.claims.push(clause.clone());
}
}
}
}
#[derive(Clone, Debug)]
struct Rules {
default: Requirement,
tools: Vec<(String, Requirement)>,
strictest: Requirement,
always_scoped: bool,
body_limit: usize,
static_bypasses: bool,
}
impl Rules {
fn with_strictest(mut self) -> Self {
let mut all = self.default.clone();
for (_, requirement) in &self.tools {
all.absorb(requirement);
}
self.strictest = all;
self.always_scoped = !self.default.is_empty()
&& self
.tools
.iter()
.all(|(_, requirement)| !requirement.is_empty());
self
}
fn for_tool(&self, name: &str) -> &Requirement {
self.tools
.iter()
.find(|(tool, _)| tool == name)
.map_or(&self.default, |(_, requirement)| requirement)
}
fn for_body(&self, body: &[u8]) -> Requirement {
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 = Requirement::default();
if needs_default {
all.absorb(&self.default);
}
for ((_, requirement), needed) in self.tools.iter().zip(&needs_tool) {
if *needed {
all.absorb(requirement);
}
}
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),
#[error(
"the requirement on claim {claim:?} needs a non-blank claim name and at least one \
value, none of them blank"
)]
#[non_exhaustive]
InvalidClaimRequirement {
claim: String,
},
}
fn checked_claim(
claim: impl Into<String>,
values: impl IntoIterator<Item = impl Into<String>>,
) -> Result<ClaimClause, McpScopesError> {
let claim = claim.into();
let mut any_of: Vec<String> = Vec::new();
let mut blank = claim.trim().is_empty();
for value in values {
let value = value.into();
blank |= value.trim().is_empty();
if !any_of.contains(&value) {
any_of.push(value);
}
}
if blank || any_of.is_empty() {
return Err(McpScopesError::InvalidClaimRequirement { claim });
}
Ok(ClaimClause { claim, any_of })
}
#[allow(clippy::new_without_default)]
impl McpToolScopes {
pub fn new() -> Self {
Self {
rules: Arc::new(Rules {
default: Requirement::default(),
tools: Vec::new(),
strictest: Requirement::default(),
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 = 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| {
let claims = rules
.tools
.iter()
.position(|(tool, _)| *tool == name)
.map(|i| rules.tools.remove(i).1.claims)
.unwrap_or_default();
rules.tools.push((name, Requirement { scopes, claims }));
}))
}
pub fn default_claim(
self,
claim: impl Into<String>,
values: impl IntoIterator<Item = impl Into<String>>,
) -> Self {
self.try_default_claim(claim, values)
.unwrap_or_else(|e| panic!("McpToolScopes::default_claim: {e}"))
}
pub fn try_default_claim(
self,
claim: impl Into<String>,
values: impl IntoIterator<Item = impl Into<String>>,
) -> Result<Self, McpScopesError> {
let clause = checked_claim(claim, values)?;
Ok(self.update(|rules| {
if !rules.default.claims.contains(&clause) {
rules.default.claims.push(clause);
}
}))
}
pub fn tool_claim(
self,
name: impl Into<String>,
claim: impl Into<String>,
values: impl IntoIterator<Item = impl Into<String>>,
) -> Self {
self.try_tool_claim(name, claim, values)
.unwrap_or_else(|e| panic!("McpToolScopes::tool_claim: {e}"))
}
pub fn try_tool_claim(
self,
name: impl Into<String>,
claim: impl Into<String>,
values: impl IntoIterator<Item = impl Into<String>>,
) -> Result<Self, McpScopesError> {
let name = name.into();
let clause = checked_claim(claim, values)?;
Ok(self.update(|rules| {
let i = match rules.tools.iter().position(|(tool, _)| *tool == name) {
Some(i) => i,
None => {
rules.tools.push((name, Requirement::default()));
rules.tools.len() - 1
}
};
let claims = &mut rules.tools[i].1.claims;
if !claims.contains(&clause) {
claims.push(clause);
}
}))
}
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).scopes
}
pub fn claims_for_tool(&self, tool: &str) -> Vec<(&str, &[String])> {
self.rules
.for_tool(tool)
.claims
.iter()
.map(|clause| (clause.claim.as_str(), &clause.any_of[..]))
.collect()
}
}
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.scopes,
"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.scopes,
&required.claims,
rules.static_bypasses,
) {
ScopeVerdict::Pass => inner.call(Request::from_parts(parts, body)).await,
ScopeVerdict::Refuse(rejection) => Ok(scope_refusal_with_claims(
&parts,
&rejection,
&required.scopes,
&required.claims,
"McpToolScopes",
)),
ScopeVerdict::NoLayer => Ok(scope_refusal_with_claims(
&parts,
&TokenRejection::Missing,
&required.scopes,
&required.claims,
"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)
}
};
}
const RAW_VALUE_TOKEN: &str = "$serde_json::private::RawValue";
const NUMBER_TOKEN: &str = "$serde_json::private::Number";
struct Validate {
token_key: bool,
}
fn validate_map<'de, A: MapAccess<'de>>(mut map: A) -> Result<bool, A::Error> {
let mut token_key = false;
while let Some(key) = map.next_key::<Key>()? {
token_key |= matches!(key, Key::SerdeJsonToken);
token_key |= map.next_value::<Validate>()?.token_key;
}
Ok(token_key)
}
fn validate_seq<'de, A: SeqAccess<'de>>(mut seq: A) -> Result<bool, A::Error> {
let mut token_key = false;
while let Some(element) = seq.next_element::<Validate>()? {
token_key |= element.token_key;
}
Ok(token_key)
}
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 { token_key: false });
fn visit_map<A: MapAccess<'de>>(self, map: A) -> Result<Validate, A::Error> {
Ok(Validate {
token_key: validate_map(map)?,
})
}
fn visit_seq<A: SeqAccess<'de>>(self, seq: A) -> Result<Validate, A::Error> {
Ok(Validate {
token_key: validate_seq(seq)?,
})
}
}
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, seq: A) -> Result<Message, A::Error> {
validate_seq(seq)?;
Ok(Message::Ambiguous)
}
}
let message = deserializer.deserialize_any(V(self.tool))?;
(self.sink)(message);
Ok(())
}
}
enum Key {
Method,
Params,
Name,
SerdeJsonToken,
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,
RAW_VALUE_TOKEN | NUMBER_TOKEN => Key::SerdeJsonToken,
_ => Key::Other,
})
}
}
deserializer.deserialize_any(V)
}
}
struct IsToolsCall {
tools_call: bool,
token_key: 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 {
tools_call: false,
token_key: false
});
fn visit_str<E>(self, v: &str) -> Result<IsToolsCall, E> {
Ok(IsToolsCall {
tools_call: v == "tools/call",
token_key: false,
})
}
fn visit_map<A: MapAccess<'de>>(self, map: A) -> Result<IsToolsCall, A::Error> {
Ok(IsToolsCall {
tools_call: false,
token_key: validate_map(map)?,
})
}
fn visit_seq<A: SeqAccess<'de>>(self, seq: A) -> Result<IsToolsCall, A::Error> {
Ok(IsToolsCall {
tools_call: false,
token_key: validate_seq(seq)?,
})
}
}
deserializer.deserialize_any(V)
}
}
struct NameSeed<'s>(&'s mut dyn FnMut(&str) -> usize);
impl<'de> DeserializeSeed<'de> for NameSeed<'_> {
type Value = (Option<usize>, bool);
fn deserialize<D: Deserializer<'de>>(self, deserializer: D) -> Result<Self::Value, D::Error> {
struct V<'s>(&'s mut dyn FnMut(&str) -> usize);
impl<'de> Visitor<'de> for V<'_> {
type Value = (Option<usize>, bool);
fn expecting(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("any JSON value")
}
accept_non_string_scalars!((None, false));
fn visit_str<E>(self, v: &str) -> Result<Self::Value, E> {
Ok((Some((self.0)(v)), false))
}
fn visit_map<A: MapAccess<'de>>(self, map: A) -> Result<Self::Value, A::Error> {
Ok((None, validate_map(map)?))
}
fn visit_seq<A: SeqAccess<'de>>(self, seq: A) -> Result<Self::Value, A::Error> {
Ok((None, validate_seq(seq)?))
}
}
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>, bool);
fn deserialize<D: Deserializer<'de>>(self, deserializer: D) -> Result<Self::Value, D::Error> {
struct V<'s>(&'s mut dyn FnMut(&str) -> usize);
impl<'de> Visitor<'de> for V<'_> {
type Value = (Option<usize>, bool);
fn expecting(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("any JSON value")
}
accept_scalars!((None, false));
fn visit_map<A: MapAccess<'de>>(self, mut map: A) -> Result<Self::Value, A::Error> {
let mut name: Option<Option<usize>> = None;
let mut repeated = false;
let mut token_key = false;
while let Some(key) = map.next_key::<Key>()? {
match key {
Key::Name if name.is_none() => {
let (found, nested) = map.next_value_seed(NameSeed(&mut *self.0))?;
name = Some(found);
token_key |= nested;
}
Key::Name => {
repeated = true;
token_key |= map.next_value::<Validate>()?.token_key;
}
key => {
token_key |= matches!(key, Key::SerdeJsonToken);
token_key |= map.next_value::<Validate>()?.token_key;
}
}
}
let name = match (repeated, name) {
(false, Some(name)) => name,
_ => None,
};
Ok((name, token_key))
}
fn visit_seq<A: SeqAccess<'de>>(self, seq: A) -> Result<Self::Value, A::Error> {
Ok((None, validate_seq(seq)?))
}
}
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 tools_call: Option<bool> = None;
let mut params: Option<Option<usize>> = None;
let mut repeated = false;
let mut token_key = false;
while let Some(key) = map.next_key::<Key>()? {
match key {
Key::Method if tools_call.is_none() => {
let method = map.next_value::<IsToolsCall>()?;
tools_call = Some(method.tools_call);
token_key |= method.token_key;
}
Key::Params if params.is_none() => {
let (name, nested) = map.next_value_seed(ParamsSeed(&mut *tool))?;
params = Some(name);
token_key |= nested;
}
Key::Method | Key::Params => {
repeated = true;
token_key |= map.next_value::<Validate>()?.token_key;
}
key => {
token_key |= matches!(key, Key::SerdeJsonToken);
token_key |= map.next_value::<Validate>()?.token_key;
}
}
}
Ok(match (repeated || token_key, tools_call) {
(true, _) => Message::Ambiguous,
(false, 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()).scopes;
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\"}")
.scopes,
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"]);
}
const SMUGGLED_CALL: &str = r#"{"$serde_json::private::RawValue":"{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"tools/call\",\"params\":{\"name\":\"admin\",\"arguments\":{}}}"}"#;
#[test]
fn serde_json_token_keys_need_the_strictest_set() {
let r = rules();
let r = &r.rules;
let strictest = ["mcp:read", "mcp:write", "mcp:admin"];
let ambiguous = Classified::Messages(vec![NamedMessage::Ambiguous]);
#[cfg(feature = "axum")]
{
let seen: serde_json::Value = serde_json::from_str(SMUGGLED_CALL).unwrap();
assert_eq!(seen["method"], "tools/call");
assert_eq!(seen["params"]["name"], "admin");
}
let batch = r#"{"$serde_json::private::RawValue":"[{\"method\":\"tools/list\"},{\"method\":\"tools/call\",\"params\":{\"name\":\"admin\"}}]"}"#;
let crash = "\n\n\n\n{\"$serde_json::private::RawValue\":\"{:\"\t} \n";
let number = r#"{"$serde_json::private::Number":"1","method":"tools/call","params":{"name":"admin"}}"#;
let escaped = r#"{"$serde_json::private::RawValue":"{}"}"#;
let later = r#"{"method":"initialize","$serde_json::private::RawValue":"{}"}"#;
for body in [SMUGGLED_CALL, batch, crash, number, escaped, later] {
assert_eq!(classify(body.as_bytes()), ambiguous, "{body:?}");
assert_eq!(r.for_body(body.as_bytes()).scopes, strictest, "{body:?}");
}
let in_batch = format!(r#"[{{"method":"tools/list"}},{SMUGGLED_CALL}]"#);
assert_eq!(
classify(in_batch.as_bytes()),
Classified::Messages(vec![NamedMessage::NotToolCall, NamedMessage::Ambiguous])
);
assert_eq!(r.for_body(in_batch.as_bytes()).scopes, strictest);
for body in [
r#"{"method":"tools/call","params":{"$serde_json::private::RawValue":"{\"name\":\"admin\"}"}}"#,
r#"{"method":{"$serde_json::private::RawValue":"\"tools/call\""},"params":{"name":"admin"}}"#,
r#"{"method":"tools/call","params":{"name":{"$serde_json::private::RawValue":"\"admin\""}}}"#,
r#"{"method":"tools/call","params":{"name":"other","arguments":[{"$serde_json::private::Number":"1"}]}}"#,
r#"{"method":"tools/list","id":{"x":[{"$serde_json::private::RawValue":"1"}]}}"#,
] {
assert_eq!(classify(body.as_bytes()), ambiguous, "{body}");
assert_eq!(r.for_body(body.as_bytes()).scopes, strictest, "{body}");
}
assert_eq!(
req_scopes(
r,
r#"{"method":"tools/call","params":{"name":"write_document","arguments":{"x":"$serde_json::private::RawValue"}}}"#
),
["mcp:write"]
);
assert_eq!(
classify(
br#"{"method":"tools/call","params":{"name":"$serde_json::private::Number"}}"#
),
Classified::Messages(vec![NamedMessage::ToolCall(
"$serde_json::private::Number".into()
)])
);
assert_eq!(
req_scopes(
r,
r#"{"method":"initialize","$serde_json::private::RawValu":"{}"}"#
),
["mcp:read"]
);
}
fn req_scopes(r: &Rules, json: &str) -> Vec<String> {
r.for_body(json.as_bytes()).scopes
}
#[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.scopes, ["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);
}
fn clause(claim: &str, any_of: &[&str]) -> ClaimClause {
ClaimClause {
claim: claim.into(),
any_of: any_of.iter().map(|v| v.to_string()).collect(),
}
}
fn claim_rules() -> McpToolScopes {
McpToolScopes::new()
.default(["mcp:read"])
.default_claim("groups", ["staff"])
.tool("write_document", ["mcp:write"])
.tool_claim("write_document", "groups", ["editors", "admins"])
.tool_claim("purge", "role", ["a"])
.tool_claim("purge", "groups", ["admins"])
.tool_claim("approve", "role", ["b"])
.tool("search", [] as [&str; 0])
}
#[test]
fn claim_requirements_per_body() {
let r = claim_rules();
let r = &r.rules;
let req = |json: &str| r.for_body(json.as_bytes());
let call =
|tool: &str| format!(r#"{{"method":"tools/call","params":{{"name":"{tool}"}}}}"#);
let write = req(&call("write_document"));
assert_eq!(write.scopes, ["mcp:write"]);
assert_eq!(write.claims, [clause("groups", &["editors", "admins"])]);
let purge = req(&call("purge"));
assert!(purge.scopes.is_empty());
assert_eq!(
purge.claims,
[clause("role", &["a"]), clause("groups", &["admins"])]
);
assert!(req(&call("search")).is_empty());
for body in [call("other"), r#"{"method":"tools/list"}"#.to_string()] {
let got = req(&body);
assert_eq!(got.scopes, ["mcp:read"]);
assert_eq!(got.claims, [clause("groups", &["staff"])]);
}
assert_eq!(req("").claims, [clause("groups", &["staff"])]);
let batch = req(&format!(
"[{},{},{}]",
call("purge"),
call("approve"),
call("purge")
));
assert!(batch.scopes.is_empty());
assert_eq!(
batch.claims,
[
clause("role", &["a"]),
clause("groups", &["admins"]),
clause("role", &["b"])
]
);
let strictest = Requirement {
scopes: vec!["mcp:read".into(), "mcp:write".into()],
claims: vec![
clause("groups", &["staff"]),
clause("groups", &["editors", "admins"]),
clause("role", &["a"]),
clause("groups", &["admins"]),
clause("role", &["b"]),
],
};
assert_eq!(req("not json"), strictest);
assert_eq!(req(r#"{"method":"tools/call"}"#), strictest);
assert_eq!(req(" "), strictest);
}
#[test]
fn tool_and_tool_claim_compose_in_either_order() {
let a = McpToolScopes::new()
.tool_claim("t", "groups", ["x"])
.tool("t", ["s"]);
let b = McpToolScopes::new()
.tool("t", ["s"])
.tool_claim("t", "groups", ["x"]);
for r in [&a, &b] {
assert_eq!(r.scopes_for_tool("t"), ["s"]);
let x = ["x".to_string()];
assert_eq!(r.claims_for_tool("t"), [("groups", &x[..])]);
}
let c = McpToolScopes::new()
.tool_claim("t", "groups", ["x", "x"])
.tool_claim("t", "groups", ["x"]);
assert_eq!(c.rules.for_tool("t").claims, [clause("groups", &["x"])]);
let d = McpToolScopes::new()
.default_claim("groups", ["staff"])
.tool("w", ["s"]);
let staff = ["staff".to_string()];
assert_eq!(d.claims_for_tool("nope"), [("groups", &staff[..])]);
assert!(d.claims_for_tool("w").is_empty());
assert!(d.scopes_for_tool("nope").is_empty());
}
#[test]
fn always_scoped_counts_claims() {
assert!(
McpToolScopes::new()
.default_claim("g", ["x"])
.rules
.always_scoped
);
assert!(
McpToolScopes::new()
.default_claim("g", ["x"])
.tool_claim("t", "g", ["y"])
.rules
.always_scoped
);
assert!(
!McpToolScopes::new()
.default_claim("g", ["x"])
.tool("t", [] as [&str; 0])
.rules
.always_scoped
);
assert!(!McpToolScopes::new().rules.always_scoped);
}
#[test]
fn invalid_claim_requirements_are_refused() {
let invalid = |r: Result<McpToolScopes, McpScopesError>, name: &str| match r {
Err(McpScopesError::InvalidClaimRequirement { claim, .. }) => assert_eq!(claim, name),
other => panic!("{other:?}"),
};
invalid(McpToolScopes::new().try_default_claim("", ["x"]), "");
invalid(McpToolScopes::new().try_default_claim(" ", ["x"]), " ");
invalid(
McpToolScopes::new().try_default_claim("g", [] as [&str; 0]),
"g",
);
invalid(
McpToolScopes::new().try_tool_claim("t", "g", ["x", " "]),
"g",
);
invalid(McpToolScopes::new().try_tool_claim("t", "g", [""]), "g");
let ok = McpToolScopes::new()
.try_default_claim("g", ["x"])
.and_then(|m| m.try_tool_claim("t", "g", ["y"]))
.unwrap();
assert_eq!(ok.rules.default.claims, [clause("g", &["x"])]);
}
#[test]
#[should_panic(expected = "McpToolScopes::tool_claim")]
fn an_invalid_claim_requirement_panics() {
let _ = McpToolScopes::new().tool_claim("t", "groups", [] as [&str; 0]);
}
}