Skip to content

Commit 4a95dce

Browse files
committed
fix: generic type handling
1 parent 38f3a4b commit 4a95dce

3 files changed

Lines changed: 26 additions & 53 deletions

File tree

src/reaktiv/model.py

Lines changed: 19 additions & 53 deletions
Original file line numberDiff line numberDiff line change
@@ -178,8 +178,24 @@ def destroy(self) -> None: ...
178178
class SignalField(Generic[T]):
179179
"""Descriptor that creates one writable Signal per model instance."""
180180

181-
def __init__(self, initial_value: Callable[[], T]) -> None:
182-
self._initial_value = initial_value
181+
@overload
182+
def __init__(self, default: T, /) -> None: ...
183+
184+
@overload
185+
def __init__(self, *, factory: Callable[[], T]) -> None: ...
186+
187+
def __init__(
188+
self,
189+
*defaults: T,
190+
factory: Optional[Callable[[], T]] = None,
191+
) -> None:
192+
if len(defaults) == 1 and factory is None:
193+
default = defaults[0]
194+
self._initial_value = lambda: default
195+
elif not defaults and factory is not None:
196+
self._initial_value = factory
197+
else:
198+
raise TypeError("field() requires exactly one default value or a factory")
183199
self._cache: WeakKeyDictionary[object, Signal[T]] = WeakKeyDictionary()
184200
self._storage_name: Optional[str] = None
185201

@@ -257,54 +273,4 @@ def _set_on_instance_dict(self, instance: object, value: Signal[T]) -> None:
257273
vars(instance)[self._storage_name] = value
258274

259275

260-
class _TypedFieldFactory(Generic[T]):
261-
"""Typed field factory used by `field[T]` syntax."""
262-
263-
@overload
264-
def __call__(self, default: T, /) -> SignalField[T]: ...
265-
266-
@overload
267-
def __call__(self, *, factory: Callable[[], T]) -> SignalField[T]: ...
268-
269-
def __call__(
270-
self,
271-
*defaults: T,
272-
factory: Optional[Callable[[], T]] = None,
273-
) -> SignalField[T]:
274-
return _create_field(defaults, factory)
275-
276-
277-
class _FieldFactory:
278-
"""Declare a per-instance Signal field on a ReactiveModel."""
279-
280-
@overload
281-
def __call__(self, default: T, /) -> SignalField[T]: ...
282-
283-
@overload
284-
def __call__(self, *, factory: Callable[[], T]) -> SignalField[T]: ...
285-
286-
def __call__(
287-
self,
288-
*defaults: T,
289-
factory: Optional[Callable[[], T]] = None,
290-
) -> SignalField[T]:
291-
return _create_field(defaults, factory)
292-
293-
def __getitem__(self, value_type: type[T]) -> _TypedFieldFactory[T]:
294-
del value_type
295-
return _TypedFieldFactory()
296-
297-
298-
field = _FieldFactory()
299-
300-
301-
def _create_field(
302-
defaults: tuple[T, ...],
303-
factory: Optional[Callable[[], T]],
304-
) -> SignalField[T]:
305-
if len(defaults) == 1 and factory is None:
306-
default = defaults[0]
307-
return SignalField(lambda: default)
308-
if not defaults and factory is not None:
309-
return SignalField(factory)
310-
raise TypeError("field() requires exactly one default value or a factory")
276+
field = SignalField

tests/test_reactive_model.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
import asyncio
22
import gc
33
import weakref
4+
from typing import Optional
45

56
import pytest
67

@@ -58,6 +59,7 @@ class Profile(ReactiveModel):
5859
name = field[str]("")
5960
age = field[int](0)
6061
tags = field[list[str]](factory=list)
62+
nickname = field[Optional[str]](None)
6163

6264
def __init__(self, name: str, age: int = 0) -> None:
6365
self.observed: list[tuple[str, int]] = []
@@ -72,6 +74,7 @@ def observe(self) -> None:
7274
assert profile.name() == "Ada"
7375
assert profile.age() == 37
7476
assert profile.tags() == []
77+
assert profile.nickname() is None
7578
assert profile.observed == [("Ada", 37)]
7679

7780

tests/typing/public_api_inference.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,8 @@
66

77
from __future__ import annotations
88

9+
from typing import Optional
10+
911
from reaktiv import (
1012
Computed,
1113
ComputeSignal,
@@ -66,6 +68,7 @@ class CounterModel(ReactiveModel):
6668
count = field(1)
6769
name = field[str]("")
6870
labels = field[list[str]](factory=list)
71+
optional_name = field[Optional[str]](None)
6972

7073
@computed
7174
def doubled(self) -> int:
@@ -132,6 +135,7 @@ async def load_cached_user(
132135
model_count: Signal[int] = model.count
133136
model_name_field: Signal[str] = model.name
134137
model_labels: Signal[list[str]] = model.labels
138+
model_optional_name: Signal[Optional[str]] = model.optional_name
135139
model_doubled: ComputeSignal[int] = model.doubled
136140
model_normalized_name: ComputeSignal[str] = model.normalized_name
137141
model_normalized_name_value: str = model.normalized_name()

0 commit comments

Comments
 (0)