initial code
This commit is contained in:
@@ -0,0 +1,545 @@
|
||||
"""Tests for the private, leakage-safe baseline training contract."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import csv
|
||||
import gzip
|
||||
import importlib.util
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
import warnings
|
||||
from pathlib import Path
|
||||
from typing import Dict, Iterable, List, Mapping
|
||||
from unittest import mock
|
||||
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
||||
MODULE_PATH = PROJECT_ROOT / "scripts/train_baselines.py"
|
||||
SPEC = importlib.util.spec_from_file_location("train_baselines", MODULE_PATH)
|
||||
if SPEC is None or SPEC.loader is None: # pragma: no cover - import machinery guard
|
||||
raise RuntimeError("Could not load scripts/train_baselines.py")
|
||||
baselines = importlib.util.module_from_spec(SPEC)
|
||||
sys.modules[SPEC.name] = baselines
|
||||
SPEC.loader.exec_module(baselines)
|
||||
|
||||
|
||||
SKLEARN_AVAILABLE = importlib.util.find_spec("numpy") is not None and importlib.util.find_spec(
|
||||
"sklearn"
|
||||
) is not None
|
||||
DUCKDB_AVAILABLE = importlib.util.find_spec("duckdb") is not None
|
||||
|
||||
|
||||
def valid_record(
|
||||
episode_start: str = "2020-06-01T09:00:00",
|
||||
partition: str = "train",
|
||||
target: int = 0,
|
||||
bucket: int = 25,
|
||||
token_number: int = 1,
|
||||
) -> Dict[str, object]:
|
||||
return {
|
||||
"vehicle_token": "{:064x}".format(token_number),
|
||||
"vehicle_bucket": bucket,
|
||||
"is_vin_audit": "true" if bucket < 10 else "false",
|
||||
"episode_number": 2,
|
||||
"episode_start": episode_start,
|
||||
"first_outcome": "fail" if target else "pass",
|
||||
"target_nonpass": target,
|
||||
"eligible_returning_target": "true",
|
||||
"temporal_partition": partition,
|
||||
"vehicle_age": 10,
|
||||
"prior_episode_count": 1,
|
||||
"prior_total_attempt_count": 1,
|
||||
"prior_attempt_count": 1,
|
||||
"days_since_prior_episode": 365,
|
||||
"days_since_prior_adverse": "",
|
||||
"prior_nonpass_rate": 0,
|
||||
"public_county": "salt_lake",
|
||||
"source_era": "slc",
|
||||
"target_outcome_label_source": "overall_result",
|
||||
"target_season": "summer",
|
||||
"prior_first_outcome": "pass",
|
||||
"prior_final_outcome": "pass",
|
||||
"last_observed_make": "toyota",
|
||||
"last_observed_model": "camry",
|
||||
}
|
||||
|
||||
|
||||
def write_csv(path: Path, records: Iterable[Mapping[str, object]]) -> None:
|
||||
rows = list(records)
|
||||
if not rows:
|
||||
raise ValueError("Tests require at least one record")
|
||||
opener = gzip.open if path.name.endswith(".gz") else open
|
||||
with opener(path, "wt", encoding="utf-8", newline="") as handle:
|
||||
writer = csv.DictWriter(handle, fieldnames=list(rows[0].keys()))
|
||||
writer.writeheader()
|
||||
writer.writerows(rows)
|
||||
|
||||
|
||||
class SchemaTests(unittest.TestCase):
|
||||
def test_required_schema_is_accepted(self) -> None:
|
||||
baselines.validate_schema(list(valid_record().keys()))
|
||||
|
||||
def test_missing_required_column_fails_closed(self) -> None:
|
||||
columns = list(valid_record().keys())
|
||||
columns.remove("prior_first_outcome")
|
||||
with self.assertRaisesRegex(baselines.SchemaError, "prior_first_outcome"):
|
||||
baselines.validate_schema(columns)
|
||||
|
||||
def test_target_label_source_column_is_required(self) -> None:
|
||||
columns = list(valid_record().keys())
|
||||
columns.remove("target_outcome_label_source")
|
||||
with self.assertRaisesRegex(
|
||||
baselines.SchemaError, "target_outcome_label_source"
|
||||
):
|
||||
baselines.validate_schema(columns)
|
||||
|
||||
def test_forbidden_column_fails_even_when_not_a_feature(self) -> None:
|
||||
columns = list(valid_record().keys()) + ["vin"]
|
||||
with self.assertRaisesRegex(baselines.SchemaError, "forbidden"):
|
||||
baselines.validate_schema(columns)
|
||||
|
||||
def test_unknown_sensitive_column_fails_closed(self) -> None:
|
||||
columns = list(valid_record().keys()) + ["owner_zip_code"]
|
||||
with self.assertRaisesRegex(baselines.SchemaError, "owner_zip_code"):
|
||||
baselines.validate_schema(columns)
|
||||
|
||||
def test_duplicate_columns_fail_closed(self) -> None:
|
||||
columns = list(valid_record().keys()) + ["vehicle_token"]
|
||||
with self.assertRaisesRegex(baselines.SchemaError, "duplicate"):
|
||||
baselines.validate_schema(columns)
|
||||
|
||||
|
||||
class MartValidationTests(unittest.TestCase):
|
||||
def test_loads_csv_gzip_and_assigns_fixed_partitions(self) -> None:
|
||||
records = [
|
||||
valid_record("2022-12-31T23:59:59", "train", 0, 25, 1),
|
||||
valid_record("2023-01-01T00:00:00", "tune", 1, 25, 2),
|
||||
valid_record("2024-01-01T00:00:00", "calibrate", 0, 5, 3),
|
||||
valid_record("2025-01-01T00:00:00", "test", 1, 25, 4),
|
||||
]
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
path = Path(directory) / "mart.csv.gz"
|
||||
write_csv(path, records)
|
||||
mart = baselines.load_mart(path)
|
||||
self.assertEqual(mart.eligible_rows, 4)
|
||||
self.assertEqual(len(mart.rows_by_partition["train"]), 1)
|
||||
self.assertEqual(len(mart.rows_by_partition["tune"]), 1)
|
||||
self.assertEqual(len(mart.rows_by_partition["calibrate"]), 1)
|
||||
self.assertEqual(len(mart.rows_by_partition["locked_test"]), 1)
|
||||
self.assertTrue(mart.rows_by_partition["calibrate"][0].audit_vehicle)
|
||||
|
||||
@unittest.skipUnless(DUCKDB_AVAILABLE, "duckdb is not installed")
|
||||
def test_loads_parquet_through_duckdb_without_pyarrow(self) -> None:
|
||||
import duckdb
|
||||
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
csv_path = Path(directory) / "mart.csv"
|
||||
parquet_path = Path(directory) / "mart.parquet"
|
||||
write_csv(csv_path, [valid_record()])
|
||||
connection = duckdb.connect(database=":memory:")
|
||||
try:
|
||||
connection.execute(
|
||||
"CREATE TABLE mart AS SELECT * FROM read_csv_auto(?)", [str(csv_path)]
|
||||
)
|
||||
escaped_path = str(parquet_path).replace("'", "''")
|
||||
connection.execute(
|
||||
"COPY mart TO '{}' (FORMAT PARQUET)".format(escaped_path)
|
||||
)
|
||||
finally:
|
||||
connection.close()
|
||||
mart = baselines.load_mart(parquet_path)
|
||||
self.assertEqual(mart.eligible_rows, 1)
|
||||
self.assertEqual(len(mart.rows_by_partition["train"]), 1)
|
||||
|
||||
def test_declared_partition_must_match_timestamp(self) -> None:
|
||||
record = valid_record("2025-05-01T00:00:00", "train")
|
||||
with self.assertRaisesRegex(baselines.DataValidationError, "belongs to locked_test"):
|
||||
baselines.episode_from_record(record, 2)
|
||||
|
||||
def test_label_column_must_match_canonical_outcome(self) -> None:
|
||||
record = valid_record(target=1)
|
||||
record["target_nonpass"] = 0
|
||||
with self.assertRaisesRegex(baselines.DataValidationError, "disagrees"):
|
||||
baselines.episode_from_record(record, 2)
|
||||
|
||||
def test_audit_boolean_must_match_bucket_contract(self) -> None:
|
||||
record = valid_record(bucket=5)
|
||||
record["is_vin_audit"] = "false"
|
||||
with self.assertRaisesRegex(baselines.DataValidationError, "bucket<10"):
|
||||
baselines.episode_from_record(record, 2)
|
||||
|
||||
def test_source_era_is_required_for_audit_but_not_modeling(self) -> None:
|
||||
record = valid_record()
|
||||
record["source_era"] = ""
|
||||
with self.assertRaisesRegex(baselines.DataValidationError, "drift auditing"):
|
||||
baselines.episode_from_record(record, 2)
|
||||
|
||||
def test_utah_proxy_is_accepted_only_for_utah_source_era(self) -> None:
|
||||
record = valid_record()
|
||||
record["source_era"] = "utah"
|
||||
record["target_outcome_label_source"] = "utah_obd_proxy"
|
||||
episode = baselines.episode_from_record(record, 2)
|
||||
self.assertEqual(episode.target_outcome_label_source, "utah_obd_proxy")
|
||||
|
||||
record["source_era"] = "slc"
|
||||
with self.assertRaisesRegex(baselines.DataValidationError, "outside source_era"):
|
||||
baselines.episode_from_record(record, 2)
|
||||
|
||||
record["source_era"] = "utah"
|
||||
record["first_outcome"] = "reject"
|
||||
record["target_nonpass"] = 1
|
||||
rejected = baselines.episode_from_record(record, 2)
|
||||
self.assertEqual(rejected.target, 1)
|
||||
self.assertEqual(rejected.target_outcome_label_source, "utah_obd_proxy")
|
||||
|
||||
record["first_outcome"] = "abort"
|
||||
aborted = baselines.episode_from_record(record, 2)
|
||||
self.assertEqual(aborted.target, 1)
|
||||
self.assertEqual(aborted.target_outcome_label_source, "utah_obd_proxy")
|
||||
|
||||
def test_unapproved_target_label_source_fails_closed(self) -> None:
|
||||
record = valid_record()
|
||||
record["target_outcome_label_source"] = "inferred_four_class"
|
||||
with self.assertRaisesRegex(baselines.DataValidationError, "unapproved"):
|
||||
baselines.episode_from_record(record, 2)
|
||||
|
||||
def test_first_episode_cannot_be_eligible_returning_target(self) -> None:
|
||||
record = valid_record()
|
||||
record["episode_number"] = 1
|
||||
with self.assertRaisesRegex(baselines.DataValidationError, "not a returning"):
|
||||
baselines.episode_from_record(record, 2)
|
||||
|
||||
def test_ineligible_row_is_not_parsed_as_a_supervised_target(self) -> None:
|
||||
record = valid_record()
|
||||
record["eligible_returning_target"] = "false"
|
||||
record["first_outcome"] = ""
|
||||
self.assertIsNone(baselines.episode_from_record(record, 2))
|
||||
|
||||
def test_private_output_path_guard(self) -> None:
|
||||
accepted = PROJECT_ROOT / "artifacts/private/baselines"
|
||||
rejected = PROJECT_ROOT / "artifacts/public/baselines"
|
||||
self.assertEqual(baselines.require_private_path(accepted, "output"), accepted)
|
||||
with self.assertRaisesRegex(baselines.DataValidationError, "must remain"):
|
||||
baselines.require_private_path(rejected, "output")
|
||||
|
||||
|
||||
class MetricTests(unittest.TestCase):
|
||||
def test_perfect_ranking_metrics(self) -> None:
|
||||
targets = [0, 1, 0, 1]
|
||||
probabilities = [0.1, 0.9, 0.2, 0.8]
|
||||
metrics = baselines.binary_metrics(targets, probabilities)
|
||||
self.assertAlmostEqual(metrics["average_precision"], 1.0)
|
||||
self.assertAlmostEqual(metrics["roc_auc"], 1.0)
|
||||
self.assertAlmostEqual(metrics["brier"], 0.025)
|
||||
self.assertAlmostEqual(metrics["top_5_precision"], 1.0)
|
||||
self.assertAlmostEqual(metrics["top_5_capture"], 0.1)
|
||||
|
||||
def test_capacity_metrics_fractionally_handle_boundary_ties(self) -> None:
|
||||
targets = [0, 1] * 50
|
||||
probabilities = [0.2] * 100
|
||||
top_5 = baselines.top_capacity_metrics(targets, probabilities, 0.05)
|
||||
top_10 = baselines.top_capacity_metrics(targets, probabilities, 0.10)
|
||||
self.assertAlmostEqual(top_5["precision"], 0.5)
|
||||
self.assertAlmostEqual(top_5["capture"], 0.05)
|
||||
self.assertAlmostEqual(top_10["precision"], 0.5)
|
||||
self.assertAlmostEqual(top_10["capture"], 0.10)
|
||||
|
||||
def test_constant_predictions_have_half_auc(self) -> None:
|
||||
targets = [0, 1, 0, 1]
|
||||
probabilities = [0.25] * 4
|
||||
self.assertAlmostEqual(baselines.roc_auc(targets, probabilities), 0.5)
|
||||
self.assertAlmostEqual(baselines.average_precision(targets, probabilities), 0.5)
|
||||
|
||||
def test_calibration_bins_cover_each_row_once(self) -> None:
|
||||
targets = [0, 1, 0, 1, 1]
|
||||
probabilities = [0.1, 0.2, 0.3, 0.8, 0.9]
|
||||
bins = baselines.calibration_bins(targets, probabilities, 3)
|
||||
self.assertEqual(sum(item["count"] for item in bins), len(targets))
|
||||
self.assertEqual([item["count"] for item in bins], [2, 2, 1])
|
||||
|
||||
|
||||
class BaselineBehaviorTests(unittest.TestCase):
|
||||
def test_default_iteration_budget_is_hardened(self) -> None:
|
||||
args = baselines.parse_args(["--mart", "data/private/example.parquet"])
|
||||
self.assertEqual(args.max_iter, 5000)
|
||||
|
||||
def _episode(
|
||||
self,
|
||||
prior_outcome: object,
|
||||
target: int = 0,
|
||||
bucket: int = 25,
|
||||
episode_start: str = "2020-06-01T09:00:00",
|
||||
partition: str = "train",
|
||||
source_era: str = "slc",
|
||||
label_source: str = "overall_result",
|
||||
):
|
||||
record = valid_record(
|
||||
episode_start=episode_start,
|
||||
partition=partition,
|
||||
target=target,
|
||||
bucket=bucket,
|
||||
)
|
||||
record["prior_first_outcome"] = prior_outcome
|
||||
record["source_era"] = source_era
|
||||
record["target_outcome_label_source"] = label_source
|
||||
return baselines.episode_from_record(record, 2)
|
||||
|
||||
def test_literal_previous_outcome_and_missing_fallback(self) -> None:
|
||||
rows = [
|
||||
self._episode("pass"),
|
||||
self._episode("fail"),
|
||||
self._episode("reject"),
|
||||
self._episode(""),
|
||||
]
|
||||
probabilities, fallback_count = baselines.literal_previous_probabilities(rows, 0.2)
|
||||
self.assertEqual(probabilities, [0.0, 1.0, 1.0, 0.2])
|
||||
self.assertEqual(fallback_count, 1)
|
||||
|
||||
def test_audit_rows_are_excluded_from_fit_cohort(self) -> None:
|
||||
rows = [self._episode("pass", bucket=5), self._episode("pass", bucket=25)]
|
||||
non_audit = baselines.evaluation_cohorts(rows, include_audit_breakout=False)
|
||||
self.assertEqual(non_audit[0][0], "non_audit")
|
||||
self.assertEqual(len(non_audit[0][1]), 1)
|
||||
self.assertFalse(non_audit[0][1][0].audit_vehicle)
|
||||
|
||||
def test_locked_manifest_counts_withhold_labels_until_unlocked(self) -> None:
|
||||
locked = self._episode(
|
||||
"pass",
|
||||
target=1,
|
||||
bucket=25,
|
||||
episode_start="2025-06-01T09:00:00",
|
||||
partition="test",
|
||||
)
|
||||
locked_proxy = self._episode(
|
||||
"pass",
|
||||
target=0,
|
||||
bucket=26,
|
||||
episode_start="2025-06-01T09:00:00",
|
||||
partition="test",
|
||||
source_era="utah",
|
||||
label_source="utah_obd_proxy",
|
||||
)
|
||||
mart = baselines.MartData(
|
||||
rows_by_partition={
|
||||
"train": [],
|
||||
"tune": [],
|
||||
"calibrate": [],
|
||||
"locked_test": [locked, locked_proxy],
|
||||
},
|
||||
input_rows=2,
|
||||
ineligible_rows=0,
|
||||
eligible_rows=2,
|
||||
)
|
||||
closed, closed_eras, closed_labels = baselines.manifest_audit_counts(
|
||||
mart, evaluate_locked=False
|
||||
)
|
||||
self.assertTrue(closed["locked_test"]["labels_withheld"])
|
||||
self.assertNotIn("nonpass", closed["locked_test"])
|
||||
self.assertNotIn("nonpass", closed_eras["locked_test"]["slc"])
|
||||
self.assertEqual(
|
||||
closed_labels["locked_test"],
|
||||
{"overall_result": 1, "utah_obd_proxy": 1},
|
||||
)
|
||||
|
||||
opened, opened_eras, opened_labels = baselines.manifest_audit_counts(
|
||||
mart, evaluate_locked=True
|
||||
)
|
||||
self.assertEqual(opened["locked_test"]["nonpass"], 1)
|
||||
self.assertEqual(opened_eras["locked_test"]["slc"]["nonpass"], 1)
|
||||
self.assertEqual(
|
||||
opened_labels["locked_test"],
|
||||
{"overall_result": 1, "utah_obd_proxy": 1},
|
||||
)
|
||||
self.assertFalse(
|
||||
baselines.TARGET_CONTRACT["utah_obd_proxy_four_class_approved"]
|
||||
)
|
||||
|
||||
|
||||
@unittest.skipUnless(SKLEARN_AVAILABLE, "numpy/scikit-learn are not installed")
|
||||
class ModelingIntegrationTests(unittest.TestCase):
|
||||
@staticmethod
|
||||
def _partition_records(
|
||||
timestamp: str,
|
||||
partition: str,
|
||||
start_token: int,
|
||||
future_category: bool = False,
|
||||
) -> List[Mapping[str, object]]:
|
||||
records: List[Mapping[str, object]] = []
|
||||
for index in range(40):
|
||||
target = 1 if index % 4 == 0 else 0
|
||||
bucket = index % 100
|
||||
record = valid_record(
|
||||
episode_start=timestamp,
|
||||
partition=partition,
|
||||
target=target,
|
||||
bucket=bucket,
|
||||
token_number=start_token + index,
|
||||
)
|
||||
record["vehicle_age"] = 2 + index % 25
|
||||
record["prior_episode_count"] = 1 + index % 6
|
||||
record["prior_total_attempt_count"] = 1 + index % 10
|
||||
record["prior_attempt_count"] = 1 + index % 3
|
||||
record["days_since_prior_episode"] = 250 + index * 4
|
||||
record["days_since_prior_adverse"] = "" if index % 3 else 500 + index
|
||||
record["prior_nonpass_rate"] = (index % 4) / 4
|
||||
record["prior_first_outcome"] = "fail" if index % 5 == 0 else "pass"
|
||||
if future_category:
|
||||
record["public_county"] = "future_only_county"
|
||||
records.append(record)
|
||||
return records
|
||||
|
||||
def _synthetic_mart(self):
|
||||
records: List[Mapping[str, object]] = []
|
||||
records.extend(self._partition_records("2022-06-01", "train", 1))
|
||||
records.extend(self._partition_records("2023-06-01", "tune", 101))
|
||||
records.extend(
|
||||
self._partition_records(
|
||||
"2024-06-01", "calibrate", 201, future_category=True
|
||||
)
|
||||
)
|
||||
records.extend(self._partition_records("2025-06-01", "test", 301))
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
path = Path(directory) / "mart.csv"
|
||||
write_csv(path, records)
|
||||
return baselines.load_mart(path)
|
||||
|
||||
def test_train_only_preprocessing_and_locked_evaluation_gate(self) -> None:
|
||||
mart = self._synthetic_mart()
|
||||
trained = baselines.train_baselines(
|
||||
mart,
|
||||
c_grid=[0.1],
|
||||
min_category_count=1,
|
||||
seed=7,
|
||||
max_iter=2000,
|
||||
)
|
||||
|
||||
feature_names = set(trained.encoder.vectorizer.get_feature_names_out())
|
||||
self.assertNotIn("cat__public_county=future_only_county", feature_names)
|
||||
self.assertFalse(any("source_era" in name for name in feature_names))
|
||||
self.assertFalse(
|
||||
any("target_outcome_label_source" in name for name in feature_names)
|
||||
)
|
||||
|
||||
locked_metrics, _ = baselines.evaluate_models(
|
||||
mart, trained, evaluate_locked=False, bin_count=5
|
||||
)
|
||||
self.assertNotIn("locked_test", {row["partition"] for row in locked_metrics})
|
||||
|
||||
unlocked_metrics, _ = baselines.evaluate_models(
|
||||
mart, trained, evaluate_locked=True, bin_count=5
|
||||
)
|
||||
locked_rows = [
|
||||
row for row in unlocked_metrics if row["partition"] == "locked_test"
|
||||
]
|
||||
self.assertTrue(locked_rows)
|
||||
self.assertIn("never_fit_audit", {row["cohort"] for row in locked_rows})
|
||||
self.assertIn("logistic_platt", {row["model"] for row in locked_rows})
|
||||
|
||||
tuning = trained.tuning_results[0]
|
||||
self.assertTrue(tuning["converged"])
|
||||
self.assertTrue(tuning["finite_coefficients"])
|
||||
self.assertTrue(tuning["finite_probabilities"])
|
||||
self.assertTrue(tuning["eligible_for_selection"])
|
||||
self.assertLess(tuning["n_iter"], trained.max_iter)
|
||||
self.assertLess(trained.selected_n_iter, trained.max_iter)
|
||||
self.assertLess(trained.platt_n_iter, trained.max_iter)
|
||||
self.assertEqual(trained.platt_model.solver, "liblinear")
|
||||
self.assertEqual(trained.platt_model.penalty, "l2")
|
||||
self.assertEqual(trained.platt_model.C, 1000.0)
|
||||
self.assertEqual(
|
||||
baselines.PLATT_CALIBRATION_CONFIG,
|
||||
{
|
||||
"method": "sigmoid_platt_scaling",
|
||||
"solver": "liblinear",
|
||||
"penalty": "l2",
|
||||
"c": 1000.0,
|
||||
},
|
||||
)
|
||||
|
||||
def test_deliberately_tiny_max_iter_rejects_all_candidates(self) -> None:
|
||||
mart = self._synthetic_mart()
|
||||
with self.assertRaisesRegex(
|
||||
baselines.DataValidationError, "No logistic candidate converged"
|
||||
):
|
||||
baselines.train_baselines(
|
||||
mart,
|
||||
c_grid=[0.03, 0.1],
|
||||
min_category_count=1,
|
||||
seed=7,
|
||||
max_iter=1,
|
||||
)
|
||||
|
||||
def test_nonconverged_and_nonfinite_candidates_are_excluded(self) -> None:
|
||||
import numpy as np
|
||||
import sklearn
|
||||
from sklearn.exceptions import ConvergenceWarning
|
||||
from sklearn.feature_extraction import DictVectorizer
|
||||
|
||||
class ControlledLogistic:
|
||||
def __init__(self, C=1.0, solver="lbfgs", max_iter=100, **_kwargs):
|
||||
self.C = C
|
||||
self.solver = solver
|
||||
self.max_iter = max_iter
|
||||
|
||||
def fit(self, features, _targets):
|
||||
self.n_iter_ = np.asarray([2])
|
||||
self.coef_ = np.zeros((1, features.shape[1]))
|
||||
self.intercept_ = np.zeros(1)
|
||||
if self.solver == "saga" and self.C == 0.03:
|
||||
warnings.warn("controlled nonconvergence", ConvergenceWarning)
|
||||
if self.solver == "saga" and self.C == 0.1:
|
||||
self.coef_[0, 0] = np.nan
|
||||
return self
|
||||
|
||||
def predict_proba(self, features):
|
||||
probability = np.full(features.shape[0], 0.25)
|
||||
return np.column_stack((1.0 - probability, probability))
|
||||
|
||||
def decision_function(self, features):
|
||||
return np.zeros(features.shape[0])
|
||||
|
||||
dependencies = (
|
||||
np,
|
||||
sklearn,
|
||||
DictVectorizer,
|
||||
ControlledLogistic,
|
||||
ConvergenceWarning,
|
||||
)
|
||||
mart = self._synthetic_mart()
|
||||
with mock.patch.object(
|
||||
baselines, "require_ml_dependencies", return_value=dependencies
|
||||
):
|
||||
trained = baselines.train_baselines(
|
||||
mart,
|
||||
c_grid=[0.03, 0.1, 1.0],
|
||||
min_category_count=1,
|
||||
seed=7,
|
||||
max_iter=100,
|
||||
)
|
||||
|
||||
self.assertEqual(trained.selected_c, 1.0)
|
||||
by_c = {result["c"]: result for result in trained.tuning_results}
|
||||
self.assertFalse(by_c[0.03]["converged"])
|
||||
self.assertFalse(by_c[0.03]["eligible_for_selection"])
|
||||
self.assertFalse(by_c[0.1]["finite_coefficients"])
|
||||
self.assertFalse(by_c[0.1]["eligible_for_selection"])
|
||||
self.assertTrue(by_c[1.0]["eligible_for_selection"])
|
||||
|
||||
def test_convergence_capture_replays_unrelated_warnings(self) -> None:
|
||||
from sklearn.exceptions import ConvergenceWarning
|
||||
|
||||
class WarningModel:
|
||||
def fit(self, _features, _targets):
|
||||
warnings.warn("unrelated diagnostic", RuntimeWarning)
|
||||
warnings.warn("did not converge", ConvergenceWarning)
|
||||
|
||||
with self.assertWarnsRegex(RuntimeWarning, "unrelated diagnostic"):
|
||||
messages = baselines.fit_with_convergence_capture(
|
||||
WarningModel(), None, None, ConvergenceWarning
|
||||
)
|
||||
self.assertEqual(messages, ["did not converge"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,564 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest import mock
|
||||
|
||||
import duckdb
|
||||
|
||||
from scripts import export_dashboard_data as dashboard_export
|
||||
|
||||
|
||||
class DashboardExportTests(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
self.temporary = tempfile.TemporaryDirectory()
|
||||
self.root = Path(self.temporary.name)
|
||||
self.mart = self.root / "private_feature_mart.parquet"
|
||||
self.mart_manifest = self.root / "private_feature_mart.parquet.manifest.json"
|
||||
self.model_manifest = self.root / "model_manifest.json"
|
||||
self.model_metrics = self.root / "model_metrics.json"
|
||||
self.public_root = (self.root / "dashboard/public/data").resolve()
|
||||
self._write_mart()
|
||||
self._write_mart_manifest()
|
||||
self._write_model_files()
|
||||
|
||||
def tearDown(self) -> None:
|
||||
self.temporary.cleanup()
|
||||
|
||||
def test_development_mart_is_refused_without_preview_flag(self) -> None:
|
||||
with mock.patch.object(
|
||||
dashboard_export, "PUBLIC_DATA_ROOT", self.public_root
|
||||
):
|
||||
with self.assertRaises(dashboard_export.ManifestError):
|
||||
self._export(development_preview=False)
|
||||
self.assertFalse(self.public_root.exists())
|
||||
|
||||
def test_preview_is_suppressed_deterministic_and_contains_no_locked_data(self) -> None:
|
||||
with mock.patch.object(
|
||||
dashboard_export, "PUBLIC_DATA_ROOT", self.public_root
|
||||
):
|
||||
paths = self._export(development_preview=True)
|
||||
|
||||
expected_names = set(dashboard_export.PUBLIC_ASSET_NAMES) | {
|
||||
dashboard_export.SHA256_MANIFEST_NAME
|
||||
}
|
||||
self.assertEqual(set(paths), expected_names)
|
||||
|
||||
first_bytes = {
|
||||
name: (self.public_root / name).read_bytes()
|
||||
for name in expected_names
|
||||
}
|
||||
joined = b"\n".join(first_bytes.values()).lower()
|
||||
for forbidden in (
|
||||
b"lockedmake",
|
||||
b"secret_locked",
|
||||
b"vehicle_token",
|
||||
b"internal_event_id",
|
||||
b'"vin"',
|
||||
b'"plate"',
|
||||
b'"zip"',
|
||||
b'"station"',
|
||||
b"episode_start",
|
||||
b"target_outcome_label_source",
|
||||
b'"prediction"',
|
||||
):
|
||||
self.assertNotIn(forbidden, joined)
|
||||
|
||||
for name in expected_names:
|
||||
payload = json.loads((self.public_root / name).read_text(encoding="utf-8"))
|
||||
self.assertIs(payload["development_preview"], True)
|
||||
self.assertIs(payload["population_estimate_allowed"], False)
|
||||
|
||||
scorecards = self._asset("cohort_scorecard.json")["rows"]
|
||||
self.assertEqual(
|
||||
scorecards,
|
||||
[
|
||||
{
|
||||
"prior_make": "SAFE",
|
||||
"prior_model": "SUPPORTED",
|
||||
"support_rounded": 400,
|
||||
"nonpass_rate": 0.167,
|
||||
}
|
||||
],
|
||||
)
|
||||
overview = self._asset("overview_period_county.json")["rows"]
|
||||
self.assertEqual(len(overview), 3)
|
||||
self.assertEqual(overview[0]["support_rounded"], 800)
|
||||
self.assertNotIn("n", overview[0])
|
||||
self.assertNotIn("nonpass_count", overview[0])
|
||||
|
||||
diagnostics = self._asset("model_diagnostics.json")["rows"]
|
||||
self.assertEqual(
|
||||
{row["partition"] for row in diagnostics},
|
||||
{"train", "tune", "calibrate"},
|
||||
)
|
||||
self.assertEqual({row["model"] for row in diagnostics}, {"test_model"})
|
||||
self.assertTrue(all("episodes" not in row for row in diagnostics))
|
||||
|
||||
data_manifest = self._asset("data_manifest.json")
|
||||
self.assertEqual(data_manifest["model_versions"], ["test_model_v1"])
|
||||
self.assertRegex(data_manifest["release_id"], r"^[0-9a-f]{64}$")
|
||||
|
||||
checksum_manifest = self._asset(dashboard_export.SHA256_MANIFEST_NAME)
|
||||
for entry in checksum_manifest["files"]:
|
||||
content = (self.public_root / entry["name"]).read_bytes()
|
||||
self.assertEqual(hashlib.sha256(content).hexdigest(), entry["sha256"])
|
||||
|
||||
with mock.patch.object(
|
||||
dashboard_export, "PUBLIC_DATA_ROOT", self.public_root
|
||||
):
|
||||
self._export(development_preview=True, overwrite=True)
|
||||
second_bytes = {
|
||||
name: (self.public_root / name).read_bytes()
|
||||
for name in expected_names
|
||||
}
|
||||
self.assertEqual(first_bytes, second_bytes)
|
||||
|
||||
def test_unlocked_model_manifest_is_refused(self) -> None:
|
||||
manifest = self._read_json(self.model_manifest)
|
||||
manifest["locked_test_evaluated"] = True
|
||||
self._write_json(self.model_manifest, manifest)
|
||||
|
||||
with mock.patch.object(
|
||||
dashboard_export, "PUBLIC_DATA_ROOT", self.public_root
|
||||
):
|
||||
with self.assertRaises(dashboard_export.ManifestError):
|
||||
self._export(development_preview=True)
|
||||
self.assertFalse(self.public_root.exists())
|
||||
|
||||
def test_locked_metric_row_is_refused_even_when_manifest_is_closed(self) -> None:
|
||||
metrics = self._read_json(self.model_metrics)
|
||||
locked = dict(metrics[0])
|
||||
locked["partition"] = "locked_test"
|
||||
metrics.append(locked)
|
||||
self._write_json(self.model_metrics, metrics)
|
||||
self._refresh_metrics_digest()
|
||||
|
||||
with mock.patch.object(
|
||||
dashboard_export, "PUBLIC_DATA_ROOT", self.public_root
|
||||
):
|
||||
with self.assertRaises(dashboard_export.ManifestError):
|
||||
self._export(development_preview=True)
|
||||
self.assertFalse(self.public_root.exists())
|
||||
|
||||
def test_metrics_digest_mismatch_is_refused(self) -> None:
|
||||
metrics = self._read_json(self.model_metrics)
|
||||
metrics[0]["average_precision"] = 0.99
|
||||
self._write_json(self.model_metrics, metrics)
|
||||
|
||||
with mock.patch.object(
|
||||
dashboard_export, "PUBLIC_DATA_ROOT", self.public_root
|
||||
):
|
||||
with self.assertRaises(dashboard_export.ManifestError):
|
||||
self._export(development_preview=True)
|
||||
self.assertFalse(self.public_root.exists())
|
||||
|
||||
def test_overwrite_refuses_unexpected_file(self) -> None:
|
||||
with mock.patch.object(
|
||||
dashboard_export, "PUBLIC_DATA_ROOT", self.public_root
|
||||
):
|
||||
self._export(development_preview=True)
|
||||
unexpected = self.public_root / "private-notes.txt"
|
||||
unexpected.write_text("must not ride along\n", encoding="utf-8")
|
||||
with self.assertRaises(dashboard_export.PublicSchemaError):
|
||||
self._export(development_preview=True, overwrite=True)
|
||||
self.assertTrue(unexpected.exists())
|
||||
|
||||
def test_overwrite_refuses_directory_in_public_contract(self) -> None:
|
||||
self.public_root.mkdir(parents=True)
|
||||
(self.public_root / "data_manifest.json").mkdir()
|
||||
with mock.patch.object(
|
||||
dashboard_export, "PUBLIC_DATA_ROOT", self.public_root
|
||||
):
|
||||
with self.assertRaises(dashboard_export.PublicSchemaError):
|
||||
self._export(development_preview=True, overwrite=True)
|
||||
|
||||
def test_pre_2025_locked_partition_aliases_are_refused(self) -> None:
|
||||
for partition in dashboard_export.LOCKED_PARTITION_ALIASES:
|
||||
with self.subTest(partition=partition):
|
||||
self._move_locked_rows_before_boundary(partition)
|
||||
with mock.patch.object(
|
||||
dashboard_export, "PUBLIC_DATA_ROOT", self.public_root
|
||||
):
|
||||
with self.assertRaises(dashboard_export.ManifestError):
|
||||
self._export(development_preview=True)
|
||||
self.assertFalse(self.public_root.exists())
|
||||
|
||||
def test_public_validator_rejects_sensitive_keys_and_exact_timestamps(self) -> None:
|
||||
base = {
|
||||
"schema_version": dashboard_export.SCHEMA_VERSION,
|
||||
"development_preview": True,
|
||||
"population_estimate_allowed": False,
|
||||
}
|
||||
unsafe_key = {
|
||||
**base,
|
||||
"rows": [
|
||||
{
|
||||
"year": 2022,
|
||||
"quarter": 1,
|
||||
"public_county": "salt_lake",
|
||||
"support_rounded": 100,
|
||||
"nonpass_rate": 0.2,
|
||||
"vehicle_token": "not-public",
|
||||
}
|
||||
],
|
||||
}
|
||||
with self.assertRaises(dashboard_export.PublicSchemaError):
|
||||
dashboard_export.validate_public_asset(
|
||||
"overview_period_county.json", unsafe_key
|
||||
)
|
||||
|
||||
exact_timestamp = {
|
||||
**base,
|
||||
"rows": [
|
||||
{
|
||||
"year": 2022,
|
||||
"quarter": 1,
|
||||
"public_county": "2022-01-01T12:34:56",
|
||||
"support_rounded": 100,
|
||||
"nonpass_rate": 0.2,
|
||||
}
|
||||
],
|
||||
}
|
||||
with self.assertRaises(dashboard_export.PublicSchemaError):
|
||||
dashboard_export.validate_public_asset(
|
||||
"overview_period_county.json", exact_timestamp
|
||||
)
|
||||
|
||||
def _export(
|
||||
self,
|
||||
development_preview: bool,
|
||||
overwrite: bool = False,
|
||||
) -> dict:
|
||||
return dashboard_export.export_dashboard_data(
|
||||
mart=self.mart,
|
||||
mart_manifest=self.mart_manifest,
|
||||
bundles=[
|
||||
dashboard_export.ModelBundle(
|
||||
manifest=self.model_manifest,
|
||||
metrics=self.model_metrics,
|
||||
)
|
||||
],
|
||||
output_dir=self.public_root,
|
||||
development_preview=development_preview,
|
||||
overwrite=overwrite,
|
||||
)
|
||||
|
||||
def _write_mart(self) -> None:
|
||||
connection = duckdb.connect(":memory:")
|
||||
try:
|
||||
connection.execute(
|
||||
"""
|
||||
CREATE TABLE mart (
|
||||
vehicle_token VARCHAR,
|
||||
is_vin_audit BOOLEAN,
|
||||
episode_number BIGINT,
|
||||
episode_start TIMESTAMP,
|
||||
first_outcome VARCHAR,
|
||||
target_outcome_label_source VARCHAR,
|
||||
target_nonpass INTEGER,
|
||||
eligible_returning_target BOOLEAN,
|
||||
temporal_partition VARCHAR,
|
||||
public_county VARCHAR,
|
||||
source_era VARCHAR,
|
||||
vehicle_age INTEGER,
|
||||
last_observed_make VARCHAR,
|
||||
last_observed_model VARCHAR
|
||||
)
|
||||
"""
|
||||
)
|
||||
rows = []
|
||||
rows.extend(self._group(120, 20, "SAFE", "SUPPORTED"))
|
||||
rows.extend(self._group(99, 19, "LOW", "SUPPORT"))
|
||||
rows.extend(self._group(120, 9, "LOW", "NONPASS"))
|
||||
rows.extend(self._group(120, 111, "LOW", "PASS"))
|
||||
rows.extend(
|
||||
self._group(
|
||||
120,
|
||||
60,
|
||||
"ENTITY",
|
||||
"LOWTOTAL",
|
||||
pass_vehicle_count=40,
|
||||
nonpass_vehicle_count=40,
|
||||
)
|
||||
)
|
||||
rows.extend(
|
||||
self._group(
|
||||
120,
|
||||
20,
|
||||
"ENTITY",
|
||||
"LOWNONPASS",
|
||||
nonpass_vehicle_count=9,
|
||||
)
|
||||
)
|
||||
rows.extend(
|
||||
self._group(
|
||||
120,
|
||||
100,
|
||||
"ENTITY",
|
||||
"LOWPASS",
|
||||
pass_vehicle_count=9,
|
||||
)
|
||||
)
|
||||
rows.extend(
|
||||
self._group(
|
||||
120,
|
||||
20,
|
||||
"SAFE",
|
||||
"SUPPORTED",
|
||||
temporal_partition="tune",
|
||||
)
|
||||
)
|
||||
rows.extend(
|
||||
self._group(
|
||||
120,
|
||||
20,
|
||||
"SAFE",
|
||||
"SUPPORTED",
|
||||
temporal_partition="calibrate",
|
||||
)
|
||||
)
|
||||
for index in range(120):
|
||||
rows.append(
|
||||
(
|
||||
"secret_vehicle_token",
|
||||
False,
|
||||
2,
|
||||
"2025-01-15 12:34:56",
|
||||
"secret_locked",
|
||||
"secret_locked_source",
|
||||
999,
|
||||
True,
|
||||
"locked_test",
|
||||
"locked_county",
|
||||
"locked_source",
|
||||
5,
|
||||
"LOCKEDMAKE",
|
||||
"LOCKEDMODEL",
|
||||
)
|
||||
)
|
||||
connection.executemany(
|
||||
"INSERT INTO mart VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)",
|
||||
rows,
|
||||
)
|
||||
connection.table("mart").write_parquet(str(self.mart))
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
@staticmethod
|
||||
def _group(
|
||||
n: int,
|
||||
nonpass: int,
|
||||
make: str,
|
||||
model: str,
|
||||
temporal_partition: str = "train",
|
||||
pass_vehicle_count: int = None,
|
||||
nonpass_vehicle_count: int = None,
|
||||
) -> list:
|
||||
partition_dates = {
|
||||
"train": "2022-01-15 08:00:00",
|
||||
"tune": "2023-01-15 08:00:00",
|
||||
"calibrate": "2024-01-15 08:00:00",
|
||||
}
|
||||
pass_count = n - nonpass
|
||||
pass_vehicle_count = pass_count if pass_vehicle_count is None else pass_vehicle_count
|
||||
nonpass_vehicle_count = (
|
||||
nonpass if nonpass_vehicle_count is None else nonpass_vehicle_count
|
||||
)
|
||||
rows = []
|
||||
for index in range(n):
|
||||
adverse = index < nonpass
|
||||
class_index = index if adverse else index - nonpass
|
||||
class_vehicle_count = (
|
||||
nonpass_vehicle_count if adverse else pass_vehicle_count
|
||||
)
|
||||
outcome = "nonpass" if adverse else "pass"
|
||||
vehicle_token = "{}_{}_{}_{}_{}".format(
|
||||
make,
|
||||
model,
|
||||
temporal_partition,
|
||||
outcome,
|
||||
class_index % class_vehicle_count,
|
||||
)
|
||||
rows.append(
|
||||
(
|
||||
vehicle_token,
|
||||
False,
|
||||
2,
|
||||
partition_dates[temporal_partition],
|
||||
"fail" if adverse else "pass",
|
||||
"overall_result",
|
||||
1 if adverse else 0,
|
||||
True,
|
||||
temporal_partition,
|
||||
"salt_lake",
|
||||
"slc",
|
||||
5,
|
||||
make,
|
||||
model,
|
||||
)
|
||||
)
|
||||
return rows
|
||||
|
||||
def _write_mart_manifest(self) -> None:
|
||||
connection = duckdb.connect(":memory:")
|
||||
try:
|
||||
columns = [
|
||||
{"name": row[0], "duckdb_type": row[1]}
|
||||
for row in connection.execute(
|
||||
"DESCRIBE SELECT * FROM read_parquet(?)", [str(self.mart)]
|
||||
).fetchall()
|
||||
]
|
||||
finally:
|
||||
connection.close()
|
||||
self._write_json(
|
||||
self.mart_manifest,
|
||||
{
|
||||
"build_kind": "private_leakage_safe_inspection_feature_mart",
|
||||
"classification": "private_pseudonymized_analytical_mart",
|
||||
"source_data_kind": "vehicle_history_development_sample",
|
||||
"population_estimate_allowed": False,
|
||||
"episode_gap_days": 30,
|
||||
"columns": columns,
|
||||
"parquet_sha256": self._sha256(self.mart),
|
||||
},
|
||||
)
|
||||
|
||||
def _write_model_files(self) -> None:
|
||||
connection = duckdb.connect(":memory:")
|
||||
try:
|
||||
private_support = {
|
||||
partition: (episodes, nonpass, vehicles)
|
||||
for partition, episodes, nonpass, vehicles in connection.execute(
|
||||
"""
|
||||
SELECT temporal_partition,
|
||||
count(*)::BIGINT,
|
||||
sum(target_nonpass)::BIGINT,
|
||||
count(DISTINCT vehicle_token)::BIGINT
|
||||
FROM read_parquet(?)
|
||||
WHERE temporal_partition IN ('train', 'tune', 'calibrate')
|
||||
AND eligible_returning_target
|
||||
AND target_nonpass IN (0, 1)
|
||||
AND NOT is_vin_audit
|
||||
GROUP BY 1
|
||||
""",
|
||||
[str(self.mart)],
|
||||
).fetchall()
|
||||
}
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
rows = []
|
||||
for partition, value in (
|
||||
("train", 0.25),
|
||||
("tune", 0.24),
|
||||
("calibrate", 0.23),
|
||||
):
|
||||
episodes, nonpass, vehicles = private_support[partition]
|
||||
rows.append(
|
||||
{
|
||||
"model": "test_model",
|
||||
"partition": partition,
|
||||
"cohort": "non_audit",
|
||||
"episodes": episodes,
|
||||
"nonpass": nonpass,
|
||||
"vehicles": vehicles,
|
||||
"average_precision": value,
|
||||
"brier": 0.1,
|
||||
"log_loss": 0.3,
|
||||
"roc_auc": 0.7,
|
||||
"top_10_capture": 0.2,
|
||||
}
|
||||
)
|
||||
for model, episodes, nonpass in (
|
||||
("low_support", 99, 20),
|
||||
("low_nonpass", 120, 9),
|
||||
("low_pass", 120, 111),
|
||||
("low_vehicles", 120, 20),
|
||||
):
|
||||
rows.append(
|
||||
{
|
||||
"model": model,
|
||||
"partition": "train",
|
||||
"cohort": "non_audit",
|
||||
"episodes": episodes,
|
||||
"nonpass": nonpass,
|
||||
"vehicles": 99 if model == "low_vehicles" else episodes,
|
||||
"average_precision": 0.2,
|
||||
"brier": 0.1,
|
||||
"log_loss": 0.3,
|
||||
"roc_auc": 0.7,
|
||||
"top_10_capture": 0.2,
|
||||
}
|
||||
)
|
||||
self._write_json(self.model_metrics, rows)
|
||||
self._write_json(
|
||||
self.model_manifest,
|
||||
{
|
||||
"classification": "private_model_artifact_no_row_predictions",
|
||||
"input_sha256": self._sha256(self.mart),
|
||||
"metrics_sha256": self._sha256(self.model_metrics),
|
||||
"locked_test_evaluated": False,
|
||||
"model_version": "test_model_v1",
|
||||
},
|
||||
)
|
||||
|
||||
def _refresh_metrics_digest(self) -> None:
|
||||
manifest = self._read_json(self.model_manifest)
|
||||
manifest["metrics_sha256"] = self._sha256(self.model_metrics)
|
||||
self._write_json(self.model_manifest, manifest)
|
||||
|
||||
def _move_locked_rows_before_boundary(self, partition: str) -> None:
|
||||
rewritten = self.root / "rewritten.parquet"
|
||||
connection = duckdb.connect(":memory:")
|
||||
try:
|
||||
connection.execute(
|
||||
"CREATE TABLE rewritten AS SELECT * FROM read_parquet(?)",
|
||||
[str(self.mart)],
|
||||
)
|
||||
connection.execute(
|
||||
"""
|
||||
UPDATE rewritten
|
||||
SET temporal_partition = ?,
|
||||
episode_start = TIMESTAMP '2024-12-31 12:34:56'
|
||||
WHERE last_observed_make = 'LOCKEDMAKE'
|
||||
""",
|
||||
[partition],
|
||||
)
|
||||
connection.table("rewritten").write_parquet(str(rewritten))
|
||||
finally:
|
||||
connection.close()
|
||||
rewritten.replace(self.mart)
|
||||
|
||||
mart_manifest = self._read_json(self.mart_manifest)
|
||||
mart_manifest["parquet_sha256"] = self._sha256(self.mart)
|
||||
self._write_json(self.mart_manifest, mart_manifest)
|
||||
|
||||
model_manifest = self._read_json(self.model_manifest)
|
||||
model_manifest["input_sha256"] = self._sha256(self.mart)
|
||||
self._write_json(self.model_manifest, model_manifest)
|
||||
|
||||
def _asset(self, name: str) -> object:
|
||||
return self._read_json(self.public_root / name)
|
||||
|
||||
@staticmethod
|
||||
def _sha256(path: Path) -> str:
|
||||
return hashlib.sha256(path.read_bytes()).hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def _write_json(path: Path, value: object) -> None:
|
||||
path.write_text(
|
||||
json.dumps(value, sort_keys=True, allow_nan=False) + "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _read_json(path: Path) -> object:
|
||||
return json.loads(path.read_text(encoding="utf-8"))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,550 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import csv
|
||||
import gzip
|
||||
import hashlib
|
||||
import json
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from typing import Dict, Iterable, List, Mapping, Optional
|
||||
|
||||
import duckdb
|
||||
|
||||
from scripts import build_feature_mart as mart
|
||||
|
||||
|
||||
def token(character: str) -> str:
|
||||
return character * 64
|
||||
|
||||
|
||||
def bucket(vehicle_token: str) -> int:
|
||||
return int(vehicle_token[:8], 16) % 100
|
||||
|
||||
|
||||
def event(
|
||||
event_id: int,
|
||||
vehicle_token: str,
|
||||
event_ts: str,
|
||||
outcome: Optional[str],
|
||||
make: Optional[str] = "TOYOTA",
|
||||
model: Optional[str] = "CAMRY",
|
||||
model_year: Optional[int] = 2010,
|
||||
source_era: str = "slc",
|
||||
public_county: str = "salt_lake",
|
||||
label_source: Optional[str] = "overall_result",
|
||||
) -> Dict[str, object]:
|
||||
return {
|
||||
"internal_event_id": event_id,
|
||||
"vehicle_token": vehicle_token,
|
||||
"vehicle_bucket": bucket(vehicle_token),
|
||||
"event_ts": event_ts,
|
||||
"source_era": source_era,
|
||||
"public_county": public_county,
|
||||
"canonical_outcome": outcome,
|
||||
"outcome_label_source": label_source if outcome is not None else None,
|
||||
"program_type": "obd",
|
||||
"test_type": "OBDII",
|
||||
"observed_make": make,
|
||||
"observed_model": model,
|
||||
"observed_model_year": model_year,
|
||||
}
|
||||
|
||||
|
||||
def write_export(
|
||||
directory: Path,
|
||||
name: str,
|
||||
rows: Iterable[Mapping[str, object]],
|
||||
start: str,
|
||||
end: str,
|
||||
kind: str = mart.DEVELOPMENT_KIND,
|
||||
) -> Path:
|
||||
path = directory / name
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
materialized = list(rows)
|
||||
with gzip.open(path, "wt", encoding="utf-8", newline="") as handle:
|
||||
writer = csv.DictWriter(handle, fieldnames=mart.EXPECTED_COLUMNS)
|
||||
writer.writeheader()
|
||||
writer.writerows(materialized)
|
||||
|
||||
digest = mart.file_sha256(path)
|
||||
manifest = {
|
||||
"export_kind": kind,
|
||||
"query_version": (
|
||||
"inspection_batch_v2"
|
||||
if kind == mart.BOUNDED_KIND
|
||||
else "inspection_history_development_sample_v1"
|
||||
),
|
||||
"query_sha256": "1" * 64,
|
||||
"generated_at_utc": "2026-07-15T12:00:00+00:00",
|
||||
"source_start_inclusive": start,
|
||||
"source_end_exclusive": end,
|
||||
"vehicle_token_key_version": "v1",
|
||||
"vehicle_token_key_fingerprint": "abcdef1234567890",
|
||||
"rows": len(materialized),
|
||||
"columns": list(mart.EXPECTED_COLUMNS),
|
||||
"compressed_file_sha256": digest,
|
||||
"compressed_file_bytes": path.stat().st_size,
|
||||
"classification": (
|
||||
"private_pseudonymized_analytical_staging"
|
||||
if kind == mart.BOUNDED_KIND
|
||||
else "private_pseudonymized_development_only"
|
||||
),
|
||||
}
|
||||
if kind == mart.DEVELOPMENT_KIND:
|
||||
manifest.update(
|
||||
{
|
||||
"population_estimate_allowed": False,
|
||||
"vehicle_sample_limit": 100,
|
||||
"sample_method": "synthetic_test_fixture",
|
||||
}
|
||||
)
|
||||
manifest_path = path.with_name(path.name + ".manifest.json")
|
||||
manifest_path.write_text(json.dumps(manifest), encoding="utf-8")
|
||||
return path
|
||||
|
||||
|
||||
class FeatureMartTest(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
self._temporary_directory = tempfile.TemporaryDirectory()
|
||||
self.root = Path(self._temporary_directory.name)
|
||||
|
||||
def tearDown(self) -> None:
|
||||
self._temporary_directory.cleanup()
|
||||
|
||||
def build(self, input_path: Path) -> mart.BuildSummary:
|
||||
return mart.build_feature_mart(
|
||||
mart.BuildConfig(
|
||||
input_path=input_path,
|
||||
database_path=self.root / "warehouse.duckdb",
|
||||
output_path=self.root / "feature_mart.parquet",
|
||||
memory_limit="512MB",
|
||||
threads=1,
|
||||
allow_external_output=True,
|
||||
)
|
||||
)
|
||||
|
||||
def row_dict(
|
||||
self, connection: duckdb.DuckDBPyConnection, sql: str, parameters=None
|
||||
) -> Dict[str, object]:
|
||||
cursor = connection.execute(sql, parameters or [])
|
||||
names = [item[0] for item in cursor.description]
|
||||
row = cursor.fetchone()
|
||||
self.assertIsNotNone(row)
|
||||
return dict(zip(names, row))
|
||||
|
||||
def test_episode_boundary_quotes_deduplication_and_prior_only_features(self) -> None:
|
||||
vehicle = token("a")
|
||||
rows = [
|
||||
event(
|
||||
1,
|
||||
vehicle,
|
||||
"2015-01-01 00:00:00",
|
||||
"pass",
|
||||
make="ACME, MOTORS",
|
||||
model='MODEL "Q"',
|
||||
),
|
||||
event(
|
||||
2,
|
||||
vehicle,
|
||||
"2015-01-31 00:00:00",
|
||||
"fail",
|
||||
make="ACME, MOTORS",
|
||||
model='MODEL "Q"',
|
||||
),
|
||||
# Exact analytical duplicate with a distinct source row ID.
|
||||
event(
|
||||
3,
|
||||
vehicle,
|
||||
"2015-01-01 00:00:00",
|
||||
"pass",
|
||||
make="ACME, MOTORS",
|
||||
model='MODEL "Q"',
|
||||
),
|
||||
# Thirty days and one second after the preceding accepted event.
|
||||
event(
|
||||
4,
|
||||
vehicle,
|
||||
"2015-03-02 00:00:01",
|
||||
"reject",
|
||||
make="ACME, MOTORS",
|
||||
model='MODEL "Q"',
|
||||
),
|
||||
event(
|
||||
5,
|
||||
vehicle,
|
||||
"2016-03-02 00:00:01",
|
||||
"fail",
|
||||
make="TARGET, LEAK",
|
||||
model="CURRENT ROW",
|
||||
),
|
||||
event(
|
||||
6,
|
||||
vehicle,
|
||||
"2016-03-03 00:00:01",
|
||||
"pass",
|
||||
make="LATER LEAK",
|
||||
model="LATER ATTEMPT",
|
||||
),
|
||||
]
|
||||
source = write_export(
|
||||
self.root,
|
||||
"history.csv.gz",
|
||||
rows,
|
||||
"2010-01-01T00:00:00",
|
||||
"2020-01-01T00:00:00",
|
||||
)
|
||||
summary = self.build(source)
|
||||
self.assertEqual(summary.clean_events, 5)
|
||||
self.assertEqual(summary.episodes, 3)
|
||||
self.assertFalse(summary.population_estimate_allowed)
|
||||
|
||||
connection = duckdb.connect(str(summary.database_path), read_only=True)
|
||||
try:
|
||||
first_episode = self.row_dict(
|
||||
connection,
|
||||
"""
|
||||
SELECT attempt_count, first_outcome, final_outcome
|
||||
FROM inspection_episodes
|
||||
WHERE vehicle_token = ? AND episode_number = 1
|
||||
""",
|
||||
[vehicle],
|
||||
)
|
||||
self.assertEqual(first_episode["attempt_count"], 2)
|
||||
self.assertEqual(first_episode["first_outcome"], "pass")
|
||||
self.assertEqual(first_episode["final_outcome"], "fail")
|
||||
|
||||
target = self.row_dict(
|
||||
connection,
|
||||
"SELECT * FROM feature_mart WHERE vehicle_token = ? "
|
||||
"AND episode_number = 3",
|
||||
[vehicle],
|
||||
)
|
||||
self.assertEqual(target["target_nonpass"], 1)
|
||||
self.assertEqual(
|
||||
target["target_outcome_label_source"], "overall_result"
|
||||
)
|
||||
self.assertTrue(target["eligible_returning_target"])
|
||||
self.assertEqual(target["prior_episode_count"], 2)
|
||||
self.assertEqual(target["prior_total_attempt_count"], 3)
|
||||
self.assertEqual(target["prior_attempt_count"], 1)
|
||||
self.assertEqual(target["prior_first_outcome"], "reject")
|
||||
self.assertEqual(target["prior_final_outcome"], "reject")
|
||||
self.assertEqual(target["last_observed_make"], "ACME, MOTORS")
|
||||
self.assertEqual(target["last_observed_model"], 'MODEL "Q"')
|
||||
self.assertEqual(target["vehicle_age"], 6)
|
||||
self.assertEqual(target["temporal_partition"], "train")
|
||||
|
||||
names = {
|
||||
row[0]
|
||||
for row in connection.execute(
|
||||
"DESCRIBE SELECT * FROM feature_mart"
|
||||
).fetchall()
|
||||
}
|
||||
required = {
|
||||
"vehicle_token",
|
||||
"vehicle_bucket",
|
||||
"is_vin_audit",
|
||||
"episode_number",
|
||||
"episode_start",
|
||||
"first_outcome",
|
||||
"target_nonpass",
|
||||
"eligible_returning_target",
|
||||
"temporal_partition",
|
||||
"vehicle_age",
|
||||
"prior_episode_count",
|
||||
"prior_total_attempt_count",
|
||||
"prior_attempt_count",
|
||||
"days_since_prior_episode",
|
||||
"days_since_prior_adverse",
|
||||
"prior_nonpass_rate",
|
||||
"public_county",
|
||||
"source_era",
|
||||
"target_season",
|
||||
"prior_first_outcome",
|
||||
"prior_final_outcome",
|
||||
"last_observed_make",
|
||||
"last_observed_model",
|
||||
}
|
||||
self.assertTrue(required.issubset(names))
|
||||
self.assertTrue(
|
||||
{
|
||||
"attempt_count",
|
||||
"final_outcome",
|
||||
"eventually_passed",
|
||||
"episode_end",
|
||||
"obd_result",
|
||||
}.isdisjoint(names)
|
||||
)
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
def test_conflicting_timestamp_is_quarantined_and_null_first_stays_null(self) -> None:
|
||||
conflicted = token("b")
|
||||
unlabeled = token("c")
|
||||
rows = [
|
||||
event(10, conflicted, "2014-01-01", "pass"),
|
||||
event(11, conflicted, "2016-01-01", "pass"),
|
||||
event(12, conflicted, "2016-01-01", "fail"),
|
||||
event(13, conflicted, "2017-01-01", "pass"),
|
||||
event(20, unlabeled, "2015-01-01", "pass"),
|
||||
event(21, unlabeled, "2016-01-01", None),
|
||||
event(22, unlabeled, "2016-01-02", "pass"),
|
||||
]
|
||||
source = write_export(
|
||||
self.root,
|
||||
"history.csv.gz",
|
||||
rows,
|
||||
"2010-01-01",
|
||||
"2020-01-01",
|
||||
)
|
||||
summary = self.build(source)
|
||||
connection = duckdb.connect(str(summary.database_path), read_only=True)
|
||||
try:
|
||||
conflict_count = connection.execute(
|
||||
"SELECT count(*) FROM event_exclusions "
|
||||
"WHERE exclusion_reason = 'conflicting_same_timestamp'"
|
||||
).fetchone()[0]
|
||||
self.assertEqual(conflict_count, 2)
|
||||
target = self.row_dict(
|
||||
connection,
|
||||
"SELECT first_outcome, target_nonpass, "
|
||||
"eligible_returning_target, eligibility_exclusion_reason "
|
||||
"FROM feature_mart WHERE vehicle_token = ? "
|
||||
"AND episode_start = TIMESTAMP '2016-01-01'",
|
||||
[unlabeled],
|
||||
)
|
||||
self.assertIsNone(target["first_outcome"])
|
||||
self.assertIsNone(target["target_nonpass"])
|
||||
self.assertFalse(target["eligible_returning_target"])
|
||||
self.assertEqual(
|
||||
target["eligibility_exclusion_reason"], "unlabeled_first_outcome"
|
||||
)
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
def test_bounded_files_are_combined_before_episode_construction(self) -> None:
|
||||
vehicle = token("d")
|
||||
first = write_export(
|
||||
self.root,
|
||||
"inspection_202601.csv.gz",
|
||||
[event(30, vehicle, "2026-01-31 12:00:00", "pass")],
|
||||
"2026-01-01",
|
||||
"2026-02-01",
|
||||
kind=mart.BOUNDED_KIND,
|
||||
)
|
||||
write_export(
|
||||
self.root,
|
||||
"inspection_202602.csv.gz",
|
||||
[event(31, vehicle, "2026-02-01 12:00:00", "fail")],
|
||||
"2026-02-01",
|
||||
"2026-03-01",
|
||||
kind=mart.BOUNDED_KIND,
|
||||
)
|
||||
summary = self.build(first.parent)
|
||||
self.assertTrue(summary.population_estimate_allowed)
|
||||
connection = duckdb.connect(str(summary.database_path), read_only=True)
|
||||
try:
|
||||
result = connection.execute(
|
||||
"SELECT count(*), max(attempt_count) FROM inspection_episodes"
|
||||
).fetchone()
|
||||
self.assertEqual(result, (1, 2))
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
def test_vehicle_quality_eligibility_uses_only_prior_activity(self) -> None:
|
||||
vehicle = token("e")
|
||||
rows: List[Mapping[str, object]] = [
|
||||
event(40, vehicle, "2014-01-01 00:00:00", "pass"),
|
||||
event(41, vehicle, "2016-01-01 00:00:00", "pass"),
|
||||
]
|
||||
rows.extend(
|
||||
event(42 + index, vehicle, f"2016-03-10 0{index}:00:00", "fail")
|
||||
for index in range(5)
|
||||
)
|
||||
rows.append(event(47, vehicle, "2017-04-15 00:00:00", "pass"))
|
||||
source = write_export(
|
||||
self.root,
|
||||
"history.csv.gz",
|
||||
rows,
|
||||
"2010-01-01",
|
||||
"2020-01-01",
|
||||
)
|
||||
summary = self.build(source)
|
||||
connection = duckdb.connect(str(summary.database_path), read_only=True)
|
||||
try:
|
||||
current_busy_episode = self.row_dict(
|
||||
connection,
|
||||
"SELECT eligible_returning_target, prior_max_events_in_day "
|
||||
"FROM feature_mart WHERE vehicle_token = ? "
|
||||
"AND episode_start = TIMESTAMP '2016-03-10 00:00:00'",
|
||||
[vehicle],
|
||||
)
|
||||
self.assertTrue(current_busy_episode["eligible_returning_target"])
|
||||
self.assertEqual(current_busy_episode["prior_max_events_in_day"], 1)
|
||||
|
||||
after_busy_episode = self.row_dict(
|
||||
connection,
|
||||
"SELECT eligible_returning_target, prior_max_events_in_day, "
|
||||
"eligibility_exclusion_reason FROM feature_mart "
|
||||
"WHERE vehicle_token = ? "
|
||||
"AND episode_start = TIMESTAMP '2017-04-15 00:00:00'",
|
||||
[vehicle],
|
||||
)
|
||||
self.assertFalse(after_busy_episode["eligible_returning_target"])
|
||||
self.assertEqual(after_busy_episode["prior_max_events_in_day"], 5)
|
||||
self.assertEqual(
|
||||
after_busy_episode["eligibility_exclusion_reason"],
|
||||
"prior_daily_activity_over_4",
|
||||
)
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
def test_digest_mismatch_fails_before_build(self) -> None:
|
||||
source = write_export(
|
||||
self.root,
|
||||
"history.csv.gz",
|
||||
[event(50, token("f"), "2016-01-01", "pass")],
|
||||
"2010-01-01",
|
||||
"2020-01-01",
|
||||
)
|
||||
manifest_path = source.with_name(source.name + ".manifest.json")
|
||||
manifest = json.loads(manifest_path.read_text(encoding="utf-8"))
|
||||
manifest["compressed_file_sha256"] = "0" * 64
|
||||
manifest_path.write_text(json.dumps(manifest), encoding="utf-8")
|
||||
|
||||
with self.assertRaisesRegex(mart.MartBuildError, "SHA-256 mismatch"):
|
||||
self.build(source)
|
||||
|
||||
def test_outcome_label_source_is_validated_and_first_row_lineage_is_audit_only(
|
||||
self,
|
||||
) -> None:
|
||||
utah_vehicle = token("7")
|
||||
rows = [
|
||||
event(
|
||||
70,
|
||||
utah_vehicle,
|
||||
"2015-01-01",
|
||||
"pass",
|
||||
source_era="utah",
|
||||
public_county="utah",
|
||||
label_source="utah_obd_proxy",
|
||||
),
|
||||
event(
|
||||
71,
|
||||
utah_vehicle,
|
||||
"2016-01-01",
|
||||
"fail",
|
||||
source_era="utah",
|
||||
public_county="utah",
|
||||
label_source="utah_obd_proxy",
|
||||
),
|
||||
event(
|
||||
72,
|
||||
token("8"),
|
||||
"2016-01-01",
|
||||
"pass",
|
||||
source_era="slc",
|
||||
public_county="salt_lake",
|
||||
label_source="utah_obd_proxy",
|
||||
),
|
||||
event(
|
||||
73,
|
||||
token("9"),
|
||||
"2016-01-01",
|
||||
"pass",
|
||||
label_source=None,
|
||||
),
|
||||
event(
|
||||
74,
|
||||
token("0"),
|
||||
"2016-01-01",
|
||||
"pass",
|
||||
label_source="unknown_source",
|
||||
),
|
||||
]
|
||||
source = write_export(
|
||||
self.root,
|
||||
"history.csv.gz",
|
||||
rows,
|
||||
"2010-01-01",
|
||||
"2020-01-01",
|
||||
)
|
||||
summary = self.build(source)
|
||||
connection = duckdb.connect(str(summary.database_path), read_only=True)
|
||||
try:
|
||||
target = self.row_dict(
|
||||
connection,
|
||||
"SELECT target_outcome_label_source, source_era "
|
||||
"FROM feature_mart WHERE vehicle_token = ? "
|
||||
"AND episode_number = 2",
|
||||
[utah_vehicle],
|
||||
)
|
||||
self.assertEqual(
|
||||
target["target_outcome_label_source"], "utah_obd_proxy"
|
||||
)
|
||||
self.assertEqual(target["source_era"], "utah")
|
||||
|
||||
reasons = dict(
|
||||
connection.execute(
|
||||
"SELECT exclusion_reason, count(*) FROM event_exclusions "
|
||||
"GROUP BY exclusion_reason"
|
||||
).fetchall()
|
||||
)
|
||||
self.assertEqual(reasons["utah_proxy_non_utah_source"], 1)
|
||||
self.assertEqual(reasons["missing_outcome_label_source"], 1)
|
||||
self.assertEqual(reasons["invalid_outcome_label_source"], 1)
|
||||
|
||||
schema_names = {
|
||||
row[0]
|
||||
for row in connection.execute(
|
||||
"DESCRIBE SELECT * FROM feature_mart"
|
||||
).fetchall()
|
||||
}
|
||||
self.assertIn("target_outcome_label_source", schema_names)
|
||||
self.assertNotIn("outcome_label_source", schema_names)
|
||||
self.assertNotIn("obd_result", schema_names)
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
def test_bounded_gap_fails_closed(self) -> None:
|
||||
vehicle = token("1")
|
||||
write_export(
|
||||
self.root,
|
||||
"one.csv.gz",
|
||||
[event(60, vehicle, "2026-01-15", "pass")],
|
||||
"2026-01-01",
|
||||
"2026-02-01",
|
||||
kind=mart.BOUNDED_KIND,
|
||||
)
|
||||
write_export(
|
||||
self.root,
|
||||
"two.csv.gz",
|
||||
[event(61, vehicle, "2026-03-15", "pass")],
|
||||
"2026-03-01",
|
||||
"2026-04-01",
|
||||
kind=mart.BOUNDED_KIND,
|
||||
)
|
||||
with self.assertRaisesRegex(mart.MartBuildError, "interval gap"):
|
||||
self.build(self.root)
|
||||
|
||||
incomplete = mart.build_feature_mart(
|
||||
mart.BuildConfig(
|
||||
input_path=self.root,
|
||||
database_path=self.root / "incomplete.duckdb",
|
||||
output_path=self.root / "incomplete.parquet",
|
||||
memory_limit="512MB",
|
||||
threads=1,
|
||||
allow_gaps=True,
|
||||
allow_external_output=True,
|
||||
)
|
||||
)
|
||||
self.assertFalse(incomplete.population_estimate_allowed)
|
||||
build_manifest = json.loads(
|
||||
incomplete.manifest_path.read_text(encoding="utf-8")
|
||||
)
|
||||
self.assertFalse(build_manifest["population_estimate_allowed"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,299 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import hmac
|
||||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
import uuid
|
||||
from datetime import datetime, timedelta
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest import mock
|
||||
|
||||
from scripts import export_history_sample as history_export
|
||||
from scripts import export_inspection_batch as exporter
|
||||
|
||||
|
||||
VALID_VIN = "1ABCD23EFGH456789"
|
||||
TEST_KEY = bytes(range(32))
|
||||
|
||||
|
||||
class PrivateExportTests(unittest.TestCase):
|
||||
def test_normalize_vin_canonicalizes_case_and_whitespace(self) -> None:
|
||||
self.assertEqual(
|
||||
exporter.normalize_vin(f" {VALID_VIN.lower()}\n"),
|
||||
VALID_VIN,
|
||||
)
|
||||
|
||||
def test_normalize_vin_rejects_invalid_values_and_placeholders(self) -> None:
|
||||
invalid_vins = (
|
||||
"",
|
||||
"1ABCD23EFGH45678",
|
||||
"1ABCD23EFGH4567890",
|
||||
"1ABCD23EFGI456789",
|
||||
"11111111111111111",
|
||||
"AAAAAAAAAAAAAAAAA",
|
||||
"12345678901234567",
|
||||
"98765432109876543",
|
||||
)
|
||||
for vin in invalid_vins:
|
||||
with self.subTest(vin=vin):
|
||||
self.assertIsNone(exporter.normalize_vin(vin))
|
||||
|
||||
def test_hmac_is_deterministic_and_version_domain_separated(self) -> None:
|
||||
expected_v1 = hmac.new(
|
||||
TEST_KEY,
|
||||
b"utah-vehicle-health/v1/vin\0" + VALID_VIN.encode("ascii"),
|
||||
hashlib.sha256,
|
||||
).hexdigest()
|
||||
|
||||
token_v1 = exporter.vehicle_token(VALID_VIN, TEST_KEY, "v1")
|
||||
repeated_v1 = exporter.vehicle_token(VALID_VIN, TEST_KEY, "v1")
|
||||
token_v2 = exporter.vehicle_token(VALID_VIN, TEST_KEY, "v2")
|
||||
|
||||
self.assertEqual(token_v1, expected_v1)
|
||||
self.assertEqual(repeated_v1, token_v1)
|
||||
self.assertNotEqual(token_v2, token_v1)
|
||||
self.assertEqual(len(token_v1), 64)
|
||||
|
||||
def test_recognized_overall_result_precedes_utah_obd_proxy(self) -> None:
|
||||
self.assertEqual(
|
||||
exporter.canonicalize_outcome(
|
||||
"PASS", "FAIL", "utah", "other", "C"
|
||||
),
|
||||
("pass", "overall_result"),
|
||||
)
|
||||
self.assertEqual(
|
||||
exporter.canonicalize_outcome(
|
||||
" f ", "PASS", " UTAH ", "tsi", "TSI"
|
||||
),
|
||||
("fail", "overall_result"),
|
||||
)
|
||||
|
||||
def test_utah_only_obd_fallback_uses_recognized_values(self) -> None:
|
||||
self.assertEqual(
|
||||
exporter.canonicalize_outcome(
|
||||
"", " P ", "utah", " ObD ", " obd "
|
||||
),
|
||||
("pass", "utah_obd_proxy"),
|
||||
)
|
||||
self.assertEqual(
|
||||
exporter.canonicalize_outcome(
|
||||
"B", "REJECT", "UTAH", "obd", "OBD"
|
||||
),
|
||||
("reject", "utah_obd_proxy"),
|
||||
)
|
||||
self.assertEqual(
|
||||
exporter.canonicalize_outcome(
|
||||
None, "FAIL", "weber", "obd", "OBD"
|
||||
),
|
||||
(None, None),
|
||||
)
|
||||
|
||||
def test_utah_non_obd_program_or_test_rows_do_not_use_proxy(self) -> None:
|
||||
excluded_cases = (
|
||||
("other", "C"),
|
||||
("tsi", "TSI"),
|
||||
("obd", "C"),
|
||||
("other", "OBD"),
|
||||
("", ""),
|
||||
(None, None),
|
||||
)
|
||||
for program_type, test_type in excluded_cases:
|
||||
with self.subTest(program_type=program_type, test_type=test_type):
|
||||
self.assertEqual(
|
||||
exporter.canonicalize_outcome(
|
||||
"B", "PASS", "utah", program_type, test_type
|
||||
),
|
||||
(None, None),
|
||||
)
|
||||
|
||||
def test_unknown_blank_and_b_proxy_values_remain_unlabeled(self) -> None:
|
||||
cases = (
|
||||
(None, None, "utah", "obd", "OBD"),
|
||||
("", "", "utah", "obd", "OBD"),
|
||||
("B", "B", "utah", "obd", "OBD"),
|
||||
("UNKNOWN", "PASS", "slco", "obd", "OBD"),
|
||||
)
|
||||
for overall, obd, source, program_type, test_type in cases:
|
||||
with self.subTest(
|
||||
overall=overall,
|
||||
obd=obd,
|
||||
source=source,
|
||||
program_type=program_type,
|
||||
test_type=test_type,
|
||||
):
|
||||
self.assertEqual(
|
||||
exporter.canonicalize_outcome(
|
||||
overall, obd, source, program_type, test_type
|
||||
),
|
||||
(None, None),
|
||||
)
|
||||
|
||||
def test_hmac_configuration_rejects_malformed_or_weak_keys(self) -> None:
|
||||
invalid_keys = (
|
||||
"",
|
||||
"not-hex",
|
||||
"00" * 31,
|
||||
"00" * 33,
|
||||
"00" * 32,
|
||||
"ab" * 32,
|
||||
)
|
||||
for encoded_key in invalid_keys:
|
||||
with self.subTest(encoded_key=encoded_key[:12]):
|
||||
with self._hmac_environment(encoded_key):
|
||||
with self.assertRaises(ValueError):
|
||||
exporter.get_hmac_configuration()
|
||||
|
||||
def test_hmac_configuration_accepts_exact_random_key(self) -> None:
|
||||
with self._hmac_environment(TEST_KEY.hex(), version="key-2026.1"):
|
||||
key, version, fingerprint = exporter.get_hmac_configuration()
|
||||
|
||||
self.assertEqual(key, TEST_KEY)
|
||||
self.assertEqual(version, "key-2026.1")
|
||||
self.assertRegex(fingerprint, r"^[0-9a-f]{16}$")
|
||||
|
||||
def test_vehicle_bucket_is_stable_after_vin_normalization(self) -> None:
|
||||
row = self._source_row(VALID_VIN)
|
||||
normalized_row = exporter.private_row(row, TEST_KEY, "v1")
|
||||
lower_row = exporter.private_row(
|
||||
self._source_row(f" {VALID_VIN.lower()} "),
|
||||
TEST_KEY,
|
||||
"v1",
|
||||
)
|
||||
|
||||
self.assertIsNotNone(normalized_row)
|
||||
self.assertIsNotNone(lower_row)
|
||||
assert normalized_row is not None
|
||||
assert lower_row is not None
|
||||
token = normalized_row[1]
|
||||
bucket = normalized_row[2]
|
||||
self.assertEqual(lower_row[1], token)
|
||||
self.assertEqual(lower_row[2], bucket)
|
||||
self.assertEqual(bucket, int(str(token)[:8], 16) % 100)
|
||||
self.assertIn(bucket, range(100))
|
||||
|
||||
def test_validate_args_enforces_private_output_root(self) -> None:
|
||||
private_output = (
|
||||
exporter.PRIVATE_DATA_ROOT
|
||||
/ "inspection_batches"
|
||||
/ f"unit-{uuid.uuid4().hex}.csv.gz"
|
||||
)
|
||||
validated = exporter.validate_args(self._args(private_output))
|
||||
self.assertEqual(validated, private_output.resolve())
|
||||
|
||||
with tempfile.TemporaryDirectory() as temporary_directory:
|
||||
external_output = Path(temporary_directory) / "private.csv.gz"
|
||||
with self.assertRaises(ValueError):
|
||||
exporter.validate_args(self._args(external_output))
|
||||
|
||||
allowed = exporter.validate_args(
|
||||
self._args(external_output, allow_external_output=True)
|
||||
)
|
||||
self.assertEqual(allowed, external_output.resolve())
|
||||
|
||||
def test_declared_output_omits_raw_vin_and_raw_obd_result(self) -> None:
|
||||
self.assertNotIn("vin", exporter.OUTPUT_COLUMNS)
|
||||
self.assertNotIn("raw_vin", exporter.OUTPUT_COLUMNS)
|
||||
self.assertNotIn("obd_result", exporter.OUTPUT_COLUMNS)
|
||||
self.assertNotIn("overall_result", exporter.OUTPUT_COLUMNS)
|
||||
self.assertIn("vehicle_token", exporter.OUTPUT_COLUMNS)
|
||||
self.assertIn("outcome_label_source", exporter.OUTPUT_COLUMNS)
|
||||
|
||||
transformed = exporter.private_row(
|
||||
self._source_row(VALID_VIN), TEST_KEY, "v1"
|
||||
)
|
||||
self.assertIsNotNone(transformed)
|
||||
assert transformed is not None
|
||||
self.assertEqual(len(transformed), len(exporter.OUTPUT_COLUMNS))
|
||||
self.assertEqual(
|
||||
transformed[exporter.OUTPUT_COLUMNS.index("outcome_label_source")],
|
||||
"overall_result",
|
||||
)
|
||||
self.assertNotIn(VALID_VIN, transformed)
|
||||
|
||||
def test_both_exporters_share_the_versioned_label_sql(self) -> None:
|
||||
history_sql = history_export.make_sql(0.5, 20260715)
|
||||
|
||||
self.assertIn(exporter.OUTCOME_SELECT_SQL, exporter.EXTRACT_SQL)
|
||||
self.assertIn(exporter.OUTCOME_SELECT_SQL, history_sql)
|
||||
self.assertIn("s.obd_result", exporter.OUTCOME_SELECT_SQL)
|
||||
self.assertIn(
|
||||
"lower(btrim(s.program_type)) = 'obd'",
|
||||
exporter.OUTCOME_SELECT_SQL,
|
||||
)
|
||||
self.assertIn(
|
||||
"upper(btrim(s.test_type)) = 'OBD'",
|
||||
exporter.OUTCOME_SELECT_SQL,
|
||||
)
|
||||
self.assertNotIn("WHEN 'B'", exporter.OUTCOME_SELECT_SQL)
|
||||
self.assertEqual(exporter.QUERY_VERSION, "inspection_batch_v4")
|
||||
self.assertEqual(
|
||||
history_export.QUERY_VERSION,
|
||||
"inspection_history_development_sample_v3",
|
||||
)
|
||||
|
||||
def test_feasibility_sql_uses_labeled_utah_only_proxy(self) -> None:
|
||||
sql_path = exporter.PROJECT_ROOT / "sql/10_episode_cohort_feasibility.sql"
|
||||
sql = sql_path.read_text(encoding="utf-8")
|
||||
|
||||
self.assertIn("lower(btrim(s.county)) = 'utah'", sql)
|
||||
self.assertIn("lower(btrim(s.program_type)) = 'obd'", sql)
|
||||
self.assertIn("upper(btrim(s.test_type)) = 'OBD'", sql)
|
||||
self.assertIn("upper(btrim(s.obd_result))", sql)
|
||||
self.assertIn("AS outcome_label_source", sql)
|
||||
self.assertIn("'utah_obd_proxy'", sql)
|
||||
self.assertNotIn("WHEN 'B' THEN", sql)
|
||||
|
||||
@staticmethod
|
||||
def _source_row(vin: str) -> tuple[object, ...]:
|
||||
return (
|
||||
123,
|
||||
vin,
|
||||
datetime(2025, 1, 15, 12, 0),
|
||||
"slco",
|
||||
"salt_lake",
|
||||
"pass",
|
||||
"overall_result",
|
||||
"obd",
|
||||
"INITIAL",
|
||||
"EXAMPLE",
|
||||
"MODEL",
|
||||
2020,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _args(
|
||||
output: Path, allow_external_output: bool = False
|
||||
) -> SimpleNamespace:
|
||||
start = datetime(2025, 1, 1)
|
||||
return SimpleNamespace(
|
||||
start=start,
|
||||
end=start + timedelta(days=1),
|
||||
output=output,
|
||||
page_size=100,
|
||||
overwrite=False,
|
||||
allow_external_output=allow_external_output,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _hmac_environment(
|
||||
encoded_key: str, version: str = "v1"
|
||||
) -> mock._patch_dict:
|
||||
missing_key_file = Path(tempfile.gettempdir()) / (
|
||||
f"uvh-missing-key-{uuid.uuid4().hex}"
|
||||
)
|
||||
return mock.patch.dict(
|
||||
os.environ,
|
||||
{
|
||||
"VIN_HASH_KEY": encoded_key,
|
||||
"VIN_HASH_KEY_FILE": str(missing_key_file),
|
||||
"VIN_HASH_KEY_VERSION": version,
|
||||
},
|
||||
clear=False,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,261 @@
|
||||
"""Tests for the leakage-safe nonlinear candidate."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
import math
|
||||
import sys
|
||||
import unittest
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from typing import Dict, List
|
||||
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
||||
MODULE_PATH = PROJECT_ROOT / "scripts/train_tree_model.py"
|
||||
SPEC = importlib.util.spec_from_file_location("train_tree_model", MODULE_PATH)
|
||||
if SPEC is None or SPEC.loader is None: # pragma: no cover
|
||||
raise RuntimeError("Could not load scripts/train_tree_model.py")
|
||||
tree = importlib.util.module_from_spec(SPEC)
|
||||
sys.modules[SPEC.name] = tree
|
||||
SPEC.loader.exec_module(tree)
|
||||
baselines = tree.baselines
|
||||
|
||||
SKLEARN_AVAILABLE = (
|
||||
importlib.util.find_spec("numpy") is not None
|
||||
and importlib.util.find_spec("sklearn") is not None
|
||||
)
|
||||
|
||||
|
||||
def episode(
|
||||
number: int,
|
||||
partition: str,
|
||||
year: int,
|
||||
target: int,
|
||||
age: float,
|
||||
make: str,
|
||||
county: str = "salt_lake",
|
||||
audit: bool = False,
|
||||
) -> baselines.EpisodeRow:
|
||||
bucket = 5 if audit else 25 + (number % 70)
|
||||
numeric: Dict[str, float] = {
|
||||
"vehicle_age": age,
|
||||
"prior_episode_count": 1.0 + number % 8,
|
||||
"prior_total_attempt_count": 1.0 + number % 12,
|
||||
"prior_attempt_count": 1.0 + number % 3,
|
||||
"days_since_prior_episode": 300.0 + number % 150,
|
||||
"days_since_prior_adverse": 200.0 + number % 500,
|
||||
"prior_nonpass_rate": (number % 10) / 10.0,
|
||||
}
|
||||
categorical = {
|
||||
"public_county": county,
|
||||
"target_season": ("summer" if number % 2 else "winter"),
|
||||
"prior_first_outcome": ("fail" if number % 4 == 0 else "pass"),
|
||||
"prior_final_outcome": ("fail" if number % 5 == 0 else "pass"),
|
||||
"last_observed_make": make,
|
||||
"last_observed_model": ("model_a" if number % 3 else "model_b"),
|
||||
}
|
||||
return baselines.EpisodeRow(
|
||||
vehicle_token="{:064x}".format(number + year * 100_000),
|
||||
vehicle_bucket=bucket,
|
||||
is_vin_audit=audit,
|
||||
episode_number=2 + number % 10,
|
||||
episode_start=datetime(year, 6, 1),
|
||||
partition=partition,
|
||||
target=target,
|
||||
source_era="slc",
|
||||
target_outcome_label_source="overall_result",
|
||||
prior_first_outcome=categorical["prior_first_outcome"],
|
||||
numeric=numeric,
|
||||
categorical=categorical,
|
||||
)
|
||||
|
||||
|
||||
def partition_rows(partition: str, year: int, count: int, offset: int) -> List[object]:
|
||||
rows: List[object] = []
|
||||
for index in range(count):
|
||||
number = offset + index
|
||||
audit = index % 10 == 0
|
||||
age = float((number * 7) % 30)
|
||||
county = "utah" if number % 2 else "salt_lake"
|
||||
signal = (age > 14.0) != (county == "utah")
|
||||
target = int(signal)
|
||||
if number % 17 == 0:
|
||||
target = 1 - target
|
||||
if audit:
|
||||
make = "audit_only_make"
|
||||
elif partition == "train":
|
||||
make = ("toyota" if number % 3 else "ford")
|
||||
else:
|
||||
make = ("future_make" if number % 7 == 0 else "toyota")
|
||||
rows.append(
|
||||
episode(
|
||||
number=number,
|
||||
partition=partition,
|
||||
year=year,
|
||||
target=target,
|
||||
age=age,
|
||||
make=make,
|
||||
county=county,
|
||||
audit=audit,
|
||||
)
|
||||
)
|
||||
return rows
|
||||
|
||||
|
||||
def synthetic_mart() -> baselines.MartData:
|
||||
rows_by_partition = {
|
||||
"train": partition_rows("train", 2022, 360, 0),
|
||||
"tune": partition_rows("tune", 2023, 140, 10_000),
|
||||
"calibrate": partition_rows("calibrate", 2024, 140, 20_000),
|
||||
"locked_test": partition_rows("locked_test", 2025, 140, 30_000),
|
||||
}
|
||||
total = sum(len(rows) for rows in rows_by_partition.values())
|
||||
return baselines.MartData(
|
||||
rows_by_partition=rows_by_partition,
|
||||
input_rows=total,
|
||||
ineligible_rows=0,
|
||||
eligible_rows=total,
|
||||
)
|
||||
|
||||
|
||||
class TreeEncoderTests(unittest.TestCase):
|
||||
def test_future_categories_and_values_do_not_change_training_preprocessing(self) -> None:
|
||||
train_rows = [
|
||||
episode(1, "train", 2022, 0, 10.0, "toyota"),
|
||||
episode(2, "train", 2022, 1, 20.0, "toyota"),
|
||||
episode(3, "train", 2022, 0, 30.0, "ford"),
|
||||
]
|
||||
future = episode(4, "tune", 2023, 1, 9_999.0, "future_make")
|
||||
encoder = tree.TreeFeatureEncoder.fit(
|
||||
train_rows, min_category_count=2, max_categories=16
|
||||
)
|
||||
|
||||
self.assertEqual(encoder.numeric_medians["vehicle_age"], 20.0)
|
||||
self.assertEqual(
|
||||
encoder.category_code("last_observed_make", "toyota"),
|
||||
tree.FIRST_KNOWN_CATEGORY_CODE,
|
||||
)
|
||||
self.assertEqual(
|
||||
encoder.category_code("last_observed_make", "ford"),
|
||||
tree.RARE_CATEGORY_CODE,
|
||||
)
|
||||
self.assertEqual(
|
||||
encoder.category_code("last_observed_make", "future_make"),
|
||||
tree.UNKNOWN_CATEGORY_CODE,
|
||||
)
|
||||
self.assertNotIn(
|
||||
"future_make", encoder.seen_categories["last_observed_make"]
|
||||
)
|
||||
self.assertEqual(encoder.numeric_medians["vehicle_age"], 20.0)
|
||||
self.assertTrue(set(tree.EXCLUDED_FROM_PREDICTORS).isdisjoint(
|
||||
encoder.feature_names
|
||||
))
|
||||
self.assertNotIn("source_era", encoder.feature_names)
|
||||
self.assertNotIn("target_outcome_label_source", encoder.feature_names)
|
||||
|
||||
if SKLEARN_AVAILABLE:
|
||||
import numpy as np
|
||||
|
||||
matrix = encoder.transform([future], np)
|
||||
category_index = 2 * len(tree.NUMERIC_FEATURES) + list(
|
||||
tree.CATEGORICAL_FEATURES
|
||||
).index("last_observed_make")
|
||||
self.assertEqual(
|
||||
matrix[0, category_index], float(tree.UNKNOWN_CATEGORY_CODE)
|
||||
)
|
||||
self.assertTrue(np.isfinite(matrix).all())
|
||||
|
||||
@unittest.skipUnless(SKLEARN_AVAILABLE, "scikit-learn is not installed")
|
||||
def test_convergence_gate_rejects_iteration_limit_and_nonfinite_scores(self) -> None:
|
||||
import numpy as np
|
||||
|
||||
at_limit = SimpleNamespace(
|
||||
n_iter_=100,
|
||||
train_score_=np.asarray([-1.0]),
|
||||
validation_score_=np.asarray([-1.0]),
|
||||
)
|
||||
converged, n_iter = tree._hist_converged(at_limit, [], 100, np)
|
||||
self.assertFalse(converged)
|
||||
self.assertEqual(n_iter, 100)
|
||||
|
||||
nonfinite = SimpleNamespace(
|
||||
n_iter_=20,
|
||||
train_score_=np.asarray([-1.0, math.nan]),
|
||||
validation_score_=np.asarray([-1.0]),
|
||||
)
|
||||
converged, _ = tree._hist_converged(nonfinite, [], 100, np)
|
||||
self.assertFalse(converged)
|
||||
|
||||
|
||||
@unittest.skipUnless(SKLEARN_AVAILABLE, "scikit-learn is not installed")
|
||||
class TreeTrainingTests(unittest.TestCase):
|
||||
@classmethod
|
||||
def setUpClass(cls) -> None:
|
||||
cls.mart = synthetic_mart()
|
||||
cls.trained = tree.train_tree_model(
|
||||
mart=cls.mart,
|
||||
candidates=[tree.TreeCandidate(0.1, 15, 1.0)],
|
||||
min_category_count=3,
|
||||
max_categories=32,
|
||||
min_samples_leaf=10,
|
||||
max_iter=150,
|
||||
n_iter_no_change=8,
|
||||
calibration_max_iter=1000,
|
||||
seed=20260715,
|
||||
)
|
||||
|
||||
def test_fit_excludes_audit_and_future_categories(self) -> None:
|
||||
seen = self.trained.encoder.seen_categories["last_observed_make"]
|
||||
self.assertNotIn("audit_only_make", seen)
|
||||
self.assertNotIn("future_make", seen)
|
||||
self.assertLess(self.trained.model_n_iter, self.trained.max_iter)
|
||||
self.assertTrue(
|
||||
all(item["eligible_for_selection"] for item in self.trained.tuning_results)
|
||||
)
|
||||
|
||||
def test_locked_metrics_require_explicit_gate(self) -> None:
|
||||
closed_metrics, _closed_bins = tree.evaluate_tree_model(
|
||||
self.mart,
|
||||
self.trained,
|
||||
evaluate_locked=False,
|
||||
bin_count=5,
|
||||
)
|
||||
self.assertNotIn(
|
||||
"locked_test", {row["partition"] for row in closed_metrics}
|
||||
)
|
||||
|
||||
open_metrics, _open_bins = tree.evaluate_tree_model(
|
||||
self.mart,
|
||||
self.trained,
|
||||
evaluate_locked=True,
|
||||
bin_count=5,
|
||||
)
|
||||
locked = [
|
||||
row for row in open_metrics if row["partition"] == "locked_test"
|
||||
]
|
||||
self.assertTrue(locked)
|
||||
self.assertIn("never_fit_audit", {row["cohort"] for row in locked})
|
||||
|
||||
def test_probabilities_are_finite(self) -> None:
|
||||
probabilities = tree.tree_probabilities(
|
||||
self.trained,
|
||||
self.mart.rows_by_partition["calibrate"],
|
||||
calibrated=True,
|
||||
)
|
||||
self.assertTrue(probabilities)
|
||||
self.assertTrue(all(math.isfinite(value) for value in probabilities))
|
||||
self.assertTrue(all(0.0 <= value <= 1.0 for value in probabilities))
|
||||
|
||||
def test_cli_locked_flag_defaults_closed(self) -> None:
|
||||
closed = tree.parse_args(["--mart", "data/private/mart.parquet"])
|
||||
opened = tree.parse_args(
|
||||
["--mart", "data/private/mart.parquet", "--evaluate-locked"]
|
||||
)
|
||||
self.assertFalse(closed.evaluate_locked)
|
||||
self.assertTrue(opened.evaluate_locked)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user