diff --git a/generator/src/codegen/rust/codegen_source.rs b/generator/src/codegen/rust/codegen_source.rs index bc78f05..ecd5d60 100644 --- a/generator/src/codegen/rust/codegen_source.rs +++ b/generator/src/codegen/rust/codegen_source.rs @@ -4,16 +4,31 @@ use ::heck::*; use std::collections::HashMap; use flat_ast::RestrictionContent::{Enumeration, Length, MaxValue, MinValue}; +use crate::type_registry::TypeRegistry; + pub (crate) struct CodeSourceGenerator<'a, W: Write + 'a> { writer: &'a mut ::writer::Writer, - version: String + version: String, + registry: &'a TypeRegistry, + shared_types_path: String, } impl<'a, W: Write> CodeSourceGenerator<'a, W> { - pub fn new(writer: &'a mut ::writer::Writer, version: String) -> Self { + pub fn new(writer: &'a mut ::writer::Writer, version: String, registry: &'a TypeRegistry, shared_types_path: String) -> Self { + Self { + writer, + version, + registry, + shared_types_path, + } + } + + pub fn new_shared(writer: &'a mut ::writer::Writer, version: String, registry: &'a TypeRegistry) -> Self { Self { writer, - version + version, + registry, + shared_types_path: String::new(), } } @@ -37,6 +52,7 @@ impl<'a, W: Write> CodeSourceGenerator<'a, W> { cg!(self, r#"use bincode::de::read::Reader;"#); cg!(self, r#"use bincode::enc::write::Writer;"#); cg!(self, r#"use utils::null_string::NullTerminatedString;"#); + cg!(self, r#"use {}::*;"#, self.shared_types_path); cg!(self, r#"use crate::enums::*;"#); cg!(self, r#"use crate::types::*;"#); cg!(self, r#"use crate::dataconsts::*;"#); @@ -62,6 +78,7 @@ impl<'a, W: Write> CodeSourceGenerator<'a, W> { for content in packet.contents() { use self::PacketContent::*; match content { + Simple(simple) if self.registry.types.contains_key(simple.name()) => continue, Simple(simple) => self.simple_type(simple, &iserialize)?, _ => {} } @@ -70,6 +87,7 @@ impl<'a, W: Write> CodeSourceGenerator<'a, W> { for content in packet.contents() { use self::PacketContent::*; match content { + Complex(ref complex) if self.registry.types.contains_key(complex.name()) => continue, Complex(ref complex) => self.complex_type(complex, &iserialize)?, _ => {} }; @@ -221,6 +239,41 @@ impl<'a, W: Write> CodeSourceGenerator<'a, W> { Ok(()) } + pub fn generate_shared(&mut self) -> Result<()> { + cg!(self, "/* This file is @generated with IDL v{} */\n", self.version); + cg!(self, r#"use bincode::{{Encode, Decode, enc::Encoder, de::Decoder, error::DecodeError}};"#); + cg!(self, r#"use bincode::de::read::Reader;"#); + cg!(self, r#"use bincode::enc::write::Writer;"#); + cg!(self, r#"use utils::null_string::NullTerminatedString;"#); + cg!(self); + + let mut iserialize: HashMap = HashMap::new(); + iserialize.insert("int8_t".to_string(), "i8".to_string()); + iserialize.insert("uint8_t".to_string(), "u8".to_string()); + iserialize.insert("int16_t".to_string(), "i16".to_string()); + iserialize.insert("uint16_t".to_string(), "u16".to_string()); + iserialize.insert("int32_t".to_string(), "i32".to_string()); + iserialize.insert("uint32_t".to_string(), "u32".to_string()); + iserialize.insert("int64_t".to_string(), "i64".to_string()); + iserialize.insert("uint64_t".to_string(), "u64".to_string()); + iserialize.insert("char".to_string(), "u8".to_string()); + iserialize.insert("int".to_string(), "i32".to_string()); + iserialize.insert("unsigned int".to_string(), "u32".to_string()); + iserialize.insert("float".to_string(), "f32".to_string()); + iserialize.insert("double".to_string(), "f64".to_string()); + iserialize.insert("std::string".to_string(), "NullTerminatedString".to_string()); + + for (name, content) in &self.registry.types { + match content { + PacketContent::Simple(simple) => self.simple_type(simple, &iserialize)?, + PacketContent::Complex(complex) => self.complex_type(complex, &iserialize)?, + _ => {} + } + cg!(self); + } + Ok(()) + } + fn doc(&mut self, doc: &Option) -> Result<()> { match doc { None => (), @@ -258,16 +311,49 @@ impl<'a, W: Write> CodeSourceGenerator<'a, W> { } } + fn has_bits(&self, complex: &ComplexType) -> bool { + match complex.content() { + ComplexTypeContent::Seq(ref s) => s.elements().iter().any(|e| e.bits().is_some()), + _ => false, + } + } + fn complex_type(&mut self, complex: &ComplexType, iserialize: &HashMap) -> Result<()> { use ::flat_ast::ComplexTypeContent::*; if complex.inline() == false { + let mut bitfield_type = None; + if self.has_bits(complex) { + let mut total_bits = 0; + if let Seq(ref s) = complex.content() { + for e in s.elements() { + total_bits += e.bits().unwrap_or_else(|| { + let rust_type = iserialize.get(e.type_().trim()).map(|s| s.to_string()).unwrap_or_else(|| { + e.type_().trim().to_string() + }); + match rust_type.as_str() { + "u8" | "i8" => 8, + "u16" | "i16" => 16, + "u32" | "i32" | "f32" => 32, + "u64" | "i64" | "f64" => 64, + _ => 8, + } + }); + } + } + bitfield_type = Some(Self::get_bitfield_type(total_bits as usize)); + } + // All unions need to be outside the struct match complex.content() { Choice(ref c) => { for elem in c.elements() { if let Some(ref seq) = c.inline_seqs().get(elem.name()) { - cg!(self, r#"#[derive(Debug)]"#); - cg!(self, "struct {} {{", elem.name()); + let mut derives = "Debug, Clone, Default".to_string(); + if *self.registry.is_copy.get(elem.name()).unwrap_or(&false) { + derives += ", Copy"; + } + cg!(self, r#"#[derive({})]"#, derives); + cg!(self, "pub struct {} {{", elem.name()); self.indent(); for e in seq.elements() { self.element(e, &iserialize)?; @@ -317,8 +403,23 @@ impl<'a, W: Write> CodeSourceGenerator<'a, W> { let mut variable_names = Vec::new(); for e in seq.elements() { let name = rename_if_reserved(e.name()); - let bits = e.bits().map_or_else(|| "".to_string(), |b| format!("{}", b)); - cg!(self, "let {}_bits = Self::encode_bitfield(self.{}, {}, &mut offset);", e.name().to_snake_case(), name, bits); + let rust_type_encoding = iserialize.get(e.type_().trim()).map(|s| s.to_string()).unwrap_or_else(|| { + e.type_().trim().to_string() + }); + let bits = e.bits().map_or_else(|| { + let size = match rust_type_encoding.as_str() { + "u8" | "i8" => 8, + "u16" | "i16" => 16, + "u32" | "i32" | "f32" => 32, + "u64" | "i64" | "f64" => 64, + _ => { + debug!("Unknown type for bitfield size calculation: {}, defaulting to 8", rust_type); + 8 + } + }; + size.to_string() + }, |b| format!("{}", b)); + cg!(self, "let {}_bits = Self::encode_bitfield(self.{} as {}, {}, &mut offset);", e.name().to_snake_case(), name, rust_type, bits); variable_names.push(format!("{}_bits", e.name().to_snake_case())); } cg!(self, "{}", variable_names.join(" | ")); @@ -349,8 +450,23 @@ impl<'a, W: Write> CodeSourceGenerator<'a, W> { let mut variable_names = Vec::new(); for e in seq.elements() { let name = rename_if_reserved(e.name()); - let bits = e.bits().map_or_else(|| "".to_string(), |b| format!("{}", b)); - cg!(self, "let {} = Self::decode_bitfield(bitfield, {}, &mut offset);", name, bits); + let rust_type = iserialize.get(e.type_().trim()).map(|s| s.to_string()).unwrap_or_else(|| { + e.type_().trim().to_string() + }); + let bits = e.bits().map_or_else(|| { + let size = match rust_type.as_str() { + "u8" | "i8" => 8, + "u16" | "i16" => 16, + "u32" | "i32" | "f32" => 32, + "u64" | "i64" | "f64" => 64, + _ => { + debug!("Unknown type for bitfield size calculation: {}, defaulting to 8", rust_type); + 8 + } + }; + size.to_string() + }, |b| format!("{}", b)); + cg!(self, "let {} = Self::decode_bitfield(bitfield, {}, &mut offset) as {};", name, bits, rust_type); variable_names.push(format!("{}", name)); } cg!(self, "Ok(Self {{ {} }})", variable_names.join(", ")); @@ -365,29 +481,94 @@ impl<'a, W: Write> CodeSourceGenerator<'a, W> { } cg!(self); - cg!(self, r#"#[derive(Debug, Clone, Default)]"#); + let mut derives = "Debug, Clone, Default".to_string(); + if *self.registry.is_copy.get(complex.name()).unwrap_or(&false) { + derives += ", Copy"; + } + cg!(self, r#"#[derive({})]"#, derives); cg!(self, "pub struct {} {{", complex.name()); self.indent(); - match complex.content() { - Seq(ref s) => { + + if let Some(rust_type) = bitfield_type { + if let Seq(ref s) = complex.content() { for elem in s.elements() { self.element(elem, &iserialize)?; } - }, - Choice(ref c) => { - for elem in c.elements() { - if let Some(ref _seq) = c.inline_seqs().get(elem.name()) { - cg!(self, "{}: {},", elem.name().to_snake_case(), elem.name()); - } else { + } + } else { + match complex.content() { + Seq(ref s) => { + for elem in s.elements() { self.element(elem, &iserialize)?; } - } - }, - Empty => {} + }, + Choice(ref c) => { + for elem in c.elements() { + if let Some(ref _seq) = c.inline_seqs().get(elem.name()) { + cg!(self, "pub {}: {},", elem.name().to_snake_case(), elem.name()); + } else { + self.element(elem, &iserialize)?; + } + } + }, + Empty => {} + } } self.dedent(); cg!(self, "}}"); cg!(self); + + if let Some(rust_type) = bitfield_type { + cg!(self, "impl {} {{", complex.name()); + self.indent(); + cg!(self, "fn encode_bitfield(value: {}, size: {}, offset: &mut {}) -> {} {{", rust_type, rust_type, rust_type, rust_type); + self.indent(); + cg!(self, "let encoded = (value & ((1 << size) - 1)) << *offset;"); + cg!(self, "*offset += size; // Update offset for the next field"); + cg!(self, "encoded"); + self.dedent(); + cg!(self, "}}"); + cg!(self); + cg!(self, "fn decode_bitfield(encoded: {}, size: {}, offset: &mut {}) -> {} {{", rust_type, rust_type, rust_type, rust_type); + self.indent(); + cg!(self, "let value = (encoded >> *offset) & ((1 << size) - 1);"); + cg!(self, "*offset += size; // Update offset for the next field"); + cg!(self, "value"); + self.dedent(); + cg!(self, "}}"); + cg!(self); + cg!(self, "pub fn encode_data(&self) -> {} {{", rust_type); + self.indent(); + cg!(self, "let mut offset = 0;"); + let mut variable_names = Vec::new(); + if let Seq(ref s) = complex.content() { + for e in s.elements() { + let name = rename_if_reserved(e.name()); + let rust_type_encoding = iserialize.get(e.type_().trim()).map(|s| s.to_string()).unwrap_or_else(|| { + e.type_().trim().to_string() + }); + let bits = e.bits().map_or_else(|| { + let size = match rust_type_encoding.as_str() { + "u8" | "i8" => 8, + "u16" | "i16" => 16, + "u32" | "i32" | "f32" => 32, + "u64" | "i64" | "f64" => 64, + _ => 8, + }; + size.to_string() + }, |b| format!("{}", b)); + cg!(self, "let {}_bits = Self::encode_bitfield(self.{} as {}, {}, &mut offset);", e.name().to_snake_case(), name, rust_type, bits); + variable_names.push(format!("{}_bits", e.name().to_snake_case())); + } + } + cg!(self, "{}", variable_names.join(" | ")); + self.dedent(); + cg!(self, "}}"); + self.dedent(); + cg!(self, "}}"); + cg!(self); + } + let _ = self.complex_encode(complex, iserialize); cg!(self); let _ = self.complex_decode(complex, iserialize); @@ -395,27 +576,31 @@ impl<'a, W: Write> CodeSourceGenerator<'a, W> { Ok(()) } - fn complex_encode(&mut self, complex: &ComplexType, _iserialize: &HashMap) -> Result<()> { + fn complex_encode(&mut self, complex: &ComplexType, iserialize: &HashMap) -> Result<()> { use ::flat_ast::ComplexTypeContent::*; cg!(self, "impl Encode for {} {{", complex.name()); self.indent(); cg!(self, "fn encode(&self, encoder: &mut E) -> std::result::Result<(), bincode::error::EncodeError> {{"); self.indent(); - match complex.content() { - Seq(ref s) => { - for elem in s.elements() { - let name = rename_if_reserved(elem.name()); - cg!(self, "self.{}.encode(encoder)?;", name); - } - }, - Choice(ref c) => { - for elem in c.elements() { - let name = rename_if_reserved(elem.name()); - cg!(self, "self.{}.encode(encoder)?;", name); - } - }, - Empty => {} + if self.has_bits(complex) { + cg!(self, "self.encode_data().encode(encoder)?;"); + } else { + match complex.content() { + Seq(ref s) => { + for elem in s.elements() { + let name = rename_if_reserved(elem.name()); + cg!(self, "self.{}.encode(encoder)?;", name); + } + }, + Choice(ref c) => { + for elem in c.elements() { + let name = rename_if_reserved(elem.name()); + cg!(self, "self.{}.encode(encoder)?;", name); + } + }, + Empty => {} + } } cg!(self, "Ok(())"); @@ -433,81 +618,70 @@ impl<'a, W: Write> CodeSourceGenerator<'a, W> { cg!(self, "fn decode(decoder: &mut D) -> std::result::Result {{"); self.indent(); - let mut output_list = Vec::new(); - match complex.content() { - Seq(ref s) => { - for elem in s.elements() { - let name = rename_if_reserved(elem.name()); - let trimmed_type = elem.type_().trim().to_string(); - let mut is_rust_native = true; - let rust_type = iserialize.get(elem.type_().trim()).map(|s| s.to_string()).unwrap_or_else(|| { - debug!(r#"Type "{}" not found, outputting anyway"#, elem.type_()); - is_rust_native = false; - trimmed_type.clone() + if self.has_bits(complex) { + let mut total_bits = 0; + if let Seq(ref s) = complex.content() { + for e in s.elements() { + total_bits += e.bits().unwrap_or_else(|| { + let rust_type = iserialize.get(e.type_().trim()).map(|s| s.to_string()).unwrap_or_else(|| { + e.type_().trim().to_string() + }); + match rust_type.as_str() { + "u8" | "i8" => 8, + "u16" | "i16" => 16, + "u32" | "i32" | "f32" => 32, + "u64" | "i64" | "f64" => 64, + _ => 8, + } }); - - if let Some(ref o) = elem.occurs() { - use ::flat_ast::Occurs::*; - match o { - Unbounded => { - cg!(self, "let {} = Vec::decode(decoder)?;", name); - } - Num(n) => { - let mut type_prefix = "0"; - if "String" == rust_type { - type_prefix = ""; - } - - if false == is_rust_native { - type_prefix = ""; - if n.parse::().is_ok() { - cg!(self, "let mut {}: [{}{}; {}] = core::array::from_fn(|i| {}{}::default());", name, type_prefix, rust_type, n, type_prefix, rust_type); - } else { - cg!(self, "let mut {}: [{}{}; ({} as usize)] = core::array::from_fn(|i| {}{}::default());", name, type_prefix, rust_type, n, type_prefix, rust_type); - } - - cg!(self, "for index in 0..{} as usize {{", n); - self.indent(); - cg!(self, "{}[index] = {}::decode(decoder)?;", name, rust_type); - self.dedent(); - cg!(self, "}}"); - } else { - if n.parse::().is_ok() { - cg!(self, "let mut {} = [{}{}; {}];", name, type_prefix, rust_type, n); - } else { - cg!(self, "let mut {} = [{}{}; ({} as usize)];", name, type_prefix, rust_type, n); - } - cg!(self, "for value in &mut {} {{", name); - self.indent(); - cg!(self, "*value = {}::decode(decoder)?;", rust_type); - self.dedent(); - cg!(self, "}}"); - } - } - }; - } else { - cg!(self, "let {} = {}::decode(decoder)?;", name, rust_type); - } - output_list.push(name); } - }, - Choice(ref c) => { - for elem in c.elements() { - let name = rename_if_reserved(elem.name()); - let trimmed_type = elem.type_().trim().to_string(); - let rust_type = iserialize.get(elem.type_().trim()).map(|s| s.to_string()).unwrap_or_else(|| { - debug!(r#"Type "{}" not found, outputting anyway"#, elem.type_()); - trimmed_type.clone() + } + let rust_type = Self::get_bitfield_type(total_bits as usize); + cg!(self, "let bitfield = {}::decode(decoder)?;", rust_type); + cg!(self, "let mut offset = 0;"); + let mut variable_names = Vec::new(); + if let Seq(ref s) = complex.content() { + for e in s.elements() { + let name = rename_if_reserved(e.name()); + let rust_type_field = iserialize.get(e.type_().trim()).map(|s| s.to_string()).unwrap_or_else(|| { + e.type_().trim().to_string() }); - - cg!(self, "let {} = {}::decode(decoder)?;", name, rust_type); - output_list.push(name); + let bits = e.bits().map_or_else(|| { + let size = match rust_type_field.as_str() { + "u8" | "i8" => 8, + "u16" | "i16" => 16, + "u32" | "i32" | "f32" => 32, + "u64" | "i64" | "f64" => 64, + _ => 8, + }; + size.to_string() + }, |b| format!("{}", b)); + cg!(self, "let {} = Self::decode_bitfield(bitfield, {}, &mut offset) as {};", name, bits, rust_type_field); + variable_names.push(format!("{}", name)); } - }, - Empty => {} + } + cg!(self, "Ok(Self {{ {} }})", variable_names.join(", ")); + } else { + let mut variable_names = Vec::new(); + match complex.content() { + Seq(ref s) => { + for elem in s.elements() { + let name = rename_if_reserved(elem.name()); + cg!(self, "let {} = Decode::decode(decoder)?;", name); + variable_names.push(name); + } + }, + Choice(ref c) => { + for elem in c.elements() { + let name = rename_if_reserved(elem.name()); + cg!(self, "let {} = Decode::decode(decoder)?;", name); + variable_names.push(name); + } + }, + Empty => {} + } + cg!(self, "Ok(Self {{ {} }})", variable_names.join(", ")); } - cg!(self, "Ok(Self {{ {} }})", output_list.join(", ")); - self.dedent(); cg!(self, "}}"); self.dedent(); @@ -531,10 +705,12 @@ impl<'a, W: Write> CodeSourceGenerator<'a, W> { self.doc(elem.doc())?; if let Some(bitset) = elem.bitset() { - if bitset.start == 0 { - cg!(self, "{}: [bool; {}],", bitset.name, bitset.size); + if bitset.size > 0 { + if bitset.start == 0 { + cg!(self, "{}: [bool; {}],", bitset.name, bitset.size); + } + return Ok(()); } - return Ok(()); } let trimmed_type = elem.type_().trim().to_string(); @@ -572,7 +748,7 @@ impl<'a, W: Write> CodeSourceGenerator<'a, W> { // }; let name = rename_if_reserved(elem.name()); // cg!(self, "{}: {}{}{},", elem.name(), type_, bits, default); - cg!(self, "pub(crate) {}: {}, {}", name, type_, bits); + cg!(self, "pub {}: {}, {}", name, type_, bits); Ok(()) } @@ -595,8 +771,12 @@ impl<'a, W: Write> CodeSourceGenerator<'a, W> { if is_enum { cg!(self, r#"#[repr({})]"#, rust_type); - cg!(self, r#"#[derive(Debug, Clone)]"#); - cg!(self, "pub(crate) enum {} {{", name.to_upper_camel_case()); + let mut derives = "Debug, Clone".to_string(); + if *self.registry.is_copy.get(name).unwrap_or(&false) { + derives += ", Copy"; + } + cg!(self, r#"#[derive({})]"#, derives); + cg!(self, "pub enum {} {{", name.to_upper_camel_case()); self.indent(); for content in restrict.contents() { if let Enumeration(en) = content { @@ -605,10 +785,16 @@ impl<'a, W: Write> CodeSourceGenerator<'a, W> { } } } else { - cg!(self, r#"#[derive(Debug)]"#); + let mut derives = "Debug".to_string(); + if *self.registry.is_copy.get(name).unwrap_or(&false) { + derives += ", Copy, Clone"; + } else { + derives += ", Clone"; + } + cg!(self, r#"#[derive({})]"#, derives); cg!(self, "pub struct {} {{", name.to_upper_camel_case()); self.indent(); - cg!(self, "pub(crate) {}: {},", name.to_string().to_snake_case(), rust_type); + cg!(self, "pub {}: {},", name.to_string().to_snake_case(), rust_type); } self.dedent(); diff --git a/generator/src/codegen/rust/mod.rs b/generator/src/codegen/rust/mod.rs index 32e0301..1ef2bef 100644 --- a/generator/src/codegen/rust/mod.rs +++ b/generator/src/codegen/rust/mod.rs @@ -2,27 +2,41 @@ use std::fs::File; use std::path::PathBuf; use codegen::Codegen; use ::{flat_ast, writer}; +use crate::type_registry::TypeRegistry; mod codegen_source; -pub struct Generator { - output: PathBuf +pub struct Generator<'a> { + output: PathBuf, + registry: &'a TypeRegistry, + shared_types_path: String, } -impl Generator { - pub fn new(args: &RustArgs) -> Self { +impl<'a> Generator<'a> { + pub fn new(args: &RustArgs, registry: &'a TypeRegistry) -> Self { Self{ - output: args.output_folder.clone().into() + output: args.output_folder.clone().into(), + registry, + shared_types_path: args.shared_types_path.clone(), } } } -impl Codegen for Generator { +pub fn generate_shared(version: &str, registry: &TypeRegistry, args: &RustArgs) -> Result<(), failure::Error> { + let output: PathBuf = args.output_folder.clone().into(); + let source_output = File::create(output.to_str().unwrap().to_owned() + "/shared_types.rs")?; + let mut writer = writer::Writer::new(source_output); + let mut codegen = codegen_source::CodeSourceGenerator::new_shared(&mut writer, version.to_string(), registry); + codegen.generate_shared()?; + Ok(()) +} + +impl<'a> Codegen for Generator<'a> { fn generate(&mut self, version: &str, packet: &flat_ast::Packet) -> Result<(), failure::Error> { let source_output = File::create(self.output.to_str().unwrap().to_owned() + &format!("/{}.rs", packet.filename()))?; debug!("source {:?}", source_output); let mut writer = writer::Writer::new(source_output); - let mut codegen = codegen_source::CodeSourceGenerator::new(&mut writer, version.to_string()); + let mut codegen = codegen_source::CodeSourceGenerator::new(&mut writer, version.to_string(), self.registry, self.shared_types_path.clone()); codegen.generate(&packet)?; Ok(()) } @@ -32,11 +46,15 @@ impl Codegen for Generator { #[command(name="rust")] pub struct RustArgs { #[arg(long)] - output_folder: String + output_folder: String, + + #[arg(long, default_value = "crate::shared_types")] + shared_types_path: String, } #[cfg(test)] mod tests { + use type_registry::TypeRegistry; use crate::{flat_ast::Packet, writer::Writer}; use super::{codegen_source}; @@ -70,7 +88,8 @@ mod tests { fn call_header(packet: &Packet) -> std::io::Result { let writer = StringWriter::new(); let mut writer = Writer::new(writer); - let mut codegen = codegen_source::CodeSourceGenerator::new(&mut writer, "0".to_string()); + let registry = TypeRegistry::new(); + let mut codegen = codegen_source::CodeSourceGenerator::new(&mut writer, "0".to_string(), ®istry, "crate::shared_types".to_string()); codegen.generate(packet)?; Ok(writer.into().into()) } diff --git a/generator/src/flat_ast.rs b/generator/src/flat_ast.rs index e547671..c6d1363 100644 --- a/generator/src/flat_ast.rs +++ b/generator/src/flat_ast.rs @@ -15,7 +15,7 @@ pub enum PacketContent { Complex(ComplexType) } -#[derive(Debug)] +#[derive(Debug, Clone)] pub struct ComplexType { name: String, content: ComplexTypeContent, @@ -24,7 +24,7 @@ pub struct ComplexType { inline: bool } -#[derive(Debug)] +#[derive(Debug, Clone)] pub enum ComplexTypeContent { Seq(Sequence), Choice(Choice), @@ -43,7 +43,7 @@ pub struct Sequence { inline: bool } -#[derive(Debug)] +#[derive(Debug, Clone)] pub struct Choice { elements: Vec, doc: Option, @@ -78,26 +78,26 @@ pub struct Bitset { pub name: String } -#[derive(Debug)] +#[derive(Debug, Clone)] pub struct SimpleType { name: String, contents: Vec, doc: Option } -#[derive(Debug)] +#[derive(Debug, Clone)] pub enum SimpleTypeContent { Restriction(Restriction) } -#[derive(Debug)] +#[derive(Debug, Clone)] pub struct Restriction { base: String, doc: Option, contents: Vec } -#[derive(Debug)] +#[derive(Debug, Clone)] pub enum RestrictionContent { Enumeration(Enumeration), Length(u32), @@ -105,7 +105,7 @@ pub enum RestrictionContent { MaxValue(String) } -#[derive(Debug)] +#[derive(Debug, Clone)] pub struct Enumeration { value: String, id: i64, diff --git a/generator/src/flatten.rs b/generator/src/flatten.rs index 58fd909..14a1cd4 100644 --- a/generator/src/flatten.rs +++ b/generator/src/flatten.rs @@ -218,19 +218,31 @@ fn flatten_enum(e: &ast::Enumeration, enum_id: &mut i64) -> flat_ast::Enumeratio fn flatten_complex(c: &ast::ComplexType, ctx: &mut Context) { use flat_ast::ComplexTypeContent::*; use self::ast::ComplexTypeContent; + + let mut path = ctx.path.clone(); + path.push(c.name().clone()); + let mut ctx2 = Context { + packet: ctx.packet, + path, + complex_types: ctx.complex_types.clone(), + is_in_choice: false, + bitsets: 0, + current_bitset: None, + }; + let mut inline = false; let content = match c.content() { - ComplexTypeContent::Choice(ref c) => Choice(flatten_choice(c, ctx)), + ComplexTypeContent::Choice(ref c) => Choice(flatten_choice(c, &mut ctx2)), ComplexTypeContent::Seq(ref s) => { - let seq = flatten_seq(s, ctx); + let seq = flatten_seq(s, &mut ctx2); inline = seq.inline(); Seq(seq) }, ComplexTypeContent::Empty => Empty }; + ctx2.stop_bits(); let cot = flat_ast::ComplexType::new(c.name().clone(), content, c.doc().clone(), false, inline); ctx.add_content(flat_ast::PacketContent::Complex(cot)); - ctx.stop_bits(); } fn flatten_anon_complex(c: &ast::AnonComplexType, ctx: &mut Context, element_name: &Option) -> flat_ast::ComplexType { @@ -357,7 +369,9 @@ fn flatten_element(elem: &ast::Element, ctx: &mut Context, id: u32) -> flat_ast: } }; let bitset = if let Some(bits) = elem.bits() { - if let Some(start) = ctx.add_bits(bits) { + if ctx.is_in_choice { + None + } else if let Some(start) = ctx.add_bits(bits) { Some(flat_ast::Bitset::new(0, start, format!("bitset{}", ctx.bitsets))) } else { None diff --git a/generator/src/main.rs b/generator/src/main.rs index 15416c1..ca99491 100644 --- a/generator/src/main.rs +++ b/generator/src/main.rs @@ -10,6 +10,7 @@ mod flatten; mod writer; mod codegen; mod graph_passes; +mod type_registry; use codegen::{cpp, rust, Codegen, CodegenCommands}; @@ -41,6 +42,9 @@ fn main() -> Result<(), failure::Error> { simple_logger::init_with_level(verbose).unwrap(); + let mut registry = type_registry::TypeRegistry::new(); + let mut packets = Vec::new(); + for filename in args.inputs.iter().map(std::path::Path::new) { debug!("filename {:?}", filename); use std::fs::File; @@ -53,9 +57,20 @@ fn main() -> Result<(), failure::Error> { trace!("packet {:?}", packet); let packet = graph_passes::run(packet)?; debug!("packet {:#?}", packet); + registry.collect_from_packet(&packet)?; + packets.push(packet); + } + + registry.compute_is_copy(); + + if let CodegenCommands::RustCommand(ref rust_args) = args.command { + rust::generate_shared(VERSION, ®istry, rust_args)?; + } + + for packet in packets { let mut generator: Box = match &args.command { CodegenCommands::CppCommand(args) => Box::new(cpp::Generator::new(args)), - CodegenCommands::RustCommand(args) => Box::new(rust::Generator::new(args)), + CodegenCommands::RustCommand(args) => Box::new(rust::Generator::new(args, ®istry)), }; generator.generate(VERSION, &packet)?; info!("Generated packet {}", packet.type_()); diff --git a/generator/src/type_registry.rs b/generator/src/type_registry.rs new file mode 100644 index 0000000..f6b4d52 --- /dev/null +++ b/generator/src/type_registry.rs @@ -0,0 +1,195 @@ +use std::collections::{BTreeMap, HashMap}; +use crate::flat_ast::{Packet, PacketContent, SimpleType, ComplexType, ComplexTypeContent, Element, RestrictionContent, SimpleTypeContent, Sequence}; + +#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)] +pub enum TypeFingerprint { + Simple { + name: String, + base: String, + restrictions: Vec, + }, + Enum { + name: String, + variants: Vec<(String, i64)>, + }, + Struct { + name: String, + fields: Vec<(String, String, Option)>, // name, type, bits + }, +} + +pub struct TypeRegistry { + pub types: BTreeMap, + pub fingerprints: HashMap, + pub is_copy: HashMap, +} + +impl TypeRegistry { + pub fn new() -> Self { + let mut is_copy = HashMap::new(); + // Built-in types + is_copy.insert("int8_t".to_string(), true); + is_copy.insert("uint8_t".to_string(), true); + is_copy.insert("int16_t".to_string(), true); + is_copy.insert("uint16_t".to_string(), true); + is_copy.insert("int32_t".to_string(), true); + is_copy.insert("uint32_t".to_string(), true); + is_copy.insert("int64_t".to_string(), true); + is_copy.insert("uint64_t".to_string(), true); + is_copy.insert("char".to_string(), true); + is_copy.insert("int".to_string(), true); + is_copy.insert("unsigned int".to_string(), true); + is_copy.insert("float".to_string(), true); + is_copy.insert("double".to_string(), true); + is_copy.insert("bool".to_string(), true); + is_copy.insert("std::string".to_string(), false); + is_copy.insert("NullTerminatedString".to_string(), false); + + Self { + types: BTreeMap::new(), + fingerprints: HashMap::new(), + is_copy, + } + } + + pub fn collect_from_packet(&mut self, packet: &Packet) -> Result<(), failure::Error> { + let filename = packet.filename(); + for content in packet.contents() { + match content { + PacketContent::Simple(ref s) => self.register_simple(s, filename)?, + PacketContent::Complex(ref c) => self.register_complex(c, filename)?, + _ => {} + } + } + Ok(()) + } + + fn register_simple(&mut self, s: &SimpleType, filename: &str) -> Result<(), failure::Error> { + let fingerprint = self.fingerprint_simple(s); + if let Some(existing) = self.fingerprints.get(s.name()) { + if existing != &fingerprint { + return Err(failure::format_err!("Conflicting definitions for type {} in {}:\nExisting: {:?}\nNew: {:?}", s.name(), filename, existing, fingerprint)); + } + } else { + self.fingerprints.insert(s.name().clone(), fingerprint); + self.types.insert(s.name().clone(), PacketContent::Simple(s.clone())); + } + Ok(()) + } + + fn register_complex(&mut self, c: &ComplexType, filename: &str) -> Result<(), failure::Error> { + if c.anonymous() || c.inline() { return Ok(()); } + + let fingerprints = self.fingerprint_complex(c); + for (name, fingerprint) in fingerprints { + if let Some(existing) = self.fingerprints.get(&name) { + if existing != &fingerprint { + return Err(failure::format_err!("Conflicting definitions for type {} in {}:\nExisting: {:?}\nNew: {:?}", name, filename, existing, fingerprint)); + } + } else { + self.fingerprints.insert(name.clone(), fingerprint); + if &name == c.name() { + self.types.insert(name.clone(), PacketContent::Complex(c.clone())); + } + } + } + Ok(()) + } + + fn fingerprint_simple(&self, s: &SimpleType) -> TypeFingerprint { + let mut is_enum = false; + let mut variants = Vec::new(); + let mut base = String::new(); + for content in s.contents() { + if let SimpleTypeContent::Restriction(ref r) = content { + base = r.base().clone(); + for r_content in r.contents() { + if let RestrictionContent::Enumeration(ref e) = r_content { + is_enum = true; + variants.push((e.value().clone(), e.id())); + } + } + } + } + + if is_enum { + variants.sort(); + TypeFingerprint::Enum { name: s.name().clone(), variants } + } else { + let mut restrictions = Vec::new(); + for content in s.contents() { + if let SimpleTypeContent::Restriction(ref r) = content { + for r_content in r.contents() { + match r_content { + RestrictionContent::Length(l) => restrictions.push(format!("Length({})", l)), + RestrictionContent::MinValue(ref v) => restrictions.push(format!("Min({})", v)), + RestrictionContent::MaxValue(ref v) => restrictions.push(format!("Max({})", v)), + _ => {} + } + } + } + } + restrictions.sort(); + TypeFingerprint::Simple { name: s.name().clone(), base, restrictions } + } + } + + fn fingerprint_complex(&self, c: &ComplexType) -> Vec<(String, TypeFingerprint)> { + let mut res = Vec::new(); + match c.content() { + ComplexTypeContent::Seq(ref s) => { + res.push((c.name().clone(), self.fingerprint_sequence(c.name(), s))); + } + ComplexTypeContent::Choice(ref ch) => { + let mut fields = Vec::new(); + for elem in ch.elements() { + fields.push((elem.name().clone(), elem.type_().clone(), elem.bits())); + if let Some(seq) = ch.inline_seqs().get(elem.name()) { + res.push((elem.name().clone(), self.fingerprint_sequence(elem.name(), seq))); + } + } + res.push((c.name().clone(), TypeFingerprint::Struct { name: c.name().clone(), fields })); + } + ComplexTypeContent::Empty => { + res.push((c.name().clone(), TypeFingerprint::Struct { name: c.name().clone(), fields: Vec::new() })); + } + } + res + } + + fn fingerprint_sequence(&self, name: &str, s: &Sequence) -> TypeFingerprint { + let mut fields = Vec::new(); + for elem in s.elements() { + fields.push((elem.name().clone(), elem.type_().clone(), elem.bits())); + } + TypeFingerprint::Struct { name: name.to_string(), fields } + } + + pub fn compute_is_copy(&mut self) { + let mut changed = true; + while changed { + changed = false; + let fingerprints = self.fingerprints.clone(); + for (name, fp) in &fingerprints { + if self.is_copy.contains_key(name) { continue; } + + let current_is_copy = match fp { + TypeFingerprint::Simple { base, .. } => { + *self.is_copy.get(base).unwrap_or(&true) + }, + TypeFingerprint::Enum { .. } => true, + TypeFingerprint::Struct { fields, .. } => { + fields.iter().all(|(_, ty, _)| { + *self.is_copy.get(ty).unwrap_or(&false) + }) + } + }; + + if current_is_copy { + self.is_copy.insert(name.clone(), true); + changed = true; + } + } + } + } +}