0.1.1 暂存
This commit is contained in:
@@ -0,0 +1,287 @@
|
||||
"""
|
||||
进程内 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()
|
||||
Reference in New Issue
Block a user