import argparse
import csv
import sys
from pathlib import Path

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

from common.image_io import list_images, read_rgb
from crop.rfdetr_adapter import predict_records
from crop.utils import expand_bbox, save_fallback_crop


def crop_one_image(
    image_path: str | Path,
    output_dir: str | Path,
    checkpoint_path: str | None = None,
    model_size: str = "small",
    threshold: float = 0.5,
    low_confidence_threshold: float = 0.35,
    expand_ratio: float = 0.08,
    predictor=None,
) -> dict:
    image_path = Path(image_path)
    output_dir = Path(output_dir)
    image = read_rgb(image_path)
    crop_path = output_dir / f"{image_path.stem}_potcrop.jpg"
    try:
        records = predictor(image_path) if predictor else predict_records(image_path, model_size, checkpoint_path, threshold)
        if not records:
            raise RuntimeError("no pot_inner detection")
        best = max(records, key=lambda item: item.get("confidence", 0.0))
        bbox = tuple(int(round(v)) for v in best["bbox"])
        bbox = expand_bbox(bbox, image.size, expand_ratio)
        crop_path.parent.mkdir(parents=True, exist_ok=True)
        image.crop(bbox).save(crop_path)
        confidence = float(best.get("confidence", 0.0))
        status = 1 if confidence >= low_confidence_threshold else 2
        error_message = ""
    except Exception as exc:
        bbox = save_fallback_crop(image, crop_path)
        confidence = 0.0
        status = 0
        error_message = str(exc)
    x1, y1, x2, y2 = bbox
    return {
        "image_id": image_path.stem,
        "source_image_path": str(image_path),
        "crop_image_path": str(crop_path),
        "crop_status": status,
        "crop_confidence": confidence,
        "x1": x1,
        "y1": y1,
        "x2": x2,
        "y2": y2,
        "error_message": error_message,
    }


def main() -> int:
    parser = argparse.ArgumentParser(description="Generate pot_inner crop images with RF-DETR-Seg.")
    parser.add_argument("--input", required=True)
    parser.add_argument("--output-dir", default="data/processed_images/pot_crop")
    parser.add_argument("--report", default="reports/crop/crop_result.csv")
    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("--low-confidence-threshold", type=float, default=0.35)
    parser.add_argument("--expand-ratio", type=float, default=0.08)
    args = parser.parse_args()

    rows = [
        crop_one_image(
            image_path,
            args.output_dir,
            args.checkpoint_path,
            args.model_size,
            args.threshold,
            args.low_confidence_threshold,
            args.expand_ratio,
        )
        for image_path in list_images(args.input)
    ]
    report = Path(args.report)
    report.parent.mkdir(parents=True, exist_ok=True)
    with report.open("w", newline="", encoding="utf-8") as f:
        writer = csv.DictWriter(f, fieldnames=list(rows[0].keys()) if rows else [
            "image_id", "source_image_path", "crop_image_path", "crop_status", "crop_confidence", "x1", "y1", "x2", "y2", "error_message"
        ])
        writer.writeheader()
        writer.writerows(rows)
    print(f"wrote crop report: {report}")
    return 0


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