import asyncio import json import unittest from typing import Optional import httpx from pwo import async_test from kaya.core import HttpContext, KayaApp from kaya.core._asgi import AsgiWebSocket from kaya.forwarded import ForwardedHeadersMixin TRUSTED = ['127.0.0.1', '::1', '10.0.0.0/8'] def make_app(trusted_proxies=TRUSTED) -> KayaApp: mixins = [ForwardedHeadersMixin(trusted_proxies=trusted_proxies)] if trusted_proxies is not None else [] app = KayaApp(mixins=mixins) @app.GET('/client') async def client(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})) return app async def request(app: KayaApp, headers: Optional[dict[str, str]] = None, client: tuple[str, int] = ('127.0.0.1', 123)) -> dict: transport = httpx.ASGITransport(app=app, client=client) async with httpx.AsyncClient(transport=transport, base_url="http://127.0.0.1:80") as http_client: r = await http_client.get('/client', headers=headers) assert r.status_code == 200 return json.loads(r.text) class TrustedPeerTest(unittest.TestCase): @async_test async def test_no_headers_returns_socket_peer(self): app = make_app() result = await request(app) self.assertEqual({'host': '127.0.0.1', 'port': 123}, result) @async_test async def test_forwarded_header_with_port(self): app = make_app() result = await request(app, headers={'Forwarded': 'for=203.0.113.5:1234'}) self.assertEqual({'host': '203.0.113.5', 'port': 1234}, result) @async_test async def test_forwarded_header_bracketed_ipv6(self): app = make_app() result = await request(app, headers={'Forwarded': 'for="[2001:db8::1]:4711"'}) self.assertEqual({'host': '2001:db8::1', 'port': 4711}, result) @async_test async def test_forwarded_header_without_port_keeps_socket_port(self): app = make_app() result = await request(app, headers={'Forwarded': 'for=203.0.113.5'}) self.assertEqual({'host': '203.0.113.5', 'port': 123}, result) @async_test async def test_forwarded_unknown_entry_skipped(self): app = make_app() result = await request(app, headers={'Forwarded': 'for=unknown, for=203.0.113.5'}) self.assertEqual('203.0.113.5', result['host']) @async_test async def test_forwarded_rightmost_untrusted_wins(self): # attacker-controlled leftmost entry is skipped: the rightmost # untrusted entry (appended by the trusted edge proxy) is the client app = make_app() result = await request(app, headers={'Forwarded': 'for=1.2.3.4, for=5.6.7.8, for=10.0.0.2'}) self.assertEqual('5.6.7.8', result['host']) @async_test async def test_x_forwarded_for_spoofed_leftmost_entry_skipped(self): # XFF = ", " as appended by the proxy app = make_app() result = await request(app, headers={'X-Forwarded-For': '1.2.3.4, 5.6.7.8'}) self.assertEqual('5.6.7.8', result['host']) @async_test async def test_x_forwarded_for_all_trusted_chain_uses_leftmost(self): app = make_app() result = await request(app, headers={'X-Forwarded-For': '10.0.0.5, 10.0.0.2'}) self.assertEqual('10.0.0.5', result['host']) @async_test async def test_x_forwarded_port(self): app = make_app() result = await request(app, headers={ 'X-Forwarded-For': '203.0.113.5', 'X-Forwarded-Port': '8443', }) self.assertEqual({'host': '203.0.113.5', 'port': 8443}, result) @async_test async def test_invalid_x_forwarded_port_ignored(self): app = make_app() result = await request(app, headers={ 'X-Forwarded-For': '203.0.113.5', 'X-Forwarded-Port': 'not-a-port', }) self.assertEqual({'host': '203.0.113.5', 'port': 123}, result) @async_test async def test_forwarded_takes_precedence_over_x_forwarded_for(self): app = make_app() result = await request(app, headers={ 'Forwarded': 'for=203.0.113.5', 'X-Forwarded-For': '198.51.100.7', }) self.assertEqual('203.0.113.5', result['host']) @async_test async def test_x_forwarded_host_fallback(self): app = make_app() result = await request(app, headers={'X-Forwarded-Host': '198.51.100.7'}) self.assertEqual('198.51.100.7', result['host']) @async_test async def test_ipv6_cidr_trust(self): app = make_app(trusted_proxies=['2001:db8::/32']) result = await request(app, headers={'X-Forwarded-For': '203.0.113.5'}, client=('2001:db8::10', 9999)) self.assertEqual({'host': '203.0.113.5', 'port': 9999}, result) class UntrustedPeerTest(unittest.TestCase): @async_test async def test_untrusted_peer_ignores_forwarded_headers(self): app = make_app() result = await request(app, headers={'Forwarded': 'for=1.2.3.4', 'X-Forwarded-For': '1.2.3.4'}, client=('203.0.113.99', 4567)) self.assertEqual({'host': '203.0.113.99', 'port': 4567}, result) @async_test async def test_empty_trusted_proxies_ignores_everything(self): app = make_app(trusted_proxies=[]) result = await request(app, headers={'X-Forwarded-For': '1.2.3.4'}) self.assertEqual({'host': '127.0.0.1', 'port': 123}, result) @async_test async def test_invalid_cidr_fails_fast(self): with self.assertRaises(ValueError): ForwardedHeadersMixin(trusted_proxies=['not-a-cidr']) class OptOutTest(unittest.TestCase): @async_test async def test_without_mixin_headers_are_ignored(self): app = KayaApp() @app.GET('/client') async def client(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})) result = await request(app, headers={'Forwarded': 'for=1.2.3.4', 'X-Forwarded-For': '1.2.3.4'}) self.assertEqual({'host': '127.0.0.1', 'port': 123}, result) class WebSocketTest(unittest.TestCase): @staticmethod def _make_ws(headers): async def send(message): pass async def receive(): return {'type': 'websocket.connect'} scope = { 'type': 'websocket', 'path': '/ws', 'query_string': b'', 'scheme': 'ws', 'client': ('127.0.0.1', 12345), 'server': ('127.0.0.1', 80), 'headers': headers, } return AsgiWebSocket(scope, receive, send) @async_test async def test_websocket_trusted_peer(self): mixin = ForwardedHeadersMixin(trusted_proxies=TRUSTED) ws = self._make_ws([(b'x-forwarded-for', b'1.2.3.4, 5.6.7.8')]) wrapped = await mixin._before_websocket(ws) assert wrapped is not None self.assertEqual(('5.6.7.8', 12345), wrapped.client) @async_test async def test_websocket_untrusted_peer(self): mixin = ForwardedHeadersMixin(trusted_proxies=['10.0.0.0/8']) ws = self._make_ws([(b'x-forwarded-for', b'1.2.3.4')]) wrapped = await mixin._before_websocket(ws) self.assertIsNone(wrapped) self.assertEqual(('127.0.0.1', 12345), ws.client) class RsgiTest(unittest.TestCase): def test_rsgi_context(self): from kaya.rsgi import RsgiContext class FakeScope: scheme = 'http' method = 'GET' path = '/' query_string = '' headers = {'x-forwarded-for': '1.2.3.4, 5.6.7.8'} client = '127.0.0.1:12345' server = '127.0.0.1:80' mixin = ForwardedHeadersMixin(trusted_proxies=TRUSTED) ctx = RsgiContext(FakeScope(), object()) # type: ignore[arg-type] wrapped = asyncio.run(mixin._before_request(ctx)) assert wrapped is not None self.assertEqual(('5.6.7.8', 12345), wrapped.client) def test_rsgi_context_untrusted_peer(self): from kaya.rsgi import RsgiContext class FakeScope: scheme = 'http' method = 'GET' path = '/' query_string = '' headers = {'x-forwarded-for': '1.2.3.4'} client = '192.0.2.1:12345' server = '127.0.0.1:80' mixin = ForwardedHeadersMixin(trusted_proxies=TRUSTED) ctx = RsgiContext(FakeScope(), object()) # type: ignore[arg-type] wrapped = asyncio.run(mixin._before_request(ctx)) self.assertIsNone(wrapped) self.assertEqual(('192.0.2.1', 12345), ctx.client)