#!/usr/bin/bash
# SPDX-License-Identifier: Apache-2.0
# onednn test runner - runs gtests and optionally benchdnn smoke tests

set -euo pipefail

GTESTS_DIR=/usr/libexec/onednn/gtests
BENCHDNN=/usr/libexec/onednn/benchdnn
INPUTS_DIR=/usr/libexec/onednn/inputs

usage() {
    cat << USAGE
usage: onednn-test [OPTIONS]

run onednn validation tests

OPTIONS:
    --test=TESTS        which gtests to run (default: all)
                        'all' = all tests, '' = none, or comma-separated list
    --benchdnn=DRIVERS  which benchdnn drivers to run (default: none)
                        comma-separated list: matmul,conv,eltwise,...
    --engine=ENGINE     test engine: cpu (default) or gpu
    --filter=PATTERN    only run tests matching pattern
    --list              list available tests
    -h, --help          show this help

EXAMPLES:
    onednn-test                                    # all gtests on cpu
    onednn-test --engine=gpu                       # all gtests on gpu
    onednn-test --benchdnn=matmul,conv,eltwise     # gtests + benchdnn
    onednn-test --test='' --benchdnn=matmul        # benchdnn only
    onednn-test --filter='*matmul*'                # matmul-related gtests
USAGE
}

ENGINE="cpu"
TEST_LIST="all"
BENCHDNN_LIST=""
FILTER="*"
LIST_ONLY=0

while [[ $# -gt 0 ]]; do
    case $1 in
        --test=*)
            TEST_LIST="${1#*=}"
            shift
            ;;
        --benchdnn=*)
            BENCHDNN_LIST="${1#*=}"
            shift
            ;;
        --engine=*)
            ENGINE="${1#*=}"
            shift
            ;;
        --filter=*)
            FILTER="${1#*=}"
            shift
            ;;
        --list)
            LIST_ONLY=1
            shift
            ;;
        -h|--help)
            usage
            exit 0
            ;;
        *)
            echo "unknown option: $1"
            usage
            exit 1
            ;;
    esac
done

if [[ $LIST_ONLY -eq 1 ]]; then
    echo "gtests:"
    ls -1 "$GTESTS_DIR" | grep -v "\.so"
    if [[ -x "$BENCHDNN" ]]; then
        echo ""
        echo "benchdnn drivers:"
        ls -1 "$INPUTS_DIR"
    fi
    exit 0
fi

export DNNL_DEFAULT_FPMATH_MODE=strict

passed=0
failed=0
skipped=0

# run gtests if requested
if [[ -n "$TEST_LIST" ]]; then
    echo "running onednn gtests (engine=$ENGINE)..."
    echo ""

    for test in "$GTESTS_DIR"/*; do
        [[ -x "$test" ]] || continue
        testname=$(basename "$test")

        # check against --test= list
        if [[ "$TEST_LIST" != "all" ]]; then
            match=0
            IFS=',' read -ra TESTS <<< "$TEST_LIST"
            for pattern in "${TESTS[@]}"; do
                if [[ "$testname" == $pattern ]]; then
                    match=1
                    break
                fi
            done
            [[ $match -eq 0 ]] && continue
        fi

        # check against --filter
        case "$testname" in
            $FILTER)
                ;;
            *)
                continue
                ;;
        esac

        # skip buffer tests on cpu
        if [[ "$ENGINE" == "cpu" && "$testname" == *_buffer ]]; then
            echo "[ SKIP ] $testname (buffer tests require gpu)"
            skipped=$((skipped + 1))
            continue
        fi

        echo "[ RUN  ] $testname"
        set +e
        "$test" --gtest_color=no > /dev/null 2>&1
        result=$?
        set -e
        if [[ $result -eq 0 ]]; then
            echo "[  OK  ] $testname"
            passed=$((passed + 1))
        else
            echo "[ FAIL ] $testname"
            failed=$((failed + 1))
        fi
    done
fi

# run benchdnn if requested
if [[ -n "$BENCHDNN_LIST" && -x "$BENCHDNN" ]]; then
    echo ""
    echo "running benchdnn tests (engine=$ENGINE)..."
    echo ""

    IFS=',' read -ra DRIVERS <<< "$BENCHDNN_LIST"
    for driver in "${DRIVERS[@]}"; do
        smoke_batch="test_${driver}_smoke"
        if [[ -f "$INPUTS_DIR/$driver/$smoke_batch" ]]; then
            echo "[ RUN  ] benchdnn --$driver --batch=$smoke_batch"
            set +e
            "$BENCHDNN" --engine="$ENGINE" --"$driver" --batch="$smoke_batch" --mode=C > /dev/null 2>&1
            result=$?
            set -e
            if [[ $result -eq 0 ]]; then
                echo "[  OK  ] benchdnn $driver smoke"
                passed=$((passed + 1))
            else
                echo "[ FAIL ] benchdnn $driver smoke"
                failed=$((failed + 1))
            fi
        else
            echo "[ SKIP ] benchdnn $driver (no smoke batch found)"
            skipped=$((skipped + 1))
        fi
    done
fi

echo ""
echo "========================================"
echo "passed:  $passed"
echo "failed:  $failed"
echo "skipped: $skipped"
echo "========================================"

[[ $failed -eq 0 ]]
