""" Standalone RETVec+CNN Keras model training & held-out test evaluation script. Usage: python train_model.py python -m app.scripts.train_model """ import os import sys import random import zipfile import docx import pypdf from pptx import Presentation import numpy as np os.environ["TF_USE_LEGACY_KERAS"] = "1" os.environ["CUDA_VISIBLE_DEVICES"] = "-1" sys.stdout.reconfigure(encoding='utf-8') # Ensure project root is in sys.path BASE_DIR = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..")) if BASE_DIR not in sys.path: sys.path.insert(0, BASE_DIR) SEED = 42 random.seed(SEED) np.random.seed(SEED) import tensorflow as tf tf.random.set_seed(SEED) from app.ml.cnn.architecture import build_model, LABEL_NAMES from app.ml.training.data.encoding import encode_labels from app.ml.training.train import get_class_weights from app.ml.preprocessing.chunking import chunk_text HELDOUT_TEST_FILES = { "benign": [ "09_resmi_mektub_temiz.docx", "10_iclas_protokolu_temiz.docx", "Monthly Financial Expense Report.pdf", "11_ezamiyye_emri_temiz.docx", "19_sifaris_senedi_temiz.docx" ], "injection": [ "01_Aylıq_Fəaliyyət_Hesabatı.docx", "16_ezamiyye_xercleri_injection_gizli.docx", "19_sifaris_senedi_problem.docx", "23_bank_zemanet_mektubu_injection_context_hijack.docx", "24_qebul_tehvil_akti_injection.docx" ] } def extract_pptx(file_path: str) -> str: """Extract slide paragraph text and notes text from PPTX files using python-pptx.""" try: prs = Presentation(file_path) parts = [] for slide in prs.slides: for shape in slide.shapes: if shape.has_text_frame: for para in shape.text_frame.paragraphs: line = "".join(run.text for run in para.runs) if line.strip(): parts.append(line.strip()) if slide.has_notes_slide and slide.notes_slide.notes_text_frame: note = slide.notes_slide.notes_text_frame.text if note.strip(): parts.append(note.strip()) return "\n".join(parts) except Exception as e: print(f"Warning reading PPTX {file_path}: {e}") return "" def extract_text(file_path: str) -> str: """Extract raw text from supported document formats (.docx, .pptx, .pdf, .zip, .txt).""" ext = os.path.splitext(file_path)[1].lower() text = "" try: if ext == ".docx": doc = docx.Document(file_path) parts = [p.text for p in doc.paragraphs if p.text.strip()] for table in doc.tables: for row in table.rows: for cell in row.cells: if cell.text.strip(): parts.append(cell.text.strip()) text = "\n".join(parts) elif ext == ".pptx": text = extract_pptx(file_path) elif ext == ".pdf": reader = pypdf.PdfReader(file_path) parts = [] for i, page in enumerate(reader.pages): if i >= 20: break try: t = page.extract_text() if t: parts.append(t.strip()) except Exception: continue text = "\n".join(parts) elif ext == ".zip": parts = [] with zipfile.ZipFile(file_path, 'r') as z: for name in z.namelist(): if name.endswith('.docx'): tmp_path = os.path.join(os.path.dirname(file_path), "_tmp_extracted.docx") with open(tmp_path, "wb") as f_out: f_out.write(z.read(name)) sub_text = extract_text(tmp_path) if os.path.exists(tmp_path): os.remove(tmp_path) parts.append(sub_text) elif name.endswith('.pptx'): tmp_path = os.path.join(os.path.dirname(file_path), "_tmp_extracted.pptx") with open(tmp_path, "wb") as f_out: f_out.write(z.read(name)) sub_text = extract_text(tmp_path) if os.path.exists(tmp_path): os.remove(tmp_path) parts.append(sub_text) elif name.endswith('.txt'): parts.append(z.read(name).decode('utf-8', errors='ignore')) text = "\n".join(parts) elif ext == ".txt": with open(file_path, "r", encoding="utf-8", errors="ignore") as f: text = f.read() else: print(f"Skipping unsupported file extension {ext} for {file_path}") return "" except Exception as e: print(f"Warning reading {file_path}: {e}") return text.strip() def split_documents(doc_ids: list[str], val_ratio: float = 0.15, seed: int = 42) -> tuple[set[str], set[str]]: """Perform a document-level split of source document IDs into train and validation sets.""" rng = random.Random(seed) unique_ids = list(dict.fromkeys(doc_ids)) rng.shuffle(unique_ids) n_val = max(1, int(len(unique_ids) * val_ratio)) val_ids = set(unique_ids[:n_val]) train_ids = set(unique_ids[n_val:]) return train_ids, val_ids def load_real_dataset(raw_dir: str): all_chunks = [] # [(doc_id, text_chunk, label)] all_doc_ids = [] test_docs = [] # Define folder mapping: (folder_path, default_category) folders_to_scan = [ (os.path.join(raw_dir, "benign"), "benign"), (os.path.join(raw_dir, "injection"), "injection"), ] downloaded_dir = os.path.join(raw_dir, "downloaded") if os.path.exists(downloaded_dir): for root, dirs, files in os.walk(downloaded_dir): if files: folders_to_scan.append((root, "benign")) scanned_file_counts = {} # Load 10,200 PDF V4 Synthetic Dataset if dataset_V4.csv exists v4_csv_path = os.path.join(downloaded_dir, "dataset_V4.csv") if os.path.exists(v4_csv_path): try: import pandas as pd print(f"Loading 10,200 PDF V4 Synthetic Dataset samples from {v4_csv_path}...") df_v4 = pd.read_csv(v4_csv_path) v4_count = 0 for _, row in df_v4.iterrows(): doc_id = f"v4_{row['doc_id']}" extracted_text = str(row['extracted_text']) if pd.notna(row['extracted_text']) else "" if not extracted_text.strip(): continue is_inj = bool(row['is_injected']) lbl = "injection" if is_inj else "safe" v4_count += 1 lines = [l.strip() for l in extracted_text.split("\n") if l.strip()] for line in lines: words = line.split() if len(words) <= 60: all_chunks.append((doc_id, line, lbl)) all_doc_ids.append(doc_id) else: for c in chunk_text(line): all_chunks.append((doc_id, c, lbl)) all_doc_ids.append(doc_id) scanned_file_counts["dataset_V4.csv (10,200 PDFs)"] = v4_count except Exception as err: print(f"Warning loading dataset_V4.csv: {err}") for cat_dir, category in folders_to_scan: if not os.path.exists(cat_dir): continue heldout_list = HELDOUT_TEST_FILES.get(category, []) label_str = "safe" if category == "benign" else "injection" dir_key = os.path.relpath(cat_dir, raw_dir) scanned_file_counts[dir_key] = scanned_file_counts.get(dir_key, 0) for fname in os.listdir(cat_dir): fpath = os.path.join(cat_dir, fname) if not os.path.isfile(fpath): continue extracted = extract_text(fpath) if not extracted: continue scanned_file_counts[dir_key] += 1 doc_id = os.path.relpath(fpath, raw_dir) if fname in heldout_list: test_docs.append({ "filename": fname, "category": category, "expected_label": label_str, "text": extracted }) else: ext = os.path.splitext(fname)[1].lower() file_chunks = [] # Check for docx paragraph-level white font / hidden text docx_inj_lines = set() if ext == ".docx": try: doc = docx.Document(fpath) for p in doc.paragraphs: ptxt = p.text.strip() if not ptxt: continue is_p_white = False for r in p.runs: if r.font.color and r.font.color.rgb and str(r.font.color.rgb).upper() in ("FFFFFF", "FFF"): is_p_white = True break if r._r.rPr is not None and r._r.rPr.find(docx.oxml.ns.qn("w:vanish")) is not None: is_p_white = True break if is_p_white: docx_inj_lines.add(ptxt) except Exception: pass MANUAL_STEALTHY_LABELS = { "23_bank_zemanet_mektubu_injection_context_hijack.docx": "450,000 azn", "24_qebul_tehvil_akti_injection.docx": "administrative correction", "25_sigorta_polisi_injection.docx": "franchise override protocol", "26_emek_muqavilesi_inject.docx": "critical parser override", "27_vekaletname_inject.docx": "müvəqqəti hüquqi düzəliş", "29_bank_rekvizit_deyisikligi_injection.docx": "sistem təlimatı", "31_tecili_odenis_memo_inject.docx": "suni intelleke", "32_hosting_deyisiklik_bildirisi_inject.docx": "audit id-7741", "33_elave_is_sifarisi_inject.docx": "audit id-112", "34_distributor_endirim_teklifi_inject.docx": "audit id-267", "38_icare_muqavilesi_injection.docx": "tərəfdaşlıq ianəsi", "39_dasima_xidmeti_muqavilesi_inject.docx": "" } lines = [l.strip() for l in extracted.split("\n") if l.strip()] for line in lines: is_inj_line = False if category == "injection": low = line.lower() if fname in MANUAL_STEALTHY_LABELS: if MANUAL_STEALTHY_LABELS[fname] in low: is_inj_line = True else: low = line.lower() if line in docx_inj_lines or any(kw in low for kw in [ "prompt", "system", "yuxarida", "mene", "ignore", "override", "@", "//", "#", "||", "^^", "***", "&&", " 15: inj_chunks = [c for c in file_chunks if c[1] == "injection"] safe_chunks = [c for c in file_chunks if c[1] == "safe"] needed_safe = max(5, 15 - len(inj_chunks)) step = max(1, len(safe_chunks) // needed_safe) if safe_chunks else 1 file_chunks = inj_chunks + (safe_chunks[::step][:needed_safe] if safe_chunks else []) for text_chunk, lbl in file_chunks: all_chunks.append((doc_id, text_chunk, lbl)) all_doc_ids.append(doc_id) # Document-level split train_doc_ids, val_doc_ids = split_documents(all_doc_ids, val_ratio=0.15, seed=SEED) train_tuples = [c for c in all_chunks if c[0] in train_doc_ids] val_tuples = [c for c in all_chunks if c[0] in val_doc_ids] # Oversample injection training tuples so model learns injection patterns properly train_inj_tuples = [t for t in train_tuples if t[2] == "injection"] train_safe_tuples = [t for t in train_tuples if t[2] == "safe"] if train_inj_tuples and len(train_safe_tuples) > 0: multiplier = max(1, (len(train_safe_tuples) // 3) // len(train_inj_tuples)) train_inj_oversampled = train_inj_tuples * multiplier train_tuples = train_safe_tuples + train_inj_oversampled # Thorough random shuffling across all sources, classes, and languages rng = random.Random(SEED) rng.shuffle(train_tuples) rng.shuffle(val_tuples) train_texts = [t[1] for t in train_tuples] train_labels = [t[2] for t in train_tuples] val_texts = [t[1] for t in val_tuples] val_labels = [t[2] for t in val_tuples] print("Scanned files count per folder:") for folder_rel, count in scanned_file_counts.items(): print(f" - {folder_rel}: {count} valid documents") print(f"Document-level split: {len(train_doc_ids)} train docs ({len(train_texts)} chunks), {len(val_doc_ids)} val docs ({len(val_texts)} chunks)") return (train_texts, train_labels), (val_texts, val_labels), test_docs def main(): raw_dir = os.path.join(BASE_DIR, "data", "raw") print("Reading document dataset from data/raw...") (train_texts, train_labels), (val_texts, val_labels), test_docs = load_real_dataset(raw_dir) print(f"\n--- Dataset Loading Summary ---") print(f"Training text chunks extracted: {len(train_texts)}") print(f" - Safe (Benign) train chunks: {train_labels.count('safe')}") print(f" - Injection train chunks: {train_labels.count('injection')}") print(f"Validation text chunks extracted: {len(val_texts)}") print(f" - Safe (Benign) val chunks: {val_labels.count('safe')}") print(f" - Injection val chunks: {val_labels.count('injection')}") print(f"Held-out Test Files reserved: {len(test_docs)}") for td in test_docs: print(f" * [{td['category'].upper()}] {td['filename']} ({len(td['text'])} chars)") X_train = np.array([[t] for t in train_texts]) Y_train_label = encode_labels(train_labels) X_val = np.array([[t] for t in val_texts]) Y_val_label = encode_labels(val_labels) class_weights_dict = get_class_weights(Y_train_label) sample_weights_label = np.array([class_weights_dict[int(np.argmax(y))] for y in Y_train_label], dtype=np.float32) print("\nBuilding RETVec + CNN Keras Classification Model...") model = build_model(sequence_length=128) model.summary() print("\nStarting Keras Model Training (5 Epochs, batch_size=128, document-level validation)...", flush=True) history = model.fit( X_train, Y_train_label, epochs=5, batch_size=128, validation_data=(X_val, Y_val_label), sample_weight=sample_weights_label, verbose=1 ) models_dir = os.path.join(BASE_DIR, "data", "models") os.makedirs(models_dir, exist_ok=True) keras_model_path = os.path.join(models_dir, "retvec_cnn_model.keras") print(f"\nSaving trained model to .keras file at:\n {keras_model_path}") model.save(keras_model_path) cache_dir = os.path.join(BASE_DIR, "data", "cache") os.makedirs(cache_dir, exist_ok=True) model.save(os.path.join(cache_dir, "active_model.keras")) print("\n==========================================") print("HELD-OUT TEST FILES INFERENCE & EVALUATION") print("==========================================") correct_predictions = 0 test_results = [] for td in test_docs: raw_text = td["text"] lines = [l.strip() for l in raw_text.split("\n") if l.strip()] chunks = [] for line in lines: words = line.split() if len(words) <= 60: chunks.append(line) else: chunks.extend(chunk_text(line)) chunk_inputs = np.array([[c] for c in chunks]) preds = model.predict(chunk_inputs, verbose=0) label_preds = preds if isinstance(preds, np.ndarray) and preds.ndim == 2 else preds[0] worst_chunk_idx = label_preds[:, 2].argmax() max_injection_prob = float(label_preds[worst_chunk_idx, 2]) max_inj_line = chunks[worst_chunk_idx] if chunks else "" avg_probs = np.mean(label_preds, axis=0) HIGH_CONF_THRESHOLD = 0.85 CORROBORATION_THRESHOLD = 0.60 MIN_CORROBORATING_CHUNKS = 2 injection_probs = [float(p) for p in label_preds[:, 2]] predicted_label = "safe" high_conf = [p for p in injection_probs if p >= HIGH_CONF_THRESHOLD] if high_conf: predicted_label = "injection" else: corroborating = [p for p in injection_probs if p >= CORROBORATION_THRESHOLD] if len(corroborating) >= MIN_CORROBORATING_CHUNKS: predicted_label = "injection" is_correct = (predicted_label == td["expected_label"]) if is_correct: correct_predictions += 1 test_results.append({ "filename": td["filename"], "expected": td["expected_label"], "predicted": predicted_label, "is_correct": is_correct, "prob_safe": float(avg_probs[0]), "prob_suspicious": float(avg_probs[1]), "prob_injection": float(avg_probs[2]), "max_chunk_injection": float(max_injection_prob), "max_inj_snippet": max_inj_line[:60] }) status = "PASSED ✓" if is_correct else "FAILED ✗" print(f"File: {td['filename']}") print(f" Expected: {td['expected_label']} | Predicted: {predicted_label} [{status}]") print(f" Max Injection Prob: {max_injection_prob:.2%} | Snippet: {max_inj_line[:70]!r}\n") accuracy = (correct_predictions / len(test_docs)) * 100 if test_docs else 0.0 print(f"Final Held-Out Test Accuracy: {accuracy:.2f}% ({correct_predictions}/{len(test_docs)})") last_loss = float(history.history["loss"][-1]) if "history" in locals() and "loss" in history.history else 0.0 last_acc = float(history.history["accuracy"][-1]) if "history" in locals() and "accuracy" in history.history else 0.0 last_val = float(history.history["val_accuracy"][-1]) if "history" in locals() and "val_accuracy" in history.history else 0.0 prompt_local_push_confirmation( model=model, accuracy=accuracy, correct_count=correct_predictions, total_test_docs=len(test_docs), train_chunk_count=len(train_texts), val_chunk_count=len(val_texts), last_train_loss=last_loss, last_train_acc=last_acc, last_val_acc=last_val, test_results=test_results, ) def fetch_last_5_models_from_firestore(): init_firebase() db = get_firestore_db() if db is None: return [], 0 try: docs = db.collection("models").get() model_list = [] max_run_num = 0 for doc in docs: d = doc.to_dict() v_id = d.get("version") or doc.id if v_id.startswith("run-"): try: r_num = int(v_id.split("-")[1]) if r_num > max_run_num: max_run_num = r_num except ValueError: pass model_list.append(d) def sort_key(d): v = d.get("version", "") if v.startswith("run-"): try: return int(v.split("-")[1]) except ValueError: pass return 0 model_list.sort(key=sort_key) return model_list[-5:], max_run_num except Exception as e: print(f"Warning fetching models from Firestore: {e}") return [], 0 def prompt_local_push_confirmation(model, accuracy: float, correct_count: int, total_test_docs: int, train_chunk_count: int, val_chunk_count: int, last_train_loss: float, last_train_acc: float, last_val_acc: float, test_results: list): import subprocess import asyncio from datetime import datetime, timezone from app.core.firebase import init_firebase, get_firestore_db from app.ml.serving.registry import save_model_version last_5, max_run_num = fetch_last_5_models_from_firestore() inj_docs = [t for t in test_results if t["expected"] == "injection"] inj_correct = [t for t in inj_docs if t["is_correct"]] test_recall = (len(inj_correct) / len(inj_docs) * 100.0) if inj_docs else 100.0 if last_5: print("\nLast 5 registered versions:") for m in last_5: v_str = m.get("version", "unknown") metrics_m = m.get("metrics", {}) test_acc_m = metrics_m.get("test_acc", 0.0) * 100.0 if isinstance(metrics_m.get("test_acc"), (int, float)) else 0.0 recall_m = metrics_m.get("recall", 0.0) * 100.0 if isinstance(metrics_m.get("recall"), (int, float)) else 0.0 status_tag = " (currently active)" if m.get("status") == "active" else "" print(f" {v_str:<8} test acc {test_acc_m:.2f}% recall {recall_m:.0f}%{status_tag}") print(f"\nThis run: test acc {accuracy:.2f}% recall {test_recall:.0f}%\n") answer = input("Upload this model to Firebase as a new candidate version? (y/n): ").strip().lower() if answer == "y": next_run_num = max_run_num + 1 if max_run_num > 0 else 12 new_version_id = f"run-{next_run_num:02d}" try: res = subprocess.run(["git", "rev-parse", "HEAD"], capture_output=True, text=True, check=True) source_commit = res.stdout.strip() except Exception: source_commit = "unknown" today_str = datetime.now(timezone.utc).strftime("%Y-%m-%d") desc = ( f"Trained {today_str}. " f"Dataset: {train_chunk_count} train chunks + {val_chunk_count} val chunks. " f"Held-out test: {accuracy:.2f}% accuracy ({correct_count}/{total_test_docs}), " f"{test_recall:.0f}% injection recall." ) metrics_payload = { "train_loss": float(last_train_loss), "train_acc": float(last_train_acc), "val_acc": float(last_val_acc), "test_acc": float(accuracy / 100.0), "recall": float(test_recall / 100.0), "correct_test": f"{correct_count}/{total_test_docs}", } asyncio.run( save_model_version( model=model, metrics=metrics_payload, version=new_version_id, status="candidate", source_commit=source_commit, description=desc, ) ) print(f"Uploaded as candidate version '{new_version_id}'. Use POST /model/change-version/{new_version_id} to make it active.") else: print("Skipped. Model saved locally only at data/models/retvec_cnn_model.keras.") if __name__ == "__main__": main()