from __future__ import annotations import asyncio 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, ) -> 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 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", } 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: retry_after = int(resp.headers.get("Retry-After", "30")) logger.warning(f"[QQBot] Rate limited, retry after {retry_after}s") await asyncio.sleep(retry_after) continue elif resp.status in (401, 403): 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 def _resolve_send_url(api_base: str, chat_id: str) -> str: 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