diff --git a/mypy.ini b/mypy.ini index 4b6f3f1..f372bf5 100644 --- a/mypy.ini +++ b/mypy.ini @@ -3,4 +3,5 @@ implicit_reexport = true explicit_package_bases = true strict = true show_error_context = true -show_error_codes = true \ No newline at end of file +show_error_codes = true +disable_error_code = overload-overlap diff --git a/test/cython/test_with_unit.py b/test/cython/test_with_unit.py index f480bcb..a1490eb 100644 --- a/test/cython/test_with_unit.py +++ b/test/cython/test_with_unit.py @@ -488,9 +488,9 @@ def test_get_item() -> None: # Wrong kinds of index (unit array, slice). with pytest.raises(UnitMismatchError): - _ = u[mps] + _ = u[mps] # type: ignore[call-overload] with pytest.raises(NotTUnitsLikeError): - _ = u[1:2] + _ = u[1:2] # type: ignore[call-overload] assert u[v / v] == 10 diff --git a/test/test_value.py b/test/test_value.py index 9440cde..c909b80 100644 --- a/test/test_value.py +++ b/test/test_value.py @@ -181,6 +181,11 @@ def test_get_item() -> None: assert Value(1, '')[Value(1, '')] == 1 assert Value(1, '')[ns / s] == 10**9 + x = np.float64(3) + unit = ns + for key in ..., (), (None,), (None, None), (None, None, None): + assert Value(x, unit)[key] == (unit * x[key]) + def test_cycles() -> None: from tunits.units import cyc, rad diff --git a/tunits/core/__init__.pyi b/tunits/core/__init__.pyi index 71f66ec..54cbdd3 100644 --- a/tunits/core/__init__.pyi +++ b/tunits/core/__init__.pyi @@ -1,4 +1,5 @@ from typing import ClassVar, Sequence, Any, Callable, Iterator, overload, Generic, SupportsIndex +from types import EllipsisType import abc from attrs import frozen from typing_extensions import TypeVar @@ -319,12 +320,35 @@ class Value(Generic[NumericalT], WithUnit, np.generic, SupportsIndex): def __ge__(self, other: _NUMERICAL_TYPE_OR_ARRAY_OR_GENERIC_UNIT) -> bool: ... def __eq__(self, other: Any) -> bool: ... def __neq__(self, other: Any) -> bool: ... - def __getitem__(self, key: Any) -> NumericalT: ... def value_in_base_units(self) -> NumericalT: ... def __index__(self) -> int: ... def __pow__(self, other: Any, modulus: Any = None) -> Value: ... def sign(self) -> int: ... def dimensionless(self) -> float: ... + @overload + def __getitem__(self, other: tuple[()], /) -> Value[NumericalT]: ... + @overload + def __getitem__( + self, other: EllipsisType | tuple[EllipsisType], / + ) -> np.ndarray[tuple[()], np.dtype[Value[NumericalT]]]: ... + @overload + def __getitem__( + self, other: tuple[None] | None, / + ) -> np.ndarray[tuple[int], np.dtype[Value[NumericalT]]]: ... + @overload + def __getitem__( + self, other: tuple[None, None], / + ) -> np.ndarray[tuple[int, int], np.dtype[Value[NumericalT]]]: ... + @overload + def __getitem__( + self, other: tuple[None, None, None], / + ) -> np.ndarray[tuple[int, int, int], np.dtype[Value[NumericalT]]]: ... + @overload + def __getitem__(self: ValueType, other: ValueType | str) -> NumericalT: ... + @overload + def __getitem__( + self, other: tuple[None, ...] + ) -> np.ndarray[tuple[Any, ...], np.dtype[Value[NumericalT]]]: ... class ValueArray(Generic[ValueType2], WithUnit): value: NDArray[Any] @@ -426,7 +450,6 @@ class ValueWithDimension(abc.ABC, Value): @classmethod def is_valid(cls, v: WithUnit) -> bool: ... - def __getitem__(self: ValueType2, unit: ValueType2 | str) -> float: ... class ArrayWithDimension(abc.ABC, ValueArray[ValueType2]): @staticmethod diff --git a/tunits/core/cython/with_unit.pyx b/tunits/core/cython/with_unit.pyx index a14bc3a..446cd38 100644 --- a/tunits/core/cython/with_unit.pyx +++ b/tunits/core/cython/with_unit.pyx @@ -576,6 +576,8 @@ cdef class WithUnit: except NotTUnitsLikeError: return NotImplemented except TypeError: + if (key is Ellipsis or isinstance(key, tuple)) and not isinstance(self.value, np.ndarray): + return self.__with_value(np.float64(self.value)[key]) try: unit_val = _try_interpret_as_with_unit(str(key), True) except: