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
61 changes: 61 additions & 0 deletions python/sedona/spark/geopandas/_crs.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,61 @@
# Licensed to the Apache Software Foundation (ASF) under one
# or more contributor license agreements. See the NOTICE file
# distributed with this work for additional information
# regarding copyright ownership. The ASF licenses this file
# to you under the Apache License, Version 2.0 (the
# "License"); you may not use this file except in compliance
# with the License. You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing,
# software distributed under the License is distributed on an
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.

"""CRS metadata helpers for the distributed GeoPandas API."""

from __future__ import annotations

from typing import Any

from pyproj import CRS
from pyspark.pandas.internal import InternalField

CRS_METADATA_KEY = "sedona.geopandas.crs_wkt"
NO_CRS_OVERRIDE = object()


def read_crs_metadata(field: InternalField) -> tuple[bool, CRS | None]:
"""Return whether CRS metadata is present and its decoded value."""
metadata = field.metadata or {}
if CRS_METADATA_KEY not in metadata:
return False, None

value = metadata[CRS_METADATA_KEY]
return True, CRS.from_wkt(value) if value else None


def with_crs_metadata(field: InternalField, crs: Any | None) -> InternalField:
"""Return an InternalField carrying normalized CRS metadata."""
metadata = dict(field.metadata or {})
metadata[CRS_METADATA_KEY] = (
CRS.from_user_input(crs).to_wkt() if crs is not None else ""
)
return field.copy(metadata=metadata)


def copy_crs_metadata(
source: InternalField,
target: InternalField,
) -> InternalField:
"""Copy only Sedona CRS metadata while retaining all target metadata."""
source_metadata = source.metadata or {}
target_metadata = dict(target.metadata or {})
if CRS_METADATA_KEY in source_metadata:
target_metadata[CRS_METADATA_KEY] = source_metadata[CRS_METADATA_KEY]
else:
target_metadata.pop(CRS_METADATA_KEY, None)
return target.copy(metadata=target_metadata)
86 changes: 68 additions & 18 deletions python/sedona/spark/geopandas/geodataframe.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,10 +51,12 @@
)
from pyspark.pandas.utils import log_advice

from sedona.spark.geopandas._crs import copy_crs_metadata, with_crs_metadata
from sedona.spark.geopandas._typing import Label
from sedona.spark.geopandas.base import GeoFrame
from sedona.spark.sql import st_aggregates as sta
from sedona.spark.sql import st_constructors as stc
from sedona.spark.sql import st_functions as stf

from pandas.api.extensions import register_extension_dtype
from geopandas.geodataframe import crs_mismatch_error
Expand Down Expand Up @@ -617,15 +619,17 @@ def __getitem__(self, key: Any) -> Any:
item = pspd.DataFrame.__getitem__(self, key)

if isinstance(item, pspd.DataFrame):
# Don't specify crs=self.crs here because it might not include the geometry column.
# If it does include the geometry column, we don't need to set crs anyways.
return GeoDataFrame(item)
result = GeoDataFrame(item)
if self._geometry_column_name in result.columns:
result._geometry_column_name = self._geometry_column_name
return result
elif isinstance(item, pspd.Series):
ps_series: pspd.Series = item
try:
return sgpd.GeoSeries(ps_series)
result = sgpd.GeoSeries(ps_series)
except TypeError:
return ps_series
return result
else:
raise Exception(f"Logical Error: Unexpected type: {type(item)}")

Expand Down Expand Up @@ -713,7 +717,47 @@ def __init__(
raise ValueError(crs_mismatch_error)

if geometry:
self.set_geometry(geometry, inplace=True, crs=crs)
existing_geometry = None
if crs is not None and pd.api.types.is_hashable(geometry):
try:
candidate = self[geometry]
except (KeyError, TypeError, ValueError):
pass
else:
if isinstance(candidate, sgpd.GeoSeries):
existing_geometry = candidate

if existing_geometry is None:
self.set_geometry(geometry, inplace=True, crs=crs)
else:
# Updating the existing Spark column directly avoids index
# alignment, which would multiply rows for duplicate indexes.
from pyproj import CRS

normalized_crs = CRS.from_user_input(crs)
new_epsg = normalized_crs.to_epsg() or 0
geometry_field = with_crs_metadata(
existing_geometry._internal.data_fields[0],
normalized_crs,
)
geometry_column_name = (
existing_geometry._internal.data_spark_column_names[0]
)
geometry_field = geometry_field.copy(name=geometry_column_name)
self._update_internal_frame(
self._internal.with_new_spark_column(
existing_geometry._column_label,
stf.ST_SetSRID(
existing_geometry.spark.column,
new_epsg,
).alias(
geometry_column_name,
metadata=geometry_field.metadata,
),
field=geometry_field,
)
)
self._geometry_column_name = geometry

if geometry is None and "geometry" in self.columns:

Expand Down Expand Up @@ -771,11 +815,7 @@ def _get_geometry(self) -> sgpd.GeoSeries:
)

raise MissingGeometryColumnError(msg)
geometry = self[self._geometry_column_name]
empty_crs_source = getattr(self, "_empty_crs_source", None)
if empty_crs_source is not None:
geometry._empty_crs_source = empty_crs_source
return geometry
return self[self._geometry_column_name]

def _set_geometry(self, col):
# This check is included in the original geopandas. Note that this prevents assigning a str to the property
Expand Down Expand Up @@ -888,7 +928,8 @@ def set_geometry(
else:
frame = self.copy()

geo_column_name = self._geometry_column_name
previous_geometry_name = self._geometry_column_name
geo_column_name = previous_geometry_name
new_series = False

if geo_column_name is None:
Expand Down Expand Up @@ -970,7 +1011,6 @@ def set_geometry(
if new_series:
# Note: This casts GeoSeries back into pspd.Series, so we lose any metadata that's not serialized.
frame[geo_column_name] = level
object.__setattr__(frame, "_empty_crs_source", level)

if not inplace:
return frame
Expand Down Expand Up @@ -1090,7 +1130,14 @@ def _to_geopandas(self) -> gpd.GeoDataFrame:
else:
pd_df[col_name] = series._to_pandas()

return gpd.GeoDataFrame(pd_df, geometry=self._geometry_column_name)
result = gpd.GeoDataFrame(
pd_df,
geometry=self._geometry_column_name,
crs=self.crs if self._geometry_column_name is not None else None,
)
if self._geometry_column_name is None:
result._geometry_column_name = None
return result

def to_spark_pandas(self) -> pspd.DataFrame:
"""
Expand Down Expand Up @@ -1124,9 +1171,10 @@ def copy(self, deep=False) -> GeoDataFrame:
0 POINT (1 1) 2 3
"""
# Note: The deep parameter is a dummy parameter just as it is in PySpark pandas.
return GeoDataFrame(
result = GeoDataFrame(
pspd.DataFrame(self._internal.copy()), geometry=self.active_geometry_name
)
return result

def _safe_get_crs(self):
"""
Expand Down Expand Up @@ -2081,9 +2129,12 @@ def dissolve(
*[scol_for(aggregated_sdf, name) for name in attribute_output_names],
],
data_fields=[
InternalField(
np.dtype("object"),
aggregated_sdf.schema[geometry_output_name],
copy_crs_metadata(
self.geometry._internal.data_fields[0],
InternalField(
np.dtype("object"),
aggregated_sdf.schema[geometry_output_name],
),
),
InternalField.from_struct_field(
aggregated_sdf.schema[order_output_spark_name]
Expand Down Expand Up @@ -2133,7 +2184,6 @@ def dissolve(
result_labels.geometry_name,
)

object.__setattr__(aggregated, "_empty_crs_source", self.geometry)
return aggregated

# ============================================================================
Expand Down
Loading
Loading