import subprocess
import sys

import importlib
import subprocess
import sys

def install(package):
    subprocess.check_call([sys.executable, "-m", "pip", "install","--quiet", package])

def is_installed(package):
    try:
        importlib.import_module(package)
        return True
    except ImportError:
        return False

packages = ["rdkit", "duckdb", "torch", "scikit-learn"]
for package in packages:
    if not is_installed(package):
        install(package)
    

from typing import Literal, Optional
from math import log
import warnings
import random

import numpy as np # linear algebra
import pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)
import pickle

import gc
import os
os.environ["TOKENIZERS_PARALLELISM"] = "false"


from rdkit import Chem
from rdkit.Chem import AllChem

warnings.filterwarnings("ignore")

from torchmetrics.classification import MultilabelAveragePrecision
from torch import Tensor
import random

from sklearn.model_selection import train_test_split


def get_cv_indexes(train_data: pd.DataFrame, random_seed: Optional["int"]=42):
    if random_seed is not None:
        random.seed(random_seed)
    # Getting dict of blocks
    bbs_dicts = []

    for i in range(1,4):
        with open(f'/kaggle/input/belka-shrinking-the-dataset/train_dicts/BBs_dict_{i}.p', 'rb') as file:
            bbs_dicts.append(pickle.load(file))


    bbs1 = set(bbs_dicts[0].keys())
    bbs2 = set(bbs_dicts[1].keys())
    bbs3 = set(bbs_dicts[2].keys())

    bbs2_and_bbs3 = bbs2 & bbs3
    bbs2_not_bbs3 = bbs3 - bbs2

    bbs1_val = random.choices(list(bbs1), k = 17)
    bbs2_val = random.choices(list(bbs2_and_bbs3), k = 34)
    bbs3_val = random.choices(list(bbs2_not_bbs3), k = 2)

    bbs1_val_encd = [bbs_dicts[0][key] for key in bbs1_val]
    bbs2_val_encd = [bbs_dicts[1][key] for key in bbs2_val]
    bbs3_val_encd = [bbs_dicts[2][key] for key in bbs3_val]


    index_eval_non_shared = train_data[train_data.buildingblock1_smiles.isin(bbs1_val_encd) 
               & train_data.buildingblock2_smiles.isin(bbs2_val_encd)
              & train_data.buildingblock3_smiles.isin(bbs2_val_encd+bbs3_val_encd)].index

    index_mixshare = train_data[train_data.buildingblock1_smiles.isin(bbs1_val_encd) 
               | train_data.buildingblock2_smiles.isin(bbs2_val_encd)
              | train_data.buildingblock3_smiles.isin(bbs2_val_encd+bbs3_val_encd)].index

    index_to_drop = index_mixshare.difference(index_eval_non_shared)

    index_random = pd.Index(random.choices(train_data.index.difference(index_mixshare), k = int(len(train_data)*0.003)))
    train_index = train_data.index.difference(index_random.union(index_mixshare))
    
    
    return train_index, index_eval_non_shared, index_random



if __name__ == "__main__":
    import duckdb
    train_path = '/kaggle/input/belka-shrinking-the-dataset/train.parquet'
    test_path = '/kaggle/input/leash-BELKA/test.parquet'

    con = duckdb.connect()


    train_data = con.query(f"""SELECT *
                            FROM parquet_scan('{train_path}')
                            """).df()

    print(get_cv_indexes(train_data))