Skip to main content

helix_graph_algorithms/
transform.rs

1use std::collections::{BTreeMap, BTreeSet};
2
3use crate::{Attributes, Edge, Graph, GraphError, GraphKind, Node, NodeId};
4
5impl Graph {
6    /// Cheaply clone the immutable, reference-counted graph allocation.
7    pub fn copy(&self) -> Self {
8        self.clone()
9    }
10
11    /// Materialize directed traversal semantics without losing undirected
12    /// edges or graph metadata.
13    pub fn to_directed(&self) -> Result<Self, GraphError> {
14        if self.is_directed() {
15            return Ok(self.clone());
16        }
17        let kind = match self.kind() {
18            GraphKind::Graph => GraphKind::DiGraph,
19            GraphKind::MultiGraph => GraphKind::MultiDiGraph,
20            GraphKind::DiGraph | GraphKind::MultiDiGraph => unreachable!("checked above"),
21        };
22        let mut reverse_generations =
23            self.edges()
24                .iter()
25                .fold(BTreeMap::<String, u64>::new(), |mut generations, edge| {
26                    generations
27                        .entry(edge.id.stored_id().to_string())
28                        .and_modify(|generation| {
29                            *generation = (*generation).max(edge.id.reverse_generation());
30                        })
31                        .or_insert(edge.id.reverse_generation());
32                    generations
33                });
34        let mut edges = Vec::with_capacity(self.edge_count().saturating_mul(2));
35        for edge in self.edges() {
36            edges.push(edge.clone());
37            if edge.source != edge.target {
38                let generation = reverse_generations
39                    .get_mut(edge.id.stored_id())
40                    .expect("every edge stored ID was indexed");
41                *generation =
42                    generation
43                        .checked_add(1)
44                        .ok_or_else(|| GraphError::EdgeIdentityExhausted {
45                            stored_id: edge.id.stored_id().to_string(),
46                        })?;
47                let reverse_id =
48                    crate::EdgeId::synthesized_reverse(edge.id.stored_id(), *generation)
49                        .expect("incremented reverse generation is non-zero");
50                edges.push(Edge {
51                    id: reverse_id,
52                    source: edge.target.clone(),
53                    target: edge.source.clone(),
54                    ..edge.clone()
55                });
56            }
57        }
58        Graph::with_attributes(
59            kind,
60            self.attributes().clone(),
61            self.nodes().iter().cloned(),
62            edges,
63        )
64    }
65
66    /// Return an immutable graph with undirected traversal semantics while
67    /// preserving every original edge record and graph attribute.
68    pub fn to_undirected(&self) -> Result<Self, GraphError> {
69        if !self.is_directed() {
70            return Ok(self.clone());
71        }
72        let kind = match self.kind() {
73            GraphKind::MultiDiGraph => GraphKind::MultiGraph,
74            GraphKind::DiGraph => {
75                let mut endpoint_pairs = BTreeSet::new();
76                let has_parallel_pair = self.edges().iter().any(|edge| {
77                    let pair = if edge.source <= edge.target {
78                        (edge.source.clone(), edge.target.clone())
79                    } else {
80                        (edge.target.clone(), edge.source.clone())
81                    };
82                    !endpoint_pairs.insert(pair)
83                });
84                if has_parallel_pair {
85                    GraphKind::MultiGraph
86                } else {
87                    GraphKind::Graph
88                }
89            }
90            GraphKind::Graph | GraphKind::MultiGraph => unreachable!("checked above"),
91        };
92        Graph::with_attributes(
93            kind,
94            self.attributes().clone(),
95            self.nodes().iter().cloned(),
96            self.edges().iter().cloned(),
97        )
98    }
99
100    /// Materialize an induced subgraph containing exactly the requested nodes
101    /// and edges whose endpoints are both selected.
102    pub fn induced_subgraph(
103        &self,
104        node_ids: impl IntoIterator<Item = impl Into<NodeId>>,
105    ) -> Result<Self, GraphError> {
106        let selected = node_ids
107            .into_iter()
108            .map(Into::into)
109            .collect::<BTreeSet<_>>();
110        for node_id in &selected {
111            if !self.contains_node(node_id) {
112                return Err(GraphError::UnknownNode(node_id.clone()));
113            }
114        }
115        Graph::with_attributes(
116            self.kind(),
117            self.attributes().clone(),
118            self.nodes()
119                .iter()
120                .filter(|node| selected.contains(&node.id))
121                .cloned(),
122            self.edges()
123                .iter()
124                .filter(|edge| selected.contains(&edge.source) && selected.contains(&edge.target))
125                .cloned(),
126        )
127    }
128
129    /// Return a graph with external node IDs replaced and every endpoint
130    /// rewired. Distinct nodes may not merge implicitly.
131    pub fn relabel(&self, mapping: &BTreeMap<NodeId, NodeId>) -> Result<Self, GraphError> {
132        let mut targets = BTreeMap::<NodeId, NodeId>::new();
133        for node in self.nodes() {
134            let target = mapping
135                .get(&node.id)
136                .cloned()
137                .unwrap_or_else(|| node.id.clone());
138            if let Some(first) = targets.insert(target.clone(), node.id.clone()) {
139                return Err(GraphError::RelabelCollision {
140                    target,
141                    first,
142                    second: node.id.clone(),
143                });
144            }
145        }
146        let relabel = |node_id: &NodeId| {
147            mapping
148                .get(node_id)
149                .cloned()
150                .unwrap_or_else(|| node_id.clone())
151        };
152        Graph::with_attributes(
153            self.kind(),
154            self.attributes().clone(),
155            self.nodes().iter().map(|node| Node {
156                id: relabel(&node.id),
157                label: node.label.clone(),
158                attributes: node.attributes.clone(),
159            }),
160            self.edges().iter().map(|edge| Edge {
161                source: relabel(&edge.source),
162                target: relabel(&edge.target),
163                ..edge.clone()
164            }),
165        )
166    }
167
168    /// Compose two immutable graphs. Right-hand attributes take precedence.
169    pub fn compose(&self, right: &Self) -> Result<Self, GraphError> {
170        if self.kind() != right.kind() {
171            return Err(GraphError::KindMismatch);
172        }
173        let mut graph_attributes = self.attributes().clone();
174        graph_attributes.extend(right.attributes().clone());
175
176        let mut nodes = self
177            .nodes()
178            .iter()
179            .cloned()
180            .map(|node| (node.id.clone(), node))
181            .collect::<BTreeMap<_, _>>();
182        for right_node in right.nodes() {
183            match nodes.get_mut(&right_node.id) {
184                Some(left_node) => {
185                    if right_node.label.is_some() {
186                        left_node.label.clone_from(&right_node.label);
187                    }
188                    left_node.attributes.extend(right_node.attributes.clone());
189                }
190                None => {
191                    nodes.insert(right_node.id.clone(), right_node.clone());
192                }
193            }
194        }
195
196        let mut edges = self
197            .edges()
198            .iter()
199            .cloned()
200            .map(|edge| (edge.id.clone(), edge))
201            .collect::<BTreeMap<_, _>>();
202        for right_edge in right.edges() {
203            match edges.get_mut(&right_edge.id) {
204                Some(left_edge)
205                    if left_edge.source != right_edge.source
206                        || left_edge.target != right_edge.target =>
207                {
208                    return Err(GraphError::ConflictingEdge {
209                        edge_id: right_edge.id.clone(),
210                    });
211                }
212                Some(left_edge) => {
213                    if right_edge.graphify_key.is_some() {
214                        left_edge.graphify_key.clone_from(&right_edge.graphify_key);
215                    }
216                    if right_edge.label.is_some() {
217                        left_edge.label.clone_from(&right_edge.label);
218                    }
219                    if right_edge.weight.is_some() {
220                        left_edge.weight = right_edge.weight;
221                    }
222                    left_edge.attributes.extend(right_edge.attributes.clone());
223                }
224                None => {
225                    edges.insert(right_edge.id.clone(), right_edge.clone());
226                }
227            }
228        }
229        Graph::with_attributes(
230            self.kind(),
231            graph_attributes,
232            nodes.into_values(),
233            edges.into_values(),
234        )
235    }
236
237    /// Create an owned export DTO with independently mutable attributes.
238    pub fn export_parts(&self) -> (Attributes, Vec<Node>, Vec<Edge>) {
239        (
240            self.attributes().clone(),
241            self.nodes().to_vec(),
242            self.edges().to_vec(),
243        )
244    }
245}
246
247#[cfg(test)]
248mod tests {
249    use super::*;
250    use serde_json::json;
251
252    fn graph() -> Graph {
253        Graph::with_attributes(
254            GraphKind::DiGraph,
255            Attributes::from([("owner".to_string(), json!("left"))]),
256            [Node::new("a"), Node::new("b"), Node::new("c")],
257            [Edge::new("ab", "a", "b"), Edge::new("bc", "b", "c")],
258        )
259        .unwrap()
260    }
261
262    #[test]
263    fn induced_subgraph_keeps_only_internal_edges() {
264        let subgraph = graph()
265            .induced_subgraph(["a".to_string(), "b".to_string()])
266            .unwrap();
267        assert_eq!(subgraph.node_count(), 2);
268        assert_eq!(subgraph.edge_count(), 1);
269        assert!(subgraph.contains_edge("ab"));
270    }
271
272    #[test]
273    fn directed_conversion_duplicates_non_loops_structurally_and_preserves_attributes() {
274        let graph = Graph::with_attributes(
275            GraphKind::MultiGraph,
276            Attributes::from([("owner".to_string(), json!("graphify"))]),
277            [Node::new("a"), Node::new("b")],
278            [
279                Edge::new("edge", "a", "b")
280                    .with_graphify_key("user-key")
281                    .with_label("REL")
282                    .with_weight(2.0)
283                    .with_attributes(Attributes::from([("generation".to_string(), json!(3))])),
284                Edge::new("loop", "a", "a"),
285            ],
286        )
287        .unwrap();
288
289        let directed = graph.to_directed().unwrap();
290        assert_eq!(directed.kind(), GraphKind::MultiDiGraph);
291        assert_eq!(directed.edge_count(), 3);
292        assert_eq!(directed.attributes(), graph.attributes());
293        let reverse_id = crate::EdgeId::original("edge").reversed().unwrap();
294        let reverse = directed.edge(reverse_id.clone()).unwrap();
295        assert_eq!(reverse.source, NodeId::from("b"));
296        assert_eq!(reverse.target, NodeId::from("a"));
297        assert_eq!(reverse.graphify_key, Some("user-key".into()));
298        assert_eq!(reverse.attributes["generation"], json!(3));
299        assert!(!directed.contains_edge(crate::EdgeId::original("reverse(edge)")));
300        assert!(directed.contains_edge(reverse_id));
301        assert_eq!(directed.edge("loop").unwrap().source, "a");
302
303        let repeated = directed.to_undirected().unwrap().to_directed().unwrap();
304        assert_eq!(repeated.edge_count(), 5);
305        assert!(repeated.contains_edge(
306            crate::EdgeId::original("edge")
307                .reversed()
308                .unwrap()
309                .reversed()
310                .unwrap()
311        ));
312    }
313
314    #[test]
315    fn conversions_are_idempotent_and_reject_exhausted_reverse_identity() {
316        let directed = graph();
317        assert_eq!(directed.to_directed().unwrap(), directed);
318        let undirected = directed.to_undirected().unwrap();
319        assert_eq!(undirected.to_undirected().unwrap(), undirected);
320
321        let exhausted = Graph::new(
322            GraphKind::Graph,
323            [Node::new("a"), Node::new("b")],
324            [Edge {
325                id: crate::EdgeId::synthesized_reverse("edge", u64::MAX).unwrap(),
326                graphify_key: None,
327                source: NodeId::from("a"),
328                target: NodeId::from("b"),
329                label: None,
330                weight: None,
331                attributes: Attributes::new(),
332            }],
333        )
334        .unwrap();
335        assert_eq!(
336            exhausted.to_directed(),
337            Err(GraphError::EdgeIdentityExhausted {
338                stored_id: "edge".to_string(),
339            })
340        );
341    }
342
343    #[test]
344    fn undirected_conversion_promotes_only_lossy_digraphs_to_multigraphs() {
345        let simple = graph().to_undirected().unwrap();
346        assert_eq!(simple.kind(), GraphKind::Graph);
347        assert_eq!(simple.attributes()["owner"], json!("left"));
348
349        let reciprocal = Graph::new(
350            GraphKind::DiGraph,
351            [Node::new("a"), Node::new("b")],
352            [Edge::new("ab", "a", "b"), Edge::new("ba", "b", "a")],
353        )
354        .unwrap()
355        .to_undirected()
356        .unwrap();
357        assert_eq!(reciprocal.kind(), GraphKind::MultiGraph);
358        assert_eq!(reciprocal.edge_count(), 2);
359
360        let multi = Graph::new(
361            GraphKind::MultiDiGraph,
362            [Node::new("a"), Node::new("b")],
363            [Edge::new("ab", "a", "b")],
364        )
365        .unwrap()
366        .to_undirected()
367        .unwrap();
368        assert_eq!(multi.kind(), GraphKind::MultiGraph);
369    }
370
371    #[test]
372    fn relabel_rewires_edges_and_rejects_collisions() {
373        let graph = graph();
374        let relabeled = graph
375            .relabel(&BTreeMap::from([(NodeId::from("a"), NodeId::from("z"))]))
376            .unwrap();
377        assert_eq!(relabeled.edge("ab").unwrap().source, "z");
378        assert!(matches!(
379            graph.relabel(&BTreeMap::from([(NodeId::from("a"), NodeId::from("b"))])),
380            Err(GraphError::RelabelCollision { .. })
381        ));
382    }
383
384    #[test]
385    fn compose_applies_right_attribute_precedence() {
386        let left = graph();
387        let right = Graph::with_attributes(
388            GraphKind::DiGraph,
389            Attributes::from([("owner".to_string(), json!("right"))]),
390            [Node::new("a")
391                .with_attributes(Attributes::from([("name".to_string(), json!("Ada"))]))],
392            [],
393        )
394        .unwrap();
395        let composed = left.compose(&right).unwrap();
396        assert_eq!(composed.attributes()["owner"], json!("right"));
397        assert_eq!(composed.node("a").unwrap().attributes["name"], json!("Ada"));
398
399        let right = Graph::new(
400            GraphKind::DiGraph,
401            [
402                Node::new("a"),
403                Node::new("b"),
404                Node::new("c"),
405                Node::new("d").with_label("File"),
406            ],
407            [
408                Edge::new("ab", "a", "b")
409                    .with_graphify_key("right")
410                    .with_label("REL")
411                    .with_weight(2.0),
412                Edge::new("cd", "c", "d"),
413            ],
414        )
415        .unwrap();
416        let composed = left.compose(&right).unwrap();
417        assert_eq!(composed.node("d").unwrap().label.as_deref(), Some("File"));
418        assert_eq!(composed.edge("ab").unwrap().weight, Some(2.0));
419        assert!(composed.contains_edge("cd"));
420    }
421
422    #[test]
423    fn transformations_cover_copy_empty_unknown_direction_conflicts_and_export() {
424        let graph = graph();
425        assert_eq!(graph.copy(), graph);
426        assert!(!graph.to_undirected().unwrap().is_directed());
427        assert_eq!(
428            graph
429                .induced_subgraph(Vec::<NodeId>::new())
430                .unwrap()
431                .node_count(),
432            0
433        );
434        assert!(matches!(
435            graph.induced_subgraph(["missing".to_string()]),
436            Err(GraphError::UnknownNode(_))
437        ));
438
439        let undirected = graph.to_undirected().unwrap();
440        assert_eq!(graph.compose(&undirected), Err(GraphError::KindMismatch));
441        let conflicting = Graph::new(
442            GraphKind::DiGraph,
443            [Node::new("a"), Node::new("b"), Node::new("c")],
444            [Edge::new("ab", "b", "c")],
445        )
446        .unwrap();
447        assert!(matches!(
448            graph.compose(&conflicting),
449            Err(GraphError::ConflictingEdge { edge_id }) if edge_id == crate::EdgeId::from("ab")
450        ));
451
452        let relabeled = graph
453            .relabel(&BTreeMap::from([
454                (NodeId::from("a"), NodeId::from("b")),
455                (NodeId::from("b"), NodeId::from("a")),
456            ]))
457            .unwrap();
458        assert_eq!(relabeled.edge("ab").unwrap().source, "b");
459        let (mut attributes, mut nodes, mut edges) = graph.export_parts();
460        attributes.insert("owner".to_string(), json!("export"));
461        nodes[0].attributes.insert("local".to_string(), json!(true));
462        edges[0].attributes.insert("local".to_string(), json!(true));
463        assert_eq!(graph.attributes()["owner"], json!("left"));
464        assert!(!graph.nodes()[0].attributes.contains_key("local"));
465        assert!(!graph.edges()[0].attributes.contains_key("local"));
466    }
467}