Files
mcp_screen/lib/websocket_client.py
T

185 lines
5.9 KiB
Python

# pyright: reportMissingImports=false, reportAttributeAccessIssue=false
"""Lightweight WebSocket client for MicroPython.
This module provides a minimal, robust, memory-efficient WebSocket client
supporting HTTP upgrade handshakes and client-to-server frame masking
(RFC 6455).
"""
import usocket as socket
import ustruct as struct
import urandom as random
import ubinascii as binascii
def parse_url(url):
if not url.startswith("ws://") and not url.startswith("wss://"):
raise ValueError("URL must start with ws:// or wss://")
is_ssl = url.startswith("wss://")
url_p = url.split("://", 1)[1]
parts = url_p.split("/", 1)
host_port = parts[0]
path = "/" + parts[1] if len(parts) > 1 else "/"
if ":" in host_port:
host, port = host_port.split(":", 1)
port = int(port)
else:
host = host_port
port = 443 if is_ssl else 80
return host, port, path, is_ssl
class WebSocketClient:
def __init__(self, url, headers=None, timeout=10):
self.url = url
self.headers = headers or {}
self.timeout = timeout
self.sock = None
self.host, self.port, self.path, self.is_ssl = parse_url(url)
def connect(self):
# 1. Generate standard Sec-WebSocket-Key
raw_key = bytes([random.getrandbits(8) for _ in range(16)])
sec_key = binascii.b2a_base64(raw_key).decode('utf-8').strip()
# 2. Resolve address and connect
addr = socket.getaddrinfo(self.host, self.port)[0][-1]
self.sock = socket.socket()
self.sock.settimeout(self.timeout)
self.sock.connect(addr)
if self.is_ssl:
import ussl
self.sock = ussl.wrap_socket(self.sock)
# 3. Construct HTTP GET upgrade request
req = [
f"GET {self.path} HTTP/1.1",
f"Host: {self.host}:{self.port}",
"Upgrade: websocket",
"Connection: Upgrade",
f"Sec-WebSocket-Key: {sec_key}",
"Sec-WebSocket-Version: 13",
]
for k, v in self.headers.items():
req.append(f"{k}: {v}")
req.append("\r\n")
self.sock.write("\r\n".join(req).encode('utf-8'))
# 4. Read HTTP response status and headers
status_line = self.sock.readline()
if not status_line or b"101" not in status_line:
self.close()
raise RuntimeError(f"Handshake failed status: {status_line.decode('utf-8', 'ignore').strip()}")
while True:
line = self.sock.readline()
if not line or line == b"\r\n":
break
def send_frame(self, opcode, payload, fin=True):
"""Send a masked WebSocket frame to the server (client-to-server MUST be masked)."""
if not self.sock:
raise RuntimeError("Not connected")
# Header byte 0: FIN and Opcode
b0 = 0x80 if fin else 0
b0 |= (opcode & 0x0F)
# Header byte 1: Mask bit (always 1 for client) and Payload length
payload_len = len(payload)
if payload_len <= 125:
header = struct.pack("!BB", b0, 0x80 | payload_len)
elif payload_len <= 65535:
header = struct.pack("!BBH", b0, 0x80 | 126, payload_len)
else:
header = struct.pack("!BBQ", b0, 0x80 | 127, payload_len)
# Generate 4-byte random masking key
mask = bytes([random.getrandbits(8) for _ in range(4)])
# Apply masking key (XOR payload)
masked_payload = bytearray(payload_len)
for i in range(payload_len):
masked_payload[i] = payload[i] ^ mask[i % 4]
# Send header, mask, and masked payload
self.sock.write(header)
self.sock.write(mask)
self.sock.write(masked_payload)
def send_text(self, text):
self.send_frame(0x1, text.encode('utf-8'))
def send_binary(self, data):
self.send_frame(0x2, data)
def recv_frame(self):
"""Receive an unmasked WebSocket frame from the server (server-to-client is unmasked)."""
if not self.sock:
raise RuntimeError("Not connected")
try:
header = self.sock.read(2)
except Exception:
# Handle socket timeout or disconnect
return None, None
if not header or len(header) < 2:
return None, None
b0, b1 = header
opcode = b0 & 0x0F
masked = bool(b1 & 0x80)
payload_len = b1 & 0x7F
if payload_len == 126:
len_bytes = self.sock.read(2)
if not len_bytes or len(len_bytes) < 2:
return None, None
payload_len = struct.unpack("!H", len_bytes)[0]
elif payload_len == 127:
len_bytes = self.sock.read(8)
if not len_bytes or len(len_bytes) < 8:
return None, None
payload_len = struct.unpack("!Q", len_bytes)[0]
if masked:
mask = self.sock.read(4)
if not mask or len(mask) < 4:
return None, None
# Read actual payload
payload = b""
while len(payload) < payload_len:
needed = payload_len - len(payload)
chunk = self.sock.read(needed)
if not chunk:
break
payload += chunk
if len(payload) < payload_len:
# Socket closed prematurely
return None, None
if masked:
# Unmask payload if masked
unmasked = bytearray(payload_len)
for i in range(payload_len):
unmasked[i] = payload[i] ^ mask[i % 4]
payload = bytes(unmasked)
return opcode, payload
def close(self):
if self.sock:
try:
# Send close frame
self.send_frame(0x8, b"")
except Exception:
pass
try:
self.sock.close()
except Exception:
pass
self.sock = None