use std::collections::BTreeSet;
use serde_json::Value;
use super::super::super::FindingSink;
use super::super::http_status_for;
use super::{
META_PROTOCOL_VERSION, Match, Post, compare, designations_by_tool, header_safe, mirrors, posts,
posts_by_message, sentinel_payload,
};
use crate::context::TraceContext;
#[cfg(test)]
mod tests;
const HEADER_MISMATCH: i64 = -32020;
const UNSUPPORTED_VERSION: i64 = -32022;
const METHOD_NOT_FOUND: i64 = -32601;
fn answer_code(response: &mcp_conformance_core::trace::TraceEvent) -> Option<i64> {
response
.message_payload()?
.get("error")?
.get("code")?
.as_i64()
}
fn answer_label(code: Option<i64>) -> String {
code.map_or_else(|| "a result".to_owned(), |code| format!("error {code}"))
}
pub(in crate::checks) fn version_mismatch_rejected(
context: &TraceContext<'_>,
sink: &mut FindingSink,
) {
rejected_for(context, sink, version_mismatch_fault);
}
pub(in crate::checks) fn invalid_param_header_rejected(
context: &TraceContext<'_>,
sink: &mut FindingSink,
) {
let recognized = recognized_param_headers(context);
rejected_for(context, sink, |post| invalid_param_fault(post, &recognized));
}
pub(in crate::checks) fn header_mismatch_status(
context: &TraceContext<'_>,
sink: &mut FindingSink,
) {
for (event, _, _) in context.messages() {
if answer_code(event) != Some(HEADER_MISMATCH) {
continue;
}
let Some((status_seq, status)) = http_status_for(context, event.seq) else {
continue;
};
sink.examined();
if status != 400 {
sink.push(
Some(status_seq),
format!(
"HeaderMismatch ({HEADER_MISMATCH}) was returned with HTTP {status}, not 400"
),
);
}
}
}
fn rejected_for(
context: &TraceContext<'_>,
sink: &mut FindingSink,
fault: impl Fn(&Post<'_>) -> Option<String>,
) {
let by_message = posts_by_message(context);
for exchange in context.exchanges() {
let Some(post) = by_message.get(&exchange.request.seq) else {
continue;
};
let Some(reason) = fault(post) else {
continue;
};
sink.examined();
let code = answer_code(exchange.response);
if code != Some(HEADER_MISMATCH) {
sink.push(
Some(exchange.response.seq),
format!(
"the POST at seq {} {reason}; the server answered with {} instead of \
rejecting it with {HEADER_MISMATCH} (HeaderMismatch)",
post.seq,
answer_label(code)
),
);
}
}
}
fn recognized_param_headers(context: &TraceContext<'_>) -> BTreeSet<String> {
designations_by_tool(context)
.values()
.flatten()
.map(|designation| designation.header.clone())
.collect()
}
fn version_mismatch_fault(post: &Post<'_>) -> Option<String> {
let sent = post.headers.get("mcp-protocol-version")?;
let body = post.body_protocol_version()?;
(sent != body).then(|| {
format!("carried `MCP-Protocol-Version: {sent}` against a body `_meta` version of {body:?}")
})
}
fn invalid_param_fault(post: &Post<'_>, recognized: &BTreeSet<String>) -> Option<String> {
post.headers
.iter()
.find(|(name, value)| {
recognized.contains(*name) && sentinel_payload(value).is_none() && !header_safe(value)
})
.map(|(name, value)| {
format!("carried `{name}: {value:?}`, whose characters are not valid unencoded")
})
}
pub(in crate::checks) fn header_body_match_validated(
context: &TraceContext<'_>,
sink: &mut FindingSink,
) {
let designated = designations_by_tool(context);
let by_message = posts_by_message(context);
for exchange in context.exchanges() {
let Some(post) = by_message.get(&exchange.request.seq) else {
continue;
};
if exchange.result.is_none() {
continue; }
for mirror in mirrors(post, &designated) {
let Some(sent) = post.headers.get(&mirror.header) else {
continue;
};
sink.examined();
if compare(sent, &mirror.value) == Match::Mismatch {
sink.push(
Some(exchange.response.seq),
format!(
"the POST at seq {} carried `{}: {sent}` against `{}` = {:?}; the \
server answered it with a result instead of rejecting the mismatch",
post.seq, mirror.label, mirror.source, mirror.value
),
);
}
}
}
}
pub(in crate::checks) fn unsupported_version_error(
context: &TraceContext<'_>,
sink: &mut FindingSink,
) {
unsupported_version_shape(context, sink);
unsupported_version_answer(context, sink);
}
fn unsupported_version_shape(context: &TraceContext<'_>, sink: &mut FindingSink) {
for (event, _, _) in context.messages() {
if answer_code(event) != Some(UNSUPPORTED_VERSION) {
continue;
}
sink.examined();
let lists_versions = event
.message_payload()
.and_then(|payload| payload.get("error"))
.and_then(|error| error.get("data"))
.and_then(|data| data.get("supported"))
.and_then(Value::as_array)
.is_some_and(|supported| {
!supported.is_empty() && supported.iter().all(Value::is_string)
});
if !lists_versions {
sink.push(
Some(event.seq),
format!(
"error {UNSUPPORTED_VERSION} does not carry `data.supported` listing the \
protocol versions the server does implement"
),
);
}
}
}
pub(in crate::checks) fn unsupported_version_status(
context: &TraceContext<'_>,
sink: &mut FindingSink,
) {
let handshakes = legacy_handshake_ids(context);
for (event, _, _) in context.messages() {
if answer_code(event) != Some(UNSUPPORTED_VERSION) {
continue;
}
if event
.message_payload()
.and_then(|payload| payload.get("id"))
.is_some_and(|id| handshakes.contains(&id.to_string()))
{
continue;
}
let Some((status_seq, status)) = http_status_for(context, event.seq) else {
continue;
};
sink.examined();
if status != 400 {
sink.push(
Some(status_seq),
format!(
"UnsupportedProtocolVersionError ({UNSUPPORTED_VERSION}) was returned \
with HTTP {status}, not 400"
),
);
}
}
}
fn unsupported_version_answer(context: &TraceContext<'_>, sink: &mut FindingSink) {
let Some(supported) = declared_versions(context) else {
return;
};
for exchange in context.exchanges() {
let Some(requested) = exchange
.params
.and_then(|params| params.get("_meta")?.get(META_PROTOCOL_VERSION)?.as_str())
else {
continue;
};
if supported.contains(requested) {
continue;
}
sink.examined();
let code = answer_code(exchange.response);
if code != Some(UNSUPPORTED_VERSION) {
sink.push(
Some(exchange.response.seq),
format!(
"the request at seq {} declared protocol version {requested:?}, which the \
server's own `supportedVersions` omits; it answered with {} instead of \
{UNSUPPORTED_VERSION}",
exchange.request.seq,
answer_label(code)
),
);
}
}
}
fn legacy_handshake_ids(context: &TraceContext<'_>) -> BTreeSet<String> {
context
.messages()
.filter_map(|(event, _, _)| {
let payload = event.message_payload()?;
if payload.get("method")?.as_str()? != "initialize" {
return None;
}
Some(payload.get("id")?.to_string())
})
.collect()
}
fn declared_versions(context: &TraceContext<'_>) -> Option<BTreeSet<String>> {
context
.exchanges_for("server/discover")
.find_map(|exchange| {
let versions: BTreeSet<String> = exchange
.result?
.get("supportedVersions")?
.as_array()?
.iter()
.filter_map(|version| version.as_str().map(str::to_owned))
.collect();
(!versions.is_empty()).then_some(versions)
})
}
pub(in crate::checks) fn unknown_method_404(context: &TraceContext<'_>, sink: &mut FindingSink) {
if posts(context).is_empty() {
return;
}
for (event, _, _) in context.messages() {
if answer_code(event) != Some(METHOD_NOT_FOUND) {
continue;
}
let Some((status_seq, status)) = http_status_for(context, event.seq) else {
continue;
};
sink.examined();
if status != 404 {
sink.push(
Some(status_seq),
format!(
"`Method not found` ({METHOD_NOT_FOUND}) was returned with HTTP {status}, \
not 404"
),
);
}
}
}