use base64::engine::general_purpose::STANDARD as BASE64_STANDARD;
use base64::Engine as _;
use serde::{de, Deserialize, Deserializer, Serialize};
use serde_json::Value;
use sha2::{Digest, Sha256};
use std::collections::BTreeMap;
use crate::types::elicitation::{ElicitRequestParams, ElicitResult};
use crate::types::roots::ListRootsResult;
use crate::types::sampling::{CreateMessageParams, CreateMessageResult};
pub(crate) const INPUT_RESPONSES_KEY: &str = "inputResponses";
pub(crate) const REQUEST_STATE_KEY: &str = "requestState";
pub(crate) const INPUT_REQUESTS_KEY: &str = "inputRequests";
pub(crate) const META_KEY: &str = "_meta";
pub(crate) const RESULT_TYPE_KEY: &str = "resultType";
pub(crate) const INPUT_REQUIRED_RESULT_TYPE: &str = "input_required";
pub(crate) const COMPLETE_RESULT_TYPE: &str = "complete";
pub(crate) const TASK_RESULT_TYPE: &str = "task";
pub(crate) struct MrtrMethod {
pub method: &'static str,
pub name_key: &'static str,
pub salient: &'static [&'static str],
}
pub(crate) const MRTR_METHODS: [MrtrMethod; 3] = [
MrtrMethod {
method: CALL_TOOL_METHOD,
name_key: "name",
salient: &["name", "arguments"],
},
MrtrMethod {
method: GET_PROMPT_METHOD,
name_key: "name",
salient: &["name", "arguments"],
},
MrtrMethod {
method: READ_RESOURCE_METHOD,
name_key: "uri",
salient: &["uri"],
},
];
pub(crate) const CALL_TOOL_METHOD: &str = "tools/call";
pub(crate) const GET_PROMPT_METHOD: &str = "prompts/get";
pub(crate) const READ_RESOURCE_METHOD: &str = "resources/read";
pub(crate) const TASKS_GET_METHOD: &str = "tasks/get";
pub(crate) const TASKS_UPDATE_METHOD: &str = "tasks/update";
pub(crate) const TASKS_CANCEL_METHOD: &str = "tasks/cancel";
pub(crate) const TASK_ID_KEY: &str = "taskId";
pub(crate) const TASK_NAME_BEARING_METHODS: [(&str, &str); 3] = [
(TASKS_GET_METHOD, TASK_ID_KEY),
(TASKS_UPDATE_METHOD, TASK_ID_KEY),
(TASKS_CANCEL_METHOD, TASK_ID_KEY),
];
fn mrtr_row(method: &str) -> Option<&'static MrtrMethod> {
MRTR_METHODS.iter().find(|row| row.method == method)
}
pub(crate) fn mrtr_eligible(method: &str) -> bool {
mrtr_row(method).is_some()
}
pub(crate) fn mrtr_method_static(method: &str) -> Option<&'static str> {
Some(mrtr_row(method)?.method)
}
fn logical_name_key(method: &str) -> Option<&'static str> {
Some(mrtr_row(method)?.name_key)
}
pub(crate) fn name_bearing_key(method: &str) -> Option<&'static str> {
if let Some(key) = logical_name_key(method) {
return Some(key);
}
TASK_NAME_BEARING_METHODS
.iter()
.find(|(table_method, _)| *table_method == method)
.map(|(_, key)| *key)
}
pub(crate) fn logical_name_of(method: &str, params: &Value) -> Option<String> {
let key = name_bearing_key(method)?;
params.get(key).and_then(Value::as_str).map(str::to_string)
}
pub(crate) fn frame_routing_pair(frame: &Value) -> Option<(&str, Option<String>)> {
let method = frame.get("method")?.as_str()?;
let name = frame
.get("params")
.and_then(|params| logical_name_of(method, params));
Some((method, name))
}
pub(crate) const HEADER_SENTINEL_PREFIX: &str = "=?base64?";
pub(crate) const HEADER_SENTINEL_SUFFIX: &str = "?=";
pub(crate) const MAX_HEADER_VALUE_LEN: usize = 8192;
pub(crate) const MAX_HEADER_SENTINEL_LEN: usize = HEADER_SENTINEL_PREFIX.len()
+ MAX_HEADER_VALUE_LEN.div_ceil(3) * 4
+ HEADER_SENTINEL_SUFFIX.len();
fn header_byte_is_safe(byte: u8) -> bool {
(0x20..=0x7E).contains(&byte) && !matches!(byte, b'"' | b',' | b';' | b'\\')
}
pub(crate) fn encode_header_value(value: &str) -> String {
let passthrough = value.bytes().all(header_byte_is_safe)
&& !value.starts_with(HEADER_SENTINEL_PREFIX)
&& value.len() <= MAX_HEADER_VALUE_LEN;
if passthrough {
return value.to_string();
}
format!(
"{HEADER_SENTINEL_PREFIX}{}{HEADER_SENTINEL_SUFFIX}",
BASE64_STANDARD.encode(value.as_bytes())
)
}
pub(crate) fn decode_header_value(raw: &str) -> Option<String> {
let Some(rest) = raw.strip_prefix(HEADER_SENTINEL_PREFIX) else {
if raw.len() > MAX_HEADER_VALUE_LEN {
return None;
}
return Some(raw.to_string());
};
if raw.len() > MAX_HEADER_SENTINEL_LEN {
return None;
}
let payload = rest.strip_suffix(HEADER_SENTINEL_SUFFIX)?;
let bytes = BASE64_STANDARD.decode(payload).ok()?;
if bytes.len() > MAX_HEADER_VALUE_LEN {
return None;
}
String::from_utf8(bytes).ok()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum InputRequestKind {
#[serde(rename = "elicitation")]
Elicitation,
#[serde(rename = "sampling")]
Sampling,
#[serde(rename = "roots")]
Roots,
}
impl InputRequestKind {
#[must_use]
pub const fn wire_method(self) -> &'static str {
match self {
Self::Elicitation => "elicitation/create",
Self::Sampling => "sampling/createMessage",
Self::Roots => "roots/list",
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "method", content = "params", rename_all = "camelCase")]
pub enum InputRequest {
#[serde(rename = "elicitation/create")]
Elicitation(Box<ElicitRequestParams>),
#[serde(rename = "sampling/createMessage")]
Sampling(Box<CreateMessageParams>),
#[serde(rename = "roots/list")]
ListRoots,
}
impl InputRequest {
#[must_use]
pub const fn kind(&self) -> InputRequestKind {
match self {
Self::Elicitation(_) => InputRequestKind::Elicitation,
Self::Sampling(_) => InputRequestKind::Sampling,
Self::ListRoots => InputRequestKind::Roots,
}
}
}
#[derive(Debug, Clone, Serialize)]
#[serde(untagged)]
pub enum InputResponse {
Elicitation(Box<ElicitResult>),
Sampling(Box<CreateMessageResult>),
Roots(Box<ListRootsResult>),
}
impl InputResponse {
pub fn decode_for(kind: InputRequestKind, value: Value) -> Result<Self, serde_json::Error> {
match kind {
InputRequestKind::Elicitation => {
serde_json::from_value(value).map(|r| Self::Elicitation(Box::new(r)))
},
InputRequestKind::Sampling => {
serde_json::from_value(value).map(|r| Self::Sampling(Box::new(r)))
},
InputRequestKind::Roots => {
serde_json::from_value(value).map(|r| Self::Roots(Box::new(r)))
},
}
}
pub fn try_from_value_untagged(value: Value) -> Result<Self, serde_json::Error> {
if let Ok(decoded) = Self::decode_for(InputRequestKind::Roots, value.clone()) {
return Ok(decoded);
}
if let Ok(decoded) = Self::decode_for(InputRequestKind::Sampling, value.clone()) {
return Ok(decoded);
}
Self::decode_for(InputRequestKind::Elicitation, value).map_err(|_| {
de::Error::custom(
"inputResponses value matches none of ElicitResult, CreateMessageResult, ListRootsResult",
)
})
}
}
impl<'de> Deserialize<'de> for InputResponse {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let value = Value::deserialize(deserializer)?;
Self::try_from_value_untagged(value).map_err(de::Error::custom)
}
}
pub type InputRequests = BTreeMap<String, InputRequest>;
pub type InputResponses = BTreeMap<String, InputResponse>;
pub(crate) type InputRequestKinds = BTreeMap<String, InputRequestKind>;
#[derive(Debug, Clone, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct InputRequiredResult {
pub result_type: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub input_requests: Option<InputRequests>,
#[serde(skip_serializing_if = "Option::is_none")]
pub request_state: Option<String>,
#[serde(rename = "_meta", skip_serializing_if = "Option::is_none")]
pub meta: Option<Value>,
#[serde(skip_serializing)]
pub raw: Value,
}
impl InputRequiredResult {
#[must_use]
pub fn is_input_required(&self) -> bool {
self.result_type == INPUT_REQUIRED_RESULT_TYPE
}
}
#[derive(Deserialize)]
#[serde(rename_all = "camelCase")]
struct InputRequiredShape {
#[serde(default)]
result_type: String,
#[serde(default)]
input_requests: Option<InputRequests>,
#[serde(default)]
request_state: Option<String>,
#[serde(default, rename = "_meta")]
meta: Option<Value>,
}
impl<'de> Deserialize<'de> for InputRequiredResult {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let raw = Value::deserialize(deserializer)?;
let shape = InputRequiredShape::deserialize(&raw).map_err(de::Error::custom)?;
Ok(Self {
result_type: shape.result_type,
input_requests: shape.input_requests,
request_state: shape.request_state,
meta: shape.meta,
raw,
})
}
}
#[allow(clippy::large_enum_variant)]
#[derive(Debug, Clone)]
pub enum MrtrOutcome<T> {
Complete(T),
InputRequired(InputRequiredResult),
}
impl<T> MrtrOutcome<T> {
#[must_use]
pub fn complete(self) -> Option<T> {
match self {
Self::Complete(value) => Some(value),
Self::InputRequired(_) => None,
}
}
#[must_use]
pub fn input_required(self) -> Option<InputRequiredResult> {
match self {
Self::Complete(_) => None,
Self::InputRequired(result) => Some(result),
}
}
}
pub const MRTR_SIGNAL_META_KEY: &str = "dev.pmcp/mrtr";
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct MrtrSignal {
pub input_requests: InputRequests,
#[serde(default)]
pub continuation: Value,
}
impl MrtrSignal {
pub fn into_meta_entry(self) -> Result<(String, Value), serde_json::Error> {
Ok((
MRTR_SIGNAL_META_KEY.to_string(),
serde_json::to_value(self)?,
))
}
}
pub(crate) fn remove_mrtr_signal(result: &mut Value) -> Option<Value> {
let object = result.as_object_mut()?;
let removed = object
.get_mut(META_KEY)
.and_then(Value::as_object_mut)
.and_then(|meta| meta.remove(MRTR_SIGNAL_META_KEY))?;
if object
.get(META_KEY)
.and_then(Value::as_object)
.is_some_and(serde_json::Map::is_empty)
{
object.remove(META_KEY);
}
Some(removed)
}
pub(crate) const MAX_REQUEST_STATE_LEN: usize = 8192;
pub(crate) const MAX_INPUT_RESPONSES: usize = 64;
pub(crate) const MAX_INPUT_RESPONSE_BYTES: usize = 65_536;
pub(crate) const MAX_INPUT_RESPONSES_TOTAL_BYTES: usize = 262_144;
pub(crate) const MAX_INPUT_RESPONSE_DEPTH: usize = 32;
pub(crate) const MAX_CANONICAL_DEPTH: usize = 64;
#[derive(Debug, Clone, Default)]
pub(crate) struct MrtrRequestParams {
pub input_responses: Option<InputResponses>,
pub input_responses_raw: Option<serde_json::Map<String, Value>>,
pub request_state: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) enum MrtrParseError {
RequestStateNotAString,
RequestStateTooLong {
len: usize,
max: usize,
},
InputResponsesNotAnObject,
TooManyInputResponses {
count: usize,
max: usize,
},
InputResponseTooLarge {
key: String,
bytes: usize,
max: usize,
},
InputResponsesTotalTooLarge {
bytes: usize,
max: usize,
},
InputResponseTooDeep {
key: String,
depth: usize,
max: usize,
},
InputResponseUndecodable {
key: String,
},
}
impl std::fmt::Display for MrtrParseError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::RequestStateNotAString => write!(f, "requestState must be a string"),
Self::RequestStateTooLong { max, .. } => {
write!(f, "requestState exceeds the {max}-byte limit")
},
Self::InputResponsesNotAnObject => write!(f, "inputResponses must be an object"),
Self::TooManyInputResponses { max, .. } => {
write!(f, "inputResponses exceeds the {max}-entry limit")
},
Self::InputResponseTooLarge { max, .. } => {
write!(f, "an inputResponses entry exceeds the {max}-byte limit")
},
Self::InputResponsesTotalTooLarge { max, .. } => {
write!(f, "inputResponses exceeds the {max}-byte total limit")
},
Self::InputResponseTooDeep { max, .. } => {
write!(f, "an inputResponses entry exceeds the {max}-level depth limit")
},
Self::InputResponseUndecodable { .. } => write!(
f,
"an inputResponses entry is not a valid ElicitResult, CreateMessageResult or ListRootsResult"
),
}
}
}
impl std::error::Error for MrtrParseError {}
fn json_depth(value: &Value) -> usize {
let mut deepest = 0usize;
let mut stack = vec![(value, 1usize)];
while let Some((current, depth)) = stack.pop() {
deepest = deepest.max(depth);
if depth > MAX_INPUT_RESPONSE_DEPTH {
return depth;
}
match current {
Value::Array(items) => stack.extend(items.iter().map(|item| (item, depth + 1))),
Value::Object(entries) => {
stack.extend(entries.iter().map(|(_, item)| (item, depth + 1)));
},
_ => {},
}
}
deepest
}
fn extract_request_state(
params: &serde_json::Map<String, Value>,
) -> Result<Option<String>, MrtrParseError> {
let Some(value) = params.get(REQUEST_STATE_KEY) else {
return Ok(None);
};
let state = value
.as_str()
.ok_or(MrtrParseError::RequestStateNotAString)?;
if state.len() > MAX_REQUEST_STATE_LEN {
return Err(MrtrParseError::RequestStateTooLong {
len: state.len(),
max: MAX_REQUEST_STATE_LEN,
});
}
Ok(Some(state.to_string()))
}
pub(crate) fn check_input_response_bounds(
key: &str,
value: &Value,
) -> Result<usize, MrtrParseError> {
let depth = json_depth(value);
if depth > MAX_INPUT_RESPONSE_DEPTH {
return Err(MrtrParseError::InputResponseTooDeep {
key: key.to_string(),
depth,
max: MAX_INPUT_RESPONSE_DEPTH,
});
}
let bytes = serde_json::to_string(value).map_or(usize::MAX, |s| s.len());
if bytes > MAX_INPUT_RESPONSE_BYTES {
return Err(MrtrParseError::InputResponseTooLarge {
key: key.to_string(),
bytes,
max: MAX_INPUT_RESPONSE_BYTES,
});
}
Ok(bytes)
}
pub(crate) fn check_input_responses_map_bounds(
entries: &serde_json::Map<String, Value>,
) -> Result<(), MrtrParseError> {
if entries.len() > MAX_INPUT_RESPONSES {
return Err(MrtrParseError::TooManyInputResponses {
count: entries.len(),
max: MAX_INPUT_RESPONSES,
});
}
let mut total = 0usize;
for (key, entry) in entries {
total = total.saturating_add(check_input_response_bounds(key, entry)?);
if total > MAX_INPUT_RESPONSES_TOTAL_BYTES {
return Err(MrtrParseError::InputResponsesTotalTooLarge {
bytes: total,
max: MAX_INPUT_RESPONSES_TOTAL_BYTES,
});
}
}
Ok(())
}
#[allow(clippy::type_complexity)]
fn extract_input_responses(
params: &serde_json::Map<String, Value>,
) -> Result<Option<(InputResponses, serde_json::Map<String, Value>)>, MrtrParseError> {
let Some(value) = params.get(INPUT_RESPONSES_KEY) else {
return Ok(None);
};
let entries = value
.as_object()
.ok_or(MrtrParseError::InputResponsesNotAnObject)?;
check_input_responses_map_bounds(entries)?;
let mut decoded = InputResponses::new();
let mut raw = serde_json::Map::new();
for (key, entry) in entries {
let response = InputResponse::try_from_value_untagged(entry.clone())
.map_err(|_| MrtrParseError::InputResponseUndecodable { key: key.clone() })?;
decoded.insert(key.clone(), response);
raw.insert(key.clone(), entry.clone());
}
Ok(Some((decoded, raw)))
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) enum InputResponseTypingError {
KindMismatch {
key: String,
expected: InputRequestKind,
},
Unsolicited {
key: String,
},
}
impl std::fmt::Display for InputResponseTypingError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::KindMismatch { key, expected } => write!(
f,
"the inputResponses entry for {key:?} is not a valid response to the \
{} request the server made under that key",
expected.wire_method()
),
Self::Unsolicited { .. } => write!(
f,
"inputResponses carries an entry under a key this continuation never requested"
),
}
}
}
impl std::error::Error for InputResponseTypingError {}
pub(crate) fn retype_input_responses_for_kinds(
raw: &serde_json::Map<String, Value>,
kinds: Option<&InputRequestKinds>,
) -> Result<Option<InputResponses>, InputResponseTypingError> {
let Some(kinds) = kinds else {
return Ok(None);
};
let mut typed = InputResponses::new();
for (key, value) in raw {
let Some((sealed_key, kind)) = kinds.get_key_value(key) else {
return Err(InputResponseTypingError::Unsolicited { key: key.clone() });
};
let response = InputResponse::decode_for(*kind, value.clone()).map_err(|_| {
InputResponseTypingError::KindMismatch {
key: sealed_key.clone(),
expected: *kind,
}
})?;
typed.insert(key.clone(), response);
}
Ok(Some(typed))
}
pub(crate) fn extract_mrtr_params(params: &Value) -> Result<MrtrRequestParams, MrtrParseError> {
let Some(object) = params.as_object() else {
return Ok(MrtrRequestParams::default());
};
let (input_responses, input_responses_raw) = match extract_input_responses(object)? {
Some((typed, raw)) => (Some(typed), Some(raw)),
None => (None, None),
};
Ok(MrtrRequestParams {
input_responses,
input_responses_raw,
request_state: extract_request_state(object)?,
})
}
pub(crate) fn splice_mrtr_params(params: &mut Value, mrtr: &MrtrRequestParams) {
let Some(object) = params.as_object_mut() else {
return;
};
object.remove(INPUT_RESPONSES_KEY);
object.remove(REQUEST_STATE_KEY);
if let Some(responses) = mrtr.input_responses.as_ref() {
if let Ok(value) = serde_json::to_value(responses) {
object.insert(INPUT_RESPONSES_KEY.to_string(), value);
}
}
if let Some(state) = mrtr.request_state.as_ref() {
object.insert(REQUEST_STATE_KEY.to_string(), Value::String(state.clone()));
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct CanonicalDepthExceeded {
pub depth: usize,
pub max: usize,
}
impl std::fmt::Display for CanonicalDepthExceeded {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"request params exceed the {}-level canonicalization depth limit",
self.max
)
}
}
impl std::error::Error for CanonicalDepthExceeded {}
fn write_canonical(
value: &Value,
depth: usize,
out: &mut String,
) -> Result<(), CanonicalDepthExceeded> {
if depth > MAX_CANONICAL_DEPTH {
return Err(CanonicalDepthExceeded {
depth,
max: MAX_CANONICAL_DEPTH,
});
}
match value {
Value::Object(entries) => write_canonical_object(entries, depth, out),
Value::Array(items) => write_canonical_array(items, depth, out),
other => {
out.push_str(&other.to_string());
Ok(())
},
}
}
fn write_canonical_object(
entries: &serde_json::Map<String, Value>,
depth: usize,
out: &mut String,
) -> Result<(), CanonicalDepthExceeded> {
let mut keys: Vec<&String> = entries.keys().collect();
keys.sort_unstable();
out.push('{');
for (index, key) in keys.iter().enumerate() {
if index > 0 {
out.push(',');
}
out.push_str(&Value::String((*key).clone()).to_string());
out.push(':');
write_canonical(&entries[key.as_str()], depth + 1, out)?;
}
out.push('}');
Ok(())
}
fn write_canonical_array(
items: &[Value],
depth: usize,
out: &mut String,
) -> Result<(), CanonicalDepthExceeded> {
out.push('[');
for (index, item) in items.iter().enumerate() {
if index > 0 {
out.push(',');
}
write_canonical(item, depth + 1, out)?;
}
out.push(']');
Ok(())
}
fn salient_params(method: &str, params: &Value) -> Value {
let mut salient = serde_json::Map::new();
let keys: &[&str] = mrtr_row(method).map_or(&[], |row| row.salient);
for key in keys {
if let Some(value) = params.get(*key) {
salient.insert((*key).to_string(), value.clone());
}
}
Value::Object(salient)
}
pub(crate) fn salient_param_digest(
method: &str,
params: &Value,
) -> Result<[u8; 32], CanonicalDepthExceeded> {
let mut canonical = String::new();
write_canonical(&salient_params(method, params), 0, &mut canonical)?;
let mut hasher = Sha256::new();
hasher.update(method.as_bytes());
hasher.update([0u8]);
hasher.update(canonical.as_bytes());
let output = hasher.finalize();
let mut digest = [0u8; 32];
digest.copy_from_slice(&output);
Ok(digest)
}
#[cfg(test)]
mod kind_directed_tests {
use super::*;
use serde_json::json;
fn overlapping_answer() -> Value {
json!({
"action": "accept",
"content": { "type": "text", "text": "hello" },
"model": "attacker-chosen-model",
})
}
fn raw_map(entries: &[(&str, Value)]) -> serde_json::Map<String, Value> {
entries
.iter()
.map(|(key, value)| ((*key).to_string(), value.clone()))
.collect()
}
fn kinds_of(entries: &[(&str, InputRequestKind)]) -> InputRequestKinds {
entries
.iter()
.map(|(key, kind)| ((*key).to_string(), *kind))
.collect()
}
#[test]
fn the_untagged_decoder_still_reclassifies_the_overlapping_answer() {
let decoded = InputResponse::try_from_value_untagged(overlapping_answer())
.expect("the overlapping value decodes as something");
assert!(
matches!(decoded, InputResponse::Sampling(_)),
"Sampling is tried before Elicitation, which is the whole of D-113-O"
);
}
#[test]
fn the_literal_d113o_answer_is_typed_as_the_elicitation_it_answers() {
let typed = retype_input_responses_for_kinds(
&raw_map(&[("k", overlapping_answer())]),
Some(&kinds_of(&[("k", InputRequestKind::Elicitation)])),
)
.expect("a valid ElicitResult answered to an elicitation is not an error")
.expect("a non-None kinds map produces a kind-directed map");
assert!(
matches!(typed["k"], InputResponse::Elicitation(_)),
"kind-directed typing must follow what the SERVER asked for, not which \
overlapping shape happens to be tried first"
);
}
#[test]
fn an_answer_that_cannot_be_the_requested_kind_is_rejected_naming_the_key() {
let sampling_only = json!({
"content": { "type": "text", "text": "hello" },
"model": "attacker-chosen-model",
});
let error = retype_input_responses_for_kinds(
&raw_map(&[("k", sampling_only)]),
Some(&kinds_of(&[("k", InputRequestKind::Elicitation)])),
)
.expect_err("an answer that is not an ElicitResult must be REJECTED");
assert_eq!(
error,
InputResponseTypingError::KindMismatch {
key: "k".to_string(),
expected: InputRequestKind::Elicitation,
}
);
let rendered = error.to_string();
assert!(
rendered.contains("\"k\""),
"the message must NAME the key, which came from the sealed continuation \
and is what makes the error actionable: {rendered}"
);
assert!(
rendered.contains("elicitation/create"),
"...and the kind that was actually requested there: {rendered}"
);
assert!(
!rendered.contains("attacker-chosen-model") && !rendered.contains("hello"),
"...and must never echo the VALUE, which is attacker-controlled: {rendered}"
);
}
#[test]
fn a_correctly_shaped_answer_decodes_to_the_requested_kind() {
let raw = raw_map(&[
("ask", json!({ "action": "accept", "content": { "v": 1 } })),
(
"model",
json!({ "content": { "type": "text", "text": "hi" }, "model": "m" }),
),
("roots", json!({ "roots": [] })),
]);
let kinds = kinds_of(&[
("ask", InputRequestKind::Elicitation),
("model", InputRequestKind::Sampling),
("roots", InputRequestKind::Roots),
]);
let typed = retype_input_responses_for_kinds(&raw, Some(&kinds))
.expect("well-shaped answers are accepted")
.expect("a non-None kinds map produces a kind-directed map");
assert!(matches!(typed["ask"], InputResponse::Elicitation(_)));
assert!(matches!(typed["model"], InputResponse::Sampling(_)));
assert!(matches!(typed["roots"], InputResponse::Roots(_)));
}
#[test]
fn a_sampling_request_answered_with_an_elicitation_shape_is_rejected() {
let error = retype_input_responses_for_kinds(
&raw_map(&[("model", json!({ "action": "decline" }))]),
Some(&kinds_of(&[("model", InputRequestKind::Sampling)])),
)
.expect_err("a wrongly-shaped answer must be REJECTED");
assert_eq!(
error,
InputResponseTypingError::KindMismatch {
key: "model".to_string(),
expected: InputRequestKind::Sampling,
}
);
}
#[test]
fn an_unsolicited_key_is_rejected_without_being_echoed() {
let client_chosen = "zzz_client_chosen_key_zzz";
let error = retype_input_responses_for_kinds(
&raw_map(&[(client_chosen, json!({ "action": "accept" }))]),
Some(&kinds_of(&[(
"something_else",
InputRequestKind::Elicitation,
)])),
)
.expect_err("a key the server never asked about must be REJECTED");
assert_eq!(
error,
InputResponseTypingError::Unsolicited {
key: client_chosen.to_string(),
},
"the key is carried for programmatic use..."
);
assert!(
!error.to_string().contains(client_chosen),
"...but never rendered: it is client-chosen and bounded only by the 256 KiB \
inputResponses total, so echoing it would both amplify and poison logs — the \
same discipline MrtrParseError's Display already applies: {error}"
);
}
#[test]
fn an_empty_kinds_map_rejects_every_answer_rather_than_degrading() {
let error = retype_input_responses_for_kinds(
&raw_map(&[("k", overlapping_answer())]),
Some(&InputRequestKinds::new()),
)
.expect_err("a round that requested nothing can be answered with nothing");
assert!(matches!(
error,
InputResponseTypingError::Unsolicited { .. }
));
}
#[test]
fn an_absent_kinds_map_degrades_to_untagged_without_rejecting() {
let retyped =
retype_input_responses_for_kinds(&raw_map(&[("k", overlapping_answer())]), None)
.expect("a pre-kinds continuation must never reject");
assert!(
retyped.is_none(),
"None means \"keep what ingress guessed\" — the caller's degradation branch"
);
}
#[test]
fn answering_nothing_is_not_a_mismatch() {
let typed = retype_input_responses_for_kinds(
&serde_json::Map::new(),
Some(&kinds_of(&[("k", InputRequestKind::Elicitation)])),
)
.expect("answering nothing is not an error")
.expect("a non-None kinds map produces a map");
assert!(typed.is_empty());
}
#[test]
fn ingress_retains_the_raw_entries_verbatim() {
let params = json!({
"name": "elicit_once",
"arguments": {},
"inputResponses": { "k": overlapping_answer() },
});
let extracted = extract_mrtr_params(¶ms).expect("ingress accepts it");
let raw = extracted
.input_responses_raw
.expect("the raw entries are retained");
assert_eq!(raw["k"], overlapping_answer());
assert!(
matches!(
extracted.input_responses.expect("typed")["k"],
InputResponse::Sampling(_)
),
"the TYPED map is still the untagged guess at this layer — correcting it \
needs the continuation, which ingress has not opened yet"
);
}
#[test]
fn the_ingress_bounds_still_fire_before_the_raw_retention() {
let mut over_count = serde_json::Map::new();
for index in 0..=MAX_INPUT_RESPONSES {
over_count.insert(format!("k{index}"), json!({ "roots": [] }));
}
assert!(matches!(
extract_mrtr_params(&json!({ "inputResponses": over_count })),
Err(MrtrParseError::TooManyInputResponses {
count,
max: MAX_INPUT_RESPONSES,
}) if count == MAX_INPUT_RESPONSES + 1
));
let huge = json!({ "roots": [], "pad": "x".repeat(MAX_INPUT_RESPONSE_BYTES) });
assert!(matches!(
extract_mrtr_params(&json!({ "inputResponses": { "k": huge } })),
Err(MrtrParseError::InputResponseTooLarge { .. })
));
let mut deep = json!({ "roots": [] });
for _ in 0..=MAX_INPUT_RESPONSE_DEPTH {
deep = json!({ "n": deep });
}
assert!(matches!(
extract_mrtr_params(&json!({ "inputResponses": { "k": deep } })),
Err(MrtrParseError::InputResponseTooDeep { .. })
));
let each = MAX_INPUT_RESPONSE_BYTES / 2;
let mut many = serde_json::Map::new();
for index in 0..MAX_INPUT_RESPONSES {
many.insert(
format!("k{index}"),
json!({ "roots": [], "pad": "x".repeat(each) }),
);
}
assert!(matches!(
extract_mrtr_params(&json!({ "inputResponses": many })),
Err(MrtrParseError::InputResponsesTotalTooLarge { .. })
));
}
#[test]
fn the_sealed_kind_spelling_is_pinned() {
for (kind, expected) in [
(InputRequestKind::Elicitation, "elicitation"),
(InputRequestKind::Sampling, "sampling"),
(InputRequestKind::Roots, "roots"),
] {
assert_eq!(serde_json::to_value(kind).unwrap(), json!(expected));
assert_eq!(
serde_json::from_value::<InputRequestKind>(json!(expected)).unwrap(),
kind
);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn form_elicitation() -> ElicitRequestParams {
ElicitRequestParams::Form {
message: "What is your name?".to_string(),
requested_schema: json!({ "type": "object" }),
}
}
#[test]
fn mrtr_method_constants_match_the_table() {
assert_eq!(CALL_TOOL_METHOD, "tools/call");
assert_eq!(GET_PROMPT_METHOD, "prompts/get");
assert_eq!(READ_RESOURCE_METHOD, "resources/read");
for method in [CALL_TOOL_METHOD, GET_PROMPT_METHOD, READ_RESOURCE_METHOD] {
assert!(mrtr_eligible(method), "{method} must be in the table");
}
assert_eq!(MRTR_METHODS.len(), 3, "a new row needs a new constant");
}
#[test]
fn input_request_elicitation_wire_shape() {
let request = InputRequest::Elicitation(Box::new(form_elicitation()));
let value = serde_json::to_value(&request).unwrap();
assert_eq!(value["method"], "elicitation/create");
assert_eq!(value["params"]["message"], "What is your name?");
assert_eq!(request.kind(), InputRequestKind::Elicitation);
}
#[test]
fn input_request_deserializes_all_three_methods() {
let elicit: InputRequest = serde_json::from_value(json!({
"method": "elicitation/create",
"params": { "mode": "form", "message": "hi", "requestedSchema": {} }
}))
.unwrap();
assert_eq!(elicit.kind(), InputRequestKind::Elicitation);
let sampling: InputRequest = serde_json::from_value(json!({
"method": "sampling/createMessage",
"params": { "messages": [], "maxTokens": 16 }
}))
.unwrap();
assert_eq!(sampling.kind(), InputRequestKind::Sampling);
let roots: InputRequest =
serde_json::from_value(json!({ "method": "roots/list" })).unwrap();
assert_eq!(roots.kind(), InputRequestKind::Roots);
}
#[test]
fn input_request_rejects_unknown_method() {
let result: Result<InputRequest, _> =
serde_json::from_value(json!({ "method": "tools/list", "params": {} }));
assert!(result.is_err());
}
#[test]
fn input_request_kind_wire_methods() {
assert_eq!(
InputRequestKind::Elicitation.wire_method(),
"elicitation/create"
);
assert_eq!(
InputRequestKind::Sampling.wire_method(),
"sampling/createMessage"
);
assert_eq!(InputRequestKind::Roots.wire_method(), "roots/list");
}
fn elicit_result_value() -> Value {
json!({ "action": "accept", "content": { "user_name": "Alice" } })
}
fn sampling_result_value() -> Value {
json!({
"content": { "type": "text", "text": "hello" },
"model": "test-model"
})
}
fn roots_result_value() -> Value {
json!({ "roots": [] })
}
#[test]
fn decode_for_elicitation_accepts_elicit_rejects_sampling() {
let ok = InputResponse::decode_for(InputRequestKind::Elicitation, elicit_result_value());
assert!(matches!(ok, Ok(InputResponse::Elicitation(_))));
let err = InputResponse::decode_for(InputRequestKind::Elicitation, sampling_result_value());
assert!(
err.is_err(),
"a CreateMessageResult must not decode as an ElicitResult"
);
}
#[test]
fn decode_for_sampling_accepts_sampling_rejects_elicit() {
let ok = InputResponse::decode_for(InputRequestKind::Sampling, sampling_result_value());
assert!(matches!(ok, Ok(InputResponse::Sampling(_))));
let err = InputResponse::decode_for(InputRequestKind::Sampling, elicit_result_value());
assert!(
err.is_err(),
"an ElicitResult must not decode as a CreateMessageResult"
);
}
#[test]
fn decode_for_roots_accepts_list_roots_result() {
let ok = InputResponse::decode_for(InputRequestKind::Roots, roots_result_value());
assert!(matches!(ok, Ok(InputResponse::Roots(_))));
}
#[test]
fn input_response_serializes_as_a_bare_result_object() {
let response =
InputResponse::decode_for(InputRequestKind::Elicitation, elicit_result_value())
.unwrap();
let value = serde_json::to_value(&response).unwrap();
assert_eq!(value["action"], "accept");
assert!(
value.get("Elicitation").is_none(),
"must be untagged on the wire"
);
}
#[test]
fn untagged_decode_is_best_effort_over_the_three_shapes() {
assert!(matches!(
InputResponse::try_from_value_untagged(roots_result_value()),
Ok(InputResponse::Roots(_))
));
assert!(matches!(
InputResponse::try_from_value_untagged(sampling_result_value()),
Ok(InputResponse::Sampling(_))
));
assert!(matches!(
InputResponse::try_from_value_untagged(elicit_result_value()),
Ok(InputResponse::Elicitation(_))
));
assert!(InputResponse::try_from_value_untagged(json!({ "nope": 1 })).is_err());
}
#[test]
fn mrtr_eligible_is_exactly_three_methods() {
for method in ["tools/call", "prompts/get", "resources/read"] {
assert!(mrtr_eligible(method), "{method} must be MRTR-eligible");
}
for method in [
"tools/list",
"server/discover",
"completion/complete",
"subscriptions/listen",
"initialize",
] {
assert!(!mrtr_eligible(method), "{method} must NOT be MRTR-eligible");
}
assert_eq!(MRTR_METHODS.len(), 3);
}
#[test]
fn every_mrtr_method_row_is_completely_populated() {
for row in &MRTR_METHODS {
assert!(
!row.salient.is_empty(),
"{} is MRTR-eligible but has no salient params — its AAD digest \
would degrade to the empty object and stop binding the request",
row.method
);
assert!(
row.salient.contains(&row.name_key),
"{}'s logical-name key {:?} must be digest-salient, or a token \
minted for one name would verify against another",
row.method,
row.name_key
);
assert!(
!row.method.is_empty() && !row.name_key.is_empty(),
"{} has an empty method or name_key",
row.method
);
}
}
#[test]
fn logical_name_key_table() {
assert_eq!(logical_name_key("tools/call"), Some("name"));
assert_eq!(logical_name_key("prompts/get"), Some("name"));
assert_eq!(logical_name_key("resources/read"), Some("uri"));
assert_eq!(logical_name_key("tools/list"), None);
}
#[test]
fn tasks_methods_are_name_bearing_but_not_mrtr_eligible() {
for method in [TASKS_GET_METHOD, TASKS_UPDATE_METHOD, TASKS_CANCEL_METHOD] {
assert_eq!(
name_bearing_key(method),
Some(TASK_ID_KEY),
"{method} must route on params.taskId"
);
assert!(
!mrtr_eligible(method),
"{method} must NOT be MRTR-eligible — an MRTR_METHODS row would make \
splice_mrtr_params delete its inputResponses payload"
);
}
assert_eq!(TASK_NAME_BEARING_METHODS.len(), 3);
assert!(
TASK_NAME_BEARING_METHODS
.iter()
.all(|(_, key)| *key == TASK_ID_KEY),
"every tasks row maps to taskId"
);
}
#[test]
fn splice_mrtr_params_would_delete_a_tasks_update_payload() {
let mut params = json!({
"taskId": "abc",
"inputResponses": { "k": { "action": "accept", "content": { "answer": 1 } } },
});
splice_mrtr_params(&mut params, &MrtrRequestParams::default());
assert!(
params.get(INPUT_RESPONSES_KEY).is_none(),
"the strip is unconditional — this is why tasks/update must stay \
OUT of MRTR_METHODS, got {params}"
);
assert_eq!(
params["taskId"], "abc",
"only the MRTR fields are stripped; the routing key survives"
);
}
#[test]
fn tasks_list_and_result_are_not_name_bearing() {
assert_eq!(name_bearing_key("tasks/list"), None);
assert_eq!(name_bearing_key("tasks/result"), None);
}
#[test]
fn the_two_name_key_tables_are_disjoint() {
for (method, _) in TASK_NAME_BEARING_METHODS {
assert_eq!(
logical_name_key(method),
None,
"{method} must not also live in MRTR_METHODS"
);
}
for row in &MRTR_METHODS {
assert!(
!TASK_NAME_BEARING_METHODS
.iter()
.any(|(method, _)| *method == row.method),
"{} must not also live in the tasks table",
row.method
);
}
}
#[test]
fn name_bearing_key_still_answers_for_the_mrtr_methods() {
assert_eq!(name_bearing_key("tools/call"), Some("name"));
assert_eq!(name_bearing_key("prompts/get"), Some("name"));
assert_eq!(name_bearing_key("resources/read"), Some("uri"));
assert_eq!(name_bearing_key("tools/list"), None);
}
#[test]
fn a_non_header_safe_task_id_round_trips_through_the_shared_codec() {
let task_id = "tenant;a,b\\c \u{2713}";
let encoded = encode_header_value(task_id);
assert!(
encoded.starts_with(HEADER_SENTINEL_PREFIX),
"a task id carrying RFC 9110 delimiters must travel as a sentinel, got {encoded}"
);
assert_eq!(decode_header_value(&encoded).as_deref(), Some(task_id));
let frame = json!({
"jsonrpc": "2.0",
"id": 1,
"method": TASKS_GET_METHOD,
"params": { TASK_ID_KEY: task_id },
});
assert_eq!(
frame_routing_pair(&frame),
Some((TASKS_GET_METHOD, Some(task_id.to_string())))
);
}
#[test]
fn logical_name_of_reads_the_method_specific_key() {
assert_eq!(
logical_name_of("tools/call", &json!({ "name": "search" })),
Some("search".to_string())
);
assert_eq!(
logical_name_of("resources/read", &json!({ "uri": "mem://a" })),
Some("mem://a".to_string())
);
assert_eq!(logical_name_of("tools/list", &json!({ "name": "x" })), None);
assert_eq!(logical_name_of("tools/call", &json!({})), None);
}
#[test]
fn encode_header_value_passes_safe_ascii_through() {
assert_eq!(encode_header_value("search"), "search");
assert_eq!(encode_header_value("mem://greeting"), "mem://greeting");
}
#[test]
fn encode_header_value_passes_the_empty_string_through() {
assert_eq!(encode_header_value(""), "");
}
#[test]
fn encode_header_value_sentinel_encodes_non_ascii() {
let encoded = encode_header_value("日本語");
assert!(encoded.starts_with("=?base64?"), "got {encoded}");
assert!(encoded.ends_with("?="), "got {encoded}");
}
#[test]
fn encode_header_value_sentinel_encodes_delimiters() {
for delimiter in ["a\"b", "a,b", "a;b", "a\\b"] {
let encoded = encode_header_value(delimiter);
assert!(
encoded.starts_with("=?base64?"),
"{delimiter} must be sentinel-encoded, got {encoded}"
);
assert_eq!(decode_header_value(&encoded).as_deref(), Some(delimiter));
}
}
#[test]
fn encode_header_value_escapes_a_value_that_looks_like_a_sentinel() {
let raw = "=?base64?abc?=";
let encoded = encode_header_value(raw);
assert_ne!(encoded, raw);
assert_eq!(decode_header_value(&encoded).as_deref(), Some(raw));
}
#[test]
fn decode_header_value_round_trips_ascii_non_ascii_and_empty() {
for value in ["search", "日本語", "", "mem://a/b?c=d"] {
assert_eq!(
decode_header_value(&encode_header_value(value)),
Some(value.to_string()),
"round trip failed for {value:?}"
);
}
}
#[test]
fn decode_header_value_rejects_a_malformed_sentinel() {
assert_eq!(decode_header_value("=?base64?not-valid-b64?="), None);
assert_eq!(decode_header_value("=?base64?no-suffix"), None);
}
fn responses_fixture() -> InputResponses {
let mut map = InputResponses::new();
map.insert(
"user_name".to_string(),
InputResponse::decode_for(InputRequestKind::Elicitation, elicit_result_value())
.unwrap(),
);
map
}
#[test]
fn splice_writes_top_level_siblings_not_meta_or_arguments() {
let mut params = json!({ "name": "search", "arguments": { "q": "x" }, "_meta": {} });
splice_mrtr_params(
&mut params,
&MrtrRequestParams {
input_responses: Some(responses_fixture()),
input_responses_raw: None,
request_state: Some("opaque".to_string()),
},
);
assert_eq!(params["inputResponses"]["user_name"]["action"], "accept");
assert_eq!(params["requestState"], "opaque");
assert!(params["arguments"].get("inputResponses").is_none());
assert!(params["_meta"].get("inputResponses").is_none());
assert!(params["_meta"].get("requestState").is_none());
assert_eq!(params["name"], "search");
}
#[test]
fn splice_default_removes_stale_keys() {
let mut params = json!({
"name": "search",
"inputResponses": { "stale": { "action": "accept" } },
"requestState": "round-1-token"
});
splice_mrtr_params(&mut params, &MrtrRequestParams::default());
assert!(params.get("inputResponses").is_none());
assert!(params.get("requestState").is_none());
assert_eq!(params["name"], "search");
}
#[test]
fn splice_is_a_noop_on_a_non_object() {
let mut params = json!([1, 2, 3]);
splice_mrtr_params(
&mut params,
&MrtrRequestParams {
input_responses: None,
input_responses_raw: None,
request_state: Some("x".to_string()),
},
);
assert_eq!(params, json!([1, 2, 3]));
}
fn assert_is_default(parsed: &MrtrRequestParams) {
assert!(parsed.input_responses.is_none());
assert!(parsed.request_state.is_none());
}
#[test]
fn extract_absent_keys_is_the_default() {
let params = json!({ "name": "search", "arguments": {} });
assert_is_default(&extract_mrtr_params(¶ms).unwrap());
}
#[test]
fn extract_non_object_params_is_the_default() {
assert_is_default(&extract_mrtr_params(&json!(null)).unwrap());
assert_is_default(&extract_mrtr_params(&json!([1, 2])).unwrap());
}
#[test]
fn extract_rejects_a_non_string_request_state() {
let err = extract_mrtr_params(&json!({ "requestState": 42 })).unwrap_err();
assert_eq!(err, MrtrParseError::RequestStateNotAString);
let err = extract_mrtr_params(&json!({ "requestState": null })).unwrap_err();
assert_eq!(err, MrtrParseError::RequestStateNotAString);
}
#[test]
fn extract_rejects_an_oversized_request_state() {
let big = "x".repeat(MAX_REQUEST_STATE_LEN + 1);
let err = extract_mrtr_params(&json!({ "requestState": big })).unwrap_err();
assert_eq!(
err,
MrtrParseError::RequestStateTooLong {
len: MAX_REQUEST_STATE_LEN + 1,
max: MAX_REQUEST_STATE_LEN,
}
);
}
#[test]
fn extract_rejects_a_non_object_input_responses() {
let err = extract_mrtr_params(&json!({ "inputResponses": [] })).unwrap_err();
assert_eq!(err, MrtrParseError::InputResponsesNotAnObject);
}
#[test]
fn extract_rejects_too_many_input_responses() {
let mut entries = serde_json::Map::new();
for index in 0..=MAX_INPUT_RESPONSES {
entries.insert(format!("k{index}"), elicit_result_value());
}
let err = extract_mrtr_params(&json!({ "inputResponses": entries })).unwrap_err();
assert_eq!(
err,
MrtrParseError::TooManyInputResponses {
count: MAX_INPUT_RESPONSES + 1,
max: MAX_INPUT_RESPONSES,
}
);
}
#[test]
fn extract_rejects_an_oversized_single_input_response() {
let huge = "y".repeat(MAX_INPUT_RESPONSE_BYTES + 1);
let params = json!({
"inputResponses": { "big": { "action": "accept", "content": { "v": huge } } }
});
let err = extract_mrtr_params(¶ms).unwrap_err();
assert!(matches!(
err,
MrtrParseError::InputResponseTooLarge { ref key, .. } if key == "big"
));
}
#[test]
fn extract_rejects_an_oversized_input_responses_total() {
let chunk = "z".repeat(MAX_INPUT_RESPONSE_BYTES - 1_000);
let mut entries = serde_json::Map::new();
for index in 0..8 {
entries.insert(
format!("k{index}"),
json!({ "action": "accept", "content": { "v": chunk } }),
);
}
let err = extract_mrtr_params(&json!({ "inputResponses": entries })).unwrap_err();
assert!(matches!(
err,
MrtrParseError::InputResponsesTotalTooLarge { .. }
));
}
#[test]
fn extract_rejects_an_over_deep_input_response() {
let mut nested = json!("leaf");
for _ in 0..(MAX_INPUT_RESPONSE_DEPTH + 4) {
nested = json!({ "n": nested });
}
let params = json!({
"inputResponses": { "deep": { "action": "accept", "content": { "v": nested } } }
});
let err = extract_mrtr_params(¶ms).unwrap_err();
assert!(matches!(
err,
MrtrParseError::InputResponseTooDeep { ref key, .. } if key == "deep"
));
}
#[test]
fn extract_rejects_an_undecodable_input_response() {
let params = json!({ "inputResponses": { "bad": { "totally": "wrong" } } });
let err = extract_mrtr_params(¶ms).unwrap_err();
assert!(matches!(
err,
MrtrParseError::InputResponseUndecodable { ref key } if key == "bad"
));
}
#[test]
fn parse_error_display_never_echoes_the_offending_key() {
let err = MrtrParseError::InputResponseTooLarge {
key: "secret-key-name".to_string(),
bytes: 1,
max: 2,
};
let rendered = err.to_string();
assert!(!rendered.contains("secret-key-name"));
assert!(rendered.contains('2'));
}
#[test]
fn splice_then_extract_round_trips() {
let mut params = json!({ "name": "search" });
let original = MrtrRequestParams {
input_responses: Some(responses_fixture()),
input_responses_raw: None,
request_state: Some("token".to_string()),
};
splice_mrtr_params(&mut params, &original);
let extracted = extract_mrtr_params(¶ms).unwrap();
assert_eq!(extracted.request_state.as_deref(), Some("token"));
assert_eq!(
extracted.input_responses.as_ref().map(BTreeMap::len),
Some(1)
);
}
fn digest(method: &str, params: &Value) -> [u8; 32] {
salient_param_digest(method, params).expect("the fixture is inside the depth cap")
}
#[test]
fn salient_digest_is_stable_across_key_insertion_order() {
let mut first = serde_json::Map::new();
first.insert("name".to_string(), json!("search"));
first.insert("arguments".to_string(), json!({ "a": 1, "b": 2 }));
let mut second = serde_json::Map::new();
second.insert("arguments".to_string(), json!({ "b": 2, "a": 1 }));
second.insert("name".to_string(), json!("search"));
assert_eq!(
digest("tools/call", &Value::Object(first)),
digest("tools/call", &Value::Object(second))
);
}
#[test]
fn salient_digest_differs_on_name_uri_and_arguments() {
let base = json!({ "name": "search", "arguments": { "q": "a" } });
let other_name = json!({ "name": "delete", "arguments": { "q": "a" } });
let other_args = json!({ "name": "search", "arguments": { "q": "b" } });
assert_ne!(
digest("tools/call", &base),
digest("tools/call", &other_name)
);
assert_ne!(
digest("tools/call", &base),
digest("tools/call", &other_args)
);
assert_ne!(
digest("resources/read", &json!({ "uri": "mem://a" })),
digest("resources/read", &json!({ "uri": "mem://b" }))
);
}
#[test]
fn salient_digest_ignores_meta_input_responses_and_request_state() {
let bare = json!({ "name": "search", "arguments": {} });
let noisy = json!({
"name": "search",
"arguments": {},
"_meta": { "io.modelcontextprotocol/protocolVersion": "2026-07-28" },
"inputResponses": { "k": { "action": "accept" } },
"requestState": "token"
});
assert_eq!(digest("tools/call", &bare), digest("tools/call", &noisy));
}
#[test]
fn salient_digest_of_an_ineligible_method_is_the_empty_object_digest() {
assert_eq!(
digest("tools/list", &json!({ "name": "x", "cursor": "y" })),
digest("tools/list", &json!({}))
);
}
#[test]
fn salient_digest_binds_the_method_name() {
let params = json!({ "name": "search", "arguments": {} });
assert_ne!(
digest("tools/call", ¶ms),
digest("prompts/get", ¶ms)
);
}
fn nest(levels: usize, leaf: Value) -> Value {
let mut value = leaf;
for _ in 0..levels {
value = json!({ "n": value });
}
value
}
fn deep_call_params(levels: usize, leaf: Value) -> Value {
json!({ "name": "search", "arguments": nest(levels, leaf) })
}
#[test]
fn params_differing_only_below_the_depth_cap_can_never_share_a_digest() {
let a = deep_call_params(MAX_CANONICAL_DEPTH, json!("SECRET-A"));
let b = deep_call_params(MAX_CANONICAL_DEPTH, json!("SECRET-B"));
assert_ne!(a, b, "the fixtures must genuinely differ");
let digest_a = salient_param_digest("tools/call", &a);
let digest_b = salient_param_digest("tools/call", &b);
let show = |outcome: &Result<[u8; 32], CanonicalDepthExceeded>| match outcome {
Ok(bytes) => bytes.map(|byte| format!("{byte:02x}")).join(""),
Err(error) => format!("REFUSED ({error})"),
};
assert!(
digest_a.is_err() && digest_b.is_err(),
"over-deep params must be REFUSED, not digested.\n \
A = {}\n B = {}\n equal? = {}",
show(&digest_a),
show(&digest_b),
digest_a == digest_b
);
let shallow_a = deep_call_params(MAX_CANONICAL_DEPTH - 2, json!("SECRET-A"));
let shallow_b = deep_call_params(MAX_CANONICAL_DEPTH - 2, json!("SECRET-B"));
assert_ne!(
digest("tools/call", &shallow_a),
digest("tools/call", &shallow_b),
"the old collision pair, inside the cap, must digest DIFFERENTLY"
);
}
#[test]
fn canonical_depth_boundary_admits_the_cap_and_refuses_one_past_it() {
let mut at_cap = String::new();
assert_eq!(
write_canonical(&nest(MAX_CANONICAL_DEPTH, json!("leaf")), 0, &mut at_cap),
Ok(()),
"a value whose leaf sits exactly AT the cap must canonicalize"
);
assert!(
at_cap.contains("leaf"),
"and must render the leaf: {at_cap}"
);
let mut past_cap = String::new();
assert_eq!(
write_canonical(
&nest(MAX_CANONICAL_DEPTH + 1, json!("leaf")),
0,
&mut past_cap
),
Err(CanonicalDepthExceeded {
depth: MAX_CANONICAL_DEPTH + 1,
max: MAX_CANONICAL_DEPTH,
}),
"one level past the cap must REFUSE"
);
}
#[test]
fn the_digest_boundary_accounts_for_the_salient_wrapper_level() {
assert!(
salient_param_digest(
"tools/call",
&deep_call_params(MAX_CANONICAL_DEPTH - 1, json!("leaf"))
)
.is_ok(),
"arguments nested MAX_CANONICAL_DEPTH - 1 deep must still bind"
);
assert!(
salient_param_digest(
"tools/call",
&deep_call_params(MAX_CANONICAL_DEPTH, json!("leaf"))
)
.is_err(),
"one deeper must be refused"
);
}
#[test]
fn arrays_count_toward_the_canonical_depth_cap() {
let mut value = json!("leaf");
for _ in 0..=MAX_CANONICAL_DEPTH {
value = json!([value]);
}
let mut out = String::new();
assert!(write_canonical(&value, 0, &mut out).is_err());
}
#[test]
fn canonical_depth_error_display_names_only_the_bound() {
let rendered = CanonicalDepthExceeded {
depth: 99,
max: MAX_CANONICAL_DEPTH,
}
.to_string();
assert!(rendered.contains(&MAX_CANONICAL_DEPTH.to_string()));
assert!(!rendered.contains("99"));
}
const _: () = assert!(MAX_INPUT_RESPONSE_DEPTH < MAX_CANONICAL_DEPTH);
#[test]
fn input_responses_are_depth_bounded_at_ingress_but_arguments_are_not() {
let over_deep = nest(MAX_INPUT_RESPONSE_DEPTH + 4, json!("leaf"));
assert!(
matches!(
check_input_response_bounds("k", &over_deep),
Err(MrtrParseError::InputResponseTooDeep { .. })
),
"an over-deep inputResponses entry is rejected at ingress"
);
let params = deep_call_params(MAX_INPUT_RESPONSE_DEPTH + 4, json!("leaf"));
assert!(
extract_mrtr_params(¶ms).is_ok(),
"deep arguments are not bounded at ingress"
);
assert!(salient_param_digest("tools/call", ¶ms).is_ok());
}
proptest::proptest! {
#[test]
fn distinct_params_within_the_cap_digest_distinctly(
left in arb_shallow_json(),
right in arb_shallow_json(),
) {
let left_params = json!({ "name": "search", "arguments": left });
let right_params = json!({ "name": "search", "arguments": right });
let left_digest = salient_param_digest("tools/call", &left_params)
.expect("a bounded-depth value canonicalizes");
let right_digest = salient_param_digest("tools/call", &right_params)
.expect("a bounded-depth value canonicalizes");
let reordered = json!({ "arguments": left_params["arguments"], "name": "search" });
proptest::prop_assert_eq!(
salient_param_digest("tools/call", &reordered)
.expect("a bounded-depth value canonicalizes"),
left_digest
);
if left_params["arguments"] == right_params["arguments"] {
proptest::prop_assert_eq!(left_digest, right_digest);
} else {
proptest::prop_assert_ne!(left_digest, right_digest);
}
}
}
fn arb_shallow_json() -> impl proptest::strategy::Strategy<Value = Value> {
use proptest::prelude::*;
let leaf = prop_oneof![
Just(Value::Null),
any::<bool>().prop_map(Value::from),
any::<i32>().prop_map(Value::from),
"[a-z ]{0,8}".prop_map(Value::from),
];
leaf.prop_recursive(4, 24, 3, |inner| {
prop_oneof![
proptest::collection::vec(inner.clone(), 0..3).prop_map(Value::from),
proptest::collection::btree_map("[a-z]{1,4}", inner, 0..3).prop_map(|m| json!(m)),
]
})
}
#[test]
fn input_required_result_deserializes_the_wire_shape() {
let parsed: InputRequiredResult = serde_json::from_value(json!({
"resultType": "input_required",
"inputRequests": {},
"requestState": "abc"
}))
.unwrap();
assert!(parsed.is_input_required());
assert_eq!(parsed.result_type, "input_required");
assert!(parsed.input_requests.is_some());
assert_eq!(parsed.request_state.as_deref(), Some("abc"));
assert_eq!(parsed.raw["resultType"], "input_required");
}
#[test]
fn input_required_result_keeps_the_verbatim_result_in_raw() {
let parsed: InputRequiredResult = serde_json::from_value(json!({
"resultType": "input_required",
"requestState": "abc",
"vendorField": { "keep": true }
}))
.unwrap();
assert_eq!(parsed.raw["vendorField"]["keep"], true);
assert!(parsed.input_requests.is_none());
}
#[test]
fn input_required_result_recognises_a_completed_result() {
let parsed: InputRequiredResult =
serde_json::from_value(json!({ "resultType": "complete", "content": [] })).unwrap();
assert!(!parsed.is_input_required());
}
#[test]
fn input_required_result_serializes_the_wire_keys() {
let parsed: InputRequiredResult = serde_json::from_value(json!({
"resultType": "input_required",
"requestState": "abc"
}))
.unwrap();
let value = serde_json::to_value(&parsed).unwrap();
assert_eq!(value["resultType"], "input_required");
assert_eq!(value["requestState"], "abc");
assert!(value.get("raw").is_none());
}
#[test]
fn mrtr_outcome_constructs_and_matches() {
let complete: MrtrOutcome<crate::types::CallToolResult> =
MrtrOutcome::Complete(crate::types::CallToolResult::new(vec![]));
assert!(matches!(complete, MrtrOutcome::Complete(_)));
assert!(complete.complete().is_some());
let pending: MrtrOutcome<crate::types::CallToolResult> = MrtrOutcome::InputRequired(
serde_json::from_value(json!({
"resultType": "input_required",
"requestState": "abc"
}))
.unwrap(),
);
assert!(matches!(pending, MrtrOutcome::InputRequired(_)));
assert!(pending.clone().complete().is_none());
assert!(pending.input_required().is_some());
}
#[test]
fn mrtr_signal_uses_camel_case_wire_keys() {
let mut requests = InputRequests::new();
requests.insert(
"user_name".to_string(),
InputRequest::Elicitation(Box::new(form_elicitation())),
);
let signal = MrtrSignal {
input_requests: requests,
continuation: json!({ "step": 1 }),
};
let value = serde_json::to_value(&signal).unwrap();
assert_eq!(
value["inputRequests"]["user_name"]["method"],
"elicitation/create"
);
assert_eq!(value["continuation"]["step"], 1);
assert_eq!(MRTR_SIGNAL_META_KEY, "dev.pmcp/mrtr");
}
use proptest::prelude::*;
fn arb_json() -> impl Strategy<Value = Value> {
let leaf = prop_oneof![
Just(Value::Null),
any::<bool>().prop_map(Value::Bool),
any::<i32>().prop_map(|n| json!(n)),
".{0,16}".prop_map(Value::String),
];
leaf.prop_recursive(4, 32, 4, |inner| {
prop_oneof![
prop::collection::vec(inner.clone(), 0..4).prop_map(Value::Array),
prop::collection::hash_map("[a-zA-Z]{1,6}", inner, 0..4)
.prop_map(|map| { Value::Object(map.into_iter().collect()) }),
]
})
}
proptest! {
#[test]
fn extract_never_panics_over_arbitrary_json(value in arb_json()) {
let _ = extract_mrtr_params(&value);
}
#[test]
fn splice_then_extract_is_identity_for_bounded_request_state(
state in "[ -~]{0,256}"
) {
let mut params = json!({ "name": "search" });
let original = MrtrRequestParams {
input_responses: None,
input_responses_raw: None,
request_state: Some(state.clone()),
};
splice_mrtr_params(&mut params, &original);
let extracted = extract_mrtr_params(¶ms).unwrap();
prop_assert_eq!(extracted.request_state, Some(state));
}
#[test]
fn default_splice_leaves_no_mrtr_key(value in arb_json()) {
let mut params = if value.is_object() {
value
} else {
json!({ "wrapped": value })
};
splice_mrtr_params(&mut params, &MrtrRequestParams::default());
prop_assert!(params.get("inputResponses").is_none());
prop_assert!(params.get("requestState").is_none());
}
#[test]
fn header_value_codec_round_trips(value in ".{0,64}") {
let encoded = encode_header_value(&value);
prop_assert_eq!(decode_header_value(&encoded), Some(value));
}
#[test]
fn decode_header_value_never_panics(raw in ".{0,128}") {
let _ = decode_header_value(&raw);
}
#[test]
fn salient_digest_never_panics(value in arb_json(), method in "[a-z/]{0,20}") {
let _ = salient_param_digest(&method, &value);
}
}
}