import argparse
import csv
import json
import sys
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parents[1]))

from PIL import Image, ImageDraw

from common.image_io import list_images
from crop.rfdetr_adapter import predict_records


def draw_overlay(image_path: Path, records: list[dict], output_path: Path) -> None:
    image = Image.open(image_path).convert("RGB")
    draw = ImageDraw.Draw(image)
    for rec in records:
        x1, y1, x2, y2 = rec["bbox"]
        draw.rectangle((x1, y1, x2, y2), outline="red", width=3)
        draw.text((x1, max(0, y1 - 14)), f"{rec['confidence']:.2f}", fill="red")
    output_path.parent.mkdir(parents=True, exist_ok=True)
    image.save(output_path)


def main() -> int:
    parser = argparse.ArgumentParser(description="Run RF-DETR-Seg mask prediction for images.")
    parser.add_argument("--input", required=True)
    parser.add_argument("--checkpoint-path", default=None)
    parser.add_argument("--model-size", default="small", choices=["small", "medium"])
    parser.add_argument("--threshold", type=float, default=0.5)
    parser.add_argument("--output-dir", default="reports/crop/predictions")
    args = parser.parse_args()

    output_dir = Path(args.output_dir)
    output_dir.mkdir(parents=True, exist_ok=True)
    rows = []
    for image_path in list_images(args.input):
        records = predict_records(image_path, args.model_size, args.checkpoint_path, args.threshold)
        serializable = [{k: v for k, v in rec.items() if k != "mask"} for rec in records]
        (output_dir / f"{image_path.stem}.json").write_text(json.dumps(serializable, indent=2), encoding="utf-8")
        draw_overlay(image_path, records, output_dir / f"{image_path.stem}_overlay.jpg")
        for rec in serializable:
            rows.append({"image_id": image_path.stem, "confidence": rec["confidence"], "bbox": rec["bbox"]})

    with (output_dir / "predictions.csv").open("w", newline="", encoding="utf-8") as f:
        writer = csv.DictWriter(f, fieldnames=["image_id", "confidence", "bbox"])
        writer.writeheader()
        writer.writerows(rows)
    print(f"wrote predictions to {output_dir}")
    return 0


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