"""测试动态语义聚类.""" 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