use std::collections::BTreeMap;
use serde_json::Value;
use super::super::FindingSink;
use super::http_status_for;
use crate::context::TraceContext;
use mcp_conformance_core::trace::Direction;
mod trace_context;
#[cfg(test)]
mod tests;
pub(in crate::checks) use trace_context::trace_context_format;
const REQUIRED_REQUEST_FIELDS: &[&str] = &[
"io.modelcontextprotocol/protocolVersion",
"io.modelcontextprotocol/clientCapabilities",
];
const MISSING_CAPABILITY_CODE: i64 = -32021;
const INVALID_PARAMS: i64 = -32602;
const LEGACY_HANDSHAKE: &str = "initialize";
fn params_meta(payload: &Value) -> Option<&serde_json::Map<String, Value>> {
payload.get("params")?.get("_meta")?.as_object()
}
fn client_requests<'a>(
context: &'a TraceContext<'_>,
) -> impl Iterator<Item = (u64, Option<&'a Value>, &'a Value)> + 'a {
context.messages().filter_map(|(event, _, _)| {
if !matches!(event.direction, Direction::ClientToServer) {
return None;
}
let payload = event.message_payload()?;
payload.get("method")?;
Some((event.seq, payload.get("id"), payload))
})
}
pub(in crate::checks) fn required_request_fields(
context: &TraceContext<'_>,
sink: &mut FindingSink,
) {
for (seq, id, payload) in client_requests(context) {
if id.is_none() {
continue; }
sink.examined();
let meta = params_meta(payload);
for field in REQUIRED_REQUEST_FIELDS {
let present = meta.is_some_and(|meta| meta.contains_key(*field));
if !present {
sink.push(
Some(seq),
format!("request `_meta` is missing required field `{field}`"),
);
}
}
}
}
fn malformed_requests(context: &TraceContext<'_>) -> BTreeMap<String, u64> {
client_requests(context)
.filter_map(|(seq, id, payload)| {
let id = id?;
if payload.get("method").and_then(Value::as_str) == Some(LEGACY_HANDSHAKE) {
return None;
}
let meta = params_meta(payload);
let complete = REQUIRED_REQUEST_FIELDS
.iter()
.all(|field| meta.is_some_and(|meta| meta.contains_key(*field)));
(!complete).then(|| (id.to_string(), seq))
})
.collect()
}
pub(in crate::checks) fn missing_required_field_rejected(
context: &TraceContext<'_>,
sink: &mut FindingSink,
) {
let malformed = malformed_requests(context);
if malformed.is_empty() {
return;
}
for (event, _, _) in context.messages() {
if !matches!(event.direction, Direction::ServerToClient) {
continue;
}
let Some(payload) = event.message_payload() else {
continue;
};
let Some(id) = payload.get("id") else {
continue;
};
let Some(&request_seq) = malformed.get(&id.to_string()) else {
continue;
};
sink.examined();
match payload.get("error").and_then(|error| error.get("code")) {
Some(code) if code.as_i64() == Some(INVALID_PARAMS) => {}
Some(code) => sink.push(
Some(event.seq),
format!(
"request at seq {request_seq} was missing a required `_meta` field; \
the server answered with error code {code} rather than {INVALID_PARAMS}"
),
),
None => sink.push(
Some(event.seq),
format!(
"request at seq {request_seq} was missing a required `_meta` field; \
the server answered with a result rather than error {INVALID_PARAMS}"
),
),
}
}
}
fn http_status_for_error(
context: &TraceContext<'_>,
sink: &mut FindingSink,
code: i64,
clause: &str,
answering: Option<&BTreeMap<String, u64>>,
) {
for (event, _, _) in context.messages() {
if !matches!(event.direction, Direction::ServerToClient) {
continue;
}
let Some(payload) = event.message_payload() else {
continue;
};
let matches_code = payload
.get("error")
.and_then(|error| error.get("code"))
.and_then(Value::as_i64)
== Some(code);
if !matches_code {
continue;
}
if let Some(answering) = answering {
let answers_a_subject = payload
.get("id")
.is_some_and(|id| answering.contains_key(&id.to_string()));
if !answers_a_subject {
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!("{clause}: error {code} was returned with HTTP {status}, not 400"),
);
}
}
}
pub(in crate::checks) fn missing_required_field_http_status(
context: &TraceContext<'_>,
sink: &mut FindingSink,
) {
let malformed = malformed_requests(context);
http_status_for_error(
context,
sink,
INVALID_PARAMS,
"missing required `_meta` field",
Some(&malformed),
);
}
pub(in crate::checks) fn missing_capability_http_status(
context: &TraceContext<'_>,
sink: &mut FindingSink,
) {
http_status_for_error(
context,
sink,
MISSING_CAPABILITY_CODE,
"missing required client capability",
None,
);
}
pub(in crate::checks) fn missing_capability_error(
context: &TraceContext<'_>,
sink: &mut FindingSink,
) {
for (event, _, _) in context.messages() {
let Some(error) = event
.message_payload()
.and_then(|payload| payload.get("error"))
else {
continue;
};
if error.get("code").and_then(Value::as_i64) != Some(MISSING_CAPABILITY_CODE) {
continue;
}
sink.examined();
match error
.get("data")
.and_then(|data| data.get("requiredCapabilities"))
{
Some(required) if required.is_object() => {
if required.as_object().is_some_and(serde_json::Map::is_empty) {
sink.push(
Some(event.seq),
format!(
"error {MISSING_CAPABILITY_CODE} carries an empty \
`data.requiredCapabilities`; it must name the missing capabilities"
),
);
}
}
Some(_) => sink.push(
Some(event.seq),
format!(
"error {MISSING_CAPABILITY_CODE} has `data.requiredCapabilities` \
that is not a `ClientCapabilities` object"
),
),
None => sink.push(
Some(event.seq),
format!(
"error {MISSING_CAPABILITY_CODE} has no `data.requiredCapabilities` \
naming the missing capabilities"
),
),
}
}
}
pub(in crate::checks) fn no_undeclared_capability_reliance(
context: &TraceContext<'_>,
sink: &mut FindingSink,
) {
let mut declared: BTreeMap<String, Vec<String>> = BTreeMap::new();
for (_, id, payload) in client_requests(context) {
let Some(id) = id else { continue };
let names = params_meta(payload)
.and_then(|meta| meta.get("io.modelcontextprotocol/clientCapabilities"))
.and_then(Value::as_object)
.map(|caps| caps.keys().cloned().collect())
.unwrap_or_default();
declared.insert(id.to_string(), names);
}
for (event, _, _) in context.messages() {
if !matches!(event.direction, Direction::ServerToClient) {
continue;
}
let Some(payload) = event.message_payload() else {
continue;
};
let Some(result) = payload.get("result") else {
continue;
};
if result.get("resultType").and_then(Value::as_str) != Some("input_required") {
continue;
}
let Some(id) = payload.get("id") else {
continue;
};
let Some(declared) = declared.get(&id.to_string()) else {
continue;
};
let requests = result
.get("inputRequests")
.and_then(Value::as_object)
.map(|map| map.values().collect::<Vec<_>>())
.unwrap_or_default();
for request in requests {
let Some(method) = request.get("method").and_then(Value::as_str) else {
continue;
};
let needed = match method {
"elicitation/create" => "elicitation",
"sampling/createMessage" => "sampling",
"roots/list" => "roots",
_ => continue,
};
sink.examined();
if !declared.iter().any(|name| name == needed) {
sink.push(
Some(event.seq),
format!(
"server asked for `{method}`, which needs the `{needed}` capability, \
but the request's `clientCapabilities` did not declare it"
),
);
}
}
}
}
pub(in crate::checks) fn subscription_id_present(
context: &TraceContext<'_>,
sink: &mut FindingSink,
) {
let listening = context.messages().any(|(event, _, _)| {
event
.message_payload()
.and_then(|payload| payload.get("method"))
.and_then(Value::as_str)
== Some("subscriptions/listen")
});
if !listening {
return;
}
for (event, _, _) in context.messages() {
if !matches!(event.direction, Direction::ServerToClient) {
continue;
}
let Some(payload) = event.message_payload() else {
continue;
};
if payload.get("id").is_some() || payload.get("method").is_none() {
continue;
}
let method = payload.get("method").and_then(Value::as_str).unwrap_or("");
if method.starts_with("notifications/progress")
|| method.starts_with("notifications/message")
{
continue;
}
sink.examined();
let tagged = params_meta(payload)
.is_some_and(|meta| meta.contains_key("io.modelcontextprotocol/subscriptionId"));
if !tagged {
sink.push(
Some(event.seq),
format!(
"notification `{method}` on a subscriptions/listen stream has no \
`io.modelcontextprotocol/subscriptionId` in `_meta`"
),
);
}
}
}