weavatrix_graph/payload/
directed.rs1use 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 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 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}