blob: e22745cc4fc72ddca7a2c125e98d8f37b7028e0f [file]
#
# 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 os
import sys
from typing import IO, Any, Callable, Optional
from pyspark.accumulators import (
SpecialAccumulatorIds,
_accumulatorRegistry,
_deserialize_accumulator,
)
from pyspark.serializers import (
SpecialLengths,
read_int,
write_int,
)
from pyspark.sql.profiler import (
ProfileResultsParam,
ProfileResultsParamV2,
WorkerMemoryProfiler,
WorkerPerfProfiler,
)
from pyspark.util import (
handle_worker_exception,
start_faulthandler_periodic_traceback,
with_faulthandler,
)
from pyspark.worker_util import (
Conf,
check_python_version,
send_accumulator_updates,
setup_broadcasts,
setup_memory_limits,
setup_spark_files,
)
class RunnerConf(Conf):
@property
def profiler(self) -> Optional[str]:
return self.get("spark.sql.pyspark.dataSource.profiler", None)
def is_method_overridden(reader: Any, name: str) -> bool:
"""
Whether `reader` overrides the `DataSourceReader` method `name`, rather than inheriting the
default implementation. Used to detect pushdown methods that a reader implements while the
corresponding pushdown configuration is disabled, so that they are not silently ignored.
"""
from pyspark.sql.datasource import DataSourceReader
return getattr(getattr(reader, name), "__func__", None) is not getattr(DataSourceReader, name)
def check_pushdown_not_disabled(
reader: Any, enable_filter_pushdown: bool, enable_limit_pushdown: bool
) -> None:
"""
Raise `DATA_SOURCE_PUSHDOWN_DISABLED` if `reader` implements a pushdown method while the
corresponding pushdown configuration is disabled, so that the method is not silently ignored.
This is shared by both planning workers: `plan_data_source_read` runs it for a plain read,
and `data_source_pushdown_filters` runs it too, because a filter- or limit-only scan caches
the read info in that worker and `plan_data_source_read` never runs for such a scan.
"""
from pyspark.errors import PySparkAssertionError
for method, conf, enabled in (
("pushFilters", "spark.sql.python.filterPushdown.enabled", enable_filter_pushdown),
("pushLimit", "spark.sql.python.limitPushdown.enabled", enable_limit_pushdown),
):
if not enabled and is_method_overridden(reader, method):
raise PySparkAssertionError(
errorClass="DATA_SOURCE_PUSHDOWN_DISABLED",
messageParameters={
"type": type(reader).__name__,
"method": method,
"conf": conf,
},
)
@with_faulthandler
def worker_run(main: Callable, infile: IO, outfile: IO) -> None:
try:
check_python_version(infile)
start_faulthandler_periodic_traceback()
memory_limit_mb = int(os.environ.get("PYSPARK_PLANNER_MEMORY_MB", "-1"))
setup_memory_limits(memory_limit_mb)
setup_spark_files(infile)
setup_broadcasts(infile)
conf = RunnerConf(infile)
_accumulatorRegistry.clear()
accumulator = _deserialize_accumulator(
SpecialAccumulatorIds.SQL_UDF_PROFIER, {}, ProfileResultsParam
)
accumulator_v2 = _deserialize_accumulator(
SpecialAccumulatorIds.SQL_UDF_PROFIER_V2, {}, ProfileResultsParamV2
)
if main.__module__ == "__main__":
try:
worker_module = sys.modules["__main__"].__spec__.name # type: ignore[union-attr]
except Exception:
worker_module = "__main__"
else:
worker_module = main.__module__
worker_module = worker_module.split(".")[-1]
if conf.profiler == "perf":
with WorkerPerfProfiler(accumulator, accumulator_v2, worker_module):
main(infile, outfile)
elif conf.profiler == "memory":
with WorkerMemoryProfiler(accumulator, accumulator_v2, worker_module, main):
main(infile, outfile)
else:
main(infile, outfile)
except BaseException as e:
handle_worker_exception(e, outfile)
sys.exit(-1)
send_accumulator_updates(outfile)
# check end of stream
if read_int(infile) == SpecialLengths.END_OF_STREAM:
write_int(SpecialLengths.END_OF_STREAM, outfile)
else:
# write a different value to tell JVM to not reuse this worker
write_int(SpecialLengths.END_OF_DATA_SECTION, outfile)
sys.exit(-1)