diff --git a/.gitignore b/.gitignore index 810c1cb..5d4d5d6 100644 --- a/.gitignore +++ b/.gitignore @@ -80,4 +80,6 @@ DK4_客服Agent需求文档(修复完善版v1.4).md 客服Agent记忆模块TODO清单.md 客服Agent记忆模块开发计划.md 开发计划_用户画像置信度实时更新.md -客服Agent记忆模块接入说明.md \ No newline at end of file +客服Agent记忆模块接入说明.md +客服Agent接口清单.md +空闲会话自动归档开发计划.md \ No newline at end of file diff --git a/agent/client_agent/session.py b/agent/client_agent/session.py index fabca83..a4498f7 100644 --- a/agent/client_agent/session.py +++ b/agent/client_agent/session.py @@ -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( diff --git a/api/chat/client_agent.py b/api/chat/client_agent.py index d8a4062..f9d77de 100644 --- a/api/chat/client_agent.py +++ b/api/chat/client_agent.py @@ -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}", diff --git a/main.py b/main.py index dab8c25..cac9545 100644 --- a/main.py +++ b/main.py @@ -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) diff --git a/rag/retrieve.py b/rag/retrieve.py index 4fd41ae..98990b1 100644 --- a/rag/retrieve.py +++ b/rag/retrieve.py @@ -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)} \ No newline at end of file + return {"status": "ok", "sources": _format_candidates(candidates)} diff --git a/repositories/conversation_archive.py b/repositories/conversation_archive.py index 1a5f929..f6f8b78 100644 --- a/repositories/conversation_archive.py +++ b/repositories/conversation_archive.py @@ -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"] - diff --git a/service/client_agent/archive_lock.py b/service/client_agent/archive_lock.py new file mode 100644 index 0000000..04a054e --- /dev/null +++ b/service/client_agent/archive_lock.py @@ -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"] diff --git a/service/client_agent/archive_schedule.py b/service/client_agent/archive_schedule.py new file mode 100644 index 0000000..2c4ecc3 --- /dev/null +++ b/service/client_agent/archive_schedule.py @@ -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"] diff --git a/service/client_agent/idle_archive_worker.py b/service/client_agent/idle_archive_worker.py new file mode 100644 index 0000000..c2fa1d5 --- /dev/null +++ b/service/client_agent/idle_archive_worker.py @@ -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"] diff --git a/service/client_agent/memory_extractor.py b/service/client_agent/memory_extractor.py index 6b5c947..97be13c 100644 --- a/service/client_agent/memory_extractor.py +++ b/service/client_agent/memory_extractor.py @@ -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 diff --git a/service/client_agent/runtime.py b/service/client_agent/runtime.py index 15f911e..9f71b3b 100644 --- a/service/client_agent/runtime.py +++ b/service/client_agent/runtime.py @@ -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( diff --git a/service/memory/facade.py b/service/memory/facade.py index 0ada857..972ef7b 100644 --- a/service/memory/facade.py +++ b/service/memory/facade.py @@ -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, *, diff --git a/service/memory/interest_topic.py b/service/memory/interest_topic.py new file mode 100644 index 0000000..1d8e23c --- /dev/null +++ b/service/memory/interest_topic.py @@ -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"] diff --git a/service/memory/long_term.py b/service/memory/long_term.py index f49324f..30f59cc 100644 --- a/service/memory/long_term.py +++ b/service/memory/long_term.py @@ -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) diff --git a/service/memory/short_term.py b/service/memory/short_term.py index 528705f..18a90d1 100644 --- a/service/memory/short_term.py +++ b/service/memory/short_term.py @@ -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) diff --git a/tool/llm.py b/tool/llm.py index c11e5e9..9ca55fb 100644 --- a/tool/llm.py +++ b/tool/llm.py @@ -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)