/
githubmirror
/
spark
Обзор
Документация
Войти
/
githubmirror
/
spark
Код
Запросы
0
Пакеты
0
Релизы
0
Аналитика
Безопасность
master
python/pyspark/tests/test_install_spark.py
190 строк
7 KB
Peter Toth
[SPARK-57962][PYTHON] Guard against path traversal in install_spark tar extraction
10 июл 2026, 13:16
10 июл 2026, 13:16
d61f250
Код
Авторство
О чём код?
# # 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 io import os import re import tarfile import tempfile import unittest import urllib.request from pyspark.install import ( get_preferred_mirrors, install_spark, _extract_tar, DEFAULT_HADOOP, DEFAULT_HIVE, UNSUPPORTED_COMBINATIONS, checked_versions, checked_package_name, ) class SparkInstallationTestCase(unittest.TestCase): def get_latest_spark_version(self): if "PYSPARK_RELEASE_MIRROR" in os.environ: sites = [os.environ["PYSPARK_RELEASE_MIRROR"]] else: sites = get_preferred_mirrors() # Filter out the archive sites sites = [site for site in sites if "archive.apache.org" not in site] for site in sites: url = site + "/spark/" try: with urllib.request.urlopen(url) as response: html = response.read().decode("utf-8") versions = re.findall(r"spark-(\d+[\.-]\d+[\.-]\d+)/", html) versions = [v.replace("-", ".") for v in versions] return max(versions) except Exception: continue return None def test_install_spark(self): # Test only one case. Testing this is expensive because it needs to download # the Spark distribution. We try to get the latest version, but if we can't, # we just use a hard-coded version. spark_version = self.get_latest_spark_version() or "4.1.1" spark_version, hadoop_version, hive_version = checked_versions(spark_version, "3", "2.3") with tempfile.TemporaryDirectory(prefix="test_install_spark") as tmp_dir: install_spark( dest=tmp_dir, spark_version=spark_version, hadoop_version=hadoop_version, hive_version=hive_version, ) self.assertTrue(os.path.isdir("%s/jars" % tmp_dir)) self.assertTrue(os.path.exists("%s/bin/spark-submit" % tmp_dir)) self.assertTrue(os.path.exists("%s/RELEASE" % tmp_dir)) def test_extract_tar(self): # A benign member is extracted with the top-level package directory # stripped, while a member whose path escapes the destination # directory (zip slip) is rejected rather than extracted. package_name = "spark-4.1.1-bin-hadoop3" def make_tar(path, member_name): with tarfile.open(path, "w") as tar: data = b"content" info = tarfile.TarInfo(name=member_name) info.size = len(data) tar.addfile(info, io.BytesIO(data)) with tempfile.TemporaryDirectory(prefix="test_install_spark") as tmp_dir: # Benign archive extracts into dest with the package prefix removed. safe_tar = os.path.join(tmp_dir, "safe.tar") make_tar(safe_tar, "%s/bin/spark-submit" % package_name) safe_dest = os.path.join(tmp_dir, "safe_dest") os.makedirs(safe_dest) with tarfile.open(safe_tar, "r") as tar: _extract_tar(tar, package_name, safe_dest) self.assertTrue(os.path.exists(os.path.join(safe_dest, "bin", "spark-submit"))) # Malicious archive member escaping dest is rejected, and nothing # is written into the destination directory. evil_tar = os.path.join(tmp_dir, "evil.tar") make_tar(evil_tar, "%s/../../evil" % package_name) evil_dest = os.path.join(tmp_dir, "evil_dest") os.makedirs(evil_dest) with tarfile.open(evil_tar, "r") as tar: with self.assertRaisesRegex(ValueError, "outside of the destination"): _extract_tar(tar, package_name, evil_dest) self.assertEqual([], os.listdir(evil_dest)) def test_package_name(self): self.assertEqual( "spark-3.0.0-bin-hadoop3.2", checked_package_name("spark-3.0.0", "hadoop3.2", "hive2.3") ) spark_version, hadoop_version, hive_version = checked_versions("3.2.0", "3", "2.3") self.assertEqual( "spark-3.2.0-bin-hadoop3.2", checked_package_name(spark_version, hadoop_version, hive_version), ) spark_version, hadoop_version, hive_version = checked_versions("3.3.0", "3", "2.3") self.assertEqual( "spark-3.3.0-bin-hadoop3", checked_package_name(spark_version, hadoop_version, hive_version), ) def test_checked_versions(self): test_version = "3.0.1" # Just pick one version to test. # Positive test cases self.assertEqual( ("spark-2.4.1", "without-hadoop", "hive2.3"), checked_versions("2.4.1", "without", "2.3"), ) self.assertEqual( ("spark-3.0.1", "without-hadoop", "hive2.3"), checked_versions("spark-3.0.1", "without-hadoop", "hive2.3"), ) self.assertEqual( ("spark-3.3.0", "hadoop3", "hive2.3"), checked_versions("spark-3.3.0", "hadoop3", "hive2.3"), ) # Prerelease version (e.g. pip dev builds) self.assertEqual( ("spark-4.2.0.dev4", "hadoop3", "hive2.3"), checked_versions("4.2.0.dev4", "3", "2.3"), ) self.assertEqual( ("spark-4.2.0.dev4", "hadoop3", "hive2.3"), checked_versions("spark-4.2.0.dev4", "hadoop3", "hive2.3"), ) # Negative test cases for hadoop_version, hive_version in UNSUPPORTED_COMBINATIONS: with self.assertRaisesRegex(RuntimeError, "Hive.*should.*Hadoop"): checked_versions( spark_version=test_version, hadoop_version=hadoop_version, hive_version=hive_version, ) with self.assertRaisesRegex(RuntimeError, "Spark version should start with 'spark-'"): checked_versions( spark_version="malformed", hadoop_version=DEFAULT_HADOOP, hive_version=DEFAULT_HIVE ) with self.assertRaisesRegex(RuntimeError, "Spark distribution.*malformed.*"): checked_versions( spark_version=test_version, hadoop_version="malformed", hive_version=DEFAULT_HIVE ) with self.assertRaisesRegex(RuntimeError, "Spark distribution.*malformed.*"): checked_versions( spark_version=test_version, hadoop_version=DEFAULT_HADOOP, hive_version="malformed" ) with self.assertRaisesRegex(RuntimeError, "Spark distribution of hive1.2 is not supported"): checked_versions( spark_version=test_version, hadoop_version="hadoop3", hive_version="hive1.2" ) if __name__ == "__main__": from pyspark.testing import main main()