1use super::index::CsrIndex;
10use crate::csr::LocalNodeId;
11
12impl CsrIndex {
13 pub(crate) fn enable_weights(&mut self) {
15 self.has_weights = true;
16
17 if !self.out_targets.is_empty() {
19 self.out_weights = Some(vec![1.0; self.out_targets.len()].into());
20 }
21 if !self.in_targets.is_empty() {
22 self.in_weights = Some(vec![1.0; self.in_targets.len()].into());
23 }
24
25 for (buf, wbuf) in self
27 .buffer_out
28 .iter()
29 .zip(self.buffer_out_weights.iter_mut())
30 {
31 if wbuf.len() < buf.len() {
32 wbuf.resize(buf.len(), 1.0);
33 }
34 }
35 for (buf, wbuf) in self.buffer_in.iter().zip(self.buffer_in_weights.iter_mut()) {
36 if wbuf.len() < buf.len() {
37 wbuf.resize(buf.len(), 1.0);
38 }
39 }
40 }
41
42 pub fn has_weights(&self) -> bool {
44 self.has_weights
45 }
46
47 pub fn out_edge_weight(&self, edge_idx: usize) -> f64 {
52 self.out_weights
53 .as_ref()
54 .and_then(|ws| ws.get(edge_idx).copied())
55 .unwrap_or(1.0)
56 }
57
58 pub fn in_edge_weight(&self, edge_idx: usize) -> f64 {
60 self.in_weights
61 .as_ref()
62 .and_then(|ws| ws.get(edge_idx).copied())
63 .unwrap_or(1.0)
64 }
65
66 pub fn edge_weight(&self, src: &str, label: &str, dst: &str) -> Option<f64> {
71 let src_id = *self.node_to_id.get(src)?;
72 let dst_id = *self.node_to_id.get(dst)?;
73 let label_id = *self.label_to_id.get(label)?;
74
75 let idx = src_id as usize;
77 if idx + 1 < self.out_offsets.len() {
78 let start = self.out_offsets[idx] as usize;
79 let end = self.out_offsets[idx + 1] as usize;
80 for i in start..end {
81 let coll = self.out_collections.get(i).copied().unwrap_or(0);
82 if self.out_labels[i] == label_id
83 && self.out_targets[i] == dst_id
84 && !self
85 .deleted_edges
86 .contains(&(src_id, label_id, dst_id, coll))
87 {
88 return Some(self.out_edge_weight(i));
89 }
90 }
91 }
92
93 if idx < self.buffer_out.len() {
95 for (buf_idx, &(l, d)) in self.buffer_out[idx].iter().enumerate() {
96 if l == label_id && d == dst_id {
97 if self.has_weights {
98 return Some(
99 self.buffer_out_weights[idx]
100 .get(buf_idx)
101 .copied()
102 .unwrap_or(1.0),
103 );
104 }
105 return Some(1.0);
106 }
107 }
108 }
109
110 None
111 }
112
113 pub fn iter_out_edges_weighted(
118 &self,
119 node: LocalNodeId,
120 ) -> impl Iterator<Item = (u32, LocalNodeId, f64)> + '_ {
121 let node = node.raw(self.partition_tag);
122 let tag = self.partition_tag;
123 let idx = node as usize;
124
125 let dense_start = if idx + 1 < self.out_offsets.len() {
127 self.out_offsets[idx] as usize
128 } else {
129 0
130 };
131 let dense_end = if idx + 1 < self.out_offsets.len() {
132 self.out_offsets[idx + 1] as usize
133 } else {
134 0
135 };
136
137 let dense = (dense_start..dense_end)
138 .filter_map(move |i| {
139 let lid = self.out_labels[i];
140 let dst = self.out_targets[i];
141 let coll = self.out_collections.get(i).copied().unwrap_or(0);
142 if self.deleted_edges.contains(&(node, lid, dst, coll)) {
143 None
144 } else {
145 Some((lid, dst, self.out_edge_weight(i)))
146 }
147 })
148 .collect::<Vec<_>>();
149
150 let buffer = if idx < self.buffer_out.len() {
152 self.buffer_out[idx]
153 .iter()
154 .enumerate()
155 .map(|(buf_idx, &(lid, dst))| {
156 let w = if self.has_weights {
157 self.buffer_out_weights[idx]
158 .get(buf_idx)
159 .copied()
160 .unwrap_or(1.0)
161 } else {
162 1.0
163 };
164 (lid, dst, w)
165 })
166 .collect::<Vec<_>>()
167 } else {
168 Vec::new()
169 };
170
171 dense
172 .into_iter()
173 .chain(buffer)
174 .map(move |(lid, dst, w)| (lid, LocalNodeId::new(dst, tag), w))
175 }
176
177 pub fn iter_out_edges_weighted_raw(
181 &self,
182 node: u32,
183 ) -> impl Iterator<Item = (u32, u32, f64)> + '_ {
184 self.iter_out_edges_weighted(self.local(node))
185 .map(move |(lid, dst, w)| (lid, dst.raw(self.partition_tag), w))
186 }
187}
188
189pub fn extract_weight_from_properties(properties: &[u8]) -> f64 {
195 if properties.is_empty() {
196 return 1.0;
197 }
198 let Ok(val) = rmpv::decode::read_value(&mut &properties[..]) else {
199 return 1.0;
200 };
201 match val {
202 rmpv::Value::Map(entries) => {
203 for (k, v) in entries {
204 if let rmpv::Value::String(ref s) = k
205 && s.as_str() == Some("weight")
206 {
207 return match v {
208 rmpv::Value::F64(f) => f,
209 rmpv::Value::F32(f) => f as f64,
210 rmpv::Value::Integer(i) => i.as_f64().unwrap_or(1.0),
211 _ => 1.0,
212 };
213 }
214 }
215 1.0
216 }
217 _ => 1.0,
218 }
219}
220
221#[cfg(test)]
222mod tests {
223 use super::*;
224
225 #[test]
226 fn unweighted_graph_has_no_weight_arrays() {
227 let mut csr = CsrIndex::new();
228 csr.add_edge("a", "L", "b").unwrap();
229 assert!(!csr.has_weights());
230 assert!(csr.out_weights.is_none());
231 assert!(csr.in_weights.is_none());
232 }
233
234 #[test]
235 fn weighted_edge_basic() {
236 let mut csr = CsrIndex::new();
237 csr.add_edge_weighted("a", "ROAD", "b", 5.0).unwrap();
238 csr.add_edge_weighted("b", "ROAD", "c", 3.0).unwrap();
239 csr.add_edge("c", "ROAD", "d").unwrap(); assert!(csr.has_weights());
242 assert_eq!(csr.edge_weight("a", "ROAD", "b"), Some(5.0));
243 assert_eq!(csr.edge_weight("b", "ROAD", "c"), Some(3.0));
244 assert_eq!(csr.edge_weight("c", "ROAD", "d"), Some(1.0));
245 assert_eq!(csr.edge_weight("a", "ROAD", "c"), None);
246 }
247
248 #[test]
249 fn weighted_edges_survive_compaction() {
250 let mut csr = CsrIndex::new();
251 csr.add_edge_weighted("a", "R", "b", 2.5).unwrap();
252 csr.add_edge_weighted("b", "R", "c", 7.0).unwrap();
253 csr.add_edge("c", "R", "d").unwrap();
254
255 csr.compact().expect("no governor, cannot fail");
256
257 assert!(csr.has_weights());
258 assert_eq!(csr.edge_weight("a", "R", "b"), Some(2.5));
259 assert_eq!(csr.edge_weight("b", "R", "c"), Some(7.0));
260 assert_eq!(csr.edge_weight("c", "R", "d"), Some(1.0));
261 }
262
263 #[test]
264 fn weighted_edge_remove_keeps_weights_consistent() {
265 let mut csr = CsrIndex::new();
266 csr.add_edge_weighted("a", "R", "b", 2.0).unwrap();
267 csr.add_edge_weighted("a", "R", "c", 3.0).unwrap();
268 csr.add_edge_weighted("a", "R", "d", 4.0).unwrap();
269
270 csr.remove_edge("a", "R", "c");
271
272 assert_eq!(csr.edge_weight("a", "R", "b"), Some(2.0));
273 assert_eq!(csr.edge_weight("a", "R", "c"), None);
274 assert_eq!(csr.edge_weight("a", "R", "d"), Some(4.0));
275 }
276
277 #[test]
278 fn iter_out_edges_weighted_returns_weights() {
279 let mut csr = CsrIndex::new();
280 csr.add_edge_weighted("a", "R", "b", 2.5).unwrap();
281 csr.add_edge_weighted("a", "R", "c", 7.0).unwrap();
282 csr.compact().expect("no governor, cannot fail");
283
284 let edges: Vec<(u32, LocalNodeId, f64)> =
285 csr.iter_out_edges_weighted(csr.local(0)).collect();
286 assert_eq!(edges.len(), 2);
287
288 let weights: Vec<f64> = edges.iter().map(|e| e.2).collect();
289 assert!(weights.contains(&2.5));
290 assert!(weights.contains(&7.0));
291 }
292
293 #[test]
294 fn mixed_weighted_unweighted_backfill() {
295 let mut csr = CsrIndex::new();
296 csr.add_edge("a", "L", "b").unwrap();
297 csr.add_edge("b", "L", "c").unwrap();
298 assert!(!csr.has_weights());
299
300 csr.add_edge_weighted("c", "L", "d", 5.0).unwrap();
301 assert!(csr.has_weights());
302 assert_eq!(csr.edge_weight("a", "L", "b"), Some(1.0));
303 assert_eq!(csr.edge_weight("c", "L", "d"), Some(5.0));
304 }
305
306 #[test]
307 fn extract_weight_from_empty_properties() {
308 assert_eq!(extract_weight_from_properties(b""), 1.0);
309 }
310
311 #[test]
312 fn extract_weight_f64() {
313 let props = rmpv::Value::Map(vec![(
314 rmpv::Value::String("weight".into()),
315 rmpv::Value::F64(0.75),
316 )]);
317 let mut buf = Vec::new();
318 rmpv::encode::write_value(&mut buf, &props).unwrap();
319 assert_eq!(extract_weight_from_properties(&buf), 0.75);
320 }
321
322 #[test]
323 fn extract_weight_integer() {
324 let props = rmpv::Value::Map(vec![(
325 rmpv::Value::String("weight".into()),
326 rmpv::Value::Integer(rmpv::Integer::from(42)),
327 )]);
328 let mut buf = Vec::new();
329 rmpv::encode::write_value(&mut buf, &props).unwrap();
330 assert_eq!(extract_weight_from_properties(&buf), 42.0);
331 }
332
333 #[test]
334 fn extract_weight_missing_key() {
335 let props = rmpv::Value::Map(vec![(
336 rmpv::Value::String("color".into()),
337 rmpv::Value::String("red".into()),
338 )]);
339 let mut buf = Vec::new();
340 rmpv::encode::write_value(&mut buf, &props).unwrap();
341 assert_eq!(extract_weight_from_properties(&buf), 1.0);
342 }
343}