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:
+1
-7
@@ -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
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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。"""
|
||||
|
||||
Reference in New Issue
Block a user