#!/usr/bin/env python3
"""
API Web Sentinel COMPLETE - Utilise le VRAI scanner CLI
Supporte 24+ langages et 84+ extensions comme le .exe
"""

import sys
import os

# Charger les variables d'environnement depuis .env
from dotenv import load_dotenv
load_dotenv()

# Ajouter le module web_sentinel au path
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))

from flask import Flask, request, jsonify
from flask_cors import CORS
import sqlite3
import psycopg2
from datetime import datetime
import logging
from pathlib import Path
import tempfile

# Configuration PostgreSQL
DB_CONFIG = {
    'host': 'web-sentinel-db',
    'port': 5432,
    'database': 'websentinel_prod',
    'user': 'websentinel_user',
    'password': 'WebSentinelDB2025!'
}

# Importer le VRAI scanner CLI
from web_sentinel.checks.source_code.scanner import SourceCodeScanner, SourceScanConfig, LANGUAGE_EXTENSIONS
from web_sentinel.model import ScanRequest

# Importer le Blueprint de paiement
from web_sentinel.payment.api_routes import payment_bp
from web_sentinel.payment.config import StripeConfig

# Importer le module d'authentification
from web_sentinel.user_auth import UserManager, require_auth

# Configuration Flask
app = Flask(__name__)
# CORS : Autoriser TOUS les domaines pour debug
CORS(app, resources={r"/api/*": {"origins": "*"}})

# Logging
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)

# Enregistrer le Blueprint de paiement
app.register_blueprint(payment_bp)

# CORS manuel car Flask-CORS ne fonctionne pas (probablement Traefik qui supprime les headers)
@app.after_request
def add_cors_headers(response):
    """Ajoute les headers CORS manuellement à TOUTES les réponses"""
    response.headers['Access-Control-Allow-Origin'] = '*'
    response.headers['Access-Control-Allow-Methods'] = 'GET, POST, PUT, DELETE, OPTIONS'
    response.headers['Access-Control-Allow-Headers'] = 'Content-Type, X-API-Key, Authorization, Stripe-Signature'
    response.headers['Access-Control-Max-Age'] = '3600'
    return response

# NOTE: Les requêtes OPTIONS sont gérées automatiquement par @app.after_request
# qui ajoute les headers CORS nécessaires. Pas besoin de route catch-all.

# 🔐 Helper pour authentification PostgreSQL
def get_user_by_api_key(api_key):
    """
    Récupère un utilisateur depuis PostgreSQL par son API key
    Retourne None si non trouvé ou invalide
    """
    try:
        conn = psycopg2.connect(**DB_CONFIG)
        cursor = conn.cursor()
        cursor.execute("""
            SELECT id, email, name, tier, is_active 
            FROM users 
            WHERE api_key = %s
        """, (api_key,))
        row = cursor.fetchone()
        cursor.close()
        conn.close()
        
        if row:
            return {
                "user_id": row[0],
                "email": row[1],
                "name": row[2],
                "tier": row[3],
                "active": row[4]
            }
        return None
    except Exception as e:
        logger.error(f"Erreur auth PostgreSQL: {e}")
        return None

def save_scan_to_postgres(user_id, user_email, target, scan_type, findings, duration):
    """
    Enregistrer un scan dans PostgreSQL avec ses résultats
    
    Args:
        user_id: ID de l'utilisateur (api_key_id)
        user_email: Email de l'utilisateur (pour logs)
        target: Cible du scan (filename, URL, etc.)
        scan_type: Type de scan (source_scan, web_scan, sast, etc.)
        findings: Liste des findings (dict avec severity, title, etc.)
        duration: Durée du scan en secondes
    
    Returns:
        scan_id si succès, None si erreur
    """
    try:
        conn = psycopg2.connect(**DB_CONFIG)
        cursor = conn.cursor()
        
        # Compter les sévérités
        severity_counts = {'critical': 0, 'high': 0, 'medium': 0, 'low': 0, 'info': 0}
        for finding in findings:
            sev = finding.get('severity', 'low').lower()
            if sev in severity_counts:
                severity_counts[sev] += 1
        
        # Insérer le scan dans PostgreSQL
        cursor.execute('''
            INSERT INTO scans (api_key_id, scan_type, target, findings_count,
                             severity_critical, severity_high, severity_medium, 
                             severity_low, severity_info, duration, timestamp)
            VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, CURRENT_TIMESTAMP)
            RETURNING id
        ''', (
            user_id,
            scan_type,
            target,
            len(findings),
            severity_counts['critical'],
            severity_counts['high'],
            severity_counts['medium'],
            severity_counts['low'],
            severity_counts['info'],
            duration
        ))
        
        scan_id = cursor.fetchone()[0]
        
        conn.commit()
        cursor.close()
        conn.close()
        
        logger.info(f"✅ Scan enregistré PostgreSQL: ID={scan_id}, user={user_email}, findings={len(findings)}")
        return scan_id
        
    except Exception as e:
        logger.error(f"❌ Erreur enregistrement scan PostgreSQL: {e}", exc_info=True)
        return None

class HonestStatsManager:
    """Gestionnaire de statistiques 100% honnêtes"""
    
    def __init__(self, db_path="/app/data/stats.db"):
        self.db_path = db_path
        # Créer le répertoire si nécessaire
        os.makedirs(os.path.dirname(self.db_path), exist_ok=True)
        self.init_database()
    
    def init_database(self):
        """Initialiser la base de données des statistiques"""
        conn = sqlite3.connect(self.db_path)
        cursor = conn.cursor()
        
        # Table des scans réalisés
        cursor.execute('''
            CREATE TABLE IF NOT EXISTS scans (
                id INTEGER PRIMARY KEY AUTOINCREMENT,
                user_email TEXT,
                target TEXT,
                scan_type TEXT,
                findings_count INTEGER,
                severity_critical INTEGER DEFAULT 0,
                severity_high INTEGER DEFAULT 0,
                severity_medium INTEGER DEFAULT 0,
                severity_low INTEGER DEFAULT 0,
                duration_seconds REAL,
                created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
            )
        ''')
        
        # Table des vulnérabilités détectées
        cursor.execute('''
            CREATE TABLE IF NOT EXISTS vulnerabilities (
                id INTEGER PRIMARY KEY AUTOINCREMENT,
                scan_id INTEGER,
                vuln_type TEXT,
                severity TEXT,
                title TEXT,
                description TEXT,
                found_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
                FOREIGN KEY (scan_id) REFERENCES scans (id)
            )
        ''')
        
        conn.commit()
        conn.close()
    
    def add_scan_result(self, user_email, target, scan_type, findings, duration):
        """Enregistrer un scan avec ses résultats"""
        conn = sqlite3.connect(self.db_path)
        cursor = conn.cursor()
        
        # Compter les sévérités
        severity_counts = {'critical': 0, 'high': 0, 'medium': 0, 'low': 0}
        for finding in findings:
            sev = finding.get('severity', 'low').lower()
            if sev in severity_counts:
                severity_counts[sev] += 1
        
        cursor.execute('''
            INSERT INTO scans (user_email, target, scan_type, findings_count, 
                             severity_critical, severity_high, severity_medium, severity_low,
                             duration_seconds)
            VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
        ''', (user_email, target, scan_type, len(findings),
              severity_counts['critical'], severity_counts['high'],
              severity_counts['medium'], severity_counts['low'],
              duration))
        
        scan_id = cursor.lastrowid
        
        # Enregistrer chaque vulnérabilité
        for finding in findings:
            cursor.execute('''
                INSERT INTO vulnerabilities (scan_id, vuln_type, severity, title, description)
                VALUES (?, ?, ?, ?, ?)
            ''', (scan_id, finding.get('check', 'unknown'), finding.get('severity', 'low'),
                  finding.get('title', ''), finding.get('description', '')))
        
        conn.commit()
        conn.close()
        
        return scan_id
    
    def add_telemetry_scan(self, findings_count, scan_duration, platform="unknown", app_version="1.0.0", 
                          severity_counts=None, vulnerability_types=None, languages_detected=None, 
                          findings_by_language=None, vulnerabilities_by_lang=None):
        """
        Enregistrer un scan anonyme depuis la télémétrie.
        
        Args:
            findings_count: Nombre de findings détectés
            scan_duration: Durée du scan en secondes
            platform: Plateforme (windows-gui, linux-cli, etc.)
            app_version: Version de l'application
            severity_counts: Dict avec comptage par sévérité (ex: {"critical": 2, "high": 5})
            vulnerability_types: Dict avec comptage par type (ex: {"SQL Injection": 3, "XSS": 2})
            languages_detected: Liste des langages détectés (ex: ["python", "javascript"])
            findings_by_language: Dict avec comptage par langage (ex: {"php": 2, "javascript": 5})
            vulnerabilities_by_lang: Mapping précis { langage: { type_vuln: count } }
        """
        conn = sqlite3.connect(self.db_path)
        cursor = conn.cursor()
        
        try:
            # Préparer les comptages de sévérité
            severity_critical = severity_counts.get('critical', 0) if severity_counts else 0
            severity_high = severity_counts.get('high', 0) if severity_counts else 0
            severity_medium = severity_counts.get('medium', 0) if severity_counts else 0
            severity_low = severity_counts.get('low', 0) if severity_counts else 0
            
            # 1. Enregistrer le scan avec les sévérités
            cursor.execute('''
                INSERT INTO scans (user_email, target, scan_type, findings_count, 
                                 severity_critical, severity_high, severity_medium, severity_low,
                                 duration_seconds, created_at)
                VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP)
            ''', (
                f"telemetry-{platform}",  # Email anonyme identifiant la plateforme
                f"v{app_version}",        # Target = version pour traçage
                "telemetry",              # Type de scan
                findings_count,
                severity_critical,
                severity_high,
                severity_medium,
                severity_low,
                scan_duration
            ))
            
            scan_id = cursor.lastrowid
            
            # 2. Enregistrer les vulnérabilités avec types détaillés ET langages réels
            if vulnerabilities_by_lang and findings_count > 0:
                # 🆕 Utiliser le mapping précis { langage: { type_vuln: count } }
                for lang, vuln_counts in vulnerabilities_by_lang.items():
                    for vuln_title, count in vuln_counts.items():
                        # Déterminer la sévérité depuis severity_counts si disponible
                        severity = "medium"  # défaut
                        if severity_counts:
                            for sev in ["critical", "high", "medium", "low", "info"]:
                                if severity_counts.get(sev, 0) > 0:
                                    severity = sev
                                    break
                        
                        # Enregistrer chaque vulnérabilité avec son VRAI langage
                        for _ in range(count):
                            description = f"{vuln_title} | Language: {lang} | Platform: {platform}"
                            
                            cursor.execute('''
                                INSERT INTO vulnerabilities (scan_id, vuln_type, severity, title, description)
                                VALUES (?, ?, ?, ?, ?)
                            ''', (
                                scan_id,
                                vuln_title,  # Type réel de la vulnérabilité
                                severity,
                                vuln_title,
                                description
                            ))
            
            elif vulnerability_types and findings_count > 0:
                # Fallback si findings_by_language n'est pas fourni (ancien code)
                for vuln_title, count in vulnerability_types.items():
                    severity = "medium"
                    if severity_counts:
                        for sev in ["critical", "high", "medium", "low", "info"]:
                            if severity_counts.get(sev, 0) > 0:
                                severity = sev
                                break
                    
                    for _ in range(count):
                        detected_lang = languages_detected[0] if languages_detected else "unknown"
                        description = f"{vuln_title} | Language: {detected_lang} | Platform: {platform}"
                        
                        cursor.execute('''
                            INSERT INTO vulnerabilities (scan_id, vuln_type, severity, title, description)
                            VALUES (?, ?, ?, ?, ?)
                        ''', (
                            scan_id,
                            vuln_title,
                            severity,
                            vuln_title,
                            description
                        ))
            
            elif severity_counts and findings_count > 0:
                # Fallback si pas de vulnerability_types (ancienne version)
                for severity, count in severity_counts.items():
                    for _ in range(count):
                        cursor.execute('''
                            INSERT INTO vulnerabilities (scan_id, vuln_type, severity, title, description)
                            VALUES (?, ?, ?, ?, ?)
                        ''', (
                            scan_id,
                            "telemetry",
                            severity,
                            f"Telemetry finding ({severity})",
                            f"Anonymous finding from {platform}"
                        ))
            
            conn.commit()
            logger.info(f"✅ Télémétrie enregistrée: {findings_count} findings, {scan_duration}s, {platform}")
            
        except Exception as e:
            logger.error(f"❌ Erreur enregistrement télémétrie: {e}")
            conn.rollback()
        finally:
            conn.close()
    
    def get_real_statistics(self):
        """Récupérer les VRAIES statistiques"""
        conn = sqlite3.connect(self.db_path)
        cursor = conn.cursor()
        
        try:
            cursor.execute("SELECT COUNT(*) FROM scans")
            total_scans = cursor.fetchone()[0]
            
            cursor.execute("SELECT COUNT(*) FROM vulnerabilities")
            total_vulns = cursor.fetchone()[0]
            
            cursor.execute("SELECT COUNT(DISTINCT user_email) FROM scans")
            active_users = cursor.fetchone()[0]
            
            cursor.execute("SELECT AVG(duration_seconds) FROM scans")
            avg_duration = cursor.fetchone()[0] or 0
            
            return {
                "total_scans": total_scans,
                "total_vulnerabilities_found": total_vulns,
                "active_users": active_users,
                "average_scan_duration": round(avg_duration, 2),
                "last_updated": datetime.utcnow().isoformat()
            }
        finally:
            conn.close()
    
    def get_detailed_statistics(self):
        """Récupérer des statistiques détaillées pour la page stats"""
        conn = sqlite3.connect(self.db_path)
        cursor = conn.cursor()
        
        try:
            # Stats globales
            cursor.execute("SELECT COUNT(*) FROM scans")
            total_scans = cursor.fetchone()[0]
            
            cursor.execute("SELECT COUNT(*) FROM vulnerabilities")
            total_vulns = cursor.fetchone()[0]
            
            cursor.execute("SELECT COUNT(DISTINCT user_email) FROM scans")
            total_users = cursor.fetchone()[0]
            
            cursor.execute("SELECT AVG(duration_seconds) FROM scans")
            avg_duration = cursor.fetchone()[0] or 0
            
            # Top 10 des types de vulnérabilités
            cursor.execute("""
                SELECT vuln_type, COUNT(*) as count, severity
                FROM vulnerabilities
                GROUP BY vuln_type, severity
                ORDER BY count DESC
                LIMIT 10
            """)
            top_vulnerabilities = [
                {"type": row[0], "count": row[1], "severity": row[2]}
                for row in cursor.fetchall()
            ]
            
            # Distribution par sévérité (depuis les vulnérabilités réelles)
            cursor.execute("""
                SELECT severity, COUNT(*) as count
                FROM vulnerabilities
                WHERE vuln_type != 'telemetry'
                GROUP BY severity
                ORDER BY 
                    CASE severity
                        WHEN 'critical' THEN 1
                        WHEN 'high' THEN 2
                        WHEN 'medium' THEN 3
                        WHEN 'low' THEN 4
                        WHEN 'info' THEN 5
                        ELSE 6
                    END
            """)
            severity_distribution = {row[0]: row[1] for row in cursor.fetchall()}
            
            # Stats par sévérité consolidées (total réel depuis vulnerabilities)
            severity_totals = (
                severity_distribution.get('critical', 0),
                severity_distribution.get('high', 0),
                severity_distribution.get('medium', 0),
                severity_distribution.get('low', 0)
            )
            
            # Scans récents (7 derniers jours) - avec jours manquants remplis à 0
            cursor.execute("""
                WITH RECURSIVE dates(date) AS (
                    SELECT DATE('now', '-6 days')
                    UNION ALL
                    SELECT DATE(date, '+1 day')
                    FROM dates
                    WHERE date < DATE('now')
                )
                SELECT dates.date, COALESCE(COUNT(scans.id), 0) as count
                FROM dates
                LEFT JOIN scans ON DATE(scans.created_at) = dates.date
                GROUP BY dates.date
                ORDER BY dates.date ASC
            """)
            recent_scans = [
                {"date": row[0], "count": row[1]}
                for row in cursor.fetchall()
            ]
            
            # Langages les plus scannés (extrait depuis description des vulnérabilités)
            cursor.execute("""
                SELECT 
                    CASE 
                        WHEN description LIKE '%Language: python%' THEN 'Python'
                        WHEN description LIKE '%Language: javascript%' THEN 'JavaScript'
                        WHEN description LIKE '%Language: php%' THEN 'PHP'
                        WHEN description LIKE '%Language: java%' THEN 'Java'
                        WHEN description LIKE '%Language: typescript%' THEN 'TypeScript'
                        WHEN description LIKE '%Language: csharp%' THEN 'C#'
                        WHEN description LIKE '%Language: go%' THEN 'Go'
                        WHEN description LIKE '%Language: ruby%' THEN 'Ruby'
                        WHEN description LIKE '%Language: rust%' THEN 'Rust'
                        WHEN description LIKE '%Language: c++%' THEN 'C++'
                        WHEN description LIKE '%Language: c%' THEN 'C'
                        WHEN description LIKE '%Language: swift%' THEN 'Swift'
                        WHEN description LIKE '%Language: kotlin%' THEN 'Kotlin'
                        WHEN description LIKE '%Language: scala%' THEN 'Scala'
                        WHEN description LIKE '%Language: perl%' THEN 'Perl'
                        WHEN description LIKE '%Language: lua%' THEN 'Lua'
                        WHEN description LIKE '%Language: auto%' THEN 'Auto-détecté'
                        WHEN description LIKE '%Language: unknown%' THEN 'Inconnu'
                        ELSE 'Autre'
                    END as language,
                    COUNT(*) as count
                FROM vulnerabilities
                WHERE description LIKE '%Language:%'
                  AND vuln_type != 'telemetry'
                GROUP BY language
                ORDER BY count DESC
            """)
            language_stats = [
                {"language": row[0], "count": row[1]}
                for row in cursor.fetchall()
            ]
            
            return {
                "global": {
                    "total_scans": total_scans,
                    "total_vulnerabilities": total_vulns,
                    "total_users": total_users,
                    "average_scan_duration": round(avg_duration, 2)
                },
                "severity": {
                    "critical": severity_totals[0] or 0,
                    "high": severity_totals[1] or 0,
                    "medium": severity_totals[2] or 0,
                    "low": severity_totals[3] or 0
                },
                "top_vulnerabilities": top_vulnerabilities,
                "severity_distribution": severity_distribution,
                "recent_activity": recent_scans,
                "language_stats": language_stats,
                "last_updated": datetime.utcnow().isoformat()
            }
        finally:
            conn.close()

# Instance du gestionnaire de stats
stats_manager = HonestStatsManager()

def detect_language_from_extension(filename):
    """Détecter le langage à partir de l'extension"""
    ext = Path(filename).suffix.lower()
    for lang, extensions in LANGUAGE_EXTENSIONS.items():
        if ext in extensions or filename.lower() in extensions:
            return lang
    return None

# Route OPTIONS globale pour les requêtes preflight CORS
@app.route('/api/<path:path>', methods=['OPTIONS'])
def handle_options(path):
    """Gère les requêtes preflight OPTIONS pour CORS"""
    response = jsonify({'status': 'ok'})
    return response

@app.route('/api/v1/health', methods=['GET'])
def health_check():
    """Health check endpoint"""
    return jsonify({
        "status": "healthy",
        "api_version": "3.0",
        "scanner": "full-cli-engine",
        "supported_languages": list(LANGUAGE_EXTENSIONS.keys()),
        "total_extensions": sum(len(exts) for exts in LANGUAGE_EXTENSIONS.values()),
        "services": {
            "auth": "operational",
            "sast_scanner": "operational",
            "stats_db": "operational"
        },
        "timestamp": datetime.utcnow().isoformat()
    })

@app.route('/api/v1/stats', methods=['GET'])
def get_stats():
    """Retourner les vraies statistiques"""
    real_stats = stats_manager.get_real_statistics()
    real_stats['version'] = '3.0'
    real_stats['scanner_engine'] = 'full-cli'
    
    logger.info(f"Stats envoyées: scans={real_stats.get('total_scans', 0)}, vulns={real_stats.get('total_vulnerabilities_found', 0)}")
    
    return jsonify(real_stats)

@app.route('/api/v1/stats/detailed', methods=['GET'])
def get_detailed_stats():
    """Endpoint pour les statistiques détaillées"""
    try:
        detailed_stats = stats_manager.get_detailed_statistics()
        detailed_stats['api_version'] = '3.0'
        detailed_stats['scanner_engine'] = 'full-cli'
        
        logger.info(f"Stats détaillées envoyées: {detailed_stats['global']['total_scans']} scans, {detailed_stats['global']['total_vulnerabilities']} vulns")
        
        return jsonify(detailed_stats)
    except Exception as e:
        logger.error(f"Erreur lors de la récupération des stats détaillées: {e}")
        return jsonify({"error": str(e)}), 500

@app.route('/api/v1/modules', methods=['GET'])
def get_modules():
    """Liste des modules/langages supportés"""
    modules = []
    for lang, extensions in LANGUAGE_EXTENSIONS.items():
        modules.append({
            "name": lang,
            "description": f"Scanner {lang.upper()} SAST",
            "extensions": list(extensions),
            "enabled": True
        })
    
    return jsonify({
        "modules": modules,
        "total": len(modules),
        "scanner_version": "3.0",
        "engine": "full-cli"
    })

@app.route('/api/v1/telemetry', methods=['POST'])
def receive_telemetry():
    """
    Recevoir les statistiques anonymes de télémétrie.
    
    Body JSON attendu:
    {
        "findings_count": 10,
        "scan_duration": 2.5,
        "modules_used": ["headers", "tls"],
        "platform": "windows-gui",
        "app_version": "1.0.0",
        "severity_counts": {"critical": 2, "high": 5, "medium": 3}
    }
    """
    try:
        data = request.get_json()
        
        # Validation des données
        required_fields = ["findings_count", "scan_duration", "modules_used", "platform"]
        if not all(field in data for field in required_fields):
            return jsonify({"error": "Missing required fields"}), 400
        
        # Debug: afficher vulnerabilities_by_lang
        vulnerabilities_by_lang = data.get("vulnerabilities_by_lang")
        logger.info(f"🔍 DEBUG vulnerabilities_by_lang: {vulnerabilities_by_lang}")
        
        # Enregistrer dans les stats globales avec les détails
        stats_manager.add_telemetry_scan(
            findings_count=data["findings_count"],
            scan_duration=data["scan_duration"],
            platform=data.get("platform", "unknown"),
            app_version=data.get("app_version", "1.0.0"),
            severity_counts=data.get("severity_counts"),
            vulnerability_types=data.get("vulnerability_types"),
            languages_detected=data.get("languages_detected"),
            findings_by_language=data.get("findings_by_language"),
            vulnerabilities_by_lang=vulnerabilities_by_lang  # 🆕 Mapping précis
        )
        
        logger.info(f"📊 Télémétrie reçue: {data['findings_count']} findings, {data['scan_duration']}s, {data['platform']}")
        
        return jsonify({
            "status": "success",
            "message": "Telemetry received"
        }), 200
        
    except Exception as e:
        logger.error(f"Erreur réception télémétrie: {e}")
        return jsonify({"error": str(e)}), 500

@app.route('/api/v1/scan', methods=['POST'])
def scan_source():
    """Scanner le code source avec le VRAI scanner CLI OU enregistrer un scan déjà effectué"""
    start_time = datetime.utcnow()
    
    # Vérifier l'authentification via PostgreSQL
    api_key = request.headers.get('X-API-Key')
    logger.info(f"🔑 API Key reçue: {api_key[:30] if api_key else 'NONE'}...")
    user_info = get_user_by_api_key(api_key) if api_key else None
    logger.info(f"👤 User info: {user_info}")
    if not user_info or not user_info.get('active'):
        logger.error(f"❌ Authentification échouée pour API key: {api_key[:30] if api_key else 'NONE'}")
        return jsonify({'error': 'Invalid API key'}), 401
    
    data = request.get_json()
    if not data:
        return jsonify({'error': 'Données requises'}), 400
    
    # FORMAT 1: Scan déjà effectué avec résultats (depuis GUI v2)
    if 'findings' in data and 'target' in data:
        logger.info(f"📥 Réception scan complet: user={user_info['email']}, target={data['target']}, findings={len(data['findings'])}")
        
        # Enregistrer directement dans PostgreSQL
        scan_id = save_scan_to_postgres(
            user_info['user_id'], 
            user_info['email'], 
            data['target'],
            data.get('scan_type', 'sast'),
            data['findings'],
            data.get('duration', 0)
        )
        
        severity_breakdown = {'critical': 0, 'high': 0, 'medium': 0, 'low': 0, 'info': 0}
        for finding in data['findings']:
            sev = finding.get('severity', 'low').lower()
            if sev in severity_breakdown:
                severity_breakdown[sev] += 1
        
        logger.info(f"✅ Scan enregistré: ID={scan_id}, vulns={len(data['findings'])}")
        
        return jsonify({
            'success': True,
            'scan_id': scan_id,
            'vulnerabilities': len(data['findings']),
            'summary': severity_breakdown
        }), 201
    
    # FORMAT 2: Code source à scanner (ancien format)
    if 'code' not in data:
        return jsonify({'error': 'Code source requis (champ "code") ou résultats complets (champ "findings")'}), 400
    
    code = data['code']
    filename = data.get('filename', 'unknown.py')
    language = data.get('language', None)
    
    # Détecter le langage si non fourni
    if not language:
        language = detect_language_from_extension(filename)
        if not language:
            language = 'python'  # Défaut
    
    logger.info(f"Scan demandé: user={user_info['email']}, filename={filename}, language={language}")
    
    try:
        # Créer un fichier temporaire avec la VRAIE extension pour que le scanner CLI le détecte
        # Le scanner détecte le langage via l'extension du fichier
        file_extension = Path(filename).suffix if Path(filename).suffix else f'.{language}'
        
        # Créer le fichier temporaire
        with tempfile.NamedTemporaryFile(mode='w', suffix=file_extension, delete=False, encoding='utf-8') as tmp_file:
            tmp_file.write(code)
            tmp_path = Path(tmp_file.name)
        
        # Configurer le scanner avec le vrai moteur CLI
        config = SourceScanConfig(
            source_files=[tmp_path],
            languages={language},
            recursive=False,
            detailed_report=True,
            advanced_rules=True
        )
        
        # Créer le scanner et lancer l'analyse
        scanner = SourceCodeScanner(config)
        
        # Scanner le fichier (utilise la méthode publique scan())
        findings_list = scanner.scan()
        
        # Nettoyer le fichier temporaire (fermer d'abord sous Windows)
        try:
            if tmp_path.exists():
                os.unlink(tmp_path)
        except Exception as e:
            logger.warning(f"Impossible de supprimer le fichier temporaire: {e}")
        
        # Convertir les findings en dict
        findings = []
        for finding in findings_list:
            findings.append({
                'check': finding.check,
                'severity': finding.severity,
                'title': finding.title,
                'description': finding.description,
                'remediation': finding.remediation,
                'evidence': finding.evidence or ''
            })
        
        # Calculer la durée
        duration = (datetime.utcnow() - start_time).total_seconds()
        
        # Enregistrer dans PostgreSQL
        scan_id = save_scan_to_postgres(
            user_info['user_id'], 
            user_info['email'], 
            filename, 
            'sast',  # Type de scan : sast (Static Application Security Testing)
            findings, 
            duration
        )
        
        # Compter les sévérités
        severity_breakdown = {'critical': 0, 'high': 0, 'medium': 0, 'low': 0, 'info': 0}
        for finding in findings:
            sev = finding.get('severity', 'low').lower()
            if sev in severity_breakdown:
                severity_breakdown[sev] += 1
        
        logger.info(f"Scan terminé: ID={scan_id}, vulns={len(findings)}, durée={duration}s")
        
        return jsonify({
            'success': True,
            'filename': filename,
            'language': language,
            'scanned_by': user_info['email'],
            'scan_timestamp': datetime.utcnow().isoformat(),
            'scan_duration_seconds': round(duration, 2),
            'scanner_engine': 'full-cli',
            'summary': {
                'total_vulnerabilities': len(findings),
                'severity_breakdown': severity_breakdown
            },
            'vulnerabilities': findings
        })
    
    except Exception as e:
        logger.error(f"Erreur lors du scan: {e}", exc_info=True)
        return jsonify({
            'error': 'Scan failed',
            'message': str(e)
        }), 500

@app.route('/api/v1/user', methods=['GET'])
def get_user_info():
    """Obtenir les informations utilisateur via PostgreSQL"""
    api_key = request.headers.get('X-API-Key')
    user = get_user_by_api_key(api_key) if api_key else None
    if not user or not user.get('active'):
        return jsonify({'error': 'Invalid API key'}), 401
    
    return jsonify({
        'user': user['email'],
        'name': user['name'],
        'tier': user['tier'],
        'permissions': {
            'can_scan_source': True,
            'advanced_rules': True,
            'all_languages': True
        }
    })

# ===============================================
# ROUTES D'AUTHENTIFICATION ET GESTION UTILISATEURS
# ===============================================

user_manager = UserManager()

@app.route('/api/v1/auth/register', methods=['POST'])
def register():
    """Inscription d'un nouvel utilisateur"""
    data = request.get_json()
    
    if not data or not all(k in data for k in ['name', 'email', 'password']):
        return jsonify({'error': 'Données incomplètes (name, email, password requis)'}), 400
    
    result = user_manager.create_user(
        name=data['name'],
        email=data['email'],
        password=data['password']
    )
    
    if result['success']:
        return jsonify({
            'status': 'success',
            'token': result['token'],
            'name': result['name'],
            'email': result['email']
        }), 201
    else:
        return jsonify({'error': result['error']}), 400

@app.route('/api/v1/auth/login', methods=['POST'])
def login():
    """Connexion d'un utilisateur"""
    data = request.get_json()
    
    if not data or not all(k in data for k in ['email', 'password']):
        return jsonify({'error': 'Email et mot de passe requis'}), 400
    
    result = user_manager.authenticate_user(
        email=data['email'],
        password=data['password']
    )
    
    if result['success']:
        # Récupérer l'API key de l'utilisateur depuis la BDD
        api_key = None
        try:
            conn = psycopg2.connect(**DB_CONFIG)
            cursor = conn.cursor()
            cursor.execute("SELECT api_key FROM users WHERE email = %s", (data['email'],))
            row = cursor.fetchone()
            if row:
                api_key = row[0]
            cursor.close()
            conn.close()
        except Exception as e:
            logger.error(f"Erreur récupération API key: {e}")
        
        return jsonify({
            'status': 'success',
            'token': result['token'],
            'api_key': api_key,  # 🔑 AJOUT DE L'API KEY
            'name': result['name'],
            'email': result['email'],
            'tier': result['tier']
        })
    else:
        return jsonify({'error': result['error']}), 401

@app.route('/api/v1/auth/reset-password-request', methods=['POST', 'OPTIONS'])
def reset_password_request():
    """Demande de réinitialisation de mot de passe - Étape 1"""
    if request.method == 'OPTIONS':
        return '', 200
    
    data = request.get_json()
    
    if not data or 'email' not in data:
        return jsonify({'error': 'Email requis'}), 400
    
    email = data['email']
    
    try:
        import secrets
        import smtplib
        from email.mime.text import MIMEText
        from email.mime.multipart import MIMEMultipart
        from datetime import datetime, timedelta
        
        conn = psycopg2.connect(**DB_CONFIG)
        cursor = conn.cursor()
        
        # Vérifier si l'utilisateur existe
        cursor.execute("SELECT id, name FROM users WHERE email = %s", (email,))
        user = cursor.fetchone()
        
        if not user:
            # Ne pas révéler si l'email existe ou non (sécurité)
            return jsonify({'status': 'success', 'message': 'Si cet email existe, un code a été envoyé'}), 200
        
        user_id, name = user
        
        # Générer un code à 6 chiffres
        reset_code = ''.join([str(secrets.randbelow(10)) for _ in range(6)])
        expires_at = datetime.utcnow() + timedelta(minutes=30)  # Code valide 30 minutes
        
        # Stocker le code dans la base de données
        cursor.execute("""
            INSERT INTO password_reset_codes (user_id, code, expires_at, used)
            VALUES (%s, %s, %s, FALSE)
            ON CONFLICT (user_id) 
            DO UPDATE SET code = EXCLUDED.code, expires_at = EXCLUDED.expires_at, used = FALSE
        """, (user_id, reset_code, expires_at))
        
        conn.commit()
        cursor.close()
        conn.close()
        
        # Envoyer l'email avec le code
        try:
            smtp_config = {
                'host': os.getenv('SMTP_HOST', 'web-sentinel-smtp'),
                'port': int(os.getenv('SMTP_PORT', 25)),
                'user': os.getenv('SMTP_USER', ''),
                'password': os.getenv('SMTP_PASSWORD', ''),
                'use_ssl': os.getenv('SMTP_USE_SSL', 'False').lower() == 'true'
            }
            
            msg = MIMEMultipart('alternative')
            msg['Subject'] = 'Réinitialisation de votre mot de passe Web Sentinel'
            msg['From'] = smtp_config['user'] if smtp_config['user'] else 'noreply@web-sentinel.taaazzz-prog.fr'
            msg['To'] = email
            
            # Version texte brut (OBLIGATOIRE pour compatibilité)
            text_content = f"""
Bonjour {name},

Vous avez demandé à réinitialiser votre mot de passe Web Sentinel.

Votre code de vérification est : {reset_code}

Ce code est valide pendant 30 minutes.

Si vous n'avez pas demandé cette réinitialisation, ignorez cet email.

---
Web Sentinel - Plateforme de sécurité web
© 2025 Tous droits réservés
            """
            
            # Version HTML (pour affichage enrichi)
            html_content = f"""
            <html>
                <body style="font-family: Arial, sans-serif; padding: 20px; background-color: #f8f9fa;">
                    <div style="max-width: 600px; margin: 0 auto; background: white; padding: 30px; border-radius: 10px; box-shadow: 0 2px 10px rgba(0,0,0,0.1);">
                        <h2 style="color: #2c3e50;">🔑 Réinitialisation de mot de passe</h2>
                        <p>Bonjour <strong>{name}</strong>,</p>
                        <p>Vous avez demandé à réinitialiser votre mot de passe Web Sentinel.</p>
                        <p>Votre code de vérification est :</p>
                        <div style="background: #f8f9fa; padding: 20px; text-align: center; margin: 20px 0; border-radius: 5px;">
                            <h1 style="color: #3498db; font-size: 2.5rem; letter-spacing: 5px; margin: 0;">{reset_code}</h1>
                        </div>
                        <p><strong>Ce code est valide pendant 30 minutes.</strong></p>
                        <p style="color: #6c757d; font-size: 0.9rem;">Si vous n'avez pas demandé cette réinitialisation, ignorez cet email.</p>
                        <hr style="border: none; border-top: 1px solid #e9ecef; margin: 20px 0;">
                        <p style="color: #95a5a6; font-size: 0.8rem; text-align: center;">
                            Web Sentinel - Plateforme de sécurité web<br>
                            © 2025 Tous droits réservés
                        </p>
                    </div>
                </body>
            </html>
            """
            
            # Attacher les deux versions (texte d'abord, puis HTML)
            msg.attach(MIMEText(text_content, 'plain', 'utf-8'))
            msg.attach(MIMEText(html_content, 'html', 'utf-8'))
            
            # Envoyer via SMTP (avec ou sans SSL selon config)
            if smtp_config['use_ssl']:
                # Port 465 : SSL direct
                logger.info(f"Envoi email via SMTP_SSL ({smtp_config['host']}:{smtp_config['port']})")
                with smtplib.SMTP_SSL(smtp_config['host'], smtp_config['port'], timeout=30) as server:
                    if smtp_config['user']:
                        server.login(smtp_config['user'], smtp_config['password'])
                    server.send_message(msg)
            else:
                # Port 25/587 : SMTP normal avec STARTTLS
                logger.info(f"Envoi email via SMTP+STARTTLS ({smtp_config['host']}:{smtp_config['port']})")
                with smtplib.SMTP(smtp_config['host'], smtp_config['port'], timeout=30) as server:
                    server.starttls()  # Activer le chiffrement TLS
                    if smtp_config['user']:
                        server.login(smtp_config['user'], smtp_config['password'])
                    server.send_message(msg)
            
            logger.info(f"✅ Code de réinitialisation envoyé à {email}")
            
        except Exception as e:
            logger.error(f"Erreur envoi email: {e}")
            # Continuer même si l'email échoue (pour le développement)
        
        return jsonify({'status': 'success', 'message': 'Code envoyé par email'}), 200
        
    except Exception as e:
        logger.error(f"Erreur reset password request: {e}")
        return jsonify({'error': 'Erreur serveur'}), 500

@app.route('/api/v1/auth/reset-password-confirm', methods=['POST', 'OPTIONS'])
def reset_password_confirm():
    """Confirmation de réinitialisation de mot de passe - Étape 2"""
    if request.method == 'OPTIONS':
        return '', 200
    
    data = request.get_json()
    
    if not data or not all(k in data for k in ['email', 'code', 'new_password']):
        return jsonify({'error': 'Email, code et nouveau mot de passe requis'}), 400
    
    email = data['email']
    code = data['code']
    new_password = data['new_password']
    
    if len(new_password) < 8:
        return jsonify({'error': 'Le mot de passe doit contenir au moins 8 caractères'}), 400
    
    try:
        from datetime import datetime
        import hashlib
        import secrets
        
        conn = psycopg2.connect(**DB_CONFIG)
        cursor = conn.cursor()
        
        # Vérifier le code
        cursor.execute("""
            SELECT prc.user_id, prc.expires_at, prc.used, u.email
            FROM password_reset_codes prc
            JOIN users u ON u.id = prc.user_id
            WHERE u.email = %s AND prc.code = %s
        """, (email, code))
        
        result = cursor.fetchone()
        
        if not result:
            cursor.close()
            conn.close()
            return jsonify({'error': 'Code invalide'}), 400
        
        user_id, expires_at, used, _ = result
        
        # Vérifier si le code est expiré
        if datetime.utcnow() > expires_at:
            cursor.close()
            conn.close()
            return jsonify({'error': 'Code expiré. Demandez un nouveau code'}), 400
        
        # Vérifier si le code a déjà été utilisé
        if used:
            cursor.close()
            conn.close()
            return jsonify({'error': 'Code déjà utilisé'}), 400
        
        # Générer le nouveau hash de mot de passe (format SHA256 + salt)
        salt = secrets.token_hex(16)
        pwd_hash = hashlib.sha256((new_password + salt).encode()).hexdigest()
        password_hash = f"{salt}${pwd_hash}"
        
        # Mettre à jour le mot de passe
        cursor.execute("""
            UPDATE users SET password_hash = %s WHERE id = %s
        """, (password_hash, user_id))
        
        # Marquer le code comme utilisé
        cursor.execute("""
            UPDATE password_reset_codes SET used = TRUE WHERE user_id = %s
        """, (user_id,))
        
        conn.commit()
        cursor.close()
        conn.close()
        
        logger.info(f"Mot de passe réinitialisé pour l'utilisateur {email}")
        
        return jsonify({'status': 'success', 'message': 'Mot de passe réinitialisé avec succès'}), 200
        
    except Exception as e:
        logger.error(f"Erreur reset password confirm: {e}")
        return jsonify({'error': 'Erreur serveur'}), 500

@app.route('/api/v1/user/licenses', methods=['GET'])
@require_auth
def get_user_licenses():
    """Récupérer les licences de l'utilisateur connecté"""
    user_id = request.current_user['user_id']
    
    # Récupérer les licences
    licenses_result = user_manager.get_user_licenses(user_id)
    
    if not licenses_result['success']:
        return jsonify({'error': licenses_result['error']}), 500
    
    # Calculer les stats utilisateur (scans, vulnérabilités) depuis PostgreSQL
    scans_count = 0
    vulns_count = 0
    
    try:
        conn = psycopg2.connect(**DB_CONFIG)
        cursor = conn.cursor()
        
        # Compter les scans et vulnérabilités de l'utilisateur par api_key_id
        cursor.execute('''
            SELECT 
                COUNT(*), 
                COALESCE(SUM(severity_critical), 0) + 
                COALESCE(SUM(severity_high), 0) + 
                COALESCE(SUM(severity_medium), 0) + 
                COALESCE(SUM(severity_low), 0) + 
                COALESCE(SUM(severity_info), 0) as total_vulns
            FROM scans
            WHERE api_key_id = %s
        ''', (user_id,))
        
        result = cursor.fetchone()
        if result:
            scans_count, vulns_count = result[0] or 0, result[1] or 0
        
        cursor.close()
        conn.close()
    except Exception as e:
        logger.error(f"Erreur récupération stats PostgreSQL: {e}")
    
    return jsonify({
        'licenses': licenses_result['licenses'],
        'stats': {
            'scans': scans_count,
            'vulnerabilities': vulns_count,
            'licenses': len(licenses_result['licenses'])
        }
    })

@app.route('/api/v1/user/profile', methods=['GET'])
@require_auth
def get_user_profile():
    """Récupérer le profil de l'utilisateur connecté"""
    return jsonify({
        'user_id': request.current_user['user_id'],
        'name': request.current_user['name'],
        'email': request.current_user['email'],
        'tier': request.current_user['tier']
    })

@app.route('/api/v1/user/stats', methods=['GET'])
@require_auth
def get_user_personal_stats():
    """
    Récupérer les statistiques PERSONNELLES de l'utilisateur connecté.
    
    Différence avec /api/v1/stats (globales) :
    - /api/v1/stats : Stats anonymes de TOUS les utilisateurs (télémétrie)
    - /api/v1/user/stats : Stats AUTHENTIFIÉES de l'utilisateur connecté uniquement
    
    Retourne :
    - Nombre de scans SAST effectués par l'utilisateur
    - Répartition des vulnérabilités par sévérité
    - Historique des scans avec dates
    - Statistiques par langage scanné
    """
    user_id = request.current_user['user_id']
    user_email = request.current_user['email']
    user_tier = request.current_user['tier']
    
    try:
        conn = psycopg2.connect(**DB_CONFIG)
        cursor = conn.cursor()
        
        # 1. Stats globales de l'utilisateur
        cursor.execute('''
            SELECT 
                COUNT(*) as total_scans,
                COALESCE(SUM(findings_count), 0) as total_findings,
                COALESCE(SUM(severity_critical), 0) as critical,
                COALESCE(SUM(severity_high), 0) as high,
                COALESCE(SUM(severity_medium), 0) as medium,
                COALESCE(SUM(severity_low), 0) as low,
                COALESCE(SUM(severity_info), 0) as info,
                COALESCE(AVG(duration), 0) as avg_duration,
                MIN(timestamp) as first_scan,
                MAX(timestamp) as last_scan
            FROM scans
            WHERE api_key_id = %s
        ''', (user_id,))
        
        stats_row = cursor.fetchone()
        
        # 2. Historique des 10 derniers scans
        cursor.execute('''
            SELECT 
                id,
                target,
                scan_type,
                findings_count,
                severity_critical,
                severity_high,
                severity_medium,
                severity_low,
                severity_info,
                duration,
                timestamp
            FROM scans
            WHERE api_key_id = %s
            ORDER BY timestamp DESC
            LIMIT 10
        ''', (user_id,))
        
        recent_scans = []
        for row in cursor.fetchall():
            recent_scans.append({
                'id': row[0],
                'target': row[1],
                'scan_type': row[2],
                'findings_count': row[3],
                'severity': {
                    'critical': row[4],
                    'high': row[5],
                    'medium': row[6],
                    'low': row[7],
                    'info': row[8]
                },
                'duration': round(row[9], 2) if row[9] else 0,
                'timestamp': row[10].isoformat() if row[10] else None
            })
        
        # 3. Stats par type de scan
        cursor.execute('''
            SELECT 
                scan_type,
                COUNT(*) as count,
                COALESCE(SUM(findings_count), 0) as total_findings
            FROM scans
            WHERE api_key_id = %s
            GROUP BY scan_type
        ''', (user_id,))
        
        scans_by_type = {}
        for row in cursor.fetchall():
            scans_by_type[row[0]] = {
                'count': row[1],
                'total_findings': row[2]
            }
        
        cursor.close()
        conn.close()
        
        # Construire la réponse
        total_scans = stats_row[0] or 0
        total_vulns = (stats_row[2] or 0) + (stats_row[3] or 0) + (stats_row[4] or 0) + (stats_row[5] or 0) + (stats_row[6] or 0)
        
        # Pour SYSOP, si aucun scan, afficher au moins 1 licence
        licenses_count = 1 if user_tier == 'sysop' else 0
        
        logger.info(f"📊 Stats personnelles {user_email}: {total_scans} scans, {total_vulns} vulns")
        
        return jsonify({
            'user': {
                'email': user_email,
                'tier': user_tier
            },
            'summary': {
                'total_scans': total_scans,
                'total_vulnerabilities': total_vulns,
                'licenses_active': licenses_count,
                'avg_scan_duration': round(stats_row[7], 2) if stats_row[7] else 0,
                'first_scan': stats_row[8].isoformat() if stats_row[8] else None,
                'last_scan': stats_row[9].isoformat() if stats_row[9] else None
            },
            'severity_breakdown': {
                'critical': stats_row[2] or 0,
                'high': stats_row[3] or 0,
                'medium': stats_row[4] or 0,
                'low': stats_row[5] or 0,
                'info': stats_row[6] or 0
            },
            'scans_by_type': scans_by_type,
            'recent_scans': recent_scans
        })
        
    except Exception as e:
        logger.error(f"❌ Erreur récupération stats personnelles {user_email}: {e}", exc_info=True)
        return jsonify({'error': 'Failed to retrieve personal statistics'}), 500

if __name__ == '__main__':
    port = int(os.environ.get('PORT', 5000))
    logger.info(f"🚀 Démarrage de l'API Web Sentinel v3.0 sur le port {port}")
    logger.info(f"📦 Langages supportés: {len(LANGUAGE_EXTENSIONS)}")
    logger.info(f"📝 Extensions supportées: {sum(len(exts) for exts in LANGUAGE_EXTENSIONS.values())}")
    app.run(host='0.0.0.0', port=port, debug=False)
