from rdkit import Chem
from rdkit.Chem import Descriptors, Lipinski
import numpy as np
import pandas as pd

def filter_lipinski(df, threshold=4):
    def calculate_descriptors(df):
        mols, valid_indices = [], []

        for idx, smiles in enumerate(df['canonical_smiles']):
            mol = Chem.MolFromSmiles(smiles)
            if mol:
                mols.append(mol)
                valid_indices.append(idx)

        if not mols:
            return pd.DataFrame()  # Return empty if no valid molecules

        data = [
            [
                Descriptors.MolWt(m),
                Descriptors.MolLogP(m),
                Lipinski.NumHDonors(m),
                Lipinski.NumHAcceptors(m)
            ]
            for m in mols
        ]

        descriptors = pd.DataFrame(
            data,
            columns=["MW", "LogP", "NumHDonors", "NumHAcceptors"]
        )

        # Align with filtered input dataframe
        df_filtered = df.iloc[valid_indices].reset_index(drop=True)
        return pd.concat([df_filtered, descriptors], axis=1)

    df_combined = calculate_descriptors(df)
    if df_combined.empty:
        return pd.DataFrame()

    df_combined = df_combined.copy()

    def rules_passed(row):
        return sum([
            row["MW"] <= 500,
            row["LogP"] <= 5,
            row["NumHDonors"] < 5,
            row["NumHAcceptors"] <= 10
        ])

    df_combined["Lipinski_Pass_Count"] = df_combined.apply(rules_passed, axis=1)

    filtered = df_combined[df_combined["Lipinski_Pass_Count"] >= threshold]

    return filtered.drop(columns=["Lipinski_Pass_Count"]).reset_index(drop=True)
