""" app.py โ€” Prescription Digitization Tool Vision-first pipeline: tries llava/moondream vision model first (most accurate), falls back to OCR + Mistral text mode if no vision model available. """ import io, os import pandas as pd import streamlit as st from PIL import Image import db from preprocess import preprocess_image from ocr_engine import run_ocr from llm_extractor import extract_from_image, extract_fields, DEFAULT_MODEL, DEFAULT_HOST st.set_page_config(page_title="Prescription Digitization", page_icon="๐Ÿฅ", layout="wide") db.init_db() for k, v in { "ocr_text": "", "extraction_result": None, "ocr_engine_used": "both", "uploaded_filename": "", "image_bytes": None, }.items(): if k not in st.session_state: st.session_state[k] = v st.title("๐Ÿฅ Prescription Digitization") st.caption("Fully offline ยท Patient data never leaves this machine") tab_upload, tab_records, tab_export = st.tabs( ["๐Ÿ“ค Upload & Extract", "๐Ÿ“‹ Saved Records", "โฌ‡๏ธ Export CSV"] ) with tab_upload: uploaded = st.file_uploader( "Upload prescription image", type=["jpg", "jpeg", "png"], label_visibility="collapsed", accept_multiple_files=False, ) if uploaded: raw_bytes = uploaded.read() st.session_state.image_bytes = raw_bytes st.session_state.uploaded_filename = uploaded.name image = Image.open(io.BytesIO(raw_bytes)) col_orig, col_clean = st.columns(2) with col_orig: st.image(image, caption="Original", width=450) with col_clean: cleaned = preprocess_image(image, do_deskew=True) st.image(cleaned, caption="Preprocessed", width=450) if st.button("๐Ÿ” Extract Information", type="primary", use_container_width=True): # โ”€โ”€ Try vision model first โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€ with st.spinner("Checking for vision model (llava/moondream)โ€ฆ"): vision_result = extract_from_image(raw_bytes, host=DEFAULT_HOST) if vision_result.success: st.session_state.extraction_result = vision_result st.session_state.ocr_engine_used = vision_result.mode n = len(vision_result.records) st.success( f"โœจ Vision model extracted **{n} visit record{'s' if n>1 else ''}** " f"โ€” review below before saving." ) elif vision_result.error == "NO_VISION_MODEL": # โ”€โ”€ Fall back to OCR + Mistral โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€ st.info("No vision model found โ€” using OCR + Mistral (slower, less accurate).\n\n" "For better results: `ollama pull llava`") with st.spinner("Running OCR (Donut + TrOCR)โ€ฆ"): try: cleaned_bytes = io.BytesIO() cleaned.save(cleaned_bytes, format="JPEG") ocr_result = run_ocr(cleaned, engine="both") st.session_state.ocr_text = ocr_result.raw_text st.session_state.ocr_engine_used = ocr_result.engine except Exception as e: st.error(f"OCR failed: {e}") st.stop() with st.spinner(f"Structuring with Ollama ({DEFAULT_MODEL})โ€ฆ"): text_result = extract_fields( st.session_state.ocr_text, model=DEFAULT_MODEL, host=DEFAULT_HOST ) if text_result.success: st.session_state.extraction_result = text_result n = len(text_result.records) st.success(f"Found **{n} visit record{'s' if n>1 else ''}** โ€” review below.") elif text_result.error == "OLLAMA_NOT_RUNNING": st.error("Ollama is not running.") st.code("ollama serve") st.session_state.extraction_result = None else: st.error(f"Extraction error: {text_result.error}") st.info("Try clicking Extract Information again.") st.session_state.extraction_result = None else: st.error(f"Vision extraction error: {vision_result.error}") st.session_state.extraction_result = None if st.session_state.ocr_text: with st.expander("Raw OCR output", expanded=False): st.code(st.session_state.ocr_text, language=None) # โ”€โ”€ Review form โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€ result = st.session_state.extraction_result if result and result.records: st.divider() st.subheader("โœ๏ธ Review extracted information") meta = result.meta if meta: st.caption(f"Mode: {result.mode}") mc = st.columns(min(len(meta), 4)) for col, (k, v) in zip(mc, meta.items()): col.text_input(k.replace("_"," ").capitalize(), value=str(v), disabled=True) for i, rec in enumerate(result.records): label = f"Visit {i+1}" if rec.get("visit_date"): label += f" โ€” {rec['visit_date']}" with st.expander(label, expanded=True): with st.form(f"form_{i}"): edited_fields = {} field_items = list(rec["fields"].items()) for j in range(0, len(field_items), 3): chunk = field_items[j:j+3] cols = st.columns(len(chunk)) for col, (key, val) in zip(cols, chunk): edited_fields[key] = col.text_input( key.replace("_"," ").capitalize(), value=str(val) if val else "" ) meds = rec.get("medications", []) if meds: st.markdown("**Medications**") med_df = pd.DataFrame(meds) for c in ["drug_name","dosage","frequency","route"]: if c not in med_df.columns: med_df[c] = "" edited_meds_df = st.data_editor( med_df[["drug_name","dosage","frequency","route"]], num_rows="dynamic", use_container_width=True, key=f"meds_{i}" ) final_meds = edited_meds_df.fillna("").to_dict(orient="records") else: final_meds = [] st.caption("No medications detected โ€” add manually if needed.") if st.form_submit_button( f"โœ… Save Visit {i+1} to Database", type="primary", use_container_width=True ): clean_fields = {k: v for k, v in edited_fields.items() if v.strip()} pid = db.save_record( visit_date=rec.get("visit_date"), fields=clean_fields, medications=final_meds, meta=meta, ocr_engine=st.session_state.ocr_engine_used, source_filename=st.session_state.uploaded_filename, ) st.success(f"Visit {i+1} saved as record #{pid}.") with tab_records: records = db.fetch_prescriptions() if not records: st.info("No records saved yet.") else: st.dataframe(pd.DataFrame(records), use_container_width=True, hide_index=True) st.subheader("Medications") sel_id = st.selectbox("Select record", [r["id"] for r in records], format_func=lambda x: f"Record #{x}") if sel_id: meds = db.fetch_medications(sel_id) if meds: st.dataframe(pd.DataFrame(meds), use_container_width=True, hide_index=True) else: st.caption("No medications for this record.") st.divider() del_id = st.number_input("Delete record by ID", min_value=0, step=1, value=0) if st.button("๐Ÿ—‘๏ธ Delete record") and del_id > 0: db.delete_prescription(int(del_id)) st.success(f"Deleted #{del_id}.") st.rerun() with tab_export: flat = db.fetch_all_flat() if not flat: st.info("No data to export yet.") else: df = pd.DataFrame(flat) st.dataframe(df, use_container_width=True, hide_index=True) csv = df.to_csv(index=False).encode("utf-8") st.download_button("โฌ‡๏ธ Download CSV", data=csv, file_name="prescriptions_export.csv", mime="text/csv", use_container_width=True)