Skip to content

Commit c6da061

Browse files
committed
Stabilize runtime correctness, typing parity, and docs for release
- Async-safe @shapix.check: preserve coroutine behavior and memo lifetime - Reject mixed Scalar in shape specs (e.g. F32[N, Scalar] raises TypeError) - Fix DT64/TD64 to accept unit-qualified dtypes (datetime64[ns], etc.) - Exclude booleans from numeric ScalarLike aliases and make_scalar_like_type - Robust __version__ fallback when package metadata is unavailable - Full type-checker parity: pyright, mypy, and ty via TypeVar/TypeAliasType under TYPE_CHECKING; unified test suite across all three checkers - Backend coverage: JAX/Torch Value(...) and boolean rejection tests - Docs: async support, Scalar constraints, DT64/TD64, boolean semantics, make_scalar_like_type location, checker-agnostic language
1 parent 35e13b3 commit c6da061

23 files changed

Lines changed: 1085 additions & 429 deletions

‎CLAUDE.md‎

Lines changed: 311 additions & 40 deletions
Large diffs are not rendered by default.

‎README.md‎

Lines changed: 47 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@ Shapix turns array shape annotations into **Python objects** that beartype valid
1818

1919
- **Zero boilerplate** — works with standard `@beartype` decorators and `beartype.claw` import hooks. No custom decorator required.
2020
- **Cross-argument consistency** — named dimensions are enforced across all parameters and the return value within a single function call.
21-
- **Static type checker friendly** — under `TYPE_CHECKING`, array types resolve to proper `NDArray` / `Array` / `Tensor` aliases. Pyright sees real types.
21+
- **Static type checker friendly** — under `TYPE_CHECKING`, array types resolve to proper `NDArray` / `Array` / `Tensor` aliases. Works with pyright, mypy, and ty.
2222
- **Readable annotations** — `F32[N, C, H, W]` reads like documentation.
2323
- **Full `BeartypeConf` support** — unlike jaxtyping, shapix doesn't replace your beartype configuration.
2424
- **Thread-safe** — each thread gets independent dimension bindings.
@@ -265,6 +265,8 @@ def dot(x: F32[N], y: F32[N]) -> F32[Scalar]:
265265
return np.dot(x, y) # returns shape ()
266266
```
267267

268+
> **Note:** `Scalar` must be the only shape token. Mixed forms like `F32[N, Scalar]` or `F32[Scalar, ...]` raise `TypeError` at hint construction time.
269+
268270
### Custom dimensions
269271

270272
Create your own with `Dimension`:
@@ -328,6 +330,8 @@ from shapix.numpy import F32, I64, Shaped # and many more
328330

329331
**Additional dtypes:** `V` (void), `Str` (string), `Bytes` (bytes), `Obj` (object), `DT64` (datetime64), `TD64` (timedelta64)
330332

333+
> `DT64` and `TD64` accept unit-qualified NumPy dtypes such as `datetime64[ns]`, `datetime64[D]`, `timedelta64[ms]`, etc.
334+
331335
### JAX
332336

333337
```python
@@ -428,7 +432,9 @@ from shapix.torch import F32Like # accepts Tensor, ndarray, scalars, sequences
428432

429433
### ScalarLike types (range-validated scalars)
430434

431-
ScalarLike types validate individual scalar values with range checking — no shape, just value:
435+
ScalarLike types validate individual scalar values with range checking — no shape, just value.
436+
437+
> **Note:** Numeric scalar aliases (`I8ScalarLike`, `F32ScalarLike`, `NumScalarLike`, etc.) reject `bool` and `np.bool_` values. Python `bool` is a subclass of `int`, but shapix treats booleans as non-numeric. Use `BoolScalarLike` for boolean scalars.
432438
433439
```python
434440
from shapix.numpy import I8ScalarLike, F32ScalarLike, U8ScalarLike
@@ -701,6 +707,8 @@ def f(x: F32[N, C], y: F32[N, C]) -> F32[N, C]: ...
701707

702708
If you want a guarantee that cross-argument checking works regardless of how your code is called (by test runners, async frameworks, deep middleware stacks), `@shapix.check` removes all dependence on call-stack structure.
703709

710+
`@shapix.check` supports both sync and async functions. For async functions, the memo scope covers the full awaited execution, and `inspect.iscoroutinefunction()` is preserved on the decorated function.
711+
704712
**When you don't need it:** If you're using plain `@beartype` and your tests pass, the frame-based detection is working. Most applications never need `@shapix.check`.
705713

706714
### Manual checks with `check_context`
@@ -733,61 +741,58 @@ Shapix uses three key mechanisms:
733741

734742
3. **Thread-local storage** — Each thread gets its own memo stack via `threading.local()`, ensuring thread safety.
735743

736-
## Static type checkers (pyright / Pylance)
744+
## Static type checkers (pyright, mypy, ty)
737745

738-
Shapix's pre-defined dimension symbols (`N`, `C`, `H`, `W`, ...) work out of the box with pyright and Pylance — under `TYPE_CHECKING` they resolve to `int` type aliases, so annotations like `F32[N, C]` are fully valid type expressions.
746+
Shapix supports **pyright**, **mypy**, and **ty**. Under `TYPE_CHECKING`, pre-defined dimension symbols (`N`, `C`, `H`, `W`, …) resolve to `TypeVar` and array types resolve to `TypeAliasType`, so core annotations like `F32[N, C]` are valid type expressions across all three checkers.
739747

740-
However, some patterns produce type checker errors because they place **runtime values** where pyright expects **types**:
748+
However, some patterns are fundamentally runtime-only and produce type checker errors regardless of the checker:
741749

742-
| Pattern | Example | Pyright rule |
743-
|---------|---------|--------------|
744-
| Integer literals | `F32[N, 3, H, W]` | `reportGeneralTypeIssues` |
745-
| Unary operators | `F32[~B, C]`, `F32[+N, C]` | `reportInvalidTypeForm` |
746-
| Arithmetic | `F32[N + 2]` | `reportInvalidTypeForm` |
747-
| Custom dimensions | `F32[Vocab, Embed]` | `reportInvalidTypeForm` (or use `TYPE_CHECKING` pattern) |
750+
| Pattern | Example | Workaround |
751+
|---------|---------|------------|
752+
| Integer literals | `F32[N, 3, H, W]` | Wrap in `Dimension("3")` |
753+
| Unary operators | `F32[~B, C]`, `F32[+N, C]` | `# type: ignore` |
754+
| Arithmetic | `F32[N + 2]` | `# type: ignore` |
755+
| Custom dimensions | `F32[Vocab, Embed]` | `# type: ignore` or `TYPE_CHECKING` pattern |
756+
| `Value(...)` | `F32[Value("size")]` | `# type: ignore` |
748757

749-
### Option A — Suppress `reportInvalidTypeForm` and wrap integers (recommended)
750-
751-
Add one line to your pyright config to silence operators, custom dimensions, and arithmetic globally. Then wrap bare integer literals in `Dimension()` to shift them to the same suppressed rule:
752-
753-
```jsonc
754-
// pyrightconfig.json (recommended — Pylance always reads this)
755-
{ "reportInvalidTypeForm": false }
756-
```
758+
### Recommended pyright config
757759

758-
Or equivalently in `pyproject.toml`:
760+
Add to your `pyproject.toml` or `pyrightconfig.json` to suppress the most common shapix-related diagnostics:
759761

760762
```toml
761763
[tool.pyright]
762764
reportInvalidTypeForm = false
763765
```
764766

765-
```python
766-
from shapix import N, H, W, Dimension
767+
### Inline `# type: ignore`
768+
769+
For patterns that all three checkers reject (arithmetic dims, Value(), custom dims), use blanket `# type: ignore`:
767770

768-
# Integer literals — wrap in Dimension() to avoid reportGeneralTypeIssues
771+
```python
769772
@beartype
770-
def rgb_to_gray(x: F32[N, Dimension("3"), H, W]) -> F32[N, Dimension("1"), H, W]:
773+
def pad(x: F32[N]) -> F32[N + 2]: # type: ignore
771774
...
772775

773-
# Operators, arithmetic, custom dimensions — all covered by the suppression
774776
@beartype
775-
def pad(x: F32[N]) -> F32[N + 2]:
777+
def f(x: F32[~B, C]) -> F32[~B, C]: # type: ignore
776778
...
777779
```
778780

779-
### Option B — Inline `# type: ignore`
781+
### Custom dimensions under TYPE_CHECKING
780782

781-
If you prefer not to change your pyright config, silence individual lines:
783+
Custom dimensions created with `Dimension()` are runtime objects. To make them work with all type checkers, use the `TYPE_CHECKING` pattern:
782784

783785
```python
784-
@beartype
785-
def rgb_to_gray(x: F32[N, 3, H, W]) -> F32[N, 1, H, W]: # type: ignore[reportGeneralTypeIssues]
786-
...
786+
import typing as tp
787+
from shapix import Dimension
787788

788-
@beartype
789-
def f(x: F32[~B, C]) -> F32[~B, C]: # type: ignore[reportInvalidTypeForm]
790-
...
789+
if tp.TYPE_CHECKING:
790+
import typing as _tp
791+
Vocab = _tp.TypeVar("Vocab")
792+
Embed = _tp.TypeVar("Embed")
793+
else:
794+
Vocab = Dimension("Vocab")
795+
Embed = Dimension("Embed")
791796
```
792797

793798
## Compared to jaxtyping
@@ -835,17 +840,22 @@ def f(x: F32[~B, C]) -> F32[~B, C]: # type: ignore[reportInvalidTypeForm]
835840

836841
Most NumPy array types, plus `BF16` and `BF16Like`. NumPy-only extended-precision array aliases such as `F128` / `C256` stay in `shapix.numpy`. Both export `Like` types, `ScalarLike` types (re-exported from numpy), and `make_scalar_like_type`. JAX also exports `Tree`.
837842

838-
### Factories (`shapix`)
843+
### Factories
844+
845+
From `shapix` (root):
839846

840847
`make_array_type(array_class, dtype_spec)` — custom array type
841848
`make_array_like_type(dtype_spec, *, casting="same_kind", name="ArrayLike")` — custom Like type
842-
`make_scalar_like_type(target_dtype, *, casting="same_kind", name="ScalarLike")` — custom ScalarLike type
843849
`DtypeSpec(name, allowed)` — custom dtype specification
844850
`DtypeSpec.structured(fields)` — structured dtype specification
845851

852+
From `shapix.numpy` (requires NumPy):
853+
854+
`make_scalar_like_type(target_dtype, *, casting="same_kind", name="ScalarLike")` — custom ScalarLike type
855+
846856
### Decorators & context managers (`shapix`)
847857

848-
`@shapix.check` — explicit memo management
858+
`@shapix.check` — explicit memo management (supports both sync and async functions)
849859
`@shapix.check(conf=BeartypeConf())` — combined memo + beartype
850860
`shapix.check_context()` — context manager for manual checks
851861

‎src/shapix/__init__.py‎

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -47,16 +47,20 @@ def conv(x: F32[N, C, H, W]) -> F32[N, C, H, W]: ...
4747
:func:`make_array_like_type` — create subscriptable array-like type
4848
factories with configurable dtype casting.
4949
:func:`check` — optional decorator for explicit memo management.
50+
Supports both sync and async functions.
5051
Also supports combined mode: ``@check(conf=BeartypeConf())``.
5152
5253
Context managers
5354
:class:`check_context` — shared dimension memo for manual
5455
``is_bearable()`` checks.
5556
"""
5657

57-
from importlib.metadata import version
58+
from importlib.metadata import PackageNotFoundError, version
5859

59-
__version__ = version("shapix")
60+
try:
61+
__version__ = version("shapix")
62+
except PackageNotFoundError:
63+
__version__ = "0+unknown"
6064

6165
__all__ = [
6266
# Dimension symbols

‎src/shapix/_array_types.py‎

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -29,7 +29,7 @@
2929

3030
from beartype.vale import Is
3131

32-
from ._dimensions import Dimension, _ValueExpr
32+
from ._dimensions import Dimension, Scalar, _ValueExpr
3333
from ._dtypes import DtypeSpec
3434
from ._dtypes import extract_dtype_str as extract_dtype_str
3535
from ._memo import ShapeMemo as ShapeMemo
@@ -406,6 +406,11 @@ def make_array_like_type(
406406

407407
def _to_shape_spec(dims: tuple[object, ...]) -> tuple[DimSpec, ...]:
408408
"""Convert a tuple of user-facing dim objects to internal DimSpec."""
409+
if any(d is Scalar for d in dims) and len(dims) > 1:
410+
msg = (
411+
"Scalar must be the only shape token; mixed use like F32[N, Scalar] is invalid"
412+
)
413+
raise TypeError(msg)
409414
specs: list[DimSpec] = []
410415
for d in dims:
411416
if d is Ellipsis:

‎src/shapix/_decorator.py‎

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -54,6 +54,20 @@ def decorator(fn: Callable[P, R]) -> Callable[P, R]:
5454

5555
inner = beartype(fn, conf=conf) # type: ignore[arg-type]
5656

57+
if inspect.iscoroutinefunction(fn):
58+
59+
@functools.wraps(fn)
60+
async def async_wrapper(*args: P.args, **kwargs: P.kwargs) -> R:
61+
bound = signature.bind_partial(*args, **kwargs)
62+
bound.apply_defaults()
63+
push_memo(scope=dict(bound.arguments))
64+
try:
65+
return await inner(*args, **kwargs) # type: ignore[misc]
66+
finally:
67+
pop_memo()
68+
69+
return async_wrapper # type: ignore[return-value]
70+
5771
@functools.wraps(fn)
5872
def wrapper(*args: P.args, **kwargs: P.kwargs) -> R:
5973
bound = signature.bind_partial(*args, **kwargs)

‎src/shapix/_dimensions.py‎

Lines changed: 11 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -311,17 +311,17 @@ def __new__(cls, expr: str) -> Dimension: ...
311311

312312
def __pos__(self) -> Dimension: ...
313313

314-
Scalar: Dimension
315-
B: Dimension
316-
N: Dimension
317-
P: Dimension
318-
L: Dimension
319-
C: Dimension
320-
D: Dimension
321-
K: Dimension
322-
H: Dimension
323-
W: Dimension
324-
__: Dimension
314+
Scalar = tp.TypeVar("Scalar")
315+
B = tp.TypeVar("B")
316+
N = tp.TypeVar("N")
317+
P = tp.TypeVar("P")
318+
L = tp.TypeVar("L")
319+
C = tp.TypeVar("C")
320+
D = tp.TypeVar("D")
321+
K = tp.TypeVar("K")
322+
H = tp.TypeVar("H")
323+
W = tp.TypeVar("W")
324+
__ = tp.TypeVar("__")
325325
else:
326326

327327
class Value(_ValueExpr):

‎src/shapix/_dtypes.py‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -265,6 +265,10 @@ def extract_dtype_str(obj: object) -> str:
265265
return "bytes"
266266
if dtype_name.startswith("str"):
267267
return "str"
268+
if dtype_name.startswith("datetime64"):
269+
return "datetime64"
270+
if dtype_name.startswith("timedelta64"):
271+
return "timedelta64"
268272
return dtype_name
269273

270274
# NumPy / JAX: dtype.type.__name__ (e.g. "float32")

0 commit comments

Comments
 (0)