AEGIS-AI-Security / engine_llm.py
Om43937's picture
Update engine_llm.py
cfa06ba verified
Raw
History Blame Contribute Delete
5.97 kB
import torch
import logging
import hashlib
from transformers import pipeline
from threading import Lock
from typing import List, Dict, Tuple
# ==========================================
# 🔧 Logging
# ==========================================
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s | %(levelname)s | %(message)s"
)
logger = logging.getLogger("LLMEngine")
# ==========================================
# 🧠 LLMEngine: Zero-Trust Prompt Injection Scanner
# ==========================================
class LLMEngine:
_instance = None
_init_lock = Lock()
_infer_lock = Lock()
# ----------------------------------------
# Singleton Pattern
# ----------------------------------------
def __new__(cls):
if cls._instance is None:
with cls._init_lock:
if cls._instance is None:
cls._instance = super().__new__(cls)
cls._instance._initialized = False
return cls._instance
# ----------------------------------------
# Initialize Engine
# ----------------------------------------
def __init__(self):
if self._initialized:
return
self.device = 0 if torch.cuda.is_available() else -1
self.model_name = "ProtectAI/deberta-v3-base-prompt-injection"
self.scanner = None
self.cache: Dict[str, Dict] = {}
logger.info("LLMEngine initialized (lazy loading)")
self._initialized = True
# ----------------------------------------
# Load Model (Lazy + Warmup)
# ----------------------------------------
def _load_model(self):
if self.scanner is None:
try:
logger.info(f"Loading model {self.model_name} on device {self.device}...")
self.scanner = pipeline(
"text-classification",
model=self.model_name,
device=self.device,
truncation=True,
max_length=512
)
# 🔥 Warmup
self.scanner("Test input for warmup")
logger.info("Model ready ✅")
except Exception as e:
logger.error(f"Model load failed: {e}")
self.scanner = None
# ----------------------------------------
# Split Text Intelligently (Sentence-based)
# ----------------------------------------
def _split_text(self, text: str, min_len=5) -> List[str]:
sentences = [s.strip() for s in text.split('.') if len(s.strip()) >= min_len]
if not sentences:
return [text]
return sentences
# ----------------------------------------
# Hash Cache Key
# ----------------------------------------
def _hash(self, text: str) -> str:
return hashlib.md5(text.encode()).hexdigest()
# ----------------------------------------
# Aggregate results (Safe vs Injection)
# ----------------------------------------
def _aggregate(self, results: List[Dict]) -> Dict:
injection_score = 0.0
safe_score = 0.0
for r in results:
label = r["label"].upper()
score = r["score"]
if "INJECTION" in label or "LABEL_1" in label:
injection_score = max(injection_score, score)
elif "SAFE" in label or "LABEL_0" in label:
safe_score = max(safe_score, score)
if safe_score >= injection_score:
return {
"label": "SAFE",
"confidence": round(safe_score, 4)
}
return {
"label": "INJECTION",
"confidence": round(injection_score, 4)
}
# ----------------------------------------
# Analyze Single Input
# ----------------------------------------
def analyze(self, text: str) -> Dict:
if not text or not text.strip():
return {"label": "INVALID", "confidence": 0.0}
key = self._hash(text)
if key in self.cache:
return self.cache[key]
self._load_model()
if self.scanner is None:
return {"label": "UNKNOWN", "confidence": 0.0}
try:
chunks = self._split_text(text)
with self._infer_lock:
results = self.scanner(chunks)
final = self._aggregate(results)
final["model"] = self.model_name
final["chunks"] = len(chunks)
# Cache result
self.cache[key] = final
return final
except Exception as e:
logger.error(f"Inference failed: {e}")
return {"label": "ERROR", "confidence": 0.0}
# ----------------------------------------
# Analyze Batch Inputs
# ----------------------------------------
def analyze_batch(self, texts: List[str]) -> List[Dict]:
self._load_model()
if self.scanner is None:
return []
clean_texts = [t for t in texts if t and t.strip()]
if not clean_texts:
return []
try:
with self._infer_lock:
results = self.scanner(clean_texts)
output = []
for r, t in zip(results, clean_texts):
agg = self._aggregate([r])
agg.update({
"text": t,
"model": self.model_name
})
output.append(agg)
return output
except Exception as e:
logger.error(f"Batch analysis failed: {e}")
return []
# ----------------------------------------
# Health Check
# ----------------------------------------
def health(self) -> Dict:
return {
"model": self.model_name,
"device": "GPU" if self.device == 0 else "CPU",
"loaded": self.scanner is not None,
"cache_size": len(self.cache)
}