diff --git a/src/httpx2/httpx2/_utils.py b/src/httpx2/httpx2/_utils.py index e97f0e2c..ecd5fa7e 100644 --- a/src/httpx2/httpx2/_utils.py +++ b/src/httpx2/httpx2/_utils.py @@ -168,12 +168,23 @@ def __init__(self, pattern: str) -> None: ) url = URL(pattern) + self.network = None + prefix = url.path.lstrip("/") + if prefix: + # A CIDR-style host, e.g. "all://192.168.0.0/16" or "all://[::1]/64". + try: + self.network = ipaddress.ip_network(f"{url.host}/{prefix}", strict=False) + except ValueError: + self.network = None + self.pattern = pattern self.scheme = "" if url.scheme == "all" else url.scheme self.host = "" if url.host == "*" else url.host self.port = url.port - if not url.host or url.host == "*": + if self.network is not None: self.host_regex: typing.Pattern[str] | None = None + elif not url.host or url.host == "*": + self.host_regex = None elif url.host.startswith("*."): # *.example.com should match "www.example.com", but not "example.com" domain = re.escape(url.host[2:]) @@ -190,25 +201,44 @@ def __init__(self, pattern: str) -> None: def matches(self, other: URL) -> bool: if self.scheme and self.scheme != other.scheme: return False - if self.host and self.host_regex is not None and not self.host_regex.match(other.host): + if self.network is not None: + try: + other_address = ipaddress.ip_address(other.host) + except ValueError: + return False + if other_address not in self.network: + return False + elif self.host and self.host_regex is not None and not self.host_regex.match(other.host): return False if self.port is not None and self.port != other.port: return False return True @property - def priority(self) -> tuple[int, int, int]: + def priority(self) -> tuple[int, int, int, int, int]: """ The priority allows URLPattern instances to be sortable, so that we can match from most specific to least specific. """ + # Patterns without a network sort after any CIDR network. Smaller + # (more specific) networks should match first, compared by address + # count rather than prefix length, since an IPv4 /8 and an IPv6 /8 + # cover wildly different numbers of addresses. + has_no_network = 0 if self.network is not None else 1 + network_priority = self.network.num_addresses if self.network is not None else 0 # URLs with a port should take priority over URLs without a port. port_priority = 0 if self.port is not None else 1 # Longer hostnames should match first. host_priority = -len(self.host) # Longer schemes should match first. scheme_priority = -len(self.scheme) - return (port_priority, host_priority, scheme_priority) + return ( + has_no_network, + network_priority, + port_priority, + host_priority, + scheme_priority, + ) def __hash__(self) -> int: return hash(self.pattern) diff --git a/tests/httpx2/test_utils.py b/tests/httpx2/test_utils.py index c542edea..18195465 100644 --- a/tests/httpx2/test_utils.py +++ b/tests/httpx2/test_utils.py @@ -136,6 +136,16 @@ def test_get_environment_proxies(environment: dict[str, str], proxies: dict[str, ("http://", "https://example.com", False), ("all://", "https://example.com:123", True), ("", "https://example.com:123", True), + ("all://192.168.0.0/16", "http://192.168.5.10", True), + ("all://192.168.0.0/16", "http://192.168.5.10:8080", True), + ("all://192.168.0.0/16", "http://10.0.0.1", False), + ("all://192.168.0.0/16", "http://example.com", False), + ("all://[::1]/128", "http://[::1]", True), + ("all://[::1]/128", "http://[::1]:8080", True), + ("all://[::1]/128", "http://[::2]", False), + ("all://192.168.0.0.0/16", "http://192.168.4.5", False), # Invalid CIDR + ("all://[fe11::]/16", "http://[fe11:1234::5]", True), + ("all://192.168.0.10/16", "http://192.168.5.10", True), # host bits set ], ) def test_url_matches(pattern: str, url: str, expected: bool) -> None: @@ -149,9 +159,17 @@ def test_pattern_priority() -> None: URLPattern("http://"), URLPattern("http://example.com"), URLPattern("http://example.com:123"), + URLPattern("all://[::]/8"), # 2**120 addresses + URLPattern("all://192.168.1.0/24"), # 256 addresses + URLPattern("all://192.168.1.0/23"), # 512 addresses + URLPattern("all://[::]/126"), # 4 addresses ] random.shuffle(matchers) assert sorted(matchers) == [ + URLPattern("all://[::]/126"), + URLPattern("all://192.168.1.0/24"), + URLPattern("all://192.168.1.0/23"), + URLPattern("all://[::]/8"), URLPattern("http://example.com:123"), URLPattern("http://example.com"), URLPattern("http://"),