这是一个批量整理提交,包含以下主要改动: 1. 删除多处冗余的空行和未使用的导入 2. 修复文件末尾缺少换行符的问题 3. 调整部分模块的导入顺序与代码排版 4. 修复部分配置默认值与策略逻辑 5. 新增多个功能模块与辅助工具 6. 完善异常处理与日志记录 7. 修复速率限制、消息缓存、权限校验等逻辑bug 8. 废弃部分旧有API与配置项并添加警告提示
131 lines
4.6 KiB
Python
131 lines
4.6 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import time
|
|
|
|
import aiohttp
|
|
|
|
from yuxi.channels.exceptions import ChannelAuthenticationError
|
|
from yuxi.utils.logging_config import logger
|
|
|
|
|
|
class QQBotTokenManager:
|
|
def __init__(
|
|
self,
|
|
app_id: str,
|
|
app_secret: str,
|
|
sandbox: bool = False,
|
|
http_client: aiohttp.ClientSession | None = None,
|
|
):
|
|
self.app_id = app_id
|
|
self.app_secret = app_secret
|
|
self.sandbox = sandbox
|
|
self._http_client = http_client
|
|
self._access_token: str | None = None
|
|
self._expires_at: float | None = None
|
|
self._token_lock = asyncio.Lock()
|
|
self._refresh_in_progress: asyncio.Event | None = None
|
|
self._refresh_task: asyncio.Task | None = None
|
|
self._refresh_interval: float = 60.0
|
|
|
|
@property
|
|
def api_base(self) -> str:
|
|
if self.sandbox:
|
|
return "https://sandbox.api.sgroup.qq.com"
|
|
return "https://api.sgroup.qq.com"
|
|
|
|
@property
|
|
def access_token(self) -> str | None:
|
|
return self._access_token
|
|
|
|
@property
|
|
def expires_at(self) -> float | None:
|
|
return self._expires_at
|
|
|
|
async def get_token(self) -> str:
|
|
async with self._token_lock:
|
|
if self._is_expired():
|
|
await self._do_refresh()
|
|
return self._access_token
|
|
|
|
async def force_refresh(self) -> str:
|
|
async with self._token_lock:
|
|
await self._do_refresh()
|
|
return self._access_token
|
|
|
|
async def _do_refresh(self) -> None:
|
|
if self._refresh_in_progress is not None:
|
|
await self._refresh_in_progress.wait()
|
|
return
|
|
|
|
self._refresh_in_progress = asyncio.Event()
|
|
try:
|
|
await self._refresh()
|
|
self._refresh_in_progress.set()
|
|
except Exception:
|
|
self._refresh_in_progress.set()
|
|
raise
|
|
finally:
|
|
self._refresh_in_progress = None
|
|
|
|
async def _refresh(self) -> None:
|
|
client = self._http_client or aiohttp.ClientSession()
|
|
try:
|
|
async with client.post(
|
|
f"{self.api_base}/oauth2/token",
|
|
json={
|
|
"app_id": self.app_id,
|
|
"app_secret": self.app_secret,
|
|
},
|
|
) as resp:
|
|
if resp.status != 200:
|
|
raise ChannelAuthenticationError(f"Token refresh failed: HTTP {resp.status}")
|
|
data = await resp.json()
|
|
self._access_token = data["access_token"]
|
|
expires_in = data.get("expires_in", 7200)
|
|
self._expires_at = time.monotonic() + expires_in
|
|
logger.info(f"[QQBot] Token refreshed, expires in {expires_in}s (app_id={self.app_id[:6]}...)")
|
|
finally:
|
|
if not self._http_client:
|
|
await client.close()
|
|
|
|
def start_background_refresh(self) -> None:
|
|
if self._refresh_task is not None and not self._refresh_task.done():
|
|
return
|
|
self._refresh_task = asyncio.create_task(self._background_refresh_loop())
|
|
logger.debug(f"[QQBot] Background token refresh started (interval={self._refresh_interval}s)")
|
|
|
|
def stop_background_refresh(self) -> None:
|
|
if self._refresh_task and not self._refresh_task.done():
|
|
self._refresh_task.cancel()
|
|
self._refresh_task = None
|
|
logger.debug("[QQBot] Background token refresh stopped")
|
|
|
|
async def _background_refresh_loop(self) -> None:
|
|
while True:
|
|
try:
|
|
sleep_duration = self._refresh_interval
|
|
if self._expires_at is not None:
|
|
remaining = self._expires_at - time.monotonic() - 300
|
|
sleep_duration = max(60.0, min(self._refresh_interval, max(remaining, 0.0)))
|
|
|
|
await asyncio.sleep(sleep_duration)
|
|
async with self._token_lock:
|
|
if not self._is_expired():
|
|
continue
|
|
try:
|
|
await self._do_refresh()
|
|
except Exception as e:
|
|
logger.warning(f"[QQBot] Background token refresh failed (will retry): {e}")
|
|
except asyncio.CancelledError:
|
|
logger.debug("[QQBot] Background token refresh cancelled")
|
|
break
|
|
except Exception as e:
|
|
logger.error(f"[QQBot] Background token refresh loop error: {e}")
|
|
await asyncio.sleep(5)
|
|
|
|
def _is_expired(self) -> bool:
|
|
if self._access_token is None or self._expires_at is None:
|
|
return True
|
|
return time.monotonic() > self._expires_at - 300
|