#================================================
# Practical Image Brokering at Netrel Language Brokering
# Dr. Hager Saleh
# Rami AbuLaban - 233000631
#================================================



import pandas as pd
import numpy as np
from sklearn.model_selection import train_test_split, GridSearchCV
from sklearn.preprocessing import StandardScaler, LabelEncoder
from sklearn.linear_model import LogisticRegression
from sklearn.tree import DecisionTreeClassifier
from sklearn.ensemble import RandomForestClassifier
from sklearn.svm import SVC
from sklearn.naive_bayes import GaussianNB
from sklearn.neighbors import KNeighborsClassifier
from sklearn.metrics import confusion_matrix, accuracy_score, precision_score, recall_score, f1_score, classification_report
import matplotlib.pyplot as plt
import seaborn as sns
from sklearn.feature_extraction.text import TfidfVectorizer, CountVectorizer
import re
import nltk
from nltk.corpus import stopwords
from nltk.tokenize import word_tokenize
from nltk.stem import PorterStemmer
from nltk.stem import WordNetLemmatizer
from imblearn.over_sampling import RandomOverSampler, SMOTE
from imblearn.under_sampling import RandomUnderSampler
from collections import Counter


# Download NLTK resources (only first time)
try: # Wrap in try-except to avoid repeated downloads
    nltk.data.find('punkt')
    nltk.data.find('stopwords')
    nltk.data.find('wordnet')
except LookupError:
    nltk.download('punkt')
    nltk.download('stopwords')
    nltk.download('wordnet')

# Function to evaluate model
def evaluate_model(model_name, y_true, y_pred):
    accuracy = accuracy_score(y_true, y_pred)
    precision = precision_score(y_true, y_pred, average='weighted')
    recall = recall_score(y_true, y_pred, average='weighted')
    f1 = f1_score(y_true, y_pred, average='weighted')
    cm = confusion_matrix(y_true, y_pred)
    report = classification_report(y_true, y_pred)

    metrics = {
        'Model Name': model_name,
        'Accuracy': accuracy,
        'Precision': precision,
        'Recall': recall,
        'F1 Score': f1,
        'Classification Report': report
    }
    return metrics


# Function to apply Random forest with grid search
def train_random_forest_classifier_with_grid_search(X_train_vec, y_train, X_test_vec, y_test, evaluate_model_func):
    rf_classifier = RandomForestClassifier(random_state=42)

    param_grid = {
        'n_estimators': [100, 200, 300],
        'max_depth': [None, 10, 20, 30],
    }

    grid_search = GridSearchCV(estimator=rf_classifier, param_grid=param_grid, cv=5, scoring='accuracy', n_jobs=-1)
    grid_search.fit(X_train_vec, y_train)
    best_rf_classifier = grid_search.best_estimator_

    y_pred = best_rf_classifier.predict(X_test_vec)
    evaluation_results = evaluate_model_func('RandomForestClassifier', y_test, y_pred)

    for key, value in evaluation_results.items():
        if key == 'Classification Report':
            print(value)  # Print report separately
        else:
            print(f"{key}: {value:.4f}" if isinstance(value, float) else f"{key}: \n{value}")

    print("\nBest hyperparameters found by GridSearchCV:")
    print(grid_search.best_params_)






# Preprocessing English text
stemmer = PorterStemmer()
stop_words = set(stopwords.words('english'))

def preprocess_text(text):
    text = text.lower()
    text = re.sub(r'http\S+|www.\S+', '', text)  # Remove URLs
    text = re.sub(r'<.*?>', '', text)            # Remove HTML tags
    text = re.sub(r'[^a-zA-Z\s]', '', text)       # Remove special chars, numbers, punctuation
    tokens = word_tokenize(text)                 # Tokenization
    tokens = [word for word in tokens if word not in stop_words]  # Remove Stop Words
    stemmed_tokens = [stemmer.stem(word) for word in tokens]     # Stemming
    clean_text = ' '.join(stemmed_tokens)                       # Reconstruct text
    return clean_text




# Preprocessing English text
stemmer = PorterStemmer()
stop_words = set(stopwords.words('english'))

def preprocess_text(text):
    text = text.lower()
    text = re.sub(r'http\S+|www.\S+', '', text)  # Remove URLs
    text = re.sub(r'<.*?>', '', text)            # Remove HTML tags
    text = re.sub(r'[^a-zA-Z\s]', '', text)      # Remove special chars, numbers, punctuation
    tokens = text.split()                        # split
    tokens = [word for word in tokens if word not in stop_words]  # Remove Stop Words
    stemmed_tokens = [stemmer.stem(word) for word in tokens]     # Stemming
    clean_text = ' '.join(stemmed_tokens)                       # Reconstruct text
    return clean_text

# NOTE: استخدمت text.split() بدل word_tokenize لتجنب مشاكل punkt في أجهزة غير مهيأة





# --- Main execution ---
df = pd.read_csv("unbalanceddataset.csv")  # Replace "unbalanceddataset.csv" with your file
df["clean_text"] = df["text"].apply(preprocess_text)
df = df.dropna()


X = df['clean_text']
y = df['sentiment']

X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.30, random_state=42, stratify=y)

vectorizer = TfidfVectorizer() # Or CountVectorizer()
X_train_vec = vectorizer.fit_transform(X_train)
X_test_vec = vectorizer.transform(X_test)


# ---  Resampling examples (Uncomment as needed) ---

# # RandomOverSampler
# ros = RandomOverSampler(random_state=42)
# X_ros, y_ros = ros.fit_resample(X_train_vec, y_train)
# print("After RandomOversampling:", Counter(y_ros))
# train_random_forest_classifier_with_grid_search(X_ros, y_ros, X_test_vec, y_test, evaluate_model)

# # RandomUnderSampler
# rus = RandomUnderSampler(random_state=42)
# X_res, y_res = rus.fit_resample(X_train_vec, y_train)
# print("After undersampling:", Counter(y_res))
# train_random_forest_classifier_with_grid_search(X_res, y_res, X_test_vec, y_test, evaluate_model)

# # SMOTE
# smote = SMOTE(random_state=42)
# X_smote, y_smote = smote.fit_resample(X_train_vec, y_train)
# print("After SMOTE:", Counter(y_smote))
# train_random_forest_classifier_with_grid_search(X_smote, y_smote, X_test_vec, y_test, evaluate_model)



# --- Baseline model without resampling ---
train_random_forest_classifier_with_grid_search(X_train_vec, y_train, X_test_vec, y_test, evaluate_model)