File size: 4,442 Bytes
00138da
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""测试动态语义聚类."""

import tempfile
import sqlite3
import os
import pytest
from hermes_core.db import init_db, upsert_scope, insert_record, get_active_scopes, get_scope
from hermes_core.embedder import Embedder
from hermes_core.types import Record, Scope
from hermes_core.cluster import (
    assign_scope,
    recluster_scope,
    check_merge,
    check_split,
    DEFAULT_MATCH_THRESHOLD,
    DEFAULT_MERGE_THRESHOLD,
    DEFAULT_SPLIT_THRESHOLD,
)

pytestmark = pytest.mark.requires_embedder


@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


@pytest.fixture
def embedder():
    return Embedder()


class TestAssignScope:
    def test_creates_new_scope_when_none_exist(self, conn, embedder):
        sid, label = assign_scope(conn, embedder, "后端API开发")
        assert sid.startswith("scope_")
        assert label == "后端API开发"
        scopes = get_active_scopes(conn)
        assert len(scopes) == 1

    def test_matches_existing_scope_on_high_similarity(self, conn, embedder):
        # 先创建一个 scope
        s = Scope(id="scope_existing", label="后端服务开发",
                  centroid=embedder.encode("后端服务开发"), record_count=5)
        upsert_scope(conn, s)
        # 赋一条语义相近的
        sid, label = assign_scope(conn, embedder, "后端API开发")
        # 应该匹配到已有 scope
        assert sid == "scope_existing"

    def test_creates_new_scope_on_low_similarity(self, conn, embedder):
        s = Scope(id="scope_existing", label="后端开发",
                  centroid=embedder.encode("后端开发"), record_count=5)
        upsert_scope(conn, s)
        sid, label = assign_scope(conn, embedder, "周末去哪玩")
        assert sid != "scope_existing"
        assert len(get_active_scopes(conn)) == 2


class TestRecluster:
    def test_recluster_updates_centroid(self, conn, embedder):
        s = Scope(id="scope_a", label="后端",
                  centroid=embedder.encode("后端开发"), record_count=1)
        upsert_scope(conn, s)
        insert_record(conn, Record(
            id="r1", user_id="u1", scope_id="scope_a",
            scope_label="后端", dimensions=[]
        ))
        insert_record(conn, Record(
            id="r2", user_id="u1", scope_id="scope_a",
            scope_label="后端服务", dimensions=[]
        ))
        recluster_scope(conn, embedder, "scope_a")
        updated = get_scope(conn, "scope_a")
        # 新的 centroid 应该是两条记录 label 的平均
        assert updated.centroid is not None
        assert updated.record_count == 2


class TestCheckMerge:
    def test_similar_scopes_flagged_for_merge(self, conn, embedder):
        s1 = Scope(id="scope_1", label="Python数据分析",
                   centroid=embedder.encode("Python数据分析"), record_count=3)
        s2 = Scope(id="scope_2", label="Python数据处理",
                   centroid=embedder.encode("Python数据处理"), record_count=3)
        upsert_scope(conn, s1)
        upsert_scope(conn, s2)
        merges = check_merge(conn, embedder, threshold=0.85)
        assert len(merges) > 0

    def test_dissimilar_scopes_not_merged(self, conn, embedder):
        s1 = Scope(id="scope_1", label="后端开发",
                   centroid=embedder.encode("后端开发"), record_count=3)
        s2 = Scope(id="scope_2", label="周末去哪玩",
                   centroid=embedder.encode("周末去哪玩"), record_count=3)
        upsert_scope(conn, s1)
        upsert_scope(conn, s2)
        merges = check_merge(conn, embedder, threshold=0.85)
        assert len(merges) == 0


class TestCheckSplit:
    def test_coherent_scope_not_flagged(self, conn, embedder):
        label = "后端开发"
        s = Scope(id="scope_a", label=label,
                  centroid=embedder.encode(label), record_count=1)
        upsert_scope(conn, s)
        for i in range(5):
            insert_record(conn, Record(
                id=f"r{i}", user_id="u1", scope_id="scope_a",
                scope_label=f"后端开发-{i}", dimensions=[]
            ))
        needs_split = check_split(conn, embedder, "scope_a", threshold=0.40)
        assert not needs_split