/* MSC-P-039 choosing a statistical test. MIT License. Secondary implementation.
   The three procedures, the scenarios, the nominal level and the tolerance are
   the declared ones, and the declared recursion is reproduced exactly, so this
   program draws the same datasets as the Python, R and SPSS references. Tail
   probabilities come from SAS's own PROBT and PROBNORM rather than from the
   incomplete beta and error function of the Python reference; the two agree far
   beyond the six printed decimals, but the last bits may differ. This file is an
   independent check of the flags and of the verdict. */
%let nominal = 0.05;
%let replications = 2000;
%let per_group = 25;
%let skew_log_sd = 0.85;
%let alternative_log_shift = 0.45;
%let normality_level = 0.05;
%let level_tolerance = 0.0125;
%let lcg_seed = 20260915;
%let lcg_multiplier = 1103515245;
%let lcg_increment = 12345;
%let lcg_modulus = 2147483648;
/* SAS numbers are doubles, exact for integers only up to 2^53. Written directly,
   &lcg_multiplier * state reaches 2.4e18 and would be rounded, giving a different
   recursion. The split is exact: 16838 * 65536 + 20077 is the multiplier. */
%let lcg_high = 16838;
%let lcg_low = 20077;

filename p039csv "public/datasets/msc-p039-two-group-outcome.csv";
data _null_;
  infile p039csv obs=1 lrecl=32767 truncover;
  input;
  if strip(_infile_) ne "unit_id,group,minutes" then do;
    put "ERROR: exact schema required"; abort cancel;
  end;
run;
data p039;
  infile p039csv dsd firstobs=2 truncover lrecl=32767 end=eof;
  input @;
  if countc(_infile_, ',') ne 2 then do; put "ERROR: exactly three cells required per row"; abort cancel; end;
  length unit_id $8 group $8;
  input unit_id $ group $ minutes;
  if unit_id ne cats('U', put(_n_, z3.)) then do; put "ERROR: units must be ordered U001, U002, ... without gaps"; abort cancel; end;
  if group not in ("control", "treated") then do; put "ERROR: the group column must contain only control and treated"; abort cancel; end;
  if missing(minutes) or minutes <= 0 then do; put "ERROR: the outcome must be finite and strictly positive"; abort cancel; end;
  if eof then call symputx("units", _n_);
run;
%if &units < 20 %then %do; %put ERROR: too few sampling units for the declared comparison; %abort cancel; %end;

/* The worked file: the three procedures, run once, so that they can disagree. */
proc means data=p039 noprint nway;
  class group; var minutes;
  output out=summary(drop=_type_ _freq_) n=n mean=mean var=variance
         skewness=skewness kurtosis=kurtosis;
run;
data _null_;
  set summary end=last;
  retain n1 m1 v1 s1 k1 n2 m2 v2 s2 k2;
  if group = "control" then do; n1 = n; m1 = mean; v1 = variance; s1 = skewness; k1 = kurtosis; end;
  else do; n2 = n; m2 = mean; v2 = variance; s2 = skewness; k2 = kurtosis; end;
  if last then do;
    if n1 < 10 or n2 < 10 then do; put "ERROR: each declared group needs at least ten sampling units"; abort cancel; end;
    statistic = (m2 - m1) / sqrt(v1 / n1 + v2 / n2);
    degrees = (v1 / n1 + v2 / n2) ** 2 / ((v1 / n1) ** 2 / (n1 - 1) + (v2 / n2) ** 2 / (n2 - 1));
    put "welch_statistic=" statistic 12.6;
    put "welch_degrees_of_freedom=" degrees 12.6;
    probability = 2 * probt(-abs(statistic), degrees);
    put "welch_probability=" probability 12.6;
  end;
run;
proc npar1way data=p039 wilcoxon noprint;
  class group; var minutes;
  output out=rank_result wilcoxon;
run;
proc print data=rank_result(keep=_wil_ _z_ p2_wil) noobs; run;

/* The simulation that decides which procedure may be trusted. */
%macro scenario(label=, skewed=, shift=);
  data one_scenario;
    length label $20;
    label = "&label";
    array first[&per_group] _temporary_;
    array second[&per_group] _temporary_;
    array pooled[%eval(2 * &per_group)] _temporary_;
    array side_of[%eval(2 * &per_group)] _temporary_;
    retain state &lcg_seed spare . has_spare 0;
    welch_rejections = 0; rank_rejections = 0; pretest_rejections = 0; chose_welch = 0;
    do replication = 1 to &replications;
      do i = 1 to &per_group;
        do side = 1 to 2;
          /* Declared Box-Muller transform over the declared recursion. */
          if has_spare then do; value = spare; has_spare = 0; end;
          else do;
            high = mod(&lcg_high * state, &lcg_modulus);
            state = mod(high * 65536 + &lcg_low * state + &lcg_increment, &lcg_modulus);
            u1 = (state + 0.5) / &lcg_modulus;
            high = mod(&lcg_high * state, &lcg_modulus);
            state = mod(high * 65536 + &lcg_low * state + &lcg_increment, &lcg_modulus);
            u2 = (state + 0.5) / &lcg_modulus;
            radius = sqrt(-2 * log(u1)); angle = 2 * constant('pi') * u2;
            spare = radius * sin(angle); has_spare = 1;
            value = radius * cos(angle);
          end;
          if side = 1 then do;
            if &skewed then first[i] = exp(&skew_log_sd * value); else first[i] = value;
          end;
          else do;
            if &skewed then second[i] = exp(&skew_log_sd * value + &shift); else second[i] = value + &shift;
          end;
        end;
      end;
      link procedures;
    end;
    welch_rate = welch_rejections / &replications;
    rank_rate = rank_rejections / &replications;
    pretest_rate = pretest_rejections / &replications;
    chose_welch_share = chose_welch / &replications;
    output;
    return;

  procedures:
    /* Welch, on the two declared groups. */
    sum1 = 0; sum2 = 0;
    do i = 1 to &per_group; sum1 + first[i]; sum2 + second[i]; end;
    mean1 = sum1 / &per_group; mean2 = sum2 / &per_group;
    ss1 = 0; ss2 = 0; third1 = 0; third2 = 0; fourth1 = 0; fourth2 = 0;
    do i = 1 to &per_group;
      d1 = first[i] - mean1; d2 = second[i] - mean2;
      ss1 + d1 ** 2; ss2 + d2 ** 2;
      third1 + d1 ** 3; third2 + d2 ** 3;
      fourth1 + d1 ** 4; fourth2 + d2 ** 4;
    end;
    var1 = ss1 / (&per_group - 1); var2 = ss2 / (&per_group - 1);
    tstat = (mean2 - mean1) / sqrt(var1 / &per_group + var2 / &per_group);
    tdf = (var1 / &per_group + var2 / &per_group) ** 2
        / ((var1 / &per_group) ** 2 / (&per_group - 1) + (var2 / &per_group) ** 2 / (&per_group - 1));
    welch_p = 2 * probt(-abs(tstat), tdf);
    if welch_p < &nominal then welch_rejections + 1;
    /* Rank test, normal approximation without a continuity correction.
       The combined sample is sorted by insertion, ties receive the average rank,
       and the tie correction of the declared variance is applied. */
    total = 2 * &per_group;
    do i = 1 to &per_group;
      pooled[i] = first[i]; side_of[i] = 0;
      pooled[&per_group + i] = second[i]; side_of[&per_group + i] = 1;
    end;
    do i = 2 to total;
      key = pooled[i]; key_side = side_of[i]; j = i - 1;
      do while (j >= 1 and pooled[j] > key);
        pooled[j + 1] = pooled[j]; side_of[j + 1] = side_of[j]; j = j - 1;
      end;
      pooled[j + 1] = key; side_of[j + 1] = key_side;
    end;
    rank_sum = 0; correction = 0; i = 1;
    do while (i <= total);
      stop_at = i;
      do while (stop_at < total and pooled[stop_at + 1] = pooled[i]); stop_at = stop_at + 1; end;
      average_rank = (i + stop_at) / 2;
      do j = i to stop_at;
        if side_of[j] = 0 then rank_sum + average_rank;
      end;
      tied = stop_at - i + 1;
      correction + (tied ** 3 - tied);
      i = stop_at + 1;
    end;
    u_statistic = rank_sum - &per_group * (&per_group + 1) / 2;
    u_variance = &per_group * &per_group / 12 * ((total + 1) - correction / (total * (total - 1)));
    if u_variance <= 0 then do; put "ERROR: the declared rank test has no variance"; abort cancel; end;
    u_z = (u_statistic - &per_group * &per_group / 2) / sqrt(u_variance);
    rank_p = 2 * probnorm(-abs(u_z));
    if rank_p < &nominal then rank_rejections + 1;
    /* Jarque-Bera in each group, then the choice it dictates. */
    m2_1 = ss1 / &per_group; m2_2 = ss2 / &per_group;
    jb1 = &per_group / 6 * ((third1 / &per_group / m2_1 ** 1.5) ** 2
        + ((fourth1 / &per_group / m2_1 ** 2) - 3) ** 2 / 4);
    jb2 = &per_group / 6 * ((third2 / &per_group / m2_2 ** 1.5) ** 2
        + ((fourth2 / &per_group / m2_2 ** 2) - 3) ** 2 / 4);
    if exp(-jb1 / 2) > &normality_level and exp(-jb2 / 2) > &normality_level then do;
      chose_welch + 1;
      if welch_p < &nominal then pretest_rejections + 1;
    end;
    else if rank_p < &nominal then pretest_rejections + 1;
    return;
  run;
  proc append base=scenarios data=one_scenario force; run;
%mend scenario;

data scenarios; length label $20; stop; run;
%scenario(label=normal_null, skewed=0, shift=0);
%scenario(label=skewed_null, skewed=1, shift=0);
%scenario(label=skewed_alternative, skewed=1, shift=&alternative_log_shift);

data flags;
  length flag $4 verdict $40;
  set scenarios;
  if label in ("normal_null", "skewed_null");
  welch_flag = ifc(abs(welch_rate - &nominal) <= &level_tolerance, "PASS", "FAIL");
  rank_flag = ifc(abs(rank_rate - &nominal) <= &level_tolerance, "PASS", "FAIL");
  pretest_flag = ifc(abs(pretest_rate - &nominal) <= &level_tolerance, "PASS", "FAIL");
run;
proc sql noprint;
  select count(*) into :failing from flags where pretest_flag = "FAIL";
quit;
data verdict;
  length verdict $40;
  verdict = ifc(&failing = 0, "PRETEST_RULE_HOLDS_ITS_LEVEL", "PRETEST_RULE_DOES_NOT_HOLD_ITS_LEVEL");
run;
proc print data=scenarios noobs; run;
proc print data=flags noobs; run;
proc print data=verdict noobs; run;
