diff --git a/array_api_tests/test_operators_and_elementwise_functions.py b/array_api_tests/test_operators_and_elementwise_functions.py index 6ab7ec34..a1247459 100644 --- a/array_api_tests/test_operators_and_elementwise_functions.py +++ b/array_api_tests/test_operators_and_elementwise_functions.py @@ -2104,6 +2104,59 @@ def test_subtract(ctx, data): raise +def make_reflected_scalar_params(): + cases = ( + ("__radd__", "add", operator.add, dh.numeric_dtypes), + ("__rsub__", "subtract", operator.sub, dh.numeric_dtypes), + ("__rmul__", "multiply", operator.mul, dh.numeric_dtypes), + ("__rtruediv__", "divide", operator.truediv, dh.all_float_dtypes), + ("__rfloordiv__", "floor_divide", operator.floordiv, dh.real_dtypes), + ("__rpow__", "pow", operator.pow, dh.numeric_dtypes), + ("__rmod__", "remainder", operator.mod, dh.real_dtypes), + ("__rand__", "bitwise_and", operator.and_, dh.bool_and_all_int_dtypes), + ("__ror__", "bitwise_or", operator.or_, dh.bool_and_all_int_dtypes), + ("__rxor__", "bitwise_xor", operator.xor, dh.bool_and_all_int_dtypes), + ("__rlshift__", "bitwise_left_shift", operator.lshift, dh.all_int_dtypes), + ("__rrshift__", "bitwise_right_shift", operator.rshift, dh.all_int_dtypes), + ) + for method, func_name, op, dtypes in cases: + for dtype in dtypes: + if isinstance(dtype, xp._UndefinedStub): + continue + yield pytest.param( + method, func_name, op, dtype, + id=f"{method}-{dh.dtype_to_name[dtype]}", + ) + + +@pytest.mark.parametrize( + "method, func_name, op, dtype", list(make_reflected_scalar_params()) +) +@given(data=st.data()) +def test_reflected_scalar_operator(method, func_name, op, dtype, data): + # Small positive values avoid division by zero, invalid powers, and overflow. + if dtype == xp.bool: + elements = st.booleans() + else: + elements = st.integers(1, 3).map(dh.get_scalar_type(dtype)) + scalar_strategy = elements + if func_name == "subtract" and dtype in dh.uint_dtypes: + scalar_strategy = st.integers(3, 7) + scalar = data.draw(scalar_strategy, label="scalar") + x = data.draw( + hh.arrays( + dtype=dtype, shape=hh.shapes(min_dims=1, min_side=1), elements=elements + ), + label="x", + ) + out = op(scalar, x) + expected = getattr(xp, func_name)(xp.asarray(scalar, dtype=dtype), x) + + ph.assert_dtype(method, in_dtype=dtype, out_dtype=out.dtype) + ph.assert_result_shape(method, in_shapes=[x.shape], out_shape=out.shape) + ph.assert_array_elements(method, out=out, expected=expected) + + @given(hh.arrays(dtype=hh.all_floating_dtypes(), shape=hh.shapes())) def test_tan(x): repro_snippet = ph.format_snippet(f"xp.tan({x!r})")