新增IRC协议相关的全套工具模块,包括: - 核心协议解析与CTCP处理 - 消息发送缓存与文本 sanitize - 账号配置管理与运行时状态 - 命令处理与权限控制 - 服务发现与诊断工具 - 多账号网关与配置加载
228 lines
7.8 KiB
Python
228 lines
7.8 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)
|
|
loop = asyncio.get_running_loop()
|
|
ssl_reader = asyncio.StreamReader()
|
|
ssl_protocol = asyncio.StreamReaderProtocol(ssl_reader)
|
|
await loop.create_connection(
|
|
lambda: ssl_protocol, sock=ssl_transport, ssl=ctx, server_hostname=self._server
|
|
)
|
|
self._reader = ssl_reader
|
|
self._writer = self._writer # Keep the proxy writer for sending
|
|
|
|
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]}...")
|
|
encoded = encoded[:510]
|
|
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
|