# 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 contextlib
from types import SimpleNamespace
from typing import cast
from unittest.mock import MagicMock, patch

import pytest
import torch
from torch.distributed.device_mesh import DeviceMesh
from torch.utils.checkpoint import checkpoint

from torchtitan.config import CommConfig
from torchtitan.distributed import DistributedTopology, utils as dist_utils
from torchtitan.distributed.parallelism_context import ParallelismContext
from torchtitan.distributed.spmd_types import set_spmd_meshes, spmd_dense_sp_enabled
from torchtitan.distributed.utils import init_distributed


@pytest.mark.parametrize(
    ("pipeline_parallel_degree", "expected"), [(1, False), (2, True)]
)
def test_init_distributed_configures_pipeline_per_edge_p2p(
    monkeypatch: pytest.MonkeyPatch,
    pipeline_parallel_degree: int,
    expected: bool,
) -> None:
    monkeypatch.setenv("NGPU", "2")
    monkeypatch.setenv("FAKE_PP_RANK", "0")
    with (
        patch("torch.distributed.is_initialized", return_value=False),
        patch("torchtitan.distributed.utils.init_fake_mode"),
        patch.object(
            dist_utils,
            "dist_config",
            SimpleNamespace(pipeline_per_edge_p2p=not expected),
        ),
    ):
        init_distributed(
            CommConfig(backend="fake"),
            pipeline_parallel_degree=pipeline_parallel_degree,
        )
        assert dist_utils.dist_config.pipeline_per_edge_p2p is expected


def test_init_distributed_allows_missing_pipeline_per_edge_p2p(
    monkeypatch: pytest.MonkeyPatch,
) -> None:
    monkeypatch.setenv("NGPU", "1")
    config_without_option = SimpleNamespace()
    with (
        patch("torch.distributed.is_initialized", return_value=False),
        patch("torchtitan.distributed.utils.init_fake_mode"),
        patch.object(dist_utils, "dist_config", config_without_option),
    ):
        init_distributed(CommConfig(backend="fake"), pipeline_parallel_degree=1)

    assert not hasattr(config_without_option, "pipeline_per_edge_p2p")


def test_init_distributed_requires_pipeline_per_edge_p2p_for_pp(
    monkeypatch: pytest.MonkeyPatch,
) -> None:
    monkeypatch.setenv("NGPU", "2")
    with (
        patch("torch.distributed.is_initialized", return_value=False),
        patch.object(dist_utils, "dist_config", SimpleNamespace()),
        pytest.raises(RuntimeError, match="pipeline_per_edge_p2p"),
    ):
        init_distributed(CommConfig(backend="fake"), pipeline_parallel_degree=2)


def test_fake_pg_defaults_to_spmd_rank_zero(monkeypatch: pytest.MonkeyPatch) -> None:
    monkeypatch.setenv("NGPU", "8")
    with (
        patch("torch.distributed.is_initialized", return_value=False),
        patch("torchtitan.distributed.utils.init_fake_mode") as init_fake_mode,
    ):
        topology = init_distributed(CommConfig(backend="fake"))
    assert topology == DistributedTopology(world_size=8)
    init_fake_mode.assert_called_once_with(8, rank=0)


def test_fake_pg_rejects_out_of_range_rank(
    monkeypatch: pytest.MonkeyPatch,
) -> None:
    monkeypatch.setenv("NGPU", "8")
    monkeypatch.setenv("FAKE_PP_RANK", "4")
    with (
        patch("torch.distributed.is_initialized", return_value=False),
        patch.object(
            dist_utils,
            "dist_config",
            SimpleNamespace(pipeline_per_edge_p2p=False),
        ),
        pytest.raises(ValueError, match=r"FAKE_PP_RANK must be in \[0, 4\)"),
    ):
        init_distributed(CommConfig(backend="fake"), pipeline_parallel_degree=4)


def test_fake_pp_uses_explicit_pipeline_coordinate(
    monkeypatch: pytest.MonkeyPatch,
) -> None:
    monkeypatch.setenv("NGPU", "16")
    monkeypatch.setenv("FAKE_PP_RANK", "2")
    with (
        patch("torch.distributed.is_initialized", return_value=False),
        patch("torchtitan.distributed.utils.init_fake_mode") as init_fake_mode,
        patch.object(
            dist_utils,
            "dist_config",
            SimpleNamespace(pipeline_per_edge_p2p=False),
        ),
    ):
        topology = init_distributed(
            CommConfig(backend="fake"), pipeline_parallel_degree=4
        )

    assert topology == DistributedTopology(world_size=16)
    init_fake_mode.assert_called_once_with(16, rank=8)


def test_fake_pp_requires_pipeline_coordinate(
    monkeypatch: pytest.MonkeyPatch,
) -> None:
    monkeypatch.setenv("NGPU", "16")
    monkeypatch.delenv("FAKE_PP_RANK", raising=False)

    with (
        patch("torch.distributed.is_initialized", return_value=False),
        patch.object(
            dist_utils,
            "dist_config",
            SimpleNamespace(pipeline_per_edge_p2p=False),
        ),
        pytest.raises(ValueError, match="FAKE_PP_RANK environment variable"),
    ):
        init_distributed(CommConfig(backend="fake"), pipeline_parallel_degree=4)


def test_real_pp_fake_spmd_returns_real_pp_group(
    monkeypatch: pytest.MonkeyPatch,
) -> None:
    monkeypatch.setenv("NGPU", "16")
    monkeypatch.setenv("WORLD_SIZE", "4")
    monkeypatch.setenv("RANK", "2")
    monkeypatch.setenv("LOCAL_RANK", "0")
    process_group = torch.distributed.ProcessGroup(2, 4)
    store = MagicMock()

    with (
        patch("torch.distributed.is_initialized", return_value=False),
        patch.object(dist_utils, "init_fake_mode") as init_fake_mode,
        patch.object(
            dist_utils.dist,
            "rendezvous",
            return_value=iter([(store, 2, 4)]),
        ),
        patch.object(
            dist_utils.c10d,
            "_new_process_group_helper",
            return_value=(process_group, store),
        ) as new_process_group,
        patch.object(
            dist_utils,
            "dist_config",
            SimpleNamespace(pipeline_per_edge_p2p=False),
        ),
        patch.dict(dist_utils.c10d._world.pg_group_ranks, {}, clear=False),
    ):
        topology = init_distributed(
            CommConfig(backend="real_pp_fake_spmd"),
            pipeline_parallel_degree=4,
        )

    init_fake_mode.assert_called_once_with(16, rank=8)
    assert topology.world_size == 16
    assert topology.real_pp_group_for_fake_spmd is process_group
    assert new_process_group.call_args.kwargs["global_ranks_in_group"] == [0, 4, 8, 12]


def test_real_pp_fake_spmd_requires_one_process_per_pp_rank(
    monkeypatch: pytest.MonkeyPatch,
) -> None:
    monkeypatch.setenv("NGPU", "16")
    monkeypatch.setenv("WORLD_SIZE", "4")
    monkeypatch.setenv("RANK", "2")
    with (
        patch("torch.distributed.is_initialized", return_value=False),
        patch.object(
            dist_utils,
            "dist_config",
            SimpleNamespace(pipeline_per_edge_p2p=False),
        ),
        pytest.raises(ValueError, match="one physical process per PP rank"),
    ):
        init_distributed(
            CommConfig(backend="real_pp_fake_spmd"),
            pipeline_parallel_degree=2,
        )


def test_real_pp_fake_spmd_rejects_fake_pp_rank(
    monkeypatch: pytest.MonkeyPatch,
) -> None:
    monkeypatch.setenv("NGPU", "16")
    monkeypatch.setenv("WORLD_SIZE", "4")
    monkeypatch.setenv("RANK", "2")
    monkeypatch.setenv("FAKE_PP_RANK", "2")
    with (
        patch("torch.distributed.is_initialized", return_value=False),
        patch.object(
            dist_utils,
            "dist_config",
            SimpleNamespace(pipeline_per_edge_p2p=False),
        ),
        pytest.raises(ValueError, match="FAKE_PP_RANK is invalid"),
    ):
        init_distributed(
            CommConfig(backend="real_pp_fake_spmd"),
            pipeline_parallel_degree=4,
        )


def test_dist_sum_tensor_keeps_local_result_as_tensor():
    value = torch.tensor(3, dtype=torch.int64)

    result = dist_utils.dist_sum_tensor(value)

    assert result is value


def test_dist_sum_tensor_waits_for_distributed_result():
    value = torch.tensor(3, dtype=torch.int64)
    reduced = torch.tensor(8, dtype=torch.int64)
    mesh = cast(DeviceMesh, object())

    with (
        patch.object(dist_utils.funcol, "all_reduce", return_value=reduced) as reduce,
        patch.object(dist_utils.funcol, "wait_tensor", return_value=reduced) as wait,
    ):
        result = dist_utils.dist_sum_tensor(value, mesh)

    assert result is reduced
    reduce.assert_called_once_with(value, reduceOp="SUM", group=mesh)
    wait.assert_called_once_with(reduced)


@pytest.mark.parametrize("enable_sequence_parallel", [False, True])
def test_activate_spmd_exposes_dense_sp_state(
    enable_sequence_parallel: bool,
) -> None:
    dense_mesh = cast(DeviceMesh, object())
    parallelism_context = ParallelismContext(
        dp_replicate=1,
        dp_shard=1,
        cp=1,
        tp=2,
        pp=1,
        ep=1,
        world_size=2,
        enable_sequence_parallel=enable_sequence_parallel,
    )
    parallelism_context._single_axis_meshes["tp"] = dense_mesh

    with (
        patch.object(parallelism_context, "spmd_dense_mesh", return_value=dense_mesh),
        patch.object(parallelism_context, "spmd_sparse_mesh", return_value=None),
        patch(
            "torchtitan.distributed.spmd_types.set_current_spmd_mesh",
            return_value=contextlib.nullcontext(),
        ),
        patch(
            "torchtitan.distributed.spmd_types.spmd_dense_mesh",
            return_value=dense_mesh,
        ),
        patch(
            "spmd_types.checker.typecheck",
            return_value=contextlib.nullcontext(),
        ) as typecheck,
        parallelism_context.activate_spmd(typechecking=True),
    ):
        assert spmd_dense_sp_enabled() is enable_sequence_parallel
    typecheck.assert_called_once_with(local=False)


def test_dense_sp_state_compiles_with_checkpoint() -> None:
    dense_mesh = cast(DeviceMesh, object())
    set_spmd_meshes(
        dense_mesh=dense_mesh,
        sparse_mesh=None,
        dense_sp_enabled=True,
    )

    def checkpointed_forward(input):
        def forward(value):
            assert spmd_dense_sp_enabled()
            return value + 1

        return checkpoint(forward, input, use_reentrant=False)

    compiled_forward = torch.compile(
        checkpointed_forward,
        backend="eager",
        fullgraph=True,
    )
    input = torch.randn(2, 3, requires_grad=True)

    output = compiled_forward(input)
    output.sum().backward()

    torch.testing.assert_close(output, input + 1)
    set_spmd_meshes(
        dense_mesh=dense_mesh,
        sparse_mesh=None,
        dense_sp_enabled=False,
    )
