#!/usr/bin/env python3
"""
NixGuard FW Manager v2.1
Agent Server üzerinde çalışır - Port 9444
- FW pairing (SSH key oluşturma, şifre değiştirme)
- FW'lere komut gönderme (SSH üzerinden)
- Lisans proxy
- Aynı IP re-pair desteği (eski key temizleme)
"""

import http.server
import json
import subprocess
import os
import socketserver
import secrets
import string
from datetime import datetime
from urllib import request as urlrequest

PORT = 9444
FW_DATA_DIR = "/var/lib/nixguard-fw"
SSH_KEY_DIR = "/root/.ssh/fw_keys"

os.makedirs(FW_DATA_DIR, exist_ok=True)
os.makedirs(SSH_KEY_DIR, exist_ok=True)

fw_nodes = {}

def load_data():
    global fw_nodes
    try:
        with open(f"{FW_DATA_DIR}/nodes.json", 'r') as f:
            fw_nodes = json.load(f)
    except Exception:
        fw_nodes = {}

def save_data():
    with open(f"{FW_DATA_DIR}/nodes.json", 'w') as f:
        json.dump(fw_nodes, f, indent=2, default=str)

def log(msg):
    ts = datetime.now().strftime('%Y-%m-%d %H:%M:%S')
    print(f"[{ts}] {msg}")
    try:
        with open(f"{FW_DATA_DIR}/manager.log", 'a') as f:
            f.write(f"[{ts}] {msg}\n")
    except Exception:
        pass

def generate_password(length=20):
    chars = string.ascii_letters + string.digits + "!@#$%^&*"
    return ''.join(secrets.choice(chars) for _ in range(length))

def cleanup_old_keys_for_ip(ip):
    """Aynı IP için eski SSH key dosyalarını ve node kaydını temizle"""
    load_data()
    to_remove = []
    for nid, node in fw_nodes.items():
        if node.get("ip") == ip:
            key_path = node.get("ssh_key_path", "")
            if key_path:
                for p in [key_path, f"{key_path}.pub"]:
                    if os.path.exists(p):
                        os.remove(p)
                        log(f"Removed old key: {p}")
            to_remove.append(nid)
    for nid in to_remove:
        del fw_nodes[nid]
        log(f"Removed old node entry: {nid}")
    if to_remove:
        save_data()

def generate_ssh_key(node_id):
    key_path = f"{SSH_KEY_DIR}/fw_{node_id}"
    # Aynı node_id için eski key varsa sil
    for p in [key_path, f"{key_path}.pub"]:
        if os.path.exists(p):
            os.remove(p)
    result = subprocess.run(
        ['ssh-keygen', '-t', 'ed25519', '-f', key_path, '-N', '', '-C', f'nixguard-fw-{node_id}'],
        capture_output=True, text=True
    )
    if result.returncode != 0:
        log(f"SSH keygen failed: {result.stderr}")
        return None, None
    with open(f"{key_path}.pub", 'r') as f:
        pub_key = f.read().strip()
    return key_path, pub_key

def fw_http_request(ip, path, method='GET', data=None, headers=None, pairing_code=None):
    try:
        url = f"http://{ip}:9443{path}"
        req_data = json.dumps(data).encode() if data else None
        req = urlrequest.Request(url, data=req_data, method=method)
        req.add_header('Content-Type', 'application/json')
        if pairing_code:
            req.add_header('X-Pairing-Code', pairing_code)
        if headers:
            for k, v in headers.items():
                req.add_header(k, v)
        resp = urlrequest.urlopen(req, timeout=30)
        return json.loads(resp.read().decode())
    except Exception as e:
        return {"error": str(e)}

def ssh_command(ip, key_path, command):
    try:
        result = subprocess.run(
            ['ssh', '-i', key_path,
             '-o', 'StrictHostKeyChecking=no',
             '-o', 'UserKnownHostsFile=/dev/null',
             '-o', 'ConnectTimeout=10',
             '-o', 'BatchMode=yes',
             f'root@{ip}', command],
            capture_output=True, text=True, timeout=30
        )
        return {
            "success": result.returncode == 0,
            "stdout": result.stdout.strip(),
            "stderr": result.stderr.strip(),
            "code": result.returncode
        }
    except subprocess.TimeoutExpired:
        return {"success": False, "error": "timeout"}
    except Exception as e:
        return {"success": False, "error": str(e)}


class FWManagerHandler(http.server.BaseHTTPRequestHandler):
    def log_message(self, format, *args):
        log(f"{self.client_address[0]} {args[0]}")

    def send_json(self, code, data):
        self.send_response(code)
        self.send_header('Content-Type', 'application/json')
        self.send_header('Access-Control-Allow-Origin', '*')
        self.end_headers()
        self.wfile.write(json.dumps(data).encode())

    def read_body(self):
        length = int(self.headers.get('Content-Length', 0))
        if length > 0:
            return json.loads(self.rfile.read(length).decode())
        return {}

    def do_OPTIONS(self):
        self.send_response(200)
        self.send_header('Access-Control-Allow-Origin', '*')
        self.send_header('Access-Control-Allow-Methods', 'GET, POST, DELETE, OPTIONS')
        self.send_header('Access-Control-Allow-Headers', 'Content-Type, Authorization, X-API-Token, X-Node-ID')
        self.end_headers()

    def do_GET(self):
        path = self.path.split('?')[0]

        if path == '/health':
            self.send_json(200, {
                "status": "ok",
                "service": "nixguard-fw-manager",
                "version": "2.1",
                "nodes": len(fw_nodes),
                "uptime": datetime.now().isoformat()
            })

        elif path == '/fw/nodes':
            load_data()
            nodes_list = []
            for nid, node in fw_nodes.items():
                nodes_list.append({
                    "node_id": nid,
                    "ip": node.get("ip"),
                    "hostname": node.get("hostname"),
                    "wan_ip": node.get("wan_ip"),
                    "lan_ip": node.get("lan_ip"),
                    "status": node.get("status", "unknown"),
                    "paired_at": node.get("paired_at")
                })
            self.send_json(200, {"nodes": nodes_list})

        elif path.startswith('/fw/nodes/') and path.endswith('/status'):
            node_id = path.split('/')[3]
            load_data()
            if node_id not in fw_nodes:
                self.send_json(404, {"error": "Node bulunamadı"})
                return
            node = fw_nodes[node_id]
            key_path = node.get("ssh_key_path")
            ip = node.get("ip")
            if key_path and ip:
                result = ssh_command(ip, key_path, "uptime && hostname && cat /etc/nixguard/config 2>/dev/null | head -5")
                self.send_json(200, {
                    "node_id": node_id,
                    "reachable": result["success"],
                    "output": result.get("stdout", result.get("error", ""))
                })
            else:
                self.send_json(200, {"node_id": node_id, "reachable": False, "error": "SSH key yok"})

        elif path == '/fw/license':
            node_id = self.headers.get('X-Node-ID', '')
            api_token = self.headers.get('X-API-Token', '')
            load_data()
            if node_id not in fw_nodes:
                self.send_json(404, {"error": "Node bulunamadı"})
                return
            node = fw_nodes[node_id]
            try:
                panel_url = node.get("panel_url", "https://mon.nixcon.com.tr")
                req = urlrequest.Request(f"{panel_url}/api/v1/firewall/agent/license", method='GET')
                req.add_header('X-API-Token', api_token)
                req.add_header('X-Node-ID', node_id)
                resp = urlrequest.urlopen(req, timeout=10)
                license_data = json.loads(resp.read().decode())
                self.send_json(200, license_data)
            except Exception as e:
                log(f"License fetch failed: {e}")
                self.send_json(200, {"valid": True, "cached": True})

        else:
            self.send_json(404, {"error": "Not found"})

    def do_POST(self):
        path = self.path.split('?')[0]

        if path == '/fw/pair':
            body = self.read_body()
            ip = body.get('ip', '')
            pairing_code = body.get('pairing_code', '')
            node_id = body.get('node_id', '')
            api_token = body.get('api_token', '')
            panel_url = body.get('panel_url', 'https://mon.nixcon.com.tr')

            if not ip or not pairing_code:
                self.send_json(400, {"error": "IP ve pairing_code zorunludur"})
                return

            log(f"Pairing request: {ip} with code {pairing_code}")

            # 0. Aynı IP için eski SSH key ve node kaydını temizle
            cleanup_old_keys_for_ip(ip)
            # SSH known_hosts'tan da eski entry'yi sil
            subprocess.run(['ssh-keygen', '-R', ip], capture_output=True)

            # 1. FW'den bilgi al (pairing code ile doğrula)
            info = fw_http_request(ip, '/info', pairing_code=pairing_code)
            if info.get('error'):
                err_str = str(info['error']).lower()
                if 'connection' in err_str or 'urlopen' in err_str or 'refused' in err_str:
                    self.send_json(500, {"error": f"FW Node'a erişilemiyor ({ip}:9443). Pairing API çalışıyor mu?"})
                else:
                    self.send_json(401, {"error": info.get('error', 'Bilinmeyen hata')})
                return

            hostname = info.get('hostname', 'nixguard-fw')
            wan_ip = info.get('wan_ip', '')
            lan_ip = info.get('lan_ip', '')
            log(f"FW info: hostname={hostname}, wan={wan_ip}, lan={lan_ip}")

            # 2. SSH key oluştur (node_id bazlı)
            effective_node_id = node_id or ip.replace('.', '_')
            key_path, pub_key = generate_ssh_key(effective_node_id)
            if not key_path:
                self.send_json(500, {"error": "SSH key oluşturulamadı"})
                return
            log(f"SSH key generated: {key_path}")

            # 3. Yeni root password oluştur
            new_password = generate_password(20)

            # 4. FW'ye pair isteği gönder (SSH key + password + güvenlik ayarları)
            setup_data = {
                "pairing_code": pairing_code,
                "ssh_public_key": pub_key,
                "new_root_password": new_password,
                "node_id": effective_node_id,
                "api_token": api_token
            }

            pair_resp = fw_http_request(ip, '/pair', method='POST', data=setup_data)
            if pair_resp.get('error'):
                err_msg = str(pair_resp.get('error', ''))
                if 'Invalid' in err_msg:
                    self.send_json(401, {"error": "Geçersiz Pairing Code"})
                else:
                    self.send_json(500, {"error": f"FW pair hatası: {err_msg}"})
                return

            if not pair_resp.get('success'):
                self.send_json(500, {"error": pair_resp.get('error', 'FW pair başarısız')})
                return

            log("FW pair successful, testing SSH...")

            # 5. SSH bağlantısını test et
            import time
            time.sleep(3)
            test = ssh_command(ip, key_path, "echo 'SSH_OK'")
            ssh_ok = test.get("success", False)
            if not ssh_ok:
                log(f"SSH test failed (may need more time): {test}")

            # 6. Node verilerini kaydet
            fw_nodes[effective_node_id] = {
                "ip": ip,
                "hostname": hostname,
                "wan_ip": wan_ip,
                "lan_ip": lan_ip,
                "ssh_key_path": key_path,
                "root_password": new_password,
                "api_token": api_token,
                "panel_url": panel_url,
                "status": "online" if ssh_ok else "paired",
                "ssh_verified": ssh_ok,
                "paired_at": datetime.now().isoformat()
            }
            save_data()

            log(f"Pairing completed: {hostname} ({ip}) ssh_ok={ssh_ok}")

            self.send_json(200, {
                "success": True,
                "hostname": hostname,
                "wan_ip": wan_ip,
                "lan_ip": lan_ip,
                "root_password": new_password,
                "ssh_verified": ssh_ok,
                "message": f"Node {hostname} başarıyla eşleştirildi"
            })

        elif path.startswith('/fw/nodes/') and path.endswith('/command'):
            node_id = path.split('/')[3]
            load_data()
            if node_id not in fw_nodes:
                self.send_json(404, {"error": "Node bulunamadı"})
                return
            body = self.read_body()
            command = body.get('command', '')
            if not command:
                self.send_json(400, {"error": "command zorunludur"})
                return
            node = fw_nodes[node_id]
            result = ssh_command(node["ip"], node["ssh_key_path"], command)
            # Update last_seen
            fw_nodes[node_id]["last_seen"] = datetime.now().isoformat()
            save_data()
            self.send_json(200, result)

        else:
            self.send_json(404, {"error": "Not found"})

    def do_DELETE(self):
        path = self.path.split('?')[0]
        if path.startswith('/fw/nodes/'):
            node_id = path.split('/')[3]
            load_data()
            if node_id in fw_nodes:
                node = fw_nodes[node_id]
                key_path = node.get("ssh_key_path", "")
                if key_path:
                    for p in [key_path, f"{key_path}.pub"]:
                        if os.path.exists(p):
                            os.remove(p)
                del fw_nodes[node_id]
                save_data()
                self.send_json(200, {"success": True})
            else:
                self.send_json(404, {"error": "Node bulunamadı"})
        else:
            self.send_json(404, {"error": "Not found"})


class ThreadedHTTPServer(socketserver.ThreadingMixIn, http.server.HTTPServer):
    allow_reuse_address = True

if __name__ == '__main__':
    load_data()
    log(f"NixGuard FW Manager v2.1 starting on port {PORT}")
    log(f"Loaded {len(fw_nodes)} nodes")
    server = ThreadedHTTPServer(('0.0.0.0', PORT), FWManagerHandler)
    try:
        server.serve_forever()
    except KeyboardInterrupt:
        log("Shutting down...")
        server.shutdown()
