proxy_headers.py 7.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187
  1. from __future__ import annotations
  2. import functools
  3. import ipaddress
  4. from uvicorn._types import ASGI3Application, ASGIReceiveCallable, ASGISendCallable, Scope
  5. class ProxyHeadersMiddleware:
  6. """Middleware for handling known proxy headers
  7. This middleware can be used when a known proxy is fronting the application,
  8. and is trusted to be properly setting the `X-Forwarded-Proto` and
  9. `X-Forwarded-For` headers with the connecting client information.
  10. Modifies the `client` and `scheme` information so that they reference
  11. the connecting client, rather that the connecting proxy.
  12. References:
  13. - <https://developer.mozilla.org/en-US/docs/Web/HTTP/Headers#Proxies>
  14. - <https://developer.mozilla.org/en-US/docs/Web/HTTP/Headers/X-Forwarded-For>
  15. """
  16. def __init__(self, app: ASGI3Application, trusted_hosts: list[str] | str = "127.0.0.1") -> None:
  17. self.app = app
  18. self.trusted_hosts = _TrustedHosts(trusted_hosts)
  19. async def __call__(self, scope: Scope, receive: ASGIReceiveCallable, send: ASGISendCallable) -> None:
  20. if scope["type"] == "lifespan":
  21. return await self.app(scope, receive, send)
  22. client_addr = scope.get("client")
  23. client_host = client_addr[0] if client_addr else None
  24. if client_host in self.trusted_hosts:
  25. x_forwarded_proto_value: bytes | None = None
  26. x_forwarded_for_values: list[bytes] = []
  27. for name, value in scope["headers"]:
  28. if name == b"x-forwarded-proto":
  29. x_forwarded_proto_value = value
  30. elif name == b"x-forwarded-for":
  31. x_forwarded_for_values.append(value)
  32. if x_forwarded_proto_value is not None:
  33. x_forwarded_proto = x_forwarded_proto_value.decode("latin1").strip()
  34. if x_forwarded_proto in {"http", "https", "ws", "wss"}:
  35. if scope["type"] == "websocket":
  36. scope["scheme"] = x_forwarded_proto.replace("http", "ws")
  37. else:
  38. scope["scheme"] = x_forwarded_proto
  39. if x_forwarded_for_values:
  40. x_forwarded_for = b", ".join(x_forwarded_for_values).decode("latin1")
  41. host, port = self.trusted_hosts.get_trusted_client_address(x_forwarded_for)
  42. if host:
  43. # If the x-forwarded-for header is empty then host is an empty string.
  44. # Only set the client if we actually got something usable.
  45. # See: https://github.com/Kludex/uvicorn/issues/1068
  46. scope["client"] = (host, port)
  47. return await self.app(scope, receive, send)
  48. def _parse_raw_hosts(value: str) -> list[str]:
  49. return [item.strip() for item in value.split(",")]
  50. def _parse_host_port(value: str) -> tuple[str, int]:
  51. """Parse a forwarded host value into host and optional port.
  52. Accepts bare IPs, IPv4 `host:port`, and bracketed IPv6 `[host]:port`.
  53. Any unrecognized or malformed value is treated conservatively and returned
  54. without a port so trust checks do not silently normalize arbitrary input.
  55. """
  56. if value.startswith("["):
  57. bracket_end = value.find("]")
  58. if bracket_end == -1:
  59. return value, 0
  60. host = value[1:bracket_end]
  61. remainder = value[bracket_end + 1 :]
  62. if not remainder:
  63. return host, 0
  64. if not remainder.startswith(":"):
  65. return value, 0
  66. try:
  67. return host, int(remainder[1:])
  68. except ValueError:
  69. return host, 0
  70. if value.count(":") == 1:
  71. host, port = value.rsplit(":", 1)
  72. try:
  73. return host, int(port)
  74. except ValueError:
  75. return value, 0
  76. return value, 0
  77. class _TrustedHosts:
  78. """Container for trusted hosts and networks"""
  79. def __init__(self, trusted_hosts: list[str] | str) -> None:
  80. self.always_trust: bool = trusted_hosts in ("*", ["*"])
  81. self.trusted_literals: set[str] = set()
  82. self.trusted_hosts: set[ipaddress.IPv4Address | ipaddress.IPv6Address] = set()
  83. self.trusted_networks: set[ipaddress.IPv4Network | ipaddress.IPv6Network] = set()
  84. # Notes:
  85. # - We separate hosts from literals as there are many ways to write
  86. # an IPv6 Address so we need to compare by object.
  87. # - We don't convert IP Address to single host networks (e.g. /32 / 128) as
  88. # it more efficient to do an address lookup in a set than check for
  89. # membership in each network.
  90. # - We still allow literals as it might be possible that we receive a
  91. # something that isn't an IP Address e.g. a unix socket.
  92. if not self.always_trust:
  93. if isinstance(trusted_hosts, str):
  94. trusted_hosts = _parse_raw_hosts(trusted_hosts)
  95. for host in trusted_hosts:
  96. # Note: because we always convert invalid IP types to literals it
  97. # is not possible for the user to know they provided a malformed IP
  98. # type - this may lead to unexpected / difficult to debug behaviour.
  99. if "/" in host:
  100. # Looks like a network
  101. try:
  102. self.trusted_networks.add(ipaddress.ip_network(host))
  103. except ValueError:
  104. # Was not a valid IP Network
  105. self.trusted_literals.add(host)
  106. else:
  107. try:
  108. self.trusted_hosts.add(ipaddress.ip_address(host))
  109. except ValueError:
  110. # Was not a valid IP Address
  111. self.trusted_literals.add(host)
  112. self._trusts = functools.lru_cache(maxsize=4096)(self._compute_trust)
  113. def __contains__(self, host: str | None) -> bool:
  114. if self.always_trust:
  115. return True
  116. if not host:
  117. return False
  118. # Don't cache hosts longer than a DNS name (253); they can't be trusted and would pin huge cache keys.
  119. if len(host) > 253:
  120. return self._compute_trust(host) # pragma: no cover
  121. return self._trusts(host)
  122. def _compute_trust(self, host: str) -> bool:
  123. try:
  124. ip = ipaddress.ip_address(host)
  125. return ip in self.trusted_hosts or any(ip in net for net in self.trusted_networks)
  126. except ValueError:
  127. return host in self.trusted_literals
  128. def get_trusted_client_address(self, x_forwarded_for: str) -> tuple[str, int]:
  129. """Extract the client address from x_forwarded_for header.
  130. In general this is the first "untrusted" host in the forwarded for list.
  131. """
  132. x_forwarded_for_hosts = _parse_raw_hosts(x_forwarded_for)
  133. if self.always_trust:
  134. return _parse_host_port(x_forwarded_for_hosts[0])
  135. # Note: each proxy appends to the header list so check it in reverse order
  136. for host_port in reversed(x_forwarded_for_hosts):
  137. host, port = _parse_host_port(host_port)
  138. if host not in self:
  139. return host, port
  140. # All hosts are trusted meaning that the client was also a trusted proxy
  141. # See https://github.com/Kludex/uvicorn/issues/1068#issuecomment-855371576
  142. return _parse_host_port(x_forwarded_for_hosts[0])