Files
dpb/backend/scripts/test_bre_analysis_ocs_tc.py
T
34047007@qq.com b95053c52c init: 初始化 dpb 桃育种系统代码库
前后端 + 后端 FastAPI 全量源码、部署脚本与文档。
2026-08-06 00:17:49 +08:00

266 lines
12 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- coding: utf-8 -*-
"""OCS 最优贡献选择正式 tc 套件:TestClient 走真实 API + golden 数值断言(纯计算端点)。
fixture(镜像 e2e_ocs):8 种质(F1/F2/F3/D/E founder + A/B 全同胞 + C 半同胞),
EBV 全 directA=10, B=9.8, C=9.5, D=9.0, E=8.0),n_select=3。
经真实 HTTP 端点 POST /api/v1/bre/statistics/ocs 断言:
[1] 基本解:Σc=3、c≥0、EBV coverage=direct 5、按 c 降序、最高贡献=A(≈1.419)
[2] vs_top_ntop-n mean_ebv 复算、OCS EBV 不减、avg_kinship 更低(亲缘受控)、c'Ac≈4.641
[3] λ↑ → avg_kinship 单调非增(包络定理)
[4] λ=0 退化:满额投最高 EBV 单亲(c_A=3, mean=10, kin=1
[5] 显式告警:D/E 未解析 → kinship_status=partial + kinship_warning 含 D/E
[6] 校验:n_select>候选 / 空候选 → 409、λ<0 → 422schema 层 Pydantic ge=0 拦截)
依赖: 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 # 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.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.pedigree.model import PedigreeModel # noqa: E402
from app.api.v1.module_bre.statistics.model import PredictionModel, PredictionValueModel # noqa: E402
create_app = main.create_app
TOKEN = None
ok, fail = 0, 0
PREFIX = "TC_OCS"
tokens: dict[str, list[int]] = {
"germ": [], "ped": [], "combo": [], "target": [], "pred": [], "predval": [],
}
FIX: dict = {}
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}"}
async def _build_fixture() -> None:
engine, sf = create_async_engine_and_session()
try:
async with sf() as db:
target = TargetModel(target_name=f"{PREFIX}-T", created_id=1)
db.add(target)
await db.flush()
tokens["target"].append(target.id)
combo = CrossCombinationModel(
combination_code=f"{PREFIX}-C1", bre_target_id=target.id,
female_parent_id=None, male_parent_id=None, design_type="full_diallel", created_id=1)
db.add(combo)
await db.flush()
tokens["combo"].append(combo.id)
def mk_germ(name):
g = BreedingGermplasmModel(cultivar_name=name, can_be_female=True,
can_be_male=True, created_id=1)
db.add(g)
return g
F1, F2, F3 = mk_germ(f"{PREFIX}-F1"), mk_germ(f"{PREFIX}-F2"), mk_germ(f"{PREFIX}-F3")
D, E = mk_germ(f"{PREFIX}-D"), mk_germ(f"{PREFIX}-E")
A = mk_germ(f"{PREFIX}-A")
B = mk_germ(f"{PREFIX}-B")
C = mk_germ(f"{PREFIX}-C")
await db.flush()
for g in (F1, F2, F3, D, E, A, B, C):
tokens["germ"].append(g.id)
FIX.update({"F1": F1.id, "F2": F2.id, "F3": F3.id, "D": D.id, "E": E.id,
"A": A.id, "B": B.id, "C": C.id})
peds = []
for child, dam, sire in ((A, F1, F2), (B, F1, F2), (C, F1, F3)):
ped = PedigreeModel(combination_id=combo.id, child_code=child.cultivar_name,
dam_id=dam.id, sire_id=sire.id, generation="F1", created_id=1)
db.add(ped)
peds.append(ped)
await db.flush()
tokens["ped"] = [x.id for x in peds]
p = PredictionModel(model_name=f"{PREFIX}-pred", trait_id=None,
method="ABLUP", heritability=0.5, created_id=1)
db.add(p)
await db.flush()
tokens["pred"].append(p.id)
FIX["pred_id"] = p.id
pvs = []
for g, v in ((A, 10.0), (B, 9.8), (C, 9.5), (D, 9.0), (E, 8.0)):
pv = PredictionValueModel(prediction_id=p.id, germplasm_id=g.id, trait_id=None,
predicted_value=v, reliability=0.8, rank=1, created_id=1)
db.add(pv)
pvs.append(pv)
await db.flush()
tokens["predval"] = [pv.id for pv in pvs]
await db.commit()
finally:
await engine.dispose()
async def _cleanup() -> None:
engine, sf = create_async_engine_and_session()
try:
async with sf() as db:
if tokens["predval"]:
await db.execute(delete(PredictionValueModel).where(
PredictionValueModel.id.in_(tokens["predval"])))
if tokens["pred"]:
await db.execute(delete(PredictionModel).where(PredictionModel.id.in_(tokens["pred"])))
if tokens["ped"]:
await db.execute(delete(PedigreeModel).where(PedigreeModel.id.in_(tokens["ped"])))
if tokens["germ"]:
await db.execute(delete(BreedingGermplasmModel).where(
BreedingGermplasmModel.id.in_(tokens["germ"])))
if tokens["combo"]:
await db.execute(delete(CrossCombinationModel).where(
CrossCombinationModel.id.in_(tokens["combo"])))
if tokens["target"]:
await db.execute(delete(TargetModel).where(TargetModel.id.in_(tokens["target"])))
await db.commit()
print("[cleanup] OCS tc 数据已清")
finally:
await engine.dispose()
def main_() -> None:
asyncio.run(_build_fixture())
try:
with TestClient(create_app()) as client:
login(client)
cand = [FIX["A"], FIX["B"], FIX["C"], FIX["D"], FIX["E"]]
TOP3_MEAN = (10.0 + 9.8 + 9.5) / 3.0 # 9.7667
TOP3_KIN = (3 + 2 * (0.5 + 0.25 + 0.25)) / 9.0 # 5/9 = 0.5556
def post_ocs(n_select, lam):
r = client.post("/api/v1/bre/statistics/ocs", json={
"candidate_germplasm_ids": cand, "n_select": n_select, "lam": lam,
"prediction_id": FIX["pred_id"]}, headers=auth())
check(f"[HTTP] ocs(n={n_select},λ={lam}) 200", r.status_code == 200,
f"{r.status_code} {str(r.text)[:150]}")
b = r.json()
return b.get("data") if b.get("code") == 0 else None
# ---- [1] 基本解(λ=0.3, n_select=3----
m1 = post_ocs(3, 0.3)
if m1 is None:
return
s1 = sum(x["c"] for x in m1["contributions"])
check("[1] Σc = n_select=3", abs(s1 - 3.0) < 1e-3, f"sum={s1:.4f}")
check("[1] 所有 c ≥ 0", all(x["c"] >= 0 for x in m1["contributions"]),
[x["c"] for x in m1["contributions"]])
check("[1] EBV 源 coverage=direct 5",
m1["ebv_source"]["coverage"] == {"direct": 5, "progeny": 0, "missing": 0},
f"{m1['ebv_source']['coverage']}")
check("[1] 贡献列表按 c 降序",
all(m1["contributions"][i]["c"] >= m1["contributions"][i + 1]["c"]
for i in range(len(m1["contributions"]) - 1)))
check("[1] 最高贡献=A(≈1.419)",
abs(m1["contributions"][0]["c"] - 1.419) < 0.05
and m1["contributions"][0]["germplasm_id"] == FIX["A"],
f"{m1['contributions'][0]}")
check("[1] 每行含 name/source", all("name" in x and "source" in x
for x in m1["contributions"]))
# ---- [2] vs_top_n:亲缘受控价值 ----
vtn = m1["vs_top_n"]
check("[2] top-n mean_ebv 复算", abs(vtn["mean_ebv"] - TOP3_MEAN) < 1e-3,
f"{vtn['mean_ebv']}")
check("[2] top-n avg_kinship 复算", abs(vtn["avg_kinship"] - TOP3_KIN) < 1e-3,
f"{vtn['avg_kinship']}")
check("[2] OCS EBV 与 top-n 相当(几乎不减)",
vtn["ebv_loss"] > -0.05 and vtn["ebv_loss"] < 0.1,
f"loss={vtn['ebv_loss']}")
check("[2] OCS avg_kinship 更低(亲缘受控)",
m1["avg_kinship"] < vtn["avg_kinship"] and vtn["kinship_reduction"] > 0.02,
f"ocs={m1['avg_kinship']} top={vtn['avg_kinship']}")
check("[2] c'Ac 复算 ≈ 4.641", abs(m1["kinship_quad"] - 4.641) < 0.05,
f"{m1['kinship_quad']}")
# ---- [3] λ↑ → avg_kinship 单调非增 ----
kins = []
for lam in (0.05, 0.3, 1.0, 3.0):
mm = post_ocs(3, lam)
if mm is not None:
kins.append(mm["avg_kinship"])
check("[3] λ↑ avg_kinship 单调非增",
len(kins) == 4 and all(kins[i] >= kins[i + 1] for i in range(len(kins) - 1)),
f"{kins}")
# ---- [4] λ=0 退化:满额投最高 EBV 单亲 ----
m4 = post_ocs(3, 0.0)
if m4:
a4 = next(x for x in m4["contributions"] if x["germplasm_id"] == FIX["A"])
check("[4] λ=0 满额投 A",
len(m4["contributions"]) == 1 and abs(a4["c"] - 3.0) < 1e-6
and abs(m4["mean_ebv"] - 10.0) < 1e-6 and abs(m4["avg_kinship"] - 1.0) < 1e-6,
f"{m4['contributions']} mean={m4['mean_ebv']} kin={m4['avg_kinship']}")
# ---- [5] 显式告警(D/E 未解析 → partial----
if m1:
check("[5] kinship_status=partial + 告警列 D/E",
m1["kinship_status"] == "partial" and m1["kinship_warning"]
and f"{PREFIX}-D" in m1["kinship_warning"]
and f"{PREFIX}-E" in m1["kinship_warning"],
f"{m1['kinship_status']} {m1['kinship_warning']}")
# ---- [6] 校验 409 ----
r = client.post("/api/v1/bre/statistics/ocs", json={
"candidate_germplasm_ids": cand, "n_select": 6, "lam": 0.3,
"prediction_id": FIX["pred_id"]}, headers=auth())
check("[6] n_select>候选 → 409", r.status_code == 409, f"{r.status_code}")
r = client.post("/api/v1/bre/statistics/ocs", json={
"candidate_germplasm_ids": [], "n_select": 3, "lam": 0.3,
"prediction_id": FIX["pred_id"]}, headers=auth())
check("[6] 空候选 → 409", r.status_code == 409, f"{r.status_code}")
r = client.post("/api/v1/bre/statistics/ocs", json={
"candidate_germplasm_ids": cand, "n_select": 3, "lam": -0.1,
"prediction_id": FIX["pred_id"]}, headers=auth())
# λ<0 在 schema 层 Pydantic 校验(ge=0)即拦截 → 422e2e 直调走服务层 CustomException → 409
check("[6] λ<0 → 422schema 层校验)", r.status_code == 422, f"{r.status_code}")
finally:
asyncio.run(_cleanup())
print(f"\n===== OCS tc 套件:ok={ok} fail={fail} =====")
if __name__ == "__main__":
main_()
sys.exit(1 if fail else 0)