import csv
import os
import random
from collections import OrderedDict
from typing import List, Set, Optional, Union
import numpy as np


class Dataset:
    """
    Represents a set of training/testing data. self.train is a list of
    Examples, as is self.dev and self.test.
    """

    def __init__(self, name: str = None, include_test: bool = False):
        self.name = name
        self.splits = OrderedDict()
        self.train = []
        self.dev = []

        self.splits['train'] = self.train
        self.splits['dev'] = self.dev

        if include_test:
            self.test = []
            self.splits["test"] = self.test

    def shuffle(self) -> None:
        for split_name in self.splits:
            random.shuffle(self.splits[split_name])


class Example:
    """
    Represents a document (list of words) with a corresponding label.
    """

    def __init__(self, words: List[str], label: Optional[int]):
        self.words = words
        self.label = label


class Classifier:
    """
    If the filter_stop_words is True, self.stop_words will be a
    set of english stop words. You should filter these from your input words.
    If it is False, no filtering is required.
    You can use util.remove_stop_words().

    If use_bigrams flag is true, your methods must use bigram features instead
    of the usual bag-of-words (unigrams).
    HINT: Remember to add start and end tokens in the bigram implementation.
    HINT: When doing add-1 smoothing with bigrams, V = # unique bigrams in data.

    NOTE: Only one of stop_words and use_bigrams will be used at once.
    """

    def __init__(self,
                 filter_stop_words: bool = False,
                 use_bigrams: bool = False):
        self.filter_stop_words = filter_stop_words
        self.stop_words = set(read_file('data/english.stop')) \
            if self.filter_stop_words else None
        self.use_bigrams = use_bigrams

    def train(self, examples: List[Example]) -> None:
        raise NotImplementedError("You need to override this!")

    def classify(self, examples: List[Example],
                 return_scores: bool = False) -> Union[List[int], List[float]]:
        raise NotImplementedError("You need to override this!")


def segment_words(s: str) -> List[str]:
    """
    Splits lines into words on whitespace.

    Args:
        s (str): The line to segment.

    Returns:
        list[str]: List of segmented words in the line.
    """
    return s.split()


def read_file(file_name: str,
              encoding: str = "utf8", mode: str = "word") -> List[str]:
    """
    Reads lines or words from a file.

    Args:
        file_name (str): The filename to read from.
        encoding (str): The encoding to use for reading the file.
        mode (str): How to extract the file contents. If "word", outputs a list
        of words (as strings). If "line", outputs a list of lines (as strings).

    Returns:
        list[str]: List of segmented words or lines in the file.
    """
    outputs = []
    with open(file_name, encoding=encoding) as f:
        for line in f:
            if mode == "word":
                outputs.extend(segment_words(line))
            elif mode == "line":
                outputs.append(line)
            else:
                raise ValueError("Invalid mode: {}".format(mode))
    return outputs


def remove_stop_words(words: List[str], stop_list: Set[str]) -> List[str]:
    """
    Filters stop words and all-whitespace words.

    Args:
        words (list[str]): List of words.
        stop_list (set[str]): Set of stop words.

    Returns:
        list[str]: List of filtered words.
    """
    return [w for w in words if w not in stop_list and w.strip() != '']


def calculate_accuracy(data: List[Example], classifier: Classifier) -> float:
    """
    Calculates the classifier's accuracy on a provided dataset.

    Args:
        data (list[Example]): List of examples to evaluate on.
        classifier (Classifier): Classifier to compute accuracy for.

    Returns:
        float: Accuracy of the given classifier on the given dataset.
    """

    if len(data) == 0:
        return 0.0

    predictions = classifier.classify(data)

    correct = 0.0
    for prediction, example in zip(predictions, data):
        if example.label == prediction:
            correct += 1.0
    return correct / len(data)


def load_data(data_dir: str,
              include_test: bool = False,
              dataset_name: str = None) -> Dataset:
    """
    Loads data into a Dataset object.

    Args:
        data_dir (str): Path to the directory containing the data.
        include_test (bool): Whether to load test data (this will only be
                    available to the teaching staff in the autograder).
        dataset_name (str): Optional, name to give to the created dataset.

    Returns:
        Dataset: Dataset containing the loaded data.
    """
    dataset = Dataset(name=dataset_name, include_test=include_test)
    for split_name in dataset.splits:
        with open(os.path.join(data_dir, split_name + ".csv"),
                  newline='', mode="r") as infile:
            reader = csv.DictReader(infile, delimiter="|")
            for row in reader:
                text = row["Text"]
                label = int(row["Label"])
                example = Example(segment_words(text.rstrip('\n')),
                                  label)
                dataset.splits[split_name].append(example)

    dataset.shuffle()
    return dataset


def evaluate(classifier: Classifier, dataset: Dataset,
             limit_training_set: bool = False) -> None:
    """
    Evaluate a classifier on the training and dev sets, printing the accuracy.

    Args:
        classifier (Classifier): Classifier to evaluate on the dataset.
        dataset (Dataset): Dataset to train and evaluate on.
        limit_training_set (bool): If true, truncate training set to
        only 25% of its full size.
    """
    training_set = dataset.train[:int(0.25 * len(dataset.train))] \
        if limit_training_set else dataset.train

    classifier.train(training_set)

    for split_name in dataset.splits:
        accuracy = calculate_accuracy(training_set
                                      if split_name == "train"
                                      else dataset.splits[split_name],
                                      classifier)
        print('Accuracy ({}): {}'.format(split_name, accuracy))


########################################################
# UNIT TESTS BELOW
########################################################

def check_logistic_loss(logistic_loss_func) -> None:
    """
    Checks the logistic_loss() function implementation.
    """
    # Use multiple examples to catch "forgot np.mean" errors
    loss = logistic_loss_func(np.array([0.9, 0.1]), np.array([1, 0]))
    correct_loss = 0.10536050454671517
    
    if np.abs(loss - correct_loss) < 0.01:
        print("Logistic loss is correct.")
    else:
        print("Logistic loss is incorrect.")
        print("  Input: pred=[0.9, 0.1], true=[1, 0]")
        print("  Expected:", correct_loss)
        print("  Got:", loss)


def check_gradient_descent(gradient_descent_func,
                           tolerance: float = 1e-2) -> None:
    """
    Checks the gradient_descent() function implementation.
    """
    np.random.seed(42)
    
    X_test = np.array([[2.0, 3.0], [3.0, 2.0], [3.0, 4.0], [4.0, 3.0],
                       [0.0, 1.0], [1.0, 0.0], [0.5, 0.5], [0.0, 0.0]])
    Y_test = np.array([1, 1, 1, 1, 0, 0, 0, 0])
    
    W, b = gradient_descent_func(X_test, Y_test, 
                                 batch_size=8, alpha=0.5, 
                                 num_iterations=100, print_every=10000,
                                 epsilon=1e-10)
    
    # Pre-computed correct outputs
    correct_W = np.array([1.3621915472787667, 1.3621915472787667])
    correct_b = -3.614102491149732
    
    # Check W
    if W.shape != (2,):
        print("W has wrong shape. Expected (2,) but got", W.shape)
    elif np.linalg.norm(W - correct_W) <= tolerance:
        print("W is correct.")
    else:
        print("W is incorrect. Got", W, "expected", correct_W)
    
    # Check b
    if np.abs(b - correct_b) <= tolerance:
        print("b is correct.")
    else:
        print("b is incorrect. Got", b, "expected", correct_b)