From d8968acab83c7a91b01eee4b35828a2b05c8dd6b Mon Sep 17 00:00:00 2001 From: Tom van der Lee Date: Wed, 1 Jul 2026 21:41:38 +0200 Subject: Split each part into its own app --- proxy/__init__.py | 13 +++ proxy/endpoints.py | 63 ++++++++++++++ proxy/utils.py | 7 ++ proxy/websockets.py | 112 ++++++++++++++++++++++++ ttun_server/__init__.py | 23 ++--- ttun_server/endpoints.py | 57 ++---------- ttun_server/types.py | 1 + ttun_server/websockets.py | 217 +++++++++++++--------------------------------- 8 files changed, 268 insertions(+), 225 deletions(-) create mode 100644 proxy/__init__.py create mode 100644 proxy/endpoints.py create mode 100644 proxy/utils.py create mode 100644 proxy/websockets.py diff --git a/proxy/__init__.py b/proxy/__init__.py new file mode 100644 index 0000000..fcdb23b --- /dev/null +++ b/proxy/__init__.py @@ -0,0 +1,13 @@ +from starlette.applications import Starlette +from starlette.routing import Route, WebSocketRoute + +from proxy.endpoints import Proxy +from proxy.websockets import WebsocketProxy + +app = Starlette( + debug=True, + routes=[ + Route('/{path:path}', Proxy), + WebSocketRoute('/{path:path}', WebsocketProxy), + ] +) diff --git a/proxy/endpoints.py b/proxy/endpoints.py new file mode 100644 index 0000000..c287a88 --- /dev/null +++ b/proxy/endpoints.py @@ -0,0 +1,63 @@ +import logging +from base64 import b64encode, b64decode +from uuid import uuid4 + +from starlette.endpoints import HTTPEndpoint +from starlette.requests import Request +from starlette.responses import Response + +from proxy.utils import get_path_with_query_string +from ttun_server.proxy_queue import ProxyQueue +from ttun_server.types import HttpMessage, HttpMessageType, HttpRequestData + +logger = logging.getLogger(__name__) + + +class HeaderMapping: + def __init__(self, headers: list[tuple[str, str]]): + self._headers = headers + + def items(self): + for header in self._headers: + yield header + + +class Proxy(HTTPEndpoint): + async def dispatch(self) -> None: + request = Request(self.scope, self.receive) + + subdomain = request.path_params['subdomain'] + response = Response(content='Not Found', status_code=404) + + identifier = str(uuid4()) + response_queue = await ProxyQueue.create_for_identifier(identifier) + + try: + request_queue = await ProxyQueue.get_for_identifier(subdomain) + + logger.debug('PROXY %s%s ', subdomain, request.url) + await request_queue.enqueue( + HttpMessage( + type=HttpMessageType.request.value, + identifier=identifier, + payload=HttpRequestData( + method=request.method, + path=get_path_with_query_string(request), + headers=list(request.headers.items()), + body=b64encode(await request.body()).decode() + ) + ) + ) + + _response = await response_queue.dequeue() + payload = _response['payload'] + response = Response( + status_code=payload['status'], + headers=HeaderMapping(payload['headers']), + content=b64decode(payload['body'].encode()) + ) + except AssertionError: + pass + finally: + await response(self.scope, self.receive, self.send) + await response_queue.delete() diff --git a/proxy/utils.py b/proxy/utils.py new file mode 100644 index 0000000..2b80f43 --- /dev/null +++ b/proxy/utils.py @@ -0,0 +1,7 @@ +from starlette.requests import HTTPConnection + + +def get_path_with_query_string(connection: HTTPConnection) -> str: + path = connection.url.path + query_string = '?' + connection.scope['query_string'].decode() if connection.scope['query_string'] else '' + return f"{path}{query_string}" 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 @@ +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() diff --git a/ttun_server/__init__.py b/ttun_server/__init__.py index 6c77858..d54227a 100644 --- a/ttun_server/__init__.py +++ b/ttun_server/__init__.py @@ -2,30 +2,23 @@ import logging import os from fastapi import FastAPI -from starlette.routing import Host, Route, Router, WebSocketRoute +from starlette.routing import Host, Route, WebSocketRoute -from ttun_server.endpoints import health, proxy -from .websockets import WebsocketProxy, Tunnel +from proxy import app as proxy_app +from ttun_server.endpoints import health, base_endpoints +from ttun_server.websockets import tunnel logging.basicConfig(level=getattr(logging, os.environ.get('LOG_LEVEL', 'INFO'))) -base_router = Router(routes=[ - Route('/health/', health), - WebSocketRoute('/tunnel/', Tunnel) -]) - -server = FastAPI( +app = FastAPI( debug=True, routes=[ - Host(os.environ['TUNNEL_DOMAIN'], base_router, 'base'), - Route('/{path:path}', proxy), - WebSocketRoute('/{path:path}', WebsocketProxy) + WebSocketRoute('/tunnel/', endpoint=tunnel), + Route('/health/', endpoint=health), + Host(f'{{subdomain}}.{os.environ['TUNNEL_DOMAIN']}', app=proxy_app) ] ) -server.post() - - try: from ._version import version __version__ = version diff --git a/ttun_server/endpoints.py b/ttun_server/endpoints.py index 22dcb6d..51c17aa 100644 --- a/ttun_server/endpoints.py +++ b/ttun_server/endpoints.py @@ -1,54 +1,7 @@ -import logging -from base64 import b64decode, b64encode -from uuid import uuid4 +from fastapi import FastAPI -from starlette.background import BackgroundTask -from starlette.requests import Request -from starlette.responses import Response +base_endpoints = FastAPI() -from ttun_server.proxy_queue import ProxyQueue -from ttun_server.types import HttpRequestData, HttpMessageType, HttpMessage - -logger = logging.getLogger(__name__) - - -async def proxy(request: Request) -> Response: - [subdomain, *_] = request.headers['host'].split('.') - identifier = str(uuid4()) - response_queue = await ProxyQueue.create_for_identifier(identifier) - - try: - request_queue = await ProxyQueue.get_for_identifier(subdomain) - - logger.debug('PROXY %s%s ', subdomain, request.url) - await request_queue.enqueue( - HttpMessage( - type=HttpMessageType.request.value, - identifier=identifier, - payload=HttpRequestData( - method=request.method, - path=str(request.url).replace(str(request.base_url), '/'), - headers=list(request.headers.items()), - body=b64encode(await request.body()).decode() - ) - ) - ) - - _response = await response_queue.dequeue() - payload = _response['payload'] - return Response( - status_code=payload['status'], - headers=dict(payload['headers']), - content=b64decode(payload['body'].encode()), - background=BackgroundTask(response_queue.delete) - ) - except AssertionError: - return Response( - content='Not Found', - status_code=404, - background=BackgroundTask(response_queue.delete) - ) - - -async def health(_: Request) -> Response: - return Response(content='OK', status_code=200) +@base_endpoints.get('/health/') +async def health(): + return 'OK' diff --git a/ttun_server/types.py b/ttun_server/types.py index 8591e7d..a0c22a3 100644 --- a/ttun_server/types.py +++ b/ttun_server/types.py @@ -9,6 +9,7 @@ class HttpMessageType(Enum): class Config(TypedDict): + version: str subdomain: str client_version: str diff --git a/ttun_server/websockets.py b/ttun_server/websockets.py index e828b8d..ea9ae94 100644 --- a/ttun_server/websockets.py +++ b/ttun_server/websockets.py @@ -1,189 +1,90 @@ import asyncio -import json import logging import os -import typing -from asyncio import create_task -from base64 import b64encode, b64decode -from contextlib import asynccontextmanager -from typing import Optional from uuid import uuid4 -from starlette.endpoints import WebSocketEndpoint -from starlette.types import Scope, Receive, Send -from starlette.websockets import WebSocket +from fastapi import WebSocket, WebSocketDisconnect import ttun_server from ttun_server.proxy_queue import ProxyQueue -from ttun_server.types import Config, Message, WebsocketMessageType, \ - WebsocketConnectData, WebsocketMessage, WebsocketMessageData, WebsocketDisconnectData, MessageType +from ttun_server.types import ( + Config, + Message, + MessageType, +) logger = logging.getLogger(__name__) logger.setLevel('DEBUG') -class WebsocketProxy(WebSocketEndpoint): - encoding = 'json' - websocket_listen_task = None +async def assert_compatible_version(websocket: WebSocket, config: Config) -> None: + client_version = config.get('version', '1.0.0') + logger.debug('client_version %s', client_version) - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - self.id = str(uuid4()) + if 'git' not in client_version and ttun_server.__version__ != 'development': + [client_major, *_] = [int(i) for i in client_version.split('.')[:3]] + [server_major, *_] = [int(i) for i in ttun_server.__version__.split('.')] - @asynccontextmanager - async def proxy(self, websocket: WebSocket, message: WebsocketMessage): - [subdomain, *_] = websocket.url.hostname.split('.') + if client_major < server_major: + await websocket.close(4000, 'Your client is too old') - expect_ack = WebsocketMessageType(message['type']) == WebsocketMessageType.connect + if client_major > server_major: + await websocket.close(4001, 'Your client is too new') - 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 tunnel(websocket: WebSocket) -> None: + request_tasks: dict[str, asyncio.Task] = {} + proxy_queues: dict[str, ProxyQueue] = {} - async def listen_for_messages(self, websocket: WebSocket): - response_queue = await ProxyQueue.create_for_identifier(self.id) + await websocket.accept() + config: Config = await websocket.receive_json() - 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=websocket.path_params['path'], - headers=[ - (k.decode(), v.decode()) - for k, v - in websocket.scope['headers'] - ], - ) - ) + await assert_compatible_version(websocket, config) - async with self.proxy(websocket, message) as m: - if m is not None and WebsocketMessageType(m['type']) == WebsocketMessageType.ack: - await super().on_connect(websocket) + if 'subdomains' not in config: + config['subdomains'] = [config['subdomain']] + elif config['subdomains'] is None: + config['subdomains'] = [None] - self.websocket_listen_task = asyncio.create_task(self.listen_for_messages(websocket)) + for i, subdomain in enumerate(config['subdomains']): + if subdomain is None or await ProxyQueue.has_connection(subdomain): + config['subdomains'][i] = uuid4().hex - def callback(*args, **kwargs): - self.websocket_listen_task = None + for subdomain in config['subdomains']: + proxy_queues[subdomain] = await ProxyQueue.create_for_identifier(subdomain) - self.websocket_listen_task.add_done_callback(callback) + hostname = os.environ.get('TUNNEL_DOMAIN') + protocol = 'https' if os.environ.get('SECURE', False) else 'http' - 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() + urls = [ + f'{protocol}://{subdomain}.{hostname}' + for subdomain in config['subdomains'] + ] - message = WebsocketMessage( - type=WebsocketMessageType.message.value, - identifier=self.id, - payload=WebsocketMessageData( - body=b64encode(data_bytes).decode(), - ) - ) + await websocket.send_json({ + 'url': urls[0], + 'urls': urls, + }) - async with self.proxy(websocket, message): - pass + async def handle_requests(subdomain: str) -> None: + while request := await proxy_queues[subdomain].dequeue(): + asyncio.create_task(websocket.send_json(request), name=request['identifier']) - 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, - ) + for subdomain in config['subdomains']: + request_tasks[subdomain] = asyncio.create_task( + handle_requests(subdomain), name=subdomain ) - async with self.proxy(websocket, message): - if self.websocket_listen_task is not None: - self.websocket_listen_task.cancel() - -class Tunnel(WebSocketEndpoint): - encoding = 'json' - - def __init__(self, scope: Scope, receive: Receive, send: Send): - super().__init__(scope, receive, send) - self.request_tasks: dict[str, asyncio.Task] = {} - self.config: Optional[Config] = None - self.proxy_queues: dict[str, ProxyQueue] = {} - - async def handle_requests(self, websocket: WebSocket, subdomain: str): - while request := await self.proxy_queues[subdomain].dequeue(): - task = asyncio.create_task(websocket.send_json(request), name=request['identifier']) - - - async def on_connect(self, websocket: WebSocket) -> None: - await websocket.accept() - self.config = await websocket.receive_json() - - client_version = self.config.get('version', '1.0.0') - logger.debug('client_version %s', client_version) - - if 'git' not in client_version and ttun_server.__version__ != 'development': - [client_major, *_] = [int(i) for i in client_version.split('.')[:3]] - [server_major, *_] = [int(i) for i in ttun_server.__version__.split('.')] - - if client_major < server_major: - await websocket.close(4000, 'Your client is too old') - - if client_major > server_major: - await websocket.close(4001, 'Your client is too new') - - if 'subdomains' not in self.config: - self.config['subdomains'] = [self.config['subdomain']] - elif self.config['subdomains'] is None: - self.config['subdomains'] = [None] - - for i, subdomain in enumerate(self.config['subdomains']): - if subdomain is None or await ProxyQueue.has_connection(subdomain): - self.config['subdomains'][i] = uuid4().hex - - for subdomain in self.config['subdomains']: - self.proxy_queues[subdomain] = await ProxyQueue.create_for_identifier(subdomain) - - hostname = os.environ.get("TUNNEL_DOMAIN") - protocol = "https" if os.environ.get("SECURE", False) else "http" - - urls = [ - f'{protocol}://{subdomain}.{hostname}' - for subdomain in self.config['subdomains'] - ] - - await websocket.send_json({ - 'url': urls[0], - 'urls': urls, - }) - - for subdomain in self.config['subdomains']: - self.request_tasks[subdomain] = asyncio.create_task(self.handle_requests(websocket, subdomain), name=subdomain) - - async def on_receive(self, websocket: WebSocket, data: Message): - try: - data['type'] = MessageType(data['type']).value - response_queue = await ProxyQueue.get_for_identifier(data['identifier']) - await response_queue.enqueue(data) - except AssertionError: - pass - - async def on_disconnect(self, websocket: WebSocket, close_code: int): - for proxy_queue in self.proxy_queues.values(): + try: + while True: + data: Message = await websocket.receive_json() + try: + data['type'] = MessageType(data['type']).value + response_queue = await ProxyQueue.get_for_identifier(data['identifier']) + await response_queue.enqueue(data) + except AssertionError: + pass + except WebSocketDisconnect: + for proxy_queue in proxy_queues.values(): await proxy_queue.delete() - - for request_task in self.request_tasks.values(): + for request_task in request_tasks.values(): request_task.cancel() -- cgit v1.2.3