diff --git a/python/pyarrow/_compute.pyx b/python/pyarrow/_compute.pyx index 1c0c779f6037..b166bc9bf7cd 100644 --- a/python/pyarrow/_compute.pyx +++ b/python/pyarrow/_compute.pyx @@ -2708,7 +2708,7 @@ cdef class Expression(_Weakrefable): cdef inline CExpression unwrap(self): return self.expr - def equals(self, Expression other): + def equals(self, Expression other not None): """ Parameters ---------- diff --git a/python/pyarrow/_dataset.pyx b/python/pyarrow/_dataset.pyx index d40614a61fc0..f5666d69e1dd 100644 --- a/python/pyarrow/_dataset.pyx +++ b/python/pyarrow/_dataset.pyx @@ -4052,7 +4052,7 @@ cdef class Scanner(_Weakrefable): return reader -def get_partition_keys(Expression partition_expression): +def get_partition_keys(Expression partition_expression not None): """ Extract partition keys (equality constraints between a field and a scalar) from an expression as a dict mapping the field's name to its value. diff --git a/python/pyarrow/tests/test_compute.py b/python/pyarrow/tests/test_compute.py index 797fbc220ec3..2ccef3ac6eb3 100644 --- a/python/pyarrow/tests/test_compute.py +++ b/python/pyarrow/tests/test_compute.py @@ -4249,6 +4249,9 @@ def test_expression_construction(): with pytest.raises(TypeError): field.isin(1) + with pytest.raises(TypeError, match="Argument 'other' has incorrect type"): + field.equals(None) + with pytest.raises(pa.ArrowInvalid): field != object() diff --git a/python/pyarrow/tests/test_dataset.py b/python/pyarrow/tests/test_dataset.py index 0a94c0bd9875..707d937da201 100644 --- a/python/pyarrow/tests/test_dataset.py +++ b/python/pyarrow/tests/test_dataset.py @@ -946,6 +946,9 @@ def test_partition_keys(): null = ds.field('a').is_null() assert ds.get_partition_keys(null) == {'a': None} + with pytest.raises(TypeError, match="Argument 'partition_expression'"): + ds.get_partition_keys(None) + @pytest.mark.parquet def test_parquet_read_options(): diff --git a/python/pyarrow/tests/test_misc.py b/python/pyarrow/tests/test_misc.py index 856873f441b4..0d755e25297b 100644 --- a/python/pyarrow/tests/test_misc.py +++ b/python/pyarrow/tests/test_misc.py @@ -270,3 +270,21 @@ def test_extension_type_constructor_errors(klass): msg = f"Do not call {klass.__name__}'s constructor directly, use .* instead." with pytest.raises(TypeError, match=msg): klass() + + +@pytest.mark.processes +@pytest.mark.dataset +@pytest.mark.parametrize(("call", "argument"), [ + ("ds.get_partition_keys(None)", "partition_expression"), + ("ds.field('a').equals(None)", "other"), +]) +def test_expression_apis_reject_none_without_crashing( + call: str, argument: str) -> None: + code = f"""if 1: + import pytest + import pyarrow.dataset as ds + + with pytest.raises(TypeError, match="Argument '{argument}'"): + {call} + """ + subprocess.check_call(args=[sys.executable, "-c", code], timeout=30)