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.runtime_state import RuntimeState
from yuxi.channels.services.stats_collector import StatsCollector from yuxi.channels.services.stats_collector import StatsCollector
from yuxi.channels.services.webhook_registry import WebhookRegistry 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__ = [ __all__ = [
"ChatAbortEntry", "ChatAbortEntry",

View File

@ -3,7 +3,7 @@ from __future__ import annotations
import asyncio import asyncio
from collections.abc import Awaitable, Callable from collections.abc import Awaitable, Callable
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Any, TYPE_CHECKING from typing import TYPE_CHECKING, Any
if TYPE_CHECKING: if TYPE_CHECKING:
from yuxi.channels.models import ChannelAccountSnapshot from yuxi.channels.models import ChannelAccountSnapshot

View File

@ -32,4 +32,11 @@ class MaintenanceRunner:
await adapter._refresh_token_if_needed() await adapter._refresh_token_if_needed()
except Exception: except Exception:
logger.debug(f"Token refresh skip for {channel_id}") 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") 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 from __future__ import annotations
import asyncio import asyncio
import time
from collections import deque
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from yuxi.utils.logging_config import logger from yuxi.utils.logging_config import logger
@ -8,10 +10,27 @@ from yuxi.utils.logging_config import logger
if TYPE_CHECKING: if TYPE_CHECKING:
from yuxi.channels.services.runtime_state import RuntimeState from yuxi.channels.services.runtime_state import RuntimeState
_MAX_RESPONSE_TIME_SAMPLES = 100
class StatsCollector: class StatsCollector:
def __init__(self, state: RuntimeState): def __init__(self, state: RuntimeState):
self._state = state 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: async def run(self, interval: float = 60) -> None:
logger.info("StatsCollector started") logger.info("StatsCollector started")
@ -25,7 +44,38 @@ class StatsCollector:
logger.exception("StatsCollector error") logger.exception("StatsCollector error")
def _collect(self) -> None: 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"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,
}