pranavKHF/code_vulnerability_detection
0
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()