测出粘包 bug

👁️ 1 人浏览 💬 0 人评论 ❤️ 添加收藏

测试用 run(*cases):静默跑 unittest,交回 (跑, 挂, 错) 三个数。

本节的切包小工具:

frame(msg)        4 字节大端长度前缀 + 内容
Reader.feed(chunk) / Reader.messages()   喂进任意切碎、可能粘在一起的字节,吐出完整消息、剩下的攒着(半包 + 粘包)
recv_exactly(chunks, n)   从陆续到达的块里正好取 n 字节
split_lines(buf)          分隔符(换行)版:交回 (整行列表, 剩下的半行)

有人图省事,用「一次 recv 当一条消息」来收。测试连发两条、模拟一次 recv 收到粘在一起的字节,看这个坏收法被测出来没有:

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)

import unittest


def bad_recv_one(stream):
    # 坏收法:把一次收到的整段当成「一条消息」
    return [stream]


class T(unittest.TestCase):
    def test_two_messages(self):
        stream = frame(b"a") + frame(b"b")
        msgs = bad_recv_one(stream)
        self.assertEqual(len(msgs), 2)


print("/".join(str(x) for x in run(T)))
提交你的答案
请登录后提交答案。
去登录
代码编辑器
Ctrl + Enter 运行
本次输入:
输出:

                        
👩‍🏫
AI
💬 题目评论

全部评论