import math import uuid import numpy as np import pandas as pd import pytest from httpx import AsyncClient from sqlalchemy import select from app.db import SessionLocal from app.models import Job, JobStatus, Prediction, Sample, Variant from app.services import scoring def variant(**kw: object) -> Variant: fields: dict = {"chrom": "22", "pos": 1, "ref": "A", "alt": "G", "annotations": {}} fields.update(kw) return Variant(**fields) def test_raw_frame_sends_the_model_contract_columns() -> None: frame = scoring.raw_frame([ variant(impact="HIGH", consequence="stop_gained", gnomad_af=None, annotations={"CADD_PHRED": "35", "am_pathogenicity": "0.98"}), variant(impact="LOW", consequence="synonymous_variant", gnomad_af=0.2, annotations={}), ]) assert list(frame.columns) == scoring.RAW_COLUMNS assert frame["impact"].tolist() == ["HIGH", "LOW"] assert frame["cadd_phred"].iloc[0] == "35" assert pd.isna(frame["cadd_phred"].iloc[1]) assert math.isnan(frame["gnomad_af"].iloc[0]) class FakeModel: def __init__(self, score: float) -> None: self.score = score def predict(self, frame: pd.DataFrame) -> np.ndarray: assert list(frame.columns) == scoring.RAW_COLUMNS return np.full(len(frame), self.score) async def make_job(status: JobStatus, n_variants: int) -> uuid.UUID: async with SessionLocal() as s: sample = Sample(name=f"s-{uuid.uuid4()}", vcf_uri="gs://b/x.vcf.gz", assembly="GRCh38") job = Job(sample=sample, status=status) s.add_all([sample, job, *(variant(job=job, pos=i + 1) for i in range(n_variants))]) await s.commit() return job.id async def predictions(job_id: uuid.UUID) -> list[Prediction]: async with SessionLocal() as s: rows = await s.scalars( select(Prediction).join(Variant).where(Variant.job_id == job_id) ) return list(rows) @pytest.mark.usefixtures("db") async def test_scoring_twice_updates_instead_of_failing( client: AsyncClient, monkeypatch: pytest.MonkeyPatch ) -> None: job_id = await make_job(JobStatus.succeeded, n_variants=3) monkeypatch.setattr(scoring, "load_model", lambda: (FakeModel(0.9), "7")) r = await client.post(f"/api/predictions/score/{job_id}") assert r.status_code == 200, r.text assert r.json() == {"job_id": str(job_id), "scored": 3, "model_version": "7"} monkeypatch.setattr(scoring, "load_model", lambda: (FakeModel(0.2), "8")) r = await client.post(f"/api/predictions/score/{job_id}") assert r.status_code == 200, r.text preds = await predictions(job_id) assert len(preds) == 3 assert {(p.score, p.model_version) for p in preds} == {(0.2, "8")} @pytest.mark.usefixtures("db") async def test_scoring_unknown_job_is_404(client: AsyncClient) -> None: r = await client.post(f"/api/predictions/score/{uuid.uuid4()}") assert r.status_code == 404 @pytest.mark.usefixtures("db") async def test_scoring_unfinished_job_is_409(client: AsyncClient) -> None: job_id = await make_job(JobStatus.running, n_variants=1) r = await client.post(f"/api/predictions/score/{job_id}") assert r.status_code == 409