From 9c682cc76dbf1534e8ce420ac082c916d585836a Mon Sep 17 00:00:00 2001 From: Teun Huijben Date: Thu, 30 Jul 2026 13:50:49 -0700 Subject: [PATCH] fix add_edge_to_view leaving edge_id to -1 in SQL --- src/tracksdata/graph/_graph_view.py | 5 +++++ src/tracksdata/graph/_test/test_subgraph.py | 17 +++++++++++++++++ 2 files changed, 22 insertions(+) diff --git a/src/tracksdata/graph/_graph_view.py b/src/tracksdata/graph/_graph_view.py index 21d8a513..1be0cf55 100644 --- a/src/tracksdata/graph/_graph_view.py +++ b/src/tracksdata/graph/_graph_view.py @@ -798,6 +798,11 @@ def _add_edge_local(self, source_id: int, target_id: int) -> int: if c in df.columns ] attrs = df.filter(pl.col(DEFAULT_ATTR_KEYS.EDGE_ID) == parent_edge_id).drop(drop_cols).rows(named=True)[0] + # A rustworkx-family root shares its payload, EDGE_ID included; other backends + # hand back a plain dict, so stamp the root edge id explicitly (as `bulk_add_edges` + # does). Otherwise the local edge keeps the -1 placeholder and disagrees with + # `_edge_map_to_root`, so reads by edge id find no row. + attrs[DEFAULT_ATTR_KEYS.EDGE_ID] = parent_edge_id local_edge_id = self.rx_graph.add_edge( self._map_to_local(source_id), diff --git a/src/tracksdata/graph/_test/test_subgraph.py b/src/tracksdata/graph/_test/test_subgraph.py index f6993677..520bbeba 100644 --- a/src/tracksdata/graph/_test/test_subgraph.py +++ b/src/tracksdata/graph/_test/test_subgraph.py @@ -1648,6 +1648,23 @@ def test_add_edge_to_view_basic(graph_backend: BaseGraph) -> None: assert row["weight"].item() == 0.5 +def test_add_edge_to_view_keeps_edge_id(graph_backend: BaseGraph) -> None: + """A revived edge's local row must carry the root edge id, so reads by id work.""" + graph_backend.add_edge_attr_key("weight", pl.Float64) + + n0 = graph_backend.add_node({"t": 0}) + n1 = graph_backend.add_node({"t": 1}) + root_edge_id = graph_backend.add_edge(n0, n1, {"weight": 1.5}) + + view = graph_backend.filter().subgraph() + view.remove_edge_from_view(n0, n1) + view.add_edge_to_view(n0, n1) + + assert view.edge_id(n0, n1) == root_edge_id + assert view.edge_attrs(attr_keys=["weight"])[DEFAULT_ATTR_KEYS.EDGE_ID].to_list() == [root_edge_id] + assert view.edges[root_edge_id]["weight"] == 1.5 + + def test_add_edge_to_view_validation(graph_backend: BaseGraph) -> None: """Bad inputs raise ValueError; sync=False raises RuntimeError.""" graph_backend.add_node_attr_key("x", pl.Float64)