use crate::core::ir::{ApiSurface, EnumDef, EnumVariant, FieldDef, PrimitiveType, TypeRef};
use ahash::AHashSet;
const CONTAINER_SIZE_ESTIMATE_BYTES: u64 = 24;
const MAP_SIZE_ESTIMATE_BYTES: u64 = 48;
const UNKNOWN_NAMED_TYPE_SIZE_ESTIMATE_BYTES: u64 = 8;
const MAX_RESOLUTION_DEPTH: usize = 6;
const CLIPPY_DEFAULT_THRESHOLD_BYTES: u64 = 200;
const ESTIMATE_TRUST_MARGIN_BYTES: u64 = 200;
const EXPECT_GAP_THRESHOLD_BYTES: u64 = CLIPPY_DEFAULT_THRESHOLD_BYTES + ESTIMATE_TRUST_MARGIN_BYTES;
fn primitive_size_estimate(primitive: &PrimitiveType) -> u64 {
match primitive {
PrimitiveType::Bool | PrimitiveType::U8 | PrimitiveType::I8 => 1,
PrimitiveType::U16 | PrimitiveType::I16 => 2,
PrimitiveType::U32 | PrimitiveType::I32 | PrimitiveType::F32 => 4,
PrimitiveType::U64 | PrimitiveType::I64 | PrimitiveType::F64 | PrimitiveType::Usize | PrimitiveType::Isize => 8,
}
}
fn type_ref_size_estimate(ty: &TypeRef, api: &ApiSurface, visiting: &mut AHashSet<String>, depth: usize) -> u64 {
match ty {
TypeRef::Primitive(p) => primitive_size_estimate(p),
TypeRef::String | TypeRef::Char | TypeRef::Path | TypeRef::Json | TypeRef::Bytes => {
CONTAINER_SIZE_ESTIMATE_BYTES
}
TypeRef::Duration => 8,
TypeRef::Unit => 0,
TypeRef::Vec(_) => CONTAINER_SIZE_ESTIMATE_BYTES,
TypeRef::Map(_, _) => MAP_SIZE_ESTIMATE_BYTES,
TypeRef::Optional(inner) => type_ref_size_estimate(inner, api, visiting, depth),
TypeRef::Named(name) => named_type_size_estimate(name, api, visiting, depth),
}
}
fn named_type_size_estimate(name: &str, api: &ApiSurface, visiting: &mut AHashSet<String>, depth: usize) -> u64 {
if depth >= MAX_RESOLUTION_DEPTH || !visiting.insert(name.to_string()) {
return UNKNOWN_NAMED_TYPE_SIZE_ESTIMATE_BYTES;
}
let size = if let Some(type_def) = api.types.iter().find(|t| t.name == name) {
type_def
.fields
.iter()
.map(|field| field_size_estimate(field, api, visiting, depth + 1))
.sum()
} else if let Some(enum_def) = api.enums.iter().find(|e| e.name == name) {
enum_def
.variants
.iter()
.map(|variant| sum_variant_fields(variant, api, visiting, depth + 1))
.max()
.unwrap_or(0)
} else {
UNKNOWN_NAMED_TYPE_SIZE_ESTIMATE_BYTES
};
visiting.remove(name);
size
}
fn field_size_estimate(field: &FieldDef, api: &ApiSurface, visiting: &mut AHashSet<String>, depth: usize) -> u64 {
type_ref_size_estimate(&field.ty, api, visiting, depth)
}
fn sum_variant_fields(variant: &EnumVariant, api: &ApiSurface, visiting: &mut AHashSet<String>, depth: usize) -> u64 {
variant
.fields
.iter()
.map(|field| field_size_estimate(field, api, visiting, depth))
.sum()
}
fn variant_size_estimate(variant: &EnumVariant, api: &ApiSurface) -> u64 {
let mut visiting = AHashSet::new();
sum_variant_fields(variant, api, &mut visiting, 0)
}
pub fn enum_should_expect_large_variant_lint(enum_def: &EnumDef, api: &ApiSurface) -> bool {
if enum_def.variants.len() < 2 {
return false;
}
let mut sizes: Vec<u64> = enum_def
.variants
.iter()
.map(|v| variant_size_estimate(v, api))
.collect();
sizes.sort_unstable_by(|a, b| b.cmp(a));
sizes[0].saturating_sub(sizes[1]) > EXPECT_GAP_THRESHOLD_BYTES
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::ir::TypeDef;
fn field(name: &str, ty: TypeRef) -> FieldDef {
FieldDef {
name: name.to_string(),
ty,
..FieldDef::default()
}
}
fn variant(name: &str, fields: Vec<FieldDef>) -> EnumVariant {
EnumVariant {
name: name.to_string(),
fields,
is_tuple: false,
..EnumVariant::default()
}
}
fn enum_def(name: &str, variants: Vec<EnumVariant>) -> EnumDef {
EnumDef {
name: name.to_string(),
rust_path: format!("sample_crate::{name}"),
variants,
..EnumDef::default()
}
}
#[test]
fn type_ref_size_estimate_table() {
let api = ApiSurface::default();
let cases: &[(&str, TypeRef, u64)] = &[
("bool", TypeRef::Primitive(PrimitiveType::Bool), 1),
("u8", TypeRef::Primitive(PrimitiveType::U8), 1),
("u16", TypeRef::Primitive(PrimitiveType::U16), 2),
("u32", TypeRef::Primitive(PrimitiveType::U32), 4),
("f32", TypeRef::Primitive(PrimitiveType::F32), 4),
("u64", TypeRef::Primitive(PrimitiveType::U64), 8),
("f64", TypeRef::Primitive(PrimitiveType::F64), 8),
("usize", TypeRef::Primitive(PrimitiveType::Usize), 8),
("string", TypeRef::String, 24),
("bytes", TypeRef::Bytes, 24),
("path", TypeRef::Path, 24),
("json", TypeRef::Json, 24),
("duration", TypeRef::Duration, 8),
("unit", TypeRef::Unit, 0),
(
"vec_of_named",
TypeRef::Vec(Box::new(TypeRef::Named("Huge".into()))),
24,
),
(
"map",
TypeRef::Map(Box::new(TypeRef::String), Box::new(TypeRef::String)),
48,
),
(
"optional_u64",
TypeRef::Optional(Box::new(TypeRef::Primitive(PrimitiveType::U64))),
8,
),
("unresolved_named", TypeRef::Named("NotInSurface".into()), 8),
];
for (label, ty, expected) in cases {
let mut visiting = AHashSet::new();
let actual = type_ref_size_estimate(ty, &api, &mut visiting, 0);
assert_eq!(actual, *expected, "case `{label}`: expected {expected}, got {actual}");
}
}
#[test]
fn named_struct_recurses_into_its_own_fields() {
let mut api = ApiSurface::default();
api.types.push(TypeDef {
name: "Heavy".to_string(),
fields: vec![
field("a", TypeRef::String),
field("b", TypeRef::String),
field("c", TypeRef::Primitive(PrimitiveType::U64)),
],
..TypeDef::default()
});
let mut visiting = AHashSet::new();
let size = type_ref_size_estimate(&TypeRef::Named("Heavy".into()), &api, &mut visiting, 0);
assert_eq!(size, 24 + 24 + 8, "two Strings plus one u64");
}
#[test]
fn named_cycle_falls_back_to_unknown_estimate_instead_of_recursing_forever() {
let mut api = ApiSurface::default();
api.types.push(TypeDef {
name: "Wrapper".to_string(),
fields: vec![field("inner", TypeRef::Named("Wrapper".into()))],
..TypeDef::default()
});
let mut visiting = AHashSet::new();
let size = type_ref_size_estimate(&TypeRef::Named("Wrapper".into()), &api, &mut visiting, 0);
assert_eq!(size, UNKNOWN_NAMED_TYPE_SIZE_ESTIMATE_BYTES);
}
#[test]
fn named_enum_resolves_to_its_heaviest_variant() {
let mut api = ApiSurface::default();
api.enums.push(enum_def(
"Inner",
vec![
variant("Small", vec![field("n", TypeRef::Primitive(PrimitiveType::U8))]),
variant("Big", vec![field("s1", TypeRef::String), field("s2", TypeRef::String)]),
],
));
let mut visiting = AHashSet::new();
let size = type_ref_size_estimate(&TypeRef::Named("Inner".into()), &api, &mut visiting, 0);
assert_eq!(size, 48, "heaviest variant (two Strings) wins, not the lightest");
}
#[test]
fn flags_enum_whose_named_payload_dwarfs_its_siblings() {
let mut api = ApiSurface::default();
api.types.push(TypeDef {
name: "HeavyConfig".to_string(),
fields: (0..20)
.map(|i| field(&format!("setting_{i}"), TypeRef::String))
.collect(),
..TypeDef::default()
});
let target = enum_def(
"ModelKind",
vec![
variant("Heavy", vec![field("heavy", TypeRef::Named("HeavyConfig".into()))]),
variant("Light", vec![field("n", TypeRef::Primitive(PrimitiveType::U32))]),
variant("Plain", vec![]),
],
);
assert!(
enum_should_expect_large_variant_lint(&target, &api),
"a variant estimated at 480 bytes against 4-byte and 0-byte siblings must be flagged"
);
}
#[test]
fn does_not_flag_enum_with_similarly_sized_variants() {
let api = ApiSurface::default();
let target = enum_def(
"RequestKind",
vec![
variant("Get", vec![field("path", TypeRef::String)]),
variant(
"Post",
vec![field("path", TypeRef::String), field("body", TypeRef::String)],
),
variant("Delete", vec![field("path", TypeRef::String)]),
],
);
assert!(
!enum_should_expect_large_variant_lint(&target, &api),
"variants within ~24 bytes of each other must not be flagged"
);
}
#[test]
fn does_not_flag_gap_inside_the_trust_margin() {
let api = ApiSurface::default();
let target = enum_def(
"Borderline",
vec![
variant(
"NineStrings",
(0..9).map(|i| field(&format!("f{i}"), TypeRef::String)).collect(),
),
variant("Unit", vec![]),
],
);
assert!(!enum_should_expect_large_variant_lint(&target, &api));
}
#[test]
fn does_not_flag_single_variant_enum() {
let api = ApiSurface::default();
let target = enum_def("Solo", vec![variant("Only", vec![field("s", TypeRef::String)])]);
assert!(!enum_should_expect_large_variant_lint(&target, &api));
}
#[test]
fn does_not_flag_unit_enum() {
let api = ApiSurface::default();
let target = enum_def("Color", vec![variant("Red", vec![]), variant("Blue", vec![])]);
assert!(!enum_should_expect_large_variant_lint(&target, &api));
}
}