File size: 4,196 Bytes
15879c4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""query — 场景识别 + 偏好检索."""

from hermes_core.embedder import Embedder
from hermes_core.types import QueryResult, MatchedScope, LoRAInfo, PreferenceItem
from hermes_core.db import (
    init_db, get_active_scopes, get_active_records, get_latest_checkpoint,
)
from hermes_core.cluster import DEFAULT_MATCH_THRESHOLD
from hermes_core.recorder import record_detail


def query(user_id: str, text: str, embedder: Embedder) -> QueryResult:
    """根据用户输入匹配 scope 并返回相关偏好。

    Args:
        user_id: 用户 ID
        text: 用户输入文本
        embedder: Embedding 服务实例

    Returns:
        QueryResult: 包含 matched_scope, active_loras, related_preferences
    """
    vec = embedder.encode(text)
    conn = init_db(user_id)
    scopes = get_active_scopes(conn)

    matched = None
    alternatives = []

    for scope in scopes:
        if scope.centroid is None:
            continue
        sim = embedder.cosine_similarity(vec, scope.centroid)
        entry = MatchedScope(scope_id=scope.id, scope_label=scope.label, confidence=float(sim))
        if sim >= DEFAULT_MATCH_THRESHOLD:
            if matched is None or sim > matched.confidence:
                if matched is not None:
                    alternatives.append(matched)
                matched = entry
            else:
                alternatives.append(entry)
        elif sim > 0.3:
            alternatives.append(entry)

    # 构建 active_loras
    active_loras = []
    training_outdated = False

    # behavior lora (特殊 scope_id)
    behavior_checkpoint = get_latest_checkpoint(conn, "behavior")
    if behavior_checkpoint and behavior_checkpoint.status.value == "done":
        active_loras.append(LoRAInfo(scope_id="behavior", version=f"v{behavior_checkpoint.version}", priority=0))

    if matched is not None:
        checkpoint = get_latest_checkpoint(conn, matched.scope_id)
        if checkpoint and checkpoint.status.value == "done":
            active_loras.append(LoRAInfo(scope_id=matched.scope_id, version=f"v{checkpoint.version}", priority=1))
        # 检查是否需要训练(有记录但无 checkpoint,或记录数变化)
        records = get_active_records(conn, matched.scope_id)
        if len(records) > 0 and (checkpoint is None or checkpoint.status.value != "done"):
            training_outdated = True

    # 检索相关偏好
    related_prefs = []
    seen_keys = set()
    if matched is not None:
        records = get_active_records(conn, matched.scope_id)
        for rec in records:
            for dim in rec.dimensions:
                if dim.key not in seen_keys:
                    seen_keys.add(dim.key)
                    related_prefs.append(PreferenceItem(
                        key=dim.key, value=dim.value, source=rec.id
                    ))

    conn.close()
    return QueryResult(
        matched_scope=matched,
        alternative_scopes=alternatives,
        active_loras=active_loras,
        related_preferences=related_prefs,
        training_outdated=training_outdated,
    )


class HermesClient:
    """Agent 侧集成入口。

    用法:
        client = HermesClient(user_id="u_alex", agent_id="my-agent")
        prefs = client.query("帮我写个用户管理模块")
        # ... 推理 ...
        client.record("后端开发", [{"key": "lang", "value": "TS", "context": "默认"}])
    """

    def __init__(self, user_id: str, agent_id: str = "",
                 model_name: str = "paraphrase-multilingual-MiniLM-L12-v2"):
        self.user_id = user_id
        self.agent_id = agent_id
        self._embedder = Embedder(model_name)

    def query(self, text: str) -> QueryResult:
        return query(self.user_id, text, self._embedder)

    def record(self, scope_desc: str, dimensions: list[dict],
               source_conv: str = "", conversation_id: str = "") -> dict:
        return record_detail(
            user_id=self.user_id,
            scope_desc=scope_desc,
            dimensions=dimensions,
            embedder=self._embedder,
            source_agent=self.agent_id,
            source_conv=source_conv,
            conversation_id=conversation_id,
        )