本次提交对Twitch适配器进行了全面升级与优化: 1. 修复UTF8截断逻辑,避免越界访问 2. 重构群聊策略配置,标准化mention相关规则 3. 新增消息缓存管理器,支持通过消息ID查询已发送消息 4. 更新配置schema,新增prefer_helix_send开关和deprecated策略自动转换 5. 新增CLEARMSG和ROOMSTATE IRC消息解析,补充事件订阅支持 6. 优化令牌刷新逻辑,增加重试机制与退避策略 7. 新增Helix API聊天消息发送、删除和公告功能 8. 扩展事件订阅类型,新增直播状态、频道更新等系统事件 9. 新增reply、delete_message、announcement等动作支持,完善操作能力 10. 重构流式发送逻辑,新增进度指示器和配置项 11. 优化重连策略,增加指数退避与计数重置
144 lines
5.1 KiB
Python
144 lines
5.1 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import time
|
|
from abc import ABC, abstractmethod
|
|
from typing import Any
|
|
|
|
from yuxi.utils.logging_config import logger
|
|
|
|
from .probe import refresh_access_token, validate_token
|
|
|
|
|
|
class AuthProvider(ABC):
|
|
@abstractmethod
|
|
async def get_access_token(self) -> str | None: ...
|
|
|
|
@abstractmethod
|
|
async def get_client_id(self) -> str: ...
|
|
|
|
@abstractmethod
|
|
async def validate(self) -> bool: ...
|
|
|
|
|
|
class StaticAuthProvider(AuthProvider):
|
|
def __init__(self, client_id: str, access_token: str):
|
|
self._client_id = client_id
|
|
self._access_token = access_token
|
|
|
|
async def get_access_token(self) -> str | None:
|
|
return self._access_token or None
|
|
|
|
async def get_client_id(self) -> str:
|
|
return self._client_id
|
|
|
|
async def validate(self) -> bool:
|
|
if not self._client_id or not self._access_token:
|
|
return False
|
|
user_info = await validate_token(self._client_id, self._access_token)
|
|
return user_info is not None
|
|
|
|
|
|
class RefreshingAuthProvider(AuthProvider):
|
|
def __init__(
|
|
self,
|
|
client_id: str,
|
|
client_secret: str,
|
|
access_token: str = "",
|
|
refresh_token: str = "",
|
|
):
|
|
self._client_id = client_id
|
|
self._client_secret = client_secret
|
|
self._access_token = access_token
|
|
self._refresh_token = refresh_token
|
|
self._token_expires_at: float | None = None
|
|
self._token_obtained_at: float | None = None
|
|
self._token_expires_in: int = 0
|
|
|
|
async def get_access_token(self) -> str | None:
|
|
if self._is_expired():
|
|
refreshed = await self._refresh()
|
|
if not refreshed:
|
|
logger.warning("[TwitchAuth] Token expired and refresh failed")
|
|
return None
|
|
return self._access_token or None
|
|
|
|
async def get_client_id(self) -> str:
|
|
return self._client_id
|
|
|
|
async def validate(self) -> bool:
|
|
if not self._client_id or not self._access_token:
|
|
return False
|
|
user_info = await validate_token(self._client_id, self._access_token)
|
|
return user_info is not None
|
|
|
|
async def _refresh(self) -> bool:
|
|
if not self._client_id or not self._client_secret or not self._refresh_token:
|
|
return False
|
|
|
|
max_retries = 3
|
|
for attempt in range(max_retries):
|
|
try:
|
|
new_tokens = await refresh_access_token(self._client_id, self._client_secret, self._refresh_token)
|
|
if new_tokens is None:
|
|
if attempt < max_retries - 1:
|
|
delay = 2**attempt
|
|
logger.warning(f"[TwitchAuth] Token refresh attempt {attempt + 1} failed, retrying in {delay}s")
|
|
await asyncio.sleep(delay)
|
|
continue
|
|
return False
|
|
new_access = new_tokens.get("access_token", "")
|
|
new_refresh = new_tokens.get("refresh_token", "")
|
|
expires_in = new_tokens.get("expires_in", 0)
|
|
if new_access:
|
|
self._access_token = new_access
|
|
if new_refresh:
|
|
self._refresh_token = new_refresh
|
|
if expires_in:
|
|
self._token_obtained_at = time.time()
|
|
self._token_expires_in = expires_in
|
|
self._token_expires_at = self._token_obtained_at + expires_in - 300
|
|
logger.info(f"[TwitchAuth] Token refreshed, expires in {expires_in}s")
|
|
return True
|
|
except Exception as e:
|
|
if attempt < max_retries - 1:
|
|
delay = 2**attempt
|
|
logger.warning(
|
|
f"[TwitchAuth] Token refresh error (attempt {attempt + 1}): {e}, retrying in {delay}s"
|
|
)
|
|
await asyncio.sleep(delay)
|
|
else:
|
|
logger.error(f"[TwitchAuth] Token refresh failed after {max_retries} attempts: {e}")
|
|
return False
|
|
|
|
def _is_expired(self) -> bool:
|
|
if self._token_expires_at is None:
|
|
return False
|
|
return time.time() >= self._token_expires_at
|
|
|
|
@property
|
|
def expires_at(self) -> float | None:
|
|
return self._token_expires_at
|
|
|
|
@property
|
|
def expires_in(self) -> int:
|
|
return self._token_expires_in
|
|
|
|
|
|
def create_auth_provider(config: dict[str, Any]) -> AuthProvider:
|
|
client_id = config.get("client_id", "")
|
|
client_secret = config.get("client_secret", "")
|
|
access_token = config.get("access_token", "")
|
|
refresh_token = config.get("refresh_token", "")
|
|
|
|
if client_secret and refresh_token:
|
|
logger.info("[TwitchAuth] Using RefreshingAuthProvider (auto-refresh enabled)")
|
|
return RefreshingAuthProvider(
|
|
client_id=client_id,
|
|
client_secret=client_secret,
|
|
access_token=access_token,
|
|
refresh_token=refresh_token,
|
|
)
|
|
logger.info("[TwitchAuth] Using StaticAuthProvider (manual token rotation)")
|
|
return StaticAuthProvider(client_id=client_id, access_token=access_token)
|