feat(#64): 配置项 CRUD 存储引擎(模型超参/RAG/布局三类,文件系统版本化 JSON+校验)
This commit is contained in:
@@ -0,0 +1,191 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""配置项 CRUD 存储引擎测试(issue #64)。
|
||||
|
||||
覆盖:
|
||||
1. 三类配置 CRUD(list/get/upsert/delete);
|
||||
2. 原子写 + 持久化(重开 store 仍在);
|
||||
3. 校验规则(model_param/rag_config/layout,非法值拒绝);
|
||||
4. 快照 snapshot/restore(为 #66 提供基础);
|
||||
5. 可解释字段(meaning/reason/updated_by/updated_at 落盘)。
|
||||
"""
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
import _bootstrap # noqa: F401
|
||||
|
||||
from template_console.config_store import ( # noqa: E402
|
||||
ALLOWED_WIDGET_TYPES,
|
||||
ConfigItem,
|
||||
ConfigKind,
|
||||
ConfigStore,
|
||||
ValidationResult,
|
||||
validate_item,
|
||||
)
|
||||
|
||||
|
||||
class _TmpStore:
|
||||
def __init__(self):
|
||||
self._tmp = tempfile.mkdtemp()
|
||||
self.store = ConfigStore(self._tmp)
|
||||
|
||||
def cleanup(self):
|
||||
import shutil
|
||||
shutil.rmtree(self._tmp, ignore_errors=True)
|
||||
|
||||
|
||||
class ValidationTest(unittest.TestCase):
|
||||
"""校验规则。"""
|
||||
|
||||
def test_model_param_scalar_ok(self):
|
||||
self.assertTrue(validate_item(ConfigKind.MODEL_PARAM, "learning_rate", 0.001))
|
||||
|
||||
def test_model_param_learning_rate_range(self):
|
||||
vr = validate_item(ConfigKind.MODEL_PARAM, "learning_rate", 1.5)
|
||||
self.assertFalse(vr)
|
||||
self.assertTrue(any("learning_rate" in e for e in vr.errors))
|
||||
|
||||
def test_model_param_bad_key(self):
|
||||
vr = validate_item(ConfigKind.MODEL_PARAM, "Bad Key!", 1)
|
||||
self.assertFalse(vr)
|
||||
|
||||
def test_rag_top_k_bounds(self):
|
||||
self.assertFalse(validate_item(ConfigKind.RAG_CONFIG, "top_k", 0))
|
||||
self.assertFalse(validate_item(ConfigKind.RAG_CONFIG, "top_k", 51))
|
||||
self.assertTrue(validate_item(ConfigKind.RAG_CONFIG, "top_k", 10))
|
||||
|
||||
def test_rag_similarity_threshold(self):
|
||||
self.assertTrue(validate_item(ConfigKind.RAG_CONFIG, "similarity_threshold", 0.5))
|
||||
self.assertFalse(validate_item(ConfigKind.RAG_CONFIG, "similarity_threshold", 1.5))
|
||||
|
||||
def test_rag_sources_must_be_nonempty_list(self):
|
||||
self.assertFalse(validate_item(ConfigKind.RAG_CONFIG, "sources", []))
|
||||
self.assertFalse(validate_item(ConfigKind.RAG_CONFIG, "sources", ["", "x"]))
|
||||
self.assertTrue(validate_item(ConfigKind.RAG_CONFIG, "sources", ["sop", "gb"]))
|
||||
|
||||
def test_layout_widget_type(self):
|
||||
bad = [{"type": "unknown", "x": 0, "y": 0, "w": 1, "h": 1}]
|
||||
self.assertFalse(validate_item(ConfigKind.LAYOUT, "dashboard", bad))
|
||||
good = [{"type": "trend", "x": 0, "y": 0, "w": 6, "h": 2}]
|
||||
self.assertTrue(validate_item(ConfigKind.LAYOUT, "dashboard", good))
|
||||
|
||||
def test_layout_widget_coords_nonneg_int(self):
|
||||
bad = [{"type": "trend", "x": -1, "y": 0, "w": 1, "h": 1}]
|
||||
vr = validate_item(ConfigKind.LAYOUT, "dashboard", bad)
|
||||
self.assertFalse(vr)
|
||||
|
||||
|
||||
class CrudTest(unittest.TestCase):
|
||||
"""CRUD + 持久化。"""
|
||||
|
||||
def setUp(self):
|
||||
self.ctx = _TmpStore()
|
||||
self.store = self.ctx.store
|
||||
|
||||
def tearDown(self):
|
||||
self.ctx.cleanup()
|
||||
|
||||
def test_upsert_and_get(self):
|
||||
self.store.upsert(ConfigKind.MODEL_PARAM, "iterations", 100,
|
||||
meaning="迭代数", updated_by="li", reason="标定")
|
||||
it = self.store.get(ConfigKind.MODEL_PARAM, "iterations")
|
||||
self.assertIsNotNone(it)
|
||||
self.assertEqual(it.value, 100)
|
||||
self.assertEqual(it.updated_by, "li")
|
||||
self.assertEqual(it.reason, "标定")
|
||||
self.assertTrue(it.updated_at) # 时间戳已写
|
||||
|
||||
def test_upsert_rejects_invalid(self):
|
||||
with self.assertRaises(ValueError):
|
||||
self.store.upsert(ConfigKind.RAG_CONFIG, "top_k", 999)
|
||||
|
||||
def test_upsert_overwrites(self):
|
||||
self.store.upsert(ConfigKind.MODEL_PARAM, "lr", 0.1)
|
||||
self.store.upsert(ConfigKind.MODEL_PARAM, "lr", 0.01, reason="调小")
|
||||
it = self.store.get(ConfigKind.MODEL_PARAM, "lr")
|
||||
self.assertEqual(it.value, 0.01)
|
||||
self.assertEqual(it.reason, "调小")
|
||||
|
||||
def test_list_and_delete(self):
|
||||
self.store.upsert(ConfigKind.RAG_CONFIG, "top_k", 5)
|
||||
self.store.upsert(ConfigKind.RAG_CONFIG, "similarity_threshold", 0.6)
|
||||
self.assertEqual(len(self.store.list(ConfigKind.RAG_CONFIG)), 2)
|
||||
self.assertTrue(self.store.delete(ConfigKind.RAG_CONFIG, "top_k"))
|
||||
self.assertIsNone(self.store.get(ConfigKind.RAG_CONFIG, "top_k"))
|
||||
self.assertFalse(self.store.delete(ConfigKind.RAG_CONFIG, "nope"))
|
||||
|
||||
def test_persistence_across_reopen(self):
|
||||
self.store.upsert(ConfigKind.MODEL_PARAM, "lr", 0.001)
|
||||
# 重开一个指向同一目录的 store
|
||||
store2 = ConfigStore(self.ctx._tmp)
|
||||
it = store2.get(ConfigKind.MODEL_PARAM, "lr")
|
||||
self.assertIsNotNone(it)
|
||||
self.assertEqual(it.value, 0.001)
|
||||
|
||||
def test_json_file_is_human_readable(self):
|
||||
self.store.upsert(ConfigKind.MODEL_PARAM, "lr", 0.001, meaning="学习率")
|
||||
path = os.path.join(self.ctx._tmp, "model_params.json")
|
||||
with open(path, encoding="utf-8") as fh:
|
||||
blob = json.load(fh)
|
||||
self.assertEqual(blob["schema_version"], 1)
|
||||
self.assertEqual(blob["kind"], "model_param")
|
||||
self.assertEqual(blob["items"][0]["meaning"], "学习率")
|
||||
|
||||
|
||||
class SnapshotTest(unittest.TestCase):
|
||||
"""快照与恢复(#66 基础)。"""
|
||||
|
||||
def setUp(self):
|
||||
self.ctx = _TmpStore()
|
||||
self.store = self.ctx.store
|
||||
|
||||
def tearDown(self):
|
||||
self.ctx.cleanup()
|
||||
|
||||
def test_snapshot_captures_all_kinds(self):
|
||||
self.store.upsert(ConfigKind.MODEL_PARAM, "lr", 0.001)
|
||||
self.store.upsert(ConfigKind.RAG_CONFIG, "top_k", 8)
|
||||
snap = self.store.snapshot()
|
||||
self.assertIn("captured_at", snap)
|
||||
self.assertEqual(set(snap["kinds"].keys()),
|
||||
{"model_param", "rag_config", "layout"})
|
||||
self.assertEqual(len(snap["kinds"]["model_param"]), 1)
|
||||
|
||||
def test_restore_replicates_state(self):
|
||||
self.store.upsert(ConfigKind.MODEL_PARAM, "lr", 0.001)
|
||||
self.store.upsert(ConfigKind.RAG_CONFIG, "top_k", 8)
|
||||
snap = self.store.snapshot()
|
||||
# 清空再恢复
|
||||
self.store.delete(ConfigKind.MODEL_PARAM, "lr")
|
||||
self.store.delete(ConfigKind.RAG_CONFIG, "top_k")
|
||||
self.store.restore(snap)
|
||||
self.assertEqual(self.store.get(ConfigKind.MODEL_PARAM, "lr").value, 0.001)
|
||||
self.assertEqual(self.store.get(ConfigKind.RAG_CONFIG, "top_k").value, 8)
|
||||
|
||||
def test_item_counts(self):
|
||||
self.store.upsert(ConfigKind.LAYOUT, "dashboard",
|
||||
[{"type": "trend", "x": 0, "y": 0, "w": 6, "h": 2}])
|
||||
counts = self.store.item_counts()
|
||||
self.assertEqual(counts["layout"], 1)
|
||||
self.assertEqual(counts["model_param"], 0)
|
||||
|
||||
|
||||
class ConfigItemSerializationTest(unittest.TestCase):
|
||||
"""ConfigItem 序列化往返。"""
|
||||
|
||||
def test_roundtrip(self):
|
||||
it = ConfigItem(key="lr", value=0.1, kind=ConfigKind.MODEL_PARAM,
|
||||
meaning="学习率", updated_by="li", reason="init",
|
||||
updated_at="2026-01-01T00:00:00Z")
|
||||
d = it.to_dict()
|
||||
self.assertEqual(d["kind"], "model_param")
|
||||
it2 = ConfigItem.from_dict(d)
|
||||
self.assertEqual(it2.value, 0.1)
|
||||
self.assertEqual(it2.kind, ConfigKind.MODEL_PARAM)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user