summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--tests/test_stash.py32
-rw-r--r--ymlstash/stash.py7
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: