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")