Skip to content

Commit 1e702c2

Browse files
feat: SparkConnect support (#506)
* wip * wip * wip * wip * The first working version * WIP * Working version? * Fix tests * Fix tests * Fix CI typo * Fix typo in CI * Fix wget's verbose + GHA bug * Stop connect server * An attempt to fix a bug in GHA with a non-stopping tests * Maybe grpc/grpc#38290? * Fix broken stop-cript * Ignore errors in clean-up * Verbosity in ci tests * Typo * Fix merge-artifacts * Fix merge artifacts * Apply pre-commit rules * Add the missing method * Restore accidently deleted part of CI * Typo * Fixes from comments * Pin the pyspark version <4.0 and re-generate lock
1 parent 9cdad91 commit 1e702c2

30 files changed

Lines changed: 3380 additions & 938 deletions

‎.github/workflows/python-ci.yml‎

Lines changed: 14 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@ jobs:
88
include:
99
- spark-version: 3.5.4
1010
scala-version: 2.12.18
11-
python-version: 3.9.19
11+
python-version: 3.10.6
1212
runs-on: ubuntu-22.04
1313
env:
1414
# define Java options for both official sbt and sbt-extras
@@ -27,8 +27,6 @@ jobs:
2727
path: |
2828
~/.ivy2/cache
2929
key: sbt-ivy-cache-spark-${{ matrix.spark-version}}-scala-${{ matrix.scala-version }}
30-
- name: Assembly
31-
run: build/sbt -v ++${{ matrix.scala-version }} -Dspark.version=${{ matrix.spark-version }} "set test in assembly := {}" assembly
3230
- uses: actions/setup-python@v4
3331
with:
3432
python-version: ${{ matrix.python-version }}
@@ -42,16 +40,24 @@ jobs:
4240
- name: Build Python package and its dependencies
4341
working-directory: ./python
4442
run: |
45-
poetry build
46-
poetry install --with dev
47-
- name: Code Style
43+
poetry install --with=dev
44+
- name: Code style
4845
working-directory: ./python
4946
run: |
5047
poetry run python -m black --check graphframes
5148
poetry run python -m flake8 graphframes
5249
poetry run python -m isort --check graphframes
50+
5351
- name: Test
5452
working-directory: ./python
5553
run: |
56-
export SPARK_HOME=$(poetry run python -c "import os; from importlib.util import find_spec; spec = find_spec('pyspark'); print(os.path.join(os.path.dirname(spec.origin)))")
57-
./run-tests.sh
54+
poetry run python -m pytest
55+
56+
- name: Test SparkConnect
57+
env:
58+
SPARK_CONNECT_MODE_ENABLED: 1
59+
working-directory: ./python
60+
run: |
61+
poetry run python dev/run_connect.py
62+
poetry run python -m pytest
63+
poetry run python dev/stop_connect.py

‎.gitignore‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -36,3 +36,9 @@ python/graphframes.egg-info
3636
python/graphframes/tutorials/data
3737
python/docs/_build
3838
python/docs/_site
39+
40+
# JAR that is build during the installation
41+
python/graphframes/resources/*
42+
43+
# tmp data for spark connect
44+
tmp/*

‎buf.gen.yaml‎

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,12 @@
1+
version: v2
2+
managed:
3+
enabled: true
4+
5+
plugins:
6+
# Python API
7+
- remote: buf.build/grpc/python:v1.64.2
8+
out: python/graphframes/connect/proto
9+
- remote: buf.build/protocolbuffers/python:v27.1
10+
out: python/graphframes/connect/proto
11+
- remote: buf.build/protocolbuffers/pyi
12+
out: python/graphframes/connect/proto

‎buf.yaml‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,3 @@
1+
version: v2
2+
modules:
3+
- path: graphframes-connect/src/main/protobuf

‎build.sbt‎

Lines changed: 64 additions & 48 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,13 @@
11
import ReleaseTransformations.*
2+
import sbt.Credentials
3+
import sbt.Keys.credentials
24

35
lazy val sparkVer = sys.props.getOrElse("spark.version", "3.5.4")
46
lazy val sparkBranch = sparkVer.substring(0, 3)
57
lazy val defaultScalaVer = sparkBranch match {
68
case "3.5" => "2.12.18"
9+
case "3.4" => "2.12.17"
10+
case "3.3" => "2.12.15"
711
case _ => throw new IllegalArgumentException(s"Unsupported Spark version: $sparkVer.")
812
}
913
lazy val scalaVer = sys.props.getOrElse("scala.version", defaultScalaVer)
@@ -20,56 +24,48 @@ ThisBuild / scalaVersion := scalaVer
2024
ThisBuild / organization := "org.graphframes"
2125
ThisBuild / crossScalaVersions := Seq("2.12.18", "2.13.8")
2226

27+
lazy val commonSetting = Seq(
28+
libraryDependencies ++= Seq(
29+
"org.apache.spark" %% "spark-graphx" % sparkVer % "provided" cross CrossVersion.for3Use2_13,
30+
"org.apache.spark" %% "spark-sql" % sparkVer % "provided" cross CrossVersion.for3Use2_13,
31+
"org.apache.spark" %% "spark-mllib" % sparkVer % "provided" cross CrossVersion.for3Use2_13,
32+
"org.slf4j" % "slf4j-api" % "2.0.16",
33+
"org.scalatest" %% "scalatest" % defaultScalaTestVer % Test,
34+
"com.github.zafarkhaja" % "java-semver" % "0.10.2" % Test),
35+
credentials += Credentials(Path.userHome / ".ivy2" / ".sbtcredentials"),
36+
licenses := Seq("Apache-2.0" -> url("https://opensource.org/licenses/Apache-2.0")),
37+
Compile / scalacOptions ++= Seq("-deprecation", "-feature"),
38+
Compile / doc / scalacOptions ++= Seq(
39+
"-groups",
40+
"-implicits",
41+
"-skip-packages",
42+
Seq("org.apache.spark").mkString(":")),
43+
Test / doc / scalacOptions ++= Seq("-groups", "-implicits"),
44+
45+
// Test settings
46+
Test / fork := true,
47+
Test / parallelExecution := false,
48+
Test / javaOptions ++= Seq(
49+
"-XX:+IgnoreUnrecognizedVMOptions",
50+
"-Xmx2048m",
51+
"-XX:ReservedCodeCacheSize=384m",
52+
"-XX:MaxMetaspaceSize=384m",
53+
"--add-opens=java.base/sun.nio.ch=ALL-UNNAMED",
54+
"--add-opens=java.base/java.lang=ALL-UNNAMED",
55+
"--add-opens=java.base/java.nio=ALL-UNNAMED",
56+
"--add-opens=java.base/java.lang.invoke=ALL-UNNAMED",
57+
"--add-opens=java.base/java.util=ALL-UNNAMED"),
58+
credentials += Credentials(Path.userHome / ".ivy2" / ".sbtcredentials"))
59+
2360
lazy val root = (project in file("."))
2461
.settings(
62+
commonSetting,
2563
name := "graphframes",
26-
27-
// Replace spark-packages plugin functionality with explicit dependencies
28-
libraryDependencies ++= Seq(
29-
"org.apache.spark" %% "spark-graphx" % sparkVer % "provided" cross CrossVersion.for3Use2_13,
30-
"org.apache.spark" %% "spark-sql" % sparkVer % "provided" cross CrossVersion.for3Use2_13,
31-
"org.apache.spark" %% "spark-mllib" % sparkVer % "provided" cross CrossVersion.for3Use2_13,
32-
"org.slf4j" % "slf4j-api" % "1.7.16",
33-
"org.scalatest" %% "scalatest" % defaultScalaTestVer % Test,
34-
"com.github.zafarkhaja" % "java-semver" % "0.9.0" % Test
35-
),
36-
37-
licenses := Seq("Apache-2.0" -> url("http://opensource.org/licenses/Apache-2.0")),
38-
39-
// Modern way to set Scala options
4064
Compile / scalacOptions ++= Seq("-deprecation", "-feature"),
4165

42-
Compile / doc / scalacOptions ++= Seq(
43-
"-groups",
44-
"-implicits",
45-
"-skip-packages", Seq("org.apache.spark").mkString(":")
46-
),
47-
48-
Test / doc / scalacOptions ++= Seq("-groups", "-implicits"),
49-
50-
// Test settings
51-
Test / fork := true,
52-
Test / parallelExecution := false,
53-
54-
Test / javaOptions ++= Seq(
55-
"-XX:+IgnoreUnrecognizedVMOptions",
56-
"-Xmx2048m",
57-
"-XX:ReservedCodeCacheSize=384m",
58-
"-XX:MaxMetaspaceSize=384m",
59-
"--add-opens=java.base/sun.nio.ch=ALL-UNNAMED",
60-
"--add-opens=java.base/java.lang=ALL-UNNAMED",
61-
"--add-opens=java.base/java.nio=ALL-UNNAMED",
62-
"--add-opens=java.base/java.lang.invoke=ALL-UNNAMED",
63-
"--add-opens=java.base/java.util=ALL-UNNAMED",
64-
),
65-
6666
// Global settings
67-
Global / concurrentRestrictions := Seq(
68-
Tags.limitAll(1)
69-
),
70-
67+
Global / concurrentRestrictions := Seq(Tags.limitAll(1)),
7168
autoAPIMappings := true,
72-
7369
coverageHighlighting := false,
7470

7571
// Release settings
@@ -79,8 +75,7 @@ lazy val root = (project in file("."))
7975
commitReleaseVersion,
8076
tagRelease,
8177
setNextVersion,
82-
commitNextVersion
83-
),
78+
commitNextVersion),
8479

8580
// Assembly settings
8681
assembly / test := {}, // No tests in assembly
@@ -90,7 +85,28 @@ lazy val root = (project in file("."))
9085
case x =>
9186
val oldStrategy = (assembly / assemblyMergeStrategy).value
9287
oldStrategy(x)
93-
},
88+
})
9489

95-
credentials += Credentials(Path.userHome / ".ivy2" / ".sbtcredentials")
96-
)
90+
lazy val connect = (project in file("graphframes-connect"))
91+
.dependsOn(root)
92+
.settings(
93+
commonSetting,
94+
name := "graphframes-connect",
95+
Compile / PB.targets := Seq(PB.gens.java -> (Compile / sourceManaged).value),
96+
Compile / PB.includePaths ++= Seq(file("src/main/protobuf")),
97+
PB.protocVersion := "3.23.4", // Spark 3.5 branch
98+
libraryDependencies ++= Seq(
99+
"org.apache.spark" %% "spark-connect" % sparkVer % "provided" cross CrossVersion.for3Use2_13),
100+
101+
// Assembly and shading
102+
assembly / test := {},
103+
assembly / assemblyShadeRules := Seq(
104+
ShadeRule.rename("com.google.protobuf.**" -> "org.sparkproject.connect.protobuf.@1").inAll),
105+
assembly / assemblyMergeStrategy := {
106+
case PathList("META-INF", xs @ _*) => MergeStrategy.discard
107+
case x if x.endsWith("module-info.class") => MergeStrategy.discard
108+
case x =>
109+
val oldStrategy = (assembly / assemblyMergeStrategy).value
110+
oldStrategy(x)
111+
}
112+
)
Lines changed: 137 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,137 @@
1+
syntax = 'proto3';
2+
3+
package org.graphframes.connect.proto;
4+
5+
option java_multiple_files = true;
6+
option java_package = "org.graphframes.connect.proto";
7+
option java_generate_equals_and_hash = true;
8+
option optimize_for=SPEED;
9+
10+
11+
message GraphFramesAPI {
12+
bytes vertices = 1;
13+
bytes edges = 2;
14+
oneof method {
15+
AggregateMessages aggregate_messages = 3;
16+
BFS bfs = 4;
17+
ConnectedComponents connected_components = 5;
18+
DropIsolatedVertices drop_isolated_vertices = 6;
19+
FilterEdges filter_edges = 7;
20+
FilterVertices filter_vertices = 8;
21+
Find find = 9;
22+
LabelPropagation label_propagation = 10;
23+
PageRank page_rank = 11;
24+
ParallelPersonalizedPageRank parallel_personalized_page_rank = 12;
25+
PowerIterationClustering power_iteration_clustering = 13;
26+
Pregel pregel = 14;
27+
ShortestPaths shortest_paths = 15;
28+
StronglyConnectedComponents strongly_connected_components = 16;
29+
SVDPlusPlus svd_plus_plus = 17;
30+
TriangleCount triangle_count = 18;
31+
Triplets triplets = 19;
32+
}
33+
}
34+
35+
message ColumnOrExpression {
36+
oneof col_or_expr {
37+
bytes col = 1;
38+
string expr = 2;
39+
}
40+
}
41+
42+
message StringOrLongID {
43+
oneof id {
44+
int64 long_id = 1;
45+
string string_id = 2;
46+
}
47+
}
48+
49+
message AggregateMessages {
50+
ColumnOrExpression agg_col = 1;
51+
optional ColumnOrExpression send_to_src = 2;
52+
optional ColumnOrExpression send_to_dst = 3;
53+
}
54+
55+
message BFS {
56+
ColumnOrExpression from_expr = 1;
57+
ColumnOrExpression to_expr = 2;
58+
ColumnOrExpression edge_filter = 3;
59+
int32 max_path_length = 4;
60+
}
61+
62+
message ConnectedComponents {
63+
string algorithm = 1;
64+
int32 checkpoint_interval = 2;
65+
int32 broadcast_threshold = 3;
66+
}
67+
68+
message DropIsolatedVertices {}
69+
70+
message FilterEdges {
71+
ColumnOrExpression condition = 1;
72+
}
73+
74+
message FilterVertices {
75+
ColumnOrExpression condition = 2;
76+
}
77+
78+
message Find {
79+
string pattern = 1;
80+
}
81+
82+
message LabelPropagation {
83+
int32 max_iter = 1;
84+
}
85+
86+
message PageRank {
87+
double reset_probability = 1;
88+
optional StringOrLongID source_id = 2;
89+
optional int32 max_iter = 3;
90+
optional double tol = 4;
91+
}
92+
93+
message ParallelPersonalizedPageRank {
94+
double reset_probability = 1;
95+
repeated StringOrLongID source_ids = 2;
96+
int32 max_iter = 3;
97+
}
98+
99+
message PowerIterationClustering {
100+
int32 k = 1;
101+
int32 max_iter = 2;
102+
optional string weight_col = 3;
103+
}
104+
105+
message Pregel {
106+
ColumnOrExpression agg_msgs = 1;
107+
repeated ColumnOrExpression send_msg_to_dst = 2;
108+
repeated ColumnOrExpression send_msg_to_src = 3;
109+
int32 checkpoint_interval = 4;
110+
int32 max_iter = 5;
111+
string additional_col_name = 6;
112+
ColumnOrExpression additional_col_initial = 7;
113+
ColumnOrExpression additional_col_upd = 8;
114+
}
115+
116+
message ShortestPaths {
117+
repeated StringOrLongID landmarks = 1;
118+
}
119+
120+
message StronglyConnectedComponents {
121+
int32 max_iter = 1;
122+
}
123+
124+
message SVDPlusPlus {
125+
int32 rank = 1;
126+
int32 max_iter = 2;
127+
double min_value = 3;
128+
double max_value = 4;
129+
double gamma1 = 5;
130+
double gamma2 = 6;
131+
double gamma6 = 7;
132+
double gamma7 = 8;
133+
}
134+
135+
message TriangleCount {}
136+
137+
message Triplets {}
Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,24 @@
1+
package org.apache.spark.sql.graphframes
2+
3+
import org.graphframes.connect.proto.GraphFramesAPI
4+
5+
import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan
6+
import org.apache.spark.sql.connect.planner.SparkConnectPlanner
7+
import org.apache.spark.sql.connect.plugin.RelationPlugin
8+
9+
import com.google.protobuf
10+
11+
class GraphFramesConnect extends RelationPlugin {
12+
override def transform(
13+
relation: protobuf.Any,
14+
planner: SparkConnectPlanner): Option[LogicalPlan] = {
15+
if (relation.is(classOf[GraphFramesAPI])) {
16+
val protoCall = relation.unpack(classOf[GraphFramesAPI])
17+
// Because the plugins API is changed in spark 4.0 it makes sense to separate plugin impl from the parsing logic
18+
val result = GraphFramesConnectUtils.parseAPICall(protoCall, planner)
19+
Some(result.logicalPlan)
20+
} else {
21+
None
22+
}
23+
}
24+
}

0 commit comments

Comments
 (0)