"""Smoke test for the H1 verification fixture (synthetic; no real-data dependency).""" from __future__ import annotations import numpy as np import pandas as pd from dsl import Cohort from validate.h1 import IMMUNE_GENES, MMR_GENES, run_h1 def _synthetic_named_cohort(n_per_group: int = 60, seed: int = 0) -> Cohort: """A synthetic cohort where MSI-H tumours have low MMR and high immune expression.""" rng = np.random.default_rng(seed) other_genes = [f"GENE_{i}" for i in range(20)] feature_cols = MMR_GENES + IMMUNE_GENES + other_genes rows = [] sample_ids = [] msi = [] stages = [] ages = [] for i in range(n_per_group): row = rng.normal(loc=5.0, scale=0.5, size=len(feature_cols)) for g in MMR_GENES: row[feature_cols.index(g)] = rng.normal(2.0, 0.3) for g in IMMUNE_GENES: row[feature_cols.index(g)] = rng.normal(8.0, 0.3) rows.append(row) sample_ids.append(f"H{i}") msi.append("MSI-H") stages.append(rng.choice(["I", "II", "III", "IV"])) ages.append(float(rng.integers(40, 85))) for i in range(n_per_group): row = rng.normal(loc=5.0, scale=0.5, size=len(feature_cols)) for g in MMR_GENES: row[feature_cols.index(g)] = rng.normal(8.0, 0.3) for g in IMMUNE_GENES: row[feature_cols.index(g)] = rng.normal(2.0, 0.3) rows.append(row) sample_ids.append(f"S{i}") msi.append("MSS") stages.append(rng.choice(["I", "II", "III", "IV"])) ages.append(float(rng.integers(40, 85))) idx = pd.Index(sample_ids, name="sample_id") expr = pd.DataFrame(rows, index=idx, columns=feature_cols) clinical = pd.DataFrame({"stage": stages, "age": ages}, index=idx) labels = pd.DataFrame({"msi_status": msi, "tmb": [np.nan] * len(idx)}, index=idx) return Cohort(expression=expr, clinical=clinical, labels=labels) def test_h1_recovers_known_direction_on_synthetic_data(): cohort = _synthetic_named_cohort() res = run_h1(cohort) assert res.mmr_separates_correct_direction # MSI-H low on MMR assert res.immune_separates_correct_direction # MSI-H high on immune # With a clean synthetic signal Fit should crush it. assert res.fit.auroc > 0.95 assert res.fit.balanced_acc > 0.9 # n_used is the full synthetic cohort (no missing stage/age). assert res.n_used == 120