diff --git a/packages/kaya-core/src/kaya/core/__init__.py b/packages/kaya-core/src/kaya/core/__init__.py index 76b56e6..c6b81f7 100644 --- a/packages/kaya-core/src/kaya/core/__init__.py +++ b/packages/kaya-core/src/kaya/core/__init__.py @@ -1,6 +1,6 @@ from ._app import AbstractKayaApp, KayaApp from ._http_method import HttpMethod -from ._http_context import HttpContext +from ._http_context import HttpContext, resolve_client from ._mixin import KayaMixin from ._tree import Tree, PathIterator from ._path_handler import PathHandler, Matches @@ -13,6 +13,7 @@ __all__ = [ 'KayaApp', 'KayaMixin', 'HttpContext', + 'resolve_client', 'Tree', 'PathHandler', 'Matches', diff --git a/packages/kaya-core/src/kaya/core/_asgi.py b/packages/kaya-core/src/kaya/core/_asgi.py index 0612612..bfcb68e 100644 --- a/packages/kaya-core/src/kaya/core/_asgi.py +++ b/packages/kaya-core/src/kaya/core/_asgi.py @@ -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 +from ._http_context import HttpContext, resolve_client 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.client = scope['client'] - self.server = scope['server'] 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.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.client = scope['client'] - self.server = scope['server'] 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: message: Dict[str, Any] = {'type': 'websocket.accept'} diff --git a/packages/kaya-core/src/kaya/core/_http_context.py b/packages/kaya-core/src/kaya/core/_http_context.py index c76aac3..084c10d 100644 --- a/packages/kaya-core/src/kaya/core/_http_context.py +++ b/packages/kaya-core/src/kaya/core/_http_context.py @@ -16,6 +16,83 @@ 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]] diff --git a/packages/kaya-core/tests/test_asgi.py b/packages/kaya-core/tests/test_asgi.py index 193441b..1941c92 100644 --- a/packages/kaya-core/tests/test_asgi.py +++ b/packages/kaya-core/tests/test_asgi.py @@ -62,6 +62,11 @@ 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) @@ -190,6 +195,60 @@ 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() diff --git a/packages/kaya-core/tests/test_websocket.py b/packages/kaya-core/tests/test_websocket.py index 14120f3..4ab36f4 100644 --- a/packages/kaya-core/tests/test_websocket.py +++ b/packages/kaya-core/tests/test_websocket.py @@ -132,6 +132,37 @@ 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. diff --git a/packages/kaya-rsgi/src/kaya/rsgi/_rsgi.py b/packages/kaya-rsgi/src/kaya/rsgi/_rsgi.py index 774cbe6..97aa28e 100644 --- a/packages/kaya-rsgi/src/kaya/rsgi/_rsgi.py +++ b/packages/kaya-rsgi/src/kaya/rsgi/_rsgi.py @@ -23,7 +23,7 @@ from granian._granian import ( # type: ignore[attr-defined] ) 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 @@ -51,9 +51,10 @@ class RsgiContext(HttpContext): fun = cast(Callable[[Mapping[str, Sequence[str]], tuple[str, str]], Mapping[str, Sequence[str]]], acc) self.headers = reduce(fun, scope.headers.items(), {}) - self.client = (Maybe.of(scope.client.split(':')) - .map(lambda it: (it[0], int(it[1]))) - .or_else_throw(RuntimeError)) + self.client = resolve_client(self.headers, + Maybe.of(scope.client.split(':')) + .map(lambda it: (it[0], int(it[1]))) + .or_else_throw(RuntimeError)) self.server = (Maybe.of(scope.server.split(':')) .map(lambda it: (it[0], int(it[1]))) .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) self.headers = reduce(fun, scope.headers.items(), {}) - self.client = (Maybe.of(scope.client.split(':')) - .map(lambda it: (it[0], int(it[1]))) - .or_else_throw(RuntimeError)) + self.client = resolve_client(self.headers, + Maybe.of(scope.client.split(':')) + .map(lambda it: (it[0], int(it[1]))) + .or_else_throw(RuntimeError)) self.server = (Maybe.of(scope.server.split(':')) .map(lambda it: (it[0], int(it[1]))) .or_else_throw(RuntimeError)) diff --git a/packages/kaya-rsgi/tests/test_rsgi.py b/packages/kaya-rsgi/tests/test_rsgi.py index 096850b..9d7ae51 100644 --- a/packages/kaya-rsgi/tests/test_rsgi.py +++ b/packages/kaya-rsgi/tests/test_rsgi.py @@ -1,5 +1,5 @@ import unittest -from kaya.rsgi import RsgiWebSocket +from kaya.rsgi import RsgiContext, RsgiWebSocket class RsgiWebSocketTest(unittest.TestCase): @@ -20,3 +20,40 @@ class RsgiWebSocketTest(unittest.TestCase): RsgiWebSocket(FakeScope(), FakeProtocol()) # type: ignore[arg-type] 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)