/
githubmirror
/
spark
Обзор
Документация
Войти
/
githubmirror
/
spark
Код
Запросы
0
Пакеты
0
Релизы
0
Аналитика
Безопасность
master
python/pyspark/sql/conversion.py
2 014 строк
77 KB
Spenser Sun
[SPARK-58625][PYTHON] Support pa.ChunkedArray columns in PandasToArrowConversion
07 авг 2026, 08:44
07 авг 2026, 08:44
86d4fe8
Код
Авторство
О чём код?
# # 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. # import array import datetime import decimal import functools from typing import TYPE_CHECKING, Any, Callable, List, Optional, Sequence, Union, overload import pyspark from pyspark.errors import PySparkNotImplementedError, PySparkRuntimeError, PySparkValueError from pyspark.sql.pandas.types import ( _dedup_names, _deduplicate_field_names, _create_converter_to_pandas, to_arrow_schema, from_arrow_schema, ) from pyspark.sql.pandas.utils import require_minimum_pyarrow_version from pyspark.sql.types import ( ArrayType, BinaryType, BooleanType, ByteType, ShortType, IntegerType, LongType, DataType, FloatType, DoubleType, DecimalType, GeographyType, Geography, GeometryType, Geometry, MapType, NullType, Row, StringType, StructField, StructType, DateType, TimeType, TimestampNTZType, TimestampType, DayTimeIntervalType, YearMonthIntervalType, UserDefinedType, VariantType, VariantVal, _create_row, _has_type, ) if TYPE_CHECKING: import pyarrow as pa import pandas as pd class ArrowBatchTransformer: """ Pure functions that transform RecordBatch -> RecordBatch. They should have no side effects (no I/O, no writing to streams). """ @staticmethod def flatten_struct(batch: "pa.RecordBatch", column_index: int = 0) -> "pa.RecordBatch": """ Flatten a struct column at given index into a RecordBatch. Used by: - SQL_GROUPED_MAP_ARROW_UDF mapper - SQL_GROUPED_MAP_ARROW_ITER_UDF mapper """ import pyarrow as pa struct = batch.column(column_index) return pa.RecordBatch.from_arrays(struct.flatten(), schema=pa.schema(struct.type)) @classmethod def select_columns(cls, batch: "pa.RecordBatch", column_indices: list[int]) -> "pa.RecordBatch": """ Select a subset of columns from a RecordBatch by index. Used by: SQL_COGROUPED_MAP_ARROW_UDF handler in worker.py """ import pyarrow as pa return pa.RecordBatch.from_arrays( [batch.columns[i] for i in column_indices], [batch.schema.names[i] for i in column_indices], ) @staticmethod def wrap_struct(batch: "pa.RecordBatch") -> "pa.RecordBatch": """ Wrap a RecordBatch's columns into a single struct column. Used by: Arrow UDF mappers in worker.py to re-wrap flattened batches before serialization. """ import pyarrow as pa if batch.num_columns == 0: # When batch has no column, it should still create # an empty batch with the number of rows set. struct = pa.array([{}] * batch.num_rows) else: struct = pa.StructArray.from_arrays(batch.columns, fields=pa.struct(list(batch.schema))) return pa.RecordBatch.from_arrays([struct], ["_0"]) @classmethod def enforce_schema( cls, batch: Union["pa.RecordBatch", "pa.Table"], arrow_schema: "pa.Schema", *, arrow_cast: bool = True, safecheck: bool = True, reorder_by_name: bool = True, ) -> Union["pa.RecordBatch", "pa.Table"]: """ Enforce a target schema on an Arrow RecordBatch or Table. Parameters ---------- batch : pa.RecordBatch or pa.Table Input to transform. Output is of the same container type. arrow_schema : pa.Schema Target Arrow schema. Callers should pre-compute this once via to_arrow_schema() to avoid repeated conversion. arrow_cast : bool, default True If True, cast mismatched types to the target type. If False, raise an error on type mismatch instead of casting. safecheck : bool, default True If True, use safe casting (fails on overflow/truncation). reorder_by_name : bool, default True If True, match columns by name and reorder to the target order; any missing or extra names raise ``RESULT_COLUMN_NAMES_MISMATCH``. Output columns are renamed to target names. If False, match columns by position (ignore names) and preserve the original column names in the output. Returns ------- pa.RecordBatch or pa.Table Same container type as ``batch``, with columns matched (and possibly reordered/cast) per the target schema. Raises ------ PySparkRuntimeError ``RESULT_COLUMN_NAMES_MISMATCH`` when ``reorder_by_name=True`` and the batch has missing or extra column names. ``RESULT_COLUMN_TYPES_MISMATCH`` when any column's type does not match the target (and either ``arrow_cast=False`` or the cast itself fails). ``RESULT_COLUMN_SCHEMA_MISMATCH`` when ``reorder_by_name=False`` and the batch has a different number of columns than the target schema. """ import pyarrow as pa if batch.num_columns == 0 or len(arrow_schema) == 0: return batch # Fast path: schema already matches (ignoring metadata), no work needed if batch.schema.equals(arrow_schema, check_metadata=False): return batch target_names = [field.name for field in arrow_schema] # Step 1: pick source columns from batch to align with target schema if reorder_by_name: batch_names = [batch.schema.field(i).name for i in range(batch.num_columns)] missing = sorted(set(target_names) - set(batch_names)) extra = sorted(set(batch_names) - set(target_names)) if missing or extra: raise PySparkRuntimeError( errorClass="RESULT_COLUMN_NAMES_MISMATCH", messageParameters={ "missing": f" Missing: {', '.join(missing)}." if missing else "", "extra": f" Unexpected: {', '.join(extra)}." if extra else "", }, ) source_columns = [batch.column(name) for name in target_names] output_names = target_names else: # Positional: require exact column-count match, then take columns by # index, preserving the batch's original column names. if batch.num_columns != len(arrow_schema): raise PySparkRuntimeError( errorClass="RESULT_COLUMN_SCHEMA_MISMATCH", messageParameters={ "expected": str(len(arrow_schema)), "actual": str(batch.num_columns), }, ) source_columns = [batch.column(i) for i in range(len(arrow_schema))] output_names = [batch.schema.field(i).name for i in range(len(arrow_schema))] # Step 2: check types / cast, collect all mismatches type_mismatches = [] coerced_arrays = [] for field, arr in zip(arrow_schema, source_columns): if arr.type == field.type: coerced_arrays.append(arr) elif not arrow_cast: type_mismatches.append((field.name, field.type, arr.type)) coerced_arrays.append(arr) else: try: coerced_arrays.append(arr.cast(target_type=field.type, safe=safecheck)) except (pa.ArrowInvalid, pa.ArrowTypeError): type_mismatches.append((field.name, field.type, arr.type)) coerced_arrays.append(arr) if type_mismatches: raise PySparkRuntimeError( errorClass="RESULT_COLUMN_TYPES_MISMATCH", messageParameters={ "mismatch": ", ".join( f"column '{name}' (expected {expected}, actual {actual})" for name, expected, actual in type_mismatches ) }, ) # Preserve input container type (Table vs RecordBatch) if isinstance(batch, pa.Table): return pa.Table.from_arrays(coerced_arrays, names=output_names) return pa.RecordBatch.from_arrays(coerced_arrays, names=output_names) @classmethod def to_pandas( cls, batch: Union["pa.RecordBatch", "pa.Table"], timezone: str, schema: Optional["StructType"] = None, struct_in_pandas: str = "dict", ndarray_as_list: bool = False, prefer_int_ext_dtype: bool = False, df_for_struct: bool = False, ) -> List[Union["pd.Series", "pd.DataFrame"]]: """ Convert a RecordBatch or Table to a list of pandas Series. Parameters ---------- batch : pa.RecordBatch or pa.Table The Arrow RecordBatch or Table to convert. timezone : str Timezone for timestamp conversion. schema : StructType, optional Spark schema for type conversion. If None, types are inferred from Arrow. struct_in_pandas : str How to represent struct in pandas ("dict", "row", etc.) ndarray_as_list : bool Whether to convert ndarray as list. prefer_int_ext_dtype : bool, optional Whether to convert integers to Pandas ExtensionDType. df_for_struct : bool If True, convert struct columns to DataFrame instead of Series. Returns ------- List[Union[pd.Series, pd.DataFrame]] List of pandas Series (or DataFrame if df_for_struct=True), one for each column. """ import pandas as pd if batch.num_columns == 0: return [pd.Series([pyspark._NoValue] * batch.num_rows)] if schema is None: schema = from_arrow_schema(batch.schema) return [ ArrowArrayToPandasConversion.convert( batch.column(i), schema[i].dataType, ser_name=schema[i].name, timezone=timezone, struct_in_pandas=struct_in_pandas, ndarray_as_list=ndarray_as_list, prefer_int_ext_dtype=prefer_int_ext_dtype, df_for_struct=df_for_struct, ) for i in range(batch.num_columns) ] class PandasToArrowConversion: """ Conversion utilities from pandas data to Arrow. """ @classmethod def convert( cls, data: Union["pd.DataFrame", Sequence[Union["pd.Series", "pd.DataFrame"]]], schema: StructType, *, timezone: Optional[str] = None, safecheck: bool = True, arrow_cast: bool = False, prefers_large_types: bool = False, assign_cols_by_name: bool = False, int_to_decimal_coercion_enabled: bool = False, ignore_unexpected_complex_type_values: bool = False, is_legacy: bool = False, ) -> "pa.RecordBatch": """ Convert a pandas DataFrame or list of Series/DataFrames to an Arrow RecordBatch. Parameters ---------- data : pd.DataFrame or list of pd.Series/pd.DataFrame Input data - either a single DataFrame, or a list of Series/DataFrames (one per schema field). A list of DataFrames is used when UDFs return struct types as DataFrames (e.g., applyInPandas with state). schema : StructType Spark schema defining the types for each column timezone : str, optional Timezone for timestamp conversion safecheck : bool Whether to use safe Arrow conversion (default True) arrow_cast : bool Whether to allow Arrow casting on type mismatch (default False) prefers_large_types : bool Whether to prefer large Arrow types (default False) assign_cols_by_name : bool Whether to reorder DataFrame columns by name to match schema (default False) int_to_decimal_coercion_enabled : bool Whether to enable int to decimal coercion (default False) ignore_unexpected_complex_type_values : bool Whether to ignore unexpected complex type values in converter (default False) is_legacy : bool Whether to use the legacy pandas-to-Arrow conversion path. The legacy path uses broader Arrow exception handling (ArrowException) to allow more implicit type coercions (e.g., int->boolean, dict->struct via ArrowTypeError). The non-legacy path only catches ArrowInvalid for the cast fallback, so type mismatches like string->decimal raise immediately. (default False) Returns ------- pa.RecordBatch """ import pyarrow as pa import pandas as pd from pyspark.errors import PySparkTypeError, PySparkValueError from pyspark.sql.pandas.types import to_arrow_type, _create_converter_from_pandas # Handle empty schema (0 columns) # Use dummy column + select([]) to preserve row count (PyArrow limitation workaround) if len(schema.fields) == 0: num_rows = len(data[0]) if isinstance(data, list) and data else len(data) return pa.RecordBatch.from_pydict({"_": [None] * num_rows}).select([]) # Handle empty DataFrame (0 columns) with non-empty schema # This happens when user returns pd.DataFrame() for struct types if isinstance(data, pd.DataFrame) and len(data.columns) == 0: arrow_type = to_arrow_type( schema, timezone=timezone, prefers_large_types=prefers_large_types ) return pa.RecordBatch.from_struct_array(pa.array([{}] * len(data), arrow_type)) # Normalize input: reorder DataFrame columns by schema names if needed, # then extract columns as a list for uniform iteration. columns: List[Union["pd.Series", "pd.DataFrame"]] if isinstance(data, pd.DataFrame): if assign_cols_by_name and any(isinstance(c, str) for c in data.columns): data = data[schema.names] columns = [data.iloc[:, i] for i in range(len(schema.fields))] else: columns = list(data) def convert_column( col: Union["pd.Series", "pd.DataFrame"], field: StructField ) -> "pa.Array": """Convert a single column (Series or DataFrame) to an Arrow Array. Uses field.name for error messages instead of series.name to avoid copying the Series via rename() - a ~20% overhead on the hot path. """ if isinstance(col, pd.DataFrame): assert isinstance(field.dataType, StructType) nested_batch = cls.convert( col, field.dataType, timezone=timezone, safecheck=safecheck, arrow_cast=arrow_cast, prefers_large_types=prefers_large_types, assign_cols_by_name=assign_cols_by_name, int_to_decimal_coercion_enabled=int_to_decimal_coercion_enabled, ignore_unexpected_complex_type_values=ignore_unexpected_complex_type_values, is_legacy=is_legacy, ) # Wrap the nested RecordBatch as a single StructArray column return ArrowBatchTransformer.wrap_struct(nested_batch).column(0) series = col field_name = field.name ret_type = field.dataType if isinstance(series.dtype, pd.CategoricalDtype): series = series.astype(series.dtype.categories.dtype) arrow_type = to_arrow_type( ret_type, timezone=timezone, prefers_large_types=prefers_large_types ) series = _create_converter_from_pandas( ret_type, timezone=timezone, error_on_duplicated_field_names=False, int_to_decimal_coercion_enabled=int_to_decimal_coercion_enabled, ignore_unexpected_complex_type_values=ignore_unexpected_complex_type_values, )(series) mask = None if hasattr(series.array, "__arrow_array__") else series.isnull() if is_legacy: # Legacy pandas conversion path: broad ArrowException catch so # that both ArrowInvalid AND ArrowTypeError (e.g. dict->struct) # trigger the cast fallback. try: try: return pa.Array.from_pandas( series, mask=mask, type=arrow_type, safe=safecheck ) except pa.lib.ArrowException: # broad: includes ArrowTypeError if arrow_cast: return pa.Array.from_pandas(series, mask=mask).cast( target_type=arrow_type, safe=safecheck ) raise except pa.lib.ArrowException as e: error_msg = ( "Exception thrown when converting pandas.Series (%s) " "with name '%s' to Arrow Array (%s)." % (series.dtype, field_name, arrow_type) ) if isinstance(e, TypeError): raise PySparkTypeError(error_msg) from e if safecheck: error_msg += ( " It can be caused by overflows or other " "unsafe conversions warned by Arrow. Arrow safe " "type check can be disabled by using SQL config " "`spark.sql.execution.pandas." "convertToArrowArraySafely`." ) raise PySparkValueError(error_msg) from e else: # Non-legacy path: only ArrowInvalid triggers the cast fallback. # ArrowTypeError (e.g. string->decimal) must NOT be silently cast. try: try: return pa.Array.from_pandas( series, mask=mask, type=arrow_type, safe=safecheck ) except pa.lib.ArrowInvalid: # narrow: skip ArrowTypeError if arrow_cast: return pa.Array.from_pandas(series, mask=mask).cast( target_type=arrow_type, safe=safecheck ) raise except TypeError as e: raise PySparkTypeError( f"Cannot convert the output value of the column " f"'{field_name}' with type '{series.dtype}' to the " f"specified return type of the column: '{arrow_type}'." f" Please check if the data types match and try again." ) from e except ValueError as e: error_msg = ( f"Failed to convert the value of the column " f"'{field_name}' with type '{series.dtype}' to Arrow " f"type '{arrow_type}'." ) if safecheck: error_msg += ( " It can be caused by overflows or other unsafe " "conversions warned by Arrow. Arrow safe type " "check can be disabled by using SQL config " "`spark.sql.execution.pandas." "convertToArrowArraySafely`." ) raise PySparkValueError(error_msg) from e converted = [convert_column(col, field) for col, field in zip(columns, schema.fields)] # pa.Array.from_pandas returns a pa.ChunkedArray for a chunked arrow-backed Series # (e.g. a pyarrow-backed extension dtype), which pa.RecordBatch.from_arrays rejects. arrays = [a.combine_chunks() if isinstance(a, pa.ChunkedArray) else a for a in converted] return pa.RecordBatch.from_arrays(arrays, schema.names) class LocalDataToArrowConversion: """ Conversion from local data (except pandas DataFrame and numpy ndarray) to Arrow. """ @staticmethod def _need_converter( dataType: DataType, nullable: bool = True, ) -> bool: if not nullable: # always check the nullability return True elif isinstance(dataType, NullType): # always check the nullability return True elif isinstance(dataType, StructType): # Struct maybe rows, should convert to dict. return True elif isinstance(dataType, ArrayType): return LocalDataToArrowConversion._need_converter( dataType.elementType, dataType.containsNull ) elif isinstance(dataType, MapType): # Different from PySpark, here always needs conversion, # since an Arrow Map requires a list of tuples. return True elif isinstance(dataType, BinaryType): return True elif isinstance(dataType, (TimestampType, TimestampNTZType)): # Always truncate return True elif isinstance(dataType, DecimalType): # Convert Decimal('NaN') to None # Rescale Decimal values return True elif isinstance(dataType, StringType): # Coercion to StringType is allowed return True elif isinstance(dataType, UserDefinedType): return True elif isinstance(dataType, VariantType): return True elif isinstance(dataType, GeometryType): return True elif isinstance(dataType, GeographyType): return True else: return False @overload @staticmethod def _create_converter( dataType: DataType, nullable: bool = True, *, int_to_decimal_coercion_enabled: bool = False ) -> Callable: pass @overload @staticmethod def _create_converter( dataType: DataType, nullable: bool = True, *, none_on_identity: bool = False, int_to_decimal_coercion_enabled: bool = False, ) -> Optional[Callable]: pass @staticmethod def _create_converter( dataType: DataType, nullable: bool = True, *, none_on_identity: bool = False, int_to_decimal_coercion_enabled: bool = False, ) -> Optional[Callable]: assert dataType is not None and isinstance(dataType, DataType) assert isinstance(nullable, bool) if not LocalDataToArrowConversion._need_converter(dataType, nullable): if none_on_identity: return None else: return lambda value: value if isinstance(dataType, NullType): def convert_null(value: Any) -> Any: if value is not None: raise PySparkValueError(f"input for {dataType} must be None, but got {value}") return None return convert_null elif isinstance(dataType, StructType): field_names = dataType.fieldNames() len_field_names = len(field_names) dedup_field_names = _dedup_names(dataType.names) field_convs = [ LocalDataToArrowConversion._create_converter( field.dataType, field.nullable, none_on_identity=True, int_to_decimal_coercion_enabled=int_to_decimal_coercion_enabled, ) for field in dataType.fields ] def convert_struct(value: Any) -> Any: if value is None: if not nullable: raise PySparkValueError(f"input for {dataType} must not be None") return None else: # The `value` should be tuple, dict, or have `__dict__`. if isinstance(value, tuple): # `Row` inherits `tuple` if len(value) != len_field_names: raise PySparkValueError( errorClass="AXIS_LENGTH_MISMATCH", messageParameters={ "expected_length": str(len_field_names), "actual_length": str(len(value)), }, ) return { dedup_field_names[i]: ( field_convs[i](value[i]) # type: ignore[misc] if field_convs[i] is not None else value[i] ) for i in range(len_field_names) } elif isinstance(value, dict): return { dedup_field_names[i]: ( field_convs[i](value.get(field)) # type: ignore[misc] if field_convs[i] is not None else value.get(field) ) for i, field in enumerate(field_names) } else: assert hasattr(value, "__dict__"), f"{type(value)} {value}" value = value.__dict__ return { dedup_field_names[i]: ( field_convs[i](value.get(field)) # type: ignore[misc] if field_convs[i] is not None else value.get(field) ) for i, field in enumerate(field_names) } return convert_struct elif isinstance(dataType, ArrayType): element_conv = LocalDataToArrowConversion._create_converter( dataType.elementType, dataType.containsNull, none_on_identity=True, int_to_decimal_coercion_enabled=int_to_decimal_coercion_enabled, ) if element_conv is None: def convert_array(value: Any) -> Any: if value is None: if not nullable: raise PySparkValueError(f"input for {dataType} must not be None") return None else: assert isinstance(value, (list, array.array)) return list(value) else: def convert_array(value: Any) -> Any: if value is None: if not nullable: raise PySparkValueError(f"input for {dataType} must not be None") return None else: assert isinstance(value, (list, array.array)) return [element_conv(v) for v in value] return convert_array elif isinstance(dataType, MapType): key_conv = LocalDataToArrowConversion._create_converter( dataType.keyType, nullable=False, int_to_decimal_coercion_enabled=int_to_decimal_coercion_enabled, ) value_conv = LocalDataToArrowConversion._create_converter( dataType.valueType, dataType.valueContainsNull, none_on_identity=True, int_to_decimal_coercion_enabled=int_to_decimal_coercion_enabled, ) if value_conv is None: def convert_map(value: Any) -> Any: if value is None: if not nullable: raise PySparkValueError(f"input for {dataType} must not be None") return None else: assert isinstance(value, dict) return [(key_conv(k), v) for k, v in value.items()] else: def convert_map(value: Any) -> Any: if value is None: if not nullable: raise PySparkValueError(f"input for {dataType} must not be None") return None else: assert isinstance(value, dict) return [(key_conv(k), value_conv(v)) for k, v in value.items()] return convert_map elif isinstance(dataType, BinaryType): def convert_binary(value: Any) -> Any: if value is None: if not nullable: raise PySparkValueError(f"input for {dataType} must not be None") return None else: assert isinstance(value, (bytes, bytearray)) return bytes(value) return convert_binary elif isinstance(dataType, TimestampType): def convert_timestamp(value: Any) -> Any: if value is None: if not nullable: raise PySparkValueError(f"input for {dataType} must not be None") return None else: assert isinstance(value, datetime.datetime) return value.astimezone(datetime.timezone.utc) return convert_timestamp elif isinstance(dataType, TimestampNTZType): def convert_timestamp_ntz(value: Any) -> Any: if value is None: if not nullable: raise PySparkValueError(f"input for {dataType} must not be None") return None else: assert isinstance(value, datetime.datetime) and value.tzinfo is None return value return convert_timestamp_ntz elif isinstance(dataType, DecimalType): exp = decimal.Decimal(f"1E-{dataType.scale}") ctx = decimal.Context(prec=dataType.precision, rounding=decimal.ROUND_HALF_EVEN) def convert_decimal(value: Any) -> Any: if value is None: if not nullable: raise PySparkValueError(f"input for {dataType} must not be None") return None else: if int_to_decimal_coercion_enabled and isinstance(value, int): value = decimal.Decimal(value) assert isinstance(value, decimal.Decimal) if value.is_nan(): if not nullable: raise PySparkValueError(f"input for {dataType} must not be None") return None return value.quantize(exp, context=ctx) return convert_decimal elif isinstance(dataType, StringType): def convert_string(value: Any) -> Any: if value is None: if not nullable: raise PySparkValueError(f"input for {dataType} must not be None") return None else: if isinstance(value, bool): # To match the PySpark Classic which convert bool to string in # the JVM side (python.EvaluatePython.makeFromJava) return str(value).lower() else: return str(value) return convert_string elif isinstance(dataType, UserDefinedType): udt: UserDefinedType = dataType conv = LocalDataToArrowConversion._create_converter( udt.sqlType(), nullable=nullable, none_on_identity=True, int_to_decimal_coercion_enabled=int_to_decimal_coercion_enabled, ) if conv is None: def convert_udt(value: Any) -> Any: if value is None: if not nullable: raise PySparkValueError(f"input for {dataType} must not be None") return None else: return udt.serialize(value) else: def convert_udt(value: Any) -> Any: if value is None: if not nullable: raise PySparkValueError(f"input for {dataType} must not be None") return None else: return conv(udt.serialize(value)) return convert_udt elif isinstance(dataType, VariantType): def convert_variant(value: Any) -> Any: if value is None: if not nullable: raise PySparkValueError(f"input for {dataType} must not be None") return None elif isinstance(value, VariantVal): return VariantType().toInternal(value) else: raise PySparkValueError(errorClass="MALFORMED_VARIANT") return convert_variant elif isinstance(dataType, GeographyType): def convert_geography(value: Any) -> Any: if value is None: if not nullable: raise PySparkValueError(f"input for {dataType} must not be None") return None elif isinstance(value, Geography): return dataType.toInternal(value) else: raise PySparkValueError(errorClass="MALFORMED_GEOGRAPHY") return convert_geography elif isinstance(dataType, GeometryType): def convert_geometry(value: Any) -> Any: if value is None: if not nullable: raise PySparkValueError(f"input for {dataType} must not be None") return None elif isinstance(value, Geometry): return dataType.toInternal(value) else: raise PySparkValueError(errorClass="MALFORMED_GEOMETRY") return convert_geometry elif not nullable: def convert_other(value: Any) -> Any: if value is None: raise PySparkValueError(f"input for {dataType} must not be None") return value return convert_other else: # pragma: no cover assert False, f"Need converter for {dataType} but failed to find one." @staticmethod def convert(data: Sequence[Any], schema: StructType, use_large_var_types: bool) -> "pa.Table": require_minimum_pyarrow_version() import pyarrow as pa assert isinstance(data, list) and len(data) > 0 assert schema is not None and isinstance(schema, StructType) column_names = schema.fieldNames() len_column_names = len(column_names) def to_row(item: Any) -> tuple: if item is None: return tuple([None] * len_column_names) elif isinstance(item, tuple): # `Row` inherits `tuple` if len(item) != len_column_names: raise PySparkValueError( errorClass="AXIS_LENGTH_MISMATCH", messageParameters={ "expected_length": str(len_column_names), "actual_length": str(len(item)), }, ) return tuple(item) elif isinstance(item, dict): return tuple([item.get(col) for col in column_names]) elif isinstance(item, VariantVal): raise PySparkValueError("Rows cannot be of type VariantVal") elif hasattr(item, "__dict__"): item = item.__dict__ return tuple([item.get(col) for col in column_names]) else: if len(item) != len_column_names: raise PySparkValueError( errorClass="AXIS_LENGTH_MISMATCH", messageParameters={ "expected_length": str(len_column_names), "actual_length": str(len(item)), }, ) return tuple(item) rows = [to_row(item) for item in data] if len_column_names > 0: column_convs = [ LocalDataToArrowConversion._create_converter( field.dataType, field.nullable, none_on_identity=True, # Default to False for general data conversion int_to_decimal_coercion_enabled=False, ) for field in schema.fields ] pylist = [ [conv(row[i]) for row in rows] if conv is not None else [row[i] for row in rows] for i, conv in enumerate(column_convs) ] pa_schema = to_arrow_schema( StructType( [ StructField( field.name, _deduplicate_field_names(field.dataType), field.nullable ) for field in schema.fields ] ), timezone="UTC", prefers_large_types=use_large_var_types, ) return pa.Table.from_arrays(pylist, schema=pa_schema) else: return pa.Table.from_struct_array(pa.array([{}] * len(rows))) class ArrowTableToRowsConversion: """ Conversion from Arrow Table to Rows. """ @staticmethod @functools.cache def _should_manual_bulk() -> bool: """ Whether ``_to_pylist`` should convert nested columns manually in bulk. Internal helper for ``_to_pylist`` only; do not use externally. Returns True when the installed PyArrow still materializes one Scalar per element in ``to_pylist`` (apache/arrow#50326, fix expected in PyArrow 25.0.1 — adjust the version below if it ships in a different release) and NumPy (used for the offsets and validity buffers) is available. This method and the manual bulk paths in ``_to_pylist`` should be removed once the minimum supported PyArrow version contains the fix. """ import pyarrow as pa from pyspark.loose_version import LooseVersion if LooseVersion(pa.__version__) >= LooseVersion("25.0.1"): # Native to_pylist converts without per-element Scalars. return False try: import numpy # noqa: F401 except ImportError: return False return True @staticmethod def _to_pylist(column: Union["pa.Array", "pa.ChunkedArray"]) -> List[Any]: """ Equivalent to ``column.to_pylist()``, but converts (nested) list, struct and map columns in bulk instead of one scalar at a time. Structs become dicts (with a fallback to ``to_pylist`` for duplicate field names, which raise ``ValueError`` there) and maps become lists of ``(key, value)`` tuples, matching ``StructScalar.as_py`` and ``MapScalar.as_py`` exactly. Internal helper for the worker and ``convert`` call sites; do not use externally. ``Array.to_pylist()`` materializes one Scalar per element; for list types each row additionally allocates a C++ scalar, a Python Scalar wrapper and a Python Array wrapper for the row's values before converting elements one by one, which is several times slower than converting the flattened child values in a single pass and slicing the resulting Python list per row (see apache/arrow#50326). The values themselves are still converted by Arrow's own ``to_pylist``, so results are exactly identical: ``None`` stays ``None`` and values inside numeric lists stay Python ints, unlike a pandas round trip which would coerce them to floats/NaN. NumPy is used only for the offsets (non-null integers) and the validity bitmap (booleans), so no value coercion can occur. This method should be removed (its call sites reverting to plain ``column.to_pylist()``) once the minimum supported PyArrow version includes the fix for apache/arrow#50326. """ import pyarrow as pa if not ArrowTableToRowsConversion._should_manual_bulk(): return column.to_pylist() if isinstance(column, pa.ChunkedArray): result = [] for chunk in column.chunks: result.extend(ArrowTableToRowsConversion._to_pylist(chunk)) return result if len(column) == 0: return [] if pa.types.is_map(column.type): # Maps have the same offsets layout as lists; each row becomes a # list of (key, value) tuples, matching MapScalar.as_py. n = len(column) offsets = column.offsets.to_numpy(zero_copy_only=True).tolist() start = offsets[0] length = offsets[-1] - start keys = ArrowTableToRowsConversion._to_pylist(column.keys.slice(start, length)) items = ArrowTableToRowsConversion._to_pylist(column.items.slice(start, length)) if column.null_count == 0: return [ list( zip( keys[offsets[i] - start : offsets[i + 1] - start], items[offsets[i] - start : offsets[i + 1] - start], ) ) for i in range(n) ] valid = column.is_valid().to_numpy(zero_copy_only=False).tolist() return [ ( list( zip( keys[offsets[i] - start : offsets[i + 1] - start], items[offsets[i] - start : offsets[i + 1] - start], ) ) if valid[i] else None ) for i in range(n) ] elif pa.types.is_list(column.type) or pa.types.is_large_list(column.type): n = len(column) # List offset buffers never carry a validity bitmap, so this conversion is # always zero-copy; zero_copy_only=True asserts that invariant and would # fail loudly if a future Arrow list variant ever violated it. offsets = column.offsets.to_numpy(zero_copy_only=True).tolist() start = offsets[0] flat = ArrowTableToRowsConversion._to_pylist( column.values.slice(start, offsets[-1] - start) ) if column.null_count == 0: return [flat[offsets[i] - start : offsets[i + 1] - start] for i in range(n)] valid = column.is_valid().to_numpy(zero_copy_only=False).tolist() return [ flat[offsets[i] - start : offsets[i + 1] - start] if valid[i] else None for i in range(n) ] elif pa.types.is_struct(column.type): n = len(column) names = [column.type.field(i).name for i in range(column.type.num_fields)] if len(set(names)) != len(names): # StructScalar.as_py raises ValueError on duplicate field names; # let the generic path surface the same error. return column.to_pylist() fields = [ ArrowTableToRowsConversion._to_pylist(column.field(i)) for i in range(column.type.num_fields) ] if column.null_count == 0: if not names: return [{} for _ in range(n)] return [dict(zip(names, row)) for row in zip(*fields)] valid = column.is_valid().to_numpy(zero_copy_only=False).tolist() if not names: return [{} if m else None for m in valid] return [dict(zip(names, row)) if m else None for row, m in zip(zip(*fields), valid)] return column.to_pylist() @staticmethod def _need_converter(dataType: DataType) -> bool: if isinstance(dataType, NullType): return True elif isinstance(dataType, StructType): return True elif isinstance(dataType, ArrayType): return ArrowTableToRowsConversion._need_converter(dataType.elementType) elif isinstance(dataType, MapType): # Different from PySpark, here always needs conversion, # since the input from Arrow is a list of tuples. return True elif isinstance(dataType, BinaryType): return True elif isinstance(dataType, (TimestampType, TimestampNTZType)): # Always remove the time zone info for now return True elif isinstance(dataType, UserDefinedType): return True elif isinstance(dataType, VariantType): return True elif isinstance(dataType, GeographyType): return True elif isinstance(dataType, GeometryType): return True else: return False @overload @staticmethod def _create_converter(dataType: DataType, *, binary_as_bytes: bool = True) -> Callable: pass @overload @staticmethod def _create_converter( dataType: DataType, *, none_on_identity: bool, binary_as_bytes: bool = True ) -> Optional[Callable]: pass @staticmethod def _create_converter( dataType: DataType, *, none_on_identity: bool = False, binary_as_bytes: bool = True ) -> Optional[Callable]: assert dataType is not None and isinstance(dataType, DataType) if not ArrowTableToRowsConversion._need_converter(dataType): if none_on_identity: return None else: return lambda value: value if isinstance(dataType, NullType): return lambda value: None elif isinstance(dataType, StructType): field_names = dataType.names dedup_field_names = _dedup_names(field_names) field_convs = [ ArrowTableToRowsConversion._create_converter( f.dataType, none_on_identity=True, binary_as_bytes=binary_as_bytes ) for f in dataType.fields ] def convert_struct(value: Any) -> Any: if value is None: return None else: assert isinstance(value, dict) _values = [ ( field_convs[i](value.get(name, None)) # type: ignore[misc] if field_convs[i] is not None else value.get(name, None) ) for i, name in enumerate(dedup_field_names) ] return _create_row(field_names, _values) return convert_struct elif isinstance(dataType, ArrayType): element_conv = ArrowTableToRowsConversion._create_converter( dataType.elementType, none_on_identity=True, binary_as_bytes=binary_as_bytes ) assert element_conv is not None, ( f"_need_converter() returned True for ArrayType of {dataType.elementType}" ) def convert_array(value: Any) -> Any: if value is None: return None else: assert isinstance(value, list) return [element_conv(v) for v in value] return convert_array elif isinstance(dataType, MapType): key_conv = ArrowTableToRowsConversion._create_converter( dataType.keyType, none_on_identity=True, binary_as_bytes=binary_as_bytes ) value_conv = ArrowTableToRowsConversion._create_converter( dataType.valueType, none_on_identity=True, binary_as_bytes=binary_as_bytes ) if key_conv is None: if value_conv is None: def convert_map(value: Any) -> Any: if value is None: return None else: assert isinstance(value, list) assert all(isinstance(t, tuple) and len(t) == 2 for t in value) return dict(value) else: def convert_map(value: Any) -> Any: if value is None: return None else: assert isinstance(value, list) assert all(isinstance(t, tuple) and len(t) == 2 for t in value) return dict((t[0], value_conv(t[1])) for t in value) else: if value_conv is None: def convert_map(value: Any) -> Any: if value is None: return None else: assert isinstance(value, list) assert all(isinstance(t, tuple) and len(t) == 2 for t in value) return dict((key_conv(t[0]), t[1]) for t in value) else: def convert_map(value: Any) -> Any: if value is None: return None else: assert isinstance(value, list) assert all(isinstance(t, tuple) and len(t) == 2 for t in value) return dict((key_conv(t[0]), value_conv(t[1])) for t in value) return convert_map elif isinstance(dataType, BinaryType): def convert_binary(value: Any) -> Any: if value is None: return None else: assert isinstance(value, bytes) return value if binary_as_bytes else bytearray(value) return convert_binary elif isinstance(dataType, TimestampType): def convert_timestamp(value: Any) -> Any: if value is None: return None else: assert isinstance(value, datetime.datetime) return value.astimezone().replace(tzinfo=None) return convert_timestamp elif isinstance(dataType, TimestampNTZType): def convert_timestamp_ntz(value: Any) -> Any: if value is None: return None else: assert isinstance(value, datetime.datetime) return value return convert_timestamp_ntz elif isinstance(dataType, UserDefinedType): udt: UserDefinedType = dataType conv = ArrowTableToRowsConversion._create_converter( udt.sqlType(), none_on_identity=True, binary_as_bytes=binary_as_bytes ) if conv is None: def convert_udt(value: Any) -> Any: if value is None: return None else: return udt.deserialize(value) else: def convert_udt(value: Any) -> Any: if value is None: return None else: return udt.deserialize(conv(value)) return convert_udt elif isinstance(dataType, VariantType): def convert_variant(value: Any) -> Any: if value is None: return None elif ( isinstance(value, dict) and all(key in value for key in ["value", "metadata"]) and all(isinstance(value[key], bytes) for key in ["value", "metadata"]) ): return VariantVal(value["value"], value["metadata"]) else: raise PySparkValueError(errorClass="MALFORMED_VARIANT") return convert_variant elif isinstance(dataType, GeographyType): def convert_geography(value: Any) -> Any: if value is None: return None elif ( isinstance(value, dict) and all(key in value for key in ["wkb", "srid"]) and isinstance(value["wkb"], bytes) and isinstance(value["srid"], int) ): return Geography.fromWKB(value["wkb"], value["srid"]) else: raise PySparkValueError(errorClass="MALFORMED_GEOGRAPHY") return convert_geography elif isinstance(dataType, GeometryType): def convert_geometry(value: Any) -> Any: if value is None: return None elif ( isinstance(value, dict) and all(key in value for key in ["wkb", "srid"]) and isinstance(value["wkb"], bytes) and isinstance(value["srid"], int) ): return Geometry.fromWKB(value["wkb"], value["srid"]) else: raise PySparkValueError(errorClass="MALFORMED_GEOMETRY") return convert_geometry else: # pragma: no cover assert False, f"Need converter for {dataType} but failed to find one." @overload @staticmethod def convert(table: "pa.Table", schema: StructType) -> List[Row]: pass @overload @staticmethod def convert(table: "pa.Table", schema: StructType, *, binary_as_bytes: bool) -> List[Row]: pass @overload @staticmethod def convert( table: "pa.Table", schema: StructType, *, return_as_tuples: bool ) -> List[Row | tuple]: pass @staticmethod # type: ignore[misc] def convert( table: "pa.Table", schema: StructType, *, return_as_tuples: bool = False, binary_as_bytes: bool = True, ) -> List[Union[Row, tuple]]: require_minimum_pyarrow_version() import pyarrow as pa assert isinstance(table, pa.Table) assert schema is not None and isinstance(schema, StructType) # YearMonthIntervalType is serialized by the JVM as an Arrow YEAR_MONTH interval, which # PyArrow cannot materialize into Python values: `to_pylist()` raises an opaque # `KeyError: <Arrow type id>` from `get_array_class_from_type`. That lookup fails for an # empty column too (it resolves the array class before reading any element), so the check # below is intentionally unconditional in the row count -- it covers empty results as well, # surfacing a clean NOT_IMPLEMENTED instead of the opaque KeyError. Collecting such a value # is therefore not supported in the Spark Connect client; raise the same NOT_IMPLEMENTED # error as the classic PySpark path (YearMonthIntervalType.fromInternal). Note that, unlike # classic, PYSPARK_YM_INTERVAL_LEGACY (returning the integer months) cannot be honored here, # and an empty result raises rather than returning [] as classic would. if any(_has_type(f.dataType, YearMonthIntervalType) for f in schema.fields): raise PySparkNotImplementedError( errorClass="NOT_IMPLEMENTED", messageParameters={ "feature": "Collecting a year-month interval value in Spark Connect" }, ) fields = schema.fieldNames() if len(fields) > 0: field_converters = [ ArrowTableToRowsConversion._create_converter( f.dataType, none_on_identity=True, binary_as_bytes=binary_as_bytes ) for f in schema.fields ] columnar_data = [ ( [conv(v) for v in ArrowTableToRowsConversion._to_pylist(column)] if conv is not None else ArrowTableToRowsConversion._to_pylist(column) ) for column, conv in zip(table.columns, field_converters) ] if return_as_tuples: rows = [tuple(cols) for cols in zip(*columnar_data)] else: rows = [_create_row(fields, tuple(cols)) for cols in zip(*columnar_data)] assert len(rows) == table.num_rows, f"{len(rows)}, {table.num_rows}" return rows else: if return_as_tuples: return [tuple()] * table.num_rows else: return [_create_row(fields, tuple())] * table.num_rows class ArrowArrayConversion: @classmethod def check_conversion( cls, pa_type: "pa.DataType", check_type: Callable[["pa.DataType"], bool], ) -> bool: import pyarrow.types as types if check_type(pa_type): return True elif ( types.is_list(pa_type) or types.is_large_list(pa_type) or types.is_fixed_size_list(pa_type) or types.is_dictionary(pa_type) ): return cls.check_conversion(pa_type.value_type, check_type) elif types.is_map(pa_type): return any( cls.check_conversion(at, check_type) for at in [ pa_type.key_type, pa_type.item_type, ] ) elif types.is_struct(pa_type): return any(cls.check_conversion(field.type, check_type) for field in pa_type) else: return False @classmethod def convert_array( cls, arr: "pa.Array", check_type: Callable[["pa.DataType"], bool], convert: Callable[["pa.Array"], "pa.Array"], ) -> "pa.Array": import pyarrow as pa import pyarrow.types as types assert isinstance(arr, pa.Array) pa_type = arr.type # fastpath if not cls.check_conversion(pa_type, check_type): return arr if check_type(pa_type): converted = convert(arr) assert len(converted) == len(arr), f"array length changed: {arr} -> {converted}" return converted elif types.is_list(pa_type): return pa.ListArray.from_arrays( offsets=arr.offsets, values=cls.convert_array(arr.values, check_type, convert), ) elif types.is_large_list(pa_type): return pa.LargeListType.from_arrays( offsets=arr.offsets, values=cls.convert_array(arr.values, check_type, convert), ) elif types.is_fixed_size_list(pa_type): return pa.FixedSizeListArray.from_arrays( values=cls.convert_array(arr.values, check_type, convert), ) elif types.is_dictionary(pa_type): return pa.DictionaryArray.from_arrays( indices=arr.indices, dictionary=cls.convert_array(arr.dictionary, check_type, convert), ) elif types.is_map(pa_type): return pa.MapArray.from_arrays( offsets=arr.offsets, keys=cls.convert_array(arr.keys, check_type, convert), items=cls.convert_array(arr.items, check_type, convert), ) elif types.is_struct(pa_type): return pa.StructArray.from_arrays( arrays=[ cls.convert_array(arr.field(i), check_type, convert) for i in range(len(arr.type)) ], names=arr.type.names, ) else: # pragma: no cover assert False, f"Need converter for {pa_type} but failed to find one." @classmethod def convert( cls, arr: Union["pa.Array", "pa.ChunkedArray"], check_type: Callable[["pa.DataType"], bool], convert: Callable[["pa.Array"], "pa.Array"], ) -> Union["pa.Array", "pa.ChunkedArray"]: import pyarrow as pa assert isinstance(arr, (pa.Array, pa.ChunkedArray)) # fastpath if not cls.check_conversion(arr.type, check_type): return arr if isinstance(arr, pa.Array): return cls.convert_array(arr, check_type, convert) else: return pa.chunked_array( (cls.convert_array(a, check_type, convert) for a in arr.iterchunks()) ) @classmethod def localize_tz( cls, arr: Union["pa.Array", "pa.ChunkedArray"], ) -> Union["pa.Array", "pa.ChunkedArray"]: """ Convert Arrow timezone-aware timestamps to timezone-naive in the specified timezone. This function works on Arrow Arrays, and it recurses to convert nested types. This function is dedicated for Pandas UDF execution. Differences from _create_converter_to_pandas + _check_series_convert_timestamps_local_tz: 1, respect the timezone field in pyarrow timestamp type; 2, do not use local time at any time; 3, handle nested types in a consistent way. (_create_converter_to_pandas handles simple timestamp series with session timezone, but handles nested series with datetime.timezone.utc) Differences from _check_arrow_array_timestamps_localize: 1, respect the timezone field in pyarrow timestamp type; 2, do not handle timezone-naive timestamp; 3, do not support unit coercion which won't happen in UDF execution. Parameters ---------- arr : :class:`pyarrow.Array` Returns ------- :class:`pyarrow.Array` Notes ----- Arrow UDF (@arrow_udf/mapInArrow/etc) always preserve the original timezone, and thus doesn't need this conversion. """ import pyarrow as pa import pyarrow.types as types import pyarrow.compute as pc def check_type_func(pa_type: pa.DataType) -> bool: # match timezone-aware TimestampType return types.is_timestamp(pa_type) and pa_type.tz is not None def convert_func(arr: pa.Array) -> pa.Array: assert isinstance(arr, pa.TimestampArray) # import datetime # from zoneinfo import ZoneInfo # ts = datetime.datetime(2022, 1, 5, 15, 0, 1, tzinfo=ZoneInfo('Asia/Singapore')) # arr = pa.array([ts]) # arr[0] # <pyarrow.TimestampScalar: '2022-01-05T15:00:01.000000+0800'> # arr = pc.local_timestamp(arr) # arr[0] # <pyarrow.TimestampScalar: '2022-01-05T15:00:01.000000'> return pc.local_timestamp(arr) return cls.convert( arr, check_type=check_type_func, convert=convert_func, ) @classmethod def preprocess_time( cls, arr: Union["pa.Array", "pa.ChunkedArray"], ) -> Union["pa.Array", "pa.ChunkedArray"]: """ 1, always drop the timezone from TimestampType; 2, coerce_temporal_nanoseconds: coerce timestamp time units to nanoseconds """ import pyarrow as pa import pyarrow.types as types import pyarrow.compute as pc def check_type_func(pa_type: pa.DataType) -> bool: return types.is_timestamp(pa_type) and (pa_type.unit != "ns" or pa_type.tz is not None) def convert_func(arr: pa.Array) -> pa.Array: assert isinstance(arr, pa.TimestampArray) pa_type = arr.type if pa_type.tz is not None: arr = pc.local_timestamp(arr) if pa_type.unit != "ns": arr = pc.cast(arr, target_type=pa.timestamp("ns", tz=None)) return arr return cls.convert( arr, check_type=check_type_func, convert=convert_func, ) class ArrowArrayToPandasConversion: """ Conversion utilities for converting PyArrow Arrays and ChunkedArrays to pandas. This class provides methods to convert PyArrow columnar data structures to pandas Series or DataFrames, with support for Spark-specific type handling and conversions. The class is primarily used by PySpark's Arrow-based serializers for UDF execution, where Arrow data needs to be converted to pandas for Python UDF processing. """ @classmethod def convert( cls, arr: Union["pa.Array", "pa.ChunkedArray"], spark_type: DataType, *, ser_name: Optional[str] = None, timezone: Optional[str] = None, struct_in_pandas: str = "dict", ndarray_as_list: bool = False, prefer_int_ext_dtype: bool = False, df_for_struct: bool = False, ) -> Union["pd.Series", "pd.DataFrame"]: """ Convert a PyArrow Array or ChunkedArray to a pandas Series or DataFrame. Parameters ---------- arr : pa.Array or pa.ChunkedArray The Arrow column to convert. spark_type : DataType The target Spark type for the column to be converted to. ser_name : str The name of returned pd.Series. If not set, will try to get it from arr._name. timezone : str, optional Timezone for timestamp conversion. Required if the data contains timestamp types. struct_in_pandas : str, optional How to represent struct types in pandas. Valid values are "dict", "row", or "legacy". Default is "dict". ndarray_as_list : bool, optional Whether to convert numpy ndarrays to Python lists. Default is False. prefer_int_ext_dtype : bool, optional Whether to convert integers to Pandas ExtensionDType. df_for_struct : bool, optional If True, convert struct columns to a DataFrame with columns corresponding to struct fields instead of a Series. Default is False. Returns ------- pd.Series or pd.DataFrame Converted pandas Series. If df_for_struct is True and the type is StructType, returns a DataFrame with columns corresponding to struct fields. """ if cls._prefer_convert_numpy(spark_type, df_for_struct): return cls.convert_numpy( arr, spark_type, ser_name=ser_name, timezone=timezone, struct_in_pandas=struct_in_pandas, ndarray_as_list=ndarray_as_list, prefer_int_ext_dtype=prefer_int_ext_dtype, df_for_struct=df_for_struct, ) return cls.convert_legacy( arr, spark_type, timezone=timezone, struct_in_pandas=struct_in_pandas, ndarray_as_list=ndarray_as_list, df_for_struct=df_for_struct, ) @classmethod def convert_legacy( cls, arr: Union["pa.Array", "pa.ChunkedArray"], spark_type: DataType, *, timezone: Optional[str] = None, struct_in_pandas: Optional[str] = None, ndarray_as_list: bool = False, df_for_struct: bool = False, ) -> Union["pd.Series", "pd.DataFrame"]: """ Convert a PyArrow Array or ChunkedArray to a pandas Series or DataFrame. This is the lower-level conversion method that requires explicit Spark type specification. For a more convenient API, see :meth:`convert`. Parameters ---------- arr : pa.Array or pa.ChunkedArray The arrow column to convert. spark_type : DataType Target Spark type. Must be specified and should match the Arrow array type. timezone : str, optional The timezone to use for timestamp conversion. Required if the data contains timestamp types. struct_in_pandas : str, optional How to handle struct types in pandas. Valid values are "dict", "row", or "legacy". Required if the data contains struct types. ndarray_as_list : bool, optional Whether to convert numpy ndarrays to Python lists. Default is False. df_for_struct : bool, optional If True and spark_type is a StructType, return a DataFrame with columns corresponding to struct fields instead of a Series. Default is False. Returns ------- pd.Series or pd.DataFrame Converted pandas Series. If df_for_struct is True and spark_type is StructType, returns a DataFrame with columns corresponding to struct fields. Notes ----- This method handles date type columns specially to avoid overflow issues with datetime64[ns] intermediate representations. """ import pyarrow as pa import pandas as pd assert isinstance(arr, (pa.Array, pa.ChunkedArray)) if df_for_struct and isinstance(spark_type, StructType): import pyarrow.types as types assert types.is_struct(arr.type) assert len(spark_type.names) == len(arr.type.names), ( f"Schema mismatch: spark_type has {len(spark_type.names)} fields, " f"but arrow type has {len(arr.type.names)} fields. " f"spark_type={spark_type}, arrow_type={arr.type}" ) series = [ cls.convert_legacy( field_arr, spark_type=field.dataType, timezone=timezone, struct_in_pandas=struct_in_pandas, ndarray_as_list=ndarray_as_list, df_for_struct=False, # always False for child fields ) for field_arr, field in zip(arr.flatten(), spark_type) ] pdf = pd.concat(series, axis=1) pdf.columns = spark_type.names return pdf # Convert Arrow array to pandas Series with specific options: # - date_as_object: Convert date types to Python datetime.date objects directly # instead of datetime64[ns] to avoid overflow issues # - coerce_temporal_nanoseconds: Handle nanosecond precision timestamps correctly # - integer_object_nulls: Use object dtype for integer arrays with nulls pandas_options = { "date_as_object": True, "coerce_temporal_nanoseconds": True, "integer_object_nulls": True, } ser = arr.to_pandas(**pandas_options) converter = _create_converter_to_pandas( data_type=spark_type, nullable=True, timezone=timezone, struct_in_pandas=struct_in_pandas, error_on_duplicated_field_names=True, ndarray_as_list=ndarray_as_list, integer_object_nulls=True, ) return converter(ser) @classmethod def _prefer_convert_numpy( cls, spark_type: DataType, df_for_struct: bool, ) -> bool: supported_types = ( NullType, BinaryType, BooleanType, FloatType, DoubleType, ByteType, ShortType, IntegerType, LongType, DateType, TimeType, TimestampType, TimestampNTZType, UserDefinedType, VariantType, GeographyType, GeometryType, ) if df_for_struct and isinstance(spark_type, StructType): return all(isinstance(f.dataType, supported_types) for f in spark_type.fields) else: return isinstance(spark_type, supported_types) @classmethod def convert_numpy( cls, arr: Union["pa.Array", "pa.ChunkedArray"], spark_type: DataType, *, ser_name: Optional[str] = None, timezone: Optional[str] = None, struct_in_pandas: Optional[str] = None, ndarray_as_list: bool = False, prefer_int_ext_dtype: bool = False, df_for_struct: bool = False, ) -> Union["pd.Series", "pd.DataFrame"]: import pyarrow as pa import pandas as pd assert isinstance(arr, (pa.Array, pa.ChunkedArray)) if df_for_struct and isinstance(spark_type, StructType): import pyarrow.types as types assert types.is_struct(arr.type) assert len(spark_type.names) == len(arr.type.names), f"{spark_type} {arr.type} " return pd.concat( [ cls.convert_numpy( field_arr, spark_type=field.dataType, ser_name=field.name, timezone=timezone, struct_in_pandas=struct_in_pandas, ndarray_as_list=ndarray_as_list, prefer_int_ext_dtype=prefer_int_ext_dtype, df_for_struct=False, # always False for child fields ) for field_arr, field in zip(arr.flatten(), spark_type) ], axis=1, ) if ser_name is None: # Arrow array from batch.column(idx) contains name, # and this name will be used to rename the pandas series # returned by array.to_pandas(). # This name will be dropped after pa.compute functions. ser_name = arr._name arr = ArrowArrayConversion.preprocess_time(arr) series: pd.Series # conversion methods are selected based on benchmark python/benchmarks/bench_arrow.py if isinstance(spark_type, ByteType): if prefer_int_ext_dtype: series = arr.to_pandas(types_mapper=pd.ArrowDtype).astype(pd.Int8Dtype()) else: series = arr.to_pandas() elif isinstance(spark_type, ShortType): if prefer_int_ext_dtype: series = arr.to_pandas(types_mapper=pd.ArrowDtype).astype(pd.Int16Dtype()) else: series = arr.to_pandas() elif isinstance(spark_type, IntegerType): if prefer_int_ext_dtype: series = arr.to_pandas(types_mapper=pd.ArrowDtype).astype(pd.Int32Dtype()) else: series = arr.to_pandas() elif isinstance(spark_type, LongType): if prefer_int_ext_dtype: series = arr.to_pandas(types_mapper=pd.ArrowDtype).astype(pd.Int64Dtype()) else: series = arr.to_pandas() elif isinstance( spark_type, ( NullType, BinaryType, BooleanType, FloatType, DoubleType, DecimalType, StringType, DateType, TimeType, TimestampType, TimestampNTZType, DayTimeIntervalType, YearMonthIntervalType, ), ): series = arr.to_pandas() elif isinstance(spark_type, UserDefinedType): udt: UserDefinedType = spark_type series = arr.to_pandas() series = series.apply( lambda v: ( v if hasattr(v, "__UDT__") else udt.deserialize(v) if v is not None else None ) ) elif isinstance(spark_type, VariantType): series = arr.to_pandas() series = series.map( lambda v: VariantVal(v["value"], v["metadata"]) if v is not None else None ) elif isinstance(spark_type, GeographyType): series = arr.to_pandas() series = series.map( lambda v: Geography.fromWKB(v["wkb"], v["srid"]) if v is not None else None ) elif isinstance(spark_type, GeometryType): series = arr.to_pandas() series = series.map( lambda v: Geometry.fromWKB(v["wkb"], v["srid"]) if v is not None else None ) # elif isinstance( # spark_type, # ( # ArrayType, # MapType, # StructType, # ), # ): # TODO(SPARK-55324): Support complex types else: # pragma: no cover assert False, f"Need converter for {spark_type} but failed to find one." return series.rename(ser_name)