use mcp_conformance_core::message::MessageKind;
use mcp_conformance_core::revision::ProtocolRevision;
use mcp_conformance_core::trace::Direction;
use serde_json::Value;
use super::FindingSink;
use crate::context::{Phase, TraceContext};
pub(super) fn first_interaction_initialize(context: &TraceContext<'_>, sink: &mut FindingSink) {
let Some((event, kind, _)) = context.messages().next() else {
return;
};
sink.examined();
match (event.direction, kind) {
(Direction::ClientToServer, MessageKind::Request { method, .. })
if *method == "initialize" => {}
(Direction::ClientToServer, MessageKind::Request { method, .. }) => sink.push(
Some(event.seq),
format!("first message is a {method:?} request, expected \"initialize\""),
),
(direction, _) => sink.push(
Some(event.seq),
format!(
"first message is {} ({}), expected the client's \"initialize\" request",
describe_kind(kind),
direction_name(direction)
),
),
}
}
pub(super) fn initialize_params(context: &TraceContext<'_>, sink: &mut FindingSink) {
let Some((seq, params)) = context.initialize().request else {
return; };
sink.examined();
let Some(params) = params else {
sink.push(
Some(seq),
"initialize request has no params; protocolVersion, capabilities, and clientInfo are required".to_owned(),
);
return;
};
expect_member(
sink,
seq,
params,
"protocolVersion",
Value::is_string,
"a string",
);
expect_member(
sink,
seq,
params,
"capabilities",
Value::is_object,
"an object",
);
expect_member(
sink,
seq,
params,
"clientInfo",
Value::is_object,
"an object",
);
}
fn expect_member(
sink: &mut FindingSink,
seq: u64,
params: &Value,
member: &str,
predicate: fn(&Value) -> bool,
expected: &str,
) {
match params.get(member) {
None => sink.push(
Some(seq),
format!("initialize params lack the {member} member"),
),
Some(value) if !predicate(value) => sink.push(
Some(seq),
format!("initialize params member {member} should be {expected}"),
),
Some(_) => {}
}
}
pub(super) fn initialized_notification(context: &TraceContext<'_>, sink: &mut FindingSink) {
let init = context.initialize();
let Some((result_seq, _)) = init.result else {
return; };
sink.examined();
if init.initialized.is_none() {
sink.push(
Some(result_seq),
"the server answered initialize here, but no notifications/initialized notification follows in the trace".to_owned(),
);
}
}
pub(super) fn client_requests_before_init_response(
context: &TraceContext<'_>,
sink: &mut FindingSink,
) {
for (event, kind, phase) in context.messages() {
if event.direction != Direction::ClientToServer {
continue;
}
if !matches!(
phase,
Phase::BeforeInitialize | Phase::AwaitingInitializeResult
) {
continue;
}
sink.examined();
let MessageKind::Request { method, .. } = kind else {
continue;
};
if *method != "initialize" && *method != "ping" {
sink.push(
Some(event.seq),
format!(
"client sent a {method:?} request before the server responded to initialize"
),
);
}
}
}
pub(super) fn server_requests_before_initialized(
context: &TraceContext<'_>,
sink: &mut FindingSink,
) {
for (event, kind, phase) in context.messages() {
if event.direction != Direction::ServerToClient || phase == Phase::Ready {
continue;
}
sink.examined();
let MessageKind::Request { method, .. } = kind else {
continue;
};
if *method != "ping" {
sink.push(
Some(event.seq),
format!(
"server sent a {method:?} request before receiving the initialized notification"
),
);
}
}
}
pub(super) fn initialize_protocol_version(context: &TraceContext<'_>, sink: &mut FindingSink) {
let Some((seq, params)) = context.initialize().request else {
return; };
sink.examined();
match params.and_then(|params| params.get("protocolVersion")) {
None => sink.push(
Some(seq),
"initialize request sends no protocolVersion".to_owned(),
),
Some(Value::String(_)) => {}
Some(other) => sink.push(
Some(seq),
format!("initialize request protocolVersion is {other}, expected a version string"),
),
}
}
pub(super) fn initialize_result_version(context: &TraceContext<'_>, sink: &mut FindingSink) {
let Some((seq, result)) = context.initialize().result else {
return;
};
sink.examined();
match result.get("protocolVersion") {
None => sink.push(
Some(seq),
"initialize result lacks the protocolVersion member".to_owned(),
),
Some(Value::String(version)) => {
if version.parse::<ProtocolRevision>().is_err() {
sink.push(
Some(seq),
format!(
"initialize result protocolVersion {version:?} is not a dated revision identifier (YYYY-MM-DD)"
),
);
}
}
Some(other) => sink.push(
Some(seq),
format!("initialize result protocolVersion is {other}, expected a revision string"),
),
}
}
pub(super) fn initialize_result_shape(context: &TraceContext<'_>, sink: &mut FindingSink) {
let Some((seq, result)) = context.initialize().result else {
return;
};
sink.examined();
for (member, label) in [
("capabilities", "its capabilities"),
("serverInfo", "its implementation information (serverInfo)"),
] {
match result.get(member) {
None => sink.push(
Some(seq),
format!("initialize result lacks {label}: no {member} member"),
),
Some(value) if !value.is_object() => sink.push(
Some(seq),
format!("initialize result {member} is {value}, expected an object"),
),
Some(_) => {}
}
}
}
const fn describe_kind(kind: &MessageKind<'_>) -> &'static str {
match kind {
MessageKind::Request { .. } => "a request",
MessageKind::Notification { .. } => "a notification",
MessageKind::Result { .. } => "a result response",
MessageKind::Error { .. } => "an error response",
MessageKind::Invalid { .. } => "not a valid JSON-RPC message",
_ => "an unrecognized message kind",
}
}
const fn direction_name(direction: Direction) -> &'static str {
match direction {
Direction::ClientToServer => "client to server",
Direction::ServerToClient => "server to client",
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn describe_kind_names_every_shape_exactly() {
let request = json!({"id": 1, "method": "x"});
let notification = json!({"method": "x"});
let result = json!({"id": 1, "result": {}});
let error = json!({"id": 1, "error": {}});
let invalid = json!([]);
let cases = [
(&request, "a request"),
(¬ification, "a notification"),
(&result, "a result response"),
(&error, "an error response"),
(&invalid, "not a valid JSON-RPC message"),
];
for (payload, expected) in cases {
let kind = mcp_conformance_core::message::classify(payload);
assert_eq!(describe_kind(&kind), expected, "for {payload}");
}
}
#[test]
fn direction_name_is_exact() {
assert_eq!(
direction_name(Direction::ClientToServer),
"client to server"
);
assert_eq!(
direction_name(Direction::ServerToClient),
"server to client"
);
}
#[test]
fn initialize_params_with_wrong_types_are_flagged() {
use crate::context::TraceContext;
use crate::reader::{Limits, parse_trace};
let doc = r#"{"seq":0,"direction":"client-to-server","transport":"stdio","kind":"message","payload":{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":123,"capabilities":[],"clientInfo":"nope"}}}"#;
let events = parse_trace(doc, &Limits::default()).expect("valid trace");
let context = TraceContext::new(&events);
let findings = crate::checks::find("lifecycle.initialize-params")
.expect("check exists")
.run(&context)
.findings;
assert_eq!(findings.len(), 3, "{findings:?}");
assert!(
findings[0]
.detail
.contains("protocolVersion should be a string")
);
assert!(
findings[1]
.detail
.contains("capabilities should be an object")
);
assert!(
findings[2]
.detail
.contains("clientInfo should be an object")
);
}
#[test]
fn initialize_result_shape_demands_capability_and_serverinfo_objects() {
fn handshake_with_result(result: &str) -> String {
let request = r#"{"seq":0,"direction":"client-to-server","transport":"stdio","kind":"message","payload":{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-11-25","capabilities":{},"clientInfo":{"name":"t","version":"0"}}}}"#;
format!(
"{request}\n{{\"seq\":1,\"direction\":\"server-to-client\",\"transport\":\"stdio\",\"kind\":\"message\",\"payload\":{{\"jsonrpc\":\"2.0\",\"id\":1,\"result\":{result}}}}}"
)
}
let run = |result: &str| {
let trace = handshake_with_result(result);
let events = crate::reader::parse_trace(&trace, &crate::reader::Limits::default())
.expect("trace parses");
let context = TraceContext::new(&events);
crate::checks::find("lifecycle.initialize-result-shape")
.expect("check registered")
.run(&context)
.findings
};
assert!(
run(r#"{"protocolVersion":"2025-11-25","capabilities":{},"serverInfo":{"name":"s","version":"0"}}"#)
.is_empty()
);
let missing_caps =
run(r#"{"protocolVersion":"2025-11-25","serverInfo":{"name":"s","version":"0"}}"#);
assert_eq!(missing_caps.len(), 1, "{missing_caps:?}");
assert!(missing_caps[0].detail.contains("capabilities"));
let missing_info = run(r#"{"protocolVersion":"2025-11-25","capabilities":{}}"#);
assert_eq!(missing_info.len(), 1, "{missing_info:?}");
assert!(missing_info[0].detail.contains("serverInfo"));
let wrong = run(r#"{"capabilities":7,"serverInfo":"s"}"#);
assert_eq!(wrong.len(), 2, "{wrong:?}");
let events = crate::reader::parse_trace("", &crate::reader::Limits::default()).unwrap();
let context = TraceContext::new(&events);
assert!(
crate::checks::find("lifecycle.initialize-result-shape")
.unwrap()
.run(&context)
.findings
.is_empty()
);
}
}