diff --git a/backend/api.py b/backend/api.py index 80e2cf7..01d11bb 100644 --- a/backend/api.py +++ b/backend/api.py @@ -562,6 +562,10 @@ def handle_internal_error(e): model = joblib.load(MODEL_PATH) vectorizer = joblib.load(VECTORIZER_PATH) label_encoder = joblib.load(LABEL_ENCODER_PATH) +# Loaded here rather than further down so the URL pair can be installed in the +# serving state below and picked up by a hot reload like everything else. +url_model = joblib.load(URL_MODEL_PATH) +url_vectorizer = joblib.load(URL_VECTORIZER_PATH) from xai_service import XAIService @@ -619,6 +623,8 @@ def _load_serving_objects(): "label_encoder": fresh_label_encoder, "xai_service": fresh_xai_service, "metadata": _build_model_metadata(), + "url_model": joblib.load(URL_MODEL_PATH), + "url_vectorizer": joblib.load(URL_VECTORIZER_PATH), } @@ -629,6 +635,8 @@ def _load_serving_objects(): xai_service=xai_service, loader=_load_serving_objects, metadata=_build_model_metadata(), + url_model=url_model, + url_vectorizer=url_vectorizer, ) @@ -788,9 +796,6 @@ def get_word_of_the_day_data(): register_reload_endpoint(app) -url_model = joblib.load(URL_MODEL_PATH) -url_vectorizer = joblib.load(URL_VECTORIZER_PATH) - # All models loaded successfully; surface readiness as a scrapeable gauge (#984). metrics.set_model_loaded(True) @@ -1311,8 +1316,8 @@ def predict(): domain_analysis = analyze_text(text) if input_type == "url": - text_vector = url_vectorizer.transform([text]) - prediction = url_model.predict(text_vector) + text_vector = serving.url_vectorizer.transform([text]) + prediction = serving.url_model.predict(text_vector) final_output = URL_LABELS.get(int(prediction[0]), "unknown") if final_output == "safe" and heuristic_url_is_malicious(text): final_output = "malicious" @@ -1326,7 +1331,7 @@ def predict(): confidence_score = 95.0 decision_score = None try: - active_model = url_model if input_type == "url" else serving.model + active_model = serving.url_model if input_type == "url" else serving.model if hasattr(active_model, "predict_proba"): proba = active_model.predict_proba(text_vector) confidence_score = round(float(max(proba[0])) * 100, 2) diff --git a/backend/email_connectors/email_scanner.py b/backend/email_connectors/email_scanner.py index 17f51f0..5d09906 100644 --- a/backend/email_connectors/email_scanner.py +++ b/backend/email_connectors/email_scanner.py @@ -1,7 +1,7 @@ from email_header_analyzer import analyze_headers -from flask import current_app import numpy as np from pathlib import Path +import serving_state import sys from text_preparation import prepare_text @@ -18,14 +18,16 @@ def scan_emails_with_model(emails): 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) + # Read through the shared serving state rather than objects pinned on the + # application at startup, so a /reload-model hot-swap reaches inbox scans too + # instead of leaving them on the model the process booted with. + snapshot = serving_state.STATE.snapshot() if serving_state.STATE else None + vectorizer = getattr(snapshot, "vectorizer", None) + model = getattr(snapshot, "model", None) + label_encoder = getattr(snapshot, "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 serving state.") scanned_emails = [] spam_count = 0 diff --git a/backend/serving_state.py b/backend/serving_state.py index c118a1e..e43c182 100644 --- a/backend/serving_state.py +++ b/backend/serving_state.py @@ -17,6 +17,11 @@ the artifacts currently loaded, refreshed on every reload so ``/model-info`` and the per-prediction provenance fields report the live model. +The URL classifier pair travels in the snapshot as well. It is a separate model +from the text classifier, but it used to live in module globals that a reload +never touched, which meant a hot swap updated some prediction paths and not +others. + >>> state = ServingState( ... model="m1", vectorizer="v1", label_encoder="l1", xai_service="x1", ... metadata="meta1", @@ -57,6 +62,10 @@ class ServingSnapshot: # Provenance for the loaded artifacts (model_registry.ModelMetadata). Defaults # to None so callers/tests that build a snapshot without provenance still work. metadata: Any = None + # The URL classifier pair. Optional for the same reason as ``metadata``: + # lightweight fakes in the test suite construct snapshots without them. + url_model: Any = None + url_vectorizer: Any = None # A loader returns the freshly loaded objects (from disk) as a mapping with the @@ -77,6 +86,8 @@ def __init__( xai_service: Any, loader: Loader, metadata: Any = None, + url_model: Any = None, + url_vectorizer: Any = None, ) -> None: self._lock = threading.RLock() self._model = model @@ -85,6 +96,8 @@ def __init__( self._xai_service = xai_service self._loader = loader self._metadata = metadata + self._url_model = url_model + self._url_vectorizer = url_vectorizer self._version = 1 def snapshot(self) -> ServingSnapshot: @@ -96,6 +109,8 @@ def snapshot(self) -> ServingSnapshot: self._xai_service, self._version, self._metadata, + self._url_model, + self._url_vectorizer, ) def reload(self) -> ServingSnapshot: @@ -112,9 +127,12 @@ def reload(self) -> ServingSnapshot: self._vectorizer = fresh["vectorizer"] self._label_encoder = fresh["label_encoder"] self._xai_service = fresh["xai_service"] - # Loaders may omit "metadata" (e.g. lightweight test fakes); fall back - # to None rather than requiring every loader to supply provenance. + # Loaders may omit "metadata" and the URL pair (e.g. lightweight test + # fakes); fall back to None rather than requiring every loader to + # supply them. self._metadata = fresh.get("metadata") + self._url_model = fresh.get("url_model") + self._url_vectorizer = fresh.get("url_vectorizer") self._version += 1 return ServingSnapshot( self._model, @@ -123,6 +141,8 @@ def reload(self) -> ServingSnapshot: self._xai_service, self._version, self._metadata, + self._url_model, + self._url_vectorizer, ) @property @@ -145,6 +165,8 @@ def init_state( xai_service: Any, loader: Loader, metadata: Any = None, + url_model: Any = None, + url_vectorizer: Any = None, ) -> ServingState: """Install the process-wide serving state and return it.""" global STATE @@ -155,5 +177,7 @@ def init_state( xai_service=xai_service, loader=loader, metadata=metadata, + url_model=url_model, + url_vectorizer=url_vectorizer, ) return STATE diff --git a/backend/tests/test_inference_parity.py b/backend/tests/test_inference_parity.py index 8161e23..a02b228 100644 --- a/backend/tests/test_inference_parity.py +++ b/backend/tests/test_inference_parity.py @@ -79,13 +79,9 @@ 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, - ), + email_scanner.serving_state, + "STATE", + SimpleNamespace(snapshot=lambda: snapshot), ) monkeypatch.setattr(email_scanner, "analyze_headers", None) diff --git a/backend/tests/test_reload_coverage.py b/backend/tests/test_reload_coverage.py new file mode 100644 index 0000000..77ea600 --- /dev/null +++ b/backend/tests/test_reload_coverage.py @@ -0,0 +1,131 @@ +"""Every prediction path follows a hot reload (issue #1037). + +A reload that refreshes only some of the served objects leaves paths disagreeing +with each other, which is the same class of defect as train-serve skew. These +tests exercise the state holder directly with fakes -- no artifacts on disk -- and +assert that the URL pair and mailbox scanning both move with the swap. +""" + + +import serving_state + + +def _state(loader): + return serving_state.ServingState( + model="model-v1", + vectorizer="vectorizer-v1", + label_encoder="encoder-v1", + xai_service="xai-v1", + loader=loader, + url_model="url-model-v1", + url_vectorizer="url-vectorizer-v1", + ) + + +class TestUrlPairIsHotSwapped: + def test_initial_snapshot_carries_the_url_pair(self): + snapshot = _state(lambda: {}).snapshot() + + assert snapshot.url_model == "url-model-v1" + assert snapshot.url_vectorizer == "url-vectorizer-v1" + + def test_reload_replaces_the_url_pair(self): + state = _state( + lambda: { + "model": "model-v2", + "vectorizer": "vectorizer-v2", + "label_encoder": "encoder-v2", + "xai_service": "xai-v2", + "url_model": "url-model-v2", + "url_vectorizer": "url-vectorizer-v2", + } + ) + + snapshot = state.reload() + + assert snapshot.url_model == "url-model-v2" + assert snapshot.url_vectorizer == "url-vectorizer-v2" + assert snapshot.version == 2 + + def test_loader_may_omit_the_url_pair(self): + """Existing lightweight loaders must keep working after the extension.""" + state = _state( + lambda: { + "model": "model-v2", + "vectorizer": "vectorizer-v2", + "label_encoder": "encoder-v2", + "xai_service": "xai-v2", + } + ) + + snapshot = state.reload() + + assert snapshot.url_model is None + assert snapshot.model == "model-v2" + + +class TestMailboxScanFollowsReload: + def test_scanning_reads_the_post_reload_objects(self, monkeypatch): + from email_connectors import email_scanner + + captured = [] + + class Vectorizer: + def __init__(self, tag): + self.tag = tag + + def transform(self, texts): + captured.append(self.tag) + raise RuntimeError("stop after capturing the serving objects") + + state = serving_state.ServingState( + model="model-v1", + vectorizer=Vectorizer("v1"), + label_encoder="encoder-v1", + xai_service="xai-v1", + loader=lambda: { + "model": "model-v2", + "vectorizer": Vectorizer("v2"), + "label_encoder": "encoder-v2", + "xai_service": "xai-v2", + }, + ) + monkeypatch.setattr(email_scanner.serving_state, "STATE", state) + monkeypatch.setattr(email_scanner, "analyze_headers", None) + + email = [{"subject": "hello", "body": "world"}] + for _ in range(1): + try: + email_scanner.scan_emails_with_model(email) + except RuntimeError: + pass + + state.reload() + try: + email_scanner.scan_emails_with_model(email) + except RuntimeError: + pass + + # The second scan must have used the reloaded vectorizer, not the one the + # process started with. + assert captured == ["v1", "v2"] + + +class TestSnapshotStaysInternallyConsistent: + def test_reader_is_unaffected_by_a_later_reload(self): + state = _state( + lambda: { + "model": "model-v2", + "vectorizer": "vectorizer-v2", + "label_encoder": "encoder-v2", + "xai_service": "xai-v2", + "url_model": "url-model-v2", + "url_vectorizer": "url-vectorizer-v2", + } + ) + held = state.snapshot() + + state.reload() + + assert held.url_model == "url-model-v1" + assert held.model == "model-v1"