use std::collections::BTreeSet;
use std::collections::HashMap;
use crate::ir::EnumKind;
use crate::ir::Item;
use crate::ir::Module;
use crate::ir::RequestPayload;
use crate::ir::ResponseBody;
use crate::ir::RustType;
use crate::ir::Service;
use crate::ir::Struct;
use crate::naming::Case;
use crate::naming::to_ident;
pub fn prune_unused_models(module: &mut Module, service: &Service) {
let mut reachable: BTreeSet<String> = BTreeSet::new();
let mut worklist: Vec<String> = Vec::new();
for name in service_refs(service) {
if reachable.insert(name.clone()) {
worklist.push(name);
}
}
let index: HashMap<&str, &Item> = module.items.iter().map(|item| return (item.name(), item)).collect();
while let Some(name) = worklist.pop() {
let Some(item) = index.get(name.as_str()) else {
continue;
};
for referenced in item_refs(item) {
if reachable.insert(referenced.clone()) {
worklist.push(referenced);
}
}
}
module.items.retain(|item| return reachable.contains(item.name()));
}
fn canonical(name: &str) -> String {
return to_ident(name, Case::Pascal).logical().to_owned();
}
fn named_ref(ty: &RustType) -> Option<String> {
return match ty {
RustType::Named(name) => Some(canonical(name)),
RustType::Vec(inner) | RustType::Map(inner) | RustType::Option(inner) | RustType::Boxed(inner) => {
named_ref(inner)
}
_ => None,
};
}
fn item_refs(item: &Item) -> Vec<String> {
return match item {
Item::Struct(s) => struct_refs(s),
Item::Enum(enumeration) => match &enumeration.kind {
EnumKind::Strings(_) | EnumKind::Integers { .. } => Vec::new(),
EnumKind::Union(variants) => variants
.iter()
.filter_map(|variant| return named_ref(&variant.ty))
.collect(),
},
Item::Alias(alias) => named_ref(&alias.ty).into_iter().collect(),
};
}
fn struct_refs(s: &Struct) -> Vec<String> {
return s
.fields
.iter()
.filter_map(|field| return named_ref(&field.ty))
.chain(s.additional_properties.as_ref().and_then(named_ref))
.collect();
}
fn service_refs(service: &Service) -> Vec<String> {
let mut refs = Vec::new();
for operation in &service.operations {
for param in &operation.path_params {
refs.extend(named_ref(¶m.ty));
}
if let Some(query) = &operation.query {
refs.extend(struct_refs(query));
}
if let Some(headers) = &operation.headers {
for param in &headers.params {
refs.extend(named_ref(¶m.ty));
}
}
if let Some(cookies) = &operation.cookies {
for param in &cookies.params {
refs.extend(named_ref(¶m.ty));
}
}
if let Some(request) = &operation.request {
match request {
RequestPayload::Single(body) => refs.extend(named_ref(&body.ty)),
RequestPayload::Multipart(multipart) => {
for field in &multipart.fields {
refs.extend(named_ref(&field.ty));
}
}
RequestPayload::Negotiated(negotiated) => {
for variant in &negotiated.variants {
refs.extend(named_ref(&variant.body.ty));
}
}
}
}
for response in &operation.responses {
match &response.body {
Some(ResponseBody::Single(body)) => refs.extend(named_ref(&body.ty)),
Some(ResponseBody::Negotiated(negotiated)) => {
for variant in &negotiated.variants {
refs.extend(named_ref(&variant.body.ty));
}
}
None => {}
}
for header in &response.headers {
refs.extend(named_ref(&header.ty));
}
}
}
return refs;
}