#!/usr/bin/env python3
"""
Script de test : Vérification du système de quotas SAST
Date: 7 novembre 2025
Usage: python scripts/database/test_sast_quotas.py
"""

import os
import sys
from pathlib import Path
from datetime import datetime, timezone

# Ajouter le répertoire racine au path
root_dir = Path(__file__).parent.parent.parent
sys.path.insert(0, str(root_dir))

from dotenv import load_dotenv
load_dotenv()

# Test avec PostgreSQL
os.environ['WEB_SENTINEL_LICENSE_BACKEND'] = 'postgresql'

from web_sentinel.payment.license_store_factory import build_license_store
from web_sentinel.subscription.models import SubscriptionTier


def test_sast_quotas():
    """Teste les quotas SAST pour tous les tiers"""
    print("🧪 Test des quotas SAST")
    print("=" * 60)
    
    store = build_license_store()
    print(f"✓ Store initialisé: {type(store).__name__}")
    
    # Test 1: Créer des licences de test pour chaque tier
    print("\n📝 Test 1: Création de licences de test")
    print("-" * 60)
    
    test_licenses = {}
    tiers = ['FREE', 'PRO', 'ENTERPRISE', 'SYSOP']
    
    for tier in tiers:
        email = f"test_{tier.lower()}@websentinel-test.local"
        
        try:
            # Créer nouvelle licence
            license_data = store.create_license(
                email=email,
                tier=tier,
                max_domains=10 if tier != 'SYSOP' else 999,
                max_users=1 if tier == 'FREE' else 5,
            )
            
            test_licenses[tier] = license_data
            
            print(f"   ✓ {tier:12} : API Key = {license_data['api_key'][:30]}...")
            print(f"      allow_source_scan  = {license_data.get('allow_source_scan', 'N/A')}")
            print(f"      max_source_files   = {license_data.get('max_source_files', 'N/A')}")
            print(f"      max_source_size_mb = {license_data.get('max_source_size_mb', 'N/A')}")
            print(f"      advanced_rules     = {license_data.get('advanced_rules', 'N/A')}")
            
        except Exception as e:
            print(f"   ❌ Erreur création {tier}: {e}")
            # Continuer quand même pour les autres tiers

    
    # Test 2: Vérifier la récupération des licences
    print("\n🔍 Test 2: Récupération et vérification des quotas")
    print("-" * 60)
    
    expected_quotas = {
        'FREE': {'files': 0, 'size': 0, 'scan': False, 'rules': False},
        'PRO': {'files': 100, 'size': 20, 'scan': True, 'rules': False},
        'ENTERPRISE': {'files': 500, 'size': 100, 'scan': True, 'rules': True},
        'SYSOP': {'files': 9999, 'size': 500, 'scan': True, 'rules': True},
    }
    
    all_passed = True
    
    for tier, license_data in test_licenses.items():
        api_key = license_data['api_key']
        retrieved = store.get_license_by_api_key(api_key)
        
        if not retrieved:
            print(f"   ❌ {tier}: Licence non trouvée par API key")
            all_passed = False
            continue
        
        expected = expected_quotas[tier]
        actual_files = retrieved.get('max_source_files', 0)
        actual_size = retrieved.get('max_source_size_mb', 0)
        actual_scan = retrieved.get('allow_source_scan', False)
        actual_rules = retrieved.get('advanced_rules', False)
        
        checks = [
            ('max_source_files', actual_files, expected['files']),
            ('max_source_size_mb', actual_size, expected['size']),
            ('allow_source_scan', actual_scan, expected['scan']),
            ('advanced_rules', actual_rules, expected['rules']),
        ]
        
        tier_passed = True
        for field, actual, expected_val in checks:
            if actual == expected_val:
                print(f"   ✓ {tier:12} {field:20} = {actual} (attendu: {expected_val})")
            else:
                print(f"   ❌ {tier:12} {field:20} = {actual} (attendu: {expected_val})")
                tier_passed = False
                all_passed = False
        
        if not tier_passed:
            print(f"      ⚠️  Quotas incorrects pour {tier}")
    
    # Test 3: Vérifier la mise à jour d'une licence
    print("\n🔄 Test 3: Mise à jour des quotas")
    print("-" * 60)
    
    try:
        # Upgrade PRO vers ENTERPRISE
        pro_license = test_licenses['PRO']
        api_key = pro_license['api_key']
        
        print(f"   📝 Upgrade PRO → ENTERPRISE (API Key: {api_key[:30]}...)")
        
        updated = store.update_license(
            api_key=api_key,
            tier='ENTERPRISE',
            allow_source_scan=True,
            max_source_files=500,
            max_source_size_mb=100,
            advanced_rules=True,
        )
        
        if updated.get('max_source_files') == 500:
            print(f"   ✓ Mise à jour réussie: {updated.get('max_source_files')} fichiers")
        else:
            print(f"   ❌ Mise à jour échouée: {updated.get('max_source_files')} fichiers (attendu: 500)")
            all_passed = False
            
    except Exception as e:
        print(f"   ❌ Erreur mise à jour: {e}")
        all_passed = False
    
    # Test 4: Vérifier le nombre total de licences
    print("\n📊 Test 4: Comptage des licences")
    print("-" * 60)
    
    try:
        # Compter via requête SQL directe
        import psycopg2
        from dotenv import load_dotenv
        load_dotenv('.env.local')
        
        DATABASE_URL = os.getenv('WEB_SENTINEL_POSTGRES_URL')
        conn = psycopg2.connect(DATABASE_URL)
        cursor = conn.cursor()
        
        cursor.execute("SELECT COUNT(*) FROM api_key WHERE email LIKE 'test_%'")
        test_count = cursor.fetchone()[0]
        
        cursor.execute("SELECT COUNT(*) FROM api_key")
        total_count = cursor.fetchone()[0]
        
        cursor.close()
        conn.close()
        
        print(f"   📌 Total licences: {total_count}")
        print(f"   🧪 Licences de test: {test_count}")
        
        if test_count >= 3:  # Au moins 3 des 4 (cas où FREE échoue)
            print(f"   ✓ Au moins 3 licences de test trouvées")
        else:
            print(f"   ⚠️  Seulement {test_count} licences de test (attendu: >= 3)")
            
    except Exception as e:
        print(f"   ❌ Erreur comptage: {e}")
        all_passed = False
    
    # Résumé final
    print("\n" + "=" * 60)
    if all_passed:
        print("✅ TOUS LES TESTS SONT PASSÉS!")
        print("\n💡 Prochaines étapes:")
        print("   1. Tester via l'interface admin (http://localhost:8000)")
        print("   2. Créer/éditer des licences manuellement")
        print("   3. Vérifier les webhooks Stripe")
        print("   4. Nettoyer les licences de test si nécessaire")
        return True
    else:
        print("❌ CERTAINS TESTS ONT ÉCHOUÉ")
        print("\n🔧 Actions recommandées:")
        print("   1. Vérifier que la migration SQL a été exécutée")
        print("   2. Vérifier les logs du backend admin")
        print("   3. Exécuter: python scripts/database/migrate_sast_columns.py")
        return False


if __name__ == '__main__':
    try:
        success = test_sast_quotas()
        sys.exit(0 if success else 1)
    except Exception as e:
        print(f"\n💥 Erreur fatale: {e}")
        import traceback
        traceback.print_exc()
        sys.exit(1)
