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

1# SPDX-License-Identifier: Apache-2.0 

2"""Keyword registration, lookup, and serialization for the sensorkit data model.""" 

3 

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) 

20 

21from pydantic import BaseModel, GetCoreSchemaHandler, TypeAdapter, ValidationError 

22from pydantic_core import core_schema 

23 

24from sensorkit.common.model import ModelRegistry 

25 

26_keyword_unknown_key = "__unknown__" 

27_keyword_registry = ModelRegistry(default_tag=_keyword_unknown_key) 

28 

29Keyword = Annotated[object, _keyword_registry.discriminator()] 

30"""Pydantic type annotation matching any dynamically registered keyword type.""" 

31 

32type KeywordKind = Literal["stream", "state", "config"] 

33 

34 

35class KeywordInfo(NamedTuple): 

36 """Registration metadata for a declared keyword type.""" 

37 

38 key: str 

39 namespace: str | None 

40 kind: KeywordKind | None 

41 

42 

43_keyword_adapter = TypeAdapter(Keyword) 

44_keyword_index: dict[type, KeywordInfo] = {} 

45 

46 

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. 

55 

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) 

64 

65 if cls in _keyword_index: 

66 raise KeywordError(f"type '{cls.__name__}' already declared as '{_keyword_index[cls].key}'") 

67 

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 

72 

73 

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) 

78 

79 return _keyword_index.get(obj) 

80 

81 

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)) 

85 

86 

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) 

90 

91 

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}) 

95 

96 

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}) 

100 

101 

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 

111 

112 yield k, v 

113 

114 

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.""" 

118 

119 

120class KeywordError(Exception): 

121 """Keyword error.""" 

122 

123 

124@runtime_checkable 

125class CompositeKeyword(Protocol): 

126 """Protocol for a keyword that exports other keywords. 

127 

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 """ 

133 

134 def composed_keywords(self) -> Iterable[object]: 

135 """Yield the child keywords this keyword carries.""" 

136 

137 

138class KeywordDict(Mapping[str, Any]): 

139 """A mapping of keywords, keyed by their registered keyword key. 

140 

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 """ 

144 

145 NO_DEFAULT: ClassVar[object] = object() 

146 

147 def __init__( 

148 self, 

149 arg: Any = None, 

150 *objs: Unpack[tuple[object, ...]], 

151 ): 

152 self._data: dict[str, Any] = {} 

153 

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)}") 

169 

170 self.set(*objs) 

171 

172 def set(self, *objs: Unpack[tuple[object, ...]]): 

173 """Insert keywords.""" 

174 for obj in objs: 

175 self.set_keyword(obj) 

176 

177 def set_keyword(self, obj: object): 

178 """Insert a keyword.""" 

179 self._data[_keyword_index[type(obj)].key] = obj 

180 

181 def set_value(self, key: str, value: Any): 

182 """Insert a named non-keyword value, reachable by that name. 

183 

184 Transitional accommodation for values not (yet) modeled as keywords. 

185 """ 

186 self._data[key] = value 

187 

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) 

191 

192 @classmethod 

193 def _validate(cls, obj): 

194 if obj is None: 

195 return cls() 

196 

197 if isinstance(obj, dict): 

198 return cls(validated_items(obj)) 

199 

200 return obj 

201 

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 ) 

221 

222 def copy(self) -> Self: 

223 return self.__class__(self) 

224 

225 @override 

226 def __iter__(self): 

227 return iter(self._data) 

228 

229 @override 

230 def __len__(self): 

231 return len(self._data) 

232 

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 

241 

242 __hash__ = None 

243 

244 def __repr__(self): 

245 return f"{type(self).__name__}({self._data!r})" 

246 

247 @overload 

248 def get[M](self, cls: type[M], default: M | None = None) -> M | None: ... 

249 

250 @overload 

251 def get(self, key: str, default: Any = None) -> Any: ... 

252 

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) 

258 

259 return self._data.get(key, default) 

260 

261 @overload 

262 def pop[M](self, cls: type[M], default: Any = ...) -> M | None: ... 

263 

264 @overload 

265 def pop(self, key: str, default: Any = ...) -> Any: ... 

266 

267 def pop(self, key, default=NO_DEFAULT): 

268 if isinstance(key, type): 

269 entry = _keyword_index.get(key) 

270 

271 if entry is None: 

272 if default is self.NO_DEFAULT: 

273 raise KeyError(key.__name__) 

274 return default 

275 

276 key = entry.key 

277 

278 if default is self.NO_DEFAULT: 

279 return self._data.pop(key) 

280 

281 return self._data.pop(key, default) 

282 

283 @overload 

284 def __getitem__[M](self, cls: type[M]) -> M: ... 

285 

286 @overload 

287 def __getitem__(self, key: str) -> Any: ... 

288 

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] 

294 

295 @overload 

296 def __delitem__(self, cls: type) -> None: ... 

297 

298 @overload 

299 def __delitem__(self, key: str) -> None: ... 

300 

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]