File size: 5,016 Bytes
6754826
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
import io
import json
import uuid
from pathlib import Path

from openai import AsyncOpenAI

from config import LLM_BASE_URL, LLM_MODEL, OPENAI_API_KEY

DATA_DIR   = Path(__file__).parent / "data"
STATE_FILE = DATA_DIR / "report_state.json"

EXTRACTION_PROMPT = (
    "You are a medical lab report parser. Extract every test result from the blood test "
    "report text below. Return ONLY a JSON object of the form "
    '{"results": [{"test_name": str, "value": str, "unit": str, "reference_range": str, '
    '"flag": "low"|"normal"|"high"|"unknown", "value_numeric": number|null, '
    '"ref_low": number|null, "ref_high": number|null}]}. '
    "If the report does not explicitly mark a result as low/high, infer the flag by "
    "comparing the value to the reference range. If a field isn't present in the report, "
    "use an empty string for string fields. "
    "value_numeric is the numeric form of `value` (e.g. \"14.2\" -> 14.2), or null if the "
    "value is not a plain number. ref_low/ref_high are the numeric lower/upper bounds "
    "parsed from `reference_range` when it describes a simple numeric range, e.g. "
    "\"30-100\" -> ref_low=30, ref_high=100; \"below 150\" -> ref_low=null, ref_high=150; "
    "\"above 40\" -> ref_low=40, ref_high=null. If `reference_range` is not a simple "
    "numeric range (qualitative, multiple sub-ranges, etc.), set both ref_low and ref_high "
    "to null. Do not include any text outside the JSON object."
)

_client: AsyncOpenAI | None = None


def _get_client() -> AsyncOpenAI:
    global _client
    if _client is None:
        _client = AsyncOpenAI(api_key=OPENAI_API_KEY, base_url=LLM_BASE_URL)
    return _client


def _load_state() -> dict:
    if STATE_FILE.exists():
        try:
            return json.loads(STATE_FILE.read_text())
        except Exception:
            pass
    return {"report_id": None, "filename": None, "results": [], "summary": None}


def _save_state() -> None:
    STATE_FILE.write_text(json.dumps(report_state, indent=2))


report_state: dict = _load_state()


def parse_pdf_text(raw: bytes) -> str:
    from pypdf import PdfReader
    reader = PdfReader(io.BytesIO(raw))
    pages  = [p.extract_text() or "" for p in reader.pages]
    return "\n\n".join(p for p in pages if p.strip())


async def extract_lab_values(report_text: str, model: str | None = None) -> list[dict]:
    resp = await _get_client().chat.completions.create(
        model=model or LLM_MODEL,
        messages=[
            {"role": "system", "content": EXTRACTION_PROMPT},
            {"role": "user", "content": report_text[:12000]},
        ],
        response_format={"type": "json_object"},
    )
    try:
        data = json.loads(resp.choices[0].message.content or "{}")
    except Exception:
        data = {}
    return data.get("results", [])


def format_report_summary(results: list[dict]) -> str:
    abnormal = [r for r in results if r.get("flag") in ("low", "high")]
    normal   = [r for r in results if r.get("flag") == "normal"]

    lines = []
    if abnormal:
        lines.append("Here are the markers that are outside the normal range:")
        for r in abnormal:
            name    = r.get("test_name") or "Unknown marker"
            val_str = f"{r.get('value', '')} {r.get('unit', '')}".strip()
            line    = f"{name} is {r.get('flag')} at {val_str}".strip()
            ref     = r.get("reference_range")
            if ref:
                line += f" (reference range: {ref})"
            lines.append(line + ".")
    else:
        lines.append("All markers in this report are within the normal range.")

    if normal:
        names = ", ".join(r.get("test_name") or "Unknown marker" for r in normal)
        lines.append(f"The following markers were normal: {names}.")

    lines.append(
        "This is not a medical diagnosis — please consult your doctor about these results "
        "before making any changes to your diet or supplements."
    )
    return "\n".join(lines)


async def process_report_upload(raw: bytes, filename: str, model: str | None = None) -> dict:
    if filename.lower().endswith(".pdf"):
        text = parse_pdf_text(raw)
    else:
        text = raw.decode("utf-8", errors="replace")

    if not text.strip():
        raise ValueError("Could not extract any text from the uploaded report.")

    results = await extract_lab_values(text, model=model)
    summary = format_report_summary(results)

    report_state["report_id"] = uuid.uuid4().hex[:8]
    report_state["filename"]  = filename
    report_state["results"]   = results
    report_state["summary"]   = summary
    _save_state()

    return dict(report_state)


def get_current_report() -> dict:
    if not report_state.get("results"):
        return {"has_report": False}
    return {"has_report": True, **report_state}


def clear_report() -> None:
    report_state["report_id"] = None
    report_state["filename"]  = None
    report_state["results"]   = []
    report_state["summary"]   = None
    _save_state()