"""How We Talk to Machines — Familiarity Divide, secondary PRISM survey analysis.

Install: pip install pandas numpy requests statsmodels
Run: python reproduce_adjusted.py

Downloads the public PRISM survey from the original authors; outputs participant-level
OLS summaries only, no individual records. Eight exploratory outcomes were tested.
"""
import io
import requests
import numpy as np
import pandas as pd
import statsmodels.api as sm
from statsmodels.stats.multitest import multipletests

SOURCE = "https://huggingface.co/datasets/HannahRoseKirk/prism-alignment/resolve/main/survey.jsonl"
ATTRIBUTES = ["values", "creativity", "fluency", "factuality", "diversity", "safety", "personalisation", "helpfulness"]
FAMILIARITY = ["Not familiar at all", "Somewhat familiar", "Very familiar"]
AGES = ["18-24 years old", "25-34 years old", "35-44 years old", "45-54 years old", "55-64 years old", "65+ years old"]
COUNTRIES = ["USA", "GBR", "ZAF", "NZL", "AUS", "MEX", "CHL", "ISR", "CAN", "OTHER"]
EDUCATION = ["University Bachelors Degree", "Graduate / Professional degree", "Some University but no degree", "Completed Secondary School", "Vocational", "OTHER"]
FREQUENCY = ["Missing", "Less than one a year", "Once per month", "More than once a month", "Every week", "Every day"]

response = requests.get(SOURCE, timeout=120)
response.raise_for_status()
df = pd.read_json(io.BytesIO(response.content), lines=True)
assert len(df) == 1500, "Data size changed; inspect dataset/version before interpreting results."

# Exclude the one undisclosed age instead of adding a one-person dummy category.
df = df[df.age.isin(AGES)].copy()
assert len(df) == 1499

# Define a category baseline consistently for reproducibility. 'Missing' frequency
# means the frequency question was NOT shown; it must not be interpreted as no use.
df["country_group"] = df["location"].map(lambda z: z.get("reside_countryISO") if z.get("reside_countryISO") in COUNTRIES[:-1] else "OTHER")
df["education_group"] = df["education"].where(df.education.isin(EDUCATION[:-1]), "OTHER")
df["frequency_group"] = df["lm_frequency_use"].fillna("Missing")

# First familiarity category is the reference; group differences are high - low.
def design(include_frequency):
    out = pd.DataFrame({
        "intercept": np.ones(len(df)),
        "fam_some": (df.lm_familiarity == FAMILIARITY[1]).astype(float).to_numpy(),
        "fam_very": (df.lm_familiarity == FAMILIARITY[2]).astype(float).to_numpy(),
    }, index=df.index)
    for col, categories, prefix in [
        ("age", AGES, "age"),
        ("country_group", COUNTRIES, "country"),
        ("education_group", EDUCATION, "education"),
    ] + ([("frequency_group", FREQUENCY, "frequency")] if include_frequency else []):
        cat = pd.Categorical(df[col], categories=categories, ordered=True)
        dummies = pd.get_dummies(cat, prefix=prefix, drop_first=True, dtype=float)
        dummies.index = df.index
        out = pd.concat([out, dummies], axis=1)
    return out

rows = []
for model_name, use_frequency in [("Demographic adjusted", False), ("Demographic + frequency adjusted", True)]:
    X = design(use_frequency)
    names = list(X.columns)
    for attribute in ATTRIBUTES:
        y = df.stated_prefs.map(lambda a: float(a[attribute])).astype(float)
        model = sm.OLS(y, X).fit(cov_type="HC3")
        contrasts = {
            "Very minus not familiar": {"fam_very": 1},
            "Very minus somewhat familiar": {"fam_very": 1, "fam_some": -1},
        }
        for label, contrast in contrasts.items():
            weights = np.array([contrast.get(name, 0.0) for name in names], dtype=float)
            test = model.t_test(weights)
            lower, upper = test.conf_int(alpha=.05)[0]
            rows.append({
                "model": model_name, "attribute": attribute, "contrast": label,
                "estimate": float(test.effect), "ci_low_95": float(lower),
                "ci_high_95": float(upper), "p_value_hc3": float(test.pvalue),
                "sample_n": len(df), "parameter_count": X.shape[1],
            })

result = pd.DataFrame(rows)
for group_idx in result.groupby(["model", "contrast"]).groups.values():
    result.loc[group_idx, "q_value_bh_eight_outcomes"] = multipletests(result.loc[group_idx,"p_value_hc3"], method="fdr_bh")[1]
result.to_csv("reproduced_adjusted_results.csv", index=False)
print(result.round(4).to_string(index=False))
print("\nSaved reproduced_adjusted_results.csv")
