import json
import copy
from pathlib import Path

import numpy as np
import pandas as pd
import torch
from sklearn.utils.class_weight import compute_class_weight
from PIL import Image
from torch import nn
from torch.utils.data import DataLoader, Dataset
from torchvision import models, transforms


class ManifestImageDataset(Dataset):
    def __init__(
        self,
        manifest: str | Path,
        image_dir: str | Path,
        target_column: str,
        image_path_column: str = "crop_image_path",
        transform=None,
    ) -> None:
        self.manifest = Path(manifest)
        self.image_dir = Path(image_dir)
        self.df = pd.read_csv(self.manifest)
        self.target_column = target_column
        self.image_path_column = image_path_column if image_path_column in self.df.columns else "image_path"
        self.transform = transform or default_transforms()
        if target_column not in self.df.columns:
            raise ValueError(f"manifest is missing target column: {target_column}")
        if "split" in self.df.columns:
            self.df["split"] = self.df["split"].fillna("train")

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

    def __getitem__(self, index: int):
        row = self.df.iloc[index]
        image_path = Path(str(row[self.image_path_column]))
        if not image_path.is_absolute():
            image_path = image_path if image_path.exists() else self.image_dir / image_path.name
        image = Image.open(image_path).convert("RGB")
        return self.transform(image), int(row[self.target_column])


def default_transforms(image_size: int = 224):
    return transforms.Compose(
        [
            transforms.Resize((image_size, image_size)),
            transforms.ToTensor(),
            transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
        ]
    )


def build_model(model_name: str, num_classes: int) -> nn.Module:
    if model_name != "resnet18":
        raise ValueError("Only resnet18 is implemented in the initial scaffold.")
    model = models.resnet18(weights=models.ResNet18_Weights.DEFAULT)
    model.fc = nn.Linear(model.fc.in_features, num_classes)
    return model


def choose_device(config_device: str = "auto") -> torch.device:
    if config_device == "auto":
        return torch.device("cuda" if torch.cuda.is_available() else "cpu")
    return torch.device(config_device)


def train_classifier(
    manifest: str | Path,
    image_dir: str | Path,
    target_column: str,
    output_dir: str | Path,
    report_dir: str | Path,
    model_name: str,
    num_classes: int,
    epochs: int,
    batch_size: int,
    lr: float,
    device_name: str = "auto",
    labels: dict[int, str] | None = None,
) -> dict:
    dataset = ManifestImageDataset(manifest, image_dir, target_column)
    if len(dataset) == 0:
        raise ValueError("manifest has no rows")
    train_df = dataset.df[dataset.df["split"] == "train"].reset_index(drop=True) if "split" in dataset.df.columns else dataset.df
    val_df = dataset.df[dataset.df["split"] == "val"].reset_index(drop=True) if "split" in dataset.df.columns else dataset.df.iloc[0:0]
    train_dataset = ManifestImageDataset(manifest, image_dir, target_column, transform=dataset.transform)
    train_dataset.df = train_df
    val_dataset = ManifestImageDataset(manifest, image_dir, target_column, transform=dataset.transform)
    val_dataset.df = val_df
    if len(train_dataset) == 0:
        raise ValueError("manifest has no train rows")

    device = choose_device(device_name)
    train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=0)
    val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=0) if len(val_dataset) else None
    model = build_model(model_name, num_classes).to(device)
    optimizer = torch.optim.AdamW(model.parameters(), lr=lr)
    class_weights = compute_class_weight(
        class_weight="balanced",
        classes=np.arange(num_classes),
        y=train_df[target_column].astype(int).tolist(),
    )
    criterion = nn.CrossEntropyLoss(weight=torch.tensor(class_weights, dtype=torch.float32, device=device))

    history = []
    best_state_dict = None
    best_epoch = None
    best_val_accuracy = None
    for epoch in range(epochs):
        model.train()
        total_loss = 0.0
        correct = 0
        seen = 0
        for images, targets in train_loader:
            images = images.to(device)
            targets = targets.to(device)
            optimizer.zero_grad(set_to_none=True)
            logits = model(images)
            loss = criterion(logits, targets)
            loss.backward()
            optimizer.step()
            total_loss += float(loss.item()) * targets.size(0)
            correct += int((logits.argmax(dim=1) == targets).sum().item())
            seen += int(targets.size(0))
        val_accuracy = None
        if val_loader is not None:
            model.eval()
            val_correct = 0
            val_seen = 0
            with torch.inference_mode():
                for images, targets in val_loader:
                    images = images.to(device)
                    targets = targets.to(device)
                    logits = model(images)
                    val_correct += int((logits.argmax(dim=1) == targets).sum().item())
                    val_seen += int(targets.size(0))
            val_accuracy = val_correct / val_seen if val_seen else None
        history.append(
            {
                "epoch": epoch + 1,
                "loss": total_loss / seen,
                "accuracy": correct / seen,
                "val_accuracy": val_accuracy,
            }
        )
        score = val_accuracy if val_accuracy is not None else (correct / seen)
        if best_state_dict is None or score > best_val_accuracy:
            best_state_dict = copy.deepcopy(model.state_dict())
            best_epoch = epoch + 1
            best_val_accuracy = score

    output_dir = Path(output_dir)
    report_dir = Path(report_dir)
    output_dir.mkdir(parents=True, exist_ok=True)
    report_dir.mkdir(parents=True, exist_ok=True)
    model_path = output_dir / "model.pt"
    torch.save(
        {
            "model_state_dict": best_state_dict or model.state_dict(),
            "model_name": model_name,
            "num_classes": num_classes,
            "labels": labels or {idx: str(idx) for idx in range(num_classes)},
            "image_size": 224,
            "best_epoch": best_epoch,
            "best_val_accuracy": best_val_accuracy,
        },
        model_path,
    )
    report = {
        "device": str(device),
        "model_path": str(model_path),
        "best_epoch": best_epoch,
        "best_val_accuracy": best_val_accuracy,
        "history": history,
    }
    (report_dir / "train_history.json").write_text(json.dumps(report, indent=2), encoding="utf-8")
    return report
