diff --git a/array_api_strict/_creation_functions.py b/array_api_strict/_creation_functions.py index f47a136..9c96967 100644 --- a/array_api_strict/_creation_functions.py +++ b/array_api_strict/_creation_functions.py @@ -48,6 +48,16 @@ def _supports_buffer_protocol(obj: object) -> TypeIs[SupportsBufferProtocol]: return True +def _contains_nested_array(obj: list[object] | tuple[object, ...]) -> bool: + from ._array_object import Array + + return any( + isinstance(item, Array) + or isinstance(item, list | tuple) and _contains_nested_array(item) + for item in obj + ) + + def asarray( obj: Array | complex | NestedSequence[complex] | SupportsBufferProtocol, /, @@ -97,7 +107,7 @@ def asarray( if isinstance(obj, Array): return Array._new(np.array(obj._array, copy=copy, dtype=_np_dtype), device=device) - elif isinstance(obj, list | tuple) and any(isinstance(x, Array) for x in obj): + elif isinstance(obj, list | tuple) and _contains_nested_array(obj): raise TypeError("Nested Arrays are not allowed. Use `stack` instead.") if dtype is None and isinstance(obj, int) and (obj > 2 ** 64 or obj < -(2 ** 63)): diff --git a/array_api_strict/tests/test_creation_functions.py b/array_api_strict/tests/test_creation_functions.py index acee119..d4971a3 100644 --- a/array_api_strict/tests/test_creation_functions.py +++ b/array_api_strict/tests/test_creation_functions.py @@ -116,6 +116,33 @@ def test_asarray_nested_arrays(): asarray([1, asarray(1)]) +@pytest.mark.parametrize("outer", [list, tuple]) +@pytest.mark.parametrize("inner", [list, tuple]) +@pytest.mark.parametrize("array_value", [1, [1]]) +@pytest.mark.parametrize("depth", [2, 3]) +def test_asarray_deeply_nested_arrays(outer, inner, array_value, depth): + obj = asarray(array_value) + for _ in range(depth - 1): + obj = inner([obj]) + obj = outer([obj]) + with pytest.raises(TypeError, match="Nested Arrays are not allowed"): + asarray(obj) + + +def test_asarray_nested_array_after_scalars(): + with pytest.raises(TypeError, match="Nested Arrays are not allowed"): + asarray([[1, 2], [3, asarray(4)]]) + + +@pytest.mark.parametrize("outer", [list, tuple]) +@pytest.mark.parametrize("inner", [list, tuple]) +def test_asarray_nested_scalars(outer, inner): + obj = outer([inner([1, 2]), inner([3, 4])]) + res = asarray(obj) + assert res.shape == (2, 2) + assert all(res == asarray([[1, 2], [3, 4]])) + + def test_asarray_device_inference(): assert asarray([1, 2, 3]).device == CPU_DEVICE