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 --- ttun_server/__init__.py | 23 ++--- ttun_server/endpoints.py | 57 ++---------- ttun_server/types.py | 1 + ttun_server/websockets.py | 217 +++++++++++++--------------------------------- 4 files changed, 73 insertions(+), 225 deletions(-) (limited to 'ttun_server') 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