cheikhdeme/tp1_malware
0
1import os2import joblib3import pefile4import numpy as np5import pandas as pd6import gradio as gr7import hashlib8import traceback9from sklearn.ensemble import RandomForestClassifier10from sklearn.model_selection import train_test_split11from sklearn.metrics import accuracy_score, recall_score12 13# Chemin vers le modèle sauvegardé14MODEL_PATH = 'random_forest_model.pkl'15 16def train_and_save_model():17 """Entraîner et sauvegarder le modèle si nécessaire."""18 print("Aucun modèle trouvé. Entraînement en cours...")19 # Chargement des données20 data = pd.read_csv("DatasetmalwareExtrait.csv")21 22 # Traitement des données23 X = data.drop(['legitimate'], axis=1)24 y = data['legitimate']25 26 # Entraînement du modèle27 X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.3, random_state=42)28 model = RandomForestClassifier(29 n_estimators=196,30 random_state=42,31 criterion="gini",32 max_depth=25,33 min_samples_split=4,34 min_samples_leaf=135 )36 model.fit(X_train, y_train)37 38 # Évaluation du modèle39 y_pred = model.predict(X_test)40 accuracy = accuracy_score(y_test, y_pred)41 recall = recall_score(y_test, y_pred, average='weighted')42 43 print(f"Précision du modèle supervisé : {accuracy:.3f}")44 print(f"Rappel du modèle supervisé : {recall:.3f}")45 46 # Sauvegarde du modèle47 joblib.dump(model, MODEL_PATH)48 print(f"Modèle sauvegardé sous : {MODEL_PATH}")49 return model50 51# Chargement ou entraînement du modèle52if os.path.exists(MODEL_PATH):53 print("Chargement du modèle existant...")54 model = joblib.load(MODEL_PATH)55else:56 model = train_and_save_model()57 58# Fonctions utilitaires59def calculate_file_hash(file_path):60 """Calculer le hash SHA-256 du fichier."""61 sha256_hash = hashlib.sha256()62 with open(file_path, "rb") as f:63 for byte_block in iter(lambda: f.read(4096), b""):64 sha256_hash.update(byte_block)65 return sha256_hash.hexdigest()66 67def extract_pe_attributes(file_path):68 """Extraction avancée des attributs du fichier PE."""69 try:70 pe = pefile.PE(file_path)71 attributes = {72 'AddressOfEntryPoint': pe.OPTIONAL_HEADER.AddressOfEntryPoint,73 'MajorLinkerVersion': pe.OPTIONAL_HEADER.MajorLinkerVersion,74 'MajorImageVersion': pe.OPTIONAL_HEADER.MajorImageVersion,75 'MajorOperatingSystemVersion': pe.OPTIONAL_HEADER.MajorOperatingSystemVersion,76 'DllCharacteristics': pe.OPTIONAL_HEADER.DllCharacteristics,77 'SizeOfStackReserve': pe.OPTIONAL_HEADER.SizeOfStackReserve,78 'NumberOfSections': pe.FILE_HEADER.NumberOfSections,79 'ResourceSize': pe.OPTIONAL_HEADER.DATA_DIRECTORY[2].Size80 }81 return attributes82 except Exception as e:83 print(f"Erreur de traitement du fichier {file_path}: {str(e)}")84 return {"Erreur": str(e)}85 86def predict_malware(file):87 """Prédiction de malware avec gestion d'erreurs."""88 if model is None:89 return "Erreur : Modèle non chargé"90 91 try:92 # Extraire les attributs du fichier93 attributes = extract_pe_attributes(file.name)94 if "Erreur" in attributes:95 return attributes["Erreur"]96 97 # Convertir en DataFrame98 df = pd.DataFrame([attributes])99 100 # Prédiction101 prediction = model.predict(df)102 proba = model.predict_proba(df)[0]103 104 # Résultat avec probabilité105 if prediction[0] == 1:106 return f"🚨 MALWARE (Probabilité: {proba[1] * 100:.2f}%)"107 else:108 return f"✅ Fichier Légitime (Probabilité: {proba[0] * 100:.2f}%)"109 except Exception as e:110 return f"Erreur d'analyse : {str(e)}"111 112# Interface Gradio113demo = gr.Interface(114 fn=predict_malware,115 inputs=gr.File(file_types=['.exe', '.dll', '.sys'], label="Télécharger un fichier exécutable"),116 outputs="text",117 title="🛡️ Détecteur de Malwares",118 theme='huggingface' # Thème moderne119)120 121if __name__ == "__main__":122 demo.launch(share=True)123 