diff options
| -rw-r--r-- | tests/test_stash.py | 32 | ||||
| -rw-r--r-- | ymlstash/stash.py | 7 |
2 files changed, 34 insertions, 5 deletions
diff --git a/tests/test_stash.py b/tests/test_stash.py index 0e84eae..d6f8ab9 100644 --- a/tests/test_stash.py +++ b/tests/test_stash.py @@ -5,7 +5,7 @@ from ymlstash import YmlStash from dataclasses import dataclass from typing import ClassVar -TEST_STASH_PATH = "/tmp" +TEST_STASH_PATH = Path("/tmp") @dataclass @@ -21,9 +21,12 @@ def test_stash_path(): assert stash.path == Path(".") -def test_invalid_path(): +def test_invalid(): with pytest.raises(Exception): YmlStash(User, "/tmp/does/not/exist") + with pytest.raises(Exception) as e: + YmlStash(object, ".") + assert "is not a dataclass" in e.value.args[0] def test_stash(): @@ -55,7 +58,7 @@ def test_stash(): def test_key_field(): @dataclass class Rat: - key: ClassVar[str] = "name" + key: ClassVar[str] = "foo" with pytest.raises(Exception): YmlStash(Rat, ".") @@ -75,3 +78,26 @@ def test_key_field(): assert stash.list_keys() == ["terra", "dupe"] stash.drop() + + +def test_dump_order(): + @dataclass + class Order: + o: int + r: int + d: int + e: int + a: int + z: int + key: ClassVar[str] = "o" + + stash = YmlStash(Order, TEST_STASH_PATH) + order = Order(1, 2, 3, 4, 5, 6) + stash.save(order) + + with open(TEST_STASH_PATH / "1.yml", "r") as f: + lines = f.readlines() + keys = "".join([line[0] for line in lines]) + assert keys == "ordeaz" + + stash.drop() diff --git a/ymlstash/stash.py b/ymlstash/stash.py index 52d74e0..b49993b 100644 --- a/ymlstash/stash.py +++ b/ymlstash/stash.py @@ -1,7 +1,7 @@ import yaml import os -from dataclasses import asdict, fields +from dataclasses import asdict, fields, is_dataclass from pathlib import Path @@ -17,6 +17,9 @@ class YmlStash: if not self.path.exists(): raise Exception(f"Path {self.path} does not exist and cannot be used") + if not is_dataclass(self.model): + raise Exception(f"Class {self.model} is not a dataclass") + key_field = getattr(self.model, "key", None) field_names = [f.name for f in fields(self.model)] if key_field and key_field not in field_names: @@ -37,7 +40,7 @@ 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))) + f.write(yaml.dump(asdict(obj), default_flow_style=False, sort_keys=False)) def delete(self, key): try: |
