ForcePilot/backend/package/yuxi/channel/lifecycle/manager.py
Kris bab30f2715
Some checks failed
Deploy VitePress site to Pages / build (push) Has been cancelled
Ruff Format Check / Ruff Format & Lint (push) Has been cancelled
Deploy VitePress site to Pages / Deploy (push) Has been cancelled
feat:0715
2026-07-15 12:30:58 +08:00

404 lines
15 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""渠道生命周期管理器"""
from __future__ import annotations
import asyncio
import os
from collections.abc import Awaitable, Callable
from datetime import UTC, datetime
from typing import TYPE_CHECKING
from yuxi.channel.exceptions import ChannelTransportReconnectRequested
from yuxi.channel.lifecycle.backoff import BackoffPolicy
from yuxi.channel.lifecycle.context import ChannelLifecycleContext
from yuxi.channel.metrics import channel_active_connections, channel_reconnect_total
from yuxi.channel.plugins.protocol import (
ChannelPlugin,
Transport,
TransportState,
TransportType,
)
from yuxi.channel.plugins.registry import ChannelRegistry
from yuxi.channel.transport.qr_login_transport import create_channel_transport
from yuxi.utils.logging_config import logger
if TYPE_CHECKING:
from yuxi.channel.config import ChannelConfigManager
def _build_webhook_callback_url(channel_type: str) -> str:
"""构造渠道 Webhook 回调 URL优先使用环境变量 CHANNEL_WEBHOOK_CALLBACK_URL。"""
base_url = os.environ.get("CHANNEL_WEBHOOK_CALLBACK_URL", "http://localhost:5050").rstrip("/")
return f"{base_url}/api/channels/{channel_type}/webhook"
class ChannelInstance:
"""单个渠道账户的运行期实例,封装 config、plugin、transport 与错误状态。"""
def __init__(
self,
channel_type: str,
account_id: str,
config: dict,
plugin: ChannelPlugin,
message_handler: Callable[[bytes, str, str], Awaitable[None]] | None = None,
):
self.channel_type = channel_type
self.account_id = account_id
self.config = config
self.plugin = plugin
self.context = ChannelLifecycleContext()
self._closed = False
self._transport: Transport | None = None
self._no_transport = False
self._message_handler = message_handler
self._exit_event = asyncio.Event()
self._exit_event.clear()
self._reconnect_requested = False
self._dispatch_semaphore = asyncio.Semaphore(32)
@property
def transport(self) -> Transport | None:
return self._transport
@property
def no_transport(self) -> bool:
return self._no_transport
@property
def last_error(self) -> str | None:
return self.context.last_error
@property
def reconnect_attempts(self) -> int:
return self.context.reconnect_attempts
def is_closed(self) -> bool:
return self._closed
def record_error(self, exc: Exception) -> None:
self.context.last_error = str(exc)
self.context.reconnect_attempts += 1
async def start(self) -> None:
if self._transport is not None:
await self.stop(timeout=5.0)
transport = await create_channel_transport(self.plugin, self.config, self.account_id)
if transport is None:
self._transport = None
self._no_transport = True
return
self._transport = transport
self._no_transport = False
if self._message_handler is not None:
self._register_transport_message_handler()
await self._transport.start()
def set_message_handler(
self,
message_handler: Callable[[bytes, str, str], Awaitable[None]] | None,
) -> None:
self._message_handler = message_handler
if self._transport is not None:
self._register_transport_message_handler()
def _register_transport_message_handler(self) -> None:
if self._transport is None or self._message_handler is None:
return
async def _handle_raw(raw: bytes) -> None:
async def _dispatch() -> None:
async with self._dispatch_semaphore:
try:
await self._message_handler(raw, self.channel_type, self.account_id)
self.context.last_message_at = datetime.now(UTC)
except ChannelTransportReconnectRequested:
self._reconnect_requested = True
logger.info(
"Transport reconnect requested for %s:%s",
self.channel_type,
self.account_id,
)
if self._transport is not None:
try:
await self._transport.stop()
except Exception:
logger.exception(
"Failed to stop transport for reconnect %s:%s",
self.channel_type,
self.account_id,
)
except Exception:
logger.exception(
"Message handler error for %s:%s",
self.channel_type,
self.account_id,
)
try:
asyncio.create_task(_dispatch())
except Exception:
logger.exception(
"Failed to schedule message handler for %s:%s",
self.channel_type,
self.account_id,
)
self._transport.on_message(_handle_raw)
async def stop(self, timeout: float = 5.0) -> None:
if self._transport is None:
return
try:
await asyncio.wait_for(self._transport.stop(), timeout=timeout)
except TimeoutError:
logger.warning(
"Timeout stopping transport for %s:%s",
self.channel_type,
self.account_id,
)
finally:
self._transport = None
async def wait_for_exit(self) -> None:
"""等待 transport 断开或被关闭。Webhook 等无 transport 渠道会阻塞到 close。"""
if self._transport is None:
await self._exit_event.wait()
return
while not self._closed and self._transport is not None:
if self._transport.state in (TransportState.STOPPED, TransportState.DISCONNECTED):
break
try:
await asyncio.wait_for(self._exit_event.wait(), timeout=0.5)
break
except TimeoutError:
continue
def close(self) -> None:
self._closed = True
self._exit_event.set()
class ChannelLifecycleManager:
"""管理所有渠道账户的生命周期:并发启动、异常重连、优雅停止。"""
def __init__(
self,
registry: ChannelRegistry,
config_manager: ChannelConfigManager,
message_handler: Callable[[bytes, str, str], Awaitable[None]] | None = None,
):
self.registry = registry
self.config_manager = config_manager
self._message_handler = message_handler
self._instances: dict[str, ChannelInstance] = {}
self._tasks: dict[str, asyncio.Task] = {}
self._backoff = BackoffPolicy()
self._semaphore = asyncio.Semaphore(4)
self._key_locks: dict[str, asyncio.Lock] = {}
def set_message_handler(
self,
handler: Callable[[bytes, str, str], Awaitable[None]] | None,
) -> None:
self._message_handler = handler
for instance in self._instances.values():
instance.set_message_handler(handler)
def _instance_key(self, channel_type: str, account_id: str) -> str:
return f"{channel_type}:{account_id}"
def _get_key_lock(self, key: str) -> asyncio.Lock:
if key not in self._key_locks:
self._key_locks[key] = asyncio.Lock()
return self._key_locks[key]
def get_transport(self, channel_type: str, account_id: str) -> Transport | None:
"""获取指定账户当前运行的 transport 实例。"""
instance = self._instances.get(self._instance_key(channel_type, account_id))
return instance.transport if instance is not None else None
async def start_all(self) -> None:
configs = await self.config_manager.list_enabled_accounts()
await asyncio.gather(
*(self._start_with_semaphore(config) for config in configs),
return_exceptions=True,
)
async def _start_with_semaphore(self, config: dict) -> None:
async with self._semaphore:
await self.start_channel(config["channel_type"], config["account_id"])
async def start_channel(self, channel_type: str, account_id: str) -> None:
instance_key = self._instance_key(channel_type, account_id)
async with self._get_key_lock(instance_key):
if instance_key in self._instances:
return
plugin = self.registry.get_plugin(channel_type)
if plugin is None:
logger.warning(
"No plugin registered for channel type: %s",
channel_type,
)
return
try:
config = await self.config_manager.get_config(channel_type, account_id)
except Exception:
logger.exception(
"Failed to load config for %s",
instance_key,
)
return
instance = ChannelInstance(
channel_type,
account_id,
config,
plugin,
message_handler=self._message_handler,
)
self._instances[instance_key] = instance
self._tasks[instance_key] = asyncio.create_task(
self._run_loop(instance), name=f"channel-lifecycle:{instance_key}"
)
try:
await asyncio.wait_for(
plugin.run_startup_maintenance(config, account_id),
timeout=30.0,
)
except Exception:
logger.exception(
"run_startup_maintenance failed for %s:%s",
channel_type,
account_id,
)
if plugin.get_meta().transport_type == TransportType.WEBHOOK:
try:
callback_url = _build_webhook_callback_url(channel_type)
ok = await plugin.setup_webhook(config, callback_url)
logger.info(
"setup_webhook for %s:%s returned %s (callback_url=%s)",
channel_type,
account_id,
ok,
callback_url,
)
except Exception:
logger.exception(
"setup_webhook failed for %s:%s",
channel_type,
account_id,
)
try:
await plugin.on_channel_enabled(config, account_id)
except Exception:
logger.exception(
"on_channel_enabled failed for %s:%s",
channel_type,
account_id,
)
logger.info("Started lifecycle for %s", instance_key)
async def stop_channel(self, channel_type: str, account_id: str) -> None:
instance_key = self._instance_key(channel_type, account_id)
async with self._get_key_lock(instance_key):
instance = self._instances.get(instance_key)
if instance is None:
return
instance.close()
await instance.stop(timeout=5.0)
task = self._tasks.pop(instance_key, None)
if task is not None:
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
del self._instances[instance_key]
if instance.plugin.get_meta().transport_type == TransportType.WEBHOOK:
try:
ok = await instance.plugin.delete_webhook(instance.config)
logger.info(
"delete_webhook for %s:%s returned %s",
channel_type,
account_id,
ok,
)
except Exception:
logger.exception(
"delete_webhook failed for %s:%s",
channel_type,
account_id,
)
try:
await instance.plugin.on_channel_disabled(instance.config, account_id)
except Exception:
logger.exception(
"on_channel_disabled failed for %s:%s",
channel_type,
account_id,
)
logger.info("Stopped lifecycle for %s", instance_key)
async def stop_all(self) -> None:
instances = list(self._instances.values())
await asyncio.gather(
*(self.stop_channel(instance.channel_type, instance.account_id) for instance in instances),
return_exceptions=True,
)
async def _run_loop(self, instance: ChannelInstance) -> None:
labels = {
"channel_type": instance.channel_type,
"account_id": instance.account_id,
}
while not instance.is_closed():
connection_counted = False
try:
await instance.start()
if instance.no_transport:
break
instance.context.last_connected_at = datetime.now(UTC)
instance.context.reconnect_attempts = 0
if instance.transport is not None:
channel_active_connections.inc(labels)
connection_counted = True
await instance.wait_for_exit()
await instance.stop(timeout=5.0)
if instance.is_closed():
break
if instance._reconnect_requested:
instance._reconnect_requested = False
instance.context.reconnect_attempts = 0
continue
raise ConnectionError("transport disconnected")
except asyncio.CancelledError:
break
except Exception as exc:
instance.record_error(exc)
channel_reconnect_total.inc(labels)
delay_ms = self._backoff.compute(instance.reconnect_attempts)
logger.exception(
"Channel %s:%s error, retrying in %sms",
instance.channel_type,
instance.account_id,
delay_ms,
)
await instance.stop(timeout=5.0)
try:
await asyncio.sleep(delay_ms / 1000)
except asyncio.CancelledError:
break
finally:
if connection_counted:
channel_active_connections.dec(labels)