1use std::borrow::Cow;
2
3use serde::{Deserialize, Serialize};
4
5use super::{
6 CustomNotification, CustomRequest, Extensions, JsonObject, MetaObject, Notification,
7 NotificationMetaObject, NotificationNoParam, Request, RequestMetaObject, RequestNoParam,
8 RequestOptionalParam,
9};
10
11pub(crate) mod json_float {
15 use serde::{Deserialize, Deserializer, de::Error};
16 use serde_json::Number;
17
18 fn to_f64<E: Error>(number: Number) -> Result<f64, E> {
19 number
22 .as_f64()
23 .ok_or_else(|| E::custom(format_args!("number out of range: {number}")))
24 }
25
26 pub(crate) fn f64<'de, D: Deserializer<'de>>(deserializer: D) -> Result<f64, D::Error> {
27 to_f64(Number::deserialize(deserializer)?)
28 }
29
30 pub(crate) fn option_f64<'de, D: Deserializer<'de>>(
31 deserializer: D,
32 ) -> Result<Option<f64>, D::Error> {
33 Option::<Number>::deserialize(deserializer)?
34 .map(to_f64)
35 .transpose()
36 }
37
38 pub(crate) fn option_f32<'de, D: Deserializer<'de>>(
39 deserializer: D,
40 ) -> Result<Option<f32>, D::Error> {
41 Ok(option_f64(deserializer)?.map(|value| value as f32))
43 }
44}
45
46#[derive(Deserialize)]
52struct WithMeta<'a, P> {
53 _meta: Option<Cow<'a, JsonObject>>,
54 #[serde(flatten)]
55 _rest: P,
56}
57
58impl<P: Serialize> Serialize for WithMeta<'_, P> {
59 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
60 where
61 S: serde::Serializer,
62 {
63 use serde::ser::SerializeMap;
64
65 let mut rest_value =
67 serde_json::to_value(&self._rest).map_err(serde::ser::Error::custom)?;
68
69 let params_meta: Option<JsonObject> = rest_value
71 .as_object_mut()
72 .and_then(|obj| obj.remove("_meta"))
73 .and_then(|v| serde_json::from_value(v).ok());
74
75 let merged_meta = match (self._meta.as_deref(), params_meta) {
77 (Some(ext_meta), Some(mut params_meta)) => {
78 params_meta.extend(ext_meta.clone());
79 Some(params_meta)
80 }
81 (Some(ext_meta), None) => Some(ext_meta.clone()),
82 (None, Some(params_meta)) => Some(params_meta),
83 (None, None) => None,
84 };
85
86 let rest_obj = match rest_value {
88 serde_json::Value::Object(map) => map,
89 _ => serde_json::Map::new(),
90 };
91 let meta_count = usize::from(merged_meta.is_some());
92 let mut map = serializer.serialize_map(Some(rest_obj.len() + meta_count))?;
93
94 if let Some(meta) = &merged_meta {
95 map.serialize_entry("_meta", meta)?;
96 }
97
98 for (k, v) in &rest_obj {
99 map.serialize_entry(k, v)?;
100 }
101
102 map.end()
103 }
104}
105
106#[derive(Serialize, Deserialize)]
107struct Proxy<'a, M, P> {
108 method: M,
109 params: WithMeta<'a, P>,
110}
111
112#[derive(Serialize, Deserialize)]
113struct ProxyOptionalParam<'a, M, P> {
114 method: M,
115 params: Option<WithMeta<'a, P>>,
116}
117
118#[derive(Serialize, Deserialize)]
119struct ProxyNoParam<M> {
120 method: M,
121}
122
123fn merge_legacy_meta<'a>(
127 typed: Option<&'a JsonObject>,
128 extensions: &'a Extensions,
129) -> Option<Cow<'a, JsonObject>> {
130 let legacy = extensions.get::<MetaObject>().map(|meta| &meta.0);
131 match (typed, legacy) {
132 (Some(typed), None) => Some(Cow::Borrowed(typed)),
133 (None, Some(legacy)) => Some(Cow::Borrowed(legacy)),
134 (Some(typed), Some(legacy)) => {
135 let mut merged = legacy.clone();
136 merged.extend(
137 typed
138 .iter()
139 .map(|(key, value)| (key.clone(), value.clone())),
140 );
141 Some(Cow::Owned(merged))
142 }
143 (None, None) => None,
144 }
145}
146
147fn request_meta(extensions: &Extensions) -> Option<Cow<'_, JsonObject>> {
149 let typed = extensions.get::<RequestMetaObject>().map(|meta| &meta.0.0);
150 merge_legacy_meta(typed, extensions)
151}
152
153fn notification_meta(extensions: &Extensions) -> Option<Cow<'_, JsonObject>> {
155 let typed = extensions
156 .get::<NotificationMetaObject>()
157 .map(|meta| &meta.0.0);
158 merge_legacy_meta(typed, extensions)
159}
160
161fn extensions_with_meta<T>(meta: Option<Cow<'_, JsonObject>>) -> Extensions
163where
164 T: From<JsonObject> + Clone + Send + Sync + 'static,
165{
166 let mut extensions = Extensions::new();
167 if let Some(meta) = meta {
168 extensions.insert(T::from(meta.into_owned()));
169 }
170 extensions
171}
172
173impl<M, R> Serialize for Request<M, R>
174where
175 M: Serialize,
176 R: Serialize,
177{
178 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
179 where
180 S: serde::Serializer,
181 {
182 Proxy::serialize(
183 &Proxy {
184 method: &self.method,
185 params: WithMeta {
186 _rest: &self.params,
187 _meta: request_meta(&self.extensions),
188 },
189 },
190 serializer,
191 )
192 }
193}
194
195impl<'de, M, R> Deserialize<'de> for Request<M, R>
196where
197 M: Deserialize<'de>,
198 R: Deserialize<'de>,
199{
200 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
201 where
202 D: serde::Deserializer<'de>,
203 {
204 let body = Proxy::deserialize(deserializer)?;
205 Ok(Request {
206 extensions: extensions_with_meta::<RequestMetaObject>(body.params._meta),
207 method: body.method,
208 params: body.params._rest,
209 })
210 }
211}
212
213impl<M, R> Serialize for RequestOptionalParam<M, R>
214where
215 M: Serialize,
216 R: Serialize,
217{
218 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
219 where
220 S: serde::Serializer,
221 {
222 Proxy::serialize(
223 &Proxy {
224 method: &self.method,
225 params: WithMeta {
226 _rest: &self.params,
227 _meta: request_meta(&self.extensions),
228 },
229 },
230 serializer,
231 )
232 }
233}
234
235impl<'de, M, R> Deserialize<'de> for RequestOptionalParam<M, R>
236where
237 M: Deserialize<'de>,
238 R: Deserialize<'de>,
239{
240 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
241 where
242 D: serde::Deserializer<'de>,
243 {
244 let body = ProxyOptionalParam::<'_, _, Option<R>>::deserialize(deserializer)?;
245 let mut params = None;
246 let mut _meta = None;
247 if let Some(body_params) = body.params {
248 params = body_params._rest;
249 _meta = body_params._meta;
250 }
251 Ok(RequestOptionalParam {
252 extensions: extensions_with_meta::<RequestMetaObject>(_meta),
253 method: body.method,
254 params,
255 })
256 }
257}
258
259impl<M> Serialize for RequestNoParam<M>
260where
261 M: Serialize,
262{
263 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
264 where
265 S: serde::Serializer,
266 {
267 match request_meta(&self.extensions) {
270 Some(_meta) => Proxy::serialize(
271 &Proxy {
272 method: &self.method,
273 params: WithMeta {
274 _meta: Some(_meta),
275 _rest: (),
276 },
277 },
278 serializer,
279 ),
280 None => ProxyNoParam::serialize(
281 &ProxyNoParam {
282 method: &self.method,
283 },
284 serializer,
285 ),
286 }
287 }
288}
289
290impl<'de, M> Deserialize<'de> for RequestNoParam<M>
291where
292 M: Deserialize<'de>,
293{
294 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
295 where
296 D: serde::Deserializer<'de>,
297 {
298 let body = ProxyOptionalParam::<'_, _, Option<JsonObject>>::deserialize(deserializer)?;
299 let _meta = body.params.and_then(|params| params._meta);
300 Ok(RequestNoParam {
301 extensions: extensions_with_meta::<RequestMetaObject>(_meta),
302 method: body.method,
303 })
304 }
305}
306
307impl<M, R> Serialize for Notification<M, R>
308where
309 M: Serialize,
310 R: Serialize,
311{
312 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
313 where
314 S: serde::Serializer,
315 {
316 Proxy::serialize(
317 &Proxy {
318 method: &self.method,
319 params: WithMeta {
320 _rest: &self.params,
321 _meta: notification_meta(&self.extensions),
322 },
323 },
324 serializer,
325 )
326 }
327}
328
329impl<'de, M, R> Deserialize<'de> for Notification<M, R>
330where
331 M: Deserialize<'de>,
332 R: Deserialize<'de>,
333{
334 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
335 where
336 D: serde::Deserializer<'de>,
337 {
338 let body = ProxyOptionalParam::<'_, _, R>::deserialize(deserializer)?;
339 let (_meta, params) = match body.params {
340 Some(with_meta) => (with_meta._meta, with_meta._rest),
341 None => {
342 let empty = serde_json::Value::Object(serde_json::Map::new());
344 let r = R::deserialize(empty).map_err(serde::de::Error::custom)?;
345 (None, r)
346 }
347 };
348 Ok(Notification {
349 extensions: extensions_with_meta::<NotificationMetaObject>(_meta),
350 method: body.method,
351 params,
352 })
353 }
354}
355
356impl<M> Serialize for NotificationNoParam<M>
357where
358 M: Serialize,
359{
360 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
361 where
362 S: serde::Serializer,
363 {
364 match notification_meta(&self.extensions) {
367 Some(_meta) => Proxy::serialize(
368 &Proxy {
369 method: &self.method,
370 params: WithMeta {
371 _meta: Some(_meta),
372 _rest: (),
373 },
374 },
375 serializer,
376 ),
377 None => ProxyNoParam::serialize(
378 &ProxyNoParam {
379 method: &self.method,
380 },
381 serializer,
382 ),
383 }
384 }
385}
386
387impl<'de, M> Deserialize<'de> for NotificationNoParam<M>
388where
389 M: Deserialize<'de>,
390{
391 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
392 where
393 D: serde::Deserializer<'de>,
394 {
395 let body = ProxyOptionalParam::<'_, _, Option<JsonObject>>::deserialize(deserializer)?;
396 let _meta = body.params.and_then(|params| params._meta);
397 Ok(NotificationNoParam {
398 extensions: extensions_with_meta::<NotificationMetaObject>(_meta),
399 method: body.method,
400 })
401 }
402}
403
404impl Serialize for CustomRequest {
405 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
406 where
407 S: serde::Serializer,
408 {
409 let _meta = request_meta(&self.extensions);
410 let params = self.params.as_ref();
411
412 let params = if _meta.is_some() || params.is_some() {
413 Some(WithMeta {
414 _meta,
415 _rest: &self.params,
416 })
417 } else {
418 None
419 };
420
421 ProxyOptionalParam::serialize(
422 &ProxyOptionalParam {
423 method: &self.method,
424 params,
425 },
426 serializer,
427 )
428 }
429}
430
431impl<'de> Deserialize<'de> for CustomRequest {
432 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
433 where
434 D: serde::Deserializer<'de>,
435 {
436 let body =
437 ProxyOptionalParam::<'_, _, Option<serde_json::Value>>::deserialize(deserializer)?;
438 let mut params = None;
439 let mut _meta = None;
440 if let Some(body_params) = body.params {
441 params = body_params._rest;
442 _meta = body_params._meta;
443 }
444 Ok(CustomRequest {
445 extensions: extensions_with_meta::<RequestMetaObject>(_meta),
446 method: body.method,
447 params,
448 })
449 }
450}
451
452impl Serialize for CustomNotification {
453 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
454 where
455 S: serde::Serializer,
456 {
457 let _meta = notification_meta(&self.extensions);
458 let params = self.params.as_ref();
459
460 let params = if _meta.is_some() || params.is_some() {
461 Some(WithMeta {
462 _meta,
463 _rest: &self.params,
464 })
465 } else {
466 None
467 };
468
469 ProxyOptionalParam::serialize(
470 &ProxyOptionalParam {
471 method: &self.method,
472 params,
473 },
474 serializer,
475 )
476 }
477}
478
479impl<'de> Deserialize<'de> for CustomNotification {
480 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
481 where
482 D: serde::Deserializer<'de>,
483 {
484 let body =
485 ProxyOptionalParam::<'_, _, Option<serde_json::Value>>::deserialize(deserializer)?;
486 let mut params = None;
487 let mut _meta = None;
488 if let Some(body_params) = body.params {
489 params = body_params._rest;
490 _meta = body_params._meta;
491 }
492 Ok(CustomNotification {
493 extensions: extensions_with_meta::<NotificationMetaObject>(_meta),
494 method: body.method,
495 params,
496 })
497 }
498}
499
500#[cfg(test)]
501mod test {
502 use serde_json::json;
503
504 use crate::model::{
505 CallToolRequest, CallToolRequestParams, CustomRequest, Extensions, InitializedNotification,
506 ListToolsRequest, NotificationMetaObject, PingRequest, RequestMetaObject,
507 };
508
509 #[test]
510 fn test_deserialize_lost_tools_request() {
511 let _req: ListToolsRequest = serde_json::from_value(json!(
512 {
513 "method": "tools/list",
514 }
515 ))
516 .unwrap();
517 }
518
519 #[test]
520 fn test_no_duplicate_meta_both_sources() {
521 let mut extensions = Extensions::new();
524 let mut ext_meta = RequestMetaObject::new();
525 ext_meta.insert("traceId".to_string(), json!("abc"));
526 extensions.insert(ext_meta);
527
528 let mut params_meta = RequestMetaObject::new();
529 params_meta.insert("progressToken".to_string(), json!(1));
530
531 let req = CallToolRequest {
532 extensions,
533 method: Default::default(),
534 params: CallToolRequestParams {
535 meta: Some(params_meta),
536 name: "my_tool".into(),
537 arguments: None,
538 input_responses: None,
539 request_state: None,
540 },
541 };
542
543 let value = serde_json::to_value(&req).unwrap();
544 let params = value.get("params").unwrap();
545
546 let meta = params.get("_meta").unwrap();
548
549 assert_eq!(meta.get("traceId").unwrap(), "abc");
551 assert_eq!(meta.get("progressToken").unwrap(), 1);
552
553 let raw = serde_json::to_string(&req).unwrap();
555 assert_eq!(
556 raw.matches("\"_meta\"").count(),
557 1,
558 "Expected exactly one _meta key in serialized output, got: {}",
559 raw
560 );
561 }
562
563 #[test]
564 fn test_meta_only_from_extensions() {
565 let mut extensions = Extensions::new();
566 let mut ext_meta = RequestMetaObject::new();
567 ext_meta.insert("traceId".to_string(), json!("ext-only"));
568 extensions.insert(ext_meta);
569
570 let req = CallToolRequest {
571 extensions,
572 method: Default::default(),
573 params: CallToolRequestParams {
574 meta: None,
575 name: "my_tool".into(),
576 arguments: None,
577 input_responses: None,
578 request_state: None,
579 },
580 };
581
582 let value = serde_json::to_value(&req).unwrap();
583 let meta = value["params"]["_meta"].as_object().unwrap();
584 assert_eq!(meta.get("traceId").unwrap(), "ext-only");
585 }
586
587 #[test]
588 fn test_meta_only_from_params() {
589 let mut params_meta = RequestMetaObject::new();
590 params_meta.insert("progressToken".to_string(), json!(42));
591
592 let req = CallToolRequest {
593 extensions: Extensions::new(),
594 method: Default::default(),
595 params: CallToolRequestParams {
596 meta: Some(params_meta),
597 name: "my_tool".into(),
598 arguments: None,
599 input_responses: None,
600 request_state: None,
601 },
602 };
603
604 let value = serde_json::to_value(&req).unwrap();
605 let meta = value["params"]["_meta"].as_object().unwrap();
606 assert_eq!(meta.get("progressToken").unwrap(), 42);
607 }
608
609 #[test]
610 fn test_no_meta_emitted_when_neither_source() {
611 let req = CallToolRequest {
612 extensions: Extensions::new(),
613 method: Default::default(),
614 params: CallToolRequestParams {
615 meta: None,
616 name: "my_tool".into(),
617 arguments: None,
618 input_responses: None,
619 request_state: None,
620 },
621 };
622
623 let value = serde_json::to_value(&req).unwrap();
624 assert!(
625 value["params"].get("_meta").is_none(),
626 "Expected no _meta when neither source is populated"
627 );
628 }
629
630 #[test]
631 fn test_extensions_meta_takes_priority_on_conflict() {
632 let mut extensions = Extensions::new();
634 let mut ext_meta = RequestMetaObject::new();
635 ext_meta.insert("shared_key".to_string(), json!("from_extensions"));
636 extensions.insert(ext_meta);
637
638 let mut params_meta = RequestMetaObject::new();
639 params_meta.insert("shared_key".to_string(), json!("from_params"));
640 params_meta.insert("params_only".to_string(), json!("kept"));
641
642 let req = CallToolRequest {
643 extensions,
644 method: Default::default(),
645 params: CallToolRequestParams {
646 meta: Some(params_meta),
647 name: "my_tool".into(),
648 arguments: None,
649 input_responses: None,
650 request_state: None,
651 },
652 };
653
654 let value = serde_json::to_value(&req).unwrap();
655 let meta = value["params"]["_meta"].as_object().unwrap();
656 assert_eq!(meta.get("shared_key").unwrap(), "from_extensions");
657 assert_eq!(meta.get("params_only").unwrap(), "kept");
658 }
659
660 #[test]
661 fn test_round_trip_preserves_meta() {
662 let mut extensions = Extensions::new();
663 let mut ext_meta = RequestMetaObject::new();
664 ext_meta.insert("traceId".to_string(), json!("round-trip"));
665 extensions.insert(ext_meta);
666
667 let req = CallToolRequest {
668 extensions,
669 method: Default::default(),
670 params: CallToolRequestParams {
671 meta: None,
672 name: "my_tool".into(),
673 arguments: Some(serde_json::Map::from_iter([("x".to_string(), json!(1))])),
674 input_responses: None,
675 request_state: None,
676 },
677 };
678
679 let serialized = serde_json::to_string(&req).unwrap();
680 let deserialized: CallToolRequest = serde_json::from_str(&serialized).unwrap();
681
682 let meta = deserialized.extensions.get::<RequestMetaObject>().unwrap();
684 assert_eq!(meta.get("traceId").unwrap(), "round-trip");
685
686 assert_eq!(deserialized.params.name, "my_tool");
688 assert_eq!(
689 deserialized
690 .params
691 .arguments
692 .as_ref()
693 .unwrap()
694 .get("x")
695 .unwrap(),
696 &json!(1)
697 );
698 }
699
700 #[test]
701 fn test_custom_request_no_duplicate_meta() {
702 let mut extensions = Extensions::new();
704 let mut ext_meta = RequestMetaObject::new();
705 ext_meta.insert("traceId".to_string(), json!("custom-ext"));
706 extensions.insert(ext_meta);
707
708 let params = Some(json!({
709 "_meta": { "progressToken": 99 },
710 "foo": "bar"
711 }));
712
713 let req = CustomRequest {
714 extensions,
715 method: "custom/method".into(),
716 params,
717 };
718
719 let raw = serde_json::to_string(&req).unwrap();
720 assert_eq!(
721 raw.matches("\"_meta\"").count(),
722 1,
723 "Expected exactly one _meta key in CustomRequest output, got: {}",
724 raw
725 );
726
727 let value: serde_json::Value = serde_json::from_str(&raw).unwrap();
728 let meta = value["params"]["_meta"].as_object().unwrap();
729 assert_eq!(meta.get("traceId").unwrap(), "custom-ext");
730 assert_eq!(meta.get("progressToken").unwrap(), 99);
731 }
732
733 #[test]
734 fn test_request_no_param_meta_round_trip() {
735 let mut extensions = Extensions::new();
737 let mut meta = RequestMetaObject::new();
738 meta.insert("traceId".to_string(), json!("ping-trace"));
739 extensions.insert(meta);
740
741 let req = PingRequest {
742 method: Default::default(),
743 extensions,
744 };
745
746 let value = serde_json::to_value(&req).unwrap();
747 assert_eq!(value["params"]["_meta"]["traceId"], json!("ping-trace"));
748
749 let deserialized: PingRequest = serde_json::from_value(value).unwrap();
750 let meta = deserialized
751 .extensions
752 .get::<RequestMetaObject>()
753 .expect("meta should survive the round-trip");
754 assert_eq!(meta.get("traceId").unwrap(), &json!("ping-trace"));
755 }
756
757 #[test]
758 fn test_request_no_param_without_meta_has_no_params_key() {
759 let req = PingRequest {
760 method: Default::default(),
761 extensions: Extensions::new(),
762 };
763 let value = serde_json::to_value(&req).unwrap();
764 assert!(
765 value.get("params").is_none(),
766 "meta-less no-param requests must keep the historical wire shape: {value}"
767 );
768 }
769
770 #[test]
771 fn test_notification_no_param_meta_round_trip() {
772 let mut extensions = Extensions::new();
774 let mut meta = NotificationMetaObject::new();
775 meta.insert("traceId".to_string(), json!("init-trace"));
776 extensions.insert(meta);
777
778 let notification = InitializedNotification {
779 method: Default::default(),
780 extensions,
781 };
782
783 let value = serde_json::to_value(¬ification).unwrap();
784 assert_eq!(value["params"]["_meta"]["traceId"], json!("init-trace"));
785
786 let deserialized: InitializedNotification = serde_json::from_value(value).unwrap();
787 let meta = deserialized
788 .extensions
789 .get::<NotificationMetaObject>()
790 .expect("meta should survive the round-trip");
791 assert_eq!(meta.get("traceId").unwrap(), &json!("init-trace"));
792 }
793
794 #[test]
795 fn test_notification_no_param_without_meta_has_no_params_key() {
796 let notification = InitializedNotification {
797 method: Default::default(),
798 extensions: Extensions::new(),
799 };
800 let value = serde_json::to_value(¬ification).unwrap();
801 assert!(
802 value.get("params").is_none(),
803 "meta-less no-param notifications must keep the historical wire shape: {value}"
804 );
805 }
806
807 #[test]
808 fn test_no_param_ignores_unknown_params_fields() {
809 let _req: PingRequest =
811 serde_json::from_value(json!({"method": "ping", "params": {}})).unwrap();
812 let _req: PingRequest =
813 serde_json::from_value(json!({"method": "ping", "params": {"unknown": 1}})).unwrap();
814 let _req: PingRequest = serde_json::from_value(json!({"method": "ping"})).unwrap();
815 }
816
817 #[test]
818 fn test_legacy_meta_extension_still_serializes() {
819 let mut extensions = Extensions::new();
820 let mut legacy = crate::model::MetaObject::new();
821 legacy.insert("traceId".to_string(), json!("legacy"));
822 extensions.insert(legacy);
823
824 let req = CallToolRequest {
825 extensions,
826 method: Default::default(),
827 params: CallToolRequestParams {
828 meta: None,
829 name: "my_tool".into(),
830 arguments: None,
831 input_responses: None,
832 request_state: None,
833 },
834 };
835
836 let value = serde_json::to_value(&req).unwrap();
837 assert_eq!(value["params"]["_meta"]["traceId"], json!("legacy"));
838 }
839
840 #[test]
841 fn test_typed_meta_wins_over_legacy_extension_on_conflict() {
842 let mut extensions = Extensions::new();
843 let mut legacy = crate::model::MetaObject::new();
844 legacy.insert("shared".to_string(), json!("legacy"));
845 legacy.insert("legacy_only".to_string(), json!("kept"));
846 extensions.insert(legacy);
847 let mut typed = RequestMetaObject::new();
848 typed.insert("shared".to_string(), json!("typed"));
849 extensions.insert(typed);
850
851 let req = CallToolRequest {
852 extensions,
853 method: Default::default(),
854 params: CallToolRequestParams {
855 meta: None,
856 name: "my_tool".into(),
857 arguments: None,
858 input_responses: None,
859 request_state: None,
860 },
861 };
862
863 let value = serde_json::to_value(&req).unwrap();
864 let meta = value["params"]["_meta"].as_object().unwrap();
865 assert_eq!(meta.get("shared").unwrap(), "typed");
866 assert_eq!(meta.get("legacy_only").unwrap(), "kept");
867 }
868
869 #[test]
870 fn test_arbitrary_meta_keys_round_trip_unchanged() {
871 let input = json!({
872 "method": "tools/call",
873 "params": {
874 "_meta": {
875 "progressToken": 5,
876 "vendor.example/custom": {"nested": ["a", 1, null]},
877 "another-key": true
878 },
879 "name": "my_tool"
880 }
881 });
882 let req: CallToolRequest = serde_json::from_value(input.clone()).unwrap();
883 let output = serde_json::to_value(&req).unwrap();
884 assert_eq!(input, output);
885 }
886
887 mod json_float {
889 use std::fmt::Debug;
890
891 use rstest::rstest;
892 use serde::{Deserialize, Serialize, de::DeserializeOwned};
893 use serde_json::{Value, json};
894
895 use crate::model::{
896 Annotations, CreateMessageRequestParams, ModelPreferences, NumberSchema,
897 ProgressNotificationParam,
898 };
899
900 #[derive(Deserialize)]
903 #[serde(untagged)]
904 enum Buffered<T> {
905 Inner(T),
906 }
907
908 fn decode<T>(text: &str) -> T
911 where
912 T: DeserializeOwned + Serialize + PartialEq + Debug,
913 {
914 let direct: T = serde_json::from_str(text).unwrap();
915 let Buffered::Inner(buffered) = serde_json::from_str(text).unwrap();
916 assert_eq!(direct, buffered, "{text}");
917 let encoded = serde_json::to_string(&direct).unwrap();
918 assert_eq!(
919 serde_json::from_str::<T>(&encoded).unwrap(),
920 direct,
921 "{encoded}"
922 );
923 direct
924 }
925
926 fn rejects<T: DeserializeOwned>(text: &str) {
927 assert!(serde_json::from_str::<T>(text).is_err(), "{text}");
928 assert!(serde_json::from_str::<Buffered<T>>(text).is_err(), "{text}");
929 }
930
931 #[rstest]
932 #[case::decimal("0.6", 0.6)]
933 #[case::trailing_zero("0.60", 0.6)]
934 #[case::exponent("6e-1", 0.6)]
935 #[case::integer("1", 1.0)]
936 #[case::integral_decimal("1.0", 1.0)]
937 #[case::zero("0", 0.0)]
938 fn fields_read_every_spelling(#[case] n: &str, #[case] expected: f64) {
939 let narrowed = Some(expected as f32);
941
942 let annotations: Annotations = decode(&format!(r#"{{"priority":{n}}}"#));
943 assert_eq!(annotations.priority, narrowed);
944
945 let progress: ProgressNotificationParam = decode(&format!(
946 r#"{{"progressToken":1,"progress":{n},"total":{n}}}"#
947 ));
948 assert_eq!(progress.progress, expected);
949 assert_eq!(progress.total, Some(expected));
950
951 let params: CreateMessageRequestParams = decode(&format!(
952 r#"{{"messages":[],"maxTokens":1,"temperature":{n}}}"#
953 ));
954 assert_eq!(params.temperature, narrowed);
955
956 let preferences: ModelPreferences = decode(&format!(
957 r#"{{"costPriority":{n},"speedPriority":{n},"intelligencePriority":{n}}}"#
958 ));
959 assert_eq!(preferences.cost_priority, narrowed);
960 assert_eq!(preferences.speed_priority, narrowed);
961 assert_eq!(preferences.intelligence_priority, narrowed);
962
963 let schema: NumberSchema = decode(&format!(
964 r#"{{"type":"number","minimum":{n},"maximum":{n},"default":{n}}}"#
965 ));
966 assert_eq!(schema.minimum, Some(expected));
967 assert_eq!(schema.maximum, Some(expected));
968 assert_eq!(schema.default, Some(expected));
969 }
970
971 #[test]
972 fn decimal_narrows_to_the_nearest_f32() {
973 let annotations: Annotations = decode(r#"{"priority":0.6}"#);
974 assert_eq!(annotations.priority, Some(0.6_f32));
975 }
976
977 #[rstest]
978 #[case::missing(None)]
979 #[case::null(Some(Value::Null))]
980 fn optional_fields_read_absent_as_none(#[case] value: Option<Value>) {
981 let with = |mut base: Value, keys: &[&str]| {
983 if let Some(value) = &value {
984 for key in keys {
985 base[*key] = value.clone();
986 }
987 }
988 base.to_string()
989 };
990
991 let annotations: Annotations = decode(&with(json!({}), &["priority"]));
992 assert_eq!(annotations, Annotations::default());
993
994 let progress: ProgressNotificationParam = decode(&with(
995 json!({"progressToken": 1, "progress": 0}),
996 &["total"],
997 ));
998 assert_eq!(progress.total, None);
999
1000 let params: CreateMessageRequestParams = decode(&with(
1001 json!({"messages": [], "maxTokens": 1}),
1002 &["temperature"],
1003 ));
1004 assert_eq!(params.temperature, None);
1005
1006 let preferences: ModelPreferences = decode(&with(
1007 json!({}),
1008 &["costPriority", "speedPriority", "intelligencePriority"],
1009 ));
1010 assert_eq!(preferences, ModelPreferences::new());
1011
1012 let schema: NumberSchema = decode(&with(
1013 json!({"type": "number"}),
1014 &["minimum", "maximum", "default"],
1015 ));
1016 assert_eq!(schema, NumberSchema::new());
1017 }
1018
1019 #[test]
1020 fn progress_is_required() {
1021 rejects::<ProgressNotificationParam>(r#"{"progressToken":1}"#);
1022 rejects::<ProgressNotificationParam>(r#"{"progressToken":1,"progress":null}"#);
1023 }
1024
1025 #[rstest]
1026 #[case::string(r#""high""#)]
1027 #[case::out_of_range("1e400")]
1028 fn fields_reject_anything_but_a_finite_number(#[case] n: &str) {
1029 rejects::<Annotations>(&format!(r#"{{"priority":{n}}}"#));
1030 rejects::<ProgressNotificationParam>(&format!(
1031 r#"{{"progressToken":1,"progress":{n}}}"#
1032 ));
1033 rejects::<ProgressNotificationParam>(&format!(
1034 r#"{{"progressToken":1,"progress":0,"total":{n}}}"#
1035 ));
1036 rejects::<CreateMessageRequestParams>(&format!(
1037 r#"{{"messages":[],"maxTokens":1,"temperature":{n}}}"#
1038 ));
1039 for key in ["costPriority", "speedPriority", "intelligencePriority"] {
1040 rejects::<ModelPreferences>(&format!(r#"{{"{key}":{n}}}"#));
1041 }
1042 for key in ["minimum", "maximum", "default"] {
1043 rejects::<NumberSchema>(&format!(r#"{{"type":"number","{key}":{n}}}"#));
1044 }
1045 }
1046 }
1047}