diff --git a/src/decman/core/store.py b/src/decman/core/store.py index 0414a25..4ab0b77 100644 --- a/src/decman/core/store.py +++ b/src/decman/core/store.py @@ -17,7 +17,7 @@ class Store: if self._path.exists(): with self._path.open("rt", encoding="utf-8") as file: - self._store = json.load(file) + self._store = json.load(file, object_hook=_decode_sets) def __getitem__(self, key: str) -> typing.Any: return self._store[key] @@ -30,7 +30,7 @@ class Store: def ensure(self, key: str, default: typing.Any = None): if key not in self._store: - self._store = default + self._store[key] = default def __enter__(self) -> "Store": return self @@ -54,7 +54,7 @@ class Store: dir=self._path.parent, delete=False, ) as tmp: - json.dump(self._store, tmp, indent=2) + json.dump(self._store, tmp, cls=_SetJSONEncoder, indent=2) tmp.flush() os.fsync(tmp.fileno()) @@ -62,3 +62,17 @@ class Store: def __repr__(self) -> str: return repr(self._store) + + +class _SetJSONEncoder(json.JSONEncoder): + def default(self, obj: typing.Any) -> typing.Any: + if isinstance(obj, set): + # generic, works for any set value + return {"__type__": "set", "items": list(obj)} + return super().default(obj) + + +def _decode_sets(obj: typing.Any) -> typing.Any: + if isinstance(obj, dict) and obj.get("__type__") == "set" and "items" in obj: + return set(obj["items"]) + return obj diff --git a/tests/test_decman_core_store.py b/tests/test_decman_core_store.py index e895913..64679e5 100644 --- a/tests/test_decman_core_store.py +++ b/tests/test_decman_core_store.py @@ -101,3 +101,29 @@ def test_repr_matches_underlying_dict(tmp_path: Path) -> None: expected = repr({"foo": "bar", "number": 123}) assert repr(store) == expected + + +def test_store_persists_sets(tmp_path: Path) -> None: + path = tmp_path / "store.json" + + # initial write with sets + store = Store(path) + store["units"] = {"a.service", "b.service"} + store["user_units"] = {"alice": {"u1.service", "u2.service"}} + store.save() + + # raw JSON should be set-encoded, not fail json.dump + raw = json.loads(path.read_text(encoding="utf-8")) + assert raw["units"]["__type__"] == "set" + assert set(raw["units"]["items"]) == {"a.service", "b.service"} + + assert raw["user_units"]["alice"]["__type__"] == "set" + assert set(raw["user_units"]["alice"]["items"]) == {"u1.service", "u2.service"} + + # reloading via Store must restore actual set objects + reloaded = Store(path) + assert reloaded["units"] == {"a.service", "b.service"} + assert isinstance(reloaded["units"], set) + + assert reloaded["user_units"]["alice"] == {"u1.service", "u2.service"} + assert isinstance(reloaded["user_units"]["alice"], set)