use super::version::{Era, SUPPORTED_PROTOCOL_VERSIONS};
use super::{Implementation, ProtocolVersion};
use crate::types::capabilities::ClientCapabilities;
#[must_use]
pub(crate) fn default_accept_list() -> Vec<ProtocolVersion> {
SUPPORTED_PROTOCOL_VERSIONS
.iter()
.map(|v| ProtocolVersion((*v).to_string()))
.collect()
}
#[must_use]
pub(crate) fn normalize_accept_list(
versions: impl IntoIterator<Item = ProtocolVersion>,
) -> Vec<ProtocolVersion> {
let collected: Vec<ProtocolVersion> = versions.into_iter().collect();
if collected.is_empty() {
default_accept_list()
} else {
collected
}
}
#[must_use]
pub(crate) fn is_v2_opted_in(accept_list: &[ProtocolVersion]) -> bool {
accept_list
.iter()
.any(|v| super::version::protocol_era(v.as_str()) == Era::V2)
}
const MAX_TRACE_VALUE_LEN: usize = 8192;
#[derive(Debug, Clone)]
pub(crate) struct VerifiedContinuation {
pub state: serde_json::Value,
pub round: u8,
}
#[derive(Clone)]
#[non_exhaustive]
pub struct ProtocolContext {
pub era: Era,
pub negotiated_version: ProtocolVersion,
pub client_info: Option<Implementation>,
pub client_capabilities: Option<ClientCapabilities>,
pub(crate) mrtr: Option<crate::types::mrtr::MrtrRequestParams>,
pub(crate) mrtr_verified: Option<VerifiedContinuation>,
}
impl std::fmt::Debug for ProtocolContext {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ProtocolContext")
.field("era", &self.era)
.field("negotiated_version", &self.negotiated_version)
.field("client_info", &self.client_info)
.field("client_capabilities", &self.client_capabilities)
.field("has_input_responses", &self.input_responses().is_some())
.field(
"has_request_state_token",
&self.request_state_token().is_some(),
)
.field("has_verified_continuation", &self.mrtr_verified.is_some())
.field("mrtr_round", &self.mrtr_round())
.finish()
}
}
impl ProtocolContext {
#[must_use]
pub fn new(era: Era, negotiated_version: ProtocolVersion) -> Self {
Self {
era,
negotiated_version,
client_info: None,
client_capabilities: None,
mrtr: None,
mrtr_verified: None,
}
}
#[must_use]
pub fn with_client_info(mut self, client_info: Implementation) -> Self {
self.client_info = Some(client_info);
self
}
#[must_use]
pub fn with_client_capabilities(mut self, client_capabilities: ClientCapabilities) -> Self {
self.client_capabilities = Some(client_capabilities);
self
}
pub(crate) fn input_responses(&self) -> Option<&crate::types::mrtr::InputResponses> {
self.mrtr.as_ref()?.input_responses.as_ref()
}
pub(crate) fn mrtr_continuation(&self) -> Option<&serde_json::Value> {
Some(&self.mrtr_verified.as_ref()?.state)
}
pub(crate) fn mrtr_round(&self) -> Option<u8> {
Some(self.mrtr_verified.as_ref()?.round)
}
}
#[cfg_attr(not(feature = "streamable-http"), allow(dead_code))]
impl ProtocolContext {
#[must_use]
pub(crate) fn with_mrtr_params(mut self, mrtr: crate::types::mrtr::MrtrRequestParams) -> Self {
self.mrtr = Some(mrtr);
self
}
#[must_use]
pub(crate) fn with_verified_continuation(
mut self,
state: serde_json::Value,
round: u8,
) -> Self {
self.mrtr_verified = Some(VerifiedContinuation { state, round });
self
}
#[must_use]
pub(crate) fn without_mrtr(mut self) -> Self {
self.mrtr = None;
self.mrtr_verified = None;
self
}
pub(crate) fn request_state_token(&self) -> Option<&str> {
self.mrtr.as_ref()?.request_state.as_deref()
}
pub(crate) fn input_responses_raw(
&self,
) -> Option<&serde_json::Map<String, serde_json::Value>> {
self.mrtr.as_ref()?.input_responses_raw.as_ref()
}
#[must_use]
pub(crate) fn with_kind_directed_input_responses(
mut self,
responses: crate::types::mrtr::InputResponses,
) -> Self {
if let Some(mrtr) = self.mrtr.as_mut() {
mrtr.input_responses = Some(responses);
mrtr.input_responses_raw = None;
}
self
}
}
pub(crate) const RESERVED_PROTOCOL_VERSION_KEY: &str = "io.modelcontextprotocol/protocolVersion";
pub(crate) const RESERVED_CLIENT_INFO_KEY: &str = "io.modelcontextprotocol/clientInfo";
pub(crate) const RESERVED_CLIENT_CAPABILITIES_KEY: &str =
"io.modelcontextprotocol/clientCapabilities";
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) enum ProtocolNegotiationError {
UnsupportedVersion(String),
MalformedMeta(&'static str),
}
pub(crate) fn resolve_protocol_context(
accept_list: &[ProtocolVersion],
meta: Option<&serde_json::Value>,
) -> Result<Option<ProtocolContext>, ProtocolNegotiationError> {
let meta_obj =
match meta {
Some(value) => Some(value.as_object().ok_or(
ProtocolNegotiationError::MalformedMeta("_meta is not an object"),
)?),
None => None,
};
let negotiated_version = resolve_negotiated_version(accept_list, meta_obj)?;
let era = super::version::protocol_era(negotiated_version.as_str());
let mut ctx = ProtocolContext::new(era, negotiated_version);
if let Some(info) = parse_reserved_object::<Implementation>(
meta_obj,
RESERVED_CLIENT_INFO_KEY,
"clientInfo is not deserializable",
)? {
ctx = ctx.with_client_info(info);
}
if let Some(caps) = parse_reserved_object::<ClientCapabilities>(
meta_obj,
RESERVED_CLIENT_CAPABILITIES_KEY,
"clientCapabilities is not deserializable",
)? {
ctx = ctx.with_client_capabilities(caps);
}
Ok(Some(ctx))
}
fn resolve_negotiated_version(
accept_list: &[ProtocolVersion],
meta_obj: Option<&serde_json::Map<String, serde_json::Value>>,
) -> Result<ProtocolVersion, ProtocolNegotiationError> {
match meta_obj.and_then(|m| m.get(RESERVED_PROTOCOL_VERSION_KEY)) {
Some(raw) => {
let requested = raw.as_str().ok_or(ProtocolNegotiationError::MalformedMeta(
"protocolVersion is not a string",
))?;
if accept_list.iter().any(|v| v.as_str() == requested) {
Ok(ProtocolVersion(requested.to_string()))
} else {
Err(ProtocolNegotiationError::UnsupportedVersion(
requested.to_string(),
))
}
},
None => accept_list
.iter()
.find(|v| super::version::protocol_era(v.as_str()) == Era::V1)
.cloned()
.ok_or(ProtocolNegotiationError::UnsupportedVersion(String::new())),
}
}
fn parse_reserved_object<T: serde::de::DeserializeOwned>(
meta_obj: Option<&serde_json::Map<String, serde_json::Value>>,
key: &str,
malformed: &'static str,
) -> Result<Option<T>, ProtocolNegotiationError> {
match meta_obj.and_then(|m| m.get(key)) {
Some(raw) => serde_json::from_value::<T>(raw.clone())
.map(Some)
.map_err(|_| ProtocolNegotiationError::MalformedMeta(malformed)),
None => Ok(None),
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct TraceContext {
pub traceparent: String,
pub tracestate: Option<String>,
pub baggage: Option<String>,
}
impl TraceContext {
#[must_use]
pub fn from_meta(meta: &serde_json::Value) -> Option<Self> {
let traceparent = bounded_trace_value(meta, "traceparent")?;
let tracestate = bounded_trace_value(meta, "tracestate");
let baggage = bounded_trace_value(meta, "baggage");
Some(Self {
traceparent,
tracestate,
baggage,
})
}
}
fn bounded_trace_value(meta: &serde_json::Value, key: &str) -> Option<String> {
let value = meta.get(key)?.as_str()?;
if value.len() > MAX_TRACE_VALUE_LEN {
None
} else {
Some(value.to_string())
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn dual_accept_list() -> Vec<ProtocolVersion> {
vec![
ProtocolVersion("2025-11-25".to_string()),
ProtocolVersion("2026-07-28".to_string()),
]
}
#[test]
fn resolve_in_list_v2_signal_classifies_v2() {
let meta = json!({ RESERVED_PROTOCOL_VERSION_KEY: "2026-07-28" });
let ctx = resolve_protocol_context(&dual_accept_list(), Some(&meta))
.expect("v2 in accept-list => Ok")
.expect("resolved => Some");
assert_eq!(ctx.era, Era::V2);
assert_eq!(ctx.negotiated_version.as_str(), "2026-07-28");
}
#[test]
fn resolve_absent_signal_falls_back_to_v1() {
let ctx = resolve_protocol_context(&dual_accept_list(), None)
.expect("v1 in accept-list => Ok")
.expect("resolved => Some");
assert_eq!(ctx.era, Era::V1);
assert_eq!(ctx.negotiated_version.as_str(), "2025-11-25");
}
#[test]
fn resolve_unsupported_version_errors() {
let meta = json!({ RESERVED_PROTOCOL_VERSION_KEY: "1999-01-01" });
let err = resolve_protocol_context(&dual_accept_list(), Some(&meta))
.expect_err("version not in accept-list => Err");
assert_eq!(
err,
ProtocolNegotiationError::UnsupportedVersion("1999-01-01".to_string())
);
}
#[test]
fn resolve_v2_only_no_signal_errors() {
let v2_only = vec![ProtocolVersion("2026-07-28".to_string())];
let err = resolve_protocol_context(&v2_only, None).expect_err("v2-only + no signal => Err");
assert_eq!(
err,
ProtocolNegotiationError::UnsupportedVersion(String::new())
);
}
#[test]
fn resolve_malformed_reserved_key_errors() {
let meta = json!({ RESERVED_PROTOCOL_VERSION_KEY: 42 });
let err = resolve_protocol_context(&dual_accept_list(), Some(&meta))
.expect_err("non-string protocolVersion => Err");
assert!(matches!(err, ProtocolNegotiationError::MalformedMeta(_)));
let non_object = json!("not-an-object");
let err = resolve_protocol_context(&dual_accept_list(), Some(&non_object))
.expect_err("non-object _meta => Err");
assert!(matches!(err, ProtocolNegotiationError::MalformedMeta(_)));
let bad_info = json!({
RESERVED_PROTOCOL_VERSION_KEY: "2026-07-28",
RESERVED_CLIENT_INFO_KEY: "should-be-an-object",
});
let err = resolve_protocol_context(&dual_accept_list(), Some(&bad_info))
.expect_err("malformed clientInfo => Err");
assert!(matches!(err, ProtocolNegotiationError::MalformedMeta(_)));
}
#[test]
fn resolve_unknown_extension_key_is_ignored() {
let meta = json!({
RESERVED_PROTOCOL_VERSION_KEY: "2026-07-28",
"com.example/whatever": { "anything": [1, 2, 3] },
});
let ctx = resolve_protocol_context(&dual_accept_list(), Some(&meta))
.expect("unknown key ignored => Ok")
.expect("resolved => Some");
assert_eq!(ctx.era, Era::V2);
}
#[test]
fn resolve_populates_client_identity_when_well_formed() {
let meta = json!({
RESERVED_PROTOCOL_VERSION_KEY: "2026-07-28",
RESERVED_CLIENT_INFO_KEY: { "name": "acme-client", "version": "1.2.3" },
RESERVED_CLIENT_CAPABILITIES_KEY: {},
});
let ctx = resolve_protocol_context(&dual_accept_list(), Some(&meta))
.expect("well-formed => Ok")
.expect("resolved => Some");
let info = ctx.client_info.expect("client_info populated");
assert_eq!(info.name, "acme-client");
assert_eq!(info.version, "1.2.3");
assert!(ctx.client_capabilities.is_some());
}
#[test]
fn protocol_context_new_defaults_optionals_to_none() {
let ctx = ProtocolContext::new(Era::V2, ProtocolVersion("2026-07-28".to_string()));
assert_eq!(ctx.era, Era::V2);
assert_eq!(ctx.negotiated_version.as_str(), "2026-07-28");
assert!(ctx.client_info.is_none());
assert!(ctx.client_capabilities.is_none());
}
#[test]
fn protocol_context_without_the_mrtr_builder_has_no_mrtr() {
let ctx = ProtocolContext::new(Era::V2, ProtocolVersion("2026-07-28".to_string()));
assert!(ctx.mrtr.is_none());
assert!(ctx.mrtr_verified.is_none());
assert!(ctx.input_responses().is_none());
assert!(ctx.request_state_token().is_none());
assert!(ctx.mrtr_continuation().is_none());
assert!(ctx.mrtr_round().is_none());
let meta = json!({ RESERVED_PROTOCOL_VERSION_KEY: "2026-07-28" });
let resolved = resolve_protocol_context(&dual_accept_list(), Some(&meta))
.expect("resolves")
.expect("some");
assert!(resolved.mrtr.is_none());
}
#[test]
fn protocol_context_with_mrtr_params_round_trips() {
let mrtr = crate::types::mrtr::MrtrRequestParams {
input_responses: None,
input_responses_raw: None,
request_state: Some("opaque-token".to_string()),
};
let ctx = ProtocolContext::new(Era::V2, ProtocolVersion("2026-07-28".to_string()))
.with_mrtr_params(mrtr);
assert_eq!(ctx.request_state_token(), Some("opaque-token"));
assert!(ctx.mrtr_continuation().is_none());
assert!(ctx.mrtr_round().is_none());
}
#[test]
fn protocol_context_verified_continuation_surfaces_state_and_round() {
let ctx = ProtocolContext::new(Era::V2, ProtocolVersion("2026-07-28".to_string()))
.with_verified_continuation(json!({ "step": 2 }), 3);
assert_eq!(ctx.mrtr_continuation(), Some(&json!({ "step": 2 })));
assert_eq!(ctx.mrtr_round(), Some(3));
}
#[test]
fn protocol_context_without_mrtr_clears_every_signal() {
let mrtr = crate::types::mrtr::MrtrRequestParams {
input_responses: Some(crate::types::mrtr::InputResponses::new()),
input_responses_raw: None,
request_state: Some("opaque-token".to_string()),
};
let ctx = ProtocolContext::new(Era::V2, ProtocolVersion("2026-07-28".to_string()))
.with_mrtr_params(mrtr)
.with_verified_continuation(json!({ "step": 9 }), 4)
.without_mrtr();
assert!(ctx.input_responses().is_none());
assert!(ctx.request_state_token().is_none());
assert!(ctx.mrtr_continuation().is_none());
assert!(ctx.mrtr_round().is_none());
}
#[test]
fn protocol_context_builders_set_optional_fields() {
let ctx = ProtocolContext::new(Era::V1, ProtocolVersion("2025-11-25".to_string()))
.with_client_info(Implementation::new("acme-client", "1.2.3"))
.with_client_capabilities(ClientCapabilities::default());
assert_eq!(ctx.era, Era::V1);
let info = ctx.client_info.expect("client_info set");
assert_eq!(info.name, "acme-client");
assert_eq!(info.version, "1.2.3");
assert!(ctx.client_capabilities.is_some());
}
#[test]
fn trace_context_from_meta_extracts_all_fields() {
let meta = json!({
"traceparent": "00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01",
"tracestate": "rojo=00f067aa0ba902b7",
"baggage": "userId=alice"
});
let tc = TraceContext::from_meta(&meta).expect("traceparent present => Some");
assert_eq!(
tc.traceparent,
"00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01"
);
assert_eq!(tc.tracestate.as_deref(), Some("rojo=00f067aa0ba902b7"));
assert_eq!(tc.baggage.as_deref(), Some("userId=alice"));
}
#[test]
fn trace_context_from_meta_traceparent_only() {
let meta = json!({ "traceparent": "00-abc-def-01" });
let tc = TraceContext::from_meta(&meta).expect("traceparent present => Some");
assert_eq!(tc.traceparent, "00-abc-def-01");
assert!(tc.tracestate.is_none());
assert!(tc.baggage.is_none());
}
#[test]
fn trace_context_from_meta_absent_returns_none() {
assert!(TraceContext::from_meta(&json!({})).is_none());
assert!(TraceContext::from_meta(&json!({ "tracestate": "a=1" })).is_none());
assert!(TraceContext::from_meta(&json!({ "traceparent": 42 })).is_none());
assert!(TraceContext::from_meta(&json!("just a string")).is_none());
assert!(TraceContext::from_meta(&json!([1, 2, 3])).is_none());
assert!(TraceContext::from_meta(&json!(null)).is_none());
}
#[test]
fn trace_context_over_bound_traceparent_yields_none() {
let huge = "a".repeat(MAX_TRACE_VALUE_LEN + 1);
let meta = json!({ "traceparent": huge });
assert!(TraceContext::from_meta(&meta).is_none());
}
#[test]
fn trace_context_over_bound_tracestate_and_baggage_are_dropped() {
let huge = "b".repeat(MAX_TRACE_VALUE_LEN + 1);
let meta = json!({
"traceparent": "00-abc-def-01",
"tracestate": huge,
"baggage": huge,
});
let tc = TraceContext::from_meta(&meta).expect("in-bounds traceparent => Some");
assert_eq!(tc.traceparent, "00-abc-def-01");
assert!(tc.tracestate.is_none());
assert!(tc.baggage.is_none());
}
proptest::proptest! {
#[test]
fn from_meta_holds_invariants_over_arbitrary_meta(
has_traceparent in proptest::prelude::any::<bool>(),
traceparent in ".*",
tracestate in proptest::option::of(".*"),
baggage in proptest::option::of(".*"),
) {
let mut map = serde_json::Map::new();
if has_traceparent {
map.insert("traceparent".into(), serde_json::Value::String(traceparent.clone()));
}
if let Some(ref ts) = tracestate {
map.insert("tracestate".into(), serde_json::Value::String(ts.clone()));
}
if let Some(ref bg) = baggage {
map.insert("baggage".into(), serde_json::Value::String(bg.clone()));
}
let value = serde_json::Value::Object(map);
let result = TraceContext::from_meta(&value);
if !has_traceparent {
proptest::prop_assert!(result.is_none());
} else if traceparent.len() <= MAX_TRACE_VALUE_LEN {
let tc = result.expect("in-bounds traceparent present => Some");
proptest::prop_assert_eq!(&tc.traceparent, &traceparent);
proptest::prop_assert!(tc.traceparent.len() <= MAX_TRACE_VALUE_LEN);
if let Some(ref ts) = tc.tracestate {
proptest::prop_assert!(ts.len() <= MAX_TRACE_VALUE_LEN);
}
if let Some(ref bg) = tc.baggage {
proptest::prop_assert!(bg.len() <= MAX_TRACE_VALUE_LEN);
}
}
}
}
}