use std::collections::BTreeMap;
use std::collections::BTreeSet;
use crate::error::Error;
use crate::error::Result;
use crate::ir::EnumKind;
use crate::ir::Item;
use crate::ir::Module;
use crate::ir::RustType;
use crate::naming::Case;
use crate::naming::to_ident;
pub fn box_recursive_types(module: &mut Module) -> Result<()> {
let graph = Graph::of(module);
graph.check_alias_cycles()?;
for item in &mut module.items {
let owner = canonical(item.name());
match item {
Item::Struct(strukt) => {
for field in &mut strukt.fields {
box_held(&mut field.ty, &owner, &graph);
}
}
Item::Enum(enumeration) => {
if let EnumKind::Union(variants) = &mut enumeration.kind {
for variant in variants {
box_held(&mut variant.ty, &owner, &graph);
}
}
}
Item::Alias(_) => {}
}
}
return Ok(());
}
struct Graph {
names: Vec<String>,
nodes: BTreeMap<String, usize>,
edges: Vec<Vec<usize>>,
is_alias: Vec<bool>,
component: Vec<usize>,
}
impl Graph {
fn of(module: &Module) -> Self {
let names: Vec<String> = module.items.iter().map(|item| return canonical(item.name())).collect();
let nodes: BTreeMap<String, usize> = names
.iter()
.enumerate()
.map(|(node, name)| return (name.clone(), node))
.collect();
let mut edges = Vec::with_capacity(module.items.len());
let mut is_alias = Vec::with_capacity(module.items.len());
for item in &module.items {
let mut targets = BTreeSet::new();
match item {
Item::Struct(strukt) => {
for field in &strukt.fields {
collect_held(&field.ty, &mut targets);
}
}
Item::Enum(enumeration) => {
if let EnumKind::Union(variants) = &enumeration.kind {
for variant in variants {
collect_held(&variant.ty, &mut targets);
}
}
}
Item::Alias(alias) => collect_held(&alias.ty, &mut targets),
}
is_alias.push(matches!(item, Item::Alias(_)));
edges.push(
targets
.iter()
.filter_map(|target| return nodes.get(target).copied())
.collect(),
);
}
let component = components(&edges);
return Self {
names,
nodes,
edges,
is_alias,
component,
};
}
fn on_a_cycle(&self, from: &str, to: &str) -> bool {
let (Some(from), Some(to)) = (self.nodes.get(from), self.nodes.get(to)) else {
return false;
};
let (Some(from), Some(to)) = (self.component.get(*from), self.component.get(*to)) else {
return false;
};
return from == to;
}
fn check_alias_cycles(&self) -> Result<()> {
let mut members: BTreeMap<usize, Vec<usize>> = BTreeMap::new();
for (node, component) in self.component.iter().enumerate() {
members.entry(*component).or_default().push(node);
}
for group in members.values() {
if !self.is_cyclic(group) || !group.iter().all(|node| return self.is_alias(*node)) {
continue;
}
return Err(Error::RecursiveAlias {
cycle: self.cycle_through(group),
hint: "Give one of these schemas `type: object` with properties, so the generator emits a struct it \
can box, or break the chain of `$ref`s."
.to_owned(),
});
}
return Ok(());
}
fn is_cyclic(&self, group: &[usize]) -> bool {
if group.len() > 1 {
return true;
}
return group
.first()
.is_some_and(|node| return self.edges_of(*node).contains(node));
}
fn edges_of(&self, node: usize) -> &[usize] {
return self.edges.get(node).map_or(&[], Vec::as_slice);
}
fn is_alias(&self, node: usize) -> bool {
return self.is_alias.get(node).copied().unwrap_or(false);
}
fn name_of(&self, node: usize) -> String {
return self.names.get(node).cloned().unwrap_or_default();
}
fn cycle_through(&self, group: &[usize]) -> Vec<String> {
let Some(start) = group.first().copied() else {
return Vec::new();
};
let mut path = vec![self.name_of(start)];
let mut current = start;
for _ in 0..group.len() {
let Some(next) = self.edges_of(current).first().copied() else {
break;
};
path.push(self.name_of(next));
if next == start {
break;
}
current = next;
}
return path;
}
}
fn components(edges: &[Vec<usize>]) -> Vec<usize> {
let mut walk = Walk::over(edges.len());
let mut work: Vec<(usize, usize)> = Vec::new();
for root in 0..edges.len() {
if walk.is_open_or_done(root) {
continue;
}
walk.open(root);
work.push((root, 0));
while let Some(&(node, taken)) = work.last() {
let next = edges.get(node).and_then(|held| return held.get(taken)).copied();
if let Some(next) = next {
if let Some(entry) = work.last_mut() {
entry.1 += 1;
}
walk.step(node, next, &mut work);
continue;
}
work.pop();
if let Some(&(parent, _)) = work.last() {
walk.carry_up(parent, node);
}
walk.close(node);
}
}
return walk.component;
}
struct Walk {
order: Vec<usize>,
lowest: Vec<usize>,
open: Vec<bool>,
component: Vec<usize>,
path: Vec<usize>,
next_order: usize,
next_component: usize,
}
impl Walk {
const UNREACHED: usize = usize::MAX;
fn over(count: usize) -> Self {
return Self {
order: vec![Self::UNREACHED; count],
lowest: vec![0; count],
open: vec![false; count],
component: vec![0; count],
path: Vec::new(),
next_order: 0,
next_component: 0,
};
}
fn is_open_or_done(&self, node: usize) -> bool {
return self.order_of(node) != Self::UNREACHED;
}
fn order_of(&self, node: usize) -> usize {
return self.order.get(node).copied().unwrap_or(Self::UNREACHED);
}
fn lowest_of(&self, node: usize) -> usize {
return self.lowest.get(node).copied().unwrap_or(Self::UNREACHED);
}
fn open(&mut self, node: usize) {
if let Some(order) = self.order.get_mut(node) {
*order = self.next_order;
}
if let Some(lowest) = self.lowest.get_mut(node) {
*lowest = self.next_order;
}
if let Some(open) = self.open.get_mut(node) {
*open = true;
}
self.next_order += 1;
self.path.push(node);
}
fn step(&mut self, node: usize, next: usize, work: &mut Vec<(usize, usize)>) {
if !self.is_open_or_done(next) {
self.open(next);
work.push((next, 0));
return;
}
if self.open.get(next).copied().unwrap_or(false) {
self.lower(node, self.order_of(next));
}
}
fn carry_up(&mut self, parent: usize, node: usize) {
self.lower(parent, self.lowest_of(node));
}
fn lower(&mut self, node: usize, order: usize) {
if let Some(lowest) = self.lowest.get_mut(node) {
*lowest = (*lowest).min(order);
}
}
fn close(&mut self, node: usize) {
if self.lowest_of(node) != self.order_of(node) {
return;
}
while let Some(member) = self.path.pop() {
if let Some(open) = self.open.get_mut(member) {
*open = false;
}
if let Some(component) = self.component.get_mut(member) {
*component = self.next_component;
}
if member == node {
break;
}
}
self.next_component += 1;
}
}
fn collect_held(ty: &RustType, out: &mut BTreeSet<String>) {
match ty {
RustType::Named(name) => {
out.insert(canonical(name));
}
RustType::Option(inner) => collect_held(inner, out),
_ => {}
}
}
fn box_held(ty: &mut RustType, owner: &str, graph: &Graph) {
match ty {
RustType::Named(name) => {
let target = canonical(name);
if graph.on_a_cycle(&target, owner) {
let inner = std::mem::replace(ty, RustType::Bool);
*ty = RustType::Boxed(Box::new(inner));
}
}
RustType::Option(inner) => box_held(inner, owner, graph),
_ => {}
}
}
fn canonical(name: &str) -> String {
return to_ident(name, Case::Pascal).logical().to_owned();
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ir::Alias;
use crate::ir::Enum;
use crate::ir::Field;
use crate::ir::Struct;
use crate::ir::UnionVariant;
fn one_field(name: &str, field: &str, ty: RustType) -> Item {
return Item::Struct(Struct {
name: to_ident(name, Case::Pascal),
doc: None,
deprecated: None,
fields: vec![Field {
name: to_ident(field, Case::Snake),
rename: None,
doc: None,
deprecated: None,
ty,
required: true,
omit_empty: None,
serde_skip: false,
default: None,
constraints: None,
}],
additional_properties: None,
deny_unknown_fields: false,
});
}
fn field_type(module: &Module, name: &str) -> RustType {
for item in &module.items {
if let Item::Struct(strukt) = item
&& strukt.name.logical() == name
{
return strukt.fields[0].ty.clone();
}
}
panic!("no struct named `{name}`");
}
fn named(name: &str) -> RustType {
return RustType::Named(name.to_owned());
}
fn boxed(inner: RustType) -> RustType {
return RustType::Boxed(Box::new(inner));
}
#[test]
fn only_a_field_that_holds_its_owner_is_boxed() {
let cases: &[(&str, RustType, RustType)] = &[
("direct", named("Node"), boxed(named("Node"))),
(
"through an option",
RustType::Option(Box::new(named("Node"))),
RustType::Option(Box::new(boxed(named("Node")))),
),
(
"through a vec",
RustType::Vec(Box::new(named("Node"))),
RustType::Vec(Box::new(named("Node"))),
),
(
"through a map",
RustType::Map(Box::new(named("Node"))),
RustType::Map(Box::new(named("Node"))),
),
("a scalar", RustType::String, RustType::String),
];
for (label, input, want) in cases {
let mut module = Module {
items: vec![one_field("Node", "child", input.clone())],
};
box_recursive_types(&mut module).expect("no alias cycle in this module");
assert_eq!(field_type(&module, "Node"), *want, "self-reference {label}");
}
}
#[test]
fn both_sides_of_a_mutual_cycle_are_boxed() {
let mut module = Module {
items: vec![
one_field("Parent", "child", named("Kid")),
one_field("Kid", "parent", named("Parent")),
],
};
box_recursive_types(&mut module).expect("no alias cycle in this module");
assert_eq!(field_type(&module, "Parent"), boxed(named("Kid")));
assert_eq!(field_type(&module, "Kid"), boxed(named("Parent")));
}
#[test]
fn a_type_the_cycle_only_points_at_is_left_alone() {
let mut module = Module {
items: vec![
one_field("Node", "child", named("Node")),
one_field("Holder", "node", named("Node")),
],
};
box_recursive_types(&mut module).expect("no alias cycle in this module");
assert_eq!(field_type(&module, "Holder"), named("Node"));
}
#[test]
fn a_union_variant_that_holds_its_own_enum_is_boxed() {
let mut module = Module {
items: vec![Item::Enum(Enum {
name: to_ident("Expression", Case::Pascal),
doc: None,
deprecated: None,
kind: EnumKind::Union(vec![
UnionVariant {
name: to_ident("Text", Case::Pascal),
ty: RustType::String,
},
UnionVariant {
name: to_ident("Nested", Case::Pascal),
ty: named("Expression"),
},
]),
})],
};
box_recursive_types(&mut module).expect("no alias cycle in this module");
let Item::Enum(enumeration) = &module.items[0] else {
panic!("the item is an enum");
};
let EnumKind::Union(variants) = &enumeration.kind else {
panic!("the enum is a union");
};
assert_eq!(variants[0].ty, RustType::String);
assert_eq!(variants[1].ty, boxed(named("Expression")));
}
#[test]
fn a_cycle_through_an_alias_is_boxed_at_the_struct() {
let mut module = Module {
items: vec![
Item::Alias(Alias {
name: to_ident("Wrapper", Case::Pascal),
doc: None,
deprecated: None,
ty: named("Holder"),
}),
one_field("Holder", "wrapped", named("Wrapper")),
],
};
box_recursive_types(&mut module).expect("this cycle holds a struct, so it is not alias-only");
assert_eq!(field_type(&module, "Holder"), boxed(named("Wrapper")));
}
#[test]
fn an_alias_only_cycle_is_rejected() {
let alias = |name: &str, target: &str| {
return Item::Alias(Alias {
name: to_ident(name, Case::Pascal),
doc: None,
deprecated: None,
ty: named(target),
});
};
let mut module = Module {
items: vec![alias("Loop", "Ring"), alias("Ring", "Loop")],
};
let outcome = box_recursive_types(&mut module);
assert!(
matches!(outcome, Err(Error::RecursiveAlias { .. })),
"a cycle of aliases must be rejected, and gave: {outcome:?}",
);
}
#[test]
fn a_self_referencing_alias_is_rejected() {
let mut module = Module {
items: vec![Item::Alias(Alias {
name: to_ident("Loop", Case::Pascal),
doc: None,
deprecated: None,
ty: named("Loop"),
})],
};
assert!(matches!(
box_recursive_types(&mut module),
Err(Error::RecursiveAlias { .. })
));
}
#[test]
fn a_module_without_a_cycle_is_unchanged() {
let mut module = Module {
items: vec![
one_field("Holder", "node", named("Node")),
one_field("Node", "id", RustType::String),
],
};
let before = module.clone();
box_recursive_types(&mut module).expect("no alias cycle in this module");
assert_eq!(module, before);
}
#[test]
fn an_existing_box_breaks_the_cycle_for_every_other_edge() {
let mut module = Module {
items: vec![
one_field("Parent", "child", boxed(named("Kid"))),
one_field("Kid", "parent", named("Parent")),
],
};
box_recursive_types(&mut module).expect("no alias cycle in this module");
assert_eq!(field_type(&module, "Parent"), boxed(named("Kid")));
assert_eq!(field_type(&module, "Kid"), named("Parent"));
}
}