"""
Keyness and POS Annotation Script for Project 2025
--------------------------------------------------
This script implements a corpus-assisted discourse analysis of Project 2025
compared to U.S. party platforms (2016–2024). It produces unigram, bigram,
and trigram keyness tables with frequency statistics and effect size measures.

It also creates a POS-annotated version of the Project 2025 corpus
where each paragraph remains one line, but tokens are stored as lemma_POS.
"""

import os
import pandas as pd
import numpy as np
import spacy
from collections import Counter

# ----------------------------------------------------------------------
# PATH CONFIGURATION
# ----------------------------------------------------------------------
BASE_DIR = os.path.dirname(os.path.abspath(__file__))  # e.g. manuscript_PLOS_One/code/
DATA_RAW = os.path.join(BASE_DIR, "..", "data", "raw_data")
DATA_KEYNESS = os.path.join(BASE_DIR, "..", "data", "keyness")

os.makedirs(DATA_KEYNESS, exist_ok=True)

# ----------------------------------------------------------------------
# LOAD MODEL AND DEFINE STOPWORDS
# ----------------------------------------------------------------------
nlp = spacy.load("en_core_web_md", disable=["parser", "ner", "textcat"])
nlp.max_length = 5_000_000
STOP = nlp.Defaults.stop_words

# ----------------------------------------------------------------------
# PREPROCESSING
# ----------------------------------------------------------------------
IRREGULAR_PLURAL_OVERRIDES = {
    ("data", "datum"): "data",
    ("media", "medium"): "media",
}

def doc_to_tokens(doc):
    """Convert spaCy Doc to list of lemma_POS tokens."""
    tokens = []
    for t in doc:
        if t.is_alpha and t.text.lower() not in STOP:
            pos = t.pos_
            lemma = t.lemma_.lower()
            key = (t.text.lower(), lemma)
            if key in IRREGULAR_PLURAL_OVERRIDES:
                lemma = IRREGULAR_PLURAL_OVERRIDES[key]
            tokens.append(f"{lemma}_{pos}")
    return tokens


def tokens_to_ngrams(tokens, n):
    """Generate n-grams of size n."""
    if n == 1:
        return tokens
    return [" ".join(tokens[i:i + n]) for i in range(len(tokens) - n + 1)]


# ----------------------------------------------------------------------
# CORPUS COUNTING
# ----------------------------------------------------------------------
def count_ngrams_in_csv(csv_path, text_col="text", n=1,
                        chunksize=1000, batch_size=200, n_process=1):
    counter = Counter()
    N_tokens = 0
    for chunk in pd.read_csv(csv_path, usecols=[text_col], chunksize=chunksize):
        texts = chunk[text_col].fillna("").astype(str).tolist()
        for doc in nlp.pipe(texts, batch_size=batch_size, n_process=n_process):
            tokens = doc_to_tokens(doc)
            if n == 1:
                N_tokens += len(tokens)
            else:
                N_tokens += max(len(tokens) - n + 1, 0)
            counter.update(tokens_to_ngrams(tokens, n))
    return counter, N_tokens


# ----------------------------------------------------------------------
# STATS FUNCTIONS
# ----------------------------------------------------------------------
def _log_safe_div(x, y):
    if x == 0 or y == 0:
        return 0.0
    return x * np.log(x / y)

def compute_measures_for_union(cnt_target: Counter, cnt_ref: Counter,
                               n_target: int, n_ref: int):
    vocab = set(cnt_target) | set(cnt_ref)
    rows = []
    N = n_target + n_ref
    eps = 1e-12
    for term in vocab:
        f_t = cnt_target.get(term, 0)
        f_r = cnt_ref.get(term, 0)
        f_tot = f_t + f_r
        e_r = n_ref * f_tot / N
        e_t = n_target * f_tot / N
        ll = 2.0 * (_log_safe_div(f_r, e_r) + _log_safe_div(f_t, e_t))
        bic = ll - (1.0 * np.log(N))
        nf_t = f_t / n_target if n_target else 0.0
        nf_r = f_r / n_ref if n_ref else 0.0
        perc_diff = ((nf_t - nf_r) * 100.0) / (nf_r if nf_r > 0 else eps)
        if nf_t == 0 and nf_r == 0:
            log_ratio = 0.0
        elif nf_t == 0:
            log_ratio = np.log2(eps / nf_r)
        elif nf_r == 0:
            log_ratio = np.log2(nf_t / eps)
        else:
            log_ratio = np.log2(nf_t / nf_r)
        word_use = "equal" if np.isclose(nf_t, nf_r) else ("overuse" if nf_t > nf_r else "underuse")
        rows.append((term, f_t, f_r, ll, bic, perc_diff, log_ratio, word_use))
    return pd.DataFrame(rows, columns=[
        "term", "freq_target", "freq_reference",
        "log_likelihood", "bic", "perc_diff", "log_ratio", "word_use"
    ])


# ----------------------------------------------------------------------
# POS-ANNOTATED CORPUS CREATION
# ----------------------------------------------------------------------
def annotate_csv_with_lemma_pos(input_csv, output_csv, text_col="text",
                                batch_size=200, n_process=1):
    """
    Create a version of the corpus where each paragraph remains one line,
    but words are replaced with 'lemma_POS' tokens separated by spaces.
    """
    df_iter = pd.read_csv(input_csv, chunksize=1000)
    annotated_chunks = []

    for chunk in df_iter:
        texts = chunk[text_col].fillna("").astype(str).tolist()
        annotated_texts = []
        for doc in nlp.pipe(texts, batch_size=batch_size, n_process=n_process):
            tokens = doc_to_tokens(doc)
            annotated_texts.append(" ".join(tokens))
        chunk[text_col + "_lemmaPOS"] = annotated_texts
        annotated_chunks.append(chunk)

    annotated_df = pd.concat(annotated_chunks, ignore_index=True)
    annotated_df.to_csv(output_csv, index=False)
    print(f"Annotated corpus saved to {output_csv}")


# ----------------------------------------------------------------------
# RUNNER FUNCTION
# ----------------------------------------------------------------------
def run_keyness(target_csv, ref_csv, text_col, n, out_csv,
                chunksize=1000, batch_size=200, n_process=1):
    print(f"[n={n}] Counting target corpus…")
    cnt_t, N_t = count_ngrams_in_csv(target_csv, text_col=text_col, n=n,
                                     chunksize=chunksize, batch_size=batch_size, n_process=n_process)
    print(f"[n={n}] Counting reference corpus…")
    cnt_r, N_r = count_ngrams_in_csv(ref_csv, text_col=text_col, n=n,
                                     chunksize=chunksize, batch_size=batch_size, n_process=n_process)
    print(f"[n={n}] Computing measures for {len(set(cnt_t)|set(cnt_r)):,} terms…")
    df = compute_measures_for_union(cnt_t, cnt_r, N_t, N_r)
    df.sort_values(["log_likelihood", "log_ratio"], ascending=[False, False], inplace=True)
    df.to_csv(out_csv, index=False)
    print(f"[n={n}] Results saved to {out_csv}")


# ----------------------------------------------------------------------
# MAIN FUNCTION
# ----------------------------------------------------------------------
def main():
    # --- 1. Create POS-annotated version of Project 2025 ---
    annotate_csv_with_lemma_pos(
        os.path.join(DATA_RAW, "Project2025.csv"),
        os.path.join(DATA_KEYNESS, "Project2025_lemmaPOS.csv"),
        text_col="text",
        n_process=2
    )

    # --- 2. Run Keyness Analyses ---
    # Project 2025 vs Democrats
    run_keyness(
        os.path.join(DATA_RAW, "Project2025.csv"),
        os.path.join(DATA_RAW, "Platforms_Democrats.csv"),
        text_col="text",
        n=1,
        out_csv=os.path.join(DATA_KEYNESS, "keyness_1gram_manifesto_dem.csv"),
        n_process=2
    )
    run_keyness(
        os.path.join(DATA_RAW, "Project2025.csv"),
        os.path.join(DATA_RAW, "Platforms_Democrats.csv"),
        text_col="text",
        n=2,
        out_csv=os.path.join(DATA_KEYNESS, "keyness_2gram_manifesto_dem.csv"),
        n_process=2
    )
    run_keyness(
        os.path.join(DATA_RAW, "Project2025.csv"),
        os.path.join(DATA_RAW, "Platforms_Democrats.csv"),
        text_col="text",
        n=3,
        out_csv=os.path.join(DATA_KEYNESS, "keyness_3gram_manifesto_dem.csv"),
        n_process=2
    )

    # Project 2025 vs Republicans
    run_keyness(
        os.path.join(DATA_RAW, "Project2025.csv"),
        os.path.join(DATA_RAW, "Platforms_Republicans.csv"),
        text_col="text",
        n=1,
        out_csv=os.path.join(DATA_KEYNESS, "keyness_1gram_manifesto_rep.csv"),
        n_process=2
    )
    run_keyness(
        os.path.join(DATA_RAW, "Project2025.csv"),
        os.path.join(DATA_RAW, "Platforms_Republicans.csv"),
        text_col="text",
        n=2,
        out_csv=os.path.join(DATA_KEYNESS, "keyness_2gram_manifesto_rep.csv"),
        n_process=2
    )
    run_keyness(
        os.path.join(DATA_RAW, "Project2025.csv"),
        os.path.join(DATA_RAW, "Platforms_Republicans.csv"),
        text_col="text",
        n=3,
        out_csv=os.path.join(DATA_KEYNESS, "keyness_3gram_manifesto_rep.csv"),
        n_process=2
    )

    print("✅ All analyses complete. Results saved in data/keyness/.")


# ----------------------------------------------------------------------
if __name__ == "__main__":
    main()