Skip to content

Commit e92906f

Browse files
author
ump45nose
committed
fix(model): preserve borrowed JSON-RPC payload deserialization
Reject conflicting result/error shapes without forcing an owned serde_json::Value intermediate for from_str/from_slice. Use serde_json's RawValue newtype token so Deserialize<'de> payload types can keep borrows into the original input, with an owned fallback for from_value. Adds regression coverage for generic borrowed request/response payloads.
1 parent c594415 commit e92906f

2 files changed

Lines changed: 218 additions & 28 deletions

File tree

‎crates/rmcp/Cargo.toml‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -48,7 +48,7 @@ rustdoc-args = ["--cfg", "docsrs"]
4848
[dependencies]
4949
async-trait = { version = "0.1.89", optional = true }
5050
serde = { version = "1.0", features = ["derive", "rc"] }
51-
serde_json = "1.0"
51+
serde_json = { version = "1.0", features = ["raw_value"] }
5252
thiserror = "2"
5353
tokio = { version = "1", features = ["sync", "macros", "rt", "time"] }
5454
futures = "0.3"

‎crates/rmcp/src/model.rs‎

Lines changed: 217 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -788,6 +788,9 @@ impl<Req, Resp, Not> JsonRpcMessage<Req, Resp, Not> {
788788
}
789789
}
790790

791+
/// Token used by `serde_json` for zero-copy raw JSON value deserialization.
792+
const JSON_RAW_VALUE_TOKEN: &str = "$serde_json::private::RawValue";
793+
791794
impl<'de, Req, Resp, Noti> serde::Deserialize<'de> for JsonRpcMessage<Req, Resp, Noti>
792795
where
793796
Req: Deserialize<'de>,
@@ -798,8 +801,6 @@ where
798801
where
799802
__D: serde::Deserializer<'de>,
800803
{
801-
use serde::de::Error as _;
802-
803804
// JSON-RPC 2.0 defines four mutually exclusive message shapes. A derived
804805
// untagged enum cannot express that: a response carrying both `result`
805806
// and `error` matches the `Response` variant (the stray `error` is
@@ -808,56 +809,193 @@ where
808809
// so spec violations are rejected instead of silently resolved, while
809810
// extra extension fields stay accepted.
810811
//
811-
// Bounds stay `Deserialize<'de>` (not `DeserializeOwned`) so the public
812-
// API matches the previous `#[derive(Deserialize)]` impl. We deserialize
813-
// variants via `Deserialize::deserialize(value)` rather than
814-
// `serde_json::from_value`, which would force `DeserializeOwned`.
815-
let value = Value::deserialize(deserializer)?;
816-
let obj = match value.as_object() {
817-
Some(obj) => obj,
812+
// Use serde_json's RawValue newtype token so `from_str`/`from_slice` can
813+
// borrow the original JSON text (preserving `Deserialize<'de>` payload
814+
// borrows), while `from_value` still works via an owned buffer.
815+
deserializer.deserialize_newtype_struct(
816+
JSON_RAW_VALUE_TOKEN,
817+
JsonRpcMessageVisitor {
818+
_marker: std::marker::PhantomData,
819+
},
820+
)
821+
}
822+
}
823+
824+
struct JsonRpcMessageVisitor<Req, Resp, Noti> {
825+
_marker: std::marker::PhantomData<fn() -> (Req, Resp, Noti)>,
826+
}
827+
828+
impl<'de, Req, Resp, Noti> serde::de::Visitor<'de> for JsonRpcMessageVisitor<Req, Resp, Noti>
829+
where
830+
Req: Deserialize<'de>,
831+
Resp: Deserialize<'de>,
832+
Noti: Deserialize<'de>,
833+
{
834+
type Value = JsonRpcMessage<Req, Resp, Noti>;
835+
836+
fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
837+
formatter.write_str("a JSON-RPC 2.0 message")
838+
}
839+
840+
fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
841+
where
842+
A: serde::de::MapAccess<'de>,
843+
{
844+
use serde::de::Error as _;
845+
846+
// serde_json RawValue protocol: a single field named TOKEN whose value
847+
// is the raw JSON text (borrowed for from_str, owned for from_value).
848+
let key: Option<std::borrow::Cow<'de, str>> = map.next_key()?;
849+
match key.as_deref() {
850+
Some(JSON_RAW_VALUE_TOKEN) => {}
851+
Some(other) => {
852+
return Err(A::Error::custom(format!(
853+
"unexpected raw value key {other:?}"
854+
)));
855+
}
818856
None => {
819-
return Err(__D::Error::custom(
820-
"data did not match any variant of untagged enum JsonRpcMessage",
821-
));
857+
return Err(A::Error::invalid_type(serde::de::Unexpected::Map, &self));
822858
}
823-
};
859+
}
860+
861+
let text = map.next_value_seed(PreferBorrowedStr)?;
862+
if map.next_key::<serde::de::IgnoredAny>()?.is_some() {
863+
return Err(A::Error::custom("unexpected extra raw value field"));
864+
}
865+
866+
match text {
867+
std::borrow::Cow::Borrowed(text) => {
868+
JsonRpcMessage::<Req, Resp, Noti>::from_checked_json_text(text)
869+
.map_err(A::Error::custom)
870+
}
871+
std::borrow::Cow::Owned(text) => {
872+
let value: Value = serde_json::from_str(&text).map_err(A::Error::custom)?;
873+
JsonRpcMessage::<Req, Resp, Noti>::from_checked_owned_value(value)
874+
.map_err(A::Error::custom)
875+
}
876+
}
877+
}
878+
}
879+
880+
/// Like `Cow<'de, str>` but guarantees `visit_borrowed_str` stays borrowed.
881+
struct PreferBorrowedStr;
882+
883+
impl<'de> serde::de::DeserializeSeed<'de> for PreferBorrowedStr {
884+
type Value = std::borrow::Cow<'de, str>;
885+
886+
fn deserialize<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
887+
where
888+
D: serde::Deserializer<'de>,
889+
{
890+
struct PreferBorrowedStrVisitor;
891+
impl<'de> serde::de::Visitor<'de> for PreferBorrowedStrVisitor {
892+
type Value = std::borrow::Cow<'de, str>;
893+
894+
fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
895+
formatter.write_str("a raw JSON string")
896+
}
897+
898+
fn visit_borrowed_str<E>(self, v: &'de str) -> Result<Self::Value, E>
899+
where
900+
E: serde::de::Error,
901+
{
902+
Ok(std::borrow::Cow::Borrowed(v))
903+
}
904+
905+
fn visit_str<E>(self, v: &str) -> Result<Self::Value, E>
906+
where
907+
E: serde::de::Error,
908+
{
909+
Ok(std::borrow::Cow::Owned(v.to_owned()))
910+
}
911+
912+
fn visit_string<E>(self, v: String) -> Result<Self::Value, E>
913+
where
914+
E: serde::de::Error,
915+
{
916+
Ok(std::borrow::Cow::Owned(v))
917+
}
918+
}
919+
920+
deserializer.deserialize_str(PreferBorrowedStrVisitor)
921+
}
922+
}
923+
924+
impl<'de, Req, Resp, Noti> JsonRpcMessage<Req, Resp, Noti>
925+
where
926+
Req: Deserialize<'de>,
927+
Resp: Deserialize<'de>,
928+
Noti: Deserialize<'de>,
929+
{
930+
fn classify_object_keys(
931+
obj: &serde_json::Map<String, Value>,
932+
) -> Result<(bool, bool, bool, bool), serde_json::Error> {
824933
let has_method = obj.contains_key("method");
825934
let has_id = obj.contains_key("id");
826935
let has_result = obj.contains_key("result");
827936
let has_error = obj.contains_key("error");
828937

829938
if has_result && has_error {
830-
return Err(__D::Error::custom(
939+
return Err(serde::de::Error::custom(
831940
"invalid JSON-RPC message: both `result` and `error` are present",
832941
));
833942
}
834943
if has_method && (has_result || has_error) {
835-
return Err(__D::Error::custom(
944+
return Err(serde::de::Error::custom(
836945
"invalid JSON-RPC message: a request or notification must not carry `result` or `error`",
837946
));
838947
}
948+
Ok((has_method, has_id, has_result, has_error))
949+
}
950+
951+
fn from_checked_json_text(text: &'de str) -> Result<Self, serde_json::Error> {
952+
let value: Value = serde_json::from_str(text)?;
953+
let obj = value.as_object().ok_or_else(|| {
954+
serde::de::Error::custom(
955+
"data did not match any variant of untagged enum JsonRpcMessage",
956+
)
957+
})?;
958+
let (has_method, has_id, has_result, has_error) = Self::classify_object_keys(obj)?;
839959

840960
if has_method {
841961
if has_id {
842-
return JsonRpcRequest::<Req>::deserialize(value)
843-
.map(JsonRpcMessage::Request)
844-
.map_err(__D::Error::custom);
962+
return serde_json::from_str(text).map(JsonRpcMessage::Request);
963+
}
964+
return serde_json::from_str(text).map(JsonRpcMessage::Notification);
965+
}
966+
if has_error {
967+
return serde_json::from_str(text).map(JsonRpcMessage::Error);
968+
}
969+
if has_result {
970+
return serde_json::from_str(text).map(JsonRpcMessage::Response);
971+
}
972+
Err(serde::de::Error::custom(
973+
"data did not match any variant of untagged enum JsonRpcMessage",
974+
))
975+
}
976+
977+
fn from_checked_owned_value(value: Value) -> Result<Self, serde_json::Error> {
978+
let obj = value.as_object().ok_or_else(|| {
979+
serde::de::Error::custom(
980+
"data did not match any variant of untagged enum JsonRpcMessage",
981+
)
982+
})?;
983+
let (has_method, has_id, has_result, has_error) = Self::classify_object_keys(obj)?;
984+
985+
if has_method {
986+
if has_id {
987+
return JsonRpcRequest::<Req>::deserialize(value).map(JsonRpcMessage::Request);
845988
}
846989
return JsonRpcNotification::<Noti>::deserialize(value)
847-
.map(JsonRpcMessage::Notification)
848-
.map_err(__D::Error::custom);
990+
.map(JsonRpcMessage::Notification);
849991
}
850992
if has_error {
851-
return JsonRpcError::deserialize(value)
852-
.map(JsonRpcMessage::Error)
853-
.map_err(__D::Error::custom);
993+
return JsonRpcError::deserialize(value).map(JsonRpcMessage::Error);
854994
}
855995
if has_result {
856-
return JsonRpcResponse::<Resp>::deserialize(value)
857-
.map(JsonRpcMessage::Response)
858-
.map_err(__D::Error::custom);
996+
return JsonRpcResponse::<Resp>::deserialize(value).map(JsonRpcMessage::Response);
859997
}
860-
Err(__D::Error::custom(
998+
Err(serde::de::Error::custom(
861999
"data did not match any variant of untagged enum JsonRpcMessage",
8621000
))
8631001
}
@@ -5057,6 +5195,58 @@ mod tests {
50575195
);
50585196
}
50595197

5198+
#[test]
5199+
fn generic_borrowed_request_deserializes_from_str() {
5200+
// Regression: field-presence dispatch must not force owned Value so hard
5201+
// that generic payload types containing `&str` lose input borrows.
5202+
#[derive(Debug, Deserialize)]
5203+
struct BorrowedRequest<'a> {
5204+
method: &'a str,
5205+
}
5206+
5207+
let input = r#"{"jsonrpc":"2.0","id":1,"method":"ping"}"#;
5208+
let parsed =
5209+
serde_json::from_str::<JsonRpcMessage<BorrowedRequest<'_>, Value, Value>>(input)
5210+
.expect("generic Deserialize API accepts borrowed request strings");
5211+
match parsed {
5212+
JsonRpcMessage::Request(request) => {
5213+
assert_eq!(request.request.method, "ping");
5214+
let start = input.as_ptr() as usize;
5215+
let borrowed = request.request.method.as_ptr() as usize;
5216+
assert!(
5217+
borrowed >= start && borrowed < start + input.len(),
5218+
"method must borrow from the original input"
5219+
);
5220+
}
5221+
other => panic!("expected Request, received {other:?}"),
5222+
}
5223+
}
5224+
5225+
#[test]
5226+
fn generic_borrowed_response_deserializes_from_str() {
5227+
#[derive(Debug, Deserialize)]
5228+
struct BorrowedResponse<'a> {
5229+
status: &'a str,
5230+
}
5231+
5232+
let input = r#"{"jsonrpc":"2.0","id":1,"result":{"status":"ok"}}"#;
5233+
let parsed =
5234+
serde_json::from_str::<JsonRpcMessage<Value, BorrowedResponse<'_>, Value>>(input)
5235+
.expect("generic Deserialize API accepts borrowed response strings");
5236+
match parsed {
5237+
JsonRpcMessage::Response(response) => {
5238+
assert_eq!(response.result.status, "ok");
5239+
let start = input.as_ptr() as usize;
5240+
let borrowed = response.result.status.as_ptr() as usize;
5241+
assert!(
5242+
borrowed >= start && borrowed < start + input.len(),
5243+
"status must borrow from the original input"
5244+
);
5245+
}
5246+
other => panic!("expected Response, received {other:?}"),
5247+
}
5248+
}
5249+
50605250
#[test]
50615251
fn test_clean_shapes_and_extension_fields_still_deserialize() {
50625252
// The four clean shapes keep working.

0 commit comments

Comments
 (0)