Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 11 additions & 6 deletions backend/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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),
}


Expand All @@ -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,
)


Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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"
Expand All @@ -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)
Expand Down
16 changes: 9 additions & 7 deletions backend/email_connectors/email_scanner.py
Original file line number Diff line number Diff line change
@@ -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

Expand All @@ -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
Expand Down
28 changes: 26 additions & 2 deletions backend/serving_state.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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:
Expand All @@ -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:
Expand All @@ -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,
Expand All @@ -123,6 +141,8 @@ def reload(self) -> ServingSnapshot:
self._xai_service,
self._version,
self._metadata,
self._url_model,
self._url_vectorizer,
)

@property
Expand All @@ -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
Expand All @@ -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
10 changes: 3 additions & 7 deletions backend/tests/test_inference_parity.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down
131 changes: 131 additions & 0 deletions backend/tests/test_reload_coverage.py
Original file line number Diff line number Diff line change
@@ -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"
Loading