summaryrefslogtreecommitdiffstats
path: root/proxy/websockets.py
diff options
context:
space:
mode:
Diffstat (limited to 'proxy/websockets.py')
-rw-r--r--proxy/websockets.py112
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 @@
1import asyncio
2import json
3import logging
4import typing
5from base64 import b64decode, b64encode
6from contextlib import asynccontextmanager
7from uuid import uuid4
8
9from starlette.endpoints import WebSocketEndpoint
10from starlette.websockets import WebSocket
11
12from proxy.utils import get_path_with_query_string
13from ttun_server.proxy_queue import ProxyQueue
14from ttun_server.types import WebsocketMessage, WebsocketMessageType, WebsocketConnectData, WebsocketMessageData, \
15 WebsocketDisconnectData
16
17logger = logging.getLogger(__name__)
18logger.setLevel('DEBUG')
19
20
21class 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()