diff --git a/grand/backends/_dataframe.py b/grand/backends/_dataframe.py index 6d94bf0..da524bb 100644 --- a/grand/backends/_dataframe.py +++ b/grand/backends/_dataframe.py @@ -417,7 +417,7 @@ def get_node_predecessors(self, u: Hashable, include_metadata: bool = False): return { ( r[self._edge_df_target_column] - if r[self._edge_df_target_column] != u + if r[self._edge_df_source_column] == u else r[self._edge_df_source_column] ): self._edge_as_dict(r) for _, r in self._edge_df[ @@ -450,9 +450,9 @@ def get_node_predecessors(self, u: Hashable, include_metadata: bool = False): return iter( [ ( - row[self._edge_df_source_column] - if row[self._edge_df_target_column] != u - else row[self._edge_df_target_column] + row[self._edge_df_target_column] + if row[self._edge_df_source_column] == u + else row[self._edge_df_source_column] ) for _, row in self._edge_df[ (self._edge_df[self._edge_df_target_column] == u) diff --git a/grand/backends/test_backends.py b/grand/backends/test_backends.py index 93368e2..77d1c1e 100644 --- a/grand/backends/test_backends.py +++ b/grand/backends/test_backends.py @@ -338,11 +338,6 @@ def test_undirected_predecessors_match_neighbors(self, backend): NetworkXBackend, "NetworkX undirected predecessor support is not implemented", ) - xfail_backend( - backend, - DataFrameBackend, - "APL #77: undirected predecessors return self", - ) b = backend(directed=False, **kwargs) b.add_edge("A", "B", {"weight": 1}) @@ -589,6 +584,17 @@ def test_edge_only_nodes_are_unique_members_and_counted_from_union(self): assert not backend.has_node("missing") assert backend.get_node_count() == 3 + def test_undirected_predecessors_return_opposite_endpoint_with_metadata(self): + backend = DataFrameBackend(directed=False) + backend.add_edge("A", "B", {"weight": 1}) + backend.add_edge("C", "A", {"weight": 2}) + + assert set(backend.get_node_predecessors("A")) == {"B", "C"} + assert backend.get_node_predecessors("A", include_metadata=True) == { + "B": {"weight": 1}, + "C": {"weight": 2}, + } + def test_networkx_can_ingest_edgelist_dataframe(): backend = NetworkXBackend(directed=True)