Files
vpn-loopback-firewall/tests/test_httpmsg.py
T

92 lines
2.9 KiB
Python

"""httpmsg 单测: 请求解析 / hop-by-hop 过滤。"""
from __future__ import annotations
import unittest
from channel.httpmsg import parse_request, parse_request_head
from channel.server import _denied, _error_resp, _upstream_addr
def _get(host: str = "example.com", path: str = "/a?b=1") -> bytes:
return (
f"GET {path} HTTP/1.1\r\n"
f"Host: {host}\r\n"
"User-Agent: curl/8.0\r\n"
"Accept: */*\r\n"
"Connection: keep-alive\r\n"
"\r\n"
).encode("ascii")
def _post(body: bytes) -> bytes:
return (
b"POST /api HTTP/1.1\r\n"
b"Host: example.com\r\n"
b"Content-Type: application/json\r\n"
+ f"Content-Length: {len(body)}\r\n".encode("ascii")
+ b"\r\n"
+ body
)
class HttpMsgTest(unittest.TestCase):
def test_parse_get(self) -> None:
req = parse_request(_get())
assert req is not None
self.assertEqual(req.method, "GET")
self.assertEqual(req.path, "/a?b=1")
self.assertEqual(req.host, "example.com")
self.assertEqual(req.body, b"")
self.assertEqual(req.header("accept"), "*/*")
def test_parse_post_body(self) -> None:
req = parse_request(_post(b'{"x":1}'))
assert req is not None
self.assertEqual(req.method, "POST")
self.assertEqual(req.body, b'{"x":1}')
def test_partial_body_returns_none(self) -> None:
data = _post(b"12345")
self.assertIsNone(parse_request(data[:-3]))
def test_head_only_when_body_incomplete(self) -> None:
data = _post(b"12345")
parsed = parse_request_head(data)
assert parsed is not None
req, head_len = parsed
self.assertEqual(req.body, b"")
self.assertGreater(len(data), head_len)
def test_hop_by_hop_filtered(self) -> None:
req = parse_request(_get())
assert req is not None
clean = dict(req.without_hop_by_hop())
self.assertNotIn("Connection", clean)
self.assertNotIn("connection", clean)
self.assertNotIn("Host", clean) # host 单独处理
self.assertIn("Accept", clean)
def test_malformed_rejected(self) -> None:
self.assertIsNone(parse_request(b"NOT-HTTP\r\n\r\n"))
self.assertIsNone(parse_request(b""))
class ServerHelperTest(unittest.TestCase):
def test_denied_suffix(self) -> None:
self.assertTrue(_denied("evil.com", ["evil.com"]))
self.assertTrue(_denied("sub.evil.com", [".evil.com"]))
self.assertFalse(_denied("evil.com.cn", [".evil.com"]))
def test_error_resp_shape(self) -> None:
resp = _error_resp(502, "boom")
self.assertTrue(resp.startswith(b"HTTP/1.1 502"))
self.assertIn(b"Content-Length:", resp)
def test_upstream_addr(self) -> None:
self.assertEqual(_upstream_addr("127.0.0.1:5353"), ("127.0.0.1", 5353))
self.assertEqual(_upstream_addr("223.5.5.5"), ("223.5.5.5", 53))
if __name__ == "__main__":
unittest.main()