Skip to main content

weavatrix_graph/payload/
directed.rs

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