# 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