v0.165.0
  1import asyncio
  2import base64
  3import json
  4import logging
  5import math
  6import struct
  7import time
  8from typing import Any
  9from urllib.parse import urlparse
 10
 11import click
 12import httpx
 13import websockets
 14
 15# Bump this when making breaking changes to the WebSocket protocol.
 16# The server will reject clients with a version lower than its minimum.
 17PROTOCOL_VERSION = 3
 18
 19
 20class TunnelClient:
 21    def __init__(
 22        self, *, destination_url: str, subdomain: str, tunnel_host: str, log_level: str
 23    ) -> None:
 24        self.destination_url = destination_url
 25        self.subdomain = subdomain
 26        self.tunnel_host = tunnel_host
 27
 28        if "localhost" in tunnel_host or "127.0.0.1" in tunnel_host:
 29            self.tunnel_http_url = f"http://{subdomain}.{tunnel_host}"
 30            self.tunnel_websocket_url = (
 31                f"ws://{subdomain}.{tunnel_host}/__tunnel__?v={PROTOCOL_VERSION}"
 32            )
 33        else:
 34            self.tunnel_http_url = f"https://{subdomain}.{tunnel_host}"
 35            self.tunnel_websocket_url = (
 36                f"wss://{subdomain}.{tunnel_host}/__tunnel__?v={PROTOCOL_VERSION}"
 37            )
 38
 39        self.logger = logging.getLogger(__name__)
 40        level = getattr(logging, log_level.upper())
 41        self.logger.setLevel(level)
 42        self.logger.propagate = False
 43        handler = logging.StreamHandler()
 44        handler.setLevel(level)
 45        handler.setFormatter(logging.Formatter("%(message)s"))
 46        self.logger.addHandler(handler)
 47
 48        self.pending_requests: dict[str, dict[str, Any]] = {}
 49        self.active_streams: dict[str, asyncio.Event] = {}
 50        self.proxied_websockets: dict[str, Any] = {}
 51        self.ws_pending_queues: dict[str, asyncio.Queue[dict[str, Any]]] = {}
 52        self.stop_event = asyncio.Event()
 53
 54    async def connect(self) -> None:
 55        retry_delay = 1.0
 56        max_retry_delay = 30.0
 57        # Connection must stay up at least this long to be considered healthy
 58        # and reset the backoff. Otherwise we keep escalating, which prevents
 59        # tight reconnect loops when something else (e.g. another client
 60        # claiming the same subdomain) keeps closing us right after connect.
 61        healthy_connection_seconds = 5.0
 62        while not self.stop_event.is_set():
 63            connection_duration: float | None = None
 64            try:
 65                self.logger.debug(
 66                    f"Connecting to WebSocket URL: {self.tunnel_websocket_url}"
 67                )
 68                async with websockets.connect(
 69                    self.tunnel_websocket_url, max_size=None
 70                ) as websocket:
 71                    self.logger.debug("WebSocket connection established")
 72                    click.secho(
 73                        f"Connected to tunnel {self.tunnel_http_url}", fg="green"
 74                    )
 75                    connected_at = time.monotonic()
 76                    try:
 77                        await self.handle_messages(websocket)
 78                    finally:
 79                        connection_duration = time.monotonic() - connected_at
 80                        await self._cleanup_proxied_websockets()
 81                if self.stop_event.is_set():
 82                    break
 83                disconnect_message = "Tunnel disconnected by server."
 84            except asyncio.CancelledError:
 85                self.logger.debug("Connection cancelled")
 86                break
 87            except websockets.InvalidStatus as e:
 88                if e.response.status_code == 426:
 89                    body = e.response.body.decode() if e.response.body else ""
 90                    click.secho(
 91                        body or "Client version too old. Please upgrade plain.tunnel.",
 92                        fg="red",
 93                    )
 94                    break
 95                raise
 96            except (websockets.ConnectionClosed, ConnectionError) as e:
 97                if self.stop_event.is_set():
 98                    self.logger.debug("Stopping reconnect attempts due to shutdown")
 99                    break
100                disconnect_message = f"Connection lost: {e}."
101            except Exception as e:
102                if self.stop_event.is_set():
103                    self.logger.debug("Stopping reconnect attempts due to shutdown")
104                    break
105                disconnect_message = f"Unexpected error: {e}."
106
107            if (
108                connection_duration is not None
109                and connection_duration >= healthy_connection_seconds
110            ):
111                retry_delay = 1.0
112            click.secho(
113                f"{disconnect_message} Retrying in {retry_delay:.0f}s...",
114                fg="yellow",
115            )
116            await asyncio.sleep(retry_delay)
117            retry_delay = min(retry_delay * 2, max_retry_delay)
118
119    async def handle_messages(self, websocket: Any) -> None:
120        try:
121            async for message in websocket:
122                if isinstance(message, str):
123                    data = json.loads(message)
124                    msg_type = data.get("type")
125                    if msg_type == "ping":
126                        self.logger.debug("Received heartbeat ping, sending pong")
127                        await websocket.send(json.dumps({"type": "pong"}))
128                    elif msg_type == "request":
129                        self.logger.debug("Received request metadata from worker")
130                        await self.handle_request_metadata(websocket, data)
131                    elif msg_type == "stream-cancel":
132                        request_id = data.get("id")
133                        self.logger.debug(
134                            f"Received stream-cancel for request ID: {request_id}"
135                        )
136                        cancel_event = self.active_streams.get(request_id)
137                        if cancel_event:
138                            cancel_event.set()
139                    elif msg_type == "ws-open":
140                        self.logger.debug(f"Received ws-open for ID: {data['id']}")
141                        self.ws_pending_queues[data["id"]] = asyncio.Queue()
142                        task = asyncio.create_task(
143                            self._handle_ws_open(websocket, data)
144                        )
145                        task.add_done_callback(self._handle_task_exception)
146                    elif msg_type == "ws-message":
147                        await self._handle_ws_message(data)
148                    elif msg_type == "ws-close":
149                        await self._handle_ws_close(data)
150                    else:
151                        self.logger.warning(
152                            f"Received unknown message type: {msg_type}"
153                        )
154                elif isinstance(message, bytes):
155                    self.logger.debug("Received binary data from worker")
156                    await self.handle_request_body_chunk(websocket, message)
157                else:
158                    self.logger.warning("Received unknown message format")
159        except asyncio.CancelledError:
160            self.logger.debug("Message handling cancelled")
161        except Exception as e:
162            self.logger.error(f"Error in handle_messages: {e}")
163            raise
164
165    async def handle_request_metadata(
166        self, websocket: Any, data: dict[str, Any]
167    ) -> None:
168        request_id = data["id"]
169        has_body = data.get("has_body", False)
170        total_body_chunks = data.get("totalBodyChunks", 0)
171        self.pending_requests[request_id] = {
172            "metadata": data,
173            "body_chunks": {},
174            "has_body": has_body,
175            "total_body_chunks": total_body_chunks,
176        }
177        self.logger.debug(
178            f"Stored metadata for request ID: {request_id}, has_body: {has_body}"
179        )
180        await self.check_and_process_request(websocket, request_id)
181
182    async def handle_request_body_chunk(
183        self, websocket: Any, chunk_data: bytes
184    ) -> None:
185        (id_length,) = struct.unpack_from("<I", chunk_data, 0)
186        request_id = chunk_data[4 : 4 + id_length].decode("utf-8")
187        header_end = 4 + id_length + 8
188        chunk_index, total_chunks = struct.unpack_from("<II", chunk_data, 4 + id_length)
189        body_chunk = chunk_data[header_end:]
190
191        if request_id in self.pending_requests:
192            request = self.pending_requests[request_id]
193            request["body_chunks"][chunk_index] = body_chunk
194            self.logger.debug(
195                f"Stored body chunk {chunk_index + 1}/{total_chunks} for request ID: {request_id}"
196            )
197            await self.check_and_process_request(websocket, request_id)
198        else:
199            self.logger.warning(
200                f"Received body chunk for unknown or completed request ID: {request_id}"
201            )
202
203    async def check_and_process_request(self, websocket: Any, request_id: str) -> None:
204        request_data = self.pending_requests.get(request_id)
205        if not request_data:
206            return
207
208        has_body = request_data["has_body"]
209        total_body_chunks = request_data["total_body_chunks"]
210        body_chunks = request_data["body_chunks"]
211
212        all_chunks_received = not has_body or len(body_chunks) == total_body_chunks
213        if not all_chunks_received:
214            return
215
216        for i in range(total_body_chunks):
217            if i not in body_chunks:
218                self.logger.error(
219                    f"Missing chunk {i + 1}/{total_body_chunks} for request ID: {request_id}"
220                )
221                return
222
223        self.logger.debug(f"Processing request ID: {request_id}")
224        del self.pending_requests[request_id]
225        task = asyncio.create_task(
226            self.process_request(
227                websocket,
228                request_data["metadata"],
229                body_chunks,
230                request_id,
231            )
232        )
233        task.add_done_callback(self._handle_task_exception)
234
235    def _handle_task_exception(self, task: asyncio.Task[None]) -> None:
236        if not task.cancelled() and task.exception():
237            self.logger.error("Error processing request", exc_info=task.exception())
238
239    async def _cleanup_proxied_websockets(self) -> None:
240        """Close all proxied WebSocket connections on tunnel disconnect."""
241        for ws_id, ws in list(self.proxied_websockets.items()):
242            try:
243                await ws.close()
244            except Exception:
245                pass
246        self.proxied_websockets.clear()
247        self.ws_pending_queues.clear()
248
249    async def process_request(
250        self,
251        websocket: Any,
252        request_metadata: dict[str, Any],
253        body_chunks: dict[int, bytes],
254        request_id: str,
255    ) -> None:
256        self.logger.debug(
257            f"Processing request: {request_id} {request_metadata['method']} {request_metadata['url']}"
258        )
259
260        if request_metadata["has_body"]:
261            total_chunks = request_metadata["totalBodyChunks"]
262            body_data = b"".join(body_chunks[i] for i in range(total_chunks))
263        else:
264            body_data = None
265
266        parsed = urlparse(request_metadata["url"])
267        path = parsed.path
268        if parsed.query:
269            path = f"{path}?{parsed.query}"
270        forward_url = f"{self.destination_url}{path}"
271
272        self.logger.debug(f"Forwarding request to: {forward_url}")
273
274        async with httpx.AsyncClient(
275            follow_redirects=False, verify=False, timeout=30
276        ) as client:
277            try:
278                async with client.stream(
279                    method=request_metadata["method"],
280                    url=forward_url,
281                    headers=request_metadata["headers"],
282                    content=body_data,
283                ) as response:
284                    response_status = response.status_code
285                    response_headers = dict(response.headers)
286
287                    self.logger.info(
288                        f"{click.style(request_metadata['method'], bold=True)} {request_metadata['url']} {response_status}"
289                    )
290
291                    if self._is_streaming_response(response):
292                        await self._handle_streaming_response(
293                            websocket,
294                            response,
295                            request_id,
296                            response_status,
297                            response_headers,
298                        )
299                    else:
300                        await response.aread()
301                        await self._handle_buffered_response(
302                            websocket,
303                            response.content,
304                            request_id,
305                            response_status,
306                            response_headers,
307                        )
308            except httpx.ConnectError as e:
309                self.logger.error(f"Connection error forwarding request: {e}")
310                self.logger.info(
311                    f"{click.style(request_metadata['method'], bold=True)} {request_metadata['url']} 502"
312                )
313                await self._handle_buffered_response(
314                    websocket, b"", request_id, 502, {}
315                )
316
317    def _is_streaming_response(self, response: httpx.Response) -> bool:
318        content_type = response.headers.get("content-type", "")
319        return "text/event-stream" in content_type
320
321    async def _handle_buffered_response(
322        self,
323        websocket: Any,
324        response_body: bytes,
325        request_id: str,
326        response_status: int,
327        response_headers: dict[str, str],
328    ) -> None:
329        has_body = len(response_body) > 0
330        max_chunk_size = 1_000_000
331        total_body_chunks = (
332            math.ceil(len(response_body) / max_chunk_size) if has_body else 0
333        )
334
335        response_metadata = {
336            "type": "response",
337            "id": request_id,
338            "status": response_status,
339            "headers": list(response_headers.items()),
340            "has_body": has_body,
341            "totalBodyChunks": total_body_chunks,
342        }
343
344        self.logger.debug(
345            f"Sending response metadata for ID: {request_id}, has_body: {has_body}"
346        )
347        await websocket.send(json.dumps(response_metadata))
348
349        if has_body:
350            self.logger.debug(
351                f"Sending {total_body_chunks} body chunks for ID: {request_id}"
352            )
353            id_bytes = request_id.encode("utf-8")
354            for i in range(total_body_chunks):
355                chunk_start = i * max_chunk_size
356                chunk_end = min(chunk_start + max_chunk_size, len(response_body))
357                header = id_bytes + struct.pack("<II", i, total_body_chunks)
358                await websocket.send(header + response_body[chunk_start:chunk_end])
359                self.logger.debug(
360                    f"Sent body chunk {i + 1}/{total_body_chunks} for ID: {request_id}"
361                )
362
363    async def _handle_streaming_response(
364        self,
365        websocket: Any,
366        response: httpx.Response,
367        request_id: str,
368        response_status: int,
369        response_headers: dict[str, str],
370    ) -> None:
371        cancel_event = asyncio.Event()
372        self.active_streams[request_id] = cancel_event
373
374        stream_start = {
375            "type": "stream-start",
376            "id": request_id,
377            "status": response_status,
378            "headers": list(response_headers.items()),
379        }
380
381        self.logger.debug(f"Sending stream-start for ID: {request_id}")
382        await websocket.send(json.dumps(stream_start))
383
384        id_bytes = request_id.encode("utf-8")
385
386        try:
387            async for chunk in response.aiter_bytes():
388                if cancel_event.is_set():
389                    self.logger.debug(
390                        f"Stream cancelled by browser for request ID: {request_id}"
391                    )
392                    break
393
394                await websocket.send(id_bytes + chunk)
395            else:
396                # Only send stream-end if the loop completed naturally
397                # (not cancelled by the server via stream-cancel)
398                stream_end = {
399                    "type": "stream-end",
400                    "id": request_id,
401                }
402                self.logger.debug(f"Sending stream-end for ID: {request_id}")
403                await websocket.send(json.dumps(stream_end))
404        except Exception as e:
405            self.logger.error(f"Error streaming response for ID {request_id}: {e}")
406            stream_error = {
407                "type": "stream-error",
408                "id": request_id,
409                "error": str(e),
410            }
411            try:
412                await websocket.send(json.dumps(stream_error))
413            except Exception:
414                pass
415        finally:
416            self.active_streams.pop(request_id, None)
417
418    async def _handle_ws_open(self, tunnel_ws: Any, data: dict[str, Any]) -> None:
419        ws_id = data["id"]
420        url = data["url"]
421        parsed = urlparse(url)
422        path = parsed.path
423        if parsed.query:
424            path = f"{path}?{parsed.query}"
425
426        # Build local WebSocket URL
427        dest_parsed = urlparse(self.destination_url)
428        if dest_parsed.scheme == "https":
429            ws_scheme = "wss"
430        else:
431            ws_scheme = "ws"
432        local_ws_url = f"{ws_scheme}://{dest_parsed.netloc}{path}"
433
434        self.logger.debug(f"Opening local WebSocket for {ws_id}: {local_ws_url}")
435
436        # Forward safe browser headers (cookies, auth, origin) to the local
437        # server. Skip hop-by-hop and WebSocket handshake headers since
438        # websockets.connect generates its own (including Host from the URL).
439        skip_headers = frozenset(
440            {
441                "host",
442                "connection",
443                "upgrade",
444                "sec-websocket-key",
445                "sec-websocket-version",
446                "sec-websocket-extensions",
447                "sec-websocket-protocol",
448            }
449        )
450        forward_headers = {}
451        for name, value in data.get("headers", {}).items():
452            if name.lower() not in skip_headers:
453                forward_headers[name] = value
454
455        try:
456            import ssl
457
458            ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
459            ssl_context.check_hostname = False
460            ssl_context.verify_mode = ssl.CERT_NONE
461            local_ws = await websockets.connect(
462                local_ws_url,
463                ssl=ssl_context if ws_scheme == "wss" else None,
464                max_size=None,
465                additional_headers=forward_headers,
466            )
467        except Exception as e:
468            self.logger.error(f"Failed to connect local WebSocket for {ws_id}: {e}")
469            self.ws_pending_queues.pop(ws_id, None)
470            try:
471                await tunnel_ws.send(
472                    json.dumps(
473                        {
474                            "type": "ws-close",
475                            "id": ws_id,
476                            "code": 1011,
477                            "reason": str(e),
478                        }
479                    )
480                )
481            except Exception:
482                pass
483            return
484
485        self.proxied_websockets[ws_id] = local_ws
486
487        # Drain any messages that arrived while connecting
488        queue = self.ws_pending_queues.pop(ws_id, None)
489        if queue is not None:
490            while not queue.empty():
491                queued = queue.get_nowait()
492                try:
493                    if queued.get("binary"):
494                        await local_ws.send(base64.b64decode(queued["data"]))
495                    else:
496                        await local_ws.send(queued["data"])
497                except Exception as e:
498                    self.logger.error(
499                        f"Failed to forward queued message to local WebSocket {ws_id}: {e}"
500                    )
501
502        self.logger.info(f"WebSocket proxy opened: {ws_id} -> {local_ws_url}")
503
504        # Relay messages from local server back to the tunnel
505        try:
506            async for message in local_ws:
507                if isinstance(message, str):
508                    await tunnel_ws.send(
509                        json.dumps({"type": "ws-message", "id": ws_id, "data": message})
510                    )
511                elif isinstance(message, bytes):
512                    await tunnel_ws.send(
513                        json.dumps(
514                            {
515                                "type": "ws-message",
516                                "id": ws_id,
517                                "data": base64.b64encode(message).decode("ascii"),
518                                "binary": True,
519                            }
520                        )
521                    )
522        except websockets.ConnectionClosed:
523            pass
524        except Exception as e:
525            self.logger.error(f"Error relaying WebSocket {ws_id}: {e}")
526        finally:
527            self.proxied_websockets.pop(ws_id, None)
528            close_code = local_ws.close_code or 1000
529            close_reason = local_ws.close_reason or ""
530            try:
531                await tunnel_ws.send(
532                    json.dumps(
533                        {
534                            "type": "ws-close",
535                            "id": ws_id,
536                            "code": close_code,
537                            "reason": close_reason,
538                        }
539                    )
540                )
541            except Exception:
542                pass
543
544    async def _handle_ws_message(self, data: dict[str, Any]) -> None:
545        ws_id = data["id"]
546        local_ws = self.proxied_websockets.get(ws_id)
547        if not local_ws:
548            # Connection still being established — buffer for later
549            queue = self.ws_pending_queues.get(ws_id)
550            if queue is not None:
551                await queue.put(data)
552                return
553            self.logger.warning(f"Received ws-message for unknown WebSocket: {ws_id}")
554            return
555        try:
556            if data.get("binary"):
557                await local_ws.send(base64.b64decode(data["data"]))
558            else:
559                await local_ws.send(data["data"])
560        except Exception as e:
561            self.logger.error(
562                f"Failed to forward message to local WebSocket {ws_id}: {e}"
563            )
564
565    async def _handle_ws_close(self, data: dict[str, Any]) -> None:
566        ws_id = data["id"]
567        local_ws = self.proxied_websockets.pop(ws_id, None)
568        if not local_ws:
569            return
570        try:
571            await local_ws.close()
572        except Exception:
573            pass
574
575    def run(self) -> None:
576        try:
577            asyncio.run(self.connect())
578        except KeyboardInterrupt:
579            self.logger.debug("Received exit signal")