Skip to content
Open
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
2 changes: 1 addition & 1 deletion python/pyarrow/_compute.pyx
Original file line number Diff line number Diff line change
Expand Up @@ -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
----------
Expand Down
2 changes: 1 addition & 1 deletion python/pyarrow/_dataset.pyx
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
3 changes: 3 additions & 0 deletions python/pyarrow/tests/test_compute.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Expand Down
3 changes: 3 additions & 0 deletions python/pyarrow/tests/test_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Expand Down
18 changes: 18 additions & 0 deletions python/pyarrow/tests/test_misc.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Loading