ForcePilot/backend/package/yuxi/channels/adapters/irc/connection.py
Kris 2dd1a075f2 feat(irc): 实现完整的IRC适配器基础功能模块
新增IRC协议相关的全套工具模块,包括:
- 核心协议解析与CTCP处理
- 消息发送缓存与文本 sanitize
- 账号配置管理与运行时状态
- 命令处理与权限控制
- 服务发现与诊断工具
- 多账号网关与配置加载
2026-05-12 00:45:10 +08:00

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