genmod / opengds /gds-2.22.0-bfs-depth.patch
pbordescnil's picture
Include BFS result type in OpenGDS patch
f18e168
Raw
History Blame Contribute Delete
18.9 kB
diff --git a/algo/src/main/java/org/neo4j/gds/paths/traverse/BFS.java b/algo/src/main/java/org/neo4j/gds/paths/traverse/BFS.java
index 1ef036b..7394194 100644
--- a/algo/src/main/java/org/neo4j/gds/paths/traverse/BFS.java
+++ b/algo/src/main/java/org/neo4j/gds/paths/traverse/BFS.java
@@ -181,6 +181,10 @@ public final class BFS extends Algorithm<HugeLongArray> {
@Override
public HugeLongArray compute() {
+ return computeWithDepths().nodeIds();
+ }
+
+ public BfsResult computeWithDepths() {
progressTracker.beginSubTask(graph.relationshipCount());
// This is used to read from `traversedNodes` in chunks, updated in `BFSTask`.
@@ -259,7 +263,10 @@ public final class BFS extends Algorithm<HugeLongArray> {
nodesLengthToRetain = targetFoundIndex.longValue() + 1;
}
- var result = traversedNodes.copyOf(nodesLengthToRetain);
+ var result = new BfsResult(
+ traversedNodes.copyOf(nodesLengthToRetain),
+ weights.copyOf(nodesLengthToRetain)
+ );
progressTracker.endSubTask();
return result;
diff --git a/algo/src/main/java/org/neo4j/gds/paths/traverse/BFSTask.java b/algo/src/main/java/org/neo4j/gds/paths/traverse/BFSTask.java
index 268dc92..dc24735 100644
--- a/algo/src/main/java/org/neo4j/gds/paths/traverse/BFSTask.java
+++ b/algo/src/main/java/org/neo4j/gds/paths/traverse/BFSTask.java
@@ -164,7 +164,15 @@ class BFSTask implements Runnable {
// In case `nodeId` is encountered in a later chunk,
// the if check will be false and not added to traversedNodes again.
if (!visited.getAndSet(nodeId)) {
+ long predecessorIndex = minimumChunk.get(nodeId);
+ long predecessorNodeId = traversedNodes.get(predecessorIndex);
+ double depth = aggregatorFunction.apply(
+ predecessorNodeId,
+ nodeId,
+ weights.get(predecessorIndex)
+ );
traversedNodes.set(index, nodeId);
+ weights.set(index, depth);
index++;
nodesTraversed++;
}
diff --git a/algo/src/main/java/org/neo4j/gds/paths/traverse/BfsResult.java b/algo/src/main/java/org/neo4j/gds/paths/traverse/BfsResult.java
new file mode 100644
index 0000000..f200329
--- /dev/null
+++ b/algo/src/main/java/org/neo4j/gds/paths/traverse/BfsResult.java
@@ -0,0 +1,20 @@
+/*
+ * Copyright (c) "Neo4j"
+ * Neo4j Sweden AB [http://neo4j.com]
+ *
+ * This file is part of Neo4j.
+ *
+ * Neo4j is free software: you can redistribute it and/or modify
+ * it under the terms of the GNU General Public License as published by
+ * the Free Software Foundation, either version 3 of the License, or
+ * (at your option) any later version.
+ */
+package org.neo4j.gds.paths.traverse;
+
+import org.neo4j.gds.collections.ha.HugeDoubleArray;
+import org.neo4j.gds.collections.ha.HugeLongArray;
+
+/**
+ * Nodes visited by BFS and their minimum depth, aligned by array index.
+ */
+public record BfsResult(HugeLongArray nodeIds, HugeDoubleArray depths) {}
diff --git a/algo/src/test/java/org/neo4j/gds/paths/traverse/BFSTest.java b/algo/src/test/java/org/neo4j/gds/paths/traverse/BFSTest.java
index 89b9074..6b4816a 100644
--- a/algo/src/test/java/org/neo4j/gds/paths/traverse/BFSTest.java
+++ b/algo/src/test/java/org/neo4j/gds/paths/traverse/BFSTest.java
@@ -174,6 +174,30 @@ class BFSTest {
);
}
+ @ParameterizedTest
+ @ValueSource(ints = {1, 4})
+ void shouldReturnMinimumDepthAlongsideEveryVisitedNode(int concurrency) {
+ long source = naturalGraph.toMappedNodeId("a");
+ var result = BFS.create(
+ naturalGraph,
+ source,
+ (s, t, w) -> Result.FOLLOW,
+ new OneHopAggregator(),
+ TraversalParameters.NO_MAX_DEPTH,
+ DefaultPool.INSTANCE,
+ new Concurrency(concurrency),
+ ProgressTracker.NULL_TRACKER,
+ TerminationFlag.RUNNING_TRUE
+ ).computeWithDepths();
+
+ assertThat(result.nodeIds().toArray()).isEqualTo(
+ Stream.of("a", "b", "c", "d", "e", "f", "g")
+ .mapToLong(naturalGraph::toMappedNodeId)
+ .toArray()
+ );
+ assertThat(result.depths().toArray()).containsExactly(0, 1, 1, 2, 3, 3, 4);
+ }
+
@ParameterizedTest
@ValueSource(ints = {1, 4})
void testBfsOnLoopGraph(int concurrency) {
diff --git a/applications/algorithms/path-finding/src/main/java/org/neo4j/gds/applications/algorithms/pathfinding/PathFindingAlgorithms.java b/applications/algorithms/path-finding/src/main/java/org/neo4j/gds/applications/algorithms/pathfinding/PathFindingAlgorithms.java
index 6313df2..edf59a3 100644
--- a/applications/algorithms/path-finding/src/main/java/org/neo4j/gds/applications/algorithms/pathfinding/PathFindingAlgorithms.java
+++ b/applications/algorithms/path-finding/src/main/java/org/neo4j/gds/applications/algorithms/pathfinding/PathFindingAlgorithms.java
@@ -44,8 +44,10 @@ import org.neo4j.gds.paths.dijkstra.DijkstraFactory;
import org.neo4j.gds.paths.dijkstra.DijkstraSourceTargetParameters;
import org.neo4j.gds.paths.dijkstra.PathFindingResult;
import org.neo4j.gds.paths.traverse.BFS;
+import org.neo4j.gds.paths.traverse.BfsResult;
import org.neo4j.gds.paths.traverse.DFS;
import org.neo4j.gds.paths.traverse.ExitAndAggregation;
+import org.neo4j.gds.paths.traverse.OneHopAggregator;
import org.neo4j.gds.paths.yens.Yens;
import org.neo4j.gds.paths.yens.YensParameters;
import org.neo4j.gds.pcst.PCSTParameters;
@@ -137,6 +139,30 @@ public class PathFindingAlgorithms {
return bfs.compute();
}
+ BfsResult breadthFirstSearchWithDepths(
+ Graph graph,
+ TraversalParameters parameters,
+ ProgressTracker progressTracker,
+ TerminationFlag terminationFlag
+ ) {
+ var exitAndAggregationConditions = ExitAndAggregation.create(graph, parameters);
+ var mappedStartNodeId = graph.toMappedNodeId(parameters.sourceNode());
+
+ var bfs = BFS.create(
+ graph,
+ mappedStartNodeId,
+ exitAndAggregationConditions.exitFunction(),
+ new OneHopAggregator(),
+ parameters.maxDepth(),
+ DefaultPool.INSTANCE,
+ parameters.concurrency(),
+ progressTracker,
+ terminationFlag
+ );
+
+ return bfs.computeWithDepths();
+ }
+
public PathFindingResult deltaStepping(
Graph graph,
DeltaSteppingParameters parameters,
diff --git a/applications/algorithms/path-finding/src/main/java/org/neo4j/gds/applications/algorithms/pathfinding/PathFindingAlgorithmsBusinessFacade.java b/applications/algorithms/path-finding/src/main/java/org/neo4j/gds/applications/algorithms/pathfinding/PathFindingAlgorithmsBusinessFacade.java
index 0753c65..ebcb9cc 100644
--- a/applications/algorithms/path-finding/src/main/java/org/neo4j/gds/applications/algorithms/pathfinding/PathFindingAlgorithmsBusinessFacade.java
+++ b/applications/algorithms/path-finding/src/main/java/org/neo4j/gds/applications/algorithms/pathfinding/PathFindingAlgorithmsBusinessFacade.java
@@ -51,6 +51,7 @@ import org.neo4j.gds.paths.dijkstra.config.DijkstraBaseConfig;
import org.neo4j.gds.paths.dijkstra.config.DijkstraSourceTargetsBaseConfig;
import org.neo4j.gds.paths.traverse.BFSProgressTask;
import org.neo4j.gds.paths.traverse.BfsBaseConfig;
+import org.neo4j.gds.paths.traverse.BfsResult;
import org.neo4j.gds.paths.traverse.DFSProgressTask;
import org.neo4j.gds.paths.traverse.DfsBaseConfig;
import org.neo4j.gds.paths.yens.YensProgressTask;
@@ -141,6 +142,21 @@ public class PathFindingAlgorithmsBusinessFacade {
);
}
+ BfsResult breadthFirstSearchWithDepths(Graph graph, BfsBaseConfig configuration) {
+ var progressTracker = createProgressTracker(BFSProgressTask.create(), configuration);
+
+ return algorithmMachinery.getResult(
+ () -> algorithms.breadthFirstSearchWithDepths(
+ graph,
+ configuration.toParameters(),
+ progressTracker,
+ requestScopedDependencies.terminationFlag()
+ ),
+ progressTracker,
+ configuration.concurrency()
+ );
+ }
+
public PathFindingResult deltaStepping(Graph graph, AllShortestPathsDeltaBaseConfig configuration) {
var progressTracker = createProgressTracker(DeltaSteppingProgressTask.create(), configuration);
diff --git a/applications/algorithms/path-finding/src/main/java/org/neo4j/gds/applications/algorithms/pathfinding/PathFindingAlgorithmsStreamModeBusinessFacade.java b/applications/algorithms/path-finding/src/main/java/org/neo4j/gds/applications/algorithms/pathfinding/PathFindingAlgorithmsStreamModeBusinessFacade.java
index b14fde5..28f9141 100644
--- a/applications/algorithms/path-finding/src/main/java/org/neo4j/gds/applications/algorithms/pathfinding/PathFindingAlgorithmsStreamModeBusinessFacade.java
+++ b/applications/algorithms/path-finding/src/main/java/org/neo4j/gds/applications/algorithms/pathfinding/PathFindingAlgorithmsStreamModeBusinessFacade.java
@@ -37,6 +37,7 @@ import org.neo4j.gds.paths.dijkstra.PathFindingResult;
import org.neo4j.gds.paths.dijkstra.config.AllShortestPathsDijkstraStreamConfig;
import org.neo4j.gds.paths.dijkstra.config.ShortestPathDijkstraStreamConfig;
import org.neo4j.gds.paths.traverse.BfsStreamConfig;
+import org.neo4j.gds.paths.traverse.BfsResult;
import org.neo4j.gds.paths.traverse.DfsStreamConfig;
import org.neo4j.gds.paths.yens.config.ShortestPathYensStreamConfig;
import org.neo4j.gds.pcst.PCSTStreamConfig;
@@ -117,14 +118,14 @@ public class PathFindingAlgorithmsStreamModeBusinessFacade {
public <RESULT> Stream<RESULT> breadthFirstSearch(
GraphName graphName,
BfsStreamConfig configuration,
- StreamResultBuilder<HugeLongArray, RESULT> resultBuilder
+ StreamResultBuilder<BfsResult, RESULT> resultBuilder
) {
return convenience.processRegularAlgorithmInStreamMode(
graphName,
configuration,
BFS,
estimation::breadthFirstSearch,
- (graph, __) -> algorithms.breadthFirstSearch(graph, configuration),
+ (graph, __) -> algorithms.breadthFirstSearchWithDepths(graph, configuration),
resultBuilder
);
}
diff --git a/proc/path-finding/src/test/java/org/neo4j/gds/paths/traverse/BfsStreamProcTest.java b/proc/path-finding/src/test/java/org/neo4j/gds/paths/traverse/BfsStreamProcTest.java
index 8c3d35a..4c641a6 100644
--- a/proc/path-finding/src/test/java/org/neo4j/gds/paths/traverse/BfsStreamProcTest.java
+++ b/proc/path-finding/src/test/java/org/neo4j/gds/paths/traverse/BfsStreamProcTest.java
@@ -125,7 +125,7 @@ class BfsStreamProcTest extends BaseProcTest {
.streamMode()
.addParameter("sourceNode", source)
.addParameter("maxDepth", 2)
- .yields("sourceNode", "nodeIds");
+ .yields("sourceNode", "nodeIds", "depths");
runQueryWithRowConsumer(query, row -> {
assertEquals(row.getNumber("sourceNode").longValue(), source);
@@ -133,6 +133,7 @@ class BfsStreamProcTest extends BaseProcTest {
assertThat(nodeIds).isEqualTo(
Stream.of("a", "b", "c", "d").map(idFunction::of).collect(Collectors.toList())
);
+ assertThat(row.get("depths")).isEqualTo(List.of(0L, 1L, 1L, 2L));
});
}
@@ -175,7 +176,7 @@ class BfsStreamProcTest extends BaseProcTest {
.algo("bfs")
.streamMode()
.addParameter("sourceNode", source)
- .yields("sourceNode", "nodeIds");
+ .yields("sourceNode", "nodeIds", "depths");
runQueryWithRowConsumer(query, row -> {
assertThat(row.getNumber("sourceNode").longValue()).isEqualTo(source);
@@ -188,6 +189,7 @@ class BfsStreamProcTest extends BaseProcTest {
.map(idFunction::of)
.collect(Collectors.toList())
);
+ assertThat(row.get("depths")).isEqualTo(List.of(0L, 1L, 1L, 2L, 3L, 3L, 4L));
});
}
diff --git a/proc/sysinfo/src/main/java/org/neo4j/gds/SysInfoProc.java b/proc/sysinfo/src/main/java/org/neo4j/gds/SysInfoProc.java
index 9b729b5..bd0ca5f 100644
--- a/proc/sysinfo/src/main/java/org/neo4j/gds/SysInfoProc.java
+++ b/proc/sysinfo/src/main/java/org/neo4j/gds/SysInfoProc.java
@@ -69,6 +69,27 @@ public class SysInfoProc {
return debugValues(properties, Runtime.getRuntime(), config);
}
+ @Procedure("gds.debug.arrow")
+ @SystemProcedure
+ @Description("Returns the status of the unavailable Arrow server in OpenGDS")
+ public Stream<ArrowInfo> arrow() {
+ return Stream.of(new ArrowInfo("", false, false, java.util.List.of()));
+ }
+
+ public static final class ArrowInfo {
+ public final String listenAddress;
+ public final boolean enabled;
+ public final boolean running;
+ public final java.util.List<String> versions;
+
+ private ArrowInfo(String listenAddress, boolean enabled, boolean running, java.util.List<String> versions) {
+ this.listenAddress = listenAddress;
+ this.enabled = enabled;
+ this.running = running;
+ this.versions = versions;
+ }
+ }
+
public static final class DebugValue {
public final String key;
public final Object value;
diff --git a/proc/sysinfo/src/test/java/org/neo4j/gds/SysInfoProcTest.java b/proc/sysinfo/src/test/java/org/neo4j/gds/SysInfoProcTest.java
index f1a3ddb..a9f44c5 100644
--- a/proc/sysinfo/src/test/java/org/neo4j/gds/SysInfoProcTest.java
+++ b/proc/sysinfo/src/test/java/org/neo4j/gds/SysInfoProcTest.java
@@ -138,4 +138,19 @@ class SysInfoProcTest extends BaseProcTest {
);
assertThat(result).containsExactly(BuildInfoProperties.get().gdsVersion());
}
+
+ @Test
+ void shouldReportArrowAsDisabledForClientCompatibility() {
+ var result = runQuery(
+ "CALL gds.debug.arrow() YIELD listenAddress, enabled, running, versions "
+ + "RETURN listenAddress, enabled, running, versions",
+ cypherResult -> cypherResult.stream().findFirst().orElseThrow()
+ );
+
+ assertThat(result)
+ .containsEntry("listenAddress", "")
+ .containsEntry("enabled", false)
+ .containsEntry("running", false)
+ .containsEntry("versions", List.of());
+ }
}
diff --git a/procedures/algorithms-facade/src/main/java/org/neo4j/gds/procedures/algorithms/pathfinding/BfsStreamResultBuilder.java b/procedures/algorithms-facade/src/main/java/org/neo4j/gds/procedures/algorithms/pathfinding/BfsStreamResultBuilder.java
index 06e5426..b8e48f8 100644
--- a/procedures/algorithms-facade/src/main/java/org/neo4j/gds/procedures/algorithms/pathfinding/BfsStreamResultBuilder.java
+++ b/procedures/algorithms-facade/src/main/java/org/neo4j/gds/procedures/algorithms/pathfinding/BfsStreamResultBuilder.java
@@ -23,16 +23,17 @@ import org.neo4j.gds.api.Graph;
import org.neo4j.gds.api.GraphStore;
import org.neo4j.gds.api.NodeLookup;
import org.neo4j.gds.applications.algorithms.machinery.StreamResultBuilder;
-import org.neo4j.gds.collections.ha.HugeLongArray;
+import org.neo4j.gds.paths.traverse.BfsResult;
import org.neo4j.gds.paths.traverse.BfsStreamConfig;
import org.neo4j.graphdb.RelationshipType;
+import java.util.Arrays;
import java.util.Optional;
import java.util.stream.Stream;
import static org.neo4j.gds.procedures.algorithms.pathfinding.TraversalStreamResult.RELATIONSHIP_TYPE_NAME;
-class BfsStreamResultBuilder implements StreamResultBuilder<HugeLongArray, TraversalStreamResult> {
+class BfsStreamResultBuilder implements StreamResultBuilder<BfsResult, TraversalStreamResult> {
private final NodeLookup nodeLookup;
private final boolean pathRequested;
private final BfsStreamConfig configuration;
@@ -47,18 +48,30 @@ class BfsStreamResultBuilder implements StreamResultBuilder<HugeLongArray, Trave
public Stream<TraversalStreamResult> build(
Graph graph,
GraphStore graphStore,
- Optional<HugeLongArray> result
+ Optional<BfsResult> result
) {
//noinspection OptionalIsPresent
if (result.isEmpty()) return Stream.empty();
- return TraverseStreamComputationResultConsumer.consume(
- configuration.sourceNode(),
- result.get(),
- graph::toOriginalNodeId,
- TraversalStreamResult::new,
- PathFactoryFacade.create(pathRequested, nodeLookup, graphStore.capabilities().canWriteToLocalDatabase()),
- RelationshipType.withName(RELATIONSHIP_TYPE_NAME)
+ var bfsResult = result.get();
+ var nodeList = Arrays.stream(bfsResult.nodeIds().toArray())
+ .map(graph::toOriginalNodeId)
+ .boxed()
+ .toList();
+ var depthList = Arrays.stream(bfsResult.depths().toArray())
+ .mapToObj(depth -> (long) depth)
+ .toList();
+ var path = PathFactoryFacade
+ .create(pathRequested, nodeLookup, graphStore.capabilities().canWriteToLocalDatabase())
+ .createPath(nodeList, RelationshipType.withName(RELATIONSHIP_TYPE_NAME));
+
+ return Stream.of(
+ new TraversalStreamResult(
+ configuration.sourceNode(),
+ nodeList,
+ depthList,
+ path
+ )
);
}
}
diff --git a/procedures/facade-api/path-finding-facade-api/src/main/java/org/neo4j/gds/procedures/algorithms/pathfinding/TraversalStreamResult.java b/procedures/facade-api/path-finding-facade-api/src/main/java/org/neo4j/gds/procedures/algorithms/pathfinding/TraversalStreamResult.java
index 35a4c21..8a64fd8 100644
--- a/procedures/facade-api/path-finding-facade-api/src/main/java/org/neo4j/gds/procedures/algorithms/pathfinding/TraversalStreamResult.java
+++ b/procedures/facade-api/path-finding-facade-api/src/main/java/org/neo4j/gds/procedures/algorithms/pathfinding/TraversalStreamResult.java
@@ -23,6 +23,10 @@ import org.neo4j.graphdb.Path;
import java.util.List;
-public record TraversalStreamResult(long sourceNode, List<Long> nodeIds, Path path) {
+public record TraversalStreamResult(long sourceNode, List<Long> nodeIds, List<Long> depths, Path path) {
public static final String RELATIONSHIP_TYPE_NAME = "NEXT";
+
+ public TraversalStreamResult(long sourceNode, List<Long> nodeIds, Path path) {
+ this(sourceNode, nodeIds, List.of(), path);
+ }
}