#!/usr/bin/env bash

set -e
set -o pipefail

: "${RUN_ID:?Set RUN_ID in the AIStudio task}"
: "${WORK_ROOT:?Set WORK_ROOT to the shared-storage path mounted in every Worker}"
TARGET_STEPS="${TARGET_STEPS:-2}"

[[ "$RUN_ID" =~ ^[A-Za-z0-9][A-Za-z0-9._-]*$ ]] || {
  echo "RUN_ID contains unsupported characters" >&2
  exit 1
}
[[ "$TARGET_STEPS" =~ ^[1-9][0-9]*$ ]] && (( TARGET_STEPS >= 2 )) || {
  echo "TARGET_STEPS must be an integer greater than or equal to 2" >&2
  exit 1
}

export RLINF_COMMIT=7d07a4212ee6858cc333e1d4fab7a37256d1f839
export MODEL_REVISION=6fc93244f442ee2b5ab5c8000687ef5f7ffe1d03
export DATA_REVISION=1799c00be3f1216ab55a5cae3562d654dbfd7d82
export RLINF_ROOT="$WORK_ROOT/code/checkouts/$RLINF_COMMIT"
export MODEL_ROOT="$WORK_ROOT/models/modelscope/deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B/$MODEL_REVISION"
export DATA_FILE="$WORK_ROOT/datasets/inclusionAI/AReaL-boba-Data/$DATA_REVISION/AReaL-boba-106k.jsonl"
export VERIFY_SCRIPT="$WORK_ROOT/tools/two-worker-grpo-pipeline/verify-grpo-pipeline-v1.py"
export OUTPUT_ROOT="$WORK_ROOT/runs/two-worker-grpo-pipeline"
export OUTPUT_DIR="$OUTPUT_ROOT/$RUN_ID"
export REPO_PATH="$RLINF_ROOT"
export PYTHONPATH="$RLINF_ROOT:/opt/Megatron-LM:${PYTHONPATH:-}"
export CUDA_DEVICE_MAX_CONNECTIONS=1
export TOKENIZERS_PARALLELISM=false
export PYTHONUNBUFFERED=1
export HF_HUB_OFFLINE=1
export HF_DATASETS_OFFLINE=1
export TRANSFORMERS_OFFLINE=1
export NCCL_IB_DISABLE=1
export RLINF_CATCH_FAILURE=0
export RLINF_CODE_WORKING_DIR=0

test ! -e "$OUTPUT_DIR" || {
  echo "Refusing to reuse OUTPUT_DIR: $OUTPUT_DIR" >&2
  exit 1
}
mkdir -p "$OUTPUT_DIR"

exec > >(tee -i "$OUTPUT_DIR/driver.log") 2>&1

record_startup_exit() {
  local startup_exit_code="$?"
  trap - EXIT
  printf 'STARTUP_EXIT_CODE=%s\n' "$startup_exit_code"
  sleep 2
  exit "$startup_exit_code"
}
trap record_startup_exit EXIT

printf 'RLINF_COMMIT=%s\n' "$RLINF_COMMIT"
printf 'MODEL_REVISION=%s\n' "$MODEL_REVISION"
printf 'DATA_REVISION=%s\n' "$DATA_REVISION"
printf 'RUN_ID=%s\n' "$RUN_ID"
printf 'TARGET_STEPS=%s\n' "$TARGET_STEPS"
printf 'OUTPUT_DIR=%s\n' "$OUTPUT_DIR"
printf 'NCCL_IB_DISABLE=%s\n' "$NCCL_IB_DISABLE"

echo "=== [Phase] Checking pinned inputs ==="
findmnt -T "$WORK_ROOT"
test -w "$WORK_ROOT"
test "$(git -C "$RLINF_ROOT" rev-parse HEAD)" = "$RLINF_COMMIT"
test -z "$(git -C "$RLINF_ROOT" status --porcelain)"
test -r "$RLINF_ROOT/tests/e2e_tests/reasoning/qwen2.5-1.5b-grpo-pipeline-fsdp-sgl.yaml"
test -r "$RLINF_ROOT/examples/reasoning/main_grpo.py"
test -r "$VERIFY_SCRIPT"
test -s "$MODEL_ROOT/config.json"
test -s "$MODEL_ROOT/tokenizer_config.json"
test -s "$MODEL_ROOT/model.safetensors"
test -s "$DATA_FILE"

echo "=== [Phase] Activating the RLinf environment ==="
source /opt/venv/reason/bin/activate
python --version
ray status

echo "=== [Phase] Checking the two-Worker Ray runtime ==="
python "$VERIFY_SCRIPT" preflight \
  --expected-nodes 2 \
  --expected-gpus-per-node 2 \
  --report "$OUTPUT_DIR/preflight.json"

echo "=== [Phase] Launching the RLinf GRPO Driver ==="
OVERRIDES=(
  "hydra.job.chdir=false"
  "hydra.run.dir=$OUTPUT_DIR/hydra"
  "hydra.output_subdir=.hydra"
  "cluster.num_nodes=2"
  "runner.max_epochs=1"
  "runner.max_steps=$TARGET_STEPS"
  "runner.val_check_interval=1"
  "runner.save_interval=$TARGET_STEPS"
  "runner.output_dir=$OUTPUT_ROOT"
  "runner.experiment_name=$RUN_ID"
  "runner.logger.log_path=$OUTPUT_DIR"
  "runner.logger.experiment_name=$RUN_ID"
  "+runner.per_worker_log=true"
  "inference.model.model_path=$MODEL_ROOT"
  "inference.tokenizer.tokenizer_model=$MODEL_ROOT"
  "rollout.model.model_path=$MODEL_ROOT"
  "actor.model.model_path=$MODEL_ROOT"
  "actor.tokenizer.tokenizer_model=$MODEL_ROOT"
  "+actor.fsdp_config.save_full_model_weights=false"
  "data.train_data_paths=[$DATA_FILE]"
  "data.val_data_paths=[$DATA_FILE]"
  "data.seed=1234"
  "actor.seed=1234"
)

RLINF_EXIT_CODE=0
python "$RLINF_ROOT/examples/reasoning/main_grpo.py" \
  --config-path "$RLINF_ROOT/tests/e2e_tests/reasoning" \
  --config-name qwen2.5-1.5b-grpo-pipeline-fsdp-sgl \
  "${OVERRIDES[@]}" || RLINF_EXIT_CODE="$?"

printf 'RLINF_PYTHON_EXIT_CODE=%s\n' "$RLINF_EXIT_CODE"
if (( RLINF_EXIT_CODE != 0 )); then
  exit "$RLINF_EXIT_CODE"
fi

sleep 2
echo "=== [Phase] Checking GRPO metrics and persistent output ==="
python "$VERIFY_SCRIPT" postflight \
  --output-dir "$OUTPUT_DIR" \
  --target-steps "$TARGET_STEPS" \
  --expected-nodes 2
