summaryrefslogtreecommitdiffstats
path: root/proxy/websockets.py
blob: ef80ec21c1fedee2dba8642153c146da245953f6 (plain)
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()