主要变更: 1. 修复速率限流器使用setdefault替代重复创建令牌桶 2. 重构交互注册表匹配逻辑,优化精确匹配查找 3. 重构去重缓存逻辑,移到适配器实例方法 4. 重构发送URL解析,增加合法性校验并拆分公共方法 5. 优化流式消息处理逻辑,简化flush_controller调用 6. 重构群聊类型判断代码,简化语法 7. 修复重连管理器对None类型关闭分类的处理 8. 新增消息缓存、线程模拟器、发送初始化模块 9. 重构凭证备份与会话存储逻辑,支持环境变量指定状态目录 10. 新增配置提示与向导二维码绑定功能 11. 优化媒体上传逻辑,增加重试机制与缓存 12. 新增审批键盘模板构建函数 13. 重构消息格式处理,修正媒体发送字段与长度限制 14. 修复令牌过期时间计算,使用time.time替代monotonic 15. 新增群组激活缓冲区与用户追踪器增强功能 16. 修复换行符问题,统一文件结尾格式
218 lines
7.7 KiB
Python
218 lines
7.7 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import re
|
|
from collections.abc import Awaitable, Callable
|
|
from dataclasses import dataclass, field
|
|
|
|
import aiohttp
|
|
|
|
from yuxi.channels.exceptions import DeliveryFailedError
|
|
from yuxi.channels.models import DeliveryResult
|
|
from yuxi.utils.logging_config import logger
|
|
|
|
from .constants import DM_CHAT_PREFIX, GROUP_CHAT_PREFIX
|
|
|
|
|
|
@dataclass
|
|
class MessageSeqManager:
|
|
next_seq: int = 1
|
|
_passive_seq: int = 0
|
|
_active_seq: int = 1
|
|
_lock: asyncio.Lock = field(default_factory=asyncio.Lock)
|
|
|
|
MAX_SEQ: int = 2**31 - 1
|
|
|
|
async def acquire_active(self) -> int:
|
|
async with self._lock:
|
|
if self._active_seq > self.MAX_SEQ:
|
|
self._active_seq = 1
|
|
seq = self._active_seq
|
|
self._active_seq += 1
|
|
return seq
|
|
|
|
async def acquire_passive(self) -> int:
|
|
async with self._lock:
|
|
if self._passive_seq > self.MAX_SEQ:
|
|
self._passive_seq = 0
|
|
seq = self._passive_seq
|
|
self._passive_seq -= 1
|
|
return seq
|
|
|
|
def reset(self) -> None:
|
|
self._active_seq = 1
|
|
self._passive_seq = 0
|
|
self.next_seq = 1
|
|
|
|
def snapshot(self) -> dict:
|
|
return {
|
|
"active_seq": self._active_seq,
|
|
"passive_seq": self._passive_seq,
|
|
}
|
|
|
|
def restore(self, snapshot: dict) -> None:
|
|
self._active_seq = snapshot.get("active_seq", 1)
|
|
self._passive_seq = snapshot.get("passive_seq", 0)
|
|
self.next_seq = self._active_seq
|
|
|
|
|
|
async def send_with_retry(
|
|
http_client: aiohttp.ClientSession,
|
|
token: str,
|
|
api_base: str,
|
|
payload: dict,
|
|
chat_id: str,
|
|
config: dict | None = None,
|
|
token_refresh_cb: Callable[[], Awaitable[str]] | None = None,
|
|
on_sent: Callable[[DeliveryResult], Awaitable[None]] | None = None,
|
|
sender_headers: dict | None = None,
|
|
) -> DeliveryResult:
|
|
cfg = config or {}
|
|
max_retries = cfg.get("retry", {}).get("attempts", 3)
|
|
min_delay = cfg.get("retry", {}).get("min_delay_ms", 400) / 1000
|
|
max_delay = cfg.get("retry", {}).get("max_delay_ms", 30000) / 1000
|
|
|
|
current_token = token
|
|
token_refreshed = False
|
|
|
|
rate_limit_retries = 0
|
|
max_rate_limit_retries = cfg.get("retry", {}).get("max_rate_limit_retries", 3)
|
|
|
|
last_error = None
|
|
url = resolve_send_url(api_base, chat_id)
|
|
|
|
for attempt in range(max_retries):
|
|
try:
|
|
headers = {
|
|
"Authorization": f"QQBot {current_token}",
|
|
"Content-Type": "application/json",
|
|
}
|
|
if sender_headers:
|
|
headers.update(sender_headers)
|
|
async with http_client.post(url, json=payload, headers=headers) as resp:
|
|
if resp.status == 200:
|
|
data = await resp.json()
|
|
result = DeliveryResult(
|
|
success=True,
|
|
message_id=data.get("id") or data.get("message_id"),
|
|
)
|
|
if on_sent:
|
|
await on_sent(result)
|
|
return result
|
|
elif resp.status == 429:
|
|
rate_limit_retries += 1
|
|
if rate_limit_retries > max_rate_limit_retries:
|
|
raise DeliveryFailedError(
|
|
f"Rate limited after {max_rate_limit_retries} consecutive 429 responses"
|
|
)
|
|
retry_after = int(resp.headers.get("Retry-After", "30"))
|
|
logger.warning(
|
|
"[QQBot] Rate limited (%d/%d), retry after %ds",
|
|
rate_limit_retries,
|
|
max_rate_limit_retries,
|
|
retry_after,
|
|
)
|
|
await asyncio.sleep(retry_after)
|
|
continue
|
|
elif resp.status in (401, 403):
|
|
# 401/403 are permanent auth failures. Only retry if:
|
|
# - Status is 401 (not 403) AND
|
|
# - A token refresh callback is available AND
|
|
# - Token hasn't been refreshed yet in this retry cycle
|
|
if resp.status == 401 and token_refresh_cb and not token_refreshed:
|
|
logger.warning("[QQBot] 401 received, refreshing token and retrying")
|
|
try:
|
|
current_token = await token_refresh_cb()
|
|
token_refreshed = True
|
|
continue
|
|
except Exception as e:
|
|
logger.error(f"[QQBot] Token refresh after 401 failed: {e}")
|
|
error_body = await resp.text()
|
|
result = DeliveryResult(success=False, error=f"Auth failed ({resp.status}): {error_body}")
|
|
if on_sent:
|
|
await on_sent(result)
|
|
raise DeliveryFailedError(f"Auth failed ({resp.status}): {error_body}")
|
|
elif 400 <= resp.status < 500:
|
|
error_body = await resp.text()
|
|
result = DeliveryResult(success=False, error=f"Client error ({resp.status}): {error_body}")
|
|
if on_sent:
|
|
await on_sent(result)
|
|
raise DeliveryFailedError(f"Client error ({resp.status}): {error_body}")
|
|
else:
|
|
error_body = await resp.text()
|
|
last_error = DeliveryFailedError(f"Server error ({resp.status}): {error_body}")
|
|
|
|
except DeliveryFailedError:
|
|
raise
|
|
except Exception as e:
|
|
last_error = DeliveryFailedError(str(e))
|
|
|
|
if attempt < max_retries - 1:
|
|
delay = min(min_delay * (2**attempt), max_delay)
|
|
await asyncio.sleep(delay)
|
|
|
|
result = DeliveryResult(
|
|
success=False,
|
|
error=str(last_error) if last_error else "Max retries exceeded",
|
|
)
|
|
if on_sent:
|
|
await on_sent(result)
|
|
return result
|
|
|
|
|
|
_ALLOWED_API_HOSTS = frozenset({"api.sgroup.qq.com", "sandbox.api.sgroup.qq.com"})
|
|
_CHAT_ID_PATTERN = re.compile(r"^[a-zA-Z0-9_-]+$")
|
|
|
|
|
|
def _validate_api_base(api_base: str) -> None:
|
|
from urllib.parse import urlparse
|
|
|
|
parsed = urlparse(api_base)
|
|
host = parsed.hostname
|
|
if host not in _ALLOWED_API_HOSTS:
|
|
raise DeliveryFailedError(f"Invalid api_base host: {host}. Allowed: {sorted(_ALLOWED_API_HOSTS)}")
|
|
|
|
|
|
def _validate_chat_id(chat_id: str) -> None:
|
|
raw = chat_id
|
|
if raw.startswith(GROUP_CHAT_PREFIX):
|
|
raw = raw[len(GROUP_CHAT_PREFIX) :]
|
|
elif raw.startswith(DM_CHAT_PREFIX):
|
|
raw = raw[len(DM_CHAT_PREFIX) :]
|
|
if raw and not _CHAT_ID_PATTERN.match(raw):
|
|
raise DeliveryFailedError(f"Invalid chat_id format: {chat_id}")
|
|
|
|
|
|
def resolve_send_url(api_base: str, chat_id: str) -> str:
|
|
_validate_api_base(api_base)
|
|
_validate_chat_id(chat_id)
|
|
|
|
if chat_id.startswith(GROUP_CHAT_PREFIX):
|
|
group_openid = chat_id.replace(GROUP_CHAT_PREFIX, "")
|
|
return f"{api_base}/v2/groups/{group_openid}/messages"
|
|
elif chat_id.startswith(DM_CHAT_PREFIX):
|
|
openid = chat_id.replace(DM_CHAT_PREFIX, "")
|
|
return f"{api_base}/v2/users/{openid}/messages"
|
|
elif chat_id:
|
|
return f"{api_base}/v2/channels/{chat_id}/messages"
|
|
return f"{api_base}/v2/users/@me/messages"
|
|
|
|
|
|
def render_reply_payload(
|
|
content: str,
|
|
msg_type: int = 0,
|
|
msg_id: str = "",
|
|
chunk_index: int = 0,
|
|
total_chunks: int = 1,
|
|
) -> dict:
|
|
payload: dict = {
|
|
"content": content,
|
|
"msg_type": msg_type,
|
|
}
|
|
if msg_id:
|
|
payload["msg_id"] = msg_id
|
|
if total_chunks > 1:
|
|
payload["chunk_index"] = chunk_index
|
|
payload["total_chunks"] = total_chunks
|
|
return payload
|