Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3103389aec | ||
|
|
30160128a4 |
+1
-5
@@ -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
|
||||
@@ -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(
|
||||
|
||||
@@ -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}",
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
+6
-3
@@ -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)}
|
||||
|
||||
@@ -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"]
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
*,
|
||||
|
||||
@@ -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"]
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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。"""
|
||||
|
||||
+1
-1
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user