Files

266 lines
12 KiB
Python
Raw Permalink Normal View History

# -*- 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)