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_infra.redis_bot_loop_guard import RedisBotLoopGuard from yuxi.channel.infrastructure.cache_infra.redis_cache import RedisCache from yuxi.channel.infrastructure.cache_infra.redis_circuit_breaker import RedisCircuitBreaker from yuxi.channel.infrastructure.cache_infra.redis_rate_limiter import RedisRateLimiter from yuxi.channel.infrastructure.configuration.channel_config import ChannelConfig from yuxi.channel.infrastructure.configuration.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, )