#!/usr/bin/env Rscript

# Observation-continuity inventory and controlled event-level case deletions
# for the accepted canonical PM EBS pollock fit. Each refit retains the
# accepted parameter map and starts from the converged full-data estimates.
# These are local influence diagnostics, not alternative assessments or LOO.

.libPaths(c(file.path(getwd(), ".r-lib-rceattle-5.8.1"), .libPaths()))

suppressPackageStartupMessages({
  library(dplyr)
  library(readr)
  library(Rceattle)
  library(tidyr)
})

stopifnot(packageVersion("Rceattle") == package_version("5.8.1"))

fit_file <- file.path(
  "results", "canonical_pm", "ebs_pollock_method_fits.rds"
)
output_dir <- file.path(
  "results", "canonical_pm", "case_deletion_diagnostics"
)
checkpoint_dir <- file.path(output_dir, "checkpoints")
dir.create(checkpoint_dir, recursive = TRUE, showWarnings = FALSE)

fits <- readRDS(fit_file)
reference <- fits$nonparametric_pm
stopifnot(reference$convergence$status == "OK")

assessment_end_year <- reference$data_list$endyr
display_years <- 2014:2026

# The schedule is explicit rather than inferred from gaps in the data. It
# describes the assessment input convention and is not a commitment about
# future field operations.
component_schedule <- tibble::tribble(
  ~Component, ~Table, ~Fleet, ~First_year, ~Last_year, ~Schedule,
  "Fishery age composition", "comp_data", "Fishery", 1964L, assessment_end_year, "annual",
  "BTS biomass index", "index_data", "BTS", 1982L, assessment_end_year, "annual",
  "BTS age composition", "comp_data", "BTS", 1982L, assessment_end_year, "annual",
  "BTS age-1 index", "index_data", "BTS_1", 1982L, assessment_end_year, "annual",
  "ATS biomass index", "index_data", "ATS", 1994L, assessment_end_year, "even",
  "ATS age composition", "comp_data", "ATS", 1994L, assessment_end_year, "even",
  "ATS age-1 index", "index_data", "ATS_1", 1994L, assessment_end_year, "even",
  "AVO index", "index_data", "AVO", 2006L, assessment_end_year, "annual"
)

is_scheduled <- function(year, schedule) {
  if (schedule == "annual") return(TRUE)
  if (schedule == "even") return(year %% 2L == 0L)
  stop("Unsupported schedule: ", schedule)
}

index_weight <- function(data, fleet, year) {
  if (fleet == "BTS") {
    covariance <- data$index_cov$BTS
    key <- as.character(year)
    if (key %in% rownames(covariance)) return(1 / covariance[key, key])
  }
  row <- data$index_data |>
    filter(.data$Fleet_name == fleet, abs(.data$Year) == year)
  if (nrow(row) != 1L || !is.finite(row$Log_sd) || row$Log_sd <= 0) {
    return(NA_real_)
  }
  1 / row$Log_sd^2
}

build_inventory_row <- function(component, table, fleet, first_year,
                                last_year, schedule, year, data) {
  source_data <- data[[table]]
  available <- source_data |>
    filter(.data$Fleet_name == fleet, abs(.data$Year) == year)
  fleet_type <- data$fleet_control$Fleet_type[
    match(fleet, data$fleet_control$Fleet_name)
  ]
  scheduled <- year >= first_year && year <= last_year &&
    is_scheduled(year, schedule)

  status <- if (year < first_year || year > last_year) {
    "Outside fitted span"
  } else if (!scheduled) {
    "Not scheduled"
  } else if (nrow(available) == 0L) {
    "No observation in scheduled year"
  } else if (fleet_type == "Off" || available$Year[1] < 0) {
    "Available but excluded from likelihood"
  } else {
    "Represented in base fit"
  }

  raw_weight <- NA_real_
  weight_definition <- if (table == "comp_data") {
    "Nominal composition sample size"
  } else if (fleet == "BTS") {
    "Inverse diagonal BTS covariance"
  } else {
    "Inverse squared observation SD"
  }
  if (nrow(available) == 1L) {
    raw_weight <- if (table == "comp_data") {
      available$Sample_size[1]
    } else {
      index_weight(data, fleet, year)
    }
  }

  tibble(
    Component = component,
    Table = table,
    Fleet = fleet,
    Year = year,
    Scheduled = scheduled,
    Status = status,
    Raw_weight = raw_weight,
    Weight_definition = weight_definition
  )
}

continuity_inventory <- bind_rows(lapply(seq_len(nrow(component_schedule)), function(i) {
  row <- component_schedule[i, ]
  bind_rows(lapply(display_years, function(year) {
    build_inventory_row(
      row$Component, row$Table, row$Fleet, row$First_year, row$Last_year,
      row$Schedule, year, reference$data_list
    )
  }))
})) |>
  group_by(.data$Component) |>
  mutate(
    Relative_weight = if (sum(is.finite(.data$Raw_weight)) > 1L) {
      scales::rescale(log1p(.data$Raw_weight), to = c(0.35, 1))
    } else ifelse(is.finite(.data$Raw_weight), 0.65, NA_real_),
    Display_weight = ifelse(
      .data$Status == "Represented in base fit",
      .data$Relative_weight,
      0.22
    )
  ) |>
  ungroup()

write_csv(
  continuity_inventory,
  file.path(output_dir, "continuity_inventory.csv")
)
write_csv(
  component_schedule,
  file.path(output_dir, "continuity_schedule.csv")
)

# Selected events span the 2020--2021 revision transition and the latest fitted
# event for each active observation series. The ATS 2020 composition is retained
# because its nominal sample size is one and provides a useful low-information
# case-deletion check.
case_manifest <- tibble::tribble(
  ~Scenario, ~Table, ~Fleet, ~Year, ~Selection_reason,
  "Fishery composition 2020", "comp_data", "Fishery", 2020L, "2020--2021 transition",
  "Fishery composition 2023", "comp_data", "Fishery", 2023L, "Latest fitted event",
  "BTS biomass index 2021", "index_data", "BTS", 2021L, "2020--2021 transition",
  "BTS biomass index 2024", "index_data", "BTS", 2024L, "Latest fitted event",
  "BTS age composition 2021", "comp_data", "BTS", 2021L, "2020--2021 transition",
  "BTS age composition 2024", "comp_data", "BTS", 2024L, "Latest fitted event",
  "ATS biomass index 2024", "index_data", "ATS", 2024L, "Latest fitted event",
  "ATS age composition 2020", "comp_data", "ATS", 2020L, "Nominal sample size of one",
  "ATS age composition 2024", "comp_data", "ATS", 2024L, "Latest fitted event",
  "ATS age-1 index 2022", "index_data", "ATS_1", 2022L, "Latest fitted event",
  "AVO index 2021", "index_data", "AVO", 2021L, "2020--2021 transition",
  "AVO index 2024", "index_data", "AVO", 2024L, "Latest fitted event"
)
case_manifest <- case_manifest |>
  mutate(
    Scenario_id = gsub("[^a-z0-9]+", "_", tolower(.data$Scenario)),
    Scenario_id = gsub("(^_|_$)", "", .data$Scenario_id),
    Checkpoint = file.path(checkpoint_dir, paste0(.data$Scenario_id, ".rds"))
  ) |>
  select(
    "Scenario_id", "Scenario", "Table", "Fleet", "Year",
    "Selection_reason", "Checkpoint"
  )

write_csv(
  select(case_manifest, -"Checkpoint"),
  file.path(output_dir, "case_deletion_manifest.csv")
)

if (identical(Sys.getenv("CASE_DELETE_INVENTORY_ONLY"), "true")) {
  message("Wrote continuity inventory and case-deletion manifest to ", output_dir)
  quit(save = "no", status = 0)
}

fit_control_case <- fit_control(
  verbose = 0,
  phase = TRUE,
  bias_adjust_proc = 0,
  bias_adjust_obs = 0,
  comp_offset = 1e-3
)
M1_fun <- build_M1(updateM1 = TRUE, M1_model = "fixed")

delete_event <- function(data, table, fleet, year) {
  changed <- data
  before <- nrow(changed[[table]])
  selected <- changed[[table]]$Fleet_name == fleet &
    changed[[table]]$Year == year
  if (sum(selected) != 1L) {
    stop("Expected one row for ", table, "/", fleet, "/", year,
         "; found ", sum(selected), ".")
  }
  changed[[table]] <- changed[[table]][!selected, , drop = FALSE]
  stopifnot(nrow(changed[[table]]) == before - 1L)
  changed
}

fit_variant <- function(data_list, inits = reference$obj$env$parList()) {
  suppressWarnings(fit_mod(
    data_list = data_list,
    inits = inits,
    map = reference$map,
    file = NULL,
    estimateMode = 0,
    random_rec = FALSE,
    msmMode = 0,
    initMode = "NonEquilibrium",
    M1Fun = M1_fun,
    fit_control = fit_control_case
  ))
}

extract_trajectory <- function(fit, scenario_id, scenario) {
  years <- as.integer(colnames(fit$quantities$R))
  keep <- years <= assessment_end_year
  tibble(
    Scenario_id = scenario_id,
    Scenario = scenario,
    Year = years[keep],
    SSB = as.numeric(fit$quantities$ssb[1, keep]),
    Recruitment = as.numeric(fit$quantities$R[1, keep])
  )
}

reference_trajectory <- extract_trajectory(
  reference, "reference", "All accepted observations"
)

extract_components <- function(fit, scenario_id, scenario) {
  components <- fit$quantities$jnll_comp
  as.data.frame(as.table(components), responseName = "NLL") |>
    transmute(
      Scenario_id = scenario_id,
      Scenario = scenario,
      Component = as.character(.data$Var1),
      Fleet = as.character(.data$Var2),
      NLL = .data$NLL
    )
}

extract_result <- function(fit, manifest_row, elapsed_seconds) {
  trajectory <- extract_trajectory(
    fit, manifest_row$Scenario_id, manifest_row$Scenario
  )
  comparison <- left_join(
    trajectory,
    select(reference_trajectory, "Year",
           Reference_SSB = "SSB",
           Reference_Recruitment = "Recruitment"),
    by = "Year"
  ) |>
    mutate(
      SSB_percent_change = 100 * (.data$SSB / .data$Reference_SSB - 1),
      Recruitment_percent_change =
        100 * (.data$Recruitment / .data$Reference_Recruitment - 1)
    )
  event_row <- comparison |>
    filter(.data$Year == manifest_row$Year)
  terminal_row <- comparison |>
    filter(.data$Year == assessment_end_year)
  convergence <- fit$convergence
  max_gradient <- convergence$checks$max_gradient$data$max_gradient
  condition_number <-
    convergence$checks$hessian_conditioning$data$condition_number
  pd_hessian <- convergence$checks$max_gradient$data$pdHess

  objective <- as.numeric(fit$opt$objective)
  reference_objective <- as.numeric(reference$opt$objective)
  summary <- tibble(
    Scenario_id = manifest_row$Scenario_id,
    Scenario = manifest_row$Scenario,
    Table = manifest_row$Table,
    Fleet = manifest_row$Fleet,
    Deleted_year = manifest_row$Year,
    Selection_reason = manifest_row$Selection_reason,
    Objective = objective,
    Reference_objective = reference_objective,
    Objective_difference = objective - reference_objective,
    Convergence_status = convergence$status,
    Maximum_gradient = as.numeric(max_gradient),
    Positive_definite_Hessian = isTRUE(pd_hessian),
    Hessian_condition_number = as.numeric(condition_number),
    Event_year_SSB_percent_change = event_row$SSB_percent_change,
    Event_year_recruitment_percent_change =
      event_row$Recruitment_percent_change,
    Terminal_SSB_percent_change = terminal_row$SSB_percent_change,
    Terminal_recruitment_percent_change =
      terminal_row$Recruitment_percent_change,
    Maximum_absolute_SSB_percent_change =
      max(abs(comparison$SSB_percent_change), na.rm = TRUE),
    Maximum_absolute_recruitment_percent_change =
      max(abs(comparison$Recruitment_percent_change), na.rm = TRUE),
    Elapsed_seconds = elapsed_seconds
  )

  list(
    summary = summary,
    trajectory = comparison,
    components = extract_components(
      fit, manifest_row$Scenario_id, manifest_row$Scenario
    )
  )
}

run_case <- function(manifest_row) {
  checkpoint <- manifest_row$Checkpoint
  if (file.exists(checkpoint)) {
    message("Using checkpoint: ", manifest_row$Scenario)
    return(readRDS(checkpoint))
  }
  message("Refitting without: ", manifest_row$Scenario)
  changed <- delete_event(
    reference$data_list, manifest_row$Table, manifest_row$Fleet,
    manifest_row$Year
  )
  start <- proc.time()[["elapsed"]]
  fit <- fit_variant(changed)
  elapsed <- proc.time()[["elapsed"]] - start
  result <- extract_result(fit, manifest_row, elapsed)
  saveRDS(result, checkpoint)
  message(
    "Completed ", manifest_row$Scenario,
    ": status=", result$summary$Convergence_status,
    ", max|gradient|=", format(result$summary$Maximum_gradient, digits = 3),
    ", elapsed=", round(elapsed, 1), " s"
  )
  result
}

case_results <- lapply(seq_len(nrow(case_manifest)), function(i) {
  run_case(case_manifest[i, ])
})

case_summary <- bind_rows(lapply(case_results, `[[`, "summary"))

# Refit the cases producing the largest change in SSB and recruitment from the
# saved full-data stage-A estimates. Agreement provides a targeted alternate-
# start check for the strongest reported influence results.
validation_ids <- unique(c(
  case_summary$Scenario_id[
    which.max(case_summary$Maximum_absolute_SSB_percent_change)
  ],
  case_summary$Scenario_id[
    which.max(case_summary$Maximum_absolute_recruitment_percent_change)
  ]
))

validate_start <- function(manifest_row, primary_summary) {
  checkpoint <- file.path(
    checkpoint_dir,
    paste0(manifest_row$Scenario_id, "_stage_a_start.rds")
  )
  if (file.exists(checkpoint)) {
    message("Using alternate-start checkpoint: ", manifest_row$Scenario)
    alternate <- readRDS(checkpoint)
  } else {
    message("Alternate-start validation: ", manifest_row$Scenario)
    changed <- delete_event(
      reference$data_list, manifest_row$Table, manifest_row$Fleet,
      manifest_row$Year
    )
    start <- proc.time()[["elapsed"]]
    alternate_fit <- fit_variant(
      changed,
      inits = fits$stage_a$obj$env$parList()
    )
    elapsed <- proc.time()[["elapsed"]] - start
    alternate <- extract_result(alternate_fit, manifest_row, elapsed)
    saveRDS(alternate, checkpoint)
  }
  alt <- alternate$summary
  tibble(
    Scenario_id = manifest_row$Scenario_id,
    Scenario = manifest_row$Scenario,
    Alternate_start = "Saved full-data stage-A estimates",
    Primary_objective = primary_summary$Objective,
    Alternate_objective = alt$Objective,
    Alternate_minus_primary_objective =
      alt$Objective - primary_summary$Objective,
    Alternate_convergence_status = alt$Convergence_status,
    Alternate_maximum_gradient = alt$Maximum_gradient,
    Alternate_positive_definite_Hessian = alt$Positive_definite_Hessian,
    Terminal_SSB_change_difference_percentage_points =
      alt$Terminal_SSB_percent_change -
        primary_summary$Terminal_SSB_percent_change,
    Terminal_recruitment_change_difference_percentage_points =
      alt$Terminal_recruitment_percent_change -
        primary_summary$Terminal_recruitment_percent_change
  )
}

start_validation <- bind_rows(lapply(validation_ids, function(scenario_id) {
  manifest_row <- case_manifest |>
    filter(.data$Scenario_id == scenario_id)
  primary_summary <- case_summary |>
    filter(.data$Scenario_id == scenario_id)
  validate_start(manifest_row, primary_summary)
}))

write_csv(
  start_validation,
  file.path(output_dir, "case_deletion_start_validation.csv")
)

case_trajectories <- bind_rows(
  reference_trajectory |>
    mutate(
      Reference_SSB = .data$SSB,
      Reference_Recruitment = .data$Recruitment,
      SSB_percent_change = 0,
      Recruitment_percent_change = 0
    ),
  bind_rows(lapply(case_results, `[[`, "trajectory"))
)
case_components <- bind_rows(
  extract_components(reference, "reference", "All accepted observations"),
  bind_rows(lapply(case_results, `[[`, "components"))
)

write_csv(case_summary, file.path(output_dir, "case_deletion_summary.csv"))
write_csv(
  case_trajectories,
  file.path(output_dir, "case_deletion_trajectories.csv")
)
write_csv(
  case_components,
  file.path(output_dir, "case_deletion_objective_components.csv")
)

stopifnot(
  nrow(case_summary) == nrow(case_manifest),
  !anyDuplicated(case_summary$Scenario_id),
  all(is.finite(case_summary$Maximum_gradient)),
  all(case_summary$Positive_definite_Hessian),
  all(case_summary$Convergence_status == "OK"),
  all(start_validation$Alternate_positive_definite_Hessian),
  all(start_validation$Alternate_convergence_status == "OK")
)

message("Wrote case-deletion diagnostics to ", output_dir)
