# -*- coding: utf-8 -*- """OCS 最优贡献选择正式 tc 套件:TestClient 走真实 API + golden 数值断言(纯计算端点)。 fixture(镜像 e2e_ocs):8 种质(F1/F2/F3/D/E founder + A/B 全同胞 + C 半同胞), EBV 全 direct(A=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_n:top-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 → 422(schema 层 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)即拦截 → 422(e2e 直调走服务层 CustomException → 409) check("[6] λ<0 → 422(schema 层校验)", 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)