dbassistant.py 60 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182838485868788899091929394959697989910010110210310410510610710810911011111211311411511611711811912012112212312412512612712812913013113213313413513613713813914014114214314414514614714814915015115215315415515615715815916016116216316416516616716816917017117217317417517617717817918018118218318418518618718818919019119219319419519619719819920020120220320420520620720820921021121221321421521621721821922022122222322422522622722822923023123223323423523623723823924024124224324424524624724824925025125225325425525625725825926026126226326426526626726826927027127227327427527627727827928028128228328428528628728828929029129229329429529629729829930030130230330430530630730830931031131231331431531631731831932032132232332432532632732832933033133233333433533633733833934034134234334434534634734834935035135235335435535635735835936036136236336436536636736836937037137237337437537637737837938038138238338438538638738838939039139239339439539639739839940040140240340440540640740840941041141241341441541641741841942042142242342442542642742842943043143243343443543643743843944044144244344444544644744844945045145245345445545645745845946046146246346446546646746846947047147247347447547647747847948048148248348448548648748848949049149249349449549649749849950050150250350450550650750850951051151251351451551651751851952052152252352452552652752852953053153253353453553653753853954054154254354454554654754854955055155255355455555655755855956056156256356456556656756856957057157257357457557657757857958058158258358458558658758858959059159259359459559659759859960060160260360460560660760860961061161261361461561661761861962062162262362462562662762862963063163263363463563663763863964064164264364464564664764864965065165265365465565665765865966066166266366466566666766866967067167267367467567667767867968068168268368468568668768868969069169269369469569669769869970070170270370470570670770870971071171271371471571671771871972072172272372472572672772872973073173273373473573673773873974074174274374474574674774874975075175275375475575675775875976076176276376476576676776876977077177277377477577677777877978078178278378478578678778878979079179279379479579679779879980080180280380480580680780880981081181281381481581681781881982082182282382482582682782882983083183283383483583683783883984084184284384484584684784884985085185285385485585685785885986086186286386486586686786886987087187287387487587687787887988088188288388488588688788888989089189289389489589689789889990090190290390490590690790890991091191291391491591691791891992092192292392492592692792892993093193293393493593693793893994094194294394494594694794894995095195295395495595695795895996096196296396496596696796896997097197297397497597697797897998098198298398498598698798898999099199299399499599699799899910001001100210031004100510061007100810091010101110121013101410151016101710181019102010211022102310241025102610271028102910301031103210331034103510361037103810391040104110421043104410451046104710481049105010511052105310541055105610571058105910601061106210631064106510661067106810691070107110721073107410751076107710781079108010811082108310841085108610871088108910901091109210931094109510961097109810991100110111021103110411051106110711081109111011111112111311141115111611171118111911201121112211231124112511261127112811291130113111321133113411351136113711381139114011411142114311441145114611471148114911501151115211531154115511561157115811591160116111621163116411651166116711681169117011711172117311741175117611771178117911801181118211831184118511861187118811891190119111921193119411951196119711981199120012011202120312041205120612071208120912101211121212131214121512161217121812191220122112221223122412251226122712281229123012311232123312341235123612371238123912401241124212431244124512461247124812491250125112521253125412551256125712581259126012611262126312641265126612671268126912701271127212731274127512761277127812791280128112821283128412851286128712881289129012911292129312941295129612971298129913001301130213031304130513061307130813091310131113121313131413151316131713181319132013211322132313241325132613271328132913301331133213331334133513361337133813391340134113421343134413451346134713481349135013511352135313541355135613571358135913601361136213631364136513661367136813691370137113721373137413751376137713781379138013811382138313841385138613871388138913901391139213931394139513961397139813991400140114021403140414051406140714081409141014111412141314141415141614171418141914201421142214231424142514261427142814291430143114321433143414351436143714381439144014411442144314441445144614471448144914501451145214531454145514561457145814591460146114621463146414651466146714681469147014711472147314741475147614771478147914801481148214831484148514861487148814891490149114921493149414951496149714981499150015011502150315041505150615071508150915101511151215131514151515161517151815191520152115221523152415251526152715281529153015311532153315341535153615371538153915401541154215431544154515461547154815491550155115521553155415551556155715581559156015611562156315641565156615671568156915701571157215731574157515761577157815791580158115821583158415851586158715881589159015911592159315941595159615971598159916001601160216031604160516061607160816091610161116121613161416151616161716181619162016211622162316241625162616271628162916301631163216331634163516361637163816391640164116421643164416451646164716481649165016511652165316541655165616571658165916601661166216631664166516661667166816691670167116721673167416751676167716781679168016811682168316841685168616871688168916901691169216931694169516961697169816991700170117021703
  1. from __future__ import annotations
  2. import asyncio
  3. import hashlib
  4. import json
  5. import re
  6. from datetime import UTC, datetime, timedelta
  7. from typing import Any
  8. from uuid import uuid4
  9. from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
  10. from pymongo import ASCENDING, DESCENDING
  11. from pymongo.errors import DuplicateKeyError
  12. from wbb import BOT_PROFILE_ID, db
  13. connectionsdb = db.business_assistant_connections
  14. settingsdb = db.business_assistant_settings
  15. knowledgedb = db.business_assistant_knowledge
  16. knowledge_sourcesdb = db.business_assistant_knowledge_sources
  17. source_eventsdb = db.business_assistant_source_events
  18. knowledge_candidatesdb = db.business_assistant_knowledge_candidates
  19. conversationsdb = db.business_assistant_conversations
  20. messagesdb = db.business_assistant_messages
  21. usagedb = db.business_assistant_usage
  22. updatesdb = db.business_assistant_updates
  23. deadlettersdb = db.business_assistant_dead_letters
  24. runtimedb = db.business_assistant_runtime
  25. DEFAULT_ACCOUNT_SETTINGS: dict[str, Any] = {
  26. "assistant_enabled": False,
  27. "system_prompt": (
  28. "你是该 Telegram 账号的智能接待秘书。回答要简洁、礼貌,只能依据提供的知识条目陈述业务事实。"
  29. ),
  30. "language": "zh-CN",
  31. "tone": "professional",
  32. "account_daily_limit": 200,
  33. "customer_daily_limit": 20,
  34. "human_pause_hours": 24,
  35. "notification_destination": "owner",
  36. "ops_group_id": 0,
  37. "timezone": "Asia/Shanghai",
  38. "digest_enabled": False,
  39. "digest_time": "09:00",
  40. "handoff_message": "这个问题需要人工确认,我已经通知负责人,请稍候。",
  41. "unsupported_message": "已收到你的消息,这类内容需要人工处理,我已经通知负责人。",
  42. }
  43. VALID_NOTIFICATION_DESTINATIONS = {"owner", "ops", "both"}
  44. VALID_TONES = {"professional", "friendly", "concise"}
  45. VALID_SOURCE_TYPES = {"channel", "group", "business"}
  46. VALID_PUBLICATION_MODES = {"review", "auto"}
  47. VALID_AUTHOR_POLICIES = {"admins_or_allowlist", "admins", "allowlist", "all"}
  48. VALID_CANDIDATE_STATUSES = {"pending", "published", "rejected", "stale"}
  49. HANDOFF_TERMS = (
  50. "人工",
  51. "真人",
  52. "客服",
  53. "负责人",
  54. "转人工",
  55. "human",
  56. "agent",
  57. )
  58. SENSITIVE_TERMS = (
  59. "承诺",
  60. "保证",
  61. "投诉",
  62. "退款",
  63. "退钱",
  64. "赔偿",
  65. "律师",
  66. "起诉",
  67. "支付失败",
  68. "账号被盗",
  69. "密码",
  70. "验证码",
  71. )
  72. _index_lock = asyncio.Lock()
  73. _indexes_ready = False
  74. class AssistantDataError(ValueError):
  75. def __init__(self, code: str, message: str):
  76. super().__init__(message)
  77. self.code = code
  78. def utc_now() -> datetime:
  79. return datetime.now(UTC)
  80. def as_utc(value: datetime) -> datetime:
  81. return value.replace(tzinfo=UTC) if value.tzinfo is None else value.astimezone(UTC)
  82. def _scope(filters: dict[str, Any] | None = None, /, **values: Any) -> dict[str, Any]:
  83. return {"bot_id": BOT_PROFILE_ID, **(filters or {}), **values}
  84. def clean_text(
  85. value: Any,
  86. *,
  87. max_length: int,
  88. required: bool = False,
  89. preserve_lines: bool = False,
  90. ) -> str:
  91. text = re.sub(r"[\x00-\x08\x0b\x0c\x0e-\x1f\x7f]", "", str(value or ""))
  92. text = text.strip()
  93. if not preserve_lines:
  94. text = " ".join(text.split())
  95. if required and not text:
  96. raise AssistantDataError("required_field", "必填内容不能为空。")
  97. if len(text) > max_length:
  98. raise AssistantDataError("text_too_long", f"内容不能超过 {max_length} 个字符。")
  99. return text
  100. def _bounded_int(
  101. value: Any,
  102. name: str,
  103. *,
  104. minimum: int,
  105. maximum: int,
  106. ) -> int:
  107. try:
  108. parsed = int(value)
  109. except (TypeError, ValueError) as exc:
  110. raise AssistantDataError("invalid_setting", f"{name} 必须是整数。") from exc
  111. if not minimum <= parsed <= maximum:
  112. raise AssistantDataError(
  113. "invalid_setting",
  114. f"{name} 必须在 {minimum} 到 {maximum} 之间。",
  115. )
  116. return parsed
  117. def _normalize_string_list(value: Any, *, max_items: int, max_length: int) -> list[str]:
  118. values = value if isinstance(value, (list, tuple, set)) else str(value or "").split(",")
  119. normalized: list[str] = []
  120. seen: set[str] = set()
  121. for item in values:
  122. text = clean_text(item, max_length=max_length)
  123. key = text.casefold()
  124. if text and key not in seen:
  125. normalized.append(text)
  126. seen.add(key)
  127. if len(normalized) > max_items:
  128. raise AssistantDataError("too_many_items", f"最多允许 {max_items} 项。")
  129. return normalized
  130. def _normalize_int_list(value: Any, *, max_items: int = 100) -> list[int]:
  131. values = value if isinstance(value, (list, tuple, set)) else str(value or "").split(",")
  132. normalized: list[int] = []
  133. for item in values:
  134. if item in {None, ""}:
  135. continue
  136. try:
  137. parsed = int(item)
  138. except (TypeError, ValueError) as exc:
  139. raise AssistantDataError("invalid_source", "白名单用户 ID 必须是整数。") from exc
  140. if parsed and parsed not in normalized:
  141. normalized.append(parsed)
  142. if len(normalized) > max_items:
  143. raise AssistantDataError("too_many_items", f"最多允许 {max_items} 项。")
  144. return normalized
  145. async def ensure_assistant_indexes() -> None:
  146. global _indexes_ready
  147. if _indexes_ready:
  148. return
  149. async with _index_lock:
  150. if _indexes_ready:
  151. return
  152. await connectionsdb.create_index(
  153. [("bot_id", ASCENDING), ("connection_id", ASCENDING)], unique=True
  154. )
  155. await connectionsdb.create_index(
  156. [("bot_id", ASCENDING), ("updated_at", DESCENDING)]
  157. )
  158. await settingsdb.create_index(
  159. [("bot_id", ASCENDING), ("connection_id", ASCENDING)], unique=True
  160. )
  161. await knowledgedb.create_index(
  162. [("bot_id", ASCENDING), ("entry_id", ASCENDING)], unique=True
  163. )
  164. await knowledge_sourcesdb.create_index(
  165. [("bot_id", ASCENDING), ("source_id", ASCENDING)], unique=True
  166. )
  167. await knowledge_sourcesdb.create_index(
  168. [("bot_id", ASCENDING), ("source_key", ASCENDING)], unique=True
  169. )
  170. await knowledge_sourcesdb.create_index(
  171. [
  172. ("bot_id", ASCENDING),
  173. ("connection_id", ASCENDING),
  174. ("enabled", ASCENDING),
  175. ("updated_at", DESCENDING),
  176. ]
  177. )
  178. await source_eventsdb.create_index(
  179. [("bot_id", ASCENDING), ("event_id", ASCENDING)], unique=True
  180. )
  181. await source_eventsdb.create_index(
  182. [
  183. ("bot_id", ASCENDING),
  184. ("source_id", ASCENDING),
  185. ("event_key", ASCENDING),
  186. ],
  187. unique=True,
  188. )
  189. await source_eventsdb.create_index("expires_at", expireAfterSeconds=0)
  190. await knowledge_candidatesdb.create_index(
  191. [("bot_id", ASCENDING), ("candidate_id", ASCENDING)], unique=True
  192. )
  193. await knowledge_candidatesdb.create_index(
  194. [
  195. ("bot_id", ASCENDING),
  196. ("event_id", ASCENDING),
  197. ("item_index", ASCENDING),
  198. ],
  199. unique=True,
  200. )
  201. await knowledge_candidatesdb.create_index(
  202. [
  203. ("bot_id", ASCENDING),
  204. ("connection_id", ASCENDING),
  205. ("status", ASCENDING),
  206. ("updated_at", DESCENDING),
  207. ]
  208. )
  209. await knowledgedb.create_index(
  210. [
  211. ("bot_id", ASCENDING),
  212. ("connection_id", ASCENDING),
  213. ("enabled", ASCENDING),
  214. ("priority", DESCENDING),
  215. ]
  216. )
  217. await conversationsdb.create_index(
  218. [("bot_id", ASCENDING), ("conversation_id", ASCENDING)], unique=True
  219. )
  220. await conversationsdb.create_index(
  221. [
  222. ("bot_id", ASCENDING),
  223. ("connection_id", ASCENDING),
  224. ("chat_id", ASCENDING),
  225. ],
  226. unique=True,
  227. )
  228. await conversationsdb.create_index(
  229. [
  230. ("bot_id", ASCENDING),
  231. ("connection_id", ASCENDING),
  232. ("status", ASCENDING),
  233. ("updated_at", DESCENDING),
  234. ]
  235. )
  236. await messagesdb.create_index(
  237. [("bot_id", ASCENDING), ("message_key", ASCENDING)], unique=True
  238. )
  239. await messagesdb.create_index(
  240. [
  241. ("bot_id", ASCENDING),
  242. ("conversation_id", ASCENDING),
  243. ("created_at", ASCENDING),
  244. ]
  245. )
  246. await messagesdb.create_index("expires_at", expireAfterSeconds=0)
  247. await usagedb.create_index(
  248. [
  249. ("bot_id", ASCENDING),
  250. ("connection_id", ASCENDING),
  251. ("chat_id", ASCENDING),
  252. ("day", ASCENDING),
  253. ],
  254. unique=True,
  255. )
  256. await updatesdb.create_index(
  257. [("bot_id", ASCENDING), ("update_id", ASCENDING)], unique=True
  258. )
  259. await updatesdb.create_index("expires_at", expireAfterSeconds=0)
  260. await deadlettersdb.create_index(
  261. [("bot_id", ASCENDING), ("update_id", ASCENDING)], unique=True
  262. )
  263. await deadlettersdb.create_index(
  264. [("bot_id", ASCENDING), ("created_at", DESCENDING)]
  265. )
  266. await deadlettersdb.create_index("expires_at", expireAfterSeconds=0)
  267. await runtimedb.create_index(
  268. [("bot_id", ASCENDING), ("runtime_id", ASCENDING)], unique=True
  269. )
  270. _indexes_ready = True
  271. def _public_user(user: Any) -> dict[str, Any]:
  272. value = user if isinstance(user, dict) else {}
  273. return {
  274. "id": int(value.get("id") or 0),
  275. "username": clean_text(value.get("username"), max_length=64),
  276. "first_name": clean_text(value.get("first_name"), max_length=128),
  277. "last_name": clean_text(value.get("last_name"), max_length=128),
  278. }
  279. async def upsert_business_connection(payload: dict[str, Any]) -> dict[str, Any]:
  280. await ensure_assistant_indexes()
  281. connection_id = clean_text(payload.get("id"), max_length=128, required=True)
  282. now = utc_now()
  283. rights = payload.get("rights") if isinstance(payload.get("rights"), dict) else {}
  284. document = {
  285. "bot_id": BOT_PROFILE_ID,
  286. "connection_id": connection_id,
  287. "user": _public_user(payload.get("user")),
  288. "user_chat_id": int(payload.get("user_chat_id") or 0),
  289. "rights": {str(key): bool(value) for key, value in rights.items() if value is True},
  290. "is_enabled": bool(payload.get("is_enabled")),
  291. "connected_at": datetime.fromtimestamp(int(payload.get("date") or 0), UTC)
  292. if payload.get("date")
  293. else now,
  294. "last_event_at": now,
  295. "updated_at": now,
  296. }
  297. await connectionsdb.update_one(
  298. _scope(connection_id=connection_id),
  299. {"$set": document, "$setOnInsert": {"created_at": now}},
  300. upsert=True,
  301. )
  302. await settingsdb.update_one(
  303. _scope(connection_id=connection_id),
  304. {
  305. "$setOnInsert": {
  306. "bot_id": BOT_PROFILE_ID,
  307. "connection_id": connection_id,
  308. **DEFAULT_ACCOUNT_SETTINGS,
  309. "created_at": now,
  310. "updated_at": now,
  311. }
  312. },
  313. upsert=True,
  314. )
  315. return await get_business_connection(connection_id) or document
  316. async def get_business_connection(connection_id: str) -> dict[str, Any] | None:
  317. await ensure_assistant_indexes()
  318. return await connectionsdb.find_one(_scope(connection_id=str(connection_id)))
  319. async def touch_business_connection(
  320. connection_id: str, *, error: str = "", is_enabled: bool | None = None
  321. ) -> dict[str, Any] | None:
  322. now = utc_now()
  323. values: dict[str, Any] = {
  324. "last_event_at": now,
  325. "updated_at": now,
  326. "last_error": clean_text(error, max_length=1000),
  327. }
  328. if is_enabled is not None:
  329. values["is_enabled"] = bool(is_enabled)
  330. await connectionsdb.update_one(
  331. _scope(connection_id=str(connection_id)),
  332. {"$set": values},
  333. )
  334. return await get_business_connection(connection_id)
  335. async def list_business_connections(
  336. *, page: int = 1, page_size: int = 20
  337. ) -> tuple[list[dict[str, Any]], int]:
  338. await ensure_assistant_indexes()
  339. page = max(1, int(page))
  340. page_size = max(1, min(int(page_size), 100))
  341. filters = _scope()
  342. total = await connectionsdb.count_documents(filters)
  343. cursor = (
  344. connectionsdb.find(filters)
  345. .sort("updated_at", DESCENDING)
  346. .skip((page - 1) * page_size)
  347. .limit(page_size)
  348. )
  349. items = []
  350. async for item in cursor:
  351. item["settings"] = await get_account_settings(item["connection_id"])
  352. items.append(item)
  353. return items, total
  354. async def get_account_settings(connection_id: str) -> dict[str, Any]:
  355. await ensure_assistant_indexes()
  356. stored = await settingsdb.find_one(_scope(connection_id=str(connection_id))) or {}
  357. return {
  358. **DEFAULT_ACCOUNT_SETTINGS,
  359. **{key: value for key, value in stored.items() if key != "_id"},
  360. "connection_id": str(connection_id),
  361. }
  362. def normalize_account_settings(
  363. values: dict[str, Any], *, previous: dict[str, Any] | None = None
  364. ) -> dict[str, Any]:
  365. current = {**DEFAULT_ACCOUNT_SETTINGS, **(previous or {})}
  366. if "assistant_enabled" in values:
  367. current["assistant_enabled"] = bool(values.get("assistant_enabled"))
  368. if "system_prompt" in values:
  369. current["system_prompt"] = clean_text(
  370. values.get("system_prompt"), max_length=4000, required=True, preserve_lines=True
  371. )
  372. if "language" in values:
  373. current["language"] = clean_text(values.get("language"), max_length=32, required=True)
  374. if "tone" in values:
  375. tone = str(values.get("tone") or "")
  376. if tone not in VALID_TONES:
  377. raise AssistantDataError("invalid_setting", "接待语气无效。")
  378. current["tone"] = tone
  379. if "account_daily_limit" in values:
  380. current["account_daily_limit"] = _bounded_int(
  381. values.get("account_daily_limit"), "账号每日额度", minimum=1, maximum=100000
  382. )
  383. if "customer_daily_limit" in values:
  384. current["customer_daily_limit"] = _bounded_int(
  385. values.get("customer_daily_limit"), "客户每日额度", minimum=1, maximum=1000
  386. )
  387. if "human_pause_hours" in values:
  388. current["human_pause_hours"] = _bounded_int(
  389. values.get("human_pause_hours"), "人工暂停小时数", minimum=1, maximum=168
  390. )
  391. if "notification_destination" in values:
  392. destination = str(values.get("notification_destination") or "")
  393. if destination not in VALID_NOTIFICATION_DESTINATIONS:
  394. raise AssistantDataError("invalid_setting", "通知目的地无效。")
  395. current["notification_destination"] = destination
  396. if "ops_group_id" in values:
  397. try:
  398. current["ops_group_id"] = int(values.get("ops_group_id") or 0)
  399. except (TypeError, ValueError) as exc:
  400. raise AssistantDataError("invalid_setting", "运营群 ID 必须是整数。") from exc
  401. if "timezone" in values:
  402. timezone_name = clean_text(values.get("timezone"), max_length=64, required=True)
  403. try:
  404. ZoneInfo(timezone_name)
  405. except ZoneInfoNotFoundError as exc:
  406. raise AssistantDataError("invalid_setting", "时区名称无效。") from exc
  407. current["timezone"] = timezone_name
  408. if "digest_enabled" in values:
  409. current["digest_enabled"] = bool(values.get("digest_enabled"))
  410. if "digest_time" in values:
  411. digest_time = str(values.get("digest_time") or "")
  412. if not re.fullmatch(r"(?:[01]\d|2[0-3]):[0-5]\d", digest_time):
  413. raise AssistantDataError("invalid_setting", "简报时间必须是 HH:MM。")
  414. current["digest_time"] = digest_time
  415. for key, max_length in (("handoff_message", 1000), ("unsupported_message", 1000)):
  416. if key in values:
  417. current[key] = clean_text(values.get(key), max_length=max_length, required=True)
  418. return {key: current[key] for key in DEFAULT_ACCOUNT_SETTINGS}
  419. async def update_account_settings(
  420. connection_id: str, values: dict[str, Any]
  421. ) -> dict[str, Any]:
  422. connection = await get_business_connection(connection_id)
  423. if not connection:
  424. raise AssistantDataError("connection_not_found", "未找到 Business 连接。")
  425. previous = await get_account_settings(connection_id)
  426. normalized = normalize_account_settings(values, previous=previous)
  427. now = utc_now()
  428. await settingsdb.update_one(
  429. _scope(connection_id=str(connection_id)),
  430. {
  431. "$set": {**normalized, "updated_at": now},
  432. "$setOnInsert": {
  433. "bot_id": BOT_PROFILE_ID,
  434. "connection_id": str(connection_id),
  435. "created_at": now,
  436. },
  437. },
  438. upsert=True,
  439. )
  440. return await get_account_settings(connection_id)
  441. async def create_knowledge_entry(
  442. connection_id: str, values: dict[str, Any]
  443. ) -> dict[str, Any]:
  444. if not await get_business_connection(connection_id):
  445. raise AssistantDataError("connection_not_found", "未找到 Business 连接。")
  446. now = utc_now()
  447. entry = {
  448. "bot_id": BOT_PROFILE_ID,
  449. "entry_id": uuid4().hex,
  450. "connection_id": str(connection_id),
  451. "question": clean_text(values.get("question"), max_length=300, required=True),
  452. "aliases": _normalize_string_list(values.get("aliases"), max_items=20, max_length=200),
  453. "keywords": _normalize_string_list(values.get("keywords"), max_items=30, max_length=50),
  454. "answer": clean_text(
  455. values.get("answer"), max_length=4000, required=True, preserve_lines=True
  456. ),
  457. "tags": _normalize_string_list(values.get("tags"), max_items=20, max_length=30),
  458. "priority": _bounded_int(
  459. values.get("priority", 0), "知识优先级", minimum=-1000, maximum=1000
  460. ),
  461. "enabled": bool(values.get("enabled", True)),
  462. "source_candidate_id": clean_text(
  463. values.get("source_candidate_id"), max_length=64
  464. ),
  465. "source_references": values.get("source_references")
  466. if isinstance(values.get("source_references"), list)
  467. else [],
  468. "created_at": now,
  469. "updated_at": now,
  470. }
  471. await knowledgedb.insert_one(entry)
  472. return entry
  473. async def update_knowledge_entry(entry_id: str, values: dict[str, Any]) -> dict[str, Any]:
  474. filters = _scope(entry_id=str(entry_id))
  475. current = await knowledgedb.find_one(filters)
  476. if not current:
  477. raise AssistantDataError("knowledge_not_found", "未找到知识条目。")
  478. update: dict[str, Any] = {}
  479. if "question" in values:
  480. update["question"] = clean_text(values.get("question"), max_length=300, required=True)
  481. if "aliases" in values:
  482. update["aliases"] = _normalize_string_list(
  483. values.get("aliases"), max_items=20, max_length=200
  484. )
  485. if "keywords" in values:
  486. update["keywords"] = _normalize_string_list(
  487. values.get("keywords"), max_items=30, max_length=50
  488. )
  489. if "answer" in values:
  490. update["answer"] = clean_text(
  491. values.get("answer"), max_length=4000, required=True, preserve_lines=True
  492. )
  493. if "tags" in values:
  494. update["tags"] = _normalize_string_list(values.get("tags"), max_items=20, max_length=30)
  495. if "priority" in values:
  496. update["priority"] = _bounded_int(
  497. values.get("priority"), "知识优先级", minimum=-1000, maximum=1000
  498. )
  499. if "enabled" in values:
  500. update["enabled"] = bool(values.get("enabled"))
  501. if "source_candidate_id" in values:
  502. update["source_candidate_id"] = clean_text(
  503. values.get("source_candidate_id"), max_length=64
  504. )
  505. if "source_references" in values:
  506. references = values.get("source_references")
  507. if not isinstance(references, list):
  508. raise AssistantDataError("invalid_knowledge", "知识来源引用必须是数组。")
  509. update["source_references"] = references[:20]
  510. if not update:
  511. raise AssistantDataError("unchanged", "没有可保存的知识条目字段。")
  512. update["updated_at"] = utc_now()
  513. await knowledgedb.update_one(filters, {"$set": update})
  514. return await knowledgedb.find_one(filters) or current
  515. async def delete_knowledge_entry(entry_id: str) -> None:
  516. result = await knowledgedb.delete_one(_scope(entry_id=str(entry_id)))
  517. if not result.deleted_count:
  518. raise AssistantDataError("knowledge_not_found", "未找到知识条目。")
  519. async def list_knowledge_entries(
  520. connection_id: str,
  521. *,
  522. query: str = "",
  523. page: int = 1,
  524. page_size: int = 20,
  525. ) -> tuple[list[dict[str, Any]], int]:
  526. await ensure_assistant_indexes()
  527. filters: dict[str, Any] = _scope(connection_id=str(connection_id))
  528. normalized_query = clean_text(query, max_length=100)
  529. if normalized_query:
  530. pattern = re.escape(normalized_query)
  531. filters["$or"] = [
  532. {"question": {"$regex": pattern, "$options": "i"}},
  533. {"aliases": {"$regex": pattern, "$options": "i"}},
  534. {"keywords": {"$regex": pattern, "$options": "i"}},
  535. {"tags": {"$regex": pattern, "$options": "i"}},
  536. ]
  537. page = max(1, int(page))
  538. page_size = max(1, min(int(page_size), 100))
  539. total = await knowledgedb.count_documents(filters)
  540. cursor = (
  541. knowledgedb.find(filters)
  542. .sort([("priority", DESCENDING), ("updated_at", DESCENDING)])
  543. .skip((page - 1) * page_size)
  544. .limit(page_size)
  545. )
  546. return [item async for item in cursor], total
  547. def _knowledge_source_key(connection_id: str, source_type: str, chat_id: int) -> str:
  548. return f"{connection_id}:{source_type}:{int(chat_id)}"
  549. def normalize_knowledge_source(
  550. values: dict[str, Any], *, previous: dict[str, Any] | None = None
  551. ) -> dict[str, Any]:
  552. current = {**(previous or {}), **values}
  553. source_type = str(current.get("source_type") or "")
  554. if source_type not in VALID_SOURCE_TYPES:
  555. raise AssistantDataError("invalid_source", "知识来源类型无效。")
  556. try:
  557. chat_id = int(current.get("chat_id") or 0)
  558. except (TypeError, ValueError) as exc:
  559. raise AssistantDataError("invalid_source", "来源聊天 ID 必须是整数。") from exc
  560. if source_type == "business":
  561. chat_id = 0
  562. elif not chat_id:
  563. raise AssistantDataError("invalid_source", "频道或群组来源必须填写聊天 ID。")
  564. publication_mode = str(current.get("publication_mode") or "review")
  565. if publication_mode not in VALID_PUBLICATION_MODES:
  566. raise AssistantDataError("invalid_source", "知识发布方式无效。")
  567. author_policy = str(current.get("author_policy") or "admins_or_allowlist")
  568. if author_policy not in VALID_AUTHOR_POLICIES:
  569. raise AssistantDataError("invalid_source", "来源作者策略无效。")
  570. linked_chat_id = 0
  571. try:
  572. linked_chat_id = int(current.get("linked_chat_id") or 0)
  573. except (TypeError, ValueError) as exc:
  574. raise AssistantDataError("invalid_source", "关联讨论群 ID 必须是整数。") from exc
  575. return {
  576. "source_type": source_type,
  577. "chat_id": chat_id,
  578. "title": clean_text(
  579. current.get("title") or ("Business 日常对话" if source_type == "business" else ""),
  580. max_length=200,
  581. required=True,
  582. ),
  583. "publication_mode": publication_mode,
  584. "author_policy": author_policy,
  585. "allowed_user_ids": _normalize_int_list(current.get("allowed_user_ids")),
  586. "include_linked_chat": bool(current.get("include_linked_chat", False))
  587. if source_type == "channel"
  588. else False,
  589. "linked_chat_id": linked_chat_id if source_type == "channel" else 0,
  590. "linked_chat_title": clean_text(
  591. current.get("linked_chat_title"), max_length=200
  592. )
  593. if source_type == "channel"
  594. else "",
  595. "backfill_days": _bounded_int(
  596. current.get("backfill_days", 30), "历史回溯天数", minimum=0, maximum=365
  597. ),
  598. "backfill_limit": _bounded_int(
  599. current.get("backfill_limit", 100), "历史回溯条数", minimum=1, maximum=500
  600. ),
  601. "enabled": bool(current.get("enabled", True)),
  602. "access_status": clean_text(current.get("access_status"), max_length=64),
  603. }
  604. async def create_knowledge_source(
  605. connection_id: str, values: dict[str, Any]
  606. ) -> dict[str, Any]:
  607. await ensure_assistant_indexes()
  608. if not await get_business_connection(connection_id):
  609. raise AssistantDataError("connection_not_found", "未找到 Business 连接。")
  610. normalized = normalize_knowledge_source(values)
  611. now = utc_now()
  612. source = {
  613. "bot_id": BOT_PROFILE_ID,
  614. "source_id": uuid4().hex,
  615. "source_key": _knowledge_source_key(
  616. str(connection_id), normalized["source_type"], normalized["chat_id"]
  617. ),
  618. "connection_id": str(connection_id),
  619. **normalized,
  620. "sync_status": "idle",
  621. "last_error": "",
  622. "created_at": now,
  623. "updated_at": now,
  624. }
  625. try:
  626. await knowledge_sourcesdb.insert_one(source)
  627. except DuplicateKeyError as exc:
  628. raise AssistantDataError("source_exists", "该连接已配置相同的知识来源。") from exc
  629. return source
  630. async def get_knowledge_source(source_id: str) -> dict[str, Any] | None:
  631. await ensure_assistant_indexes()
  632. return await knowledge_sourcesdb.find_one(_scope(source_id=str(source_id)))
  633. async def update_knowledge_source(
  634. source_id: str, values: dict[str, Any]
  635. ) -> dict[str, Any]:
  636. current = await get_knowledge_source(source_id)
  637. if not current:
  638. raise AssistantDataError("source_not_found", "未找到知识来源。")
  639. for field in ("connection_id", "source_type", "chat_id"):
  640. if field in values and str(values[field]) != str(current[field]):
  641. raise AssistantDataError("immutable_source", "来源连接、类型和聊天 ID 不可修改。")
  642. normalized = normalize_knowledge_source(values, previous=current)
  643. await knowledge_sourcesdb.update_one(
  644. _scope(source_id=str(source_id)),
  645. {"$set": {**normalized, "updated_at": utc_now()}},
  646. )
  647. return await get_knowledge_source(source_id) or current
  648. async def list_knowledge_sources(
  649. connection_id: str,
  650. *,
  651. page: int = 1,
  652. page_size: int = 20,
  653. ) -> tuple[list[dict[str, Any]], int]:
  654. await ensure_assistant_indexes()
  655. filters = _scope(connection_id=str(connection_id))
  656. page = max(1, int(page))
  657. page_size = max(1, min(int(page_size), 100))
  658. total = await knowledge_sourcesdb.count_documents(filters)
  659. cursor = (
  660. knowledge_sourcesdb.find(filters)
  661. .sort("updated_at", DESCENDING)
  662. .skip((page - 1) * page_size)
  663. .limit(page_size)
  664. )
  665. return [item async for item in cursor], total
  666. async def list_sources_for_chat(
  667. chat_id: int, *, enabled_only: bool = True
  668. ) -> list[dict[str, Any]]:
  669. await ensure_assistant_indexes()
  670. values: dict[str, Any] = {
  671. "$or": [
  672. {"chat_id": int(chat_id)},
  673. {"include_linked_chat": True, "linked_chat_id": int(chat_id)},
  674. ]
  675. }
  676. if enabled_only:
  677. values["enabled"] = True
  678. filters = _scope(values)
  679. return [item async for item in knowledge_sourcesdb.find(filters)]
  680. async def list_enabled_sources_for_chat(chat_id: int) -> list[dict[str, Any]]:
  681. return await list_sources_for_chat(chat_id, enabled_only=True)
  682. async def list_enabled_business_sources(connection_id: str) -> list[dict[str, Any]]:
  683. await ensure_assistant_indexes()
  684. return [
  685. item
  686. async for item in knowledge_sourcesdb.find(
  687. _scope(
  688. connection_id=str(connection_id),
  689. source_type="business",
  690. enabled=True,
  691. )
  692. )
  693. ]
  694. async def list_business_sources(connection_id: str) -> list[dict[str, Any]]:
  695. await ensure_assistant_indexes()
  696. return [
  697. item
  698. async for item in knowledge_sourcesdb.find(
  699. _scope(connection_id=str(connection_id), source_type="business")
  700. )
  701. ]
  702. async def update_knowledge_source_sync(
  703. source_id: str,
  704. *,
  705. status: str,
  706. error: str = "",
  707. processed: int | None = None,
  708. ) -> dict[str, Any]:
  709. values: dict[str, Any] = {
  710. "sync_status": clean_text(status, max_length=32, required=True),
  711. "last_error": clean_text(error, max_length=1000),
  712. "updated_at": utc_now(),
  713. }
  714. if status == "running":
  715. values["last_sync_started_at"] = utc_now()
  716. if status in {"completed", "failed"}:
  717. values["last_sync_finished_at"] = utc_now()
  718. if processed is not None:
  719. values["last_sync_processed"] = max(0, int(processed))
  720. result = await knowledge_sourcesdb.update_one(
  721. _scope(source_id=str(source_id)), {"$set": values}
  722. )
  723. if not result.matched_count:
  724. raise AssistantDataError("source_not_found", "未找到知识来源。")
  725. return await get_knowledge_source(source_id) or {}
  726. async def touch_knowledge_source(source_id: str, *, error: str = "") -> None:
  727. await knowledge_sourcesdb.update_one(
  728. _scope(source_id=str(source_id)),
  729. {
  730. "$set": {
  731. "last_event_at": utc_now(),
  732. "last_error": clean_text(error, max_length=1000),
  733. "updated_at": utc_now(),
  734. }
  735. },
  736. )
  737. async def _disable_candidate_entries(filters: dict[str, Any]) -> None:
  738. entry_ids = [
  739. str(item.get("knowledge_entry_id") or "")
  740. async for item in knowledge_candidatesdb.find(filters)
  741. if item.get("knowledge_entry_id")
  742. ]
  743. if entry_ids:
  744. await knowledgedb.update_many(
  745. _scope({"entry_id": {"$in": entry_ids}}),
  746. {"$set": {"enabled": False, "updated_at": utc_now()}},
  747. )
  748. async def delete_knowledge_source(source_id: str) -> None:
  749. source = await get_knowledge_source(source_id)
  750. if not source:
  751. raise AssistantDataError("source_not_found", "未找到知识来源。")
  752. candidate_filters = _scope(source_id=str(source_id))
  753. await _disable_candidate_entries(candidate_filters)
  754. await knowledge_candidatesdb.update_many(
  755. candidate_filters,
  756. {"$set": {"status": "stale", "stale_reason": "source_deleted", "updated_at": utc_now()}},
  757. )
  758. await source_eventsdb.update_many(
  759. _scope(source_id=str(source_id)),
  760. {"$set": {"deleted": True, "delete_reason": "source_deleted", "updated_at": utc_now()}},
  761. )
  762. await knowledge_sourcesdb.delete_one(_scope(source_id=str(source_id)))
  763. def source_event_id_for(source_id: str, event_key: str) -> str:
  764. digest = hashlib.sha256(
  765. f"{BOT_PROFILE_ID}:{source_id}:{event_key}".encode()
  766. ).hexdigest()
  767. return digest[:32]
  768. async def record_source_event(
  769. source: dict[str, Any],
  770. *,
  771. event_key: str,
  772. content: str,
  773. metadata: dict[str, Any] | None = None,
  774. ) -> tuple[dict[str, Any], bool]:
  775. await ensure_assistant_indexes()
  776. cleaned_content = clean_text(
  777. content, max_length=20000, required=True, preserve_lines=True
  778. )
  779. cleaned_key = clean_text(event_key, max_length=300, required=True)
  780. event_metadata = metadata if isinstance(metadata, dict) else {}
  781. content_hash = hashlib.sha256(
  782. json.dumps(
  783. {"content": cleaned_content, "metadata": event_metadata},
  784. ensure_ascii=False,
  785. sort_keys=True,
  786. default=str,
  787. ).encode()
  788. ).hexdigest()
  789. filters = _scope(source_id=str(source["source_id"]), event_key=cleaned_key)
  790. current = await source_eventsdb.find_one(filters)
  791. if current and current.get("content_hash") == content_hash and not current.get("deleted"):
  792. return current, False
  793. now = utc_now()
  794. event_id = str(current.get("event_id")) if current else source_event_id_for(
  795. str(source["source_id"]), cleaned_key
  796. )
  797. version = int(current.get("version") or 0) + 1 if current else 1
  798. document = {
  799. "bot_id": BOT_PROFILE_ID,
  800. "event_id": event_id,
  801. "source_id": str(source["source_id"]),
  802. "connection_id": str(source["connection_id"]),
  803. "event_key": cleaned_key,
  804. "content": cleaned_content,
  805. "content_hash": content_hash,
  806. "metadata": event_metadata,
  807. "version": version,
  808. "deleted": False,
  809. "extraction_status": "pending",
  810. "extraction_error": "",
  811. "updated_at": now,
  812. "expires_at": now + timedelta(days=30),
  813. }
  814. await source_eventsdb.update_one(
  815. filters,
  816. {"$set": document, "$setOnInsert": {"created_at": now}},
  817. upsert=True,
  818. )
  819. return await source_eventsdb.find_one(filters) or document, True
  820. async def get_source_event(event_id: str) -> dict[str, Any] | None:
  821. await ensure_assistant_indexes()
  822. return await source_eventsdb.find_one(_scope(event_id=str(event_id)))
  823. async def get_source_event_by_key(
  824. source_id: str, event_key: str
  825. ) -> dict[str, Any] | None:
  826. await ensure_assistant_indexes()
  827. return await source_eventsdb.find_one(
  828. _scope(source_id=str(source_id), event_key=str(event_key))
  829. )
  830. async def update_source_event_extraction(
  831. event_id: str, *, status: str, error: str = "", item_count: int = 0
  832. ) -> None:
  833. await source_eventsdb.update_one(
  834. _scope(event_id=str(event_id)),
  835. {
  836. "$set": {
  837. "extraction_status": clean_text(status, max_length=32, required=True),
  838. "extraction_error": clean_text(error, max_length=1000),
  839. "extracted_item_count": max(0, int(item_count)),
  840. "extracted_at": utc_now(),
  841. "updated_at": utc_now(),
  842. }
  843. },
  844. )
  845. async def invalidate_source_event_candidates(
  846. event_id: str, *, reason: str
  847. ) -> None:
  848. filters = _scope(event_id=str(event_id))
  849. await _disable_candidate_entries(filters)
  850. await knowledge_candidatesdb.update_many(
  851. filters,
  852. {
  853. "$set": {
  854. "status": "stale",
  855. "stale_reason": clean_text(reason, max_length=100),
  856. "updated_at": utc_now(),
  857. }
  858. },
  859. )
  860. async def mark_source_event_deleted(source_id: str, event_key: str) -> bool:
  861. event = await source_eventsdb.find_one(
  862. _scope(source_id=str(source_id), event_key=str(event_key))
  863. )
  864. if not event:
  865. return False
  866. await source_eventsdb.update_one(
  867. _scope(event_id=str(event["event_id"])),
  868. {"$set": {"deleted": True, "updated_at": utc_now()}},
  869. )
  870. await invalidate_source_event_candidates(
  871. str(event["event_id"]), reason="source_message_deleted"
  872. )
  873. return True
  874. def _normalize_candidate_item(item: dict[str, Any]) -> dict[str, Any]:
  875. try:
  876. confidence = float(item.get("confidence") or 0)
  877. except (TypeError, ValueError):
  878. confidence = 0
  879. return {
  880. "question": clean_text(item.get("question"), max_length=300, required=True),
  881. "aliases": _normalize_string_list(item.get("aliases"), max_items=20, max_length=200),
  882. "keywords": _normalize_string_list(item.get("keywords"), max_items=30, max_length=50),
  883. "answer": clean_text(
  884. item.get("answer"), max_length=4000, required=True, preserve_lines=True
  885. ),
  886. "tags": _normalize_string_list(item.get("tags"), max_items=20, max_length=30),
  887. "confidence": max(0.0, min(confidence, 1.0)),
  888. }
  889. async def replace_event_candidates(
  890. source: dict[str, Any],
  891. event: dict[str, Any],
  892. items: list[dict[str, Any]],
  893. *,
  894. auto_publish: bool,
  895. ) -> list[dict[str, Any]]:
  896. await invalidate_source_event_candidates(
  897. str(event["event_id"]), reason="source_message_changed"
  898. )
  899. existing = {
  900. int(item.get("item_index") or 0): item
  901. async for item in knowledge_candidatesdb.find(
  902. _scope(event_id=str(event["event_id"]))
  903. )
  904. }
  905. results: list[dict[str, Any]] = []
  906. metadata = event.get("metadata") if isinstance(event.get("metadata"), dict) else {}
  907. for index, raw_item in enumerate(items[:5]):
  908. normalized = _normalize_candidate_item(raw_item)
  909. current = existing.get(index) or {}
  910. now = utc_now()
  911. candidate_id = str(current.get("candidate_id") or uuid4().hex)
  912. candidate = {
  913. "bot_id": BOT_PROFILE_ID,
  914. "candidate_id": candidate_id,
  915. "connection_id": str(source["connection_id"]),
  916. "source_id": str(source["source_id"]),
  917. "event_id": str(event["event_id"]),
  918. "item_index": index,
  919. **normalized,
  920. "status": "pending",
  921. "stale_reason": "",
  922. "knowledge_entry_id": str(current.get("knowledge_entry_id") or ""),
  923. "source_snapshot": {
  924. "title": str(source.get("title") or ""),
  925. "source_type": str(source.get("source_type") or ""),
  926. "chat_id": int(metadata.get("chat_id") or source.get("chat_id") or 0),
  927. "message_id": int(metadata.get("message_id") or 0),
  928. "author_id": int(metadata.get("author_id") or 0),
  929. },
  930. "updated_at": now,
  931. }
  932. await knowledge_candidatesdb.update_one(
  933. _scope(event_id=str(event["event_id"]), item_index=index),
  934. {"$set": candidate, "$setOnInsert": {"created_at": now}},
  935. upsert=True,
  936. )
  937. if auto_publish:
  938. candidate = await publish_knowledge_candidate(candidate_id)
  939. else:
  940. candidate = await get_knowledge_candidate(candidate_id) or candidate
  941. results.append(candidate)
  942. await update_source_event_extraction(
  943. str(event["event_id"]), status="completed", item_count=len(results)
  944. )
  945. return results
  946. async def get_knowledge_candidate(candidate_id: str) -> dict[str, Any] | None:
  947. await ensure_assistant_indexes()
  948. return await knowledge_candidatesdb.find_one(
  949. _scope(candidate_id=str(candidate_id))
  950. )
  951. async def publish_knowledge_candidate(candidate_id: str) -> dict[str, Any]:
  952. candidate = await get_knowledge_candidate(candidate_id)
  953. if not candidate:
  954. raise AssistantDataError("candidate_not_found", "未找到知识候选。")
  955. if candidate.get("status") == "stale":
  956. raise AssistantDataError("candidate_stale", "来源已变更或删除,不能发布该候选。")
  957. if not await get_knowledge_source(str(candidate.get("source_id") or "")):
  958. await knowledge_candidatesdb.update_one(
  959. _scope(candidate_id=str(candidate_id)),
  960. {
  961. "$set": {
  962. "status": "stale",
  963. "stale_reason": "source_deleted",
  964. "updated_at": utc_now(),
  965. }
  966. },
  967. )
  968. raise AssistantDataError("source_not_found", "知识来源已删除,不能发布该候选。")
  969. reference = {
  970. **(candidate.get("source_snapshot") or {}),
  971. "source_id": str(candidate.get("source_id") or ""),
  972. "event_id": str(candidate.get("event_id") or ""),
  973. }
  974. values = {
  975. key: candidate.get(key)
  976. for key in ("question", "aliases", "keywords", "answer", "tags")
  977. }
  978. values.update(
  979. {
  980. "priority": round(float(candidate.get("confidence") or 0) * 100),
  981. "enabled": True,
  982. "source_candidate_id": str(candidate_id),
  983. "source_references": [reference],
  984. }
  985. )
  986. entry_id = str(candidate.get("knowledge_entry_id") or "")
  987. if entry_id:
  988. try:
  989. entry = await update_knowledge_entry(entry_id, values)
  990. except AssistantDataError as exc:
  991. if exc.code != "knowledge_not_found":
  992. raise
  993. entry = await create_knowledge_entry(str(candidate["connection_id"]), values)
  994. else:
  995. entry = await create_knowledge_entry(str(candidate["connection_id"]), values)
  996. await knowledge_candidatesdb.update_one(
  997. _scope(candidate_id=str(candidate_id)),
  998. {
  999. "$set": {
  1000. "status": "published",
  1001. "knowledge_entry_id": str(entry["entry_id"]),
  1002. "published_at": utc_now(),
  1003. "updated_at": utc_now(),
  1004. }
  1005. },
  1006. )
  1007. return await get_knowledge_candidate(candidate_id) or candidate
  1008. async def reject_knowledge_candidate(candidate_id: str) -> dict[str, Any]:
  1009. candidate = await get_knowledge_candidate(candidate_id)
  1010. if not candidate:
  1011. raise AssistantDataError("candidate_not_found", "未找到知识候选。")
  1012. entry_id = str(candidate.get("knowledge_entry_id") or "")
  1013. if entry_id:
  1014. await knowledgedb.update_one(
  1015. _scope(entry_id=entry_id),
  1016. {"$set": {"enabled": False, "updated_at": utc_now()}},
  1017. )
  1018. await knowledge_candidatesdb.update_one(
  1019. _scope(candidate_id=str(candidate_id)),
  1020. {
  1021. "$set": {
  1022. "status": "rejected",
  1023. "rejected_at": utc_now(),
  1024. "updated_at": utc_now(),
  1025. }
  1026. },
  1027. )
  1028. return await get_knowledge_candidate(candidate_id) or candidate
  1029. async def list_knowledge_candidates(
  1030. connection_id: str,
  1031. *,
  1032. status: str = "",
  1033. source_id: str = "",
  1034. query: str = "",
  1035. page: int = 1,
  1036. page_size: int = 20,
  1037. ) -> tuple[list[dict[str, Any]], int]:
  1038. await ensure_assistant_indexes()
  1039. filters: dict[str, Any] = _scope(connection_id=str(connection_id))
  1040. if status:
  1041. if status not in VALID_CANDIDATE_STATUSES:
  1042. raise AssistantDataError("invalid_status", "知识候选状态无效。")
  1043. filters["status"] = status
  1044. if source_id:
  1045. filters["source_id"] = str(source_id)
  1046. normalized_query = clean_text(query, max_length=100)
  1047. if normalized_query:
  1048. pattern = re.escape(normalized_query)
  1049. filters["$or"] = [
  1050. {"question": {"$regex": pattern, "$options": "i"}},
  1051. {"answer": {"$regex": pattern, "$options": "i"}},
  1052. {"keywords": {"$regex": pattern, "$options": "i"}},
  1053. ]
  1054. page = max(1, int(page))
  1055. page_size = max(1, min(int(page_size), 100))
  1056. total = await knowledge_candidatesdb.count_documents(filters)
  1057. cursor = (
  1058. knowledge_candidatesdb.find(filters)
  1059. .sort("updated_at", DESCENDING)
  1060. .skip((page - 1) * page_size)
  1061. .limit(page_size)
  1062. )
  1063. return [item async for item in cursor], total
  1064. def _search_text(value: str) -> str:
  1065. return re.sub(r"[^\w\u3400-\u9fff]+", "", value.casefold())
  1066. async def match_knowledge(
  1067. connection_id: str, query: str, *, limit: int = 8
  1068. ) -> list[dict[str, Any]]:
  1069. await ensure_assistant_indexes()
  1070. normalized_query = _search_text(query)
  1071. if not normalized_query:
  1072. return []
  1073. entries = [
  1074. item
  1075. async for item in knowledgedb.find(
  1076. _scope(connection_id=str(connection_id), enabled=True)
  1077. )
  1078. ]
  1079. scored: list[tuple[int, dict[str, Any]]] = []
  1080. for entry in entries:
  1081. score = int(entry.get("priority") or 0)
  1082. phrases = [entry.get("question", ""), *(entry.get("aliases") or [])]
  1083. for phrase in phrases:
  1084. normalized_phrase = _search_text(str(phrase))
  1085. if not normalized_phrase:
  1086. continue
  1087. if normalized_query == normalized_phrase:
  1088. score += 1000
  1089. elif normalized_phrase in normalized_query:
  1090. score += 300 + min(len(normalized_phrase), 100)
  1091. elif normalized_query in normalized_phrase and len(normalized_query) >= 4:
  1092. score += 120
  1093. for keyword in entry.get("keywords") or []:
  1094. normalized_keyword = _search_text(str(keyword))
  1095. if normalized_keyword and normalized_keyword in normalized_query:
  1096. score += 80
  1097. if score > int(entry.get("priority") or 0):
  1098. scored.append((score, entry))
  1099. scored.sort(key=lambda item: (item[0], item[1].get("priority", 0)), reverse=True)
  1100. return [{**entry, "match_score": score} for score, entry in scored[: max(1, limit)]]
  1101. def classify_handoff(text: str) -> str | None:
  1102. normalized = text.casefold()
  1103. if any(term in normalized for term in HANDOFF_TERMS):
  1104. return "customer_requested_human"
  1105. if any(term in normalized for term in SENSITIVE_TERMS):
  1106. return "sensitive_request"
  1107. return None
  1108. def conversation_id_for(connection_id: str, chat_id: int) -> str:
  1109. digest = hashlib.sha256(f"{BOT_PROFILE_ID}:{connection_id}:{int(chat_id)}".encode()).hexdigest()
  1110. return digest[:32]
  1111. async def get_or_create_conversation(
  1112. connection_id: str,
  1113. chat_id: int,
  1114. *,
  1115. customer: dict[str, Any] | None = None,
  1116. ) -> dict[str, Any]:
  1117. await ensure_assistant_indexes()
  1118. conversation_id = conversation_id_for(connection_id, chat_id)
  1119. now = utc_now()
  1120. update: dict[str, Any] = {"updated_at": now, "last_message_at": now}
  1121. if customer:
  1122. update["customer"] = _public_user(customer)
  1123. await conversationsdb.update_one(
  1124. _scope(conversation_id=conversation_id),
  1125. {
  1126. "$set": update,
  1127. "$setOnInsert": {
  1128. "bot_id": BOT_PROFILE_ID,
  1129. "conversation_id": conversation_id,
  1130. "connection_id": str(connection_id),
  1131. "chat_id": int(chat_id),
  1132. "status": "auto",
  1133. "summary": "",
  1134. "handoff_reason": "",
  1135. "created_at": now,
  1136. },
  1137. },
  1138. upsert=True,
  1139. )
  1140. return await conversationsdb.find_one(_scope(conversation_id=conversation_id)) or {}
  1141. async def get_conversation(conversation_id: str) -> dict[str, Any] | None:
  1142. await ensure_assistant_indexes()
  1143. return await conversationsdb.find_one(
  1144. _scope(conversation_id=str(conversation_id))
  1145. )
  1146. async def get_conversation_by_chat(
  1147. connection_id: str, chat_id: int
  1148. ) -> dict[str, Any] | None:
  1149. await ensure_assistant_indexes()
  1150. return await conversationsdb.find_one(
  1151. _scope(connection_id=str(connection_id), chat_id=int(chat_id))
  1152. )
  1153. async def append_conversation_message(
  1154. conversation_id: str,
  1155. *,
  1156. direction: str,
  1157. telegram_message_id: int,
  1158. text: str,
  1159. sender_id: int = 0,
  1160. metadata: dict[str, Any] | None = None,
  1161. ) -> dict[str, Any]:
  1162. await ensure_assistant_indexes()
  1163. if direction not in {"incoming", "assistant", "human"}:
  1164. raise AssistantDataError("invalid_direction", "会话消息方向无效。")
  1165. now = utc_now()
  1166. message_key = f"{conversation_id}:{direction}:{int(telegram_message_id)}"
  1167. document = {
  1168. "bot_id": BOT_PROFILE_ID,
  1169. "message_key": message_key,
  1170. "conversation_id": str(conversation_id),
  1171. "direction": direction,
  1172. "telegram_message_id": int(telegram_message_id),
  1173. "sender_id": int(sender_id or 0),
  1174. "text": clean_text(text, max_length=12000, preserve_lines=True),
  1175. "metadata": metadata or {},
  1176. "created_at": now,
  1177. "expires_at": now + timedelta(days=30),
  1178. }
  1179. try:
  1180. await messagesdb.insert_one(document)
  1181. except DuplicateKeyError:
  1182. return await messagesdb.find_one(_scope(message_key=message_key)) or document
  1183. await conversationsdb.update_one(
  1184. _scope(conversation_id=str(conversation_id)),
  1185. {"$set": {"last_message_at": now, "updated_at": now}},
  1186. )
  1187. return document
  1188. async def recent_conversation_messages(
  1189. conversation_id: str, *, limit: int = 12
  1190. ) -> list[dict[str, Any]]:
  1191. cursor = (
  1192. messagesdb.find(_scope(conversation_id=str(conversation_id)))
  1193. .sort([("created_at", DESCENDING), ("_id", DESCENDING)])
  1194. .limit(max(1, min(int(limit), 50)))
  1195. )
  1196. items = [item async for item in cursor]
  1197. items.reverse()
  1198. return items
  1199. async def set_conversation_handoff(
  1200. conversation_id: str,
  1201. reason: str,
  1202. *,
  1203. summary: str = "",
  1204. ) -> dict[str, Any]:
  1205. now = utc_now()
  1206. update: dict[str, Any] = {
  1207. "status": "handoff",
  1208. "handoff_reason": clean_text(reason, max_length=200, required=True),
  1209. "paused_until": None,
  1210. "handoff_at": now,
  1211. "updated_at": now,
  1212. }
  1213. if summary:
  1214. update["summary"] = clean_text(summary, max_length=4000, preserve_lines=True)
  1215. result = await conversationsdb.find_one_and_update(
  1216. _scope(conversation_id=str(conversation_id)),
  1217. {"$set": update},
  1218. return_document=True,
  1219. )
  1220. if not result:
  1221. raise AssistantDataError("conversation_not_found", "未找到客户会话。")
  1222. return result
  1223. async def pause_conversation_for_human(
  1224. conversation_id: str, *, hours: int
  1225. ) -> dict[str, Any]:
  1226. now = utc_now()
  1227. result = await conversationsdb.find_one_and_update(
  1228. _scope(conversation_id=str(conversation_id)),
  1229. {
  1230. "$set": {
  1231. "status": "human_paused",
  1232. "handoff_reason": "human_reply_detected",
  1233. "paused_until": now + timedelta(hours=max(1, min(int(hours), 168))),
  1234. "updated_at": now,
  1235. }
  1236. },
  1237. return_document=True,
  1238. )
  1239. if not result:
  1240. raise AssistantDataError("conversation_not_found", "未找到客户会话。")
  1241. return result
  1242. async def resume_conversation(conversation_id: str) -> dict[str, Any]:
  1243. result = await conversationsdb.find_one_and_update(
  1244. _scope(conversation_id=str(conversation_id)),
  1245. {
  1246. "$set": {
  1247. "status": "auto",
  1248. "handoff_reason": "",
  1249. "paused_until": None,
  1250. "updated_at": utc_now(),
  1251. }
  1252. },
  1253. return_document=True,
  1254. )
  1255. if not result:
  1256. raise AssistantDataError("conversation_not_found", "未找到客户会话。")
  1257. return result
  1258. async def pause_conversation(conversation_id: str) -> dict[str, Any]:
  1259. return await set_conversation_handoff(conversation_id, "admin_paused")
  1260. async def close_conversation(conversation_id: str) -> dict[str, Any]:
  1261. result = await conversationsdb.find_one_and_update(
  1262. _scope(conversation_id=str(conversation_id)),
  1263. {"$set": {"status": "closed", "closed_at": utc_now(), "updated_at": utc_now()}},
  1264. return_document=True,
  1265. )
  1266. if not result:
  1267. raise AssistantDataError("conversation_not_found", "未找到客户会话。")
  1268. return result
  1269. async def clear_conversation(conversation_id: str) -> dict[str, Any]:
  1270. current = await get_conversation(conversation_id)
  1271. if not current:
  1272. raise AssistantDataError("conversation_not_found", "未找到客户会话。")
  1273. await messagesdb.delete_many(_scope(conversation_id=str(conversation_id)))
  1274. await conversationsdb.update_one(
  1275. _scope(conversation_id=str(conversation_id)),
  1276. {
  1277. "$set": {
  1278. "summary": "",
  1279. "handoff_reason": "",
  1280. "status": "auto",
  1281. "paused_until": None,
  1282. "cleared_at": utc_now(),
  1283. "updated_at": utc_now(),
  1284. }
  1285. },
  1286. )
  1287. return await get_conversation(conversation_id) or current
  1288. async def update_conversation_summary(
  1289. conversation_id: str, summary: str
  1290. ) -> None:
  1291. cleaned = clean_text(summary, max_length=4000, preserve_lines=True)
  1292. if cleaned:
  1293. await conversationsdb.update_one(
  1294. _scope(conversation_id=str(conversation_id)),
  1295. {"$set": {"summary": cleaned, "updated_at": utc_now()}},
  1296. )
  1297. async def list_conversations(
  1298. *,
  1299. connection_id: str = "",
  1300. status: str = "",
  1301. page: int = 1,
  1302. page_size: int = 20,
  1303. ) -> tuple[list[dict[str, Any]], int]:
  1304. await ensure_assistant_indexes()
  1305. filters: dict[str, Any] = _scope()
  1306. if connection_id:
  1307. filters["connection_id"] = str(connection_id)
  1308. if status:
  1309. if status not in {"auto", "handoff", "human_paused", "closed"}:
  1310. raise AssistantDataError("invalid_status", "会话状态无效。")
  1311. filters["status"] = status
  1312. page = max(1, int(page))
  1313. page_size = max(1, min(int(page_size), 100))
  1314. total = await conversationsdb.count_documents(filters)
  1315. cursor = (
  1316. conversationsdb.find(filters)
  1317. .sort("updated_at", DESCENDING)
  1318. .skip((page - 1) * page_size)
  1319. .limit(page_size)
  1320. )
  1321. return [item async for item in cursor], total
  1322. async def conversation_detail(conversation_id: str) -> dict[str, Any]:
  1323. conversation = await get_conversation(conversation_id)
  1324. if not conversation:
  1325. raise AssistantDataError("conversation_not_found", "未找到客户会话。")
  1326. conversation["messages"] = await recent_conversation_messages(
  1327. conversation_id, limit=50
  1328. )
  1329. return conversation
  1330. def _usage_day(now: datetime, timezone_name: str) -> str:
  1331. try:
  1332. zone = ZoneInfo(timezone_name)
  1333. except ZoneInfoNotFoundError:
  1334. zone = ZoneInfo("Asia/Shanghai")
  1335. return as_utc(now).astimezone(zone).date().isoformat()
  1336. async def reserve_ai_usage(
  1337. connection_id: str,
  1338. chat_id: int,
  1339. settings: dict[str, Any],
  1340. *,
  1341. now: datetime | None = None,
  1342. ) -> tuple[bool, str]:
  1343. await ensure_assistant_indexes()
  1344. current_time = as_utc(now or utc_now())
  1345. day = _usage_day(current_time, str(settings.get("timezone") or "Asia/Shanghai"))
  1346. account_limit = int(settings.get("account_daily_limit") or 200)
  1347. customer_limit = int(settings.get("customer_daily_limit") or 20)
  1348. account_filter = _scope(
  1349. connection_id=str(connection_id), chat_id=0, day=day
  1350. )
  1351. customer_filter = _scope(
  1352. connection_id=str(connection_id), chat_id=int(chat_id), day=day
  1353. )
  1354. for filters, usage_scope in (
  1355. (account_filter, "account"),
  1356. (customer_filter, "customer"),
  1357. ):
  1358. await usagedb.update_one(
  1359. filters,
  1360. {
  1361. "$setOnInsert": {
  1362. **filters,
  1363. "scope": usage_scope,
  1364. "call_count": 0,
  1365. "created_at": current_time,
  1366. "updated_at": current_time,
  1367. }
  1368. },
  1369. upsert=True,
  1370. )
  1371. account = await usagedb.find_one_and_update(
  1372. {**account_filter, "call_count": {"$lt": account_limit}},
  1373. {"$inc": {"call_count": 1}, "$set": {"updated_at": current_time}},
  1374. return_document=True,
  1375. )
  1376. if not account:
  1377. return False, "account_daily_limit"
  1378. customer = await usagedb.find_one_and_update(
  1379. {**customer_filter, "call_count": {"$lt": customer_limit}},
  1380. {"$inc": {"call_count": 1}, "$set": {"updated_at": current_time}},
  1381. return_document=True,
  1382. )
  1383. if not customer:
  1384. await usagedb.update_one(
  1385. {**account_filter, "call_count": {"$gt": 0}},
  1386. {"$inc": {"call_count": -1}, "$set": {"updated_at": current_time}},
  1387. )
  1388. return False, "customer_daily_limit"
  1389. return True, ""
  1390. async def usage_metrics(*, connection_id: str = "", day: str = "") -> dict[str, Any]:
  1391. await ensure_assistant_indexes()
  1392. match = _scope()
  1393. if connection_id:
  1394. match["connection_id"] = str(connection_id)
  1395. if day:
  1396. match["day"] = str(day)
  1397. calls = 0
  1398. customers: set[tuple[str, int]] = set()
  1399. customer_match = {
  1400. **match,
  1401. "scope": {"$ne": "account"},
  1402. "chat_id": {"$ne": 0},
  1403. "call_count": {"$gt": 0},
  1404. }
  1405. async for item in usagedb.find(customer_match):
  1406. calls += int(item.get("call_count") or 0)
  1407. customers.add((str(item.get("connection_id")), int(item.get("chat_id") or 0)))
  1408. conversation_filters = _scope()
  1409. if connection_id:
  1410. conversation_filters["connection_id"] = str(connection_id)
  1411. return {
  1412. "ai_calls": calls,
  1413. "customers": len(customers),
  1414. "conversations": await conversationsdb.count_documents(conversation_filters),
  1415. "handoffs": await conversationsdb.count_documents(
  1416. {**conversation_filters, "status": "handoff"}
  1417. ),
  1418. "human_paused": await conversationsdb.count_documents(
  1419. {**conversation_filters, "status": "human_paused"}
  1420. ),
  1421. }
  1422. async def claim_update(update_id: int) -> bool:
  1423. await ensure_assistant_indexes()
  1424. now = utc_now()
  1425. try:
  1426. await updatesdb.insert_one(
  1427. {
  1428. "bot_id": BOT_PROFILE_ID,
  1429. "update_id": int(update_id),
  1430. "status": "processing",
  1431. "attempts": 1,
  1432. "created_at": now,
  1433. "updated_at": now,
  1434. "expires_at": now + timedelta(days=7),
  1435. }
  1436. )
  1437. return True
  1438. except DuplicateKeyError:
  1439. filters = _scope(update_id=int(update_id))
  1440. current = await updatesdb.find_one(filters) or {}
  1441. if current.get("status") == "done":
  1442. return False
  1443. await updatesdb.update_one(
  1444. filters,
  1445. {"$inc": {"attempts": 1}, "$set": {"updated_at": now}},
  1446. )
  1447. return True
  1448. async def mark_update_done(update_id: int) -> None:
  1449. await updatesdb.update_one(
  1450. _scope(update_id=int(update_id)),
  1451. {"$set": {"status": "done", "updated_at": utc_now(), "last_error": ""}},
  1452. )
  1453. async def mark_update_failed(update_id: int, error: str) -> int:
  1454. await updatesdb.update_one(
  1455. _scope(update_id=int(update_id)),
  1456. {"$set": {"status": "failed", "updated_at": utc_now(), "last_error": error[:1000]}},
  1457. )
  1458. current = await updatesdb.find_one(_scope(update_id=int(update_id))) or {}
  1459. return int(current.get("attempts") or 1)
  1460. async def dead_letter_update(update: dict[str, Any], error: str) -> None:
  1461. update_id = int(update.get("update_id") or 0)
  1462. now = utc_now()
  1463. await deadlettersdb.update_one(
  1464. _scope(update_id=update_id),
  1465. {
  1466. "$set": {
  1467. "bot_id": BOT_PROFILE_ID,
  1468. "payload": update,
  1469. "error": error[:2000],
  1470. "updated_at": now,
  1471. "expires_at": now + timedelta(days=30),
  1472. },
  1473. "$setOnInsert": {"created_at": now},
  1474. },
  1475. upsert=True,
  1476. )
  1477. await mark_update_done(update_id)
  1478. async def runtime_status() -> dict[str, Any]:
  1479. await ensure_assistant_indexes()
  1480. stored = await runtimedb.find_one(
  1481. _scope(runtime_id="business_assistant")
  1482. ) or {}
  1483. return {
  1484. "runtime_id": "business_assistant",
  1485. "polling_state": "stopped",
  1486. "business_mode_supported": False,
  1487. "webhook_conflict": False,
  1488. "webhook_url": "",
  1489. "model_configured": False,
  1490. "last_error": "",
  1491. **{key: value for key, value in stored.items() if key != "_id"},
  1492. }
  1493. async def update_runtime_status(values: dict[str, Any]) -> dict[str, Any]:
  1494. await ensure_assistant_indexes()
  1495. insert_values = {
  1496. "bot_id": BOT_PROFILE_ID,
  1497. "runtime_id": "business_assistant",
  1498. "created_at": utc_now(),
  1499. }
  1500. if "offset" not in values:
  1501. insert_values["offset"] = 0
  1502. await runtimedb.update_one(
  1503. _scope(runtime_id="business_assistant"),
  1504. {
  1505. "$set": {**values, "updated_at": utc_now()},
  1506. "$setOnInsert": insert_values,
  1507. },
  1508. upsert=True,
  1509. )
  1510. return await runtime_status()
  1511. async def load_update_offset() -> int:
  1512. return int((await runtime_status()).get("offset") or 0)
  1513. async def save_update_offset(offset: int) -> None:
  1514. await update_runtime_status({"offset": int(offset), "last_poll_at": utc_now()})
  1515. async def due_digest_connections(now: datetime | None = None) -> list[dict[str, Any]]:
  1516. current = as_utc(now or utc_now())
  1517. results: list[dict[str, Any]] = []
  1518. async for settings in settingsdb.find(
  1519. _scope(digest_enabled=True, assistant_enabled=True)
  1520. ):
  1521. try:
  1522. local = current.astimezone(ZoneInfo(str(settings.get("timezone"))))
  1523. except ZoneInfoNotFoundError:
  1524. local = current.astimezone(ZoneInfo("Asia/Shanghai"))
  1525. day = local.date().isoformat()
  1526. if local.strftime("%H:%M") != str(settings.get("digest_time") or "09:00"):
  1527. continue
  1528. if settings.get("last_digest_day") == day:
  1529. continue
  1530. connection = await get_business_connection(settings["connection_id"])
  1531. if connection and connection.get("is_enabled"):
  1532. results.append({"connection": connection, "settings": settings, "day": day})
  1533. return results
  1534. async def mark_digest_sent(connection_id: str, day: str) -> None:
  1535. await settingsdb.update_one(
  1536. _scope(connection_id=str(connection_id)),
  1537. {"$set": {"last_digest_day": str(day), "last_digest_at": utc_now()}},
  1538. )