hermes-association-engine / tests /test_cluster.py
feiertu's picture
Upload tests/test_cluster.py with huggingface_hub
00138da verified
Raw
History Blame Contribute Delete
4.44 kB
"""测试动态语义聚类."""
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