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")
|