feiertu's picture
Upload tests/test_db.py with huggingface_hub
a6898ea verified
Raw
History Blame Contribute Delete
6.51 kB
"""测试数据库层."""
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