use prost::Message;
use crate::LoaderError;
#[allow(clippy::all, missing_docs, non_snake_case)]
pub mod onnx {
include!(concat!(env!("OUT_DIR"), "/onnx.rs"));
}
pub use onnx::ModelProto;
pub const FILE_DESCRIPTOR_SET: &[u8] =
include_bytes!(concat!(env!("OUT_DIR"), "/onnx_descriptor.bin"));
pub fn decode_model(bytes: &[u8]) -> Result<ModelProto, LoaderError> {
ModelProto::decode(bytes).map_err(|e| LoaderError::ProtobufParse(e.to_string()))
}
pub fn textproto_to_binary(text: &str) -> Result<Vec<u8>, LoaderError> {
use prost_reflect::{DescriptorPool, DynamicMessage};
use std::sync::OnceLock;
static POOL: OnceLock<DescriptorPool> = OnceLock::new();
let pool = POOL.get_or_init(|| {
DescriptorPool::decode(FILE_DESCRIPTOR_SET)
.expect("the generated ONNX descriptor set must be valid")
});
let descriptor = pool
.get_message_by_name("onnx.ModelProto")
.expect("the ONNX descriptor set must define onnx.ModelProto");
let dynamic = DynamicMessage::parse_text_format(descriptor, text)
.map_err(|e| LoaderError::TextProtoParse(e.to_string()))?;
Ok(dynamic.encode_to_vec())
}