use super::{Hyperedge, NodeId, OpenHypergraph};
use crate::strict::vec::FiniteFunction;
#[derive(Debug, Clone, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum WithSpider<A> {
Operation(A),
Spider,
}
impl<O: Clone + PartialEq, A: Clone> OpenHypergraph<O, A> {
pub fn spiderize(self) -> Result<OpenHypergraph<O, WithSpider<A>>, FiniteFunction> {
let nodes: Vec<NodeId> = (0..self.hypergraph.nodes.len()).map(NodeId).collect();
self.spiderize_nodes(&nodes)
}
pub fn spiderize_nodes(
mut self,
nodes: &[NodeId],
) -> Result<OpenHypergraph<O, WithSpider<A>>, FiniteFunction> {
let input_node_count = self.hypergraph.nodes.len();
for node in nodes {
assert!(
node.0 < input_node_count,
"node id {:?} is out of bounds",
node
);
}
let remapped_nodes;
let nodes = if self.hypergraph.is_strict() {
nodes
} else {
let quotient = self.quotient()?;
remapped_nodes = nodes
.iter()
.map(|node| NodeId(quotient.table[node.0]))
.collect::<Vec<_>>();
remapped_nodes.as_slice()
};
assert_eq!(
self.hypergraph.edges.len(),
self.hypergraph.adjacency.len(),
"malformed hypergraph: edges and adjacency lengths differ"
);
let spiders = rewrite_occurrences(
nodes,
&mut self.hypergraph.nodes,
&mut self.hypergraph.adjacency,
&mut self.sources,
&mut self.targets,
);
let mut result = self.map_edges(WithSpider::Operation);
for (left_spider, right_spider) in spiders {
result.new_edge(WithSpider::Spider, left_spider);
result.new_edge(WithSpider::Spider, right_spider);
}
Ok(result)
}
}
fn new_occurrence<O: Clone>(nodes: &mut Vec<O>, node: NodeId) -> NodeId {
let occurrence = NodeId(nodes.len());
nodes.push(nodes[node.0].clone());
occurrence
}
fn rewrite_occurrences<O: Clone>(
selected: &[NodeId],
nodes: &mut Vec<O>,
adjacency: &mut [Hyperedge],
sources: &mut [NodeId],
targets: &mut [NodeId],
) -> Vec<(Hyperedge, Hyperedge)> {
let node_count = nodes.len();
let mut spiders: Vec<Option<(Hyperedge, Hyperedge)>> = (0..node_count).map(|_| None).collect();
for &node in selected {
spiders[node.0].get_or_insert_with(|| {
(
Hyperedge {
sources: vec![],
targets: vec![node],
},
Hyperedge {
sources: vec![node],
targets: vec![],
},
)
});
}
for adjacency in adjacency {
for node in &mut adjacency.sources {
let original = *node;
if let Some((left_spider, _)) = spiders[original.0].as_mut() {
let occurrence = new_occurrence(nodes, original);
left_spider.targets.push(occurrence);
*node = occurrence;
}
}
for node in &mut adjacency.targets {
let original = *node;
if let Some((_, right_spider)) = spiders[original.0].as_mut() {
let occurrence = new_occurrence(nodes, original);
right_spider.sources.push(occurrence);
*node = occurrence;
}
}
}
for node in sources {
let original = *node;
if let Some((left_spider, _)) = spiders[original.0].as_mut() {
let occurrence = new_occurrence(nodes, original);
left_spider.sources.push(occurrence);
*node = occurrence;
}
}
for node in targets {
let original = *node;
if let Some((_, right_spider)) = spiders[original.0].as_mut() {
let occurrence = new_occurrence(nodes, original);
right_spider.targets.push(occurrence);
*node = occurrence;
}
}
spiders.into_iter().flatten().collect()
}