From 61abf7f2d20c424389d9b132fbfb5d2418e4cef6 Mon Sep 17 00:00:00 2001 From: pavsoss Date: Sat, 1 Aug 2026 14:14:34 -0400 Subject: [PATCH] Record the preparation contract in model provenance Artifacts carried no record of how text must be prepared for them, so a future change to preparation could silently desynchronize from a model trained under the old rules. Training now emits the provenance sidecar the registry has always looked for, and a mismatch is reported when the artifacts are loaded. --- backend/api.py | 93 ++++++++++++----- backend/bulk_predict.py | 6 +- backend/email_connectors/email_scanner.py | 55 +++++----- backend/model_registry.py | 20 +++- backend/retrain.py | 71 +++++++++++-- backend/tests/test_inference_parity.py | 96 ++++++++++++++++++ backend/tests/test_preparation_provenance.py | 101 +++++++++++++++++++ backend/tests/test_retrain.py | 2 +- backend/tests/test_text_preparation.py | 51 ++++++++++ backend/text_preparation.py | 47 +++++++++ 10 files changed, 478 insertions(+), 64 deletions(-) create mode 100644 backend/tests/test_inference_parity.py create mode 100644 backend/tests/test_preparation_provenance.py create mode 100644 backend/tests/test_text_preparation.py create mode 100644 backend/text_preparation.py diff --git a/backend/api.py b/backend/api.py index 248dee0f..80e2cf71 100644 --- a/backend/api.py +++ b/backend/api.py @@ -43,6 +43,7 @@ configure_rate_limiting, rate_limit) import requests from routes.analytics import analytics_bp, record_scan +from text_preparation import PREPARATION_VERSION, prepare_text from utils.spamSeverity import calculate_spam_severity @@ -583,11 +584,23 @@ def handle_internal_error(e): def _build_model_metadata(): """Fingerprint the currently-on-disk classifier artifacts (issue #1007).""" - return model_registry.build_metadata( + metadata = model_registry.build_metadata( model_path=str(MODEL_PATH), vectorizer_path=str(VECTORIZER_PATH), label_encoder_path=str(LABEL_ENCODER_PATH), ) + # Surfaced loudly rather than fatally: a contract mismatch degrades accuracy + # but the model still answers, and refusing to boot would take the API down + # over a metadata disagreement an operator may already be mid-way through + # resolving with a retrain. + if not metadata.preparation_matches(PREPARATION_VERSION): + app.logger.warning( + "served model was trained under text-preparation contract %s but %s " + "is in force; retrain to restore train-serve parity", + metadata.preparation_version, + PREPARATION_VERSION, + ) + return metadata def _load_serving_objects(): @@ -826,6 +839,7 @@ def heuristic_url_is_malicious(url): tld = host.rsplit(".", 1)[-1] if "." in host else "" return tld in SUSPICIOUS_TLDS + MAX_MESSAGE_LENGTH = int(os.getenv("MAX_MESSAGE_LENGTH", 5000)) MAX_MESSAGE_LENGTH = settings.max_message_length @@ -1214,23 +1228,32 @@ def predict(): if text is None or (isinstance(text, str) and not text.strip()): with open(LOG_FILE, "a") as f: - f.write(f"WARNING: No text provided at {__import__('datetime').datetime.now()}\n") + f.write( + f"WARNING: No text provided at {__import__('datetime').datetime.now()}\n" + ) return jsonify({"error": "No text provided"}), 400 if not isinstance(text, str): - return jsonify({ - "error": f"'text' must be a string, got {type(text).__name__}" - }), 400 - + return ( + jsonify( + {"error": f"'text' must be a string, got {type(text).__name__}"} + ), + 400, + ) # Maximum-length validation before any vectorization/inference work. if len(text) > MAX_MESSAGE_LENGTH: - return jsonify({ - "error": ( - f"'text' exceeds maximum length of {MAX_MESSAGE_LENGTH} " - f"characters (got {len(text)})" - ) - }), 400 + return ( + jsonify( + { + "error": ( + f"'text' exceeds maximum length of {MAX_MESSAGE_LENGTH} " + f"characters (got {len(text)})" + ) + } + ), + 400, + ) # Read the live serving objects through the shared state so a # POST /reload-model hot-swap is picked up here without a restart # (#973). One snapshot per request keeps the model, vectorizer and @@ -1250,7 +1273,7 @@ def predict(): # request also repopulates the cache). cache_options = {"type": input_type} cache_key = predict_cache.make_cache_key( - normalizer.normalize(text), serving.version, cache_options + prepare_text(text), serving.version, cache_options ) cache_bypass = _cache_bypass_requested() if not cache_bypass: @@ -1294,7 +1317,9 @@ def predict(): if final_output == "safe" and heuristic_url_is_malicious(text): final_output = "malicious" else: - text_vector = serving.vectorizer.transform([text]) + # Prepared after translation, so the string handed to the vectorizer + # is the one the model was trained on regardless of source language. + text_vector = serving.vectorizer.transform([prepare_text(text)]) prediction = serving.model.predict(text_vector) final_output = serving.label_encoder.inverse_transform(prediction)[0] @@ -2283,14 +2308,16 @@ def imap_status(): if not conn_row: return jsonify({"connected": False}) - return jsonify({ - "connected": True, - "host": conn_row["host"], - "imap_username": conn_row["imap_username"], - "scan_interval_minutes": conn_row["scan_interval_minutes"], - "consent_given_at": conn_row["consent_given_at"], - "last_scan_at": conn_row["last_scan_at"], - }) + return jsonify( + { + "connected": True, + "host": conn_row["host"], + "imap_username": conn_row["imap_username"], + "scan_interval_minutes": conn_row["scan_interval_minutes"], + "consent_given_at": conn_row["consent_given_at"], + "last_scan_at": conn_row["last_scan_at"], + } + ) @app.route("/imap/schedule", methods=["PUT"]) @@ -2304,14 +2331,26 @@ def imap_schedule(): scan_interval_minutes = data.get("scan_interval_minutes") if scan_interval_minutes not in imap_store.ALLOWED_INTERVALS: - return jsonify({"error": f"scan_interval_minutes must be one of {imap_store.ALLOWED_INTERVALS}"}), 400 + return ( + jsonify( + { + "error": f"scan_interval_minutes must be one of {imap_store.ALLOWED_INTERVALS}" + } + ), + 400, + ) if not imap_store.get_connection(username): return jsonify({"error": "No connected inbox found for this account"}), 404 imap_store.update_schedule(username, scan_interval_minutes) _schedule_user_job(username, scan_interval_minutes) - return jsonify({"message": "Scan schedule updated", "scan_interval_minutes": scan_interval_minutes}) + return jsonify( + { + "message": "Scan schedule updated", + "scan_interval_minutes": scan_interval_minutes, + } + ) @app.route("/imap/disconnect", methods=["POST"]) @@ -2349,7 +2388,11 @@ def imap_scan_now(): try: emails = imap_connector.fetch_imap_emails( - conn_row["host"], conn_row["port"], conn_row["imap_username"], password, limit=50 + conn_row["host"], + conn_row["port"], + conn_row["imap_username"], + password, + limit=50, ) scan_results = scan_emails_with_model(emails) imap_store.save_scan_results(username, scan_results["emails"]) diff --git a/backend/bulk_predict.py b/backend/bulk_predict.py index a4c35d54..89b1a0e3 100644 --- a/backend/bulk_predict.py +++ b/backend/bulk_predict.py @@ -187,7 +187,11 @@ def _predict_batch(messages, snapshot): Raises on any transform/predict failure; callers isolate failures by retrying the offending batch one row at a time. """ - text_vectors = snapshot.vectorizer.transform(messages) + # Scored on the prepared form for parity with training and /predict; the + # rows echoed back below stay verbatim so a result still matches its input. + text_vectors = snapshot.vectorizer.transform( + [prepare_text(msg) for msg in messages] + ) predictions = snapshot.model.predict(text_vectors) final_outputs = snapshot.label_encoder.inverse_transform(predictions) decisions = snapshot.model.decision_function(text_vectors) diff --git a/backend/email_connectors/email_scanner.py b/backend/email_connectors/email_scanner.py index 72ded053..17f51f07 100644 --- a/backend/email_connectors/email_scanner.py +++ b/backend/email_connectors/email_scanner.py @@ -1,36 +1,41 @@ -from flask import current_app +from email_header_analyzer import analyze_headers +from flask import current_app import numpy as np +from pathlib import Path import sys -from pathlib import Path -from email_header_analyzer import analyze_headers - +from text_preparation import prepare_text + try: # Import standard headers analyzer if available - + sys.path.insert(0, str(Path(__file__).resolve().parents[1])) except ImportError: analyze_headers = None + def scan_emails_with_model(emails): """Classifies a list of fetched emails using the active machine learning model. - + Optionally appends header analysis results (risk_score, trust_level) if headers exist. """ vectorizer = getattr(current_app, "vectorizer", None) model = getattr(current_app, "model", None) label_encoder = getattr(current_app, "label_encoder", None) - + if not model or not vectorizer or not label_encoder: - raise ValueError("ML model dependencies are not loaded in the Flask application.") - + raise ValueError( + "ML model dependencies are not loaded in the Flask application." + ) + scanned_emails = [] spam_count = 0 safe_count = 0 - - # Extract email subjects and bodies for batch vectorization - texts = [f"{e['subject']}. {e['body']}" for e in emails] + + # Extract email subjects and bodies for batch vectorization, prepared with + # the shared contract so a scanned inbox agrees with /predict on the same + # message rather than scoring obfuscated text the model never saw. + texts = [prepare_text(f"{e['subject']}. {e['body']}") for e in emails] if texts: - text_vectors = vectorizer.transform(texts) predictions = model.predict(text_vectors) final_outputs = label_encoder.inverse_transform(predictions) @@ -38,29 +43,28 @@ def scan_emails_with_model(emails): else: final_outputs = [] decisions = [] - + for i, (e, pred) in enumerate(zip(emails, final_outputs)): pred_str = str(pred) # Classify as spam if not explicitly 'ham' or 'safe' is_spam = pred_str.lower() not in ("ham", "safe") - + if is_spam: spam_count += 1 else: safe_count += 1 - - + dec_score = float(np.max(np.abs(decisions[i]))) prob = 1.0 / (1.0 + np.exp(-dec_score)) conf_score = round(prob * 100, 2) - + if conf_score >= 80: conf_level = "high" elif conf_score >= 60: conf_level = "medium" else: conf_level = "low" - + email_result = { "id": e.get("id"), "subject": e.get("subject", "No Subject"), @@ -71,7 +75,7 @@ def scan_emails_with_model(emails): "confidence": round(conf_score / 100.0, 4), "confidence_score": conf_score, "decision_score": dec_score, - "confidence_level": conf_level + "confidence_level": conf_level, } # Phishing integration preparation (optional header analysis) has_header_risk = False @@ -83,7 +87,7 @@ def scan_emails_with_model(emails): has_header_risk = True except Exception: pass - + # Fallback: check sender's domain metadata if header analysis wasn't performed/successful if not has_header_risk and e.get("sender"): try: @@ -92,10 +96,11 @@ def scan_emails_with_model(emails): sender_domain = sender_val.split("@")[-1].lower().strip(" >") if sender_domain: import domain_checker + domain_analysis = domain_checker.analyze_domain(sender_domain) risk_val = domain_analysis.get("risk_score", 0) email_result["risk_score"] = risk_val - + if risk_val <= 20: email_result["trust_level"] = "Trusted" elif risk_val <= 60: @@ -104,12 +109,12 @@ def scan_emails_with_model(emails): email_result["trust_level"] = "High Risk" except Exception: pass - + scanned_emails.append(email_result) - + return { "total_scanned": len(emails), "spam_count": spam_count, "safe_count": safe_count, - "emails": scanned_emails + "emails": scanned_emails, } diff --git a/backend/model_registry.py b/backend/model_registry.py index 1b39b796..edda77e7 100644 --- a/backend/model_registry.py +++ b/backend/model_registry.py @@ -8,8 +8,9 @@ This module fingerprints those artifacts. :func:`build_metadata` reads each ``.pkl`` and captures its SHA-256, size and mtime, and -- when a -``model_card.json`` sits next to the model -- folds in the human-authored -provenance fields (``trained_at``, ``metrics``, ``labels``). The immutable +``model_card.json`` sits next to the model -- folds in the provenance fields +``retrain.py`` records there (``trained_at``, ``metrics``, ``labels`` and the +``preparation_version`` the artifacts were trained under). The immutable :class:`ModelMetadata` it returns is stored alongside the serving objects in ``serving_state`` and surfaced at ``GET /model-info``; its :attr:`ModelMetadata.short_checksum` tags predictions and reload audit logs. @@ -87,6 +88,19 @@ class ModelMetadata: trained_at: str | None = None metrics: dict | None = None labels: list | None = None + preparation_version: str | None = None + + def preparation_matches(self, serving_version: str) -> bool: + """Whether these artifacts were trained under ``serving_version``. + + An unrecorded version (``None``) counts as a match: artifacts predating + the model card carry no claim about their preparation, and refusing to + serve them would break existing deployments over missing metadata rather + than over a known conflict. + """ + if self.preparation_version is None: + return True + return self.preparation_version == serving_version @property def short_checksum(self) -> str: @@ -111,6 +125,7 @@ def to_dict(self) -> dict: "trained_at": self.trained_at, "metrics": self.metrics, "labels": self.labels, + "preparation_version": self.preparation_version, } @@ -142,6 +157,7 @@ def build_metadata( trained_at=card.get("trained_at"), metrics=card.get("metrics"), labels=card.get("labels"), + preparation_version=card.get("preparation_version"), ) diff --git a/backend/retrain.py b/backend/retrain.py index c81d3730..e8cc7de5 100644 --- a/backend/retrain.py +++ b/backend/retrain.py @@ -12,8 +12,8 @@ 1. Loads the original training dataset (DATASET_PATH env var, default: dataset.csv) 2. Loads feedback_store.csv (the corrected labels submitted via /feedback) 3. Merges them into one training set (feedback's `correct_label` becomes the label) - 4. Encodes labels once with a single LabelEncoder and normalizes text with the - same normalizer api.py uses at inference time. + 4. Encodes labels once with a single LabelEncoder and prepares text with the + shared contract every inference path applies. 5. Fits the vectorizer + LinearSVC ONCE on a held-out train split to report an honest accuracy, then refits ONCE on the full combined data for the artifacts actually written to disk. @@ -21,7 +21,10 @@ - linear_svm_model.pkl - tfidf_vectorizer.pkl - label_encoder.pkl - 7. Triggers a live model reload only AFTER a successful save. + 7. Writes model_card.json alongside them, recording when the model was + trained, how it scored, its label set, and the text-preparation contract + it was trained under. + 8. Triggers a live model reload only AFTER a successful save. Run this from the backend/ directory: cd backend @@ -46,7 +49,7 @@ from sklearn.preprocessing import LabelEncoder from sklearn.svm import LinearSVC -from utils.text_normalizer import normalizer +from text_preparation import prepare_text VALID_LABELS = {"ham", "spam", "smishing"} @@ -170,13 +173,13 @@ def train( ): """Deterministic training pipeline. - Text is normalized with the same normalizer api.py applies at inference, so - the vectorizer vocabulary matches what serving will see. Labels are encoded - ONCE and the encoded integers are used for every fit -- no raw string labels - leak into any model. The held-out fit and the production fit each happen - exactly once. + Text goes through the shared preparation contract that every inference path + also applies, so the vectorizer vocabulary matches what serving will see. + Labels are encoded ONCE and the encoded integers are used for every fit -- no + raw string labels leak into any model. The held-out fit and the production + fit each happen exactly once. """ - normalized = combined["text"].apply(normalizer.normalize) + normalized = combined["text"].apply(prepare_text) label_encoder = LabelEncoder() y = label_encoder.fit_transform(combined["label"]) @@ -227,6 +230,36 @@ def save_artifacts( print(f"Saved: {label_encoder_path}") +def write_model_card( + result, + *, + model_path=MODEL_PATH, + card_path=None, +): + """Emit the provenance sidecar the registry reads for ``GET /model-info``. + + Records the preparation contract the artifacts were trained under so a later + mismatch between trained and serving text handling is detectable instead of + silently degrading predictions. Written after the artifacts, so a card can + never describe a model that failed to persist. + """ + path = card_path or os.path.join( + os.path.dirname(os.path.abspath(model_path)), MODEL_CARD_FILENAME + ) + card = { + "trained_at": datetime.now(timezone.utc).isoformat(), + "metrics": { + "holdout_accuracy": round(float(result.holdout.accuracy), 4), + "training_rows": result.n_rows, + }, + "labels": [str(label) for label in result.label_encoder.classes_], + "preparation_version": PREPARATION_VERSION, + } + _atomic_write_json(card, path) + print(f"Saved: {path}") + return path + + def backup_existing_files(): """Copy existing .pkl files to a timestamped backup folder before overwriting.""" timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") @@ -316,6 +349,7 @@ def main(argv=None): backup_existing_files() save_artifacts(result) + write_model_card(result) print("\nRetraining complete. Triggering live model reload...") trigger_model_reload() @@ -356,6 +390,23 @@ def _evaluate_holdout( ) +def _atomic_write_json(payload, path): + """Write JSON through a temp file in the destination directory, then replace, + so a reader never sees a half-written card.""" + directory = os.path.dirname(os.path.abspath(path)) + os.makedirs(directory, exist_ok=True) + fd, tmp_path = tempfile.mkstemp(dir=directory, suffix=".tmp") + os.close(fd) + try: + with open(tmp_path, "w", encoding="utf-8") as fh: + json.dump(payload, fh, indent=2, sort_keys=True) + os.replace(tmp_path, path) + except BaseException: + if os.path.exists(tmp_path): + os.remove(tmp_path) + raise + + def _atomic_joblib_dump(obj, path): """joblib.dump to a temp file in the destination directory, then os.replace so readers never observe a half-written artifact.""" diff --git a/backend/tests/test_inference_parity.py b/backend/tests/test_inference_parity.py new file mode 100644 index 00000000..8161e23e --- /dev/null +++ b/backend/tests/test_inference_parity.py @@ -0,0 +1,96 @@ +"""Cross-path parity for the text-preparation contract (issue #1037). + +Bulk scoring and mailbox scanning must hand the vectorizer the same canonical +string ``/predict`` does, otherwise the same message is classified differently +depending on how it arrived. These tests capture what actually reaches +``transform`` using recording doubles, so they assert the contract without +loading real model artifacts. +""" + +from types import SimpleNamespace + +import numpy as np +import pytest + +import bulk_predict +from text_preparation import prepare_text + +# Subject/body pairs whose obfuscated and plain spellings must reduce alike. +OBFUSCATED = "F r e e\u200b \u0440rize" +PLAIN = "Free prize" + + +class RecordingVectorizer: + """Captures the texts handed to ``transform`` and returns a dummy matrix.""" + + def __init__(self): + self.seen = [] + + def transform(self, texts): + self.seen = list(texts) + return np.zeros((len(self.seen), 1)) + + +class ConstantModel: + def predict(self, vectors): + return np.zeros(vectors.shape[0], dtype=int) + + def decision_function(self, vectors): + return np.zeros((vectors.shape[0], 1)) + + +class PassthroughEncoder: + def inverse_transform(self, predictions): + return ["ham"] * len(predictions) + + +@pytest.fixture +def snapshot(): + return SimpleNamespace( + vectorizer=RecordingVectorizer(), + model=ConstantModel(), + label_encoder=PassthroughEncoder(), + version=1, + ) + + +class TestBulkPrediction: + def test_rows_are_scored_in_prepared_form(self, snapshot): + bulk_predict._predict_batch([OBFUSCATED], snapshot) + + assert snapshot.vectorizer.seen == [prepare_text(OBFUSCATED)] + + def test_obfuscated_and_plain_rows_score_identically(self, snapshot): + bulk_predict._predict_batch([OBFUSCATED, PLAIN], snapshot) + + seen = snapshot.vectorizer.seen + assert seen[0] == seen[1] + + def test_returned_row_echoes_the_original_text(self, snapshot): + results = bulk_predict._predict_batch([OBFUSCATED], snapshot) + + # The caller uploaded this row and must recognise it in the response, + # even though a different string was scored. + assert results[0]["message"] == OBFUSCATED + + +class TestMailboxScanning: + def test_scanned_emails_are_prepared_before_scoring(self, snapshot, monkeypatch): + from email_connectors import email_scanner + + monkeypatch.setattr( + email_scanner, + "current_app", + SimpleNamespace( + vectorizer=snapshot.vectorizer, + model=snapshot.model, + label_encoder=snapshot.label_encoder, + ), + ) + monkeypatch.setattr(email_scanner, "analyze_headers", None) + + email_scanner.scan_emails_with_model( + [{"subject": "F r e e", "body": "\u0440rize inside"}] + ) + + assert snapshot.vectorizer.seen == [prepare_text("F r e e. \u0440rize inside")] diff --git a/backend/tests/test_preparation_provenance.py b/backend/tests/test_preparation_provenance.py new file mode 100644 index 00000000..9eeef95a --- /dev/null +++ b/backend/tests/test_preparation_provenance.py @@ -0,0 +1,101 @@ +"""Preparation contract recorded in, and read back from, model provenance (#1037). + +Covers both ends of the sidecar: what ``retrain.py`` writes after a successful +save, and what ``model_registry`` reports for the artifacts being served. +""" + +import json +from types import SimpleNamespace + +import model_registry +import retrain +from text_preparation import PREPARATION_VERSION + + +def _artifact(name): + return model_registry.ArtifactInfo( + path=name, sha256="ab" * 32, size_bytes=1, mtime=0.0 + ) + + +def _metadata(**overrides): + fields = { + "model": _artifact("m.pkl"), + "vectorizer": _artifact("v.pkl"), + "label_encoder": _artifact("l.pkl"), + } + fields.update(overrides) + return model_registry.ModelMetadata(**fields) + + +class TestPreparationMatching: + def test_same_version_matches(self): + assert _metadata(preparation_version="1").preparation_matches("1") + + def test_different_version_does_not_match(self): + assert not _metadata(preparation_version="1").preparation_matches("2") + + def test_unrecorded_version_is_treated_as_compatible(self): + # Artifacts predating the card make no claim; they must stay servable. + assert _metadata().preparation_matches("2") + + def test_version_is_reported_in_the_payload(self): + assert ( + _metadata(preparation_version="1").to_dict()["preparation_version"] == "1" + ) + + +class TestCardIsReadBack: + def test_registry_surfaces_the_recorded_version(self, tmp_path): + for name in ("m.pkl", "v.pkl", "l.pkl"): + (tmp_path / name).write_bytes(b"x") + (tmp_path / model_registry.MODEL_CARD_FILENAME).write_text( + json.dumps({"preparation_version": "7"}) + ) + + metadata = model_registry.build_metadata( + model_path=str(tmp_path / "m.pkl"), + vectorizer_path=str(tmp_path / "v.pkl"), + label_encoder_path=str(tmp_path / "l.pkl"), + ) + + assert metadata.preparation_version == "7" + + +class TestCardEmission: + def test_training_records_the_contract_in_force(self, tmp_path): + result = SimpleNamespace( + holdout=SimpleNamespace(accuracy=0.9375), + label_encoder=SimpleNamespace(classes_=["ham", "spam"]), + n_rows=120, + ) + + path = retrain.write_model_card( + result, card_path=str(tmp_path / model_registry.MODEL_CARD_FILENAME) + ) + card = json.loads(open(path, encoding="utf-8").read()) + + assert card["preparation_version"] == PREPARATION_VERSION + assert card["metrics"]["holdout_accuracy"] == 0.9375 + assert card["metrics"]["training_rows"] == 120 + assert card["labels"] == ["ham", "spam"] + assert card["trained_at"] + + def test_card_is_readable_by_the_registry(self, tmp_path): + result = SimpleNamespace( + holdout=SimpleNamespace(accuracy=1.0), + label_encoder=SimpleNamespace(classes_=["ham"]), + n_rows=10, + ) + for name in ("m.pkl", "v.pkl", "l.pkl"): + (tmp_path / name).write_bytes(b"x") + + retrain.write_model_card(result, model_path=str(tmp_path / "m.pkl")) + metadata = model_registry.build_metadata( + model_path=str(tmp_path / "m.pkl"), + vectorizer_path=str(tmp_path / "v.pkl"), + label_encoder_path=str(tmp_path / "l.pkl"), + ) + + assert metadata.preparation_version == PREPARATION_VERSION + assert metadata.preparation_matches(PREPARATION_VERSION) diff --git a/backend/tests/test_retrain.py b/backend/tests/test_retrain.py index 55de778a..3705193f 100644 --- a/backend/tests/test_retrain.py +++ b/backend/tests/test_retrain.py @@ -124,7 +124,7 @@ def test_labels_are_encoded_consistently_everywhere(trained): assert set(trained.label_encoder.classes_) == retrain.VALID_LABELS sample = trained.vectorizer.transform( - [retrain.normalizer.normalize("free prize claim now")] + [retrain.prepare_text("free prize claim now")] ) encoded_pred = trained.model.predict(sample) decoded = trained.label_encoder.inverse_transform(encoded_pred)[0] diff --git a/backend/tests/test_text_preparation.py b/backend/tests/test_text_preparation.py new file mode 100644 index 00000000..870d0525 --- /dev/null +++ b/backend/tests/test_text_preparation.py @@ -0,0 +1,51 @@ +"""Parity coverage for the shared text-preparation contract (issue #1037). + +The value of the contract is that training and serving cannot disagree, so these +tests assert the property that matters -- obfuscated input reduces to the same +canonical string as its plain equivalent -- rather than pinning the normalizer's +internal steps, which belong to its own tests. +""" + +import text_preparation +from text_preparation import prepare_text + + +class TestCanonicalForm: + def test_zero_width_characters_are_stripped(self): + assert prepare_text("Free\u200b Prize\u200d") == prepare_text("Free Prize") + + def test_cyrillic_homoglyphs_fold_to_latin(self): + # "claim" spelled with Cyrillic es, a and i. + assert prepare_text("\u0441l\u0430\u0456m") == prepare_text("claim") + + def test_spaced_out_words_are_rejoined(self): + assert prepare_text("f r e e money") == prepare_text("free money") + + def test_repeated_whitespace_collapses(self): + assert prepare_text("win a prize") == prepare_text("win a prize") + + def test_already_canonical_text_is_unchanged(self): + assert prepare_text("claim your free prize") == "claim your free prize" + + def test_preparation_is_idempotent(self): + once = prepare_text("F r e e\u200b m\u043en\u0435y") + assert prepare_text(once) == once + + +class TestNonStringInput: + def test_none_passes_through(self): + assert prepare_text(None) is None + + def test_empty_string_passes_through(self): + assert prepare_text("") == "" + + def test_non_string_passes_through(self): + assert prepare_text(42) == 42 + + +class TestTrainServeParity: + def test_training_and_inference_share_one_entry_point(self): + """Both regimes must import the same callable, not two copies of it.""" + import retrain + + assert retrain.prepare_text is text_preparation.prepare_text diff --git a/backend/text_preparation.py b/backend/text_preparation.py new file mode 100644 index 00000000..9df8cb9f --- /dev/null +++ b/backend/text_preparation.py @@ -0,0 +1,47 @@ +"""The single text-preparation contract shared by training and inference. + +``retrain.py`` fits the TF-IDF vocabulary on text that has been run through +:data:`~utils.text_normalizer.normalizer`, so anything that reaches +``vectorizer.transform`` at serving time must be prepared the same way or the +model is scoring a different alphabet than the one it learned. Homoglyph +substitutions, zero-width joiners and spaced-out words -- exactly the evasions +the normalizer exists to undo -- otherwise survive into the vectorizer and fall +out of vocabulary. + +Every producer of model input calls :func:`prepare_text`; nothing calls the +normalizer directly. Routing both regimes through one function is what makes the +parity checkable rather than a convention that drifts. + +>>> prepare_text("Free\\u200b Prize") +'Free Prize' +>>> prepare_text("\\u0441laim now") +'claim now' +>>> prepare_text("f r e e money") +'free money' +>>> prepare_text("") +'' +>>> prepare_text(None) is None +True +""" + +from utils.text_normalizer import normalizer + +__all__ = ["PREPARATION_VERSION", "prepare_text"] + +# Identifies the canonical form :func:`prepare_text` produces. Bump it whenever a +# change to the preparation steps alters the output for any input: artifacts +# trained under an older version learned a different vocabulary, and serving them +# under the newer one silently reintroduces the train-serve skew this contract +# exists to prevent. Recorded in the model card at training time and compared +# against the served artifacts when they are loaded. +PREPARATION_VERSION = "1" + + +def prepare_text(text): + """Return ``text`` in the canonical form the model was trained on. + + Non-string input is handed back untouched: callers upstream of validation + (bulk rows, mailbox payloads) can pass ``None`` or a stray numeric cell, and + a preparation step is the wrong place to decide that is an error. + """ + return normalizer.normalize(text)