288 lines
10 KiB
Python
288 lines
10 KiB
Python
"""
|
|
进程内 NL 附属数据(会话 / 消息 / 收藏),字段与 web 端 ApiEnvelope、SessionRow、MessageRow 对齐。
|
|
仅用于本地或 demo 联调,重启后数据丢失。
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import uuid
|
|
from datetime import datetime, timezone
|
|
from typing import Any, Dict, List, Optional, Tuple
|
|
|
|
|
|
def utc_ts() -> str:
|
|
return datetime.now(timezone.utc).strftime("%Y-%m-%d %H:%M:%S")
|
|
|
|
|
|
class LiteNlStore:
|
|
def __init__(self) -> None:
|
|
self._lock = asyncio.Lock()
|
|
self._sessions: Dict[Tuple[str, str], List[Dict[str, Any]]] = {}
|
|
self._messages: Dict[Tuple[str, str, str], List[Dict[str, Any]]] = {}
|
|
self._msg_counters: Dict[Tuple[str, str, str], int] = {}
|
|
self._fav: Dict[Tuple[str, str], Dict[str, List[Any]]] = {}
|
|
|
|
@staticmethod
|
|
def _visitor_key(vid: Optional[str]) -> str:
|
|
return (vid or "").strip()
|
|
|
|
def _scope(self, user_id: Optional[str], visitor_biz_id: Optional[str]) -> Tuple[str, str]:
|
|
uid = (user_id or "anonymous").strip() or "anonymous"
|
|
return uid, self._visitor_key(visitor_biz_id)
|
|
|
|
async def create_session(
|
|
self,
|
|
user_id: Optional[str],
|
|
visitor_biz_id: Optional[str],
|
|
title: Optional[str],
|
|
) -> Dict[str, Any]:
|
|
async with self._lock:
|
|
sk = self._scope(user_id, visitor_biz_id)
|
|
sid = str(uuid.uuid4())
|
|
now = utc_ts()
|
|
row: Dict[str, Any] = {
|
|
"session_id": sid,
|
|
"title": (title or "").strip() or None,
|
|
"user_id": sk[0],
|
|
"visitor_biz_id": visitor_biz_id or None,
|
|
"created_at": now,
|
|
"updated_at": now,
|
|
}
|
|
self._sessions.setdefault(sk, []).insert(0, row)
|
|
self._messages[(*sk, sid)] = []
|
|
self._msg_counters[(*sk, sid)] = 0
|
|
return {"session_id": sid, "title": row["title"]}
|
|
|
|
async def list_sessions(
|
|
self,
|
|
user_id: Optional[str],
|
|
visitor_biz_id: Optional[str],
|
|
limit: int,
|
|
offset: int,
|
|
) -> Dict[str, Any]:
|
|
async with self._lock:
|
|
sk = self._scope(user_id, visitor_biz_id)
|
|
items = list(self._sessions.get(sk, []))
|
|
items.sort(key=lambda r: str(r.get("updated_at") or ""), reverse=True)
|
|
sl = items[offset : offset + limit]
|
|
return {"items": sl, "limit": limit, "offset": offset}
|
|
|
|
def _bump_session(self, sk: Tuple[str, str], session_id: str, user_first_line: str) -> None:
|
|
now = utc_ts()
|
|
for r in self._sessions.get(sk, []):
|
|
if r.get("session_id") != session_id:
|
|
continue
|
|
r["updated_at"] = now
|
|
if not (r.get("title") or "").strip() and user_first_line.strip():
|
|
r["title"] = user_first_line.strip()[:80]
|
|
break
|
|
|
|
async def append_exchange(
|
|
self,
|
|
user_id: Optional[str],
|
|
visitor_biz_id: Optional[str],
|
|
session_id: str,
|
|
user_text: str,
|
|
assistant_content_json: str,
|
|
) -> None:
|
|
async with self._lock:
|
|
sk = self._scope(user_id, visitor_biz_id)
|
|
sk3 = (*sk, session_id)
|
|
if sk3 not in self._messages:
|
|
self._messages[sk3] = []
|
|
self._msg_counters[sk3] = 0
|
|
now = utc_ts()
|
|
|
|
def next_id() -> int:
|
|
n = self._msg_counters.get(sk3, 0) + 1
|
|
self._msg_counters[sk3] = n
|
|
return n
|
|
|
|
self._messages[sk3].append(
|
|
{
|
|
"id": next_id(),
|
|
"role": "user",
|
|
"content": user_text,
|
|
"created_at": now,
|
|
"llm_total_tokens": None,
|
|
}
|
|
)
|
|
self._messages[sk3].append(
|
|
{
|
|
"id": next_id(),
|
|
"role": "assistant",
|
|
"content": assistant_content_json,
|
|
"created_at": now,
|
|
"llm_total_tokens": None,
|
|
}
|
|
)
|
|
self._bump_session(sk, session_id, user_text)
|
|
|
|
async def get_messages(
|
|
self,
|
|
user_id: Optional[str],
|
|
visitor_biz_id: Optional[str],
|
|
session_id: str,
|
|
limit: int,
|
|
offset: int,
|
|
) -> Optional[Dict[str, Any]]:
|
|
async with self._lock:
|
|
sk = self._scope(user_id, visitor_biz_id)
|
|
sk3 = (*sk, session_id)
|
|
if sk3 not in self._messages:
|
|
found = any(
|
|
r.get("session_id") == session_id for r in self._sessions.get(sk, [])
|
|
)
|
|
if not found:
|
|
return None
|
|
self._messages[sk3] = []
|
|
self._msg_counters[sk3] = 0
|
|
rows = list(self._messages[sk3])
|
|
rows = rows[offset : offset + limit]
|
|
return {"session_id": session_id, "items": rows, "limit": limit, "offset": offset}
|
|
|
|
async def update_session_title(
|
|
self,
|
|
user_id: Optional[str],
|
|
visitor_biz_id: Optional[str],
|
|
session_id: str,
|
|
title: Optional[str],
|
|
) -> Optional[Dict[str, Any]]:
|
|
async with self._lock:
|
|
sk = self._scope(user_id, visitor_biz_id)
|
|
for r in self._sessions.get(sk, []):
|
|
if r.get("session_id") == session_id:
|
|
r["title"] = title
|
|
r["updated_at"] = utc_ts()
|
|
return {"session_id": session_id, "title": r.get("title")}
|
|
return None
|
|
|
|
async def delete_session(
|
|
self,
|
|
user_id: Optional[str],
|
|
visitor_biz_id: Optional[str],
|
|
session_id: str,
|
|
) -> Optional[Dict[str, Any]]:
|
|
async with self._lock:
|
|
sk = self._scope(user_id, visitor_biz_id)
|
|
lst = self._sessions.get(sk, [])
|
|
sk3 = (*sk, session_id)
|
|
n = len(self._messages.get(sk3, []))
|
|
new_lst = [x for x in lst if x.get("session_id") != session_id]
|
|
if len(new_lst) == len(lst):
|
|
return None
|
|
self._sessions[sk] = new_lst
|
|
self._messages.pop(sk3, None)
|
|
self._msg_counters.pop(sk3, None)
|
|
return {"session_id": session_id, "deleted_messages": n}
|
|
|
|
async def patch_message(
|
|
self,
|
|
user_id: Optional[str],
|
|
visitor_biz_id: Optional[str],
|
|
session_id: str,
|
|
message_id: int,
|
|
content: str,
|
|
) -> Optional[Dict[str, Any]]:
|
|
async with self._lock:
|
|
sk3 = (*self._scope(user_id, visitor_biz_id), session_id)
|
|
for m in self._messages.get(sk3, []):
|
|
if m.get("id") == message_id:
|
|
m["content"] = content
|
|
m["created_at"] = utc_ts()
|
|
return {
|
|
"id": message_id,
|
|
"session_id": session_id,
|
|
"role": m.get("role"),
|
|
"content": content,
|
|
"created_at": m.get("created_at"),
|
|
"llm_total_tokens": m.get("llm_total_tokens"),
|
|
}
|
|
return None
|
|
|
|
async def get_favorites_grouped(
|
|
self, user_id: Optional[str], visitor_biz_id: Optional[str]
|
|
) -> Dict[str, List[Any]]:
|
|
async with self._lock:
|
|
sk = self._scope(user_id, visitor_biz_id)
|
|
g = self._fav.get(sk) or {"sql": [], "function": [], "report": []}
|
|
return {"sql": list(g["sql"]), "function": list(g["function"]), "report": list(g["report"])}
|
|
|
|
async def add_favorite(
|
|
self, user_id: Optional[str], visitor_biz_id: Optional[str], body: Dict[str, Any]
|
|
) -> Any:
|
|
async with self._lock:
|
|
sk = self._scope(user_id, visitor_biz_id)
|
|
g = self._fav.setdefault(sk, {"sql": [], "function": [], "report": []})
|
|
fav_type = str(body.get("fav_type") or "sql")
|
|
fid = f"{uuid.uuid4().hex[:12]}"
|
|
if fav_type == "sql":
|
|
row = {
|
|
"id": fid,
|
|
"name": str(body.get("name") or ""),
|
|
"desc": str(body.get("desc") or ""),
|
|
"sql": str(body.get("sql") or ""),
|
|
}
|
|
if body.get("sql_explain"):
|
|
row["sql_explain"] = str(body["sql_explain"])
|
|
g["sql"].insert(0, row)
|
|
return row
|
|
if fav_type == "function":
|
|
row = {
|
|
"id": fid,
|
|
"name": str(body.get("name") or ""),
|
|
"desc": str(body.get("desc") or ""),
|
|
"path": str(body.get("path") or ""),
|
|
}
|
|
g["function"].insert(0, row)
|
|
return row
|
|
row = {
|
|
"id": fid,
|
|
"name": str(body.get("name") or ""),
|
|
"desc": str(body.get("desc") or ""),
|
|
"reportPath": str(body.get("reportPath") or ""),
|
|
"params": str(body.get("params") or ""),
|
|
}
|
|
g["report"].insert(0, row)
|
|
return row
|
|
|
|
async def patch_favorite(
|
|
self,
|
|
user_id: Optional[str],
|
|
visitor_biz_id: Optional[str],
|
|
fav_id: str,
|
|
patch: Dict[str, Any],
|
|
) -> bool:
|
|
async with self._lock:
|
|
sk = self._scope(user_id, visitor_biz_id)
|
|
g = self._fav.get(sk)
|
|
if not g:
|
|
return False
|
|
for bucket in ("sql", "function", "report"):
|
|
lst = g.get(bucket, [])
|
|
for i, it in enumerate(lst):
|
|
if str(it.get("id")) == fav_id:
|
|
lst[i] = {**it, **patch}
|
|
return True
|
|
return False
|
|
|
|
async def delete_favorite(
|
|
self, user_id: Optional[str], visitor_biz_id: Optional[str], fav_id: str
|
|
) -> bool:
|
|
async with self._lock:
|
|
sk = self._scope(user_id, visitor_biz_id)
|
|
g = self._fav.get(sk)
|
|
if not g:
|
|
return False
|
|
removed = False
|
|
for bucket in ("sql", "function", "report"):
|
|
before = len(g.get(bucket, []))
|
|
g[bucket] = [x for x in g.get(bucket, []) if str(x.get("id")) != fav_id]
|
|
if len(g[bucket]) < before:
|
|
removed = True
|
|
return removed
|
|
|
|
|
|
lite_nl_store = LiteNlStore()
|