# 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 os
import pickle
import tempfile
import unittest
from dataclasses import dataclass, field
from types import SimpleNamespace
from unittest.mock import MagicMock, patch

import torch

from torchtitan.experiments.graph_trainer.configs import EpOverlapConfig
from torchtitan.experiments.graph_trainer.storage import DiskStorageAdapter


class TestDiskStorageAdapter(unittest.TestCase):
    def test_save_load_roundtrip(self):
        with tempfile.TemporaryDirectory() as tmpdir:
            storage = DiskStorageAdapter(tmpdir)
            data = pickle.dumps({"hello": "world", "values": [1, 2, 3]})

            path = storage.save("test_key", data)
            self.assertTrue(os.path.exists(path))
            self.assertTrue(storage.exists("test_key"))

            loaded = storage.load("test_key")
            self.assertEqual(data, loaded)

    def test_load_nonexistent_raises(self):
        with tempfile.TemporaryDirectory() as tmpdir:
            storage = DiskStorageAdapter(tmpdir)
            with self.assertRaises(FileNotFoundError):
                storage.load("nonexistent")

    def test_exists_false_for_missing(self):
        with tempfile.TemporaryDirectory() as tmpdir:
            storage = DiskStorageAdapter(tmpdir)
            self.assertFalse(storage.exists("missing"))

    def test_save_creates_subdirs(self):
        with tempfile.TemporaryDirectory() as tmpdir:
            nested = os.path.join(tmpdir, "a", "b", "c")
            storage = DiskStorageAdapter(nested)
            data = b"test"
            path = storage.save("key", data)
            self.assertTrue(os.path.exists(path))
            self.assertEqual(storage.load("key"), data)

    def test_delete_existing(self):
        with tempfile.TemporaryDirectory() as tmpdir:
            storage = DiskStorageAdapter(tmpdir)
            storage.save("key", b"data")
            self.assertTrue(storage.exists("key"))
            storage.delete("key")
            self.assertFalse(storage.exists("key"))

    def test_delete_nonexistent_noop(self):
        with tempfile.TemporaryDirectory() as tmpdir:
            storage = DiskStorageAdapter(tmpdir)
            storage.delete("nonexistent")

    def test_save_overwrites_existing(self):
        with tempfile.TemporaryDirectory() as tmpdir:
            storage = DiskStorageAdapter(tmpdir)
            storage.save("key", b"original")
            storage.save("key", b"updated")
            self.assertEqual(storage.load("key"), b"updated")

    def test_path_traversal_rejected(self):
        with tempfile.TemporaryDirectory() as tmpdir:
            storage = DiskStorageAdapter(tmpdir)
            with self.assertRaises(ValueError):
                storage.save("../../escape", b"data")


@dataclass
class _StubCompileConfig:
    passes: list = field(default_factory=list)
    memory_policy: str = "default"
    full_recompute_save_ops: str = ""
    ep_overlap: EpOverlapConfig = field(default_factory=EpOverlapConfig)


@dataclass
class _StubParallelismContext:
    world_size: int = 8
    dp_replicate: int = 1
    dp_shard: int = 2
    cp: int = 1
    tp: int = 2
    pp: int = 2
    ep: int = 1


def _make_stub_model(params=None, buffers=None):
    """
    Build a mock model with controlled named_parameters() and
    named_buffers() for deterministic fingerprint testing.
    """
    if params is None:
        params = [
            ("layer.weight", torch.zeros(4, 4)),
            ("layer.bias", torch.zeros(4)),
        ]
    if buffers is None:
        buffers = [("running_mean", torch.zeros(4))]

    model = MagicMock()
    # Use side_effect (not return_value) so each call produces a
    # fresh iterator — just like real nn.Module methods. A single
    # return_value=iter(...) would be exhausted after the first call.
    model.named_parameters.side_effect = lambda: iter(params)
    model.named_buffers.side_effect = lambda: iter(buffers)
    return model


class TestPrecompileMain(unittest.TestCase):
    def test_validates_memory_policy_after_model_setup(self):
        from torchtitan.experiments.graph_trainer import precompile_main

        events = []
        compile_config = SimpleNamespace()
        config = SimpleNamespace(compile=compile_config)
        config_loader = MagicMock()
        config_loader.load.return_value = config
        setup_result = (
            object(),
            object(),
            compile_config,
            object(),
            object(),
            object(),
        )

        def common_setup(_config):
            events.append("setup")
            return setup_result

        def validate(actual_compile_config):
            self.assertIs(actual_compile_config, compile_config)
            self.assertEqual(events, ["setup"])
            events.append("validate")

        def precompile(*_args):
            self.assertEqual(events, ["setup", "validate"])
            events.append("precompile")

        with (
            patch.object(precompile_main, "ConfigLoader", return_value=config_loader),
            patch.object(precompile_main, "_common_setup", side_effect=common_setup),
            patch.object(
                precompile_main,
                "validate_memory_policy_config",
                side_effect=validate,
            ),
            patch.object(
                precompile_main,
                "_precompile_aot_fx_trace",
                side_effect=precompile,
            ),
            patch.object(precompile_main.dist, "destroy_process_group"),
        ):
            precompile_main.main()

        self.assertEqual(events, ["setup", "validate", "precompile"])


class TestConfigFingerprint(unittest.TestCase):
    def test_deterministic(self):
        from torchtitan.experiments.graph_trainer.precompile import (
            compute_config_fingerprint,
        )

        cfg = _StubCompileConfig()
        dims = _StubParallelismContext()

        fp1 = compute_config_fingerprint(_make_stub_model(), cfg, dims)
        fp2 = compute_config_fingerprint(_make_stub_model(), cfg, dims)
        self.assertEqual(fp1, fp2)
        self.assertEqual(len(fp1), 16)

    def test_memory_policy_save_ops_sensitivity(self):
        from torchtitan.experiments.graph_trainer.precompile import (
            compute_config_fingerprint,
        )

        dims = _StubParallelismContext()
        cfg_a = _StubCompileConfig(memory_policy="full")
        cfg_b = _StubCompileConfig(
            memory_policy="full",
            full_recompute_save_ops=("layers.*.moe.router.gate :: aten.mm.dtype"),
        )

        fp_a = compute_config_fingerprint(_make_stub_model(), cfg_a, dims)
        fp_b = compute_config_fingerprint(_make_stub_model(), cfg_b, dims)
        self.assertNotEqual(fp_a, fp_b)

    def test_model_shape_sensitivity(self):
        from torchtitan.experiments.graph_trainer.precompile import (
            compute_config_fingerprint,
        )

        cfg = _StubCompileConfig()
        dims = _StubParallelismContext()

        model_a = _make_stub_model(params=[("w", torch.zeros(4, 4))], buffers=[])
        model_b = _make_stub_model(params=[("w", torch.zeros(8, 8))], buffers=[])
        fp_a = compute_config_fingerprint(model_a, cfg, dims)
        fp_b = compute_config_fingerprint(model_b, cfg, dims)
        self.assertNotEqual(fp_a, fp_b)

    def test_parallelism_sensitivity(self):
        from torchtitan.experiments.graph_trainer.precompile import (
            compute_config_fingerprint,
        )

        cfg = _StubCompileConfig()
        model = _make_stub_model()

        dims_tp2 = _StubParallelismContext(tp=2)
        dims_tp4 = _StubParallelismContext(tp=4)
        fp_tp2 = compute_config_fingerprint(model, cfg, dims_tp2)
        fp_tp4 = compute_config_fingerprint(_make_stub_model(), cfg, dims_tp4)
        self.assertNotEqual(fp_tp2, fp_tp4)

    def test_compile_config_sensitivity(self):
        from torchtitan.experiments.graph_trainer.precompile import (
            compute_config_fingerprint,
        )

        dims = _StubParallelismContext()

        cfg_a = _StubCompileConfig(passes=["pass_a"])
        cfg_b = _StubCompileConfig(passes=["pass_a", "pass_b"])
        fp_a = compute_config_fingerprint(_make_stub_model(), cfg_a, dims)
        fp_b = compute_config_fingerprint(_make_stub_model(), cfg_b, dims)
        self.assertNotEqual(fp_a, fp_b)

        cfg_graph_batch = _StubCompileConfig(ep_overlap=EpOverlapConfig(enabled=True))
        cfg_graph_seq = _StubCompileConfig(
            ep_overlap=EpOverlapConfig(
                enabled=True,
                chunk_dim="seq",
                module_fqn="layers.*.moe",
            ),
        )
        fp_graph_batch = compute_config_fingerprint(
            _make_stub_model(), cfg_graph_batch, dims
        )
        fp_graph_seq = compute_config_fingerprint(
            _make_stub_model(), cfg_graph_seq, dims
        )
        self.assertNotEqual(fp_graph_batch, fp_graph_seq)

    def test_pass_order_sensitive(self):
        from torchtitan.experiments.graph_trainer.precompile import (
            compute_config_fingerprint,
        )

        dims = _StubParallelismContext()

        cfg_ab = _StubCompileConfig(passes=["a", "b"])
        cfg_ba = _StubCompileConfig(passes=["b", "a"])
        fp_ab = compute_config_fingerprint(_make_stub_model(), cfg_ab, dims)
        fp_ba = compute_config_fingerprint(_make_stub_model(), cfg_ba, dims)
        self.assertNotEqual(fp_ab, fp_ba)


class TestPrecompileLossSetup(unittest.TestCase):
    def test_chunked_loss_setup_matches_trainer_boundary(self):
        from torchtitan.experiments.graph_trainer.chunked_loss import (
            ChunkedLossWrapperWithParamGrads,
        )
        from torchtitan.experiments.graph_trainer.precompile_main import (
            _prepare_loss_for_precompile,
        )

        lm_head = torch.nn.Linear(2, 3)
        model = SimpleNamespace(lm_head=lm_head, _skip_lm_head=False)
        loss_fn = ChunkedLossWrapperWithParamGrads.Config().build()

        _prepare_loss_for_precompile(model, loss_fn)

        self.assertIs(loss_fn.lm_head, lm_head)
        self.assertTrue(model._skip_lm_head)


class TestPrecompiledFxTraceArtifact(unittest.TestCase):
    def test_loaded_artifact_supports_traced_execution(self):
        from torchtitan.experiments.graph_trainer.make_fx_tracer import (
            minimal_fx_tracer,
            run_traced,
        )
        from torchtitan.experiments.graph_trainer.precompile import (
            flatten_runtime_inputs,
            PrecompiledFxTraceArtifact,
        )

        model = torch.nn.Linear(3, 2, dtype=torch.float64)
        inputs = torch.randn(4, 3, dtype=torch.float64)

        def forward(value):
            return model(value)

        traced = minimal_fx_tracer(forward, module=model)(inputs)
        example_inputs = flatten_runtime_inputs(model, (inputs,), {})
        loaded = PrecompiledFxTraceArtifact.from_traced_result(traced).to_traced_result(
            example_inputs
        )
        run = run_traced(loaded, module=model)

        self.assertTrue(torch.equal(model(inputs), run(inputs)))

    def test_rejects_trainer_owned_gradient_state(self):
        from torchtitan.experiments.graph_trainer.make_fx_tracer import (
            minimal_fx_tracer,
        )
        from torchtitan.experiments.graph_trainer.precompile import (
            PrecompiledFxTraceArtifact,
        )

        inputs = torch.randn(2, 3)
        traced = minimal_fx_tracer(
            lambda _state, value: value.sum(),
            graph_state={"accumulator": torch.zeros_like(inputs)},
        )(inputs)

        with self.assertRaisesRegex(ValueError, "trainer-owned gradient state"):
            PrecompiledFxTraceArtifact.from_traced_result(traced)

    @unittest.skipUnless(torch.cuda.is_available(), "CUDA required")
    def test_standalone_inductor_precompile(self):
        from torchtitan.experiments.graph_trainer.inductor_passes import (
            standalone_inductor_compilation_pass,
        )
        from torchtitan.experiments.graph_trainer.make_fx_tracer import (
            minimal_fx_tracer,
            run_traced,
        )
        from torchtitan.experiments.graph_trainer.precompile import (
            flatten_runtime_inputs,
            precompile_fx_trace_load,
            precompile_fx_trace_save,
        )

        model = torch.nn.Sequential(
            torch.nn.RMSNorm(8),
            torch.nn.Linear(8, 4),
        ).cuda()

        def train_step(x, unused):
            loss = model(x).square().sum()
            return loss, *torch.autograd.grad(loss, tuple(model.parameters()))

        x = torch.randn(2, 8, device="cuda")
        unused = torch.randn(1, device="cuda")
        expected = train_step(x, unused)
        traced = minimal_fx_tracer(train_step, module=model)(x, unused)
        traced.gm = standalone_inductor_compilation_pass(
            traced.gm, traced.example_inputs
        )

        with tempfile.TemporaryDirectory() as tmpdir:
            storage = DiskStorageAdapter(tmpdir)
            precompile_fx_trace_save(traced, storage)
            example_inputs = flatten_runtime_inputs(model, (x, unused), {})
            loaded = precompile_fx_trace_load(
                storage,
                expected_fingerprint="",
                example_inputs=example_inputs,
            )

        self.assertEqual(len(loaded.example_inputs), len(example_inputs))
        self.assertTrue(
            all(
                loaded_input is example_input
                for loaded_input, example_input in zip(
                    loaded.example_inputs, example_inputs, strict=True
                )
            )
        )

        with patch(
            "torch._inductor.standalone_compile",
            side_effect=AssertionError("precompiled graph must not compile again"),
        ):
            for _ in range(2):
                actual = run_traced(loaded, module=model)(x, unused)
                for actual_tensor, expected_tensor in zip(
                    actual, expected, strict=True
                ):
                    torch.testing.assert_close(actual_tensor, expected_tensor)

    def test_artifact_pickle_roundtrip(self):
        from torchtitan.experiments.graph_trainer.make_fx_tracer import SubclassLayout
        from torchtitan.experiments.graph_trainer.precompile import (
            PrecompiledFxTraceArtifact,
        )

        flat_vals, spec = torch.utils._pytree.tree_flatten({"a": torch.zeros(2)})
        artifact = PrecompiledFxTraceArtifact(
            serialized_gm=b"fake_serialized_data",
            state_fqns=["w1", "w2"],
            num_flat_inputs=4,
            input_subclass_layouts={
                0: SubclassLayout(1, None),
                1: SubclassLayout(1, None),
            },
            num_flat_outputs=2,
            output_subclass_layouts={0: SubclassLayout(1, None)},
            output_spec=spec,
            tensor_input_indices=[0, 1, 2, 3],
            config_fingerprint="test_fp_123",
        )

        data = pickle.dumps(artifact)
        loaded = pickle.loads(data)

        self.assertEqual(loaded.serialized_gm, artifact.serialized_gm)
        self.assertEqual(len(loaded.state_fqns), 2)
        self.assertEqual(loaded.num_flat_inputs, 4)
        self.assertEqual(len(loaded.input_subclass_layouts), 2)
        self.assertEqual(loaded.num_flat_outputs, 2)
        self.assertEqual(loaded.config_fingerprint, "test_fp_123")
        self.assertEqual(loaded.num_optimizer_state_inputs, 0)
        self.assertEqual(loaded.num_runtime_mesh_inputs, 0)

    def test_artifact_pickle_with_blockmask_treespec(self):
        """Verify artifact pickles when user_inputs_spec contains BlockMask.

        BlockMask's pytree context stores a _MaskModWrapper holding the
        mask_mod closure, which is not picklable. The artifact must not
        serialize user_inputs_spec (the mask_mod is already compiled into
        standalone Inductor HOPs baked into serialized_gm).
        """
        from torch.nn.attention.flex_attention import create_block_mask

        from torchtitan.experiments.graph_trainer.common_utils import (
            maybe_register_blockmask_pytree_node,
        )
        from torchtitan.experiments.graph_trainer.make_fx_tracer import TracedResult
        from torchtitan.experiments.graph_trainer.precompile import (
            PrecompiledFxTraceArtifact,
        )
        from torchtitan.models.common.attention import get_causal_mask_mod

        maybe_register_blockmask_pytree_node()

        mask_mod = get_causal_mask_mod()
        block_mask = create_block_mask(mask_mod, B=1, H=1, Q_LEN=128, KV_LEN=128)

        # Build a user_inputs_spec that includes BlockMask — this is what
        # minimal_fx_tracer produces when FlexInnerAttention is configured.
        _, blockmask_spec = torch.utils._pytree.tree_flatten(
            ((torch.zeros(2),), {"attention_metadata": block_mask})
        )

        # Sanity: the raw TreeSpec itself is NOT picklable (the bug).
        with self.assertRaises((TypeError, AttributeError)):
            pickle.dumps(blockmask_spec)

        # Build a TracedResult with the unpicklable spec, then create
        # the artifact via from_traced_result — this must succeed because
        # user_inputs_spec is excluded from serialization.
        gm = torch.fx.GraphModule(torch.nn.Module(), torch.fx.Graph())
        dummy_spec = torch.utils._pytree.tree_flatten(((), {}))[1]
        traced_result = TracedResult(
            gm=gm,
            example_inputs=(),
            num_flat_inputs=0,
            input_subclass_layouts={},
            user_inputs_spec=blockmask_spec,
            tensor_input_indices=[],
            num_flat_outputs=0,
            output_subclass_layouts={},
            output_spec=dummy_spec,
            state_fqns=[],
        )

        artifact = PrecompiledFxTraceArtifact.from_traced_result(traced_result)
        data = pickle.dumps(artifact)
        loaded = pickle.loads(data)
        self.assertEqual(loaded.serialized_gm, artifact.serialized_gm)

    def test_fx_trace_save_load_fingerprint_mismatch(self):
        from torchtitan.experiments.graph_trainer.precompile import (
            _FX_TRACE_ARTIFACT_KEY,
            precompile_fx_trace_load,
            PrecompiledFxTraceArtifact,
        )

        flat_vals, spec = torch.utils._pytree.tree_flatten({"a": torch.zeros(2)})
        artifact = PrecompiledFxTraceArtifact(
            serialized_gm=b"fake",
            state_fqns=["w"],
            num_flat_inputs=2,
            input_subclass_layouts={},
            num_flat_outputs=1,
            output_subclass_layouts={},
            output_spec=spec,
            tensor_input_indices=[0, 1],
            config_fingerprint="old_fp",
        )
        with tempfile.TemporaryDirectory() as tmpdir:
            storage = DiskStorageAdapter(tmpdir)
            storage.save(_FX_TRACE_ARTIFACT_KEY, pickle.dumps(artifact))

            with self.assertRaisesRegex(ValueError, "fingerprint mismatch"):
                precompile_fx_trace_load(
                    storage,
                    expected_fingerprint="new_fp",
                    example_inputs=(),
                )


class TestCudaGraphPass(unittest.TestCase):
    """Test cuda_graph_pass behavior."""

    def test_non_graphmodule_raises(self):
        """cuda_graph_pass rejects non-GraphModule callables (e.g.
        OutputCode from full_inductor_compilation)."""
        from torchtitan.experiments.graph_trainer.passes import cuda_graph_pass

        def plain_fn(*args):
            return args

        with self.assertRaisesRegex(TypeError, "requires a GraphModule"):
            cuda_graph_pass(plain_fn, (torch.zeros(4),), static_input_indices=[0])

    def test_graphmodule_wraps_forward(self):
        """cuda_graph_pass wraps gm.forward with CUDAGraphWrapper."""
        from torchtitan.experiments.graph_trainer.passes import cuda_graph_pass

        gm = torch.fx.GraphModule(torch.nn.Module(), torch.fx.Graph())
        example_inputs = (torch.zeros(4),)

        with patch(
            "torchtitan.experiments.graph_trainer.cuda_graph.CUDAGraphWrapper"
        ) as MockWrapper:
            mock_instance = MagicMock()
            MockWrapper.return_value = mock_instance
            result = cuda_graph_pass(gm, example_inputs, static_input_indices=[0])
            self.assertIs(result, gm)
            self.assertIs(gm.forward, mock_instance)


class TestCudaGraphFingerprintConsistency(unittest.TestCase):
    """Test that save and load paths produce the same fingerprint.

    Both paths compute the fingerprint from the original (unmodified)
    compile_config — CUDA graph stripping in precompile_main happens
    AFTER fingerprint computation, so no manual filtering is needed.
    """

    def test_cuda_graph_included_in_fingerprint(self):
        """CUDA graph in passes should produce a different fingerprint
        than without CUDA graph — no filtering is applied."""
        from torchtitan.experiments.graph_trainer.precompile import (
            compute_config_fingerprint,
        )

        dims = _StubParallelismContext()

        cfg_with = _StubCompileConfig(
            passes=["full_inductor_compilation", "cuda_graph"]
        )
        cfg_without = _StubCompileConfig(passes=["full_inductor_compilation"])

        fp_with = compute_config_fingerprint(_make_stub_model(), cfg_with, dims)
        fp_without = compute_config_fingerprint(_make_stub_model(), cfg_without, dims)

        self.assertNotEqual(fp_with, fp_without)

    def test_same_config_produces_same_fingerprint(self):
        """Both save and load paths use the same unmodified config,
        so the fingerprint is identical."""
        from torchtitan.experiments.graph_trainer.precompile import (
            compute_config_fingerprint,
        )

        dims = _StubParallelismContext()

        cfg = _StubCompileConfig(passes=["full_inductor_compilation", "cuda_graph"])

        fp1 = compute_config_fingerprint(_make_stub_model(), cfg, dims)
        fp2 = compute_config_fingerprint(_make_stub_model(), cfg, dims)

        self.assertEqual(fp1, fp2)


if __name__ == "__main__":
    unittest.main()
