supervisor.py 9.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237
  1. from __future__ import annotations
  2. import asyncio
  3. import os
  4. import re
  5. import secrets
  6. import socket
  7. import sys
  8. from dataclasses import dataclass
  9. from datetime import UTC, datetime
  10. from pathlib import Path
  11. from typing import Any
  12. from aiohttp import ClientError, ClientSession, ClientTimeout
  13. from wbb.admin.bot_config import (
  14. BotConfigError,
  15. get_bot_profile_secrets,
  16. telegram_config_status,
  17. )
  18. @dataclass
  19. class BotProcess:
  20. bot_id: str
  21. port: int
  22. process: asyncio.subprocess.Process
  23. state: str = "starting"
  24. started_at: str = ""
  25. error: str = ""
  26. expected_stop: bool = False
  27. monitor_task: asyncio.Task[None] | None = None
  28. readiness_task: asyncio.Task[None] | None = None
  29. class BotSupervisor:
  30. def __init__(
  31. self,
  32. config_path: str | Path,
  33. *,
  34. project_root: str | Path,
  35. control_port: int | None = None,
  36. ) -> None:
  37. self.config_path = Path(config_path).resolve()
  38. self.project_root = Path(project_root).resolve()
  39. self.control_port = int(
  40. control_port or os.environ.get("ADMIN_WEB_PORT", 8088)
  41. )
  42. self.internal_token = secrets.token_urlsafe(32)
  43. self._processes: dict[str, BotProcess] = {}
  44. self._lock = asyncio.Lock()
  45. def runtimes(self) -> dict[str, dict[str, Any]]:
  46. values: dict[str, dict[str, Any]] = {}
  47. for bot_id, item in self._processes.items():
  48. state = item.state
  49. if item.process.returncode is not None and state not in {"failed", "stopped"}:
  50. state = "stopped" if item.expected_stop else "failed"
  51. values[bot_id] = {
  52. "state": state,
  53. "pid": item.process.pid,
  54. "port": item.port,
  55. "started_at": item.started_at,
  56. "error": item.error,
  57. }
  58. return values
  59. def endpoint_for(self, bot_id: str) -> str | None:
  60. item = self._processes.get(str(bot_id))
  61. if not item or item.state != "running" or item.process.returncode is not None:
  62. return None
  63. return f"http://127.0.0.1:{item.port}"
  64. async def start_enabled(self) -> None:
  65. status = telegram_config_status(self.config_path)
  66. for profile in status["bots"]:
  67. if profile["enabled"] and profile["ready_to_connect"]:
  68. await self.start(profile["bot_id"])
  69. async def start(self, bot_id: str) -> dict[str, Any]:
  70. async with self._lock:
  71. existing = self._processes.get(str(bot_id))
  72. if existing and existing.process.returncode is None:
  73. return self.runtimes()[str(bot_id)]
  74. profile = get_bot_profile_secrets(self.config_path, bot_id)
  75. if not profile.get("enabled", True):
  76. raise BotConfigError("bot_disabled", "该机器人已停用。")
  77. if not profile.get("api_id") or not profile.get("api_hash"):
  78. raise BotConfigError(
  79. "telegram_api_required",
  80. "请先配置 Telegram 应用编号和应用密钥。",
  81. )
  82. port = self._available_port()
  83. environment = self._worker_environment(profile, port)
  84. process = await asyncio.create_subprocess_exec(
  85. sys.executable,
  86. str(self.project_root / "launcher.py"),
  87. "--worker",
  88. str(bot_id),
  89. "--port",
  90. str(port),
  91. cwd=self.project_root,
  92. env=environment,
  93. )
  94. item = BotProcess(
  95. bot_id=str(bot_id),
  96. port=port,
  97. process=process,
  98. started_at=datetime.now(UTC).isoformat().replace("+00:00", "Z"),
  99. )
  100. self._processes[str(bot_id)] = item
  101. item.monitor_task = asyncio.create_task(self._monitor(item))
  102. item.readiness_task = asyncio.create_task(self._wait_until_ready(item))
  103. return self.runtimes()[str(bot_id)]
  104. async def stop(self, bot_id: str) -> dict[str, Any]:
  105. async with self._lock:
  106. item = self._processes.get(str(bot_id))
  107. if not item or item.process.returncode is not None:
  108. return {"state": "stopped"}
  109. item.expected_stop = True
  110. item.state = "stopping"
  111. item.process.terminate()
  112. try:
  113. await asyncio.wait_for(item.process.wait(), timeout=10)
  114. except TimeoutError:
  115. item.process.kill()
  116. await item.process.wait()
  117. item.state = "stopped"
  118. return self.runtimes()[str(bot_id)]
  119. async def restart(self, bot_id: str) -> dict[str, Any]:
  120. await self.stop(bot_id)
  121. return await self.start(bot_id)
  122. async def reconcile(self, bot_id: str) -> dict[str, Any]:
  123. profile = get_bot_profile_secrets(self.config_path, bot_id)
  124. if profile.get("enabled", True) and profile.get("api_id") and profile.get("api_hash"):
  125. return await self.restart(bot_id)
  126. return await self.stop(bot_id)
  127. async def close(self) -> None:
  128. for bot_id in list(self._processes):
  129. await self.stop(bot_id)
  130. async def _monitor(self, item: BotProcess) -> None:
  131. return_code = await item.process.wait()
  132. if item.readiness_task and not item.readiness_task.done():
  133. item.readiness_task.cancel()
  134. if item.expected_stop:
  135. item.state = "stopped"
  136. return
  137. item.state = "failed"
  138. item.error = f"Bot 工作进程退出,状态码 {return_code}。"
  139. async def _wait_until_ready(self, item: BotProcess) -> None:
  140. timeout = ClientTimeout(total=2)
  141. async with ClientSession(timeout=timeout) as session:
  142. for _ in range(45):
  143. if item.process.returncode is not None:
  144. return
  145. try:
  146. async with session.get(
  147. f"http://127.0.0.1:{item.port}/api/admin/v1/health"
  148. ) as response:
  149. if response.status == 200:
  150. item.state = "running"
  151. item.error = ""
  152. return
  153. except (ClientError, TimeoutError):
  154. pass
  155. await asyncio.sleep(1)
  156. if item.process.returncode is None:
  157. item.state = "failed"
  158. item.error = "Bot 已连接,但内部管理接口未在限定时间内就绪。"
  159. def _worker_environment(self, profile: dict[str, Any], port: int) -> dict[str, str]:
  160. environment = os.environ.copy()
  161. assistant = (
  162. profile.get("business_assistant")
  163. if isinstance(profile.get("business_assistant"), dict)
  164. else {}
  165. )
  166. environment.update(
  167. {
  168. "WBB_ADMIN_BOOTSTRAP": "0",
  169. "WBB_SUPERVISOR_MODE": "0",
  170. "WBB_BOT_WORKER": "1",
  171. "WBB_BOT_PROFILE_ID": str(profile["bot_id"]),
  172. "WBB_BOT_DATABASE": self._worker_database_name(profile["bot_id"]),
  173. "WBB_BOT_PERMISSIONS": ",".join(profile.get("permissions", [])),
  174. "WBB_SUPERVISOR_INTERNAL_URL": (
  175. f"http://127.0.0.1:{self.control_port}"
  176. ),
  177. "WBB_INTERNAL_TOKEN": self.internal_token,
  178. "BOT_TOKEN": str(profile["bot_token"]),
  179. "API_ID": str(profile["api_id"]),
  180. "API_HASH": str(profile["api_hash"]),
  181. "USERBOT_ENABLED": "0",
  182. "SESSION_STRING": "",
  183. "PHONE_NUMBER": "",
  184. "SUDO_USERS_ID": " ".join(str(item) for item in profile["sudo_users_id"]),
  185. "LOG_GROUP_ID": str(profile["log_group_id"]),
  186. "GBAN_LOG_GROUP_ID": str(profile["gban_log_group_id"]),
  187. "MESSAGE_DUMP_CHAT": str(profile["message_dump_chat"]),
  188. "ADMIN_WEB_ENABLED": "1",
  189. "ADMIN_WEB_HOST": "127.0.0.1",
  190. "ADMIN_WEB_PORT": str(port),
  191. "BOT_PROFILES_PATH": str(self.config_path),
  192. "BUSINESS_ASSISTANT_OPENAI_BASE_URL": str(assistant.get("base_url") or ""),
  193. "BUSINESS_ASSISTANT_OPENAI_API_KEY": str(assistant.get("api_key") or ""),
  194. "BUSINESS_ASSISTANT_OPENAI_MODEL": str(assistant.get("model") or ""),
  195. "BUSINESS_ASSISTANT_OPENAI_TIMEOUT_SECONDS": str(
  196. assistant.get("timeout_seconds") or 30
  197. ),
  198. "BUSINESS_ASSISTANT_OPENAI_MAX_OUTPUT_TOKENS": str(
  199. assistant.get("max_output_tokens") or 600
  200. ),
  201. }
  202. )
  203. return environment
  204. @staticmethod
  205. def _worker_database_name(bot_id: str) -> str:
  206. storage_key = re.sub(r"[^A-Za-z0-9_-]", "_", str(bot_id))[:48]
  207. if not storage_key:
  208. raise BotConfigError(
  209. "invalid_bot_id",
  210. "机器人编号无法用于创建独立数据空间。",
  211. )
  212. return f"wbb_bot_{storage_key}"
  213. @staticmethod
  214. def _available_port() -> int:
  215. with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
  216. sock.bind(("127.0.0.1", 0))
  217. return int(sock.getsockname()[1])