diff options
| author | Yuval Adam <_@yuv.al> | 2024-05-09 12:34:10 +0200 |
|---|---|---|
| committer | Yuval Adam <_@yuv.al> | 2024-05-09 12:34:10 +0200 |
| commit | d089e3acb34cf3214c4c1a52b3142798de5f92f2 (patch) | |
| tree | 94dd8a58bfe4dff6c71323b43226eeeac654c44e | |
| parent | eaab19e310edd682af0231fe06f9a647bfcd5834 (diff) | |
Add support for filtering null values
| -rw-r--r-- | tests/test_stash.py | 21 | ||||
| -rw-r--r-- | ymlstash/stash.py | 8 |
2 files changed, 26 insertions, 3 deletions
diff --git a/tests/test_stash.py b/tests/test_stash.py index d6f8ab9..3ab86a3 100644 --- a/tests/test_stash.py +++ b/tests/test_stash.py @@ -3,7 +3,7 @@ import pytest from pathlib import Path from ymlstash import YmlStash from dataclasses import dataclass -from typing import ClassVar +from typing import ClassVar, Optional TEST_STASH_PATH = Path("/tmp") @@ -101,3 +101,22 @@ def test_dump_order(): assert keys == "ordeaz" stash.drop() + + +def test_null_values(): + @dataclass + class Nullable: + name: str + nullval: Optional[str] = None + key: ClassVar[str] = "name" + + stash = YmlStash(Nullable, TEST_STASH_PATH, filter_none=True) + goo = Nullable("goo") + stash.save(goo) + + with open(TEST_STASH_PATH / "goo.yml", "r") as f: + lines = f.readlines() + for line in lines: + assert "nullval" not in line + + stash.drop() diff --git a/ymlstash/stash.py b/ymlstash/stash.py index b49993b..b8f95d5 100644 --- a/ymlstash/stash.py +++ b/ymlstash/stash.py @@ -6,11 +6,12 @@ from pathlib import Path class YmlStash: - def __init__(self, model, path, file_suffix="yml", unsafe=False): + def __init__(self, model, path, file_suffix="yml", unsafe=False, filter_none=False): self.model = model self.path = Path(path) self.file_suffix = f".{file_suffix}" self.yaml_loader = yaml.SafeLoader + self.filter_none = filter_none self._validate() def _validate(self): @@ -40,7 +41,10 @@ class YmlStash: raise Exception("Cannot save object without a key or key field") key = getattr(obj, key_field) with open(self._get_path(key), "w") as f: - f.write(yaml.dump(asdict(obj), default_flow_style=False, sort_keys=False)) + obj_dict = asdict(obj) + if self.filter_none: + obj_dict = {k: v for (k, v) in obj_dict.items() if v is not None} + f.write(yaml.dump(obj_dict, default_flow_style=False, sort_keys=False)) def delete(self, key): try: |
