ai-code-scanner-ui / combined_detector.py
mubi-613's picture
Update model path to Hugging Face Hub in app and combined detector
5eee832
Raw
History Blame Contribute Delete
6.23 kB
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!")