diff --git a/src/models/emotion_detection/hf_loader.py b/src/models/emotion_detection/hf_loader.py index 46a29f08f..b68468731 100644 --- a/src/models/emotion_detection/hf_loader.py +++ b/src/models/emotion_detection/hf_loader.py @@ -1,17 +1,17 @@ #!/usr/bin/env python3 from __future__ import annotations -from dataclasses import dataclass -from typing import Dict, Optional import os -import tempfile -import tarfile import shutil +import tarfile +import tempfile +from dataclasses import dataclass +from typing import Dict, Optional -import torch import requests -from transformers import AutoConfig, AutoTokenizer, AutoModelForSequenceClassification +import torch from huggingface_hub import snapshot_download +from transformers import AutoConfig, AutoModelForSequenceClassification, AutoTokenizer @dataclass @@ -71,7 +71,9 @@ def predict(self, text: str, threshold: float = 0.5) -> Dict: headers["Authorization"] = f"Bearer {self.token}" payload = {"inputs": text} try: - resp = requests.post(self.endpoint_url, json=payload, headers=headers, timeout=30) + resp = requests.post( + self.endpoint_url, json=payload, headers=headers, timeout=30 + ) resp.raise_for_status() data = resp.json() # data can be [[{"label":..., "score":...}, ...]] or {"error":...} @@ -79,12 +81,21 @@ def predict(self, text: str, threshold: float = 0.5) -> Dict: items = data[0] if data and isinstance(data[0], list) else data else: items = [] - emotions = {item.get("label", f"L{i}"): float(item.get("score", 0.0)) for i, item in enumerate(items)} + emotions = { + item.get("label", f"L{i}"): float(item.get("score", 0.0)) + for i, item in enumerate(items) + } if emotions: primary_label, confidence = max(emotions.items(), key=lambda kv: kv[1]) else: primary_label, confidence = "neutral", 1.0 - intensity = "high" if confidence >= 0.75 else ("moderate" if confidence >= 0.4 else "low") + + if confidence >= 0.75: + intensity = "high" + elif confidence >= 0.4: + intensity = "moderate" + else: + intensity = "low" return { "emotions": emotions, "primary_emotion": primary_label, @@ -95,20 +106,30 @@ def predict(self, text: str, threshold: float = 0.5) -> Dict: return {} -def _wrap_local_model(local_dir: str, token: Optional[str] = None, force_multi_label: Optional[bool] = None) -> HFEmotionDetector: +def _wrap_local_model( + local_dir: str, + token: Optional[str] = None, + force_multi_label: Optional[bool] = None, +) -> HFEmotionDetector: cfg = AutoConfig.from_pretrained(local_dir, token=token) tok = AutoTokenizer.from_pretrained(local_dir, token=token, use_fast=True) mdl = AutoModelForSequenceClassification.from_pretrained(local_dir, token=token) - id2label = getattr(cfg, "id2label", None) or {i: str(i) for i in range(cfg.num_labels)} + id2label = getattr(cfg, "id2label", None) or { + i: str(i) for i in range(cfg.num_labels) + } if force_multi_label is not None: multi_label = bool(force_multi_label) else: problem_type = getattr(cfg, "problem_type", None) multi_label = problem_type == "multi_label_classification" - return HFEmotionDetector(model=mdl, tokenizer=tok, id2label=id2label, multi_label=multi_label) + return HFEmotionDetector( + model=mdl, tokenizer=tok, id2label=id2label, multi_label=multi_label + ) -def load_hf_emotion_model(model_id: str, token: Optional[str] = None, force_multi_label: Optional[bool] = None) -> HFEmotionDetector: +def load_hf_emotion_model( + model_id: str, token: Optional[str] = None, force_multi_label: Optional[bool] = None +) -> HFEmotionDetector: return _wrap_local_model(model_id, token=token, force_multi_label=force_multi_label) @@ -120,7 +141,9 @@ def load_emotion_model_multi_source( endpoint_url: Optional[str] = None, force_multi_label: Optional[bool] = None, ) -> object: - """Try multiple sources to load the emotion model. Returns an object with .predict(text, threshold). + """Try multiple sources to load the emotion model. + + Returns an object with .predict(text, threshold). Priority: 1) Explicit local_dir if provided and exists @@ -132,14 +155,18 @@ def load_emotion_model_multi_source( # 1) Local directory if local_dir and os.path.isdir(local_dir): try: - return _wrap_local_model(local_dir, token=token, force_multi_label=force_multi_label) + return _wrap_local_model( + local_dir, token=token, force_multi_label=force_multi_label + ) except Exception: pass # 2) HF Hub direct if model_id: try: - return load_hf_emotion_model(model_id, token=token, force_multi_label=force_multi_label) + return load_hf_emotion_model( + model_id, token=token, force_multi_label=force_multi_label + ) except Exception: pass @@ -147,17 +174,23 @@ def load_emotion_model_multi_source( if model_id: try: cache_base = os.getenv("HF_HOME", "/var/tmp/hf-cache") - snap_dir = snapshot_download(repo_id=model_id, token=token, cache_dir=cache_base) - return _wrap_local_model(snap_dir, token=token, force_multi_label=force_multi_label) + snap_dir = snapshot_download( + repo_id=model_id, token=token, cache_dir=cache_base + ) + return _wrap_local_model( + snap_dir, token=token, force_multi_label=force_multi_label + ) except Exception: pass # 4) Archive URL if archive_url: try: - cache_dir = os.path.join(os.getenv("XDG_CACHE_HOME", "/var/tmp/hf-cache"), "model-archives") + cache_base = os.getenv("XDG_CACHE_HOME", "/var/tmp/hf-cache") + cache_dir = os.path.join(cache_base, "model-archives") os.makedirs(cache_dir, exist_ok=True) - archive_path = os.path.join(cache_dir, os.path.basename(archive_url.split("?")[0])) + archive_name = os.path.basename(archive_url.split("?")[0]) + archive_path = os.path.join(cache_dir, archive_name) # Download if not exists if not os.path.exists(archive_path): r = requests.get(archive_url, timeout=60) @@ -171,17 +204,24 @@ def load_emotion_model_multi_source( tar.extractall(path=extract_dir) elif archive_path.endswith(".zip"): import zipfile + with zipfile.ZipFile(archive_path, "r") as zf: zf.extractall(path=extract_dir) else: # Unknown archive, try treating as directory pass # Try load from extracted directory (assume single top-level) - candidates = [extract_dir] + [os.path.join(extract_dir, d) for d in os.listdir(extract_dir)] + candidates = [extract_dir] + [ + os.path.join(extract_dir, d) for d in os.listdir(extract_dir) + ] for cand in candidates: - if os.path.isdir(cand) and os.path.exists(os.path.join(cand, "config.json")): + if os.path.isdir(cand) and os.path.exists( + os.path.join(cand, "config.json") + ): try: - det = _wrap_local_model(cand, token=token, force_multi_label=force_multi_label) + det = _wrap_local_model( + cand, token=token, force_multi_label=force_multi_label + ) return det except Exception: continue @@ -198,4 +238,4 @@ def load_emotion_model_multi_source( pass # Exhausted all sources - raise RuntimeError("Could not load emotion model from any source") \ No newline at end of file + raise RuntimeError("Could not load emotion model from any source") diff --git a/src/models/emotion_detection/training_pipeline.py b/src/models/emotion_detection/training_pipeline.py index 077db6605..6267981ef 100644 --- a/src/models/emotion_detection/training_pipeline.py +++ b/src/models/emotion_detection/training_pipeline.py @@ -1,36 +1,28 @@ #!/usr/bin/env python3 -""" -Training Pipeline for BERT Emotion Detection. +"""Training Pipeline for BERT Emotion Detection. -This module provides a comprehensive training pipeline for the BERT-based -emotion detection model with advanced features like focal loss, temperature -scaling, and ensemble methods. +This module provides a comprehensive training pipeline for the BERT-based emotion +detection model with advanced features like focal loss, temperature scaling, and +ensemble methods. """ import json import logging import time from pathlib import Path -from typing import Any, Dict, List, Optional, Tuple, Union +from typing import Any, Dict, List, Optional, Union import numpy as np import torch import torch.nn.functional as F from torch.optim import AdamW from torch.utils.data import DataLoader -from transformers import ( - AutoTokenizer, - get_linear_schedule_with_warmup, -) - -from .bert_classifier import ( - create_bert_emotion_classifier, - evaluate_emotion_classifier, -) -from .dataset_loader import ( - create_goemotions_loader, - GoEmotionsDataset, -) +from transformers import AutoTokenizer, get_linear_schedule_with_warmup + +from ...utils import count_model_params + +from .bert_classifier import create_bert_emotion_classifier, evaluate_emotion_classifier +from .dataset_loader import GoEmotionsDataset, create_goemotions_loader # Configure logging # G004: Logging f-strings temporarily allowed for development @@ -98,7 +90,7 @@ def __init__( else: self.device = torch.device(device) - logger.info("Using device: {self.device}") + logger.info("Using device: %s", self.device) self.output_dir.mkdir(parents=True, exist_ok=True) @@ -109,6 +101,14 @@ def __init__( self.scheduler = None self.tokenizer = None + # Dataset attributes + self.train_dataset = None + self.val_dataset = None + self.test_dataset = None + self.train_dataloader = None + self.val_dataloader = None + self.test_dataloader = None + self.best_score = 0.0 self.patience_counter = 0 self.training_history = [] @@ -142,7 +142,8 @@ def prepare_data(self, dev_mode: bool = False) -> Dict[str, Any]: test_labels = datasets["test"]["labels"] if dev_mode: - logger.info("🔧 DEVELOPMENT MODE: Using 5% of dataset for faster training") + dev_msg = "🔧 DEVELOPMENT MODE: Using 5% of dataset for faster training" + logger.info(dev_msg) train_size = len(train_texts) dev_size = int(train_size * 0.05) # Reduced from 10% to 5% @@ -157,16 +158,23 @@ def prepare_data(self, dev_mode: bool = False) -> Dict[str, Any]: val_labels = [val_labels[i] for i in val_indices] original_batch_size = self.batch_size - self.batch_size = min(128, self.batch_size * 8) # Much larger batch size - logger.info( - "🔧 DEVELOPMENT MODE: Using {len(train_texts)} training examples, batch_size={self.batch_size} (was {original_batch_size})" + # Increase batch size for dev mode + self.batch_size = min(128, self.batch_size * 8) + dev_msg = ( + "🔧 DEVELOPMENT MODE: Using %d training examples, " + "batch_size=%d (was %d)" ) + logger.info(dev_msg, len(train_texts), self.batch_size, original_batch_size) self.train_dataset = GoEmotionsDataset( train_texts, train_labels, self.tokenizer, self.max_length ) - self.val_dataset = GoEmotionsDataset(val_texts, val_labels, self.tokenizer, self.max_length) - self.test_dataset = GoEmotionsDataset(test_texts, test_labels, self.tokenizer, self.max_length) + self.val_dataset = GoEmotionsDataset( + val_texts, val_labels, self.tokenizer, self.max_length + ) + self.test_dataset = GoEmotionsDataset( + test_texts, test_labels, self.tokenizer, self.max_length + ) self.train_dataloader = DataLoader( self.train_dataset, @@ -188,8 +196,10 @@ def prepare_data(self, dev_mode: bool = False) -> Dict[str, Any]: ) logger.info( - "Prepared datasets - Train: {len(self.train_dataset)}, " - "Val: {len(self.val_dataset)}, Test: {len(self.test_dataset)}" + "Prepared datasets - Train: %d, Val: %d, Test: %d", + len(self.train_dataset), + len(self.val_dataset), + len(self.test_dataset), ) return datasets @@ -208,19 +218,28 @@ def initialize_model(self, class_weights: Optional[np.ndarray] = None) -> None: freeze_bert_layers=self.freeze_initial_layers, ) - logger.info("🔍 DEBUG: Loss Function Analysis") - logger.info(" Loss function type: {type(self.loss_fn).__name__}") + if logger.isEnabledFor(logging.DEBUG): + logger.debug("Loss Function Analysis") + logger.debug(" Loss function type: %s", type(self.loss_fn).__name__) - if hasattr(self.loss_fn, "class_weights") and self.loss_fn.class_weights is not None: + if ( + hasattr(self.loss_fn, "class_weights") + and self.loss_fn.class_weights is not None + ): weights = self.loss_fn.class_weights - logger.info(" Class weights shape: {weights.shape}") - logger.info(" Class weights min: {weights.min().item():.6f}") - logger.info(" Class weights max: {weights.max().item():.6f}") - logger.info(" Class weights mean: {weights.mean().item():.6f}") - - if weights.min() <= 0: - logger.error("❌ CRITICAL: Class weights contain zero or negative values!") - if weights.max() > 100: + if logger.isEnabledFor(logging.DEBUG): + logger.debug( + " Class weights shape: %s", getattr(weights, "shape", None) + ) + logger.debug(" Class weights min: %.6f", weights.min().item()) + logger.debug(" Class weights mean: %.6f", weights.mean().item()) + logger.debug(" Class weights max: %.6f", weights.max().item()) + + if weights.min().item() <= 0: + logger.error( + "❌ CRITICAL: Class weights contain zero or negative values!" + ) + if weights.max().item() > 100: logger.error("❌ CRITICAL: Class weights contain very large values!") else: logger.info(" No class weights used") @@ -242,9 +261,10 @@ def initialize_model(self, class_weights: Optional[np.ndarray] = None) -> None: ) logger.info( - "Model initialized with {self.model.count_parameters():,} trainable parameters" + "Model initialized with %s trainable parameters", + format(count_model_params(self.model, only_trainable=True), ",d"), ) - logger.info("Total training steps: {total_steps}") + logger.info("Total training steps: %d", total_steps) def load_model(self, checkpoint_path: str) -> None: """Load a trained model from checkpoint. @@ -252,7 +272,7 @@ def load_model(self, checkpoint_path: str) -> None: Args: checkpoint_path: Path to the model checkpoint file """ - logger.info("Loading model from checkpoint: {checkpoint_path}") + logger.info("Loading model from checkpoint: %s", checkpoint_path) checkpoint = torch.load(checkpoint_path, map_location=self.device) @@ -277,150 +297,145 @@ def train_epoch(self, epoch: int) -> Dict[str, float]: """ self.model.train() + # Handle progressive unfreezing + self._handle_progressive_unfreezing(epoch) + + # Setup training total_loss = 0.0 num_batches = len(self.train_dataloader) start_time = time.time() + val_frequency = max(500, num_batches // 5) + logger.info("🔧 Validation frequency: every %d batches", val_frequency) + # Train on all batches + for batch_idx, batch in enumerate(self.train_dataloader): + batch_loss = self._train_single_batch(batch, batch_idx, epoch, num_batches) + total_loss += batch_loss + + # Log progress periodically + if batch_idx < 5 or (batch_idx + 1) % 100 == 0: + self._log_progress(epoch, batch_idx, num_batches, total_loss) + + # Check for early stopping + maybe_metrics = self._maybe_validate_and_early_stop( + batch_idx, + epoch, + num_batches, + total_loss, + self.scheduler.get_last_lr()[0], + val_frequency, + start_time, + ) + if maybe_metrics is not None: + return maybe_metrics + + # Return epoch metrics + return self._create_epoch_metrics(epoch, total_loss, num_batches, start_time) + + def _handle_progressive_unfreezing(self, epoch: int) -> None: + """Handle progressive unfreezing of BERT layers. + + Args: + epoch: Current epoch number + """ if epoch in self.unfreeze_schedule: layers_to_unfreeze = 2 # Unfreeze 2 layers at a time self.model.unfreeze_bert_layers(layers_to_unfreeze) - logger.info( - "Epoch {epoch}: Applied progressive unfreezing", extra={"format_args": True} - ) + logger.info("Epoch %d: Applied progressive unfreezing", epoch) - val_frequency = max(500, num_batches // 5) - logger.info("🔧 Validation frequency: every {val_frequency} batches") + def _train_single_batch( + self, + batch: Dict[str, torch.Tensor], + batch_idx: int, + epoch: int, + num_batches: int, + ) -> float: + """Train on a single batch. - for batch_idx, batch in enumerate(self.train_dataloader): - input_ids = batch["input_ids"].to(self.device) - attention_mask = batch["attention_mask"].to(self.device) - labels = batch["labels"].to(self.device) - - if batch_idx == 0: - logger.info("🔍 DEBUG: Data Distribution Analysis") - logger.info(" Labels shape: {labels.shape}") - logger.info(" Labels dtype: {labels.dtype}") - logger.info(" Labels min: {labels.min().item()}") - logger.info(" Labels max: {labels.max().item()}") - logger.info(" Labels mean: {labels.float().mean().item():.6f}") - logger.info(" Labels sum: {labels.sum().item()}") - logger.info(" Non-zero labels: {(labels > 0).sum().item()}") - logger.info(" Total labels: {labels.numel()}") - - if labels.sum() == 0: - logger.error("❌ CRITICAL: All labels are zero!") - elif labels.sum() == labels.numel(): - logger.error("❌ CRITICAL: All labels are one!") - - for i in range(min(10, labels.shape[1])): # First 10 classes - class_count = labels[:, i].sum().item() - if class_count > 0: - logger.info(" Class {i}: {class_count} positive samples") - - self.optimizer.zero_grad() - - logits = self.model(input_ids, attention_mask) - - if batch_idx == 0: - logger.info("🔍 DEBUG: Model Output Analysis") - logger.info(" Logits shape: {logits.shape}") - logger.info(" Logits min: {logits.min().item():.6f}") - logger.info(" Logits max: {logits.max().item():.6f}") - logger.info(" Logits mean: {logits.mean().item():.6f}") - logger.info(" Logits std: {logits.std().item():.6f}") - - if torch.isnan(logits).any(): - logger.error("❌ CRITICAL: NaN values in logits!") - if torch.isinf(logits).any(): - logger.error("❌ CRITICAL: Inf values in logits!") - - predictions = torch.sigmoid(logits) - logger.info(" Predictions min: {predictions.min().item():.6f}") - logger.info(" Predictions max: {predictions.max().item():.6f}") - logger.info(" Predictions mean: {predictions.mean().item():.6f}") - - loss = self.loss_fn(logits, labels) - - if batch_idx == 0: - logger.info("🔍 DEBUG: Loss Analysis") - logger.info(" Raw loss: {loss.item():.8f}") - - bce_manual = F.binary_cross_entropy_with_logits( - logits, labels.float(), reduction="mean" - ) - logger.info(" Manual BCE loss: {bce_manual.item():.8f}") - - if abs(loss.item()) < 1e-10: - logger.error("❌ CRITICAL: Loss is effectively zero!") - logger.error(" This indicates a serious training issue!") - - for i in range(min(5, logits.shape[1])): - class_logits = logits[:, i] - class_labels = labels[:, i].float() - class_loss = F.binary_cross_entropy_with_logits( - class_logits, class_labels, reduction="mean" - ) - logger.info(" Class {i} loss: {class_loss.item():.8f}") - - loss.backward() - - if batch_idx == 0: - logger.info("🔍 DEBUG: Gradient Analysis") - total_norm = 0 - param_count = 0 - for p in self.model.parameters(): - if p.grad is not None: - param_norm = p.grad.data.norm(2) - total_norm += param_norm.item() ** 2 - param_count += 1 - - if param_count > 0: - total_norm = total_norm ** (1.0 / 2) - logger.info(" Gradient norm before clipping: {total_norm:.6f}") - - if total_norm > 10: - logger.warning("⚠️ WARNING: Large gradient norm detected!") - if total_norm < 1e-6: - logger.warning("⚠️ WARNING: Very small gradient norm detected!") - - clip_norm = torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm=1.0) - - if batch_idx == 0: - logger.info(" Gradient norm after clipping: {clip_norm:.6f}") - - self.optimizer.step() - self.scheduler.step() - - total_loss += loss.item() - - if batch_idx < 5 or (batch_idx + 1) % 100 == 0: # First 5 batches + every 100 - avg_loss = total_loss / (batch_idx + 1) - current_lr = self.scheduler.get_last_lr()[0] - - logger.info( - "Epoch {epoch}, Batch {batch_idx + 1}/{num_batches}, " - "Loss: {avg_loss:.8f}, LR: {current_lr:.2e}" - ) + Args: + batch: Input batch data + batch_idx: Current batch index + epoch: Current epoch number + num_batches: Total number of batches - if avg_loss < 1e-8: - logger.error("❌ CRITICAL: Average loss is suspiciously small: {avg_loss:.8f}") - if avg_loss > 100: - logger.error("❌ CRITICAL: Average loss is suspiciously large: {avg_loss:.8f}") + Returns: + float: Loss value for this batch + """ + # Move data to device + input_ids = batch["input_ids"].to(self.device) + attention_mask = batch["attention_mask"].to(self.device) + labels = batch["labels"].to(self.device) + + # Log debug information for first batch + if batch_idx == 0: + self._log_batch_debug_info(labels, None, None) # logits not available yet + + # Forward pass + self.optimizer.zero_grad() + logits = self.model(input_ids, attention_mask) + loss = self.loss_fn(logits, labels) + + # Log debug information for first batch + if batch_idx == 0: + self._log_batch_debug_info(labels, logits, loss) + + # Backward pass + loss.backward() + + # Gradient clipping + clip_norm = torch.nn.utils.clip_grad_norm_( + self.model.parameters(), max_norm=1.0 + ) - if (batch_idx + 1) % val_frequency == 0: - logger.info("🔍 Validating at batch {batch_idx + 1}...") - self.validate(epoch) + # Log gradient stats for first batch + if batch_idx == 0: + self._log_gradient_stats_before() + self._log_gradient_stats_after(clip_norm) - if self.should_stop_early(): - logger.info("🛑 Early stopping triggered at batch {batch_idx + 1}") - return { - "epoch": epoch, - "train_loss": total_loss / (batch_idx + 1), - "epoch_time": time.time() - start_time, - "learning_rate": self.scheduler.get_last_lr()[0], - "early_stopped": True, - } + # Update parameters + self.optimizer.step() + self.scheduler.step() + + return loss.item() + + def _log_batch_debug_info( + self, + labels: torch.Tensor, + logits: Optional[torch.Tensor], + loss: Optional[torch.Tensor], + ) -> None: + """Log debug information for the first batch. + + Args: + labels: Ground truth labels + logits: Model output logits (None for first call) + loss: Computed loss (None for first call) + """ + if not logger.isEnabledFor(logging.DEBUG): + return + + self._log_data_distribution(labels) + + if logits is not None: + self._log_model_output(logits) + + if loss is not None: + self._log_loss_analysis(loss, logits, labels) + + def _create_epoch_metrics( + self, epoch: int, total_loss: float, num_batches: int, start_time: float + ) -> Dict[str, float]: + """Create metrics dictionary for completed epoch. + + Args: + epoch: Current epoch number + total_loss: Total loss for the epoch + num_batches: Number of batches in epoch + start_time: Start time of epoch + Returns: + Dictionary with epoch metrics + """ epoch_time = time.time() - start_time avg_loss = total_loss / num_batches @@ -431,10 +446,198 @@ def train_epoch(self, epoch: int) -> Dict[str, float]: "learning_rate": self.scheduler.get_last_lr()[0], } - logger.info("Epoch {epoch} completed - Loss: {avg_loss:.4f}, Time: {epoch_time:.1f}s") + logger.info( + "Epoch %d completed - Loss: %.4f, Time: %.1fs", + epoch, + avg_loss, + epoch_time, + ) return metrics + @staticmethod + def _log_data_distribution(labels: torch.Tensor) -> None: + """Log data distribution analysis for debugging. + + Args: + labels: Ground truth labels tensor + """ + logger.info("🔍 DEBUG: Data Distribution Analysis") + logger.info(" Labels shape: %s", labels.shape) + logger.info(" Labels dtype: %s", labels.dtype) + logger.info(" Labels min: %s", labels.min().item()) + logger.info(" Labels max: %s", labels.max().item()) + logger.info(" Labels mean: %.6f", labels.float().mean().item()) + logger.info(" Labels sum: %s", labels.sum().item()) + logger.info(" Non-zero labels: %s", (labels > 0).sum().item()) + logger.info(" Total labels: %s", labels.numel()) + + if labels.sum() == 0: + logger.error("❌ CRITICAL: All labels are zero!") + elif labels.sum() == labels.numel(): + logger.error("❌ CRITICAL: All labels are one!") + + # Iterate over first 10 classes + max_classes = min(10, labels.shape[1]) + for i in range(max_classes): + class_count = labels[:, i].sum().item() + if class_count > 0: + logger.info(" Class %d: %d positive samples", i, int(class_count)) + + @staticmethod + def _log_model_output(logits: torch.Tensor) -> None: + """Log model output analysis for debugging. + + Args: + logits: Model output logits tensor + """ + logger.info("🔍 DEBUG: Model Output Analysis") + logger.info(" Logits shape: %s", logits.shape) + logger.info(" Logits min: %.6f", logits.min().item()) + logger.info(" Logits max: %.6f", logits.max().item()) + logger.info(" Logits mean: %.6f", logits.mean().item()) + logger.info(" Logits std: %.6f", logits.std().item()) + if torch.isnan(logits).any(): + logger.error("❌ CRITICAL: NaN values in logits!") + if torch.isinf(logits).any(): + logger.error("❌ CRITICAL: Inf values in logits!") + predictions = torch.sigmoid(logits) + logger.info(" Predictions min: %.6f", predictions.min().item()) + logger.info(" Predictions max: %.6f", predictions.max().item()) + logger.info(" Predictions mean: %.6f", predictions.mean().item()) + + @staticmethod + def _log_loss_analysis( + loss: torch.Tensor, logits: torch.Tensor, labels: torch.Tensor + ) -> None: + """Log detailed loss analysis for debugging. + + Args: + loss: Computed loss tensor + logits: Model output logits + labels: Ground truth labels + """ + logger.info("🔍 DEBUG: Loss Analysis") + logger.info(" Raw loss: %.8f", loss.item()) + bce_manual = F.binary_cross_entropy_with_logits( + logits, labels.float(), reduction="mean" + ) + logger.info(" Manual BCE loss: %.8f", bce_manual.item()) + if abs(loss.item()) < 1e-10: + logger.error("❌ CRITICAL: Loss is effectively zero!") + logger.error(" This indicates a serious training issue!") + for i in range(min(5, logits.shape[1])): + class_logits = logits[:, i] + class_labels = labels[:, i].float() + class_loss = F.binary_cross_entropy_with_logits( + class_logits, class_labels, reduction="mean" + ) + logger.info(" Class %d loss: %.8f", i, class_loss.item()) + + def _log_gradient_stats_before(self) -> None: + """Log gradient statistics before gradient clipping. + + Analyzes gradient norms across all model parameters and logs statistics for + debugging purposes. + """ + logger.info("🔍 DEBUG: Gradient Analysis") + total_norm = 0.0 + param_count = 0 + for p in self.model.parameters(): + if p.grad is not None: + param_norm = p.grad.data.norm(2) + total_norm += param_norm.item() ** 2 + param_count += 1 + if param_count > 0: + total_norm = total_norm**0.5 + logger.info(" Gradient norm before clipping: %.6f", total_norm) + if total_norm > 10: + logger.warning("⚠️ WARNING: Large gradient norm detected!") + if total_norm < 1e-6: + logger.warning("⚠️ WARNING: Very small gradient norm detected!") + + @staticmethod + def _log_gradient_stats_after(clip_norm: Union[float, torch.Tensor]) -> None: + """Log gradient statistics after gradient clipping. + + Args: + clip_norm: Gradient norm value after clipping + """ + if not isinstance(clip_norm, (int, float)): + clip_val = float(clip_norm) + else: + clip_val = clip_norm + logger.info(" Gradient norm after clipping: %.6f", clip_val) + + def _log_progress( + self, epoch: int, batch_idx: int, num_batches: int, total_loss: float + ) -> None: + """Log training progress information. + + Args: + epoch: Current epoch number + batch_idx: Current batch index + num_batches: Total number of batches in epoch + total_loss: Cumulative loss for current epoch + """ + avg_loss = total_loss / (batch_idx + 1) + current_lr = self.scheduler.get_last_lr()[0] + logger.info( + "Epoch %d, Batch %d/%d, Loss: %.8f, LR: %.2e", + epoch, + batch_idx + 1, + num_batches, + avg_loss, + current_lr, + ) + if avg_loss < 1e-8: + logger.error( + "❌ CRITICAL: Average loss is suspiciously small: %.8f", avg_loss + ) + if avg_loss > 100: + logger.error( + "❌ CRITICAL: Average loss is suspiciously large: %.8f", avg_loss + ) + + def _maybe_validate_and_early_stop( + self, + batch_idx: int, + epoch: int, + num_batches: int, + total_loss: float, + current_lr: float, + val_frequency: int, + start_time: float, + ) -> Optional[Dict[str, Any]]: + """Check if validation should be performed and handle early stopping. + + Args: + batch_idx: Current batch index + epoch: Current epoch number + num_batches: Total number of batches in epoch + total_loss: Total loss for current epoch + current_lr: Current learning rate + val_frequency: Frequency of validation + start_time: Start time of training + + Returns: + Dictionary with early stopping metrics if stopping, None otherwise + """ + if (batch_idx + 1) % val_frequency != 0: + return None + logger.info("🔍 Validating at batch %d...", batch_idx + 1) + self.validate(epoch) + if self.should_stop_early(): + logger.info("🛑 Early stopping triggered at batch %d", batch_idx + 1) + return { + "epoch": epoch, + "train_loss": total_loss / (batch_idx + 1), + "epoch_time": time.time() - start_time, + "learning_rate": current_lr, + "early_stopped": True, + } + return None + def validate(self, epoch: int) -> Dict[str, float]: """Validate model performance. @@ -444,7 +647,7 @@ def validate(self, epoch: int) -> Dict[str, float]: Returns: Dictionary with validation metrics """ - logger.info("Validating model at epoch {epoch}...") + logger.info("Validating model at epoch %d...", epoch) val_metrics = evaluate_emotion_classifier( self.model, self.val_dataloader, self.device, threshold=0.2 @@ -459,11 +662,13 @@ def validate(self, epoch: int) -> Dict[str, float]: if self.save_best_only: self.save_checkpoint(epoch, val_metrics, is_best=True) - logger.info("New best model saved! Macro F1: {current_score:.4f}") + logger.info("New best model saved! Macro F1: %.4f", current_score) else: self.patience_counter += 1 logger.info( - "No improvement. Patience: {self.patience_counter}/{self.early_stopping_patience}" + "No improvement. Patience: %d/%d", + self.patience_counter, + self.early_stopping_patience, ) return val_metrics @@ -472,7 +677,9 @@ def should_stop_early(self) -> bool: """Check if training should stop early.""" return self.patience_counter >= self.early_stopping_patience - def save_checkpoint(self, epoch: int, metrics: Dict[str, float], is_best: bool = False) -> None: + def save_checkpoint( + self, epoch: int, metrics: Dict[str, float], is_best: bool = False + ) -> None: """Save model checkpoint. Args: @@ -498,10 +705,10 @@ def save_checkpoint(self, epoch: int, metrics: Dict[str, float], is_best: bool = if is_best: checkpoint_path = self.output_dir / "best_model.pt" else: - checkpoint_path = self.output_dir / "checkpoint_epoch_{epoch}.pt" + checkpoint_path = self.output_dir / f"checkpoint_epoch_{epoch}.pt" torch.save(checkpoint, checkpoint_path) - logger.info("Checkpoint saved: {checkpoint_path}") + logger.info("Checkpoint saved: %s", checkpoint_path) def train(self) -> Dict[str, Any]: """Complete training pipeline. @@ -526,7 +733,7 @@ def train(self) -> Dict[str, Any]: self.training_history.append(epoch_metrics) if self.should_stop_early(): - logger.info(f"Early stopping at epoch {epoch}") + logger.info("Early stopping at epoch %d", epoch) break else: self.training_history.append(train_metrics) @@ -556,7 +763,7 @@ def convert_numpy_types(obj): serializable_history = convert_numpy_types(self.training_history) with Path(history_path).open("w") as f: json.dump(serializable_history, f, indent=2) - logger.info("Training history saved to {history_path}") + logger.info("Training history saved to %s", history_path) except Exception: logger.exception("Failed to save training history") simplified_history = [] @@ -576,7 +783,7 @@ def convert_numpy_types(obj): with Path(history_path).open("w") as f: json.dump(simplified_history, f, indent=2) - logger.info("Simplified training history saved to {history_path}") + logger.info("Simplified training history saved to %s", history_path) results = { "final_test_metrics": test_metrics, @@ -587,9 +794,9 @@ def convert_numpy_types(obj): } logger.info("✅ Training completed!") - logger.info(f"Best validation Macro F1: {self.best_score:.4f}") - logger.info(f"Final test Macro F1: {test_metrics['macro_f1']:.4f}") - logger.info(f"Final test Micro F1: {test_metrics['micro_f1']:.4f}") + logger.info("Best validation Macro F1: %.4f", self.best_score) + logger.info("Final test Macro F1: %.4f", test_metrics["macro_f1"]) + logger.info("Final test Micro F1: %.4f", test_metrics["micro_f1"]) return results @@ -602,9 +809,8 @@ def train_emotion_detection_model( learning_rate: float = 2e-6, # Reduced from 2e-5 to 2e-6 for debugging num_epochs: int = 3, device: Optional[str] = None, - dev_mode: bool = True, # Enable development mode by default - debug_mode: bool = True, # Enable debugging by default - ) -> Dict[str, Any]: + dev_mode: bool = False, +) -> Dict[str, Any]: """Convenient function to train emotion detection model with default settings. Args: @@ -615,8 +821,7 @@ def train_emotion_detection_model( learning_rate: Learning rate for optimization num_epochs: Number of training epochs device: Device to use for training (auto-detect if None) - dev_mode: Enable development mode with smaller dataset - debug_mode: Enable debugging mode with enhanced logging + dev_mode: If True, use a small subset of data for quicker iterations Returns: Dictionary containing training results and metrics diff --git a/src/security/jwt_manager.py b/src/security/jwt_manager.py index fad3d47f3..17dd608b3 100644 --- a/src/security/jwt_manager.py +++ b/src/security/jwt_manager.py @@ -1,5 +1,4 @@ -""" -JWT-based Authentication Manager for SAMO Deep Learning API +"""JWT-based Authentication Manager for SAMO Deep Learning API. This module provides comprehensive JWT token management including: - Access and refresh token creation @@ -11,9 +10,8 @@ import logging import os -import time from datetime import datetime, timedelta -from typing import Dict, List, Optional, Union, Any +from typing import Any, Dict, List, Optional import jwt from pydantic import BaseModel, Field @@ -27,8 +25,10 @@ ACCESS_TOKEN_EXPIRE_MINUTES = 30 REFRESH_TOKEN_EXPIRE_DAYS = 7 + class TokenPayload(BaseModel): - """Token payload structure""" + """Token payload structure.""" + user_id: str = Field(..., description="User identifier") username: str = Field(..., description="Username") email: str = Field(..., description="User email") @@ -37,17 +37,21 @@ class TokenPayload(BaseModel): iat: Optional[int] = Field(None, description="Issued at timestamp") type: Optional[str] = Field(None, description="Token type, e.g., 'refresh'") + class TokenResponse(BaseModel): - """Token response structure""" + """Token response structure.""" + access_token: str = Field(..., description="Access token") refresh_token: str = Field(..., description="Refresh token") token_type: str = Field(default="bearer", description="Token type") expires_in: int = Field(..., description="Access token expiration in seconds") + # Token pair is returned as a plain dict + class JWTManager: - """Comprehensive JWT token management system""" + """Comprehensive JWT token management system.""" def __init__(self, secret_key: str = SECRET_KEY, algorithm: str = ALGORITHM): self.secret_key = secret_key @@ -56,19 +60,19 @@ def __init__(self, secret_key: str = SECRET_KEY, algorithm: str = ALGORITHM): self.blacklisted_tokens: dict = {} def create_access_token(self, user_data: Dict[str, Any]) -> str: - """Create a new access token""" + """Create a new access token.""" payload = { "user_id": user_data["user_id"], "username": user_data["username"], "email": user_data["email"], "permissions": user_data.get("permissions", []), "exp": datetime.utcnow() + timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES), - "iat": datetime.utcnow() + "iat": datetime.utcnow(), } return jwt.encode(payload, self.secret_key, algorithm=self.algorithm) def create_refresh_token(self, user_data: Dict[str, Any]) -> str: - """Create a new refresh token""" + """Create a new refresh token.""" payload = { "user_id": user_data["user_id"], "username": user_data["username"], @@ -76,7 +80,7 @@ def create_refresh_token(self, user_data: Dict[str, Any]) -> str: "permissions": user_data.get("permissions", []), "exp": datetime.utcnow() + timedelta(days=REFRESH_TOKEN_EXPIRE_DAYS), "iat": datetime.utcnow(), - "type": "refresh" + "type": "refresh", } return jwt.encode(payload, self.secret_key, algorithm=self.algorithm) @@ -87,11 +91,11 @@ def create_token_pair(self, user_data: Dict[str, Any]) -> TokenResponse: return TokenResponse( access_token=access_token, refresh_token=refresh_token, - expires_in=ACCESS_TOKEN_EXPIRE_MINUTES * 60 + expires_in=ACCESS_TOKEN_EXPIRE_MINUTES * 60, ) def verify_token(self, token: str) -> Optional[TokenPayload]: - """Verify and decode a token""" + """Verify and decode a token.""" try: if token in self.blacklisted_tokens: return None @@ -102,14 +106,14 @@ def verify_token(self, token: str) -> Optional[TokenPayload]: logger.warning(f"Token expired: {token[:10]}...") return None except jwt.InvalidTokenError as e: - logger.warning(f"Invalid token: {str(e)}") + logger.warning(f"Invalid token: {e!s}") return None except Exception as e: - logger.error(f"Token verification error: {str(e)}") + logger.error(f"Token verification error: {e!s}") return None def refresh_access_token(self, refresh_token: str) -> Optional[str]: - """Refresh an access token using a valid refresh token""" + """Refresh an access token using a valid refresh token.""" payload = self.verify_token(refresh_token) if not payload or getattr(payload, "type", None) != "refresh": return None @@ -118,36 +122,39 @@ def refresh_access_token(self, refresh_token: str) -> Optional[str]: "user_id": payload.user_id, "username": payload.username, "email": payload.email, - "permissions": payload.permissions + "permissions": payload.permissions, } return self.create_access_token(user_data) def blacklist_token(self, token: str) -> bool: - """Add a token to the blacklist""" + """Add a token to the blacklist.""" try: payload = jwt.decode(token, self.secret_key, algorithms=[self.algorithm]) - exp_datetime = datetime.fromtimestamp(payload["exp"]) if payload.get("exp") else None + exp_timestamp = payload.get("exp") + exp_datetime = ( + datetime.fromtimestamp(exp_timestamp) if exp_timestamp else None + ) self.blacklisted_tokens[token] = exp_datetime return True except jwt.InvalidTokenError: return False def is_token_blacklisted(self, token: str) -> bool: - """Check if a token is blacklisted""" + """Check if a token is blacklisted.""" return token in self.blacklisted_tokens def get_user_permissions(self, token: str) -> List[str]: - """Extract user permissions from token""" + """Extract user permissions from token.""" payload = self.verify_token(token) return payload.permissions if payload else [] def has_permission(self, token: str, required_permission: str) -> bool: - """Check if user has a specific permission""" + """Check if user has a specific permission.""" permissions = self.get_user_permissions(token) return required_permission in permissions def cleanup_expired_tokens(self) -> int: - """Clean up expired tokens from blacklist""" + """Clean up expired tokens from blacklist.""" initial_count = len(self.blacklisted_tokens) current_time = datetime.utcnow() @@ -161,5 +168,6 @@ def cleanup_expired_tokens(self) -> int: self.blacklisted_tokens.pop(token, None) return initial_count - len(self.blacklisted_tokens) + # Global JWT manager instance jwt_manager = JWTManager() diff --git a/src/security_headers.py b/src/security_headers.py index b4a579794..3b1eb0b2a 100644 --- a/src/security_headers.py +++ b/src/security_headers.py @@ -1,25 +1,27 @@ #!/usr/bin/env python3 -""" -🛡️ Security Headers Middleware +"""🛡️ Security Headers Middleware ============================== Flask middleware for adding security headers and implementing security policies. """ -import logging -from typing import Dict, List, Optional, Callable -from dataclasses import dataclass -from flask import Flask, request, Response, g -import time import hashlib +import logging +import os import secrets +import time +from dataclasses import dataclass +from typing import Dict, List + import yaml -import os +from flask import Flask, Response, g, request logger = logging.getLogger(__name__) + @dataclass class SecurityHeadersConfig: """Security headers configuration.""" + enable_csp: bool = True enable_hsts: bool = True enable_x_frame_options: bool = True @@ -40,9 +42,9 @@ class SecurityHeadersConfig: ua_suspicious_score_threshold: int = 4 # Score threshold for suspicious UAs ua_blocking_enabled: bool = False # Whether to block suspicious UAs (vs just log) + class SecurityHeadersMiddleware: - """ - Flask middleware for adding security headers and implementing security policies. + """Flask middleware for adding security headers and implementing security policies. Features: - Content Security Policy (CSP) @@ -63,9 +65,16 @@ def __init__(self, app: Flask, config: SecurityHeadersConfig): # Load CSP from YAML config if available self.csp_policy = None try: - with open(os.path.join(os.path.dirname(__file__), '../configs/security.yaml'), 'r') as f: + config_path = os.path.join( + os.path.dirname(__file__), "../configs/security.yaml" + ) + with open(config_path) as f: security_config = yaml.safe_load(f) - self.csp_policy = security_config.get('security_headers', {}).get('headers', {}).get('Content-Security-Policy') + self.csp_policy = ( + security_config.get("security_headers", {}) + .get("headers", {}) + .get("Content-Security-Policy") + ) except Exception as e: logger.warning(f"Could not load CSP from config: {e}") @@ -80,17 +89,41 @@ def _before_request(self): """Process request before handling.""" # Generate request ID for correlation if self.config.enable_request_id: - g.request_id = hashlib.sha256( - f"{time.time()}:{request.remote_addr}:{secrets.token_hex(8)}".encode() - ).hexdigest() + request_data = f"{time.time()}:{request.remote_addr}:{secrets.token_hex(8)}" + g.request_id = hashlib.sha256(request_data.encode()).hexdigest() # Generate correlation ID if self.config.enable_correlation_id: - g.correlation_id = request.headers.get('X-Correlation-ID', g.request_id) + g.correlation_id = request.headers.get("X-Correlation-ID", g.request_id) + + # Check for high-risk requests and block them + if self.config.ua_blocking_enabled: + user_agent = request.headers.get("User-Agent", "") + ua_analysis = self._analyze_user_agent_enhanced(user_agent) + + if ua_analysis["risk_level"] in ["high", "very_high"]: + # Store blocking information for after_request logging + g.security_patterns = [ + f"BLOCKED: High-risk user agent - {ua_analysis['category']} " + f"(score: {ua_analysis['score']})" + ] + g.block_reason = f"User agent risk level: {ua_analysis['risk_level']}" + g.ua_analysis = ua_analysis + + # Return 403 Forbidden response + from flask import make_response + response = make_response( + "Access Forbidden - High-risk user agent detected", 403 + ) + response.headers["Content-Type"] = "text/plain" + return response # Log security-relevant request information self._log_security_info() + # Explicit return None for consistency (PEP8 compliance) + return None + def _after_request(self, response: Response) -> Response: """Process response after handling.""" # Add security headers @@ -102,6 +135,14 @@ def _after_request(self, response: Response) -> Response: # Log security-relevant response information self._log_response_security(response) + # Log blocking information if request was blocked + if hasattr(g, 'security_patterns'): + logger.warning("Request blocked: %s", g.security_patterns) + if hasattr(g, 'block_reason'): + logger.warning("Block reason: %s", g.block_reason) + if hasattr(g, 'ua_analysis'): + logger.warning("User agent analysis: %s", g.ua_analysis) + return response def _add_security_headers(self, response: Response): @@ -109,48 +150,50 @@ def _add_security_headers(self, response: Response): # Content Security Policy if self.config.enable_content_security_policy: csp_policy = self._build_csp_policy() - response.headers['Content-Security-Policy'] = csp_policy + response.headers["Content-Security-Policy"] = csp_policy # HTTP Strict Transport Security if self.config.enable_strict_transport_security: - response.headers['Strict-Transport-Security'] = 'max-age=31536000; includeSubDomains; preload' + hsts_value = "max-age=31536000; includeSubDomains; preload" + response.headers["Strict-Transport-Security"] = hsts_value # X-Frame-Options if self.config.enable_x_frame_options: - response.headers['X-Frame-Options'] = 'DENY' + response.headers["X-Frame-Options"] = "DENY" # X-Content-Type-Options if self.config.enable_x_content_type_options: - response.headers['X-Content-Type-Options'] = 'nosniff' + response.headers["X-Content-Type-Options"] = "nosniff" # X-XSS-Protection if self.config.enable_x_xss_protection: - response.headers['X-XSS-Protection'] = '1; mode=block' + response.headers["X-XSS-Protection"] = "1; mode=block" # Referrer Policy if self.config.enable_referrer_policy: - response.headers['Referrer-Policy'] = 'strict-origin-when-cross-origin' + referrer_policy = "strict-origin-when-cross-origin" + response.headers["Referrer-Policy"] = referrer_policy # Permissions Policy if self.config.enable_permissions_policy: permissions_policy = self._build_permissions_policy() - response.headers['Permissions-Policy'] = permissions_policy + response.headers["Permissions-Policy"] = permissions_policy # Cross-Origin Embedder Policy if self.config.enable_cross_origin_embedder_policy: - response.headers['Cross-Origin-Embedder-Policy'] = 'require-corp' + response.headers["Cross-Origin-Embedder-Policy"] = "require-corp" # Cross-Origin Opener Policy if self.config.enable_cross_origin_opener_policy: - response.headers['Cross-Origin-Opener-Policy'] = 'same-origin' + response.headers["Cross-Origin-Opener-Policy"] = "same-origin" # Cross-Origin Resource Policy if self.config.enable_cross_origin_resource_policy: - response.headers['Cross-Origin-Resource-Policy'] = 'same-origin' + response.headers["Cross-Origin-Resource-Policy"] = "same-origin" # Origin-Agent-Cluster if self.config.enable_origin_agent_cluster: - response.headers['Origin-Agent-Cluster'] = '?1' + response.headers["Origin-Agent-Cluster"] = "?1" def _build_csp_policy(self) -> str: """Return CSP policy from config, or a secure default if not set.""" @@ -195,40 +238,40 @@ def _build_permissions_policy(self) -> str: "sync-xhr=()", "usb=()", "web-share=()", - "xr-spatial-tracking=()" + "xr-spatial-tracking=()", ] return ", ".join(policies) def _add_correlation_headers(self, response: Response): """Add request correlation headers.""" - if hasattr(g, 'request_id'): - response.headers['X-Request-ID'] = g.request_id + if hasattr(g, "request_id"): + response.headers["X-Request-ID"] = g.request_id - if hasattr(g, 'correlation_id'): - response.headers['X-Correlation-ID'] = g.correlation_id + if hasattr(g, "correlation_id"): + response.headers["X-Correlation-ID"] = g.correlation_id def _log_security_info(self): """Log security-relevant request information.""" security_info = { - 'timestamp': time.time(), - 'request_id': getattr(g, 'request_id', None), - 'correlation_id': getattr(g, 'correlation_id', None), - 'method': request.method, - 'path': request.path, - 'remote_addr': request.remote_addr, - 'user_agent': request.headers.get('User-Agent', ''), - 'content_type': request.headers.get('Content-Type', ''), - 'content_length': request.headers.get('Content-Length', ''), - 'referer': request.headers.get('Referer', ''), - 'origin': request.headers.get('Origin', ''), - 'x_forwarded_for': request.headers.get('X-Forwarded-For', ''), - 'x_real_ip': request.headers.get('X-Real-IP', ''), + "timestamp": time.time(), + "request_id": getattr(g, "request_id", None), + "correlation_id": getattr(g, "correlation_id", None), + "method": request.method, + "path": request.path, + "remote_addr": request.remote_addr, + "user_agent": request.headers.get("User-Agent", ""), + "content_type": request.headers.get("Content-Type", ""), + "content_length": request.headers.get("Content-Length", ""), + "referer": request.headers.get("Referer", ""), + "origin": request.headers.get("Origin", ""), + "x_forwarded_for": request.headers.get("X-Forwarded-For", ""), + "x_real_ip": request.headers.get("X-Real-IP", ""), } # Log suspicious patterns suspicious_patterns = self._detect_suspicious_patterns() if suspicious_patterns: - security_info['suspicious_patterns'] = suspicious_patterns + security_info["suspicious_patterns"] = suspicious_patterns logger.warning(f"Security warning: {suspicious_patterns}") logger.info(f"Security audit: {security_info}") @@ -236,7 +279,12 @@ def _log_security_info(self): def _analyze_user_agent_enhanced(self, user_agent: str) -> dict: """Enhanced user agent analysis with scoring and detailed categorization.""" if not user_agent: - return {"score": 0, "category": "empty", "patterns": [], "risk_level": "low"} + return { + "score": 0, + "category": "empty", + "patterns": [], + "risk_level": "low", + } score = 0 patterns = [] @@ -244,29 +292,74 @@ def _analyze_user_agent_enhanced(self, user_agent: str) -> dict: # Legitimate bot whitelist (negative scoring) legitimate_bots = [ - 'googlebot', 'bingbot', 'slurp', 'duckduckbot', 'facebookexternalhit', - 'twitterbot', 'linkedinbot', 'whatsapp', 'telegrambot', 'discordbot', - 'slackbot', 'github-camo', 'github-actions', 'vercel', 'netlify', - 'uptimerobot', 'pingdom', 'statuscake', 'monitor', 'healthcheck' + "googlebot", + "bingbot", + "slurp", + "duckduckbot", + "facebookexternalhit", + "twitterbot", + "linkedinbot", + "whatsapp", + "telegrambot", + "discordbot", + "slackbot", + "github-camo", + "github-actions", + "vercel", + "netlify", + "uptimerobot", + "pingdom", + "statuscake", + "monitor", + "healthcheck", ] # High-risk patterns (score +3 each) high_risk_patterns = [ - 'sqlmap', 'nikto', 'nmap', 'scanner', 'grabber', 'harvester', - 'exploit', 'vulnerability', 'penetration', 'security', 'audit' + "sqlmap", + "nikto", + "nmap", + "scanner", + "grabber", + "harvester", + "exploit", + "vulnerability", + "penetration", + "security", + "audit", ] # Medium-risk patterns (score +2 each) medium_risk_patterns = [ - 'headless', 'phantom', 'selenium', 'webdriver', 'automated', - 'testing', 'script', 'python-requests', 'curl', 'wget', - 'httrack', 'scraper', 'crawler', 'spider', 'bot' + "headless", + "phantom", + "selenium", + "webdriver", + "automated", + "testing", + "script", + "python-requests", + "curl", + "wget", + "httrack", + "scraper", + "crawler", + "spider", + "bot", ] # Low-risk patterns (score +1 each) low_risk_patterns = [ - 'indexer', 'feed', 'rss', 'aggregator', 'monitor', 'checker', - 'validator', 'linter', 'checker', 'analyzer' + "indexer", + "feed", + "rss", + "aggregator", + "monitor", + "checker", + "validator", + "linter", + "checker", + "analyzer", ] # Check legitimate bots first (negative scoring) @@ -298,13 +391,18 @@ def _analyze_user_agent_enhanced(self, user_agent: str) -> dict: logger.debug(f"Low-risk UA pattern detected: {pattern}") # Bonus for suspicious combinations - if any(pattern in ua_lower for pattern in ['bot', 'crawler', 'spider']) and any(pattern in ua_lower for pattern in ['python', 'curl', 'wget', 'script']): + bot_patterns = ["bot", "crawler", "spider"] + script_patterns = ["python", "curl", "wget", "script"] + + if any(pattern in ua_lower for pattern in bot_patterns) and any( + pattern in ua_lower for pattern in script_patterns + ): score += 2 patterns.append("suspicious_combination") logger.debug("Suspicious UA combination detected") # Check for missing or generic user agents - if user_agent in ['', 'null', 'undefined', 'unknown', 'anonymous']: + if user_agent in ["", "null", "undefined", "unknown", "anonymous"]: score += 2 patterns.append("missing_generic_ua") logger.debug("Missing or generic user agent detected") @@ -331,7 +429,7 @@ def _analyze_user_agent_enhanced(self, user_agent: str) -> dict: "category": category, "patterns": patterns, "risk_level": risk_level, - "user_agent": user_agent[:100] # Truncate for logging + "user_agent": user_agent[:100], # Truncate for logging } def _detect_suspicious_patterns(self) -> List[str]: @@ -340,10 +438,10 @@ def _detect_suspicious_patterns(self) -> List[str]: # Check for suspicious headers suspicious_headers = [ - 'X-Forwarded-Host', - 'X-Original-URL', - 'X-Rewrite-URL', - 'X-Custom-IP-Authorization' + "X-Forwarded-Host", + "X-Original-URL", + "X-Rewrite-URL", + "X-Custom-IP-Authorization", ] for header in suspicious_headers: @@ -352,8 +450,16 @@ def _detect_suspicious_patterns(self) -> List[str]: # Check for suspicious query parameters suspicious_params = [ - 'cmd', 'exec', 'system', 'eval', 'script', - 'union', 'select', 'insert', 'update', 'delete' + "cmd", + "exec", + "system", + "eval", + "script", + "union", + "select", + "insert", + "update", + "delete", ] for param in suspicious_params: @@ -362,39 +468,44 @@ def _detect_suspicious_patterns(self) -> List[str]: # Enhanced user agent analysis if self.config.enable_enhanced_ua_analysis: - user_agent = request.headers.get('User-Agent', '') + user_agent = request.headers.get("User-Agent", "") ua_analysis = self._analyze_user_agent_enhanced(user_agent) if ua_analysis["score"] >= self.config.ua_suspicious_score_threshold: - patterns.append(f"Suspicious user agent: {ua_analysis['category']} (score: {ua_analysis['score']})") + ua_msg = ( + f"Suspicious user agent: {ua_analysis['category']} " + f"(score: {ua_analysis['score']})" + ) + patterns.append(ua_msg) # Log detailed analysis logger.warning(f"User agent analysis: {ua_analysis}") - # Optionally block based on configuration - if self.config.ua_blocking_enabled and ua_analysis["risk_level"] in ["high", "very_high"]: - patterns.append("BLOCKED: High-risk user agent") + # Note: High-risk user agents are now blocked in _before_request + # This is just for logging and pattern detection return patterns def _log_response_security(self, response: Response): """Log security-relevant response information.""" security_info = { - 'timestamp': time.time(), - 'request_id': getattr(g, 'request_id', None), - 'correlation_id': getattr(g, 'correlation_id', None), - 'status_code': response.status_code, - 'content_type': response.headers.get('Content-Type', ''), - 'content_length': response.headers.get('Content-Length', ''), - 'security_headers': { - 'csp': response.headers.get('Content-Security-Policy', ''), - 'hsts': response.headers.get('Strict-Transport-Security', ''), - 'x_frame_options': response.headers.get('X-Frame-Options', ''), - 'x_content_type_options': response.headers.get('X-Content-Type-Options', ''), - 'x_xss_protection': response.headers.get('X-XSS-Protection', ''), - 'referrer_policy': response.headers.get('Referrer-Policy', ''), - 'permissions_policy': response.headers.get('Permissions-Policy', ''), - } + "timestamp": time.time(), + "request_id": getattr(g, "request_id", None), + "correlation_id": getattr(g, "correlation_id", None), + "status_code": response.status_code, + "content_type": response.headers.get("Content-Type", ""), + "content_length": response.headers.get("Content-Length", ""), + "security_headers": { + "csp": response.headers.get("Content-Security-Policy", ""), + "hsts": response.headers.get("Strict-Transport-Security", ""), + "x_frame_options": response.headers.get("X-Frame-Options", ""), + "x_content_type_options": response.headers.get( + "X-Content-Type-Options", "" + ), + "x_xss_protection": response.headers.get("X-XSS-Protection", ""), + "referrer_policy": response.headers.get("Referrer-Policy", ""), + "permissions_policy": response.headers.get("Permissions-Policy", ""), + }, } logger.info(f"Response security: {security_info}") @@ -406,19 +517,29 @@ def get_security_stats(self) -> Dict: "enable_csp": self.config.enable_content_security_policy, "enable_hsts": self.config.enable_strict_transport_security, "enable_x_frame_options": self.config.enable_x_frame_options, - "enable_x_content_type_options": self.config.enable_x_content_type_options, + "enable_x_content_type_options": ( + self.config.enable_x_content_type_options + ), "enable_x_xss_protection": self.config.enable_x_xss_protection, "enable_referrer_policy": self.config.enable_referrer_policy, "enable_permissions_policy": self.config.enable_permissions_policy, - "enable_cross_origin_embedder_policy": self.config.enable_cross_origin_embedder_policy, - "enable_cross_origin_opener_policy": self.config.enable_cross_origin_opener_policy, - "enable_cross_origin_resource_policy": self.config.enable_cross_origin_resource_policy, + "enable_cross_origin_embedder_policy": ( + self.config.enable_cross_origin_embedder_policy + ), + "enable_cross_origin_opener_policy": ( + self.config.enable_cross_origin_opener_policy + ), + "enable_cross_origin_resource_policy": ( + self.config.enable_cross_origin_resource_policy + ), "enable_origin_agent_cluster": self.config.enable_origin_agent_cluster, "enable_request_id": self.config.enable_request_id, "enable_correlation_id": self.config.enable_correlation_id, "enable_enhanced_ua_analysis": self.config.enable_enhanced_ua_analysis, - "ua_suspicious_score_threshold": self.config.ua_suspicious_score_threshold, + "ua_suspicious_score_threshold": ( + self.config.ua_suspicious_score_threshold + ), "ua_blocking_enabled": self.config.ua_blocking_enabled, }, - "csp_nonce": self._csp_nonce + "csp_nonce": self._csp_nonce, } diff --git a/src/utils.py b/src/utils.py new file mode 100644 index 000000000..509717b8f --- /dev/null +++ b/src/utils.py @@ -0,0 +1,20 @@ +#!/usr/bin/env python3 +"""Utility functions for the SAMO-DL project.""" + +import torch +from typing import Union + + +def count_model_params(model: torch.nn.Module, only_trainable: bool = False) -> int: + """Count the number of parameters in a PyTorch model. + + Args: + model: PyTorch model to count parameters for + only_trainable: If True, only count trainable parameters + + Returns: + int: Number of parameters (total or trainable only) + """ + if only_trainable: + return sum(p.numel() for p in model.parameters() if p.requires_grad) + return sum(p.numel() for p in model.parameters())