Resolve client address from Forwarded and X-Forwarded-* headers
This commit is contained in:
@@ -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',
|
||||||
|
|||||||
@@ -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]]
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
@@ -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,7 +51,8 @@ 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,
|
||||||
|
Maybe.of(scope.client.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))
|
||||||
self.server = (Maybe.of(scope.server.split(':'))
|
self.server = (Maybe.of(scope.server.split(':'))
|
||||||
@@ -134,7 +135,8 @@ 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,
|
||||||
|
Maybe.of(scope.client.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))
|
||||||
self.server = (Maybe.of(scope.server.split(':'))
|
self.server = (Maybe.of(scope.server.split(':'))
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user