115 lines
3.9 KiB
Python
115 lines
3.9 KiB
Python
|
|
import asyncio
|
||
|
|
import logging
|
||
|
|
import zlib
|
||
|
|
from collections.abc import Callable, Awaitable
|
||
|
|
|
||
|
|
from yuxi.channel.extensions.minecraft.protocol import (
|
||
|
|
read_varint,
|
||
|
|
write_varint,
|
||
|
|
write_packet_frame,
|
||
|
|
)
|
||
|
|
from yuxi.channel.extensions.minecraft.types import ConnectionState, McPacket
|
||
|
|
|
||
|
|
logger = logging.getLogger(__name__)
|
||
|
|
|
||
|
|
|
||
|
|
class McClient:
|
||
|
|
def __init__(self, host: str, port: int):
|
||
|
|
self.host = host
|
||
|
|
self.port = port
|
||
|
|
self.reader: asyncio.StreamReader | None = None
|
||
|
|
self.writer: asyncio.StreamWriter | None = None
|
||
|
|
self.state: ConnectionState = ConnectionState.DISCONNECTED
|
||
|
|
self.compression_threshold: int = -1
|
||
|
|
self._running = False
|
||
|
|
self._on_packet: Callable[[McPacket], Awaitable[None]] | None = None
|
||
|
|
self._send_lock = asyncio.Lock()
|
||
|
|
|
||
|
|
@property
|
||
|
|
def is_connected(self) -> bool:
|
||
|
|
return self.writer is not None and not self.writer.is_closing()
|
||
|
|
|
||
|
|
async def connect(self) -> None:
|
||
|
|
self.reader, self.writer = await asyncio.open_connection(self.host, self.port)
|
||
|
|
logger.info("TCP connected to %s:%d", self.host, self.port)
|
||
|
|
self.state = ConnectionState.HANDSHAKE
|
||
|
|
|
||
|
|
async def disconnect(self) -> None:
|
||
|
|
self._running = False
|
||
|
|
if self.writer:
|
||
|
|
self.writer.close()
|
||
|
|
try:
|
||
|
|
await self.writer.wait_closed()
|
||
|
|
except Exception:
|
||
|
|
pass
|
||
|
|
self.writer = None
|
||
|
|
self.reader = None
|
||
|
|
self.state = ConnectionState.DISCONNECTED
|
||
|
|
|
||
|
|
def set_packet_handler(self, handler: Callable[[McPacket], Awaitable[None]]) -> None:
|
||
|
|
self._on_packet = handler
|
||
|
|
|
||
|
|
async def send_raw(self, data: bytes) -> None:
|
||
|
|
if not self.writer:
|
||
|
|
raise ConnectionError("Not connected")
|
||
|
|
self.writer.write(data)
|
||
|
|
await self.writer.drain()
|
||
|
|
|
||
|
|
async def send_packet(self, packet_id: int, data: bytes = b"") -> None:
|
||
|
|
frame = write_packet_frame(packet_id, data)
|
||
|
|
|
||
|
|
if self.compression_threshold >= 0:
|
||
|
|
uncompressed_length = len(frame)
|
||
|
|
if uncompressed_length >= self.compression_threshold:
|
||
|
|
compressed_data = zlib.compress(frame)
|
||
|
|
frame = write_varint(uncompressed_length) + compressed_data
|
||
|
|
else:
|
||
|
|
frame = write_varint(0) + frame
|
||
|
|
|
||
|
|
async with self._send_lock:
|
||
|
|
await self.send_raw(frame)
|
||
|
|
|
||
|
|
async def recv_packet(self) -> McPacket:
|
||
|
|
if not self.reader:
|
||
|
|
raise ConnectionError("Not connected")
|
||
|
|
|
||
|
|
raw_length = bytearray()
|
||
|
|
while True:
|
||
|
|
byte = await self.reader.readexactly(1)
|
||
|
|
raw_length.append(byte[0])
|
||
|
|
if not (byte[0] & 0x80):
|
||
|
|
break
|
||
|
|
if len(raw_length) > 5:
|
||
|
|
raise ValueError("Packet length VarInt too large")
|
||
|
|
|
||
|
|
packet_length, _ = read_varint(bytes(raw_length))
|
||
|
|
packet_data = await self.reader.readexactly(packet_length)
|
||
|
|
|
||
|
|
if self.compression_threshold >= 0:
|
||
|
|
data_length_val, d_len = read_varint(packet_data)
|
||
|
|
if data_length_val == 0:
|
||
|
|
payload = packet_data[d_len:]
|
||
|
|
else:
|
||
|
|
payload = zlib.decompress(packet_data[d_len:])
|
||
|
|
else:
|
||
|
|
payload = packet_data
|
||
|
|
|
||
|
|
packet_id, id_len = read_varint(payload)
|
||
|
|
return McPacket(packet_id=packet_id, data=payload[id_len:])
|
||
|
|
|
||
|
|
async def run_recv_loop(self) -> None:
|
||
|
|
self._running = True
|
||
|
|
while self._running and self.reader:
|
||
|
|
try:
|
||
|
|
packet = await self.recv_packet()
|
||
|
|
if self._on_packet:
|
||
|
|
await self._on_packet(packet)
|
||
|
|
except asyncio.IncompleteReadError:
|
||
|
|
logger.warning("MC connection closed by server")
|
||
|
|
break
|
||
|
|
except ConnectionError:
|
||
|
|
break
|
||
|
|
except Exception:
|
||
|
|
logger.exception("Error in MC recv loop")
|
||
|
|
break
|