From fc26e4f4eafa61b28442c51471f058b6697b5a79 Mon Sep 17 00:00:00 2001 From: boulea7 Date: Mon, 28 Sep 2026 17:20:11 +0800 Subject: [PATCH] TST: cover reflected scalar operators --- ...est_operators_and_elementwise_functions.py | 44 +++++++++++++++++++ 1 file changed, 44 insertions(+) diff --git a/array_api_tests/test_operators_and_elementwise_functions.py b/array_api_tests/test_operators_and_elementwise_functions.py index 6ab7ec34..91ed7378 100644 --- a/array_api_tests/test_operators_and_elementwise_functions.py +++ b/array_api_tests/test_operators_and_elementwise_functions.py @@ -2104,6 +2104,50 @@ def test_subtract(ctx, data): raise +def make_reflected_scalar_params(): + cases = ( + ("__radd__", "add", operator.add, dh.numeric_dtypes, 7, (2, 3)), + ("__rsub__", "subtract", operator.sub, dh.numeric_dtypes, 7, (2, 3)), + ("__rmul__", "multiply", operator.mul, dh.numeric_dtypes, 7, (2, 3)), + ("__rtruediv__", "divide", operator.truediv, dh.all_float_dtypes, 6, (2, 4)), + ("__rfloordiv__", "floor_divide", operator.floordiv, dh.real_dtypes, 7, (2, 3)), + ("__rpow__", "pow", operator.pow, dh.numeric_dtypes, 3, (1, 2)), + ("__rmod__", "remainder", operator.mod, dh.real_dtypes, 7, (2, 3)), + ("__rand__", "bitwise_and", operator.and_, dh.bool_and_all_int_dtypes, 6, (1, 2)), + ("__ror__", "bitwise_or", operator.or_, dh.bool_and_all_int_dtypes, 6, (1, 2)), + ("__rxor__", "bitwise_xor", operator.xor, dh.bool_and_all_int_dtypes, 6, (1, 2)), + ("__rlshift__", "bitwise_left_shift", operator.lshift, dh.all_int_dtypes, 3, (1, 2)), + ("__rrshift__", "bitwise_right_shift", operator.rshift, dh.all_int_dtypes, 8, (1, 2)), + ) + for method, func_name, op, dtypes, scalar, values in cases: + for dtype in dtypes: + if isinstance(dtype, xp._UndefinedStub): + continue + if dtype == xp.bool: + operand, elements = True, (False, True) + else: + scalar_type = dh.get_scalar_type(dtype) + operand = scalar_type(scalar) + elements = tuple(scalar_type(value) for value in values) + yield pytest.param( + method, func_name, op, dtype, operand, elements, + id=f"{method}-{dh.dtype_to_name[dtype]}", + ) + + +@pytest.mark.parametrize( + "method, func_name, op, dtype, scalar, values", make_reflected_scalar_params() +) +def test_reflected_scalar_operator(method, func_name, op, dtype, scalar, values): + x = xp.asarray(values, dtype=dtype) + 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})")