185 lines
5.9 KiB
Python
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
|