File size: 6,512 Bytes
a6898ea
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
"""测试数据库层."""

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