Skip to main content

nodedb_graph/csr/
weights.rs

1// SPDX-License-Identifier: Apache-2.0
2
3//! Edge weight management for the CSR index.
4//!
5//! Optional `f64` weight per edge stored in parallel arrays. `None` when the
6//! graph is entirely unweighted (zero memory overhead). Populated from the
7//! `"weight"` edge property at insertion time. Unweighted edges default to 1.0.
8
9use super::index::CsrIndex;
10use crate::csr::LocalNodeId;
11
12impl CsrIndex {
13    /// Enable weight tracking. Backfills existing buffer entries with 1.0.
14    pub(crate) fn enable_weights(&mut self) {
15        self.has_weights = true;
16
17        // Backfill existing dense arrays with 1.0.
18        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        // Backfill existing buffer entries with 1.0.
26        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    /// Whether this CSR has any weighted edges.
43    pub fn has_weights(&self) -> bool {
44        self.has_weights
45    }
46
47    /// Get the weight of the i-th outbound edge from a node (dense CSR only).
48    ///
49    /// `edge_idx` is the absolute index into `out_targets`/`out_weights`.
50    /// Returns 1.0 for unweighted graphs.
51    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    /// Get the weight of the i-th inbound edge to a node (dense CSR only).
59    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    /// Get the weight of a specific outbound edge from `src` to `dst` via `label`.
67    ///
68    /// Checks both dense and buffer. Returns 1.0 if the edge exists but has
69    /// no weight, or `None` if the edge doesn't exist.
70    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        // Check dense CSR.
76        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        // Check buffer.
94        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    /// Iterate outbound edges of a node with weights: `(label_id, dst_id, weight)`.
114    ///
115    /// Yields from both dense CSR and buffer, excluding deleted edges.
116    /// Weights are 1.0 for unweighted graphs.
117    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        // Dense edges with weights.
126        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        // Buffer edges with weights.
151        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    /// Raw u32 variant of `iter_out_edges_weighted`. In-partition
178    /// algorithm use only — see [`Self::iter_out_edges_raw`] for the
179    /// safety rationale.
180    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
189/// Extract the `"weight"` property from MessagePack-encoded edge properties.
190///
191/// Returns 1.0 if properties are empty, malformed, or don't contain a
192/// `"weight"` key. Handles F64, F32, and integer weight values; other
193/// numeric types default to 1.0.
194pub 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(); // unweighted → 1.0
240
241        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}