Skip to main content

onnx_std/json/
mod.rs

1//! Canonical descriptor-driven ONNX protobuf JSON interchange.
2//!
3//! The generated descriptor for the crate's vendored `onnx.proto3` is the only
4//! field table. This keeps JSON and protobuf TextFormat automatically complete
5//! and consistent for every message, field, oneof, and enum in the bound spec.
6
7use prost_reflect::DynamicMessage;
8
9use crate::{Error, Model, Result};
10
11/// Serialize a model using protobuf's canonical JSON mapping.
12pub 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
17/// Parse a canonical ONNX protobuf JSON document.
18pub 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}