1use prost_reflect::DynamicMessage;
8
9use crate::{Error, Model, Result};
10
11pub fn to_json(model: &Model) -> Result<String> {
13 let dynamic = crate::proto_serde::to_dynamic(&model.to_proto()?)?;
14 serde_json::to_string_pretty(&dynamic).map_err(json_error)
15}
16
17pub fn from_json(source: &str) -> Result<Model> {
19 let mut deserializer = serde_json::Deserializer::from_str(source);
20 let dynamic = DynamicMessage::deserialize(crate::proto_serde::descriptor(), &mut deserializer)
21 .map_err(json_error)?;
22 deserializer.end().map_err(json_error)?;
23 Model::from_proto(crate::proto_serde::from_dynamic(&dynamic)?)
24}
25
26fn json_error(error: impl std::fmt::Display) -> Error {
27 Error::Json(error.to_string())
28}
29
30#[cfg(test)]
31mod tests {
32 use onnx_runtime_ir::{DataType, Graph, Node, NodeId, static_shape};
33
34 use super::*;
35
36 #[test]
37 fn simple_model_round_trips() {
38 let mut graph = Graph::new();
39 graph.opset_imports.insert(String::new(), 21);
40 let input = graph.create_named_value("X", DataType::Float32, static_shape([2, 3]));
41 let output = graph.create_named_value("Y", DataType::Float32, static_shape([2, 3]));
42 graph.add_input(input);
43 graph.add_output(output);
44 graph.insert_node(Node::new(
45 NodeId(0),
46 "Identity",
47 vec![Some(input)],
48 vec![output],
49 ));
50
51 let json = to_json(&Model::new(graph)).unwrap();
52 let decoded = from_json(&json).unwrap();
53 assert_eq!(to_json(&decoded).unwrap(), json);
54 assert_eq!(decoded.graph.num_nodes(), 1);
55 }
56
57 #[test]
58 fn rejects_unknown_fields() {
59 let error = match from_json(r#"{"unknownOnnxField": true}"#) {
60 Ok(_) => panic!("unknown field must be rejected"),
61 Err(error) => error,
62 };
63 assert!(error.to_string().contains("unknownOnnxField"));
64 }
65}