summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorYuval Adam <_@yuv.al>2024-05-09 12:34:10 +0200
committerYuval Adam <_@yuv.al>2024-05-09 12:34:10 +0200
commitd089e3acb34cf3214c4c1a52b3142798de5f92f2 (patch)
tree94dd8a58bfe4dff6c71323b43226eeeac654c44e
parenteaab19e310edd682af0231fe06f9a647bfcd5834 (diff)
Add support for filtering null values
-rw-r--r--tests/test_stash.py21
-rw-r--r--ymlstash/stash.py8
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: