新增了插件相关的完整领域模型、应用服务、基础设施实现,包括: 1. 插件状态、注册模式、来源等基础枚举和数据结构 2. 插件清单解析、发现、加载工具类 3. 插件注册表领域服务和内存存储实现 4. 插件相关的命令、查询、事件定义 5. 插件REST API接口和DTO映射 6. 集成了原有通道适配器到插件系统 7. 新增内置插件注册和自动发现能力
698 lines
28 KiB
Python
698 lines
28 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_channel_meta
|
|
from yuxi.channel.domain.port.agent_port import AgentPort
|
|
from yuxi.channel.domain.port.bot_loop_guard_port import BotLoopGuardPort
|
|
from yuxi.channel.domain.port.cache_port import CachePort
|
|
from yuxi.channel.application.service.plugin_registry_app_service import PluginRegistryAppService
|
|
from yuxi.channel.domain.model.plugin_registry.plugin_registry import PluginRegistry
|
|
from yuxi.channel.domain.port.channel_adapter_port import ChannelAdapterPort
|
|
from yuxi.channel.domain.port.channel_request_verifier_port import ChannelRequestVerifierPort
|
|
from yuxi.channel.domain.port.circuit_breaker_port import CircuitBreakerPort
|
|
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.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_agent_config_lookup import PgAgentConfigLookup
|
|
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)
|
|
verifiers: dict[str, ChannelRequestVerifierPort] = 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
|
|
_plugin_registry: PluginRegistry | None = field(default=None, repr=False, compare=False)
|
|
plugin_registry_impl: object | None = None
|
|
plugin_app_service: PluginRegistryAppService | None = None
|
|
|
|
def get_adapters_snapshot(self) -> dict[str, ChannelAdapterPort]:
|
|
if self._plugin_registry:
|
|
return self._plugin_registry.list_adapters()
|
|
return dict(self.adapters)
|
|
|
|
def get_adapter(self, channel_type: str) -> ChannelAdapterPort | None:
|
|
if self._plugin_registry:
|
|
return self._plugin_registry.get_adapter(channel_type)
|
|
return self.adapters.get(channel_type)
|
|
|
|
def require(self, attr_name: str) -> object:
|
|
value = getattr(self, attr_name, None)
|
|
if value is None:
|
|
raise HTTPException(status_code=503, detail=f"{attr_name} not initialized")
|
|
return value
|
|
|
|
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.get_adapters_snapshot().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,
|
|
plugin_registry_impl: object | None = None,
|
|
) -> 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,
|
|
plugin_registry_impl=plugin_registry_impl,
|
|
)
|
|
tracer.end()
|
|
|
|
ChannelContainerFactory._inject_bot_ids(pipeline, adapters)
|
|
|
|
for warning in ChannelContainerFactory._validate_channel_consistency(adapters, config):
|
|
logger.warning(warning)
|
|
|
|
tracer.begin("verifiers")
|
|
verifiers = ChannelContainerFactory._build_verifiers(config, adapters)
|
|
tracer.end()
|
|
|
|
tracer.begin("ws_connections")
|
|
ws_manager = await ChannelContainerFactory._build_ws_connections(
|
|
adapters,
|
|
pipeline,
|
|
tracer,
|
|
metrics=infra.metrics,
|
|
registry=plugin_registry_impl.registry if plugin_registry_impl else None,
|
|
)
|
|
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,
|
|
registry=plugin_registry_impl.registry if plugin_registry_impl else None,
|
|
auth_service=auth_service,
|
|
)
|
|
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,
|
|
verifiers=verifiers,
|
|
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 plugin_registry_impl is not None:
|
|
container._plugin_registry = plugin_registry_impl.registry
|
|
container.plugin_registry_impl = plugin_registry_impl
|
|
container.plugin_app_service = PluginRegistryAppService(
|
|
registry_impl=plugin_registry_impl,
|
|
event_publisher=infra.event_publisher,
|
|
)
|
|
|
|
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,
|
|
*,
|
|
plugin_registry_impl: object | None = None,
|
|
) -> tuple[dict[str, ChannelAdapterPort], SseEndpoint]:
|
|
import yuxi.channel.channels # noqa: F401
|
|
|
|
from yuxi.channel.channels._registry import get_registered_channels
|
|
|
|
sse_endpoint = SseEndpoint(metrics=infra.metrics)
|
|
await sse_endpoint.start()
|
|
|
|
infra_bundle: dict[str, object] = {
|
|
"sse_push": sse_endpoint,
|
|
"hook_mappings": config.get_channel_config("hooks").get("mappings", []),
|
|
}
|
|
|
|
adapters: dict[str, ChannelAdapterPort] = {}
|
|
|
|
if plugin_registry_impl is not None:
|
|
from yuxi.channel.domain.service.plugin_registry_domain_service import PluginRegistryDomainService
|
|
|
|
registry: PluginRegistry = plugin_registry_impl.registry
|
|
|
|
for record in registry.list_records():
|
|
if not record.source.enabled:
|
|
continue
|
|
|
|
default_config = record.get_default_config()
|
|
override = config.get_channel_config(record.channel_type)
|
|
merged = {**default_config, **override}
|
|
|
|
meta = get_channel_meta(record.channel_type)
|
|
for dep in meta.get("infra_dependencies", []):
|
|
if dep in infra_bundle and dep not in merged:
|
|
merged[dep] = infra_bundle[dep]
|
|
|
|
env_mapping = meta.get("env_mapping", {})
|
|
if env_mapping:
|
|
env_overrides = ChannelContainerFactory._apply_env_overrides_generic(env_mapping)
|
|
merged.update(env_overrides)
|
|
|
|
record.set_config(merged)
|
|
PluginRegistryDomainService.transition_to_configured(record)
|
|
|
|
adapter = await registry.activate(record.plugin_id, **merged)
|
|
assert isinstance(adapter, ChannelAdapterPort), (
|
|
f"{type(adapter).__name__} does not implement ChannelAdapterPort"
|
|
)
|
|
|
|
adapters = registry.list_adapters()
|
|
else:
|
|
for name, cls in get_registered_channels().items():
|
|
meta = get_channel_meta(name)
|
|
default_config = cls.get_default_config()
|
|
override = config.get_channel_config(name)
|
|
merged = {**default_config, **override}
|
|
|
|
for dep in meta.get("infra_dependencies", []):
|
|
if dep in infra_bundle and dep not in merged:
|
|
merged[dep] = infra_bundle[dep]
|
|
|
|
env_mapping = meta.get("env_mapping", {})
|
|
if env_mapping:
|
|
env_overrides = ChannelContainerFactory._apply_env_overrides_generic(env_mapping)
|
|
merged.update(env_overrides)
|
|
|
|
adapter = cls(**merged)
|
|
assert isinstance(adapter, ChannelAdapterPort), f"{cls.__name__} does not implement ChannelAdapterPort"
|
|
await adapter.open()
|
|
adapters[name] = adapter
|
|
|
|
return adapters, sse_endpoint
|
|
|
|
@staticmethod
|
|
def _apply_env_overrides_generic(env_mapping: dict[str, str]) -> dict:
|
|
overrides: dict = {}
|
|
for env_var, config_key in env_mapping.items():
|
|
value = os.getenv(env_var, "")
|
|
if value:
|
|
overrides[config_key] = value
|
|
return overrides
|
|
|
|
@staticmethod
|
|
def _validate_channel_consistency(adapters: dict, config: ChannelConfig) -> list[str]:
|
|
from yuxi.channel.domain.model.shared.channel_type import ChannelType
|
|
|
|
warnings_list: list[str] = []
|
|
valid_types = {t.value for t in ChannelType}
|
|
for name in adapters:
|
|
if name not in valid_types:
|
|
warnings_list.append(f"adapter '{name}' not in ChannelType enum")
|
|
if not config.get_channel_config(name):
|
|
warnings_list.append(f"no config section for channel '{name}' in YAML")
|
|
return warnings_list
|
|
|
|
@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
|
|
def _build_verifiers(
|
|
config: ChannelConfig,
|
|
adapters: dict[str, ChannelAdapterPort],
|
|
) -> dict[str, ChannelRequestVerifierPort]:
|
|
from yuxi.channel.channels._registry import get_channel_meta
|
|
|
|
verifiers: dict[str, ChannelRequestVerifierPort] = {}
|
|
for name in adapters:
|
|
meta = get_channel_meta(name)
|
|
verifier_factory = meta.get("verifier_factory")
|
|
if verifier_factory:
|
|
channel_config = config.get_channel_config(name)
|
|
verifiers[name] = verifier_factory(channel_config)
|
|
|
|
for name, verifier in verifiers.items():
|
|
if not verifier.enabled:
|
|
logger.warning("channel '%s' is running without request verification", name)
|
|
|
|
return verifiers
|
|
|
|
@staticmethod
|
|
async def _build_ws_connections(
|
|
adapters: dict[str, ChannelAdapterPort],
|
|
pipeline: Pipeline,
|
|
tracer: StartupTracer,
|
|
metrics: MetricsPort | None = None,
|
|
*,
|
|
registry: PluginRegistry | None = None,
|
|
) -> WsConnectionManager:
|
|
ws_manager = WsConnectionManager(metrics=metrics, registry=registry)
|
|
if not registry:
|
|
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,
|
|
registry: PluginRegistry | None = None,
|
|
auth_service: AuthService | None = None,
|
|
) -> _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, 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,
|
|
default_agent_config_id=default_agent_config_id,
|
|
)
|
|
|
|
delivery_service = DeliveryService(
|
|
adapters=adapters,
|
|
registry=registry,
|
|
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, registry=registry)
|
|
|
|
binding_service = BindingService(binding_repo, PgAgentConfigLookup(db_session_factory))
|
|
config_service = ConfigService(
|
|
config_data=config.raw_data,
|
|
config_reload=infra.config_reload,
|
|
pipeline=pipeline,
|
|
channel_config=config,
|
|
auth_service=auth_service,
|
|
)
|
|
|
|
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,
|
|
)
|