diff --git a/openwrt/Makefile b/openwrt/Makefile index d8606c7a..650fb800 100644 --- a/openwrt/Makefile +++ b/openwrt/Makefile @@ -31,6 +31,7 @@ define Package/ua2f +(PACKAGE_nftables-json||PACKAGE_nftables-nojson):kmod-nft-tproxy \ +PACKAGE_firewall:iptables-mod-conntrack-extra \ +PACKAGE_firewall:iptables-mod-filter \ + +PACKAGE_firewall:iptables-mod-u32 \ +PACKAGE_firewall:iptables-mod-nfqueue \ +PACKAGE_firewall:iptables-mod-tproxy endef @@ -91,6 +92,8 @@ define Package/ua2f/install $(INSTALL_DIR) $(1)/etc/config $(1)/etc/init.d $(INSTALL_BIN) ./files/ua2f.config $(1)/etc/config/ua2f $(INSTALL_BIN) ./files/ua2f.init $(1)/etc/init.d/ua2f + $(INSTALL_DIR) $(1)/usr/share/ua2f + $(INSTALL_DATA) ./files/ua2f.firewall $(1)/usr/share/ua2f/firewall.sh $(if $(UA2F_COVERAGE_BUILD),rm -rf $(UA2F_COVERAGE_DIR)) $(if $(UA2F_COVERAGE_BUILD),$(INSTALL_DIR) $(UA2F_COVERAGE_DIR)/objects/CMakeFiles/ua2f.dir) diff --git a/openwrt/files/ua2f.firewall b/openwrt/files/ua2f.firewall new file mode 100644 index 00000000..0b48a40d --- /dev/null +++ b/openwrt/files/ua2f.firewall @@ -0,0 +1,109 @@ +#!/bin/sh +# Pure rule generators: sourcing this file does not read or change the firewall. +# Experimental empty-ACK bypass. Keep SYN/FIN/RST queued, never change a mark, +# and leave IP options, fragments and IPv6 extension headers on the slow path. +# Doff is in 32-bit words. Enumerating it avoids dynamic-length arithmetic and +# nft >= 1.0.3 integer-concatenation requirements on older OpenWrt installations. + +# Print one complete u32 expression for each legal TCP header length. +# The caller must select the IP family and TCP ORIGINAL direction externally. +ua2f_empty_ack_u32() { + case "${1:-}" in + 4|6) ;; + *) return 1 ;; + esac + for ua2f_doff in 5 6 7 8 9 10 11 12 13 14 15; do + if [ "$1" = 4 ]; then + printf '%s\n' "0&0xFFFF=$((20 + 4 * ua2f_doff)) && 0>>24&0xF=5 && 4&0x3FFF=0 && 6&0xFF=6 && 32>>28=$ua2f_doff && 33>>24&0x17=0x10 && $((16 + 4 * ua2f_doff))&0=0" + else + printf '%s\n' "4>>16=$((4 * ua2f_doff)) && 4>>8&0xFF=6 && 52>>28=$ua2f_doff && 53>>24&0x17=0x10 && $((36 + 4 * ua2f_doff))&0=0" + fi + done +} + +# Print rules for the existing inet postrouting base chain, immediately before +# its queue statement. Exact skb length also excludes truncated/GRO oddities. +ua2f_empty_ack_nft() { + for ua2f_doff in 5 6 7 8 9 10 11 12 13 14 15; do + printf '%s\n' "meta length $((20 + 4 * ua2f_doff)) ct direction original ip protocol tcp ip hdrlength 5 ip frag-off & 0x3fff == 0 tcp flags & 0x17 == 0x10 tcp doff $ua2f_doff ip length $((20 + 4 * ua2f_doff)) counter return comment \"!ua2f: empty ACK\";" + printf '%s\n' "meta length $((40 + 4 * ua2f_doff)) ct direction original ip6 nexthdr tcp tcp flags & 0x17 == 0x10 tcp doff $ua2f_doff ip6 length $((4 * ua2f_doff)) counter return comment \"!ua2f: empty ACK\";" + done +} + +# Explicit command prefix, e.g. 4 iptables -t mangle -A ua2f. Never uses eval. +# No calls occur unless the init script's opt-in config enables this function. +ua2f_add_empty_ack_iptables() { + [ "$#" -ge 2 ] || return 1 + ua2f_family="$1" + shift + ua2f_rules="$(ua2f_empty_ack_u32 "$ua2f_family")" || return 1 + while IFS= read -r ua2f_expr; do + "$@" -p tcp -m conntrack --ctdir ORIGINAL -m u32 --u32 "$ua2f_expr" -j RETURN || return 1 + done <>16=20:60' ;; + *) return 1 ;; + esac +} + +# Pure complete candidate-tail generator: one rule per line, TAB-separated argv. +# Consumers must split on TAB, not whitespace, so each u32 expression stays one +# argument. No shell evaluation is needed. Family, first queue and last queue +# are explicit; equal endpoints select --queue-num, otherwise --queue-balance. +# Install AFTER all existing bypass/mark rules. The only skipped work is failed, +# side-effect-free empty-ACK matching; neither queue changes the connection mark. +ua2f_empty_ack_queue_iptables() { + [ "$#" -eq 3 ] || return 1 + ua2f_length="$(ua2f_empty_ack_length_u32 "$1")" || return 1 + case "$2" in ""|*[!0-9]*) return 1 ;; esac + case "$3" in ""|*[!0-9]*) return 1 ;; esac + [ "$2" -le 65535 ] && [ "$3" -le 65535 ] && [ "$2" -le "$3" ] || return 1 + if [ "$2" -eq "$3" ]; then + ua2f_queue="$(printf '%s\t%s\t%s\t%s\t%s' -j NFQUEUE --queue-num "$2" --queue-bypass)" + else + ua2f_queue="$(printf '%s\t%s\t%s\t%s\t%s' -j NFQUEUE --queue-balance "$2:$3" --queue-bypass)" + fi + printf '%s\t%s\t%s\t%s\t%s\t%s\n' -m u32 '!' --u32 "$ua2f_length" "$ua2f_queue" + ua2f_rules="$(ua2f_empty_ack_u32 "$1")" || return 1 + while IFS= read -r ua2f_expr; do + printf '%s\t%s\t%s\t%s\t%s\t%s\t%s\t%s\t%s\t%s\t%s\n' \ + -p tcp -m conntrack --ctdir ORIGINAL -m u32 --u32 "$ua2f_expr" "$(printf '%s\t%s' -j RETURN)" + done <&2 + if [ "$nfqueue_workers" -gt 1 ]; then + $IPT_M -A ua2f -j NFQUEUE --queue-balance "10010:$nfqueue_end" --queue-bypass + else + $IPT_M -A ua2f -j NFQUEUE --queue-num 10010 --queue-bypass + fi fi fi $IPT_M -A POSTROUTING -p tcp -m conntrack --ctdir ORIGINAL -j ua2f @@ -484,10 +494,13 @@ setup_firewall() { [ "$disable_connmark" -eq "1" ] || $IPT6_M -A ua2f -p tcp --dport 80 -j CONNMARK --set-mark 44 [ "$disable_connmark" -eq "1" ] || $IPT6_M -A ua2f -m connmark --mark 43 -j RETURN # 不处理标记为非 http 的流 [ "$handle_mmtls" -eq "1" ] || $IPT6_M -A ua2f -p tcp --dport 80 -m string --string "/mmtls/" --algo bm -j RETURN # 不处理微信的mmtls - if [ "$nfqueue_workers" -gt 1 ]; then - $IPT6_M -A ua2f -j NFQUEUE --queue-balance "10010:$nfqueue_end" --queue-bypass - else - $IPT6_M -A ua2f -j NFQUEUE --queue-num 10010 --queue-bypass + if [ "$bypass_empty_ack" -eq "0" ] || ! ua2f_add_empty_ack_queue_iptables 6 10010 "$nfqueue_end" ip6tables -t mangle -A ua2f; then + [ "$bypass_empty_ack" -eq "0" ] || echo "UA2F: empty-ACK IPv6 optimization unavailable; keeping NFQUEUE fallback" >&2 + if [ "$nfqueue_workers" -gt 1 ]; then + $IPT6_M -A ua2f -j NFQUEUE --queue-balance "10010:$nfqueue_end" --queue-bypass + else + $IPT6_M -A ua2f -j NFQUEUE --queue-num 10010 --queue-bypass + fi fi fi $IPT6_M -A POSTROUTING -p tcp -m conntrack --ctdir ORIGINAL -j ua2f diff --git a/scripts/test_empty_ack_netns.py b/scripts/test_empty_ack_netns.py new file mode 100644 index 00000000..b8d17709 --- /dev/null +++ b/scripts/test_empty_ack_netns.py @@ -0,0 +1,208 @@ +#!/usr/bin/env python3 +"""Opt-in, isolated packet-path checks for the exact empty-ACK rule generators. + +Run only inside a disposable outer netns. This does not alter host routes, +sysctls, or security settings. IPv4/IPv6 and nft/iptables share one HTTP probe. +""" +from __future__ import annotations + +import argparse +import json +import os +from pathlib import Path +import socket +import subprocess +import sys +import time + +import benchmark as bench + + +def nft_family_matches(rule: dict, family: int) -> bool: + """Identify generated predicates by their explicit IP-family header read.""" + protocol = "ip" if family == 4 else "ip6" + def includes_header(value): + if isinstance(value, dict): + if value.get("payload", {}).get("protocol") == protocol: + return True + return any(includes_header(child) for child in value.values()) + return isinstance(value, list) and any(includes_header(child) for child in value) + return includes_header(rule.get("expr", [])) + + +def nft_ruleset(interface: str, rules: str) -> str: + # nft requires a newline/semicolon between chain declarations and closing + # blocks. Keep declarations on separate lines, including the final braces. + return ("table inet ua_empty_test {\n" + " chain prerouting {\n" + " type filter hook prerouting priority mangle; policy accept;\n" + f' iifname "{interface}" tcp dport 18080 ct direction original jump inspect;\n' + " }\n" + " chain inspect {\n" + rules.rstrip() + "\n" + " counter queue num 10010;\n" + " }\n" + "}\n") + + +def probe(host: str, port: int) -> None: + """Split a header, pipeline requests, preserve a POST body, then half-close.""" + body = b"User-Agent: preserve this body\r\n" * 2048 + cases = [("GET", f"/first?padding=65536", "split-user-agent", b"")] + cases += [("GET", f"/pipeline-{i}?padding=1024", f"pipeline-{i}", b"") for i in range(16)] + cases += [("POST", "/body", "post-user-agent", body), + ("GET", "/after-post?padding=65536", "after-post", b""), + ("GET", "/without-ua", None, b"")] + encoded = [] + for method, path, agent, payload in cases: + headers = f"{method} {path} HTTP/1.1\r\nHost: test\r\nConnection: keep-alive\r\nContent-Length: {len(payload)}\r\n" + if agent is not None: + headers += f"User-Agent: {agent}\r\n" + encoded.append(headers.encode() + b"\r\n" + payload) + with socket.create_connection((host, port), timeout=10) as connection: + split = encoded[0].index(b"split-user-agent") + 6 + connection.sendall(encoded[0][:split]) + time.sleep(.02) + connection.sendall(encoded[0][split:]) + stream = connection.makefile("rb") + for index, (method, path, agent, payload) in enumerate(cases): + status = stream.readline() + if not status.startswith(b"HTTP/1.1 200 "): + raise AssertionError(f"unexpected response: {status!r}") + headers = {} + while True: + line = stream.readline() + if line == b"\r\n": + break + if not line: + raise AssertionError("truncated response headers") + key, value = line.split(b":", 1) + headers[key.lower()] = value.strip() + size = int(headers[b"content-length"]) + raw = stream.read(size) + if len(raw) != size: + raise AssertionError("truncated response body") + response = json.loads(raw) + expected = None if agent is None else "F" * len(agent) + if response["user_agent"] != expected or response["body"] != payload.decode(): + raise AssertionError(f"request rewrite/body mismatch: {path}") + if response["method"] != method or response["path"] != path: + raise AssertionError("request framing changed") + expected_padding = 65536 if "65536" in path else (1024 if "padding=1024" in path else 0) + if response["padding"] != "x" * expected_padding: + raise AssertionError("response content changed") + if index == 0: + # Let response ACKs and classification settle before any later + # request exists; then pipeline sixteen requests and a POST. + time.sleep(.1) + connection.sendall(b"".join(encoded[1:18])) + elif index == 17: + # A truly subsequent keep-alive request after the complete POST + # response, rather than one already queued in the first write. + time.sleep(.1) + connection.sendall(b"".join(encoded[18:])) + connection.shutdown(socket.SHUT_WR) + if stream.read(1): + raise AssertionError("unexpected bytes after final response") + print(json.dumps({"ok": True, "host": host, "requests": len(cases), + "split_header": True, "pipeline_requests": 16, + "post_bytes": len(body), "later_keep_alive": True, "half_close": True})) + + +def run(binary: Path, helper: Path, output: Path) -> None: + bench.require_root() + bench.require_commands(["ip", "iptables", "ip6tables", "nft"]) + if os.readlink("/proc/self/ns/net") == os.readlink("/proc/1/ns/net"): + raise SystemExit("Use a disposable outer network namespace") + output.mkdir(parents=True, exist_ok=True) + suffix = str(os.getpid())[-6:] + ns = bench.Netns(f"uae-{suffix}", f"uae{suffix}h", f"uae{suffix}c", "10.250.0.1", "10.250.0.2", 24) + endpoint = Path(__file__).with_name("http_test_endpoint.py") + server = target = None + results = [] + try: + bench.setup_netns(ns) + bench.run_cmd(["ip", "-6", "addr", "add", "fd42:250::1/64", "dev", ns.host_if, "nodad"]) + bench.run_cmd(["ip", "netns", "exec", ns.name, "ip", "-6", "addr", "add", "fd42:250::2/64", "dev", ns.ns_if, "nodad"]) + with (output / "origin.log").open("wb") as log: + server = bench.ProcessHandle(subprocess.Popen( + [sys.executable, str(endpoint), "serve", "--port", "18080", "--ipv6"], + stdout=log, stderr=subprocess.STDOUT, start_new_session=True), output / "origin.log") + bench.wait_for_port("127.0.0.1", 18080, 5) + env = {key: value for key, value in os.environ.items() if not key.startswith("UA2F_")} + env.update(UA2F_NFQUEUE_WORKERS="1", UA2F_PROXY_WORKERS="1") + log_path = output / "ua2f.log" + with log_path.open("wb") as log: + target = bench.ProcessHandle(subprocess.Popen( + [str(binary), "--mode", "NFQUEUE", "--listen-port", "10010"], + cwd=binary.parent, env=env, stdout=log, stderr=subprocess.STDOUT, + start_new_session=True), log_path) + time.sleep(.5) + if target.proc.poll() is not None: + raise RuntimeError(f"UA2F exited early: {target.proc.returncode}") + for backend in ("iptables", "nft"): + if backend == "iptables": + for family, executable in ((4, "iptables"), (6, "ip6tables")): + bench.run_cmd([executable, "-t", "mangle", "-N", "UA_EMPTY_TEST"]) + bench.run_cmd([executable, "-t", "mangle", "-A", "PREROUTING", "-i", ns.host_if, + "-p", "tcp", "--dport", "18080", "-m", "conntrack", "--ctdir", "ORIGINAL", "-j", "UA_EMPTY_TEST"]) + rules = bench.run_cmd(["sh", "-c", '. "$1"; ua2f_empty_ack_queue_iptables "$2" 10010 10010', + "sh", str(helper), str(family)]).stdout.splitlines() + if len(rules) != 13: + raise AssertionError("candidate did not emit its thirteen-rule tail") + for rule in rules: + bench.run_cmd([executable, "-t", "mangle", "-A", "UA_EMPTY_TEST", *rule.split("\t")]) + else: + rules = bench.run_cmd(["sh", "-c", '. "$1"; ua2f_empty_ack_nft', "sh", str(helper)]).stdout + ruleset = nft_ruleset(ns.host_if, rules) + (output / "candidate.nft").write_text(ruleset) + subprocess.run(["nft", "-f", "-"], input=ruleset, text=True, check=True) + for family, host in ((4, ns.server_ip), (6, "fd42:250::1")): + result = bench.run_cmd(["ip", "netns", "exec", ns.name, sys.executable, + str(Path(__file__).resolve()), "--probe", host]) + results.append(dict(json.loads(result.stdout), backend=backend, family=family)) + # Preserve successful probes even if a later backend fails. + (output / "probes.json").write_text(json.dumps(results, indent=2)) + print(json.dumps(results[-1]), flush=True) + if backend == "iptables": + for executable in ("iptables", "ip6tables"): + snapshot = bench.run_cmd([executable, "-t", "mangle", "-L", "UA_EMPTY_TEST", "-nvx"]) + (output / f"{executable}-counters.txt").write_text(snapshot.stdout) + matches = [int(line.split()[0]) for line in snapshot.stdout.splitlines() + if len(line.split()) > 2 and line.split()[2] == "RETURN"] + if sum(matches) == 0: + raise AssertionError(f"{executable}: candidate rules never matched an ACK") + bench.run_cmd([executable, "-t", "mangle", "-D", "PREROUTING", "-i", ns.host_if, + "-p", "tcp", "--dport", "18080", "-m", "conntrack", "--ctdir", "ORIGINAL", "-j", "UA_EMPTY_TEST"]) + bench.run_cmd([executable, "-t", "mangle", "-F", "UA_EMPTY_TEST"]) + bench.run_cmd([executable, "-t", "mangle", "-X", "UA_EMPTY_TEST"]) + else: + snapshot = bench.run_cmd(["nft", "-j", "list", "table", "inet", "ua_empty_test"]) + (output / "nft-counters.json").write_text(snapshot.stdout) + rules = [item["rule"] for item in json.loads(snapshot.stdout)["nftables"] if "rule" in item] + for family in (4, 6): + matches = [expr["counter"]["packets"] for rule in rules + if rule.get("comment") == "!ua2f: empty ACK" and nft_family_matches(rule, family) + for expr in rule["expr"] if "counter" in expr] + if sum(matches) == 0: + raise AssertionError(f"nft IPv{family}: candidate rules never matched an ACK") + bench.run_cmd(["nft", "delete", "table", "inet", "ua_empty_test"]) + (output / "summary.json").write_text(json.dumps({"status": "passed", "results": results}, indent=2)) + finally: + bench.stop_process(target) + bench.stop_process(server) + bench.cleanup_netns(ns) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--binary", type=Path) + parser.add_argument("--helper", type=Path) + parser.add_argument("--output-dir", type=Path) + parser.add_argument("--probe") + arguments = parser.parse_args() + if arguments.probe: + probe(arguments.probe, 18080) + elif arguments.binary and arguments.helper and arguments.output_dir: + run(arguments.binary.resolve(), arguments.helper.resolve(), arguments.output_dir.resolve()) + else: + parser.error("specify --binary/--helper/--output-dir or --probe") diff --git a/scripts/test_firewall_rules.py b/scripts/test_firewall_rules.py new file mode 100644 index 00000000..fd76f5d3 --- /dev/null +++ b/scripts/test_firewall_rules.py @@ -0,0 +1,315 @@ +#!/usr/bin/env python3 +"""Offline semantic tests. Never calls nft/iptables or creates a network namespace.""" +from functools import lru_cache +import itertools +import os +from pathlib import Path +import re +import struct +import subprocess +import unittest + +ROOT = Path(__file__).resolve().parents[1] +HELPER = ROOT / "openwrt/files/ua2f.firewall" + + +def generate(function, *args): + return subprocess.check_output( + ["sh", "-c", '. "$1"; shift; "$@"', "sh", str(HELPER), function, *map(str, args)], + text=True, + ).splitlines() + + +@lru_cache(maxsize=None) +def parse_u32(expression): + """Compile only the u32 operations emitted by this helper, including ranges.""" + tests = [] + for test in expression.split("&&"): + location, expected = test.strip().split("=") + tokens = re.findall(r"0x[0-9a-fA-F]+|[0-9]+|>>|&", location) + bounds = [int(value.strip(), 0) for value in expected.split(":")] + tests.append((int(tokens[0], 0), + tuple((op, int(raw, 0)) for op, raw in zip(tokens[1::2], tokens[2::2])), + bounds[0], bounds[-1])) + return tuple(tests) + + +def u32_matches(expression, packet): + """Independent interpreter; a failed packet read is a non-match.""" + for offset, operations, low, high in parse_u32(expression): + if offset + 4 > len(packet): + return False + value = int.from_bytes(packet[offset:offset + 4], "big") + for op, operand in operations: + value = value >> operand if op == ">>" else value & operand + if not low <= value <= high: + return False + return True + + +def packet(family, doff=5, payload=b"", flags=0x10, ihl=5, frag=0, nexthdr=6): + tcp = bytearray(max(20, 4 * doff)) + tcp[12] = doff << 4 + tcp[13] = flags + if family == 4: + ip = bytearray(max(20, ihl * 4)) + ip[0] = 0x40 | ihl + ip[9] = nexthdr + struct.pack_into("!HH", ip, 2, len(ip) + len(tcp) + len(payload), 0) + struct.pack_into("!H", ip, 6, frag) + else: + ip = bytearray(40) + ip[0] = 0x60 + ip[6] = nexthdr + struct.pack_into("!H", ip, 4, len(tcp) + len(payload)) + return bytes(ip + tcp + payload) + + +def mocked_init(*, enabled=1, workers=1, marks=1, intranet=0, tls=0, mmtls=0, + backend="iptables", fail_expr="", fail_queue=False): + """Exercise actual init control flow with print-only shell command functions.""" + source = (ROOT / "openwrt/files/ua2f.init").read_text().replace( + ". /usr/share/ua2f/firewall.sh", '. "$TEST_HELPER"') + mocks = r''' +record() { + printf '%s' "$1" + shift + printf '\t%s' "$@" + printf '\n' + case " $* " in *' -D '*) return 1 ;; esac + for arg in "$@"; do + [ -z "$TEST_FAIL_EXPR" ] || [ "$arg" != "$TEST_FAIL_EXPR" ] || return 1 + [ "$TEST_FAIL_QUEUE" != 1 ] || [ "$arg" != NFQUEUE ] || return 1 + done + if [ "$TEST_FAIL_QUEUE" = 2 ] && [ "$1" = -t ] && [ "${5:-}" = -j ] && [ "${6:-}" = NFQUEUE ]; then + return 1 + fi + return 0 +} +iptables() { record iptables "$@"; } +ip6tables() { record ip6tables "$@"; } +nft() { if [ "$1" = -f- ]; then cat; else record nft "$@"; fi; } +config_load() { :; } +config_get_bool() { + case "$1" in + handle_fw) export "$1=1" ;; + bypass_empty_ack) export "$1=$TEST_ENABLED" ;; + disable_connmark) export "$1=$TEST_DISABLE_MARKS" ;; + handle_intranet) export "$1=$TEST_INTRANET" ;; + handle_tls) export "$1=$TEST_TLS" ;; + handle_mmtls) export "$1=$TEST_MMTLS" ;; + *) return 1 ;; + esac +} +config_get() { + case "$1" in + mode) export "$1=NFQUEUE" ;; + listen_port) export "$1=10010" ;; + nfqueue_workers) export "$1=$TEST_WORKERS" ;; + *) return 1 ;; + esac +} +''' + env = dict(os.environ, TEST_HELPER=str(HELPER), TEST_ENABLED=str(enabled), + TEST_WORKERS=str(workers), TEST_DISABLE_MARKS=str(1-marks), + TEST_INTRANET=str(intranet), TEST_TLS=str(tls), TEST_MMTLS=str(mmtls), + TEST_BACKEND=backend, TEST_FAIL_EXPR=fail_expr, TEST_FAIL_QUEUE=str(int(fail_queue))) + run = subprocess.run(["sh", "-c", mocks + source + r''' +HAS_IPT6=mock +if [ "$TEST_BACKEND" = nft ]; then HAS_NFT=mock; else HAS_NFT=; fi +setup_firewall +'''], env=env, text=True, capture_output=True, check=True) + return run.stdout.splitlines(), run.stderr + + +class EmptyAckRules(unittest.TestCase): + @classmethod + def setUpClass(cls): + cls.rules = {family: generate("ua2f_empty_ack_u32", family) for family in (4, 6)} + cls.nft = generate("ua2f_empty_ack_nft") + cls.lengths = {family: generate("ua2f_empty_ack_length_u32", family)[0] for family in (4, 6)} + cls.tails = {family: [row.split("\t") for row in generate("ua2f_empty_ack_queue_iptables", family, 10010, 10010)] + for family in (4, 6)} + + def bypass(self, family, data): + return u32_matches(self.lengths[family], data) and any( + u32_matches(expr, data) for expr in self.rules[family]) + + def test_early_queue_cannot_skip_a_valid_empty_ack(self): + # Every exact RETURN implies the coarse length guard, including failed + # reads. The early queue can only bypass predicates that would all fail. + for family in (4, 6): + coarse, = parse_u32(self.lengths[family]) + for expr in self.rules[family]: + precise = parse_u32(expr)[0] + self.assertEqual(precise[:2], coarse[:2]) + self.assertEqual(precise[2], precise[3]) + self.assertLessEqual(coarse[2], precise[2]) + self.assertLessEqual(precise[3], coarse[3]) + tail = self.tails[family] + self.assertEqual(tail[0], ["-m", "u32", "!", "--u32", self.lengths[family], *tail[-1]]) + self.assertEqual(tail[-1], ["-j", "NFQUEUE", "--queue-num", "10010", "--queue-bypass"]) + self.assertEqual(tail[1:-1], [ + ["-p", "tcp", "-m", "conntrack", "--ctdir", "ORIGINAL", "-m", "u32", "--u32", expr, "-j", "RETURN"] + for expr in self.rules[family]]) + + def test_exact_rule_shape(self): + self.assertEqual(len(self.nft), 22) + for family in (4, 6): + self.assertEqual(len(self.rules[family]), 11) + for line in self.nft: + self.assertIn("ct direction original", line) + self.assertIn("tcp flags & 0x17 == 0x10", line) + self.assertIn('counter return comment "!ua2f: empty ACK";', line) + self.assertNotIn("mark", line) + for doff in range(5, 16): + self.assertIn(f"tcp doff {doff} ip length {20+4*doff}", self.nft[2*(doff-5)]) + self.assertTrue(self.nft[2*(doff-5)].startswith(f"meta length {20+4*doff} ")) + self.assertIn("ip hdrlength 5 ip frag-off & 0x3fff == 0", self.nft[2*(doff-5)]) + self.assertIn(f"tcp doff {doff} ip6 length {4*doff}", self.nft[2*(doff-5)+1]) + self.assertTrue(self.nft[2*(doff-5)+1].startswith(f"meta length {40+4*doff} ")) + self.assertIn("ip6 nexthdr tcp", self.nft[2*(doff-5)+1]) + + def test_nft_ruleset_separates_chains(self): + from test_empty_ack_netns import nft_ruleset + generated = nft_ruleset("test0", "\n".join(self.nft)) + self.assertIn(" }\n chain inspect {\n", generated) + self.assertTrue(generated.endswith(" }\n}\n")) + self.assertNotIn("} chain", generated) + self.assertNotIn("} }", generated) + self.assertEqual(generated.count("counter queue num 10010;"), 1) + for rule in self.nft: + self.assertEqual(generated.count(rule), 1) + + def test_all_flags_header_lengths_and_payloads(self): + # Exhaust all eight TCP flag bits, valid/invalid doff and adjacent lengths. + for family, doff, flags, size in itertools.product((4, 6), range(16), range(256), (0, 1, 2, 4, 100)): + want = doff >= 5 and size == 0 and flags & 0x17 == 0x10 + self.assertEqual(self.bypass(family, packet(family, doff, b"x"*size, flags)), want, + (family, doff, flags, size)) + + def test_ipv4_options_and_fragmentation_fall_back(self): + for ihl, doff, frag in itertools.product(range(16), range(16), (0, 0x4000, 0x2000, 1, 0x3fff)): + want = ihl == 5 and doff >= 5 and frag & 0x3fff == 0 + self.assertEqual(self.bypass(4, packet(4, doff, ihl=ihl, frag=frag)), want, (ihl, doff, frag)) + + def test_ipv6_extensions_and_other_protocols_fall_back(self): + for family, nh, doff in itertools.product((4, 6), range(256), range(5, 16)): + self.assertEqual(self.bypass(family, packet(family, doff, nexthdr=nh)), nh == 6) + + def test_truncation_invalid_lengths_and_jumbo_lengths_fall_back(self): + for family, doff in itertools.product((4, 6), range(5, 16)): + data = packet(family, doff) + for size in range(len(data)): + self.assertFalse(self.bypass(family, data[:size]), (family, doff, size)) + for length in (0, 1, 19, 65535): + malformed = bytearray(data) + struct.pack_into("!H", malformed, 2 if family == 4 else 4, length) + self.assertFalse(self.bypass(family, malformed)) + + def test_queue_tail_parameters_and_print_only_installer(self): + # The installer must execute the TSV generator verbatim, including arg + # boundaries, alternate queue IDs, balance ranges and queue-bypass flags. + code = r'''. "$1" +mock() { printf '%s' "$1"; shift; printf '\t%s' "$@"; printf '\n'; } +ua2f_add_empty_ack_queue_iptables "$2" "$3" "$4" mock -t mangle -A 'space in chain' +''' + for family, first, last in itertools.product((4, 6), (0, 10010, 65520), (0, 1, 15)): + last += first + rows = generate("ua2f_empty_ack_queue_iptables", family, first, last) + self.assertEqual(len(rows), 13) + self.assertEqual(rows[0].split("\t")[-5:], rows[-1].split("\t")) + self.assertEqual(rows[-1].split("\t"), ["-j", "NFQUEUE", "--queue-num" if first == last else "--queue-balance", + str(first) if first == last else f"{first}:{last}", "--queue-bypass"]) + actual = subprocess.check_output(["sh", "-c", code, "sh", str(HELPER), str(family), str(first), str(last)], text=True).splitlines() + self.assertEqual(actual, ["-t\tmangle\t-A\tspace in chain\t" + row for row in rows]) + + def test_invalid_queue_configuration_is_rejected_before_emission(self): + for family, first, last in ((7, 10010, 10010), (4, -1, 1), (6, 1, 65536), (4, 2, 1), + (4, "", 1), (4, 1, ""), (6, "x", 1), (4, "1:2", 3), + (4, "1;echo bad", 2), (6, "1\t2", 3)): + args = ["sh", "-c", '. "$1"; ua2f_empty_ack_queue_iptables "$2" "$3" "$4"', "sh", str(HELPER), str(family), str(first), str(last)] + run = subprocess.run(args, text=True, capture_output=True, check=False) + self.assertNotEqual(run.returncode, 0, (family, first, last)) + self.assertEqual(run.stdout, "") + + def test_actual_init_preserves_mark_and_bypass_order(self): + for workers, marks, intranet, tls, mmtls in itertools.product((1, 4, 16), (0, 1), (0, 1), (0, 1), (0, 1)): + baseline, err = mocked_init(enabled=0, workers=workers, marks=marks, intranet=intranet, tls=tls, mmtls=mmtls) + self.assertEqual(err, "") + optimized, err = mocked_init(enabled=1, workers=workers, marks=marks, intranet=intranet, tls=tls, mmtls=mmtls) + self.assertEqual(err, "") + for family, command in ((4, "iptables"), (6, "ip6tables")): + prefix = f"{command}\t-t\tmangle\t-A\tua2f\t" + before = [line.removeprefix(prefix) for line in baseline if line.startswith(prefix)] + after = [line.removeprefix(prefix) for line in optimized if line.startswith(prefix)] + # All rules preceding the original queue stay byte-for-byte in + # order; only its location is replaced by the generated tail. + tail = generate("ua2f_empty_ack_queue_iptables", family, 10010, 10010+workers-1) + self.assertEqual(after, before[:-1] + tail) + self.assertEqual(before[-1], tail[-1]) + if marks: + set44 = next(i for i, line in enumerate(after) if "--set-mark\t44" in line) + ret43 = next(i for i, line in enumerate(after) if "--mark\t43" in line) + self.assertLess(set44, ret43) + self.assertLess(ret43, len(before)-1) + else: + self.assertFalse(any("CONNMARK" in line or "connmark" in line for line in after)) + # Hook direction/family gates and all other chain operations unchanged. + is_tail = lambda line: any(line.startswith(f"{cmd}\t-t\tmangle\t-A\tua2f\t") for cmd in ("iptables", "ip6tables")) + self.assertEqual([line for line in baseline if not is_tail(line)], [line for line in optimized if not is_tail(line)]) + self.assertIn("iptables\t-t\tmangle\t-A\tPOSTROUTING\t-p\ttcp\t-m\tconntrack\t--ctdir\tORIGINAL\t-j\tua2f", optimized) + self.assertIn("ip6tables\t-t\tmangle\t-A\tPOSTROUTING\t-p\ttcp\t-m\tconntrack\t--ctdir\tORIGINAL\t-j\tua2f", optimized) + + def test_partial_install_always_attempts_original_fallback(self): + for family, workers in itertools.product((4, 6), (1, 16)): + command = "iptables" if family == 4 else "ip6tables" + prefix = f"{command}\t-t\tmangle\t-A\tua2f\t" + tail = generate("ua2f_empty_ack_queue_iptables", family, 10010, 10010+workers-1) + # Failure at every u32 rule, including the early guard, leaves a + # prefix of safe rules followed by the identical original fallback. + for failed_index in range(12): + failed = tail[failed_index].split("\t") + expr = failed[failed.index("--u32") + 1] + lines, err = mocked_init(workers=workers, fail_expr=expr) + actual = [line.removeprefix(prefix) for line in lines if line.startswith(prefix)] + self.assertTrue(actual[-1] == tail[-1]) + self.assertEqual(actual[-failed_index-2:], tail[:failed_index+1] + [tail[-1]]) + self.assertIn(f"empty-ACK IPv{family} optimization unavailable; keeping NFQUEUE fallback", err) + # A queue target failure also still attempts the old unconditional rule; + # the test does not pretend an unavailable NFQUEUE target can be fixed. + lines, err = mocked_init(fail_queue=True) + self.assertIn("keeping NFQUEUE fallback", err) + self.assertEqual(sum(line.endswith("-j\tNFQUEUE\t--queue-num\t10010\t--queue-bypass") for line in lines), 4) + lines, err = mocked_init(fail_queue=2) + self.assertIn("keeping NFQUEUE fallback", err) + for command in ("iptables", "ip6tables"): + prefix = f"{command}\t-t\tmangle\t-A\tua2f\t" + actual = [line.removeprefix(prefix) for line in lines if line.startswith(prefix)] + self.assertEqual(actual[-1], actual[-2]) + self.assertEqual(actual[-1], "-j\tNFQUEUE\t--queue-num\t10010\t--queue-bypass") + + def test_nft_path_is_unchanged(self): + for workers, marks in itertools.product((1, 16), (0, 1)): + disabled, err = mocked_init(enabled=0, workers=workers, marks=marks, backend="nft") + self.assertEqual(err, "") + enabled, err = mocked_init(enabled=1, workers=workers, marks=marks, backend="nft") + self.assertEqual(err, "") + # nft still inserts only the original 22 predicates, no new queue. + self.assertEqual([line for line in enabled if "!ua2f: empty ACK" not in line and line.strip()], + [line for line in disabled if line.strip()]) + emitted = [line.strip() for line in enabled if "!ua2f: empty ACK" in line] + self.assertEqual(emitted, self.nft) + + def test_opt_in_only_and_queue_order(self): + source = (ROOT / "openwrt/files/ua2f.init").read_text() + self.assertIn('config_get_bool bypass_empty_ack "firewall" "bypass_empty_ack" "0"', source) + self.assertEqual(source.count('! ua2f_add_empty_ack_queue_iptables 4 10010 "$nfqueue_end" iptables'), 1) + self.assertEqual(source.count('! ua2f_add_empty_ack_queue_iptables 6 10010 "$nfqueue_end" ip6tables'), 1) + self.assertLess(source.index("ct mark 43 counter return", source.index("setup_firewall()")), source.index("|| ua2f_empty_ack_nft")) + self.assertLess(source.index("|| ua2f_empty_ack_nft"), source.index("ct direction original counter $nfqueue_expr")) + self.assertIn("iptables-mod-u32", (ROOT / "openwrt/Makefile").read_text()) + + +if __name__ == "__main__": + unittest.main() diff --git a/src/handler.c b/src/handler.c index 071cc6a9..a53fcdb2 100644 --- a/src/handler.c +++ b/src/handler.c @@ -18,13 +18,13 @@ #include #include #include -#include #include #include #include #include static char *replacement_user_agent_string = NULL; +static size_t replacement_user_agent_string_length = 0; static bool replacement_user_agent_cleanup_registered = false; static const struct mark_op MARK_NONE = {false, 0}; @@ -40,13 +40,17 @@ bool use_conntrack = false; static void destroy_handler(void) { free(replacement_user_agent_string); replacement_user_agent_string = NULL; + replacement_user_agent_string_length = 0; } void init_handler() { init_not_http_cache(60); destroy_handler(); - replacement_user_agent_string = malloc(UA2F_MAX_USER_AGENT_LENGTH); + // libmnl may implement MNL_SOCKET_BUFFER_SIZE with sysconf(_SC_PAGESIZE). + // Capacity is fixed for this allocation; do not query it again per UA. + replacement_user_agent_string_length = UA2F_MAX_USER_AGENT_LENGTH; + replacement_user_agent_string = malloc(replacement_user_agent_string_length); assert(replacement_user_agent_string != NULL && "Failed to allocate user agent string"); if (!replacement_user_agent_cleanup_registered) { atexit(destroy_handler); @@ -56,12 +60,12 @@ void init_handler() { #ifdef UA2F_ENABLE_UCI if (config.use_custom_ua) { - memset(replacement_user_agent_string, ' ', UA2F_MAX_USER_AGENT_LENGTH); + memset(replacement_user_agent_string, ' ', replacement_user_agent_string_length); size_t custom_ua_len = strlen(config.custom_ua); - if (custom_ua_len > UA2F_MAX_USER_AGENT_LENGTH) { + if (custom_ua_len > replacement_user_agent_string_length) { syslog(LOG_WARNING, "Config user agent string is too long, truncating to %zu bytes", - (size_t)UA2F_MAX_USER_AGENT_LENGTH); - custom_ua_len = UA2F_MAX_USER_AGENT_LENGTH; + (size_t)replacement_user_agent_string_length); + custom_ua_len = replacement_user_agent_string_length; } memcpy(replacement_user_agent_string, config.custom_ua, custom_ua_len); syslog(LOG_INFO, "Using config user agent string: %.*s", (int)custom_ua_len, replacement_user_agent_string); @@ -76,12 +80,12 @@ void init_handler() { #ifdef UA2F_USE_CUSTOM_UA if (!ua_set) { - memset(replacement_user_agent_string, ' ', UA2F_MAX_USER_AGENT_LENGTH); + memset(replacement_user_agent_string, ' ', replacement_user_agent_string_length); size_t custom_ua_len = strlen(UA2F_CUSTOM_UA); - if (custom_ua_len > UA2F_MAX_USER_AGENT_LENGTH) { + if (custom_ua_len > replacement_user_agent_string_length) { syslog(LOG_WARNING, "Embed user agent string is too long, truncating to %zu bytes", - (size_t)UA2F_MAX_USER_AGENT_LENGTH); - custom_ua_len = UA2F_MAX_USER_AGENT_LENGTH; + (size_t)replacement_user_agent_string_length); + custom_ua_len = replacement_user_agent_string_length; } memcpy(replacement_user_agent_string, UA2F_CUSTOM_UA, custom_ua_len); syslog(LOG_INFO, "Using embed user agent string: %.*s", (int)custom_ua_len, replacement_user_agent_string); @@ -90,7 +94,7 @@ void init_handler() { #endif if (!ua_set) { - memset(replacement_user_agent_string, 'F', UA2F_MAX_USER_AGENT_LENGTH); + memset(replacement_user_agent_string, 'F', replacement_user_agent_string_length); syslog(LOG_INFO, "Custom user agent string not set, using default F-string."); } @@ -99,32 +103,7 @@ void init_handler() { const char *get_replacement_user_agent_string() { return replacement_user_agent_string; } -size_t get_replacement_user_agent_string_length() { return UA2F_MAX_USER_AGENT_LENGTH; } - -static const char *replacement_chunk(size_t replacement_offset, size_t ua_len, char **owned) { - *owned = NULL; - - if (replacement_offset <= UA2F_MAX_USER_AGENT_LENGTH && ua_len <= UA2F_MAX_USER_AGENT_LENGTH - replacement_offset) { - return replacement_user_agent_string + replacement_offset; - } - - char *buf = malloc(ua_len); - if (buf == NULL) { - return NULL; - } - memset(buf, ' ', ua_len); - - if (replacement_offset < UA2F_MAX_USER_AGENT_LENGTH) { - size_t available = UA2F_MAX_USER_AGENT_LENGTH - replacement_offset; - if (available > ua_len) { - available = ua_len; - } - memcpy(buf, replacement_user_agent_string + replacement_offset, available); - } - - *owned = buf; - return buf; -} +size_t get_replacement_user_agent_string_length() { return replacement_user_agent_string_length; } void add_to_cache(const struct nf_packet *pkt) { const struct addr_port target = { @@ -485,6 +464,14 @@ void handle_packet(const struct packet_io *io, void *io_ctx, const struct nf_pac if (parse_ret == 0 && ua_count > 0) { for (size_t i = 0; i < ua_count; i++) { ua_entries_copy[i] = *session_ua_entry_const(session, i); + // Validate every span before changing any bytes. Keep a failed + // stream closed even if the same continuation is retransmitted. + if (ua_entries_copy[i].offset > tcp_payload_len || + ua_entries_copy[i].len > tcp_payload_len - ua_entries_copy[i].offset) { + session->ua_allocation_failed = true; + parse_ret = HTTP_PARSER_NO_MEMORY; + break; + } } } session_state_unlock(session); @@ -494,7 +481,7 @@ void handle_packet(const struct packet_io *io, void *io_ctx, const struct nf_pac // mistaken for a new, non-HTTP stream and bypass rewriting. session_release(session); session = NULL; - syslog(LOG_ERR, "Failed to allocate User-Agent entries, dropping packet"); + syslog(LOG_ERR, "Failed to allocate or validate User-Agent entries, dropping packet"); SEND_VERDICT(NF_DROP, MARK_NONE, NULL); goto end; } @@ -518,45 +505,44 @@ void handle_packet(const struct packet_io *io, void *io_ctx, const struct nf_pac session_release(session); session = NULL; - // Mangle UA entries (using copied data, session lock released) + // Replacements have exactly the original length. Copy each span directly, + // then checksum once per packet rather than rescanning the entire packet + // once per User-Agent (quadratic for large pipelined/duplicate batches). for (size_t i = 0; i < ua_count; i++) { - const size_t ua_offset = ua_entries_copy[i].offset; - const size_t ua_len = ua_entries_copy[i].len; - const size_t replacement_offset = ua_entries_copy[i].replacement_offset; - if (ua_offset > UINT_MAX || ua_len > UINT_MAX) { - syslog(LOG_WARNING, "Skipping too-large user agent mangle entry"); - continue; + const struct ua_mangle_entry *entry = &ua_entries_copy[i]; + size_t available = 0; + if (entry->replacement_offset < replacement_user_agent_string_length) { + available = replacement_user_agent_string_length - entry->replacement_offset; + if (available > entry->len) { + available = entry->len; + } + memcpy((char *)tcp_payload + entry->offset, + replacement_user_agent_string + entry->replacement_offset, available); } - char *owned_replacement = NULL; - const char *replacement = replacement_chunk(replacement_offset, ua_len, &owned_replacement); - if (replacement == NULL) { - syslog(LOG_ERR, "Failed to allocate replacement chunk"); - goto end; + if (available < entry->len) { + memset((char *)tcp_payload + entry->offset + available, ' ', entry->len - available); } + } + if (ua_count > 0) { + // Preserve the length normalization performed by nfq_tcp_mangle_*. + // The verdict API uses the explicit packet pointer, not pktb_mangled. if (type == IPV4) { - if (!nfq_tcp_mangle_ipv4(pkt_buff, (unsigned int)ua_offset, (unsigned int)ua_len, replacement, - (unsigned int)ua_len)) { - free(owned_replacement); - syslog(LOG_ERR, "Failed to mangle ipv4 packet"); - goto end; - } + struct iphdr *ip_hdr = nfq_ip_get_hdr(pkt_buff); + ip_hdr->tot_len = htons((uint16_t)pktb_len(pkt_buff)); + nfq_ip_set_checksum(ip_hdr); + nfq_tcp_compute_checksum_ipv4(tcp_hdr, ip_hdr); } else { - if (!nfq_tcp_mangle_ipv6(pkt_buff, (unsigned int)ua_offset, (unsigned int)ua_len, replacement, - (unsigned int)ua_len)) { - free(owned_replacement); - syslog(LOG_ERR, "Failed to mangle ipv6 packet"); - goto end; - } + struct ip6_hdr *ip_hdr = nfq_ip6_get_hdr(pkt_buff); + ip_hdr->ip6_plen = htons((uint16_t)(pktb_len(pkt_buff) - sizeof(*ip_hdr))); + nfq_tcp_compute_checksum_ipv6(tcp_hdr, ip_hdr); } - free(owned_replacement); - } - - if (ua_count > 0) { count_user_agent_packet(); } - SEND_VERDICT(NF_ACCEPT, (ct_ok && new_session) ? MARK_HTTP : MARK_NONE, pkt_buff); + // An unchanged payload is already queued in the kernel. Returning it again + // adds an unnecessary userspace/netlink copy, especially for upload bodies. + SEND_VERDICT(NF_ACCEPT, (ct_ok && new_session) ? MARK_HTTP : MARK_NONE, ua_count > 0 ? pkt_buff : NULL); end: if (ua_entries_copy != ua_entries_inline) { diff --git a/src/proxy.c b/src/proxy.c index 87b70701..fddf731c 100644 --- a/src/proxy.c +++ b/src/proxy.c @@ -530,7 +530,7 @@ static void rewrite_user_agent_entries(uint8_t *buf, size_t len, const struct ht if (replacement == NULL) { return; } - const size_t replacement_len = UA2F_MAX_USER_AGENT_LENGTH; + const size_t replacement_len = get_replacement_user_agent_string_length(); for (size_t i = 0; i < session->ua_entry_count; i++) { const struct ua_mangle_entry *entry = session_ua_entry_const(session, i); @@ -541,14 +541,17 @@ static void rewrite_user_agent_entries(uint8_t *buf, size_t len, const struct ht continue; } - memset(buf + offset, ' ', ua_len); + size_t available = 0; if (replacement_offset < replacement_len) { - size_t available = replacement_len - replacement_offset; + available = replacement_len - replacement_offset; if (available > ua_len) { available = ua_len; } memcpy(buf + offset, replacement + replacement_offset, available); } + if (available < ua_len) { + memset(buf + offset + available, ' ', ua_len - available); + } } } diff --git a/test/handler_test.cc b/test/handler_test.cc index 0ea626ec..3a2ba122 100644 --- a/test/handler_test.cc +++ b/test/handler_test.cc @@ -13,6 +13,69 @@ extern "C" { #include "mock_packet_io.h" #include "packet_builder.h" +namespace { + +// Compute Internet checksums from bytes, independently of libnetfilter_queue. +uint16_t read_network_u16(const uint8_t *data) { + return static_cast((static_cast(data[0]) << 8) | data[1]); +} + +void write_network_u16(uint8_t *data, uint16_t value) { + data[0] = static_cast(value >> 8); + data[1] = static_cast(value); +} + +uint32_t checksum_sum(const uint8_t *data, size_t len, uint32_t sum = 0) { + while (len >= 2) { + sum += read_network_u16(data); + data += 2; + len -= 2; + } + if (len != 0) { + sum += static_cast(data[0]) << 8; + } + while (sum >> 16) { + sum = (sum & 0xffffU) + (sum >> 16); + } + return sum; +} + +uint32_t tcp_checksum_sum(const std::vector &packet, int ip_version, size_t tcp_offset) { + const size_t tcp_len = packet.size() - tcp_offset; + const size_t address_offset = ip_version == IPV4 ? 12 : 8; + const size_t address_len = ip_version == IPV4 ? 8 : 32; + uint32_t sum = checksum_sum(packet.data() + address_offset, address_len); + sum += IPPROTO_TCP + static_cast(tcp_len); + return checksum_sum(packet.data() + tcp_offset, tcp_len, sum); +} + +void set_packet_checksums(std::vector &packet, int ip_version, size_t tcp_offset) { + if (ip_version == IPV4) { + write_network_u16(packet.data() + 2, static_cast(packet.size())); + write_network_u16(packet.data() + 10, 0); + write_network_u16(packet.data() + 10, + static_cast(~checksum_sum(packet.data(), tcp_offset))); + } else { + write_network_u16(packet.data() + 4, static_cast(packet.size() - 40)); + } + write_network_u16(packet.data() + tcp_offset + 16, 0); + write_network_u16(packet.data() + tcp_offset + 16, + static_cast(~tcp_checksum_sum(packet, ip_version, tcp_offset))); +} + +void expect_valid_packet_checksums(const std::vector &packet, int ip_version, size_t tcp_offset) { + ASSERT_GE(packet.size(), tcp_offset + 20); + if (ip_version == IPV4) { + EXPECT_EQ(read_network_u16(packet.data() + 2), packet.size()); + EXPECT_EQ(checksum_sum(packet.data(), tcp_offset), 0xffffU); + } else { + EXPECT_EQ(read_network_u16(packet.data() + 4), packet.size() - 40); + } + EXPECT_EQ(tcp_checksum_sum(packet, ip_version, tcp_offset), 0xffffU); +} + +} // namespace + class HandlerTest : public ::testing::Test { protected: mock_io_context mock_ctx; @@ -74,7 +137,7 @@ TEST_F(HandlerTest, HttpGetWithUserAgent) { EXPECT_EQ(payload_str.find("Mozilla/5.0"), std::string::npos); } -// 2. HTTP GET without User-Agent → NF_ACCEPT, mangled data present (packet still sent back) +// 2. HTTP GET without User-Agent → NF_ACCEPT without redundant replacement data TEST_F(HandlerTest, HttpGetWithoutUserAgent) { const char *req = "GET / HTTP/1.1\r\nHost: example.com\r\n\r\n"; auto pkt = make_http_packet(req); @@ -83,6 +146,7 @@ TEST_F(HandlerTest, HttpGetWithoutUserAgent) { ASSERT_EQ(mock_ctx.verdicts.size(), 1u); EXPECT_EQ(mock_ctx.verdicts[0].verdict, NF_ACCEPT); + EXPECT_TRUE(mock_ctx.verdicts[0].mangled_data.empty()); } // 3. Non-HTTP traffic → NF_ACCEPT, no mangling @@ -520,3 +584,184 @@ TEST_F(HandlerTest, AllocationFailureKeepsSessionClosedForRetransmittedFragments session_wrunlock(); } } + + +TEST_F(HandlerTest, BodyOnlyVerdictOmitsPayloadAndPreservesSession) { + use_conntrack = true; + const char *header = "POST /upload HTTP/1.1\r\nUser-Agent: Original\r\nContent-Length: 12\r\n\r\n"; + auto first = make_http_packet_ct(header); + handle_packet(&mock_packet_io, &mock_ctx, &first); + ASSERT_EQ(mock_ctx.verdicts.size(), 1u); + ASSERT_FALSE(mock_ctx.verdicts[0].mangled_data.empty()); + const auto key = session_key_from_connid(100); + session_wrlock(); + auto *session = session_find(&key); + session_wrunlock(); + ASSERT_NE(session, nullptr); + session_state_lock(session); + session->last_active = time(nullptr) - 301; + session_state_unlock(session); + + auto body = make_http_packet_ct("User-Agent: ", 2); + handle_packet(&mock_packet_io, &mock_ctx, &body); + ASSERT_EQ(mock_ctx.verdicts.size(), 2u); + EXPECT_EQ(mock_ctx.verdicts[1].verdict, NF_ACCEPT); + EXPECT_TRUE(mock_ctx.verdicts[1].mangled_data.empty()); + EXPECT_FALSE(mock_ctx.verdicts[1].mark.should_set); + session_wrlock(); + EXPECT_EQ(session_find(&key), session); + EXPECT_EQ(session_cleanup_expired(300), 0); + session_wrunlock(); + + auto next = make_http_packet_ct("GET / HTTP/1.1\r\nUser-Agent: Next\r\n\r\n", 3); + handle_packet(&mock_packet_io, &mock_ctx, &next); + ASSERT_EQ(mock_ctx.verdicts.size(), 3u); + const auto payload = extract_tcp_payload(mock_ctx.verdicts[2].mangled_data, IPV4); + EXPECT_EQ(std::string(payload.begin(), payload.end()), "GET / HTTP/1.1\r\nUser-Agent: FFFF\r\n\r\n"); +} + +TEST_F(HandlerTest, NoUserAgentStillMarksNewHttpConnection) { + use_conntrack = true; + auto packet = make_http_packet_ct("GET / HTTP/1.1\r\nHost: example.com\r\n\r\n"); + handle_packet(&mock_packet_io, &mock_ctx, &packet); + ASSERT_EQ(mock_ctx.verdicts.size(), 1u); + EXPECT_EQ(mock_ctx.verdicts[0].verdict, NF_ACCEPT); + EXPECT_TRUE(mock_ctx.verdicts[0].mangled_data.empty()); + EXPECT_TRUE(mock_ctx.verdicts[0].mark.should_set); + EXPECT_EQ(mock_ctx.verdicts[0].mark.mark, static_cast(CONNMARK_HTTP)); +} + +TEST_F(HandlerTest, PipelinedOddLengthUserAgentsHaveValidIpv4Checksums) { + std::string request; + std::string expected; + for (size_t i = 0; i < 33; ++i) { + request += "GET / HTTP/1.1\r\nUser-Agent: OddUA\r\n\r\n"; + expected += "GET / HTTP/1.1\r\nUser-Agent: FFFFF\r\n\r\n"; + } + ASSERT_EQ(request.size() % 2, 1u); + auto raw = build_ipv4_tcp_packet(htonl(0x0a000001), htonl(0x0a000002), + 12345, 80, request.data(), request.size()); + set_packet_checksums(raw, IPV4, 20); + auto packet = make_nf_packet(raw, 1, IPV4); + handle_packet(&mock_packet_io, &mock_ctx, &packet); + + ASSERT_EQ(mock_ctx.verdicts.size(), 1u); + ASSERT_EQ(mock_ctx.verdicts[0].verdict, NF_ACCEPT); + const auto &rewritten = mock_ctx.verdicts[0].mangled_data; + ASSERT_EQ(rewritten.size(), raw.size()); + expect_valid_packet_checksums(rewritten, IPV4, 20); + const auto payload = extract_tcp_payload(rewritten, IPV4); + EXPECT_EQ(std::string(payload.begin(), payload.end()), expected); +} + +TEST_F(HandlerTest, DuplicateOddLengthUserAgentsHaveValidIpv6Checksum) { + std::string request = "GET / HTTP/1.1\r\n"; + std::string expected = request; + for (size_t i = 0; i < 33; ++i) { + request += "User-Agent: OddUA\r\n"; + expected += "User-Agent: FFFFF\r\n"; + } + request += "\r\n"; + expected += "\r\n"; + ASSERT_EQ(request.size() % 2, 1u); + struct in6_addr src = IN6ADDR_LOOPBACK_INIT; + struct in6_addr dst = IN6ADDR_LOOPBACK_INIT; + dst.s6_addr[15] = 2; + auto raw = build_ipv6_tcp_packet(src, dst, 12345, 80, request.data(), request.size()); + set_packet_checksums(raw, IPV6, 40); + auto packet = make_nf_packet(raw, 1, IPV6); + handle_packet(&mock_packet_io, &mock_ctx, &packet); + + ASSERT_EQ(mock_ctx.verdicts.size(), 1u); + ASSERT_EQ(mock_ctx.verdicts[0].verdict, NF_ACCEPT); + const auto &rewritten = mock_ctx.verdicts[0].mangled_data; + ASSERT_EQ(rewritten.size(), raw.size()); + expect_valid_packet_checksums(rewritten, IPV6, 40); + const auto payload = extract_tcp_payload(rewritten, IPV6); + EXPECT_EQ(std::string(payload.begin(), payload.end()), expected); +} + +TEST_F(HandlerTest, Ipv4AndTcpOptionsSurviveUserAgentRewriteWithValidChecksums) { + const std::string request = "GET / HTTP/1.1\r\nUser-Agent: OddUA\r\n\r\n"; + auto raw = build_ipv4_tcp_packet(htonl(0x0a000001), htonl(0x0a000002), + 12345, 80, request.data(), request.size()); + // Four IPv4 NOP options; TCP NOP, NOP, Timestamp options (12 bytes). + raw.insert(raw.begin() + 20, {1, 1, 1, 1}); + raw[0] = 0x46; + const std::vector tcp_options = {1, 1, 8, 10, 0, 0, 0, 1, 0, 0, 0, 2}; + raw.insert(raw.begin() + 44, tcp_options.begin(), tcp_options.end()); + raw[24 + 12] = static_cast((8U << 4) | (raw[24 + 12] & 0x0fU)); + set_packet_checksums(raw, IPV4, 24); + auto packet = make_nf_packet(raw, 1, IPV4); + handle_packet(&mock_packet_io, &mock_ctx, &packet); + + ASSERT_EQ(mock_ctx.verdicts.size(), 1u); + ASSERT_EQ(mock_ctx.verdicts[0].verdict, NF_ACCEPT); + const auto &rewritten = mock_ctx.verdicts[0].mangled_data; + ASSERT_EQ(rewritten.size(), raw.size()); + expect_valid_packet_checksums(rewritten, IPV4, 24); + EXPECT_EQ(std::vector(rewritten.begin() + 20, rewritten.begin() + 24), + (std::vector{1, 1, 1, 1})); + EXPECT_EQ(std::vector(rewritten.begin() + 44, rewritten.begin() + 56), tcp_options); + const auto payload = extract_tcp_payload(rewritten, IPV4); + EXPECT_EQ(std::string(payload.begin(), payload.end()), + "GET / HTTP/1.1\r\nUser-Agent: FFFFF\r\n\r\n"); +} + +TEST_F(HandlerTest, Ipv6AtomicFragmentHeaderSurvivesRewriteWithValidTcpChecksum) { + const std::string request = "GET / HTTP/1.1\r\nUser-Agent: OddUA\r\n\r\n"; + struct in6_addr src = IN6ADDR_LOOPBACK_INIT; + struct in6_addr dst = IN6ADDR_LOOPBACK_INIT; + dst.s6_addr[15] = 2; + auto raw = build_ipv6_tcp_packet(src, dst, 12345, 80, request.data(), request.size()); + // A complete atomic fragment: offset zero, M flag clear, identification zero. + const std::vector fragment_header = {IPPROTO_TCP, 0, 0, 0, 0, 0, 0, 0}; + raw.insert(raw.begin() + 40, fragment_header.begin(), fragment_header.end()); + raw[6] = IPPROTO_FRAGMENT; + set_packet_checksums(raw, IPV6, 48); + auto packet = make_nf_packet(raw, 1, IPV6); + handle_packet(&mock_packet_io, &mock_ctx, &packet); + + ASSERT_EQ(mock_ctx.verdicts.size(), 1u); + ASSERT_EQ(mock_ctx.verdicts[0].verdict, NF_ACCEPT); + const auto &rewritten = mock_ctx.verdicts[0].mangled_data; + ASSERT_EQ(rewritten.size(), raw.size()); + expect_valid_packet_checksums(rewritten, IPV6, 48); + EXPECT_EQ(std::vector(rewritten.begin() + 40, rewritten.begin() + 48), fragment_header); + EXPECT_EQ(std::string(rewritten.begin() + 68, rewritten.end()), + "GET / HTTP/1.1\r\nUser-Agent: FFFFF\r\n\r\n"); +} + +TEST_F(HandlerTest, SplitLongUserAgentPadsOnlyBeyondReplacementCapacity) { + const size_t capacity = get_replacement_user_agent_string_length(); + const size_t first_value_len = 60000; + ASSERT_GT(capacity, first_value_len); + const std::string prefix = "GET / HTTP/1.1\r\nUser-Agent: "; + const std::string first = prefix + std::string(first_value_len, 'A'); + const size_t replacement_remaining = capacity - first_value_len; + const std::string second(replacement_remaining + 17, 'B'); + const std::string third = "Original-tail\r\n\r\n"; + const std::string expected[] = { + prefix + std::string(get_replacement_user_agent_string(), first_value_len), + std::string(get_replacement_user_agent_string() + first_value_len, replacement_remaining) + + std::string(17, ' '), + std::string(std::strlen("Original-tail"), ' ') + "\r\n\r\n", + }; + const std::string fragments[] = {first, second, third}; + + for (size_t i = 0; i < 3; ++i) { + auto raw = build_ipv4_tcp_packet(htonl(0x0a000001), htonl(0x0a000002), + 12345, 80, fragments[i].data(), fragments[i].size()); + ASSERT_LE(raw.size(), 65535u); + set_packet_checksums(raw, IPV4, 20); + auto packet = make_nf_packet(raw, static_cast(i + 1), IPV4); + handle_packet(&mock_packet_io, &mock_ctx, &packet); + ASSERT_EQ(mock_ctx.verdicts.size(), i + 1); + ASSERT_EQ(mock_ctx.verdicts[i].verdict, NF_ACCEPT); + const auto &rewritten = mock_ctx.verdicts[i].mangled_data; + ASSERT_EQ(rewritten.size(), raw.size()); + expect_valid_packet_checksums(rewritten, IPV4, 20); + const auto payload = extract_tcp_payload(rewritten, IPV4); + EXPECT_EQ(std::string(payload.begin(), payload.end()), expected[i]); + } +}