ForcePilot/backend/package/yuxi/channel/container.py
Kris 06e9ae397a refactor(channel/container): 修复缓存绑定仓储的导入路径
调整了CachingBindingRepository的导入目录,将其从cache仓库迁移到persistence仓库
2026-05-31 22:20:23 +08:00

609 lines
23 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 import (
AgentPort,
AuthenticationPort,
BotLoopGuardPort,
CachePort,
ChannelAdapterPort,
ChannelRequestVerifierPort,
CircuitBreakerPort,
ConfigReloadPort,
ContentFilterPort,
EventPublisherPort,
MetricsPort,
QueuePort,
RateLimitPort,
SignatureVerifyPort,
)
from yuxi.channel.domain.repository import (
BindingRepositoryPort,
MessageLogRepositoryPort,
MessageRepositoryPort,
OutboxRepositoryPort,
SessionRepositoryPort,
)
from yuxi.channel.domain.service.pipeline import Pipeline
from yuxi.channel.infrastructure.agent.agent_adapter import AgentAdapter
from yuxi.channel.infrastructure.persistence.repository.caching_binding_repository import CachingBindingRepository
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.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
verifiers: dict[str, ChannelRequestVerifierPort] = field(default_factory=dict)
@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
verifiers: dict[str, ChannelRequestVerifierPort] = field(default_factory=dict)
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
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,
auth_service,
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,
verifiers=workers.verifiers,
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
encrypt_key = os.getenv("FEISHU_ENCRYPT_KEY", "")
if encrypt_key:
overrides["encrypt_key"] = encrypt_key
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,
auth_service: AuthService,
*,
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
raw_binding_repo = PgBindingRepository(db_session_factory)
binding_repo: BindingRepositoryPort = (
CachingBindingRepository(raw_binding_repo, infra.cache_port) if infra.cache_port else raw_binding_repo
)
session_repo = PgSessionRepository(db_session_factory, infra.cache_port)
outbox_repo = PgOutboxRepository(db_session_factory, infra.cache_port)
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,
agent_port=effective_agent_port,
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, infra.cache_port, metrics=infra.metrics)
binding_service = BindingService(binding_repo)
verifiers = ChannelContainerFactory._build_verifiers(config, auth_service)
config_service = ConfigService(
config_data=config.raw_data,
config_reload=infra.config_reload,
pipeline=pipeline,
channel_config=config,
auth_service=auth_service,
verifiers=verifiers,
)
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,
verifiers=verifiers,
)
@staticmethod
def _build_verifiers(
config: ChannelConfig,
auth_service: AuthenticationPort,
) -> dict[str, ChannelRequestVerifierPort]:
from yuxi.channel.channels.feishu.verifier import FeishuRequestVerifier
from yuxi.channel.channels.hooks.verifier import HooksRequestVerifier
from yuxi.channel.channels.web.verifier import WebRequestVerifier
verifiers: dict[str, ChannelRequestVerifierPort] = {}
feishu_verifier = FeishuRequestVerifier(
verification_token=config.feishu_verification_token or "",
encrypt_key=config.feishu_encrypt_key or "",
)
verifiers["feishu"] = feishu_verifier
verifiers["hooks"] = HooksRequestVerifier()
web_verifier = WebRequestVerifier(auth_service=auth_service, allow_anonymous=config.allow_anonymous)
verifiers["web"] = web_verifier
return verifiers