Files
dpb/backend/scripts/test_bre_analysis_ssgblup_tc.py
T

288 lines
13 KiB
Python
Raw Normal View History

# -*- coding: utf-8 -*-
"""ssGBLUP 正式 tc 套件:TestClient 走真实 API + golden 数值断言(H⁻¹ 拼接判别)。
fixture:5 株表型 + 基因型数据集(树0/1/3 直连、树2 名称兜底、树4 全多等位无有效剂量)。
经真实 HTTP 端点 POST /api/v1/bre/statistics/gblup/run 落库后,ORM 断言:
[1] GBLUP method=GBLUP、h²∈(0,1)、5 条 EBV、非基因型株(树4) EBV=0
[2] ssGBLUPmethod=ssGBLUP、h²∈(0,1)、5 条 EBV、非基因型株(树4) EBV≠0
(H⁻¹ 拼接生效的判别核心——若 H⁻¹=A⁻¹+[0;G⁻¹-A22⁻¹] 拼接错,非基因型亲属
EBV 会为 0 或爆炸,正是「算偏但当正常结果输出」的同款静默失效)
[3] 非基因型亲本(g 前缀)在 ssGBLUP 下 EBV 有限(H 进 MME,非 A22 全零)
依赖: Redis + PG 正常(TestClient 走真实 lifespan)。运行后自动清理。
"""
import os
os.environ["ENVIRONMENT"] = "dev"
os.environ["PYTHONUTF8"] = "1"
import sys, asyncio # noqa: E402
sys.path.insert(0, r"d:\dpb\dpb\backend")
import main # noqa: E402
from fastapi.testclient import TestClient # noqa: E402
from sqlalchemy import delete, select # noqa: E402
from app.core.database import create_async_engine_and_session # noqa: E402
from app.api.v1.module_system.user.model import UserModel # noqa: E402 (注册 mapper)
from app.api.v1.module_bre.trait.model import TraitModel # noqa: E402
from app.api.v1.module_bre.target.model import TargetModel # noqa: E402
from app.api.v1.module_bre.germplasm.model import BreedingGermplasmModel # noqa: E402
from app.api.v1.module_bre.cross_combination.model import CrossCombinationModel # noqa: E402
from app.api.v1.module_bre.tree.model import TreeModel # noqa: E402
from app.api.v1.module_bre.trait_observation.model import TraitObservationModel # noqa: E402
from app.api.v1.module_bre.genotype_dataset.model import GenotypingDatasetModel # noqa: E402
from app.api.v1.module_bre.genotype_sample.model import GenotypeSampleModel # noqa: E402
from app.api.v1.module_bre.genotype_call.model import GenotypeCallModel # noqa: E402
from app.api.v1.module_bre.marker.model import MarkerModel # noqa: E402
from app.api.v1.module_bre.statistics.model import ( # noqa: E402
PredictionModel,
PredictionValueModel,
StatisticsJobModel,
)
create_app = main.create_app
TOKEN = None
ok, fail = 0, 0
PREFIX = "TCSSGB"
tokens: dict[str, list[int]] = {
"combo": [], "tree": [], "obs": [],
"marker": [], "dataset": [], "sample": [], "call": [],
}
TRAIT_CODE = f"tR_{PREFIX}"
pred_ids: list[int] = []
FIX = {}
def check(name, cond, detail=""):
global ok, fail
if cond:
ok += 1
print(f" [ok] {name} {detail}")
else:
fail += 1
print(f" [FAIL] {name} {detail}")
def login(client):
global TOKEN
d = {"username": "super", "password": "123456", "grant_type": "password", "login_type": "PC端"}
r = client.post("/api/v1/system/auth/login", data=d)
b = r.json()
if r.status_code == 200 and b.get("code") == 0:
TOKEN = b["data"]["access_token"]
return
key = client.get("/api/v1/system/auth/captcha/get").json()["data"]["key"]
client.post("/api/v1/system/auth/captcha/slider/complete", json={"captcha_key": key})
d["captcha_key"] = key
r = client.post("/api/v1/system/auth/login", data=d)
b = r.json()
assert r.status_code == 200 and b.get("code") == 0, f"LOGIN FAIL {r.status_code} {b}"
TOKEN = b["data"]["access_token"]
def auth():
return {"Authorization": f"Bearer {TOKEN}"}
GENOS = {
0: [0, 0, 0, 1, 1, 1, 2, 2, 2, 1],
1: [1, 1, 0, 0, 1, 2, 2, 0, 1, 1],
2: [0, 1, 2, 1, 0, 1, 0, 1, 2, 0],
3: [1, 2, 1, 2, 1, 0, 1, 0, 0, 1],
}
GT = {0: "0/0", 1: "0/1", 2: "1/1"}
async def _build_fixture() -> None:
engine, sf = create_async_engine_and_session()
try:
async with sf() as db:
trait = TraitModel(trait_code=TRAIT_CODE, trait_name=f"含糖量{PREFIX}", data_type="numeric",
unit="%", is_core="1", direction="desc", into_ebv="1",
default_h2=0.5, created_id=1)
db.add(trait)
await db.flush()
FIX["trait_id"] = trait.id
target = TargetModel(target_name=f"目标{PREFIX}", created_id=1)
db.add(target)
await db.flush()
fm = BreedingGermplasmModel(cultivar_name=f"FM_{PREFIX}", can_be_female=True, created_id=1)
mm = BreedingGermplasmModel(cultivar_name=f"MM_{PREFIX}", can_be_male=True, created_id=1)
db.add_all([fm, mm])
await db.flush()
combo = CrossCombinationModel(
combination_code=f"C_{PREFIX}", bre_target_id=target.id,
female_parent_id=fm.id, male_parent_id=mm.id, design_type="full_diallel", created_id=1)
db.add(combo)
await db.flush()
tokens["combo"].append(combo.id)
trees = []
for i in range(5):
t = TreeModel(combination_id=combo.id, tree_no=f"{PREFIX}-T{i:02d}", status="alive",
stage="seedling", generation="F1", planted_date="2023-03-10", created_id=1)
db.add(t)
trees.append(t)
await db.flush()
tokens["tree"] = [t.id for t in trees]
FIX["tree_ids"] = [t.id for t in trees]
obs_rows = []
for i, t in enumerate(trees):
o = TraitObservationModel(tree_id=t.id, combination_id=combo.id, trait_id=trait.id,
value_numeric=12.0 + 0.8 * i, evaluate_year=2025, created_id=1)
db.add(o)
obs_rows.append(o)
await db.flush()
tokens["obs"] = [o.id for o in obs_rows]
ds = GenotypingDatasetModel(dataset_name=f"GS_{PREFIX}", platform="SSR",
purpose="GS", created_id=1)
db.add(ds)
await db.flush()
FIX["dataset_id"] = ds.id
tokens["dataset"].append(ds.id)
markers = []
for j in range(10):
m = MarkerModel(marker_name=f"MK{j}_{PREFIX}", marker_type="SNP",
chromosome=str(j % 5), position=j * 100, created_id=1)
db.add(m)
markers.append(m)
await db.flush()
tokens["marker"] = [m.id for m in markers]
FIX["markers"] = markers
samples = []
for i in (0, 1, 3):
s = GenotypeSampleModel(sample_name=f"{PREFIX}-S{i}", dataset_id=ds.id,
source_type="tree", source_id=trees[i].id, created_id=1)
db.add(s)
samples.append(s)
s_fb = GenotypeSampleModel(sample_name=trees[2].tree_no, dataset_id=ds.id,
source_type=None, source_id=None, created_id=1)
db.add(s_fb)
samples.append(s_fb)
s_bad = GenotypeSampleModel(sample_name=f"{PREFIX}-B4", dataset_id=ds.id,
source_type="tree", source_id=trees[4].id, created_id=1)
db.add(s_bad)
samples.append(s_bad)
await db.flush()
tokens["sample"] = [s.id for s in samples]
calls = []
tree_samp = {trees[0].id: samples[0], trees[1].id: samples[1], trees[3].id: samples[2]}
tree_samp[trees[2].id] = s_fb
genos_by_tid = {trees[i].id: GENOS[i] for i in range(4)}
for tid, smp in tree_samp.items():
for j, m in enumerate(markers):
calls.append(GenotypeCallModel(sample_id=smp.id, marker_id=m.id,
allele=GT[genos_by_tid[tid][j]], created_id=1))
for j, m in enumerate(markers):
calls.append(GenotypeCallModel(sample_id=s_bad.id, marker_id=m.id,
allele="1/2", created_id=1))
db.add_all(calls)
await db.flush()
tokens["call"] = [c.id for c in calls]
await db.commit()
finally:
await engine.dispose()
async def _verify_and_cleanup() -> None:
engine, sf = create_async_engine_and_session()
try:
async with sf() as db:
# ---- [1] GBLUP ----
p1 = await db.get(PredictionModel, pred_ids[0])
check("[1] method=GBLUP", p1.method == "GBLUP", p1.method)
check("[1] h²∈(0,1)", p1.heritability is not None and 0 < p1.heritability < 1,
f"{p1.heritability}")
vals1 = (await db.execute(select(PredictionValueModel).where(
PredictionValueModel.prediction_id == pred_ids[0]))).scalars().all()
check("[1] 5 条 EBV", len(vals1) == 5, f"{len(vals1)}")
v4_1 = next(v for v in vals1 if v.tree_id == FIX["tree_ids"][4])
check("[1] 非基因型株(树4) GBLUP EBV=0", float(v4_1.predicted_value) == 0.0,
f"{v4_1.predicted_value}")
check("[1] 全部 EBV 非 NaN", all(float(v.predicted_value) == float(v.predicted_value)
for v in vals1))
# ---- [2] ssGBLUP ----
p2 = await db.get(PredictionModel, pred_ids[1])
check("[2] method=ssGBLUP", p2.method == "ssGBLUP", p2.method)
check("[2] h²∈(0,1)", p2.heritability is not None and 0 < p2.heritability < 1,
f"{p2.heritability}")
vals2 = (await db.execute(select(PredictionValueModel).where(
PredictionValueModel.prediction_id == pred_ids[1]))).scalars().all()
check("[2] 5 条 EBV", len(vals2) == 5, f"{len(vals2)}")
v4_2 = next(v for v in vals2 if v.tree_id == FIX["tree_ids"][4])
check("[2] 非基因型株(树4) ssGBLUP EBV≠0H⁻¹ 拼接生效)",
float(v4_2.predicted_value) != 0.0, f"{v4_2.predicted_value}")
# ---- [3] 非基因型亲本 EBV 有限 ----
for pv in vals2:
assert float(pv.predicted_value) == float(pv.predicted_value), "EBV NaN"
check("[3] ssGBLUP 全部 EBV 有限", all(abs(float(v.predicted_value)) < 1e6 for v in vals2))
# ---- cleanup ----
if pred_ids:
await db.execute(delete(PredictionValueModel).where(
PredictionValueModel.prediction_id.in_(pred_ids)))
await db.execute(delete(StatisticsJobModel).where(
StatisticsJobModel.result_ref.in_(pred_ids)))
await db.execute(delete(PredictionModel).where(PredictionModel.id.in_(pred_ids)))
if tokens["call"]:
await db.execute(delete(GenotypeCallModel).where(
GenotypeCallModel.id.in_(tokens["call"])))
if tokens["sample"]:
await db.execute(delete(GenotypeSampleModel).where(
GenotypeSampleModel.id.in_(tokens["sample"])))
if tokens["dataset"]:
await db.execute(delete(GenotypingDatasetModel).where(
GenotypingDatasetModel.id.in_(tokens["dataset"])))
if tokens["marker"]:
await db.execute(delete(MarkerModel).where(MarkerModel.id.in_(tokens["marker"])))
if tokens["obs"]:
await db.execute(delete(TraitObservationModel).where(
TraitObservationModel.id.in_(tokens["obs"])))
if tokens["tree"]:
await db.execute(delete(TreeModel).where(TreeModel.id.in_(tokens["tree"])))
if tokens["combo"]:
await db.execute(delete(CrossCombinationModel).where(
CrossCombinationModel.id.in_(tokens["combo"])))
await db.execute(delete(BreedingGermplasmModel).where(
BreedingGermplasmModel.cultivar_name.in_([f"FM_{PREFIX}", f"MM_{PREFIX}"])))
await db.execute(delete(TraitModel).where(TraitModel.trait_code == TRAIT_CODE))
await db.execute(delete(TargetModel).where(TargetModel.target_name == f"目标{PREFIX}"))
await db.commit()
print(f"[cleanup] ssGBLUP tc 数据已清(批次 {len(pred_ids)} 个)")
finally:
await engine.dispose()
def main_() -> None:
asyncio.run(_build_fixture())
try:
with TestClient(create_app()) as client:
login(client)
for method in ("gblup", "ssgblup"):
r = client.post("/api/v1/bre/statistics/gblup/run", json={
"dataset_id": FIX["dataset_id"], "trait_id": FIX["trait_id"],
"trait_code": TRAIT_CODE, "year": None, "method": method, "maf_min": 0.05,
}, headers=auth())
check(f"[HTTP] gblup/run {method} 200", r.status_code == 200,
f"{r.status_code} {str(r.text)[:150]}")
body = r.json()
pid = body.get("data") if body.get("code") == 0 else None
check(f"[HTTP] {method} 返回 id", isinstance(pid, int), f"pid={pid}")
if isinstance(pid, int):
pred_ids.append(pid)
finally:
asyncio.run(_verify_and_cleanup())
print(f"\n===== ssGBLUP tc 套件:ok={ok} fail={fail} =====")
if __name__ == "__main__":
main_()
sys.exit(1 if fail else 0)