business_assistant.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345
  1. from __future__ import annotations
  2. import asyncio
  3. from contextlib import suppress
  4. from datetime import UTC, datetime, timedelta
  5. from typing import Any
  6. from pyrogram import filters
  7. from pyrogram.enums import ChatMemberStatus, ChatType
  8. import wbb
  9. from wbb import SUDOERS, app
  10. from wbb.services.business_assistant import BusinessAssistantRuntime, classify_handoff
  11. from wbb.utils.dbassistant import (
  12. AssistantDataError,
  13. get_account_settings,
  14. get_business_connection,
  15. get_conversation,
  16. get_knowledge_source,
  17. list_enabled_sources_for_chat,
  18. list_sources_for_chat,
  19. mark_source_event_deleted,
  20. resume_conversation,
  21. touch_knowledge_source,
  22. update_knowledge_source_sync,
  23. )
  24. __MODULE__ = "智能接待"
  25. __HELP__ = "Telegram Business 智能接待由 Web 管理后台配置。"
  26. _runtime: BusinessAssistantRuntime | None = None
  27. _source_sync_tasks: dict[str, asyncio.Task[None]] = {}
  28. async def start_business_assistant_runtime() -> None:
  29. global _runtime
  30. if _runtime is not None:
  31. return
  32. session = getattr(wbb, "aiohttpsession", None)
  33. token = str(getattr(wbb, "BOT_TOKEN", ""))
  34. if not session or not token:
  35. wbb.log.error("智能接待未启动:Telegram Bot 会话或令牌不可用。")
  36. return
  37. _runtime = BusinessAssistantRuntime(token=token, session=session)
  38. status = await _runtime.start()
  39. if status.get("polling_state") == "running":
  40. wbb.log.info("Telegram Business 智能接待已启动。")
  41. else:
  42. wbb.log.error(f"Telegram Business 智能接待未运行:{status.get('last_error') or '未知原因'}")
  43. async def stop_business_assistant_runtime() -> None:
  44. global _runtime
  45. tasks = [task for task in _source_sync_tasks.values() if not task.done()]
  46. for task in tasks:
  47. task.cancel()
  48. if tasks:
  49. await asyncio.gather(*tasks, return_exceptions=True)
  50. _source_sync_tasks.clear()
  51. if _runtime is None:
  52. return
  53. await _runtime.stop()
  54. _runtime = None
  55. def get_business_assistant_runtime() -> BusinessAssistantRuntime | None:
  56. return _runtime
  57. async def validate_knowledge_source(values: dict[str, Any]) -> dict[str, Any]:
  58. source_type = str(values.get("source_type") or "")
  59. if source_type == "business":
  60. return {
  61. **values,
  62. "chat_id": 0,
  63. "title": str(values.get("title") or "Business 日常对话"),
  64. "linked_chat_id": 0,
  65. "linked_chat_title": "",
  66. "access_status": "business_connected",
  67. }
  68. try:
  69. chat_id = int(values.get("chat_id") or 0)
  70. except (TypeError, ValueError) as exc:
  71. raise AssistantDataError("invalid_source", "来源聊天 ID 必须是整数。") from exc
  72. if not chat_id:
  73. raise AssistantDataError("invalid_source", "请填写频道或群组聊天 ID。")
  74. try:
  75. chat = await app.get_chat(chat_id)
  76. member = await app.get_chat_member(chat_id, int(wbb.BOT_ID))
  77. except Exception as exc:
  78. raise AssistantDataError(
  79. "source_access_denied", "机器人无法访问该频道或群组,请先将机器人加入。"
  80. ) from exc
  81. chat_type = getattr(chat, "type", None)
  82. if source_type == "channel" and chat_type != ChatType.CHANNEL:
  83. raise AssistantDataError("source_type_mismatch", "该聊天不是频道。")
  84. if source_type == "group" and chat_type not in {ChatType.GROUP, ChatType.SUPERGROUP}:
  85. raise AssistantDataError("source_type_mismatch", "该聊天不是群组。")
  86. if source_type not in {"channel", "group"}:
  87. raise AssistantDataError("invalid_source", "知识来源类型无效。")
  88. if source_type == "channel" and member.status not in {
  89. ChatMemberStatus.OWNER,
  90. ChatMemberStatus.ADMINISTRATOR,
  91. }:
  92. raise AssistantDataError(
  93. "source_access_denied", "机器人必须是频道管理员才能持续采集频道消息。"
  94. )
  95. linked_chat = getattr(chat, "linked_chat", None)
  96. return {
  97. **values,
  98. "chat_id": chat_id,
  99. "title": str(getattr(chat, "title", None) or values.get("title") or chat_id),
  100. "linked_chat_id": int(getattr(linked_chat, "id", 0) or 0),
  101. "linked_chat_title": str(getattr(linked_chat, "title", "") or ""),
  102. "access_status": str(getattr(member.status, "value", member.status)),
  103. }
  104. def _message_text(message: Any) -> str:
  105. return str(getattr(message, "text", None) or getattr(message, "caption", None) or "").strip()
  106. async def _trusted_source_author(
  107. source: dict[str, Any], message: Any, *, effective_source_type: str
  108. ) -> bool:
  109. if effective_source_type == "channel":
  110. return True
  111. author = getattr(message, "from_user", None)
  112. author_id = int(getattr(author, "id", 0) or 0)
  113. policy = str(source.get("author_policy") or "admins_or_allowlist")
  114. allowlisted = author_id in {int(item) for item in source.get("allowed_user_ids") or []}
  115. if policy == "all":
  116. return True
  117. if policy == "allowlist":
  118. return allowlisted
  119. sender_chat_id = int(getattr(getattr(message, "sender_chat", None), "id", 0) or 0)
  120. is_admin = sender_chat_id == int(message.chat.id)
  121. if author_id:
  122. with suppress(Exception):
  123. member = await app.get_chat_member(int(message.chat.id), author_id)
  124. is_admin = member.status in {
  125. ChatMemberStatus.OWNER,
  126. ChatMemberStatus.ADMINISTRATOR,
  127. }
  128. if policy == "admins":
  129. return is_admin
  130. return allowlisted or is_admin
  131. async def _ingest_source_message(source: dict[str, Any], message: Any) -> bool:
  132. runtime = get_business_assistant_runtime()
  133. if runtime is None:
  134. return False
  135. text = _message_text(message)
  136. author = getattr(message, "from_user", None)
  137. if (
  138. not text
  139. or bool(getattr(message, "outgoing", False))
  140. or bool(getattr(author, "is_bot", False))
  141. or classify_handoff(text) == "sensitive_request"
  142. ):
  143. return False
  144. chat_id = int(message.chat.id)
  145. linked_discussion = bool(
  146. source.get("include_linked_chat")
  147. and int(source.get("linked_chat_id") or 0) == chat_id
  148. )
  149. if linked_discussion and bool(getattr(message, "is_automatic_forward", False)):
  150. return False
  151. effective_source_type = "discussion" if linked_discussion else str(source["source_type"])
  152. trusted = await _trusted_source_author(
  153. source, message, effective_source_type=effective_source_type
  154. )
  155. author_id = int(getattr(author, "id", 0) or 0)
  156. message_id = int(getattr(message, "id", 0) or 0)
  157. if not message_id:
  158. return False
  159. try:
  160. await runtime.ingestion.ingest(
  161. source,
  162. event_key=f"telegram:{chat_id}:{message_id}",
  163. content=text,
  164. metadata={
  165. "effective_source_type": effective_source_type,
  166. "chat_id": chat_id,
  167. "message_id": message_id,
  168. "author_id": author_id,
  169. "trusted_author": trusted,
  170. "message_date": getattr(message, "date", None),
  171. },
  172. auto_publish=bool(
  173. source.get("publication_mode") == "auto" and trusted
  174. ),
  175. )
  176. await touch_knowledge_source(str(source["source_id"]))
  177. return True
  178. except Exception as exc:
  179. await touch_knowledge_source(str(source["source_id"]), error=str(exc))
  180. wbb.log.error(f"知识来源采集失败:{source['source_id']} / {exc}")
  181. return False
  182. async def process_knowledge_source_message(message: Any) -> None:
  183. chat_id = int(getattr(getattr(message, "chat", None), "id", 0) or 0)
  184. if not chat_id:
  185. return
  186. for source in await list_enabled_sources_for_chat(chat_id):
  187. await _ingest_source_message(source, message)
  188. async def _run_source_sync(source_id: str) -> None:
  189. processed = 0
  190. try:
  191. source = await get_knowledge_source(source_id)
  192. if not source:
  193. return
  194. await update_knowledge_source_sync(source_id, status="running", processed=0)
  195. cutoff = datetime.now(UTC) - timedelta(days=int(source.get("backfill_days") or 0))
  196. chat_ids = [int(source.get("chat_id") or 0)]
  197. if source.get("include_linked_chat") and source.get("linked_chat_id"):
  198. chat_ids.append(int(source["linked_chat_id"]))
  199. limit = int(source.get("backfill_limit") or 100)
  200. for chat_id in dict.fromkeys(item for item in chat_ids if item):
  201. remaining = limit - processed
  202. if remaining <= 0:
  203. break
  204. async for message in app.get_chat_history(chat_id, limit=remaining):
  205. message_date = getattr(message, "date", None)
  206. if isinstance(message_date, datetime):
  207. aware_date = (
  208. message_date.replace(tzinfo=UTC)
  209. if message_date.tzinfo is None
  210. else message_date.astimezone(UTC)
  211. )
  212. if int(source.get("backfill_days") or 0) and aware_date < cutoff:
  213. break
  214. if await _ingest_source_message(source, message):
  215. processed += 1
  216. await update_knowledge_source_sync(
  217. source_id, status="completed", processed=processed
  218. )
  219. except asyncio.CancelledError:
  220. raise
  221. except Exception as exc:
  222. with suppress(Exception):
  223. await update_knowledge_source_sync(
  224. source_id, status="failed", error=str(exc), processed=processed
  225. )
  226. wbb.log.error(f"知识来源历史同步失败:{source_id} / {exc}")
  227. async def schedule_knowledge_source_sync(source_id: str) -> dict[str, Any]:
  228. if get_business_assistant_runtime() is None:
  229. raise AssistantDataError(
  230. "assistant_runtime_unavailable", "智能接待运行时尚未启动。"
  231. )
  232. source = await get_knowledge_source(source_id)
  233. if not source:
  234. raise AssistantDataError("source_not_found", "未找到知识来源。")
  235. current = _source_sync_tasks.get(source_id)
  236. if current and not current.done():
  237. return source
  238. queued = await update_knowledge_source_sync(source_id, status="queued", processed=0)
  239. task = asyncio.create_task(
  240. _run_source_sync(source_id), name=f"business-knowledge-sync-{source_id}"
  241. )
  242. _source_sync_tasks[source_id] = task
  243. task.add_done_callback(lambda _task: _source_sync_tasks.pop(source_id, None))
  244. return queued
  245. async def cancel_knowledge_source_sync(source_id: str) -> None:
  246. task = _source_sync_tasks.pop(source_id, None)
  247. if task and not task.done():
  248. task.cancel()
  249. await asyncio.gather(task, return_exceptions=True)
  250. async def regenerate_knowledge_candidate(candidate_id: str) -> list[dict[str, Any]]:
  251. runtime = get_business_assistant_runtime()
  252. if runtime is None:
  253. raise AssistantDataError(
  254. "assistant_runtime_unavailable", "智能接待运行时尚未启动。"
  255. )
  256. return await runtime.ingestion.regenerate_candidate(candidate_id)
  257. @app.on_message(filters.channel | filters.group, group=-30)
  258. async def collect_knowledge_source_message(_, message):
  259. await process_knowledge_source_message(message)
  260. @app.on_edited_message(filters.channel | filters.group, group=-30)
  261. async def collect_edited_knowledge_source_message(_, message):
  262. await process_knowledge_source_message(message)
  263. @app.on_deleted_messages(filters.channel | filters.group, group=-30)
  264. async def collect_deleted_knowledge_source_messages(_, messages):
  265. for message in messages:
  266. chat_id = int(getattr(getattr(message, "chat", None), "id", 0) or 0)
  267. message_id = int(getattr(message, "id", 0) or 0)
  268. if not chat_id or not message_id:
  269. continue
  270. for source in await list_sources_for_chat(chat_id, enabled_only=False):
  271. await mark_source_event_deleted(
  272. str(source["source_id"]), f"telegram:{chat_id}:{message_id}"
  273. )
  274. async def _ops_group_admin_allowed(query, settings: dict) -> bool:
  275. ops_group_id = int(settings.get("ops_group_id") or 0)
  276. if not ops_group_id or not query.message or int(query.message.chat.id) != ops_group_id:
  277. return False
  278. with suppress(Exception):
  279. member = await app.get_chat_member(ops_group_id, query.from_user.id)
  280. return member.status in {
  281. ChatMemberStatus.OWNER,
  282. ChatMemberStatus.ADMINISTRATOR,
  283. }
  284. return False
  285. @app.on_callback_query(filters.regex(r"^ba:resume:[a-f0-9]{32}$"), group=-20)
  286. async def resume_business_assistant_callback(_, query):
  287. conversation_id = str(query.data).rsplit(":", 1)[-1]
  288. conversation = await get_conversation(conversation_id)
  289. if not conversation:
  290. return await query.answer("会话不存在或已清理。", show_alert=True)
  291. connection = await get_business_connection(conversation["connection_id"])
  292. if not connection:
  293. return await query.answer("Business 连接不存在。", show_alert=True)
  294. settings = await get_account_settings(connection["connection_id"])
  295. owner_id = int((connection.get("user") or {}).get("id") or 0)
  296. allowed = (
  297. query.from_user.id == owner_id
  298. or query.from_user.id in SUDOERS
  299. or await _ops_group_admin_allowed(query, settings)
  300. )
  301. if not allowed:
  302. return await query.answer("只有账号本人或运营群管理员可以恢复。", show_alert=True)
  303. await resume_conversation(conversation_id)
  304. with suppress(Exception):
  305. await query.message.edit_reply_markup(None)
  306. await query.answer("已恢复自动回复。", show_alert=True)