use std::{fmt::Write, fs, path::Path};
use heck::{ToSnakeCase, ToUpperCamelCase};
use protox::prost_reflect::{DescriptorPool, FieldDescriptor, Kind, MessageDescriptor};
const HEADER: &str = "\
// @generated by `make generate-proto` from proto/monty/v1/monty.proto — DO NOT EDIT.
#![allow(clippy::allow_attributes, clippy::pedantic, clippy::use_self, clippy::absolute_paths, missing_docs)]
";
const ORACLE_HEADER: &str = "\
// @generated by `make generate-proto` from proto/monty/v1/monty.proto — DO NOT EDIT.
#![allow(clippy::allow_attributes, clippy::pedantic, clippy::use_self, clippy::absolute_paths, missing_docs, dead_code)]
";
fn main() {
let manifest_dir = Path::new(env!("CARGO_MANIFEST_DIR"));
let proto_dir = manifest_dir.join("proto");
let proto_file = proto_dir.join("monty/v1/monty.proto");
let descriptors = protox::compile([&proto_file], [&proto_dir]).expect("failed to compile monty.proto");
let pool = DescriptorPool::from_file_descriptor_set(descriptors.clone()).expect("invalid schema");
let out_dir = manifest_dir.join("src/generated");
prost_build::Config::new()
.out_dir(&out_dir)
.prost_path("crate::budgeted_prost")
.extern_path(".monty.v1.Arena", "crate::WireArena")
.extern_path(".monty.v1.FunctionCall", "crate::WireFunctionCall")
.extern_path(".monty.v1.Indexes", "crate::WireIndexes")
.extern_path(".monty.v1.NodePairs", "crate::WireNodePairs")
.extern_path(".monty.v1.NamedTupleNode", "crate::WireNamedTuple")
.compile_fds(descriptors.clone())
.expect("failed to generate Rust code from monty.proto");
let generated = out_dir.join("monty.v1.rs");
check_allocation_forms(&fs::read_to_string(&generated).expect("generated file missing"));
prepend_header(&generated, HEADER);
let oracle_dir = manifest_dir.join("tests/oracle");
fs::create_dir_all(&oracle_dir).expect("failed to create tests/oracle");
prost_build::Config::new()
.out_dir(&oracle_dir)
.compile_fds(descriptors)
.expect("failed to generate oracle Rust code from monty.proto");
prepend_header(&oracle_dir.join("monty.v1.rs"), ORACLE_HEADER);
generate_repeated_tests(&pool, &oracle_dir.join("repeated_fields.rs"));
}
fn generate_repeated_tests(pool: &DescriptorPool, path: &Path) {
let mut source = String::from(
"// @generated by `make generate-proto` — DO NOT EDIT.\n\n/// Exercises every repeated field in the schema.\n#[test]\nfn every_schema_repeated_field_is_budgeted() {\n",
);
for message in pool.all_messages() {
assert!(!message.is_map_entry(), "protobuf maps need a budgeted runtime adapter");
for field in message.fields().filter(FieldDescriptor::is_list) {
let (wire_type, payload) = match field.kind() {
Kind::Message(element) => {
let payload: &[u8] = match element.full_name() {
"monty.v1.MontyNode" => &[0x12, 0],
_ => &[],
};
("LengthDelimited", payload)
}
Kind::Enum(_) => panic!(
"{}: prost's repeated-enum accessors require infallible push; extend budgeted_prost first",
field.full_name()
),
Kind::String | Kind::Bytes => ("LengthDelimited", &[][..]),
Kind::Float | Kind::Fixed32 | Kind::Sfixed32 => ("ThirtyTwoBit", &[0; 4][..]),
Kind::Double | Kind::Fixed64 | Kind::Sfixed64 => ("SixtyFourBit", &[0; 8][..]),
_ => ("Varint", &[0][..]),
};
writeln!(
source,
" check_repeated(\n \"{}\",\n {},\n WireType::{wire_type},\n &{payload:?},\n |message: &{}| &message.{},\n );",
field.full_name(),
field.number(),
message_rust_path(&message),
if matches!(message.full_name(), "monty.v1.Arena" | "monty.v1.Indexes" | "monty.v1.NodePairs") {
"0".to_owned()
} else {
field.name().to_snake_case()
},
)
.expect("write to string");
}
}
source.push_str("}\n");
fs::write(path, source).expect("failed to write repeated field fixtures");
}
fn message_rust_path(message: &MessageDescriptor) -> String {
match message.full_name() {
"monty.v1.Arena" => "WireArena".to_owned(),
"monty.v1.FunctionCall" => "WireFunctionCall".to_owned(),
"monty.v1.Indexes" => "WireIndexes".to_owned(),
"monty.v1.NodePairs" => "WireNodePairs".to_owned(),
"monty.v1.NamedTupleNode" => "WireNamedTuple".to_owned(),
_ => {
let mut parents = Vec::new();
let mut parent = message.parent_message();
while let Some(message) = parent {
parents.push(message.name().to_snake_case());
parent = message.parent_message();
}
parents.reverse();
parents.insert(0, "pb".to_owned());
parents.push(message.name().to_upper_camel_case());
parents.join("::")
}
}
}
fn check_allocation_forms(source: &str) {
for unsupported in [
"::boxed::",
"::collections::",
"::bytes::Bytes",
"#[prost(group",
"#[prost(map",
"#[prost(btree_map",
] {
assert!(
!source.contains(unsupported),
"unbudgeted protobuf allocation form {unsupported}: extend budgeted_prost before adding this field"
);
}
}
fn prepend_header(generated: &Path, header: &str) {
let body = fs::read_to_string(generated).expect("generated file missing");
fs::write(generated, format!("{header}{body}")).expect("failed to write generated file");
println!("regenerated {}", generated.display());
}