feat(#64): 配置项 CRUD 存储引擎(模型超参/RAG/布局三类,文件系统版本化 JSON+校验)

This commit is contained in:
2026-08-05 05:29:53 +08:00
parent 94c2782e75
commit b0ae378150
2 changed files with 482 additions and 0 deletions
@@ -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()