1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
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()
|