#!/bin/bash

# MIT No Attribution
# Copyright 2020 Amazon.com, Inc. or its affiliates. All Rights Reserved.

set -euo pipefail

slurm_command="$(basename "$0")"
cluster_config="/opt/parallelcluster/shared/cluster-config.yaml"
reserved_idle_cost_center="idle"

cost_center=""
cluster_name=""
declare -a passthrough_args=()

urlencode() {
    python3 -c 'from urllib.parse import quote; import sys; print(quote(sys.argv[1], safe=""))' "$1"
}

aws_console_base_url() {
    local console_region="$1"
    if [ -n "${console_region}" ]; then
        printf "https://%s.console.aws.amazon.com" "${console_region}"
    else
        printf "https://console.aws.amazon.com"
    fi
}

budget_monitor_url() {
    local console_base
    console_base="$(aws_console_base_url "${region:-}")"
    printf "%s/billing/home#/budgets/details?name=%s" \
        "${console_base}" "$(urlencode "$1")"
}

cost_center_monitor_url() {
    local console_region="${cost_center_region:-${region:-}}"
    local console_base
    console_base="$(aws_console_base_url "${console_region}")"
    printf "%s/dynamodbv2/home?region=%s#item-explorer?table=%s&cost_center=%s" \
        "${console_base}" \
        "$(urlencode "${console_region}")" \
        "$(urlencode "${cost_center_table:-}")" \
        "$(urlencode "$1")"
}

fail() {
    local message="$1"
    echo "ERROR: ${message}" >&2
    if [ -n "${cluster_name:-}" ]; then
        echo "Cluster budget: ${cluster_name}" >&2
        echo "Cluster budget monitor: $(budget_monitor_url "${cluster_name}")" >&2
    fi
    if [ -n "${cost_center:-}" ]; then
        echo "Cost center: ${cost_center}" >&2
        echo "Cost-center report: $(cost_center_monitor_url "${cost_center}")" >&2
    fi
    exit 1
}

info() {
    echo "DYEC sbatch enforcement: $*" >&2
}

budget_override_enabled() {
    local variable_name="$1"
    local value="${!variable_name-}"
    if [ -z "${!variable_name+x}" ]; then
        return 1
    fi
    value="$(printf '%s' "$value" | tr '[:upper:]' '[:lower:]' | tr -d '[:space:]')"
    [ "$value" != "0" ] && [ "$value" != "false" ]
}

budget_warning_or_fail() {
    local variable_name="$1"
    local message="$2"
    if budget_override_enabled "$variable_name"; then
        echo "WARNING: ${message} Continuing because ${variable_name} is explicitly enabled." >&2
        return 0
    fi
    fail "$message"
}

while [ "$#" -gt 0 ]; do
    case "$1" in
        --comment)
            if [ -n "${cost_center}" ]; then
                fail "multiple --comment values are not allowed"
            fi
            shift
            if [ "$#" -eq 0 ] || [ -z "${1:-}" ]; then
                fail "--comment requires a cost-center value"
            fi
            cost_center="$1"
            ;;
        --comment=*)
            if [ -n "${cost_center}" ]; then
                fail "multiple --comment values are not allowed"
            fi
            cost_center="${1#--comment=}"
            if [ -z "${cost_center}" ]; then
                fail "--comment requires a cost-center value"
            fi
            ;;
        --mem|--mem=*|--mem-per-cpu|--mem-per-cpu=*|--mem-per-gpu|--mem-per-gpu=*|--mem-per-tres|--mem-per-tres=*)
            fail "Slurm memory placement is disabled; remove '${1}'. Select capacity with --partition and CPU/thread count."
            ;;
        *)
            passthrough_args+=("$1")
            ;;
    esac
    shift
done

if [ -z "${cost_center}" ]; then
    echo "ERROR: Please specify a cost center with --comment <cost-center>." >&2
    exit 1
fi

if [[ "${cost_center}" =~ [[:space:]\|] ]]; then
    fail "--comment cost center must not contain whitespace or '|'"
fi

if ! [[ "${cost_center}" =~ ^[A-Za-z0-9][A-Za-z0-9_.:@+-]{0,127}$ ]]; then
    fail "--comment cost center must start with a letter or number and contain only letters, numbers, '.', '_', ':', '@', '+', or '-'"
fi

if [ "${cost_center}" = "${reserved_idle_cost_center}" ]; then
    fail "'${reserved_idle_cost_center}' is reserved for no-job allocation and cannot be submitted"
fi

if [ ! -f "${cluster_config}" ]; then
    fail "cluster config not found at ${cluster_config}"
fi

cluster_tag_value() {
    local key="$1"
    awk -v key="${key}" '$0 ~ "Key: " key "$" {getline; print $2; exit}' "${cluster_config}" |
        sed -e 's/^"//' -e 's/"$//' -e "s/^'//" -e "s/'$//"
}

cluster_name="$(cluster_tag_value "aws-parallelcluster-clustername")"
if [ -z "${cluster_name}" ]; then
    fail "aws-parallelcluster-clustername tag is missing from ${cluster_config}"
fi

enforce_budget="$(
    cluster_tag_value "aws-parallelcluster-enforce-budget"
)"
cost_center_region="$(cluster_tag_value "aws-parallelcluster-cost-center-region")"
cost_center_table="$(cluster_tag_value "aws-parallelcluster-cost-center-table")"

if [ -z "${cost_center_region}" ]; then
    fail "aws-parallelcluster-cost-center-region tag is missing from ${cluster_config}"
fi
if [ -z "${cost_center_table}" ]; then
    fail "aws-parallelcluster-cost-center-table tag is missing from ${cluster_config}"
fi

case "${enforce_budget}" in
    skip|false|False|0|no|No)
        info "cluster AWS Budget check skipped by cluster config '${enforce_budget}' for cluster '${cluster_name}'."
        ;;
    true|True|1|yes|Yes|enforce|enforced)
        enforce_budget="true"
        ;;
    "")
        fail "aws-parallelcluster-enforce-budget tag is missing from ${cluster_config}"
        ;;
    *)
        fail "unknown aws-parallelcluster-enforce-budget value '${enforce_budget}'"
        ;;
esac

current_user="${USER:-$(id -un 2>/dev/null || true)}"
if [ -z "${current_user}" ]; then
    fail "USER is not set and id -un failed; cannot validate cost-center membership"
fi
current_groups="$(id -Gn "${current_user}" 2>/dev/null | tr ' ' ',' || true)"

region="$(
    awk -F'=' '$1 == "cfn_region" {print $2; exit}' /etc/parallelcluster/cfnconfig 2>/dev/null ||
        true
)"
if [ -z "${region}" ]; then
    fail "region is not set in /etc/parallelcluster/cfnconfig; cannot verify budget"
fi

AWS_ACCOUNT_ID="$(aws sts get-caller-identity --query "Account" --output text 2>/dev/null || true)"
if [ -z "${AWS_ACCOUNT_ID}" ] || [ "${AWS_ACCOUNT_ID}" = "None" ]; then
    fail "unable to resolve AWS account id with sts get-caller-identity"
fi

if [ "${enforce_budget}" = "true" ]; then
    budget_values="$(
        aws budgets describe-budget \
            --account-id "${AWS_ACCOUNT_ID}" \
            --budget-name "${cluster_name}" \
            --region "${region}" \
            --query 'Budget.[BudgetLimit.Amount,CalculatedSpend.ActualSpend.Amount,BudgetLimit.Unit]' \
            --output text 2>&1
    )" || fail "no readable AWS Budget named '${cluster_name}' in account ${AWS_ACCOUNT_ID}: ${budget_values}"

    read -r total_budget used_budget budget_unit _extra <<< "${budget_values}"
    if [ -z "${total_budget:-}" ] || [ -z "${used_budget:-}" ] || [ "${total_budget}" = "None" ] || [ "${used_budget}" = "None" ]; then
        fail "AWS Budget '${cluster_name}' did not include BudgetLimit.Amount and CalculatedSpend.ActualSpend.Amount"
    fi

    if ! awk -v total="${total_budget}" -v used="${used_budget}" 'BEGIN { exit !((total + 0) > 0 && used >= 0) }'; then
        fail "AWS Budget '${cluster_name}' has invalid numeric values total='${total_budget}' used='${used_budget}'"
    fi

    percent_used="$(awk -v total="${total_budget}" -v used="${used_budget}" 'BEGIN { printf "%.2f", (used / total) * 100 }')"

    info "cluster '${cluster_name}' budget in region '${region}': total=${total_budget} ${budget_unit:-USD}, used=${used_budget} ${budget_unit:-USD}, percent=${percent_used}%."
    info "cluster budget monitor: $(budget_monitor_url "${cluster_name}")"

    if awk -v percent="${percent_used}" 'BEGIN { exit !(percent >= 100) }'; then
        budget_warning_or_fail "DAY_PASS_ON_BUDGET_EXCEEDED" "AWS Budget '${cluster_name}' is exhausted: ${percent_used}% used (${used_budget}/${total_budget} ${budget_unit:-USD})"
    fi
fi

cost_center_json="$(
    aws dynamodb get-item \
        --table-name "${cost_center_table}" \
        --region "${cost_center_region}" \
        --key "{\"cost_center\":{\"S\":\"${cost_center}\"}}" \
        --consistent-read \
        --output json 2>&1
)" || fail "unable to read cost-center registry '${cost_center_table}' in ${cost_center_region}: ${cost_center_json}"

registry_validation="$(
    COST_CENTER_JSON="${cost_center_json}" \
    COST_CENTER="${cost_center}" \
    CURRENT_USER="${current_user}" \
    CURRENT_GROUPS="${current_groups}" \
    python3 <<'PY'
import json
import os
import re
import sys
from datetime import datetime, timezone

raw_payload = os.environ["COST_CENTER_JSON"]
if not raw_payload.strip():
    print("cost-center registry lookup returned empty JSON")
    sys.exit(14)
try:
    payload = json.loads(raw_payload)
except json.JSONDecodeError as exc:
    print(f"cost-center registry lookup returned invalid JSON: {exc.msg}")
    sys.exit(14)
if not isinstance(payload, dict):
    print("cost-center registry lookup returned JSON that is not an object")
    sys.exit(14)
item = payload.get("Item") or {}
if not item:
    print(f"cost center '{os.environ.get('COST_CENTER', '')}' does not exist")
    sys.exit(10)
status = item.get("status", {}).get("S", "")
if status != "active":
    print(f"cost center is not active: status={status or '<missing>'}")
    sys.exit(11)
active_until = item.get("active_until", {}).get("S", "")
if active_until:
    if not re.fullmatch(r"\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}Z", active_until):
        print("cost center active_until must be exactly YYYY-MM-DDTHH:MM:SSZ in UTC")
        sys.exit(12)
    try:
        expires = datetime.strptime(active_until, "%Y-%m-%dT%H:%M:%SZ").replace(
            tzinfo=timezone.utc
        )
    except ValueError:
        print(f"cost center active_until is not a valid UTC datetime: {active_until}")
        sys.exit(12)
    if datetime.now(timezone.utc) >= expires:
        print(f"cost center expired at active_until={active_until}")
        sys.exit(15)
allowed_users = set(item.get("allowed_users", {}).get("SS", []))
allowed_groups = set(item.get("allowed_groups", {}).get("SS", []))
current_user = os.environ["CURRENT_USER"]
current_groups = {g for g in os.environ.get("CURRENT_GROUPS", "").split(",") if g}
if "*" not in allowed_users and current_user not in allowed_users and not (allowed_groups & current_groups):
    print(f"user '{current_user}' is not authorized for this cost center")
    sys.exit(13)
print(f"active_until={active_until or '<none>'}")
PY
)" || fail "${registry_validation}"

info "cost center '${cost_center}' admission ok: ${registry_validation}."
info "cost-center report: $(cost_center_monitor_url "${cost_center}")"

exec "/opt/slurm/sbin/${slurm_command}" --comment="${cost_center}" --export=ALL "${passthrough_args[@]}"
