use std::collections::HashSet;
use anyhow::Context;
use prost::Message;
use prost_reflect::{DescriptorPool, DynamicMessage};
use prost_types::FileDescriptorSet;
use substrait::proto::{FILE_DESCRIPTOR_SET, Plan};
pub fn build_descriptor_pool(extra_descriptors: &[&[u8]]) -> anyhow::Result<DescriptorPool> {
let mut fds = FileDescriptorSet::decode(FILE_DESCRIPTOR_SET)
.context("failed to decode substrait core descriptor")?;
let mut seen: HashSet<String> = fds.file.iter().map(|f| f.name().to_owned()).collect();
for blob in extra_descriptors {
let extra =
FileDescriptorSet::decode(*blob).context("failed to decode extra descriptor")?;
for f in extra.file {
if seen.insert(f.name().to_owned()) {
fds.file.push(f);
}
}
}
DescriptorPool::decode(fds.encode_to_vec().as_slice())
.context("failed to build descriptor pool")
}
pub fn parse_json(json: &str, pool: &DescriptorPool) -> anyhow::Result<Plan> {
if let Ok(plan) = serde_json::from_str::<Plan>(json) {
return Ok(plan);
}
let plan_desc = pool
.get_message_by_name("substrait.Plan")
.context("substrait.Plan not found in descriptor pool")?;
let dyn_msg =
DynamicMessage::deserialize(plan_desc, &mut serde_json::Deserializer::from_str(json))
.context("failed to parse JSON as substrait.Plan")?;
Plan::decode(dyn_msg.encode_to_vec().as_slice())
.context("failed to decode Plan from dynamic message bytes")
}