import numpy as np
from sklearn.metrics.cluster import adjusted_mutual_info_score
from sklearn.metrics import matthews_corrcoef

_RANKING_METRICS = {'youden', 'patn', 'ap', 'maxf1', 'auc'}
_PRED_METRICS = {'ami', 'mcc'}
_ALL_METRICS = _RANKING_METRICS | _PRED_METRICS


def _fast_auc_from_sorted(tp_cum, fp_cum, n_pos, n_neg):
    tpr = np.concatenate(([0.0], tp_cum / n_pos))
    fpr = np.concatenate(([0.0], fp_cum / n_neg))
    _trapz = getattr(np, 'trapezoid', None) or np.trapz
    return _trapz(tpr, fpr)


def _fast_auc_rank(y_true, scores, n_pos, n_neg):
    from scipy.stats import rankdata
    ranks = rankdata(scores, method='average')
    sum_pos_ranks = ranks[y_true == 1].sum()
    return (sum_pos_ranks - n_pos * (n_pos + 1) / 2.0) / (n_pos * n_neg)


def get_indices(bin_labels, scores, bin_preds=None, metrics=('auc'), auc_method='trapz'):
    if isinstance(metrics, str):
        metrics = {metrics}
    metrics = set(metrics)
    if 'all' in metrics:
        metrics = set(_ALL_METRICS)

    res = {}
    ranking_requested = metrics & _RANKING_METRICS

    if ranking_requested:
        bin_labels_arr = np.asarray(bin_labels)
        scores = np.asarray(scores, dtype=np.float64)

        valid = ~np.isnan(scores)
        if not valid.all():
            y = bin_labels_arr[valid]
            s = scores[valid]
        else:
            y = bin_labels_arr
            s = scores

        m = y.size
        n_pos = int(y.sum())
        n_neg = m - n_pos

        if m == 0 or n_pos == 0 or n_neg == 0:
            for k in ranking_requested:
                if k == 'youden':
                    res['youden_fpr'] = res['youden_fnr'] = np.nan
                elif k == 'patn':
                    res['Patn'] = res['adj_Patn'] = np.nan
                elif k == 'ap':
                    res['ap'] = res['adj_ap'] = np.nan
                elif k == 'maxf1':
                    res['maxf1'] = res['adj_maxf1'] = np.nan
                elif k == 'auc':
                    res['auc'] = np.nan
        else:
            base_rate = n_pos / m

            auc_needs_sort = 'auc' in ranking_requested and auc_method == 'trapz'
            need_order = bool(ranking_requested & {'youden', 'patn', 'ap', 'maxf1'}) or auc_needs_sort

            order = y_s = tp_cum = None
            if need_order:
                order = np.argsort(s)[::-1]
                y_s = y[order]
                tp_cum = np.cumsum(y_s)

            need_fp = ('youden' in ranking_requested) or auc_needs_sort
            fp_cum = None
            ranks_full = None
            if need_fp or 'maxf1' in ranking_requested:
                ranks_full = np.arange(1, m + 1)
            if need_fp:
                fp_cum = ranks_full - tp_cum

            if 'youden' in ranking_requested:
                fpr = fp_cum / n_neg
                fnr = (n_pos - tp_cum) / n_pos
                idx_youden = np.argmin(fpr + fnr)
                res['youden_fpr'] = fpr[idx_youden]
                res['youden_fnr'] = fnr[idx_youden]

            if 'patn' in ranking_requested:
                Patn = tp_cum[n_pos - 1] / n_pos
                res['Patn'] = Patn
                res['adj_Patn'] = (Patn - base_rate) / (1 - base_rate)

            if 'ap' in ranking_requested:
                pos_idx = np.flatnonzero(y_s)
                ap = np.sum(tp_cum[pos_idx] / (pos_idx + 1)) / n_pos
                res['ap'] = ap
                res['adj_ap'] = (ap - base_rate) / (1 - base_rate)

            if 'maxf1' in ranking_requested:
                f1_curve = 2 * tp_cum / (ranks_full + n_pos)
                maxf1 = f1_curve.max()
                res['maxf1'] = maxf1
                res['adj_maxf1'] = (maxf1 - base_rate) / (1 - base_rate)

            if 'auc' in ranking_requested:
                if auc_method == 'rank':
                    res['auc'] = _fast_auc_rank(y, s, n_pos, n_neg)
                else:
                    res['auc'] = _fast_auc_from_sorted(tp_cum, fp_cum, n_pos, n_neg)

    if 'ami' in metrics or 'mcc' in metrics:
        if bin_preds is None:
            raise ValueError("bin_preds required!")
        bin_labels_arr = np.asarray(bin_labels)
        bin_preds = np.asarray(bin_preds)
        if 'ami' in metrics:
            res['ami'] = adjusted_mutual_info_score(bin_labels_arr, bin_preds)
        if 'mcc' in metrics:
            res['mcc'] = matthews_corrcoef(bin_labels_arr, bin_preds)

    return res
