diff --git a/asn-compiler/src/generator/asn/types/base/real.rs b/asn-compiler/src/generator/asn/types/base/real.rs index 8becf16..36a5eb5 100644 --- a/asn-compiler/src/generator/asn/types/base/real.rs +++ b/asn-compiler/src/generator/asn/types/base/real.rs @@ -12,7 +12,8 @@ impl Asn1ResolvedReal { let type_name = generator.to_type_ident(name); let vis = generator.get_visibility_tokens(); - let dir = generator.generate_derive_tokens(); + // `f64` is not `Eq`, so never derive `Eq` for a `REAL`. + let dir = generator.generate_derive_tokens_skip_eq(&type_name.to_string(), true); Ok(quote! { #dir diff --git a/asn-compiler/src/generator/asn/types/constructed/choice.rs b/asn-compiler/src/generator/asn/types/constructed/choice.rs index fa5e617..ae5f716 100644 --- a/asn-compiler/src/generator/asn/types/constructed/choice.rs +++ b/asn-compiler/src/generator/asn/types/constructed/choice.rs @@ -49,7 +49,10 @@ impl ResolvedConstructedType { }; let vis = generator.get_visibility_tokens(); - let dir = generator.generate_derive_tokens(); + let dir = generator.generate_derive_tokens_skip_eq( + &type_name.to_string(), + generator.constructed_type_has_real(self), + ); let struct_tokens = ResolvedConstructedType::generate_struct_tokens_for_asn_choice_type( &type_name, diff --git a/asn-compiler/src/generator/asn/types/constructed/seq.rs b/asn-compiler/src/generator/asn/types/constructed/seq.rs index fb3aa91..2eb15f8 100644 --- a/asn-compiler/src/generator/asn/types/constructed/seq.rs +++ b/asn-compiler/src/generator/asn/types/constructed/seq.rs @@ -83,7 +83,10 @@ impl ResolvedConstructedType { ty_tokens.extend(quote! { , optional_fields = #optflds }); } - let dir = generator.generate_derive_tokens(); + let dir = generator.generate_derive_tokens_skip_eq( + &type_name.to_string(), + generator.constructed_type_has_real(self), + ); Ok(quote! { #dir #[asn(#ty_tokens)] diff --git a/asn-compiler/src/generator/asn/types/constructed/seqof.rs b/asn-compiler/src/generator/asn/types/constructed/seqof.rs index 7e55cce..2364c11 100644 --- a/asn-compiler/src/generator/asn/types/constructed/seqof.rs +++ b/asn-compiler/src/generator/asn/types/constructed/seqof.rs @@ -41,7 +41,10 @@ impl ResolvedConstructedType { )?; let vis = generator.get_visibility_tokens(); - let dir = generator.generate_derive_tokens(); + let dir = generator.generate_derive_tokens_skip_eq( + &seq_of_type_ident.to_string(), + generator.constructed_type_has_real(self), + ); Ok(quote! { #dir diff --git a/asn-compiler/src/generator/asn/types/int.rs b/asn-compiler/src/generator/asn/types/int.rs index 9c5f6c7..414c481 100644 --- a/asn-compiler/src/generator/asn/types/int.rs +++ b/asn-compiler/src/generator/asn/types/int.rs @@ -74,7 +74,10 @@ impl ResolvedSetType { let ty_elements = self.generate_aux_types(generator)?; let vis = generator.get_visibility_tokens(); - let dir = generator.generate_derive_tokens(); + let dir = generator.generate_derive_tokens_skip_eq( + &ty_ident.to_string(), + generator.set_type_has_real(self), + ); Ok(quote! { #dir @@ -97,7 +100,10 @@ impl ResolvedSetType { let ty_elements = self.generate_aux_types(generator)?; let vis = generator.get_visibility_tokens(); - let dir = generator.generate_derive_tokens(); + let dir = generator.generate_derive_tokens_skip_eq( + &ty_ident.to_string(), + generator.set_type_has_real(self), + ); let set_ty = quote! { #dir diff --git a/asn-compiler/src/generator/int.rs b/asn-compiler/src/generator/int.rs index 21b094a..7b59964 100644 --- a/asn-compiler/src/generator/int.rs +++ b/asn-compiler/src/generator/int.rs @@ -1,6 +1,6 @@ //! Code Generation module -use std::collections::HashMap; +use std::collections::{HashMap, HashSet}; use anyhow::Result; use heck::{ToShoutySnakeCase, ToSnakeCase}; @@ -11,7 +11,13 @@ use lazy_static::lazy_static; use crate::resolver::Resolver; -use crate::resolver::asn::structs::{types::Asn1ResolvedType, values::Asn1ResolvedValue}; +use crate::resolver::asn::structs::{ + types::{ + base::ResolvedBaseType, constructed::ResolvedConstructedType, Asn1ResolvedType, + ResolvedSetType, + }, + values::Asn1ResolvedValue, +}; /// Supported Codecs #[derive(clap::ValueEnum, Clone, Debug, PartialEq, Eq, Hash)] @@ -98,6 +104,10 @@ pub(crate) struct Generator { // Derives pub(crate) derives: Vec, + + // Names of the resolved types that contain a `REAL` (directly or through a + // reference). `f64` is not `Eq`, so `Eq` must not be derived for these types. + pub(crate) real_types: HashSet, } impl Generator { @@ -109,6 +119,7 @@ impl Generator { visibility: visibility.clone(), codecs, derives, + real_types: HashSet::new(), } } @@ -127,8 +138,12 @@ impl Generator { } } + // Find out the types that have a `REAL` in them before generating any code. + let resolved_types = resolver.get_resolved_types(); + self.real_types = Self::find_types_containing_real(&resolved_types); + // Now get the types - for (k, t) in resolver.get_resolved_types() { + for (k, t) in resolved_types { let item = Asn1ResolvedType::generate_for_type(k, t, self)?; if let Some(it) = item { items.push(it) @@ -223,6 +238,13 @@ impl Generator { } pub(crate) fn generate_derive_tokens(&self) -> TokenStream { + self.generate_derive_tokens_skip_eq("", false) + } + + // Same as `generate_derive_tokens`, but leaves out `Eq` when `skip_eq` is true. Used for + // types that contain a `REAL` (`f64`), which can only be `PartialEq`. `name` is the name of + // the type being generated, used for the warning when `Eq` is asked for but is skipped. + pub(crate) fn generate_derive_tokens_skip_eq(&self, name: &str, skip_eq: bool) -> TokenStream { let mut tokens = vec![]; for codec in &self.codecs { let codec_token = CODEC_TOKENS.get(codec).unwrap(); @@ -231,9 +253,22 @@ impl Generator { for derive in &self.derives { if derive == &Derive::All { - for derive_token in DERIVE_TOKENS.values() { + for (d, derive_token) in DERIVE_TOKENS.iter() { + if skip_eq && d == &Derive::Eq { + log::warn!( + "Not deriving `Eq` for type `{}` as it contains a `REAL`.", + name + ); + continue; + } tokens.push(derive_token.to_string()); } + } else if skip_eq && derive == &Derive::Eq { + log::warn!( + "Not deriving `Eq` for type `{}` as it contains a `REAL`.", + name + ); + continue; } else { let derive_token = DERIVE_TOKENS.get(derive).unwrap(); tokens.push(derive_token.to_string()); @@ -246,6 +281,72 @@ impl Generator { let derive_token_stream: TokenStream = derive_token_string.parse().unwrap(); derive_token_stream } + + // A type can refer to another type that contains a `REAL`, which in turn can be referred + // by some other type and so on. So we keep going over all the types till no new type + // gets added to the set. + fn find_types_containing_real(types: &[(&String, &Asn1ResolvedType)]) -> HashSet { + let mut found = HashSet::new(); + loop { + let mut changed = false; + for (name, ty) in types { + if !found.contains(*name) && Self::type_has_real(ty, &found) { + found.insert((*name).clone()); + changed = true; + } + } + if !changed { + break; + } + } + found + } + + fn type_has_real(ty: &Asn1ResolvedType, real_types: &HashSet) -> bool { + match ty { + Asn1ResolvedType::Base(ResolvedBaseType::Real(..)) => true, + Asn1ResolvedType::Base(..) => false, + Asn1ResolvedType::Reference(ref r) => real_types.contains(r), + Asn1ResolvedType::Constructed(ref c) => Self::constructed_has_real(c, real_types), + Asn1ResolvedType::Set(ref s) => Self::set_has_real(s, real_types), + } + } + + fn constructed_has_real(c: &ResolvedConstructedType, real_types: &HashSet) -> bool { + match c { + ResolvedConstructedType::Choice { + root_components, + additions, + .. + } => root_components + .iter() + .chain(additions.iter().flatten()) + .any(|comp| Self::type_has_real(&comp.ty, real_types)), + ResolvedConstructedType::Sequence { + components, + additions, + .. + } => components + .iter() + .chain(additions.iter().flatten()) + .any(|comp| Self::type_has_real(&comp.component.ty, real_types)), + ResolvedConstructedType::SequenceOf { ty, .. } => Self::type_has_real(ty, real_types), + } + } + + fn set_has_real(s: &ResolvedSetType, real_types: &HashSet) -> bool { + s.types + .values() + .any(|(_, ty)| Self::type_has_real(ty, real_types)) + } + + pub(crate) fn constructed_type_has_real(&self, c: &ResolvedConstructedType) -> bool { + Self::constructed_has_real(c, &self.real_types) + } + + pub(crate) fn set_type_has_real(&self, s: &ResolvedSetType) -> bool { + Self::set_has_real(s, &self.real_types) + } } fn capitalize_first(input: &str) -> String { diff --git a/asn-compiler/tests/mod.rs b/asn-compiler/tests/mod.rs index 6524978..90acb0c 100644 --- a/asn-compiler/tests/mod.rs +++ b/asn-compiler/tests/mod.rs @@ -170,4 +170,82 @@ AS-Config ::= SEQUENCE { assert!(result.is_ok(), "{:#?}", result.err().unwrap()); } + // Returns the derives generated for the item ` ` (eg. `struct Foo`). + fn get_derives_for(code: &str, kind: &str, name: &str) -> Vec { + let code: String = code.chars().filter(|c| !c.is_whitespace()).collect(); + let item = format!("pub{}{}", kind, name); + let item_pos = code + .find(&format!("{}{{", item)) + .or_else(|| code.find(&format!("{}(", item))) + .unwrap_or_else(|| panic!("`{} {}` not found in generated code", kind, name)); + let before = &code[..item_pos]; + let derive_start = before.rfind("#[derive(").unwrap() + "#[derive(".len(); + let derive_end = derive_start + before[derive_start..].find(")]").unwrap(); + before[derive_start..derive_end] + .split(',') + .map(|d| d.to_string()) + .collect() + } + + #[test] + fn no_eq_derive_for_types_with_real_110() { + let module_name = "NoEqDeriveForReal"; + let test_no = 5; + let module_header = super::get_module_header(module_name, test_no); + + let definitions = r#" +Measurement ::= SEQUENCE { + value MeasValue, + count INTEGER (0..255) +} + +MeasValue ::= CHOICE { + realValue REAL, + intValue INTEGER (0..255) +} + +MeasList ::= SEQUENCE (SIZE (1..16)) OF Measurement + +Counter ::= SEQUENCE { + count INTEGER (0..255) +}"#; + + let definitions = super::get_module_definitions(definitions); + let module_str = format!("{} {}", module_header, definitions); + + let output = std::env::temp_dir().join("hampi_test_no_eq_derive_for_real.rs"); + let mut compiler = Asn1Compiler::new( + output.to_str().unwrap(), + &Visibility::Public, + vec![Codec::Aper], + vec![Derive::Debug, Derive::Eq, Derive::PartialEq], + ); + compiler.set_rustfmt_generated_code(false); + let result = compiler.compile_string(&module_str, false); + assert!(result.is_ok(), "{:#?}", result.err().unwrap()); + + let code = std::fs::read_to_string(&output).unwrap(); + let _ = std::fs::remove_file(&output); + + // Types that have a REAL in them, directly or through a reference, should not get `Eq`. + for (kind, name) in [ + ("enum", "MeasValue"), + ("struct", "Measurement"), + ("struct", "MeasList"), + ] { + let derives = get_derives_for(&code, kind, name); + assert!( + !derives.contains(&"Eq".to_string()), + "`{} {}` should not derive Eq: {:?}", + kind, + name, + derives + ); + assert!(derives.contains(&"PartialEq".to_string())); + } + + // Types without REAL still get `Eq`. + let derives = get_derives_for(&code, "struct", "Counter"); + assert!(derives.contains(&"Eq".to_string()), "{:?}", derives); + } }