Develop #25

Merged
ouyangyang_0626 merged 69 commits from develop into master 2026-09-14 19:28:54 +08:00
16 changed files with 589 additions and 23 deletions
Showing only changes of commit 30160128a4 - Show all commits
+3 -1
View File
@@ -80,4 +80,6 @@ DK4_客服Agent需求文档(修复完善版v1.4).md
客服Agent记忆模块TODO清单.md
客服Agent记忆模块开发计划.md
开发计划_用户画像置信度实时更新.md
客服Agent记忆模块接入说明.md
客服Agent记忆模块接入说明.md
客服Agent接口清单.md
空闲会话自动归档开发计划.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
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)
+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"]
+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
+62 -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,46 @@ 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:
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(
customer_id=customer_id,
session_id=session_id,
@@ -124,8 +163,24 @@ class MemoryAwareClientAgent:
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(
+7
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):
@@ -140,6 +143,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"]
+5 -1
View File
@@ -46,7 +46,11 @@ class LongTermMemoryService:
existing = await repo.merge_evidence(existing)
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"},
)
values["memory_type"] = memory.memory_type.value
values["source"] = memory.source.value
confidence_result, confidence_warning = self._calculate_confidence(memory)
+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 -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)