1. 简化message_actions.py中获取适配器的逻辑 2. 新增适配器合法性校验工具方法 3. 新增会话映射过期清理功能 4. 重构渠道状态机与基础适配器实现 5. 统一渠道操作异常处理逻辑 6. 新增凭证状态查询与刷新接口 7. 优化健康检查与自动重连逻辑 8. 新增统计数据缓存与批量查询优化 9. 修复部分数据库操作的异常处理逻辑
1379 lines
56 KiB
Python
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())
|