25 changed files with 741 additions and 90 deletions
+1 -5
View File
@@ -76,8 +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
+53 -2
View File
@@ -7,6 +7,7 @@ import uuid
from inspect import isawaitable
from agent.customer_agent.session import SessionOwnershipError
from service.client_agent.archive_schedule import SessionArchiveSchedule
async def _config(config_getter, key: str, default):
@@ -23,16 +24,41 @@ class ClientSessionService:
self.redis = redis
self.config_getter = config_getter
self.clock = clock
self.archive_schedule = SessionArchiveSchedule(redis)
async def _archive_grace(self) -> int:
"""读取 Redis 消息归档保护缓冲时间,默认 60 秒。"""
return await _config(
self.config_getter, "agent.customer.session.archive_grace", 60
)
async def create_session(self, customer_id: int) -> str:
"""Create a session whose Redis owner value is the current customer."""
session_id = uuid.uuid4().hex
now = self.clock()
ttl = await _config(self.config_getter, "agent.customer.session.ttl", 1800)
await self.redis.set(f"session:{session_id}", str(customer_id), ex=ttl)
max_lifetime = await _config(
self.config_getter, "agent.customer.session.max_lifetime", 86400
)
grace = await self._archive_grace()
effective_ttl = min(ttl, max_lifetime)
await self.redis.set(
f"session:{session_id}", str(customer_id), ex=effective_ttl + grace
)
await self.redis.rpush(f"session:{session_id}:messages", "")
await self.redis.expire(f"session:{session_id}:messages", ttl)
await self.redis.expire(
f"session:{session_id}:messages", effective_ttl + grace
)
await self.archive_schedule.register(
session_id=session_id,
created_at=now,
last_activity_at=now,
archive_due_at=now + ttl,
absolute_expire_at=now + max_lifetime,
)
return session_id
async def verify_session_ownership(self, session_id: str, *, customer_id: int) -> None:
"""Reject missing sessions and sessions owned by another customer."""
owner = await self.redis.get(f"session:{session_id}")
@@ -43,6 +69,31 @@ class ClientSessionService:
if str(owner) != str(customer_id):
raise SessionOwnershipError(403, "无权访问该会话")
async def touch_session(self, session_id: str, *, customer_id: int) -> None:
"""刷新所有权 TTL 和自动归档截止时间,但不突破最长生命周期。"""
await self.verify_session_ownership(session_id, customer_id=customer_id)
now = self.clock()
meta = await self.redis.hgetall(
self.archive_schedule.meta_key(session_id)
)
absolute_expire_at = float(meta.get("absolute_expire_at", now))
remaining = int(absolute_expire_at - now)
if remaining <= 0:
raise SessionOwnershipError(404, "会话不存在或已过期")
ttl = await _config(self.config_getter, "agent.customer.session.ttl", 1800)
effective_ttl = min(ttl, remaining)
grace = await self._archive_grace()
await self.redis.expire(f"session:{session_id}", effective_ttl + grace)
await self.archive_schedule.touch(
session_id=session_id,
last_activity_at=now,
archive_due_at=now + effective_ttl,
)
async def forget_session(self, session_id: str) -> None:
"""移除会话的自动归档调度记录。"""
await self.archive_schedule.remove(session_id)
async def consume_chat_quota(self, session_id: str, *, customer_id: int) -> int | None:
"""Apply the existing per-session rate limit with a customer-scoped key."""
window = await _config(
+2 -2
View File
@@ -59,8 +59,8 @@ class AnonymousSessionService:
await self.redis.zremrangebyscore(key, 0, (now - window) * 1000)
count = await self.redis.zcard(key)
if count >= maximum:
entries = self.redis.sorted_sets.get(key, [])
oldest = min((score for score, _ in entries), default=now * 1000)
entries = await self.redis.zrange(key, 0, 0, withscores=True)
oldest = entries[0][1] if entries else now * 1000
return max(1, int((oldest / 1000 + window) - now))
await self.redis.zadd(key, {uuid.uuid4().hex: now * 1000})
await self.redis.expire(key, window)
+2
View File
@@ -65,6 +65,7 @@ async def chat(
headers={"Retry-After": str(retry_after)},
content=response.model_dump(),
)
await runtime.session_service.touch_session(session_id, customer_id=user.id)
result = await runtime.agent.handle(
session_id,
query,
@@ -114,6 +115,7 @@ async def end_session(
)
if warnings:
return success({"session_id": session_id, "archived": False, "warnings": warnings})
await runtime.session_service.forget_session(session_id)
if hasattr(runtime.redis, "delete"):
await runtime.redis.delete(
f"session:{session_id}",
+15
View File
@@ -1,5 +1,6 @@
"""应用入口:四库异步生命周期 + 中间件 + 全局异常 + 路由装配。"""
from contextlib import asynccontextmanager
import asyncio
import uvicorn
from fastapi import FastAPI
@@ -12,6 +13,7 @@ from service.customer_agent.bootstrap import (
build_default_runtime,
)
from service.client_agent.bootstrap import build_default_runtime as build_client_runtime
from service.client_agent.idle_archive_worker import IdleArchiveWorker
from utils.exceptions import register_exception_handlers
from utils.logger import setup_logging
from utils.request_id import RequestIdMiddleware
@@ -23,8 +25,18 @@ async def lifespan(app: FastAPI):
await ensure_collections()
app.state.customer_agent_runtime = build_default_runtime()
app.state.client_agent_runtime = build_client_runtime()
app.state.client_agent_archive_worker = IdleArchiveWorker(
redis=app.state.client_agent_runtime.redis,
memory_service=app.state.client_agent_runtime.memory_service,
)
await app.state.client_agent_archive_worker.recover_due_index()
app.state.client_agent_archive_task = asyncio.create_task(
app.state.client_agent_archive_worker.run()
)
app.state.knowledge_upload_service = build_default_knowledge_upload_service()
yield
await app.state.client_agent_archive_worker.stop()
await app.state.client_agent_archive_task
await database.dispose()
@@ -38,3 +50,6 @@ app.include_router(api_router)
@app.get("/")
async def root():
return {"message": "智能公募基金系统 API", "docs": "/docs"}
if __name__ == '__main__':
uvicorn.run(app, host="127.0.0.1", port=8000)
-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)
+36 -17
View File
@@ -2,6 +2,7 @@
from __future__ import annotations
import logging
import re
from enum import StrEnum
@@ -12,32 +13,51 @@ class Intent(StrEnum):
GUIDE_PURCHASE = "guide_purchase"
WANT_ADVISOR = "want_advisor"
NL2SQL_REQUEST = "nl2sql_request"
COMPLAIN = "complain"
KNOWLEDGE_QA = "knowledge_qa"
COMPANY_INFO = "company_info"
CHITCHAT = "chitchat"
OFF_TOPIC = "off_topic"
NO_MATCH = "no_match"
INTENT_VALUES = frozenset(item.value for item in Intent)
# 按长度降序匹配,避免短标签误命中长标签的一部分
_INTENT_PATTERN = re.compile(
"|".join(re.escape(value) for value in sorted(INTENT_VALUES, key=len, reverse=True))
)
INTENT_SYSTEM_PROMPT = (
"你是华夏科技(一家基金代销金融机构)智能客服的意图分类器。"
"请根据用户输入,从以下选项中选择最匹配的意图,并仅输出对应的英文标签(不输出任何其他内容):\n"
"- guide_purchase: 用户询问如何购买基金、开户、注册等引导类问题\n"
"- want_advisor: 用户希望获得个性化基金推荐或投资顾问服务\n"
"- knowledge_qa: 用户询问基金相关的知识性问题,如净值、费率、风险、申赎规则等\n"
"- company_info: 用户询问华夏科技公司本身的信息,如公司全称、成立时间、牌照、总部地址、"
"客服电话、服务时间、官网、投诉渠道等\n"
"- nl2sql_request: 用户要求查询具体数据或账户信息\n"
"- chitchat: 普通寒暄,如问候、致谢、告别、询问你是谁/你能做什么、在吗等一两句话的闲聊\n"
"- off_topic: 用户要求你实质性地处理与金融、基金、公司业务无关的事情,"
"如写代码、讲笑话、写作文、问天气、聊政治、情感咨询、做数学题等\n"
"- no_match: 无法归入以上任何一类\n"
"注意:寒暄性质的一两句话归为 chitchat;一旦用户提出金融之外的实质性请求,归为 off_topic。\n"
"仅输出一个小写英文标签,不要输出解释、标点或换行。"
)
def parse_intent(raw: str | None) -> Intent:
"""从模型原始输出中提取第一个合法标签,提取不到则返回 NO_MATCH。"""
if not raw:
return Intent.NO_MATCH
match = _INTENT_PATTERN.search(raw.lower())
return Intent(match.group(0)) if match else Intent.NO_MATCH
async def intent_recognize(query: str, *, llm_client) -> Intent:
if not query or not query.strip():
return Intent.NO_MATCH
messages = [
{
"role": "system",
"content": (
"你是一个意图分类器。请根据用户输入,从以下选项中选择最匹配的意图,"
"并仅输出对应的英文标签(不输出任何其他内容):\n"
"- guide_purchase: 用户询问如何购买基金、开户、注册等引导类问题\n"
"- want_advisor: 用户希望获得个性化基金推荐或投资顾问服务\n"
"- knowledge_qa: 用户询问基金相关的知识性问题,如净值、费率、风险等\n"
"- nl2sql_request: 用户要求查询具体数据或账户信息\n"
"- complain: 用户表达不满或投诉\n"
"- no_match: 以上都不匹配\n"
"仅输出一个小写英文标签,不要输出解释、标点或换行。"
),
},
{"role": "system", "content": INTENT_SYSTEM_PROMPT},
{"role": "user", "content": query},
]
try:
@@ -45,5 +65,4 @@ async def intent_recognize(query: str, *, llm_client) -> Intent:
except Exception:
logger.exception("intent recognition failed")
return Intent.NO_MATCH
value = raw.strip().strip('`').lower()
return Intent(value) if value in INTENT_VALUES else Intent.NO_MATCH
return parse_intent(raw)
+6 -3
View File
@@ -90,9 +90,12 @@ def _format_candidates(candidates) -> list[dict]:
distance = hit.get("distance")
raw_score = hit.get("score")
if distance is not None:
score = 1.0 - distance
# All project Milvus collections use COSINE. For COSINE,
# Milvus returns a similarity score: larger means more relevant.
# Do not invert it as if it were an L2 distance.
score = float(distance)
else:
score = raw_score
score = float(raw_score) if raw_score is not None else None
if score is None or score < threshold:
continue
sources.append(
@@ -160,4 +163,4 @@ async def retrieve_with_status(
except Exception:
logger.exception("Milvus retrieval failed")
return {"status": "milvus_unavailable", "sources": []}
return {"status": "ok", "sources": _format_candidates(candidates)}
return {"status": "ok", "sources": _format_candidates(candidates)}
+6 -1
View File
@@ -31,6 +31,12 @@ class ConversationArchiveRepo(BaseRepository):
if not rows:
return 0
session_ids = {row.get("session_id") for row in rows}
if len(session_ids) != 1 or None in session_ids:
raise ValueError("archive_batch 只能接收同一会话的完整消息")
if any(not row.get("message_id") for row in rows):
raise ValueError("archive_batch 的 message_id 不能为空")
session_id = rows[0]["session_id"]
message_ids = [row["message_id"] for row in rows]
existing = await self.db.scalars(
@@ -65,4 +71,3 @@ class ConversationArchiveRepo(BaseRepository):
__all__ = ["ConversationArchiveRepo"]
+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)
+53
View File
@@ -0,0 +1,53 @@
"""Redis 会话归档分布式锁。"""
from __future__ import annotations
import uuid
class SessionArchiveLock:
"""保证同一会话同一时刻只被一个归档任务处理。"""
KEY = "session:archive:lock:{session_id}"
DEFAULT_TTL = 120
RELEASE_SCRIPT = """
if redis.call('get', KEYS[1]) == ARGV[1] then
return redis.call('del', KEYS[1])
end
return 0
"""
def __init__(self, redis, *, default_ttl: int = DEFAULT_TTL):
self.redis = redis
self.default_ttl = default_ttl
@classmethod
def key(cls, session_id: str) -> str:
"""生成会话级锁 Key。"""
return cls.KEY.format(session_id=session_id)
async def acquire(self, session_id: str, *, ttl: int | None = None) -> str | None:
"""尝试获取锁,成功返回随机 token,失败返回 None。"""
lock_ttl = self.default_ttl if ttl is None else ttl
if lock_ttl <= 0:
raise ValueError("锁 TTL 必须为正整数")
token = uuid.uuid4().hex
acquired = await self.redis.set(
self.key(session_id), token, nx=True, ex=lock_ttl
)
return token if acquired else None
async def release(self, session_id: str, token: str) -> bool:
"""仅当 token 属于当前持有者时释放锁。"""
if not token:
return False
result = await self.redis.eval(
self.RELEASE_SCRIPT,
1,
self.key(session_id),
token,
)
return bool(result)
__all__ = ["SessionArchiveLock"]
+132
View File
@@ -0,0 +1,132 @@
"""Redis client_agent 会话自动归档调度数据结构。"""
from __future__ import annotations
class SessionArchiveSchedule:
"""管理空闲会话的归档截止时间,不负责实际归档。"""
DUE_KEY = "session:archive:due"
META_KEY = "session:{session_id}:meta"
def __init__(self, redis):
self.redis = redis
@classmethod
def meta_key(cls, session_id: str) -> str:
"""生成与短期记忆共用的会话元数据 Key。"""
return cls.META_KEY.format(session_id=session_id)
async def register(
self,
*,
session_id: str,
created_at: float,
last_activity_at: float,
archive_due_at: float,
absolute_expire_at: float,
) -> None:
"""登记新会话的归档时间和生命周期信息。"""
await self.redis.hset(
self.meta_key(session_id),
mapping={
"created_at": str(created_at),
"last_activity_at": str(last_activity_at),
"archive_due_at": str(archive_due_at),
"absolute_expire_at": str(absolute_expire_at),
"archive_status": "pending",
"archive_retry_count": "0",
},
)
await self.redis.expire(
self.meta_key(session_id),
max(1, int(absolute_expire_at - created_at)),
)
await self.redis.zadd(self.DUE_KEY, {session_id: archive_due_at})
async def touch(
self,
*,
session_id: str,
last_activity_at: float,
archive_due_at: float,
) -> None:
"""更新会话最近活动时间和下一次归档截止时间。"""
await self.redis.hset(
self.meta_key(session_id),
mapping={
"last_activity_at": str(last_activity_at),
"archive_due_at": str(archive_due_at),
"archive_status": "pending",
},
)
await self.redis.zadd(self.DUE_KEY, {session_id: archive_due_at})
async def list_due(self, *, now: float, limit: int = 100) -> list[str]:
"""返回归档截止时间已到的会话。"""
if limit <= 0:
return []
values = await self.redis.zrangebyscore(
self.DUE_KEY, 0, now, start=0, num=limit
)
return [value.decode() if isinstance(value, bytes) else str(value) for value in values]
async def get_meta(self, session_id: str) -> dict:
"""读取会话归档元数据。"""
return await self.redis.hgetall(self.meta_key(session_id))
async def reschedule(self, session_id: str, *, due_at: float) -> None:
"""更新下一次扫描时间,保留当前调度记录。"""
await self.redis.zadd(self.DUE_KEY, {session_id: due_at})
async def rebuild_due_index(self, *, limit: int = 10000) -> int:
"""从会话元数据重建丢失的到期索引。"""
rebuilt = 0
async for raw_key in self.redis.scan_iter(match="session:*:meta", count=100):
key = raw_key.decode() if isinstance(raw_key, bytes) else raw_key
prefix = "session:"
suffix = ":meta"
if not key.startswith(prefix) or not key.endswith(suffix):
continue
session_id = key[len(prefix) : -len(suffix)]
meta = await self.redis.hgetall(key)
due_at = meta.get("archive_due_at")
if due_at is None or meta.get("archive_status") == "success":
continue
await self.redis.zadd(self.DUE_KEY, {session_id: float(due_at)})
rebuilt += 1
if rebuilt >= limit:
break
return rebuilt
async def mark_processing(self, session_id: str) -> None:
"""标记会话进入归档处理状态。"""
await self.redis.hset(
self.meta_key(session_id), mapping={"archive_status": "processing"}
)
async def mark_retry(
self,
session_id: str,
*,
retry_count: int,
error: str,
next_retry_at: float,
) -> None:
"""记录归档失败信息和下一次重试时间。"""
await self.redis.hset(
self.meta_key(session_id),
mapping={
"archive_status": "retry",
"archive_retry_count": str(retry_count),
"last_archive_error": error[:500],
"next_retry_at": str(next_retry_at),
},
)
async def remove(self, session_id: str) -> None:
"""移除归档调度记录,不删除会话消息或短期记忆元数据。"""
await self.redis.zrem(self.DUE_KEY, session_id)
__all__ = ["SessionArchiveSchedule"]
+144
View File
@@ -0,0 +1,144 @@
"""client_agent 空闲会话自动归档后台任务。"""
from __future__ import annotations
import asyncio
import logging
import time
from .archive_lock import SessionArchiveLock
from .archive_schedule import SessionArchiveSchedule
logger = logging.getLogger("client_agent.idle_archive_worker")
class IdleArchiveWorker:
"""扫描到期会话并调用统一 MemoryService 完成归档。"""
RETRY_BASE_DELAY = 30
RETRY_MAX_DELAY = 15 * 60
def __init__(
self,
*,
redis,
memory_service,
interval: int = 10,
batch_size: int = 100,
lock_ttl: int = 120,
clock=time.time,
):
self.redis = redis
self.memory_service = memory_service
self.interval = interval
self.batch_size = batch_size
self.clock = clock
self.schedule = SessionArchiveSchedule(redis)
self.lock = SessionArchiveLock(redis, default_ttl=lock_ttl)
self._stop_event = asyncio.Event()
async def archive_due_sessions(self) -> int:
"""处理一批到期会话,返回成功归档数量。"""
session_ids = await self.schedule.list_due(
now=self.clock(), limit=self.batch_size
)
archived_count = 0
for session_id in session_ids:
token = await self.lock.acquire(session_id)
if token is None:
continue
try:
if await self._archive_one(session_id):
archived_count += 1
finally:
await self.lock.release(session_id, token)
return archived_count
async def recover_due_index(self) -> int:
"""应用启动时从会话元数据恢复归档调度索引。"""
try:
return await self.schedule.rebuild_due_index()
except Exception:
logger.exception("failed to rebuild idle archive due index")
return 0
async def _archive_one(self, session_id: str) -> bool:
"""归档单个会话;失败时保留消息并推迟下一次扫描。"""
owner = await self.redis.get(f"session:{session_id}")
if owner is None:
await self.schedule.remove(session_id)
return False
owner = owner.decode() if isinstance(owner, bytes) else owner
meta = await self.schedule.get_meta(session_id)
retry_count = int(meta.get("archive_retry_count", 0) or 0)
await self.schedule.mark_processing(session_id)
try:
warnings = await self.memory_service.close_session(
customer_id=int(owner),
session_id=session_id,
)
if warnings:
await self._schedule_retry(
session_id,
retry_count=retry_count + 1,
error=";".join(warnings),
)
return False
await self.redis.delete(f"session:{session_id}")
await self.schedule.remove(session_id)
return True
except Exception as exc:
logger.exception("idle session archive failed: session_id=%s", session_id)
await self._schedule_retry(
session_id,
retry_count=retry_count + 1,
error=f"{type(exc).__name__}: {exc}",
)
return False
async def _schedule_retry(
self, session_id: str, *, retry_count: int, error: str
) -> None:
"""按指数退避登记下一次归档尝试。"""
delay = min(
self.RETRY_BASE_DELAY * (2 ** max(0, retry_count - 1)),
self.RETRY_MAX_DELAY,
)
next_retry_at = self.clock() + delay
await self.schedule.mark_retry(
session_id,
retry_count=retry_count,
error=error,
next_retry_at=next_retry_at,
)
await self.schedule.reschedule(session_id, due_at=next_retry_at)
async def retry_now(self, session_id: str) -> None:
"""提供手工补偿入口,立即把指定会话放回扫描队列。"""
await self.redis.hset(
self.schedule.meta_key(session_id),
mapping={"archive_status": "pending"},
)
await self.schedule.reschedule(session_id, due_at=self.clock())
async def run(self) -> None:
"""启动后台循环,异常不会击穿主应用。"""
while not self._stop_event.is_set():
try:
await self.archive_due_sessions()
except asyncio.CancelledError:
raise
except Exception:
logger.exception("idle archive scan failed")
try:
await asyncio.wait_for(self._stop_event.wait(), timeout=self.interval)
except asyncio.TimeoutError:
continue
async def stop(self) -> None:
"""请求后台循环停止。"""
self._stop_event.set()
__all__ = ["IdleArchiveWorker"]
+58 -6
View File
@@ -16,9 +16,12 @@ class DialogueMemoryExtractor:
你是客服记忆提取器,只提取用户明确表达或稳定陈述的客户信息。
只返回 JSON 数组,不要输出 Markdown 或解释文字。
每项必须包含:tag、content、memory_type、source。
对于用户反复询问的产品或投资主题,也可以提取兴趣信号,并额外返回 signal_type=interest_query。
兴趣信号必须使用 CUSTOMER_PREFERENCE 和 dialogue_inferred,tag 使用稳定、可归一化的英文主题名。
兴趣主题只允许以下五类:conservative_interest、steady_interest、balanced_interest、enterprising_interest、aggressive_interest。
memory_type 只能是 PROFILE_FACT、PROFILE_CANDIDATE、CUSTOMER_PREFERENCE、INVESTMENT_GOAL、SERVICE_FACT。
source 只能是 dialogue_confirmed、dialogue_stated、dialogue_inferred。
客服知识问题、产品政策、寒暄、客服回复内容不要提取。
单次客服知识问题不要作为长期记忆保存;如果问题反映出客户对某个产品或投资主题的关注,使用 signal_type=interest_query 表示兴趣信号。
不确定的信息使用 dialogue_inferred,无法形成客户画像的信息不要提取。
""".strip()
@@ -37,8 +40,56 @@ source 只能是 dialogue_confirmed、dialogue_stated、dialogue_inferred。
),
}
)
response = await self.llm_client.chat(prompt, temperature=0, max_tokens=800)
return self._parse(response)
interest_signal = self._interest_fallback(query)
try:
response = await self.llm_client.chat(prompt, temperature=0, max_tokens=800)
result = self._parse(response)
except Exception:
if interest_signal:
return [interest_signal]
raise
if interest_signal and not any(
item.get("signal_type") == "interest_query" for item in result
):
result.append(interest_signal)
return result
@staticmethod
def _interest_fallback(query: str) -> dict[str, str] | None:
"""为产品兴趣问题提供确定性兜底,避免依赖 LLM 输出可选字段。"""
text = query.strip().lower()
if not text or not any(
marker in text
for marker in ("基金", "理财", "投资", "产品", "fund", "investment")
):
return None
if not any(
marker in text
for marker in ("哪些", "什么", "怎么选", "推荐", "适合", "比较", "了解", "有哪些", "what", "which", "how")
):
return None
risk_topics = (
(("保守", "conservative"), "conservative_interest", "用户关注保守型投资产品"),
(("稳健", "steady", "moderate"), "steady_interest", "用户关注稳健型投资产品"),
(("平衡", "balanced"), "balanced_interest", "用户关注平衡型投资产品"),
(("进取", "enterprising"), "enterprising_interest", "用户关注进取型投资产品"),
(("激进", "aggressive", "高风险", "high risk"), "aggressive_interest", "用户关注激进型投资产品"),
)
for markers, topic_tag, topic_content in risk_topics:
if any(marker in text for marker in markers):
tag = topic_tag
content = topic_content
break
else:
return None
return {
"tag": tag,
"content": content,
"memory_type": "CUSTOMER_PREFERENCE",
"source": "dialogue_inferred",
"signal_type": "interest_query",
}
@staticmethod
def _parse(response: str) -> list[dict[str, str]]:
@@ -60,14 +111,15 @@ source 只能是 dialogue_confirmed、dialogue_stated、dialogue_inferred。
continue
if item["memory_type"] not in valid_types or item["source"] not in valid_sources:
continue
result.append(
{
candidate = {
"tag": str(item["tag"])[:64],
"content": str(item["content"])[:512],
"memory_type": item["memory_type"],
"source": item["source"],
}
)
if item.get("signal_type") == "interest_query":
candidate["signal_type"] = "interest_query"
result.append(candidate)
return result
+59 -7
View File
@@ -3,6 +3,7 @@
from __future__ import annotations
import contextvars
import logging
import uuid
from types import SimpleNamespace
@@ -22,6 +23,8 @@ _active_messages = contextvars.ContextVar("client_agent_messages", default=None)
_active_warnings = contextvars.ContextVar("client_agent_memory_warnings", default=None)
_active_memory_context = contextvars.ContextVar("client_agent_memory_context", default=None)
logger = logging.getLogger(__name__)
class MemoryConversationContext:
"""将现有客服 Agent 的上下文接口桥接到 MemoryService 短期记忆。"""
@@ -90,7 +93,14 @@ class MemoryAwareClientAgent:
try:
result = await self.agent.handle(session_id, query, trace_id=trace_id)
if self.extractor is not None:
await self._save_candidates(customer_id, session_id, query, memory_context, warnings)
await self._save_candidates(
customer_id,
session_id,
query,
memory_context,
warnings,
trace_id=trace_id,
)
result["memory_warnings"] = list(warnings)
return result
finally:
@@ -99,17 +109,42 @@ class MemoryAwareClientAgent:
_active_warnings.reset(warnings_token)
_active_memory_context.reset(context_token)
async def _save_candidates(self, customer_id, session_id, query, context, warnings):
"""提取并保存候选客户记忆,任何失败都只写入 warning。"""
async def _save_candidates(
self,
customer_id,
session_id,
query,
context,
warnings,
*,
trace_id: str | None = None,
):
"""提取并保存候选客户记忆,单个阶段或候选失败不影响客服回答。"""
try:
candidates = await self.extractor.extract(
query,
context={
"profile": context.customer_profile,
"memories": [item.model_dump(mode="json") for item in context.long_term_memories],
"memories": [self._memory_to_dict(item) for item in context.long_term_memories],
},
)
for candidate in candidates:
except Exception as exc:
logger.exception(
"client memory extraction failed: trace_id=%s customer_id=%s session_id=%s",
trace_id,
customer_id,
session_id,
)
warnings.append(f"memory_extraction_failed:{type(exc).__name__}")
return
for index, candidate in enumerate(candidates):
try:
evidence_count = 1
if candidate.get("signal_type") == "interest_query":
# 兴趣主题首次出现即保存为候选,后续由长期记忆按精确内容合并证据。
candidate["memory_type"] = "CUSTOMER_PREFERENCE"
candidate["source"] = "dialogue_inferred"
memory = MemoryUnitDTO(
customer_id=customer_id,
session_id=session_id,
@@ -118,14 +153,31 @@ 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,
memory=memory,
)
warnings.extend(save_warnings)
except Exception as exc:
warnings.append(f"memory_candidate_extract_failed:{type(exc).__name__}")
except Exception as exc:
logger.exception(
"client memory candidate save failed: trace_id=%s customer_id=%s session_id=%s index=%s",
trace_id,
customer_id,
session_id,
index,
)
warnings.append(f"memory_save_failed:{type(exc).__name__}")
@staticmethod
def _memory_to_dict(item) -> dict:
"""将召回记忆兼容转换为可序列化字典。"""
if hasattr(item, "model_dump"):
return item.model_dump(mode="json")
if isinstance(item, dict):
return dict(item)
raise TypeError(f"unsupported memory context item: {type(item).__name__}")
def build_client_runtime(
+37 -3
View File
@@ -54,19 +54,28 @@ class AnonymousCustomerAgent:
intent = await _maybe_await(self.intent_recognize(query))
sources = []
if intent is Intent.GUIDE_PURCHASE:
if intent == Intent.GUIDE_PURCHASE:
answer = await _config(
self.config_getter,
"agent.customer.template.guide_purchase",
"请前往开户页面办理。",
)
elif intent is Intent.WANT_ADVISOR:
elif intent == Intent.WANT_ADVISOR:
answer = await _config(
self.config_getter,
"agent.customer.template.guide_advisor",
"如需基金推荐,请联系投资顾问。",
)
elif intent is Intent.KNOWLEDGE_QA:
elif intent == Intent.OFF_TOPIC:
answer = await _config(
self.config_getter,
"agent.customer.template.off_topic",
"我是华夏科技的智能客服,只能解答基金与公司业务相关的问题,"
"您可以问我基金知识、开户流程或公司信息~",
)
elif intent == Intent.CHITCHAT:
answer = await self._chitchat(session_id)
elif intent in (Intent.KNOWLEDGE_QA, Intent.COMPANY_INFO):
try:
sources = await _maybe_await(self.rag_retrieve(query, None))
except Exception:
@@ -112,6 +121,31 @@ class AnonymousCustomerAgent:
"trace_id": trace_id,
}
async def _chitchat(self, session_id: str) -> str:
"""带对话历史调用 LLM 做受限闲聊,失败时退回固定话术。"""
messages = await self.context.get(session_id)
prompt = [
{
"role": "system",
"content": (
"你是华夏科技(基金代销金融机构)的智能客服助手。"
"用户正在与你寒暄,请用一两句话简短、友好地回应,"
"并自然地引导用户咨询基金知识、开户流程或公司信息。"
"不得推荐任何基金产品,不得谈论具体收益,不得回答金融之外的实质性问题。"
),
},
*messages,
]
try:
return await _maybe_await(self.generate_answer(prompt))
except Exception:
return await _config(
self.config_getter,
"agent.customer.template.chitchat_fallback",
"您好,我是华夏科技的智能客服,很高兴为您服务!"
"您可以问我基金知识、开户流程或公司信息~",
)
@staticmethod
def _contains_sensitive_input(query: str) -> bool:
return bool(
+15
View File
@@ -13,6 +13,7 @@ from .customer_relation import CustomerRelationMemory
from .customer_product import CustomerProductMemory
from .context_builder import build_customer_memory_context
from .long_term import LongTermMemoryService
from .interest_topic import InterestTopicTracker
from .profile import CustomerProfileMemory
from .schemas import CustomerMemoryContext, MemoryUnitDTO, ShortTermMessage
from .short_term import ShortTermMemory
@@ -34,6 +35,7 @@ class MemoryService:
long_term: LongTermMemoryService | None = None,
archiver: ConversationArchiver | None = None,
rank_tool: FinalConfidenceRankTool | None = None,
interest_tracker: InterestTopicTracker | None = None,
):
self.session_factory = session_factory or get_session_factory()
self.short_term = short_term or ShortTermMemory()
@@ -44,6 +46,7 @@ class MemoryService:
self.long_term = long_term or LongTermMemoryService()
self.archiver = archiver or ConversationArchiver(short_term=self.short_term)
self.rank_tool = rank_tool or FinalConfidenceRankTool()
self.interest_tracker = interest_tracker or InterestTopicTracker(self.short_term.redis)
@asynccontextmanager
async def _db(self):
@@ -67,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:
@@ -140,6 +151,10 @@ class MemoryService:
async with self._db() as db:
return await self.long_term.save(db, memory)
async def record_interest_signal(self, *, customer_id: int, tag: str) -> tuple[int, bool]:
"""记录一次重复兴趣主题信号,达到阈值后允许写入长期记忆。"""
return await self.interest_tracker.record(customer_id, tag)
async def close_session(
self,
*,
+38
View File
@@ -0,0 +1,38 @@
"""客服对话中的重复兴趣主题计数。"""
from __future__ import annotations
import hashlib
from datetime import datetime
class InterestTopicTracker:
"""使用 Redis 统计客户当天对同一主题的重复关注次数。"""
KEY_PREFIX = "customer:memory:interest"
WINDOW_SECONDS = 24 * 60 * 60
DEFAULT_THRESHOLD = 3
def __init__(self, redis, *, threshold: int = DEFAULT_THRESHOLD):
self.redis = redis
self.threshold = threshold
@classmethod
def key(cls, customer_id: int, tag: str, now: datetime | None = None) -> str:
"""生成按客户、自然日和主题隔离的 Redis 计数 Key。"""
day = (now or datetime.now()).strftime("%Y%m%d")
digest = hashlib.sha256(tag.strip().lower().encode("utf-8")).hexdigest()[:16]
return f"{cls.KEY_PREFIX}:{customer_id}:{day}:{digest}"
async def record(self, customer_id: int, tag: str) -> tuple[int, bool]:
"""记录一次主题关注,返回累计次数和是否达到长期记忆阈值。"""
if not tag or not tag.strip():
raise ValueError("兴趣主题不能为空")
key = self.key(customer_id, tag)
count = await self.redis.incr(key)
if count == 1:
await self.redis.expire(key, self.WINDOW_SECONDS)
return int(count), int(count) >= self.threshold
__all__ = ["InterestTopicTracker"]
+28 -5
View File
@@ -39,14 +39,26 @@ 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:
values = memory.model_dump(mode="json", exclude={"id", "milvus_id", "graph_node_id"})
# final_score 是召回重排阶段的临时分数,不属于 MySQL 主体字段。
values = memory.model_dump(
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)
@@ -112,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,
)
@@ -143,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)
@@ -178,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
+4 -1
View File
@@ -184,7 +184,10 @@ class ShortTermMemory:
idle_ttl = await _config(
self.config_getter, "agent.customer.session.ttl", 1800
)
ttl = min(idle_ttl, remaining)
archive_grace = await _config(
self.config_getter, "agent.customer.session.archive_grace", 60
)
ttl = min(idle_ttl + archive_grace, remaining)
await self.redis.expire(self.message_key(session_id), ttl)
await self.redis.expire(self.meta_key(session_id), remaining)
-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。"""
+1 -1
View File
@@ -204,7 +204,7 @@ class LLMClient:
payload = {
"model": backend.embed_model,
"input": texts,
"dimensions": self.cfg.embed_dimensions,
# "dimensions": self.cfg.embed_dimensions,
}
async with backend.client(self.cfg.timeout) as client:
r = await client.post(url, headers=backend.headers, json=payload)