Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 16 additions & 2 deletions python/sedona/spark/geopandas/geoseries.py
Original file line number Diff line number Diff line change
Expand Up @@ -737,8 +737,22 @@ def __init__(

pd_series = pd_series.astype(object)

# Initialize the parent class PySpark Series with the pandas Series.
super().__init__(data=pd_series)
if (
not pd_series.empty
and pd_series.isna().iloc[0]
and any(isinstance(value, BaseGeometry) for value in pd_series)
):
# Spark 3.5 infers object UDTs from the first value. WKB lets
# leading missing values retain geometry type without reordering.
wkb_series = (
gpd.GeoSeries(pd_series.where(pd_series.notna(), None))
.to_wkb(include_srid=True)
.rename(pd_series.name)
)
ps_series = pspd.Series(wkb_series).spark.transform(stc.ST_GeomFromWKB)
super().__init__(data=ps_series._anchor, index=ps_series._col_label)
else:
super().__init__(data=pd_series)

# Ensure we're storing geometry types.
if (
Expand Down
86 changes: 86 additions & 0 deletions python/tests/geopandas/test_geoseries.py
Original file line number Diff line number Diff line change
Expand Up @@ -677,6 +677,92 @@ def test_constructor(self, obj):
sgpd_series = sgpd.GeoSeries(obj)
assert isinstance(sgpd_series, sgpd.GeoSeries)

@pytest.mark.parametrize(
"wrap",
[list, tuple, np.asarray, pd.Series, gpd.GeoSeries, gpd.array.from_shapely],
ids=["list", "tuple", "numpy", "pandas", "geopandas", "geometry_array"],
)
def test_constructor_leading_null_local_inputs(self, wrap):
from geopandas.testing import assert_geoseries_equal

_ = self.spark
values = [None, Point(1, 0), None, Point()]
result = GeoSeries(wrap(values))

assert_geoseries_equal(
result.to_geopandas(), gpd.GeoSeries(values), check_index_type=False
)
assert result.crs is None

@pytest.mark.parametrize(
"missing", [None, np.nan, pd.NA, pd.NaT, np.datetime64("NaT")]
)
def test_constructor_leading_missing_values(self, missing):
from geopandas.testing import assert_geoseries_equal

_ = self.spark
result = GeoSeries([missing, Point(1, 0), missing])

assert_geoseries_equal(
result.to_geopandas(),
gpd.GeoSeries([None, Point(1, 0), None]),
check_index_type=False,
)

@pytest.mark.parametrize(
"name, inherited_crs, crs",
[
(None, None, None),
("geometry", "EPSG:4326", None),
(("geometry", "shape"), "EPSG:4326", "EPSG:3857"),
],
)
def test_constructor_leading_null_metadata(self, name, inherited_crs, crs):
from geopandas.testing import assert_geoseries_equal

_ = self.spark
index = pd.MultiIndex.from_tuples(
[("b", 2), ("a", 1), ("b", 2)], names=["letter", "number"]
)
values = [None, Point(1, 0, 2), Point(3, 4)]
local = gpd.GeoSeries(values, index=index, name=name, crs=inherited_crs)
result = GeoSeries(local, crs=crs)
expected = gpd.GeoSeries(
values, index=index, name=name, crs=crs or inherited_crs
)

assert result.name == name
assert_geoseries_equal(result.to_geopandas(), expected)
assert_geoseries_equal(
local, gpd.GeoSeries(values, index=index, name=name, crs=inherited_crs)
)

@pytest.mark.parametrize("values", [[], [None, None], [Point(1, 0), None]])
def test_constructor_null_and_empty_controls(self, values):
from geopandas.testing import assert_geoseries_equal

_ = self.spark
assert_geoseries_equal(
GeoSeries(values).to_geopandas(),
gpd.GeoSeries(values),
check_index_type=False,
)

def test_constructor_leading_null_preserves_embedded_srid(self):
from shapely import wkb
from sedona.spark.sql.types import GeometryType

_ = self.spark
point = wkb.loads(wkb.dumps(Point(1, 2, 3), srid=4326))
result = GeoSeries([None, point])

assert result.spark.data_type == GeometryType()
assert result.crs is None
assert result.spark.transform(stf.ST_SRID).to_pandas().dropna().tolist() == [
4326
]
assert result.spark.transform(stf.ST_Z).to_pandas().dropna().tolist() == [3]

def test_constructor_pandas_on_spark(self):
obj = ps.Series([Point(x, x) for x in range(3)])
sgpd_series = GeoSeries(obj)
Expand Down
Loading