import os
import pandas as pd
from PIL import Image
import requests
from io import BytesIO
from concurrent.futures import ThreadPoolExecutor, as_completed
from tqdm.notebook import tqdm
import shutil
from sklearn.model_selection import train_test_split
import yaml
import random

# Function to download and save a single image
def download_image(args):
    filename, url, output_dir = args
    try:
        response = requests.get(url, timeout=10)
        if response.status_code == 200:
            image = Image.open(BytesIO(response.content)).convert('RGB')
            image = image.resize((256, 256))
            image.save(os.path.join(output_dir, filename))
        else:
            return f"Failed to download {filename}: Status code {response.status_code}"
    except Exception as e:
        return f"Error downloading {filename}: {e}"

# Function to download and preprocess images in parallel
def download_images(csv_file_path, output_dir):
    csv_data = pd.read_csv(csv_file_path, encoding='ISO-8859-1')
    filenames = csv_data.iloc[:, 1].apply(lambda x: x.strip('"').strip())
    urls = csv_data.iloc[:, 2]

    os.makedirs(output_dir, exist_ok=True)

    tasks = [(filename, url, output_dir) for filename, url in zip(filenames, urls)]

    with ThreadPoolExecutor(max_workers=8) as executor:
        futures = [executor.submit(download_image, task) for task in tasks]
        for future in tqdm(as_completed(futures), total=len(futures)):
            print(future.result())

# Function to organize images into directories based on their labels
def organize_images(txt_file_path, image_dir, output_dir):
    with open(txt_file_path, 'r') as file:
        txt_data = file.readlines()

    categories = {line.split()[0]: line.split()[1].strip() for line in txt_data}

    for category in set(categories.values()):
        os.makedirs(os.path.join(output_dir, category), exist_ok=True)

    for filename, category in categories.items():
        src_path = os.path.join(image_dir, filename)
        dest_path = os.path.join(output_dir, category, filename)
        if os.path.exists(src_path):
            shutil.move(src_path, dest_path)

# Mapping from old class IDs to new class names
id_to_class = {
    4: "Children's Books",
    10: 'Engineering & Transportation',
    9: 'Christian Books & Bibles',
    26: 'Sports & Outdoors',
    11: 'Health, Fitness & Dieting',
    16: 'Medical Books',
    23: 'Science & Math',
    29: 'Travel',
    2: 'Business & Money',
    7: 'Cookbooks, Food & Wine',
    19: 'Politics & Social Sciences',
    8: 'Crafts, Hobbies & Home',
    21: 'Religion & Spirituality',
    15: 'Literature & Fiction',
    13: 'Humor & Entertainment',
    14: 'Law',
    6: 'Computers & Technology',
    28: 'Test Preparation',
    1: 'Biographies & Memoirs',
    0: 'Arts & Photography',
    18: 'Parenting & Relationships',
    22: 'Romance',
    12: 'History',
    5: 'Comics & Graphic Novels',
    20: 'Reference',
    27: 'Teen & Young Adult',
    25: 'Self-Help',
    3: 'Calendars',
    24: 'Science Fiction & Fantasy',
    17: 'Mystery, Thriller & Suspense'
}
# Mapping from new class names to merged classes, excluding "Comics & Graphic Novels"
# Mapping from new class names to merged classes, excluding "Comics & Graphic Novels"
merge_mapping = {
    'Fiction': ['15', '24', '17','27'],  # Fiction
    'Religion & History': ['9', '21','12','1'],  # Religion
}


def merge_classes(data_dir, merge_mapping, id_to_class):
    for new_class, old_classes in merge_mapping.items():
        new_class_dir = os.path.join(data_dir, new_class)
        os.makedirs(new_class_dir, exist_ok=True)
        for old_class_id in old_classes:
            old_class_name = id_to_class[int(old_class_id)]
            old_class_dir = os.path.join(data_dir, old_class_name)
            if os.path.exists(old_class_dir):
                for file_name in os.listdir(old_class_dir):
                    shutil.move(os.path.join(old_class_dir, file_name), new_class_dir)
                shutil.rmtree(old_class_dir)

def remove_hidden_dirs(data_dir):
    for root, dirs, files in os.walk(data_dir):
        for dir_name in dirs:
            if dir_name.startswith('.'):
                shutil.rmtree(os.path.join(root, dir_name))

def balance_dataset(data_dir):
    # Get the list of all classes
    classes = [d for d in os.listdir(data_dir) if os.path.isdir(os.path.join(data_dir, d))]
    class_counts = {cls: len([f for f in os.listdir(os.path.join(data_dir, cls)) if os.path.isfile(os.path.join(data_dir, cls, f))]) for cls in classes}
    max_count = max(class_counts.values())

    for cls in classes:
        class_dir = os.path.join(data_dir, cls)
        images = [f for f in os.listdir(class_dir) if os.path.isfile(os.path.join(class_dir, f))]
        num_images = len(images)
        if num_images < max_count:
            # Over-sample
            images_to_add = random.choices(images, k=max_count - num_images)
            for img in images_to_add:
                shutil.copy(os.path.join(class_dir, img), os.path.join(class_dir, f"copy_{img}"))
        elif num_images > max_count:
            # Under-sample
            images_to_remove = random.sample(images, num_images - max_count)
            for img in images_to_remove:
                os.remove(os.path.join(class_dir, img))

def create_validation_set(data_dir, val_size=0.2):
    classes = [d for d in os.listdir(data_dir) if os.path.isdir(os.path.join(data_dir, d))]
    for cls in classes:
        class_dir = os.path.join(data_dir, cls)
        images = [f for f in os.listdir(class_dir) if os.path.isfile(os.path.join(class_dir, f))]
        train_images, val_images = train_test_split(images, test_size=val_size, random_state=42)

        val_dir = os.path.join('/content/datasets/imagenet/val', cls)
        os.makedirs(val_dir, exist_ok=True)

        for img in val_images:
            shutil.move(os.path.join(class_dir, img), os.path.join(val_dir, img))

# Paths
train_csv_file_path = '/content/drive/MyDrive/data/book30-listing-train.csv'
train_txt_file_path = '/content/drive/MyDrive/data/bookcover30-labels-train.txt'
test_csv_file_path = '/content/drive/MyDrive/data/book30-listing-test.csv'
test_txt_file_path = '/content/drive/MyDrive/data/bookcover30-labels-test.txt'

# Directories for downloading and organizing images
local_train_images_dir = '/content/train_images'
local_test_images_dir = '/content/test_images'
train_dir = '/content/datasets/imagenet/train'
test_dir = '/content/datasets/imagenet/test'
val_dir = '/content/datasets/imagenet/val'

# Download images
download_images(train_csv_file_path, local_train_images_dir)
download_images(test_csv_file_path, local_test_images_dir)

# Organize images into categorized directories
organize_images(train_txt_file_path, local_train_images_dir, train_dir)
organize_images(test_txt_file_path, local_test_images_dir, test_dir)

# Rename directories from ID to class name
for old_id, new_name in id_to_class.items():
    old_train_dir = os.path.join(train_dir, str(old_id))
    new_train_dir = os.path.join(train_dir, new_name)
    if os.path.exists(old_train_dir):
        os.rename(old_train_dir, new_train_dir)

    old_test_dir = os.path.join(test_dir, str(old_id))
    new_test_dir = os.path.join(test_dir, new_name)
    if os.path.exists(old_test_dir):
        os.rename(old_test_dir, new_test_dir)

# Merge classes
merge_classes(train_dir, merge_mapping, id_to_class)
merge_classes(test_dir, merge_mapping, id_to_class)

# Remove unwanted classes
unwanted_classes = ['Calendars', 'Reference', 'Arts & Photography', 'Biographies & Memoirs', 'Travel', 'Humor & Entertainment','Children\'s Books', 'Cookbooks, Food & Wine'
,'Crafts, Hobbies & Home','Self-Help','Health, Fitness & Dieting','Parenting & Relationships','Sports & Outdoors']
for cls in unwanted_classes:
    shutil.rmtree(os.path.join(train_dir, cls), ignore_errors=True)
    shutil.rmtree(os.path.join(test_dir, cls), ignore_errors=True)


# Create validation set from training data
create_validation_set('/content/datasets/imagenet/train')

# Remove hidden directories
remove_hidden_dirs(train_dir)
remove_hidden_dirs(val_dir)
remove_hidden_dirs(test_dir)


# Balance the dataset
balance_dataset(train_dir)
balance_dataset(val_dir)
balance_dataset(test_dir)

data_yaml_content = """
path: /content/datasets  # Base path for the dataset
train: /content/datasets/imagenet/train  # Path to training data
val: /content/datasets/imagenet/val  # Path to validation data
test: /content/datasets/imagenet/test  # Path to testing data (optional)

nc: 10  # Number of classes
names:  # List of class names
  - Art & Entertainment
  - Practical Skills
  - Romance
  - Sports & Outdoors
  - Comics & Graphic Novels
  - Religion
  - Fiction
  - School and Professions
  - Personal Development
  - Children's Books
"""

# Write the content to data.yaml file
with open('/content/datasets/data.yaml', 'w') as file:
    file.write(data_yaml_content)

