Resolve client address from Forwarded and X-Forwarded-* headers

This commit is contained in:
2026-09-04 15:25:18 +08:00
parent 9a68d10868
commit 59f1a8227f
7 changed files with 221 additions and 14 deletions
+2 -1
View File
@@ -1,6 +1,6 @@
from ._app import AbstractKayaApp, KayaApp from ._app import AbstractKayaApp, KayaApp
from ._http_method import HttpMethod from ._http_method import HttpMethod
from ._http_context import HttpContext from ._http_context import HttpContext, resolve_client
from ._mixin import KayaMixin from ._mixin import KayaMixin
from ._tree import Tree, PathIterator from ._tree import Tree, PathIterator
from ._path_handler import PathHandler, Matches from ._path_handler import PathHandler, Matches
@@ -13,6 +13,7 @@ __all__ = [
'KayaApp', 'KayaApp',
'KayaMixin', 'KayaMixin',
'HttpContext', 'HttpContext',
'resolve_client',
'Tree', 'Tree',
'PathHandler', 'PathHandler',
'Matches', 'Matches',
+5 -5
View File
@@ -16,7 +16,7 @@ from typing import (
from pwo import Maybe from pwo import Maybe
from pathlib import Path from pathlib import Path
from ._http_method import HttpMethod from ._http_method import HttpMethod
from ._http_context import HttpContext from ._http_context import HttpContext, resolve_client
from ._websocket import WebSocket, WebSocketMessage from ._websocket import WebSocket, WebSocketMessage
from ._types import StrOrStrings from ._types import StrOrStrings
from ._types.asgi import HTTPScope, WebSocketScope from ._types.asgi import HTTPScope, WebSocketScope
@@ -84,9 +84,9 @@ class AsgiContext(HttpContext):
self.query_string = scope['query_string'].decode() self.query_string = scope['query_string'].decode()
self.method = HttpMethod(scope['method']) self.method = HttpMethod(scope['method'])
self.scheme = scope.get('scheme', 'http') self.scheme = scope.get('scheme', 'http')
self.client = scope['client']
self.server = scope['server']
self.headers = decode_headers(scope['headers']) self.headers = decode_headers(scope['headers'])
self.client = resolve_client(self.headers, scope['client'])
self.server = scope['server']
self.request_body = request_body_iterator self.request_body = request_body_iterator
self.session = (scope.get('state') or {}).get('kaya_session') self.session = (scope.get('state') or {}).get('kaya_session')
@@ -169,9 +169,9 @@ class AsgiWebSocket(WebSocket):
self.path = scope['path'] self.path = scope['path']
self.query_string = scope['query_string'].decode() self.query_string = scope['query_string'].decode()
self.scheme = scope.get('scheme', 'ws') self.scheme = scope.get('scheme', 'ws')
self.client = scope['client']
self.server = scope['server']
self.headers = decode_headers(scope['headers']) self.headers = decode_headers(scope['headers'])
self.client = resolve_client(self.headers, scope['client'])
self.server = scope['server']
async def accept(self, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None: async def accept(self, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None:
message: Dict[str, Any] = {'type': 'websocket.accept'} message: Dict[str, Any] = {'type': 'websocket.accept'}
@@ -16,6 +16,83 @@ from ._http_method import HttpMethod
from ._types.base import StrOrStrings 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): class HttpContext(ABC):
pathsend: bool pathsend: bool
receive: Callable[[], Awaitable[Any]] receive: Callable[[], Awaitable[Any]]
+59
View File
@@ -62,6 +62,11 @@ class AsgiTest(unittest.TestCase):
async def handle_request(ctx: HttpContext, _: List[str]) -> None: async def handle_request(ctx: HttpContext, _: List[str]) -> None:
await ctx.stream_body(200, (chunk async for chunk in ctx.request_body)) 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_test
async def test_hello(self): async def test_hello(self):
transport = httpx.ASGITransport(app=self.app) transport = httpx.ASGITransport(app=self.app)
@@ -190,6 +195,60 @@ class AsgiTest(unittest.TestCase):
'employee_id': 101325 'employee_id': 101325
}, response) }, 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_test
async def test_nested_param_routes(self): async def test_nested_param_routes(self):
app = KayaApp() app = KayaApp()
@@ -132,6 +132,37 @@ class WebSocketTest(unittest.TestCase):
self.assertEqual(1, len(sent_messages)) self.assertEqual(1, len(sent_messages))
self.assertEqual({'type': 'websocket.accept'}, sent_messages[0]) 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_test
async def test_websocket_scope_without_scheme(self): async def test_websocket_scope_without_scheme(self):
# Daphne omits the optional `scheme` key from websocket scopes. # Daphne omits the optional `scheme` key from websocket scopes.
+9 -7
View File
@@ -23,7 +23,7 @@ from granian._granian import ( # type: ignore[attr-defined]
) )
from pwo import Maybe from pwo import Maybe
from kaya.core import AbstractKayaApp, HttpContext, HttpMethod, WebSocket, WebSocketMessage from kaya.core import AbstractKayaApp, HttpContext, HttpMethod, WebSocket, WebSocketMessage, resolve_client
from kaya.core._types import StrOrStrings from kaya.core._types import StrOrStrings
@@ -51,9 +51,10 @@ class RsgiContext(HttpContext):
fun = cast(Callable[[Mapping[str, Sequence[str]], tuple[str, str]], Mapping[str, Sequence[str]]], acc) fun = cast(Callable[[Mapping[str, Sequence[str]], tuple[str, str]], Mapping[str, Sequence[str]]], acc)
self.headers = reduce(fun, scope.headers.items(), {}) self.headers = reduce(fun, scope.headers.items(), {})
self.client = (Maybe.of(scope.client.split(':')) self.client = resolve_client(self.headers,
.map(lambda it: (it[0], int(it[1]))) Maybe.of(scope.client.split(':'))
.or_else_throw(RuntimeError)) .map(lambda it: (it[0], int(it[1])))
.or_else_throw(RuntimeError))
self.server = (Maybe.of(scope.server.split(':')) self.server = (Maybe.of(scope.server.split(':'))
.map(lambda it: (it[0], int(it[1]))) .map(lambda it: (it[0], int(it[1])))
.or_else_throw(RuntimeError)) .or_else_throw(RuntimeError))
@@ -134,9 +135,10 @@ class RsgiWebSocket(WebSocket):
fun = cast(Callable[[Mapping[str, Sequence[str]], tuple[str, str]], Mapping[str, Sequence[str]]], acc) fun = cast(Callable[[Mapping[str, Sequence[str]], tuple[str, str]], Mapping[str, Sequence[str]]], acc)
self.headers = reduce(fun, scope.headers.items(), {}) self.headers = reduce(fun, scope.headers.items(), {})
self.client = (Maybe.of(scope.client.split(':')) self.client = resolve_client(self.headers,
.map(lambda it: (it[0], int(it[1]))) Maybe.of(scope.client.split(':'))
.or_else_throw(RuntimeError)) .map(lambda it: (it[0], int(it[1])))
.or_else_throw(RuntimeError))
self.server = (Maybe.of(scope.server.split(':')) self.server = (Maybe.of(scope.server.split(':'))
.map(lambda it: (it[0], int(it[1]))) .map(lambda it: (it[0], int(it[1])))
.or_else_throw(RuntimeError)) .or_else_throw(RuntimeError))
+38 -1
View File
@@ -1,5 +1,5 @@
import unittest import unittest
from kaya.rsgi import RsgiWebSocket from kaya.rsgi import RsgiContext, RsgiWebSocket
class RsgiWebSocketTest(unittest.TestCase): class RsgiWebSocketTest(unittest.TestCase):
@@ -20,3 +20,40 @@ class RsgiWebSocketTest(unittest.TestCase):
RsgiWebSocket(FakeScope(), FakeProtocol()) # type: ignore[arg-type] RsgiWebSocket(FakeScope(), FakeProtocol()) # type: ignore[arg-type]
self.assertIn('Granian was not configured for websockets', str(ctx.exception)) self.assertIn('Granian was not configured for websockets', str(ctx.exception))
class RsgiContextTest(unittest.TestCase):
@staticmethod
def _make_context(headers):
class FakeScope:
scheme = 'http'
method = 'GET'
path = '/'
query_string = ''
client = '127.0.0.1:12345'
server = '127.0.0.1:80'
def __init__(self, headers):
self.headers = headers
return RsgiContext(FakeScope(headers), object()) # type: ignore[arg-type]
def test_forwarded_header(self):
ctx = self._make_context({'forwarded': 'for=203.0.113.5:1234'})
self.assertEqual(('203.0.113.5', 1234), ctx.client)
def test_x_forwarded_headers(self):
ctx = self._make_context({
'x-forwarded-for': '203.0.113.5, 70.41.3.18',
'x-forwarded-port': '8443',
})
self.assertEqual(('203.0.113.5', 8443), ctx.client)
def test_x_forwarded_host_fallback(self):
ctx = self._make_context({'x-forwarded-host': '198.51.100.7'})
self.assertEqual(('198.51.100.7', 12345), ctx.client)
def test_socket_peer_fallback(self):
ctx = self._make_context({})
self.assertEqual(('127.0.0.1', 12345), ctx.client)