Merge pull request 'feat:客户agent以及记忆模块优化' (#14) from develop_feature_memory into develop

Reviewed-on: #14
This commit was merged in pull request #14.
This commit is contained in:
2026-09-12 20:08:18 +08:00
10 changed files with 86 additions and 54 deletions
+1 -7
View File
@@ -76,10 +76,4 @@ data/output/
DK2_客服Agent模块完整开发计划(v1.1).md
DK4_客服Agent需求文档(修复完善版v1.4).md
客服Agent模块 · TODO List(最终修复版v1.2).md
三层记忆模块需求规格说明书.md
客服Agent记忆模块TODO清单.md
客服Agent记忆模块开发计划.md
开发计划_用户画像置信度实时更新.md
客服Agent记忆模块接入说明.md
客服Agent接口清单.md
空闲会话自动归档开发计划.md
客服Agent长期记忆与置信度最终开发文档.md
-2
View File
@@ -35,7 +35,6 @@ class MemoryUnit(Base):
confidence_reason: Mapped[str | None] = mapped_column(String(255))
confidence_update_time: Mapped[datetime | None] = mapped_column(DateTime)
evidence_count: Mapped[int] = mapped_column(Integer, server_default="0")
conflict_count: Mapped[int] = mapped_column(Integer, server_default="0")
recall_count: Mapped[int] = mapped_column(Integer, server_default="0")
memory_version: Mapped[int] = mapped_column(Integer, server_default="1")
update_time: Mapped[datetime | None] = mapped_column(DateTime)
@@ -53,4 +52,3 @@ class MemoryUnit(Base):
last_sync_error: Mapped[str | None] = mapped_column(String(500))
next_retry_at: Mapped[datetime | None] = mapped_column(DateTime)
last_synced_at: Mapped[datetime | None] = mapped_column(DateTime)
+41 -6
View File
@@ -28,15 +28,14 @@ class MemoryUnitRepo(BaseRepository):
return obj
async def find_exact(
self, customer_id: int, memory_type: str, tag: str, content: str
self, customer_id: int, memory_type: str, tag: str
) -> MemoryUnit | None:
"""按客户、类型、标签和内容精确查找去重对象。"""
"""按客户、类型和标签查找当前有效记忆。"""
return await self.db.scalar(
select(MemoryUnit).where(
MemoryUnit.customer_id == customer_id,
MemoryUnit.memory_type == memory_type,
MemoryUnit.tag == tag,
MemoryUnit.content == content,
MemoryUnit.status.not_in(INACTIVE_STATUSES),
)
)
@@ -70,16 +69,52 @@ class MemoryUnitRepo(BaseRepository):
return list((await self.db.scalars(statement)).all())
async def merge_evidence(
self, memory: MemoryUnit, *, evidence_count: int = 1, conflict_count: int = 0
self,
memory: MemoryUnit,
*,
content: str | None = None,
source: str | None = None,
evidence_count: int = 1,
evidence_ref: list[dict] | None = None,
) -> MemoryUnit:
"""合并证据和冲突计数,不改变客户隔离范围。"""
"""更新最新内容并合并证据,不改变客户隔离范围。"""
if content:
memory.content = content
if source:
memory.source = source
memory.evidence_count = (memory.evidence_count or 0) + evidence_count
memory.conflict_count = (memory.conflict_count or 0) + conflict_count
if evidence_ref:
memory.evidence_ref = [*(memory.evidence_ref or []), *evidence_ref]
memory.last_verified_at = datetime.now()
memory.update_time = datetime.now()
await self.db.commit()
await self.db.refresh(memory)
return memory
async def list_for_confidence_refresh(
self, customer_id: int, limit: int = 500
) -> list[MemoryUnit]:
"""返回需要重新计算时间衰减置信度的有效记忆。"""
statement = (
select(MemoryUnit)
.where(
MemoryUnit.customer_id == customer_id,
MemoryUnit.status.in_(ACTIVE_STATUSES),
)
.order_by(MemoryUnit.id)
.limit(limit)
)
return list((await self.db.scalars(statement)).all())
async def update_confidence(self, memory_id: int, **values) -> None:
"""只更新置信度字段,不改变记忆内容和证据。"""
memory = await self.db.get(MemoryUnit, memory_id)
if memory is None:
return
for key, value in values.items():
setattr(memory, key, value)
await self.db.commit()
async def update_sync_status(self, memory_id: int, **values) -> None:
"""更新向量/图谱索引 ID 和同步状态。"""
memory = await self.db.get(MemoryUnit, memory_id)
+3 -6
View File
@@ -140,13 +140,9 @@ class MemoryAwareClientAgent:
for index, candidate in enumerate(candidates):
try:
evidence_count = 1
if candidate.get("signal_type") == "interest_query":
count, reached = await self.memory_service.record_interest_signal(
customer_id=customer_id,
tag=candidate["tag"],
)
if not reached:
continue
# 兴趣主题首次出现即保存为候选,后续由长期记忆按精确内容合并证据。
candidate["memory_type"] = "CUSTOMER_PREFERENCE"
candidate["source"] = "dialogue_inferred"
memory = MemoryUnitDTO(
@@ -157,6 +153,7 @@ class MemoryAwareClientAgent:
content=candidate["content"],
source=candidate["source"],
evidence_ref=[{"session_id": session_id, "query": query}],
evidence_count=evidence_count,
)
_, save_warnings = await self.memory_service.save_memory(
customer_id=customer_id,
+8
View File
@@ -70,6 +70,14 @@ class MemoryService:
warnings.extend(self.short_term.last_warnings)
async with self._db() as db:
try:
refresh_confidence = getattr(self.long_term, "refresh_confidence", None)
if refresh_confidence is not None:
await refresh_confidence(
db, customer_id, limit=max(limit, 10)
)
except Exception as exc:
warnings.append(f"confidence_refresh_failed:{type(exc).__name__}")
profile, profile_warnings = await self.profile.get(db, customer_id)
warnings.extend(profile_warnings)
try:
+23 -4
View File
@@ -39,11 +39,16 @@ class LongTermMemoryService:
memory.customer_id,
memory.memory_type.value,
memory.tag,
memory.content,
)
warnings: list[str] = []
if existing is not None:
existing = await repo.merge_evidence(existing)
existing = await repo.merge_evidence(
existing,
content=memory.content,
source=memory.source.value,
evidence_count=max(memory.evidence_count or 1, 1),
evidence_ref=memory.evidence_ref,
)
entity = existing
else:
# final_score 是召回重排阶段的临时分数,不属于 MySQL 主体字段。
@@ -51,6 +56,9 @@ class LongTermMemoryService:
mode="json",
exclude={"id", "milvus_id", "graph_node_id", "final_score"},
)
now = datetime.now()
values["last_verified_at"] = values.get("last_verified_at") or now
values["update_time"] = values.get("update_time") or now
values["memory_type"] = memory.memory_type.value
values["source"] = memory.source.value
confidence_result, confidence_warning = self._calculate_confidence(memory)
@@ -116,7 +124,6 @@ class LongTermMemoryService:
tag=memory.tag,
source=source,
evidence_count=memory.evidence_count or 0,
conflict_count=memory.conflict_count or 0,
age_days=age_days,
memory_type=memory_type,
)
@@ -147,6 +154,19 @@ class LongTermMemoryService:
)
return [self._to_dto(entity) for entity in entities], []
async def refresh_confidence(
self, db, customer_id: int, *, limit: int = 500
) -> int:
"""按当前时间刷新有效记忆的置信度并持久化结果。"""
repo = self.repository_factory(db)
entities = await repo.list_for_confidence_refresh(customer_id, limit=limit)
refreshed = 0
for entity in entities:
confidence_result, _ = self._calculate_confidence(entity)
await repo.update_confidence(entity.id, **confidence_result)
refreshed += 1
return refreshed
async def retry_pending(self, db, *, limit: int = 100) -> dict[str, int]:
"""重试 MySQL 中缺少外部索引或同步失败的记忆。"""
entities = await self.repository_factory(db).list_pending_sync(limit)
@@ -182,7 +202,6 @@ class LongTermMemoryService:
"confidence_reason": getattr(entity, "confidence_reason", None),
"confidence_update_time": getattr(entity, "confidence_update_time", None),
"evidence_count": entity.evidence_count or 0,
"conflict_count": entity.conflict_count or 0,
"recall_count": entity.recall_count or 0,
"status": entity.status,
"valid_from": entity.valid_from,
-1
View File
@@ -80,7 +80,6 @@ class MemoryUnitDTO(BaseModel):
confidence_update_time: datetime | None = None
final_score: float | None = Field(default=None, ge=0.0, le=1.0)
evidence_count: int = Field(default=0, ge=0)
conflict_count: int = Field(default=0, ge=0)
recall_count: int = Field(default=0, ge=0)
status: MemoryStatus = MemoryStatus.CANDIDATE
valid_from: datetime | None = None
-1
View File
@@ -444,7 +444,6 @@ CREATE TABLE IF NOT EXISTS memory_unit (
confidence_reason VARCHAR(255) NULL COMMENT '本次置信度评估原因',
confidence_update_time DATETIME NULL COMMENT '置信度最近计算时间',
evidence_count INT NOT NULL DEFAULT 0 COMMENT '证据数',
conflict_count INT NOT NULL DEFAULT 0 COMMENT '冲突数',
recall_count INT NOT NULL DEFAULT 0 COMMENT '召回数',
memory_version INT NOT NULL DEFAULT 1 COMMENT '记忆版本号',
update_time DATETIME NULL,
+9 -15
View File
@@ -6,7 +6,7 @@ from typing import Any
class BaseConfidenceCalcTool:
"""根据来源、证据、冲突和时间计算单条记忆的长期置信度。"""
"""根据来源、证据和时间计算单条客服记忆的长期置信度。"""
SOURCE_INITIAL = {
"dialogue_confirmed": 0.75,
@@ -20,23 +20,21 @@ class BaseConfidenceCalcTool:
"SERVICE_FACT": 0.75,
}
DEFAULT_THRESHOLD = 0.80
VERSION = "confidence-v1"
VERSION = "confidence-v2"
def calc(
self,
tag: str,
source: str,
evidence_count: int,
conflict_count: int,
age_days: int,
) -> float:
"""计算基础置信度分数,返回范围为 0 到 1 的浮点数。"""
self._validate(tag, source, evidence_count, conflict_count, age_days)
self._validate(tag, source, evidence_count, age_days)
base = self.SOURCE_INITIAL[source]
gain = min(evidence_count * 0.05, 0.30)
penalty = min(conflict_count * 0.10, 0.50)
decay = max(0.80, 1 - age_days / 365 * 0.20)
return max(0.0, min(1.0, (base + gain - penalty) * decay))
return max(0.0, min(1.0, (base + gain) * decay))
def evaluate(
self,
@@ -44,13 +42,12 @@ class BaseConfidenceCalcTool:
tag: str,
source: str,
evidence_count: int,
conflict_count: int,
age_days: int,
memory_type: str | None = None,
threshold: float | None = None,
) -> dict[str, Any]:
"""返回可供记忆模块保存的完整置信度评估结果。"""
score = self.calc(tag, source, evidence_count, conflict_count, age_days)
score = self.calc(tag, source, evidence_count, age_days)
if threshold is None:
threshold = self.MEMORY_THRESHOLDS.get(
memory_type or "", self.DEFAULT_THRESHOLD
@@ -64,11 +61,10 @@ class BaseConfidenceCalcTool:
"confidence": score,
"status": status,
"evidence_count": evidence_count,
"conflict_count": conflict_count,
"age_days": age_days,
"threshold": threshold,
"confidence_reason": self._reason(
source, evidence_count, conflict_count, age_days
source, evidence_count, age_days
),
"confidence_version": self.VERSION,
}
@@ -87,7 +83,6 @@ class BaseConfidenceCalcTool:
tag: str,
source: str,
evidence_count: int,
conflict_count: int,
age_days: int,
) -> None:
"""校验工具输入,避免非法计数污染记忆分数。"""
@@ -97,18 +92,17 @@ class BaseConfidenceCalcTool:
raise ValueError(f"不支持的客服对话来源: {source}")
for name, value in (
("evidence_count", evidence_count),
("conflict_count", conflict_count),
("age_days", age_days),
):
if not isinstance(value, int) or isinstance(value, bool) or value < 0:
raise ValueError(f"{name} 必须是非负整数")
@staticmethod
def _reason(source: str, evidence_count: int, conflict_count: int, age_days: int) -> str:
def _reason(source: str, evidence_count: int, age_days: int) -> str:
"""生成便于审计和排查的评分原因。"""
return (
f"来源={source}; 支持证据={evidence_count}; 冲突证据={conflict_count}; "
f"存在天数={age_days}; 采用证据增益、冲突惩罚和时间衰减"
f"来源={source}; 支持证据={evidence_count}; "
f"存在天数={age_days}; 采用证据增益和时间衰减"
)
+1 -12
View File
@@ -14,8 +14,7 @@ class FinalConfidenceRankTool:
"semantic": 0.30,
"timeliness": 0.20,
"accuracy": 0.20,
"base": 0.25,
"conflict": 0.05,
"base": 0.30,
}
INVALID_STATUSES = frozenset({"expired", "rejected", "archived"})
@@ -39,13 +38,11 @@ class FinalConfidenceRankTool:
timeliness = self._calc_timeliness(unit.get("age_days", 0))
accuracy = self._bounded(unit.get("historical_accuracy", 0.5), 0.5)
base = self._bounded(unit.get("confidence", 0.5), 0.5)
conflict_penalty = min(self._non_negative_int(unit.get("conflict_count", 0)), 5) * 0.1
final_score = (
self.WEIGHTS["semantic"] * semantic
+ self.WEIGHTS["timeliness"] * timeliness
+ self.WEIGHTS["accuracy"] * accuracy
+ self.WEIGHTS["base"] * base
- self.WEIGHTS["conflict"] * conflict_penalty
)
unit["final_score"] = max(0.0, min(1.0, final_score))
ranked.append(unit)
@@ -78,14 +75,6 @@ class FinalConfidenceRankTool:
return default
return max(0.0, min(1.0, value))
@staticmethod
def _non_negative_int(value: Any) -> int:
"""将冲突次数转换为非负整数。"""
try:
return max(0, int(value))
except (TypeError, ValueError):
return 0
@staticmethod
def _calc_timeliness(age_days: Any) -> float:
"""按每年 20% 计算平滑时效分,最低保留 0.8。"""