Browse Source

feat(memory): speaker-aware conversation memory via Zvec

- Save each conversation turn (query+reply) indexed by speaker_id
- Load historical context on reconnection for the same speaker
- Text embedding via char n-gram hash projection
- Speaker context + memory injected as system messages
- AI renamed to 狄诺尼试验员
wenhongquan 3 weeks ago
parent
commit
f7486333df
2 changed files with 136 additions and 13 deletions
  1. 111 0
      asr_agent/voiceprint.py
  2. 25 13
      asr_agent/worker.py

+ 111 - 0
asr_agent/voiceprint.py

@@ -216,3 +216,114 @@ def get_store() -> VoiceprintStore:
     if _store is None:
     if _store is None:
         _store = VoiceprintStore()
         _store = VoiceprintStore()
     return _store
     return _store
+
+
+# ═══════════════════════════════════════════════════════════════
+#  Conversation Memory — speaker-aware chat history in Zvec
+# ═══════════════════════════════════════════════════════════════
+
+_MEM_COLLECTION = None
+_MEM_DIM = 64  # small projection for topic similarity
+
+
+def _get_memory_collection():
+    global _MEM_COLLECTION
+    if _MEM_COLLECTION is not None:
+        return _MEM_COLLECTION
+    try:
+        import zvec
+        mem_path = os.path.join(VP_DB_PATH, "conversations")
+        os.makedirs(mem_path, exist_ok=True)
+
+        schema = zvec.CollectionSchema(
+            name="conversations",
+            fields=[
+                zvec.FieldSchema(name="speaker_name", data_type=zvec.DataType.STRING),
+                zvec.FieldSchema(name="speaker_id", data_type=zvec.DataType.STRING),
+                zvec.FieldSchema(name="user_query", data_type=zvec.DataType.STRING),
+                zvec.FieldSchema(name="reply", data_type=zvec.DataType.STRING),
+                zvec.FieldSchema(name="timestamp", data_type=zvec.DataType.FLOAT64),
+            ],
+            vectors=[
+                zvec.VectorSchema(
+                    name="embedding",
+                    data_type=zvec.DataType.VECTOR_FP32,
+                    dimension=_MEM_DIM,
+                    index_param=zvec.HnswIndexParam(metric_type=zvec.MetricType.COSINE),
+                ),
+            ],
+        )
+        _MEM_COLLECTION = zvec.create_and_open(path=mem_path, schema=schema)
+        _MEM_COLLECTION.optimize()
+        logger.info("Conversation memory DB opened at %s (%d entries)",
+                    mem_path, _MEM_COLLECTION.stats.row_count)
+    except Exception:
+        logger.exception("Failed to open conversation memory DB")
+        return None
+    return _MEM_COLLECTION
+
+
+def _text_embedding(text: str) -> list[float]:
+    """Simple text embedding via character n-gram hash projection."""
+    vec = np.zeros(_MEM_DIM, dtype=np.float32)
+    for i in range(len(text) - 2):
+        h = hash(text[i:i + 3]) % _MEM_DIM
+        vec[h] += 1.0
+    norm = np.linalg.norm(vec) + 1e-8
+    return (vec / norm).tolist()
+
+
+def save_conversation(speaker_id: str, speaker_name: str, user_query: str, reply: str):
+    """Save a conversation turn to persistent memory."""
+    col = _get_memory_collection()
+    if col is None:
+        return
+    try:
+        import zvec
+        full_text = user_query + " " + reply
+        doc_id = f"{speaker_id}_{int(time.time() * 1000)}"
+        col.insert(zvec.Doc(
+            id=doc_id,
+            vectors={"embedding": _text_embedding(full_text)},
+            fields={
+                "speaker_name": speaker_name,
+                "speaker_id": speaker_id,
+                "user_query": user_query,
+                "reply": reply,
+                "timestamp": time.time(),
+            },
+        ))
+    except Exception:
+        logger.exception("Failed to save conversation")
+
+
+def load_conversation_context(speaker_id: str | None, speaker_name: str | None, limit: int = 5) -> str:
+    """Retrieve recent conversation history for a speaker. Returns formatted context string."""
+    if not speaker_id:
+        return ""
+    col = _get_memory_collection()
+    if col is None:
+        return ""
+    try:
+        import zvec
+        result = col.query(
+            queries=zvec.Query(
+                field_name="embedding",
+                vector=[0.0] * _MEM_DIM,  # dummy vector for scalar-only query
+            ),
+            topk=limit,
+            filter=f'speaker_id == "{speaker_id}"',
+        )
+        if not result:
+            return ""
+        history = []
+        for doc in reversed(result):
+            q = doc.fields.get("user_query", "")
+            r = doc.fields.get("reply", "")
+            if q and r:
+                history.append(f"用户: {q}\n助手: {r}")
+        if history:
+            return "以下是之前和该用户的对话记录,请参考上下文回应:\n" + "\n".join(history)
+    except Exception:
+        logger.exception("Failed to load conversation context")
+    return ""

+ 25 - 13
asr_agent/worker.py

@@ -37,12 +37,13 @@ from audio_processor import AudioBuffer
 from qwen_engine import QwenASREngine
 from qwen_engine import QwenASREngine
 
 
 try:
 try:
-    from voiceprint import get_store
+    from voiceprint import get_store, save_conversation, load_conversation_context
     _VP_ENABLED = True
     _VP_ENABLED = True
 except ImportError:
 except ImportError:
     _VP_ENABLED = False
     _VP_ENABLED = False
-    def get_store():
-        return None
+    def get_store(): return None
+    def save_conversation(*a): pass
+    def load_conversation_context(*a): return ""
 
 
 logger = logging.getLogger("worker")
 logger = logging.getLogger("worker")
 
 
@@ -65,7 +66,7 @@ LLM_TEMPERATURE = 0.7
 LLM_TIMEOUT     = 25
 LLM_TIMEOUT     = 25
 
 
 SYSTEM_PROMPT = (
 SYSTEM_PROMPT = (
-    "你是友好的中文语音助手,名字叫小智。"
+    "你是狄诺尼试验员,一名专业的工程试验助手。"
     "回答简洁自然,2-3句即可,用口语化中文。"
     "回答简洁自然,2-3句即可,用口语化中文。"
     "不要使用 Markdown、代码块、表格或特殊符号。"
     "不要使用 Markdown、代码块、表格或特殊符号。"
     "不要输出括号注释、不要使用英文缩写。"
     "不要输出括号注释、不要使用英文缩写。"
@@ -97,21 +98,25 @@ def _token(room, ident):
 #  LLM (streaming, <think> filtering)
 #  LLM (streaming, <think> filtering)
 # ═══════════════════════════════════════════════════════════════
 # ═══════════════════════════════════════════════════════════════
 
 
-async def _llm_stream(prompt: str, hist: list[dict], speaker: str | None = None):
+async def _llm_stream(prompt: str, hist: list[dict], speaker: str | None = None, speaker_id: str | None = None):
     """Stream LLM tokens, dispatching based on LLM_PROVIDER env."""
     """Stream LLM tokens, dispatching based on LLM_PROVIDER env."""
     if LLM_PROVIDER == "mimo":
     if LLM_PROVIDER == "mimo":
-        async for x in _llm_stream_mimo(prompt, hist, speaker):
+        async for x in _llm_stream_mimo(prompt, hist, speaker, speaker_id):
             yield x
             yield x
         return
         return
-    async for x in _llm_stream_vllm(prompt, hist, speaker):
+    async for x in _llm_stream_vllm(prompt, hist, speaker, speaker_id):
         yield x
         yield x
 
 
 
 
-async def _llm_stream_vllm(prompt: str, hist: list[dict], speaker: str | None = None):
+async def _llm_stream_vllm(prompt: str, hist: list[dict], speaker: str | None = None, speaker_id: str | None = None):
     """vLLM streaming via OpenAI-compatible SSE."""
     """vLLM streaming via OpenAI-compatible SSE."""
     msgs = [{"role": "system", "content": SYSTEM_PROMPT}]
     msgs = [{"role": "system", "content": SYSTEM_PROMPT}]
     if speaker:
     if speaker:
         msgs.append({"role": "system", "content": f"[内部上下文] 当前说话人: {speaker}。请根据对方身份自然地调整回应风格,但不要在回复中主动提及说话人的名字。"})
         msgs.append({"role": "system", "content": f"[内部上下文] 当前说话人: {speaker}。请根据对方身份自然地调整回应风格,但不要在回复中主动提及说话人的名字。"})
+    # Conversation memory context
+    mem_ctx = load_conversation_context(speaker_id, speaker) if speaker_id else ""
+    if mem_ctx:
+        msgs.append({"role": "system", "content": mem_ctx})
     msgs.extend(hist)
     msgs.extend(hist)
     msgs.append({"role": "user", "content": prompt})
     msgs.append({"role": "user", "content": prompt})
 
 
@@ -147,7 +152,7 @@ async def _llm_stream_vllm(prompt: str, hist: list[dict], speaker: str | None =
     yield "", True
     yield "", True
 
 
 
 
-async def _llm_stream_mimo(prompt: str, hist: list[dict], speaker: str | None = None):
+async def _llm_stream_mimo(prompt: str, hist: list[dict], speaker: str | None = None, speaker_id: str | None = None):
     """Mimo v2.5 LLM streaming via OpenAI-compatible SSE."""
     """Mimo v2.5 LLM streaming via OpenAI-compatible SSE."""
     if not MIMO_KEY:
     if not MIMO_KEY:
         yield "", True
         yield "", True
@@ -155,6 +160,9 @@ async def _llm_stream_mimo(prompt: str, hist: list[dict], speaker: str | None =
     msgs = [{"role": "system", "content": SYSTEM_PROMPT}]
     msgs = [{"role": "system", "content": SYSTEM_PROMPT}]
     if speaker:
     if speaker:
         msgs.append({"role": "system", "content": f"[内部上下文] 当前说话人: {speaker}。请根据对方身份自然地调整回应风格,但不要在回复中主动提及说话人的名字。"})
         msgs.append({"role": "system", "content": f"[内部上下文] 当前说话人: {speaker}。请根据对方身份自然地调整回应风格,但不要在回复中主动提及说话人的名字。"})
+    mem_ctx = load_conversation_context(speaker_id, speaker) if speaker_id else ""
+    if mem_ctx:
+        msgs.append({"role": "system", "content": mem_ctx})
     msgs.extend(hist)
     msgs.extend(hist)
     msgs.append({"role": "user", "content": prompt})
     msgs.append({"role": "user", "content": prompt})
 
 
@@ -631,21 +639,21 @@ class Worker:
 
 
         # ── Voiceprint: identify or auto-register ──
         # ── Voiceprint: identify or auto-register ──
         speaker_name = None
         speaker_name = None
+        speaker_id = None
         if _VP_ENABLED:
         if _VP_ENABLED:
             store = get_store()
             store = get_store()
             if store and store.enabled:
             if store and store.enabled:
                 loop = asyncio.get_running_loop()
                 loop = asyncio.get_running_loop()
-                _id, speaker_name, _sim = await loop.run_in_executor(None, store.identify, all_audio)
+                speaker_id, speaker_name, _sim = await loop.run_in_executor(None, store.identify, all_audio)
                 if speaker_name:
                 if speaker_name:
                     logger.info("VP: identified %s (sim=%.3f)", speaker_name, _sim)
                     logger.info("VP: identified %s (sim=%.3f)", speaker_name, _sim)
                     await self._send({"type": "speaker", "name": speaker_name, "confidence": round(_sim, 3)})
                     await self._send({"type": "speaker", "name": speaker_name, "confidence": round(_sim, 3)})
                 else:
                 else:
-                    # Auto-register: first-time speaker
                     auto_id = f"user_{int(time.time() * 1000) % 1000000:06d}"
                     auto_id = f"user_{int(time.time() * 1000) % 1000000:06d}"
                     auto_name = f"用户{auto_id[-3:]}"
                     auto_name = f"用户{auto_id[-3:]}"
                     ok = await loop.run_in_executor(None, store.register, auto_id, auto_name, all_audio)
                     ok = await loop.run_in_executor(None, store.register, auto_id, auto_name, all_audio)
                     if ok:
                     if ok:
-                        speaker_name = auto_name
+                        speaker_id, speaker_name = auto_id, auto_name
                         logger.info("VP: auto-registered %s", auto_name)
                         logger.info("VP: auto-registered %s", auto_name)
                         await self._send({"type": "vp_registered", "name": auto_name, "user_id": auto_id})
                         await self._send({"type": "vp_registered", "name": auto_name, "user_id": auto_id})
 
 
@@ -655,7 +663,7 @@ class Worker:
         tts_segments: list[str] = []
         tts_segments: list[str] = []
 
 
         try:
         try:
-            async for delta, is_final in _llm_stream(txt, self.hist, speaker_name):
+            async for delta, is_final in _llm_stream(txt, self.hist, speaker_name, speaker_id):
                 if delta:
                 if delta:
                     reply_full += delta
                     reply_full += delta
                     await self._send({"type": "reply_partial", "text": reply_full, "seq": 0})
                     await self._send({"type": "reply_partial", "text": reply_full, "seq": 0})
@@ -693,6 +701,10 @@ class Worker:
                 {"role": "assistant", "content": reply_full},
                 {"role": "assistant", "content": reply_full},
             ])
             ])
             self.hist[:] = self.hist[-20:]
             self.hist[:] = self.hist[-20:]
+            # Save conversation to persistent memory
+            if speaker_id and _VP_ENABLED:
+                loop = asyncio.get_running_loop()
+                await loop.run_in_executor(None, save_conversation, speaker_id, speaker_name or "unknown", txt, reply_full)
 
 
     async def _send(self, msg: dict):
     async def _send(self, msg: dict):
         try:
         try: