Skip to content
Merged
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
3 changes: 2 additions & 1 deletion mypy.ini
Original file line number Diff line number Diff line change
Expand Up @@ -3,4 +3,5 @@ implicit_reexport = true
explicit_package_bases = true
strict = true
show_error_context = true
show_error_codes = true
show_error_codes = true
disable_error_code = overload-overlap
4 changes: 2 additions & 2 deletions test/cython/test_with_unit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
5 changes: 5 additions & 0 deletions test/test_value.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
27 changes: 25 additions & 2 deletions tunits/core/__init__.pyi
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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]
Expand Down Expand Up @@ -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
Expand Down
2 changes: 2 additions & 0 deletions tunits/core/cython/with_unit.pyx
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
Loading