""" 进程内 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()