Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion asn-compiler/src/generator/asn/types/base/real.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
5 changes: 4 additions & 1 deletion asn-compiler/src/generator/asn/types/constructed/choice.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
5 changes: 4 additions & 1 deletion asn-compiler/src/generator/asn/types/constructed/seq.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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)]
Expand Down
5 changes: 4 additions & 1 deletion asn-compiler/src/generator/asn/types/constructed/seqof.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
10 changes: 8 additions & 2 deletions asn-compiler/src/generator/asn/types/int.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down
109 changes: 105 additions & 4 deletions asn-compiler/src/generator/int.rs
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
//! Code Generation module

use std::collections::HashMap;
use std::collections::{HashMap, HashSet};

use anyhow::Result;
use heck::{ToShoutySnakeCase, ToSnakeCase};
Expand All @@ -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)]
Expand Down Expand Up @@ -98,6 +104,10 @@ pub(crate) struct Generator {

// Derives
pub(crate) derives: Vec<Derive>,

// 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<String>,
}

impl Generator {
Expand All @@ -109,6 +119,7 @@ impl Generator {
visibility: visibility.clone(),
codecs,
derives,
real_types: HashSet::new(),
}
}

Expand All @@ -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)
Expand Down Expand Up @@ -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();
Expand All @@ -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 {
Comment thread
rizwan3659 marked this conversation as resolved.
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 {
Comment thread
rizwan3659 marked this conversation as resolved.
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());
Expand All @@ -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<String> {
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<String>) -> 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<String>) -> 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<String>) -> 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 {
Expand Down
78 changes: 78 additions & 0 deletions asn-compiler/tests/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -170,4 +170,82 @@ AS-Config ::= SEQUENCE {

assert!(result.is_ok(), "{:#?}", result.err().unwrap());
}
// Returns the derives generated for the item `<kind> <name>` (eg. `struct Foo`).
fn get_derives_for(code: &str, kind: &str, name: &str) -> Vec<String> {
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);
}
}
Loading