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