dbassistant.py 36 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896897898899900901902903904905906907908909910911912913914915916917918919920921922923924925926927928929930931932933934935936937938939940941942943944945946947948949950951952953954955956957958959960961962963964965966967968969970971972973974975976977978979980981982983984985986987988989990991992993994995996997998999100010011002100310041005100610071008100910101011101210131014101510161017101810191020102110221023102410251026102710281029103010311032103310341035103610371038103910401041104210431044104510461047
  1. from __future__ import annotations
  2. import asyncio
  3. import hashlib
  4. import re
  5. from datetime import UTC, datetime, timedelta
  6. from typing import Any
  7. from uuid import uuid4
  8. from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
  9. from pymongo import ASCENDING, DESCENDING
  10. from pymongo.errors import DuplicateKeyError
  11. from wbb import BOT_PROFILE_ID, db
  12. connectionsdb = db.business_assistant_connections
  13. settingsdb = db.business_assistant_settings
  14. knowledgedb = db.business_assistant_knowledge
  15. conversationsdb = db.business_assistant_conversations
  16. messagesdb = db.business_assistant_messages
  17. usagedb = db.business_assistant_usage
  18. updatesdb = db.business_assistant_updates
  19. deadlettersdb = db.business_assistant_dead_letters
  20. runtimedb = db.business_assistant_runtime
  21. DEFAULT_ACCOUNT_SETTINGS: dict[str, Any] = {
  22. "assistant_enabled": False,
  23. "system_prompt": (
  24. "你是该 Telegram 账号的智能接待秘书。回答要简洁、礼貌,只能依据提供的知识条目陈述业务事实。"
  25. ),
  26. "language": "zh-CN",
  27. "tone": "professional",
  28. "account_daily_limit": 200,
  29. "customer_daily_limit": 20,
  30. "human_pause_hours": 24,
  31. "notification_destination": "owner",
  32. "ops_group_id": 0,
  33. "timezone": "Asia/Shanghai",
  34. "digest_enabled": False,
  35. "digest_time": "09:00",
  36. "handoff_message": "这个问题需要人工确认,我已经通知负责人,请稍候。",
  37. "unsupported_message": "已收到你的消息,这类内容需要人工处理,我已经通知负责人。",
  38. }
  39. VALID_NOTIFICATION_DESTINATIONS = {"owner", "ops", "both"}
  40. VALID_TONES = {"professional", "friendly", "concise"}
  41. HANDOFF_TERMS = (
  42. "人工",
  43. "真人",
  44. "客服",
  45. "负责人",
  46. "转人工",
  47. "human",
  48. "agent",
  49. )
  50. SENSITIVE_TERMS = (
  51. "承诺",
  52. "保证",
  53. "投诉",
  54. "退款",
  55. "退钱",
  56. "赔偿",
  57. "律师",
  58. "起诉",
  59. "支付失败",
  60. "账号被盗",
  61. "密码",
  62. "验证码",
  63. )
  64. _index_lock = asyncio.Lock()
  65. _indexes_ready = False
  66. class AssistantDataError(ValueError):
  67. def __init__(self, code: str, message: str):
  68. super().__init__(message)
  69. self.code = code
  70. def utc_now() -> datetime:
  71. return datetime.now(UTC)
  72. def as_utc(value: datetime) -> datetime:
  73. return value.replace(tzinfo=UTC) if value.tzinfo is None else value.astimezone(UTC)
  74. def _scope(filters: dict[str, Any] | None = None, /, **values: Any) -> dict[str, Any]:
  75. return {"bot_id": BOT_PROFILE_ID, **(filters or {}), **values}
  76. def clean_text(
  77. value: Any,
  78. *,
  79. max_length: int,
  80. required: bool = False,
  81. preserve_lines: bool = False,
  82. ) -> str:
  83. text = re.sub(r"[\x00-\x08\x0b\x0c\x0e-\x1f\x7f]", "", str(value or ""))
  84. text = text.strip()
  85. if not preserve_lines:
  86. text = " ".join(text.split())
  87. if required and not text:
  88. raise AssistantDataError("required_field", "必填内容不能为空。")
  89. if len(text) > max_length:
  90. raise AssistantDataError("text_too_long", f"内容不能超过 {max_length} 个字符。")
  91. return text
  92. def _bounded_int(
  93. value: Any,
  94. name: str,
  95. *,
  96. minimum: int,
  97. maximum: int,
  98. ) -> int:
  99. try:
  100. parsed = int(value)
  101. except (TypeError, ValueError) as exc:
  102. raise AssistantDataError("invalid_setting", f"{name} 必须是整数。") from exc
  103. if not minimum <= parsed <= maximum:
  104. raise AssistantDataError(
  105. "invalid_setting",
  106. f"{name} 必须在 {minimum} 到 {maximum} 之间。",
  107. )
  108. return parsed
  109. def _normalize_string_list(value: Any, *, max_items: int, max_length: int) -> list[str]:
  110. values = value if isinstance(value, (list, tuple, set)) else str(value or "").split(",")
  111. normalized: list[str] = []
  112. seen: set[str] = set()
  113. for item in values:
  114. text = clean_text(item, max_length=max_length)
  115. key = text.casefold()
  116. if text and key not in seen:
  117. normalized.append(text)
  118. seen.add(key)
  119. if len(normalized) > max_items:
  120. raise AssistantDataError("too_many_items", f"最多允许 {max_items} 项。")
  121. return normalized
  122. async def ensure_assistant_indexes() -> None:
  123. global _indexes_ready
  124. if _indexes_ready:
  125. return
  126. async with _index_lock:
  127. if _indexes_ready:
  128. return
  129. await connectionsdb.create_index(
  130. [("bot_id", ASCENDING), ("connection_id", ASCENDING)], unique=True
  131. )
  132. await connectionsdb.create_index(
  133. [("bot_id", ASCENDING), ("updated_at", DESCENDING)]
  134. )
  135. await settingsdb.create_index(
  136. [("bot_id", ASCENDING), ("connection_id", ASCENDING)], unique=True
  137. )
  138. await knowledgedb.create_index(
  139. [("bot_id", ASCENDING), ("entry_id", ASCENDING)], unique=True
  140. )
  141. await knowledgedb.create_index(
  142. [
  143. ("bot_id", ASCENDING),
  144. ("connection_id", ASCENDING),
  145. ("enabled", ASCENDING),
  146. ("priority", DESCENDING),
  147. ]
  148. )
  149. await conversationsdb.create_index(
  150. [("bot_id", ASCENDING), ("conversation_id", ASCENDING)], unique=True
  151. )
  152. await conversationsdb.create_index(
  153. [
  154. ("bot_id", ASCENDING),
  155. ("connection_id", ASCENDING),
  156. ("chat_id", ASCENDING),
  157. ],
  158. unique=True,
  159. )
  160. await conversationsdb.create_index(
  161. [
  162. ("bot_id", ASCENDING),
  163. ("connection_id", ASCENDING),
  164. ("status", ASCENDING),
  165. ("updated_at", DESCENDING),
  166. ]
  167. )
  168. await messagesdb.create_index(
  169. [("bot_id", ASCENDING), ("message_key", ASCENDING)], unique=True
  170. )
  171. await messagesdb.create_index(
  172. [
  173. ("bot_id", ASCENDING),
  174. ("conversation_id", ASCENDING),
  175. ("created_at", ASCENDING),
  176. ]
  177. )
  178. await messagesdb.create_index("expires_at", expireAfterSeconds=0)
  179. await usagedb.create_index(
  180. [
  181. ("bot_id", ASCENDING),
  182. ("connection_id", ASCENDING),
  183. ("chat_id", ASCENDING),
  184. ("day", ASCENDING),
  185. ],
  186. unique=True,
  187. )
  188. await updatesdb.create_index(
  189. [("bot_id", ASCENDING), ("update_id", ASCENDING)], unique=True
  190. )
  191. await updatesdb.create_index("expires_at", expireAfterSeconds=0)
  192. await deadlettersdb.create_index(
  193. [("bot_id", ASCENDING), ("update_id", ASCENDING)], unique=True
  194. )
  195. await deadlettersdb.create_index(
  196. [("bot_id", ASCENDING), ("created_at", DESCENDING)]
  197. )
  198. await deadlettersdb.create_index("expires_at", expireAfterSeconds=0)
  199. await runtimedb.create_index(
  200. [("bot_id", ASCENDING), ("runtime_id", ASCENDING)], unique=True
  201. )
  202. _indexes_ready = True
  203. def _public_user(user: Any) -> dict[str, Any]:
  204. value = user if isinstance(user, dict) else {}
  205. return {
  206. "id": int(value.get("id") or 0),
  207. "username": clean_text(value.get("username"), max_length=64),
  208. "first_name": clean_text(value.get("first_name"), max_length=128),
  209. "last_name": clean_text(value.get("last_name"), max_length=128),
  210. }
  211. async def upsert_business_connection(payload: dict[str, Any]) -> dict[str, Any]:
  212. await ensure_assistant_indexes()
  213. connection_id = clean_text(payload.get("id"), max_length=128, required=True)
  214. now = utc_now()
  215. rights = payload.get("rights") if isinstance(payload.get("rights"), dict) else {}
  216. document = {
  217. "bot_id": BOT_PROFILE_ID,
  218. "connection_id": connection_id,
  219. "user": _public_user(payload.get("user")),
  220. "user_chat_id": int(payload.get("user_chat_id") or 0),
  221. "rights": {str(key): bool(value) for key, value in rights.items() if value is True},
  222. "is_enabled": bool(payload.get("is_enabled")),
  223. "connected_at": datetime.fromtimestamp(int(payload.get("date") or 0), UTC)
  224. if payload.get("date")
  225. else now,
  226. "last_event_at": now,
  227. "updated_at": now,
  228. }
  229. await connectionsdb.update_one(
  230. _scope(connection_id=connection_id),
  231. {"$set": document, "$setOnInsert": {"created_at": now}},
  232. upsert=True,
  233. )
  234. await settingsdb.update_one(
  235. _scope(connection_id=connection_id),
  236. {
  237. "$setOnInsert": {
  238. "bot_id": BOT_PROFILE_ID,
  239. "connection_id": connection_id,
  240. **DEFAULT_ACCOUNT_SETTINGS,
  241. "created_at": now,
  242. "updated_at": now,
  243. }
  244. },
  245. upsert=True,
  246. )
  247. return await get_business_connection(connection_id) or document
  248. async def get_business_connection(connection_id: str) -> dict[str, Any] | None:
  249. await ensure_assistant_indexes()
  250. return await connectionsdb.find_one(_scope(connection_id=str(connection_id)))
  251. async def touch_business_connection(
  252. connection_id: str, *, error: str = "", is_enabled: bool | None = None
  253. ) -> dict[str, Any] | None:
  254. now = utc_now()
  255. values: dict[str, Any] = {
  256. "last_event_at": now,
  257. "updated_at": now,
  258. "last_error": clean_text(error, max_length=1000),
  259. }
  260. if is_enabled is not None:
  261. values["is_enabled"] = bool(is_enabled)
  262. await connectionsdb.update_one(
  263. _scope(connection_id=str(connection_id)),
  264. {"$set": values},
  265. )
  266. return await get_business_connection(connection_id)
  267. async def list_business_connections(
  268. *, page: int = 1, page_size: int = 20
  269. ) -> tuple[list[dict[str, Any]], int]:
  270. await ensure_assistant_indexes()
  271. page = max(1, int(page))
  272. page_size = max(1, min(int(page_size), 100))
  273. filters = _scope()
  274. total = await connectionsdb.count_documents(filters)
  275. cursor = (
  276. connectionsdb.find(filters)
  277. .sort("updated_at", DESCENDING)
  278. .skip((page - 1) * page_size)
  279. .limit(page_size)
  280. )
  281. items = []
  282. async for item in cursor:
  283. item["settings"] = await get_account_settings(item["connection_id"])
  284. items.append(item)
  285. return items, total
  286. async def get_account_settings(connection_id: str) -> dict[str, Any]:
  287. await ensure_assistant_indexes()
  288. stored = await settingsdb.find_one(_scope(connection_id=str(connection_id))) or {}
  289. return {
  290. **DEFAULT_ACCOUNT_SETTINGS,
  291. **{key: value for key, value in stored.items() if key != "_id"},
  292. "connection_id": str(connection_id),
  293. }
  294. def normalize_account_settings(
  295. values: dict[str, Any], *, previous: dict[str, Any] | None = None
  296. ) -> dict[str, Any]:
  297. current = {**DEFAULT_ACCOUNT_SETTINGS, **(previous or {})}
  298. if "assistant_enabled" in values:
  299. current["assistant_enabled"] = bool(values.get("assistant_enabled"))
  300. if "system_prompt" in values:
  301. current["system_prompt"] = clean_text(
  302. values.get("system_prompt"), max_length=4000, required=True, preserve_lines=True
  303. )
  304. if "language" in values:
  305. current["language"] = clean_text(values.get("language"), max_length=32, required=True)
  306. if "tone" in values:
  307. tone = str(values.get("tone") or "")
  308. if tone not in VALID_TONES:
  309. raise AssistantDataError("invalid_setting", "接待语气无效。")
  310. current["tone"] = tone
  311. if "account_daily_limit" in values:
  312. current["account_daily_limit"] = _bounded_int(
  313. values.get("account_daily_limit"), "账号每日额度", minimum=1, maximum=100000
  314. )
  315. if "customer_daily_limit" in values:
  316. current["customer_daily_limit"] = _bounded_int(
  317. values.get("customer_daily_limit"), "客户每日额度", minimum=1, maximum=1000
  318. )
  319. if "human_pause_hours" in values:
  320. current["human_pause_hours"] = _bounded_int(
  321. values.get("human_pause_hours"), "人工暂停小时数", minimum=1, maximum=168
  322. )
  323. if "notification_destination" in values:
  324. destination = str(values.get("notification_destination") or "")
  325. if destination not in VALID_NOTIFICATION_DESTINATIONS:
  326. raise AssistantDataError("invalid_setting", "通知目的地无效。")
  327. current["notification_destination"] = destination
  328. if "ops_group_id" in values:
  329. try:
  330. current["ops_group_id"] = int(values.get("ops_group_id") or 0)
  331. except (TypeError, ValueError) as exc:
  332. raise AssistantDataError("invalid_setting", "运营群 ID 必须是整数。") from exc
  333. if "timezone" in values:
  334. timezone_name = clean_text(values.get("timezone"), max_length=64, required=True)
  335. try:
  336. ZoneInfo(timezone_name)
  337. except ZoneInfoNotFoundError as exc:
  338. raise AssistantDataError("invalid_setting", "时区名称无效。") from exc
  339. current["timezone"] = timezone_name
  340. if "digest_enabled" in values:
  341. current["digest_enabled"] = bool(values.get("digest_enabled"))
  342. if "digest_time" in values:
  343. digest_time = str(values.get("digest_time") or "")
  344. if not re.fullmatch(r"(?:[01]\d|2[0-3]):[0-5]\d", digest_time):
  345. raise AssistantDataError("invalid_setting", "简报时间必须是 HH:MM。")
  346. current["digest_time"] = digest_time
  347. for key, max_length in (("handoff_message", 1000), ("unsupported_message", 1000)):
  348. if key in values:
  349. current[key] = clean_text(values.get(key), max_length=max_length, required=True)
  350. return {key: current[key] for key in DEFAULT_ACCOUNT_SETTINGS}
  351. async def update_account_settings(
  352. connection_id: str, values: dict[str, Any]
  353. ) -> dict[str, Any]:
  354. connection = await get_business_connection(connection_id)
  355. if not connection:
  356. raise AssistantDataError("connection_not_found", "未找到 Business 连接。")
  357. previous = await get_account_settings(connection_id)
  358. normalized = normalize_account_settings(values, previous=previous)
  359. now = utc_now()
  360. await settingsdb.update_one(
  361. _scope(connection_id=str(connection_id)),
  362. {
  363. "$set": {**normalized, "updated_at": now},
  364. "$setOnInsert": {
  365. "bot_id": BOT_PROFILE_ID,
  366. "connection_id": str(connection_id),
  367. "created_at": now,
  368. },
  369. },
  370. upsert=True,
  371. )
  372. return await get_account_settings(connection_id)
  373. async def create_knowledge_entry(
  374. connection_id: str, values: dict[str, Any]
  375. ) -> dict[str, Any]:
  376. if not await get_business_connection(connection_id):
  377. raise AssistantDataError("connection_not_found", "未找到 Business 连接。")
  378. now = utc_now()
  379. entry = {
  380. "bot_id": BOT_PROFILE_ID,
  381. "entry_id": uuid4().hex,
  382. "connection_id": str(connection_id),
  383. "question": clean_text(values.get("question"), max_length=300, required=True),
  384. "aliases": _normalize_string_list(values.get("aliases"), max_items=20, max_length=200),
  385. "keywords": _normalize_string_list(values.get("keywords"), max_items=30, max_length=50),
  386. "answer": clean_text(
  387. values.get("answer"), max_length=4000, required=True, preserve_lines=True
  388. ),
  389. "tags": _normalize_string_list(values.get("tags"), max_items=20, max_length=30),
  390. "priority": _bounded_int(
  391. values.get("priority", 0), "知识优先级", minimum=-1000, maximum=1000
  392. ),
  393. "enabled": bool(values.get("enabled", True)),
  394. "created_at": now,
  395. "updated_at": now,
  396. }
  397. await knowledgedb.insert_one(entry)
  398. return entry
  399. async def update_knowledge_entry(entry_id: str, values: dict[str, Any]) -> dict[str, Any]:
  400. filters = _scope(entry_id=str(entry_id))
  401. current = await knowledgedb.find_one(filters)
  402. if not current:
  403. raise AssistantDataError("knowledge_not_found", "未找到知识条目。")
  404. update: dict[str, Any] = {}
  405. if "question" in values:
  406. update["question"] = clean_text(values.get("question"), max_length=300, required=True)
  407. if "aliases" in values:
  408. update["aliases"] = _normalize_string_list(
  409. values.get("aliases"), max_items=20, max_length=200
  410. )
  411. if "keywords" in values:
  412. update["keywords"] = _normalize_string_list(
  413. values.get("keywords"), max_items=30, max_length=50
  414. )
  415. if "answer" in values:
  416. update["answer"] = clean_text(
  417. values.get("answer"), max_length=4000, required=True, preserve_lines=True
  418. )
  419. if "tags" in values:
  420. update["tags"] = _normalize_string_list(values.get("tags"), max_items=20, max_length=30)
  421. if "priority" in values:
  422. update["priority"] = _bounded_int(
  423. values.get("priority"), "知识优先级", minimum=-1000, maximum=1000
  424. )
  425. if "enabled" in values:
  426. update["enabled"] = bool(values.get("enabled"))
  427. if not update:
  428. raise AssistantDataError("unchanged", "没有可保存的知识条目字段。")
  429. update["updated_at"] = utc_now()
  430. await knowledgedb.update_one(filters, {"$set": update})
  431. return await knowledgedb.find_one(filters) or current
  432. async def delete_knowledge_entry(entry_id: str) -> None:
  433. result = await knowledgedb.delete_one(_scope(entry_id=str(entry_id)))
  434. if not result.deleted_count:
  435. raise AssistantDataError("knowledge_not_found", "未找到知识条目。")
  436. async def list_knowledge_entries(
  437. connection_id: str,
  438. *,
  439. query: str = "",
  440. page: int = 1,
  441. page_size: int = 20,
  442. ) -> tuple[list[dict[str, Any]], int]:
  443. await ensure_assistant_indexes()
  444. filters: dict[str, Any] = _scope(connection_id=str(connection_id))
  445. normalized_query = clean_text(query, max_length=100)
  446. if normalized_query:
  447. pattern = re.escape(normalized_query)
  448. filters["$or"] = [
  449. {"question": {"$regex": pattern, "$options": "i"}},
  450. {"aliases": {"$regex": pattern, "$options": "i"}},
  451. {"keywords": {"$regex": pattern, "$options": "i"}},
  452. {"tags": {"$regex": pattern, "$options": "i"}},
  453. ]
  454. page = max(1, int(page))
  455. page_size = max(1, min(int(page_size), 100))
  456. total = await knowledgedb.count_documents(filters)
  457. cursor = (
  458. knowledgedb.find(filters)
  459. .sort([("priority", DESCENDING), ("updated_at", DESCENDING)])
  460. .skip((page - 1) * page_size)
  461. .limit(page_size)
  462. )
  463. return [item async for item in cursor], total
  464. def _search_text(value: str) -> str:
  465. return re.sub(r"[^\w\u3400-\u9fff]+", "", value.casefold())
  466. async def match_knowledge(
  467. connection_id: str, query: str, *, limit: int = 8
  468. ) -> list[dict[str, Any]]:
  469. await ensure_assistant_indexes()
  470. normalized_query = _search_text(query)
  471. if not normalized_query:
  472. return []
  473. entries = [
  474. item
  475. async for item in knowledgedb.find(
  476. _scope(connection_id=str(connection_id), enabled=True)
  477. )
  478. ]
  479. scored: list[tuple[int, dict[str, Any]]] = []
  480. for entry in entries:
  481. score = int(entry.get("priority") or 0)
  482. phrases = [entry.get("question", ""), *(entry.get("aliases") or [])]
  483. for phrase in phrases:
  484. normalized_phrase = _search_text(str(phrase))
  485. if not normalized_phrase:
  486. continue
  487. if normalized_query == normalized_phrase:
  488. score += 1000
  489. elif normalized_phrase in normalized_query:
  490. score += 300 + min(len(normalized_phrase), 100)
  491. elif normalized_query in normalized_phrase and len(normalized_query) >= 4:
  492. score += 120
  493. for keyword in entry.get("keywords") or []:
  494. normalized_keyword = _search_text(str(keyword))
  495. if normalized_keyword and normalized_keyword in normalized_query:
  496. score += 80
  497. if score > int(entry.get("priority") or 0):
  498. scored.append((score, entry))
  499. scored.sort(key=lambda item: (item[0], item[1].get("priority", 0)), reverse=True)
  500. return [{**entry, "match_score": score} for score, entry in scored[: max(1, limit)]]
  501. def classify_handoff(text: str) -> str | None:
  502. normalized = text.casefold()
  503. if any(term in normalized for term in HANDOFF_TERMS):
  504. return "customer_requested_human"
  505. if any(term in normalized for term in SENSITIVE_TERMS):
  506. return "sensitive_request"
  507. return None
  508. def conversation_id_for(connection_id: str, chat_id: int) -> str:
  509. digest = hashlib.sha256(f"{BOT_PROFILE_ID}:{connection_id}:{int(chat_id)}".encode()).hexdigest()
  510. return digest[:32]
  511. async def get_or_create_conversation(
  512. connection_id: str,
  513. chat_id: int,
  514. *,
  515. customer: dict[str, Any] | None = None,
  516. ) -> dict[str, Any]:
  517. await ensure_assistant_indexes()
  518. conversation_id = conversation_id_for(connection_id, chat_id)
  519. now = utc_now()
  520. update: dict[str, Any] = {"updated_at": now, "last_message_at": now}
  521. if customer:
  522. update["customer"] = _public_user(customer)
  523. await conversationsdb.update_one(
  524. _scope(conversation_id=conversation_id),
  525. {
  526. "$set": update,
  527. "$setOnInsert": {
  528. "bot_id": BOT_PROFILE_ID,
  529. "conversation_id": conversation_id,
  530. "connection_id": str(connection_id),
  531. "chat_id": int(chat_id),
  532. "status": "auto",
  533. "summary": "",
  534. "handoff_reason": "",
  535. "created_at": now,
  536. },
  537. },
  538. upsert=True,
  539. )
  540. return await conversationsdb.find_one(_scope(conversation_id=conversation_id)) or {}
  541. async def get_conversation(conversation_id: str) -> dict[str, Any] | None:
  542. await ensure_assistant_indexes()
  543. return await conversationsdb.find_one(
  544. _scope(conversation_id=str(conversation_id))
  545. )
  546. async def get_conversation_by_chat(
  547. connection_id: str, chat_id: int
  548. ) -> dict[str, Any] | None:
  549. await ensure_assistant_indexes()
  550. return await conversationsdb.find_one(
  551. _scope(connection_id=str(connection_id), chat_id=int(chat_id))
  552. )
  553. async def append_conversation_message(
  554. conversation_id: str,
  555. *,
  556. direction: str,
  557. telegram_message_id: int,
  558. text: str,
  559. sender_id: int = 0,
  560. metadata: dict[str, Any] | None = None,
  561. ) -> dict[str, Any]:
  562. await ensure_assistant_indexes()
  563. if direction not in {"incoming", "assistant", "human"}:
  564. raise AssistantDataError("invalid_direction", "会话消息方向无效。")
  565. now = utc_now()
  566. message_key = f"{conversation_id}:{direction}:{int(telegram_message_id)}"
  567. document = {
  568. "bot_id": BOT_PROFILE_ID,
  569. "message_key": message_key,
  570. "conversation_id": str(conversation_id),
  571. "direction": direction,
  572. "telegram_message_id": int(telegram_message_id),
  573. "sender_id": int(sender_id or 0),
  574. "text": clean_text(text, max_length=12000, preserve_lines=True),
  575. "metadata": metadata or {},
  576. "created_at": now,
  577. "expires_at": now + timedelta(days=30),
  578. }
  579. try:
  580. await messagesdb.insert_one(document)
  581. except DuplicateKeyError:
  582. return await messagesdb.find_one(_scope(message_key=message_key)) or document
  583. await conversationsdb.update_one(
  584. _scope(conversation_id=str(conversation_id)),
  585. {"$set": {"last_message_at": now, "updated_at": now}},
  586. )
  587. return document
  588. async def recent_conversation_messages(
  589. conversation_id: str, *, limit: int = 12
  590. ) -> list[dict[str, Any]]:
  591. cursor = (
  592. messagesdb.find(_scope(conversation_id=str(conversation_id)))
  593. .sort([("created_at", DESCENDING), ("_id", DESCENDING)])
  594. .limit(max(1, min(int(limit), 50)))
  595. )
  596. items = [item async for item in cursor]
  597. items.reverse()
  598. return items
  599. async def set_conversation_handoff(
  600. conversation_id: str,
  601. reason: str,
  602. *,
  603. summary: str = "",
  604. ) -> dict[str, Any]:
  605. now = utc_now()
  606. update: dict[str, Any] = {
  607. "status": "handoff",
  608. "handoff_reason": clean_text(reason, max_length=200, required=True),
  609. "paused_until": None,
  610. "handoff_at": now,
  611. "updated_at": now,
  612. }
  613. if summary:
  614. update["summary"] = clean_text(summary, max_length=4000, preserve_lines=True)
  615. result = await conversationsdb.find_one_and_update(
  616. _scope(conversation_id=str(conversation_id)),
  617. {"$set": update},
  618. return_document=True,
  619. )
  620. if not result:
  621. raise AssistantDataError("conversation_not_found", "未找到客户会话。")
  622. return result
  623. async def pause_conversation_for_human(
  624. conversation_id: str, *, hours: int
  625. ) -> dict[str, Any]:
  626. now = utc_now()
  627. result = await conversationsdb.find_one_and_update(
  628. _scope(conversation_id=str(conversation_id)),
  629. {
  630. "$set": {
  631. "status": "human_paused",
  632. "handoff_reason": "human_reply_detected",
  633. "paused_until": now + timedelta(hours=max(1, min(int(hours), 168))),
  634. "updated_at": now,
  635. }
  636. },
  637. return_document=True,
  638. )
  639. if not result:
  640. raise AssistantDataError("conversation_not_found", "未找到客户会话。")
  641. return result
  642. async def resume_conversation(conversation_id: str) -> dict[str, Any]:
  643. result = await conversationsdb.find_one_and_update(
  644. _scope(conversation_id=str(conversation_id)),
  645. {
  646. "$set": {
  647. "status": "auto",
  648. "handoff_reason": "",
  649. "paused_until": None,
  650. "updated_at": utc_now(),
  651. }
  652. },
  653. return_document=True,
  654. )
  655. if not result:
  656. raise AssistantDataError("conversation_not_found", "未找到客户会话。")
  657. return result
  658. async def pause_conversation(conversation_id: str) -> dict[str, Any]:
  659. return await set_conversation_handoff(conversation_id, "admin_paused")
  660. async def close_conversation(conversation_id: str) -> dict[str, Any]:
  661. result = await conversationsdb.find_one_and_update(
  662. _scope(conversation_id=str(conversation_id)),
  663. {"$set": {"status": "closed", "closed_at": utc_now(), "updated_at": utc_now()}},
  664. return_document=True,
  665. )
  666. if not result:
  667. raise AssistantDataError("conversation_not_found", "未找到客户会话。")
  668. return result
  669. async def clear_conversation(conversation_id: str) -> dict[str, Any]:
  670. current = await get_conversation(conversation_id)
  671. if not current:
  672. raise AssistantDataError("conversation_not_found", "未找到客户会话。")
  673. await messagesdb.delete_many(_scope(conversation_id=str(conversation_id)))
  674. await conversationsdb.update_one(
  675. _scope(conversation_id=str(conversation_id)),
  676. {
  677. "$set": {
  678. "summary": "",
  679. "handoff_reason": "",
  680. "status": "auto",
  681. "paused_until": None,
  682. "cleared_at": utc_now(),
  683. "updated_at": utc_now(),
  684. }
  685. },
  686. )
  687. return await get_conversation(conversation_id) or current
  688. async def update_conversation_summary(
  689. conversation_id: str, summary: str
  690. ) -> None:
  691. cleaned = clean_text(summary, max_length=4000, preserve_lines=True)
  692. if cleaned:
  693. await conversationsdb.update_one(
  694. _scope(conversation_id=str(conversation_id)),
  695. {"$set": {"summary": cleaned, "updated_at": utc_now()}},
  696. )
  697. async def list_conversations(
  698. *,
  699. connection_id: str = "",
  700. status: str = "",
  701. page: int = 1,
  702. page_size: int = 20,
  703. ) -> tuple[list[dict[str, Any]], int]:
  704. await ensure_assistant_indexes()
  705. filters: dict[str, Any] = _scope()
  706. if connection_id:
  707. filters["connection_id"] = str(connection_id)
  708. if status:
  709. if status not in {"auto", "handoff", "human_paused", "closed"}:
  710. raise AssistantDataError("invalid_status", "会话状态无效。")
  711. filters["status"] = status
  712. page = max(1, int(page))
  713. page_size = max(1, min(int(page_size), 100))
  714. total = await conversationsdb.count_documents(filters)
  715. cursor = (
  716. conversationsdb.find(filters)
  717. .sort("updated_at", DESCENDING)
  718. .skip((page - 1) * page_size)
  719. .limit(page_size)
  720. )
  721. return [item async for item in cursor], total
  722. async def conversation_detail(conversation_id: str) -> dict[str, Any]:
  723. conversation = await get_conversation(conversation_id)
  724. if not conversation:
  725. raise AssistantDataError("conversation_not_found", "未找到客户会话。")
  726. conversation["messages"] = await recent_conversation_messages(
  727. conversation_id, limit=50
  728. )
  729. return conversation
  730. def _usage_day(now: datetime, timezone_name: str) -> str:
  731. try:
  732. zone = ZoneInfo(timezone_name)
  733. except ZoneInfoNotFoundError:
  734. zone = ZoneInfo("Asia/Shanghai")
  735. return as_utc(now).astimezone(zone).date().isoformat()
  736. async def reserve_ai_usage(
  737. connection_id: str,
  738. chat_id: int,
  739. settings: dict[str, Any],
  740. *,
  741. now: datetime | None = None,
  742. ) -> tuple[bool, str]:
  743. await ensure_assistant_indexes()
  744. current_time = as_utc(now or utc_now())
  745. day = _usage_day(current_time, str(settings.get("timezone") or "Asia/Shanghai"))
  746. account_limit = int(settings.get("account_daily_limit") or 200)
  747. customer_limit = int(settings.get("customer_daily_limit") or 20)
  748. account_filter = _scope(
  749. connection_id=str(connection_id), chat_id=0, day=day
  750. )
  751. customer_filter = _scope(
  752. connection_id=str(connection_id), chat_id=int(chat_id), day=day
  753. )
  754. for filters, usage_scope in (
  755. (account_filter, "account"),
  756. (customer_filter, "customer"),
  757. ):
  758. await usagedb.update_one(
  759. filters,
  760. {
  761. "$setOnInsert": {
  762. **filters,
  763. "scope": usage_scope,
  764. "call_count": 0,
  765. "created_at": current_time,
  766. "updated_at": current_time,
  767. }
  768. },
  769. upsert=True,
  770. )
  771. account = await usagedb.find_one_and_update(
  772. {**account_filter, "call_count": {"$lt": account_limit}},
  773. {"$inc": {"call_count": 1}, "$set": {"updated_at": current_time}},
  774. return_document=True,
  775. )
  776. if not account:
  777. return False, "account_daily_limit"
  778. customer = await usagedb.find_one_and_update(
  779. {**customer_filter, "call_count": {"$lt": customer_limit}},
  780. {"$inc": {"call_count": 1}, "$set": {"updated_at": current_time}},
  781. return_document=True,
  782. )
  783. if not customer:
  784. await usagedb.update_one(
  785. {**account_filter, "call_count": {"$gt": 0}},
  786. {"$inc": {"call_count": -1}, "$set": {"updated_at": current_time}},
  787. )
  788. return False, "customer_daily_limit"
  789. return True, ""
  790. async def usage_metrics(*, connection_id: str = "", day: str = "") -> dict[str, Any]:
  791. await ensure_assistant_indexes()
  792. match = _scope()
  793. if connection_id:
  794. match["connection_id"] = str(connection_id)
  795. if day:
  796. match["day"] = str(day)
  797. calls = 0
  798. customers: set[tuple[str, int]] = set()
  799. customer_match = {
  800. **match,
  801. "scope": {"$ne": "account"},
  802. "chat_id": {"$ne": 0},
  803. "call_count": {"$gt": 0},
  804. }
  805. async for item in usagedb.find(customer_match):
  806. calls += int(item.get("call_count") or 0)
  807. customers.add((str(item.get("connection_id")), int(item.get("chat_id") or 0)))
  808. conversation_filters = _scope()
  809. if connection_id:
  810. conversation_filters["connection_id"] = str(connection_id)
  811. return {
  812. "ai_calls": calls,
  813. "customers": len(customers),
  814. "conversations": await conversationsdb.count_documents(conversation_filters),
  815. "handoffs": await conversationsdb.count_documents(
  816. {**conversation_filters, "status": "handoff"}
  817. ),
  818. "human_paused": await conversationsdb.count_documents(
  819. {**conversation_filters, "status": "human_paused"}
  820. ),
  821. }
  822. async def claim_update(update_id: int) -> bool:
  823. await ensure_assistant_indexes()
  824. now = utc_now()
  825. try:
  826. await updatesdb.insert_one(
  827. {
  828. "bot_id": BOT_PROFILE_ID,
  829. "update_id": int(update_id),
  830. "status": "processing",
  831. "attempts": 1,
  832. "created_at": now,
  833. "updated_at": now,
  834. "expires_at": now + timedelta(days=7),
  835. }
  836. )
  837. return True
  838. except DuplicateKeyError:
  839. filters = _scope(update_id=int(update_id))
  840. current = await updatesdb.find_one(filters) or {}
  841. if current.get("status") == "done":
  842. return False
  843. await updatesdb.update_one(
  844. filters,
  845. {"$inc": {"attempts": 1}, "$set": {"updated_at": now}},
  846. )
  847. return True
  848. async def mark_update_done(update_id: int) -> None:
  849. await updatesdb.update_one(
  850. _scope(update_id=int(update_id)),
  851. {"$set": {"status": "done", "updated_at": utc_now(), "last_error": ""}},
  852. )
  853. async def mark_update_failed(update_id: int, error: str) -> int:
  854. await updatesdb.update_one(
  855. _scope(update_id=int(update_id)),
  856. {"$set": {"status": "failed", "updated_at": utc_now(), "last_error": error[:1000]}},
  857. )
  858. current = await updatesdb.find_one(_scope(update_id=int(update_id))) or {}
  859. return int(current.get("attempts") or 1)
  860. async def dead_letter_update(update: dict[str, Any], error: str) -> None:
  861. update_id = int(update.get("update_id") or 0)
  862. now = utc_now()
  863. await deadlettersdb.update_one(
  864. _scope(update_id=update_id),
  865. {
  866. "$set": {
  867. "bot_id": BOT_PROFILE_ID,
  868. "payload": update,
  869. "error": error[:2000],
  870. "updated_at": now,
  871. "expires_at": now + timedelta(days=30),
  872. },
  873. "$setOnInsert": {"created_at": now},
  874. },
  875. upsert=True,
  876. )
  877. await mark_update_done(update_id)
  878. async def runtime_status() -> dict[str, Any]:
  879. await ensure_assistant_indexes()
  880. stored = await runtimedb.find_one(
  881. _scope(runtime_id="business_assistant")
  882. ) or {}
  883. return {
  884. "runtime_id": "business_assistant",
  885. "polling_state": "stopped",
  886. "business_mode_supported": False,
  887. "webhook_conflict": False,
  888. "webhook_url": "",
  889. "model_configured": False,
  890. "last_error": "",
  891. **{key: value for key, value in stored.items() if key != "_id"},
  892. }
  893. async def update_runtime_status(values: dict[str, Any]) -> dict[str, Any]:
  894. await ensure_assistant_indexes()
  895. insert_values = {
  896. "bot_id": BOT_PROFILE_ID,
  897. "runtime_id": "business_assistant",
  898. "created_at": utc_now(),
  899. }
  900. if "offset" not in values:
  901. insert_values["offset"] = 0
  902. await runtimedb.update_one(
  903. _scope(runtime_id="business_assistant"),
  904. {
  905. "$set": {**values, "updated_at": utc_now()},
  906. "$setOnInsert": insert_values,
  907. },
  908. upsert=True,
  909. )
  910. return await runtime_status()
  911. async def load_update_offset() -> int:
  912. return int((await runtime_status()).get("offset") or 0)
  913. async def save_update_offset(offset: int) -> None:
  914. await update_runtime_status({"offset": int(offset), "last_poll_at": utc_now()})
  915. async def due_digest_connections(now: datetime | None = None) -> list[dict[str, Any]]:
  916. current = as_utc(now or utc_now())
  917. results: list[dict[str, Any]] = []
  918. async for settings in settingsdb.find(
  919. _scope(digest_enabled=True, assistant_enabled=True)
  920. ):
  921. try:
  922. local = current.astimezone(ZoneInfo(str(settings.get("timezone"))))
  923. except ZoneInfoNotFoundError:
  924. local = current.astimezone(ZoneInfo("Asia/Shanghai"))
  925. day = local.date().isoformat()
  926. if local.strftime("%H:%M") != str(settings.get("digest_time") or "09:00"):
  927. continue
  928. if settings.get("last_digest_day") == day:
  929. continue
  930. connection = await get_business_connection(settings["connection_id"])
  931. if connection and connection.get("is_enabled"):
  932. results.append({"connection": connection, "settings": settings, "day": day})
  933. return results
  934. async def mark_digest_sent(connection_id: str, day: str) -> None:
  935. await settingsdb.update_one(
  936. _scope(connection_id=str(connection_id)),
  937. {"$set": {"last_digest_day": str(day), "last_digest_at": utc_now()}},
  938. )