把服务端写坏会怎样

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

(每道题开头都有同一段:上面的内存版 UDP。)

贯穿全条的内存版 UDP(判题机不联网;接口和真 socket 一字不差,真机上把 net.socket() 换成 socket.socket(AF_INET, SOCK_DGRAM) 就是真的):

net = FakeNet()          一个内存版的网络:每个 (地址, 端口) 一个信箱
s = net.socket()         s.bind((host, port))  端口 0 由系统挑;已被占用报 Address already in use
s.sendto(data, addr)     超过 65507 字节报 Message too long;没 bind 就发会自动挑端口;对面没人听 → 悄悄丢(connect 过的 socket 下次 recv 报 Connection refused)
s.recvfrom(bufsize)      交回 (data, 来源地址);一次一个完整数据报;缓冲小了就截断;信箱空:设了超时报 timed out,没设就永远等(内存版抛 RuntimeError 提醒)
s.settimeout(秒) / s.connect(addr) + send / recv / s.close()
net.drop_next(k)         接下来 k 个包丢掉(发送方毫不知情)

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

回显服务端有个 bug:回信时把来源地址的端口写死成 9000。用同样的测试跑,看结果:

import collections

MAX_DATAGRAM = 65507


class FakeNet:
    """内存版的网络:每个 (地址, 端口) 一个信箱,信箱里是一个个完整的数据报。"""
    def __init__(self):
        self.boxes = {}
        self.next_port = 40000
        self.drops = 0
        self.sent = 0

    def socket(self):
        return FakeSock(self)

    def drop_next(self, k):
        self.drops = k


class FakeSock:
    def __init__(self, net):
        self.net = net
        self.addr = None
        self.timeout = None
        self.peer = None
        self.refused = False

    def bind(self, addr):
        host, port = addr
        if port == 0:
            port = self.net.next_port
            self.net.next_port += 1
        if (host, port) in self.net.boxes:
            raise OSError(98, "Address already in use")
        self.addr = (host, port)
        self.net.boxes[self.addr] = collections.deque()

    def getsockname(self):
        return self.addr

    def settimeout(self, t):
        self.timeout = t

    def sendto(self, data, addr):
        if len(data) > MAX_DATAGRAM:
            raise OSError(90, "Message too long")
        if self.addr is None:
            self.bind(("127.0.0.1", 0))          # 没 bind 就发:系统自动挑一个端口
        self.net.sent += 1
        if self.net.drops > 0:
            self.net.drops -= 1
            return len(data)                     # 丢了——发送方毫不知情
        if addr in self.net.boxes:
            self.net.boxes[addr].append((bytes(data), self.addr))
        else:
            self.refused = True                  # 没人听:Linux 会回一个 ICMP 不可达,只有 connect 过的 socket 才看得到
        return len(data)

    def recvfrom(self, bufsize):
        if self.refused and self.peer is not None:
            self.refused = False
            raise ConnectionRefusedError(111, "Connection refused")
        box = self.net.boxes.get(self.addr)
        if not box:
            if self.timeout is None:
                raise RuntimeError("信箱是空的又没设超时:真的 socket 会在这里永远等下去")
            raise TimeoutError("timed out")
        data, frm = box.popleft()
        return data[:bufsize], frm               # 缓冲区小了就截断,多出来的部分丢掉

    def connect(self, addr):
        self.peer = addr

    def send(self, data):
        return self.sendto(data, self.peer)

    def recv(self, bufsize):
        return self.recvfrom(bufsize)[0]

    def close(self):
        if self.addr in self.net.boxes:
            del self.net.boxes[self.addr]
        self.addr = None


def deliver(msgs, plan):
    """msgs: 按发送顺序的数据报列表;plan: 每个报一个字符:. 正常到达  x 丢掉  d 到两次  s 和后一个交换顺序。交回接收方看到的顺序。"""
    out = []
    i = 0
    while i < len(msgs):
        p = plan[i] if i < len(plan) else "."
        if p == "x":
            pass
        elif p == "d":
            out += [msgs[i], msgs[i]]
        elif p == "s" and i + 1 < len(msgs):
            out += [msgs[i + 1], msgs[i]]
            i += 1
        else:
            out.append(msgs[i])
        i += 1
    return out


def seq_gaps(seqs):
    """收到的序号里,从 1 到最大值之间缺了哪些。"""
    got = set(seqs)
    return [k for k in range(1, max(seqs) + 1) if k not in got]


def dedupe(seqs):
    seen = set()
    out = []
    for s in seqs:
        if s not in seen:
            seen.add(s)
            out.append(s)
    return out


def reorder(items):
    """items: (序号, 内容),按序号排好交回内容。"""
    return [x for _, x in sorted(items)]


def stop_and_wait(payloads, plan):
    """发送方每个报带序号,收到 ack 才发下一个;没 ack 就重发。plan 作用在「每一次发送」上(含重发)。
    交回 (接收方按序收到的内容, 一共发了几次)。"""
    sends = 0
    got = []
    k = 0
    for seq, p in enumerate(payloads, 1):
        while True:
            act = plan[k] if k < len(plan) else "."
            k += 1
            sends += 1
            if act == "x":
                continue                          # 丢了:等超时、重发
            if not got or got[-1][0] != seq:
                got.append((seq, p))              # 重复到达(重发的那份也到了)就按序号丢掉
            break
    return [p for _, p in got], sends


IP_HDR = 20
UDP_HDR = 8


def max_payload(mtu):
    """一个不分片的 UDP 数据报最多装多少字节的应用数据。"""
    return mtu - IP_HDR - UDP_HDR


def fragments(size, mtu):
    """应用数据 size 字节,走 MTU 为 mtu 的链路会被 IP 切成几片(每片的数据长度是 8 的倍数,最后一片除外)。"""
    total = size + UDP_HDR
    per = (mtu - IP_HDR) // 8 * 8
    return (total + per - 1) // per


def loss_prob(n, p):
    """n 片各以概率 p 丢,整个数据报丢的概率。"""
    return round(1 - (1 - p) ** n, 4)


def encode_name(name):
    """www.example.com → b"\x03www\x07example\x03com\x00" """
    out = b""
    for label in name.split("."):
        out += bytes([len(label)]) + label.encode()
    return out + b"\x00"


def build_query(qid, name):
    """最小的 DNS 查询:12 字节头 + 问题(名字 + 类型 A + 类 IN)。"""
    header = qid.to_bytes(2, "big") + b"\x01\x00" + b"\x00\x01" + b"\x00\x00" * 3
    return header + encode_name(name) + b"\x00\x01" + b"\x00\x01"


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)


def bad_echo(srv):
    data, frm = srv.recvfrom(65535)
    srv.sendto(data, ("127.0.0.1", 9000))


class TestBad(unittest.TestCase):
    def test_echo(self):
        net = FakeNet()
        srv = net.socket()
        srv.bind(("127.0.0.1", 0))
        c = net.socket()
        c.settimeout(0.1)
        c.sendto(b"hi", srv.getsockname())
        bad_echo(srv)
        self.assertEqual(c.recvfrom(65535)[0], b"hi")


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

                        
👩‍🏫
AI
💬 题目评论

全部评论