# =============================================================================
# S1_analysis_script.R
# Do Knowledge Economy Pillars Move Together?
# Bayesian Parallel LGCM + Conditional Extension (Teacher Support)
# N = 52 countries, 2012-2022 (PISA cycles)
#
# Authors: Ergün Kara
# Affiliation: CHERI, University of Aberdeen
# R version: 4.5.2 | Date: 2025-04-27
#
# STRUCTURE:
#   STEP 01 – Packages
#   STEP 02 – PISA data (learningtower)
#   STEP 03 – World Bank R&D data
#   STEP 04 – ILO employment data
#   STEP 05 – Data harmonisation & panel assembly
#   STEP 06 – Z-score standardisation (2012 anchor)
#   STEP 07 – Unconditional single-domain models (frequentist)
#   STEP 08 – Bayesian unconditional single-domain LGCMs
#   STEP 09 – Bayesian parallel LGCM (primary analysis)
#   STEP 10 – Posterior diagnostics & cross-domain correlations
#   STEP 11 – Conditional LGCM: instructional quality predictors
#   STEP 12 – Bayesian conditional LGCMs
#   STEP 13 – Figures
#   STEP 14 – Tables
#   STEP 15 – Session info
#
# OUTPUTS (relative to project root, set via here::here()):
#   data/processed/panel_standardised.rds  – primary analysis panel (208 × 15)
#   data/processed/panel_enriched.rds      – with instructional quality (208 × 31)
#   data/processed/pisa_delta_2012_2022.rds – country-level changes (59 × 19)
#   models/                                 – all fitted model objects (.rds)
#   output/figures/                         – PNG figures
#   output/tables/                          – CSV tables
# =============================================================================


# ─────────────────────────────────────────────────────────────────────────────
# STEP 01 – PACKAGES
# ─────────────────────────────────────────────────────────────────────────────
pkgs <- c(
  "here",           # reproducible file paths (Müller, 2025)
  "tidyverse",      # data wrangling & ggplot2 (Wickham et al., 2019)
  "learningtower",  # PISA microdata 2000–2022 (Wang et al., 2024)
  "httr",           # World Bank API calls (Wickham, 2023)
  "Rilostat",       # ILO ILOSTAT API (Bescond, 2026)
  "countrycode",    # ISO-3166 harmonisation (Arel-Bundock et al., 2018)
  "haven",          # PISA SPSS files (Wickham et al., 2023)
  "lme4",           # frequentist mixed-effects models (Bates et al., 2015)
  "performance",    # R² and model fit (Lüdecke et al., 2021)
  "brms",           # Bayesian multilevel models (Bürkner, 2017, 2018)
  "rstan",          # HMC-NUTS backend (Stan Dev Team, 2025)
  "patchwork"       # multi-panel figures (Pedersen, 2025)
)

for (p in pkgs) {
  if (!requireNamespace(p, quietly=TRUE)) install.packages(p)
  library(p, character.only=TRUE)
}

rstan_options(auto_write = TRUE)
options(mc.cores = 4)
set.seed(2025)

# directories
dir.create(here("data/raw"),       recursive=TRUE, showWarnings=FALSE)
dir.create(here("data/processed"), recursive=TRUE, showWarnings=FALSE)
dir.create(here("models"),         recursive=TRUE, showWarnings=FALSE)
dir.create(here("output/figures"), recursive=TRUE, showWarnings=FALSE)
dir.create(here("output/tables"),  recursive=TRUE, showWarnings=FALSE)


# ─────────────────────────────────────────────────────────────────────────────
# STEP 02 – PISA DATA (learningtower)
# ─────────────────────────────────────────────────────────────────────────────
# Load student-level PISA data for four cycles and compute weighted country means.
# 2015 irregularities for specific countries are treated as missing at random
# (see OECD PISA 2015 Technical Report, 2017, for documentation).

cat("Loading PISA data via learningtower...\n")

pisa_raw <- load_student("all") |>
  filter(year %in% c(2012, 2015, 2018, 2022)) |>
  select(year, country, math, read, science, student_weight = starts_with("w_fstu")) |>
  rename(w = student_weight)

# Countries with documented 2015 administration problems
flag_2015 <- c("AUS","CAN","DNK","HKG","IRQ","JOR","KAZ","KWT","MDA","SVN")

pisa_country <- pisa_raw |>
  mutate(
    pisa_2015_flag = (year == 2015 & country %in% flag_2015),
    math    = if_else(pisa_2015_flag, NA_real_, math),
    read    = if_else(pisa_2015_flag, NA_real_, read),
    science = if_else(pisa_2015_flag, NA_real_, science)
  ) |>
  group_by(year, country) |>
  summarise(
    pisa_math     = weighted.mean(math,    w, na.rm=TRUE),
    pisa_read     = weighted.mean(read,    w, na.rm=TRUE),
    pisa_science  = weighted.mean(science, w, na.rm=TRUE),
    pisa_2015_flag = any(pisa_2015_flag),
    .groups = "drop"
  ) |>
  rename(iso3c = country) |>
  mutate(iso3c = countrycode(iso3c, "iso3c", "iso3c", warn=FALSE)) |>
  filter(!is.na(iso3c))

cat("  PISA: ", n_distinct(pisa_country$iso3c), "countries,",
    n_distinct(pisa_country$year), "waves\n")


# ─────────────────────────────────────────────────────────────────────────────
# STEP 03 – WORLD BANK R&D DATA
# ─────────────────────────────────────────────────────────────────────────────
# Gross domestic expenditure on R&D as % of GDP
# WB indicator: GB.XPD.RSDV.GD.ZS

cat("Fetching World Bank R&D data...\n")

wb_url <- paste0(
  "https://api.worldbank.org/v2/country/all/indicator/GB.XPD.RSDV.GD.ZS",
  "?date=2012:2022&format=json&per_page=20000"
)
resp <- GET(wb_url)
wb_json <- content(resp, as="parsed")[[2]]

rd_raw <- bind_rows(lapply(wb_json, function(x) {
  tibble(
    iso3c = x$countryiso3code,
    year  = as.integer(x$date),
    rd_pct_gdp = if (is.null(x$value)) NA_real_ else as.numeric(x$value)
  )
})) |>
  filter(!is.na(rd_pct_gdp), year %in% c(2012, 2015, 2018, 2022)) |>
  mutate(iso3c = countrycode(iso3c, "iso3c", "iso3c", warn=FALSE)) |>
  filter(!is.na(iso3c))

cat("  R&D: ", n_distinct(rd_raw$iso3c), "countries\n")


# ─────────────────────────────────────────────────────────────────────────────
# STEP 04 – ILO EMPLOYMENT DATA (Rilostat)
# ─────────────────────────────────────────────────────────────────────────────
# Employment-to-population ratio, prime-age adults 25–54, both sexes
# ILO dataset: EMP_DWAP_SEX_AGE_RT_A

cat("Fetching ILO employment data via Rilostat...\n")

emp_raw <- get_ilostat(
  id    = "EMP_DWAP_SEX_AGE_RT_A",
  timefrom = 2012, timeto = 2022
)

emp_country <- emp_raw |>
  filter(
    sex.label  == "Sex: Total",
    classif1.label == "Age (Aggregate bands): 25-54"
  ) |>
  select(ref_area, time, obs_value) |>
  rename(iso3c = ref_area, year = time, emp_ratio = obs_value) |>
  mutate(
    year  = as.integer(year),
    iso3c = countrycode(iso3c, "iso3c", "iso3c", warn=FALSE)
  ) |>
  filter(!is.na(iso3c), year %in% c(2012, 2015, 2018, 2022))

cat("  Employment: ", n_distinct(emp_country$iso3c), "countries\n")


# ─────────────────────────────────────────────────────────────────────────────
# STEP 05 – PANEL ASSEMBLY AND HARMONISATION
# ─────────────────────────────────────────────────────────────────────────────
cat("Assembling panel...\n")

# Country metadata
meta <- tibble(iso3c = unique(pisa_country$iso3c)) |>
  mutate(
    country_name = countrycode(iso3c, "iso3c", "country.name"),
    region       = countrycode(iso3c, "iso3c", "region"),
    oecd_member  = iso3c %in% c(
      "AUS","AUT","BEL","CAN","CHL","COL","CRI","CZE","DNK","EST","FIN","FRA",
      "DEU","GRC","HUN","ISL","IRL","ISR","ITA","JPN","KOR","LVA","LTU","LUX",
      "MEX","NLD","NZL","NOR","POL","PRT","SVK","SVN","ESP","SWE","CHE","TUR",
      "GBR","USA"
    )
  )

# Full join across three sources
panel_raw <- pisa_country |>
  left_join(rd_raw,     by = c("iso3c","year")) |>
  left_join(emp_country, by = c("iso3c","year")) |>
  left_join(meta,        by = "iso3c") |>
  mutate(time_metric = case_when(
    year == 2012 ~ 0L,
    year == 2015 ~ 3L,
    year == 2018 ~ 6L,
    year == 2022 ~ 10L
  ))

# Restrict to countries with ≥3 complete observations across all 3 domains
complete_flag <- panel_raw |>
  group_by(iso3c) |>
  summarise(
    n_pisa = sum(!is.na(pisa_math)),
    n_rd   = sum(!is.na(rd_pct_gdp)),
    n_emp  = sum(!is.na(emp_ratio))
  ) |>
  filter(n_pisa >= 3, n_rd >= 3, n_emp >= 3)

panel_analytic <- panel_raw |>
  filter(iso3c %in% complete_flag$iso3c) |>
  arrange(iso3c, year)

cat("  Analytic panel:", nrow(panel_analytic), "rows,",
    n_distinct(panel_analytic$iso3c), "countries\n")
saveRDS(panel_analytic, here("data/processed/panel_analytical.rds"))


# ─────────────────────────────────────────────────────────────────────────────
# STEP 06 – Z-SCORE STANDARDISATION (2012 ANCHOR)
# ─────────────────────────────────────────────────────────────────────────────
# All outcomes z-scored using the cross-country mean and SD from 2012 only.
# This preserves 2012 as the baseline reference, so:
#   - intercepts = 2012 position relative to sample mean (in SDs)
#   - slopes = annual change in those same SD units

baseline <- panel_analytic |>
  filter(year == 2012) |>
  summarise(
    pisa_m = mean(pisa_math,  na.rm=TRUE), pisa_s = sd(pisa_math,  na.rm=TRUE),
    rd_m   = mean(rd_pct_gdp, na.rm=TRUE), rd_s   = sd(rd_pct_gdp, na.rm=TRUE),
    emp_m  = mean(emp_ratio,  na.rm=TRUE), emp_s  = sd(emp_ratio,  na.rm=TRUE)
  )

panel_z <- panel_analytic |>
  mutate(
    pisa_z = (pisa_math  - baseline$pisa_m) / baseline$pisa_s,
    rd_z   = (rd_pct_gdp - baseline$rd_m)   / baseline$rd_s,
    emp_z  = (emp_ratio  - baseline$emp_m)  / baseline$emp_s
  )

saveRDS(panel_z, here("data/processed/panel_standardised.rds"))
write.csv(panel_z, here("output/tables/panel_standardised_data.csv"),
          row.names=FALSE, fileEncoding="UTF-8")
cat("  Saved panel_standardised.rds\n")


# ─────────────────────────────────────────────────────────────────────────────
# STEP 07 – UNCONDITIONAL FREQUENTIST LGCMs (single domain, lme4)
# ─────────────────────────────────────────────────────────────────────────────
# These serve as sanity checks and to establish baseline AIC/R² for each domain.

cat("\nFitting frequentist unconditional models...\n")

d_pisa <- filter(panel_z, !is.na(pisa_z))
d_rd   <- filter(panel_z, !is.na(rd_z))
d_emp  <- filter(panel_z, !is.na(emp_z))

ctrl <- lmerControl(optimizer = "bobyqa")

freq_pisa <- lmer(pisa_z ~ time_metric + (time_metric | iso3c), data=d_pisa, REML=TRUE, control=ctrl)
freq_rd   <- lmer(rd_z   ~ time_metric + (time_metric | iso3c), data=d_rd,   REML=TRUE, control=ctrl)
freq_emp  <- lmer(emp_z  ~ time_metric + (time_metric | iso3c), data=d_emp,  REML=TRUE, control=ctrl)

cat("  PISA  AIC:", round(AIC(freq_pisa),1), "| R2m:", round(r2(freq_pisa)$R2_marginal,3), "\n")
cat("  R&D   AIC:", round(AIC(freq_rd),1),   "| R2m:", round(r2(freq_rd)$R2_marginal,3), "\n")
cat("  Emp   AIC:", round(AIC(freq_emp),1),   "| R2m:", round(r2(freq_emp)$R2_marginal,3), "\n")

saveRDS(freq_pisa, here("models/m0_uncond.rds"))


# ─────────────────────────────────────────────────────────────────────────────
# STEP 08 – BAYESIAN UNCONDITIONAL SINGLE-DOMAIN LGCMs
# ─────────────────────────────────────────────────────────────────────────────
# Prior specification (Gelman et al., 2013):
#   intercept     ~ Normal(0, 1)
#   slopes        ~ Normal(0, 0.5)
#   random SDs    ~ Exponential(2)
#   correlations  ~ LKJ(2)
#   residual SD   ~ Exponential(2)

cat("\nFitting Bayesian single-domain LGCMs...\n")

priors_base <- c(
  prior(normal(0, 1),   class = "Intercept"),
  prior(normal(0, 0.5), class = "b"),
  prior(exponential(2), class = "sd"),
  prior(lkj(2),         class = "cor"),
  prior(exponential(2), class = "sigma")
)

fit_pisa <- brm(
  formula = pisa_z ~ time_metric + (time_metric | iso3c),
  data    = d_pisa,
  prior   = priors_base,
  chains  = 4, iter = 4000, warmup = 2000, cores = 4, seed = 2025,
  backend = "rstan",
  file    = here("models/fit_pisa_lgcm_done"),
  silent  = 2
)

fit_rd <- brm(
  formula = rd_z ~ time_metric + (time_metric | iso3c),
  data    = d_rd,
  prior   = priors_base,
  chains  = 4, iter = 4000, warmup = 2000, cores = 4, seed = 2025,
  backend = "rstan",
  file    = here("models/fit_rd_lgcm_done"),
  silent  = 2
)

fit_emp <- brm(
  formula = emp_z ~ time_metric + (time_metric | iso3c),
  data    = d_emp,
  prior   = priors_base,
  chains  = 4, iter = 4000, warmup = 2000, cores = 4, seed = 2025,
  backend = "rstan",
  file    = here("models/fit_emp_lgcm_done"),
  silent  = 2
)

cat("  PISA  max Rhat:", round(max(rhat(fit_pisa), na.rm=TRUE), 4), "\n")
cat("  R&D   max Rhat:", round(max(rhat(fit_rd),   na.rm=TRUE), 4), "\n")
cat("  Emp   max Rhat:", round(max(rhat(fit_emp),  na.rm=TRUE), 4), "\n")


# ─────────────────────────────────────────────────────────────────────────────
# STEP 09 – BAYESIAN PARALLEL LGCM (primary analysis)
# ─────────────────────────────────────────────────────────────────────────────
# The multivariate brms specification with | p | notation estimates cross-equation
# random-effect correlations simultaneously.
# set_rescor(FALSE) suppresses residual correlations (data sources are independent).

cat("\nFitting Bayesian parallel LGCM...\n")

# Use only countries with complete data across all three domains
d_parallel <- panel_z |>
  filter(!is.na(pisa_z), !is.na(rd_z), !is.na(emp_z)) |>
  group_by(iso3c) |>
  filter(n() >= 3) |>
  ungroup()

cat("  Parallel LGCM sample:", n_distinct(d_parallel$iso3c), "countries,",
    nrow(d_parallel), "obs\n")

# Multivariate formula using mvbf()
bf_pisa <- bf(pisa_z ~ time_metric + (time_metric | p | iso3c))
bf_rd   <- bf(rd_z   ~ time_metric + (time_metric | p | iso3c))
bf_emp  <- bf(emp_z  ~ time_metric + (time_metric | p | iso3c))

priors_parallel <- c(
  # intercepts per response
  prior(normal(0, 1),   class = "Intercept", resp = "pisaz"),
  prior(normal(0, 1),   class = "Intercept", resp = "rdz"),
  prior(normal(0, 1),   class = "Intercept", resp = "empz"),
  # slopes per response
  prior(normal(0, 0.5), class = "b",         resp = "pisaz"),
  prior(normal(0, 0.5), class = "b",         resp = "rdz"),
  prior(normal(0, 0.5), class = "b",         resp = "empz"),
  # random-effect SDs
  prior(exponential(2), class = "sd",        resp = "pisaz"),
  prior(exponential(2), class = "sd",        resp = "rdz"),
  prior(exponential(2), class = "sd",        resp = "empz"),
  # shared 6×6 RE correlation matrix — LKJ(2) moderate shrinkage
  prior(lkj(2),         class = "cor"),
  # residual SDs
  prior(exponential(2), class = "sigma",     resp = "pisaz"),
  prior(exponential(2), class = "sigma",     resp = "rdz"),
  prior(exponential(2), class = "sigma",     resp = "empz")
)

fit_parallel <- brm(
  formula = mvbf(bf_pisa, bf_rd, bf_emp, rescor = FALSE),
  data    = d_parallel,
  prior   = priors_parallel,
  chains  = 4, iter = 4000, warmup = 2000, cores = 4, seed = 2025,
  backend = "rstan",
  file    = here("models/fit_parallel_lgcm_done"),
  silent  = 2
)

cat("  Parallel LGCM max Rhat:", round(max(rhat(fit_parallel), na.rm=TRUE), 4), "\n")


# ─────────────────────────────────────────────────────────────────────────────
# STEP 10 – POSTERIOR DIAGNOSTICS & CROSS-DOMAIN CORRELATIONS
# ─────────────────────────────────────────────────────────────────────────────
cat("\n=== PARALLEL LGCM RESULTS ===\n")

# Fixed effects
cat("\nPopulation-average slopes:\n")
print(round(fixef(fit_parallel)[grepl("time_metric", rownames(fixef(fit_parallel))),], 4))

# Cross-domain slope correlations
draws <- as_draws_df(fit_parallel)

cor_map <- c(
  "PISA ~ R&D  (slope)"  = "cor_iso3c__pisaz_time_metric__rdz_time_metric",
  "PISA ~ Emp  (slope)"  = "cor_iso3c__pisaz_time_metric__empz_time_metric",
  "R&D  ~ Emp  (slope)"  = "cor_iso3c__rdz_time_metric__empz_time_metric",
  "PISA ~ R&D  (int)"    = "cor_iso3c__pisaz_Intercept__rdz_Intercept",
  "PISA ~ Emp  (int)"    = "cor_iso3c__pisaz_Intercept__empz_Intercept",
  "R&D  ~ Emp  (int)"    = "cor_iso3c__rdz_Intercept__empz_Intercept",
  "R&D  int~slope"       = "cor_iso3c__rdz_Intercept__rdz_time_metric",
  "Emp  int~slope"       = "cor_iso3c__empz_Intercept__empz_time_metric"
)

cat("\nRandom-effect correlations:\n")
cor_results <- lapply(names(cor_map), function(nm) {
  v <- draws[[cor_map[nm]]]
  tibble(
    label  = nm,
    r      = round(mean(v), 3),
    q2.5   = round(quantile(v, .025), 3),
    q97.5  = round(quantile(v, .975), 3),
    P_pos  = round(mean(v > 0), 3)
  )
})
print(bind_rows(cor_results))

# Turkey random effects
cat("\nTurkey random effects:\n")
re <- ranef(fit_parallel, summary=TRUE)
tur_domains <- c("pisaz_Intercept","pisaz_time_metric",
                 "rdz_Intercept",  "rdz_time_metric",
                 "empz_Intercept", "empz_time_metric")
for (d in tur_domains) {
  v <- re$iso3c["TUR",,d]
  cat(sprintf("  %-25s  %.3f  [%.3f, %.3f]\n", d,
              v["Estimate"], v["Q2.5"], v["Q97.5"]))
}

# Posterior predictive check
pp_check(fit_parallel, resp="pisaz", ndraws=50)


# ─────────────────────────────────────────────────────────────────────────────
# STEP 11 – CONDITIONAL LGCM: INSTRUCTIONAL QUALITY (frequentist, lme4)
# ─────────────────────────────────────────────────────────────────────────────
cat("\n=== CONDITIONAL LGCMs ===\n")

# --- Build enriched panel with TEACHSUP, DISCLIM, COGACT ---
# PISA 2022 student file: extract TEACHSUP, DISCLIM, BELONG, ESCS, COGACT items
# PISA 2012 student file: extract TEACHSUP, BELONG, ESCS

pisa_raw_dir <- here("data/pisa_raw")

# Helper: extract country-level weighted means from SAV file
extract_pisa_vars <- function(path, vars_select) {
  read_sav(path, col_select = all_of(vars_select)) |>
    zap_labels() |>
    group_by(CNT) |>
    summarise(across(-W_FSTUWT, ~ weighted.mean(.x, W_FSTUWT, na.rm=TRUE)),
              n_stu = n(), .groups="drop") |>
    mutate(iso3c = countrycode(CNT, "iso3c", "iso3c", warn=FALSE)) |>
    filter(!is.na(iso3c)) |>
    select(-CNT)
}

vars_2022 <- c("CNT","W_FSTUWT","TEACHSUP","DISCLIM","BELONG","ESCS",
               paste0("ST283Q0",1:9,"JA"), paste0("ST285Q0",1:9,"JA"))
vars_2012 <- c("CNT","W_FSTUWT","TEACHSUP","BELONG","ESCS")

iq22 <- extract_pisa_vars(file.path(pisa_raw_dir,"CY08MSP_STU_QQQ.SAV"), vars_2022) |>
  mutate(year = 2022,
         cogact = rowMeans(across(matches("ST283|ST285")), na.rm=TRUE)) |>
  select(iso3c, year, teachsup, disclim, belong, escs, cogact, n_stu)

iq12 <- extract_pisa_vars(file.path(pisa_raw_dir,"CY6_MS_CMB_STU_QQQ.sav"), vars_2012) |>
  mutate(year = 2012) |>
  select(iso3c, year, teachsup, belong, escs, n_stu)

# Wide format: one row per country
ts12 <- iq12 |> select(iso3c, teachsup_12=teachsup, belong_12=belong, escs_12=escs)
ts22 <- iq22 |> select(iso3c, teachsup_22=teachsup, disclim_22=disclim,
                        cogact_22=cogact, belong_22=belong, escs_22=escs)

ts_wide <- left_join(ts12, ts22, by="iso3c") |>
  mutate(d_ts = teachsup_22 - teachsup_12)

# Join to panel_z
panel_enriched <- panel_z |>
  left_join(ts_wide, by="iso3c") |>
  mutate(
    ts12_z      = as.numeric(scale(teachsup_12)),
    ts22_z      = as.numeric(scale(teachsup_22)),
    d_ts_z      = as.numeric(scale(d_ts)),
    disclim22_z = as.numeric(scale(disclim_22)),
    cogact22_z  = as.numeric(scale(cogact_22))
  )

saveRDS(panel_enriched, here("data/processed/panel_enriched.rds"))
write.csv(panel_enriched, here("output/tables/panel_enriched_data.csv"),
          row.names=FALSE, fileEncoding="UTF-8")

# Descriptive: country-level change in teacher support
both_cnt <- intersect(
  ts12$iso3c[!is.na(ts12$teachsup_12)],
  ts22$iso3c[!is.na(ts22$teachsup_22)]
)
delta <- ts_wide |>
  filter(iso3c %in% both_cnt) |>
  left_join(
    panel_enriched |>
      filter(year == 2012) |>
      select(iso3c, pisa_math_12=pisa_math, oecd_member),
    by = "iso3c"
  ) |>
  left_join(
    panel_enriched |>
      filter(year == 2022) |>
      select(iso3c, pisa_math_22=pisa_math),
    by = "iso3c"
  ) |>
  mutate(d_math = pisa_math_22 - pisa_math_12)

saveRDS(delta, here("data/processed/pisa_delta_2012_2022.rds"))
write.csv(delta, here("data/processed/pisa_delta_2012_2022.csv"),
          row.names=FALSE, fileEncoding="UTF-8")

ok <- !is.na(delta$d_ts)
cat("r(Δteachsup, Δmath):", round(cor(delta$d_ts[ok], delta$d_math[ok]), 3), "\n")
cat("OECD avg Δteachsup:", round(mean(delta$d_ts[delta$oecd_member & ok], na.rm=TRUE), 3), "\n")
cat("Turkey Δteachsup:", round(delta$d_ts[delta$iso3c=="TUR"], 3),
    "| Δmath:", round(delta$d_math[delta$iso3c=="TUR"], 1), "\n")

# ── Conditional frequentist models ──────────────────────────────────────────

d_m1 <- filter(panel_enriched, !is.na(pisa_z), !is.na(ts12_z))
d_m2 <- filter(panel_enriched, !is.na(pisa_z), !is.na(d_ts_z))
d_m3 <- filter(panel_enriched, !is.na(pisa_z), !is.na(ts12_z), !is.na(disclim22_z))

# M0: unconditional on restricted sample
m0_cond <- lmer(pisa_z ~ time_metric + (time_metric | iso3c),
                data=d_m1, REML=TRUE, control=ctrl)

# M1: level of teacher support 2012
m1 <- lmer(pisa_z ~ time_metric * ts12_z + (time_metric | iso3c),
           data=d_m1, REML=TRUE, control=ctrl)

# M2: change in teacher support
m2 <- lmer(pisa_z ~ time_metric * d_ts_z + (time_metric | iso3c),
           data=d_m2, REML=TRUE, control=ctrl)

# M3: teacher support + disciplinary climate
m3 <- lmer(pisa_z ~ time_metric * ts12_z + time_metric * disclim22_z +
             (time_metric | iso3c),
           data=d_m3, REML=TRUE, control=ctrl)

# Results summary
for (nm in c("m0_cond","m1","m2","m3")) {
  m   <- get(nm)
  fe  <- summary(m)$coefficients
  r2v <- r2(m)
  cat(sprintf("\n%s | AIC=%.1f | R2m=%.3f | R2c=%.3f\n",
              nm, AIC(m), r2v$R2_marginal, r2v$R2_conditional))
  ci <- confint(m, method="Wald", parm="beta_")
  print(round(cbind(fe, ci), 4))
}

saveRDS(m0_cond, here("models/m0_uncond.rds"))
saveRDS(m1,      here("models/m1_cond_ts12.rds"))
saveRDS(m2,      here("models/m2_cond_dts.rds"))
saveRDS(m3,      here("models/m3_cond_full.rds"))


# ─────────────────────────────────────────────────────────────────────────────
# STEP 12 – BAYESIAN CONDITIONAL LGCMs
# ─────────────────────────────────────────────────────────────────────────────
cat("\nFitting Bayesian conditional models...\n")

priors_cond <- c(
  prior(normal(0, 1),   class = "Intercept"),
  prior(normal(0, 0.5), class = "b"),
  prior(exponential(2), class = "sd"),
  prior(lkj(2),         class = "cor"),
  prior(exponential(2), class = "sigma")
)

# BM0: unconditional (restricted sample, for LOO comparison with BM1)
bm0 <- brm(
  pisa_z ~ time_metric + (time_metric | iso3c),
  data=d_m1, prior=priors_cond,
  chains=4, iter=4000, warmup=2000, cores=4, seed=2025,
  backend="rstan",
  file=here("models/fit_b_m0c_done"), silent=2
)

# BM1: time × TEACHSUP_2012
bm1 <- brm(
  pisa_z ~ time_metric * ts12_z + (time_metric | iso3c),
  data=d_m1, prior=priors_cond,
  chains=4, iter=4000, warmup=2000, cores=4, seed=2025,
  backend="rstan",
  file=here("models/fit_b_m1_ts12"), silent=2
)

# BM2: time × ΔTEACHSUP (6,000 iterations for better convergence)
bm2 <- brm(
  pisa_z ~ time_metric * d_ts_z + (time_metric | iso3c),
  data=d_m2, prior=priors_cond,
  chains=4, iter=6000, warmup=3000, cores=4, seed=2025,
  backend="rstan",
  file=here("models/fit_b_m2_dts_6k"), silent=2
)

# BM3: time × TS + time × DISCLIM22
bm3 <- brm(
  pisa_z ~ time_metric * ts12_z + time_metric * disclim22_z +
    (time_metric | iso3c),
  data=d_m3, prior=priors_cond,
  chains=4, iter=4000, warmup=2000, cores=4, seed=2025,
  backend="rstan",
  file=here("models/fit_b_m3_full"), silent=2
)

# Convergence
for (nm in c("bm0","bm1","bm2","bm3")) {
  m  <- get(nm)
  rh <- round(max(rhat(m), na.rm=TRUE), 4)
  cat(nm, "max Rhat:", rh, "\n")
}

# Posterior probabilities for interaction terms
draws1 <- as_draws_df(bm1)
draws2 <- as_draws_df(bm2)
draws3 <- as_draws_df(bm3)
cat("BM1 P(time:ts12_z > 0):",    round(mean(draws1$`b_time_metric:ts12_z` > 0),    3), "\n")
cat("BM2 P(time:d_ts_z > 0):",    round(mean(draws2$`b_time_metric:d_ts_z` > 0),    3), "\n")
cat("BM3 P(time:ts12_z > 0):",    round(mean(draws3$`b_time_metric:ts12_z` > 0),    3), "\n")
cat("BM3 P(time:disclim22_z > 0):",round(mean(draws3$`b_time_metric:disclim22_z` > 0), 3), "\n")

# LOO comparison: BM0 vs BM1 (same N)
loo_bm0 <- loo(bm0, cores=2)
loo_bm1 <- loo(bm1, cores=2)
cat("\nLOO comparison BM0 vs BM1:\n")
print(loo_compare(loo_bm0, loo_bm1))

# Turkey posterior (BM1)
cat("\nTurkey random effects (BM1):\n")
print(round(ranef(bm1, summary=TRUE)$iso3c["TUR",,], 4))


# ─────────────────────────────────────────────────────────────────────────────
# STEP 13 – FIGURES
# ─────────────────────────────────────────────────────────────────────────────
cat("\nGenerating figures...\n")

# ── Figure 1: Country-level trajectories (3-panel) ──────────────────────────
oecd_avg <- panel_z |>
  group_by(year) |>
  summarise(
    pisa_z = mean(pisa_z[oecd_member], na.rm=TRUE),
    rd_z   = mean(rd_z[oecd_member],   na.rm=TRUE),
    emp_z  = mean(emp_z[oecd_member],  na.rm=TRUE),
    .groups="drop"
  )

make_traj_panel <- function(dat, yvar, title, oecd_dat) {
  ggplot(dat, aes(year, .data[[yvar]], group=iso3c)) +
    geom_line(aes(
      colour = case_when(
        iso3c == "TUR" ~ "Turkey (OECD)",
        oecd_member    ~ "OECD average",
        TRUE           ~ "Non-OECD"
      ),
      linewidth = case_when(
        iso3c == "TUR" ~ 1.2,
        oecd_member    ~ 0.35,
        TRUE           ~ 0.25
      ),
      alpha = case_when(
        iso3c == "TUR" ~ 1,
        oecd_member    ~ 0.55,
        TRUE           ~ 0.3
      )
    )) +
    geom_line(data=oecd_dat, aes(year, .data[[yvar]], group=1),
              colour="black", linewidth=1.1, linetype="solid") +
    geom_point(data=oecd_dat, aes(year, .data[[yvar]]),
               colour="black", size=2.5) +
    scale_colour_manual(
      values=c("Turkey (OECD)"="red","OECD average"="grey50","Non-OECD"="grey80"),
      name=NULL
    ) +
    scale_linewidth_identity() +
    scale_alpha_identity() +
    scale_x_continuous(breaks=c(2012,2015,2018,2022)) +
    labs(title=title, x=NULL, y="z-score (2012 anchor)") +
    theme_minimal(base_size=10) +
    theme(legend.position="bottom", panel.grid.minor=element_blank())
}

fig_traj <- make_traj_panel(panel_z, "pisa_z", "A: PISA Mathematics", oecd_avg) |>
  { \(p1) make_traj_panel(panel_z, "rd_z", "B: R&D Expenditure (% GDP)", oecd_avg) |>
      { \(p2) make_traj_panel(panel_z, "emp_z", "C: Employment-to-Population (25-54)", oecd_avg) |>
          { \(p3) (p1 | p2 | p3) +
              plot_annotation(
                title   = "Country-level trajectories across three knowledge economy domains, 2012-2022",
                subtitle= "Turkey (red) vs. OECD average (black); z-scores standardised to 2012 sample mean/SD",
                caption = "Sources: PISA (learningtower), World Bank (R&D), ILO (ILOSTAT). N = 52 countries."
              ) }() }() }()

ggsave(here("output/figures/fig_country_trajectories.png"),
       fig_traj, width=13, height=5.5, dpi=300)
cat("  fig_country_trajectories.png saved\n")

# ── Figure 2: Teacher support slopegraph + scatter ───────────────────────────
pisa_panel_all <- readRDS(here("data/processed/pisa_teachsup_panel_2012_2022.rds"))
both_cnt2      <- names(which(table(pisa_panel_all$iso3c)==2))
sl   <- pisa_panel_all[pisa_panel_all$iso3c %in% both_cnt2 & !is.na(pisa_panel_all$teachsup),]
sl$hl <- ifelse(sl$iso3c=="TUR","Turkey",ifelse(sl$oecd,"OECD","Non-OECD"))
lbl  <- c("TUR","FIN","KOR","EST","POL","DEU","JPN","GRC","CAN","NOR")

delta_plot <- readRDS(here("data/processed/pisa_delta_2012_2022.rds"))
ok2  <- !is.na(delta_plot$d_ts)
lm_d <- lm(d_math ~ d_ts, data=delta_plot[ok2,])
r_v  <- round(cor(delta_plot$d_ts[ok2], delta_plot$d_math[ok2]),2)
delta_plot$hl <- ifelse(delta_plot$iso3c=="TUR","Turkey",
                  ifelse(delta_plot$oecd_member,"OECD","Non-OECD"))
cv <- c("Turkey"="red","OECD"="#1565C0","Non-OECD"="#90A4AE")
sv <- c("Turkey"=4,    "OECD"=2.5,      "Non-OECD"=1.8)

p_slope <- ggplot(sl, aes(factor(year), teachsup, group=iso3c,
                           colour=hl, alpha=hl, linewidth=hl)) +
  geom_line() + geom_point(aes(size=hl)) +
  geom_text(data=sl[sl$year==2022 & sl$iso3c %in% lbl,],
            aes(label=iso3c, colour=hl), hjust=-0.15, size=2.7, fontface="bold") +
  scale_colour_manual(values=c("Turkey"="red","OECD"="#78909C","Non-OECD"="#CFD8DC"), name=NULL) +
  scale_alpha_manual(values=c("Turkey"=1,"OECD"=0.5,"Non-OECD"=0.2), name=NULL) +
  scale_linewidth_manual(values=c("Turkey"=1.6,"OECD"=0.55,"Non-OECD"=0.35), name=NULL) +
  scale_size_manual(values=c("Turkey"=3,"OECD"=1.5,"Non-OECD"=0.8), name=NULL) +
  labs(title="A: Teacher Support Index 2012 vs. 2022", x=NULL,
       y="Teacher Support (OECD-standardised)") +
  theme_minimal(base_size=11) +
  theme(legend.position="bottom", panel.grid.minor=element_blank())

p_scatter <- ggplot(delta_plot[ok2,], aes(d_ts, d_math)) +
  geom_hline(yintercept=0, linetype="dashed", colour="grey60", linewidth=0.4) +
  geom_vline(xintercept=0, linetype="dashed", colour="grey60", linewidth=0.4) +
  geom_abline(intercept=coef(lm_d)[1], slope=coef(lm_d)[2],
              colour="grey30", linewidth=0.7, alpha=0.6) +
  geom_point(aes(colour=hl, size=hl), alpha=0.85) +
  geom_text(data=delta_plot[ok2 & delta_plot$hl!="Non-OECD",],
            aes(label=iso3c, colour=hl), vjust=-0.85, size=2.7, fontface="bold") +
  scale_colour_manual(values=cv, name=NULL) +
  scale_size_manual(values=sv, name=NULL) +
  annotate("text", x=0.2, y=min(delta_plot$d_math[ok2])*0.85,
           label=paste0("r = ",r_v), size=3.5, colour="grey30") +
  labs(title="B: Change in Teacher Support vs. Change in PISA Mathematics",
       x="Δ Teacher Support (2022 − 2012)", y="Δ PISA Mathematics (points)") +
  theme_minimal(base_size=11) +
  theme(legend.position="bottom", panel.grid.minor=element_blank())

fig2 <- (p_slope | p_scatter) +
  plot_annotation(
    title   = "Instructional Quality and PISA Mathematics Change, 2012-2022",
    subtitle= "N=57 countries. Turkey in red. Near-zero correlation between teacher support change and math change.",
    caption = "Sources: PISA 2012 & 2022, OECD. Weighted country means. TEACHSUP = OECD teacher support index."
  )
ggsave(here("output/figures/fig_teachsup_2012_2022.png"),
       fig2, width=13, height=6.5, dpi=300)
cat("  fig_teachsup_2012_2022.png saved\n")

# ── Figure 3: Conditional LGCM predicted trajectories + RE scatter ───────────
m1_loaded <- readRDS(here("models/m1_cond_ts12.rds"))
fe_m1     <- fixef(m1_loaded)
times     <- seq(0, 10, 0.5)
ts_levels <- c("High (+1.5 SD)"=1.5, "Average"=0.0, "Low (-1.5 SD)"=-1.5)

traj <- do.call(rbind, lapply(names(ts_levels), function(g) {
  ts <- ts_levels[g]
  data.frame(time_metric=times, group=g,
             pisa_z=fe_m1[1]+fe_m1[2]*times+fe_m1[3]*ts+fe_m1[4]*times*ts)
}))
traj$group <- factor(traj$group, levels=c("High (+1.5 SD)","Average","Low (-1.5 SD)"))
col_g <- c("High (+1.5 SD)"="#1565C0","Average"="#546E7A","Low (-1.5 SD)"="#EF6C00")

p_traj <- ggplot(traj, aes(time_metric, pisa_z, colour=group, linetype=group)) +
  geom_line(linewidth=1.1) +
  scale_colour_manual(values=col_g, name="2012 Teacher Support") +
  scale_linetype_manual(values=c("solid","dashed","dotdash"),
                        name="2012 Teacher Support") +
  scale_x_continuous(breaks=c(0,3,6,10), labels=c("2012","2015","2018","2022")) +
  labs(title="C: Predicted PISA Trajectories by 2012 Teacher Support Level",
       subtitle="Fixed-effect predictions from M1 (time × TEACHSUP_2012)",
       x=NULL, y="PISA Mathematics (z-score)") +
  theme_minimal(base_size=11) +
  theme(legend.position="bottom", panel.grid.minor=element_blank())

re_df <- as.data.frame(ranef(m1_loaded)$iso3c)
re_df$iso3c    <- rownames(re_df)
re_df$re_slope <- re_df[["time_metric"]]
pe_12 <- panel_enriched[panel_enriched$year==2012 & !is.na(panel_enriched$ts12_z),
                         c("iso3c","ts12_z","oecd_member")]
sc    <- merge(re_df[,c("iso3c","re_slope")], pe_12, by="iso3c")
sc$hl <- ifelse(sc$iso3c=="TUR","Turkey",
          ifelse(sc$oecd_member,"OECD","Non-OECD"))
r_re  <- round(cor(sc$ts12_z, sc$re_slope), 2)
lm_sc <- lm(re_slope ~ ts12_z, data=sc)

p_re <- ggplot(sc, aes(ts12_z, re_slope, colour=hl, size=hl)) +
  geom_hline(yintercept=0, linetype="dashed", colour="grey60") +
  geom_abline(intercept=coef(lm_sc)[1], slope=coef(lm_sc)[2],
              colour="grey30", linewidth=0.7, alpha=0.6) +
  geom_point(alpha=0.85) +
  geom_text(data=sc[sc$hl!="Non-OECD",],
            aes(label=iso3c), vjust=-0.85, size=2.7, fontface="bold") +
  scale_colour_manual(values=c("Turkey"="red","OECD"="#1565C0","Non-OECD"="#90A4AE"),
                      name=NULL) +
  scale_size_manual(values=c("Turkey"=4,"OECD"=2.5,"Non-OECD"=1.8), name=NULL) +
  annotate("text", x=1.2, y=min(sc$re_slope)*0.85,
           label=paste0("r = ",r_re), size=3.5, colour="grey30") +
  labs(title="D: 2012 Teacher Support vs. Country-Specific PISA Growth Rate",
       subtitle="Country random slope deviation from population average",
       x="Teacher Support 2012 (z-score)", y="Country RE Slope") +
  theme_minimal(base_size=11) +
  theme(legend.position="bottom", panel.grid.minor=element_blank())

fig3 <- (p_traj | p_re) +
  plot_annotation(
    title   = "Conditional LGCM: Teacher Support as Predictor of PISA Trajectories",
    subtitle= "Left: fixed-effect predicted trajectories. Right: 2012 teacher support vs. country growth rates.",
    caption = "M1: PISA_z ~ time_metric × TEACHSUP_2012 + (time_metric | country). N = 49 countries."
  )
ggsave(here("output/figures/fig_cond_lgcm.png"),
       fig3, width=13, height=6.5, dpi=300)
cat("  fig_cond_lgcm.png saved\n")


# ─────────────────────────────────────────────────────────────────────────────
# STEP 14 – TABLES (CSV)
# ─────────────────────────────────────────────────────────────────────────────
cat("\nExporting tables...\n")

# Table: fixed effects all conditional models
fe_table <- bind_rows(
  lapply(list(m0_cond=m0_cond, m1=m1, m2=m2, m3=m3), function(m) {
    fe  <- as.data.frame(summary(m)$coefficients)
    ci  <- as.data.frame(confint(m, method="Wald", parm="beta_"))
    r2v <- r2(m)
    fe$param <- rownames(fe)
    fe$lower <- ci[rownames(fe),"2.5 %"]
    fe$upper <- ci[rownames(fe),"97.5 %"]
    fe$AIC   <- round(AIC(m),1)
    fe$R2m   <- round(r2v$R2_marginal, 3)
    fe$R2c   <- round(r2v$R2_conditional, 3)
    fe
  }), .id="model"
)
write.csv(fe_table, here("output/tables/freq_conditional_models.csv"),
          row.names=FALSE)
cat("  freq_conditional_models.csv saved\n")

# Table: Bayesian posteriors BM0 and BM1
bm_table <- bind_rows(
  as.data.frame(fixef(bm0)) |> mutate(param=rownames(fixef(bm0)), model="BM0"),
  as.data.frame(fixef(bm1)) |> mutate(param=rownames(fixef(bm1)), model="BM1")
) |> mutate(across(where(is.numeric), ~round(.,4)))
write.csv(bm_table, here("output/tables/bayes_conditional_models.csv"),
          row.names=FALSE)
cat("  bayes_conditional_models.csv saved\n")

# Table: cross-domain correlations
write.csv(bind_rows(cor_results),
          here("output/tables/cross_domain_correlations.csv"),
          row.names=FALSE)
cat("  cross_domain_correlations.csv saved\n")


# ─────────────────────────────────────────────────────────────────────────────
# STEP 15 – SESSION INFO
# ─────────────────────────────────────────────────────────────────────────────
sink(here("output/tables/session_info.txt"))
cat("=== SESSION INFO ===\n\n")
cat("Date:", format(Sys.time(), "%Y-%m-%d %H:%M"), "\n\n")
cat("R version:", R.Version()$version.string, "\n\n")
cat("Key packages:\n")
for (p in pkgs) {
  v <- tryCatch(as.character(packageVersion(p)), error=function(e) "not installed")
  cat(sprintf("  %-20s %s\n", p, v))
}
cat("\n")
print(sessionInfo())
sink()

cat("\n=== ANALYSIS COMPLETE ===\n")
cat("All outputs saved to: ", here("output"), "\n")
cat("All models saved to: ",  here("models"), "\n")
