import argparse
import json
import random
import shutil
from pathlib import Path


def split_counts(total: int, valid_ratio: float, test_ratio: float) -> tuple[int, int, int]:
    if total <= 0:
        return 0, 0, 0
    valid_count = max(1, round(total * valid_ratio)) if valid_ratio > 0 and total >= 3 else 0
    test_count = max(1, round(total * test_ratio)) if test_ratio > 0 and total >= 3 else 0
    if valid_count + test_count >= total:
        test_count = max(0, min(test_count, total - 2))
        valid_count = max(0, min(valid_count, total - test_count - 1))
    train_count = total - valid_count - test_count
    return train_count, valid_count, test_count


def subset_coco(data: dict, images: list[dict]) -> dict:
    image_ids = {image["id"] for image in images}
    return {
        "info": data.get("info", {}),
        "licenses": data.get("licenses", []),
        "categories": data.get("categories", []),
        "images": images,
        "annotations": [ann for ann in data.get("annotations", []) if ann.get("image_id") in image_ids],
    }


def write_split(source_dir: Path, output_dir: Path, split: str, data: dict, images: list[dict]) -> None:
    split_dir = output_dir / split
    split_dir.mkdir(parents=True, exist_ok=True)
    for image in images:
        source_image = source_dir / image["file_name"]
        if not source_image.exists():
            raise FileNotFoundError(f"missing image referenced by COCO json: {source_image}")
        shutil.copy2(source_image, split_dir / image["file_name"])
    (split_dir / "_annotations.coco.json").write_text(
        json.dumps(subset_coco(data, images), ensure_ascii=False, indent=2),
        encoding="utf-8",
    )


def main() -> int:
    parser = argparse.ArgumentParser(description="Create train/valid/test COCO segmentation splits from one COCO split.")
    parser.add_argument("--source-dir", default="data/soup_valid_area_dataset/train")
    parser.add_argument("--output-dir", default="data/soup_valid_area_dataset_split")
    parser.add_argument("--valid-ratio", type=float, default=0.2)
    parser.add_argument("--test-ratio", type=float, default=0.1)
    parser.add_argument("--seed", type=int, default=42)
    args = parser.parse_args()

    source_dir = Path(args.source_dir)
    output_dir = Path(args.output_dir)
    ann_path = source_dir / "_annotations.coco.json"
    data = json.loads(ann_path.read_text(encoding="utf-8"))
    images = list(data.get("images", []))
    random.Random(args.seed).shuffle(images)
    train_count, valid_count, test_count = split_counts(len(images), args.valid_ratio, args.test_ratio)
    splits = {
        "train": images[:train_count],
        "valid": images[train_count : train_count + valid_count],
        "test": images[train_count + valid_count : train_count + valid_count + test_count],
    }
    for split, split_images in splits.items():
        write_split(source_dir, output_dir, split, data, split_images)
        print(f"{split}: images={len(split_images)} annotations={len(subset_coco(data, split_images)['annotations'])}")
    return 0


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