diff --git a/backend/package/yuxi/channels/services/__init__.py b/backend/package/yuxi/channels/services/__init__.py index bc3ba66e..1f80e314 100644 --- a/backend/package/yuxi/channels/services/__init__.py +++ b/backend/package/yuxi/channels/services/__init__.py @@ -4,7 +4,7 @@ from yuxi.channels.services.maintenance import MaintenanceRunner 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, WsLogEntry +from yuxi.channels.services.ws_logger import WsLogEntry, WsLogger __all__ = [ "ChatAbortEntry", diff --git a/backend/package/yuxi/channels/services/context.py b/backend/package/yuxi/channels/services/context.py index 5657eb66..3a9a0e38 100644 --- a/backend/package/yuxi/channels/services/context.py +++ b/backend/package/yuxi/channels/services/context.py @@ -3,7 +3,7 @@ from __future__ import annotations import asyncio from collections.abc import Awaitable, Callable from dataclasses import dataclass, field -from typing import Any, TYPE_CHECKING +from typing import TYPE_CHECKING, Any if TYPE_CHECKING: from yuxi.channels.models import ChannelAccountSnapshot diff --git a/backend/package/yuxi/channels/services/maintenance.py b/backend/package/yuxi/channels/services/maintenance.py index 164ac6a6..86edff2f 100644 --- a/backend/package/yuxi/channels/services/maintenance.py +++ b/backend/package/yuxi/channels/services/maintenance.py @@ -32,4 +32,11 @@ class MaintenanceRunner: await adapter._refresh_token_if_needed() except Exception: logger.debug(f"Token refresh skip for {channel_id}") + + if self._manager._state_store is not None: + try: + await self._manager._state_store.cleanup_expired() + except Exception: + logger.debug("Plugin state cleanup failed") + logger.debug("Maintenance cycle complete") diff --git a/backend/package/yuxi/channels/services/plugin_state_store.py b/backend/package/yuxi/channels/services/plugin_state_store.py new file mode 100644 index 00000000..7bae960a --- /dev/null +++ b/backend/package/yuxi/channels/services/plugin_state_store.py @@ -0,0 +1,172 @@ +from __future__ import annotations + +import time +from abc import ABC, abstractmethod +from datetime import UTC, datetime, timedelta +from typing import Any + +from sqlalchemy import delete, select + +from yuxi.storage.postgres.manager import pg_manager +from yuxi.storage.postgres.models_channels import ChannelPluginState +from yuxi.utils.logging_config import logger + + +def _utc_now() -> datetime: + return datetime.fromtimestamp(time.time(), tz=UTC) + + +class PluginStateStore(ABC): + """渠道插件状态统一存储抽象""" + + @abstractmethod + async def get(self, channel_id: str, key: str, namespace: str = "default") -> Any | None: + """读取状态值,过期返回 None""" + ... + + @abstractmethod + async def set( + self, + channel_id: str, + key: str, + value: Any, + namespace: str = "default", + ttl_seconds: int | None = None, + ) -> None: + """写入状态值,可选 TTL""" + ... + + @abstractmethod + async def delete(self, channel_id: str, key: str, namespace: str = "default") -> None: + """删除状态""" + ... + + @abstractmethod + async def list_keys(self, channel_id: str, namespace: str = "default") -> list[str]: + """列出命名空间下所有 key""" + ... + + @abstractmethod + async def cleanup_expired(self) -> int: + """清理过期状态,返回清理数量""" + ... + + @abstractmethod + async def get_all(self, channel_id: str) -> dict[str, dict[str, Any]]: + """获取渠道所有状态(用于管理面板)""" + ... + + +class PostgresPluginStateStore(PluginStateStore): + def __init__(self): + pass + + def _session(self): + return pg_manager.get_async_session_context() + + async def get(self, channel_id: str, key: str, namespace: str = "default") -> Any | None: + async with self._session() as db: + stmt = select(ChannelPluginState).where( + ChannelPluginState.channel_id == channel_id, + ChannelPluginState.namespace == namespace, + ChannelPluginState.entry_key == key, + ) + result = await db.execute(stmt) + row = result.scalar_one_or_none() + if row is None: + return None + if row.expires_at and row.expires_at < _utc_now(): + await db.delete(row) + await db.commit() + return None + return row.value_json + + async def set( + self, + channel_id: str, + key: str, + value: Any, + namespace: str = "default", + ttl_seconds: int | None = None, + ) -> None: + expires_at = _utc_now() + timedelta(seconds=ttl_seconds) if ttl_seconds is not None else None + + async with self._session() as db: + stmt = select(ChannelPluginState).where( + ChannelPluginState.channel_id == channel_id, + ChannelPluginState.namespace == namespace, + ChannelPluginState.entry_key == key, + ) + result = await db.execute(stmt) + row = result.scalar_one_or_none() + + if row: + row.value_json = value + row.updated_at = _utc_now() + row.expires_at = expires_at + else: + row = ChannelPluginState( + channel_id=channel_id, + namespace=namespace, + entry_key=key, + value_json=value, + expires_at=expires_at, + ) + db.add(row) + await db.commit() + + async def delete(self, channel_id: str, key: str, namespace: str = "default") -> None: + async with self._session() as db: + await db.execute( + delete(ChannelPluginState).where( + ChannelPluginState.channel_id == channel_id, + ChannelPluginState.namespace == namespace, + ChannelPluginState.entry_key == key, + ) + ) + await db.commit() + + async def list_keys(self, channel_id: str, namespace: str = "default") -> list[str]: + async with self._session() as db: + stmt = select(ChannelPluginState.entry_key).where( + ChannelPluginState.channel_id == channel_id, + ChannelPluginState.namespace == namespace, + ) + result = await db.execute(stmt) + return [row[0] for row in result.all()] + + async def cleanup_expired(self) -> int: + async with self._session() as db: + result = await db.execute( + delete(ChannelPluginState).where( + ChannelPluginState.expires_at.isnot(None), + ChannelPluginState.expires_at < _utc_now(), + ) + ) + await db.commit() + count = result.rowcount + if count > 0: + logger.info(f"[PluginState] Cleaned up {count} expired entries") + return count + + async def get_all(self, channel_id: str) -> dict[str, dict[str, Any]]: + result: dict[str, dict[str, Any]] = {} + async with self._session() as db: + stmt = select(ChannelPluginState).where( + ChannelPluginState.channel_id == channel_id, + ) + rows = (await db.execute(stmt)).scalars().all() + for row in rows: + if row.expires_at and row.expires_at < _utc_now(): + await db.delete(row) + continue + ns = row.namespace + if ns not in result: + result[ns] = {} + result[ns][row.entry_key] = { + "value": row.value_json, + "expires_at": row.expires_at.isoformat() if row.expires_at else None, + "updated_at": row.updated_at.isoformat() if row.updated_at else None, + } + await db.commit() + return result diff --git a/backend/package/yuxi/channels/services/stats_collector.py b/backend/package/yuxi/channels/services/stats_collector.py index 298b4992..350be5aa 100644 --- a/backend/package/yuxi/channels/services/stats_collector.py +++ b/backend/package/yuxi/channels/services/stats_collector.py @@ -1,6 +1,8 @@ from __future__ import annotations import asyncio +import time +from collections import deque from typing import TYPE_CHECKING from yuxi.utils.logging_config import logger @@ -8,10 +10,27 @@ from yuxi.utils.logging_config import logger if TYPE_CHECKING: from yuxi.channels.services.runtime_state import RuntimeState +_MAX_RESPONSE_TIME_SAMPLES = 100 + class StatsCollector: def __init__(self, state: RuntimeState): self._state = state + self._response_times: deque[float] = deque(maxlen=_MAX_RESPONSE_TIME_SAMPLES) + self._request_count_local: int = 0 + self._error_count_local: int = 0 + self._last_collect_at = time.monotonic() + + def record_request(self) -> None: + self._state.request_count += 1 + self._request_count_local += 1 + + def record_error(self) -> None: + self._state.error_count += 1 + self._error_count_local += 1 + + def record_response_time(self, ms: float) -> None: + self._response_times.append(ms) async def run(self, interval: float = 60) -> None: logger.info("StatsCollector started") @@ -25,7 +44,38 @@ class StatsCollector: logger.exception("StatsCollector error") def _collect(self) -> None: - logger.debug( + now = time.monotonic() + elapsed = now - self._last_collect_at + self._last_collect_at = now + + rps = self._request_count_local / elapsed if elapsed > 0 else 0 + eps = self._error_count_local / elapsed if elapsed > 0 else 0 + + avg_rt = sum(self._response_times) / len(self._response_times) if self._response_times else 0 + p95_rt = 0.0 + if len(self._response_times) >= 20: + sorted_times = sorted(self._response_times) + p95_idx = int(len(sorted_times) * 0.95) + p95_rt = sorted_times[p95_idx] if p95_idx < len(sorted_times) else sorted_times[-1] + + error_rate = self._error_count_local / self._request_count_local if self._request_count_local else 0 + + logger.info( f"Stats: channels={self._state.active_channels}, " - f"requests={self._state.request_count}, errors={self._state.error_count}" + f"rps={rps:.1f}, eps={eps:.1f}, errors={self._error_count_local}, " + f"error_rate={error_rate:.2%}, " + f"avg_rt={avg_rt:.0f}ms, p95_rt={p95_rt:.0f}ms" ) + + self._request_count_local = 0 + self._error_count_local = 0 + + def get_summary(self) -> dict: + avg_rt = sum(self._response_times) / len(self._response_times) if self._response_times else 0 + return { + "active_channels": self._state.active_channels, + "total_requests": self._state.request_count, + "total_errors": self._state.error_count, + "avg_response_time_ms": round(avg_rt, 1), + "phase": self._state.phase, + }