refactor(channel-services): 整理导入顺序并新增插件状态清理功能

1. 调整多个服务文件的导入顺序以符合规范
2. 新增Postgres插件状态存储实现,支持CRUD和过期清理
3. 在维护任务中添加插件过期状态自动清理逻辑
4. 新增统计收集器,支持请求/错误计数、响应时间统计和指标上报
This commit is contained in:
Kris 2026-05-13 16:21:09 +08:00
parent f8a985aad8
commit fb8bd90329
5 changed files with 233 additions and 4 deletions

View File

@ -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",

View File

@ -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

View File

@ -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")

View File

@ -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

View File

@ -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,
}