import argparse
import csv
import random
from pathlib import Path


LABEL_TO_ID = {"10": 0, "11": 1, "12": 2, "13": 3, "14": 4}


def list_images(folder: Path) -> list[Path]:
    return sorted(
        path
        for path in folder.iterdir()
        if path.is_file() and path.suffix.lower() in {".jpg", ".jpeg", ".png", ".bmp", ".webp", ".tif", ".tiff"}
    )


def assign_splits(image_paths: list[Path], seed: int, val_ratio: float) -> list[tuple[Path, str]]:
    shuffled = list(image_paths)
    random.Random(seed).shuffle(shuffled)
    val_count = max(1, round(len(shuffled) * val_ratio)) if len(shuffled) >= 5 and val_ratio > 0 else 0
    val_names = {path.name for path in shuffled[:val_count]}
    return [(path, "val" if path.name in val_names else "train") for path in image_paths]


def main() -> int:
    parser = argparse.ArgumentParser(description="Create a density-classification manifest from extracted soup-image folders.")
    parser.add_argument("--input-root", required=True)
    parser.add_argument("--output", default="data/manifests/density_extracted_manifest.csv")
    parser.add_argument("--val-ratio", type=float, default=0.2)
    parser.add_argument("--seed", type=int, default=42)
    args = parser.parse_args()

    input_root = Path(args.input_root)
    rows: list[dict] = []
    for label_text, label_id in LABEL_TO_ID.items():
        folder = input_root / label_text
        if not folder.exists():
            raise FileNotFoundError(f"missing folder: {folder}")
        labeled = assign_splits(list_images(folder), args.seed, args.val_ratio)
        for image_path, split in labeled:
            rows.append(
                {
                    "image_id": image_path.stem,
                    "image_path": str(image_path.resolve()),
                    "label_text": label_text,
                    "label_id": label_id,
                    "split": split,
                }
            )

    output_path = Path(args.output)
    output_path.parent.mkdir(parents=True, exist_ok=True)
    with output_path.open("w", newline="", encoding="utf-8") as f:
        writer = csv.DictWriter(f, fieldnames=["image_id", "image_path", "label_text", "label_id", "split"])
        writer.writeheader()
        writer.writerows(rows)
    print(f"wrote manifest: {output_path} rows={len(rows)}")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
