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/websockets.py | 112 ++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 112 insertions(+) create mode 100644 proxy/websockets.py (limited to 'proxy/websockets.py') 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