feat:新增投顾agent和nl2sqlagent
This commit is contained in:
@@ -0,0 +1,82 @@
|
||||
"""NL2SQL 会话上下文的 Redis 存储与 Prompt 格式化。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
|
||||
logger = logging.getLogger("nl2sql.session_context")
|
||||
|
||||
|
||||
class SessionContextStore:
|
||||
"""按员工隔离并限制大小的短期查询上下文。"""
|
||||
|
||||
def __init__(self, redis, *, ttl: int = 1800, max_messages: int = 8, max_chars: int = 4000):
|
||||
self.redis = redis
|
||||
self.ttl = ttl
|
||||
self.max_messages = max_messages
|
||||
self.max_chars = max_chars
|
||||
|
||||
def _key(self, session_id: str) -> str:
|
||||
return f"nl2sql:session:{session_id}:context"
|
||||
|
||||
def _owner_key(self, session_id: str) -> str:
|
||||
return f"nl2sql:session:{session_id}:owner"
|
||||
|
||||
async def _belongs_to(self, user_id: int, session_id: str) -> bool:
|
||||
owner_key = self._owner_key(session_id)
|
||||
owner = await self.redis.get(owner_key)
|
||||
if owner is None:
|
||||
return bool(await self.redis.set(owner_key, str(user_id), nx=True, ex=self.ttl))
|
||||
if isinstance(owner, bytes):
|
||||
owner = owner.decode()
|
||||
return str(owner) == str(user_id)
|
||||
|
||||
async def load(self, user_id: int, session_id: str | None) -> list[dict[str, str]]:
|
||||
"""读取当前员工的上下文,异常或跨员工访问时返回空列表。"""
|
||||
if not session_id:
|
||||
return []
|
||||
try:
|
||||
if not await self._belongs_to(user_id, session_id):
|
||||
return []
|
||||
raw = await self.redis.get(self._key(session_id))
|
||||
if not raw:
|
||||
return []
|
||||
if isinstance(raw, bytes):
|
||||
raw = raw.decode()
|
||||
value = json.loads(raw)
|
||||
return value if isinstance(value, list) else []
|
||||
except Exception: # noqa: BLE001 上下文故障不阻断查询
|
||||
logger.warning("NL2SQL 会话上下文读取失败", exc_info=True)
|
||||
return []
|
||||
|
||||
async def append(self, user_id: int, session_id: str | None, question: str, status: str) -> bool:
|
||||
"""追加问题和状态摘要,不保存 SQL、结果行或敏感数据。"""
|
||||
if not session_id:
|
||||
return False
|
||||
try:
|
||||
if not await self._belongs_to(user_id, session_id):
|
||||
return False
|
||||
messages = await self.load(user_id, session_id)
|
||||
messages.append({"role": "user", "content": question[:1000]})
|
||||
messages.append({"role": "assistant", "content": status[:100]})
|
||||
messages = messages[-self.max_messages :]
|
||||
while len(json.dumps(messages, ensure_ascii=False)) > self.max_chars and messages:
|
||||
messages.pop(0)
|
||||
await self.redis.set(self._key(session_id), json.dumps(messages, ensure_ascii=False), ex=self.ttl)
|
||||
return True
|
||||
except Exception: # noqa: BLE001 上下文故障不阻断查询
|
||||
logger.warning("NL2SQL 会话上下文写入失败", exc_info=True)
|
||||
return False
|
||||
|
||||
|
||||
def build_conversation_context(messages: list[dict[str, Any]], *, max_chars: int = 2000) -> str:
|
||||
"""构造有界 Prompt 上下文,只保留角色和文本内容。"""
|
||||
lines: list[str] = []
|
||||
for message in messages:
|
||||
role = str(message.get("role", ""))[:20]
|
||||
content = str(message.get("content", ""))[:500]
|
||||
if role and content:
|
||||
lines.append(f"{role}: {content}")
|
||||
return "\n".join(lines)[-max_chars:]
|
||||
Reference in New Issue
Block a user