From 41284f0bb1882cffc396e2d115821345277279e4 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=AC=A7=E9=98=B3=E6=B4=8B?= <2443479321@qq.com> Date: Sun, 13 Sep 2026 21:22:54 +0800 Subject: [PATCH] =?UTF-8?q?fix:=E8=AE=B0=E5=BF=86=E6=9E=B6=E6=9E=84?= =?UTF-8?q?=E4=BC=98=E5=8C=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- repositories/memory_unit.py | 30 ++++++ scripts/apply_memory_unit_upgrade.py | 138 +++++++++++++++++++++++++++ service/client_agent/runtime.py | 69 +++++++++++++- service/memory/facade.py | 2 +- service/memory/long_term.py | 46 ++++++++- service/memory/milvus_memory.py | 34 ++++++- service/memory/schemas.py | 1 + sql/memory_unit_upgrade_20260913.sql | 45 +++++++++ 8 files changed, 354 insertions(+), 11 deletions(-) create mode 100644 scripts/apply_memory_unit_upgrade.py create mode 100644 sql/memory_unit_upgrade_20260913.sql diff --git a/repositories/memory_unit.py b/repositories/memory_unit.py index 61dee88..6f6ebb8 100644 --- a/repositories/memory_unit.py +++ b/repositories/memory_unit.py @@ -68,6 +68,36 @@ class MemoryUnitRepo(BaseRepository): ) return list((await self.db.scalars(statement)).all()) + async def list_for_customer_by_ids( + self, + customer_id: int, + ids: list[int], + *, + memory_type: str | None = None, + tag: str | None = None, + now: datetime | None = None, + ) -> list[MemoryUnit]: + """按 ID 集合取回本人有效记忆,用于向量召回命中后的主体回表。""" + if not ids: + return [] + now = now or datetime.now() + conditions = [ + MemoryUnit.customer_id == customer_id, + MemoryUnit.id.in_(ids), + MemoryUnit.status.in_(ACTIVE_STATUSES), + (MemoryUnit.valid_until.is_(None) | (MemoryUnit.valid_until > now)), + ] + if memory_type: + conditions.append(MemoryUnit.memory_type == memory_type) + if tag: + conditions.append(MemoryUnit.tag == tag) + statement = ( + select(MemoryUnit) + .where(*conditions) + .order_by(MemoryUnit.update_time.desc(), MemoryUnit.id.desc()) + ) + return list((await self.db.scalars(statement)).all()) + async def merge_evidence( self, memory: MemoryUnit, diff --git a/scripts/apply_memory_unit_upgrade.py b/scripts/apply_memory_unit_upgrade.py new file mode 100644 index 0000000..10f9b4a --- /dev/null +++ b/scripts/apply_memory_unit_upgrade.py @@ -0,0 +1,138 @@ +"""幂等执行 memory_unit 表结构升级(对齐 model/memory_unit.py)。 + +用法: + python scripts/apply_memory_unit_upgrade.py + +对应 SQL 版本见 sql/memory_unit_upgrade_20260913.sql。 +重复执行安全:ADD COLUMN 前检查 information_schema,MODIFY/UPDATE 本身幂等。 +""" +from __future__ import annotations + +import asyncio +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from sqlalchemy import text + +from config.database import mysql as mysql_db +from config.database.mysql import get_session_factory +from config.settings import settings + +TABLE = "memory_unit" + +ADD_COLUMNS = { + "session_id": "session_id VARCHAR(64) NULL COMMENT '产生记忆的会话ID'", + "agent_run_id": "agent_run_id VARCHAR(64) NULL COMMENT '产生记忆的Agent运行ID'", + "evidence_ref": "evidence_ref JSON NULL COMMENT '证据引用列表(会话/消息溯源)'", + "historical_accuracy": "historical_accuracy DECIMAL(5,2) NOT NULL DEFAULT 0.50 COMMENT '历史准确率'", + "confidence_version": "confidence_version VARCHAR(32) NULL COMMENT '置信度算法版本'", + "confidence_reason": "confidence_reason VARCHAR(255) NULL COMMENT '置信度评分原因'", + "confidence_update_time": "confidence_update_time DATETIME NULL COMMENT '置信度更新时间'", + "memory_version": "memory_version INT NOT NULL DEFAULT 1 COMMENT '记忆版本号'", + "last_verified_at": "last_verified_at DATETIME NULL COMMENT '最近验证时间'", + "milvus_id": "milvus_id VARCHAR(128) NULL COMMENT 'Milvus向量主键'", + "graph_node_id": "graph_node_id VARCHAR(128) NULL COMMENT 'Neo4j图谱节点ID'", + "milvus_sync_status": "milvus_sync_status VARCHAR(16) NOT NULL DEFAULT 'pending' COMMENT '向量同步状态'", + "neo4j_sync_status": "neo4j_sync_status VARCHAR(16) NOT NULL DEFAULT 'pending' COMMENT '图谱同步状态'", + "sync_retry_count": "sync_retry_count INT NOT NULL DEFAULT 0 COMMENT '同步重试次数'", + "last_sync_error": "last_sync_error VARCHAR(500) NULL COMMENT '最近同步错误'", + "next_retry_at": "next_retry_at DATETIME NULL COMMENT '下次重试时间'", + "last_synced_at": "last_synced_at DATETIME NULL COMMENT '最近成功同步时间'", +} + + +async def columns_of(session, database: str) -> dict[str, str]: + rows = ( + await session.execute( + text( + "SELECT COLUMN_NAME, DATA_TYPE FROM information_schema.columns " + "WHERE TABLE_SCHEMA = :d AND TABLE_NAME = :t" + ), + {"d": database, "t": TABLE}, + ) + ).mappings().all() + return {str(r["COLUMN_NAME"]): str(r["DATA_TYPE"]).lower() for r in rows} + + +async def main() -> None: + async with get_session_factory()() as session: + db = settings.mysql.database + cols = await columns_of(session, db) + + added = [] + for name, ddl in ADD_COLUMNS.items(): + if name in cols: + print(f"[skip] 列已存在: {name}") + continue + await session.execute(text(f"ALTER TABLE {TABLE} ADD COLUMN {ddl}")) + added.append(name) + print(f"[ok] 新增列: {name}") + await session.commit() + + cols = await columns_of(session, db) + + if cols.get("valid_from") == "date" or cols.get("valid_until") == "date": + await session.execute( + text( + f"ALTER TABLE {TABLE} " + "MODIFY COLUMN valid_from DATETIME NULL COMMENT '生效起始时间', " + "MODIFY COLUMN valid_until DATETIME NULL COMMENT '失效时间'" + ) + ) + await session.commit() + print("[ok] valid_from/valid_until: DATE -> DATETIME") + else: + print("[skip] valid_from/valid_until 已是 DATETIME") + + result = await session.execute( + text(f"UPDATE {TABLE} SET status = 'candidate' WHERE status = 'active'") + ) + await session.commit() + print(f"[ok] status active->candidate: {result.rowcount} 行") + + await session.execute( + text( + f"ALTER TABLE {TABLE} MODIFY COLUMN status VARCHAR(16) NOT NULL " + "DEFAULT 'candidate' COMMENT '记忆状态(candidate/confirmed/expired/rejected/archived)'" + ) + ) + await session.commit() + print("[ok] status 默认值改为 candidate") + + result = await session.execute( + text( + f"UPDATE {TABLE} SET memory_type = 'SERVICE_FACT' " + "WHERE memory_type IS NULL OR memory_type = 'FACT'" + ) + ) + await session.commit() + print(f"[ok] memory_type NULL/FACT -> SERVICE_FACT: {result.rowcount} 行") + + await session.execute( + text( + f"ALTER TABLE {TABLE} MODIFY COLUMN memory_type VARCHAR(32) NOT NULL " + "COMMENT '记忆业务类型'" + ) + ) + await session.commit() + print("[ok] memory_type 收紧为 NOT NULL") + + final = await columns_of(session, db) + missing = [name for name in ADD_COLUMNS if name not in final] + if missing: + print(f"[fail] 仍有缺失列: {missing}") + raise SystemExit(1) + null_types = ( + await session.execute( + text(f"SELECT COUNT(*) AS n FROM {TABLE} WHERE memory_type IS NULL") + ) + ).scalar() + print(f"[done] 迁移完成,共新增 {len(added)} 列;memory_type NULL 残留: {null_types}") + + await mysql_db.dispose() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/service/client_agent/runtime.py b/service/client_agent/runtime.py index 3449bf3..0d14f54 100644 --- a/service/client_agent/runtime.py +++ b/service/client_agent/runtime.py @@ -2,6 +2,7 @@ from __future__ import annotations +import asyncio import contextvars import logging import uuid @@ -72,13 +73,20 @@ class MemoryConversationContext: class MemoryAwareClientAgent: - """在现有客服 Agent 外包裹记忆召回、候选保存和降级处理。""" + """在现有客服 Agent 外包裹记忆召回、候选保存和降级处理。 - def __init__(self, *, agent, memory_service, context, extractor=None): + 候选记忆保存默认放入后台任务执行(background_saves=True),不阻塞 + 客服响应;测试或需要确定性顺序的场景可设为 False 改回同步执行。 + """ + + def __init__(self, *, agent, memory_service, context, extractor=None, + background_saves: bool = True): self.agent = agent self.memory_service = memory_service self.context = context self.extractor = extractor + self.background_saves = background_saves + self._pending_saves: set[asyncio.Task] = set() async def handle(self, session_id: str, query: str, *, trace_id: str, customer_id: int) -> dict: """执行记忆召回、客服回答、消息写入和候选记忆保存。""" @@ -107,7 +115,7 @@ class MemoryAwareClientAgent: result = await self.agent.handle( session_id, query, trace_id=trace_id, customer_id=customer_id ) - if self.extractor is not None: + if self.extractor is not None and not self.background_saves: await self._save_candidates( customer_id, session_id, @@ -117,6 +125,15 @@ class MemoryAwareClientAgent: trace_id=trace_id, ) result["memory_warnings"] = list(warnings) + if self.extractor is not None and self.background_saves: + self._spawn_save( + customer_id, + session_id, + query, + memory_context, + warnings, + trace_id=trace_id, + ) return result finally: _active_customer.reset(message_token) @@ -124,6 +141,35 @@ class MemoryAwareClientAgent: _active_warnings.reset(warnings_token) _active_memory_context.reset(context_token) + def _spawn_save( + self, + customer_id, + session_id, + query, + context, + warnings, + *, + trace_id: str | None = None, + ) -> None: + """把候选记忆保存放入后台任务;任务异常已自捕获,不击穿响应。""" + task = asyncio.create_task( + self._save_candidates( + customer_id, + session_id, + query, + context, + warnings, + trace_id=trace_id, + ) + ) + self._pending_saves.add(task) + task.add_done_callback(self._pending_saves.discard) + + async def wait_for_pending_saves(self) -> None: + """等待全部后台保存完成,供测试与优雅退出使用。""" + if self._pending_saves: + await asyncio.gather(*list(self._pending_saves), return_exceptions=True) + async def _save_candidates( self, customer_id, @@ -157,9 +203,24 @@ class MemoryAwareClientAgent: try: evidence_count = 1 if candidate.get("signal_type") == "interest_query": - # 兴趣主题首次出现即保存为候选,后续由长期记忆按精确内容合并证据。 + # 兴趣主题先计数:达到阈值才固化为长期记忆,避免单次关注污染画像。 candidate["memory_type"] = "CUSTOMER_PREFERENCE" candidate["source"] = "dialogue_inferred" + try: + _, reached = await self.memory_service.record_interest_signal( + customer_id=customer_id, tag=candidate["tag"] + ) + except Exception as exc: + logger.exception( + "client interest signal failed: trace_id=%s customer_id=%s tag=%s", + trace_id, + customer_id, + candidate.get("tag"), + ) + warnings.append(f"interest_signal_failed:{type(exc).__name__}") + continue + if not reached: + continue memory = MemoryUnitDTO( customer_id=customer_id, session_id=session_id, diff --git a/service/memory/facade.py b/service/memory/facade.py index 868558d..98f2112 100644 --- a/service/memory/facade.py +++ b/service/memory/facade.py @@ -97,7 +97,7 @@ class MemoryService: warnings.append(f"customer_product_recall_failed:{type(exc).__name__}") try: memories, memory_warnings = await self.long_term.recall( - db, customer_id, limit=max(limit, 10) + db, customer_id, limit=max(limit, 10), query=query ) warnings.extend(memory_warnings) except Exception as exc: diff --git a/service/memory/long_term.py b/service/memory/long_term.py index 5106a3c..6b1f90e 100644 --- a/service/memory/long_term.py +++ b/service/memory/long_term.py @@ -51,10 +51,10 @@ class LongTermMemoryService: ) entity = existing else: - # final_score 是召回重排阶段的临时分数,不属于 MySQL 主体字段。 + # final_score/semantic_similarity 是召回重排阶段的临时分数,不属于 MySQL 主体字段。 values = memory.model_dump( mode="json", - exclude={"id", "milvus_id", "graph_node_id", "final_score"}, + exclude={"id", "milvus_id", "graph_node_id", "final_score", "semantic_similarity"}, ) now = datetime.now() values["last_verified_at"] = values.get("last_verified_at") or now @@ -147,12 +147,50 @@ class LongTermMemoryService: memory_type: str | None = None, tag: str | None = None, limit: int = 100, + query: str | None = None, ) -> tuple[list[MemoryUnitDTO], list[str]]: - """按客户、类型、标签和有效期召回主体记忆。""" + """召回主体记忆;带 query 时叠加 Milvus 语义召回并标注相似度。""" entities = await self.repository_factory(db).list_for_customer( customer_id, memory_type=memory_type, tag=tag, limit=limit ) - return [self._to_dto(entity) for entity in entities], [] + dtos = [self._to_dto(entity) for entity in entities] + if query is None or not str(query).strip(): + return dtos, [] + + warnings: list[str] = [] + try: + vector = self.embedder(str(query)) + if isawaitable(vector): + vector = await vector + hits = await self.milvus_store.search(vector, customer_id, limit=max(limit, 1)) + except Exception as exc: + return dtos, [f"semantic_recall_failed:{type(exc).__name__}"] + + similarities: dict[int, float] = {} + for hit in hits: + raw_id = str(hit.get("memory_id", "")) + if raw_id.lstrip("-").isdigit(): + similarities[int(raw_id)] = float(hit.get("distance") or 0.0) + + matched: set[int] = set() + for dto in dtos: + if dto.id in similarities: + dto.semantic_similarity = similarities[dto.id] + matched.add(dto.id) + + missing_ids = [mid for mid in similarities if mid not in matched] + if missing_ids: + try: + extra_entities = await self.repository_factory(db).list_for_customer_by_ids( + customer_id, missing_ids, memory_type=memory_type, tag=tag + ) + for entity in extra_entities: + dto = self._to_dto(entity) + dto.semantic_similarity = similarities.get(dto.id) + dtos.append(dto) + except Exception as exc: + warnings.append(f"semantic_recall_fetch_failed:{type(exc).__name__}") + return dtos, warnings async def refresh_confidence( self, db, customer_id: int, *, limit: int = 500 diff --git a/service/memory/milvus_memory.py b/service/memory/milvus_memory.py index 40a041a..3016f19 100644 --- a/service/memory/milvus_memory.py +++ b/service/memory/milvus_memory.py @@ -2,6 +2,8 @@ from __future__ import annotations +from typing import Any + from pymilvus import AsyncMilvusClient, DataType from config.database.milvus import client as configured_client @@ -87,15 +89,43 @@ class MilvusMemoryStore: return memory_id async def search(self, vector: list[float], customer_id: int, *, limit: int = 10) -> list[dict]: - """按客户 ID 过滤向量查询结果。""" + """按客户 ID 过滤向量查询,返回归一化的命中列表。""" await self.ensure_collection() - return await self.client.search( + raw = await self.client.search( collection_name=self.collection_name, data=[vector], limit=limit, filter=f"customer_id == {int(customer_id)}", output_fields=["memory_id", "customer_id", "memory_type", "tag", "content", "status"], ) + return self._normalize_hits(raw) + + @staticmethod + def _normalize_hits(raw: Any) -> list[dict]: + """将 pymilvus 返回结构收敛为 memory_id + distance 的扁平列表。""" + hits: list[dict] = [] + for batch in raw or []: + for hit in batch or []: + if not isinstance(hit, dict): + continue + entity = hit.get("entity") or {} + memory_id = entity.get("memory_id") or hit.get("id") + if memory_id is None: + continue + try: + distance = float(hit.get("distance", 0.0)) + except (TypeError, ValueError): + distance = 0.0 + hits.append( + { + "memory_id": str(memory_id), + "distance": max(0.0, min(1.0, distance)), + "tag": entity.get("tag"), + "content": entity.get("content"), + "status": entity.get("status"), + } + ) + return hits async def delete(self, memory_id: int | str) -> None: """删除一条客户记忆向量。""" diff --git a/service/memory/schemas.py b/service/memory/schemas.py index 47a909a..d3f3be2 100644 --- a/service/memory/schemas.py +++ b/service/memory/schemas.py @@ -79,6 +79,7 @@ class MemoryUnitDTO(BaseModel): confidence_reason: str | None = Field(default=None, max_length=255) confidence_update_time: datetime | None = None final_score: float | None = Field(default=None, ge=0.0, le=1.0) + semantic_similarity: float | None = Field(default=None, ge=0.0, le=1.0) evidence_count: int = Field(default=0, ge=0) recall_count: int = Field(default=0, ge=0) status: MemoryStatus = MemoryStatus.CANDIDATE diff --git a/sql/memory_unit_upgrade_20260913.sql b/sql/memory_unit_upgrade_20260913.sql new file mode 100644 index 0000000..022a7b9 --- /dev/null +++ b/sql/memory_unit_upgrade_20260913.sql @@ -0,0 +1,45 @@ +-- memory_unit 表结构升级:对齐 model/memory_unit.py(34 列) +-- 日期:2026-09-13 执行方式:scripts/apply_memory_unit_upgrade.py(幂等)或本文件手工执行 +-- 背景:DB 为旧版 20 列结构,缺 session_id/evidence_ref/置信度/同步状态等 17 列, +-- 导致客服 Agent 长期记忆保存与召回抛 OperationalError(memory_warnings 来源)。 + +-- 1) 新增 17 列 +ALTER TABLE memory_unit + ADD COLUMN session_id VARCHAR(64) NULL COMMENT '产生记忆的会话ID', + ADD COLUMN agent_run_id VARCHAR(64) NULL COMMENT '产生记忆的Agent运行ID', + ADD COLUMN evidence_ref JSON NULL COMMENT '证据引用列表(会话/消息溯源)', + ADD COLUMN historical_accuracy DECIMAL(5,2) NOT NULL DEFAULT 0.50 COMMENT '历史准确率', + ADD COLUMN confidence_version VARCHAR(32) NULL COMMENT '置信度算法版本', + ADD COLUMN confidence_reason VARCHAR(255) NULL COMMENT '置信度评分原因', + ADD COLUMN confidence_update_time DATETIME NULL COMMENT '置信度更新时间', + ADD COLUMN memory_version INT NOT NULL DEFAULT 1 COMMENT '记忆版本号', + ADD COLUMN last_verified_at DATETIME NULL COMMENT '最近验证时间', + ADD COLUMN milvus_id VARCHAR(128) NULL COMMENT 'Milvus向量主键', + ADD COLUMN graph_node_id VARCHAR(128) NULL COMMENT 'Neo4j图谱节点ID', + ADD COLUMN milvus_sync_status VARCHAR(16) NOT NULL DEFAULT 'pending' COMMENT '向量同步状态', + ADD COLUMN neo4j_sync_status VARCHAR(16) NOT NULL DEFAULT 'pending' COMMENT '图谱同步状态', + ADD COLUMN sync_retry_count INT NOT NULL DEFAULT 0 COMMENT '同步重试次数', + ADD COLUMN last_sync_error VARCHAR(500) NULL COMMENT '最近同步错误', + ADD COLUMN next_retry_at DATETIME NULL COMMENT '下次重试时间', + ADD COLUMN last_synced_at DATETIME NULL COMMENT '最近成功同步时间'; + +-- 2) 类型对齐:DATE 无法承载时间语义,扩为 DATETIME +ALTER TABLE memory_unit + MODIFY COLUMN valid_from DATETIME NULL COMMENT '生效起始时间', + MODIFY COLUMN valid_until DATETIME NULL COMMENT '失效时间'; + +-- 3) 状态语义迁移:旧默认 active → 新枚举 candidate,并更新表默认值 +UPDATE memory_unit SET status = 'candidate' WHERE status = 'active'; +ALTER TABLE memory_unit + MODIFY COLUMN status VARCHAR(16) NOT NULL DEFAULT 'candidate' COMMENT '记忆状态(candidate/confirmed/expired/rejected/archived)'; + +-- 4) memory_type 旧枚举归并:NULL 与 'FACT' → 'SERVICE_FACT',再收紧 NOT NULL +UPDATE memory_unit SET memory_type = 'SERVICE_FACT' + WHERE memory_type IS NULL OR memory_type = 'FACT'; +ALTER TABLE memory_unit + MODIFY COLUMN memory_type VARCHAR(32) NOT NULL COMMENT '记忆业务类型'; + +-- 说明: +-- - id/customer_id 为 BIGINT UNSIGNED,与 ORM BigInteger 兼容,不做变更; +-- - 遗留列 conflict_count/dimension/polarity 无 ORM 映射,保留不动; +-- - NUMERIC(5,2) 在 MySQL 中即 DECIMAL(5,2),探针报告的差异为别名,非真实差异。