use std::collections::{HashMap, HashSet};
use zerodds_idl::ast::{
Annotation, AnnotationParams, ConstrTypeDecl, Definition, Identifier, Member, ScopedName,
Specification, StructDcl, StructDef, TypeDecl,
};
use zerodds_idl::errors::Span;
use zerodds_idl::preprocessor::{OpenSplicePragma, PragmaKeylist};
const SYNTH: Span = Span { start: 0, end: 0 };
type KeyMap = HashMap<String, HashSet<String>>;
#[must_use]
pub fn collect_key_pragmas(keylists: &[PragmaKeylist], opensplice: &[OpenSplicePragma]) -> KeyMap {
let mut map: KeyMap = HashMap::new();
for kl in keylists {
let set = map.entry(kl.type_name.clone()).or_default();
for k in &kl.keys {
set.insert(k.clone());
}
}
for p in opensplice {
match p {
OpenSplicePragma::DataKey {
type_name, fields, ..
} => {
let set = map.entry(type_name.clone()).or_default();
for f in fields {
set.insert(f.clone());
}
}
OpenSplicePragma::Cats {
type_name, keys, ..
} => {
let set = map.entry(type_name.clone()).or_default();
for k in keys {
set.insert(k.clone());
}
}
OpenSplicePragma::DataType { .. } | OpenSplicePragma::GenEquality { .. } => {}
}
}
map
}
pub fn apply_key_pragmas(spec: &mut Specification, keys: &KeyMap) -> usize {
if keys.is_empty() {
return 0;
}
let mut patched = 0;
for def in &mut spec.definitions {
patched += patch_definition(def, &[], keys);
}
patched
}
fn patch_definition(def: &mut Definition, scope: &[String], keys: &KeyMap) -> usize {
match def {
Definition::Module(m) => {
let mut inner_scope = scope.to_vec();
inner_scope.push(m.name.text.clone());
let mut n = 0;
for d in &mut m.definitions {
n += patch_definition(d, &inner_scope, keys);
}
n
}
Definition::Type(TypeDecl::Constr(ConstrTypeDecl::Struct(StructDcl::Def(s)))) => {
patch_struct(s, scope, keys)
}
_ => 0,
}
}
fn patch_struct(s: &mut StructDef, scope: &[String], keys: &KeyMap) -> usize {
let local = s.name.text.clone();
let scoped = {
let mut parts = scope.to_vec();
parts.push(local.clone());
parts.join("::")
};
let Some(key_fields) = keys.get(&scoped).or_else(|| keys.get(&local)) else {
return 0;
};
let mut patched = 0;
for member in &mut s.members {
for decl in &member.declarators {
if key_fields.contains(&decl.name().text) && !has_key_annotation(member) {
member.annotations.push(make_key_annotation());
patched += 1;
}
}
}
patched
}
fn has_key_annotation(member: &Member) -> bool {
member
.annotations
.iter()
.any(|a| a.name.parts.last().is_some_and(|p| p.text == "key"))
}
fn make_key_annotation() -> Annotation {
Annotation {
name: ScopedName::single(Identifier::new("key", SYNTH)),
params: AnnotationParams::None,
span: SYNTH,
}
}