Coverage for src/metaclass_registry/cache.py: 84%

178 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-10-02 00:58 +0000

1""" 

2Generic caching system for plugin registries. 

3 

4Provides unified caching for both function registries (Pattern B) and 

5metaclass registries (Pattern A), eliminating code duplication and 

6ensuring consistent cache behavior across the codebase. 

7 

8Architecture: 

9- RegistryCacheManager: Generic cache manager for any registry type 

10- Supports version validation, age-based invalidation, mtime checking 

11- JSON-based serialization with custom serializers/deserializers 

12- XDG-compliant cache locations 

13 

14Usage: 

15 # For function registries 

16 cache_mgr = RegistryCacheManager( 

17 cache_name="scikit_image_functions", 

18 version_getter=lambda: skimage.__version__, 

19 serializer=serialize_function_metadata, 

20 deserializer=deserialize_function_metadata 

21 ) 

22 

23 # For metaclass registries 

24 cache_mgr = RegistryCacheManager( 

25 cache_name="microscope_handlers", 

26 version_getter=lambda: openhcs.__version__, 

27 serializer=serialize_plugin_class, 

28 deserializer=deserialize_plugin_class 

29 ) 

30""" 

31 

32import json 

33import logging 

34import time 

35from collections.abc import Callable, Hashable 

36from dataclasses import dataclass 

37from enum import Enum, EnumMeta 

38from pathlib import Path 

39from typing import Any, Generic, TypeVar 

40 

41logger = logging.getLogger(__name__) 

42 

43 

44def get_cache_file_path(cache_name: str, *, create: bool = True) -> Path: 

45 """ 

46 Get XDG-compliant cache file path. 

47 

48 Args: 

49 cache_name: Name of the cache file 

50 

51 Returns: 

52 Path to cache file in XDG cache directory 

53 """ 

54 # Use XDG_CACHE_HOME if set, otherwise default to ~/.cache 

55 import os 

56 

57 from . import _home 

58 

59 cache_home_value = os.environ.get("XDG_CACHE_HOME") 

60 cache_home = ( 

61 Path(cache_home_value) if cache_home_value else Path(_home.get_home_dir()) / ".cache" 

62 ) 

63 

64 # Create metaclass-registry subdirectory 

65 cache_dir = cache_home / "metaclass-registry" 

66 if create: 

67 cache_dir.mkdir(parents=True, exist_ok=True) 

68 

69 return cache_dir / cache_name 

70 

71 

72T = TypeVar("T") # Generic type for cached items 

73 

74 

75class CacheKeyCodec: 

76 """JSON-safe, reversible representation for registry keys.""" 

77 

78 PRIMITIVE_KIND = "primitive" 

79 ENUM_KIND = "enum" 

80 TUPLE_KIND = "tuple" 

81 

82 @classmethod 

83 def encode(cls, key: Any) -> dict[str, Any]: 

84 if isinstance(key, (str, int, float, bool)) or key is None: 

85 return {"kind": cls.PRIMITIVE_KIND, "value": key} 

86 if isinstance(key, Enum): 

87 return { 

88 "kind": cls.ENUM_KIND, 

89 "module": key.__class__.__module__, 

90 "class_name": key.__class__.__name__, 

91 "member_name": key.name, 

92 } 

93 if isinstance(key, tuple): 

94 return { 

95 "kind": cls.TUPLE_KIND, 

96 "items": [cls.encode(item) for item in key], 

97 } 

98 raise TypeError(f"Unsupported registry cache key type: {type(key)!r}") 

99 

100 @classmethod 

101 def decode(cls, payload: dict[str, Any]) -> Hashable: 

102 kind = payload["kind"] 

103 try: 

104 decoder = CACHE_KEY_DECODERS[kind] 

105 except KeyError as exc: 

106 raise ValueError(f"Unsupported registry cache key kind: {kind!r}") from exc 

107 return decoder(payload) 

108 

109 

110def _decode_primitive_cache_key(payload: dict[str, Any]) -> Hashable: 

111 value = payload["value"] 

112 if not isinstance(value, Hashable): 

113 raise TypeError(f"Cached primitive key {value!r} is not hashable") 

114 return value 

115 

116 

117def _decode_enum_cache_key(payload: dict[str, Any]) -> Enum: 

118 import importlib 

119 

120 enum_module = importlib.import_module(payload["module"]) 

121 enum_type = vars(enum_module)[payload["class_name"]] 

122 if not isinstance(enum_type, EnumMeta): 

123 raise TypeError(f"Cached enum owner {enum_type!r} is not an Enum type") 

124 return enum_type[payload["member_name"]] 

125 

126 

127def _decode_tuple_cache_key(payload: dict[str, Any]) -> tuple[Any, ...]: 

128 return tuple(CacheKeyCodec.decode(item) for item in payload["items"]) 

129 

130 

131CACHE_KEY_DECODERS: dict[str, Callable[[dict[str, Any]], Hashable]] = { 

132 CacheKeyCodec.PRIMITIVE_KIND: _decode_primitive_cache_key, 

133 CacheKeyCodec.ENUM_KIND: _decode_enum_cache_key, 

134 CacheKeyCodec.TUPLE_KIND: _decode_tuple_cache_key, 

135} 

136 

137 

138@dataclass 

139class CacheConfig: 

140 """Configuration for registry caching behavior.""" 

141 

142 max_age_days: int = 7 # Maximum cache age before invalidation 

143 check_mtimes: bool = False # Check file modification times 

144 cache_version: str = "1.0" # Cache format version 

145 

146 

147class RegistryCacheManager(Generic[T]): 

148 """ 

149 Generic cache manager for plugin registries. 

150 

151 Handles caching, validation, and reconstruction of registry data 

152 with support for version checking, age-based invalidation, and 

153 custom serialization. 

154 

155 Type Parameters: 

156 T: Type of items being cached (e.g., FunctionMetadata, Type[Plugin]) 

157 """ 

158 

159 def __init__( 

160 self, 

161 cache_name: str, 

162 version_getter: Callable[[], str], 

163 serializer: Callable[[T], dict[str, Any]], 

164 deserializer: Callable[[dict[str, Any]], T], 

165 config: CacheConfig | None = None, 

166 file_mtimes_getter: Callable[[], dict[str, float]] | None = None, 

167 ): 

168 """ 

169 Initialize cache manager. 

170 

171 Args: 

172 cache_name: Name for the cache file (e.g., "microscope_handlers") 

173 version_getter: Function that returns current version string 

174 serializer: Function to serialize item to JSON-compatible dict 

175 deserializer: Function to deserialize dict back to item 

176 config: Optional cache configuration 

177 file_mtimes_getter: Optional authoritative inventory of source files used by 

178 discovery. When provided, added files invalidate the cache as well as 

179 modified and removed files. 

180 """ 

181 self.cache_name = cache_name 

182 self.version_getter = version_getter 

183 self.serializer = serializer 

184 self.deserializer = deserializer 

185 self.config = config or CacheConfig() 

186 self.file_mtimes_getter = file_mtimes_getter 

187 self._cache_path = get_cache_file_path(f"{cache_name}.json") 

188 

189 def load_cache(self) -> dict[Hashable, T] | None: 

190 """ 

191 Load cached items with validation. 

192 

193 Returns: 

194 Dictionary of cached items, or None if cache is invalid 

195 """ 

196 if not self._cache_path.exists(): 

197 logger.debug(f"No cache found for {self.cache_name}") 

198 return None 

199 

200 try: 

201 with open(self._cache_path) as f: 

202 cache_data = json.load(f) 

203 except json.JSONDecodeError: 

204 logger.warning(f"Corrupt cache file {self._cache_path}, rebuilding") 

205 self._cache_path.unlink(missing_ok=True) 

206 return None 

207 

208 # Validate cache version 

209 if cache_data.get("cache_version") != self.config.cache_version: 

210 logger.debug(f"Cache version mismatch for {self.cache_name}") 

211 return None 

212 

213 # Validate library/package version 

214 cached_version = cache_data.get("version", "unknown") 

215 current_version = self.version_getter() 

216 if cached_version != current_version: 

217 logger.info( 

218 f"{self.cache_name} version changed " 

219 f"({cached_version} → {current_version}) - cache invalid" 

220 ) 

221 return None 

222 

223 # Validate cache age 

224 cache_timestamp = cache_data.get("timestamp", 0) 

225 cache_age_days = (time.time() - cache_timestamp) / (24 * 3600) 

226 if cache_age_days > self.config.max_age_days: 

227 logger.debug( 

228 f"Cache for {self.cache_name} is {cache_age_days:.1f} days old - rebuilding" 

229 ) 

230 return None 

231 

232 # Validate file mtimes if configured 

233 if self.config.check_mtimes and "file_mtimes" in cache_data: 

234 current_mtimes = ( 

235 self.file_mtimes_getter() if self.file_mtimes_getter is not None else None 

236 ) 

237 if not self._validate_mtimes( 

238 cache_data["file_mtimes"], 

239 current_mtimes=current_mtimes, 

240 ): 

241 logger.debug(f"File modifications detected for {self.cache_name}") 

242 return None 

243 

244 # Deserialize items 

245 items: dict[Hashable, T] = {} 

246 item_entries = cache_data.get("item_entries") 

247 if item_entries is None: 

248 item_entries = [ 

249 {"key": CacheKeyCodec.encode(key), "item": item_data} 

250 for key, item_data in cache_data.get("items", {}).items() 

251 ] 

252 for entry in item_entries: 

253 key = CacheKeyCodec.decode(entry["key"]) 

254 item_data = entry["item"] 

255 try: 

256 items[key] = self.deserializer(item_data) 

257 except Exception as e: 

258 logger.warning(f"Failed to deserialize {key} from cache: {e}") 

259 return None # Invalidate entire cache on any deserialization error 

260 

261 logger.debug(f"✅ Loaded {len(items)} items from {self.cache_name} cache") 

262 return items 

263 

264 def save_cache( 

265 self, 

266 items: dict[Hashable, T], 

267 file_mtimes: dict[str, float] | None = None, 

268 ) -> None: 

269 """ 

270 Save items to cache. 

271 

272 Args: 

273 items: Dictionary of items to cache 

274 file_mtimes: Optional dict of file paths to modification times 

275 """ 

276 cache_data: dict[str, Any] = { 

277 "cache_version": self.config.cache_version, 

278 "version": self.version_getter(), 

279 "timestamp": time.time(), 

280 "item_entries": [], 

281 } 

282 

283 # Add file mtimes if provided 

284 if file_mtimes: 

285 cache_data["file_mtimes"] = file_mtimes 

286 

287 # Serialize items 

288 for key, item in items.items(): 

289 try: 

290 cache_data["item_entries"].append( 

291 { 

292 "key": CacheKeyCodec.encode(key), 

293 "item": self.serializer(item), 

294 } 

295 ) 

296 except Exception as e: 

297 logger.warning(f"Failed to serialize {key} for cache: {e}") 

298 

299 # Save to disk 

300 try: 

301 self._cache_path.parent.mkdir(parents=True, exist_ok=True) 

302 with open(self._cache_path, "w") as f: 

303 json.dump(cache_data, f, indent=2) 

304 logger.debug(f"💾 Saved {len(items)} items to {self.cache_name} cache") 

305 except Exception as e: 

306 logger.warning(f"Failed to save {self.cache_name} cache: {e}") 

307 

308 def clear_cache(self) -> None: 

309 """Clear the cache file.""" 

310 if self._cache_path.exists(): 

311 self._cache_path.unlink() 

312 logger.info(f"🧹 Cleared {self.cache_name} cache") 

313 

314 def _validate_mtimes( 

315 self, 

316 cached_mtimes: dict[str, float], 

317 *, 

318 current_mtimes: dict[str, float] | None = None, 

319 ) -> bool: 

320 """ 

321 Validate that file modification times haven't changed. 

322 

323 Args: 

324 cached_mtimes: Dictionary of file paths to cached mtimes 

325 current_mtimes: Optional complete current inventory. When present, its paths 

326 must exactly match the cached discovery inventory. 

327 

328 Returns: 

329 True if all mtimes match, False if any file changed 

330 """ 

331 if current_mtimes is not None and current_mtimes.keys() != cached_mtimes.keys(): 

332 return False 

333 

334 for file_path, cached_mtime in cached_mtimes.items(): 

335 path = Path(file_path) 

336 if not path.exists(): 

337 return False # File was deleted 

338 

339 current_mtime = ( 

340 current_mtimes[file_path] if current_mtimes is not None else path.stat().st_mtime 

341 ) 

342 if abs(current_mtime - cached_mtime) > 1.0: # 1 second tolerance 

343 return False # File was modified 

344 

345 return True 

346 

347 

348# Serializers for metaclass registries (Pattern A) 

349 

350 

351def serialize_plugin_class(plugin_class: type) -> dict[str, Any]: 

352 """ 

353 Serialize a plugin class to JSON-compatible dict. 

354 

355 Args: 

356 plugin_class: Plugin class to serialize 

357 

358 Returns: 

359 Dictionary with module and class name 

360 """ 

361 return { 

362 "module": plugin_class.__module__, 

363 "class_name": plugin_class.__name__, 

364 "qualname": plugin_class.__qualname__, 

365 } 

366 

367 

368def deserialize_plugin_class(data: dict[str, Any]) -> type: 

369 """ 

370 Deserialize a plugin class from JSON-compatible dict. 

371 

372 Args: 

373 data: Dictionary with module and class name 

374 

375 Returns: 

376 Reconstructed plugin class 

377 

378 Raises: 

379 ImportError: If module cannot be imported 

380 AttributeError: If class not found in module 

381 """ 

382 import importlib 

383 

384 declaration: Any = importlib.import_module(data["module"]) 

385 for owner_name in data["qualname"].split("."): 

386 if owner_name == "<locals>": 

387 raise AttributeError("Function-local classes cannot be restored from a registry cache") 

388 declaration = getattr(declaration, owner_name) 

389 if not isinstance(declaration, type): 

390 raise TypeError(f"Cached plugin declaration {data['qualname']!r} is not a class") 

391 return declaration 

392 

393 

394def get_package_file_mtimes( 

395 package_path: str, 

396 *, 

397 recursive: bool = True, 

398) -> dict[str, float]: 

399 """ 

400 Get modification times for all Python files in a package. 

401 

402 Args: 

403 package_path: Package path (e.g., "openhcs.microscopes") 

404 recursive: Whether nested package files participate in discovery. 

405 

406 Returns: 

407 Dictionary mapping file paths to modification times 

408 """ 

409 import importlib 

410 import importlib.util 

411 from pathlib import Path 

412 

413 try: 

414 importlib.import_module(package_path) 

415 spec = importlib.util.find_spec(package_path) 

416 if spec is None: 

417 return {} 

418 if spec.submodule_search_locations: 

419 package_dirs = [Path(location) for location in spec.submodule_search_locations] 

420 elif spec.origin is not None: 

421 package_dirs = [Path(spec.origin).parent] 

422 else: 

423 return {} 

424 

425 mtimes = {} 

426 for package_dir in package_dirs: 

427 source_files = package_dir.rglob("*.py") if recursive else package_dir.glob("*.py") 

428 for py_file in source_files: 

429 mtimes[str(py_file)] = py_file.stat().st_mtime 

430 

431 return mtimes 

432 except Exception as e: 

433 logger.warning(f"Failed to get mtimes for {package_path}: {e}") 

434 return {}