一个文本帧收发
本节的 WebSocket 小工具:
accept_key(key) 握手:base64(sha1(key + 魔法串))——服务端要回的 Sec-WebSocket-Accept
encode_text(text, mask=None) / decode_text(frame) 一个最小文本帧:FIN+opcode 0x1、长度 < 126;mask 给了是客户端帧(要异或)客户端帧要带掩码,服务端帧不带。编一个带掩码的 "Hi",再解开看内容;服务端帧解一次对比:
import collections
class Net:
"""内存版的网络:端口 -> 监听者。判题机不联网,接口和真 socket 一样,真机上把这些换成 socket.socket() 就是真的。"""
def __init__(self):
self.listeners = {}
class Endpoint:
"""一条连接的一头:自己的收缓冲 rbuf,写就写进对面的 rbuf。"""
def __init__(self, name):
self.name = name
self.rbuf = b""
self.peer = None
self.peer_closed = False
self.closed = False
def send(self, data):
if self.closed or self.peer is None:
raise BrokenPipeError("连接已关")
self.peer.rbuf += bytes(data)
return len(data)
def sendall(self, data):
self.send(data)
def recv(self, bufsize):
if self.rbuf:
out, self.rbuf = self.rbuf[:bufsize], self.rbuf[bufsize:]
return out # 字节流:给你缓冲里现有的,最多 bufsize,不保证是「一条消息」
if self.peer_closed:
return b"" # 对面关了,读到头
raise BlockingIOError("暂时没有数据(真 socket 会在这里阻塞等)")
def close(self):
self.closed = True
if self.peer is not None:
self.peer.peer_closed = True
class Listener:
def __init__(self):
self.backlog = collections.deque()
def accept(self):
if not self.backlog:
raise BlockingIOError("暂时没有新连接(真 socket 会在 accept 阻塞)")
return self.backlog.popleft() # 交回服务端那一头的 Endpoint
def listen(net, port):
lis = Listener()
net.listeners[port] = lis
return lis
def connect(net, port):
"""建一对相连的端点:客户端这头交回给调用方,服务端那头塞进监听队列等 accept。"""
if port not in net.listeners:
raise ConnectionRefusedError(111, "Connection refused")
cli = Endpoint("client")
srv = Endpoint("server")
cli.peer = srv
srv.peer = cli
net.listeners[port].backlog.append(srv)
return cli
def frame(msg):
"""给一条消息加上 4 字节大端长度前缀。"""
return len(msg).to_bytes(4, "big") + msg
class Reader:
"""把「任意切碎、可能粘在一起」的字节,还原成一条条完整消息。喂多少无所谓,攒着,够一条吐一条。"""
def __init__(self):
self.buf = b""
def feed(self, chunk):
self.buf += chunk
def messages(self):
out = []
while len(self.buf) >= 4:
n = int.from_bytes(self.buf[:4], "big")
if len(self.buf) - 4 < n:
break # 半包:长度都不够,等下一块
out.append(self.buf[4:4 + n])
self.buf = self.buf[4 + n:] # 粘包:切走这一条,剩下的留着
return out
def recv_exactly(chunks, n):
"""chunks 是一串陆续到达的字节块;正好取 n 字节交回(凑不满就把有的都交回)。剩下的块和半块不管。"""
buf = b""
for c in chunks:
buf += c
if len(buf) >= n:
break
return buf[:n]
def split_lines(buf):
"""分隔符(换行)版拆包:交回 (完整的整行列表, 剩下的半行)。"""
parts = buf.split(b"\n")
return parts[:-1], parts[-1]
def req(*parts):
"""把一条请求编成一行字节:用空格连起来、末尾加换行。"""
return (" ".join(parts) + "\n").encode()
def decode(line):
"""把一行(不含换行)拆成 (命令, 参数列表)。"""
bits = line.decode().split(" ")
return bits[0], bits[1:]
def handle(store, line):
"""执行一条命令,交回要回给对方的一行字节。PUT k v -> OK;GET k -> 值 or NONE;DEL k -> OK/NONE;别的 -> ERR。"""
cmd, args = decode(line)
if cmd == "PUT" and len(args) == 2:
store[args[0]] = args[1]
return b"OK\n"
if cmd == "GET" and len(args) == 1:
return (store.get(args[0], "NONE") + "\n").encode()
if cmd == "DEL" and len(args) == 1:
return (b"OK\n" if store.pop(args[0], None) is not None else b"NONE\n")
return b"ERR\n"
import base64
import hashlib
GUID = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11"
def accept_key(key):
"""WebSocket 握手:把客户端给的 Sec-WebSocket-Key 拼上魔法串、SHA-1、Base64——这就是服务端要回的 Sec-WebSocket-Accept。"""
return base64.b64encode(hashlib.sha1((key + GUID).encode()).digest()).decode()
def encode_text(text, mask=None):
"""一个最小的文本帧:FIN=1、opcode=0x1,长度 < 126。mask 给了就是客户端帧(要异或),不给是服务端帧。"""
body = text.encode()
out = bytes([0x81])
if mask is None:
out += bytes([len(body)]) + body
else:
out += bytes([0x80 | len(body)]) + bytes(mask)
out += bytes(b ^ mask[i % 4] for i, b in enumerate(body))
return out
def decode_text(frame):
"""解一个最小文本帧,交回里面的文字。"""
length = frame[1] & 0x7F
if frame[1] & 0x80:
mask = frame[2:6]
body = frame[6:6 + length]
return bytes(b ^ mask[i % 4] for i, b in enumerate(body)).decode()
return frame[2:2 + length].decode()
import io
import unittest
def run(*cases):
suite = unittest.TestSuite()
for c in cases:
suite.addTests(unittest.defaultTestLoader.loadTestsFromTestCase(c))
r = unittest.TextTestRunner(stream=io.StringIO()).run(suite)
return r.testsRun, len(r.failures), len(r.errors)
f = encode_text("Hi", mask=[1, 2, 3, 4])
svr = encode_text("Ok")
print(f.hex() + "/" + decode_text(f) + "/" + decode_text(svr))
全部评论