1use mcp_conformance_core::message::{MessageKind, classify};
11use mcp_conformance_core::trace::{Direction, TraceEvent};
12use serde_json::Value;
13
14mod pairing;
15
16pub use pairing::Exchange;
17
18#[derive(Debug, Clone, Copy, PartialEq, Eq)]
20#[non_exhaustive]
21pub enum Phase {
22 BeforeInitialize,
24 AwaitingInitializeResult,
26 AfterInitializeSuccess,
29 AfterInitializeError,
31 Ready,
33}
34
35#[derive(Debug, Clone, Copy, Default)]
37#[non_exhaustive]
38pub struct InitializeExchange<'a> {
39 pub request: Option<(u64, Option<&'a Value>)>,
41 pub result: Option<(u64, &'a Value)>,
43 pub initialized: Option<u64>,
45}
46
47#[derive(Debug)]
49pub struct TraceContext<'a> {
50 events: &'a [TraceEvent],
51 kinds: Vec<Option<MessageKind<'a>>>,
52 phases: Vec<Phase>,
53 pairs: Vec<Option<usize>>,
54 init: InitializeExchange<'a>,
55 final_phase: Phase,
56}
57
58impl<'a> TraceContext<'a> {
59 #[must_use]
61 pub fn new(events: &'a [TraceEvent]) -> Self {
62 let kinds: Vec<Option<MessageKind<'a>>> = events
63 .iter()
64 .map(|event| event.message_payload().map(classify))
65 .collect();
66
67 let mut phases = Vec::with_capacity(events.len());
68 let mut tracker = LifecycleTracker::start();
69 for (event, kind) in events.iter().zip(&kinds) {
70 phases.push(tracker.phase);
71 if let Some(kind) = kind {
72 tracker.step(event, kind);
73 }
74 }
75
76 let pairs = pairing::pair_responses(events, &kinds);
77
78 Self {
79 events,
80 kinds,
81 phases,
82 pairs,
83 init: tracker.init,
84 final_phase: tracker.phase,
85 }
86 }
87
88 #[must_use]
90 pub const fn events(&self) -> &'a [TraceEvent] {
91 self.events
92 }
93
94 pub fn messages(&self) -> impl Iterator<Item = (&'a TraceEvent, &MessageKind<'a>, Phase)> + '_ {
97 self.events
98 .iter()
99 .zip(&self.kinds)
100 .zip(&self.phases)
101 .filter_map(|((event, kind), phase)| kind.as_ref().map(|kind| (event, kind, *phase)))
102 }
103
104 #[must_use]
106 pub const fn initialize(&self) -> &InitializeExchange<'a> {
107 &self.init
108 }
109
110 #[must_use]
112 pub fn server_capabilities(&self) -> Option<&'a Value> {
113 self.init
114 .result
115 .and_then(|(_, result)| result.get("capabilities"))
116 }
117
118 #[must_use]
120 pub fn client_capabilities(&self) -> Option<&'a Value> {
121 self.init
122 .request
123 .and_then(|(_, params)| params?.get("capabilities"))
124 }
125
126 #[must_use]
128 pub const fn final_phase(&self) -> Phase {
129 self.final_phase
130 }
131}
132
133struct LifecycleTracker<'a> {
135 phase: Phase,
136 init: InitializeExchange<'a>,
137 initialize_id: Option<&'a Value>,
138}
139
140impl<'a> LifecycleTracker<'a> {
141 const fn start() -> Self {
142 Self {
143 phase: Phase::BeforeInitialize,
144 init: InitializeExchange {
145 request: None,
146 result: None,
147 initialized: None,
148 },
149 initialize_id: None,
150 }
151 }
152
153 fn step(&mut self, event: &'a TraceEvent, kind: &MessageKind<'a>) {
154 match (self.phase, event.direction, kind) {
155 (
156 Phase::BeforeInitialize,
157 Direction::ClientToServer,
158 MessageKind::Request { method, id },
159 ) if *method == "initialize" => {
160 self.initialize_id = Some(id);
161 self.init.request = Some((
162 event.seq,
163 event
164 .message_payload()
165 .and_then(|payload| payload.get("params")),
166 ));
167 self.phase = Phase::AwaitingInitializeResult;
168 }
169 (
170 Phase::AwaitingInitializeResult,
171 Direction::ServerToClient,
172 MessageKind::Result { id: Some(id) },
173 ) if Some(*id) == self.initialize_id => {
174 self.init.result = event
175 .message_payload()
176 .and_then(|payload| payload.get("result"))
177 .map(|result| (event.seq, result));
178 self.phase = Phase::AfterInitializeSuccess;
179 }
180 (
181 Phase::AwaitingInitializeResult,
182 Direction::ServerToClient,
183 MessageKind::Error { id: Some(id), .. },
184 ) if Some(*id) == self.initialize_id => {
185 self.phase = Phase::AfterInitializeError;
186 }
187 (
188 Phase::AfterInitializeSuccess,
189 Direction::ClientToServer,
190 MessageKind::Notification { method },
191 ) if *method == "notifications/initialized" => {
192 self.init.initialized = Some(event.seq);
193 self.phase = Phase::Ready;
194 }
195 _ => {}
196 }
197 }
198}
199
200#[cfg(test)]
201#[allow(clippy::unwrap_used)]
202mod tests {
203 use super::*;
204 use crate::reader::{Limits, parse_trace};
205
206 fn happy_path() -> Vec<TraceEvent> {
207 let doc = r#"{"seq":0,"direction":"client-to-server","transport":"stdio","kind":"lifecycle","event":"transport-open"}
208{"seq":1,"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"}}}}
209{"seq":2,"direction":"server-to-client","transport":"stdio","kind":"message","payload":{"jsonrpc":"2.0","id":1,"result":{"protocolVersion":"2025-11-25","capabilities":{},"serverInfo":{"name":"s","version":"0"}}}}
210{"seq":3,"direction":"client-to-server","transport":"stdio","kind":"message","payload":{"jsonrpc":"2.0","method":"notifications/initialized"}}
211{"seq":4,"direction":"client-to-server","transport":"stdio","kind":"message","payload":{"jsonrpc":"2.0","id":2,"method":"tools/list"}}"#;
212 parse_trace(doc, &Limits::default()).unwrap()
213 }
214
215 #[test]
216 fn tracks_phases_through_initialization() {
217 let events = happy_path();
218 let context = TraceContext::new(&events);
219 let phases: Vec<Phase> = context.messages().map(|(_, _, phase)| phase).collect();
220 assert_eq!(
221 phases,
222 vec![
223 Phase::BeforeInitialize,
224 Phase::AwaitingInitializeResult,
225 Phase::AfterInitializeSuccess,
226 Phase::Ready,
227 ]
228 );
229 assert_eq!(context.final_phase(), Phase::Ready);
230 }
231
232 #[test]
233 fn records_initialize_exchange() {
234 let events = happy_path();
235 let context = TraceContext::new(&events);
236 let init = context.initialize();
237 assert_eq!(init.request.unwrap().0, 1);
238 assert!(init.request.unwrap().1.is_some());
239 assert_eq!(init.result.unwrap().0, 2);
240 assert_eq!(init.initialized, Some(3));
241 }
242
243 #[test]
244 fn initialize_error_blocks_ready() {
245 let doc = r#"{"seq":1,"direction":"client-to-server","transport":"stdio","kind":"message","payload":{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}}
246{"seq":2,"direction":"server-to-client","transport":"stdio","kind":"message","payload":{"jsonrpc":"2.0","id":1,"error":{"code":-32602,"message":"Unsupported protocol version"}}}
247{"seq":3,"direction":"client-to-server","transport":"stdio","kind":"message","payload":{"jsonrpc":"2.0","method":"notifications/initialized"}}"#;
248 let events = parse_trace(doc, &Limits::default()).unwrap();
249 let context = TraceContext::new(&events);
250 assert_eq!(context.initialize().initialized, None);
253 assert_eq!(context.final_phase(), Phase::AfterInitializeError);
254 }
255
256 #[test]
257 fn empty_trace_has_no_exchange() {
258 let context = TraceContext::new(&[]);
259 assert!(context.initialize().request.is_none());
260 assert_eq!(context.final_phase(), Phase::BeforeInitialize);
261 assert_eq!(context.server_capabilities(), None);
262 assert_eq!(context.client_capabilities(), None);
263 }
264
265 #[test]
266 fn capability_accessors_read_their_declaration_surfaces() {
267 use serde_json::json;
268 let events = happy_path();
269 let context = TraceContext::new(&events);
270 assert_eq!(context.client_capabilities(), Some(&json!({})));
272 assert_eq!(context.server_capabilities(), Some(&json!({})));
273
274 let doc = r#"{"seq":0,"direction":"client-to-server","transport":"stdio","kind":"message","payload":{"jsonrpc":"2.0","id":1,"method":"initialize"}}
276{"seq":1,"direction":"server-to-client","transport":"stdio","kind":"message","payload":{"jsonrpc":"2.0","id":1,"error":{"code":-32603,"message":"x"}}}"#;
277 let events = parse_trace(doc, &Limits::default()).unwrap();
278 let context = TraceContext::new(&events);
279 assert_eq!(context.client_capabilities(), None);
280 assert_eq!(context.server_capabilities(), None);
281 }
282
283 #[test]
284 fn responses_with_unrelated_ids_do_not_complete_initialization() {
285 for body in [
288 r#"{"jsonrpc":"2.0","id":99,"result":{}}"#,
289 r#"{"jsonrpc":"2.0","id":99,"error":{"code":-32600,"message":"x"}}"#,
290 ] {
291 let response = format!(
292 r#"{{"seq":1,"direction":"server-to-client","transport":"stdio","kind":"message","payload":{body}}}"#
293 );
294 let doc = format!(
295 "{}\n{response}",
296 r#"{"seq":0,"direction":"client-to-server","transport":"stdio","kind":"message","payload":{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}}"#,
297 );
298 let events = parse_trace(&doc, &Limits::default()).unwrap();
299 let context = TraceContext::new(&events);
300 assert!(context.initialize().result.is_none(), "{body}");
301 assert_eq!(
302 context.final_phase(),
303 Phase::AwaitingInitializeResult,
304 "{body}"
305 );
306 }
307 }
308
309 #[test]
310 fn only_the_initialized_notification_makes_the_session_ready() {
311 let doc = r#"{"seq":0,"direction":"client-to-server","transport":"stdio","kind":"message","payload":{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}}
312{"seq":1,"direction":"server-to-client","transport":"stdio","kind":"message","payload":{"jsonrpc":"2.0","id":1,"result":{}}}
313{"seq":2,"direction":"client-to-server","transport":"stdio","kind":"message","payload":{"jsonrpc":"2.0","method":"notifications/cancelled"}}"#;
314 let events = parse_trace(doc, &Limits::default()).unwrap();
315 let context = TraceContext::new(&events);
316 assert_eq!(context.initialize().initialized, None);
317 assert_eq!(context.final_phase(), Phase::AfterInitializeSuccess);
318 }
319
320 mod state_machine_properties {
323 use super::*;
324 use proptest::prelude::*;
325 use serde_json::json;
326
327 fn arbitrary_event(seq: u64, choice: u8, direction_bit: bool) -> TraceEvent {
331 let payload = match choice % 7 {
332 0 => json!({"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}),
333 1 => json!({"jsonrpc":"2.0","id":1,"result":{"protocolVersion":"2025-11-25"}}),
334 2 => json!({"jsonrpc":"2.0","id":1,"error":{"code":-32602,"message":"x"}}),
335 3 => json!({"jsonrpc":"2.0","method":"notifications/initialized"}),
336 4 => json!({"jsonrpc":"2.0","id":99,"result":{}}),
337 5 => json!({"jsonrpc":"2.0","id":2,"method":"tools/list"}),
338 _ => json!({"jsonrpc":"2.0","method":"notifications/cancelled"}),
339 };
340 let direction = if direction_bit {
341 "client-to-server"
342 } else {
343 "server-to-client"
344 };
345 serde_json::from_value(json!({
346 "seq": seq,
347 "direction": direction,
348 "transport": "stdio",
349 "kind": "message",
350 "payload": payload,
351 }))
352 .unwrap()
353 }
354
355 const fn edge_is_legal(from: Phase, to: Phase) -> bool {
357 matches!(
358 (from, to),
359 (
360 Phase::BeforeInitialize,
361 Phase::BeforeInitialize | Phase::AwaitingInitializeResult
362 ) | (
363 Phase::AwaitingInitializeResult,
364 Phase::AwaitingInitializeResult
365 | Phase::AfterInitializeSuccess
366 | Phase::AfterInitializeError
367 ) | (
368 Phase::AfterInitializeSuccess,
369 Phase::AfterInitializeSuccess | Phase::Ready
370 ) | (Phase::AfterInitializeError, Phase::AfterInitializeError)
371 | (Phase::Ready, Phase::Ready)
372 )
373 }
374
375 proptest! {
376 #[test]
377 fn invariants_hold_for_arbitrary_sequences(
378 moves in proptest::collection::vec((any::<u8>(), any::<bool>()), 0..32)
379 ) {
380 let events: Vec<TraceEvent> = moves
381 .iter()
382 .enumerate()
383 .map(|(index, (choice, direction))| {
384 arbitrary_event(index as u64, *choice, *direction)
385 })
386 .collect();
387 let context = TraceContext::new(&events);
388
389 let phases: Vec<Phase> =
391 context.messages().map(|(_, _, phase)| phase).collect();
392 prop_assert_eq!(phases.len(), events.len());
393 for window in phases.windows(2) {
394 prop_assert!(
395 edge_is_legal(window[0], window[1]),
396 "illegal edge {:?} -> {:?}",
397 window[0],
398 window[1]
399 );
400 }
401 if let Some(last) = phases.last() {
402 prop_assert!(
403 edge_is_legal(*last, context.final_phase()),
404 "illegal final edge {:?} -> {:?}",
405 last,
406 context.final_phase()
407 );
408 }
409
410 let init = context.initialize();
412 if init.result.is_some() || init.initialized.is_some() {
413 prop_assert!(init.request.is_some());
414 }
415 if init.initialized.is_some() {
416 prop_assert!(init.result.is_some());
417 prop_assert_eq!(context.final_phase(), Phase::Ready);
418 }
419 if context.final_phase() == Phase::Ready {
420 prop_assert!(init.initialized.is_some());
421 }
422 }
423 }
424 }
425}