266 lines
12 KiB
Python
266 lines
12 KiB
Python
# -*- 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)
|