Spaces:
Running
Running
| import torch | |
| from transformers import AutoTokenizer, AutoModelForSequenceClassification | |
| from rule_detector import RuleBasedCodeDetector | |
| from typing import Dict, List, Any | |
| class CombinedCodeDetector: | |
| def __init__(self): | |
| print("Loading Combined Detector...") | |
| # 1. NEW MODEL PATH (Points to Hugging Face Hub) | |
| self.model_path = "mubi-613/ai-code-security-scanner" | |
| # 2. Loading ML Model from Hugging Face | |
| # We replace "enhanced_saved_model" with self.model_path | |
| self.ml_tokenizer = AutoTokenizer.from_pretrained(self.model_path) | |
| self.ml_model = AutoModelForSequenceClassification.from_pretrained(self.model_path) | |
| self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| self.ml_model.to(self.device) | |
| self.ml_model.eval() | |
| #Loading rule-based detector | |
| self.rule_detector = RuleBasedCodeDetector() | |
| print(f"Combined detector loaded Successfully from {self.model_path}!") | |
| def ml_analysis(self, code: str) -> Dict: | |
| # ML-based analysis using fine-tuned CodeBERT | |
| try: | |
| inputs = self.ml_tokenizer( | |
| code, | |
| return_tensors = "pt", | |
| truncation = True, | |
| max_length = 256, | |
| padding = True | |
| ) | |
| inputs = {k: v.to(self.device) for k, v in inputs.items()} | |
| with torch.no_grad(): | |
| outputs = self.ml_model(**inputs) | |
| probabilities = torch.nn.functional.softmax(outputs.logits, dim=-1) | |
| prediction_idx = torch.argmax(probabilities, dim=-1).item() | |
| confidence = probabilities[0][prediction_idx].item() | |
| return { | |
| "prediction" :"vulnerable" if prediction_idx == 1 else "safe", | |
| "confidence" : float(confidence), | |
| "safe_prob" : float(probabilities[0][0]), | |
| "vulnerable_prob" : float(probabilities[0][1]), | |
| "success" : True | |
| } | |
| except Exception as e: | |
| return{ | |
| "prediction": "error", | |
| "error": str(e), | |
| "success": False | |
| } | |
| def combined_analysis(self, code: str) -> Dict[str, Any]: | |
| """Combine rule-based and ML analysis""" | |
| # Get rule-based results | |
| rule_results = self.rule_detector.analyze(code) | |
| # Get ML results (if no critical syntax errors) | |
| ml_results = {} | |
| if not any(i["severity"] == "CRITICAL" for i in rule_results["issues"]): | |
| ml_results = self.ml_analysis(code) | |
| # Combine issues | |
| all_issues = rule_results["issues"].copy() | |
| # Add ML prediction as issue if confident about vulnerability | |
| if ml_results.get("success") and ml_results["prediction"] == "vulnerable": | |
| if ml_results["confidence"] > 0.8: # High confidence | |
| all_issues.append({ | |
| "type": "ml_detected_issue", | |
| "severity": "MEDIUM", | |
| "message": f"AI detected potential issue (confidence: {ml_results['confidence']:.1%})", | |
| "line": "N/A", | |
| "fix": "Review with security expert", | |
| "detector": "ml_model" | |
| }) | |
| # Calculate combined score | |
| severity_weights = {'CRITICAL': 1.0, 'HIGH': 0.7, 'MEDIUM': 0.4, 'LOW': 0.1} | |
| rule_weight = sum( | |
| severity_weights.get(issue.get('severity', 'LOW'), 0.1) | |
| for issue in rule_results["issues"] | |
| ) | |
| # Adjust with ML confidence | |
| ml_adjustment = 0 | |
| if ml_results.get("success"): | |
| if ml_results["prediction"] == "vulnerable": | |
| ml_adjustment = ml_results["confidence"] * 0.5 | |
| else: | |
| ml_adjustment = -ml_results["confidence"] * 0.3 | |
| total_weight = rule_weight + ml_adjustment | |
| security_score = min(100, max(0, 100 - (total_weight * 15))) # Adjusted scaling | |
| return { | |
| "issues": all_issues, | |
| "security_score": round(security_score, 1), | |
| "issue_count": len(all_issues), | |
| "ml_analysis": ml_results, | |
| "detectors_used": ["rule_based", "ml_model"] if ml_results.get("success") else ["rule_based"], | |
| "summary": { | |
| "critical": sum(1 for i in all_issues if i["severity"] == "CRITICAL"), | |
| "high": sum(1 for i in all_issues if i["severity"] == "HIGH"), | |
| "medium": sum(1 for i in all_issues if i["severity"] == "MEDIUM"), | |
| "low": sum(1 for i in all_issues if i["severity"] == "LOW"), | |
| } | |
| } | |
| # Test the combined detector | |
| if __name__ == "__main__": | |
| detector = CombinedCodeDetector() | |
| test_cases = [ | |
| """def get_user(user_id):\n query = f"SELECT * FROM users WHERE id = {user_id}"\n return query""", | |
| """api_key = os.getenv("API_KEY")""", | |
| """def test()\n print("hello")""", # Syntax error | |
| """import pickle\ndata = pickle.loads(user_data)""", | |
| ] | |
| for i, code in enumerate(test_cases, 1): | |
| print(f"\n{'='*60}") | |
| print(f"TEST CASE {i}") | |
| print(f"{'='*60}") | |
| print(f"Code:\n{code}") | |
| result = detector.combined_analysis(code) | |
| print(f"\n📊 Results:") | |
| print(f"Security Score: {result['security_score']}/100") | |
| print(f"Issues Found: {result['issue_count']}") | |
| print(f"Detectors Used: {', '.join(result['detectors_used'])}") | |
| if result['ml_analysis'].get('success'): | |
| ml = result['ml_analysis'] | |
| print(f"ML Prediction: {ml['prediction'].upper()} ({ml['confidence']:.1%} confidence)") | |
| if result['issues']: | |
| print("\n🔍 Issues:") | |
| for issue in result['issues']: | |
| print(f" [{issue['severity']}] {issue['message']} (Line {issue['line']})") | |
| print(f" Fix: {issue['fix']}") | |
| print(f" Detector: {issue['detector']}") | |
| else: | |
| print("\n✅ No issues found!") |