ahmedtaha100's picture
Upload app.py with huggingface_hub
e3dafe3 verified
Raw
History Blame
10.9 kB
from __future__ import annotations
import csv
import hashlib
import io
import json
import random
import re
import time
from pathlib import Path
import streamlit as st
from github_storage import (
is_configured as github_configured,
)
DATA_DIR = Path("data")
PAIRS_FILE = DATA_DIR / "pairs.json"
RESULTS_DIR = Path("results")
AUDIT_LOG = RESULTS_DIR / "audit_log.csv"
FIELDNAMES = [
"pair_index",
"reviewer_id",
"clinically_plausible",
"pathology_preserved",
"quality_score",
"comments",
"timestamp",
]
USE_GITHUB = github_configured()
def _load_reviewer_tokens() -> dict[str, str]:
import os
tokens: dict[str, str] = {}
for key, val in os.environ.items():
if key.startswith("REVIEWER_TOKEN_"):
reviewer_id = key.replace("REVIEWER_TOKEN_", "").lower()
tokens[val] = reviewer_id
try:
reviewer_secrets = st.secrets.get("reviewers", {})
for reviewer_id, token in reviewer_secrets.items():
tokens[str(token)] = str(reviewer_id)
except (KeyError, FileNotFoundError):
pass
return tokens
def _authenticate(token: str) -> str | None:
valid_tokens = _load_reviewer_tokens()
if not valid_tokens:
if re.match(r"^[a-zA-Z0-9_-]{2,30}$", token):
return token
return None
return valid_tokens.get(token)
def _safe_reviewer_id(reviewer_id: str) -> str:
return hashlib.sha256(reviewer_id.encode()).hexdigest()[:12]
def _safe_filename_id(value: str) -> str:
safe = re.sub(r"[^a-zA-Z0-9_-]+", "_", value).strip("_")
return safe or _safe_reviewer_id(value)
def _reviews_to_csv(reviews: dict[int, dict]) -> str:
output = io.StringIO()
writer = csv.DictWriter(output, fieldnames=FIELDNAMES, extrasaction="ignore")
writer.writeheader()
for idx in sorted(reviews.keys()):
writer.writerow(reviews[idx])
return output.getvalue()
def _strip_dicom_metadata(file_path: Path) -> bytes | None:
suffix = file_path.suffix.lower()
if suffix in (".png", ".jpg", ".jpeg", ".bmp", ".tiff", ".tif"):
return file_path.read_bytes()
if suffix in (".dcm", ".dicom"):
try:
import pydicom
from PIL import Image
ds = pydicom.dcmread(str(file_path))
pixel_array = ds.pixel_data
if hasattr(ds, "pixel_array"):
pixel_array = ds.pixel_array
else:
return None
img = Image.fromarray(pixel_array)
if img.mode not in ("L", "RGB"):
img = img.convert("L")
buf = io.BytesIO()
img.save(buf, format="PNG")
return buf.getvalue()
except Exception:
return None
return file_path.read_bytes() if file_path.exists() else None
def load_pairs() -> list[dict]:
if USE_GITHUB:
import github_storage
pairs = github_storage.load_pairs()
if pairs:
return pairs
if not PAIRS_FILE.exists():
st.error(f"Pairs file not found: {PAIRS_FILE}")
st.stop()
with open(PAIRS_FILE) as f:
return json.load(f)
def get_results_path(reviewer_id: str) -> Path:
RESULTS_DIR.mkdir(parents=True, exist_ok=True)
safe_id = _safe_reviewer_id(reviewer_id)
return RESULTS_DIR / f"reviews_{safe_id}.csv"
def load_existing_reviews(reviewer_id: str) -> dict[int, dict]:
if USE_GITHUB:
import github_storage
return github_storage.load_existing_reviews(reviewer_id)
path = get_results_path(reviewer_id)
reviews: dict[int, dict] = {}
if path.exists():
with open(path) as f:
for row in csv.DictReader(f):
reviews[int(row["pair_index"])] = row
return reviews
def save_review(reviewer_id: str, pair_index: int, review: dict) -> None:
if USE_GITHUB:
import github_storage
github_storage.save_review(reviewer_id, pair_index, review)
return
path = get_results_path(reviewer_id)
existing = load_existing_reviews(reviewer_id)
is_update = pair_index in existing
existing[pair_index] = review
with open(path, "w", newline="") as f:
f.write(_reviews_to_csv(existing))
append_audit_log(reviewer_id, pair_index, "update" if is_update else "submit")
def append_audit_log(reviewer_id: str, pair_index: int, action: str) -> None:
if USE_GITHUB:
import github_storage
github_storage.append_audit_log(reviewer_id, pair_index, action)
return
RESULTS_DIR.mkdir(parents=True, exist_ok=True)
write_header = not AUDIT_LOG.exists()
with open(AUDIT_LOG, "a", newline="") as f:
writer = csv.writer(f)
if write_header:
writer.writerow(["timestamp", "reviewer_id", "pair_index", "action"])
writer.writerow([
time.strftime("%Y-%m-%d %H:%M:%S UTC", time.gmtime()),
reviewer_id,
pair_index,
action,
])
def get_shuffled_order(pairs: list[dict], reviewer_id: str) -> list[int]:
seed = int(hashlib.sha256(reviewer_id.encode()).hexdigest(), 16) % (2**32)
indices = list(range(len(pairs)))
random.Random(seed).shuffle(indices)
return indices
def get_lr_swap(pair_index: int, reviewer_id: str) -> bool:
combined = f"{reviewer_id}_{pair_index}"
return int(hashlib.md5(combined.encode()).hexdigest(), 16) % 2 == 0
def render_login() -> str | None:
st.title("SpineFairBench — Radiologist Validation")
st.markdown(
"You are reviewing pairs of spine X-rays for clinical realism. "
"For each pair you will see two images side by side. "
"Please answer the three questions below each pair honestly and to "
"the best of your clinical judgment."
)
st.markdown("---")
token = st.text_input("Enter your access token to begin", type="password")
if token.strip():
reviewer_id = _authenticate(token.strip())
if reviewer_id:
return reviewer_id
st.error("Invalid access token.")
return None
def render_image(path: Path, label: str) -> None:
st.subheader(label)
if not path.exists():
st.warning(f"Image not found: {path.name}")
return
image_bytes = _strip_dicom_metadata(path)
if image_bytes:
st.image(image_bytes, use_container_width=True)
else:
st.warning(f"Could not load image: {path.name}")
def main() -> None:
st.set_page_config(page_title="SpineFairBench Validation", layout="wide")
if "reviewer_id" not in st.session_state:
st.session_state.reviewer_id = None
if "position" not in st.session_state:
st.session_state.position = 0
if st.session_state.reviewer_id is None:
result = render_login()
if result:
st.session_state.reviewer_id = result
append_audit_log(result, -1, "login")
st.rerun()
return
reviewer_id: str = st.session_state.reviewer_id
pairs = load_pairs()
order = get_shuffled_order(pairs, reviewer_id)
reviews = load_existing_reviews(reviewer_id)
total = len(pairs)
reviewed = len(reviews)
pos = st.session_state.position
st.sidebar.markdown(f"**Reviewer:** {reviewer_id}")
st.sidebar.metric("Progress", f"{reviewed} / {total}")
st.sidebar.progress(reviewed / total if total > 0 else 0)
st.sidebar.download_button(
"Export Results CSV",
data=_reviews_to_csv(reviews),
file_name=f"validation_{_safe_filename_id(reviewer_id)}.csv",
mime="text/csv",
disabled=reviewed == 0,
)
if st.sidebar.button("Log out"):
append_audit_log(reviewer_id, -1, "logout")
st.session_state.reviewer_id = None
st.session_state.position = 0
st.rerun()
return
if pos >= total:
st.balloons()
st.success(f"All {total} pairs reviewed. Thank you!")
return
pair_index = order[pos]
pair = pairs[pair_index]
col_prev, col_counter, col_next = st.columns([1, 2, 1])
with col_prev:
if st.button("← Previous", disabled=pos == 0):
st.session_state.position = max(0, pos - 1)
st.rerun()
with col_counter:
status = " ✓" if pair_index in reviews else ""
st.markdown(f"**Pair {pos + 1} of {total}{status}**")
with col_next:
if st.button("Skip →", disabled=pos >= total - 1):
st.session_state.position = min(total - 1, pos + 1)
st.rerun()
source_path = DATA_DIR / pair.get("source_path", "")
generated_path = DATA_DIR / pair.get("generated_path", "")
swap = get_lr_swap(pair_index, reviewer_id)
if swap:
left_path, right_path = generated_path, source_path
else:
left_path, right_path = source_path, generated_path
col_a, col_b = st.columns(2)
with col_a:
render_image(left_path, "Image A")
with col_b:
render_image(right_path, "Image B")
st.markdown("---")
existing = reviews.get(pair_index, {})
plausible = st.radio(
"1. Do both images appear clinically plausible as real spine X-rays?",
["Yes", "No"],
index=0 if existing.get("clinically_plausible", "Yes") == "Yes" else 1,
horizontal=True,
key=f"plausible_{pair_index}",
)
preserved = st.radio(
"2. Is the pathology consistent across both images?",
["Yes", "No", "Uncertain"],
index=["Yes", "No", "Uncertain"].index(
existing.get("pathology_preserved", "Yes")
),
horizontal=True,
key=f"preserved_{pair_index}",
)
quality = st.radio(
"3. Overall pair quality",
["1", "2", "3", "4", "5"],
index=int(existing.get("quality_score", "3")) - 1,
horizontal=True,
key=f"quality_{pair_index}",
captions=["Poor", "Below average", "Acceptable", "Good", "Excellent"],
)
comments = st.text_area(
"Comments (optional)",
value=existing.get("comments", ""),
key=f"comments_{pair_index}",
height=80,
)
if st.button("Submit & Next", type="primary", use_container_width=True):
review = {
"pair_index": pair_index,
"reviewer_id": reviewer_id,
"clinically_plausible": plausible,
"pathology_preserved": preserved,
"quality_score": quality,
"comments": comments,
"timestamp": time.strftime("%Y-%m-%d %H:%M:%S UTC", time.gmtime()),
}
try:
save_review(reviewer_id, pair_index, review)
except Exception as e:
st.error(f"Failed to save review: {e}. Please try again.")
return
if pos < total - 1:
st.session_state.position = pos + 1
else:
st.session_state.position = total
st.rerun()
if __name__ == "__main__":
main()