# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.

import subprocess
import sys
from pathlib import Path
from types import SimpleNamespace

import pytest

from torchtitan.components.checkpointer import CheckpointManager
from torchtitan.models.common.attention import VarlenInnerAttention
from torchtitan_recipes.tests.models.llama3 import llama3_debugmodel
from torchtitan_recipes.tests.suites.features import (
    llama3_debugmodel_default,
    llama3_debugmodel_hf_checkpoint_load,
    muse_glimmer_debugmodel_fsdp2_per_group_cuda_graph,
)
from torchtitan_recipes.tests.suites.models import llama3_debugmodel_fsdp2_tp2_pp2

from tests.integration_tests import (
    get_importable_config_module,
    IntegrationTestDefinition,
    validate_fake_pg_compatibility,
)
from tests.integration_tests.b200 import build_b200_tests_list
from tests.integration_tests.features import build_features_test_list
from tests.integration_tests.flux import build_flux_test_list
from tests.integration_tests.h100 import build_h100_tests_list
from tests.integration_tests.models import build_model_tests_list
from tests.integration_tests.run_tests import _parse_test_suites, run_single_test


def test_hf_checkpoint_load_path_comes_from_test_config(monkeypatch) -> None:
    test_output_dir = "/tmp/model_only_hf_checkpoint"
    monkeypatch.setenv("TORCHTITAN_TEST_OUTPUT_DIR", test_output_dir)

    config = llama3_debugmodel_hf_checkpoint_load()

    assert config.checkpointer.initial_load_path == (
        f"{test_output_dir}/hf_checkpoint/step-10/"
    )


def test_spmd_typechecking_config_disables_local_compile() -> None:
    config = llama3_debugmodel_default()

    assert config.debug.spmd_typechecking
    assert config.model.local_compile_regions == []
    config.__post_init__()


def test_integration_run_exports_test_output_dir(monkeypatch, tmp_path: Path) -> None:
    captured_env = None

    def fake_run_cmd(cmd, timeout=None, env=None):
        nonlocal captured_env
        captured_env = env
        return subprocess.CompletedProcess(cmd, 0, stdout="")

    monkeypatch.setattr("tests.integration_tests.run_tests._run_cmd", fake_run_cmd)
    test = IntegrationTestDefinition(
        configs=[llama3_debugmodel],
        test_name="output_dir_test",
        ngpu=1,
    )

    run_single_test(test, str(tmp_path))

    assert captured_env is not None
    assert captured_env["TORCHTITAN_TEST_OUTPUT_DIR"] == str(
        tmp_path / "output_dir_test"
    )


def test_config_module_resolves_python_m_entrypoint(monkeypatch) -> None:
    def config_fn():
        return llama3_debugmodel()

    monkeypatch.setattr(config_fn, "__module__", "__main__")
    monkeypatch.setattr(
        sys.modules["__main__"],
        "__spec__",
        SimpleNamespace(name="tests.integration_tests.example"),
    )

    assert get_importable_config_module(config_fn) == "tests.integration_tests.example"


def test_numerics_run_uses_seed_config(monkeypatch, tmp_path: Path) -> None:
    captured_command = None

    def seed_config():
        return llama3_debugmodel()

    def fake_run(command, **kwargs):
        nonlocal captured_command
        captured_command = command
        return subprocess.CompletedProcess(command, 0, stdout="")

    golden_path = tmp_path / "golden.txt"
    golden_path.write_text("# step loss\n1 1.0\n")
    monkeypatch.setattr("tests.integration_tests.run_tests.subprocess.run", fake_run)
    test = IntegrationTestDefinition(
        configs=[llama3_debugmodel],
        test_name="seed_config_test",
        ngpu=1,
        golden_numerics_path=str(golden_path),
        loss_compare_seed_config=seed_config,
    )

    run_single_test(test, str(tmp_path))

    assert captured_command is not None
    assert f"--seed-module={seed_config.__module__}" in captured_command
    assert f"--seed-config={seed_config.__name__}" in captured_command


def test_llama3_pp_numerics_has_one_microbatch_per_stage() -> None:
    config = llama3_debugmodel_fsdp2_tp2_pp2()

    assert (
        config.parallelism.num_pp_microbatches
        >= config.parallelism.pipeline_parallel_degree
    )


def test_split_backward_pp_cases_exercise_varlen_cuda_graphs() -> None:
    tests_by_name = {test.test_name: test for test in build_features_test_list()}

    for test_name in ("pp_looped_zero_bubble", "pp_zbv", "pp_custom_csv"):
        test = tests_by_name[test_name]
        config = test.configs[0]()
        assert not test.disabled
        assert not config.training.disable_cuda_graphs
        assert isinstance(
            config.model.layers[0].attention.inner_attention,
            VarlenInnerAttention.Config,
        )


def test_per_group_cuda_graph_integration_supports_fake_pg() -> None:
    tests_by_name = {test.test_name: test for test in build_features_test_list()}
    test = tests_by_name["fsdp_per_group_cuda_graph"]
    config = test.configs[0]()

    assert test.configs == [muse_glimmer_debugmodel_fsdp2_per_group_cuda_graph]
    assert not test.use_real_pg
    assert test.ngpu == 2
    assert config.parallelism.data_parallel_shard_degree == 2
    assert config.training.cuda_graph_per_accumulation_group
    assert config.training.steps == 10
    assert config.training.num_tokens_per_train_step == (
        3 * test.ngpu * config.training.num_tokens_per_microbatch_per_dp_rank
    )


def test_llama3_debug_config_defaults_to_short_context() -> None:
    config = llama3_debugmodel()

    assert config.model.max_context_length == 2048
    assert config.training.max_context_length == 2048


def test_parse_multiple_integration_test_suites() -> None:
    assert _parse_test_suites("features,models,h100,b200") == (
        "features",
        "models",
        "h100",
        "b200",
    )


def test_h100_tests_are_registered_in_separate_suite() -> None:
    h100_tests = build_h100_tests_list()
    assert {test.test_name for test in h100_tests} == {
        "deepseek_v3_fsdp+hybridep",
        "dist_gemm",
        "fsdp_symm_mem",
        "kimi_k3_mm_allgather_kv_cp",
        "kimi_k3_mm_ulysses_cp",
        "qwen3_fsdp+deepep",
        "qwen3_5_moe_lora",
    }
    assert all(not hasattr(test, "use_h100") for test in build_features_test_list())
    assert all(not hasattr(test, "use_h100") for test in build_model_tests_list())


def test_b200_tests_are_registered_in_separate_suite() -> None:
    assert {test.test_name for test in build_b200_tests_list()} == {
        "kimi_k3_fsdp2_tp2_ep2_pp2_vpp4",
        "kimi_k3_mm",
        "kimi_k3_mm_muon",
        "dist_moe_eager_fsdp_ep_cudagraph",
        "dist_moe_eager_fsdp_ep_pp_cudagraph",
        "mxfp8_linear_fsdp",
        "nvfp4_linear_fsdp",
    }
    assert "kimi_k3_mm" not in {test.test_name for test in build_model_tests_list()}


def test_specialized_moe_backends_have_ep_coverage() -> None:
    specialized_names = {
        "deepseek_v3_fsdp+hybridep",
        "qwen3_fsdp+deepep",
    }
    h100_model_tests = [
        test for test in build_h100_tests_list() if test.test_name in specialized_names
    ]

    for test in h100_model_tests:
        config = test.configs[0]()
        assert config.parallelism.expert_parallel_degree > 1


def test_models_select_fake_and_real_pg_cases() -> None:
    model_tests = build_model_tests_list()
    fake_pg_model_tests = {
        test.test_name for test in model_tests if not test.use_real_pg
    }
    real_pg_model_tests = {test.test_name for test in model_tests if test.use_real_pg}

    assert {
        "deepseek_v3_fsdp+ep",
        "qwen3_moe_fsdp+tp+cp+ep_param_groups",
        "kimi_k2_5_muon_fsdp+ep",
        "muse_glimmer_text_fsdp",
        "muse_glimmer_mm_fsdp+tp+sp",
    } <= fake_pg_model_tests
    assert {
        "deepseek_v3_fsdp+cp+pp+ep",
        "deepseek_v4_fsdp+tp+ep",
    } <= real_pg_model_tests


def test_flux_fake_pg_filters_real_collective_cases() -> None:
    flux_tests = build_flux_test_list()
    fake_pg_tests = {test.test_name for test in flux_tests if not test.use_real_pg}

    assert fake_pg_tests == set()


@pytest.mark.parametrize(
    ("test_name", "incompatibility"),
    [
        ("checkpoint", "checkpointing"),
        ("pipeline_parallel", "pipeline parallelism"),
    ],
)
def test_fake_pg_incompatible_test_requires_explicit_marker(
    test_name: str, incompatibility: str
) -> None:
    config = llama3_debugmodel(seq_len=2048)
    if test_name == "checkpoint":
        config.checkpointer = CheckpointManager.Config()
    elif test_name == "pipeline_parallel":
        config.parallelism.pipeline_parallel_degree = 2

    test = IntegrationTestDefinition(configs=[llama3_debugmodel], test_name=test_name)

    with pytest.raises(ValueError, match=incompatibility):
        validate_fake_pg_compatibility(test, config)

    test.use_real_pg = True
    validate_fake_pg_compatibility(test, config)
