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 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()