import argparse
import os
import sys
import logging
from typing import List, Tuple

import torch
from PIL import Image
from torchvision import transforms

from src.models import create_model


def parse_args() -> argparse.Namespace:
    p = argparse.ArgumentParser(description="推理脚本（单图/目录）")
    p.add_argument("--checkpoint", type=str, required=True, help="模型检查点路径")
    p.add_argument("--image", type=str, default=None, help="单张图片路径")
    p.add_argument("--image-dir", type=str, default=None, help="图片目录")
    p.add_argument("--img-size", type=int, default=224, help="输入图像边长")
    p.add_argument("--out", type=str, default=None, help="可选：将结果写入 CSV 路径")
    p.add_argument("--log-file", type=str, default=None, help="可选日志文件路径")
    return p.parse_args()


def _build_tf(img_size: int) -> transforms.Compose:
    try:
        from torchvision.transforms import InterpolationMode
        bicubic = InterpolationMode.BICUBIC
    except Exception:
        bicubic = 3  # type: ignore
    return transforms.Compose([
        transforms.Resize(int(img_size * 1.14), interpolation=bicubic),
        transforms.CenterCrop(img_size),
        transforms.ToTensor(),
        transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),
    ])


def _list_images(path: str) -> List[str]:
    exts = {".jpg", ".jpeg", ".png", ".bmp", ".gif"}
    files: List[str] = []
    for root, _, names in os.walk(path):
        for n in names:
            if os.path.splitext(n)[1].lower() in exts:
                files.append(os.path.join(root, n))
    files.sort()
    return files


def main() -> None:
    args = parse_args()
    handlers = [logging.StreamHandler(sys.stdout)]
    if args.log_file:
        os.makedirs(os.path.dirname(args.log_file) or ".", exist_ok=True)
        handlers.append(logging.FileHandler(args.log_file, mode="w"))
    logging.basicConfig(level=logging.INFO, format="%(asctime)s | %(levelname)s | %(message)s", handlers=handlers)

    if not args.image and not args.image_dir:
        raise ValueError("需要指定 --image 或 --image-dir")

    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    ckpt = torch.load(args.checkpoint, map_location=device)
    class_names = tuple(ckpt.get("class_names", ("cats", "dogs")))
    num_classes = len(class_names)
    logging.info("加载检查点: %s | 类别: %s", args.checkpoint, class_names)

    model = create_model(num_classes=num_classes)
    model.load_state_dict(ckpt["model_state"], strict=True)
    model.to(device)
    model.eval()

    tf = _build_tf(args.img_size)

    paths: List[str]
    if args.image_dir:
        paths = _list_images(args.image_dir)
    else:
        paths = [args.image]
    if not paths:
        logging.info("未找到任何图片")
        return

    results: List[Tuple[str, str, float]] = []
    with torch.no_grad():
        for p in paths:
            img = Image.open(p).convert("RGB")
            x = tf(img).unsqueeze(0).to(device)
            logits = model(x)
            prob = torch.softmax(logits, dim=1)[0]
            score, pred = prob.max(dim=0)
            label = class_names[int(pred)]
            results.append((p, label, float(score)))

    # 输出
    for p, label, score in results:
        logging.info("%s\t%s\t%.4f", p, label, score)

    if args.out:
        import csv
        os.makedirs(os.path.dirname(args.out) or ".", exist_ok=True)
        with open(args.out, "w", newline="") as f:
            w = csv.writer(f)
            w.writerow(["path", "label", "score"])
            for r in results:
                w.writerow(r)
        logging.info("已保存: %s", args.out)


if __name__ == "__main__":
    main()
