Team Ai
Apppublic

pranavKHF/code_vulnerability_detection

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes
inference.py156 linesDownload Raw Back to src
1"""Inference module for vulnerability detection2Load trained models and make predictions"""3 4import torch5from transformers import RobertaTokenizer6from pathlib import Path7import sys8sys.path.append(str(Path(__file__).parent.parent.parent))9 10from src.model import VulnerabilityCodeT511 12class VulnerabilityDetector:13    def __init__(self, model_path="models/best_model_clean.pt",14                 model_name="Salesforce/codet5-base", max_length=256):15        16        ### CHANGED FOR DEPLOYMENT17        self.device = torch.device('cpu')18        self.max_length = max_length19 20        self.tokenizer = RobertaTokenizer.from_pretrained(model_name)21 22        self.model = VulnerabilityCodeT5(model_name=model_name, num_labels=2)23 24        state_dict = torch.load(model_path, map_location=self.device)25        self.model.load_state_dict(state_dict)26        self.model.to(self.device)27        self.model.eval()28 29 30        print("Model Loaded Successfully")31 32        self.labels = {33            0: "Safe Code",34            1: "Vulnerable Code"35        }36 37    def predict(self, code_snippet):38        """Predict Vulnerability of Code Snippet39        40        Args : 41        code_snippet: String Containing source code42        43        Returns:44          dict with predictions, confidence and label45        46        """47        inputs = self.tokenizer(48            code_snippet, 49            max_length=256,50            padding='max_length',51            truncation=True,52            return_tensors='pt'53        )54 55        input_ids = inputs['input_ids'].to(self.device)56        attention_mask = inputs['attention_mask'].to(self.device)57 58        with torch.no_grad():59 60            predictions, probs = self.model.predict(input_ids, attention_mask)61 62        pred_label = predictions[0].item()63        confidence = probs[0][pred_label].item()64 65        return {66            'prediction': pred_label,67            'label': self.labels[pred_label],68            'confidence': confidence,69            'probabilities':{70                'safe': probs[0][0].item(),71                'vulnerable': probs[0][1].item()72            }73        }74    75    def analyze_batch(self, code_snippets):76        """Analyze multiple code snippets at once"""77        return [self.predict(code) for code in code_snippets]78    79def test_inference():80    detector = VulnerabilityDetector()81    82    83 84 85    test_cases = [86        {87            "name": "Safe Bounded Copy",88            "code": """void copy_input(const char *input) {89        char buffer[32];90        strncpy(buffer, input, sizeof(buffer) - 1);91        buffer[sizeof(buffer) - 1] = '\\0';92    }"""93        },94        {95            "name": "Safe fgets Input",96            "code": """void read_input() {97        char buffer[64];98        if (fgets(buffer, sizeof(buffer), stdin) != NULL) {99            printf("%s", buffer);100        }101    }"""102        },103        {104            "name": "Safe malloc usage",105            "code": """void allocate() {106        char *buf = (char *)malloc(128);107        if (buf == NULL) {108            return;109        }110        strcpy(buf, "safe");111        free(buf);112    }"""113        },114        {115            "name": "Stack Buffer Overflow",116            "code": """void copy_input(char *input) {117        char buffer[8];118        strcpy(buffer, input);119    }"""120        },121        {122            "name": "Integer Overflow",123            "code": """void allocate(int size) {124        char *buf = (char *)malloc(size * sizeof(char));125        if (buf == NULL) return;126        memset(buf, 'A', size + 10);127    }"""128        },129        {130            "name": "Use After Free",131            "code": """void uaf() {132        char *buf = (char *)malloc(16);133        free(buf);134        strcpy(buf, "UAF");135    }"""136        }137    ]138 139 140    print("\n" + "="*60)141    print("Testing Vulnerability Detection")142    print("="*60)143    144    for test in test_cases:145        print(f"\nTest: {test['name']}")146        print(f"Code: {test['code'][:60]}...")147        148        result = detector.predict(test['code'])149        150        print(f"Prediction: {result['label']}")151        print(f"Confidence: {result['confidence']:.2%}")152        print(f"   - Safe: {result['probabilities']['safe']:.2%}")153        print(f"   - Vulnerable: {result['probabilities']['vulnerable']:.2%}")154 155if __name__ == "__main__":156    test_inference()