"""Tests for the Learn-phase "wow" additions to DE round 2 (dee.server._de_round2_library): pool_deltas (the before/after re-rank + plain-language reasons) and global_prior (cross-user aggregate surfacing). Exercises the real top_percentile_pool / evolve / variants_to_dataframe / active_learning / aggregate pipeline on a tiny synthetic pool — only the ESM-2 scorer itself and the stored cross-user aggregate are mocked, so this proves the new wiring (not just the isolated math already covered by test_active_learning.py / test_aggregate.py).""" import pandas as pd import pytest from dee import server from dee.core.aggregate import GlobalPrior _SETTINGS = { "model": "small", "host": "e_coli", "percentile": 85.0, "k": 5, "min_mutations": 1, "max_mutations": 2, "restarts": 1, "steps": 50, "seed": 1, } def _fake_scores_df(n=5): # wt_aa is always 'A' -> labels are A1G, A2G, ... A5G. return pd.DataFrame({ "position": list(range(n)), "wt_aa": ["A"] * n, "mut_aa": ["G"] * n, "delta_ll": [float(i) - 2 for i in range(n)], # -2, -1, 0, 1, 2 }) def _label(i): return f"A{i + 1}G" @pytest.fixture(autouse=True) def _mock_scorer(monkeypatch): monkeypatch.setattr(server._scoring, "get_scorer", lambda *a, **kw: "the-scorer") monkeypatch.setattr(server._scoring, "score_guarded", lambda scorer, protein: _fake_scores_df()) monkeypatch.setattr(server, "top_percentile_pool", lambda df, percentile: df) def test_pool_deltas_shape_sorted_and_reasoned_when_learned(monkeypatch): monkeypatch.setattr(server, "_load_global_prior", lambda: GlobalPrior(effects={}, n_users={}, n_obs={})) # Enough varied measurements to clear MIN_MEASUREMENTS with real spread. measurements = [ ([_label(0)], 1.0), ([_label(1)], 3.0), ([_label(2)], 5.0), ([_label(3)], 7.0), ([_label(4)], 9.0), ] rows, info = server._de_round2_library("A" * 5, _SETTINGS, measurements) assert info["learned"] is True deltas = info["pool_deltas"] assert 1 <= len(deltas) <= server._POOL_DELTAS_MAX # Sorted descending by adjusted_score. scores = [d["adjusted_score"] for d in deltas] assert scores == sorted(scores, reverse=True) for d in deltas: assert set(d.keys()) == {"label", "prior_score", "adjusted_score", "delta", "n_measured", "reason"} assert isinstance(d["reason"], str) and d["reason"] assert d["delta"] == pytest.approx(d["adjusted_score"] - d["prior_score"]) # Every measured mutation should show n_measured >= 1. measured_labels = {_label(i) for i in range(5)} assert all(d["n_measured"] >= 1 for d in deltas if d["label"] in measured_labels) assert rows # a real variant table came back def test_pool_deltas_all_zero_delta_when_not_enough_signal(monkeypatch): monkeypatch.setattr(server, "_load_global_prior", lambda: GlobalPrior(effects={}, n_users={}, n_obs={})) # Only 1 measurement — below MIN_MEASUREMENTS -> honest fallback, no change. rows, info = server._de_round2_library("A" * 5, _SETTINGS, [([_label(0)], 3.0)]) assert info["learned"] is False assert all(d["delta"] == 0.0 for d in info["pool_deltas"]) assert rows def test_global_prior_absent_reports_not_applied(monkeypatch): monkeypatch.setattr(server, "_load_global_prior", lambda: GlobalPrior(effects={}, n_users={}, n_obs={})) _, info = server._de_round2_library("A" * 5, _SETTINGS, [([_label(0)], 3.0)]) assert info["global_prior"] == {"applied": False, "substitution_types": 0} def test_global_prior_present_blends_and_reports_applied(monkeypatch): # A field-wide prior that says A>G substitutions tend to be strongly # favorable — should nudge prior_score upward vs the no-prior case, and # be honestly reported (substitution_types == 1, the one key present). gp = GlobalPrior(effects={("A", "G"): 2.0}, n_users={("A", "G"): 5}, n_obs={("A", "G"): 12}) monkeypatch.setattr(server, "_load_global_prior", lambda: gp) _, info = server._de_round2_library("A" * 5, _SETTINGS, [([_label(0)], 3.0)]) assert info["global_prior"] == {"applied": True, "substitution_types": 1} # Not enough measurements to learn, but the global-prior nudge still shows # up in prior_score (baseline moved even though round 2 fell back). unpatched_gp = GlobalPrior(effects={}, n_users={}, n_obs={}) monkeypatch.setattr(server, "_load_global_prior", lambda: unpatched_gp) _, info_no_gp = server._de_round2_library("A" * 5, _SETTINGS, [([_label(0)], 3.0)]) by_label = {d["label"]: d["prior_score"] for d in info["pool_deltas"]} by_label_no_gp = {d["label"]: d["prior_score"] for d in info_no_gp["pool_deltas"]} assert by_label[_label(0)] > by_label_no_gp[_label(0)] def test_round2_route_exposes_pool_deltas_and_global_prior(monkeypatch): import types app = server.create_app() app.config.update(TESTING=True) client = app.test_client() monkeypatch.setattr(server._auth, "get_auth", lambda: types.SimpleNamespace(anonymous=False, user_id="u1", email="x@y.z", plan="free")) monkeypatch.setattr(server._auth, "cleanup_expired_de_outcomes_async", lambda uid: None) monkeypatch.setattr(server, "_load_global_prior", lambda: GlobalPrior(effects={}, n_users={}, n_obs={})) r = client.post("/api/de/round2", json={ "wt_protein": "A" * 5, "measurements": [ {"mutations": _label(0), "measured_value": 1.0}, {"mutations": _label(1), "measured_value": 3.0}, {"mutations": _label(2), "measured_value": 5.0}, {"mutations": _label(3), "measured_value": 7.0}, ], "settings": _SETTINGS, }) assert r.status_code == 200 body = r.get_json() assert body["ok"] is True assert body["round"] == 2 assert "pool_deltas" in body["surrogate"] assert "global_prior" in body["surrogate"] # New: round 2 carries the epistasis-in-the-loop block and per-variant # confidence. (The string-scorer fixture has no log_probs, so the # interaction re-rank degrades gracefully to applied=False — the point of # this assertion is that the WIRING is present and never crashes the run.) assert "epistasis" in body["surrogate"] assert set(body["surrogate"]["epistasis"]) == {"applied", "n_clash", "n_analyzed"} non_wt = [v for v in body["variants"] if v.get("Variant_ID") != "WT"] assert non_wt and all("Confidence" in v for v in non_wt) def test_result_route_defaults_global_prior_when_never_set(): app = server.create_app() app.config.update(TESTING=True) client = app.test_client() job = server.JobState(job_id="j1", status="done", wt_identifier="WT", wt_protein="ACDEFG", variants=[]) with server._JOBS_LOCK: server._JOBS["j1"] = job r = client.get("/api/result/j1") assert r.status_code == 200 assert r.get_json()["global_prior"] == {"applied": False, "substitution_types": 0} def test_result_route_surfaces_global_prior_when_set(): app = server.create_app() app.config.update(TESTING=True) client = app.test_client() job = server.JobState(job_id="j2", status="done", wt_identifier="WT", wt_protein="ACDEFG", variants=[], global_prior_info={"applied": True, "substitution_types": 7}) with server._JOBS_LOCK: server._JOBS["j2"] = job r = client.get("/api/result/j2") assert r.status_code == 200 assert r.get_json()["global_prior"] == {"applied": True, "substitution_types": 7}