diff options
Diffstat (limited to 'proxy/websockets.py')
| -rw-r--r-- | proxy/websockets.py | 112 |
1 files changed, 112 insertions, 0 deletions
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 @@ | |||
| 1 | import asyncio | ||
| 2 | import json | ||
| 3 | import logging | ||
| 4 | import typing | ||
| 5 | from base64 import b64decode, b64encode | ||
| 6 | from contextlib import asynccontextmanager | ||
| 7 | from uuid import uuid4 | ||
| 8 | |||
| 9 | from starlette.endpoints import WebSocketEndpoint | ||
| 10 | from starlette.websockets import WebSocket | ||
| 11 | |||
| 12 | from proxy.utils import get_path_with_query_string | ||
| 13 | from ttun_server.proxy_queue import ProxyQueue | ||
| 14 | from ttun_server.types import WebsocketMessage, WebsocketMessageType, WebsocketConnectData, WebsocketMessageData, \ | ||
| 15 | WebsocketDisconnectData | ||
| 16 | |||
| 17 | logger = logging.getLogger(__name__) | ||
| 18 | logger.setLevel('DEBUG') | ||
| 19 | |||
| 20 | |||
| 21 | class WebsocketProxy(WebSocketEndpoint): | ||
| 22 | encoding = 'json' | ||
| 23 | websocket_listen_task = None | ||
| 24 | |||
| 25 | def __init__(self, *args, **kwargs): | ||
| 26 | super().__init__(*args, **kwargs) | ||
| 27 | self.id = str(uuid4()) | ||
| 28 | |||
| 29 | @asynccontextmanager | ||
| 30 | async def proxy(self, websocket: WebSocket, message: WebsocketMessage): | ||
| 31 | [subdomain, *_] = websocket.url.hostname.split('.') | ||
| 32 | |||
| 33 | expect_ack = WebsocketMessageType(message['type']) == WebsocketMessageType.connect | ||
| 34 | |||
| 35 | try: | ||
| 36 | request_queue = await ProxyQueue.get_for_identifier(subdomain) | ||
| 37 | await request_queue.enqueue(message) | ||
| 38 | |||
| 39 | if expect_ack: | ||
| 40 | response_queue = await ProxyQueue.create_for_identifier(message["identifier"]) | ||
| 41 | yield await response_queue.dequeue() | ||
| 42 | await response_queue.delete() | ||
| 43 | else: | ||
| 44 | yield | ||
| 45 | except AssertionError: | ||
| 46 | yield None | ||
| 47 | |||
| 48 | async def listen_for_messages(self, websocket: WebSocket): | ||
| 49 | response_queue = await ProxyQueue.create_for_identifier(self.id) | ||
| 50 | |||
| 51 | while True: | ||
| 52 | message: WebsocketMessage = await response_queue.dequeue() | ||
| 53 | logger.debug(message) | ||
| 54 | await websocket.send_text(b64decode(message['payload']['body'].encode()).decode()) | ||
| 55 | |||
| 56 | async def on_connect(self, websocket: WebSocket) -> None: | ||
| 57 | message = WebsocketMessage( | ||
| 58 | type=WebsocketMessageType.connect.value, | ||
| 59 | identifier=self.id, | ||
| 60 | payload=WebsocketConnectData( | ||
| 61 | path=get_path_with_query_string(websocket), | ||
| 62 | headers=[ | ||
| 63 | (k.decode(), v.decode()) | ||
| 64 | for k, v | ||
| 65 | in websocket.scope['headers'] | ||
| 66 | ], | ||
| 67 | ) | ||
| 68 | ) | ||
| 69 | |||
| 70 | async with self.proxy(websocket, message) as m: | ||
| 71 | if m is not None and WebsocketMessageType(m['type']) == WebsocketMessageType.ack: | ||
| 72 | await super().on_connect(websocket) | ||
| 73 | |||
| 74 | self.websocket_listen_task = asyncio.create_task(self.listen_for_messages(websocket)) | ||
| 75 | |||
| 76 | def callback(*args, **kwargs): | ||
| 77 | self.websocket_listen_task = None | ||
| 78 | |||
| 79 | self.websocket_listen_task.add_done_callback(callback) | ||
| 80 | |||
| 81 | async def on_receive(self, websocket: WebSocket, data: typing.Any) -> None: | ||
| 82 | match data: | ||
| 83 | case dict(): | ||
| 84 | data_bytes = json.dumps(data).encode() | ||
| 85 | case bytes(): | ||
| 86 | data_bytes = data | ||
| 87 | case _: | ||
| 88 | data_bytes = data.encode() | ||
| 89 | |||
| 90 | message = WebsocketMessage( | ||
| 91 | type=WebsocketMessageType.message.value, | ||
| 92 | identifier=self.id, | ||
| 93 | payload=WebsocketMessageData( | ||
| 94 | body=b64encode(data_bytes).decode(), | ||
| 95 | ) | ||
| 96 | ) | ||
| 97 | |||
| 98 | async with self.proxy(websocket, message): | ||
| 99 | pass | ||
| 100 | |||
| 101 | async def on_disconnect(self, websocket: WebSocket, close_code: int) -> None: | ||
| 102 | message = WebsocketMessage( | ||
| 103 | type=WebsocketMessageType.disconnect.value, | ||
| 104 | identifier=self.id, | ||
| 105 | payload=WebsocketDisconnectData( | ||
| 106 | close_code=close_code, | ||
| 107 | ) | ||
| 108 | ) | ||
| 109 | |||
| 110 | async with self.proxy(websocket, message): | ||
| 111 | if self.websocket_listen_task is not None: | ||
| 112 | self.websocket_listen_task.cancel() | ||
