"""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()