Refactor forwarded header handling into opt-in kaya-forwarded package with trusted CIDRs
CI / Build Pip package (push) Successful in 3m59s

This commit is contained in:
2026-09-05 15:11:05 +08:00
parent 850d5dda35
commit aa35e2d30c
17 changed files with 643 additions and 222 deletions
+1 -2
View File
@@ -1,6 +1,6 @@
from ._app import AbstractKayaApp, KayaApp
from ._http_method import HttpMethod
from ._http_context import HttpContext, resolve_client
from ._http_context import HttpContext
from ._mixin import KayaMixin
from ._tree import Tree, PathIterator
from ._path_handler import PathHandler, Matches
@@ -13,7 +13,6 @@ __all__ = [
'KayaApp',
'KayaMixin',
'HttpContext',
'resolve_client',
'Tree',
'PathHandler',
'Matches',
+5 -5
View File
@@ -16,7 +16,7 @@ from typing import (
from pwo import Maybe
from pathlib import Path
from ._http_method import HttpMethod
from ._http_context import HttpContext, resolve_client
from ._http_context import HttpContext
from ._websocket import WebSocket, WebSocketMessage
from ._types import StrOrStrings
from ._types.asgi import HTTPScope, WebSocketScope
@@ -84,9 +84,9 @@ class AsgiContext(HttpContext):
self.query_string = scope['query_string'].decode()
self.method = HttpMethod(scope['method'])
self.scheme = scope.get('scheme', 'http')
self.headers = decode_headers(scope['headers'])
self.client = resolve_client(self.headers, scope['client'])
self.client = scope['client']
self.server = scope['server']
self.headers = decode_headers(scope['headers'])
self.request_body = request_body_iterator
self.session = (scope.get('state') or {}).get('kaya_session')
@@ -169,9 +169,9 @@ class AsgiWebSocket(WebSocket):
self.path = scope['path']
self.query_string = scope['query_string'].decode()
self.scheme = scope.get('scheme', 'ws')
self.headers = decode_headers(scope['headers'])
self.client = resolve_client(self.headers, scope['client'])
self.client = scope['client']
self.server = scope['server']
self.headers = decode_headers(scope['headers'])
async def accept(self, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None:
message: Dict[str, Any] = {'type': 'websocket.accept'}
@@ -16,83 +16,6 @@ from ._http_method import HttpMethod
from ._types.base import StrOrStrings
def _split_host_port(value: str) -> Tuple[str, Optional[int]]:
if value.startswith('['):
# bracketed IPv6 address, optionally followed by :port
closing = value.find(']')
if closing == -1:
return value, None
host = value[1:closing]
rest = value[closing + 1:]
if rest.startswith(':'):
try:
return host, int(rest[1:])
except ValueError:
return host, None
return host, None
if value.count(':') == 1:
host, _, port_str = value.rpartition(':')
try:
return host, int(port_str)
except ValueError:
return value, None
return value, None
def _first_header_value(headers: Mapping[str, Sequence[str]], name: str) -> Optional[str]:
values = headers.get(name)
if not values:
return None
first = values[0].split(',')[0].strip()
return first or None
def resolve_client(headers: Mapping[str, Sequence[str]],
client: Optional[Tuple[str, int]]) -> Optional[Tuple[str, int]]:
"""
Resolve the client (host, port) pair, honoring forwarded headers.
Resolution order:
1. the ``for=`` parameter of the first entry of the RFC 7239 ``Forwarded`` header
(a ``:port`` suffix, if present, also populates the port)
2. the first entry of ``X-Forwarded-For``
3. the first entry of ``X-Forwarded-Host``
4. the socket peer address (``client``), returned unchanged
In the ``X-Forwarded-*`` cases the port is taken from ``X-Forwarded-Port``
when present and valid, otherwise the socket port is kept.
"""
socket_port = client[1] if client is not None else 0
forwarded_values = headers.get('forwarded')
if forwarded_values:
for raw_value in forwarded_values:
first_entry = raw_value.split(',')[0]
for param in first_entry.split(';'):
key, sep, value = param.partition('=')
if sep and key.strip().lower() == 'for':
for_value = value.strip().strip('"')
if for_value and for_value.lower() != 'unknown':
forwarded_host, forwarded_port = _split_host_port(for_value)
if forwarded_host:
return (forwarded_host,
forwarded_port if forwarded_port is not None else socket_port)
for header_name in ('x-forwarded-for', 'x-forwarded-host'):
host = _first_header_value(headers, header_name)
if host is not None:
port = socket_port
x_forwarded_port = _first_header_value(headers, 'x-forwarded-port')
if x_forwarded_port is not None:
try:
port = int(x_forwarded_port)
except ValueError:
pass
return host, port
return client
class HttpContext(ABC):
pathsend: bool
receive: Callable[[], Awaitable[Any]]
-59
View File
@@ -62,11 +62,6 @@ class AsgiTest(unittest.TestCase):
async def handle_request(ctx: HttpContext, _: List[str]) -> None:
await ctx.stream_body(200, (chunk async for chunk in ctx.request_body))
@self.app.GET('/client-ip')
async def handle_request(ctx: HttpContext) -> None:
host, port = ctx.client if ctx.client is not None else (None, None)
await ctx.send_str(200, json.dumps({'host': host, 'port': port}))
@async_test
async def test_hello(self):
transport = httpx.ASGITransport(app=self.app)
@@ -195,60 +190,6 @@ class AsgiTest(unittest.TestCase):
'employee_id': 101325
}, response)
@async_test
async def test_client_ip_forwarded(self):
transport = httpx.ASGITransport(app=self.app)
async with httpx.AsyncClient(transport=transport, base_url="http://127.0.0.1:80") as client:
# socket peer, no forwarded headers
r = await client.get("/client-ip")
socket_client = json.loads(r.text)
self.assertEqual('127.0.0.1', socket_client['host'])
# RFC 7239 Forwarded header, with port
r = await client.get("/client-ip", headers={'Forwarded': 'for=203.0.113.5:1234'})
self.assertEqual({'host': '203.0.113.5', 'port': 1234}, json.loads(r.text))
# RFC 7239 Forwarded header, bracketed IPv6 with port
r = await client.get("/client-ip", headers={'Forwarded': 'for="[2001:db8::1]:4711"'})
self.assertEqual({'host': '2001:db8::1', 'port': 4711}, json.loads(r.text))
# RFC 7239 Forwarded header without port keeps the socket port
r = await client.get("/client-ip", headers={'Forwarded': 'for=203.0.113.5'})
self.assertEqual({'host': '203.0.113.5', 'port': socket_client['port']}, json.loads(r.text))
# Forwarded with for=unknown falls through to X-Forwarded-For
r = await client.get("/client-ip", headers={
'Forwarded': 'for=unknown',
'X-Forwarded-For': '198.51.100.7',
})
self.assertEqual('198.51.100.7', json.loads(r.text)['host'])
# Forwarded takes precedence over X-Forwarded-For
r = await client.get("/client-ip", headers={
'Forwarded': 'for=203.0.113.5',
'X-Forwarded-For': '198.51.100.7',
})
self.assertEqual('203.0.113.5', json.loads(r.text)['host'])
# X-Forwarded-For: first entry of the chain, port from X-Forwarded-Port
r = await client.get("/client-ip", headers={
'X-Forwarded-For': '203.0.113.5, 70.41.3.18',
'X-Forwarded-Port': '8443',
})
self.assertEqual({'host': '203.0.113.5', 'port': 8443}, json.loads(r.text))
# X-Forwarded-Host fallback
r = await client.get("/client-ip", headers={'X-Forwarded-Host': '198.51.100.7'})
self.assertEqual('198.51.100.7', json.loads(r.text)['host'])
# invalid X-Forwarded-Port is ignored, socket port is kept
r = await client.get("/client-ip", headers={
'X-Forwarded-For': '203.0.113.5',
'X-Forwarded-Port': 'not-a-port',
})
self.assertEqual({'host': '203.0.113.5', 'port': socket_client['port']}, json.loads(r.text))
@async_test
async def test_nested_param_routes(self):
app = KayaApp()
@@ -132,37 +132,6 @@ class WebSocketTest(unittest.TestCase):
self.assertEqual(1, len(sent_messages))
self.assertEqual({'type': 'websocket.accept'}, sent_messages[0])
@async_test
async def test_client_forwarded_header(self):
async def send(message):
pass
async def receive():
return {'type': 'websocket.connect'}
scope = {
'type': 'websocket',
'path': '/echo',
'query_string': b'',
'scheme': 'ws',
'client': ('127.0.0.1', 12345),
'server': ('127.0.0.1', 80),
'headers': [(b'forwarded', b'for=203.0.113.5:1234')],
}
ws = AsgiWebSocket(scope, receive, send)
self.assertEqual(('203.0.113.5', 1234), ws.client)
scope['headers'] = [
(b'x-forwarded-for', b'203.0.113.5, 70.41.3.18'),
(b'x-forwarded-port', b'8443'),
]
ws = AsgiWebSocket(scope, receive, send)
self.assertEqual(('203.0.113.5', 8443), ws.client)
scope['headers'] = []
ws = AsgiWebSocket(scope, receive, send)
self.assertEqual(('127.0.0.1', 12345), ws.client)
@async_test
async def test_websocket_scope_without_scheme(self):
# Daphne omits the optional `scheme` key from websocket scopes.