Skip to main content

weavatrix_graph/payload/
undirected.rs

1use crate::Vec;
2use crate::{
3    EdgeEndpoints, EdgeIndex, GraphError, IndexUndirectedGraphView, NodeIndex, Result,
4    UndirectedGraphView, UndirectedTopology,
5};
6use serde::{Deserialize, Deserializer, Serialize, de::Error as _};
7
8#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
9pub struct UndirectedPayloadGraph<NodePayload, EdgePayload> {
10    nodes: Vec<NodePayload>,
11    edges: Vec<EdgePayload>,
12    topology: UndirectedTopology,
13}
14
15impl<NodePayload, EdgePayload> UndirectedPayloadGraph<NodePayload, EdgePayload> {
16    /// Builds an undirected graph with arbitrary node and edge payloads.
17    ///
18    /// # Errors
19    ///
20    /// Returns an error for invalid compact endpoints or capacity overflow.
21    pub fn try_from_edges(
22        nodes: Vec<NodePayload>,
23        edges: impl IntoIterator<Item = (EdgeEndpoints, EdgePayload)>,
24    ) -> Result<Self> {
25        let (endpoints, edges): (Vec<_>, Vec<_>) = edges.into_iter().unzip();
26        let topology = UndirectedTopology::try_from_edges(nodes.len(), endpoints)?;
27        Ok(Self {
28            nodes,
29            edges,
30            topology,
31        })
32    }
33
34    /// Reattaches payload vectors to a prebuilt undirected topology.
35    ///
36    /// # Errors
37    ///
38    /// Returns an error when either payload count differs from the topology.
39    pub fn try_from_parts(
40        topology: UndirectedTopology,
41        nodes: Vec<NodePayload>,
42        edges: Vec<EdgePayload>,
43    ) -> Result<Self> {
44        validate_count("node", topology.node_count(), nodes.len())?;
45        validate_count("edge", topology.edge_count(), edges.len())?;
46        Ok(Self {
47            nodes,
48            edges,
49            topology,
50        })
51    }
52
53    #[must_use]
54    pub const fn topology(&self) -> &UndirectedTopology {
55        &self.topology
56    }
57
58    #[must_use]
59    pub fn nodes(&self) -> &[NodePayload] {
60        &self.nodes
61    }
62
63    #[must_use]
64    pub fn edges(&self) -> &[EdgePayload] {
65        &self.edges
66    }
67
68    #[must_use]
69    pub fn node(&self, index: NodeIndex) -> Option<&NodePayload> {
70        self.nodes.get(index.index())
71    }
72
73    #[must_use]
74    pub fn node_mut(&mut self, index: NodeIndex) -> Option<&mut NodePayload> {
75        self.nodes.get_mut(index.index())
76    }
77
78    #[must_use]
79    pub fn edge(&self, index: EdgeIndex) -> Option<&EdgePayload> {
80        self.edges.get(index.index())
81    }
82
83    #[must_use]
84    pub fn edge_mut(&mut self, index: EdgeIndex) -> Option<&mut EdgePayload> {
85        self.edges.get_mut(index.index())
86    }
87
88    #[must_use]
89    pub fn into_parts(self) -> (UndirectedTopology, Vec<NodePayload>, Vec<EdgePayload>) {
90        (self.topology, self.nodes, self.edges)
91    }
92}
93
94impl<NodePayload, EdgePayload> UndirectedGraphView
95    for UndirectedPayloadGraph<NodePayload, EdgePayload>
96{
97    type Node = NodeIndex;
98    type Edge = EdgeIndex;
99
100    fn node_count(&self) -> usize {
101        self.topology.node_count()
102    }
103
104    fn edge_count(&self) -> usize {
105        self.topology.edge_count()
106    }
107
108    fn contains_node(&self, node: Self::Node) -> bool {
109        self.topology.contains_node(node)
110    }
111
112    fn contains_edge(&self, edge: Self::Edge) -> bool {
113        self.topology.contains_edge(edge)
114    }
115
116    fn node_indices(&self) -> impl Iterator<Item = Self::Node> + '_ {
117        self.topology.node_indices()
118    }
119
120    fn edge_indices(&self) -> impl Iterator<Item = Self::Edge> + '_ {
121        self.topology.edge_indices()
122    }
123
124    fn edge_endpoints(&self, edge: Self::Edge) -> Option<EdgeEndpoints<Self::Node>> {
125        self.topology.edge_endpoints(edge)
126    }
127
128    fn incident_edges(
129        &self,
130        node: Self::Node,
131    ) -> impl DoubleEndedIterator<Item = Self::Edge> + ExactSizeIterator + '_ {
132        self.topology.incident_edges(node)
133    }
134}
135
136impl<NodePayload, EdgePayload> IndexUndirectedGraphView
137    for UndirectedPayloadGraph<NodePayload, EdgePayload>
138{
139    fn node_bound(&self) -> usize {
140        self.topology.node_bound()
141    }
142
143    fn edge_bound(&self) -> usize {
144        self.topology.edge_bound()
145    }
146
147    fn node_slot(node: Self::Node) -> usize {
148        node.index()
149    }
150
151    fn edge_slot(edge: Self::Edge) -> usize {
152        edge.index()
153    }
154
155    fn incident_edge_at(&self, node: Self::Node, offset: usize) -> Option<Self::Edge> {
156        self.topology.incident_edge_at(node, offset)
157    }
158}
159
160#[derive(Deserialize)]
161struct PayloadWire<NodePayload, EdgePayload> {
162    nodes: Vec<NodePayload>,
163    edges: Vec<EdgePayload>,
164    topology: UndirectedTopology,
165}
166
167impl<'de, NodePayload, EdgePayload> Deserialize<'de>
168    for UndirectedPayloadGraph<NodePayload, EdgePayload>
169where
170    NodePayload: Deserialize<'de>,
171    EdgePayload: Deserialize<'de>,
172{
173    fn deserialize<D>(deserializer: D) -> core::result::Result<Self, D::Error>
174    where
175        D: Deserializer<'de>,
176    {
177        let wire = PayloadWire::deserialize(deserializer)?;
178        Self::try_from_parts(wire.topology, wire.nodes, wire.edges).map_err(D::Error::custom)
179    }
180}
181
182fn validate_count(category: &'static str, expected: usize, actual: usize) -> Result<()> {
183    if expected == actual {
184        Ok(())
185    } else {
186        Err(GraphError::PayloadCountMismatch {
187            category,
188            expected,
189            actual,
190        })
191    }
192}