这是一个批量整理提交,包含以下主要改动: 1. 删除多处冗余的空行和未使用的导入 2. 修复文件末尾缺少换行符的问题 3. 调整部分模块的导入顺序与代码排版 4. 修复部分配置默认值与策略逻辑 5. 新增多个功能模块与辅助工具 6. 完善异常处理与日志记录 7. 修复速率限制、消息缓存、权限校验等逻辑bug 8. 废弃部分旧有API与配置项并添加警告提示
228 lines
7.6 KiB
Python
228 lines
7.6 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import ssl
|
|
import time
|
|
from collections.abc import Callable
|
|
from typing import Any
|
|
|
|
from yuxi.utils.logging_config import logger
|
|
|
|
|
|
class IRCConnection:
|
|
PING_TIMEOUT = 30
|
|
RECONNECT_DELAY = 5
|
|
MAX_RECONNECT_DELAY = 300
|
|
|
|
def __init__(self, config: dict[str, Any]):
|
|
self._server = config.get("server", "irc.libera.chat")
|
|
self._port = config.get("port", 6697)
|
|
self._use_tls = config.get("use_tls", True)
|
|
self._connect_timeout = config.get("connect_timeout", 30.0)
|
|
self._proxy = config.get("proxy")
|
|
self._last_pong = time.monotonic()
|
|
self._reconnect_attempts = 0
|
|
self._reader: asyncio.StreamReader | None = None
|
|
self._writer: asyncio.StreamWriter | None = None
|
|
self._keepalive_task: asyncio.Task | None = None
|
|
self._on_reconnect: Callable[[], None] | None = None
|
|
|
|
@property
|
|
def reader(self) -> asyncio.StreamReader | None:
|
|
return self._reader
|
|
|
|
@property
|
|
def writer(self) -> asyncio.StreamWriter | None:
|
|
return self._writer
|
|
|
|
def on_reconnect(self, callback: Callable[[], None]) -> None:
|
|
self._on_reconnect = callback
|
|
|
|
async def connect(self) -> None:
|
|
if self._proxy:
|
|
await self._connect_via_proxy()
|
|
elif self._use_tls:
|
|
ctx = ssl.create_default_context()
|
|
ctx.check_hostname = True
|
|
ctx.verify_mode = ssl.CERT_REQUIRED
|
|
self._reader, self._writer = await asyncio.open_connection(
|
|
host=self._server,
|
|
port=self._port,
|
|
ssl=ctx,
|
|
server_hostname=self._server,
|
|
)
|
|
else:
|
|
self._reader, self._writer = await asyncio.open_connection(
|
|
host=self._server,
|
|
port=self._port,
|
|
)
|
|
self._reconnect_attempts = 0
|
|
self._last_pong = time.monotonic()
|
|
logger.info(f"IRC TCP connected: {self._server}:{self._port} (TLS={self._use_tls})")
|
|
|
|
async def _connect_via_proxy(self) -> None:
|
|
proxy = self._proxy or {}
|
|
proxy_host = proxy.get("host", "")
|
|
proxy_port = proxy.get("port", 1080)
|
|
|
|
if not proxy_host:
|
|
raise ConnectionError("Proxy host is required")
|
|
|
|
self._reader, self._writer = await asyncio.open_connection(
|
|
host=proxy_host,
|
|
port=proxy_port,
|
|
)
|
|
|
|
connect_line = f"CONNECT {self._server}:{self._port} HTTP/1.1\r\nHost: {self._server}:{self._port}\r\n\r\n"
|
|
self._writer.write(connect_line.encode())
|
|
await self._writer.drain()
|
|
|
|
response = await asyncio.wait_for(self._reader.readline(), timeout=self._connect_timeout)
|
|
response_str = response.decode("utf-8", errors="replace").strip()
|
|
|
|
if not response_str.startswith("HTTP/1.") or "200" not in response_str:
|
|
raise ConnectionError(f"Proxy CONNECT failed: {response_str}")
|
|
|
|
while True:
|
|
line = await asyncio.wait_for(self._reader.readline(), timeout=self._connect_timeout)
|
|
if line.strip() == b"":
|
|
break
|
|
|
|
if self._use_tls:
|
|
ctx = ssl.create_default_context()
|
|
ctx.check_hostname = True
|
|
ctx.verify_mode = ssl.CERT_REQUIRED
|
|
transport = self._writer.get_extra_info("socket")
|
|
if transport is None:
|
|
raise ConnectionError("Failed to get underlying socket for TLS upgrade")
|
|
self._writer.write(b"")
|
|
await self._writer.drain()
|
|
ssl_transport = ctx.wrap_socket(transport, server_hostname=self._server)
|
|
self._reader, self._writer = await asyncio.open_connection(
|
|
sock=ssl_transport,
|
|
)
|
|
|
|
logger.info(f"IRC TCP connected via proxy {proxy_host}:{proxy_port} -> {self._server}:{self._port}")
|
|
|
|
async def disconnect(self) -> None:
|
|
if self._keepalive_task and not self._keepalive_task.done():
|
|
self._keepalive_task.cancel()
|
|
try:
|
|
await self._keepalive_task
|
|
except asyncio.CancelledError:
|
|
pass
|
|
self._keepalive_task = None
|
|
|
|
if self._writer:
|
|
try:
|
|
self._writer.close()
|
|
except Exception:
|
|
pass
|
|
self._writer = None
|
|
self._reader = None
|
|
|
|
def send_line(self, line: str) -> None:
|
|
if self._writer is None:
|
|
raise ConnectionError("IRC writer not available")
|
|
|
|
encoded = line.encode("utf-8")
|
|
if len(encoded) > 510:
|
|
logger.warning(f"IRC message truncated from {len(encoded)} to 510 bytes: {line[:80]}...")
|
|
cut = 510
|
|
while cut > 0 and (encoded[cut - 1] & 0xC0) == 0x80:
|
|
cut -= 1
|
|
if cut == 0:
|
|
cut = max(1, 510)
|
|
encoded = encoded[:cut]
|
|
self._writer.write(encoded + b"\r\n")
|
|
|
|
async def drain(self) -> None:
|
|
if self._writer:
|
|
await self._writer.drain()
|
|
|
|
async def read_line(self) -> str | None:
|
|
if self._reader is None:
|
|
return None
|
|
line = await self._reader.readline()
|
|
if not line:
|
|
return None
|
|
return line.decode("utf-8", errors="replace").rstrip("\r\n")
|
|
|
|
def mark_pong(self) -> None:
|
|
self._last_pong = time.monotonic()
|
|
|
|
async def start_keepalive(self) -> None:
|
|
if self._keepalive_task and not self._keepalive_task.done():
|
|
return
|
|
self._keepalive_task = asyncio.create_task(self._keepalive_loop())
|
|
|
|
async def _keepalive_loop(self) -> None:
|
|
while True:
|
|
await asyncio.sleep(self.PING_TIMEOUT // 2)
|
|
try:
|
|
if self._writer:
|
|
self._writer.write(f"PING :{self._server}\r\n".encode())
|
|
await self._writer.drain()
|
|
except (ConnectionError, OSError) as e:
|
|
logger.warning(f"IRC keepalive PING failed: {e}")
|
|
elapsed = time.monotonic() - self._last_pong
|
|
if elapsed > self.PING_TIMEOUT:
|
|
logger.warning(f"IRC PONG timeout ({elapsed:.0f}s), triggering reconnect...")
|
|
asyncio.create_task(self._reconnect())
|
|
break
|
|
|
|
async def _reconnect(self) -> None:
|
|
delay = min(
|
|
self.RECONNECT_DELAY * (2**self._reconnect_attempts),
|
|
self.MAX_RECONNECT_DELAY,
|
|
)
|
|
self._reconnect_attempts += 1
|
|
logger.info(f"IRC reconnecting in {delay}s (attempt #{self._reconnect_attempts})")
|
|
await asyncio.sleep(delay)
|
|
|
|
await self.disconnect()
|
|
try:
|
|
await self.connect()
|
|
if self._on_reconnect:
|
|
self._on_reconnect()
|
|
except Exception as e:
|
|
logger.error(f"IRC reconnection failed: {e}")
|
|
|
|
|
|
async def resolve_srv_record(server: str) -> tuple[str, int] | None:
|
|
try:
|
|
import socket
|
|
|
|
loop = asyncio.get_running_loop()
|
|
answers = await loop.getaddrinfo(
|
|
f"_ircs._tcp.{server}",
|
|
None,
|
|
family=socket.AF_INET,
|
|
type=socket.SOCK_STREAM,
|
|
)
|
|
if answers:
|
|
addr = answers[0]
|
|
host = addr[4][0] if len(addr) > 3 else server
|
|
return host, 6697
|
|
except Exception:
|
|
pass
|
|
|
|
try:
|
|
import socket
|
|
|
|
loop = asyncio.get_running_loop()
|
|
answers = await loop.getaddrinfo(
|
|
f"_irc._tcp.{server}",
|
|
None,
|
|
family=socket.AF_INET,
|
|
type=socket.SOCK_STREAM,
|
|
)
|
|
if answers:
|
|
addr = answers[0]
|
|
host = addr[4][0] if len(addr) > 3 else server
|
|
return host, 6667
|
|
except Exception:
|
|
pass
|
|
|
|
return None
|