把服务端写坏会怎样
(每道题开头都有同一段:上面的内存版 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))
全部评论