From d8968acab83c7a91b01eee4b35828a2b05c8dd6b Mon Sep 17 00:00:00 2001 From: Tom van der Lee Date: Wed, 1 Jul 2026 21:41:38 +0200 Subject: Split each part into its own app --- proxy/__init__.py | 13 ++++++ proxy/endpoints.py | 63 +++++++++++++++++++++++++++++ proxy/utils.py | 7 ++++ proxy/websockets.py | 112 ++++++++++++++++++++++++++++++++++++++++++++++++++++ 4 files changed, 195 insertions(+) create mode 100644 proxy/__init__.py create mode 100644 proxy/endpoints.py create mode 100644 proxy/utils.py create mode 100644 proxy/websockets.py (limited to 'proxy') diff --git a/proxy/__init__.py b/proxy/__init__.py new file mode 100644 index 0000000..fcdb23b --- /dev/null +++ b/proxy/__init__.py @@ -0,0 +1,13 @@ +from starlette.applications import Starlette +from starlette.routing import Route, WebSocketRoute + +from proxy.endpoints import Proxy +from proxy.websockets import WebsocketProxy + +app = Starlette( + debug=True, + routes=[ + Route('/{path:path}', Proxy), + WebSocketRoute('/{path:path}', WebsocketProxy), + ] +) diff --git a/proxy/endpoints.py b/proxy/endpoints.py new file mode 100644 index 0000000..c287a88 --- /dev/null +++ b/proxy/endpoints.py @@ -0,0 +1,63 @@ +import logging +from base64 import b64encode, b64decode +from uuid import uuid4 + +from starlette.endpoints import HTTPEndpoint +from starlette.requests import Request +from starlette.responses import Response + +from proxy.utils import get_path_with_query_string +from ttun_server.proxy_queue import ProxyQueue +from ttun_server.types import HttpMessage, HttpMessageType, HttpRequestData + +logger = logging.getLogger(__name__) + + +class HeaderMapping: + def __init__(self, headers: list[tuple[str, str]]): + self._headers = headers + + def items(self): + for header in self._headers: + yield header + + +class Proxy(HTTPEndpoint): + async def dispatch(self) -> None: + request = Request(self.scope, self.receive) + + subdomain = request.path_params['subdomain'] + response = Response(content='Not Found', status_code=404) + + identifier = str(uuid4()) + response_queue = await ProxyQueue.create_for_identifier(identifier) + + try: + request_queue = await ProxyQueue.get_for_identifier(subdomain) + + logger.debug('PROXY %s%s ', subdomain, request.url) + await request_queue.enqueue( + HttpMessage( + type=HttpMessageType.request.value, + identifier=identifier, + payload=HttpRequestData( + method=request.method, + path=get_path_with_query_string(request), + headers=list(request.headers.items()), + body=b64encode(await request.body()).decode() + ) + ) + ) + + _response = await response_queue.dequeue() + payload = _response['payload'] + response = Response( + status_code=payload['status'], + headers=HeaderMapping(payload['headers']), + content=b64decode(payload['body'].encode()) + ) + except AssertionError: + pass + finally: + await response(self.scope, self.receive, self.send) + await response_queue.delete() diff --git a/proxy/utils.py b/proxy/utils.py new file mode 100644 index 0000000..2b80f43 --- /dev/null +++ b/proxy/utils.py @@ -0,0 +1,7 @@ +from starlette.requests import HTTPConnection + + +def get_path_with_query_string(connection: HTTPConnection) -> str: + path = connection.url.path + query_string = '?' + connection.scope['query_string'].decode() if connection.scope['query_string'] else '' + return f"{path}{query_string}" diff --git a/proxy/websockets.py b/proxy/websockets.py new file mode 100644 index 0000000..ef80ec2 --- /dev/null +++ b/proxy/websockets.py @@ -0,0 +1,112 @@ +import asyncio +import json +import logging +import typing +from base64 import b64decode, b64encode +from contextlib import asynccontextmanager +from uuid import uuid4 + +from starlette.endpoints import WebSocketEndpoint +from starlette.websockets import WebSocket + +from proxy.utils import get_path_with_query_string +from ttun_server.proxy_queue import ProxyQueue +from ttun_server.types import WebsocketMessage, WebsocketMessageType, WebsocketConnectData, WebsocketMessageData, \ + WebsocketDisconnectData + +logger = logging.getLogger(__name__) +logger.setLevel('DEBUG') + + +class WebsocketProxy(WebSocketEndpoint): + encoding = 'json' + websocket_listen_task = None + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.id = str(uuid4()) + + @asynccontextmanager + async def proxy(self, websocket: WebSocket, message: WebsocketMessage): + [subdomain, *_] = websocket.url.hostname.split('.') + + expect_ack = WebsocketMessageType(message['type']) == WebsocketMessageType.connect + + try: + request_queue = await ProxyQueue.get_for_identifier(subdomain) + await request_queue.enqueue(message) + + if expect_ack: + response_queue = await ProxyQueue.create_for_identifier(message["identifier"]) + yield await response_queue.dequeue() + await response_queue.delete() + else: + yield + except AssertionError: + yield None + + async def listen_for_messages(self, websocket: WebSocket): + response_queue = await ProxyQueue.create_for_identifier(self.id) + + while True: + message: WebsocketMessage = await response_queue.dequeue() + logger.debug(message) + await websocket.send_text(b64decode(message['payload']['body'].encode()).decode()) + + async def on_connect(self, websocket: WebSocket) -> None: + message = WebsocketMessage( + type=WebsocketMessageType.connect.value, + identifier=self.id, + payload=WebsocketConnectData( + path=get_path_with_query_string(websocket), + headers=[ + (k.decode(), v.decode()) + for k, v + in websocket.scope['headers'] + ], + ) + ) + + async with self.proxy(websocket, message) as m: + if m is not None and WebsocketMessageType(m['type']) == WebsocketMessageType.ack: + await super().on_connect(websocket) + + self.websocket_listen_task = asyncio.create_task(self.listen_for_messages(websocket)) + + def callback(*args, **kwargs): + self.websocket_listen_task = None + + self.websocket_listen_task.add_done_callback(callback) + + async def on_receive(self, websocket: WebSocket, data: typing.Any) -> None: + match data: + case dict(): + data_bytes = json.dumps(data).encode() + case bytes(): + data_bytes = data + case _: + data_bytes = data.encode() + + message = WebsocketMessage( + type=WebsocketMessageType.message.value, + identifier=self.id, + payload=WebsocketMessageData( + body=b64encode(data_bytes).decode(), + ) + ) + + async with self.proxy(websocket, message): + pass + + async def on_disconnect(self, websocket: WebSocket, close_code: int) -> None: + message = WebsocketMessage( + type=WebsocketMessageType.disconnect.value, + identifier=self.id, + payload=WebsocketDisconnectData( + close_code=close_code, + ) + ) + + async with self.proxy(websocket, message): + if self.websocket_listen_task is not None: + self.websocket_listen_task.cancel() -- cgit v1.2.3