"""测试数据库层.""" import os import sqlite3 import tempfile import pytest from hermes_core.types import Record, Scope, TrainingRun, RecordState, TrainingStatus from hermes_core.db import ( init_db, insert_record, get_records_by_scope, get_active_records, get_record, update_record_state, update_record_label, upsert_scope, get_scope, get_active_scopes, get_scopes_needing_training, insert_training_run, update_training_run, get_latest_checkpoint, get_max_dims, set_max_dims, ) @pytest.fixture def conn(): """创建临时数据库用于测试。""" tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) conn = sqlite3.connect(tmp.name) init_db(conn) yield conn conn.close() try: os.unlink(tmp.name) except PermissionError: pass # Windows may hold file lock briefly class TestRecords: def test_insert_and_get_record(self, conn): rec = Record( id="rec_001", user_id="u1", scope_id="scope_a", scope_label="后端开发", dimensions=[], confidence=0.8, ) insert_record(conn, rec) got = get_record(conn, "rec_001") assert got is not None assert got.id == "rec_001" assert got.confidence == 0.8 def test_insert_upserts_on_duplicate(self, conn): rec1 = Record(id="rec_001", user_id="u1", scope_id="scope_a", scope_label="后端", dimensions=[], confidence=0.5) insert_record(conn, rec1) rec2 = Record(id="rec_001", user_id="u1", scope_id="scope_a", scope_label="后端", dimensions=[], confidence=0.9, occurrences=2) insert_record(conn, rec2) got = get_record(conn, "rec_001") assert got.confidence == 0.9 assert got.occurrences == 2 def test_get_records_by_scope(self, conn): for i in range(3): rec = Record(id=f"rec_{i}", user_id="u1", scope_id="scope_a", scope_label="后端", dimensions=[]) insert_record(conn, rec) records = get_records_by_scope(conn, "scope_a") assert len(records) == 3 def test_get_active_records_filters_rejected(self, conn): rec1 = Record(id="rec_1", user_id="u1", scope_id="scope_a", scope_label="后端", dimensions=[], state=RecordState.active) rec2 = Record(id="rec_2", user_id="u1", scope_id="scope_a", scope_label="后端", dimensions=[], state=RecordState.rejected) insert_record(conn, rec1) insert_record(conn, rec2) active = get_active_records(conn, "scope_a") assert len(active) == 1 assert active[0].id == "rec_1" def test_update_record_state(self, conn): rec = Record(id="rec_001", user_id="u1", scope_id="scope_a", scope_label="后端", dimensions=[], state=RecordState.active) insert_record(conn, rec) update_record_state(conn, "rec_001", RecordState.rejected) got = get_record(conn, "rec_001") assert got.state == RecordState.rejected def test_update_record_label(self, conn): rec = Record(id="rec_001", user_id="u1", scope_id="scope_a", scope_label="后端", dimensions=[]) insert_record(conn, rec) update_record_label(conn, "rec_001", "scope_b", "系统编程") got = get_record(conn, "rec_001") assert got.scope_id == "scope_b" assert got.scope_label == "系统编程" class TestScopes: def test_upsert_and_get_scope(self, conn): scope = Scope(id="scope_a", label="后端开发", centroid=[0.1, 0.2]) upsert_scope(conn, scope) got = get_scope(conn, "scope_a") assert got is not None assert got.label == "后端开发" def test_upsert_scope_merges(self, conn): s1 = Scope(id="scope_a", label="后端", centroid=[0.1], record_count=5) upsert_scope(conn, s1) s2 = Scope(id="scope_a", label="后端开发", centroid=[0.2], record_count=10, coherence=0.9) upsert_scope(conn, s2) got = get_scope(conn, "scope_a") assert got.label == "后端开发" assert got.record_count == 10 assert got.coherence == 0.9 def test_get_active_scopes(self, conn): s1 = Scope(id="scope_a", label="活跃", status="active") s2 = Scope(id="scope_b", label="归档", status="archived") upsert_scope(conn, s1) upsert_scope(conn, s2) scopes = get_active_scopes(conn) assert len(scopes) == 1 assert scopes[0].id == "scope_a" def test_get_scopes_needing_training(self, conn): s1 = Scope(id="scope_a", label="a", needs_training=True) s2 = Scope(id="scope_b", label="b", needs_training=False) s3 = Scope(id="scope_c", label="c", needs_training=True, status="archived") upsert_scope(conn, s1) upsert_scope(conn, s2) upsert_scope(conn, s3) needing = get_scopes_needing_training(conn) assert len(needing) == 1 assert needing[0].id == "scope_a" class TestTrainingRuns: def test_insert_and_get_latest_checkpoint(self, conn): run = TrainingRun(id="run_1", scope_id="scope_a", version=1, status=TrainingStatus.done, checkpoint_path="/tmp/v1") insert_training_run(conn, run) latest = get_latest_checkpoint(conn, "scope_a") assert latest is not None assert latest.version == 1 def test_update_training_run(self, conn): run = TrainingRun(id="run_1", scope_id="scope_a", version=1) insert_training_run(conn, run) update_training_run(conn, "run_1", status=TrainingStatus.training, progress={"epoch": 1, "loss": 0.5}) # verify by re-reading import json cur = conn.execute("SELECT status, progress FROM training_runs WHERE id='run_1'") row = cur.fetchone() assert row[0] == "training" assert json.loads(row[1]) == {"epoch": 1, "loss": 0.5} class TestDimensionConstraints: def test_get_set_max_dims(self, conn): max_dims = get_max_dims(conn, "scope_a") assert max_dims is None # 无记录时返回 None set_max_dims(conn, "scope_a", "u1", 3) assert get_max_dims(conn, "scope_a") == 3 set_max_dims(conn, "scope_a", "u1", 5) assert get_max_dims(conn, "scope_a") == 5