Files

288 lines
10 KiB
Python
Raw Permalink Normal View History

2026-04-14 10:28:22 +08:00
"""
进程内 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()