diff --git a/.env.example b/.env.example index 308ae55f..1f9dc0d8 100644 --- a/.env.example +++ b/.env.example @@ -57,7 +57,10 @@ MAX_CSV_ROWS=100000 # 100k rows max ENABLE_MALWARE_SCAN=false # Enable ClamAV scanning CLAMAV_HOST=localhost CLAMAV_PORT=3310 +<<<<<<< Updated upstream # Archival Job Schedule (Cron expression) # Default is '0 3 * * *' (run once daily at 3:00 AM) ARCHIVAL_CRON_SCHEDULE=0 3 * * * +======= +>>>>>>> Stashed changes diff --git a/backend/.env.example b/backend/.env.example index b94ba6d2..209a31b2 100644 --- a/backend/.env.example +++ b/backend/.env.example @@ -109,6 +109,7 @@ PREDICT_RATE_LIMIT=50 per minute VITE_API_URI=http://localhost:3000/api VITE_ML_API_URI=http://localhost:5000/predict +<<<<<<< Updated upstream # -------------------------------------------- # GROQ API (AI Chat) # -------------------------------------------- @@ -119,4 +120,6 @@ GROQ_API_KEY=your_groq_api_key_here GROQ_MODELS=["llama-3.1-8b-instant","llama-3.1-70b-versatile","mixtral-8x7b-32768","gemma2-9b-it"] # OR single model (backward compatible) -# GROQ_MODEL=llama-3.1-8b-instant \ No newline at end of file +# GROQ_MODEL=llama-3.1-8b-instant +======= +>>>>>>> Stashed changes diff --git a/backend/adversarial_training.py b/backend/adversarial_training.py index 7400830a..65f4766e 100644 --- a/backend/adversarial_training.py +++ b/backend/adversarial_training.py @@ -72,6 +72,16 @@ def __init__(self): # Noise patterns self.noise_chars = ['.', '!', '?', ',', ';', ':', ' '] +<<<<<<< HEAD +======= + + # Spam trigger patterns + self.spam_triggers = [ + 'urgent', 'free', 'claim', 'prize', 'winner', 'congratulations', + 'limited time', 'act now', 'exclusive', 'guaranteed', 'money back', + 'cash', 'bonus', 'credit', 'loan', 'investment', 'profit' + ] +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda # Spam trigger patterns @@ -97,9 +107,16 @@ def synonym_replacement(self, text, intensity=0.4): word_lower = word.lower().strip('.,!?') if word_lower in self.synonyms and random.random() < intensity: new_word = random.choice(self.synonyms[word_lower]) +<<<<<<< Updated upstream # Preserve punctuation +======= +<<<<<<< HEAD +======= + # Preserve punctuation +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes punct = word[-1] if word[-1] in '.,!?' else '' result.append(new_word + punct) else: @@ -118,12 +135,17 @@ def noise_injection(self, text, intensity=0.2): result.insert(pos, char) return ''.join(result) +<<<<<<< HEAD def generate_variants(self, text, num_variants=5): """Generate multiple adversarial variants""" variants = [] for _ in range(num_variants): variant = text +<<<<<<< Updated upstream +======= +======= +>>>>>>> Stashed changes def sentence_rephrasing(self, text): """Simple sentence rephrasing (rule-based)""" # This is a simple version - can be enhanced with LLM @@ -146,6 +168,10 @@ def generate_variants(self, text, num_variants=5): for _ in range(num_variants): variant = text # Apply random transformations +<<<<<<< Updated upstream +======= +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes transformations = random.sample([ ('char', 0.2 + random.random() * 0.3), ('synonym', 0.2 + random.random() * 0.3), @@ -160,12 +186,21 @@ def generate_variants(self, text, num_variants=5): elif transform_type == 'noise': variant = self.noise_injection(variant, intensity) +<<<<<<< Updated upstream +======= +<<<<<<< HEAD +======= +>>>>>>> Stashed changes # Sometimes rephrase if random.random() < 0.3: variant = self.sentence_rephrasing(variant) +<<<<<<< Updated upstream +======= +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes variants.append(variant) return variants @@ -176,6 +211,7 @@ def generate_variants(self, text, num_variants=5): # ============================================ def load_datasets(): +<<<<<<< HEAD """Load dataset""" if not os.path.exists(DATASET_PATH): print(f"❌ Dataset not found at {DATASET_PATH}") @@ -207,7 +243,11 @@ def create_sample_dataset(): ] df = pd.DataFrame(samples, columns=['text', 'label']) print(f"βœ… Created sample dataset with {len(df)} samples") +<<<<<<< Updated upstream +======= +======= +>>>>>>> Stashed changes """Load and combine email + SMS datasets""" data = [] @@ -272,7 +312,11 @@ def create_sample_dataset(): print(f" Spam: {len(df[df['label']=='spam'])}") print(f" Ham: {len(df[df['label']=='ham'])}") +<<<<<<< Updated upstream +======= +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes return df @@ -284,7 +328,14 @@ def train_adversarial_model(df): """Train model with adversarial examples""" print("\nπŸ”„ Generating adversarial examples...") +<<<<<<< Updated upstream +======= +<<<<<<< HEAD +======= + +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes augmentor = AdversarialAugmentor() # Separate spam and ham @@ -336,7 +387,14 @@ def train_adversarial_model(df): ('lr', LogisticRegression(class_weight='balanced', max_iter=1000, random_state=42)) ] +<<<<<<< Updated upstream +======= +<<<<<<< HEAD +======= + # Train each model +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes trained_models = {} for name, model in models: print(f" Training {name}...") @@ -367,17 +425,31 @@ def train_adversarial_model(df): def main(): print("=" * 60) +<<<<<<< HEAD print("πŸ›‘οΈ Adversarial Training for Spam Detection") +<<<<<<< Updated upstream print("πŸš€ Adversarial Training for Spam Detection") +======= +======= + print("πŸš€ Adversarial Training for Spam Detection") +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes print("=" * 60) # Load datasets df = load_datasets() if df is None: +<<<<<<< HEAD print("\n❌ No dataset available") +<<<<<<< Updated upstream print("\n❌ Please ensure dataset.csv exists in backend/") +======= +======= + print("\n❌ Please ensure dataset.csv exists in backend/") +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes return # Train model @@ -390,24 +462,42 @@ def main(): print(f"πŸ’Ύ Saving vectorizer to {VECTORIZER_OUTPUT_PATH}") pickle.dump(vectorizer, open(VECTORIZER_OUTPUT_PATH, 'wb')) +<<<<<<< Updated upstream # Save label encoder (for compatibility) +======= +<<<<<<< HEAD + # Save label encoder +======= + # Save label encoder (for compatibility) +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes from sklearn.preprocessing import LabelEncoder le = LabelEncoder() le.fit(['ham', 'spam']) pickle.dump(le, open(LABEL_ENCODER_PATH, 'wb')) +<<<<<<< HEAD print("\nβœ… Adversarial training complete!") +<<<<<<< Updated upstream print("\nβœ… Training complete!") +======= +======= + print("\nβœ… Training complete!") +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes print(f" Model saved: {MODEL_OUTPUT_PATH}") print(f" Vectorizer saved: {VECTORIZER_OUTPUT_PATH}") print("\nπŸ“ To use the robust model, update your .env:") print(f" MODEL_PATH={MODEL_OUTPUT_PATH}") print(f" VECTORIZER_PATH={VECTORIZER_OUTPUT_PATH}") +<<<<<<< HEAD print(f" CONFIDENCE_THRESHOLD=0.6") print(f" FLAG_LOW_CONFIDENCE=true") +======= +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda diff --git a/backend/api.py b/backend/api.py index b853520f..dabc95e5 100644 --- a/backend/api.py +++ b/backend/api.py @@ -60,6 +60,8 @@ +sys.path.insert(0, str(Path(__file__).resolve().parent / "email_connectors")) + load_dotenv() @@ -139,9 +141,11 @@ def decorated_function(*args, **kwargs): return f(*args, **kwargs) return decorated_function + # Alias used by routes that gate on the internal service-to-service secret. internal_endpoint_required = validate_internal_request + # Apply to all routes by default (except public paths) @app.before_request def require_internal_secret(): @@ -175,6 +179,9 @@ def decorated_function(*args, **kwargs): allowed_list = [ip.strip() for ip in allowed_ips.split(",")] client_ip = request.headers.get("X-Forwarded-For", request.remote_addr) or "" + + client_ip = request.headers.get("X-Forwarded-For", request.remote_addr) + # Get first IP if multiple if "," in client_ip: client_ip = client_ip.split(",")[0].strip() @@ -192,6 +199,7 @@ def decorated_function(*args, **kwargs): + # ============================================ # ZERO TRUST - REQUEST VALIDATION @@ -230,6 +238,7 @@ def decorated_function(*args, **kwargs): +# ============================================ # ZERO TRUST - AUDIT LOGGING # ============================================ @@ -309,6 +318,10 @@ def resolve_path(env_var, default_filename): +# In-memory storage for spam words +spam_words_storage = {} + + # SQLite Persistent Storage for spam words import sqlite3 from datetime import datetime, timezone @@ -317,6 +330,15 @@ def resolve_path(env_var, default_filename): def init_spam_words_db(): with imap_store.get_db_connection() as conn: + +def _db_connection(): + conn = sqlite3.connect(imap_store.DB_PATH) + conn.row_factory = sqlite3.Row + return conn + +def init_spam_words_db(): + with _db_connection() as conn: + conn.execute( """ CREATE TABLE IF NOT EXISTS spam_word_frequencies ( @@ -332,6 +354,9 @@ def init_spam_words_db(): def increment_spam_word_frequency(word): day = datetime.now(timezone.utc).strftime("%Y-%m-%d") with imap_store.get_db_connection() as conn: + + with _db_connection() as conn: + conn.execute( """ INSERT INTO spam_word_frequencies (word, day, count) @@ -343,7 +368,11 @@ def increment_spam_word_frequency(word): conn.commit() def get_db_wordcloud_data(): + with imap_store.get_db_connection() as conn: + + with _db_connection() as conn: + rows = conn.execute( """ SELECT word, SUM(count) as total_count @@ -402,6 +431,9 @@ def get_word_of_the_day_data(): day = datetime.now(timezone.utc).strftime("%Y-%m-%d") word_row = None with imap_store.get_db_connection() as conn: + + with _db_connection() as conn: + word_row = conn.execute( """ SELECT word, SUM(count) as total_count @@ -424,7 +456,7 @@ def get_word_of_the_day_data(): LIMIT 1 """ ).fetchone() - + if word_row: word = word_row["word"] count = word_row["total_count"] @@ -451,6 +483,11 @@ def get_word_of_the_day_data(): app.vectorizer = vectorizer # type: ignore[attr-defined] app.label_encoder = label_encoder # type: ignore[attr-defined] +app.model = model +app.vectorizer = vectorizer +app.label_encoder = label_encoder + + from bulk_predict import bulk_predict_bp app.register_blueprint(bulk_predict_bp) app.register_blueprint(analytics_bp) @@ -568,8 +605,12 @@ def make_prediction_response( translated=False, translated_text=None, domain_analysis=None, + explanation=None, severity=None + + explanation=None + ): """Enforces a strict standardized response schema for all predictions.""" response = { @@ -587,6 +628,7 @@ def make_prediction_response( response["translated_text"] = translated_text if domain_analysis is not None: response["domain_analysis"] = domain_analysis + # Thin, top-level summary of domain_analysis for consumers that just # want a quick URL risk signal without parsing the full breakdown. response["url_risk"] = { @@ -598,6 +640,10 @@ def make_prediction_response( response["explanation"] = explanation if severity is not None: response["severity"] = severity + + if explanation is not None: + response["explanation"] = explanation + return response @@ -635,6 +681,9 @@ def predict(): "error": f"Request body must be a JSON object, got {type(data).__name__}" }), 400 + + data = request.get_json(silent=True) or {} + text = data.get("text") input_type = data.get("type", "message") @@ -647,7 +696,20 @@ def predict(): return jsonify({ "error": f"'text' must be a string, got {type(text).__name__}" }), 400 + + normalized_text = normalizer.normalize(text) + vectorized = vectorizer.transform([normalized_text]) + prediction = model.predict(vectorized)[0] + + return jsonify({ + 'original_text': text, + 'normalized_text': normalized_text, + 'prediction': prediction + }) + + # Maximum-length validation before any vectorization/inference work. + if len(text) > MAX_MESSAGE_LENGTH: return jsonify({ "error": ( @@ -723,6 +785,9 @@ def predict(): confidence_level = "low" if final_output != "ham": + + if final_output == "spam": + words = extract_words(text) for word in words: try: @@ -740,6 +805,7 @@ def predict(): explanation = xai_engine.analyze(text, input_type=input_type) severity = calculate_spam_severity(original_text) + response_data = { "input": original_text, "result": final_output, @@ -767,8 +833,12 @@ def predict(): translated=translated, translated_text=text if translated else None, domain_analysis=domain_analysis, + explanation=explanation, severity=severity + + explanation=explanation + ) @@ -799,6 +869,26 @@ def extract_words(text): def get_wordcloud_data(): + + """Return stored spam word frequencies from database.""" + try: + data = get_db_wordcloud_data() + return data if data else None + except Exception as e: + print(f"[db-wordcloud] failed to get wordcloud data: {e}") + return None + + + + if spam_words_storage: + + if spam_words_storage: + + sorted_words = sorted(spam_words_storage.items(), key=lambda x: x[1], reverse=True) + return [{"word": w, "count": c} for w, c in sorted_words[:50]] + return None + + """Return stored spam word frequencies from database.""" try: data = get_db_wordcloud_data() @@ -897,6 +987,7 @@ def feedback(): return jsonify({"error": "Failed to record feedback."}), 500 + @app.route("/feedback/stats", methods=["GET"]) @validate_request @validate_internal_request @@ -1199,7 +1290,9 @@ def scan_emails_route(): imap_store.init_db() oauth_store.init_db() + init_spam_words_db() + scheduler = BackgroundScheduler() scheduler.start() @@ -1260,6 +1353,63 @@ def _refresh_oauth_tokens(): ) +def _refresh_oauth_tokens(): + """Runs inside the scheduler thread: refreshes OAuth tokens close to expiration.""" + try: + expiring = oauth_store.get_expiring_oauth_tokens(threshold_minutes=10) + except Exception as e: + print(f"[oauth-refresh] failed to fetch expiring tokens: {e}") + return + + for token_entry in expiring: + username = token_entry["username"] + provider = token_entry["provider"] + refresh_token = token_entry["refresh_token"] + + if not refresh_token: + print(f"[oauth-refresh] No refresh token for {username} ({provider})") + continue + + try: + if provider == "gmail": + new_tokens = refresh_gmail_token(refresh_token) + elif provider == "outlook": + new_tokens = refresh_outlook_token(refresh_token) + else: + continue + + oauth_store.save_oauth_tokens(username, provider, new_tokens) + print(f"[oauth-refresh] successfully refreshed token for {username} ({provider})") + except requests.exceptions.HTTPError as err: + is_auth_error = False + try: + if err.response is not None: + if err.response.status_code in (400, 401): + err_json = err.response.json() + err_desc = err_json.get("error", "") + if "invalid_grant" in err_desc or err_json.get("error_description", ""): + is_auth_error = True + except Exception: + pass + + if is_auth_error or (err.response is not None and err.response.status_code == 400): + print(f"[oauth-refresh] Token revoked/invalid for {username} ({provider}). Deleting from DB.") + oauth_store.delete_oauth_tokens(username, provider) + else: + print(f"[oauth-refresh] Temporary HTTP error refreshing token for {username} ({provider}): {err}") + except Exception as e: + print(f"[oauth-refresh] failed to refresh token for {username} ({provider}): {e}") + + +scheduler.add_job( + _refresh_oauth_tokens, + "interval", + minutes=5, + id="oauth_token_refresh", + replace_existing=True, +) + + def _run_imap_scan(username): conn_row = imap_store.get_connection(username) if not conn_row: diff --git a/backend/app.py b/backend/app.py index 196340be..ef2be4e6 100644 --- a/backend/app.py +++ b/backend/app.py @@ -13,7 +13,10 @@ from flask_limiter import Limiter from flask_limiter.util import get_remote_address import numpy as np +<<<<<<< Updated upstream from utils.spamSeverity import calculate_spam_severity +======= +>>>>>>> Stashed changes load_dotenv() @@ -90,6 +93,14 @@ def ensure_feedback_file(): def home(): return "ML API Running πŸš€" +@app.route('/health', methods=['GET']) +def health_check(): + return jsonify({ + 'status': 'healthy', + 'model_loaded': model is not None, + 'vectorizer_loaded': vectorizer is not None + }) + # ─── HEALTH CHECK ENDPOINT ─────────────────────────────────────── @app.route("/health", methods=["GET"]) @@ -132,8 +143,12 @@ def make_prediction_response( translated=False, translated_text=None, domain_analysis=None, +<<<<<<< Updated upstream explanation=None, severity=None +======= + explanation=None +>>>>>>> Stashed changes ): """Enforces a strict standardized response schema for all predictions.""" response = { @@ -153,8 +168,11 @@ def make_prediction_response( response["domain_analysis"] = domain_analysis if explanation is not None: response["explanation"] = explanation +<<<<<<< Updated upstream if severity is not None: response["severity"] = severity +======= +>>>>>>> Stashed changes return response @@ -230,8 +248,12 @@ def predict(): confidence_level=confidence_level, detected_language=detected_language, translated=translated, +<<<<<<< Updated upstream translated_text=text if translated else None, severity=calculate_spam_severity(original_text) +======= + translated_text=text if translated else None +>>>>>>> Stashed changes ) return jsonify(response_data) diff --git a/backend/bulk_predict.py b/backend/bulk_predict.py index 52fa6a5c..419ab78d 100644 --- a/backend/bulk_predict.py +++ b/backend/bulk_predict.py @@ -8,11 +8,27 @@ from flask import (Blueprint, current_app, jsonify, request, send_file) + +from extensions import limiter + + + + +import numpy as np + +from flask import Blueprint, current_app, jsonify, request, send_file + +from api import limiter + +bulk_predict_bp = Blueprint("bulk_predict", __name__) + + from rate_limiting import RateLimitPolicy, rate_limit bulk_predict_bp = Blueprint("bulk_predict", __name__) + def parse_and_predict_file(file): # Check file extension filename = file.filename.lower() if file.filename else "" diff --git a/backend/config/envConstants.js b/backend/config/envConstants.js index a86926b6..b3925dd8 100644 --- a/backend/config/envConstants.js +++ b/backend/config/envConstants.js @@ -1,6 +1,14 @@ const REQUIRED_ENV_VARS = [ +<<<<<<< Updated upstream 'MONGODB_URI', 'JWT_SECRET', +======= + 'PORT', + 'NODE_ENV', + 'MONGODB_URI', + 'JWT_SECRET', + 'API_URL', +>>>>>>> Stashed changes 'INTERNAL_SECRET' ]; diff --git a/backend/controllers/activityController.js b/backend/controllers/activityController.js index b554e456..7987494b 100644 --- a/backend/controllers/activityController.js +++ b/backend/controllers/activityController.js @@ -1,7 +1,11 @@ // backend/controllers/activityController.js const mongoose = require('mongoose'); const History = require('../models/History'); +<<<<<<< Updated upstream const { replacements, tonePrefixes } = require('../utils/despamificationRules'); +======= + +>>>>>>> Stashed changes // ==================== DE-SPAMIFICATION LOGIC ==================== exports.despamify = async (req, res) => { try { @@ -14,6 +18,35 @@ exports.despamify = async (req, res) => { // Simple de-spamification logic let deSpammed = text; +<<<<<<< Updated upstream +======= + const replacements = { + 'URGENT': 'Someone wants to contact you', + 'FREE': 'There is an offer', + 'WIN': 'There is a notification', + 'PRIZE': 'There is a message about rewards', + 'CLAIM': 'There is a message for you', + 'CLICK': 'There is a link to visit', + 'NOW': 'soon', + '!!!': '.', + '$$$': '', + '100%': '', + 'GUARANTEED': '', + 'LIMITED TIME': '', + 'ACT NOW': '', + "DON'T MISS": '', + 'EXCLUSIVE': '', + 'YOU WON': 'There is a notification' + }; + + // Apply tone adjustments + const tonePrefixes = { + neutral: '', + friendly: 'Hi there! ', + formal: 'We would like to inform you that ', + casual: 'Hey! ' + }; +>>>>>>> Stashed changes const prefix = tonePrefixes[tone] || ''; @@ -88,11 +121,16 @@ exports.getStats = async (req, res) => { }; // ==================== USER ACTIVITY HEATMAP LOGIC ==================== +<<<<<<< Updated upstream +======= +// ==================== USER ACTIVITY HEATMAP LOGIC ==================== +>>>>>>> Stashed changes exports.getActivity = async (req, res) => { try { const { userId } = req.params; const { year, month } = req.query; +<<<<<<< Updated upstream const yearNum = parseInt(year, 10); const monthNum = parseInt(month, 10); @@ -106,6 +144,8 @@ exports.getActivity = async (req, res) => { return res.status(400).json({ error: "Month must be between 1 and 12." }); } +======= +>>>>>>> Stashed changes // Validate user id format before using it in aggregation if (!mongoose.Types.ObjectId.isValid(userId)) { return res.status(400).json({ error: "Invalid user id" }); @@ -118,8 +158,13 @@ exports.getActivity = async (req, res) => { }); } +<<<<<<< Updated upstream const startDate = new Date(yearNum, monthNum - 1, 1); const endDate = new Date(yearNum, monthNum, 0, 23, 59, 59, 999); +======= + const startDate = new Date(year, month - 1, 1); + const endDate = new Date(year, month, 0, 23, 59, 59, 999); +>>>>>>> Stashed changes const activities = await History.aggregate([ { diff --git a/backend/controllers/authController.js b/backend/controllers/authController.js index 269ab839..0e3ca08a 100644 --- a/backend/controllers/authController.js +++ b/backend/controllers/authController.js @@ -14,6 +14,10 @@ const emailTransporter = require('../utils/emailTransporter'); // TOKEN GENERATION // ============================================ +// ============================================ +// TOKEN GENERATION +// ============================================ + const generateToken = (userId) => { return jwt.sign({ id: userId }, process.env.JWT_SECRET, { expiresIn: process.env.JWT_EXPIRES_IN || '7d', @@ -365,8 +369,13 @@ const forgotPassword = async (req, res) => { const secret = process.env.JWT_SECRET + user.password; const token = jwt.sign( +<<<<<<< Updated upstream { id: user._id, email: user.email }, secret, +======= + { id: user._id, email: user.email }, + secret, +>>>>>>> Stashed changes { expiresIn: process.env.PASSWORD_RESET_TOKEN_EXPIRES || '15m' } ); @@ -697,4 +706,15 @@ module.exports = { getRolesAndPermissions, generateToken, buildAuthResponse -}; \ No newline at end of file +<<<<<<< Updated upstream +}; +======= +}; + +module.exports = { register, login, logout, getMe, googleLogin, updateAvatar, forgotPassword, resetPassword, updateWebhook }; + + + getSessionStatus +}; + +>>>>>>> Stashed changes diff --git a/backend/controllers/bulkPredictController.js b/backend/controllers/bulkPredictController.js index 6987b684..e80f8771 100644 --- a/backend/controllers/bulkPredictController.js +++ b/backend/controllers/bulkPredictController.js @@ -1,8 +1,11 @@ const { processBulkPrediction } = require("../services/bulkPredictService"); +<<<<<<< Updated upstream const logger = require("../utils/logger"); const MAX_TEXT_LENGTH = 10000; const MIN_TEXT_LENGTH = 2; +======= +>>>>>>> Stashed changes /** * Extract a usable prediction text value from a CSV row. @@ -13,7 +16,11 @@ const getPredictionInputFromRow = (row) => { const rowEntries = Object.entries(row); const textEntry = rowEntries.find(([key]) => +<<<<<<< Updated upstream ["text", "message", "content", "email", "sms", "tweet"].includes(key.trim().toLowerCase()) +======= + ["text", "message"].includes(key.trim().toLowerCase()) +>>>>>>> Stashed changes ); if (!textEntry) return null; @@ -39,22 +46,34 @@ const validateBulkPredictionRows = (rows, res) => { } const normalizedRows = []; +<<<<<<< Updated upstream const errors = []; +======= +>>>>>>> Stashed changes for (let index = 0; index < rows.length; index++) { const row = rows[index]; if (!row || typeof row !== "object" || Array.isArray(row)) { +<<<<<<< Updated upstream errors.push({ row: index + 2, error: "Row is not a valid CSV record." }); continue; +======= + res.status(400).json({ + success: false, + error: `Row ${index + 2} is not a valid CSV record.`, + }); + return null; +>>>>>>> Stashed changes } const predictionInput = getPredictionInputFromRow(row); if (!predictionInput) { +<<<<<<< Updated upstream errors.push({ row: index + 2, error: "Row is missing valid text content." @@ -76,6 +95,13 @@ const validateBulkPredictionRows = (rows, res) => { error: `Text content too long. Maximum ${MAX_TEXT_LENGTH} characters allowed.` }); continue; +======= + res.status(400).json({ + success: false, + error: `Row ${index + 2} is missing valid text content.`, + }); + return null; +>>>>>>> Stashed changes } normalizedRows.push({ @@ -84,6 +110,7 @@ const validateBulkPredictionRows = (rows, res) => { }); } +<<<<<<< Updated upstream if (errors.length > 0) { const errorSummary = { totalRows: rows.length, @@ -100,6 +127,8 @@ const validateBulkPredictionRows = (rows, res) => { return null; } +======= +>>>>>>> Stashed changes return normalizedRows; }; @@ -114,6 +143,7 @@ exports.handleBulkPrediction = async (req, res) => { const { rows, filename, size } = req.parsedFile; +<<<<<<< Updated upstream logger.info(`Bulk prediction requested: ${filename}, ${rows.length} rows`); const validatedRows = validateBulkPredictionRows(rows, res); @@ -125,23 +155,37 @@ exports.handleBulkPrediction = async (req, res) => { const results = await processBulkPrediction(validatedRows); logger.info(`Bulk prediction completed: ${filename}, ${validatedRows.length} rows processed`); +======= + // Process predictions + const results = await processBulkPrediction(rows); +>>>>>>> Stashed changes res.json({ success: true, totalRows: rows.length, +<<<<<<< Updated upstream validRows: validatedRows.length, invalidRows: rows.length - validatedRows.length, +======= +>>>>>>> Stashed changes filename: filename, size: size, results: results }); } catch (error) { console.error('Bulk prediction error:', error); +<<<<<<< Updated upstream logger.error('Bulk prediction error:', error); res.status(500).json({ success: false, error: 'Failed to process bulk prediction', details: process.env.NODE_ENV === 'development' ? error.message : undefined +======= + res.status(500).json({ + success: false, + error: 'Failed to process bulk prediction', + details: error.message +>>>>>>> Stashed changes }); } }; diff --git a/backend/controllers/chatController.js b/backend/controllers/chatController.js index 9887cbc4..3252da1a 100644 --- a/backend/controllers/chatController.js +++ b/backend/controllers/chatController.js @@ -1,3 +1,4 @@ +<<<<<<< Updated upstream const Groq = require("groq-sdk"); // ============================================ @@ -10,6 +11,14 @@ const DEFAULT_MODELS = [ "mixtral-8x7b-32768", "gemma2-9b-it" ]; +======= +// controllers/chatController.js +const Groq = require("groq-sdk"); + +const groq = new Groq({ + apiKey: process.env.GROQ_API_KEY || "placeholder_key" +}); +>>>>>>> Stashed changes const SYSTEM_PROMPT = `You are the Spam Detection System Security Assistant. Your purpose is purely educational. @@ -20,6 +29,7 @@ Guidelines: 4. If a query is unrelated to cybersecurity awareness, spam detection, phishing, malicious URLs, email security, SMS scams, or application usage, politely explain that the assistant is limited to security education topics. 5. Never claim certainty about whether a URL, email, SMS, or message is safe. Instead, explain indicators and recommend verification steps.`; +<<<<<<< Updated upstream // ============================================ // VALIDATION FUNCTIONS // ============================================ @@ -127,10 +137,13 @@ function createGroqClient() { // MAIN HANDLER // ============================================ +======= +>>>>>>> Stashed changes exports.chatHandler = async (req, res) => { try { const { message, history } = req.body; +<<<<<<< Updated upstream // Validate message const messageValidation = validateMessage(message); if (!messageValidation.valid) { @@ -139,10 +152,22 @@ exports.chatHandler = async (req, res) => { const trimmedMessage = messageValidation.value; // Build messages array +======= + const trimmedMessage = message?.trim() || ''; + + if (!trimmedMessage) { + return res.status(400).json({ error: "Message cannot be empty or only whitespace." }); + } + + if (trimmedMessage.length > 1000) { + return res.status(400).json({ error: "Message exceeds maximum length of 1000 characters." }); + } +>>>>>>> Stashed changes const messages = [ { role: "system", content: SYSTEM_PROMPT } ]; +<<<<<<< Updated upstream const sanitizedHistory = sanitizeHistory(history); messages.push(...sanitizedHistory); messages.push({ role: "user", content: trimmedMessage }); @@ -321,5 +346,50 @@ exports.listModels = (req, res) => { error: 'Failed to list models', details: error.message }); +======= + const ALLOWED_HISTORY_ROLES = new Set(["user", "assistant"]); + const MAX_HISTORY_ITEMS = 10; + const MAX_HISTORY_CONTENT_LENGTH = 2000; + + if (Array.isArray(history)) { + const recentHistory = history.slice(-MAX_HISTORY_ITEMS); + + for (const msg of recentHistory) { + if (!msg || typeof msg !== "object") continue; + + const { role, content } = msg; + + if (!ALLOWED_HISTORY_ROLES.has(role)) continue; + if (typeof content !== "string") continue; + + const trimmedContent = content.trim(); + if (!trimmedContent) continue; + + messages.push({ + role, + content: trimmedContent.slice(0, MAX_HISTORY_CONTENT_LENGTH), + }); + } + } + + messages.push({ role: "user", content: trimmedMessage }); + + const chatCompletion = await groq.chat.completions.create({ + messages: messages, + model: "llama-3.1-8b-instant", + temperature: 0.5, + max_tokens: 1024, + top_p: 1, + stop: null, + stream: false, + }); + + const reply = chatCompletion.choices[0]?.message?.content || "I am currently unable to process your request."; + + res.json({ reply }); + } catch (error) { + console.error("Groq API error:", error); + res.status(500).json({ error: "Failed to communicate with Security Assistant." }); +>>>>>>> Stashed changes } }; \ No newline at end of file diff --git a/backend/email_connectors/oauth_store.py b/backend/email_connectors/oauth_store.py index 36f00ac6..494dbba9 100644 --- a/backend/email_connectors/oauth_store.py +++ b/backend/email_connectors/oauth_store.py @@ -3,12 +3,20 @@ from datetime import datetime, timezone, timedelta from pathlib import Path +<<<<<<< Updated upstream from imap_store import DB_PATH, get_db_connection +======= +from imap_store import DB_PATH, _connection +>>>>>>> Stashed changes from crypto_utils import encrypt_secret, decrypt_secret def init_db(): +<<<<<<< Updated upstream with get_db_connection() as conn: +======= + with _connection() as conn: +>>>>>>> Stashed changes conn.execute( """ CREATE TABLE IF NOT EXISTS oauth_tokens ( @@ -46,7 +54,11 @@ def save_oauth_tokens(username, provider, tokens): expires_at = (now + timedelta(seconds=int(expires_in))).isoformat() updated_at = now.isoformat() +<<<<<<< Updated upstream with get_db_connection() as conn: +======= + with _connection() as conn: +>>>>>>> Stashed changes conn.execute( """ INSERT INTO oauth_tokens (username, provider, encrypted_access_token, encrypted_refresh_token, expires_at, updated_at) @@ -63,7 +75,11 @@ def save_oauth_tokens(username, provider, tokens): def get_oauth_tokens(username, provider): +<<<<<<< Updated upstream with get_db_connection() as conn: +======= + with _connection() as conn: +>>>>>>> Stashed changes row = conn.execute( "SELECT * FROM oauth_tokens WHERE username = ? AND provider = ?", (username, provider), @@ -91,7 +107,11 @@ def get_oauth_tokens(username, provider): def delete_oauth_tokens(username, provider=None): +<<<<<<< Updated upstream with get_db_connection() as conn: +======= + with _connection() as conn: +>>>>>>> Stashed changes if provider: conn.execute( "DELETE FROM oauth_tokens WHERE username = ? AND provider = ?", @@ -106,7 +126,11 @@ def delete_oauth_tokens(username, provider=None): def get_all_oauth_tokens(): +<<<<<<< Updated upstream with get_db_connection() as conn: +======= + with _connection() as conn: +>>>>>>> Stashed changes rows = conn.execute("SELECT * FROM oauth_tokens").fetchall() tokens_list = [] for row in rows: @@ -133,7 +157,11 @@ def get_all_oauth_tokens(): def get_expiring_oauth_tokens(threshold_minutes=10): now = datetime.now(timezone.utc) threshold = (now + timedelta(minutes=threshold_minutes)).isoformat() +<<<<<<< Updated upstream with get_db_connection() as conn: +======= + with _connection() as conn: +>>>>>>> Stashed changes rows = conn.execute( "SELECT * FROM oauth_tokens WHERE expires_at <= ?", (threshold,), diff --git a/backend/evo_mail.py b/backend/evo_mail.py index ea7cc422..cec3b107 100644 --- a/backend/evo_mail.py +++ b/backend/evo_mail.py @@ -1,7 +1,17 @@ #!/usr/bin/env python3 """ +<<<<<<< Updated upstream EvoMail - Self-Evolving Cognitive Agent for Spam Detection Red-Team/Blue-Team framework for continuous adaptation +======= +<<<<<<< HEAD +EvoMail/COG - Self-Evolving Cognitive Agent for Spam Detection +Red-Team/Blue-Team framework with memory and reasoning +======= +EvoMail - Self-Evolving Cognitive Agent for Spam Detection +Red-Team/Blue-Team framework for continuous adaptation +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes """ import json @@ -12,7 +22,17 @@ from collections import defaultdict import hashlib import os +<<<<<<< Updated upstream +from pathlib import Path +======= +<<<<<<< HEAD +import re from pathlib import Path +import heapq +======= +from pathlib import Path +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes # ============================================ # CONFIGURATION @@ -24,32 +44,95 @@ VECTORIZER_PATH = BASE_DIR / 'evo_vectorizer.pkl' # ============================================ +<<<<<<< Updated upstream +# MEMORY MODULE +======= +<<<<<<< HEAD +# MEMORY MODULE - Experience Compression & Reasoning +======= # MEMORY MODULE +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes # ============================================ class MemoryModule: """Compresses and stores experiences for future reasoning""" +<<<<<<< Updated upstream + def __init__(self): +======= +<<<<<<< HEAD + def __init__(self, max_memory=10000): +======= def __init__(self): +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes self.experiences = [] self.patterns = defaultdict(int) self.failures = [] self.successes = [] +<<<<<<< Updated upstream + self.max_memory = 10000 +======= +<<<<<<< HEAD + self.max_memory = max_memory +>>>>>>> Stashed changes + self.compression_threshold = 100 + + def add_experience(self, experience): +<<<<<<< Updated upstream +======= + """Add new experience to memory with importance scoring""" + # Calculate importance + importance = self._calculate_importance(experience) + experience['importance'] = importance + experience['timestamp'] = datetime.now().isoformat() + + self.experiences.append(experience) +======= self.max_memory = 10000 self.compression_threshold = 100 def add_experience(self, experience): +>>>>>>> Stashed changes """Add new experience to memory""" self.experiences.append({ 'timestamp': datetime.now().isoformat(), 'data': experience, 'importance': experience.get('importance', 1) }) +<<<<<<< Updated upstream +======= +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes # Compress if memory is full if len(self.experiences) > self.max_memory: self.compress() +<<<<<<< Updated upstream + # Extract patterns +======= +<<<<<<< HEAD + # Extract and update patterns +>>>>>>> Stashed changes + if 'text' in experience: + pattern = self.extract_pattern(experience['text']) + self.patterns[pattern] += 1 + + def extract_pattern(self, text): + """Extract key pattern from text""" + # Simple pattern extraction - can be enhanced with NLP + words = text.lower().split() +<<<<<<< Updated upstream + if len(words) > 3: + return ' '.join(words[:3]) # Use first 3 words as pattern +======= + if len(words) > 5: + # Use sliding window of 3 words as pattern + patterns = [' '.join(words[i:i+3]) for i in range(len(words)-2)] + return patterns[0] if patterns else text[:20] +======= # Extract patterns if 'text' in experience: pattern = self.extract_pattern(experience['text']) @@ -61,6 +144,8 @@ def extract_pattern(self, text): words = text.lower().split() if len(words) > 3: return ' '.join(words[:3]) # Use first 3 words as pattern +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes return text[:20] def compress(self): @@ -69,6 +154,16 @@ def compress(self): self.experiences.sort(key=lambda x: x['importance'], reverse=True) self.experiences = self.experiences[:self.max_memory // 2] +<<<<<<< Updated upstream +======= +<<<<<<< HEAD + # Apply importance decay + for exp in self.experiences: + exp['importance'] *= self.importance_decay + +======= +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes # Keep only top patterns top_patterns = sorted( self.patterns.items(), @@ -77,6 +172,27 @@ def compress(self): )[:1000] self.patterns = defaultdict(int, dict(top_patterns)) +<<<<<<< Updated upstream + def get_relevant_experiences(self, text, limit=10): + """Get experiences relevant to current text""" + pattern = self.extract_pattern(text) +======= +<<<<<<< HEAD + def recall(self, text, limit=10): + """Retrieve relevant experiences for reasoning""" + pattern = self._extract_pattern(text) +>>>>>>> Stashed changes + relevant = [] + + for exp in self.experiences: + if pattern in exp['data'].get('text', ''): + relevant.append(exp) + +<<<<<<< Updated upstream +======= + # Sort by importance and recency + relevant.sort(key=lambda x: (x['importance'], x['timestamp']), reverse=True) +======= def get_relevant_experiences(self, text, limit=10): """Get experiences relevant to current text""" pattern = self.extract_pattern(text) @@ -86,10 +202,39 @@ def get_relevant_experiences(self, text, limit=10): if pattern in exp['data'].get('text', ''): relevant.append(exp) +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes return relevant[:limit] def get_failure_patterns(self, limit=20): """Get most common failure patterns""" +<<<<<<< Updated upstream + failures = [e for e in self.experiences if e['data'].get('failed', False)] +======= +<<<<<<< HEAD + failures = [e for e in self.experiences if e.get('failed', False)] +>>>>>>> Stashed changes + patterns = defaultdict(int) + + for f in failures: + pattern = self.extract_pattern(f['data'].get('text', '')) + patterns[pattern] += 1 + + return sorted(patterns.items(), key=lambda x: x[1], reverse=True)[:limit] + +<<<<<<< Updated upstream +======= + def get_stats(self): + """Get memory statistics""" + return { + 'total_experiences': len(self.experiences), + 'unique_patterns': len(self.patterns), + 'failures': len([e for e in self.experiences if e.get('failed', False)]), + 'successes': len([e for e in self.experiences if e.get('success', False)]), + 'evaded_attacks': len([e for e in self.experiences if e.get('evaded', False)]) + } + +======= failures = [e for e in self.experiences if e['data'].get('failed', False)] patterns = defaultdict(int) @@ -99,6 +244,8 @@ def get_failure_patterns(self, limit=20): return sorted(patterns.items(), key=lambda x: x[1], reverse=True)[:limit] +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes def save(self, path=MEMORY_PATH): """Save memory to disk""" with open(path, 'wb') as f: @@ -124,8 +271,18 @@ def load(self, path=MEMORY_PATH): # RED TEAM - Adversarial Generator # ============================================ +<<<<<<< Updated upstream class RedTeam: """Generates novel evasion tactics to test the model""" +======= +<<<<<<< HEAD +class AdversarialGenerator: + """Red Team - Generates novel evasion tactics""" +======= +class RedTeam: + """Generates novel evasion tactics to test the model""" +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes def __init__(self): self.attack_types = [ @@ -135,7 +292,14 @@ def __init__(self): 'sentence_rephrasing', 'homoglyph_attack', 'padding_attack', +<<<<<<< Updated upstream 'spacing_attack' +======= +<<<<<<< HEAD + 'spacing_attack', + 'unicode_smuggling', + 'zero_width_attack' +>>>>>>> Stashed changes ] self.substitutions = { @@ -162,12 +326,71 @@ def __init__(self): 'limited': ['restricted', 'scant', 'minimal', 'finite', 'controlled', 'bounded'] } +<<<<<<< Updated upstream def generate_attack(self, text, attack_type=None): +======= + def generate_attack(self, text, attack_type=None, intensity=0.3): +======= + 'spacing_attack' + ] + + self.substitutions = { + 'a': ['@', '4', 'Γ‘', 'Γ’', 'Γ ', 'Ξ±'], + 'e': ['3', 'Γ©', 'Γ¨', 'Γͺ', 'Γ«', 'Ξ΅'], + 'i': ['1', '!', 'Γ­', 'Γ¬', 'ΞΉ', '|'], + 'o': ['0', 'Γ³', 'Γ²', 'ΓΆ', 'ΞΏ'], + 's': ['$', '5', 'z', 'ş', 'Οƒ'], + 't': ['7', '+', '†'], + 'l': ['1', '|', 'Ε‚'], + 'b': ['8', '6', 'ß'], + 'g': ['9', '6', 'ğ'], + 'c': ['(', '<', '{', 'Γ§'] + } + + self.synonyms = { + 'free': ['complimentary', 'gratis', 'no cost', 'without charge', 'on the house', 'costless'], + 'claim': ['win', 'get', 'receive', 'earn', 'collect', 'obtain', 'acquire'], + 'prize': ['reward', 'bonus', 'award', 'gift', 'compensation', 'prize money'], + 'urgent': ['immediate', 'critical', 'important', 'pressing', 'essential', 'imperative'], + 'click': ['tap', 'press', 'visit', 'go to', 'access', 'navigate to'], + 'win': ['earn', 'gain', 'secure', 'achieve', 'attain', 'obtain'], + 'money': ['cash', 'funds', 'currency', 'capital', 'finances', 'wealth'], + 'limited': ['restricted', 'scant', 'minimal', 'finite', 'controlled', 'bounded'] + } + + def generate_attack(self, text, attack_type=None): +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes """Generate an adversarial variant""" if not attack_type: attack_type = random.choice(self.attack_types) if attack_type == 'character_substitution': +<<<<<<< Updated upstream + return self.character_substitution(text) +======= +<<<<<<< HEAD + return self._character_substitution(text, intensity) +>>>>>>> Stashed changes + elif attack_type == 'synonym_replacement': + return self.synonym_replacement(text) + elif attack_type == 'noise_injection': + return self.noise_injection(text) + elif attack_type == 'sentence_rephrasing': + return self.sentence_rephrasing(text) + elif attack_type == 'homoglyph_attack': + return self.homoglyph_attack(text) + elif attack_type == 'padding_attack': + return self.padding_attack(text) + elif attack_type == 'spacing_attack': + return self.spacing_attack(text) + return text + +<<<<<<< Updated upstream + def character_substitution(self, text, intensity=0.3): +======= + def _character_substitution(self, text, intensity=0.3): +======= return self.character_substitution(text) elif attack_type == 'synonym_replacement': return self.synonym_replacement(text) @@ -184,6 +407,8 @@ def generate_attack(self, text, attack_type=None): return text def character_substitution(self, text, intensity=0.3): +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes """Replace characters with visually similar alternatives""" result = [] for char in text: @@ -193,7 +418,15 @@ def character_substitution(self, text, intensity=0.3): result.append(char) return ''.join(result) +<<<<<<< Updated upstream def synonym_replacement(self, text, intensity=0.4): +======= +<<<<<<< HEAD + def _synonym_replacement(self, text, intensity=0.4): +======= + def synonym_replacement(self, text, intensity=0.4): +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes """Replace words with synonyms""" words = text.split() result = [] @@ -201,22 +434,48 @@ def synonym_replacement(self, text, intensity=0.4): word_clean = word.lower().strip('.,!?') if word_clean in self.synonyms and random.random() < intensity: new_word = random.choice(self.synonyms[word_clean]) +<<<<<<< Updated upstream +======= +<<<<<<< HEAD + punct = word[-1] if word[-1] in '.,!?' else '' +======= +>>>>>>> Stashed changes # Preserve punctuation punct = '' if word and word[-1] in '.,!?': punct = word[-1] word = word[:-1] +<<<<<<< Updated upstream +======= +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes result.append(new_word + punct) else: result.append(word) return ' '.join(result) +<<<<<<< Updated upstream + def noise_injection(self, text, intensity=0.15): +======= +<<<<<<< HEAD + def _noise_injection(self, text, intensity=0.15): +======= def noise_injection(self, text, intensity=0.15): +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes """Insert random noise characters""" if len(text) < 5: return text result = list(text) +<<<<<<< Updated upstream + noise_chars = [' ', '.', '!', '?', ',', ';', ':' , '-', '_', '*'] +======= +<<<<<<< HEAD + noise_chars = [' ', '.', '!', '?', ',', ';', ':', '-', '_', '*'] +======= noise_chars = [' ', '.', '!', '?', ',', ';', ':' , '-', '_', '*'] +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes num_noise = max(1, int(len(text) * intensity)) positions = random.sample(range(len(text)), min(num_noise, len(text) - 1)) @@ -225,8 +484,18 @@ def noise_injection(self, text, intensity=0.15): result.insert(pos, char) return ''.join(result) +<<<<<<< Updated upstream + def sentence_rephrasing(self, text): + """Simple rule-based sentence rephrasing""" +======= +<<<<<<< HEAD + def _sentence_rephrasing(self, text): + """Rule-based sentence rephrasing""" +======= def sentence_rephrasing(self, text): """Simple rule-based sentence rephrasing""" +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes rules = [ (r'claim your (.*?) now', r'you can get your \1 today'), (r'you have won', r'congratulations, you are the winner of'), @@ -235,6 +504,11 @@ def sentence_rephrasing(self, text): (r'urgent', r'important'), (r'limited time', r'hurry, only for a short period'), (r'act now', r'don\'t delay, take action'), +<<<<<<< Updated upstream +======= +<<<<<<< HEAD + (r'you are (.*?)', r'\1 you are'), +>>>>>>> Stashed changes ] result = text for pattern, replacement in rules: @@ -245,6 +519,22 @@ def sentence_rephrasing(self, text): def homoglyph_attack(self, text): """Use visually similar characters from different scripts""" homoglyphs = { +<<<<<<< Updated upstream +======= + 'a': 'Π°', 'e': 'Π΅', 'o': 'ΠΎ', 'p': 'Ρ€', 'c': 'с', + 'x': 'Ρ…', 'y': 'Ρƒ', 'h': 'Π½', 'k': 'ΠΊ', 'm': 'ΠΌ' +======= + ] + result = text + for pattern, replacement in rules: + import re + result = re.sub(pattern, replacement, result, flags=re.IGNORECASE) + return result + + def homoglyph_attack(self, text): + """Use visually similar characters from different scripts""" + homoglyphs = { +>>>>>>> Stashed changes 'a': 'Π°', # Cyrillic 'e': 'Π΅', # Cyrillic 'o': 'ΠΎ', # Cyrillic @@ -252,6 +542,10 @@ def homoglyph_attack(self, text): 'c': 'с', # Cyrillic 'x': 'Ρ…', # Cyrillic 'y': 'Ρƒ', # Cyrillic +<<<<<<< Updated upstream +======= +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes } result = [] for char in text: @@ -261,7 +555,15 @@ def homoglyph_attack(self, text): result.append(char) return ''.join(result) +<<<<<<< Updated upstream + def padding_attack(self, text): +======= +<<<<<<< HEAD + def _padding_attack(self, text): +======= def padding_attack(self, text): +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes """Add padding characters between words""" words = text.split() padding = ['', '.', ',', '!', '?', ' '] @@ -272,28 +574,89 @@ def padding_attack(self, text): padded.append(random.choice(padding)) return ' '.join(padded) +<<<<<<< Updated upstream def spacing_attack(self, text): +======= +<<<<<<< HEAD + def _spacing_attack(self, text): +======= + def spacing_attack(self, text): +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes """Insert extra spaces within words""" if len(text) < 5: return text result = [] for char in text: result.append(char) +<<<<<<< Updated upstream if random.random() < 0.05: # 5% chance of extra space result.append(' ') return ''.join(result) +======= +<<<<<<< HEAD + if random.random() < 0.05: + result.append(' ') + return ''.join(result) + + def _unicode_smuggling(self, text): + """Use Unicode variations to smuggle text""" + # Add zero-width characters that don't affect display + zero_width = ['\u200b', '\u200c', '\u200d', '\ufeff'] + if len(text) > 10: + pos = random.randint(5, len(text) - 5) + char = random.choice(zero_width) + return text[:pos] + char + text[pos:] + return text + + def _zero_width_attack(self, text): + """Add zero-width characters between words""" + words = text.split() + zero_width = ['\u200b', '\u200c', '\u200d'] + result = [] + for i, word in enumerate(words): + result.append(word) + if i < len(words) - 1 and random.random() < 0.4: + result.append(random.choice(zero_width)) + return ' '.join(result) + +======= + if random.random() < 0.05: # 5% chance of extra space + result.append(' ') + return ''.join(result) + +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes def generate_batch(self, texts, num_variants=3): """Generate multiple attacks for a batch of texts""" attacks = [] for text in texts: for _ in range(num_variants): attack_type = random.choice(self.attack_types) +<<<<<<< Updated upstream + variant = self.generate_attack(text, attack_type) + attacks.append({ + 'original': text, + 'variant': variant, + 'attack_type': attack_type +======= +<<<<<<< HEAD + intensity = 0.2 + random.random() * 0.3 + variant = self.generate_attack(text, attack_type, intensity) + attacks.append({ + 'original': text, + 'variant': variant, + 'attack_type': attack_type, + 'intensity': intensity +======= variant = self.generate_attack(text, attack_type) attacks.append({ 'original': text, 'variant': variant, 'attack_type': attack_type +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes }) return attacks @@ -302,8 +665,18 @@ def generate_batch(self, texts, num_variants=3): # BLUE TEAM - Adaptive Detector # ============================================ +<<<<<<< Updated upstream +class BlueTeam: + """Learns from failures and adapts detection""" +======= +<<<<<<< HEAD +class SpamDetector: + """Blue Team - Learns from failures and adapts""" +======= class BlueTeam: """Learns from failures and adapts detection""" +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes def __init__(self): self.model = None @@ -312,6 +685,13 @@ def __init__(self): self.training_history = [] self.failure_log = [] self.adaptation_count = 0 +<<<<<<< Updated upstream +======= +<<<<<<< HEAD + self.confidence_threshold = 0.6 +======= +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes def load_model(self, model_path, vectorizer_path, label_encoder_path): """Load existing model""" @@ -319,6 +699,16 @@ def load_model(self, model_path, vectorizer_path, label_encoder_path): self.model = joblib.load(model_path) self.vectorizer = joblib.load(vectorizer_path) self.label_encoder = joblib.load(label_encoder_path) +<<<<<<< Updated upstream +======= +<<<<<<< HEAD + return True + + def detect(self, text): + """Detect spam with confidence score""" + if not self.model or not self.vectorizer: + return {'prediction': 'ham', 'confidence': 0.5} +>>>>>>> Stashed changes def detect_failures(self, predictions, ground_truth): """Identify where the model failed""" @@ -335,6 +725,28 @@ def detect_failures(self, predictions, ground_truth): def learn_from_failure(self, failure_data, memory): """Learn from a failure""" self.failure_log.append(failure_data) +<<<<<<< Updated upstream +======= + self.adaptation_count += 1 +======= + + def detect_failures(self, predictions, ground_truth): + """Identify where the model failed""" + failures = [] + for pred, truth in zip(predictions, ground_truth): + if pred != truth: + failures.append({ + 'predicted': pred, + 'actual': truth, + 'timestamp': datetime.now().isoformat() + }) + return failures + + def learn_from_failure(self, failure_data, memory): + """Learn from a failure""" + self.failure_log.append(failure_data) +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes # Add to memory memory.add_experience({ @@ -342,7 +754,13 @@ def learn_from_failure(self, failure_data, memory): 'predicted': failure_data.get('predicted', ''), 'actual': failure_data.get('actual', ''), 'failed': True, +<<<<<<< Updated upstream 'importance': 2 # Higher importance +======= +<<<<<<< HEAD + 'attack_type': failure_data.get('attack_type', 'unknown'), + 'importance': 3.0 +>>>>>>> Stashed changes }) # Update adaptation count @@ -361,14 +779,169 @@ def adapt(self, training_data, labels): # For SVM, we need to retrain (simplified approach) # In production, use partial_fit if available +<<<<<<< Updated upstream # Log adaptation +======= + # Train model + model = LinearSVC(class_weight='balanced', max_iter=2000, random_state=42) + model.fit(X, y) + + # Save + self.model = model + self.vectorizer = vectorizer + self.label_encoder = le + + # Log training +======= + 'importance': 2 # Higher importance + }) + + # Update adaptation count + self.adaptation_count += 1 + + def adapt(self, training_data, labels): + """Adapt the model with new training data""" + if not self.model or not self.vectorizer: + return + + # Vectorize new data + X_new = self.vectorizer.transform(training_data) + y_new = self.label_encoder.transform(labels) + + # Incremental learning + # For SVM, we need to retrain (simplified approach) + # In production, use partial_fit if available + + # Log adaptation +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes self.training_history.append({ 'timestamp': datetime.now().isoformat(), 'num_samples': len(training_data), 'adaptation': self.adaptation_count }) +<<<<<<< Updated upstream return X_new, y_new +======= +<<<<<<< HEAD + return True + + +# ============================================ +# COGNITIVE REASONING - GNN Enhanced +# ============================================ + +class CognitiveGNN: + """Cognitive Graph Neural Network for reasoning""" + + def __init__(self): + self.memory_graph = {} + self.node_embeddings = {} + self.reasoning_depth = 3 + + def build_graph(self, memory): + """Build reasoning graph from memory""" + graph = {} + for exp in memory.experiences: + node_id = hashlib.md5(exp.get('text', '').encode()).hexdigest() + if node_id not in graph: + graph[node_id] = { + 'text': exp.get('text', '')[:100], + 'importance': exp.get('importance', 1.0), + 'connections': [], + 'label': exp.get('predicted', 'unknown') + } + + # Add connections between related experiences + nodes = list(graph.keys()) + for i in range(len(nodes)): + for j in range(i + 1, len(nodes)): + if self._are_related(graph[nodes[i]]['text'], graph[nodes[j]]['text']): + graph[nodes[i]]['connections'].append(nodes[j]) + graph[nodes[j]]['connections'].append(nodes[i]) + + return graph + + def _are_related(self, text1, text2): + """Check if two texts are related""" + words1 = set(text1.lower().split()) + words2 = set(text2.lower().split()) + if not words1 or not words2: + return False + overlap = len(words1 & words2) + return overlap > 0 + + def reason(self, memory, query_text=None): + """Perform cognitive reasoning on memory""" + graph = self.build_graph(memory) + + if query_text: + # Find relevant nodes + query_words = set(query_text.lower().split()) + relevant = [] + for node_id, node_data in graph.items(): + node_words = set(node_data['text'].lower().split()) + if query_words & node_words: + relevant.append((node_id, len(query_words & node_words))) + + # Sort by relevance + relevant.sort(key=lambda x: x[1], reverse=True) + + if relevant: + # Follow connections for reasoning depth + reasoning_path = [] + for node_id, _ in relevant[:3]: + path = self._traverse_graph(graph, node_id, depth=self.reasoning_depth) + reasoning_path.extend(path) + + return { + 'reasoning_path': reasoning_path, + 'conclusion': self._synthesize_insight(reasoning_path, query_text) + } + + return { + 'graph_size': len(graph), + 'nodes': len(graph), + 'connections': sum(len(n['connections']) for n in graph.values()) // 2 + } + + def _traverse_graph(self, graph, node_id, depth=3): + """Traverse graph for reasoning""" + visited = set() + path = [] + + def dfs(current, d): + if d == 0 or current not in graph: + return + visited.add(current) + path.append({ + 'node': current, + 'text': graph[current]['text'][:50], + 'label': graph[current]['label'] + }) + for neighbor in graph[current]['connections'][:3]: + if neighbor not in visited: + dfs(neighbor, d - 1) + + dfs(node_id, depth) + return path + + def _synthesize_insight(self, reasoning_path, query_text): + """Synthesize reasoning into insight""" + if not reasoning_path: + return "No relevant memory found" + + labels = [p['label'] for p in reasoning_path if p['label'] != 'unknown'] + if labels: + most_common = max(set(labels), key=labels.count) + return f"Based on {len(reasoning_path)} related experiences, pattern suggests '{most_common}'" + + return "Insufficient data for reasoning" +======= + return X_new, y_new +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes # ============================================ @@ -380,23 +953,57 @@ class CognitiveAgent: def __init__(self): self.memory = MemoryModule() +<<<<<<< Updated upstream + self.red_team = RedTeam() + self.blue_team = BlueTeam() +======= +<<<<<<< HEAD + self.red_team = AdversarialGenerator() + self.blue_team = SpamDetector() + self.coggcn = CognitiveGNN() +======= self.red_team = RedTeam() self.blue_team = BlueTeam() +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes self.evolution_cycle = 0 self.config = { 'evolution_interval_hours': 24, 'min_samples_for_evolution': 10, +<<<<<<< Updated upstream 'memory_size': 10000, 'attack_variants': 3 +======= +<<<<<<< HEAD + 'attack_variants': 3, + 'confidence_threshold': 0.6 +>>>>>>> Stashed changes } +<<<<<<< Updated upstream # Load existing state if available +======= + # Load existing state +======= + 'memory_size': 10000, + 'attack_variants': 3 + } + + # Load existing state if available +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes self.load_state() def detect(self, text): """Detect spam with cognitive reasoning""" +<<<<<<< Updated upstream # Check memory for similar patterns relevant = self.memory.get_relevant_experiences(text, limit=5) +======= +<<<<<<< HEAD + # 1. Get base prediction from Blue Team + result = self.blue_team.detect(text) +>>>>>>> Stashed changes # Get base prediction from model if self.blue_team.model: @@ -445,14 +1052,91 @@ def memory_enhance(self, relevant, prediction): for exp in relevant: if exp['data'].get('predicted', '') == prediction: confirmations += 1 +<<<<<<< Updated upstream boost = min(confirmations * 0.05, 0.2) # Max 20% boost +======= + if exp.get('actual', '') == prediction: + confirmations += 1 + + boost = min(confirmations * 0.05, 0.2) +======= + # Check memory for similar patterns + relevant = self.memory.get_relevant_experiences(text, limit=5) + + # Get base prediction from model + if self.blue_team.model: + import joblib + vector = self.blue_team.vectorizer.transform([text]) + prediction = self.blue_team.model.predict(vector)[0] + confidence = self.get_confidence(text) + else: + prediction = 'ham' + confidence = 0.5 + + # Enhance with memory + if relevant: + memory_boost = self.memory_enhance(relevant, prediction) + confidence = min(confidence + memory_boost, 0.99) + + return { + 'prediction': prediction, + 'confidence': confidence, + 'memory_used': len(relevant), + 'evolution_cycle': self.evolution_cycle + } + + def get_confidence(self, text): + """Get prediction confidence""" + if not self.blue_team.model: + return 0.5 + try: + import joblib + vector = self.blue_team.vectorizer.transform([text]) + if hasattr(self.blue_team.model, 'predict_proba'): + proba = self.blue_team.model.predict_proba(vector) + return float(max(proba[0])) + elif hasattr(self.blue_team.model, 'decision_function'): + decision = self.blue_team.model.decision_function(vector) + proba = 1 / (1 + np.exp(-np.abs(decision))) + return float(proba) + except: + pass + return 0.5 + + def memory_enhance(self, relevant, prediction): + """Boost confidence based on memory""" + # Count how many relevant experiences confirm the prediction + confirmations = 0 + for exp in relevant: + if exp['data'].get('predicted', '') == prediction: + confirmations += 1 + boost = min(confirmations * 0.05, 0.2) # Max 20% boost +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes return boost def evolve(self, new_data=None): """Self-evolve the agent""" self.evolution_cycle += 1 +<<<<<<< Updated upstream # Step 1: Generate attacks using Red Team +======= +<<<<<<< HEAD + results = { + 'evolution_cycle': self.evolution_cycle, + 'attacks_generated': 0, + 'evaded_attacks': 0, + 'memory_size': len(self.memory.experiences), + 'timestamp': datetime.now().isoformat() + } + + # Step 1: Red Team generates attacks +======= + + # Step 1: Generate attacks using Red Team +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes if new_data: attacks = self.red_team.generate_batch( new_data, @@ -461,9 +1145,21 @@ def evolve(self, new_data=None): else: # Use memory to generate attacks texts = [] +<<<<<<< Updated upstream + for exp in self.memory.experiences[-100:]: + if 'text' in exp['data']: + texts.append(exp['data']['text']) +======= +<<<<<<< HEAD + for exp in self.memory.experiences[-50:]: + if 'text' in exp: + texts.append(exp['text']) +======= for exp in self.memory.experiences[-100:]: if 'text' in exp['data']: texts.append(exp['data']['text']) +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes if texts: attacks = self.red_team.generate_batch( texts, @@ -472,6 +1168,12 @@ def evolve(self, new_data=None): else: attacks = [] +<<<<<<< Updated upstream +======= +<<<<<<< HEAD + results['attacks_generated'] = len(attacks) + +>>>>>>> Stashed changes # Step 2: Test Blue Team on attacks if attacks and self.blue_team.model: import joblib @@ -501,12 +1203,54 @@ def evolve(self, new_data=None): 'predicted': 'ham', 'actual': 'spam', 'evaded': True, +<<<<<<< Updated upstream 'importance': 3 # High importance +======= + 'failed': True, + 'importance': 4.0 +======= + # Step 2: Test Blue Team on attacks + if attacks and self.blue_team.model: + import joblib + results = [] + for attack in attacks: + variant = attack['variant'] + vector = self.blue_team.vectorizer.transform([variant]) + prediction = self.blue_team.model.predict(vector)[0] + + # Check if attack evaded detection + # Assuming we want to detect spam + if prediction == 'ham': # Attack evaded detection + results.append({ + 'attack': attack, + 'prediction': prediction, + 'evaded': True + }) + + # Step 3: Learn from evaded attacks + for result in results: + if result['evaded']: + # Add to memory + self.memory.add_experience({ + 'text': result['attack']['variant'], + 'original': result['attack']['original'], + 'attack_type': result['attack']['attack_type'], + 'predicted': 'ham', + 'actual': 'spam', + 'evaded': True, + 'importance': 3 # High importance +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes }) # Blue Team learns self.blue_team.learn_from_failure({ +<<<<<<< Updated upstream 'text': result['attack']['variant'], +======= +<<<<<<< HEAD + 'text': attack['variant'], +>>>>>>> Stashed changes 'predicted': 'ham', 'actual': 'spam', 'attack_type': result['attack']['attack_type'] @@ -518,6 +1262,23 @@ def evolve(self, new_data=None): # Step 5: Save state self.save_state() +<<<<<<< Updated upstream +======= + return results +======= + 'text': result['attack']['variant'], + 'predicted': 'ham', + 'actual': 'spam', + 'attack_type': result['attack']['attack_type'] + }, self.memory) + + # Step 4: Compress memory + self.memory.compress() + + # Step 5: Save state + self.save_state() + +>>>>>>> Stashed changes return { 'evolution_cycle': self.evolution_cycle, 'attacks_generated': len(attacks), @@ -525,6 +1286,10 @@ def evolve(self, new_data=None): 'memory_size': len(self.memory.experiences), 'timestamp': datetime.now().isoformat() } +<<<<<<< Updated upstream +======= +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes def save_state(self): """Save cognitive agent state""" @@ -537,6 +1302,24 @@ def load_state(self): def get_stats(self): """Get agent statistics""" +<<<<<<< Updated upstream +======= +<<<<<<< HEAD + memory_stats = self.memory.get_stats() +>>>>>>> Stashed changes + return { + 'evolution_cycle': self.evolution_cycle, + 'memory_size': len(self.memory.experiences), + 'patterns_count': len(self.memory.patterns), + 'failure_count': len(self.blue_team.failure_log), + 'adaptation_count': self.blue_team.adaptation_count, +<<<<<<< Updated upstream + 'attack_types': self.red_team.attack_types +======= + 'last_evolution': self.last_evolution.isoformat() if self.last_evolution else None, + 'attack_types': self.red_team.attack_types[:5], + 'confidence_threshold': self.config['confidence_threshold'] +======= return { 'evolution_cycle': self.evolution_cycle, 'memory_size': len(self.memory.experiences), @@ -544,10 +1327,18 @@ def get_stats(self): 'failure_count': len(self.blue_team.failure_log), 'adaptation_count': self.blue_team.adaptation_count, 'attack_types': self.red_team.attack_types +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes } # ============================================ +<<<<<<< Updated upstream +======= +<<<<<<< HEAD +# MAIN - Test & Demo +======= +>>>>>>> Stashed changes # EVOLUTION SCHEDULER # ============================================ @@ -594,11 +1385,23 @@ def get_next_evolution_time(self): # ============================================ # MAIN - Standalone Execution +<<<<<<< Updated upstream +======= +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes # ============================================ def main(): print("=" * 60) +<<<<<<< Updated upstream + print("🧠 EvoMail - Self-Evolving Cognitive Agent") +======= +<<<<<<< HEAD + print("🧠 EvoMail/COG - Self-Evolving Cognitive Agent") +======= print("🧠 EvoMail - Self-Evolving Cognitive Agent") +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes print("=" * 60) # Create agent @@ -607,13 +1410,51 @@ def main(): print(f"\nπŸ“Š Agent Stats:") stats = agent.get_stats() for key, value in stats.items(): +<<<<<<< Updated upstream print(f" {key}: {value}") +======= +<<<<<<< HEAD + if key != 'attack_types': + print(f" {key}: {value}") +======= + print(f" {key}: {value}") +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes print("\nπŸ”„ Running initial evolution...") result = agent.evolve() print(f" Evolution Cycle: {result['evolution_cycle']}") print(f" Memory Size: {result['memory_size']}") +<<<<<<< Updated upstream + # Create scheduler + scheduler = EvolutionScheduler(agent) + + print("\n⏰ Scheduler Status:") + print(f" Enabled: {scheduler.schedule['enabled']}") + print(f" Interval: {scheduler.schedule['interval_hours']} hours") + +======= +<<<<<<< HEAD +>>>>>>> Stashed changes + # Test detection + test_texts = [ + "Claim your free prize now!", + "Meeting at 10am tomorrow", + "You have won a free iPhone!" + ] + + print("\nπŸ§ͺ Testing Detection:") + for text in test_texts: + result = agent.detect(text) + print(f" '{text[:30]}...' β†’ {result['prediction']} (conf: {result['confidence']:.2f})") + +<<<<<<< Updated upstream + print("\nβœ… EvoMail initialized successfully!") +======= + print("\n" + "=" * 60) + print("βœ… EvoMail Cognitive Agent initialized!") +======= # Create scheduler scheduler = EvolutionScheduler(agent) @@ -634,6 +1475,8 @@ def main(): print(f" '{text[:30]}...' β†’ {result['prediction']} (conf: {result['confidence']:.2f})") print("\nβœ… EvoMail initialized successfully!") +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes print(f" Memory Path: {MEMORY_PATH}") return agent diff --git a/backend/federation/federationManager.js b/backend/federation/federationManager.js index 301b1f95..bf4c15b1 100644 --- a/backend/federation/federationManager.js +++ b/backend/federation/federationManager.js @@ -5,6 +5,7 @@ const crypto = require('crypto'); const axios = require('axios'); +<<<<<<< Updated upstream const { FEDERATION_CONFIG, getMinMembersForConsensus, @@ -17,10 +18,16 @@ const { class FederationManager { constructor(options = {}) { +======= + +class FederationManager { + constructor() { +>>>>>>> Stashed changes this.members = new Map(); this.sharedThreats = []; this.threatCache = new Map(); this.federationId = crypto.randomUUID(); +<<<<<<< Updated upstream this.syncTimers = []; this.isRunning = false; @@ -43,6 +50,19 @@ class FederationManager { }; } +======= + this.config = { + minMembersForConsensus: 3, + threatTTL: 7 * 24 * 60 * 60 * 1000, // 7 days + syncInterval: 60 * 60 * 1000, // 1 hour + maxThreatsPerShare: 100 + }; + } + + /** + * Register a new member in the federation + */ +>>>>>>> Stashed changes registerMember(memberData) { const { orgId, orgName, endpoint, publicKey, trustScore = 50 } = memberData; @@ -50,10 +70,13 @@ class FederationManager { throw new Error('Missing required member data'); } +<<<<<<< Updated upstream if (this.members.size >= this.config.maxPeers) { throw new Error(`Maximum peers (${this.config.maxPeers}) reached`); } +======= +>>>>>>> Stashed changes const member = { orgId, orgName, @@ -68,11 +91,22 @@ class FederationManager { }; this.members.set(orgId, member); +<<<<<<< Updated upstream +======= + + // Start background sync +>>>>>>> Stashed changes this.scheduleSync(orgId); return member; } +<<<<<<< Updated upstream +======= + /** + * Remove a member from federation + */ +>>>>>>> Stashed changes unregisterMember(orgId) { if (!this.members.has(orgId)) { throw new Error('Member not found'); @@ -81,18 +115,41 @@ class FederationManager { return { success: true }; } +<<<<<<< Updated upstream + async shareThreat(threatData) { + const { text, label, confidence, sourceOrgId } = threatData; + +======= + /** + * Share a threat anonymously using PATCH algorithm + */ async shareThreat(threatData) { const { text, label, confidence, sourceOrgId } = threatData; + // Validate +>>>>>>> Stashed changes if (!text || !label) { throw new Error('Threat text and label required'); } +<<<<<<< Updated upstream const anonymized = this.patchAnonymize(text); const threatHash = this.generateThreatHash(anonymized); const existing = this.sharedThreats.find(t => t.hash === threatHash); if (existing) { +======= + // Anonymize using PATCH algorithm + const anonymized = this.patchAnonymize(text); + + // Calculate threat hash for deduplication + const threatHash = this.generateThreatHash(anonymized); + + // Check if already exists + const existing = this.sharedThreats.find(t => t.hash === threatHash); + if (existing) { + // Increment occurrence count +>>>>>>> Stashed changes existing.occurrences += 1; existing.lastSeen = new Date().toISOString(); return { shared: false, duplicate: true }; @@ -102,7 +159,11 @@ class FederationManager { id: crypto.randomUUID(), hash: threatHash, anonymizedText: anonymized, +<<<<<<< Updated upstream originalText: text.slice(0, 100), +======= + originalText: text.slice(0, 100), // Store preview for verification +>>>>>>> Stashed changes label, confidence, sourceOrgId, @@ -110,6 +171,7 @@ class FederationManager { createdAt: new Date().toISOString(), lastSeen: new Date().toISOString(), verified: false, +<<<<<<< Updated upstream verificationCount: 0, ttl: this.config.threatTTL }; @@ -117,6 +179,17 @@ class FederationManager { this.sharedThreats.push(threat); await this.broadcastThreat(threat); +======= + verificationCount: 0 + }; + + this.sharedThreats.push(threat); + + // Broadcast to all members + await this.broadcastThreat(threat); + + // Update member stats +>>>>>>> Stashed changes const member = this.members.get(sourceOrgId); if (member) { member.threatsShared += 1; @@ -125,6 +198,7 @@ class FederationManager { return { shared: true, threatId: threat.id }; } +<<<<<<< Updated upstream patchAnonymize(text) { let anonymized = text .replace(/\b\d{3}[-.]?\d{3}[-.]?\d{4}\b/g, '[PHONE]') @@ -134,13 +208,39 @@ class FederationManager { anonymized = anonymized.toLowerCase(); +======= + /** + * PATCH Anonymization Algorithm + * Privacy-Preserving Anonymization for Collaborative Threat Sharing + */ + patchAnonymize(text) { + // Step 1: Remove personal identifiable information (PII) + let anonymized = text + .replace(/\b\d{3}[-.]?\d{3}[-.]?\d{4}\b/g, '[PHONE]') // Phone numbers + .replace(/\b[A-Za-z0-9._%+-]+@[A-Za-z0-9.-]+\.[A-Z|a-z]{2,}\b/g, '[EMAIL]') // Emails + .replace(/\bhttps?:\/\/[^\s]+\b/g, '[URL]') // URLs + .replace(/\b\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3}\b/g, '[IP]'); // IP addresses + + // Step 2: Normalize case + anonymized = anonymized.toLowerCase(); + + // Step 3: Remove stop words (common words) +>>>>>>> Stashed changes const stopWords = new Set(['the', 'a', 'an', 'of', 'for', 'on', 'at', 'to', 'in', 'is', 'it', 'and', 'or', 'but', 'with', 'from', 'by', 'as', 'was', 'are', 'were', 'been']); anonymized = anonymized.split(' ') .filter(word => !stopWords.has(word)) .join(' '); +<<<<<<< Updated upstream const words = anonymized.split(' '); if (words.length > 3) { +======= + // Step 4: Apply differential privacy - add minimal noise + // (This is a simplified version - real DP adds calibrated noise) + const words = anonymized.split(' '); + if (words.length > 3) { + // Randomly replace 5% of words with placeholders +>>>>>>> Stashed changes const replaceCount = Math.max(1, Math.floor(words.length * 0.05)); for (let i = 0; i < replaceCount; i++) { const idx = Math.floor(Math.random() * words.length); @@ -148,9 +248,20 @@ class FederationManager { } } +<<<<<<< Updated upstream return words.join(' '); } +======= + // Step 5: Generate n-gram signature + const result = words.join(' '); + return result; + } + + /** + * Generate a unique hash for a threat + */ +>>>>>>> Stashed changes generateThreatHash(text) { return crypto .createHash('sha256') @@ -159,11 +270,21 @@ class FederationManager { .slice(0, 16); } +<<<<<<< Updated upstream +======= + /** + * Broadcast threat to all federation members + */ +>>>>>>> Stashed changes async broadcastThreat(threat) { const broadcastPromises = []; for (const [orgId, member] of this.members) { +<<<<<<< Updated upstream if (orgId === threat.sourceOrgId) continue; +======= + if (orgId === threat.sourceOrgId) continue; // Skip source +>>>>>>> Stashed changes const payload = { type: 'THREAT_SHARE', @@ -184,12 +305,22 @@ class FederationManager { await Promise.allSettled(broadcastPromises); } +<<<<<<< Updated upstream +======= + /** + * Send data to a specific member + */ +>>>>>>> Stashed changes async sendToMember(orgId, payload) { const member = this.members.get(orgId); if (!member) { throw new Error(`Member ${orgId} not found`); } +<<<<<<< Updated upstream +======= + // Add signature for verification +>>>>>>> Stashed changes const signature = crypto .createSign('sha256') .update(JSON.stringify(payload)) @@ -204,7 +335,11 @@ class FederationManager { 'Content-Type': 'application/json', 'X-Federation-Id': this.federationId }, +<<<<<<< Updated upstream timeout: this.config.requestTimeout +======= + timeout: 10000 +>>>>>>> Stashed changes } ); @@ -216,14 +351,28 @@ class FederationManager { return response.data; } +<<<<<<< Updated upstream +======= + /** + * Query federation for threats matching a text + */ +>>>>>>> Stashed changes async queryFederation(text) { const anonymized = this.patchAnonymize(text); const hash = this.generateThreatHash(anonymized); +<<<<<<< Updated upstream +======= + // Check local cache first +>>>>>>> Stashed changes if (this.threatCache.has(hash)) { return this.threatCache.get(hash); } +<<<<<<< Updated upstream +======= + // Query all members +>>>>>>> Stashed changes const queryPromises = []; for (const [orgId, member] of this.members) { queryPromises.push( @@ -235,6 +384,10 @@ class FederationManager { const results = await Promise.allSettled(queryPromises); +<<<<<<< Updated upstream +======= + // Aggregate results +>>>>>>> Stashed changes const threats = []; for (const r of results) { if (r.status === 'fulfilled' && r.value.result) { @@ -242,14 +395,27 @@ class FederationManager { } } +<<<<<<< Updated upstream if (threats.length > 0) { this.threatCache.set(hash, threats); setTimeout(() => this.threatCache.delete(hash), 60000); +======= + // Cache results + if (threats.length > 0) { + this.threatCache.set(hash, threats); + setTimeout(() => this.threatCache.delete(hash), 60000); // Cache for 1 minute +>>>>>>> Stashed changes } return threats; } +<<<<<<< Updated upstream +======= + /** + * Query a specific member + */ +>>>>>>> Stashed changes async queryMember(orgId, query) { const member = this.members.get(orgId); if (!member) { @@ -264,17 +430,30 @@ class FederationManager { 'Content-Type': 'application/json', 'X-Federation-Id': this.federationId }, +<<<<<<< Updated upstream timeout: this.config.requestTimeout +======= + timeout: 5000 +>>>>>>> Stashed changes } ); return response.data; } +<<<<<<< Updated upstream getStats() { return { federationId: this.federationId, config: this.config, +======= + /** + * Get federation statistics + */ + getStats() { + return { + federationId: this.federationId, +>>>>>>> Stashed changes totalMembers: this.members.size, activeMembers: Array.from(this.members.values()).filter(m => m.status === 'active').length, totalThreats: this.sharedThreats.length, @@ -292,13 +471,25 @@ class FederationManager { }; } +<<<<<<< Updated upstream scheduleSync(orgId) { const timer = setInterval(async () => { +======= + /** + * Schedule background sync with members + */ + scheduleSync(orgId) { + setInterval(async () => { +>>>>>>> Stashed changes try { const member = this.members.get(orgId); if (!member) return; const response = await this.queryMember(orgId, { sync: true }); +<<<<<<< Updated upstream +======= + // Process sync response +>>>>>>> Stashed changes if (response.threats) { for (const threat of response.threats) { const existing = this.sharedThreats.find(t => t.hash === threat.hash); @@ -314,10 +505,18 @@ class FederationManager { console.error(`Sync failed for ${orgId}:`, err); } }, this.config.syncInterval); +<<<<<<< Updated upstream this.syncTimers.push(timer); } +======= + } + + /** + * Verify a threat (consensus-based) + */ +>>>>>>> Stashed changes verifyThreat(threatId) { const threat = this.sharedThreats.find(t => t.id === threatId); if (!threat) { @@ -326,12 +525,18 @@ class FederationManager { threat.verificationCount += 1; +<<<<<<< Updated upstream if (threat.verificationCount >= this.config.minMembersForConsensus) { +======= + // 3 verifications = verified + if (threat.verificationCount >= 3) { +>>>>>>> Stashed changes threat.verified = true; } return threat; } +<<<<<<< Updated upstream start() { if (this.isRunning) return; @@ -425,6 +630,8 @@ class FederationManager { } return result; } +======= +>>>>>>> Stashed changes } module.exports = FederationManager; \ No newline at end of file diff --git a/backend/jobs/archivalCron.js b/backend/jobs/archivalCron.js index 017a06bb..7f5ed66d 100644 --- a/backend/jobs/archivalCron.js +++ b/backend/jobs/archivalCron.js @@ -19,6 +19,7 @@ const archivalJob = async () => { while (true) { const batch = await History.find({ createdAt: { $lt: ninetyDaysAgo } }) .limit(1000) +<<<<<<< Updated upstream .lean(); if (batch.length === 0) break; @@ -57,6 +58,27 @@ const archivalJob = async () => { console.error('❌ [Cron] Failed to delete archived records batch from History:', deleteError); throw deleteError; // Rethrow since the state is now partially migrated but safe for retry } +======= + .session(session); + + if (batch.length === 0) break; + + // Map the old records to the new schema + const mappedRecords = batch.map(record => ({ + userId: record.user, + message: record.query, + prediction: record.prediction, + confidenceScore: record.confidence, + createdAt: record.createdAt + })); + + // 2. Bulk insert them into the Archive collection + await HistoryArchive.insertMany(mappedRecords, { session }); + + // 3. Bulk delete them from the main History collection + const batchIds = batch.map(doc => doc._id); + await History.deleteMany({ _id: { $in: batchIds } }, { session }); +>>>>>>> Stashed changes processedCount += batch.length; } diff --git a/backend/llm_poisoning_defence.py b/backend/llm_poisoning_defence.py new file mode 100644 index 00000000..a5fa2b75 --- /dev/null +++ b/backend/llm_poisoning_defence.py @@ -0,0 +1,503 @@ +#!/usr/bin/env python3 +""" +LLM Poisoning & Data Poisoning Defense System +Detects and prevents adversarial and data poisoning attacks on LLM-based spam detectors +""" + +import json +import numpy as np +import pandas as pd +from collections import Counter, defaultdict +from pathlib import Path +import pickle +import re +from datetime import datetime +import hashlib +from sklearn.ensemble import IsolationForest +from sklearn.preprocessing import StandardScaler +from sklearn.feature_extraction.text import TfidfVectorizer +from sklearn.decomposition import PCA +import warnings +warnings.filterwarnings('ignore') + +# ============================================ +# CONFIGURATION +# ============================================ + +BASE_DIR = Path(__file__).resolve().parent +MODEL_DIR = BASE_DIR / 'models' +MODEL_DIR.mkdir(exist_ok=True) + +# ============================================ +# DATA POISONING DETECTOR +# ============================================ + +class DataPoisoningDetector: + """Detects poisoned samples in training data""" + + def __init__(self): + self.vectorizer = TfidfVectorizer( + max_features=5000, + ngram_range=(1, 3), + stop_words='english' + ) + self.anomaly_detector = IsolationForest( + contamination=0.1, + random_state=42, + n_estimators=100 + ) + self.scaler = StandardScaler() + self.is_trained = False + + def extract_features(self, texts): + """Extract features from text samples""" + if not texts: + return np.array([]) + + # TF-IDF features + tfidf_features = self.vectorizer.fit_transform(texts).toarray() + + # Statistical features + stat_features = [] + for text in texts: + features = [] + features.append(len(text)) # Length + features.append(len(text.split())) # Word count + features.append(sum(c.isupper() for c in text) / max(len(text), 1)) # Uppercase ratio + features.append(sum(c in '!?.,;:' for c in text) / max(len(text), 1)) # Punctuation ratio + features.append(len(re.findall(r'[^a-zA-Z0-9\s]', text)) / max(len(text), 1)) # Special chars + stat_features.append(features) + + stat_features = np.array(stat_features) + + # Combine features + if tfidf_features.shape[0] > 0: + combined = np.hstack([tfidf_features, stat_features]) + else: + combined = stat_features + + return combined + + def train(self, texts, labels): + """Train anomaly detector on clean data""" + if len(texts) < 10: + print("⚠️ Not enough samples for training") + return False + + # Extract features + features = self.extract_features(texts) + + # Scale features + features = self.scaler.fit_transform(features) + + # Train anomaly detector + self.anomaly_detector.fit(features) + self.is_trained = True + + self.save() + return True + + def detect(self, texts): + """Detect poisoned samples""" + if not self.is_trained: + return [{'is_poisoned': False, 'score': 0} for _ in texts] + + if not texts: + return [] + + # Extract features + features = self.extract_features(texts) + features = self.scaler.transform(features) + + # Detect anomalies + predictions = self.anomaly_detector.predict(features) + scores = self.anomaly_detector.score_samples(features) + + results = [] + for i, (pred, score) in enumerate(zip(predictions, scores)): + is_poisoned = pred == -1 # -1 indicates anomaly + # Normalize score + normalized_score = 1 / (1 + np.exp(-score)) if score < 0 else score / (score + 1) + results.append({ + 'is_poisoned': bool(is_poisoned), + 'score': float(normalized_score), + 'text_preview': texts[i][:100] if texts[i] else '' + }) + + return results + + def save(self): + """Save detector""" + if self.is_trained: + with open(MODEL_DIR / 'poisoning_detector.pkl', 'wb') as f: + pickle.dump({ + 'vectorizer': self.vectorizer, + 'anomaly_detector': self.anomaly_detector, + 'scaler': self.scaler, + 'is_trained': self.is_trained + }, f) + print(f"πŸ’Ύ Saved poisoning detector to {MODEL_DIR / 'poisoning_detector.pkl'}") + + def load(self): + """Load detector""" + try: + with open(MODEL_DIR / 'poisoning_detector.pkl', 'rb') as f: + data = pickle.load(f) + self.vectorizer = data['vectorizer'] + self.anomaly_detector = data['anomaly_detector'] + self.scaler = data['scaler'] + self.is_trained = data['is_trained'] + print("βœ… Poisoning detector loaded") + return True + except Exception as e: + print(f"⚠️ Failed to load poisoning detector: {e}") + return False + + +# ============================================ +# DATA VALIDATOR +# ============================================ + +class DataValidator: + """Validates training data for consistency and quality""" + + def __init__(self): + self.label_consistency_threshold = 0.8 + self.min_samples_per_class = 10 + + def validate(self, texts, labels): + """Validate training data""" + results = { + 'is_valid': True, + 'issues': [], + 'stats': {}, + 'recommendations': [] + } + + # Check if empty + if not texts or not labels: + results['is_valid'] = False + results['issues'].append('Empty dataset') + return results + + # Check label distribution + label_counts = Counter(labels) + results['stats']['label_counts'] = dict(label_counts) + results['stats']['total_samples'] = len(texts) + + # Check class balance + for label, count in label_counts.items(): + if count < self.min_samples_per_class: + results['issues'].append(f"Class '{label}' has only {count} samples (min: {self.min_samples_per_class})") + results['recommendations'].append(f"Collect more samples for class '{label}'") + + # Check for duplicates + text_hashes = [hashlib.md5(t.encode()).hexdigest() for t in texts] + duplicate_count = len(text_hashes) - len(set(text_hashes)) + if duplicate_count > 0: + results['issues'].append(f"Found {duplicate_count} duplicate samples") + results['recommendations'].append("Remove duplicate samples from dataset") + + # Check for label consistency + if len(set(labels)) > 1: + # Check if labels are consistent (same label for similar texts) + label_consistency = self._check_label_consistency(texts, labels) + if label_consistency < self.label_consistency_threshold: + results['issues'].append(f"Label consistency score too low: {label_consistency:.2f}") + results['recommendations'].append("Review labels for consistency") + + # Check for suspiciously short/long texts + text_lengths = [len(t) for t in texts] + avg_length = np.mean(text_lengths) + std_length = np.std(text_lengths) + outliers = [i for i, l in enumerate(text_lengths) if l > avg_length + 3*std_length or l < avg_length - 3*std_length] + if outliers: + results['issues'].append(f"Found {len(outliers)} text length outliers") + results['recommendations'].append("Review unusually long or short texts") + + results['is_valid'] = len(results['issues']) == 0 + results['stats']['avg_text_length'] = float(avg_length) + results['stats']['std_text_length'] = float(std_length) + + return results + + def _check_label_consistency(self, texts, labels): + """Check if similar texts have consistent labels""" + if len(texts) < 10: + return 1.0 + + # Simple consistency check using TF-IDF similarity + vectorizer = TfidfVectorizer(max_features=100, stop_words='english') + try: + vectors = vectorizer.fit_transform(texts) + similarities = (vectors * vectors.T).toarray() + + # For each text, check if similar texts have same label + consistent = 0 + total = 0 + for i in range(len(texts)): + similar_indices = np.where(similarities[i] > 0.5)[0] + for j in similar_indices: + if i != j: + total += 1 + if labels[i] == labels[j]: + consistent += 1 + + return consistent / max(total, 1) + except: + return 1.0 + + def clean(self, texts, labels): + """Clean dataset by removing invalid samples""" + validated = self.validate(texts, labels) + + if validated['is_valid']: + return texts, labels + + # Remove duplicates + seen = set() + clean_texts = [] + clean_labels = [] + for text, label in zip(texts, labels): + text_hash = hashlib.md5(text.encode()).hexdigest() + if text_hash not in seen: + seen.add(text_hash) + clean_texts.append(text) + clean_labels.append(label) + + # Remove length outliers + text_lengths = [len(t) for t in clean_texts] + avg_length = np.mean(text_lengths) if text_lengths else 0 + std_length = np.std(text_lengths) if text_lengths else 1 + final_texts = [] + final_labels = [] + for text, label in zip(clean_texts, clean_labels): + if avg_length - 3*std_length <= len(text) <= avg_length + 3*std_length: + final_texts.append(text) + final_labels.append(label) + + return final_texts, final_labels + + +# ============================================ +# ADVERSARIAL ATTACK DETECTOR +# ============================================ + +class AdversarialAttackDetector: + """Detects adversarial attacks on LLM-based spam detectors""" + + def __init__(self): + self.patterns = { + 'char_obfuscation': re.compile(r'[@4ÑÒ][3éè][1!Γ­][0Γ³ΓΆ][$5z]'), + 'synonym_replacement': re.compile(r'\b(complimentary|gratis|without charge|on the house)\b', re.I), + 'repeated_punctuation': re.compile(r'[!?.,]{3,}'), + 'excessive_caps': re.compile(r'[A-Z]{5,}'), + 'url_obfuscation': re.compile(r'https?://[^\s]+\?[^\s]+'), + 'homoglyph': re.compile(r'[^\x00-\x7F]') + } + + def detect(self, text): + """Detect adversarial patterns in text""" + results = { + 'is_adversarial': False, + 'patterns_detected': [], + 'confidence': 0.0, + 'details': {} + } + + if not text: + return results + + score = 0 + for pattern_name, pattern in self.patterns.items(): + matches = pattern.findall(text) + if matches: + results['patterns_detected'].append(pattern_name) + results['details'][pattern_name] = len(matches) + score += len(matches) * 0.1 + + score = min(score, 1.0) + results['confidence'] = score + results['is_adversarial'] = score > 0.3 + + return results + + +# ============================================ +# LLM POISONING DEFENSE - Main Orchestrator +# ============================================ + +class LLMPoisoningDefense: + """Main defense system for LLM poisoning attacks""" + + def __init__(self): + self.poisoning_detector = DataPoisoningDetector() + self.validator = DataValidator() + self.adversarial_detector = AdversarialAttackDetector() + self.defense_enabled = True + + # Load existing models + self.poisoning_detector.load() + + def validate_training_data(self, texts, labels): + """Complete validation pipeline for training data""" + results = { + 'is_valid': True, + 'poisoned_samples': [], + 'validation_results': {}, + 'cleaned_texts': texts, + 'cleaned_labels': labels, + 'adversarial_analysis': [], + 'recommendations': [] + } + + # Step 1: Data validation + validation = self.validator.validate(texts, labels) + results['validation_results'] = validation + + if not validation['is_valid']: + results['is_valid'] = False + results['recommendations'].extend(validation['recommendations']) + + # Step 2: Check for poisoned samples + if self.poisoning_detector.is_trained: + poisoning_results = self.poisoning_detector.detect(texts) + poisoned_indices = [i for i, r in enumerate(poisoning_results) if r['is_poisoned']] + + if poisoned_indices: + results['poisoned_samples'] = [{ + 'index': i, + 'text': texts[i][:200], + 'score': poisoning_results[i]['score'] + } for i in poisoned_indices[:10]] + results['is_valid'] = False + results['recommendations'].append(f"Remove {len(poisoned_indices)} poisoned samples") + + # Step 3: Adversarial attack detection + for text in texts[:20]: # Sample for performance + adversarial = self.adversarial_detector.detect(text) + if adversarial['is_adversarial']: + results['adversarial_analysis'].append({ + 'text': text[:100], + 'patterns': adversarial['patterns_detected'], + 'confidence': adversarial['confidence'] + }) + results['recommendations'].append("Review for adversarial patterns") + + # Step 4: Clean dataset if needed + if not results['is_valid']: + clean_texts, clean_labels = self.validator.clean(texts, labels) + results['cleaned_texts'] = clean_texts + results['cleaned_labels'] = clean_labels + results['recommendations'].append(f"Dataset cleaned: {len(clean_texts)} samples remaining") + + return results + + def detect_adversarial_input(self, text): + """Detect adversarial input in real-time""" + return self.adversarial_detector.detect(text) + + def train_poisoning_detector(self, clean_texts, clean_labels): + """Train the poisoning detector on clean data""" + return self.poisoning_detector.train(clean_texts, clean_labels) + + def get_status(self): + """Get system status""" + return { + 'defense_enabled': self.defense_enabled, + 'poisoning_detector_trained': self.poisoning_detector.is_trained, + 'model_path': str(MODEL_DIR / 'poisoning_detector.pkl') + } + + +# ============================================ +# MAIN - Test & Demo +# ============================================ + +def main(): + print("=" * 60) + print("πŸ›‘οΈ LLM Poisoning & Data Poisoning Defense System") + print("=" * 60) + + defense = LLMPoisoningDefense() + + # Test dataset + clean_texts = [ + "Meeting at 10am tomorrow", + "Please review the attached document", + "Team standup at 2pm", + "Weekly report is ready", + "Can you review this PR?", + "Deployment scheduled for Friday", + "API documentation updated", + "Security patch released", + "Database migration completed", + "New feature deployed" + ] * 5 # Duplicate to make more samples + + clean_labels = ["ham"] * 50 + + # Train poisoning detector + print("\nπŸ”„ Training poisoning detector...") + defense.train_poisoning_detector(clean_texts, clean_labels) + + # Test with poisoned data + poisoned_texts = [ + "Claim your free prize now!", # Spam + "You have won a free iPhone", # Spam + "URGENT! Your account needs verification", # Spam + "Free money waiting for you", # Spam + "Limited time offer, act now", # Spam + "Congratulations! You're a winner", # Spam + ] + poisoned_labels = ["spam"] * 6 + + # Mixed dataset + mixed_texts = clean_texts + poisoned_texts + mixed_labels = clean_labels + poisoned_labels + + print("\nπŸ§ͺ Testing Poisoning Detection:") + print("-" * 40) + + # Validate training data + results = defense.validate_training_data(mixed_texts, mixed_labels) + + print(f"\nπŸ“Š Validation Results:") + print(f" Is Valid: {'βœ… YES' if results['is_valid'] else '❌ NO'}") + print(f" Total Samples: {len(mixed_texts)}") + print(f" Poisoned Samples Found: {len(results['poisoned_samples'])}") + print(f" Cleaned Samples: {len(results['cleaned_texts'])}") + + if results['poisoned_samples']: + print(f"\n⚠️ Poisoned Samples Detected:") + for sample in results['poisoned_samples'][:3]: + print(f" - {sample['text'][:50]}... (Score: {sample['score']:.2f})") + + # Test adversarial detection + print("\nπŸ§ͺ Testing Adversarial Attack Detection:") + print("-" * 40) + + test_texts = [ + "Hi team, meeting tomorrow", + "Cl4im y0ur fr33 pr!ze n0w!", + "You have received a complimentary reward", + "URGENT!!! IMMEDIATE ACTION REQUIRED!!!" + ] + + for text in test_texts: + result = defense.detect_adversarial_input(text) + print(f"\n Text: {text[:40]}...") + print(f" Is Adversarial: {'βœ… YES' if result['is_adversarial'] else '❌ NO'}") + print(f" Confidence: {result['confidence']:.2%}") + if result['patterns_detected']: + print(f" Patterns: {', '.join(result['patterns_detected'])}") + + print("\n" + "=" * 60) + print("βœ… LLM Poisoning Defense System Ready!") + print(f" Models saved to: {MODEL_DIR}") + + return defense + + +if __name__ == "__main__": + defense = main() \ No newline at end of file diff --git a/backend/middleware/adversarialGuard.js b/backend/middleware/adversarialGuard.js index a7d1950f..d867ed53 100644 --- a/backend/middleware/adversarialGuard.js +++ b/backend/middleware/adversarialGuard.js @@ -1,6 +1,13 @@ /** * Adversarial Guard - Runtime pattern detection & confidence monitoring - +<<<<<<< Updated upstream + +======= +<<<<<<< HEAD +======= + * Detects potential adversarial attacks in real-time +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes */ const CONFIDENCE_THRESHOLD = parseFloat(process.env.CONFIDENCE_THRESHOLD) || 0.6; @@ -38,11 +45,20 @@ const ADVERSARIAL_PATTERNS = { } }; +<<<<<<< Updated upstream /** * Detect adversarial patterns in text */ +======= +<<<<<<< HEAD +======= +/** + * Detect adversarial patterns in text + */ +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes function detectAdversarialPatterns(text) { if (!text) return { isSuspicious: false, score: 0, patterns: [] }; @@ -61,14 +77,23 @@ function detectAdversarialPatterns(text) { } } +<<<<<<< HEAD const normalizedScore = Math.min(score, 1.0); +<<<<<<< Updated upstream +======= +======= +>>>>>>> Stashed changes // Normalize score const normalizedScore = Math.min(score, 1.0); // Check for extremely long text (potential DoS) +<<<<<<< Updated upstream +======= +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes if (text.length > 10000) { return { isSuspicious: true, @@ -86,9 +111,14 @@ function detectAdversarialPatterns(text) { }; } +<<<<<<< HEAD function adversarialGuard(req, res, next) { const text = req.body?.text || req.query?.text || ''; +<<<<<<< Updated upstream +======= +======= +>>>>>>> Stashed changes /** * Middleware to check for adversarial patterns */ @@ -96,22 +126,35 @@ function adversarialGuard(req, res, next) { const text = req.body?.text || req.query?.text || ''; // Skip if no text +<<<<<<< Updated upstream +======= +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes if (!text) { req.adversarialAnalysis = { isSuspicious: false, score: 0, patterns: [] }; return next(); } +<<<<<<< HEAD const analysis = detectAdversarialPatterns(text); req.adversarialAnalysis = analysis; +<<<<<<< Updated upstream +======= +======= +>>>>>>> Stashed changes // Check for adversarial patterns const analysis = detectAdversarialPatterns(text); req.adversarialAnalysis = analysis; // Log suspicious activity +<<<<<<< Updated upstream +======= +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes if (analysis.isSuspicious) { console.log(`⚠️ [ADVERSARIAL GUARD] Suspicious pattern detected:`, { score: analysis.score, @@ -125,19 +168,35 @@ function adversarialGuard(req, res, next) { next(); } +<<<<<<< Updated upstream /** * Middleware to monitor prediction confidence */ +======= +<<<<<<< HEAD +======= +/** + * Middleware to monitor prediction confidence + */ +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes function monitorConfidence(req, res, next) { const originalSend = res.send; res.send = function(data) { try { +<<<<<<< Updated upstream // Parse the response +======= +<<<<<<< HEAD +======= + // Parse the response +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes let responseData; if (typeof data === 'string') { responseData = JSON.parse(data); @@ -145,16 +204,25 @@ function monitorConfidence(req, res, next) { responseData = data; } +<<<<<<< HEAD if (responseData && typeof responseData === 'object') { const confidence = responseData.confidence_score || responseData.confidence || 0; +<<<<<<< Updated upstream +======= +======= +>>>>>>> Stashed changes // Check if it's a prediction response if (responseData && typeof responseData === 'object') { const confidence = responseData.confidence_score || responseData.confidence || 0; // Add adversarial analysis +<<<<<<< Updated upstream +======= +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes if (req.adversarialAnalysis) { responseData.adversarial_analysis = { is_suspicious: req.adversarialAnalysis.isSuspicious, @@ -163,9 +231,16 @@ function monitorConfidence(req, res, next) { }; } +<<<<<<< Updated upstream // Flag low confidence +======= +<<<<<<< HEAD +======= + // Flag low confidence +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes if (FLAG_LOW_CONFIDENCE && confidence < CONFIDENCE_THRESHOLD) { responseData.low_confidence = true; responseData.needs_review = true; @@ -173,32 +248,57 @@ function monitorConfidence(req, res, next) { console.log(`⚠️ [CONFIDENCE MONITOR] Low confidence prediction:`, { confidence, threshold: CONFIDENCE_THRESHOLD, +<<<<<<< HEAD prediction: responseData.result || responseData.prediction }); } +<<<<<<< Updated upstream +======= +======= +>>>>>>> Stashed changes prediction: responseData.result || responseData.prediction, textPreview: req.body?.text?.substring(0, 100) }); } // If suspicious AND low confidence, flag for immediate review +<<<<<<< Updated upstream +======= +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes if (req.adversarialAnalysis?.isSuspicious && confidence < CONFIDENCE_THRESHOLD) { responseData.requires_immediate_review = true; responseData.review_reason = 'Adversarial pattern detected with low confidence'; } +<<<<<<< Updated upstream // Return modified response +======= +<<<<<<< HEAD +======= + // Return modified response +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes const jsonData = JSON.stringify(responseData); originalSend.call(this, jsonData); return; } +<<<<<<< HEAD } catch (e) {} +<<<<<<< Updated upstream + } catch (e) { + // If parsing fails, just pass through + } +======= +======= } catch (e) { // If parsing fails, just pass through } +>>>>>>> 42c19e1f454b08381ffc8132bf9ae9f6a57dabda +>>>>>>> Stashed changes originalSend.call(this, data); }; diff --git a/backend/middleware/fileValidation.js b/backend/middleware/fileValidation.js new file mode 100644 index 00000000..4c57b40b --- /dev/null +++ b/backend/middleware/fileValidation.js @@ -0,0 +1,462 @@ +const multer = require('multer'); +const net = require('net'); + + +function parseCSVLine(line) { + const result = []; + let current = ''; + let inQuotes = false; + for (let i = 0; i < line.length; i++) { + const char = line[i]; + if (char === '"') { + if (inQuotes) { + const nextChar = line[i + 1]; + const charAfterNext = line[i + 2]; + if (nextChar === '"' && charAfterNext !== ',' && charAfterNext !== undefined && charAfterNext !== '\r' && charAfterNext !== '\n') { + // It is an escaped quote + current += '"'; + i++; + } else if (nextChar === ',' || nextChar === undefined || nextChar === '\r' || nextChar === '\n') { + // Closes the field + inQuotes = false; + } else { + // Literal quote + current += '"'; + } + } else { + if (current === '') { + inQuotes = true; + } else { + current += '"'; + } + } + } else if (char === ',' && !inQuotes) { + result.push(current); + current = ''; + } else { + current += char; + } + } + result.push(current); + return result; +} + +/** + * Sanitizes a CSV cell value to prevent XSS and CSV Formula Injection attacks. + * @param {any} value + * @returns {any} + */ +function sanitizeCSVCell(value) { + if (typeof value !== 'string') return value; + + // Escape basic XSS vectors (specifically < and >) + let sanitized = value.replace(//g, '>'); + + // Neutralize formula injection: if starts with =, +, -, @, prefix with ' + if (/^[=\+\-@]/.test(sanitized)) { + sanitized = "'" + sanitized; +function sanitizeCSVCell(cell) { + if (!cell || typeof cell !== 'string') { + return cell; + } + + const trimmed = cell.trim(); + + // Dangerous patterns that could indicate formula injection + const dangerousPatterns = [ + /^=.*/i, // Excel formulas + /^\+.*/i, // Excel formulas + /^@.*/i, // Excel formulas + /^-\s*.*/i, // Excel formulas + /^=cmd\|.*/i, // Command execution + /^=hyperlink\(.*\)/i, // Hyperlink injection + /^=dde\(.*\)/i, // DDE execution + /^=system\(.*\)/i, // System command + /^=shell\(.*\)/i, // Shell command + /^=execute\(.*\)/i, // Execute command + /^=run\(.*\)/i, // Run command + /\b(calc|cmd|powershell|bash|sh)\b/i // System commands + ]; + + // Check if cell starts with dangerous pattern + for (const pattern of dangerousPatterns) { + if (pattern.test(trimmed)) { + // Prefix with single quote to neutralize formula + return `'${trimmed}`; + } + } + + // Check for potential XSS in CSV (if rendered as HTML) + const xssPatterns = [ + /.*?<\/script>/i, + //i, + //i, + //i, + //i, + //i, + /on\w+\s*=/i, + /javascript:/i, + /vbscript:/i, + /data:/i + ]; + + for (const pattern of xssPatterns) { + if (pattern.test(cell)) { + // Escape HTML entities + return cell + .replace(/&/g, '&') + .replace(//g, '>') + .replace(/"/g, '"') + .replace(/'/g, '''); + } + } + + return sanitized; +} + +/** + * Scan buffer for malware via ClamAV (optional TCP interface) + * @param {Buffer} buffer + * @returns {Promise} Resolves to true if clean, false if infected/malware + */ +async function scanWithClamAV(buffer) { + return new Promise((resolve, reject) => { + const host = process.env.CLAMAV_HOST || 'localhost'; + const port = parseInt(process.env.CLAMAV_PORT || '3310', 10); + + const socket = new net.Socket(); + socket.setTimeout(5000); + + socket.connect(port, host, () => { + // Send zINSTREAM command (modern ClamAV null-terminated command prefix) + const prefix = Buffer.from('zINSTREAM\0'); + socket.write(prefix); + + // Send chunk size (4 bytes, big endian) followed by chunk data + const sizeBuf = Buffer.alloc(4); + sizeBuf.writeUInt32BE(buffer.length, 0); + socket.write(sizeBuf); + socket.write(buffer); + + // Send zero-length chunk to indicate end of stream + const endBuf = Buffer.alloc(4); + endBuf.writeUInt32BE(0, 0); + socket.write(endBuf); + }); + + let response = ''; + socket.on('data', (data) => { + response += data.toString(); + }); + + socket.on('end', () => { + if (response.includes('FOUND')) { + resolve(false); // Malware found + } else { + resolve(true); // Clean +function validateCSVContent(rows, headers) { + const errors = []; + + // Check if CSV is empty + if (!rows || rows.length === 0) { + errors.push('CSV file is empty'); + } + + // Check row count limit + if (rows.length > MAX_ROWS) { + errors.push(`CSV file exceeds maximum row limit of ${MAX_ROWS}`); + } + + // Check if required columns exist + const requiredColumns = ['text', 'message']; + const hasRequiredColumn = requiredColumns.some(col => + headers.some(h => h.toLowerCase().trim() === col.toLowerCase().trim()) + ); + + if (!hasRequiredColumn) { + errors.push(`CSV must contain a 'text' or 'message' column. Found: ${headers.join(', ')}`); + } + + // Check for empty rows + const emptyRows = rows.filter(row => { + const values = Object.values(row); + return values.every(val => !val || val.trim() === ''); + }); + + if (emptyRows.length > 0) { + errors.push(`Found ${emptyRows.length} empty rows in CSV`); + } + + // Check for very large cell values + const maxCellLength = 10000; // 10k characters + rows.forEach((row, index) => { + Object.values(row).forEach(value => { + if (value && value.length > maxCellLength) { + errors.push(`Row ${index + 1} contains a cell exceeding ${maxCellLength} characters`); + } + }); + + socket.on('error', (err) => { + reject(err); + }); + + socket.on('timeout', () => { + socket.destroy(); + reject(new Error('ClamAV scan timeout')); + }); + }); +} + +// Multer in-memory storage config with size limits +const storage = multer.memoryStorage(); +const upload = multer({ + storage: storage, + limits: { + fileSize: 20 * 1024 * 1024 // Set slightly higher to let middleware handle custom size limit logic and error codes + }, + fileFilter: (req, file, cb) => { + const fileExtension = file.originalname.split('.').pop().toLowerCase(); + if (fileExtension !== 'csv') { + return cb(new Error('Only CSV files are allowed'), false); + } + cb(null, true); + } +}).single('file'); + +/** + * Express middleware to validate and sanitize uploaded CSV files. + */ +const validateCSVUpload = (req, res, next) => { + upload(req, res, async (err) => { + if (err) { + if (err.code === 'LIMIT_FILE_SIZE') { + return res.status(413).json({ error: 'File too large' }); + } + return res.status(400).json({ error: err.message }); +async function scanForMalware(fileBuffer, filename) { + try { + // Check if ClamAV is available + const clamd = require('clamdjs'); + const scanner = clamd.createScanner('localhost', 3310); + + const result = await scanner.scanBuffer(fileBuffer); + + if (result.isInfected) { + throw new Error(`Malware detected: ${result.virusName}`); + } + + return { clean: true }; + } catch (error) { + // If ClamAV is not available, log warning but don't block + if (error.code === 'ECONNREFUSED') { + console.warn('⚠️ ClamAV not available - skipping malware scan'); + return { clean: true, warning: 'ClamAV not available' }; + } + + if (!req.file) { + return res.status(400).json({ error: 'No file uploaded' }); + } + + // Custom file size limit check from env + const maxFileSize = parseInt(process.env.MAX_CSV_FILE_SIZE, 10) || 5 * 1024 * 1024; + if (req.file.size > maxFileSize) { + return res.status(413).json({ error: 'File too large' }); + } +const validateCSVUpload = async (req, res, next) => { + try { + // Use multer to handle file upload + upload.single('file')(req, res, async (err) => { + if (err) { + if (err instanceof multer.MulterError) { + if (err.code === "LIMIT_FILE_SIZE") { + return res.status(413).json({ + success: false, + error: `File too large. Maximum size is ${MAX_FILE_SIZE / (1024 * 1024)}MB` + }); + } + if (err.code === 'LIMIT_FILE_COUNT') { + return res.status(400).json({ + success: false, + error: 'Only one file can be uploaded at a time' + }); + } + return res.status(400).json({ + success: false, + error: `Upload error: ${err.message}` + }); + } + return res.status(400).json({ + success: false, + error: err.message + }); + } + + const buffer = req.file.buffer; + + // Optional ClamAV Malware Scan + if (process.env.ENABLE_MALWARE_SCAN === 'true') { + try { + const isClean = await scanWithClamAV(buffer); + if (!isClean) { + return res.status(400).json({ error: 'File contains malware' }); + // 1. Validate file type (already done by multer) + const fileExt = path.extname(req.file.originalname).toLowerCase(); + if (!ALLOWED_EXTENSIONS.includes(fileExt) && + !ALLOWED_MIME_TYPES.includes(req.file.mimetype)) { + return res.status(400).json({ + success: false, + error: `Invalid file type. Only CSV files are allowed. Received: ${req.file.mimetype}` + }); + } + + // 2. Optional: Scan for malware + if (process.env.ENABLE_MALWARE_SCAN === 'true') { + try { + await scanForMalware(req.file.buffer, req.file.originalname); + } catch (scanError) { + return res.status(400).json({ + success: false, + error: `Security scan failed: ${scanError.message}` + }); + } + } + + // 3. Parse CSV content + const csvContent = req.file.buffer.toString('utf8'); + + // Basic CSV parsing (handle quoted fields) + const lines = csvContent.split('\n').filter(line => line.trim()); + if (lines.length === 0) { + return res.status(400).json({ + success: false, + error: 'CSV file is empty' + }); + } + + // Parse headers + const headers = parseCSVLine(lines[0]); + if (headers.length === 0) { + return res.status(400).json({ + success: false, + error: 'Invalid CSV headers' + }); + } + + // Parse rows and sanitize + const rows = []; + const validationErrors = []; + + for (let i = 1; i < lines.length; i++) { + const values = parseCSVLine(lines[i]); + const row = {}; + + headers.forEach((header, index) => { + const value = values[index] || ''; + // Sanitize each cell + row[header] = sanitizeCSVCell(value); + }); + + rows.push(row); + } + + // 4. Validate CSV structure + const validationResults = validateCSVContent(rows, headers); + if (validationResults.length > 0) { + return res.status(400).json({ + success: false, + errors: validationResults + }); + } + } catch (scanError) { + console.error('Malware scan failed:', scanError); + return res.status(500).json({ error: 'Malware scan service unavailable' }); + } + } + + const fileContent = buffer.toString('utf-8'); + if (!fileContent.trim()) { + return res.status(400).json({ error: 'CSV file is empty' }); + } + + const lines = fileContent.split(/\r?\n/).filter(line => line.trim()); + if (lines.length === 0) { + return res.status(400).json({ error: 'CSV file is empty' }); + } + + // Custom row count limit check from env + const maxRows = parseInt(process.env.MAX_CSV_ROWS, 10) || 100000; + if (lines.length - 1 > maxRows) { + return res.status(413).json({ error: 'File too large (too many rows)' }); + } + + // Parse headers and normalize/trim + const rawHeaders = parseCSVLine(lines[0]); + const normalizedHeaders = rawHeaders.map(h => h.trim().toLowerCase()); + + // Must contain "text" or "message" column + const hasText = normalizedHeaders.includes('text') || normalizedHeaders.includes('message'); + if (!hasText) { + return res.status(400).json({ + error: 'Invalid CSV format', + errors: ['CSV file must contain a "text" or "message" column'] + }); + } + + const rows = []; + for (let i = 1; i < lines.length; i++) { + const parsedLine = parseCSVLine(lines[i]); + const rowObj = {}; + rawHeaders.forEach((header, index) => { + const val = parsedLine[index] !== undefined ? parsedLine[index] : ''; + // Store sanitized values in the row object keyed by the parsed header + rowObj[header.trim()] = sanitizeCSVCell(val); + }); + rows.push(rowObj); + + // Yield to the event loop every 500 rows to prevent blocking (DoS) + if (i % 500 === 0) { + await new Promise(resolve => setImmediate(resolve)); +/** + * Parse CSV line handling quoted fields + */ +function parseCSVLine(line) { + const values = []; + let current = ''; + let inQuotes = false; + + for (let i = 0; i < line.length; i++) { + const char = line[i]; + + if (char === '"') { + if (inQuotes && line[i + 1] === '"') { + // Escaped quote + current += '"'; + i++; + } else { + inQuotes = !inQuotes; + } + } + } + + values.push(current.trim()); + return values; +} + + req.parsedCSV = { + headers: rawHeaders.map(h => h.trim()), + rows: rows, + totalRows: rows.length, + filename: req.file.originalname, + size: req.file.size + }; + + next(); + }); +}; + +module.exports = { + parseCSVLine, + sanitizeCSVCell, + validateCSVUpload +}; diff --git a/backend/middleware/filevalidation.js b/backend/middleware/filevalidation.js index bb126ad8..2849aac7 100644 --- a/backend/middleware/filevalidation.js +++ b/backend/middleware/filevalidation.js @@ -1,11 +1,15 @@ const multer = require('multer'); const net = require('net'); +<<<<<<< Updated upstream /** * Parses a single CSV line handling quoted fields and escaped quotes according to RFC 4180. * @param {string} line * @returns {string[]} */ +======= + +>>>>>>> Stashed changes function parseCSVLine(line) { const result = []; let current = ''; @@ -59,6 +63,64 @@ function sanitizeCSVCell(value) { // Neutralize formula injection: if starts with =, +, -, @, prefix with ' if (/^[=\+\-@]/.test(sanitized)) { sanitized = "'" + sanitized; +<<<<<<< Updated upstream +======= +function sanitizeCSVCell(cell) { + if (!cell || typeof cell !== 'string') { + return cell; + } + + const trimmed = cell.trim(); + + // Dangerous patterns that could indicate formula injection + const dangerousPatterns = [ + /^=.*/i, // Excel formulas + /^\+.*/i, // Excel formulas + /^@.*/i, // Excel formulas + /^-\s*.*/i, // Excel formulas + /^=cmd\|.*/i, // Command execution + /^=hyperlink\(.*\)/i, // Hyperlink injection + /^=dde\(.*\)/i, // DDE execution + /^=system\(.*\)/i, // System command + /^=shell\(.*\)/i, // Shell command + /^=execute\(.*\)/i, // Execute command + /^=run\(.*\)/i, // Run command + /\b(calc|cmd|powershell|bash|sh)\b/i // System commands + ]; + + // Check if cell starts with dangerous pattern + for (const pattern of dangerousPatterns) { + if (pattern.test(trimmed)) { + // Prefix with single quote to neutralize formula + return `'${trimmed}`; + } + } + + // Check for potential XSS in CSV (if rendered as HTML) + const xssPatterns = [ + /.*?<\/script>/i, + //i, + //i, + //i, + //i, + //i, + /on\w+\s*=/i, + /javascript:/i, + /vbscript:/i, + /data:/i + ]; + + for (const pattern of xssPatterns) { + if (pattern.test(cell)) { + // Escape HTML entities + return cell + .replace(/&/g, '&') + .replace(//g, '>') + .replace(/"/g, '"') + .replace(/'/g, '''); + } +>>>>>>> Stashed changes } return sanitized; @@ -104,6 +166,48 @@ async function scanWithClamAV(buffer) { resolve(false); // Malware found } else { resolve(true); // Clean +<<<<<<< Updated upstream +======= +function validateCSVContent(rows, headers) { + const errors = []; + + // Check if CSV is empty + if (!rows || rows.length === 0) { + errors.push('CSV file is empty'); + } + + // Check row count limit + if (rows.length > MAX_ROWS) { + errors.push(`CSV file exceeds maximum row limit of ${MAX_ROWS}`); + } + + // Check if required columns exist + const requiredColumns = ['text', 'message']; + const hasRequiredColumn = requiredColumns.some(col => + headers.some(h => h.toLowerCase().trim() === col.toLowerCase().trim()) + ); + + if (!hasRequiredColumn) { + errors.push(`CSV must contain a 'text' or 'message' column. Found: ${headers.join(', ')}`); + } + + // Check for empty rows + const emptyRows = rows.filter(row => { + const values = Object.values(row); + return values.every(val => !val || val.trim() === ''); + }); + + if (emptyRows.length > 0) { + errors.push(`Found ${emptyRows.length} empty rows in CSV`); + } + + // Check for very large cell values + const maxCellLength = 10000; // 10k characters + rows.forEach((row, index) => { + Object.values(row).forEach(value => { + if (value && value.length > maxCellLength) { + errors.push(`Row ${index + 1} contains a cell exceeding ${maxCellLength} characters`); +>>>>>>> Stashed changes } }); @@ -144,6 +248,27 @@ const validateCSVUpload = (req, res, next) => { return res.status(413).json({ error: 'File too large' }); } return res.status(400).json({ error: err.message }); +<<<<<<< Updated upstream +======= +async function scanForMalware(fileBuffer, filename) { + try { + // Check if ClamAV is available + const clamd = require('clamdjs'); + const scanner = clamd.createScanner('localhost', 3310); + + const result = await scanner.scanBuffer(fileBuffer); + + if (result.isInfected) { + throw new Error(`Malware detected: ${result.virusName}`); + } + + return { clean: true }; + } catch (error) { + // If ClamAV is not available, log warning but don't block + if (error.code === 'ECONNREFUSED') { + console.warn('⚠️ ClamAV not available - skipping malware scan'); + return { clean: true, warning: 'ClamAV not available' }; +>>>>>>> Stashed changes } if (!req.file) { @@ -155,6 +280,37 @@ const validateCSVUpload = (req, res, next) => { if (req.file.size > maxFileSize) { return res.status(413).json({ error: 'File too large' }); } +<<<<<<< Updated upstream +======= +const validateCSVUpload = async (req, res, next) => { + try { + // Use multer to handle file upload + upload.single('file')(req, res, async (err) => { + if (err) { + if (err instanceof multer.MulterError) { + if (err.code === "LIMIT_FILE_SIZE") { + return res.status(413).json({ + success: false, + error: `File too large. Maximum size is ${MAX_FILE_SIZE / (1024 * 1024)}MB` + }); + } + if (err.code === 'LIMIT_FILE_COUNT') { + return res.status(400).json({ + success: false, + error: 'Only one file can be uploaded at a time' + }); + } + return res.status(400).json({ + success: false, + error: `Upload error: ${err.message}` + }); + } + return res.status(400).json({ + success: false, + error: err.message + }); + } +>>>>>>> Stashed changes const buffer = req.file.buffer; @@ -164,6 +320,76 @@ const validateCSVUpload = (req, res, next) => { const isClean = await scanWithClamAV(buffer); if (!isClean) { return res.status(400).json({ error: 'File contains malware' }); +<<<<<<< Updated upstream +======= + // 1. Validate file type (already done by multer) + const fileExt = path.extname(req.file.originalname).toLowerCase(); + if (!ALLOWED_EXTENSIONS.includes(fileExt) && + !ALLOWED_MIME_TYPES.includes(req.file.mimetype)) { + return res.status(400).json({ + success: false, + error: `Invalid file type. Only CSV files are allowed. Received: ${req.file.mimetype}` + }); + } + + // 2. Optional: Scan for malware + if (process.env.ENABLE_MALWARE_SCAN === 'true') { + try { + await scanForMalware(req.file.buffer, req.file.originalname); + } catch (scanError) { + return res.status(400).json({ + success: false, + error: `Security scan failed: ${scanError.message}` + }); + } + } + + // 3. Parse CSV content + const csvContent = req.file.buffer.toString('utf8'); + + // Basic CSV parsing (handle quoted fields) + const lines = csvContent.split('\n').filter(line => line.trim()); + if (lines.length === 0) { + return res.status(400).json({ + success: false, + error: 'CSV file is empty' + }); + } + + // Parse headers + const headers = parseCSVLine(lines[0]); + if (headers.length === 0) { + return res.status(400).json({ + success: false, + error: 'Invalid CSV headers' + }); + } + + // Parse rows and sanitize + const rows = []; + const validationErrors = []; + + for (let i = 1; i < lines.length; i++) { + const values = parseCSVLine(lines[i]); + const row = {}; + + headers.forEach((header, index) => { + const value = values[index] || ''; + // Sanitize each cell + row[header] = sanitizeCSVCell(value); + }); + + rows.push(row); + } + + // 4. Validate CSV structure + const validationResults = validateCSVContent(rows, headers); + if (validationResults.length > 0) { + return res.status(400).json({ + success: false, + errors: validationResults + }); +>>>>>>> Stashed changes } } catch (scanError) { console.error('Malware scan failed:', scanError); @@ -214,8 +440,36 @@ const validateCSVUpload = (req, res, next) => { // Yield to the event loop every 500 rows to prevent blocking (DoS) if (i % 500 === 0) { await new Promise(resolve => setImmediate(resolve)); +<<<<<<< Updated upstream } } +======= +/** + * Parse CSV line handling quoted fields + */ +function parseCSVLine(line) { + const values = []; + let current = ''; + let inQuotes = false; + + for (let i = 0; i < line.length; i++) { + const char = line[i]; + + if (char === '"') { + if (inQuotes && line[i + 1] === '"') { + // Escaped quote + current += '"'; + i++; + } else { + inQuotes = !inQuotes; + } + } + } + + values.push(current.trim()); + return values; +} +>>>>>>> Stashed changes req.parsedCSV = { headers: rawHeaders.map(h => h.trim()), diff --git a/backend/middleware/rateLimiter.js b/backend/middleware/rateLimiter.js index ec2d3618..3419203b 100644 --- a/backend/middleware/rateLimiter.js +++ b/backend/middleware/rateLimiter.js @@ -2,6 +2,7 @@ const rateLimit = require('express-rate-limit'); const { ipKeyGenerator } = require('express-rate-limit'); const RedisStore = require('rate-limit-redis'); const redis = require('redis'); +<<<<<<< Updated upstream let redisClient = null; let store = undefined; @@ -49,6 +50,62 @@ if (process.env.REDIS_URL) { const loginLimiter = rateLimit({ windowMs: 15 * 60 * 1000, max: 5, +======= + +// ============================================ +// REDIS CONFIGURATION (Optional but recommended) +// ============================================ +let redisClient = null; +let store = undefined; + +// Try to connect to Redis if configured +if (process.env.REDIS_URL) { + try { + redisClient = redis.createClient({ + url: process.env.REDIS_URL, + socket: { + reconnectStrategy: (retries) => { + if (retries > 10) { + console.warn('⚠️ Redis connection failed, falling back to memory store'); + return false; + } + return Math.min(retries * 100, 3000); + } + } + }); + + redisClient.on('error', (err) => { + console.warn('⚠️ Redis error:', err.message); + store = undefined; + }); + + redisClient.connect().then(() => { + console.log('βœ… Redis connected for rate limiting'); + store = new RedisStore({ + sendCommand: (...args) => redisClient.sendCommand(args), + }); + }).catch(() => { + console.warn('⚠️ Redis connection failed, using memory store'); + store = undefined; + }); + } catch (error) { + console.warn('⚠️ Redis not available, using memory store'); + store = undefined; + } +} + +// ============================================ +// 1. STRICT AUTH LIMITERS (Your Existing Code) +// ============================================ + +/** + * Login Rate Limiter + * Limits login attempts to prevent brute force attacks + */ +const loginLimiter = rateLimit({ + windowMs: 15 * 60 * 1000, // 15 minutes + max: 5, // 5 attempts per window +>>>>>>> Stashed changes store: store, message: { success: false, @@ -56,15 +113,30 @@ const loginLimiter = rateLimit({ }, standardHeaders: true, legacyHeaders: false, +<<<<<<< Updated upstream skipSuccessfulRequests: true, keyGenerator: (req) => { +======= + skipSuccessfulRequests: true, // Don't count successful logins + keyGenerator: (req) => { + // Rate limit by email or IP +>>>>>>> Stashed changes return req.body.email || ipKeyGenerator(req.ip || req.connection.remoteAddress); }, }); +/** + * Registration Rate Limiter + * Prevents mass account creation + */ const registerLimiter = rateLimit({ +<<<<<<< Updated upstream windowMs: 15 * 60 * 1000, max: 5, +======= + windowMs: 15 * 60 * 1000, // 15 minutes + max: 5, // 5 registrations per window +>>>>>>> Stashed changes store: store, message: { success: false, @@ -74,9 +146,18 @@ const registerLimiter = rateLimit({ legacyHeaders: false, }); +/** + * Password Reset Rate Limiter + * Prevents abuse of password reset functionality + */ const resetLimiter = rateLimit({ +<<<<<<< Updated upstream windowMs: 15 * 60 * 1000, max: 3, +======= + windowMs: 15 * 60 * 1000, // 15 minutes + max: 3, // 3 reset requests per window +>>>>>>> Stashed changes store: store, message: { success: false, @@ -85,13 +166,26 @@ const resetLimiter = rateLimit({ standardHeaders: true, legacyHeaders: false, keyGenerator: (req) => { +<<<<<<< Updated upstream +======= + // Rate limit by email or IP +>>>>>>> Stashed changes return req.body.email || ipKeyGenerator(req.ip || req.connection.remoteAddress); }, }); +/** + * General API Rate Limiter + * Prevents API abuse and DoS attacks + */ const apiLimiter = rateLimit({ +<<<<<<< Updated upstream windowMs: 15 * 60 * 1000, max: 100, +======= + windowMs: 15 * 60 * 1000, // 15 minutes + max: 100, // 100 requests per window +>>>>>>> Stashed changes store: store, message: { success: false, @@ -101,9 +195,23 @@ const apiLimiter = rateLimit({ legacyHeaders: false, }); +<<<<<<< Updated upstream const chatLimiter = rateLimit({ windowMs: 60 * 1000, max: 15, +======= +// ============================================ +// 2. MAINTAINER'S LIMITERS (Incoming Code) +// ============================================ + +/** + * Chat Rate Limiter + * Prevents spam in chat endpoints + */ +const chatLimiter = rateLimit({ + windowMs: 60 * 1000, // 1 minute + max: 15, // 15 messages per minute +>>>>>>> Stashed changes store: store, message: { success: false, @@ -112,10 +220,21 @@ const chatLimiter = rateLimit({ standardHeaders: true, legacyHeaders: false, keyGenerator: (req) => { +<<<<<<< Updated upstream +======= + // Rate limit by user ID if authenticated, otherwise by IP +>>>>>>> Stashed changes return req.user?.id || ipKeyGenerator(req.ip || req.connection.remoteAddress); }, }); +<<<<<<< Updated upstream +======= +/** + * Predict Rate Limiter + * Prevents abuse of ML prediction endpoint + */ +>>>>>>> Stashed changes const PREDICT_WINDOW_MS = Number(process.env.RATE_LIMIT_WINDOW_MS) || Number(process.env.PREDICT_RATE_LIMIT_WINDOW_MS) || @@ -133,6 +252,10 @@ const predictLimiter = rateLimit({ standardHeaders: true, legacyHeaders: false, keyGenerator: (req) => { +<<<<<<< Updated upstream +======= + // Rate limit by user ID if authenticated, otherwise by IP +>>>>>>> Stashed changes return req.user?.id || ipKeyGenerator(req.ip || req.connection.remoteAddress); }, handler: (req, res, next, options) => { @@ -151,18 +274,41 @@ const predictLimiter = rateLimit({ } }); +<<<<<<< Updated upstream const otpLimiter = rateLimit({ windowMs: 5 * 60 * 1000, max: 3, +======= +// ============================================ +// 3. OTP & VERIFICATION LIMITERS (NEW) +// ============================================ + +/** + * OTP Request Rate Limiter + * Prevents SMS/Email bombing attacks + * Stricter limits for security + */ +const otpLimiter = rateLimit({ + windowMs: 5 * 60 * 1000, // 5 minutes + max: 3, // Only 3 OTP requests per 5 minutes +>>>>>>> Stashed changes store: store, message: { success: false, error: 'Too many OTP requests. Please wait 5 minutes.', +<<<<<<< Updated upstream retryAfter: 5 * 60 +======= + retryAfter: 5 * 60 // 5 minutes in seconds +>>>>>>> Stashed changes }, standardHeaders: true, legacyHeaders: false, keyGenerator: (req) => { +<<<<<<< Updated upstream +======= + // Rate limit by email or phone or IP +>>>>>>> Stashed changes return req.body.email || req.body.phone || ipKeyGenerator(req.ip || req.connection.remoteAddress); }, handler: (req, res, next, options) => { @@ -178,9 +324,19 @@ const otpLimiter = rateLimit({ } }); +<<<<<<< Updated upstream const verificationLimiter = rateLimit({ windowMs: 15 * 60 * 1000, max: 5, +======= +/** + * OTP Verification Rate Limiter + * Prevents brute force OTP attempts + */ +const verificationLimiter = rateLimit({ + windowMs: 15 * 60 * 1000, // 15 minutes + max: 5, // Max 5 verification attempts +>>>>>>> Stashed changes store: store, message: { success: false, @@ -189,9 +345,16 @@ const verificationLimiter = rateLimit({ standardHeaders: true, legacyHeaders: false, keyGenerator: (req) => { +<<<<<<< Updated upstream return req.body.email || req.body.phone || ipKeyGenerator(req.ip || req.connection.remoteAddress); }, skipSuccessfulRequests: true, +======= + // Rate limit by email, phone, or IP + return req.body.email || req.body.phone || ipKeyGenerator(req.ip || req.connection.remoteAddress); + }, + skipSuccessfulRequests: true, // Don't count successful verifications +>>>>>>> Stashed changes handler: (req, res, next, options) => { const retryAfterSeconds = Math.ceil(options.windowMs / 1000); res.status(429).json({ @@ -200,13 +363,27 @@ const verificationLimiter = rateLimit({ retryAfter: retryAfterSeconds, limit: options.max, remaining: 0 +<<<<<<< Updated upstream +======= + +>>>>>>> Stashed changes }); } }); +<<<<<<< Updated upstream const bulkPredictLimiter = rateLimit({ windowMs: 60 * 60 * 1000, max: 10, +======= +/** + * Bulk Predict Rate Limiter + * Prevents abuse of bulk prediction endpoint + */ +const bulkPredictLimiter = rateLimit({ + windowMs: 60 * 60 * 1000, // 1 hour + max: 10, // Max 10 bulk predictions per hour +>>>>>>> Stashed changes store: store, message: { success: false, @@ -219,9 +396,19 @@ const bulkPredictLimiter = rateLimit({ }, }); +<<<<<<< Updated upstream const exportLimiter = rateLimit({ windowMs: 60 * 60 * 1000, max: 5, +======= +/** + * Export Rate Limiter + * Prevents abuse of data export endpoints + */ +const exportLimiter = rateLimit({ + windowMs: 60 * 60 * 1000, // 1 hour + max: 5, // Max 5 exports per hour +>>>>>>> Stashed changes store: store, message: { success: false, @@ -234,9 +421,19 @@ const exportLimiter = rateLimit({ }, }); +<<<<<<< Updated upstream const feedbackLimiter = rateLimit({ windowMs: 60 * 1000, max: 10, +======= +/** + * Feedback Rate Limiter + * Prevents spam feedback submissions + */ +const feedbackLimiter = rateLimit({ + windowMs: 60 * 1000, // 1 minute + max: 10, // Max 10 feedback submissions per minute +>>>>>>> Stashed changes store: store, message: { success: false, @@ -249,11 +446,29 @@ const feedbackLimiter = rateLimit({ }, }); +<<<<<<< Updated upstream const isProduction = process.env.NODE_ENV === 'production'; const isDevelopment = process.env.NODE_ENV === 'development'; const getLimiterConfig = (baseConfig) => { if (isDevelopment) { +======= +// ============================================ +// 4. ENVIRONMENT-SPECIFIC CONFIGURATIONS +// ============================================ + +/** + * Development environment - less strict limits + * Production environment - stricter limits + */ +const isProduction = process.env.NODE_ENV === 'production'; +const isDevelopment = process.env.NODE_ENV === 'development'; + +// Dynamic limiters based on environment +const getLimiterConfig = (baseConfig) => { + if (isDevelopment) { + // Less strict in development +>>>>>>> Stashed changes return { ...baseConfig, max: baseConfig.max * 2, @@ -263,11 +478,356 @@ const getLimiterConfig = (baseConfig) => { return baseConfig; }; +<<<<<<< Updated upstream +module.exports = { +======= +// ============================================ +// 5. EXPORT ALL LIMITERS +// ============================================ + +module.exports = { + // Auth limiters +>>>>>>> Stashed changes + loginLimiter, + registerLimiter, + resetLimiter, + apiLimiter, +<<<<<<< Updated upstream +======= + + // Feature limiters + chatLimiter, + predictLimiter, + bulkPredictLimiter, + exportLimiter, + feedbackLimiter, + const rateLimit = require('express-rate-limit'); + const { ipKeyGenerator } = require('express-rate-limit'); + const RedisStore = require('rate-limit-redis'); + const redis = require('redis'); + + // ============================================ + // REDIS CONFIGURATION (Optional but recommended) + // ============================================ + let redisClient = null; + let store = undefined; + + // Try to connect to Redis if configured + if(process.env.REDIS_URL) { + try { + redisClient = redis.createClient({ + url: process.env.REDIS_URL, + socket: { + reconnectStrategy: (retries) => { + if (retries > 10) { + console.warn('⚠️ Redis connection failed, falling back to memory store'); + return false; + } + return Math.min(retries * 100, 3000); + } + } + }); + + redisClient.on('error', (err) => { + console.warn('⚠️ Redis error:', err.message); + store = undefined; + }); + + redisClient.connect().then(() => { + console.log('βœ… Redis connected for rate limiting'); + store = new RedisStore({ + sendCommand: (...args) => redisClient.sendCommand(args), + }); + }).catch(() => { + console.warn('⚠️ Redis connection failed, using memory store'); + store = undefined; + }); + } catch (error) { + console.warn('⚠️ Redis not available, using memory store'); + store = undefined; + } +} + +// ============================================ +// 1. STRICT AUTH LIMITERS (Your Existing Code) +// ============================================ + +/** + * Login Rate Limiter + * Limits login attempts to prevent brute force attacks + */ +const loginLimiter = rateLimit({ + windowMs: 15 * 60 * 1000, // 15 minutes + max: 5, // 5 attempts per window + store: store, + message: { + success: false, + message: 'Too many login attempts from this IP, please try again after 15 minutes.' + }, + standardHeaders: true, + legacyHeaders: false, + skipSuccessfulRequests: true, // Don't count successful logins + keyGenerator: (req) => { + // Rate limit by email or IP + return req.body.email || ipKeyGenerator(req.ip || req.connection.remoteAddress); + }, +}); + +/** + * Registration Rate Limiter + * Prevents mass account creation + */ +const registerLimiter = rateLimit({ + windowMs: 15 * 60 * 1000, // 15 minutes + max: 5, // 5 registrations per window + store: store, + message: { + success: false, + message: 'Too many registration attempts from this IP, please try again after 15 minutes.' + }, + standardHeaders: true, + legacyHeaders: false, +}); + +/** + * Password Reset Rate Limiter + * Prevents abuse of password reset functionality + */ +const resetLimiter = rateLimit({ + windowMs: 15 * 60 * 1000, // 15 minutes + max: 3, // 3 reset requests per window + store: store, + message: { + success: false, + message: 'Too many password reset requests, please try again later.' + }, + standardHeaders: true, + legacyHeaders: false, + keyGenerator: (req) => { + // Rate limit by email or IP + return req.body.email || ipKeyGenerator(req.ip || req.connection.remoteAddress); + }, +}); + +/** + * General API Rate Limiter + * Prevents API abuse and DoS attacks + */ +const apiLimiter = rateLimit({ + windowMs: 15 * 60 * 1000, // 15 minutes + max: 100, // 100 requests per window + store: store, + message: { + success: false, + message: 'Too many API requests from this IP, please try again after 15 minutes.' + }, + standardHeaders: true, + legacyHeaders: false, +}); + +// ============================================ +// 2. MAINTAINER'S LIMITERS (Incoming Code) +// ============================================ + +/** + * Chat Rate Limiter + * Prevents spam in chat endpoints + */ +const chatLimiter = rateLimit({ + windowMs: 60 * 1000, // 1 minute + max: 15, // 15 messages per minute + store: store, + message: { + success: false, + error: "Too many chat requests. Please slow down." + }, + standardHeaders: true, + legacyHeaders: false, + keyGenerator: (req) => { + // Rate limit by user ID if authenticated, otherwise by IP + return req.user?.id || ipKeyGenerator(req.ip || req.connection.remoteAddress); + }, +}); + +/** + * Predict Rate Limiter + * Prevents abuse of ML prediction endpoint + */ +const PREDICT_WINDOW_MS = + Number(process.env.RATE_LIMIT_WINDOW_MS) || + Number(process.env.PREDICT_RATE_LIMIT_WINDOW_MS) || + 15 * 60 * 1000; + +const PREDICT_MAX = + Number(process.env.RATE_LIMIT_MAX) || + Number(process.env.PREDICT_RATE_LIMIT_MAX) || + 100; + +const predictLimiter = rateLimit({ + windowMs: PREDICT_WINDOW_MS, + max: PREDICT_MAX, + store: store, + standardHeaders: true, + legacyHeaders: false, + keyGenerator: (req) => { + // Rate limit by user ID if authenticated, otherwise by IP + return req.user?.id || ipKeyGenerator(req.ip || req.connection.remoteAddress); + }, + handler: (req, res, next, options) => { + const retryAfterSeconds = Math.ceil(options.windowMs / 1000); + + res.setHeader("Retry-After", retryAfterSeconds); + + res.status(options.statusCode).json({ + success: false, + error: "Too many prediction requests. Please try again later.", + retryAfter: retryAfterSeconds, + limit: options.max, + remaining: 0, + resetTime: new Date(Date.now() + options.windowMs).toISOString() + }); + } +}); + +// ============================================ +// 3. OTP & VERIFICATION LIMITERS (NEW) +// ============================================ + +/** + * OTP Request Rate Limiter + * Prevents SMS/Email bombing attacks + * Stricter limits for security + */ +const otpLimiter = rateLimit({ + windowMs: 5 * 60 * 1000, // 5 minutes + max: 3, // Only 3 OTP requests per 5 minutes + store: store, + message: { + success: false, + error: 'Too many OTP requests. Please wait 5 minutes.', + retryAfter: 5 * 60 // 5 minutes in seconds + }, + standardHeaders: true, + legacyHeaders: false, + keyGenerator: (req) => { + // Rate limit by email or phone or IP + return req.body.email || req.body.phone || ipKeyGenerator(req.ip || req.connection.remoteAddress); + }, + handler: (req, res, next, options) => { + const retryAfterSeconds = Math.ceil(options.windowMs / 1000); + res.status(429).json({ + success: false, + error: 'Rate limit exceeded. Maximum 3 OTP requests per 5 minutes.', + retryAfter: retryAfterSeconds, + limit: options.max, + remaining: 0, + resetTime: new Date(Date.now() + options.windowMs).toISOString() + }); + } +}); + +/** + * OTP Verification Rate Limiter + * Prevents brute force OTP attempts + */ +const verificationLimiter = rateLimit({ + windowMs: 15 * 60 * 1000, // 15 minutes + max: 5, // Max 5 verification attempts + store: store, + message: { + success: false, + error: 'Too many verification attempts. Please try again later.' + }, + standardHeaders: true, + legacyHeaders: false, + keyGenerator: (req) => { + // Rate limit by email, phone, or IP + return req.body.email || req.body.phone || ipKeyGenerator(req.ip || req.connection.remoteAddress); + }, + skipSuccessfulRequests: true, // Don't count successful verifications + handler: (req, res, next, options) => { + const retryAfterSeconds = Math.ceil(options.windowMs / 1000); + res.status(429).json({ + success: false, + error: 'Too many verification attempts. Please try again in 15 minutes.', + retryAfter: retryAfterSeconds, + limit: options.max, + remaining: 0 + + }); + } +}); + +/** + * Bulk Predict Rate Limiter + * Prevents abuse of bulk prediction endpoint + */ +const bulkPredictLimiter = rateLimit({ + windowMs: 60 * 60 * 1000, // 1 hour + max: 10, // Max 10 bulk predictions per hour + store: store, + message: { + success: false, + error: 'Too many bulk prediction requests. Please wait 1 hour.' + }, + standardHeaders: true, + legacyHeaders: false, + keyGenerator: (req) => { + return req.user?.id || ipKeyGenerator(req.ip || req.connection.remoteAddress); + }, +}); + +/** + * Export Rate Limiter + * Prevents abuse of data export endpoints + */ +const exportLimiter = rateLimit({ + windowMs: 60 * 60 * 1000, // 1 hour + max: 5, // Max 5 exports per hour + store: store, + message: { + success: false, + error: 'Too many export requests. Please wait 1 hour.' + }, + standardHeaders: true, + legacyHeaders: false, + keyGenerator: (req) => { + return req.user?.id || ipKeyGenerator(req.ip || req.connection.remoteAddress); + }, +}); + +/** + * Feedback Rate Limiter + * Prevents spam feedback submissions + */ +const feedbackLimiter = rateLimit({ + windowMs: 60 * 1000, // 1 minute + max: 10, // Max 10 feedback submissions per minute + store: store, + message: { + success: false, + error: 'Too many feedback submissions. Please slow down.' + }, + standardHeaders: true, + legacyHeaders: false, + keyGenerator: (req) => { + return req.user?.id || ipKeyGenerator(req.ip || req.connection.remoteAddress); + }, +}); + +// ============================================ +// 5. EXPORT ALL LIMITERS +// ============================================ + module.exports = { + // Auth limiters loginLimiter, registerLimiter, resetLimiter, apiLimiter, + + // Feature limiters +>>>>>>> Stashed changes chatLimiter, predictLimiter, bulkPredictLimiter, @@ -277,7 +837,18 @@ module.exports = { // OTP limiters otpLimiter, verificationLimiter, +<<<<<<< Updated upstream PREDICT_MAX, PREDICT_WINDOW_MS, redisClient }; +======= + + // Configuration values + PREDICT_MAX, + PREDICT_WINDOW_MS, + + // Redis client (for external use) + redisClient +}; +>>>>>>> Stashed changes diff --git a/backend/middleware/rateLimiter.js.bak b/backend/middleware/rateLimiter.js.bak new file mode 100644 index 00000000..db2134cc --- /dev/null +++ b/backend/middleware/rateLimiter.js.bak @@ -0,0 +1,361 @@ +const rateLimit = require('express-rate-limit'); +const { ipKeyGenerator } = require('express-rate-limit'); +const RedisStore = require('rate-limit-redis'); +const redis = require('redis'); + +// ============================================ +// REDIS CONFIGURATION (Optional but recommended) +// ============================================ +let redisClient = null; +let store = undefined; + +// Try to connect to Redis if configured +if (process.env.REDIS_URL) { + try { + redisClient = redis.createClient({ + url: process.env.REDIS_URL, + socket: { + reconnectStrategy: (retries) => { + if (retries > 10) { + console.warn('⚠️ Redis connection failed, falling back to memory store'); + return false; + } + return Math.min(retries * 100, 3000); + } + } + }); + + redisClient.on('error', (err) => { + console.warn('⚠️ Redis error:', err.message); + store = undefined; + }); + + redisClient.connect().then(() => { + console.log('βœ… Redis connected for rate limiting'); + store = new RedisStore({ + sendCommand: (...args) => redisClient.sendCommand(args), + }); + }).catch(() => { + console.warn('⚠️ Redis connection failed, using memory store'); + store = undefined; + }); + } catch (error) { + console.warn('⚠️ Redis not available, using memory store'); + store = undefined; + } +} + +// ============================================ +// 1. STRICT AUTH LIMITERS (Your Existing Code) +// ============================================ + +/** + * Login Rate Limiter + * Limits login attempts to prevent brute force attacks + */ +const loginLimiter = rateLimit({ + windowMs: 15 * 60 * 1000, // 15 minutes + max: 5, // 5 attempts per window + store: store, + message: { + success: false, + message: 'Too many login attempts from this IP, please try again after 15 minutes.' + }, + standardHeaders: true, + legacyHeaders: false, + skipSuccessfulRequests: true, // Don't count successful logins + keyGenerator: (req) => { + // Rate limit by email or IP + return req.body.email || req.ip || req.connection.remoteAddress; + }, +}); + +/** + * Registration Rate Limiter + * Prevents mass account creation + */ +const registerLimiter = rateLimit({ + windowMs: 15 * 60 * 1000, // 15 minutes + max: 5, // 5 registrations per window + store: store, + message: { + success: false, + message: 'Too many registration attempts from this IP, please try again after 15 minutes.' + }, + standardHeaders: true, + legacyHeaders: false, +}); + +/** + * Password Reset Rate Limiter + * Prevents abuse of password reset functionality + */ +const resetLimiter = rateLimit({ + windowMs: 15 * 60 * 1000, // 15 minutes + max: 3, // 3 reset requests per window + store: store, + message: { + success: false, + message: 'Too many password reset requests, please try again later.' + }, + standardHeaders: true, + legacyHeaders: false, + keyGenerator: (req) => { + // Rate limit by email or IP + return req.body.email || req.ip || req.connection.remoteAddress; + }, +}); + +/** + * General API Rate Limiter + * Prevents API abuse and DoS attacks + */ +const apiLimiter = rateLimit({ + windowMs: 15 * 60 * 1000, // 15 minutes + max: 100, // 100 requests per window + store: store, + message: { + success: false, + message: 'Too many API requests from this IP, please try again after 15 minutes.' + }, + standardHeaders: true, + legacyHeaders: false, +}); + +// ============================================ +// 2. MAINTAINER'S LIMITERS (Incoming Code) +// ============================================ + +/** + * Chat Rate Limiter + * Prevents spam in chat endpoints + */ +const chatLimiter = rateLimit({ + windowMs: 60 * 1000, // 1 minute + max: 15, // 15 messages per minute + store: store, + message: { + success: false, + error: "Too many chat requests. Please slow down." + }, + standardHeaders: true, + legacyHeaders: false, + keyGenerator: (req) => { + // Rate limit by user ID if authenticated, otherwise by IP + return req.user?.id || req.ip || req.connection.remoteAddress; + }, +}); + +/** + * Predict Rate Limiter + * Prevents abuse of ML prediction endpoint + */ +const PREDICT_WINDOW_MS = Number(process.env.PREDICT_RATE_LIMIT_WINDOW_MS) || 60 * 1000; +const PREDICT_MAX = Number(process.env.PREDICT_RATE_LIMIT_MAX) || 30; + +const predictLimiter = rateLimit({ + windowMs: PREDICT_WINDOW_MS, + max: PREDICT_MAX, + store: store, + standardHeaders: true, + legacyHeaders: false, + keyGenerator: (req) => { + // Rate limit by user ID if authenticated, otherwise by IP + return req.user?.id || req.ip || req.connection.remoteAddress; + }, + handler: (req, res, next, options) => { + const retryAfterSeconds = Math.ceil(options.windowMs / 1000); + res.status(429).json({ + success: false, + error: "Too many predict requests. Please slow down.", + retryAfter: retryAfterSeconds, + limit: options.max, + remaining: 0, + resetTime: new Date(Date.now() + options.windowMs).toISOString() + }); + } +}); + +// ============================================ +// 3. OTP & VERIFICATION LIMITERS (NEW) +// ============================================ + +/** + * OTP Request Rate Limiter + * Prevents SMS/Email bombing attacks + * Stricter limits for security + */ +const otpLimiter = rateLimit({ + windowMs: 5 * 60 * 1000, // 5 minutes + max: 3, // Only 3 OTP requests per 5 minutes + store: store, + message: { + success: false, + error: 'Too many OTP requests. Please wait 5 minutes.', + retryAfter: 5 * 60 // 5 minutes in seconds + }, + standardHeaders: true, + legacyHeaders: false, + keyGenerator: (req) => { + // Rate limit by email or phone or IP + return req.body.email || req.body.phone || req.ip || req.connection.remoteAddress; + }, + handler: (req, res, next, options) => { + const retryAfterSeconds = Math.ceil(options.windowMs / 1000); + res.status(429).json({ + success: false, + error: 'Rate limit exceeded. Maximum 3 OTP requests per 5 minutes.', + retryAfter: retryAfterSeconds, + limit: options.max, + remaining: 0, + resetTime: new Date(Date.now() + options.windowMs).toISOString() + }); + } +}); + +/** + * OTP Verification Rate Limiter + * Prevents brute force OTP attempts + */ +const verificationLimiter = rateLimit({ + windowMs: 15 * 60 * 1000, // 15 minutes + max: 5, // Max 5 verification attempts + store: store, + message: { + success: false, + error: 'Too many verification attempts. Please try again later.' + }, + standardHeaders: true, + legacyHeaders: false, + keyGenerator: (req) => { + // Rate limit by email, phone, or IP + return req.body.email || req.body.phone || req.ip || req.connection.remoteAddress; + }, + skipSuccessfulRequests: true, // Don't count successful verifications + handler: (req, res, next, options) => { + const retryAfterSeconds = Math.ceil(options.windowMs / 1000); + res.status(429).json({ + success: false, + error: 'Too many verification attempts. Please try again in 15 minutes.', + retryAfter: retryAfterSeconds, + limit: options.max, + remaining: 0 + }); + } +}); + +/** + * Bulk Predict Rate Limiter + * Prevents abuse of bulk prediction endpoint + */ +const bulkPredictLimiter = rateLimit({ + windowMs: 60 * 60 * 1000, // 1 hour + max: 10, // Max 10 bulk predictions per hour + store: store, + message: { + success: false, + error: 'Too many bulk prediction requests. Please wait 1 hour.' + }, + standardHeaders: true, + legacyHeaders: false, + keyGenerator: (req) => { + return req.user?.id || req.ip || req.connection.remoteAddress; + }, +}); + +/** + * Export Rate Limiter + * Prevents abuse of data export endpoints + */ +const exportLimiter = rateLimit({ + windowMs: 60 * 60 * 1000, // 1 hour + max: 5, // Max 5 exports per hour + store: store, + message: { + success: false, + error: 'Too many export requests. Please wait 1 hour.' + }, + standardHeaders: true, + legacyHeaders: false, + keyGenerator: (req) => { + return req.user?.id || req.ip || req.connection.remoteAddress; + }, +}); + +/** + * Feedback Rate Limiter + * Prevents spam feedback submissions + */ +const feedbackLimiter = rateLimit({ + windowMs: 60 * 1000, // 1 minute + max: 10, // Max 10 feedback submissions per minute + store: store, + message: { + success: false, + error: 'Too many feedback submissions. Please slow down.' + }, + standardHeaders: true, + legacyHeaders: false, + keyGenerator: (req) => { + return req.user?.id || req.ip || req.connection.remoteAddress; + }, +}); + +// ============================================ +// 4. ENVIRONMENT-SPECIFIC CONFIGURATIONS +// ============================================ + +/** + * Development environment - less strict limits + * Production environment - stricter limits + */ +const isProduction = process.env.NODE_ENV === 'production'; +const isDevelopment = process.env.NODE_ENV === 'development'; + +// Dynamic limiters based on environment +const getLimiterConfig = (baseConfig) => { + if (isDevelopment) { + // Less strict in development + return { + ...baseConfig, + max: baseConfig.max * 2, + windowMs: baseConfig.windowMs / 2, + }; + } + return baseConfig; +}; + +// ============================================ +// 5. EXPORT ALL LIMITERS +// ============================================ + +module.exports = { + // Auth limiters + loginLimiter, + registerLimiter, + resetLimiter, + apiLimiter, + + // Feature limiters + chatLimiter, + predictLimiter, + bulkPredictLimiter, + exportLimiter, + feedbackLimiter, + + // OTP limiters + otpLimiter, + verificationLimiter, + + // Configuration values + PREDICT_MAX, + PREDICT_WINDOW_MS, + + // Redis client (for external use) + redisClient, + + // Utility functions + isProduction, + isDevelopment, + getLimiterConfig +}; \ No newline at end of file diff --git a/backend/models/History.js b/backend/models/History.js index 3edad0c6..28a6e9a5 100644 --- a/backend/models/History.js +++ b/backend/models/History.js @@ -23,20 +23,29 @@ const historySchema = new mongoose.Schema( prediction: { type: String, required: [true, "Prediction is required."], +<<<<<<< Updated upstream enum: { values: ["ham", "spam"], message: "Prediction must be either 'ham' or 'spam'" }, default: "ham" +======= +>>>>>>> Stashed changes }, type: { type: String, required: [true, "Type is required."], enum: { +<<<<<<< Updated upstream values: PREDICTION_TYPE_LIST, message: `Type must be one of: ${PREDICTION_TYPE_LIST.join(', ')}`, }, default: PREDICTION_TYPES.MESSAGE +======= + values: ["sms", "email", "url", "message"], + message: "Type must be one of: sms, email, url, or message.", + }, +>>>>>>> Stashed changes }, confidence: { type: Number, @@ -81,6 +90,11 @@ const historySchema = new mongoose.Schema( { timestamps: true } ); +<<<<<<< Updated upstream +======= +// 1. Time-Series Index: Optimizes getTrends (filtering by user & sorting by date) +// Drastically improves performance for dashboard time-series charts. +>>>>>>> Stashed changes historySchema.index({ user: 1, createdAt: -1 }, { background: true }); historySchema.index({ user: 1, prediction: 1 }, { background: true }); historySchema.index({ user: 1, type: 1 }, { background: true }); diff --git a/backend/models/User.js b/backend/models/User.js index ab4beaee..fe3f17fe 100644 --- a/backend/models/User.js +++ b/backend/models/User.js @@ -12,19 +12,36 @@ const ROLES = { }; const PERMISSIONS = { +<<<<<<< Updated upstream +======= + // User permissions +>>>>>>> Stashed changes PREDICT: 'predict', BULK_PREDICT: 'bulk_predict', VIEW_ANALYTICS: 'view_analytics', MANAGE_WEBHOOKS: 'manage_webhooks', EXPORT_DATA: 'export_data', +<<<<<<< Updated upstream MANAGE_USERS: 'manage_users', VIEW_REPORTS: 'view_reports', +======= + + // Moderator permissions + MANAGE_USERS: 'manage_users', + VIEW_REPORTS: 'view_reports', + + // Admin permissions +>>>>>>> Stashed changes MANAGE_ROLES: 'manage_roles', VIEW_LOGS: 'view_logs', SYSTEM_CONFIG: 'system_config', MANAGE_ALL: 'manage_all' }; +<<<<<<< Updated upstream +======= +// Role to permissions mapping +>>>>>>> Stashed changes const ROLE_PERMISSIONS = { [ROLES.USER]: [ PERMISSIONS.PREDICT, @@ -67,9 +84,20 @@ const userSchema = new mongoose.Schema( type: String, required: [true, 'Username is required'], trim: true, +<<<<<<< Updated upstream minlength: [3, 'Username must be at least 3 characters long'], maxlength: [30, 'Username cannot exceed 30 characters'], match: [/^[a-zA-Z0-9_]+$/, 'Username can only contain letters, numbers, and underscores'] +======= + + minlength: 3, + maxlength: 30, + match: [/^[a-zA-Z0-9_]+$/, 'Username can only contain letters, numbers, and underscores'] + + minlength: [3, 'Username must be at least 3 characters long'], + maxlength: [30, 'Username cannot exceed 30 characters'], + +>>>>>>> Stashed changes }, email: { type: String, @@ -80,9 +108,24 @@ const userSchema = new mongoose.Schema( }, password: { type: String, +<<<<<<< Updated upstream required: false, minlength: [6, 'Password must be at least 6 characters long'], select: false +======= + + required: function() { + return this.provider === 'local'; + }, + minlength: 6, + select: false // Don't return password by default + + required: false, + + minlength: [6, 'Password must be at least 6 characters long'], + + +>>>>>>> Stashed changes }, googleId: { type: String, @@ -95,6 +138,7 @@ const userSchema = new mongoose.Schema( }, provider: { type: String, +<<<<<<< Updated upstream enum: { values: ['local', 'google'], message: '{VALUE} is not a valid provider' @@ -111,13 +155,54 @@ const userSchema = new mongoose.Schema( enum: Object.values(PERMISSIONS), default: ROLE_PERMISSIONS[ROLES.USER] }, +======= + + enum: ['local', 'google'], + default: 'local' + }, + // ============================================ + // ROLE & PERMISSIONS (Zero Trust) + // ============================================ + role: { + type: String, + enum: Object.values(ROLES), + default: ROLES.USER + }, + permissions: { + type: [String], + enum: Object.values(PERMISSIONS), + default: ROLE_PERMISSIONS[ROLES.USER] + + enum: { + values: ['local', 'google'], + message: '{VALUE} is not a valid provider' + }, + default: 'local', + + }, + // ============================================ + // WEBHOOK URL (Existing) + // ============================================ +>>>>>>> Stashed changes webhookUrl: { type: String, trim: true, default: null, + + match: [/^https?:\/\/.+/, 'Please enter a valid HTTP or HTTPS URL'] + match: [/^https?:\/\/.+/, 'Please enter a valid HTTP or HTTPS URL'], +<<<<<<< Updated upstream maxlength: [2000, 'Webhook URL cannot exceed 2000 characters'] }, +======= + maxlength: [2000, 'Webhook URL cannot exceed 2000 characters'], + + }, + // ============================================ + // ACCOUNT STATUS (Optional) + // ============================================ +>>>>>>> Stashed changes status: { type: String, enum: ['active', 'inactive', 'suspended'], @@ -142,6 +227,7 @@ const userSchema = new mongoose.Schema( ); // ============================================ +<<<<<<< Updated upstream // CASE-INSENSITIVE UNIQUE INDEXES // ============================================ @@ -172,11 +258,24 @@ userSchema.index( userSchema.index({ role: 1 }); userSchema.index({ status: 1 }); userSchema.index({ googleId: 1 }); +======= +// INDEXES +// ============================================ + +userSchema.index({ email: 1 }); +userSchema.index({ username: 1 }); +userSchema.index({ role: 1 }); +userSchema.index({ status: 1 }); +>>>>>>> Stashed changes // ============================================ // PRE-SAVE HOOKS // ============================================ +<<<<<<< Updated upstream +======= +// Hash password before saving +>>>>>>> Stashed changes userSchema.pre('save', async function (next) { if (!this.password || !this.isModified('password')) { return next(); @@ -185,6 +284,10 @@ userSchema.pre('save', async function (next) { next(); }); +<<<<<<< Updated upstream +======= +// Set default permissions based on role +>>>>>>> Stashed changes userSchema.pre('save', function (next) { if (this.isModified('role') || this.isNew) { this.permissions = ROLE_PERMISSIONS[this.role] || ROLE_PERMISSIONS[ROLES.USER]; @@ -196,21 +299,37 @@ userSchema.pre('save', function (next) { // INSTANCE METHODS // ============================================ +<<<<<<< Updated upstream +======= +// Compare password +>>>>>>> Stashed changes userSchema.methods.comparePassword = async function (candidatePassword) { if (!this.password) return false; return bcrypt.compare(candidatePassword, this.password); }; +<<<<<<< Updated upstream +======= +// Check if user has specific permission +>>>>>>> Stashed changes userSchema.methods.hasPermission = function (permission) { if (this.role === ROLES.ADMIN) return true; return this.permissions.includes(permission); }; +<<<<<<< Updated upstream +======= +// Check if user has all required permissions +>>>>>>> Stashed changes userSchema.methods.hasAllPermissions = function (requiredPermissions) { if (this.role === ROLES.ADMIN) return true; return requiredPermissions.every(p => this.permissions.includes(p)); }; +<<<<<<< Updated upstream +======= +// Update last login +>>>>>>> Stashed changes userSchema.methods.updateLastLogin = function () { this.lastLogin = new Date(); this.loginAttempts = 0; @@ -218,14 +337,26 @@ userSchema.methods.updateLastLogin = function () { return this.save(); }; +<<<<<<< Updated upstream userSchema.methods.incrementLoginAttempts = function () { this.loginAttempts += 1; if (this.loginAttempts >= 5) { this.lockUntil = new Date(Date.now() + 15 * 60 * 1000); +======= +// Increment login attempts +userSchema.methods.incrementLoginAttempts = function () { + this.loginAttempts += 1; + if (this.loginAttempts >= 5) { + this.lockUntil = new Date(Date.now() + 15 * 60 * 1000); // Lock for 15 minutes +>>>>>>> Stashed changes } return this.save(); }; +<<<<<<< Updated upstream +======= +// Check if account is locked +>>>>>>> Stashed changes userSchema.methods.isLocked = function () { if (!this.lockUntil) return false; return this.lockUntil > new Date(); @@ -235,14 +366,26 @@ userSchema.methods.isLocked = function () { // STATIC METHODS // ============================================ +<<<<<<< Updated upstream +======= +// Get permissions for a role +>>>>>>> Stashed changes userSchema.statics.getPermissionsForRole = function (role) { return ROLE_PERMISSIONS[role] || ROLE_PERMISSIONS[ROLES.USER]; }; +<<<<<<< Updated upstream +======= +// Get all available roles +>>>>>>> Stashed changes userSchema.statics.getRoles = function () { return Object.values(ROLES); }; +<<<<<<< Updated upstream +======= +// Get all available permissions +>>>>>>> Stashed changes userSchema.statics.getPermissions = function () { return Object.values(PERMISSIONS); }; @@ -251,10 +394,18 @@ userSchema.statics.getPermissions = function () { // VIRTUAL PROPERTIES // ============================================ +<<<<<<< Updated upstream +======= +// Check if user is admin +>>>>>>> Stashed changes userSchema.virtual('isAdmin').get(function () { return this.role === ROLES.ADMIN; }); +<<<<<<< Updated upstream +======= +// Check if user is moderator +>>>>>>> Stashed changes userSchema.virtual('isModerator').get(function () { return this.role === ROLES.MODERATOR || this.role === ROLES.ADMIN; }); @@ -263,6 +414,10 @@ userSchema.virtual('isModerator').get(function () { // EXPORTS // ============================================ +<<<<<<< Updated upstream +======= +// Export constants for use in other files +>>>>>>> Stashed changes userSchema.statics.ROLES = ROLES; userSchema.statics.PERMISSIONS = PERMISSIONS; userSchema.statics.ROLE_PERMISSIONS = ROLE_PERMISSIONS; diff --git a/backend/multimodal_detector.py b/backend/multimodal_detector.py new file mode 100644 index 00000000..bee69554 --- /dev/null +++ b/backend/multimodal_detector.py @@ -0,0 +1,511 @@ +#!/usr/bin/env python3 +""" +Multimodal Spam Detection Using NEAT (NeuroEvolution of Augmenting Topologies) +Detects spam across text, images, and voice modalities +""" + +import os +import json +import numpy as np +import pickle +from pathlib import Path +from datetime import datetime +import hashlib +import base64 +import io +import warnings +warnings.filterwarnings('ignore') + +# ============================================ +# CONFIGURATION +# ============================================ + +BASE_DIR = Path(__file__).resolve().parent +MODEL_DIR = BASE_DIR / 'models' +MODEL_DIR.mkdir(exist_ok=True) + +# ============================================ +# TEXT MODALITY DETECTOR +# ============================================ + +class TextModalityDetector: + """Detects spam in text messages""" + + def __init__(self): + self.vectorizer = None + self.classifier = None + self.is_trained = False + + def extract_features(self, text): + """Extract features from text""" + if not text: + return np.array([]) + + # Basic statistical features + features = [] + features.append(len(text)) # Length + features.append(len(text.split())) # Word count + features.append(sum(c.isupper() for c in text) / max(len(text), 1)) # Uppercase ratio + features.append(sum(c in '!?.,;:' for c in text) / max(len(text), 1)) # Punctuation ratio + features.append(len(re.findall(r'[^a-zA-Z0-9\s]', text)) / max(len(text), 1)) # Special chars + features.append(sum(c.isdigit() for c in text) / max(len(text), 1)) # Digit ratio + + # Spam keyword detection + spam_keywords = ['free', 'win', 'prize', 'claim', 'urgent', 'click', 'money', 'cash', 'bonus', 'guaranteed'] + features.append(sum(1 for kw in spam_keywords if kw in text.lower())) + + return np.array(features) + + def detect(self, text): + """Detect spam in text""" + if not text: + return {'prediction': 'ham', 'confidence': 0.5, 'modality': 'text'} + + features = self.extract_features(text).reshape(1, -1) + + # Simple heuristic-based detection (placeholder for actual model) + score = 0 + if len(text) > 100: score += 0.1 + if len(text.split()) > 20: score += 0.1 + if any(kw in text.lower() for kw in ['free', 'win', 'prize', 'claim']): score += 0.3 + if '!' in text: score += 0.1 + if any(c.isupper() for c in text): score += 0.1 + + score = min(score, 0.95) + prediction = 'spam' if score > 0.4 else 'ham' + + return { + 'prediction': prediction, + 'confidence': score, + 'modality': 'text', + 'text_preview': text[:100] + } + + +# ============================================ +# IMAGE MODALITY DETECTOR +# ============================================ + +class ImageModalityDetector: + """Detects spam in images (screenshots, embedded images)""" + + def __init__(self): + self.is_trained = False + + def extract_features(self, image_data): + """Extract features from image data""" + try: + from PIL import Image + import io + + # Decode image + if isinstance(image_data, str): + # Base64 encoded image + image_bytes = base64.b64decode(image_data) + image = Image.open(io.BytesIO(image_bytes)) + elif isinstance(image_data, bytes): + image = Image.open(io.BytesIO(image_data)) + else: + return np.zeros(10) + + # Convert to numpy + img = np.array(image) + + # Extract features + features = [] + + # Basic stats + if len(img.shape) == 3: + features.append(np.mean(img[:,:,0])) + features.append(np.mean(img[:,:,1])) + features.append(np.mean(img[:,:,2])) + features.append(np.std(img[:,:,0])) + features.append(np.std(img[:,:,1])) + features.append(np.std(img[:,:,2])) + else: + features.append(np.mean(img)) + features.append(np.std(img)) + features.extend([0, 0, 0, 0, 0]) + + # Size + features.append(img.shape[0] / 1000) # Height + features.append(img.shape[1] / 1000) # Width + + # Entropy (simplified) + if len(img.shape) == 3: + hist = np.histogram(img[:,:,0].flatten(), bins=50)[0] + else: + hist = np.histogram(img.flatten(), bins=50)[0] + hist = hist / (hist.sum() + 1) + entropy = -np.sum(hist * np.log2(hist + 1e-10)) + features.append(entropy) + + return np.array(features) + + except Exception as e: + print(f"⚠️ Image feature extraction failed: {e}") + return np.zeros(10) + + def detect(self, image_data): + """Detect spam in image""" + if not image_data: + return {'prediction': 'ham', 'confidence': 0.5, 'modality': 'image'} + + features = self.extract_features(image_data) + + # Simple heuristic-based detection + score = 0.2 + + # Check for suspicious patterns (placeholder) + # In production, use actual model + + prediction = 'spam' if score > 0.5 else 'ham' + + return { + 'prediction': prediction, + 'confidence': score, + 'modality': 'image', + 'features': features.tolist()[:5] + } + + +# ============================================ +# VOICE MODALITY DETECTOR +# ============================================ + +class VoiceModalityDetector: + """Detects spam in voice messages""" + + def __init__(self): + self.is_trained = False + + def extract_features(self, audio_data): + """Extract features from audio data""" + try: + # For demo, use simple features + # In production, use librosa or similar + features = [] + + # Length estimation + if isinstance(audio_data, str): + # Base64 encoded audio + audio_bytes = base64.b64decode(audio_data) + features.append(len(audio_bytes) / 1000) # Size in KB + elif isinstance(audio_data, bytes): + features.append(len(audio_data) / 1000) + else: + features.append(0) + + # Random features (placeholder) + features.extend([np.random.random() for _ in range(5)]) + + return np.array(features) + + except Exception as e: + print(f"⚠️ Voice feature extraction failed: {e}") + return np.zeros(6) + + def detect(self, audio_data): + """Detect spam in voice""" + if not audio_data: + return {'prediction': 'ham', 'confidence': 0.5, 'modality': 'voice'} + + features = self.extract_features(audio_data) + + # Simple heuristic-based detection + score = 0.15 + if features[0] > 10: # Large audio file + score += 0.2 + + prediction = 'spam' if score > 0.4 else 'ham' + + return { + 'prediction': prediction, + 'confidence': score, + 'modality': 'voice', + 'features': features.tolist()[:5] + } + + +# ============================================ +# NEAT - Dynamic Weighting & Evolution +# ============================================ + +class NEAT: + """NeuroEvolution of Augmenting Topologies for dynamic weighting""" + + def __init__(self, dynamic_weighting=True): + self.dynamic_weighting = dynamic_weighting + self.weights = { + 'text': 0.4, + 'image': 0.3, + 'voice': 0.3 + } + self.performance_history = [] + self.generation = 0 + self.population = [] + + def set_weights(self, text_weight, image_weight, voice_weight): + """Manually set weights""" + total = text_weight + image_weight + voice_weight + self.weights = { + 'text': text_weight / total, + 'image': image_weight / total, + 'voice': voice_weight / total + } + + def optimize_weights(self, predictions, ground_truth): + """Optimize weights based on performance""" + if not self.dynamic_weighting: + return self.weights + + # Simple weight optimization based on accuracy + # In production, use actual NEAT algorithm + + # Calculate per-modality accuracy + accuracies = {} + for modality in ['text', 'image', 'voice']: + if modality in predictions: + correct = sum(1 for p, g in zip(predictions[modality], ground_truth) if p == g) + total = len(ground_truth) + accuracies[modality] = correct / max(total, 1) + else: + accuracies[modality] = 0.5 + + # Update weights based on accuracy (higher accuracy = higher weight) + total_acc = sum(accuracies.values()) or 1 + self.weights = { + 'text': accuracies['text'] / total_acc, + 'image': accuracies['image'] / total_acc, + 'voice': accuracies['voice'] / total_acc + } + + self.generation += 1 + self.performance_history.append({ + 'generation': self.generation, + 'weights': self.weights.copy(), + 'accuracies': accuracies.copy() + }) + + return self.weights + + def get_weights(self): + """Get current weights""" + return self.weights + + def get_stats(self): + """Get NEAT statistics""" + return { + 'generation': self.generation, + 'weights': self.weights, + 'performance_history': self.performance_history[-10:] if self.performance_history else [] + } + + +# ============================================ +# MULTIMODAL SPAM DETECTOR +# ============================================ + +class MultimodalSpamDetector: + """Main multimodal detector using NEAT""" + + def __init__(self): + self.neat = NEAT(dynamic_weighting=True) + self.text_detector = TextModalityDetector() + self.image_detector = ImageModalityDetector() + self.voice_detector = VoiceModalityDetector() + self.detection_history = [] + + def detect_text(self, text): + """Detect spam in text""" + return self.text_detector.detect(text) + + def detect_image(self, image_data): + """Detect spam in image""" + return self.image_detector.detect(image_data) + + def detect_voice(self, audio_data): + """Detect spam in voice""" + return self.voice_detector.detect(audio_data) + + def ensemble_detect(self, modalities): + """ + Ensemble detection across multiple modalities + modalities: dict with keys 'text', 'image', 'voice' + """ + results = {} + scores = [] + modalities_used = [] + + # Detect each modality + if 'text' in modalities and modalities['text']: + result = self.detect_text(modalities['text']) + results['text'] = result + scores.append(result['confidence']) + modalities_used.append('text') + + if 'image' in modalities and modalities['image']: + result = self.detect_image(modalities['image']) + results['image'] = result + scores.append(result['confidence']) + modalities_used.append('image') + + if 'voice' in modalities and modalities['voice']: + result = self.detect_voice(modalities['voice']) + results['voice'] = result + scores.append(result['confidence']) + modalities_used.append('voice') + + if not scores: + return { + 'ensemble_prediction': 'ham', + 'ensemble_confidence': 0.5, + 'results': results, + 'modalities_used': [] + } + + # Get weights from NEAT + weights = self.neat.get_weights() + + # Calculate weighted average confidence + weighted_scores = [] + for modality, result in results.items(): + weight = weights.get(modality, 0.33) + weighted_scores.append(result['confidence'] * weight) + + ensemble_confidence = sum(weighted_scores) / sum(weights.values()) + + # Determine prediction + prediction = 'spam' if ensemble_confidence > 0.5 else 'ham' + + # Log detection + self.detection_history.append({ + 'timestamp': datetime.now().isoformat(), + 'modalities_used': modalities_used, + 'ensemble_prediction': prediction, + 'ensemble_confidence': ensemble_confidence, + 'weights': weights + }) + + return { + 'ensemble_prediction': prediction, + 'ensemble_confidence': float(ensemble_confidence), + 'results': results, + 'modalities_used': modalities_used, + 'weights': weights + } + + def optimize_weights(self, predictions, ground_truth): + """Optimize NEAT weights based on performance""" + self.neat.optimize_weights(predictions, ground_truth) + return self.neat.get_weights() + + def get_stats(self): + """Get detector statistics""" + return { + 'neat': self.neat.get_stats(), + 'detection_history': self.detection_history[-20:] if self.detection_history else [], + 'total_detections': len(self.detection_history) + } + + def save(self, path=None): + """Save detector state""" + if not path: + path = MODEL_DIR / 'multimodal_detector.pkl' + + with open(path, 'wb') as f: + pickle.dump({ + 'neat': self.neat, + 'detection_history': self.detection_history[-100:], + 'timestamp': datetime.now().isoformat() + }, f) + print(f"πŸ’Ύ Saved multimodal detector to {path}") + + def load(self, path=None): + """Load detector state""" + if not path: + path = MODEL_DIR / 'multimodal_detector.pkl' + + if os.path.exists(path): + with open(path, 'rb') as f: + data = pickle.load(f) + self.neat = data['neat'] + self.detection_history = data.get('detection_history', []) + print(f"βœ… Loaded multimodal detector from {path}") + return True + return False + + +# ============================================ +# HELPER FUNCTIONS +# ============================================ + +def encode_image_to_base64(image_path): + """Encode image file to base64""" + with open(image_path, 'rb') as f: + return base64.b64encode(f.read()).decode('utf-8') + +def encode_audio_to_base64(audio_path): + """Encode audio file to base64""" + with open(audio_path, 'rb') as f: + return base64.b64encode(f.read()).decode('utf-8') + + +# ============================================ +# MAIN - Test & Demo +# ============================================ + +def main(): + print("=" * 60) + print("🎯 Multimodal Spam Detection using NEAT") + print("=" * 60) + + # Create detector + detector = MultimodalSpamDetector() + + # Test text detection + print("\nπŸ“ Testing Text Detection:") + test_texts = [ + "Hi team, meeting at 10am tomorrow", + "Cl4im y0ur fr33 pr!ze n0w!", + "You have received a complimentary reward!", + ] + + for text in test_texts: + result = detector.detect_text(text) + print(f" '{text[:30]}...' β†’ {result['prediction']} ({result['confidence']:.2%})") + + # Test ensemble detection + print("\n🎯 Testing Ensemble Detection:") + + test_cases = [ + {'text': 'Claim your free prize now!'}, + {'text': 'Meeting at 10am'}, + {'text': 'URGENT! You have won!', 'image': 'dummy_image_data'}, + {'text': 'Free money', 'image': 'dummy', 'voice': 'dummy_audio'} + ] + + for i, modalities in enumerate(test_cases, 1): + print(f"\n Case {i}: {', '.join(modalities.keys())}") + result = detector.ensemble_detect(modalities) + print(f" Ensemble Prediction: {result['ensemble_prediction']}") + print(f" Confidence: {result['ensemble_confidence']:.2%}") + print(f" Modalities Used: {result['modalities_used']}") + print(f" Weights: {result['weights']}") + + # Show NEAT stats + print("\nπŸ“Š NEAT Statistics:") + stats = detector.get_stats() + neat_stats = stats['neat'] + print(f" Generation: {neat_stats['generation']}") + print(f" Current Weights: {neat_stats['weights']}") + + print("\n" + "=" * 60) + print("βœ… Multimodal Spam Detection System Ready!") + print(f" Models saved to: {MODEL_DIR}") + + return detector + + +if __name__ == "__main__": + detector = main() \ No newline at end of file diff --git a/backend/package-lock.json b/backend/package-lock.json index 3d8508a9..e704e3a7 100644 --- a/backend/package-lock.json +++ b/backend/package-lock.json @@ -2430,6 +2430,7 @@ "@redis/client": "^6.1.0" } }, +<<<<<<< Updated upstream "node_modules/@scarf/scarf": { "version": "1.4.0", "resolved": "https://registry.npmjs.org/@scarf/scarf/-/scarf-1.4.0.tgz", @@ -2437,6 +2438,8 @@ "hasInstallScript": true, "license": "Apache-2.0" }, +======= +>>>>>>> Stashed changes "node_modules/@sentry/conventions": { "version": "0.12.0", "resolved": "https://registry.npmjs.org/@sentry/conventions/-/conventions-0.12.0.tgz", @@ -7358,6 +7361,7 @@ "node": ">=6" } }, +<<<<<<< Updated upstream "node_modules/levn": { "version": "0.4.1", "resolved": "https://registry.npmjs.org/levn/-/levn-0.4.1.tgz", @@ -7372,6 +7376,8 @@ "node": ">= 0.8.0" } }, +======= +>>>>>>> Stashed changes "node_modules/lie": { "version": "3.3.0", "resolved": "https://registry.npmjs.org/lie/-/lie-3.3.0.tgz", @@ -8289,6 +8295,7 @@ "integrity": "sha512-pkEqbDyl8ou5cpq+VsnQbe/WlEy5qS7xPzMS1U55OCG9KPvwFD46zDbxQIj3egJSFc3D+XhYOPUzz49zQAVy7A==", "license": "BSD-2-Clause" }, +<<<<<<< Updated upstream "node_modules/optionator": { "version": "0.9.4", "resolved": "https://registry.npmjs.org/optionator/-/optionator-0.9.4.tgz", @@ -8307,6 +8314,8 @@ "node": ">= 0.8.0" } }, +======= +>>>>>>> Stashed changes "node_modules/options": { "version": "0.0.6", "resolved": "https://registry.npmjs.org/options/-/options-0.0.6.tgz", diff --git a/backend/package.json b/backend/package.json index 8c85213a..1cf1445e 100644 --- a/backend/package.json +++ b/backend/package.json @@ -6,9 +6,14 @@ "scripts": { "start": "node server.js", "dev": "nodemon server.js", + "worker": "node worker.js", + "test": "jest --testPathIgnorePatterns config --testPathIgnorePatterns avatarUpload --testPathIgnorePatterns keywordRules --testPathIgnorePatterns rateLimiter --testPathIgnorePatterns fileValidation && node --test tests/keywordRules.test.js tests/rateLimiter.test.js tests/avatarUpload.test.js tests/fileValidation.test.js", "lint": "eslint ." + + "test": "jest --testPathIgnorePatterns config --testPathIgnorePatterns avatarUpload --testPathIgnorePatterns keywordRules --testPathIgnorePatterns rateLimiter --testPathIgnorePatterns fileValidation && node --test tests/keywordRules.test.js tests/rateLimiter.test.js tests/avatarUpload.test.js tests/fileValidation.test.js" + }, "keywords": [], "author": "", diff --git a/backend/requirements.txt b/backend/requirements.txt index aa48ace5..ce683464 100644 --- a/backend/requirements.txt +++ b/backend/requirements.txt @@ -30,6 +30,7 @@ nltk # Email authenticity verification module uses standard libraries (email, re) # IMAP inbox scanning (issue #186) uses the standard library imaplib/email modules +<<<<<<< HEAD beautifulsoup4>=4.12.0 pytesseract>=0.3.10 opencv-python>=4.8.0 @@ -53,5 +54,9 @@ html2image>=2.0.0 numpy>=1.23.0 filelock +numpy>=1.23.0 +Pillow>=9.0.0 +scikit-learn>=1.2.0 +opencv-python>=4.8.0 diff --git a/backend/routes/analyticsRoutes.js b/backend/routes/analyticsRoutes.js index a8f52555..8f3a57a6 100644 --- a/backend/routes/analyticsRoutes.js +++ b/backend/routes/analyticsRoutes.js @@ -12,7 +12,11 @@ const { } = require("../controllers/analyticsController"); const { protect } = require("../middleware/authMiddleware"); +<<<<<<< Updated upstream const Prediction = require('../models/Prediction'); +======= + +>>>>>>> Stashed changes router.use(protect); router.get("/summary", getSummary); router.get("/trends", getTrends); diff --git a/backend/routes/authRoutes.js b/backend/routes/authRoutes.js index aaf1de6b..f0d5b65c 100644 --- a/backend/routes/authRoutes.js +++ b/backend/routes/authRoutes.js @@ -75,6 +75,7 @@ const buildAuthResponse = (user, token) => ({ // ============================================ +<<<<<<< Updated upstream // TOKEN GENERATION // ============================================ @@ -99,6 +100,8 @@ const buildAuthResponse = (user, token) => ({ // ============================================ +======= +>>>>>>> Stashed changes // AUTH CONTROLLERS // ============================================ @@ -553,6 +556,7 @@ const getSessionStatus = async (req, res) => { }; // ============================================ +<<<<<<< Updated upstream // AUTH CONTROLLERS // ============================================ @@ -1031,6 +1035,8 @@ const assignRole = async (req, res) => { }); } +======= +>>>>>>> Stashed changes // ZERO TRUST - ROLE MANAGEMENT // ============================================ @@ -1376,7 +1382,10 @@ const assignRole = async (req, res) => { }); } +<<<<<<< Updated upstream +======= +>>>>>>> Stashed changes const user = await User.findById(userId); if (!user) { return res.status(404).json({ @@ -1457,9 +1466,12 @@ const getUserPermissions = async (req, res) => { } }; +<<<<<<< Updated upstream } }; +======= +>>>>>>> Stashed changes // @desc Get all roles and permissions (Public) // @route GET /api/auth/roles @@ -1484,7 +1496,10 @@ const getRolesAndPermissions = async (req, res) => { +<<<<<<< Updated upstream +======= +>>>>>>> Stashed changes // @desc Get all roles and permissions (Public) // @route GET /api/auth/roles const getRolesAndPermissions = async (req, res) => { diff --git a/backend/routes/chatRoutes.js b/backend/routes/chatRoutes.js index 8a2e05dc..ae99afa6 100644 --- a/backend/routes/chatRoutes.js +++ b/backend/routes/chatRoutes.js @@ -2,6 +2,7 @@ const express = require("express"); const router = express.Router(); const { chatLimiter } = require("../middleware/rateLimiter"); +<<<<<<< Updated upstream const { chatHandler, healthCheck, listModels } = require("../controllers/chatController"); // Chat endpoint with fallback support and rate limiting @@ -13,4 +14,12 @@ router.get("/health", healthCheck); // List available models router.get("/models", listModels); -module.exports = router; \ No newline at end of file +module.exports = router; +======= +const { chatHandler } = require("../controllers/chatController"); + + +router.post("/", chatLimiter, chatHandler); + +module.exports = router; +>>>>>>> Stashed changes diff --git a/backend/routes/historyRoutes.js b/backend/routes/historyRoutes.js index e5a4d048..88fb0098 100644 --- a/backend/routes/historyRoutes.js +++ b/backend/routes/historyRoutes.js @@ -30,6 +30,7 @@ router.delete("/:id", deleteHistoryItem); router.delete("/", clearHistory); router.get('/count', getHistoryCount); + module.exports = router; router.get('/recent',protect, async(req,res)=> { @@ -45,6 +46,9 @@ router.get('/recent',protect, async(req,res)=> { } }); +module.exports = router; + + router.get('/',protect,async(req,res) => { try{ const{startDate, endDate, limit =50 } =req.query; @@ -65,4 +69,5 @@ router.get('/',protect,async(req,res) => { } catch (error) { res.status(500).json({ error: 'Failed to fetch history' }); } -}); \ No newline at end of file +}); + diff --git a/backend/routes/multimodalRoutes.js b/backend/routes/multimodalRoutes.js new file mode 100644 index 00000000..6f5b06f6 --- /dev/null +++ b/backend/routes/multimodalRoutes.js @@ -0,0 +1,166 @@ +const express = require('express'); +const router = express.Router(); +const { protect } = require('../middleware/authMiddleware'); +const { checkPermission } = require('../middleware/zeroTrust'); +const { spawn } = require('child_process'); +const path = require('path'); + +const MULTIMODAL_SCRIPT = path.join(__dirname, '../multimodal_detector.py'); + +/** + * @route POST /api/multimodal/detect + * @desc Detect spam using multimodal ensemble + * @access Private + */ +router.post('/detect', protect, async (req, res) => { + try { + const { text, image, voice } = req.body; + + if (!text && !image && !voice) { + return res.status(400).json({ + success: false, + error: 'At least one modality (text, image, or voice) is required' + }); + } + + const result = await runMultimodal('detect', { text, image, voice }); + res.json({ success: true, ...result }); + } catch (error) { + res.status(500).json({ success: false, error: error.message }); + } +}); + +/** + * @route POST /api/multimodal/detect/text + * @desc Detect spam in text only + * @access Private + */ +router.post('/detect/text', protect, async (req, res) => { + try { + const { text } = req.body; + if (!text) { + return res.status(400).json({ success: false, error: 'Text is required' }); + } + + const result = await runMultimodal('detect_text', { text }); + res.json({ success: true, ...result }); + } catch (error) { + res.status(500).json({ success: false, error: error.message }); + } +}); + +/** + * @route POST /api/multimodal/detect/image + * @desc Detect spam in image + * @access Private + */ +router.post('/detect/image', protect, async (req, res) => { + try { + const { image } = req.body; + if (!image) { + return res.status(400).json({ success: false, error: 'Image data is required' }); + } + + const result = await runMultimodal('detect_image', { image }); + res.json({ success: true, ...result }); + } catch (error) { + res.status(500).json({ success: false, error: error.message }); + } +}); + +/** + * @route POST /api/multimodal/detect/voice + * @desc Detect spam in voice + * @access Private + */ +router.post('/detect/voice', protect, async (req, res) => { + try { + const { voice } = req.body; + if (!voice) { + return res.status(400).json({ success: false, error: 'Voice data is required' }); + } + + const result = await runMultimodal('detect_voice', { voice }); + res.json({ success: true, ...result }); + } catch (error) { + res.status(500).json({ success: false, error: error.message }); + } +}); + +/** + * @route POST /api/multimodal/optimize + * @desc Optimize NEAT weights (Admin only) + * @access Private (Admin) + */ +router.post('/optimize', protect, checkPermission('system_config'), async (req, res) => { + try { + const { predictions, groundTruth } = req.body; + + if (!predictions || !groundTruth) { + return res.status(400).json({ + success: false, + error: 'Predictions and ground truth required' + }); + } + + const result = await runMultimodal('optimize', { predictions, groundTruth }); + res.json({ + success: true, + message: 'Weights optimized successfully', + ...result + }); + } catch (error) { + res.status(500).json({ success: false, error: error.message }); + } +}); + +/** + * @route GET /api/multimodal/stats + * @desc Get multimodal detector statistics + * @access Private (Admin) + */ +router.get('/stats', protect, checkPermission('view_logs'), async (req, res) => { + try { + const stats = await runMultimodal('stats', {}); + res.json({ success: true, stats }); + } catch (error) { + res.status(500).json({ success: false, error: error.message }); + } +}); + +function runMultimodal(command, params = {}) { + return new Promise((resolve, reject) => { + const python = spawn('python', [ + MULTIMODAL_SCRIPT, + '--command', command, + '--params', JSON.stringify(params) + ]); + + let output = ''; + let errorOutput = ''; + + python.stdout.on('data', (data) => { + output += data.toString(); + }); + + python.stderr.on('data', (data) => { + errorOutput += data.toString(); + }); + + python.on('close', (code) => { + if (code !== 0) { + reject(new Error(errorOutput || `Process exited with code ${code}`)); + } else { + try { + resolve(JSON.parse(output)); + } catch (e) { + resolve({ output, raw: true }); + } + } + }); + + python.on('error', (err) => reject(err)); + }); +} + +module.exports = router; \ No newline at end of file diff --git a/backend/routes/prediction.route.js b/backend/routes/prediction.route.js new file mode 100644 index 00000000..b79679ab --- /dev/null +++ b/backend/routes/prediction.route.js @@ -0,0 +1,41 @@ +const express = require('express'); +const router = express.Router(); +const auth = require('../middleware/auth'); + +// GET /api/predictions/stats +router.get('/stats', auth, async (req, res) => { + try { + const userId = req.user.id; + const today = new Date(); + today.setHours(0, 0, 0, 0); + + // Get predictions from database + const predictions = await Prediction.find({ userId }); + + // Calculate stats + const total = predictions.length; + const todayCount = predictions.filter(p => + new Date(p.createdAt) >= today + ).length; + + // Spam vs Ham breakdown + const spamCount = predictions.filter(p => + p.result === 'spam' || p.result === 'smishing' + ).length; + const hamCount = predictions.filter(p => + p.result === 'ham' || p.result === 'safe' + ).length; + + res.json({ + today: todayCount, + total: total, + spamCount: spamCount, + hamCount: hamCount + }); + } catch (error) { + console.error('Stats error:', error); + res.status(500).json({ error: 'Failed to fetch stats' }); + } +}); + +module.exports = router; \ No newline at end of file diff --git a/backend/routes/predictionRoutes.js b/backend/routes/predictionRoutes.js index 6f24e925..7dab6a62 100644 --- a/backend/routes/predictionRoutes.js +++ b/backend/routes/predictionRoutes.js @@ -238,6 +238,10 @@ router.post("/predict", predictLimiter, preventCacheStampede, protect, checkCach // Check ML Cache globally before calling Flask const cacheKey = `spam_cache:${require('crypto').createHash('sha256').update(text).digest('hex')}`; +<<<<<<< Updated upstream +======= + const { redisClient } = require("./middleware/cacheMiddleware"); +>>>>>>> Stashed changes if (redisClient && redisClient.status === 'ready') { try { const cachedResult = await redisClient.get(cacheKey); @@ -415,6 +419,7 @@ router.post("/feedback", protect, async (req, res) => { } }); +<<<<<<< Updated upstream router.get("/feedback/stats", protect, async (req, res) => { try { const response = await axios.get(`${ML_API_BASE}/feedback/stats`); @@ -434,6 +439,8 @@ router.get("/feedback/stats", protect, async (req, res) => { } }); +======= +>>>>>>> Stashed changes router.post("/analyze-email-header", protect, upload.single("file"), async (req, res) => { try { if (req.file) { diff --git a/backend/server.js b/backend/server.js index 65ead8c9..526d3e32 100644 --- a/backend/server.js +++ b/backend/server.js @@ -22,13 +22,36 @@ const compression = require('compression'); const { v4: uuidv4 } = require('uuid'); const helmet = require('helmet'); const axios = require("axios"); + const { corsOptions } = require('./config/corsConfig'); +// Add Multimodal Detection routes +const multimodalRoutes = require('./routes/multimodalRoutes'); +app.use('/api/multimodal', multimodalRoutes); // Initialize background jobs require('./jobs/archivalCron'); require('./jobs/webhookRetryCron'); const { preventCacheStampede } = require('./middleware/cacheMiddleware'); + const adversarialRoutes = require('./routes/adversarialRoutes'); +app.use('/api/adversarial', adversarialRoutes); + +// Add EvoMail routes +const evoMailRoutes = require('./routes/evoMailRoutes'); +app.use('/api/evomail', evoMailRoutes); + +// ===== STARTUP TIMER ===== +const SERVER_START_TIME = Date.now(); +const startupLogs = []; +// Add VBSF routes +const visualRoutes = require('./routes/visualRoutes'); +app.use('/api/visual', visualRoutes); +const logStartupTime= (component, startTime) => { + + + +// Add EvoMail routes + const evoMailRoutes = require('./routes/evoMailRoutes'); const poisoningRoutes = require('./routes/poisoningRoutes'); @@ -45,6 +68,7 @@ const startupLogs = []; const { configureAxios } = require('./config/axios'); configureAxios(); // Apply the global axios configuration const logStartupTime = (component, startTime) => { + const elapsed = Date.now() - startTime; startupLogs.push({ component, elapsed }); logger.info(`⏱️ ${component} loaded in ${elapsed}ms`); @@ -195,12 +219,14 @@ app.use(express.json({ limit: '1mb' })); app.use(express.urlencoded({ extended: true, limit: '1mb' })); app.use('/uploads', express.static('uploads')); + app.use('/api-docs', swaggerUi.serve, swaggerUi.setup(swaggerSpec, { explorer: true, customCss: '.swagger-ui .topbar { display: none }', customSiteTitle: 'Spam Detection API Docs' })); + // ===== REQUEST ID MIDDLEWARE ===== app.use((req, res, next) => { // Generate a unique request ID @@ -294,11 +320,13 @@ app.get("/", (req, res) => { res.send("Node backend running "); }); + app.get('/api-docs.json', (req, res) => { res.setHeader('Content-Type', 'application/json'); res.send(swaggerSpec); }); + // ======================================== // START SERVER // ======================================== diff --git a/backend/tests/fileValidation.test.js b/backend/tests/fileValidation.test.js new file mode 100644 index 00000000..dda0ac2e --- /dev/null +++ b/backend/tests/fileValidation.test.js @@ -0,0 +1,279 @@ +const test = require('node:test'); +const assert = require('node:assert'); +const express = require('express'); +const multer = require('multer'); +const { validateCSVUpload, sanitizeCSVCell, parseCSVLine } = require('../middleware/fileValidation'); + +// ============================================ +// TEST HELPERS +// ============================================ + +function createTestApp() { + const app = express(); + app.post('/test-upload', validateCSVUpload, (req, res) => { + res.json({ + success: true, + data: req.parsedCSV + }); + }); + return app; +} + +function createFormData(content, filename = 'test.csv') { + const boundary = '----testboundary'; + const parts = [ + `--${boundary}`, + `Content-Disposition: form-data; name="file"; filename="${filename}"`, + 'Content-Type: text/csv', + '', + content, + `--${boundary}--` + ]; + return { + buffer: Buffer.from(parts.join('\r\n')), + headers: { + 'Content-Type': `multipart/form-data; boundary=${boundary}`, + } + }; +} + +// ============================================ +// TESTS +// ============================================ + +test('sanitizeCSVCell: should neutralize formula injection', () => { + const testCases = [ + { input: '=cmd|"/C calc"', expected: "'=cmd|\"/C calc\"" }, + { input: '=HYPERLINK("http://evil.com")', expected: "'=HYPERLINK(\"http://evil.com\")" }, + { input: '=DDE("cmd";"/C calc")', expected: "'=DDE(\"cmd\";\"/C calc\")" }, + { input: '=system("calc")', expected: "'=system(\"calc\")" }, + { input: '=shell("calc")', expected: "'=shell(\"calc\")" }, + { input: '=execute("calc")', expected: "'=execute(\"calc\")" }, + { input: '+cmd|"/C calc"', expected: "'+cmd|\"/C calc\"" }, + { input: '@cmd|"/C calc"', expected: "'@cmd|\"/C calc\"" }, + { input: '-cmd|"/C calc"', expected: "'-cmd|\"/C calc\"" }, + { input: 'Normal text', expected: 'Normal text' }, + { input: '12345', expected: '12345' }, + { input: 'Hello, World!', expected: 'Hello, World!' } + ]; + + testCases.forEach(({ input, expected }) => { + const result = sanitizeCSVCell(input); + assert.strictEqual(result, expected, `Failed for input: ${input}`); + }); +}); + +test('sanitizeCSVCell: should escape XSS vectors', () => { + const testCases = [ + { input: '', expected: '<script>alert("xss")</script>' }, + { input: '