import tempfile
from pathlib import Path
from typing import Callable

import pypdfium2 as pdfium


def render_pdf_as_slides(
    pdf_source,
    output_dir: Path,
    source_name: str,
    describe_fn: Callable[[Path], str | None],
) -> str:
    """Render each page of a PDF as a JPEG, optionally describe with LLM, return markdown.

    pdf_source: file path (str/Path) or raw bytes.
    describe_fn: callable(img_path) -> alt_text or None; pass None to skip descriptions.
    """
    slides_dir = output_dir / f".{source_name}"
    slides_dir.mkdir(exist_ok=True)

    if isinstance(pdf_source, (str, Path)):
        doc = pdfium.PdfDocument(str(pdf_source))
    else:
        doc = pdfium.PdfDocument(pdf_source)

    md_chunks = []
    with tempfile.TemporaryDirectory() as tmp:
        tmp_path = Path(tmp)
        for i, page in enumerate(doc):
            bitmap = page.render(scale=2)
            pil_img = bitmap.to_pil().convert("RGB")
            tmp_img = tmp_path / f"slide-{i:03d}.jpg"
            pil_img.save(tmp_img, format="JPEG", quality=85)

            dest = slides_dir / f"slide-{i + 1:03d}.jpg"
            dest.write_bytes(tmp_img.read_bytes())

            alt = f"Slide {i + 1}"
            if describe_fn:
                alt = describe_fn(dest) or alt

            rel_path = dest.relative_to(output_dir)
            md_chunks.append(f"## Slide {i + 1}\n\n![{alt}]({rel_path})\n")

    return "\n".join(md_chunks)
