250 lines
9.5 KiB
Python
250 lines
9.5 KiB
Python
"""多渠道网关启动与关闭封装。
|
||
|
||
将 `lifespan.py` 中分散的渠道组件初始化、组装与销毁逻辑集中到此处,
|
||
降低应用启动器与渠道子系统内部实现的耦合。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import asyncio
|
||
from typing import TYPE_CHECKING
|
||
|
||
from yuxi.channel.config import ChannelConfigManager
|
||
from yuxi.channel.lifecycle import ChannelLifecycleManager
|
||
from yuxi.channel.message.dedupe import MessageDeduper
|
||
from yuxi.channel.message.dispatcher import ChannelGateway
|
||
from yuxi.channel.middlewares.inbound import (
|
||
CreateRunMiddleware,
|
||
DedupeMiddleware,
|
||
RouteMiddleware,
|
||
SecurityMiddleware,
|
||
SessionMiddleware,
|
||
TransactionMiddleware,
|
||
)
|
||
from yuxi.channel.middlewares.outbound import (
|
||
BuildMessageMiddleware,
|
||
ChunkMiddleware,
|
||
DowngradeMiddleware,
|
||
EnrichMiddleware,
|
||
FormatMiddleware,
|
||
MediaUploadMiddleware,
|
||
SendMiddleware,
|
||
StatusUpdateMiddleware,
|
||
)
|
||
from yuxi.channel.middlewares.registry import (
|
||
InboundMiddlewareRegistry,
|
||
OutboundMiddlewareRegistry,
|
||
)
|
||
from yuxi.channel.outbound import OutboundDispatcher
|
||
from yuxi.channel.outbound.retry import compensate_channel_messages
|
||
from yuxi.channel.plugins.loader import load_plugins
|
||
from yuxi.channel.plugins.registry import get_registry
|
||
from yuxi.channel.routing.router import BindingRouter
|
||
from yuxi.channel.security.allowlist import AllowlistChecker
|
||
from yuxi.channel.security.bot_loop import BotLoopDetector
|
||
from yuxi.channel.security.identity import IdentityLinkResolver
|
||
from yuxi.channel.security.pairing import PairingChecker, PairingManager
|
||
from yuxi.channel.security.policy import SecurityPolicy
|
||
from yuxi.channel.security.rate_limit import RateLimiter
|
||
from yuxi.channel.security.registry import SecurityCheckerRegistry
|
||
from yuxi.channel.session.manager import SessionManager
|
||
from yuxi.services.run_queue_service import get_redis_client
|
||
from yuxi.utils.logging_config import logger
|
||
|
||
if TYPE_CHECKING:
|
||
from fastapi import FastAPI
|
||
|
||
|
||
class ChannelGatewayBootstrap:
|
||
"""多渠道网关启动器:负责加载插件、组装网关、启停生命周期与后台补偿任务。"""
|
||
|
||
def __init__(self) -> None:
|
||
self._lifecycle_manager: ChannelLifecycleManager | None = None
|
||
self._outbound_dispatcher: OutboundDispatcher | None = None
|
||
self._gateway: ChannelGateway | None = None
|
||
self._security_checker_registry: SecurityCheckerRegistry | None = None
|
||
self._inbound_registry: InboundMiddlewareRegistry | None = None
|
||
self._outbound_registry: OutboundMiddlewareRegistry | None = None
|
||
self._compensate_task: asyncio.Task[None] | None = None
|
||
self._started: bool = False
|
||
|
||
async def start(self, app: FastAPI) -> None:
|
||
"""加载插件、组装并启动多渠道网关,将关键组件挂载到 ``app.state``。
|
||
|
||
启动过程中任意步骤失败都会触发回滚(``stop``),避免组件泄漏。
|
||
重复调用会被忽略。
|
||
"""
|
||
if self._started:
|
||
logger.warning("Channel gateway already started, ignoring duplicate start")
|
||
return
|
||
|
||
load_plugins()
|
||
registry = get_registry()
|
||
config_manager = ChannelConfigManager()
|
||
|
||
outbound_registry = OutboundMiddlewareRegistry()
|
||
self._outbound_registry = outbound_registry
|
||
|
||
dispatcher = OutboundDispatcher(
|
||
registry=registry,
|
||
config_manager=config_manager,
|
||
outbound_registry=outbound_registry,
|
||
)
|
||
self._outbound_dispatcher = dispatcher
|
||
|
||
outbound_registry.register(BuildMessageMiddleware())
|
||
outbound_registry.register(MediaUploadMiddleware())
|
||
outbound_registry.register(DowngradeMiddleware())
|
||
outbound_registry.register(ChunkMiddleware())
|
||
outbound_registry.register(FormatMiddleware())
|
||
outbound_registry.register(EnrichMiddleware())
|
||
outbound_registry.register(SendMiddleware())
|
||
outbound_registry.register(
|
||
StatusUpdateMiddleware(update_status=dispatcher._finalize_dispatch)
|
||
)
|
||
|
||
inbound_registry = InboundMiddlewareRegistry()
|
||
self._inbound_registry = inbound_registry
|
||
|
||
router = BindingRouter()
|
||
session_manager = SessionManager(router)
|
||
redis = await get_redis_client()
|
||
|
||
identity_resolver = await IdentityLinkResolver.create()
|
||
|
||
security_checker_registry = SecurityCheckerRegistry()
|
||
security_checker_registry.register(AllowlistChecker())
|
||
security_checker_registry.register(PairingChecker(PairingManager()))
|
||
security_checker_registry.register(BotLoopDetector())
|
||
security_checker_registry.register(RateLimiter(redis))
|
||
self._security_checker_registry = security_checker_registry
|
||
|
||
security_policy = SecurityPolicy(security_checker_registry, identity_resolver)
|
||
|
||
lifecycle_manager = ChannelLifecycleManager(registry, config_manager)
|
||
self._lifecycle_manager = lifecycle_manager
|
||
|
||
dedupe = MessageDeduper()
|
||
|
||
gateway = ChannelGateway(
|
||
registry=registry,
|
||
config_manager=config_manager,
|
||
session_manager=session_manager,
|
||
binding_router=router,
|
||
outbound_dispatcher=dispatcher,
|
||
security_policy=security_policy,
|
||
lifecycle_manager=lifecycle_manager,
|
||
inbound_registry=inbound_registry,
|
||
dedupe=dedupe,
|
||
)
|
||
self._gateway = gateway
|
||
|
||
inbound_registry.register(DedupeMiddleware(dedupe))
|
||
inbound_registry.register(SecurityMiddleware(security_policy, dedupe))
|
||
inbound_registry.register(TransactionMiddleware(dedupe))
|
||
inbound_registry.register(SessionMiddleware(session_manager))
|
||
inbound_registry.register(RouteMiddleware(router))
|
||
inbound_registry.register(
|
||
CreateRunMiddleware(
|
||
create_run=gateway.create_agent_run,
|
||
record_message=session_manager.record_message,
|
||
)
|
||
)
|
||
|
||
# 必须在 lifecycle 启动前注册 handler,否则 WebSocket/Polling 类渠道
|
||
# 在 start_all() 与 set_message_handler() 之间收到的消息会丢失。
|
||
lifecycle_manager.set_message_handler(lambda raw, ct, aid: gateway.on_transport_message(raw, ct, aid))
|
||
|
||
try:
|
||
await security_checker_registry.start_all()
|
||
await inbound_registry.start_all()
|
||
await lifecycle_manager.start_all()
|
||
await dispatcher.start()
|
||
await gateway.start_config_change_listener()
|
||
|
||
app.state.channel_lifecycle_manager = lifecycle_manager
|
||
app.state.channel_outbound_dispatcher = dispatcher
|
||
app.state.channel_gateway = gateway
|
||
|
||
self._compensate_task = asyncio.create_task(self._compensate_loop())
|
||
self._started = True
|
||
except Exception:
|
||
logger.exception("Failed to start channel gateway, rolling back")
|
||
await self.stop()
|
||
raise
|
||
|
||
async def stop(self) -> None:
|
||
"""按依赖顺序关闭多渠道网关相关组件与后台任务。
|
||
|
||
重复调用或 ``start`` 未成功时调用均为安全空操作;
|
||
若 ``start`` 中途失败,也会按实际已创建的组件进行清理。
|
||
"""
|
||
if (
|
||
not self._started
|
||
and self._lifecycle_manager is None
|
||
and self._gateway is None
|
||
and self._outbound_dispatcher is None
|
||
and self._security_checker_registry is None
|
||
):
|
||
return
|
||
|
||
self._started = False
|
||
|
||
if self._security_checker_registry is not None:
|
||
try:
|
||
await self._security_checker_registry.stop_all()
|
||
except Exception:
|
||
logger.exception("Failed to stop security checker registry")
|
||
finally:
|
||
self._security_checker_registry = None
|
||
|
||
if self._inbound_registry is not None:
|
||
try:
|
||
await self._inbound_registry.stop_all()
|
||
except Exception:
|
||
logger.exception("Failed to stop inbound middleware registry")
|
||
finally:
|
||
self._inbound_registry = None
|
||
|
||
if self._compensate_task is not None:
|
||
try:
|
||
self._compensate_task.cancel()
|
||
try:
|
||
await self._compensate_task
|
||
except asyncio.CancelledError:
|
||
pass
|
||
except Exception:
|
||
logger.exception("Failed to stop channel compensate task")
|
||
finally:
|
||
self._compensate_task = None
|
||
|
||
# gateway.stop() 内部已会关闭 lifecycle,避免重复停止。
|
||
if self._gateway is not None:
|
||
try:
|
||
await self._gateway.stop()
|
||
except Exception:
|
||
logger.exception("Failed to stop channel gateway")
|
||
elif self._lifecycle_manager is not None:
|
||
try:
|
||
await self._lifecycle_manager.stop_all()
|
||
except Exception:
|
||
logger.exception("Failed to stop channel lifecycle manager")
|
||
|
||
if self._outbound_dispatcher is not None:
|
||
try:
|
||
await self._outbound_dispatcher.stop()
|
||
except Exception:
|
||
logger.exception("Failed to stop channel outbound dispatcher")
|
||
|
||
self._gateway = None
|
||
self._lifecycle_manager = None
|
||
self._outbound_dispatcher = None
|
||
|
||
async def _compensate_loop(self) -> None:
|
||
while True:
|
||
await asyncio.sleep(60)
|
||
try:
|
||
await compensate_channel_messages()
|
||
except Exception:
|
||
logger.exception("Channel compensate loop error")
|