Coverage for core / src / sensorkit / common / keyword.py: 83%
138 statements
« prev ^ index » next coverage.py v7.13.5, created at 2026-09-02 00:03 +0000
« prev ^ index » next coverage.py v7.13.5, created at 2026-09-02 00:03 +0000
1# SPDX-License-Identifier: Apache-2.0
2"""Keyword registration, lookup, and serialization for the sensorkit data model."""
4import functools
5from collections.abc import Iterable, Mapping
6from functools import partial
7from typing import (
8 Annotated,
9 Any,
10 ClassVar,
11 Literal,
12 NamedTuple,
13 Protocol,
14 Self,
15 Unpack,
16 overload,
17 override,
18 runtime_checkable,
19)
21from pydantic import BaseModel, GetCoreSchemaHandler, TypeAdapter, ValidationError
22from pydantic_core import core_schema
24from sensorkit.common.model import ModelRegistry
26_keyword_unknown_key = "__unknown__"
27_keyword_registry = ModelRegistry(default_tag=_keyword_unknown_key)
29Keyword = Annotated[object, _keyword_registry.discriminator()]
30"""Pydantic type annotation matching any dynamically registered keyword type."""
32type KeywordKind = Literal["stream", "state", "config"]
35class KeywordInfo(NamedTuple):
36 """Registration metadata for a declared keyword type."""
38 key: str
39 namespace: str | None
40 kind: KeywordKind | None
43_keyword_adapter = TypeAdapter(Keyword)
44_keyword_index: dict[type, KeywordInfo] = {}
47def declare_keyword[M](
48 cls: type[M] | None = None,
49 *,
50 key: str | None = None,
51 ns: str | None = None,
52 kind: KeywordKind = "stream",
53) -> partial[type[M]] | type[M]:
54 """Decorator to register a type as a Keyword.
56 Args:
57 cls: The type to register.
58 key: The keyword key. Defaults to the class name.
59 ns: Optional namespace, forwarded to the model registry.
60 kind: The keyword variant; can be `config`, `state`, or `stream` (the default).
61 """
62 if cls is None:
63 return functools.partial(declare_keyword, key=key, ns=ns, kind=kind)
65 if cls in _keyword_index:
66 raise KeywordError(f"type '{cls.__name__}' already declared as '{_keyword_index[cls].key}'")
68 key = key or cls.__name__
69 _keyword_registry.add(cls, tag=key, namespace=ns)
70 _keyword_index[cls] = KeywordInfo(key=key, namespace=ns, kind=kind)
71 return cls
74def get_keyword_info(obj: type | object) -> KeywordInfo | None:
75 """Return the `KeywordInfo` for a registered keyword type or instance, or `None` if not registered."""
76 if not isinstance(obj, type):
77 obj = type(obj)
79 return _keyword_index.get(obj)
82def is_keyword(key: str) -> bool:
83 """Return whether `key` is declared as a keyword in any namespace."""
84 return bool(_keyword_registry.get_namespaces(key))
87def dump_keyword_json(obj: object):
88 """Serialize a keyword object to JSON bytes using the shared keyword type adapter."""
89 return _keyword_adapter.dump_json(obj)
92def validate_keyword(key: str, data: Any):
93 """Validate and deserialize `data` as the keyword type identified by `key`."""
94 return _keyword_adapter.validate_python(data, context={ModelRegistry.DISCRIMINATOR_CONTEXT: key})
97def validate_keyword_json(key: str, json: bytes):
98 """Validate and deserialize a JSON byte string as the keyword type identified by `key`."""
99 return _keyword_adapter.validate_json(json, context={ModelRegistry.DISCRIMINATOR_CONTEXT: key})
102def validated_items(dct: dict[str, object]) -> Iterable[tuple[str, object]]:
103 """Yield `(key, value)` pairs from `dct`, deserializing dict values as keywords where possible."""
104 for k, v in dct.items():
105 if isinstance(v, dict):
106 try:
107 yield k, validate_keyword(k, v)
108 continue
109 except ValidationError:
110 pass
112 yield k, v
115@declare_keyword(key=_keyword_unknown_key)
116class UnknownKeyword(BaseModel, extra="allow"):
117 """Fallback keyword type that accepts any extra fields for unrecognized keyword keys."""
120class KeywordError(Exception):
121 """Keyword error."""
124@runtime_checkable
125class CompositeKeyword(Protocol):
126 """Protocol for a keyword that exports other keywords.
128 Composite keywords may only compose a set of keywords that are a pure projection over
129 its own fields. Expansion runs wherever a composite is read into a `Context`, including
130 on the consumer side after deserialization, so anything not carried in the keyword's own
131 serialized fields would simply be absent there.
132 """
134 def composed_keywords(self) -> Iterable[object]:
135 """Yield the child keywords this keyword carries."""
138class KeywordDict(Mapping[str, Any]):
139 """A mapping of keywords, keyed by their registered keyword key.
141 Values are inserted through `set` (keyword objects, keyed by their type) or `set_value` (a
142 named non-keyword value). Some, but not all, of the `MutableMapping` interface is supported.
143 """
145 NO_DEFAULT: ClassVar[object] = object()
147 def __init__(
148 self,
149 arg: Any = None,
150 *objs: Unpack[tuple[object, ...]],
151 ):
152 self._data: dict[str, Any] = {}
154 match arg:
155 case None:
156 pass
157 case _ if type(arg) in _keyword_index:
158 self.set(arg)
159 case KeywordDict():
160 self.update(arg)
161 case Mapping():
162 for key, value in arg.items():
163 self.set_value(key, value)
164 case Iterable():
165 for key, value in arg:
166 self.set_value(key, value)
167 case _:
168 raise RuntimeError(f"KeywordDict arg has invalid type {type(arg)}")
170 self.set(*objs)
172 def set(self, *objs: Unpack[tuple[object, ...]]):
173 """Insert keywords."""
174 for obj in objs:
175 self.set_keyword(obj)
177 def set_keyword(self, obj: object):
178 """Insert a keyword."""
179 self._data[_keyword_index[type(obj)].key] = obj
181 def set_value(self, key: str, value: Any):
182 """Insert a named non-keyword value, reachable by that name.
184 Transitional accommodation for values not (yet) modeled as keywords.
185 """
186 self._data[key] = value
188 def update(self, other: Mapping[str, Any]):
189 """Copy every entry from another mapping into this one, overwriting on key collision."""
190 self._data.update(other._data if isinstance(other, KeywordDict) else other)
192 @classmethod
193 def _validate(cls, obj):
194 if obj is None:
195 return cls()
197 if isinstance(obj, dict):
198 return cls(validated_items(obj))
200 return obj
202 @classmethod
203 def __get_pydantic_core_schema__(cls, _source_type: Any, handler: GetCoreSchemaHandler):
204 return core_schema.json_or_python_schema(
205 json_schema=core_schema.chain_schema(
206 [
207 core_schema.nullable_schema(core_schema.dict_schema()),
208 core_schema.no_info_plain_validator_function(
209 function=cls._validate,
210 ),
211 ]
212 ),
213 python_schema=core_schema.no_info_plain_validator_function(
214 function=cls._validate,
215 ),
216 serialization=core_schema.plain_serializer_function_ser_schema(
217 lambda v: dict(v) if v is not None else None,
218 return_schema=core_schema.nullable_schema(core_schema.dict_schema()),
219 ),
220 )
222 def copy(self) -> Self:
223 return self.__class__(self)
225 @override
226 def __iter__(self):
227 return iter(self._data)
229 @override
230 def __len__(self):
231 return len(self._data)
233 def __eq__(self, other: object) -> bool:
234 match other:
235 case KeywordDict():
236 return self._data == other._data
237 case Mapping():
238 return self._data == dict(other)
239 case _:
240 return NotImplemented
242 __hash__ = None
244 def __repr__(self):
245 return f"{type(self).__name__}({self._data!r})"
247 @overload
248 def get[M](self, cls: type[M], default: M | None = None) -> M | None: ...
250 @overload
251 def get(self, key: str, default: Any = None) -> Any: ...
253 @override
254 def get(self, key, default=None):
255 if isinstance(key, type):
256 entry = _keyword_index.get(key)
257 return default if entry is None else self._data.get(entry.key, default)
259 return self._data.get(key, default)
261 @overload
262 def pop[M](self, cls: type[M], default: Any = ...) -> M | None: ...
264 @overload
265 def pop(self, key: str, default: Any = ...) -> Any: ...
267 def pop(self, key, default=NO_DEFAULT):
268 if isinstance(key, type):
269 entry = _keyword_index.get(key)
271 if entry is None:
272 if default is self.NO_DEFAULT:
273 raise KeyError(key.__name__)
274 return default
276 key = entry.key
278 if default is self.NO_DEFAULT:
279 return self._data.pop(key)
281 return self._data.pop(key, default)
283 @overload
284 def __getitem__[M](self, cls: type[M]) -> M: ...
286 @overload
287 def __getitem__(self, key: str) -> Any: ...
289 @override
290 def __getitem__(self, key, /):
291 if isinstance(key, type):
292 return self._data[_keyword_index[key].key]
293 return self._data[key]
295 @overload
296 def __delitem__(self, cls: type) -> None: ...
298 @overload
299 def __delitem__(self, key: str) -> None: ...
301 def __delitem__(self, key, /):
302 if isinstance(key, type):
303 del self._data[_keyword_index[key].key]
304 else:
305 del self._data[key]