ForcePilot/backend/package/yuxi/channels/manager.py
Kris 6ca611fead refactor(channel): 重构并新增多项渠道管理功能
1. 简化message_actions.py中获取适配器的逻辑
2. 新增适配器合法性校验工具方法
3. 新增会话映射过期清理功能
4. 重构渠道状态机与基础适配器实现
5. 统一渠道操作异常处理逻辑
6. 新增凭证状态查询与刷新接口
7. 优化健康检查与自动重连逻辑
8. 新增统计数据缓存与批量查询优化
9. 修复部分数据库操作的异常处理逻辑
2026-05-14 09:24:50 +08:00

1379 lines
56 KiB
Python

from __future__ import annotations
import asyncio
import importlib
import threading
import time
from collections import defaultdict, deque
from dataclasses import dataclass
from typing import Any
from sqlalchemy import Integer, func, select
from yuxi.channels.auth.secret_manager import SecretManager
from yuxi.channels.base import BaseChannelAdapter
from yuxi.channels.exceptions import ChannelException, ChannelTimeoutError
from yuxi.channels.infra.broadcast import EventBroadcaster
from yuxi.channels.infra.circuit_breaker import CircuitBreaker, CircuitBreakerOpenError, CircuitState
from yuxi.channels.infra.config_watcher import ConfigWatcher
from yuxi.channels.models import ChannelStatus
from yuxi.channels.registry import _BUILTIN_ADAPTERS, ChannelRegistry, _load_builtin_adapters
from yuxi.channels.router import MessageRouter
from yuxi.channels.session_mapper import SessionMapper
from yuxi.channels.services.context import GatewayRequestContext
from yuxi.channels.services.doctor import ConfigDoctor, DiagnosisIssue
from yuxi.channels.services.maintenance import MaintenanceRunner
from yuxi.channels.services.plugin_state_store import PostgresPluginStateStore
from yuxi.channels.services.runtime_state import RuntimeState
from yuxi.channels.services.stats_collector import StatsCollector
from yuxi.channels.services.webhook_registry import WebhookRegistry
from yuxi.channels.services.ws_logger import WsLogger
from yuxi.storage.postgres.manager import pg_manager
from yuxi.storage.postgres.models_channels import ChannelConfig, ChannelMsgRecord
from yuxi.utils.datetime_utils import utc_now_naive as _utc_now
from yuxi.utils.logging_config import logger
HEALTH_CHECK_INTERVAL = 60
HEALTH_CHECK_MAX_AGE = 65
CONFIG_WATCH_INTERVAL = 30.0
STATS_CACHE_TTL = 60.0
_DEFAULT_RATE_LIMIT = {
"enabled": True,
"max_per_minute": 60,
"max_per_hour": 1000,
"max_per_day": 10000,
"burst_size": 10,
"cooldown": 0,
}
FRAMEWORK_CONFIG_KEYS = {"enabled", "connect_timeout", "rate_limit", "display_name"}
def _validate_channel_config(channel_id: str, config: dict[str, Any]) -> None:
adapter_cls = _BUILTIN_ADAPTERS.get(channel_id)
if not adapter_cls or not hasattr(adapter_cls, "default_config"):
return
default_config = adapter_cls.default_config()
valid_keys = set(default_config.keys()) | FRAMEWORK_CONFIG_KEYS
invalid_keys = [k for k in config if k not in valid_keys]
if invalid_keys:
raise ValueError(f"无效的配置项: {', '.join(invalid_keys)}。有效配置项: {', '.join(sorted(valid_keys))}")
for key, default_val in default_config.items():
is_credential = any(k in key for k in ("token", "secret", "key", "app_id"))
if is_credential and key not in config and not default_val:
raise ValueError(f"Missing required credential field: '{key}'")
for key in config:
if key in default_config:
expected_type = type(default_config[key])
if not isinstance(config[key], expected_type) and config[key] is not None:
raise ValueError(
f"Type mismatch for '{key}': expected {expected_type.__name__}, got {type(config[key]).__name__}"
)
@dataclass
class RateLimitResult:
allowed: bool
retry_after_seconds: int
remaining: int
limit: int
window: str
class ChannelManager:
SENSITIVE_KEY_PATTERNS = ("token", "secret", "key", "password", "api_key", "app_id", "app_secret")
@staticmethod
def _mask_sensitive_config(config: dict) -> dict:
if not config:
return config
return {
k: ("***" if any(pattern in k.lower() for pattern in ChannelManager.SENSITIVE_KEY_PATTERNS) else v)
for k, v in config.items()
}
def __init__(self, registry: ChannelRegistry | None = None, router: MessageRouter | None = None):
self._registry = registry or ChannelRegistry()
self._router = router or MessageRouter(channel_manager=self)
self._adapters: dict[str, BaseChannelAdapter] = {}
self._circuit_breakers: dict[str, CircuitBreaker] = defaultdict(lambda: CircuitBreaker())
self._health_tasks: dict[str, asyncio.Task] = {}
self._initialized = False
self._phase: str = "not_started"
self._rate_limiters: dict[str, deque[float]] = {}
self._rate_limit_locks: dict[str, asyncio.Lock] = {}
self._restart_locks: dict[str, asyncio.Lock] = {}
self._channels_config: dict[str, dict[str, Any]] = {}
self._dynamic_channel_ids: set[str] = set()
self._watcher: ConfigWatcher | None = None
self._doctor: ConfigDoctor | None = None
self._broadcaster: EventBroadcaster | None = None
self._ctx: GatewayRequestContext | None = None
self.runtime_state = RuntimeState()
self._ws_handlers: dict[str, Any] = {}
self._scheduled_tasks: list[asyncio.Task] = []
self._ws_logger: WsLogger | None = None
self._maintenance_runner: MaintenanceRunner | None = None
self._stats_collector: StatsCollector | None = None
self._webhook_registry: WebhookRegistry | None = None
self._state_store: PostgresPluginStateStore | None = None
self._ws_broadcast: Any = None
self._prev_statuses: dict[str, str] = {}
self._config_lock: asyncio.Lock = asyncio.Lock()
self._cached_health: dict[str, tuple[float, Any]] = {}
self._stats_cache: dict[str, tuple[float, dict]] = {}
self._batch_stats_cache: tuple[float, dict[str, dict]] = (0, {})
self._start_stop_locks: dict[str, asyncio.Lock] = {}
self._health_failure_counts: dict[str, int] = defaultdict(int)
@property
def phase(self) -> str:
return self._phase
@property
def context(self) -> GatewayRequestContext | None:
return self._ctx
@property
def broadcaster(self) -> EventBroadcaster | None:
return self._broadcaster
@property
def watcher(self) -> ConfigWatcher | None:
return self._watcher
@property
def doctor(self) -> ConfigDoctor | None:
return self._doctor
def _register_all_adapters(self) -> None:
from yuxi.channels.message_actions import ActionRegistry
_load_builtin_adapters()
for channel_id, adapter_cls in _BUILTIN_ADAPTERS.items():
self._registry.register(channel_id, adapter_cls)
ActionRegistry.register_adapter(adapter_cls)
_registerer_map: dict[str, str] = {
"wechat": "yuxi.channels.adapters.wechat.message_actions",
"telegram": "yuxi.channels.adapters.telegram.message_actions",
"whatsapp": "yuxi.channels.adapters.whatsapp.message_actions",
"feishu": "yuxi.channels.adapters.feishu.message_actions",
"discord": "yuxi.channels.adapters.discord.message_actions",
"slack": "yuxi.channels.adapters.slack.message_actions",
"signal": "yuxi.channels.adapters.signal.message_actions",
"zalo_oa": "yuxi.channels.adapters.zalo_oa.message_actions",
}
for channel_id, module_path in _registerer_map.items():
try:
mod = importlib.import_module(module_path)
register_fn = getattr(mod, f"register_{channel_id}_actions", None)
if register_fn:
register_fn()
except Exception:
logger.exception(f"Failed to register actions for {channel_id}")
async def load_config(self) -> None:
if self._phase not in ("not_started",):
return
self._register_all_adapters()
from yuxi import config as conf
self._channels_config = dict(getattr(conf, "channels", {}))
self._merge_dynamic_channels()
await self._recover_channels_from_db()
self._phase = "config_loaded"
logger.info("ChannelManager: config loaded")
async def prepare_bootstrap(self) -> None:
if self._phase not in ("config_loaded",):
return
async with pg_manager.get_async_session_context() as db:
await self._ensure_virtual_department(db)
await self._ensure_default_agent_config(db)
self._phase = "bootstrapped"
logger.info("ChannelManager: bootstrap prepared")
async def start_channels(self) -> None:
if self._phase not in ("config_loaded", "bootstrapped"):
return
for channel_id in self._registry.list_channels():
channel_conf = self._channels_config.get(channel_id, {})
if channel_conf.get("enabled", False):
try:
await self.start_channel(channel_id, channel_conf)
except Exception:
logger.exception(f"Failed to start channel {channel_id}")
self._phase = "channels_started"
logger.info(f"ChannelManager: channels started: {list(self._adapters.keys())}")
async def start_subscriptions(self) -> None:
if self._phase not in ("channels_started",):
return
self._phase = "fully_running"
logger.info("ChannelManager: subscriptions started, fully running")
async def initialize(self) -> None:
if self._initialized:
return
self._register_all_adapters()
from yuxi import config as conf
self._channels_config = dict(getattr(conf, "channels", {}))
self._merge_dynamic_channels()
for channel_id in self._registry.list_channels():
channel_conf = self._channels_config.get(channel_id, {})
if channel_conf.get("enabled", False):
try:
await self.start_channel(channel_id, channel_conf)
except Exception:
logger.exception(f"Failed to start channel {channel_id}")
async with pg_manager.get_async_session_context() as db:
await self._ensure_virtual_department(db)
await self._ensure_default_agent_config(db)
self._initialized = True
self._phase = "fully_running"
logger.info(f"ChannelManager initialized with channels: {list(self._adapters.keys())}")
async def shutdown(self) -> None:
for channel_id in list(self._adapters.keys()):
try:
await self.stop_channel(channel_id)
except Exception:
logger.exception(f"Failed to stop channel {channel_id} during shutdown")
for task in self._health_tasks.values():
task.cancel()
for task in self._scheduled_tasks:
task.cancel()
if self._watcher:
await self._watcher.stop()
self._initialized = False
logger.info("ChannelManager shutdown complete")
async def startup(self) -> None:
if self._initialized:
return
try:
await self._stage_load_config()
await self._stage_prepare_bootstrap()
await self._stage_start_early_runtime()
await self._stage_init_channels()
await self._stage_create_runtime_state()
await self._stage_start_runtime_services()
await self._stage_activate_scheduled_services()
await self._stage_attach_ws_handlers()
await self._stage_start_event_subscriptions()
self._initialized = True
self._phase = "fully_running"
logger.info("ChannelManager: 9-stage startup complete")
except Exception:
logger.exception(f"ChannelManager startup failed at phase: {self._phase}")
raise
async def _stage_load_config(self) -> None:
if self._phase not in ("not_started",):
return
self._register_all_adapters()
from yuxi import config as conf
self._channels_config = dict(getattr(conf, "channels", {}))
self._merge_dynamic_channels()
await self._recover_channels_from_db()
self._phase = "config_loaded"
logger.info("ChannelManager: [1/10] config loaded")
async def _stage_prepare_bootstrap(self) -> None:
if self._phase not in ("config_loaded",):
return
async with pg_manager.get_async_session_context() as db:
await self._ensure_virtual_department(db)
await self._ensure_default_agent_config(db)
self._phase = "bootstrapped"
logger.info("ChannelManager: [2/10] bootstrap prepared")
async def _stage_start_early_runtime(self) -> None:
if self._phase not in ("bootstrapped",):
return
self._broadcaster = EventBroadcaster()
self._ctx = GatewayRequestContext(
runtime_config=self._channels_config,
start_channel=self._ctx_start_channel,
stop_channel=self._ctx_stop_channel,
mark_channel_logged_out=self._ctx_mark_channel_logged_out,
broadcast_fn=self._broadcaster.broadcast,
node_send_to_session_fn=self._broadcaster.node_send_to_session,
)
self._phase = "early_runtime"
logger.info("ChannelManager: [3/9] early runtime started")
async def _stage_init_channels(self) -> None:
if self._phase not in ("early_runtime",):
return
for channel_id in self._registry.list_channels():
channel_conf = self._channels_config.get(channel_id, {})
if channel_conf.get("enabled", False):
try:
await self.start_channel(channel_id, channel_conf)
except Exception:
logger.exception(f"Failed to start channel {channel_id}")
self._phase = "channels_started"
logger.info(f"ChannelManager: [4/9] channels started: {list(self._adapters.keys())}")
async def _stage_create_runtime_state(self) -> None:
if self._phase not in ("channels_started",):
return
self.runtime_state = RuntimeState(
started_at=time.monotonic(),
active_channels=len(self._adapters),
phase=self._phase,
node_id="forcepilot-gateway",
main_node=True,
services=["doctor", "maintenance", "stats", "webhooks"],
)
self._state_store = PostgresPluginStateStore()
self._ws_logger = WsLogger(max_entries=1000)
self._phase = "runtime_state_created"
logger.info("ChannelManager: [5/9] runtime state created")
async def _stage_start_runtime_services(self) -> None:
if self._phase not in ("runtime_state_created",):
return
self._doctor = ConfigDoctor(self)
self._maintenance_runner = MaintenanceRunner(self, self.runtime_state)
self._stats_collector = StatsCollector(self.runtime_state)
self._webhook_registry = WebhookRegistry(self)
self._phase = "runtime_services_started"
logger.info("ChannelManager: [6/9] runtime services started (doctor+maintenance+stats+webhooks)")
async def _stage_activate_scheduled_services(self) -> None:
if self._phase not in ("runtime_services_started",):
return
self._watcher = ConfigWatcher(self)
await self._watcher.watch(interval=CONFIG_WATCH_INTERVAL)
self._scheduled_tasks.append(asyncio.create_task(self._maintenance_runner.run()))
self._scheduled_tasks.append(asyncio.create_task(self._stats_collector.run()))
self._scheduled_tasks.append(asyncio.create_task(self._webhook_registry.run()))
self._scheduled_tasks.append(asyncio.create_task(self._cleanup_expired_states_loop()))
self._phase = "scheduled_services_active"
logger.info("ChannelManager: [7/9] scheduled services activated (watcher+maintenance+stats+webhooks)")
async def _stage_attach_ws_handlers(self) -> None:
if self._phase not in ("scheduled_services_active",):
return
if self._broadcaster:
self._broadcaster.subscribe_callback("channel.status_change", self._on_channel_status_change)
self._broadcaster.subscribe_callback("tick", self._on_tick)
self._broadcaster.subscribe_callback("chat", self._on_chat_event)
self._broadcaster.subscribe_callback("channel.logout", self._on_channel_logout)
self._broadcaster.subscribe_callback("config.reload", self._on_config_reload)
self._phase = "ws_handlers_attached"
logger.info("ChannelManager: [8/9] websocket handlers attached")
async def _stage_start_event_subscriptions(self) -> None:
if self._phase not in ("ws_handlers_attached",):
return
self.runtime_state.services = ["doctor", "maintenance", "stats", "webhooks", "watcher", "ws"]
self._phase = "subscriptions_started"
logger.info("ChannelManager: [9/9] event subscriptions started")
def set_ws_broadcast(self, cb) -> None:
self._ws_broadcast = cb
async def _push_channel_status_to_ws(self, channel_id: str, status: str, health: dict | None = None) -> None:
if not self._ws_broadcast:
return
payload: dict[str, Any] = {"channel_id": channel_id, "status": status}
if health:
payload["health"] = health
try:
await self._ws_broadcast(
{
"type": "channel_status",
"payload": payload,
"timestamp": _now_iso(),
}
)
except Exception:
pass
async def _on_channel_status_change(self, event: str, payload: Any) -> None:
logger.debug(f"channel.status_change: {event}")
if isinstance(payload, dict):
channel_id = payload.get("channel_id")
status = payload.get("status")
health = payload.get("health")
if channel_id and status:
await self._push_channel_status_to_ws(channel_id, status, health)
async def _on_tick(self, event: str, payload: Any) -> None:
pass # TODO: implement periodic tick logic (health checks, stats flush, etc.)
async def _on_chat_event(self, event: str, payload: Any) -> None:
logger.debug(f"WS chat event: {event}")
async def _on_channel_logout(self, event: str, payload: Any) -> None:
channel_id = payload.get("channel_id") if isinstance(payload, dict) else None
account_id = payload.get("account_id") if isinstance(payload, dict) else None
logger.info(f"Channel logout: {channel_id}/{account_id}")
async def _on_config_reload(self, event: str, payload: Any) -> None:
await self.reload_config_now()
async def _ctx_start_channel(self, channel_id: str, config: dict) -> None:
await self.start_channel(channel_id, config)
async def _ctx_stop_channel(self, channel_id: str) -> None:
await self.stop_channel(channel_id)
async def _ctx_mark_channel_logged_out(self, channel_id: str, account_id: str) -> None:
logger.info(f"Channel {channel_id} account {account_id} marked as logged out")
async def diagnose_channels(self) -> list[DiagnosisIssue]:
if self._doctor is None:
return []
return await self._doctor.diagnose()
async def auto_fix_channel(self, issue: DiagnosisIssue) -> bool:
if self._doctor is None:
return False
return await self._doctor.auto_fix(issue)
async def reload_config_now(self, changed_keys: list[str] | None = None) -> dict[str, Any]:
if self._watcher is None:
return {}
return await self._watcher.reload_now(changed_keys)
async def start_channel(self, channel_id: str, config: dict[str, Any] | None = None) -> None:
lock = self._start_stop_locks.setdefault(channel_id, asyncio.Lock())
async with lock:
return await self._start_channel_impl(channel_id, config)
async def _start_channel_impl(self, channel_id: str, config: dict[str, Any] | None = None) -> None:
if channel_id in self._adapters:
logger.warning(f"Channel {channel_id} already running")
return
adapter_cls = self._registry.get(channel_id)
if not adapter_cls:
raise ValueError(f"No adapter registered for channel {channel_id}")
config = config or self._channels_config.get(channel_id, {})
_validate_channel_config(channel_id, config)
adapter = adapter_cls(config=config)
adapter._state_store = self._state_store
adapter.on_message(self._handle_inbound_message)
pre_connect_result = await adapter.pre_connect()
if pre_connect_result:
qr_url = pre_connect_result.get("qr_url")
if qr_url:
logger.info(f"Channel {channel_id} requires QR scan: {qr_url}")
connect_timeout = config.get("connect_timeout", 30)
try:
await asyncio.wait_for(adapter.connect(), timeout=connect_timeout)
except TimeoutError:
raise ChannelTimeoutError(f"Channel {channel_id} connection timed out after {connect_timeout}s")
try:
if channel_id not in self._health_tasks:
self._health_tasks[channel_id] = asyncio.create_task(self._health_check_loop(channel_id))
except Exception:
logger.exception(f"Failed to create health check task for {channel_id}")
self._adapters[channel_id] = adapter
current_status = self._adapter_status(adapter)
self._prev_statuses[channel_id] = current_status
if self._broadcaster:
await self._broadcaster.broadcast(
"channel.status_change",
{
"channel_id": channel_id,
"status": current_status,
"health": None,
},
)
logger.info(f"Channel {channel_id} started")
async def stop_channel(self, channel_id: str) -> None:
lock = self._start_stop_locks.setdefault(channel_id, asyncio.Lock())
async with lock:
return await self._stop_channel_impl(channel_id)
async def _stop_channel_impl(self, channel_id: str) -> None:
adapter = self._adapters.get(channel_id)
if not adapter:
return
task = self._health_tasks.pop(channel_id, None)
if task:
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
try:
await adapter.disconnect()
except Exception:
logger.exception(f"Error disconnecting channel {channel_id}")
finally:
self._adapters.pop(channel_id, None)
self._prev_statuses.pop(channel_id, None)
self._health_failure_counts.pop(channel_id, None)
if channel_id in self._circuit_breakers:
self._circuit_breakers[channel_id] = CircuitBreaker(channel_id=channel_id)
if self._broadcaster:
try:
await self._broadcaster.broadcast(
"channel.status_change",
{
"channel_id": channel_id,
"status": ChannelStatus.DISCONNECTED.value,
"health": None,
},
)
except Exception:
logger.exception(f"Error broadcasting status change for channel {channel_id}")
logger.info(f"Channel {channel_id} stopped")
async def restart_channel(self, channel_id: str, timeout: float = 30.0) -> None:
logger.info(f"Restarting channel {channel_id}...")
lock = self._restart_locks.setdefault(channel_id, asyncio.Lock())
async with lock:
config = self._channels_config.get(channel_id, {})
try:
await asyncio.wait_for(self.stop_channel(channel_id), timeout=timeout)
except TimeoutError:
raise ChannelException(
f"Stop channel {channel_id} timed out after {timeout}s",
retryable=True,
retry_after_ms=int(timeout * 1000),
)
await self._start_channel_with_retry(channel_id, config)
logger.info(f"Channel {channel_id} restarted")
async def _start_channel_with_retry(
self, channel_id: str, config: dict[str, Any] | None = None, max_retries: int = 3
) -> None:
last_error: Exception | None = None
for attempt in range(max_retries):
try:
await self.start_channel(channel_id, config)
return
except Exception as e:
last_error = e
if attempt < max_retries - 1:
delay = 2**attempt
logger.warning(f"Restart attempt {attempt + 1}/{max_retries} failed for channel {channel_id}: {e}")
await asyncio.sleep(delay)
raise RuntimeError(f"Failed to restart channel {channel_id} after {max_retries} attempts") from last_error
async def register_channel(
self, channel_id: str, config: dict[str, Any] | None = None, registered_by: str | None = None
) -> None:
adapter_cls = self._registry.get(channel_id)
if not adapter_cls:
raise ValueError(f"Unknown channel type: '{channel_id}'")
_validate_channel_config(channel_id, config or {})
async with self._config_lock:
if channel_id in self._channels_config:
raise ValueError(f"Channel '{channel_id}' is already registered")
from yuxi.channels.message_actions import ActionRegistry
ActionRegistry.register_adapter(adapter_cls)
merged_config = {**(config or {}), "enabled": True}
self._channels_config[channel_id] = merged_config
self._dynamic_channel_ids.add(channel_id)
async with pg_manager.get_async_session_context() as db:
existing = await db.execute(select(ChannelConfig).where(ChannelConfig.channel_id == channel_id))
if existing.scalar_one_or_none():
raise ValueError(f"Channel '{channel_id}' is already registered")
db_config = ChannelConfig(
channel_id=channel_id,
config_json=merged_config,
enabled=True,
registered_by=registered_by,
)
db.add(db_config)
await db.commit()
logger.info(
"AUDIT: channel=%s action=register user=%s at=%s config_keys=%s",
channel_id,
registered_by,
_utc_now().isoformat(),
list((config or {}).keys()),
)
logger.info(f"Channel '{channel_id}' registered by {registered_by}")
async def unregister_channel(self, channel_id: str) -> None:
from yuxi.channels.message_actions import ActionRegistry
async with self._config_lock:
await self.stop_channel(channel_id)
if channel_id in self._channels_config:
self._channels_config[channel_id]["enabled"] = False
self._dynamic_channel_ids.discard(channel_id)
self._circuit_breakers.pop(channel_id, None)
ActionRegistry.deregister(channel_id)
if self._state_store:
try:
await self._state_store.delete_by_channel(channel_id)
except Exception:
logger.exception(f"Failed to cleanup plugin state for channel {channel_id}")
async with pg_manager.get_async_session_context() as db:
result = await db.execute(select(ChannelConfig).where(ChannelConfig.channel_id == channel_id))
db_config = result.scalar_one_or_none()
if db_config:
db_config.enabled = False
db_config.updated_at = _utc_now()
else:
db_config = ChannelConfig(
channel_id=channel_id,
config_json=self._channels_config.get(channel_id, {}),
enabled=False,
)
db.add(db_config)
await db.commit()
logger.info(
"AUDIT: channel=%s action=unregister at=%s",
channel_id,
_utc_now().isoformat(),
)
logger.info(f"Channel {channel_id} unregistered")
async def send_outbound(self, channel_id: str, response) -> None:
adapter = self._adapters.get(channel_id)
if not adapter:
raise ChannelException(f"Channel {channel_id} not found", retryable=False)
cb = self._circuit_breakers[channel_id]
try:
await cb.call(lambda: adapter.send(response))
except CircuitBreakerOpenError:
raise ChannelException(
f"Channel {channel_id} temporarily unavailable",
retryable=True,
retry_after_ms=int(cb.recovery_timeout * 1000),
)
async def get_channel_status(self, channel_id: str | None = None, include_stats: bool = False) -> dict:
if channel_id:
return await self._get_single_channel_status(channel_id, stats=None if include_stats else {})
channel_ids = self._registry.list_channels()
if not channel_ids:
return {"channels": {}}
batch_stats = await self._get_batch_channel_stats(channel_ids)
tasks = [self._get_single_channel_status(cid, stats=batch_stats.get(cid)) for cid in channel_ids]
results = await asyncio.gather(*tasks, return_exceptions=True)
all_channels = {}
success_count = 0
failed_count = 0
errors_dict: dict[str, str] = {}
for cid, info in zip(channel_ids, results):
if isinstance(info, Exception):
failed_count += 1
errors_dict[cid] = str(info)
logger.warning(f"获取渠道 {cid} 状态异常: {info}")
all_channels[cid] = {
"channel_id": cid,
"channel_type": None,
"display_name": None,
"enabled": False,
"status": "error",
"total_messages": 0,
"today_messages": 0,
"active_connections": 0,
"health": None,
}
continue
success_count += 1
stats = info.get("stats") or {}
all_channels[cid] = {
"channel_id": cid,
"channel_type": info.get("channel_type"),
"display_name": info.get("display_name"),
"enabled": info.get("enabled", False),
"status": info.get("status", "not_found"),
"total_messages": stats.get("total_messages", 0),
"today_messages": stats.get("today_messages", 0),
"active_connections": 1 if info.get("status") == ChannelStatus.CONNECTED.value else 0,
"health": info.get("health"),
}
return {
"channels": all_channels,
"_meta": {
"total_channels": len(channel_ids),
"successful": success_count,
"failed": failed_count,
"errors": errors_dict,
},
}
async def update_channel_config(
self, channel_id: str, config_updates: dict[str, Any], user_id: str | None = None
) -> dict:
_validate_channel_config(channel_id, config_updates)
adapter = self._adapters.get(channel_id)
async with pg_manager.get_async_session_context() as db:
result = await db.execute(select(ChannelConfig).where(ChannelConfig.channel_id == channel_id))
db_config = result.scalar_one_or_none()
if db_config:
existing = db_config.config_json or {}
db_config.config_json = {**existing, **config_updates}
db_config.updated_at = _utc_now()
else:
db_config = ChannelConfig(
channel_id=channel_id,
config_json=config_updates,
enabled=True,
)
db.add(db_config)
await db.commit()
async with self._config_lock:
if channel_id not in self._channels_config:
self._channels_config[channel_id] = {}
self._channels_config[channel_id].update(config_updates)
if adapter:
try:
await adapter.reload_config(config_updates)
except Exception:
logger.exception(f"Failed to reload config for running adapter {channel_id}")
needs_restart = adapter is None
logger.info(
"AUDIT: channel=%s action=config_update user=%s at=%s keys=%s",
channel_id,
user_id or "unknown",
_utc_now().isoformat(),
list(config_updates.keys()),
)
return {
"channel_id": channel_id,
"config_updated": True,
"changed_keys": list(config_updates.keys()),
"needs_restart": needs_restart,
"adapter_status": "running" if adapter else "stopped",
}
def _merge_dynamic_channels(self) -> None:
for channel_id in list(self._dynamic_channel_ids):
if channel_id in self._adapters:
if channel_id not in self._channels_config:
adapter = self._adapters[channel_id]
self._channels_config[channel_id] = getattr(adapter, "config", {"enabled": True})
else:
self._dynamic_channel_ids.discard(channel_id)
async def _recover_channels_from_db(self) -> None:
try:
async with pg_manager.get_async_session_context() as db:
stmt = select(ChannelConfig).where(
ChannelConfig.enabled.is_(True),
)
result = await db.execute(stmt)
enabled_rows = result.scalars().all()
for row in enabled_rows:
channel_id = row.channel_id
db_config = row.config_json or {}
if channel_id in self._channels_config:
self._channels_config[channel_id].update(db_config)
else:
self._channels_config[channel_id] = {"enabled": True, **db_config}
self._dynamic_channel_ids.add(channel_id)
stmt_disabled = select(ChannelConfig).where(
ChannelConfig.enabled.is_(False),
)
result_disabled = await db.execute(stmt_disabled)
disabled_rows = result_disabled.scalars().all()
for row in disabled_rows:
channel_id = row.channel_id
if channel_id in self._channels_config:
self._channels_config[channel_id]["enabled"] = False
if enabled_rows:
logger.info(
"ChannelManager: recovered %d enabled + %d disabled channels from DB",
len(enabled_rows),
len(disabled_rows),
)
except Exception:
logger.exception("Failed to recover channels from DB")
async def test_channel(self, channel_id: str) -> dict:
adapter = self._adapters.get(channel_id)
if not adapter:
return {"channel_id": channel_id, "test_result": "failure", "error": "Channel not running"}
cb = self._circuit_breakers.get(channel_id)
if cb and cb.state == CircuitState.OPEN:
return {
"channel_id": channel_id,
"test_result": "failure",
"error": f"Circuit breaker OPEN, retry after {cb.recovery_timeout}s",
}
try:
start = time.monotonic()
health = await asyncio.wait_for(adapter.health_check(), timeout=15.0)
latency_ms = (time.monotonic() - start) * 1000
status_map = {"healthy": "success", "degraded": "degraded", "unhealthy": "failure"}
test_result = status_map.get(health.status, "failure")
if cb:
if health.status == "healthy":
await cb.record_success()
else:
await cb.record_failure()
result = {
"channel_id": channel_id,
"test_result": test_result,
"latency_ms": round(latency_ms, 1),
"health": health.model_dump(),
"details": health.metadata or {},
}
if test_result != "success":
result["error"] = health.last_error or f"Health status: {health.status}"
return result
except TimeoutError:
logger.warning(f"Test channel {channel_id}: health_check timeout")
if cb:
await cb.record_failure()
return {"channel_id": channel_id, "test_result": "failure", "error": "Health check timeout after 15s"}
except Exception as e:
logger.error(f"Test channel {channel_id} failed: {e}")
if cb:
await cb.record_failure()
return {"channel_id": channel_id, "test_result": "failure", "error": str(e)}
def is_registered(self, channel_id: str) -> bool:
return channel_id in self._channels_config
def is_enabled(self, channel_id: str) -> bool:
if channel_id not in self._channels_config:
return False
return self._channels_config[channel_id].get("enabled", True)
def is_available(self, channel_id: str) -> bool:
return self._registry.is_available(channel_id)
def is_running(self, channel_id: str) -> bool:
return channel_id in self._adapters
async def _cleanup_expired_states_loop(self) -> None:
CLEANUP_INTERVAL = 300
while True:
await asyncio.sleep(CLEANUP_INTERVAL)
try:
count = await self._state_store.cleanup_expired()
if count > 0:
logger.info(f"[Cleanup] Removed {count} expired state entries")
except asyncio.CancelledError:
break
except Exception:
logger.warning("Failed to cleanup expired state entries", exc_info=True)
try:
async with pg_manager.get_async_session_context() as db:
removed = await SessionMapper._cleanup_expired_mappings(db)
if removed > 0:
logger.info(f"[Cleanup] Removed {removed} expired thread mappings")
except asyncio.CancelledError:
break
except Exception:
logger.warning("Failed to cleanup expired thread mappings", exc_info=True)
async def check_rate_limit(self, key: str, max_req: int, window_seconds: int) -> bool:
lock = self._rate_limit_locks.setdefault(key, asyncio.Lock())
async with lock:
now = time.monotonic()
history = self._rate_limiters.setdefault(key, deque(maxlen=max_req))
while history and now - history[0] > window_seconds:
history.popleft()
if len(history) >= max_req:
return False
history.append(now)
return True
async def check_per_channel_rate_limit(
self, channel_id: str, user_id: str | None = None, action: str = "message"
) -> RateLimitResult:
config = self._channels_config.get(channel_id, {})
rate_config = config.get("rate_limit", _DEFAULT_RATE_LIMIT)
if not rate_config.get("enabled", True):
return RateLimitResult(allowed=True, retry_after_seconds=0, remaining=-1, limit=-1, window="unlimited")
burst_size = rate_config.get("burst_size", 10)
max_per_minute = rate_config.get("max_per_minute", 60) + burst_size
uid = user_id or "anonymous"
minute_key = f"rl:{channel_id}:minute:{uid}"
allowed = await self.check_rate_limit(minute_key, max_per_minute, 60)
if not allowed:
remaining = 0
return RateLimitResult(
allowed=False, retry_after_seconds=60, remaining=remaining, limit=max_per_minute, window="minute"
)
history = self._rate_limiters.get(minute_key, deque())
remaining = max_per_minute - len(history)
return RateLimitResult(
allowed=True, retry_after_seconds=0, remaining=remaining, limit=max_per_minute, window="minute"
)
async def _handle_inbound_message(self, message) -> None:
try:
await self._router.route_inbound(message)
except Exception:
logger.exception("Error handling inbound message")
async def _health_check_loop(self, channel_id: str) -> None:
HEALTH_FAILURE_THRESHOLD = 3
while channel_id in self._adapters:
await asyncio.sleep(HEALTH_CHECK_INTERVAL)
adapter = self._adapters.get(channel_id)
if not adapter:
break
prev_status = self._prev_statuses.get(channel_id)
current_status = self._adapter_status(adapter)
cb = self._circuit_breakers.get(channel_id)
health = None
try:
health = await adapter.health_check()
self._cached_health[channel_id] = (time.monotonic(), health)
self._health_failure_counts[channel_id] = 0
if cb and health.status == "healthy":
try:
await cb.record_success()
except Exception:
pass
elif cb:
try:
await cb.record_failure()
except Exception:
pass
except asyncio.CancelledError:
break
except Exception as e:
logger.warning(f"Health check failed for {channel_id}: {e}")
if cb:
try:
await cb.record_failure()
except Exception:
pass
self._health_failure_counts[channel_id] += 1
if self._health_failure_counts[channel_id] >= HEALTH_FAILURE_THRESHOLD:
logger.warning(
f"Health check failed {self._health_failure_counts[channel_id]} consecutive times "
f"for {channel_id}, triggering auto-reconnect"
)
self._health_failure_counts[channel_id] = 0
try:
await self.restart_channel(channel_id)
except Exception as reconnect_error:
logger.error(f"Auto-reconnect failed for {channel_id}: {reconnect_error}")
continue
if current_status != prev_status and self._broadcaster:
self._prev_statuses[channel_id] = current_status
health_dict = health.model_dump() if health else None
await self._broadcaster.broadcast(
"channel.status_change",
{
"channel_id": channel_id,
"status": current_status,
"health": health_dict,
},
)
async def _get_batch_channel_stats(self, channel_ids: list[str]) -> dict[str, dict]:
if not channel_ids:
return {}
now = time.monotonic()
cache_ts, cached = self._batch_stats_cache
if now - cache_ts < STATS_CACHE_TTL and cached:
return {cid: cached.get(cid, self._empty_stats()) for cid in channel_ids}
try:
from yuxi.utils.datetime_utils import utc_now_naive
utc_now = utc_now_naive()
today_start = utc_now.replace(hour=0, minute=0, second=0, microsecond=0)
async with pg_manager.get_async_session_context() as session:
totals_result = await session.execute(
select(
ChannelMsgRecord.channel_id,
func.count().label("total"),
func.sum(
(ChannelMsgRecord.status == "success").cast(Integer),
).label("success_count"),
func.sum(
(ChannelMsgRecord.status == "error").cast(Integer),
).label("error_count"),
)
.where(ChannelMsgRecord.channel_id.in_(channel_ids))
.group_by(ChannelMsgRecord.channel_id)
)
totals = {
r.channel_id: (r.total or 0, r.success_count or 0, r.error_count or 0) for r in totals_result.all()
}
today_result = await session.execute(
select(
ChannelMsgRecord.channel_id,
func.count().label("today_count"),
)
.where(
ChannelMsgRecord.channel_id.in_(channel_ids),
ChannelMsgRecord.created_at >= today_start,
)
.group_by(ChannelMsgRecord.channel_id)
)
today_counts = {r.channel_id: r.today_count for r in today_result.all()}
stats_map = {}
for cid in channel_ids:
if cid in totals:
total, success, error = totals[cid]
stats_map[cid] = {
"total_messages": total,
"today_messages": today_counts.get(cid, 0),
"success_count": int(success),
"error_count": int(error),
"success_rate": round(success / total, 3) if total > 0 else 0,
}
else:
stats_map[cid] = self._empty_stats()
self._batch_stats_cache = (now, stats_map)
return {cid: stats_map.get(cid, self._empty_stats()) for cid in channel_ids}
except Exception:
logger.warning(f"批量查询渠道统计失败 (channel_ids 数量: {len(channel_ids)}):", exc_info=True)
return {cid: self._empty_stats() for cid in channel_ids}
async def _get_single_channel_status(self, channel_id: str, stats: dict | None = None) -> dict:
adapter = self._adapters.get(channel_id)
if not adapter:
adapter_cls = self._registry.get(channel_id)
if adapter_cls:
caps = (
adapter_cls.capabilities.model_dump()
if hasattr(adapter_cls, "capabilities")
else {
"text_chunk_limit": adapter_cls.text_chunk_limit,
"supports_markdown": adapter_cls.supports_markdown,
"supports_streaming": adapter_cls.supports_streaming,
"max_media_size_mb": adapter_cls.max_media_size_mb,
}
)
channel_type = adapter_cls.channel_type.value
saved_config = self._channels_config.get(channel_id, {})
return {
"channel_id": channel_id,
"channel_type": channel_type,
"display_name": saved_config.get("display_name"),
"enabled": saved_config.get("enabled", False),
"status": ChannelStatus.DISABLED.value,
"config": SecretManager.redact_config(saved_config) if saved_config else {"enabled": False},
"capabilities": caps,
"health": None,
"health_error": None,
"circuit_state": "unknown",
"credential": {"has_credential": False, "is_expired": False, "source": "adapter_not_running"},
"stats": stats if stats is not None else None,
}
return {"channel_id": channel_id, "status": "not_found"}
health = None
health_error = None
cached_entry = self._cached_health.get(channel_id)
if cached_entry and (time.monotonic() - cached_entry[0]) < HEALTH_CHECK_MAX_AGE:
health = cached_entry[1].model_dump()
else:
try:
health_result = await asyncio.wait_for(
adapter.health_check(),
timeout=5.0,
)
health = health_result.model_dump()
self._cached_health[channel_id] = (time.monotonic(), health_result)
except TimeoutError:
health_error = "health_check_timeout"
logger.warning(f"渠道 {channel_id} 健康检查超时 (5s)")
except Exception:
health_error = "health_check_failed"
logger.warning(f"渠道 {channel_id} 健康检查异常", exc_info=True)
caps = (
type(adapter).capabilities.model_dump()
if hasattr(type(adapter), "capabilities")
else {
"text_chunk_limit": type(adapter).text_chunk_limit,
"supports_markdown": type(adapter).supports_markdown,
"supports_streaming": type(adapter).supports_streaming,
"max_media_size_mb": type(adapter).max_media_size_mb,
}
)
circuit_state = self._circuit_breakers.get(channel_id)
cb_state = circuit_state.state.value if circuit_state else "unknown"
channel_type = adapter.channel_type.value
adapter_config = getattr(adapter, "config", {})
display_name = adapter_config.get("display_name")
enabled = adapter_config.get("enabled", False)
return {
"channel_id": channel_id,
"channel_type": channel_type,
"display_name": display_name,
"enabled": enabled,
"status": self._adapter_status(adapter),
"config": SecretManager.redact_config(adapter_config),
"capabilities": caps,
"health": health,
"health_error": health_error,
"circuit_state": cb_state,
"credential": await self._get_credential_summary(channel_id, adapter),
"stats": stats if stats is not None else await self._get_channel_stats(channel_id),
}
async def _get_channel_stats(self, channel_id: str) -> dict:
now = time.monotonic()
cached = self._stats_cache.get(channel_id)
if cached and now - cached[0] < STATS_CACHE_TTL:
return cached[1]
try:
from yuxi.utils.datetime_utils import utc_now_naive
utc_now = utc_now_naive()
today_start = utc_now.replace(hour=0, minute=0, second=0, microsecond=0)
async with pg_manager.get_async_session_context() as session:
result = await session.execute(
select(
func.count().label("total"),
func.sum(
(ChannelMsgRecord.status == "success").cast(Integer),
).label("success_count"),
func.sum(
(ChannelMsgRecord.status == "error").cast(Integer),
).label("error_count"),
).where(ChannelMsgRecord.channel_id == channel_id)
)
row = result.one()
total = row.total or 0
success = row.success_count or 0
error = row.error_count or 0
today_result = await session.execute(
select(func.count()).where(
ChannelMsgRecord.channel_id == channel_id,
ChannelMsgRecord.created_at >= today_start,
)
)
today = today_result.scalar() or 0
stats = {
"total_messages": total,
"today_messages": today,
"success_count": int(success),
"error_count": int(error),
"success_rate": round(success / total, 3) if total > 0 else 0,
}
self._stats_cache[channel_id] = (now, stats)
return stats
except Exception:
return self._empty_stats()
@staticmethod
def _empty_stats() -> dict:
return {
"total_messages": 0,
"today_messages": 0,
"success_count": 0,
"error_count": 0,
"success_rate": 0,
}
def _adapter_status(self, adapter: BaseChannelAdapter) -> str:
_status = getattr(adapter, "_status", None)
if _status is None:
return "unknown"
return _status.value if hasattr(_status, "value") else str(_status)
async def _get_credential_summary(self, channel_id: str, adapter: BaseChannelAdapter) -> dict:
try:
secret_manager = getattr(adapter, "_secret_manager", None)
if secret_manager and hasattr(secret_manager, "get_credential_status"):
return await secret_manager.get_credential_status()
except Exception:
logger.warning(f"获取渠道 {channel_id} 凭证状态失败", exc_info=True)
config = self._channels_config.get(channel_id, {})
has_credential = any(
any(pattern in k.lower() for pattern in self.SENSITIVE_KEY_PATTERNS) and v for k, v in config.items()
)
return {
"has_credential": has_credential,
"is_expired": False,
"source": "config" if has_credential else "none",
}
async def get_credential_status(self, channel_id: str) -> dict:
adapter = self._adapters.get(channel_id)
if adapter:
return await self._get_credential_summary(channel_id, adapter)
config = self._channels_config.get(channel_id, {})
has_credential = any(
any(pattern in k.lower() for pattern in self.SENSITIVE_KEY_PATTERNS) and v for k, v in config.items()
)
return {
"has_credential": has_credential,
"is_expired": False,
"source": "config" if has_credential else "none",
}
async def refresh_credential(self, channel_id: str) -> dict:
adapter = self._adapters.get(channel_id)
if adapter and hasattr(adapter, "refresh_credential"):
result = await adapter.refresh_credential()
return result
if adapter:
logger.info(f"渠道 {channel_id} 适配器不支持凭证刷新,尝试重连")
await self.restart_channel(channel_id)
return {"refreshed": True, "method": "restart", "message": "通过重启渠道刷新凭证"}
raise ValueError(f"Channel '{channel_id}' is not running, cannot refresh credential")
async def _ensure_virtual_department(self, db) -> None:
from yuxi.storage.postgres.models_business import Department
result = await db.execute(select(Department).where(Department.id == -1))
if result.scalar_one_or_none() is None:
dept = Department(
id=-1,
name="\u6e20\u9053\u7f51\u5173",
description="\u591a\u6e20\u9053\u7f51\u5173\u865a\u62df\u90e8\u95e8",
)
db.add(dept)
await db.commit()
logger.info("Created virtual department (id=-1)")
async def _ensure_default_agent_config(self, db) -> None:
from yuxi.repositories.agent_config_repository import AgentConfigRepository
repo = AgentConfigRepository(db)
await repo.get_or_create_default(
department_id=-1,
agent_id="ChatbotAgent",
created_by="system",
)
logger.info("Ensured default agent config for ChatbotAgent")
_channel_manager: ChannelManager | None = None
_channel_manager_lock = threading.Lock()
def get_channel_manager() -> ChannelManager:
global _channel_manager
if _channel_manager is None:
with _channel_manager_lock:
if _channel_manager is None:
_channel_manager = ChannelManager()
return _channel_manager
def _now_iso() -> str:
from yuxi.utils.datetime_utils import format_utc_datetime, utc_now_naive
return format_utc_datetime(utc_now_naive())