ForcePilot/backend/package/yuxi/channel/container.py
Kris 9a8a27bf36 feat(channel): 新增渠道网关模块完整实现
本次提交新增了完整的多渠道消息网关系统,包括:
1. 支持飞书、钉钉、Web、Hook 四种渠道的适配器与配置
2. 领域模型层:消息、会话、绑定、出箱等核心实体
3. 应用服务层:管道、中间件、DTO 与业务逻辑
4. 基础设施层:持久化、过滤器、队列等端口实现
5. 接口层:REST API、SSE、WebSocket 通信端点
6. 前端页面与路由配置,添加渠道管理菜单
7. 新增相关依赖包与 docker-compose 部署配置
2026-05-30 21:53:09 +08:00

573 lines
22 KiB
Python

from __future__ import annotations
import asyncio
import logging
import os
import time
from dataclasses import dataclass, field
import redis.asyncio as aioredis
from fastapi import HTTPException, Request
from yuxi.channel.application.pipeline.builder import build_inbound_pipeline
from yuxi.channel.application.service.auth_service import AuthService
from yuxi.channel.application.service.binding_service import BindingService
from yuxi.channel.application.service.config_service import ConfigService
from yuxi.channel.application.service.delivery_service import DeliveryService
from yuxi.channel.application.service.dispatch_service import DispatchService
from yuxi.channel.application.service.inbound_service import InboundService
from yuxi.channel.application.service.session_resolver import SessionResolver
from yuxi.channel.channels._registry import get_registered_channels
from yuxi.channel.domain.port.agent_port import AgentPort
from yuxi.channel.domain.port.cache_port import CachePort
from yuxi.channel.domain.port.channel_adapter_port import ChannelAdapterPort
from yuxi.channel.domain.port.config_reload_port import ConfigReloadPort
from yuxi.channel.domain.port.content_filter_port import ContentFilterPort
from yuxi.channel.domain.port.event_publisher_port import EventPublisherPort
from yuxi.channel.domain.port.metrics_port import MetricsPort
from yuxi.channel.domain.port.queue_port import QueuePort
from yuxi.channel.domain.port.rate_limit_port import RateLimitPort
from yuxi.channel.domain.port.signature_verify_port import SignatureVerifyPort
from yuxi.channel.domain.port.bot_loop_guard_port import BotLoopGuardPort
from yuxi.channel.domain.port.circuit_breaker_port import CircuitBreakerPort
from yuxi.channel.domain.repository.binding_repository import BindingRepositoryPort
from yuxi.channel.domain.repository.message_log_repository import MessageLogRepositoryPort
from yuxi.channel.domain.repository.message_repository import MessageRepositoryPort
from yuxi.channel.domain.repository.outbox_repository import OutboxRepositoryPort
from yuxi.channel.domain.repository.session_repository import SessionRepositoryPort
from yuxi.channel.domain.service.pipeline import Pipeline
from yuxi.channel.infrastructure.agent.agent_adapter import AgentAdapter
from yuxi.channel.infrastructure.cache.redis_bot_loop_guard import RedisBotLoopGuard
from yuxi.channel.infrastructure.cache.redis_cache import RedisCache
from yuxi.channel.infrastructure.cache.redis_circuit_breaker import RedisCircuitBreaker
from yuxi.channel.infrastructure.cache.redis_rate_limiter import RedisRateLimiter
from yuxi.channel.infrastructure.config.channel_config import ChannelConfig
from yuxi.channel.infrastructure.config.redis_config_reload import RedisConfigReload
from yuxi.channel.infrastructure.content_filter.composite_content_filter import CompositeContentFilter
from yuxi.channel.infrastructure.content_filter.llm_content_filter import LlmContentFilter
from yuxi.channel.infrastructure.content_filter.redis_content_filter import RedisContentFilter
from yuxi.channel.infrastructure.messaging.redis_pubsub_publisher import RedisPubSubPublisher
from yuxi.channel.infrastructure.messaging.redis_pubsub_subscriber import RedisPubSubSubscriber
from yuxi.channel.infrastructure.messaging.redis_stream_queue import RedisStreamQueue
from yuxi.channel.infrastructure.metrics.prometheus_metrics import PrometheusMetricsAdapter
from yuxi.channel.infrastructure.persistence.repository.pg_binding_repository import PgBindingRepository
from yuxi.channel.infrastructure.persistence.repository.pg_message_log_repository import PgMessageLogRepository
from yuxi.channel.infrastructure.persistence.repository.pg_message_repository import PgMessageRepository
from yuxi.channel.infrastructure.persistence.repository.pg_outbox_repository import PgOutboxRepository
from yuxi.channel.infrastructure.persistence.repository.pg_session_repository import PgSessionRepository
from yuxi.channel.infrastructure.security.hmac_signature_verifier import HmacSignatureVerifier
from yuxi.channel.interfaces.sse.endpoint import SseEndpoint
from yuxi.channel.interfaces.websocket.manager import WsConnectionManager
from yuxi.channel.worker.outbox_retry import OutboxRetryWorker
from yuxi.channel.worker.pool import WorkerPool, WorkerPoolConfig
from yuxi.channel.worker.session_factory import WorkerSessionFactory
from yuxi.repositories.agent_config_repository import AgentConfigRepository
from yuxi.storage.postgres.manager import pg_manager
logger = logging.getLogger(__name__)
@dataclass
class StartupTracer:
_events: list[dict] = field(default_factory=list)
_start_time: float = 0.0
_current_phase: str = ""
def begin(self, phase: str) -> None:
self._start_time = time.monotonic()
self._current_phase = phase
def end(self) -> None:
duration = (time.monotonic() - self._start_time) * 1000
self._events.append(
{
"name": self._current_phase,
"status": "ok",
"duration_ms": round(duration, 1),
}
)
def mark(self, name: str, status: str, *, duration_ms: float = 0.0) -> None:
self._events.append({"name": name, "status": status, "duration_ms": duration_ms})
@property
def total_ms(self) -> float:
return sum(e["duration_ms"] for e in self._events)
def to_dict(self) -> list[dict]:
return list(self._events)
@dataclass
class _InfraBundle:
queue_port: QueuePort
event_publisher: EventPublisherPort
content_filter: ContentFilterPort
config_reload: ConfigReloadPort
cache_port: CachePort
rate_limit_port: RateLimitPort
bot_loop_guard: BotLoopGuardPort
circuit_breaker: CircuitBreakerPort
metrics: MetricsPort
signature_verifier: SignatureVerifyPort
@dataclass
class _WorkerBundle:
dispatch_service: DispatchService
binding_service: BindingService
config_service: ConfigService
worker_pool: WorkerPool
outbox_worker: OutboxRetryWorker
session_factory: WorkerSessionFactory
binding_repo: BindingRepositoryPort
session_repo: SessionRepositoryPort
outbox_repo: OutboxRepositoryPort
message_repo: MessageRepositoryPort
message_log_repo: MessageLogRepositoryPort
@dataclass
class ChannelContainer:
pipeline: Pipeline
inbound_service: InboundService
dispatch_service: DispatchService
delivery_service: DeliveryService
session_resolver: SessionResolver
binding_service: BindingService
config_service: ConfigService
auth_service: AuthService
channel_config: ChannelConfig
worker_pool: WorkerPool
outbox_worker: OutboxRetryWorker
session_factory: WorkerSessionFactory
adapters: dict[str, ChannelAdapterPort] = field(default_factory=dict)
sse_endpoint: SseEndpoint | None = None
ws_manager: WsConnectionManager | None = None
binding_repo: BindingRepositoryPort | None = None
session_repo: SessionRepositoryPort | None = None
outbox_repo: OutboxRepositoryPort | None = None
message_repo: MessageRepositoryPort | None = None
message_log_repo: MessageLogRepositoryPort | None = None
queue_port: QueuePort | None = None
event_publisher: EventPublisherPort | None = None
content_filter: ContentFilterPort | None = None
config_reload: ConfigReloadPort | None = None
cache_port: CachePort | None = None
rate_limit_port: RateLimitPort | None = None
bot_loop_guard: BotLoopGuardPort | None = None
circuit_breaker: CircuitBreakerPort | None = None
metrics: MetricsPort | None = None
signature_verifier: SignatureVerifyPort | None = None
startup_tracer: StartupTracer = field(default_factory=StartupTracer)
redis: aioredis.Redis | None = None
pubsub_subscriber: RedisPubSubSubscriber | None = None
async def shutdown(self) -> None:
async def _safe(coro, label: str):
try:
await coro
except Exception:
logger.exception("shutdown %s failed", label)
if self.event_publisher:
try:
from yuxi.channel.domain.event.gateway_shutdown import GatewayShutdown
await self.event_publisher.publish(GatewayShutdown())
except Exception:
logger.debug("gateway shutdown event publish failed")
if self.sse_endpoint:
await _safe(self.sse_endpoint.broadcast_shutdown(), "sse_broadcast")
if self.pubsub_subscriber:
await _safe(self.pubsub_subscriber.stop(), "pubsub_subscriber")
if self.worker_pool:
await _safe(self.worker_pool.stop(), "worker_pool")
if self.outbox_worker:
await _safe(self.outbox_worker.stop(), "outbox_worker")
if self.ws_manager:
await _safe(self.ws_manager.stop_all(), "ws_connections")
if self.sse_endpoint:
await _safe(self.sse_endpoint.stop(), "sse_endpoint")
for name, adapter in self.adapters.items():
await _safe(adapter.close(), f"adapter:{name}")
if self.redis:
await _safe(self.redis.aclose(), "redis")
logger.info("channel shutdown complete")
def setup_channel(app: object, container: ChannelContainer) -> None:
from fastapi import FastAPI
if isinstance(app, FastAPI):
app.state.channel = container
def get_channel(request: Request) -> ChannelContainer:
container = getattr(request.app.state, "channel", None)
if not container:
raise HTTPException(status_code=503, detail="channel not initialized")
return container
async def _resolve_agent_id(session_factory, agent_config_id: int) -> str:
async with session_factory() as session:
repo = AgentConfigRepository(session)
config = await repo.get_by_id(config_id=agent_config_id)
return config.agent_id if config else "chatbot"
class ChannelContainerFactory:
@staticmethod
async def create(
redis_url: str,
agent_port: AgentPort | None = None,
*,
config_yaml_path: str = "channel_config.yaml",
mq_workers: int = 4,
mq_max_concurrent: int = 20,
default_agent_config_id: int = 1,
) -> ChannelContainer:
tracer = StartupTracer()
tracer.begin("config")
config = ChannelContainerFactory._build_config(config_yaml_path)
tracer.end()
tracer.begin("redis")
redis = ChannelContainerFactory._connect_redis(redis_url)
tracer.end()
tracer.begin("infra")
infra = ChannelContainerFactory._build_infra(redis, config)
tracer.end()
tracer.begin("auth")
auth_service = ChannelContainerFactory._build_auth_service(infra, config)
tracer.end()
tracer.begin("pipeline")
pipeline = ChannelContainerFactory._build_pipeline(infra, config, auth_service)
tracer.end()
tracer.begin("adapters")
adapters, sse_endpoint = await ChannelContainerFactory._build_adapters(config, infra)
tracer.end()
ChannelContainerFactory._inject_bot_ids(pipeline, adapters)
tracer.begin("ws_connections")
ws_manager = await ChannelContainerFactory._build_ws_connections(
adapters, pipeline, tracer, metrics=infra.metrics
)
tracer.end()
tracer.begin("workers")
workers = ChannelContainerFactory._build_workers(
infra,
adapters,
agent_port,
config,
pipeline,
redis,
default_agent_config_id=default_agent_config_id,
mq_workers=mq_workers,
mq_max_concurrent=mq_max_concurrent,
)
await workers.worker_pool.start()
await workers.outbox_worker.start()
tracer.end()
inbound_service = InboundService(
pipeline,
message_log_repo=workers.message_log_repo,
event_publisher=infra.event_publisher,
)
pubsub_subscriber = RedisPubSubSubscriber(redis, sse_endpoint)
await pubsub_subscriber.start()
container = ChannelContainer(
pipeline=pipeline,
adapters=adapters,
inbound_service=inbound_service,
dispatch_service=workers.dispatch_service,
delivery_service=workers.dispatch_service.delivery_service,
session_resolver=workers.dispatch_service.session_resolver,
binding_service=workers.binding_service,
config_service=workers.config_service,
auth_service=auth_service,
channel_config=config,
worker_pool=workers.worker_pool,
outbox_worker=workers.outbox_worker,
session_factory=workers.session_factory,
sse_endpoint=sse_endpoint,
ws_manager=ws_manager,
binding_repo=workers.binding_repo,
session_repo=workers.session_repo,
outbox_repo=workers.outbox_repo,
message_repo=workers.message_repo,
message_log_repo=workers.message_log_repo,
queue_port=infra.queue_port,
event_publisher=infra.event_publisher,
content_filter=infra.content_filter,
config_reload=infra.config_reload,
cache_port=infra.cache_port,
rate_limit_port=infra.rate_limit_port,
bot_loop_guard=infra.bot_loop_guard,
circuit_breaker=infra.circuit_breaker,
metrics=infra.metrics,
signature_verifier=infra.signature_verifier,
startup_tracer=tracer,
redis=redis,
pubsub_subscriber=pubsub_subscriber,
)
if infra.config_reload:
async def _on_config_change(config: dict):
await container.config_service.reload()
asyncio.create_task(infra.config_reload.watch(_on_config_change))
logger.info(
"channel initialized: pipeline=%d middlewares, workers=%d, startup=%.1fms",
len(pipeline.middlewares),
mq_workers,
tracer.total_ms,
)
return container
@staticmethod
def _build_config(yaml_path: str) -> ChannelConfig:
return ChannelConfig(yaml_path)
@staticmethod
def _connect_redis(redis_url: str) -> aioredis.Redis:
return aioredis.from_url(redis_url, decode_responses=False)
@staticmethod
def _build_infra(redis: aioredis.Redis, config: ChannelConfig) -> _InfraBundle:
cache_port = RedisCache(redis)
rate_limit_port = RedisRateLimiter(redis)
bot_loop_guard = RedisBotLoopGuard(cache_port)
circuit_breaker = RedisCircuitBreaker(cache_port)
queue_port = RedisStreamQueue(redis)
event_publisher = RedisPubSubPublisher(redis)
redis_filter = RedisContentFilter(redis)
llm_filter = LlmContentFilter(fallback=redis_filter)
content_filter = CompositeContentFilter([llm_filter])
config_reload = RedisConfigReload(redis)
metrics = PrometheusMetricsAdapter()
signature_verifier = HmacSignatureVerifier(secret=os.getenv("WEBHOOK_SECRET", ""))
return _InfraBundle(
queue_port=queue_port,
event_publisher=event_publisher,
content_filter=content_filter,
config_reload=config_reload,
cache_port=cache_port,
rate_limit_port=rate_limit_port,
bot_loop_guard=bot_loop_guard,
circuit_breaker=circuit_breaker,
metrics=metrics,
signature_verifier=signature_verifier,
)
@staticmethod
def _build_auth_service(infra: _InfraBundle, config: ChannelConfig) -> AuthService:
return AuthService(
infra.rate_limit_port,
token=config.auth_token,
password=config.auth_password,
max_attempts=config.max_auth_attempts,
lockout_seconds=config.lockout_seconds,
)
@staticmethod
def _build_pipeline(
infra: _InfraBundle,
config: ChannelConfig,
auth_service: AuthService,
) -> Pipeline:
return build_inbound_pipeline(
auth_service=auth_service,
cache_port=infra.cache_port,
rate_limit_port=infra.rate_limit_port,
queue_port=infra.queue_port,
signature_verifier=infra.signature_verifier,
metrics=infra.metrics,
channel_config=config.raw_data,
)
@staticmethod
async def _build_adapters(
config: ChannelConfig,
infra: _InfraBundle,
) -> tuple[dict[str, ChannelAdapterPort], SseEndpoint]:
import yuxi.channel.channels # noqa: F401
sse_endpoint = SseEndpoint(metrics=infra.metrics)
await sse_endpoint.start()
adapters: dict[str, ChannelAdapterPort] = {}
for name, cls in get_registered_channels().items():
default_config = cls.get_default_config()
override = config.get_channel_config(name)
merged = {**default_config, **override}
if name == "web":
merged["sse_push"] = sse_endpoint
if name == "hooks":
hook_mappings = config.hook_mappings
if hook_mappings:
merged["mappings"] = hook_mappings
env_overrides = ChannelContainerFactory._apply_env_overrides(name, merged)
merged.update(env_overrides)
adapter = cls(**merged)
await adapter.open()
adapters[name] = adapter
return adapters, sse_endpoint
@staticmethod
def _apply_env_overrides(channel_name: str, config: dict) -> dict:
overrides: dict = {}
if channel_name == "feishu":
app_id = os.getenv("FEISHU_APP_ID", "")
app_secret = os.getenv("FEISHU_APP_SECRET", "")
if app_id:
overrides["app_id"] = app_id
if app_secret:
overrides["app_secret"] = app_secret
elif channel_name == "dingtalk":
client_id = os.getenv("DINGTALK_CLIENT_ID", "")
client_secret = os.getenv("DINGTALK_CLIENT_SECRET", "")
if client_id:
overrides["client_id"] = client_id
if client_secret:
overrides["client_secret"] = client_secret
return overrides
@staticmethod
def _inject_bot_ids(pipeline: Pipeline, adapters: dict[str, ChannelAdapterPort]) -> None:
from yuxi.channel.application.pipeline.middlewares.mention_gate_middleware import MentionGateMiddleware
for middleware in pipeline.middlewares:
if isinstance(middleware, MentionGateMiddleware):
for name, adapter in adapters.items():
bot_id = getattr(adapter, "bot_id", None)
if bot_id:
middleware.set_bot_id(bot_id)
break
@staticmethod
async def _build_ws_connections(
adapters: dict[str, ChannelAdapterPort],
pipeline: Pipeline,
tracer: StartupTracer,
metrics: MetricsPort | None = None,
) -> WsConnectionManager:
ws_manager = WsConnectionManager(metrics=metrics)
ws_manager.set_adapters(adapters)
for name, adapter in adapters.items():
ws = adapter.ws_connection
if ws:
ws_manager.register(ws)
tracer.mark(f"ws_{name}", "registered")
inbound_service = InboundService(pipeline)
await ws_manager.start_all(asyncio.get_running_loop(), inbound_service)
return ws_manager
@staticmethod
def _build_workers(
infra: _InfraBundle,
adapters: dict[str, ChannelAdapterPort],
agent_port: AgentPort | None,
config: ChannelConfig,
pipeline: Pipeline,
redis: aioredis.Redis,
*,
default_agent_config_id: int = 1,
mq_workers: int = 4,
mq_max_concurrent: int = 20,
) -> _WorkerBundle:
session_factory = WorkerSessionFactory(pg_manager)
db_session_factory = pg_manager.get_async_session_context
binding_repo = PgBindingRepository(db_session_factory, infra.cache_port)
session_repo = PgSessionRepository(
db_session_factory,
infra.cache_port,
agent_id_resolver=lambda aid: _resolve_agent_id(db_session_factory, aid),
)
outbox_repo = PgOutboxRepository(db_session_factory, redis)
message_repo = PgMessageRepository(db_session_factory)
message_log_repo = PgMessageLogRepository(db_session_factory)
effective_agent_port = agent_port or AgentAdapter(db_session_factory)
session_resolver = SessionResolver(
session_repo=session_repo,
binding_repo=binding_repo,
default_agent_config_id=default_agent_config_id,
)
delivery_service = DeliveryService(
adapters=adapters,
message_repo=message_repo,
outbox_repo=outbox_repo,
event_publisher=infra.event_publisher,
agent_port=effective_agent_port,
message_log_repo=message_log_repo,
metrics=infra.metrics,
)
dispatch_service = DispatchService(
content_filter=infra.content_filter,
bot_loop_guard=infra.bot_loop_guard,
circuit_breaker=infra.circuit_breaker,
session_resolver=session_resolver,
delivery_service=delivery_service,
message_log_repo=message_log_repo,
metrics=infra.metrics,
)
worker_pool = WorkerPool(
queue_port=infra.queue_port,
dispatch_fn=dispatch_service.dispatch,
config=WorkerPoolConfig(num_workers=mq_workers, max_concurrent=mq_max_concurrent),
)
outbox_worker = OutboxRetryWorker(outbox_repo, adapters, redis, metrics=infra.metrics)
binding_service = BindingService(binding_repo)
config_service = ConfigService(
config_data=config.raw_data,
config_reload=infra.config_reload,
pipeline=pipeline,
channel_config=config,
)
return _WorkerBundle(
dispatch_service=dispatch_service,
binding_service=binding_service,
config_service=config_service,
worker_pool=worker_pool,
outbox_worker=outbox_worker,
session_factory=session_factory,
binding_repo=binding_repo,
session_repo=session_repo,
outbox_repo=outbox_repo,
message_repo=message_repo,
message_log_repo=message_log_repo,
)