import os
import logging
import random
from typing import Dict, Tuple, List

from PIL import Image
import torch
from torch.utils.data import DataLoader
from torchvision import transforms


IMAGENET_MEAN = (0.485, 0.456, 0.406)
IMAGENET_STD = (0.229, 0.224, 0.225)


def _build_transforms(img_size: int) -> Tuple[transforms.Compose, transforms.Compose]:
    try:
        from torchvision.transforms import InterpolationMode
        bicubic = InterpolationMode.BICUBIC
    except Exception:  # pragma: no cover
        bicubic = 3  # type: ignore

    train_tf = transforms.Compose([
        transforms.RandomResizedCrop(img_size, scale=(0.8, 1.0), interpolation=bicubic),
        transforms.RandomHorizontalFlip(),
        transforms.ToTensor(),
        transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),
    ])
    val_tf = transforms.Compose([
        transforms.Resize(int(img_size * 1.14), interpolation=bicubic),
        transforms.CenterCrop(img_size),
        transforms.ToTensor(),
        transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),
    ])
    return train_tf, val_tf


def _create_from_kaggle_flat(data_dir: str, img_size: int) -> Tuple[Dict[str, torch.utils.data.Dataset], Tuple[str, ...]]:
    train_root = os.path.join(data_dir, "train")
    if not os.path.isdir(train_root):
        raise FileNotFoundError(f"未找到数据集目录: {data_dir}")

    def _glob(prefix: str) -> List[str]:
        return sorted(
            os.path.join(train_root, n)
            for n in os.listdir(train_root)
            if n.startswith(prefix) and n.lower().endswith(".jpg")
        )

    cats = _glob("cat.")
    dogs = _glob("dog.")
    if not cats or not dogs:
        raise RuntimeError(f"期望 Kaggle 布局 (cat.*.jpg/dog.*.jpg) 于: {train_root}")

    per_class = {"train": 1000, "val": 500, "test": 500}
    seed = 42
    logging.info("检测到 Kaggle 扁平布局：将从 train 中切分 1000/500/500（每类），随机种子=42")

    def _sample(items: List[str], counts: Dict[str, int], seed_offset: int = 0) -> Dict[str, List[str]]:
        total = sum(counts.values())
        if len(items) < total:
            raise ValueError(f"样本不足：需要 {total}，实际 {len(items)}")
        rnd = random.Random(seed + seed_offset)
        picks = rnd.sample(items, total)
        res: Dict[str, List[str]] = {}
        i = 0
        for k in ("train", "val", "test"):
            n = counts[k]
            res[k] = picks[i:i+n]
            i += n
        return res

    c_splits = _sample(cats, per_class, 0)
    d_splits = _sample(dogs, per_class, 1)
    for sp in ("train", "val", "test"):
        logging.info(
            f"{sp}: cats={len(c_splits[sp])} 张, dogs={len(d_splits[sp])} 张, 合计={len(c_splits[sp])+len(d_splits[sp])} 张"
        )

    train_tf, val_tf = _build_transforms(img_size)

    class PathsDataset(torch.utils.data.Dataset):
        def __init__(self, paths: List[str], label: int, tf: transforms.Compose) -> None:
            self.paths = paths
            self.label = label
            self.tf = tf

        def __len__(self) -> int:
            return len(self.paths)

        def __getitem__(self, idx: int):
            p = self.paths[idx]
            x = Image.open(p).convert("RGB")
            x = self.tf(x)
            y = self.label
            return x, y

    datasets_map: Dict[str, torch.utils.data.Dataset] = {}
    for sp in ("train", "val", "test"):
        tf = train_tf if sp == "train" else val_tf
        ds_c = PathsDataset(c_splits[sp], 0, tf)
        ds_d = PathsDataset(d_splits[sp], 1, tf)
        datasets_map[sp] = torch.utils.data.ConcatDataset([ds_c, ds_d])

    return datasets_map, ("cats", "dogs")


def create_dataloaders(
    data_dir: str,
    img_size: int = 224,
    batch_size: int = 32,
    workers: int = 4,
) -> Tuple[Dict[str, DataLoader], Tuple[str, ...]]:
    # 仅支持 Kaggle 扁平布局（dataset/train 下 cat.*.jpg 与 dog.*.jpg）
    splits, class_names = _create_from_kaggle_flat(data_dir, img_size)

    loaders: Dict[str, DataLoader] = {}
    pin = torch.cuda.is_available()
    persistent = workers > 0
    for sp, ds in splits.items():
        loaders[sp] = DataLoader(
            ds,
            batch_size=batch_size,
            shuffle=(sp == "train"),
            num_workers=workers,
            pin_memory=pin,
            persistent_workers=persistent,
        )
    return loaders, class_names
