# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

import json
import logging
import os
from dataclasses import MISSING, Field, asdict, dataclass, field
from datetime import timedelta
from pathlib import Path
from types import SimpleNamespace
from typing import cast
from unittest.mock import patch

import pydantic
import pytest
import torch
from huggingface_hub import ResolvedRevision
from pydantic import ValidationError

import vllm.config.vllm as vllm_config_module
import vllm.envs as envs
from vllm.compilation.backends import VllmBackend
from vllm.config import (
    AttentionConfig,
    CacheConfig,
    CompilationConfig,
    DeviceConfig,
    EngramConfig,
    HiSparseConfig,
    KernelConfig,
    KVTransferConfig,
    ModelConfig,
    ObservabilityConfig,
    ParallelConfig,
    PoolerConfig,
    SchedulerConfig,
    SpeculativeConfig,
    VllmConfig,
    WatermarkConfig,
    update_config,
)
from vllm.config.compilation import CompilationMode, CUDAGraphMode, PassConfig
from vllm.config.kernel import IrOpPriorityConfig
from vllm.config.load import LoadConfig
from vllm.config.mamba import MambaBackendEnum
from vllm.config.speculative import _validate_qwen3_omni_dspark
from vllm.config.utils import get_field
from vllm.config.vllm import OPTIMIZATION_LEVEL_TO_CONFIG, OptimizationLevel
from vllm.platforms import current_platform
from vllm.transformers_utils.config import (
    _patch_hf_transformers_nested_rope_validation,
    get_pooling_config,
    try_get_dense_modules,
)
from vllm.v1.attention.backend import AttentionCGSupport

DEVICE_TYPE = current_platform.device_type


def test_nested_rope_validation_patch_preserves_flat_rope_parameters(monkeypatch):
    calls = []

    def original_validate_rope(config, *args, **kwargs):
        calls.append(config)

    from transformers import PreTrainedConfig

    monkeypatch.setattr(PreTrainedConfig, "validate_rope", original_validate_rope)
    _patch_hf_transformers_nested_rope_validation()

    nested_rope_parameters = {
        "full_attention": {"rope_type": "default"},
        "original_max_position_embeddings": 32768,
    }
    PreTrainedConfig.validate_rope(
        SimpleNamespace(rope_parameters=nested_rope_parameters)
    )
    assert nested_rope_parameters == {"full_attention": {"rope_type": "default"}}

    flat_rope_parameters = {
        "rope_type": "linear",
        "factor": 8.0,
        "rope_theta": 500000.0,
    }
    PreTrainedConfig.validate_rope(
        SimpleNamespace(rope_parameters=flat_rope_parameters)
    )
    assert flat_rope_parameters == {
        "rope_type": "linear",
        "factor": 8.0,
        "rope_theta": 500000.0,
    }
    assert len(calls) == 2


def test_dspark_adaptive_verification_separates_graph_cache():
    config = object.__new__(SpeculativeConfig)
    config.method = "dspark"
    config.draft_model_config = None
    config.enable_adaptive_verification = False
    fixed_hash = config.compute_hash()
    config.enable_adaptive_verification = True
    assert config.compute_hash() != fixed_hash


def _write_json(path: Path, value: object) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    path.write_text(json.dumps(value), encoding="utf-8")


@pytest.mark.skipif(not current_platform.is_rocm(), reason="ROCm-specific test")
@pytest.mark.parametrize(
    ("is_mm_prefix_lm", "is_multimodal_model", "expected"),
    [
        pytest.param(True, True, True, id="multimodal-prefix-lm"),
        pytest.param(False, True, False, id="multimodal-causal"),
        pytest.param(True, False, False, id="text-prefix-lm"),
        pytest.param(None, True, False, id="missing-model-config"),
    ],
)
def test_rocm_mm_prefix_lm_disables_chunked_mm_input(
    is_mm_prefix_lm: bool | None,
    is_multimodal_model: bool,
    expected: bool,
) -> None:
    from vllm.platforms.rocm import RocmPlatform

    config = SimpleNamespace(
        compilation_config=SimpleNamespace(cudagraph_mode=CUDAGraphMode.NONE),
        parallel_config=SimpleNamespace(
            prefill_context_parallel_size=1,
            worker_cls="test-worker",
        ),
        model_config=(
            None
            if is_mm_prefix_lm is None
            else SimpleNamespace(is_mm_prefix_lm=is_mm_prefix_lm)
        ),
        scheduler_config=SimpleNamespace(
            is_multimodal_model=is_multimodal_model,
            disable_chunked_mm_input=False,
        ),
    )

    RocmPlatform.check_and_update_config(config)

    assert config.scheduler_config.disable_chunked_mm_input is expected


def _sampling_replay_config(
    *,
    return_sampling_mask: bool = True,
    use_v2_model_runner: bool = True,
    speculative_method: str | None = None,
    rejection_sample_method: str = "standard",
    adaptive: bool = False,
    is_diffusion: bool = False,
    logits_processors: list[str] | None = None,
    logprobs_mode: str = "processed_logprobs",
):
    speculative_config = None
    if speculative_method is not None:
        speculative_config = SimpleNamespace(
            method=speculative_method,
            enable_adaptive_verification=adaptive,
            rejection_sample_method=rejection_sample_method,
        )
    return SimpleNamespace(
        model_config=SimpleNamespace(
            return_sampling_mask=return_sampling_mask,
            is_diffusion=is_diffusion,
            logits_processors=logits_processors or [],
            logprobs_mode=logprobs_mode,
        ),
        use_v2_model_runner=use_v2_model_runner,
        speculative_config=speculative_config,
    )


@pytest.mark.parametrize(
    ("config", "message"),
    [
        (_sampling_replay_config(), None),
        *(
            (_sampling_replay_config(speculative_method=method), None)
            for method in (
                "mtp",
                "eagle",
                "eagle3",
                "dflash",
                "dspark",
                "draft_model",
            )
        ),
        *(
            (
                _sampling_replay_config(
                    speculative_method="mtp",
                    rejection_sample_method=rejection_sample_method,
                ),
                None,
            )
            for rejection_sample_method in ("standard", "block", "synthetic")
        ),
        (
            _sampling_replay_config(
                return_sampling_mask=False, speculative_method="dspark", adaptive=True
            ),
            None,
        ),
        (
            _sampling_replay_config(speculative_method="dspark", adaptive=True),
            "requires fixed verification boundaries",
        ),
        (
            _sampling_replay_config(is_diffusion=True),
            "does not support diffusion models",
        ),
        (
            _sampling_replay_config(logits_processors=["custom"]),
            "does not support custom logits processors",
        ),
        (
            _sampling_replay_config(logprobs_mode="raw_logprobs"),
            "requires logprobs_mode='processed_logprobs'",
        ),
        (
            _sampling_replay_config(use_v2_model_runner=False),
            "requires Model Runner V2",
        ),
    ],
)
def test_sampling_replay_config(config, message):
    if message is None:
        VllmConfig._verify_sampling_replay_config(config)
    else:
        with pytest.raises(ValueError, match=message):
            VllmConfig._verify_sampling_replay_config(config)


def test_kda_recoverssm_derivation_is_revalidated():
    config = SimpleNamespace(
        cache_config=SimpleNamespace(
            use_replayssm=True,
            use_kda_recoverssm=False,
            mamba_cache_mode="none",
            replayssm_buffer_len=16,
        ),
        num_speculative_tokens=3,
        model_config=SimpleNamespace(
            supports_replayssm=True,
            architecture="KimiLinearForCausalLM",
        ),
        mamba_config=SimpleNamespace(
            backend=MambaBackendEnum.TRITON,
            enable_stochastic_rounding=False,
        ),
        parallel_config=SimpleNamespace(pipeline_parallel_size=1),
        kv_transfer_config=None,
        use_v2_model_runner=True,
    )

    VllmConfig.validate_mamba_cached_kernel(config)
    assert config.cache_config.use_replayssm
    assert config.cache_config.use_kda_recoverssm

    config.cache_config.mamba_cache_mode = "align"
    VllmConfig.validate_mamba_cached_kernel(config)
    config.use_v2_model_runner = False
    with pytest.raises(ValueError, match="VLLM_USE_V2_MODEL_RUNNER=1"):
        VllmConfig.validate_mamba_cached_kernel(config)
    config.use_v2_model_runner = True
    config.cache_config.mamba_cache_mode = "none"

    config.model_config.architecture = "NemotronHForCausalLM"
    config.mamba_config.backend = MambaBackendEnum.FLASHINFER
    VllmConfig.validate_mamba_cached_kernel(config)
    assert not config.cache_config.use_kda_recoverssm

    config.model_config.architecture = "KimiLinearForCausalLM"
    config.parallel_config.pipeline_parallel_size = 2
    with pytest.raises(ValueError, match="pipeline_parallel_size=1"):
        VllmConfig.validate_mamba_cached_kernel(config)


def test_mamba_cache_mode_all_is_rejected():
    """The removed 'all' mode must fail validation instead of being ignored."""
    with pytest.raises(ValidationError, match="mamba_cache_mode"):
        CacheConfig(mamba_cache_mode="all")


def test_per_request_spec_decode_metrics_requires_spec_decode():
    # The flag only makes sense with speculative decoding configured; enabling
    # it without --speculative-config should fail fast rather than silently
    # produce no metrics.
    for level in ("summary", "detailed"):
        with pytest.raises(ValueError, match="speculative"):
            VllmConfig(
                observability_config=ObservabilityConfig(
                    per_request_spec_decode_metrics=level
                )
            )


@pytest.mark.parametrize(
    "kv_transfer_config",
    [
        KVTransferConfig(
            kv_connector="NixlConnector",
            kv_role="kv_both",
        ),
        KVTransferConfig(
            kv_connector="MultiConnector",
            kv_role="kv_both",
            kv_connector_extra_config={
                "connectors": [
                    {
                        "kv_connector": "NixlConnector",
                        "kv_role": "kv_both",
                    },
                    {
                        "kv_connector": "OffloadingConnector",
                        "kv_role": "kv_both",
                    },
                ]
            },
        ),
    ],
)
def test_pd_dcp_interleave_size_is_adjusted_to_block_size(
    caplog, disable_log_dedup, kv_transfer_config
):
    config = VllmConfig(
        cache_config=CacheConfig(block_size=16),
        device_config=DeviceConfig(device="cpu"),
        parallel_config=ParallelConfig(
            tensor_parallel_size=2,
            decode_context_parallel_size=2,
            cp_kv_cache_interleave_size=3,
            distributed_executor_backend="mp",
        ),
        kv_transfer_config=kv_transfer_config,
    )

    kv_cache_config = SimpleNamespace(
        kv_cache_groups=[SimpleNamespace(kv_cache_spec=SimpleNamespace(block_size=16))]
    )
    with caplog.at_level(logging.INFO):
        config.adjust_dcp_kv_cache_interleave_size(kv_cache_config)

    assert config.parallel_config.cp_kv_cache_interleave_size == 16
    assert "automatically adjusted from 3 to block_size 16" in caplog.text


def test_kv_offloading_does_not_adjust_dcp_interleave_size():
    config = VllmConfig(
        cache_config=CacheConfig(block_size=16),
        device_config=DeviceConfig(device="cpu"),
        parallel_config=ParallelConfig(
            tensor_parallel_size=2,
            decode_context_parallel_size=2,
            cp_kv_cache_interleave_size=1,
            distributed_executor_backend="mp",
        ),
        kv_transfer_config=KVTransferConfig(
            kv_connector="OffloadingConnector",
            kv_role="kv_both",
        ),
    )

    kv_cache_config = SimpleNamespace(
        kv_cache_groups=[SimpleNamespace(kv_cache_spec=SimpleNamespace(block_size=16))]
    )
    config.adjust_dcp_kv_cache_interleave_size(kv_cache_config)

    assert config.parallel_config.cp_kv_cache_interleave_size == 1


def test_kv_offloading_does_not_skip_dcp_interleave_validation():
    config = SimpleNamespace(
        cache_config=SimpleNamespace(
            block_size=16,
            mamba_cache_mode="none",
        ),
        parallel_config=SimpleNamespace(
            decode_context_parallel_size=2,
            cp_kv_cache_interleave_size=3,
        ),
        scheduler_config=SimpleNamespace(disable_chunked_mm_input=False),
        kv_transfer_config=KVTransferConfig(
            kv_connector="OffloadingConnector",
            kv_role="kv_both",
        ),
    )

    with pytest.raises(AssertionError, match="divisible by"):
        VllmConfig.validate_block_size(config)


def test_nixl_dcp_check_skipped_for_submodel_config():
    # with_hf_config() builds a submodule view of the config (e.g. a
    # multimodal model's text stack) whose architecture list is empty.
    # Re-running VllmConfig validation on it is unsafe: the NIXL DCP check
    # reads use_mla, which resolves the architecture registry.
    model_config = ModelConfig("Qwen/Qwen2-VL-2B-Instruct", max_model_len=2048)
    vllm_config = VllmConfig(
        model_config=model_config,
        device_config=DeviceConfig(device="cpu"),
        kv_transfer_config=KVTransferConfig(
            kv_connector="NixlConnector",
            kv_role="kv_both",
        ),
    )
    submodel_config = vllm_config.with_hf_config(model_config.hf_text_config)
    assert submodel_config.model_config.is_submodel_config
    assert submodel_config.model_config.architectures == []
    assert submodel_config.model_config.use_mla is False


def test_nixl_dcp_check_rejects_non_mla_model_with_dcp(monkeypatch):
    # Pretend the model has a single KV head so the DCP feasibility checks
    # pass and the MLA-only assert is what actually fires.
    monkeypatch.setattr(ModelConfig, "get_total_num_kv_heads", lambda self: 1)
    with pytest.raises(ValidationError, match="only supported for MLA models"):
        VllmConfig(
            model_config=ModelConfig("Qwen/Qwen3-0.6B", max_model_len=2048),
            device_config=DeviceConfig(device="cpu"),
            parallel_config=ParallelConfig(
                tensor_parallel_size=2,
                decode_context_parallel_size=2,
                distributed_executor_backend="mp",
            ),
            kv_transfer_config=KVTransferConfig(
                kv_connector="NixlConnector",
                kv_role="kv_both",
            ),
        )


def test_compile_config_repr_succeeds():
    # setup: VllmBackend mutates the config object
    config = VllmConfig()
    backend = VllmBackend(config)
    backend.configure_post_pass()

    # test that repr(config) succeeds
    val = repr(config)
    assert "VllmConfig" in val
    assert "inductor_passes" in val


@pytest.mark.parametrize(
    ("env_value", "expected"),
    [
        (None, None),
        ("0", False),
        ("1", True),
    ],
)
def test_v2_model_runner_env_tri_state(monkeypatch, env_value, expected):
    if env_value is None:
        monkeypatch.delenv("VLLM_USE_V2_MODEL_RUNNER", raising=False)
    else:
        monkeypatch.setenv("VLLM_USE_V2_MODEL_RUNNER", env_value)

    assert envs.VLLM_USE_V2_MODEL_RUNNER is expected


def test_hisparse_requires_v2_model_runner():
    config = object.__new__(VllmConfig)
    config.attention_config = AttentionConfig(hisparse_config=HiSparseConfig())

    with patch.object(envs, "VLLM_USE_V2_MODEL_RUNNER", None):
        assert config.use_v2_model_runner
    with (
        patch.object(envs, "VLLM_USE_V2_MODEL_RUNNER", False),
        pytest.raises(ValueError, match="requires Model Runner V2"),
    ):
        _ = config.use_v2_model_runner


def test_hisparse_rejects_decode_context_parallelism(monkeypatch):
    monkeypatch.setattr(current_platform, "is_cuda", lambda: True)
    monkeypatch.setattr(current_platform, "device_count", lambda: 2)
    with pytest.raises(ValueError, match="decode context parallelism"):
        VllmConfig(
            attention_config=AttentionConfig(hisparse_config=HiSparseConfig()),
            parallel_config=ParallelConfig(
                tensor_parallel_size=2,
                decode_context_parallel_size=2,
            ),
        )


def test_hisparse_rejects_pipeline_parallelism(monkeypatch):
    monkeypatch.setattr(current_platform, "is_cuda", lambda: True)
    monkeypatch.setattr(current_platform, "device_count", lambda: 2)
    with pytest.raises(ValueError, match="pipeline parallelism"):
        VllmConfig(
            attention_config=AttentionConfig(hisparse_config=HiSparseConfig()),
            parallel_config=ParallelConfig(pipeline_parallel_size=2),
        )


def test_hisparse_rejects_disabled_hybrid_kv_cache_manager(monkeypatch):
    monkeypatch.setattr(current_platform, "is_cuda", lambda: True)
    monkeypatch.setattr("vllm.config.vllm.HAS_TRITON", True)
    with pytest.raises(ValueError, match="requires the hybrid KV cache manager"):
        VllmConfig(
            attention_config=AttentionConfig(hisparse_config=HiSparseConfig()),
            scheduler_config=SchedulerConfig(
                max_model_len=2048,
                is_encoder_decoder=False,
                disable_hybrid_kv_cache_manager=True,
            ),
        )


def test_hisparse_rejects_disabled_full_isl_reservation(monkeypatch):
    monkeypatch.setattr(current_platform, "is_cuda", lambda: True)
    with pytest.raises(ValueError, match="requires --scheduler-reserve-full-isl"):
        VllmConfig(
            attention_config=AttentionConfig(hisparse_config=HiSparseConfig()),
            scheduler_config=SchedulerConfig(
                max_model_len=2048,
                is_encoder_decoder=False,
                scheduler_reserve_full_isl=False,
            ),
        )


def test_hisparse_rejects_non_cuda(monkeypatch):
    monkeypatch.setattr(current_platform, "is_cuda", lambda: False)
    with pytest.raises(ValueError, match="requires NVIDIA CUDA"):
        VllmConfig(attention_config=AttentionConfig(hisparse_config=HiSparseConfig()))


@pytest.mark.parametrize(
    "kv_transfer_config",
    [
        KVTransferConfig(
            kv_connector="HiSparseConnector",
            kv_role="kv_both",
            kv_connector_extra_config={"host_pool_gib": 128},
        ),
        KVTransferConfig(
            kv_connector="MultiConnector",
            kv_role="kv_both",
            kv_connector_extra_config={
                "connectors": [
                    {
                        "kv_connector": "OffloadingConnector",
                        "kv_role": "kv_both",
                        "kv_connector_extra_config": {"cpu_bytes_to_use": 1 << 30},
                    },
                    {
                        "kv_connector": "HiSparseConnector",
                        "kv_role": "kv_both",
                        "kv_connector_extra_config": {"host_pool_gib": 128},
                    },
                ]
            },
        ),
        KVTransferConfig(
            kv_connector="MultiConnector",
            kv_role="kv_consumer",
            kv_connector_extra_config={
                "connectors": [
                    {"kv_connector": "NixlConnector", "kv_role": "kv_consumer"},
                    {
                        "kv_connector": "HiSparseConnector",
                        "kv_role": "kv_both",
                        "kv_connector_extra_config": {"host_pool_gib": 128},
                    },
                ]
            },
        ),
    ],
    ids=["standalone", "multi-connector", "pd-decode"],
)
def test_hisparse_connector_implies_attention_config(monkeypatch, kv_transfer_config):
    monkeypatch.setattr(current_platform, "is_cuda", lambda: True)
    monkeypatch.setattr("vllm.config.vllm.HAS_TRITON", True)
    config = VllmConfig(
        kv_transfer_config=kv_transfer_config,
        # Skip the HMA auto-detect block: it imports the connector class,
        # which pulls the platform's compiled attention extensions.
        scheduler_config=SchedulerConfig(
            max_model_len=2048,
            is_encoder_decoder=False,
            disable_hybrid_kv_cache_manager=False,
        ),
    )
    assert isinstance(config.attention_config.hisparse_config, HiSparseConfig)


def test_hisparse_connector_preserves_explicit_attention_config(monkeypatch):
    monkeypatch.setattr(current_platform, "is_cuda", lambda: True)
    monkeypatch.setattr("vllm.config.vllm.HAS_TRITON", True)
    config = VllmConfig(
        attention_config=AttentionConfig(
            hisparse_config=HiSparseConfig(device_buffer_size=512)
        ),
        kv_transfer_config=KVTransferConfig(
            kv_connector="HiSparseConnector",
            kv_role="kv_both",
            kv_connector_extra_config={"host_pool_gib": 128},
        ),
        scheduler_config=SchedulerConfig(
            max_model_len=2048,
            is_encoder_decoder=False,
            disable_hybrid_kv_cache_manager=False,
        ),
    )
    assert config.attention_config.hisparse_config.device_buffer_size == 512


def test_hisparse_connector_without_cuda_still_rejected(monkeypatch):
    monkeypatch.setattr(current_platform, "is_cuda", lambda: False)
    with pytest.raises(ValueError, match="requires NVIDIA CUDA"):
        VllmConfig(
            kv_transfer_config=KVTransferConfig(
                kv_connector="HiSparseConnector",
                kv_role="kv_both",
                kv_connector_extra_config={"host_pool_gib": 128},
            )
        )


def test_no_hisparse_connector_keeps_attention_config_unset(monkeypatch):
    monkeypatch.setattr(current_platform, "is_cuda", lambda: True)
    config = VllmConfig(
        kv_transfer_config=KVTransferConfig(
            kv_connector="OffloadingConnector", kv_role="kv_both"
        )
    )
    assert config.attention_config.hisparse_config is None


def test_rocm_keeps_compiled_deepseek_defaults(monkeypatch):
    """ROCm keeps the DSA models (DeepSeek V3.2/V4, GLM-5.2) on their compiled
    MRV1 paths and off breakable cudagraphs by default."""
    from vllm.config.vllm import (
        ROCM_DEFAULT_MRV1_ARCHITECTURES,
        default_breakable_cudagraph_architectures,
    )
    from vllm.platforms import current_platform

    monkeypatch.setattr(current_platform, "is_rocm", lambda: True)
    # The lookup is lru_cached against a fixed platform.
    default_breakable_cudagraph_architectures.cache_clear()
    try:
        assert "DeepseekV32ForCausalLM" in ROCM_DEFAULT_MRV1_ARCHITECTURES
        assert "DeepseekV4ForCausalLM" in ROCM_DEFAULT_MRV1_ARCHITECTURES
        assert "GlmMoeDsaForCausalLM" in ROCM_DEFAULT_MRV1_ARCHITECTURES

        breakable_architectures = default_breakable_cudagraph_architectures()
        assert "DeepseekV32ForCausalLM" not in breakable_architectures
        assert "DeepseekV32MTPModel" not in breakable_architectures
        assert "GlmMoeDsaForCausalLM" not in breakable_architectures
        # V4.1 cannot torch.compile and the ROCm sparse SWA backend only
        # supports uniform-batch CUDA graphs, so it must opt in or the
        # default FULL_AND_PIECEWISE serve path cannot start.
        assert "DeepseekV41ForCausalLM" in breakable_architectures

        # The carve-out takes effect via the runner-selection property
        # (warning_once args must be hashable for its lru_cache).
        monkeypatch.delenv("VLLM_USE_V2_MODEL_RUNNER", raising=False)
        config = SimpleNamespace(
            model_config=SimpleNamespace(architectures=["DeepseekV32ForCausalLM"]),
            attention_config=AttentionConfig(),
        )
        config._get_v1_model_runner_unsupported_features = lambda: []
        assert VllmConfig.use_v2_model_runner.fget(config) is False
    finally:
        default_breakable_cudagraph_architectures.cache_clear()


def test_rocm_mrv1_default_yields_to_v1_unsupported_config(monkeypatch):
    """The ROCm V1 default is a speed preference, not a capability claim.

    DSpark runs only on V2, so pinning DeepSeek V4 to V1 would fail config
    validation instead of serving it. With nothing V1 refuses, it still holds.
    """
    from vllm.platforms import current_platform

    monkeypatch.setattr(current_platform, "is_rocm", lambda: True)
    monkeypatch.setattr(vllm_config_module, "HAS_TRITON", True)
    monkeypatch.delenv("VLLM_USE_V2_MODEL_RUNNER", raising=False)

    config = SimpleNamespace(
        model_config=SimpleNamespace(
            architectures=["DeepseekV4ForCausalLM"], is_diffusion=False
        ),
        attention_config=AttentionConfig(),
        parallel_config=SimpleNamespace(
            prefill_context_parallel_size=1,
            pipeline_parallel_size=1,
            enable_batch_sharded_sampling=False,
        ),
        scheduler_config=SimpleNamespace(async_scheduling=False),
        speculative_config=None,
    )
    config._dflash_needs_multi_kv_group = lambda: False
    config._is_dflash_candidate_draft = lambda: False
    config._get_v2_model_runner_unsupported_features = lambda: []
    # The real predicate, so the test also pins where dspark lands in it.
    config._get_v1_model_runner_unsupported_features = lambda: (
        VllmConfig._get_v1_model_runner_unsupported_features(config)
    )

    assert VllmConfig.use_v2_model_runner.fget(config) is False

    config.speculative_config = SimpleNamespace(
        method="dspark", enable_adaptive_verification=False
    )
    assert VllmConfig.use_v2_model_runner.fget(config) is True

    # Yielding is not the same as selecting V2: the later checks still run, so
    # a config neither runner can serve lands on V1 and fails validation there.
    config._get_v2_model_runner_unsupported_features = lambda: ["sequence parallelism"]
    assert VllmConfig.use_v2_model_runner.fget(config) is False


@pytest.mark.parametrize(
    ("model", "architecture"),
    [
        ("nvidia/GLM-5.2-NVFP4", "GlmMoeDsaForCausalLM"),
        ("zai-org/GLM-5.2-FP8", "GlmMoeDsaForCausalLM"),
        ("nvidia/DeepSeek-V3.2-NVFP4", "DeepseekV32ForCausalLM"),
    ],
)
@pytest.mark.parametrize("with_mtp", [False, True], ids=["no-mtp", "mtp"])
def test_dsa_models_default_to_mrv2_and_breakable_cudagraph(
    monkeypatch, model, architecture, with_mtp
):
    from vllm.compilation.breakable_cudagraph import (
        is_breakable_cudagraph_enabled,
    )
    from vllm.config.vllm import default_breakable_cudagraph_architectures
    from vllm.platforms import current_platform

    monkeypatch.delenv("VLLM_USE_BREAKABLE_CUDAGRAPH", raising=False)
    monkeypatch.delenv("VLLM_USE_V2_MODEL_RUNNER", raising=False)
    monkeypatch.setattr(vllm_config_module, "HAS_TRITON", True)
    monkeypatch.setattr(current_platform, "is_rocm", lambda: False)
    default_breakable_cudagraph_architectures.cache_clear()

    model_config = SimpleNamespace(
        model=model,
        architectures=[architecture],
        runner_type="generate",
        is_moe=True,
        is_hybrid=False,
        is_attention_free=False,
        is_diffusion=False,
    )
    config = SimpleNamespace(
        model_config=model_config,
        attention_config=AttentionConfig(),
        speculative_config=SimpleNamespace(method="mtp") if with_mtp else None,
        parallel_config=SimpleNamespace(prefill_context_parallel_size=1),
        compilation_config=CompilationConfig(
            cudagraph_mode=CUDAGraphMode.FULL_AND_PIECEWISE
        ),
    )
    config._dflash_needs_multi_kv_group = lambda: False
    config._get_v2_model_runner_unsupported_features = lambda: []
    config._uses_breakable_cudagraph_by_default = lambda: (
        VllmConfig._uses_breakable_cudagraph_by_default(config)
    )

    try:
        assert VllmConfig.use_v2_model_runner.fget(config)
        assert VllmConfig._maybe_enable_breakable_cudagraph(config)
        assert is_breakable_cudagraph_enabled()
        assert config.compilation_config.mode == CompilationMode.NONE
        assert config.compilation_config.cudagraph_mode.has_piecewise_cudagraphs()
    finally:
        os.environ.pop("VLLM_USE_BREAKABLE_CUDAGRAPH", None)
        default_breakable_cudagraph_architectures.cache_clear()


@pytest.mark.parametrize(
    ("architecture", "is_rocm", "expected"),
    [
        ("DeepseekV32ForCausalLM", False, True),
        ("DeepseekV32ForCausalLM", True, False),
        ("DeepseekV32MTPModel", False, True),
        ("DeepseekV32MTPModel", True, False),
        ("GlmMoeDsaForCausalLM", False, True),
        ("GlmMoeDsaForCausalLM", True, False),
        ("Qwen4ExpForCausalLM", False, True),
        ("Qwen4ExpForCausalLM", True, False),
        ("Qwen4ExpForConditionalGeneration", False, True),
        ("Qwen4ExpForConditionalGeneration", True, False),
        ("Qwen4ExpMTP", False, True),
        ("Qwen4ExpMTP", True, False),
    ],
)
def test_breakable_cudagraph_platform_default(
    monkeypatch, architecture, is_rocm, expected
):
    from vllm.config.vllm import default_breakable_cudagraph_architectures
    from vllm.platforms import current_platform

    monkeypatch.delenv("VLLM_USE_BREAKABLE_CUDAGRAPH", raising=False)
    monkeypatch.setattr(current_platform, "is_rocm", lambda: is_rocm)
    default_breakable_cudagraph_architectures.cache_clear()
    config = SimpleNamespace(
        model_config=SimpleNamespace(architectures=[architecture]),
        compilation_config=CompilationConfig(),
    )
    config._uses_breakable_cudagraph_by_default = lambda: (
        VllmConfig._uses_breakable_cudagraph_by_default(config)
    )

    try:
        assert VllmConfig._maybe_enable_breakable_cudagraph(config) is expected
        if expected:
            assert config.compilation_config.mode == CompilationMode.NONE
    finally:
        os.environ.pop("VLLM_USE_BREAKABLE_CUDAGRAPH", None)
        default_breakable_cudagraph_architectures.cache_clear()


@pytest.mark.parametrize(
    "case,hidden,heads,intermediate,tp,expected",
    [
        ("bf16", 2048, (16, 8), 6144, 1, True),
        ("bi-off", 2048, (16, 8), 6144, 1, False),
        ("opt-out", 2048, (16, 8), 6144, 1, False),
        ("eager", 2048, (16, 8), 6144, 1, False),
        ("sp", 2048, (16, 8), 6144, 1, False),
        ("fp16", 2048, (16, 8), 6144, 1, False),
        ("quantized", 2048, (16, 8), 6144, 1, False),
        ("untuned", 4096, (32, 32), 11008, 1, False),
        ("tp4-tuned", 4096, (8, 2), 12288, 4, True),
        ("list-intermediate", 2048, (16, 8), [2048, 4096], 1, False),
        ("no-table", 2048, None, 6144, 1, False),
    ],
)
def test_batch_invariant_breakable_cudagraph(
    monkeypatch, case, hidden, heads, intermediate, tp, expected
):
    from vllm.config.vllm import default_breakable_cudagraph_architectures
    from vllm.model_executor.determinism import batch_invariant_configs as bi_configs

    monkeypatch.setenv("VLLM_BATCH_INVARIANT", "0" if case == "bi-off" else "1")
    monkeypatch.delenv("VLLM_USE_BREAKABLE_CUDAGRAPH", raising=False)
    if case == "opt-out":
        monkeypatch.setenv("VLLM_USE_BREAKABLE_CUDAGRAPH", "0")
    monkeypatch.setattr(current_platform, "is_cuda", lambda: True)
    monkeypatch.setattr(current_platform, "is_rocm", lambda: False)
    table = None if case == "no-table" else {(12288, 2048): None, (6144, 4096): None}
    monkeypatch.setattr(current_platform, "get_device_capability", lambda: None)
    monkeypatch.setattr(bi_configs, "_get_tuned_matmul_arch_family", lambda cap: "test")
    monkeypatch.setattr(
        bi_configs,
        "_BATCH_INVARIANT_MATMUL_TUNED_CONFIGS",
        {"test": table} if table else {},
    )
    default_breakable_cudagraph_architectures.cache_clear()
    config = object.__new__(VllmConfig)
    config.model_config = SimpleNamespace(
        architectures=["Qwen3ForCausalLM"],
        enforce_eager=case == "eager",
        dtype=torch.float16 if case == "fp16" else torch.bfloat16,
        quantization="fp8" if case == "quantized" else None,
        hf_text_config=SimpleNamespace(intermediate_size=intermediate),
        get_hidden_size=lambda: hidden,
        get_head_size=lambda: 128,
        get_num_attention_heads=lambda pc: heads[0],
        get_num_kv_heads=lambda pc: heads[1],
    )
    config.parallel_config = SimpleNamespace(tensor_parallel_size=tp)
    config.compilation_config = (
        CompilationConfig(pass_config=PassConfig(enable_sp=True))
        if case == "sp"
        else CompilationConfig()
    )
    try:
        assert config._maybe_enable_breakable_cudagraph() is expected
        if expected:
            assert config.compilation_config.mode == CompilationMode.NONE
    finally:
        os.environ.pop("VLLM_USE_BREAKABLE_CUDAGRAPH", None)
        default_breakable_cudagraph_architectures.cache_clear()


@pytest.mark.parametrize(
    ("model_type", "expected_architecture"),
    [
        ("deepseek_v32", "DeepseekV32MTPModel"),
        ("glm_moe_dsa", "DeepseekV32MTPModel"),
        ("deepseek_v3", "DeepSeekMTPModel"),
    ],
)
def test_dsa_models_select_matching_mtp(model_type, expected_architecture):
    from transformers import PreTrainedConfig

    hf_config = PreTrainedConfig(
        architectures=["DeepseekV32ForCausalLM"],
        num_nextn_predict_layers=1,
    )
    hf_config.model_type = model_type

    SpeculativeConfig.hf_config_override(hf_config)

    assert hf_config.architectures == [expected_architecture]


def test_v2_model_runner_supports_extract_hidden_states():
    config = VllmConfig()
    config.speculative_config = cast(
        SpeculativeConfig,
        SimpleNamespace(
            method="extract_hidden_states",
            parallel_drafting=False,
            enable_adaptive_verification=False,
        ),
    )

    assert config._get_v2_model_runner_unsupported_features() == []


def test_v2_model_runner_supports_custom_logits_processors():
    config = VllmConfig()
    config.model_config = cast(
        ModelConfig, SimpleNamespace(logits_processors=["a.b:C"])
    )

    assert config._get_v2_model_runner_unsupported_features() == []


@pytest.mark.parametrize("architecture", ["DFlash2DraftModel", "LiLiCorrDraftModel"])
def test_dflash_candidate_draft_forces_v2_model_runner(architecture):
    """A DFlash2 draft must reach the V2 speculator, the only one that runs its
    candidate selector; on V1 it would draft as DFlash1 without raising."""

    def config(method, architectures):
        return SimpleNamespace(
            speculative_config=SimpleNamespace(
                method=method,
                draft_model_config=SimpleNamespace(architectures=architectures),
            )
        )

    assert VllmConfig._is_dflash_candidate_draft(config("dflash", [architecture]))
    assert not VllmConfig._is_dflash_candidate_draft(
        config("dflash", ["DFlashDraftModel"])
    )
    assert not VllmConfig._is_dflash_candidate_draft(config("eagle", [architecture]))
    assert not VllmConfig._is_dflash_candidate_draft(
        SimpleNamespace(speculative_config=None)
    )
    assert not VllmConfig._is_dflash_candidate_draft(
        SimpleNamespace(
            speculative_config=SimpleNamespace(method="dflash", draft_model_config=None)
        )
    )


@pytest.mark.parametrize(
    ("use_v2_model_runner", "expected_capture_sizes"),
    [
        (False, [4, 8, 12, 16]),
        (True, list(range(1, 17))),
    ],
)
def test_resolve_cudagraph_mode_adjusts_spec_decode_sizes_only_for_v1(
    use_v2_model_runner,
    expected_capture_sizes,
):
    compilation_config = CompilationConfig(
        cudagraph_mode=CUDAGraphMode.FULL_AND_PIECEWISE,
        cudagraph_capture_sizes=list(range(1, 17)),
    )
    compilation_config.max_cudagraph_capture_size = 16
    compilation_config.post_init_cudagraph_sizes()

    cudagraph_mode = compilation_config.resolve_cudagraph_mode_and_sizes(
        AttentionCGSupport.ALWAYS,
        "FakeAttentionBackend",
        uniform_decode_query_len=4,
        use_v2_model_runner=use_v2_model_runner,
        tensor_parallel_size=1,
    )

    assert cudagraph_mode == CUDAGraphMode.FULL_AND_PIECEWISE
    assert compilation_config.cudagraph_capture_sizes == expected_capture_sizes


@pytest.mark.parametrize(
    ("mode", "piecewise_capture_available", "attention_support", "expected"),
    [
        ("PIECEWISE", False, "ALWAYS", "NONE"),
        ("FULL_AND_PIECEWISE", False, "ALWAYS", "FULL_DECODE_ONLY"),
        ("FULL_DECODE_ONLY", False, "ALWAYS", "FULL_DECODE_ONLY"),
        ("FULL_DECODE_ONLY", False, "NEVER", "NONE"),
        ("FULL_AND_PIECEWISE", True, "UNIFORM_BATCH", "FULL_AND_PIECEWISE"),
        ("FULL", True, "UNIFORM_BATCH", "FULL_DECODE_ONLY"),
    ],
)
def test_resolve_cudagraph_mode_uses_loaded_piecewise_provider(
    mode, piecewise_capture_available, attention_support, expected
):
    compilation_config = CompilationConfig(
        mode=CompilationMode.VLLM_COMPILE,
        cudagraph_mode=CUDAGraphMode[mode],
        use_inductor_graph_partition=True,
    )

    resolved = compilation_config.resolve_cudagraph_mode_and_sizes(
        AttentionCGSupport[attention_support],
        "FakeAttentionBackend",
        piecewise_capture_available=piecewise_capture_available,
    )

    assert resolved.name == expected
    assert compilation_config.cudagraph_mode == resolved


@pytest.mark.skipif(
    not current_platform.is_cuda_alike(), reason="Requires CUDA graph support"
)
@pytest.mark.parametrize(
    "engine_kwargs",
    [
        {"runner": "pooling", "convert": "embed"},
        pytest.param(
            {"prefill_context_parallel_size": 2, "tensor_parallel_size": 2},
            marks=pytest.mark.skipif(
                not current_platform.is_rocm(), reason="ROCm PCP graph restriction"
            ),
        ),
    ],
)
def test_late_piecewise_restrictions_without_compilation(monkeypatch, engine_kwargs):
    """Late compatibility overrides must not restore unavailable piecewise graphs."""
    from vllm.engine.arg_utils import EngineArgs

    monkeypatch.setenv("VLLM_USE_BREAKABLE_CUDAGRAPH", "0")
    monkeypatch.setenv("VLLM_USE_V2_MODEL_RUNNER", "1")
    config = EngineArgs(
        model="facebook/opt-125m",
        compilation_config=CompilationConfig(mode=CompilationMode.NONE),
        **engine_kwargs,
    ).create_engine_config()

    assert config.compilation_config.cudagraph_mode == CUDAGraphMode.NONE
    assert config.compilation_config.cudagraph_capture_sizes == []
    assert config.compilation_config.max_cudagraph_capture_size == 0


def test_resolve_cudagraph_mode_skips_mamba_block_check_while_profiling():
    """Cudagraph memory profiling uses a minimal KV cache, so the Mamba
    block-count guard must only fire for the real cache sizing."""
    kv_cache_config = SimpleNamespace(has_mamba_layers=True, num_blocks=4)

    compilation_config = CompilationConfig(
        cudagraph_mode=CUDAGraphMode.FULL_AND_PIECEWISE,
    )
    with pytest.raises(ValueError, match="exceeds available Mamba cache blocks"):
        compilation_config.resolve_cudagraph_mode_and_sizes(
            AttentionCGSupport.ALWAYS,
            "FakeAttentionBackend",
            uniform_decode_query_len=1,
            use_v2_model_runner=True,
            tensor_parallel_size=1,
            kv_cache_config=kv_cache_config,
            max_num_reqs=256,
        )

    compilation_config = CompilationConfig(
        cudagraph_mode=CUDAGraphMode.FULL_AND_PIECEWISE,
    )
    cudagraph_mode = compilation_config.resolve_cudagraph_mode_and_sizes(
        AttentionCGSupport.ALWAYS,
        "FakeAttentionBackend",
        uniform_decode_query_len=1,
        use_v2_model_runner=True,
        tensor_parallel_size=1,
        kv_cache_config=kv_cache_config,
        max_num_reqs=256,
        is_profiling=True,
    )
    assert cudagraph_mode == CUDAGraphMode.FULL_AND_PIECEWISE


@pytest.mark.parametrize(
    ("graph_mode", "should_raise"),
    [
        (CUDAGraphMode.FULL_DECODE_ONLY, False),
        (CUDAGraphMode.FULL_AND_PIECEWISE, False),
        (CUDAGraphMode.NONE, True),
        (CUDAGraphMode.PIECEWISE, True),
    ],
)
def test_adaptive_verification_requires_full_cudagraphs(graph_mode, should_raise):
    config = SimpleNamespace(
        speculative_config=SimpleNamespace(enable_adaptive_verification=True),
        lora_config=None,
        compilation_config=CompilationConfig(cudagraph_mode=graph_mode),
        parallel_config=SimpleNamespace(pipeline_parallel_size=1),
    )

    if should_raise:
        with pytest.raises(ValueError, match="requires full CUDA graphs"):
            VllmConfig._validate_adaptive_verification(config)
    else:
        VllmConfig._validate_adaptive_verification(config)


@pytest.mark.parametrize(
    ("model_config", "expected"),
    [
        (
            SimpleNamespace(
                model="Qwen/Qwen3-32B",
                architectures=["Qwen3ForCausalLM"],
                runner_type="generate",
                is_moe=False,
                is_quantized=False,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="Qwen/Qwen2-7B-Instruct",
                architectures=["Qwen2ForCausalLM"],
                runner_type="generate",
                is_moe=False,
                is_quantized=False,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="meta-llama/Llama-3.2-1B",
                architectures=["LlamaForCausalLM"],
                runner_type="generate",
                is_moe=False,
                is_quantized=False,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="mistralai/Mistral-7B-v0.1",
                architectures=["MistralForCausalLM"],
                runner_type="generate",
                is_moe=False,
                is_quantized=False,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="facebook/opt-125m",
                architectures=["OPTForCausalLM"],
                runner_type="generate",
                is_moe=False,
                is_quantized=False,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="google/gemma-2-2b",
                architectures=["Gemma2ForCausalLM"],
                runner_type="generate",
                is_moe=False,
                is_quantized=False,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="deepseek-ai/DeepSeek-V2-Lite-Chat",
                architectures=["DeepseekV2ForCausalLM"],
                runner_type="generate",
                is_moe=True,
                is_quantized=False,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="deepseek-ai/DeepSeek-V2-Chat",
                architectures=["DeepseekV2ForCausalLM"],
                runner_type="generate",
                is_moe=True,
                is_quantized=False,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="deepseek-ai/DeepSeek-V3",
                architectures=["DeepseekV3ForCausalLM"],
                runner_type="generate",
                is_moe=True,
                is_quantized=False,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="deepseek-ai/DeepSeek-V4-Flash",
                architectures=["DeepseekV4ForCausalLM"],
                runner_type="generate",
                is_moe=True,
                is_quantized=True,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="Qwen/Qwen1.5-MoE-A2.7B",
                architectures=["Qwen2MoeForCausalLM"],
                runner_type="generate",
                is_moe=True,
                is_quantized=False,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="Qwen/Qwen1.5-MoE-A2.7B-Chat",
                architectures=["Qwen2MoeForCausalLM"],
                runner_type="generate",
                is_moe=True,
                is_quantized=False,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="ibm-research/PowerMoE-3b",
                architectures=["GraniteMoeForCausalLM"],
                runner_type="generate",
                is_moe=True,
                is_quantized=False,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="thinkingmachines/Inkling",
                architectures=["InklingForCausalLM"],
                runner_type="generate",
                is_moe=True,
                is_quantized=False,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="thinkingmachines/Inkling",
                architectures=["InklingForConditionalGeneration"],
                runner_type="generate",
                is_moe=True,
                is_quantized=False,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="mistralai/Mixtral-8x7B-Instruct-v0.1",
                architectures=["MixtralForCausalLM"],
                runner_type="generate",
                is_moe=True,
                is_quantized=False,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="Qwen/Qwen3-1.7B-FP8",
                architectures=["Qwen3ForCausalLM"],
                runner_type="generate",
                is_moe=False,
                is_quantized=True,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="Qwen/Qwen3.5-4B",
                architectures=["Qwen3_5ForConditionalGeneration"],
                runner_type="generate",
                is_moe=False,
                is_quantized=False,
                is_hybrid=True,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="state-spaces/mamba-130m-hf",
                architectures=["MambaForCausalLM"],
                runner_type="generate",
                is_moe=False,
                is_quantized=False,
                is_attention_free=True,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="sentence-transformers/all-MiniLM-L6-v2",
                architectures=["BertModel"],
                runner_type="pooling",
                is_multimodal_model=False,
                is_moe=False,
                is_quantized=False,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="Qwen/Qwen3-Embedding-0.6B",
                architectures=["Qwen3ForCausalLM"],
                runner_type="pooling",
                is_multimodal_model=False,
                is_moe=False,
                is_quantized=False,
            ),
            True,
        ),
        (
            SimpleNamespace(
                model="TomoroAI/tomoro-colqwen3-embed-4b",
                architectures=["ColQwen3"],
                runner_type="pooling",
                is_multimodal_model=True,
                is_moe=False,
                is_quantized=False,
            ),
            True,
        ),
    ],
)
def test_models_default_to_v2_model_runner(model_config, expected, monkeypatch):
    from vllm.platforms import current_platform

    # The expectations below are the platform-independent defaults; ROCm's
    # DeepSeek carve-out is covered by test_rocm_keeps_compiled_deepseek_defaults.
    monkeypatch.delenv("VLLM_USE_V2_MODEL_RUNNER", raising=False)
    monkeypatch.setattr(vllm_config_module, "HAS_TRITON", True)
    monkeypatch.setattr(current_platform, "is_rocm", lambda: False)
    config = SimpleNamespace(
        model_config=model_config,
        attention_config=AttentionConfig(),
    )
    config._get_v2_model_runner_unsupported_features = lambda: []

    assert VllmConfig.use_v2_model_runner.fget(config) is expected


def test_v1_model_runner_rejects_v2_only_features():
    config = SimpleNamespace(
        parallel_config=ParallelConfig(
            prefill_context_parallel_size=2,
            distributed_executor_backend="mp",
        ),
        scheduler_config=SchedulerConfig.default_factory(async_scheduling=False),
        speculative_config=None,
        model_config=None,
    )
    config._dflash_needs_multi_kv_group = lambda: False
    config._is_dflash_candidate_draft = lambda: False
    config._get_v1_model_runner_unsupported_features = lambda: (
        VllmConfig._get_v1_model_runner_unsupported_features(config)
    )

    with pytest.raises(ValueError, match="prefill context parallel"):
        VllmConfig._validate_v1_model_runner(config)


def test_batch_sharded_sampling_rejects_return_sampling_mask():
    """The batch-sharded gather drops sampling masks, so the combination must
    fail loudly instead of returning ``sampling_mask=None``."""
    config = SimpleNamespace(
        parallel_config=SimpleNamespace(
            enable_batch_sharded_sampling=True, tensor_parallel_size=2
        ),
        scheduler_config=SimpleNamespace(max_num_seqs=8),
        model_config=SimpleNamespace(max_logprobs=20, return_sampling_mask=True),
        speculative_config=None,
    )

    with pytest.raises(ValueError, match="sampling masks"):
        VllmConfig._validate_batch_sharded_sampling(config)


@pytest.mark.skip_global_cleanup
def test_with_hf_config_populates_missing_architectures_from_causal_lm_mapping(
    monkeypatch,
):
    monkeypatch.setattr(
        vllm_config_module,
        "replace",
        lambda self, **kwargs: SimpleNamespace(**kwargs),
    )
    cfg = SimpleNamespace(
        model_config=SimpleNamespace(
            is_multimodal_model=False,
            hf_config=SimpleNamespace(),
            get_model_arch_config=lambda: "arch-config",
        )
    )
    hf_config = SimpleNamespace(model_type="mistral", architectures=None)

    updated = VllmConfig.with_hf_config(cfg, hf_config)

    assert updated.model_config.hf_config.architectures == ["MistralForCausalLM"]
    assert hf_config.architectures is None


@pytest.mark.skip_global_cleanup
def test_with_hf_config_preserves_explicit_architectures_override(monkeypatch):
    monkeypatch.setattr(
        vllm_config_module,
        "replace",
        lambda self, **kwargs: SimpleNamespace(**kwargs),
    )
    cfg = SimpleNamespace(
        model_config=SimpleNamespace(
            is_multimodal_model=False,
            hf_config=SimpleNamespace(),
            get_model_arch_config=lambda: "arch-config",
        )
    )
    hf_config = SimpleNamespace(model_type="mistral", architectures=None)

    updated = VllmConfig.with_hf_config(
        cfg,
        hf_config,
        architectures=["Ministral3ForCausalLM"],
    )

    assert updated.model_config.hf_config.architectures == ["Ministral3ForCausalLM"]


@pytest.mark.skip_global_cleanup
def test_with_hf_config_leaves_unknown_model_type_without_architectures(
    monkeypatch,
):
    monkeypatch.setattr(
        vllm_config_module,
        "replace",
        lambda self, **kwargs: SimpleNamespace(**kwargs),
    )
    cfg = SimpleNamespace(
        model_config=SimpleNamespace(
            is_multimodal_model=False,
            hf_config=SimpleNamespace(),
            get_model_arch_config=lambda: "arch-config",
        )
    )
    hf_config = SimpleNamespace(
        model_type="not_a_real_model",
        architectures=None,
    )

    updated = VllmConfig.with_hf_config(cfg, hf_config)

    assert updated.model_config.hf_config.architectures is None


@pytest.mark.parametrize(
    "checkpoint_tensors,tied",
    [
        # The checkpoint has an lm_head of its own, so it must win over the config
        (["model.embed_tokens.weight", "lm_head.weight"], False),
        (["model.embed_tokens.weight"], True),
        # Contents unknown (not safetensors), so the config must be left alone
        ([], True),
    ],
)
def test_maybe_untie_word_embeddings(tmp_path, checkpoint_tensors, tied):
    import torch
    from safetensors.torch import save_file

    if checkpoint_tensors:
        save_file(
            {name: torch.zeros(2, 2) for name in checkpoint_tensors},
            tmp_path / "model.safetensors",
        )

    text_config = SimpleNamespace(tie_word_embeddings=True)
    model_config = SimpleNamespace(
        model=str(tmp_path),
        revision=None,
        hf_config=SimpleNamespace(
            tie_word_embeddings=True,
            get_text_config=lambda: text_config,
        ),
        word_embeddings_untied_by_checkpoint=False,
    )

    ModelConfig.maybe_untie_word_embeddings(model_config)

    # Both levels must agree, since different callers read different ones
    assert model_config.hf_config.tie_word_embeddings is tied
    assert text_config.tie_word_embeddings is tied
    assert model_config.word_embeddings_untied_by_checkpoint is not tied


def test_async_scheduling_with_pipeline_parallelism_is_allowed():
    cfg = VllmConfig(
        scheduler_config=SchedulerConfig(
            max_model_len=8192,
            is_encoder_decoder=False,
            async_scheduling=True,
        ),
        parallel_config=ParallelConfig(
            pipeline_parallel_size=2,
            distributed_executor_backend="mp",
            nnodes=2,
        ),
    )
    assert cfg.scheduler_config.async_scheduling is True


def test_v1_model_runner_drops_async_scheduling_with_pipeline_parallelism(monkeypatch):
    """PP>1 must stay buildable whenever the V1 model runner is selected, such
    as the external_launcher fallback; async scheduling is dropped instead."""
    monkeypatch.setenv("VLLM_USE_V2_MODEL_RUNNER", "0")

    cfg = VllmConfig(
        scheduler_config=SchedulerConfig(
            max_model_len=8192,
            is_encoder_decoder=False,
        ),
        parallel_config=ParallelConfig(
            pipeline_parallel_size=2,
            distributed_executor_backend="mp",
            nnodes=2,
        ),
    )
    assert cfg.scheduler_config.async_scheduling is False


def test_v1_model_runner_rejects_pipeline_parallelism_with_async_scheduling():
    """Only the async combination desyncs the grammar FSM (#45014), so plain
    PP>1 must stay usable on the V1 model runner."""
    config = SimpleNamespace(
        parallel_config=ParallelConfig(
            pipeline_parallel_size=2,
            distributed_executor_backend="mp",
        ),
        scheduler_config=SchedulerConfig.default_factory(async_scheduling=False),
        speculative_config=None,
        model_config=None,
    )
    config._dflash_needs_multi_kv_group = lambda: False
    config._is_dflash_candidate_draft = lambda: False

    assert VllmConfig._get_v1_model_runner_unsupported_features(config) == []

    config.scheduler_config.async_scheduling = True
    unsupported = VllmConfig._get_v1_model_runner_unsupported_features(config)
    assert "pipeline parallelism with async scheduling" in unsupported


def test_data_parallel_rpc_port_has_fixed_default():
    assert ParallelConfig().data_parallel_rpc_port == 29550


def test_all2all_backend_has_portable_default():
    assert ParallelConfig().all2all_backend == "allgather_reducescatter"


def test_dp_group_uses_configured_timeout_without_current_config(monkeypatch):
    monkeypatch.setattr(vllm_config_module, "_current_vllm_config", None)
    config = ParallelConfig(cpu_distributed_timeout_seconds=30)
    with (
        patch(
            "vllm.distributed.utils.rendezvous",
            return_value=iter([(torch.distributed.HashStore(), 0, 1)]),
        ),
        patch("vllm.distributed.utils.init_gloo_process_group") as init_group,
    ):
        config.stateless_init_dp_group()
    assert init_group.call_args.kwargs["timeout"] == timedelta(seconds=30)


@pytest.mark.parametrize(
    "dp_size, across_dp, expected",
    [(1, False, 4), (1, True, 4), (2, False, 4), (2, True, 8)],
)
def test_engram_tensor_parallel_size(dp_size: int, across_dp: bool, expected: int):
    parallel = ParallelConfig(
        tensor_parallel_size=4,
        data_parallel_size=dp_size,
        distributed_executor_backend="mp",
    )
    config = EngramConfig(embedding_across_dp=across_dp)
    assert config.get_parallel_size(parallel) == expected


@pytest.mark.parametrize("option", ["embedding_across_dp", "dp_shared_memory"])
def test_engram_rejects_elastic_cross_dp(option):
    parallel = ParallelConfig(
        tensor_parallel_size=4,
        data_parallel_size=2,
        distributed_executor_backend="mp",
    )
    parallel.enable_elastic_ep = True
    with pytest.raises(ValueError, match=f"{option}.*elastic EP"):
        EngramConfig(**{option: True}).verify_parallel_config(parallel)


def test_engram_dp_shared_memory_requires_cpu_offload():
    with pytest.raises(ValueError, match="requires cpu_offload"):
        EngramConfig(cpu_offload=False, dp_shared_memory=True)


@pytest.mark.skip_global_cleanup
@pytest.mark.parametrize(
    "cpu_offload,use_thp,dp_shared_memory,dp_size,elastic_ep,expected",
    [
        (True, False, None, 2, False, True),
        (False, False, None, 2, False, False),
        (True, False, None, 1, False, False),
        (True, False, None, 2, True, False),
        # use_thp backs private tables; sharing would silently ignore it.
        (True, True, None, 2, False, False),
        (True, True, False, 2, False, False),
    ],
)
def test_engram_dp_shared_memory_defaults_when_supported(
    cpu_offload, use_thp, dp_shared_memory, dp_size, elastic_ep, expected
):
    """Resolve sharing only where supported, preserving an explicit false."""
    parallel = ParallelConfig(data_parallel_size=dp_size)
    parallel.enable_elastic_ep = elastic_ep
    config = EngramConfig(
        cpu_offload=cpu_offload,
        use_thp=use_thp,
        dp_shared_memory=dp_shared_memory,
    )
    config.resolve_dp_shared_memory(parallel)
    assert config.dp_shared_memory is expected


@pytest.mark.skip_global_cleanup
def test_engram_thp_rejects_explicit_shared_memory():
    with pytest.raises(
        ValueError, match="use_thp requires cpu_offload=True and dp_shared_memory=False"
    ):
        EngramConfig(use_thp=True, dp_shared_memory=True)


@pytest.mark.parametrize(
    "dp_size,load_format,multithread,error",
    [
        (1, "auto", False, "requires data_parallel_size > 1"),
        (2, "auto", False, None),
        (2, "safetensors", True, None),
    ],
)
def test_engram_dp_shared_memory_config_validation(
    monkeypatch, dp_size, load_format, multithread, error
):
    """Reject invalid shared configs before distributed init; allow threaded loads."""
    monkeypatch.setattr(current_platform, "is_cuda_alike", lambda: True)
    config = cast(
        VllmConfig,
        SimpleNamespace(
            model_config=SimpleNamespace(
                architecture="DeepseekV41ForCausalLM",
                hf_text_config=SimpleNamespace(engram_layer_ids=[1]),
            ),
            speculative_config=None,
            engram_config=EngramConfig(cpu_offload=True, dp_shared_memory=True),
            parallel_config=ParallelConfig(data_parallel_size=dp_size),
            load_config=LoadConfig(
                load_format=load_format,
                model_loader_extra_config={"enable_multithread_load": multithread},
            ),
        ),
    )
    if error:
        with pytest.raises(ValueError, match=error):
            VllmConfig._resolve_and_verify_engram_config(config)
    else:
        VllmConfig._resolve_and_verify_engram_config(config)


@pytest.mark.parametrize(
    "architecture, ple_layers, platform, supported",
    [
        ("DeepseekV41ForCausalLM", [1], "cuda", True),
        ("DeepseekV41ForCausalLM", [1], "rocm", True),
        ("DeepseekV41ForCausalLM", [], "cuda", False),
        ("DeepseekV41ForCausalLM", [1], "cpu", True),
        ("Qwen4ExpForCausalLM", [1], "cuda", True),
        ("Qwen4ExpForCausalLM", [1], "rocm", True),
        ("Qwen4ExpForConditionalGeneration", [1], "cuda", True),
        ("Qwen4ExpForConditionalGeneration", [1], "rocm", True),
        ("Qwen4ExpForCausalLM", [], "cuda", False),
        ("Qwen4ExpForCausalLM", None, "cuda", False),
        ("Qwen4ExpForCausalLM", [1], "cpu", True),
        ("LlamaForCausalLM", [1], "cuda", False),
        ("Qwen4ExpMTP", [], "cuda", False),
        (None, None, "cuda", False),
    ],
)
def test_engram_model_support(
    monkeypatch, architecture, ple_layers, platform, supported
):
    """A similarly named HF field must not enable unsupported implementations."""
    monkeypatch.setattr(current_platform, "is_cuda", lambda: platform == "cuda")
    monkeypatch.setattr(
        current_platform, "is_cuda_alike", lambda: platform in ("cuda", "rocm")
    )
    model = (
        cast(
            ModelConfig,
            SimpleNamespace(
                architecture=architecture,
                hf_text_config=SimpleNamespace(
                    ple_layer_ids=ple_layers, engram_layer_ids=ple_layers
                ),
            ),
        )
        if architecture is not None
        else None
    )
    config = EngramConfig(cpu_offload=False, embedding_across_dp=False)
    if supported:
        config.verify_model_config(model)
    else:
        with pytest.raises(ValueError, match="requires a model with supported Engram"):
            config.verify_model_config(model)

    resolved = cast(
        VllmConfig,
        SimpleNamespace(
            model_config=model,
            speculative_config=None,
            parallel_config=ParallelConfig(),
            load_config=LoadConfig(load_format="dummy"),
            engram_config=None,
        ),
    )
    VllmConfig._resolve_and_verify_engram_config(resolved)
    assert (resolved.engram_config is not None) == supported
    if supported:
        assert resolved.engram_config.cpu_offload is True


def test_engram_cpu_offload_default():
    assert EngramConfig().cpu_offload is True


def test_engram_config_defaults_to_none():
    config = VllmConfig()
    assert config.engram_config is None
    assert config.compute_hash()


@pytest.mark.parametrize("explicit", [False, True])
@pytest.mark.parametrize("enable_dbo,ubatch_size", [(True, 0), (False, 2)])
def test_deepseek_engram_rejects_microbatching(explicit, enable_dbo, ubatch_size):
    """Reject overlapping Engram staging even without an explicit EngramConfig."""
    config = cast(
        VllmConfig,
        SimpleNamespace(
            model_config=SimpleNamespace(
                architecture="DeepseekV41ForCausalLM",
                hf_text_config=SimpleNamespace(engram_layer_ids=[1]),
            ),
            speculative_config=None,
            engram_config=EngramConfig(cpu_offload=False) if explicit else None,
            parallel_config=ParallelConfig(
                enable_dbo=enable_dbo, ubatch_size=ubatch_size
            ),
        ),
    )
    with pytest.raises(
        ValueError, match="Engram does not support DBO or microbatching"
    ):
        VllmConfig._resolve_and_verify_engram_config(config)


def test_engram_explicit_config_requires_supported_model():
    """Explicit all-false settings still opt into model validation."""
    with pytest.raises(ValueError, match="requires a model with supported Engram"):
        VllmConfig(engram_config=EngramConfig(cpu_offload=False))


@pytest.mark.parametrize("target_has_ple", [False, True])
@pytest.mark.parametrize("explicit", [False, True])
def test_engram_draft_config_validates_target(monkeypatch, target_has_ple, explicit):
    """MTP may inherit cross-DP sharding without having its own PLE layers."""
    monkeypatch.setattr(current_platform, "is_cuda_alike", lambda: True)
    target = SimpleNamespace(
        architecture="Qwen4ExpForCausalLM",
        hf_text_config=SimpleNamespace(ple_layer_ids=[1] if target_has_ple else []),
    )
    draft = SimpleNamespace(architecture="Qwen4ExpMTP")
    config = cast(
        VllmConfig,
        SimpleNamespace(
            model_config=draft,
            speculative_config=SimpleNamespace(
                draft_model_config=draft, target_model_config=target
            ),
            engram_config=EngramConfig(embedding_across_dp=True) if explicit else None,
            parallel_config=ParallelConfig(),
            load_config=LoadConfig(load_format="dummy"),
        ),
    )
    if explicit and not target_has_ple:
        with pytest.raises(ValueError, match="requires a model with supported Engram"):
            VllmConfig._resolve_and_verify_engram_config(config)
    else:
        VllmConfig._resolve_and_verify_engram_config(config)
        assert (config.engram_config is not None) == target_has_ple
        if target_has_ple:
            assert config.engram_config.cpu_offload is True


def test_engram_hash_tracks_execution_options():
    configs = [
        EngramConfig(cpu_offload=True),
        EngramConfig(cpu_offload=False),
        EngramConfig(cpu_offload=True, embedding_across_dp=True),
        EngramConfig(cpu_offload=True, dp_shared_memory=True),
    ]
    assert len({config.compute_hash() for config in configs}) == len(configs)


@pytest.mark.parametrize("port", [1, 29550, 65535])
def test_data_parallel_rpc_port_accepts_valid_ports(port: int):
    assert ParallelConfig(data_parallel_rpc_port=port).data_parallel_rpc_port == port


@pytest.mark.parametrize("port", [-1, 0, 65536])
def test_data_parallel_rpc_port_rejects_invalid_ports(port: int):
    with pytest.raises(ValidationError):
        ParallelConfig(data_parallel_rpc_port=port)


def test_reconfigure_for_independent_dp_rank_on_multinode_dense_model():
    parallel_config = ParallelConfig(
        tensor_parallel_size=8,
        data_parallel_size=2,
        data_parallel_size_local=1,
        data_parallel_rank=1,
        distributed_executor_backend="mp",
        nnodes=2,
        node_rank=1,
    )

    assert parallel_config.nnodes_within_dp == 1
    assert parallel_config.node_rank_within_dp == 0

    parallel_config.reconfigure_for_independent_dp_rank()

    assert parallel_config.data_parallel_size == 1
    assert parallel_config.data_parallel_size_local == 1
    assert parallel_config.data_parallel_rank == 0
    assert parallel_config.data_parallel_index == 1
    assert parallel_config.nnodes == 1
    assert parallel_config.node_rank == 0
    assert parallel_config.world_size == 8


@pytest.mark.parametrize("data_parallel_size", [4, 8])
def test_nnodes_within_dp_when_replicas_outnumber_nodes(data_parallel_size):
    """A replica that fits on one node spans one node, never zero.

    External LB pins ``data_parallel_size_local`` to 1, so the ratio rounds
    down to 0 as soon as there are more DP replicas than nodes. Anything
    dividing by ``nnodes_within_dp`` then raises ZeroDivisionError.
    """
    parallel_config = ParallelConfig(
        data_parallel_size=data_parallel_size,
        data_parallel_size_local=1,
        data_parallel_external_lb=True,
        distributed_executor_backend="mp",
        nnodes=2,
        node_rank=1,
    )

    assert parallel_config.nnodes_within_dp == 1
    assert parallel_config.node_rank_within_dp == 0
    assert parallel_config.local_world_size == parallel_config.world_size


@pytest.mark.parametrize("nnodes", [4, 9])
def test_nnodes_within_dp_rejects_uneven_internal_lb(nnodes):
    parallel_config = ParallelConfig(
        data_parallel_size=8,
        data_parallel_size_local=1,
        distributed_executor_backend="mp",
        nnodes=nnodes,
    )

    with pytest.raises(ValueError, match="Invalid data parallel configuration"):
        _ = parallel_config.nnodes_within_dp


def test_draft_model_enables_async_scheduling_by_default():
    parallel_config = ParallelConfig(distributed_executor_backend="uni")
    model_config = ModelConfig("Qwen/Qwen3-0.6B", max_model_len=2048)
    speculative_config = SpeculativeConfig(
        method="draft_model",
        model="Qwen/Qwen3-0.6B",
        num_speculative_tokens=3,
        target_model_config=model_config,
        target_parallel_config=parallel_config,
    )
    cfg = VllmConfig(
        model_config=model_config,
        scheduler_config=SchedulerConfig(
            max_model_len=2048,
            is_encoder_decoder=False,
        ),
        parallel_config=parallel_config,
        speculative_config=speculative_config,
    )

    assert cfg.scheduler_config.async_scheduling is True


def test_dflash_allows_async_scheduling(tmp_path: Path):
    """DFlash must stay async-schedulable: it was the only parallel-drafting
    method excluded from the allowlist."""
    from transformers import LlamaConfig

    common = dict(
        hidden_size=128,
        intermediate_size=256,
        num_hidden_layers=2,
        num_attention_heads=4,
        num_key_value_heads=2,
        vocab_size=256,
        max_position_embeddings=2048,
    )
    target_path = tmp_path / "target"
    draft_path = tmp_path / "draft"
    _write_json(
        target_path / "config.json",
        LlamaConfig(architectures=["LlamaForCausalLM"], **common).to_dict(),
    )
    _write_json(
        draft_path / "config.json",
        dict(
            common,
            architectures=["DFlash2DraftModel"],
            model_type="qwen3",
            num_hidden_layers=1,
            layer_types=["sliding_attention"],
            sliding_window=128,
            dflash_config=dict(
                block_size=4,
                mask_token_id=255,
                target_layer_ids=[0],
                conv_group_size=2,
                conv_kernel_size=2,
                selector_rank=8,
                selector_top_k=4,
            ),
        ),
    )
    parallel_config = ParallelConfig(distributed_executor_backend="uni")
    model_config = ModelConfig(
        model=str(target_path), tokenizer_mode="skip", max_model_len=2048
    )
    speculative_config = SpeculativeConfig(
        method="dflash",
        model=str(draft_path),
        num_speculative_tokens=3,
        target_model_config=model_config,
        target_parallel_config=parallel_config,
    )
    cfg = VllmConfig(
        model_config=model_config,
        scheduler_config=SchedulerConfig(
            max_model_len=2048,
            is_encoder_decoder=False,
        ),
        parallel_config=parallel_config,
        speculative_config=speculative_config,
    )

    assert speculative_config.method == "dflash"
    assert cfg.scheduler_config.async_scheduling is True


@pytest.mark.skip_global_cleanup
@pytest.mark.parametrize("tp_size", [1, 2])
@pytest.mark.parametrize("target_ep", [False, True], ids=["ep-off", "ep-on"])
@pytest.mark.parametrize(
    ("method", "draft_is_moe"),
    [
        pytest.param("draft_model", False, id="dense-draft"),
        pytest.param("eagle", False, id="eagle"),
        pytest.param("eagle3", False, id="eagle3"),
        pytest.param("draft_model", True, id="moe-draft"),
        pytest.param("mtp", True, id="mtp"),
        pytest.param("dspark", True, id="dspark"),
    ],
)
def test_draft_inherits_ep_only_for_moe(
    tmp_path: Path,
    tp_size: int,
    target_ep: bool,
    method: str,
    draft_is_moe: bool,
):
    """Validate final draft configs without loading weights or mocking validation."""
    from transformers import LlamaConfig, MixtralConfig

    from vllm.transformers_utils.configs.deepseek_v4 import DeepseekV4Config

    common = dict(
        hidden_size=128,
        intermediate_size=256,
        num_hidden_layers=2,
        num_attention_heads=4,
        num_key_value_heads=2,
        vocab_size=128,
        max_position_embeddings=2048,
    )
    if method in ("mtp", "dspark"):
        target_hf_config = DeepseekV4Config(
            architectures=["DeepseekV4ForCausalLM"],
            n_routed_experts=4,
            num_experts_per_tok=2,
            num_nextn_predict_layers=1,
            compress_ratios=[1, 1],
            head_dim=32,
            **common,
        )
    else:
        target_hf_config = MixtralConfig(
            architectures=["MixtralForCausalLM"], num_local_experts=4, **common
        )
    target_path = tmp_path / "target"
    draft_path = tmp_path / "draft"
    _write_json(target_path / "config.json", target_hf_config.to_dict())
    _write_json(
        draft_path / "config.json",
        LlamaConfig(architectures=["LlamaForCausalLM"], **common).to_dict(),
    )
    target_model_config = ModelConfig(
        model=str(target_path), tokenizer_mode="skip", max_model_len=2048
    )
    target_parallel_config = ParallelConfig(
        tensor_parallel_size=tp_size,
        enable_expert_parallel=target_ep,
        distributed_executor_backend="mp",
    )
    target_model_config.verify_with_parallel_config(target_parallel_config)
    speculative_config = SpeculativeConfig(
        method=method,
        model=str(target_path if draft_is_moe else draft_path),
        num_speculative_tokens=1,
        target_model_config=target_model_config,
        target_parallel_config=target_parallel_config,
    )

    assert speculative_config.method == method
    assert speculative_config.draft_model_config.is_moe is draft_is_moe
    assert speculative_config.draft_parallel_config.enable_expert_parallel is (
        target_ep and draft_is_moe
    )
    assert speculative_config.draft_parallel_config.tensor_parallel_size == tp_size
    assert target_parallel_config.enable_expert_parallel is target_ep
    assert target_parallel_config.tensor_parallel_size == tp_size


@pytest.mark.skip_global_cleanup
@pytest.mark.parametrize("target_ep", [False, True], ids=["ep-off", "ep-on"])
@pytest.mark.parametrize("pass_none", [False, True], ids=["omitted", "explicit-none"])
def test_draft_parallel_config_preserves_ep_without_model(
    target_ep: bool, pass_none: bool
):
    """Legacy callers without draft model information keep EP inheritance."""
    target_parallel_config = ParallelConfig(
        tensor_parallel_size=2,
        enable_expert_parallel=target_ep,
        distributed_executor_backend="mp",
    )
    if pass_none:
        draft_parallel_config = SpeculativeConfig.create_draft_parallel_config(
            target_parallel_config, 2, draft_model_config=None
        )
    else:
        draft_parallel_config = SpeculativeConfig.create_draft_parallel_config(
            target_parallel_config, 2
        )

    assert draft_parallel_config.enable_expert_parallel is target_ep
    assert draft_parallel_config.tensor_parallel_size == 2
    assert target_parallel_config.enable_expert_parallel is target_ep


@pytest.mark.parametrize(
    ("method", "parallel_drafting", "expected_slots"),
    [
        pytest.param("eagle3", False, 0, id="eagle3"),
        pytest.param("eagle3", True, 7, id="p-eagle"),
        pytest.param("dflash", True, 8, id="dflash"),
        pytest.param("dspark", True, 7, id="dspark"),
        pytest.param("mtp", False, 0, id="mtp"),
        pytest.param("ngram", False, 0, id="ngram"),
        pytest.param("draft_model", False, 1, id="draft-model"),
        pytest.param("draft_model", True, 8, id="pard"),
    ],
)
def test_max_num_new_slots_for_drafting(method, parallel_drafting, expected_slots):
    speculative_config = SpeculativeConfig(
        model="ngram",
        num_speculative_tokens=8,
    )
    speculative_config.method = method
    speculative_config.parallel_drafting = parallel_drafting

    assert speculative_config.max_num_new_slots_for_drafting == expected_slots


@dataclass
class _TestConfigFields:
    a: int
    b: dict = field(default_factory=dict)
    c: str = "default"


def test_get_field():
    b = get_field(_TestConfigFields, "b")
    assert isinstance(b, Field)
    assert b.default is MISSING
    assert b.default_factory is dict

    c = get_field(_TestConfigFields, "c")
    assert isinstance(c, Field)
    assert c.default == "default"
    assert c.default_factory is MISSING


@dataclass
class _TestNestedConfig:
    a: _TestConfigFields = field(default_factory=lambda: _TestConfigFields(a=0))


@dataclass
class _TestDerivedConfigFields(_TestConfigFields):
    pass


def test_update_config():
    # Simple update
    config1 = _TestConfigFields(a=0)
    new_config1 = update_config(config1, {"a": 42})
    assert new_config1.a == 42
    # Nonexistent field
    with pytest.raises(ValueError, match=r"_TestConfigFields\.nonexistent"):
        new_config1 = update_config(config1, {"nonexistent": 1})
    # Nested update with dataclass
    config2 = _TestNestedConfig()
    new_inner_config = _TestConfigFields(a=1, c="new_value")
    new_config2 = update_config(config2, {"a": new_inner_config})
    assert new_config2.a == new_inner_config
    # Declared field type, not the live value's subtype, defines valid overrides
    config_with_derived = _TestNestedConfig(a=_TestDerivedConfigFields(a=0))
    new_config2 = update_config(config_with_derived, {"a": new_inner_config})
    assert new_config2.a is new_inner_config
    # Nested update with unrelated dataclass
    with pytest.raises(ValueError, match=r"_TestNestedConfig\.a"):
        update_config(config2, {"a": _TestNestedConfig()})
    # Nested update with dict
    config3 = _TestNestedConfig()
    new_config3 = update_config(config3, {"a": {"c": "new_value"}})
    assert new_config3.a.c == "new_value"
    # Nested update with invalid type
    with pytest.raises(ValueError, match=r"_TestNestedConfig\.a"):
        update_config(config3, {"a": "new_value"})
    # Invalid nested field preserves its full path
    with pytest.raises(ValueError, match=r"_TestNestedConfig\.a\.nonexistent"):
        update_config(config3, {"a": {"nonexistent": 1}})


@pytest.mark.parametrize(
    ("model_id", "expected_runner_type", "expected_convert_type"),
    [
        ("distilbert/distilgpt2", "generate", "none"),
        ("intfloat/multilingual-e5-small", "pooling", "none"),
        ("jason9693/Qwen2.5-1.5B-apeach", "pooling", "classify"),
        ("cross-encoder/ms-marco-MiniLM-L-6-v2", "pooling", "none"),
        ("Qwen/Qwen2.5-Math-RM-72B", "pooling", "none"),
        ("openai/whisper-small", "generate", "none"),
    ],
)
def test_auto_runner(model_id, expected_runner_type, expected_convert_type):
    config = ModelConfig(model_id, runner="auto")

    assert config.runner_type == expected_runner_type
    assert config.convert_type == expected_convert_type


@pytest.mark.parametrize(
    ("model_id", "expected_runner_type", "expected_convert_type"),
    [
        ("distilbert/distilgpt2", "pooling", "embed"),
        ("intfloat/multilingual-e5-small", "pooling", "none"),
        ("jason9693/Qwen2.5-1.5B-apeach", "pooling", "classify"),
        ("cross-encoder/ms-marco-MiniLM-L-6-v2", "pooling", "none"),
        ("Qwen/Qwen2.5-Math-RM-72B", "pooling", "none"),
        ("openai/whisper-small", "pooling", "embed"),
    ],
)
def test_pooling_runner(model_id, expected_runner_type, expected_convert_type):
    config = ModelConfig(model_id, runner="pooling")

    assert config.runner_type == expected_runner_type
    assert config.convert_type == expected_convert_type


@pytest.mark.parametrize(
    ("model_id", "expected_runner_type", "expected_convert_type"),
    [
        ("Qwen/Qwen2.5-1.5B-Instruct", "draft", "none"),
    ],
)
def test_draft_runner(model_id, expected_runner_type, expected_convert_type):
    config = ModelConfig(model_id, runner="draft")

    assert config.runner_type == expected_runner_type
    assert config.convert_type == expected_convert_type


MODEL_IDS_EXPECTED = [
    ("Qwen/Qwen1.5-7B", 32768),
    ("mistralai/Mistral-7B-v0.1", 4096),
    ("mistralai/Mistral-7B-Instruct-v0.2", 32768),
]


@pytest.mark.parametrize("model_id_expected", MODEL_IDS_EXPECTED)
def test_disable_sliding_window(model_id_expected):
    model_id, expected = model_id_expected
    model_config = ModelConfig(model_id, disable_sliding_window=True)
    assert model_config.max_model_len == expected


@pytest.mark.skipif(
    current_platform.is_rocm(), reason="Xformers backend is not supported on ROCm."
)
def test_get_pooling_config():
    model_id = "sentence-transformers/all-MiniLM-L12-v2"
    model_config = ModelConfig(model_id)

    assert model_config.pooler_config is not None
    assert model_config.pooler_config.use_activation
    assert model_config.pooler_config.seq_pooling_type == "MEAN"
    assert model_config.pooler_config.tok_pooling_type == "ALL"


@pytest.mark.parametrize(
    ("pooling_module_type", "normalize_module_type", "pooling_config", "expected"),
    [
        (
            "sentence_transformers.models.Pooling",
            "sentence_transformers.models.Normalize",
            {"pooling_mode_mean_tokens": True},
            {"use_activation": True, "seq_pooling_type": "MEAN"},
        ),
        (
            "sentence_transformers.sentence_transformer.modules.pooling.Pooling",
            "sentence_transformers.sentence_transformer.modules.normalize.Normalize",
            {"pooling_mode": "lasttoken"},
            {"use_activation": True, "seq_pooling_type": "LAST"},
        ),
        (
            "sentence_transformers.sentence_transformer.modules.pooling.Pooling",
            "sentence_transformers.base.modules.normalize.Normalize",
            {"pooling_mode": "mean"},
            {"use_activation": True, "seq_pooling_type": "MEAN"},
        ),
    ],
)
def test_get_pooling_config_supports_sentence_transformers_schemas(
    tmp_path,
    caplog,
    pooling_module_type,
    normalize_module_type,
    pooling_config,
    expected,
):
    modules = [
        {"idx": 1, "name": "1", "path": "1_Pooling", "type": pooling_module_type},
        {
            "idx": 2,
            "name": "2",
            "path": "2_Normalize",
            "type": normalize_module_type,
        },
    ]
    _write_json(tmp_path / "modules.json", modules)
    _write_json(tmp_path / "1_Pooling" / "config.json", pooling_config)

    with caplog.at_level(logging.WARNING):
        config = get_pooling_config(str(tmp_path))

    assert config == expected
    assert "Unable to determine Sentence Transformers pooling type" not in caplog.text


def test_get_pooling_config_warns_when_pooling_mode_is_unknown(tmp_path, caplog):
    modules = [
        {
            "idx": 1,
            "name": "1",
            "path": "1_Pooling",
            "type": "sentence_transformers.models.Pooling",
        }
    ]
    _write_json(tmp_path / "modules.json", modules)
    _write_json(
        tmp_path / "1_Pooling" / "config.json",
        {"pooling_mode": "unsupported"},
    )

    with caplog.at_level(logging.WARNING):
        config = get_pooling_config(str(tmp_path))

    assert config == {"use_activation": False}
    assert "Unable to determine Sentence Transformers pooling type" in caplog.text


@pytest.mark.parametrize(
    "dense_module_type",
    [
        "sentence_transformers.models.Dense",
        "sentence_transformers.base.modules.dense.Dense",
    ],
)
def test_get_dense_modules_supports_sentence_transformers_schemas(
    tmp_path, dense_module_type
):
    modules = [{"idx": 2, "name": "2", "path": "2_Dense", "type": dense_module_type}]
    dense_config = {
        "in_features": 768,
        "out_features": 3072,
        "bias": False,
        "activation_function": "torch.nn.modules.linear.Identity",
    }
    _write_json(tmp_path / "modules.json", modules)
    _write_json(tmp_path / "2_Dense" / "config.json", dense_config)

    assert try_get_dense_modules(str(tmp_path)) == [
        {**dense_config, "folder": "2_Dense"}
    ]


@pytest.mark.skipif(
    current_platform.is_rocm(), reason="Xformers backend is not supported on ROCm."
)
def test_get_pooling_config_from_args():
    model_id = "sentence-transformers/all-MiniLM-L12-v2"
    pooler_config = PoolerConfig(seq_pooling_type="CLS", use_activation=False)
    model_config = ModelConfig(model_id, pooler_config=pooler_config)

    assert asdict(model_config.pooler_config) == asdict(pooler_config)


@pytest.mark.parametrize(
    ("model_id", "default_pooling_type", "pooling_type"),
    [
        ("tomaarsen/Qwen3-Reranker-0.6B-seq-cls", "LAST", "LAST"),  # LLM
        ("intfloat/e5-small", "CLS", "MEAN"),  # BertModel
    ],
)
def test_default_seq_pooling_type(model_id, default_pooling_type, pooling_type):
    model_config = ModelConfig(model_id)
    assert model_config._model_info.default_seq_pooling_type == default_pooling_type
    assert model_config.pooler_config.seq_pooling_type == pooling_type


@pytest.mark.parametrize(
    ("model_id", "default_pooling_type", "pooling_type"),
    [
        ("Qwen/Qwen2.5-Math-RM-72B", "ALL", "ALL"),  # reward
        ("Qwen/Qwen2.5-Math-PRM-7B", "STEP", "STEP"),  # step reward
    ],
)
def test_default_tok_pooling_type(model_id, default_pooling_type, pooling_type):
    model_config = ModelConfig(model_id)
    assert model_config._model_info.default_tok_pooling_type == default_pooling_type
    assert model_config.pooler_config.tok_pooling_type == pooling_type


@pytest.mark.parametrize(
    ("model_id", "expected_is_moe_model"),
    [
        ("RedHatAI/Qwen3-8B-speculator.eagle3", False),
        ("RedHatAI/Llama-3.1-8B-Instruct-NVFP4", False),
        ("RedHatAI/Llama-3.2-1B-FP8", False),
        ("RedHatAI/Mistral-Small-24B-Instruct-2501-quantized.w8a8", False),
        ("RedHatAI/gpt-oss-20b", True),
        ("RedHatAI/DeepSeek-V2.5-1210-FP8", True),
        ("RedHatAI/Llama-4-Scout-17B-16E-Instruct", True),
        ("RedHatAI/Mixtral-8x7B-Instruct-v0.1", True),
    ],
)
def test_moe_model_detection(model_id, expected_is_moe_model):
    model_config = ModelConfig(model_id)
    # Just check that is_moe field exists and is a boolean
    assert model_config.is_moe == expected_is_moe_model


@pytest.mark.parametrize(
    ("model_id", "quantized"),
    [
        ("RedHatAI/Qwen3-8B-speculator.eagle3", False),
        ("RedHatAI/Llama-3.1-8B-Instruct-NVFP4", True),
        ("RedHatAI/Llama-3.2-1B-FP8", True),
        ("RedHatAI/Mistral-Small-24B-Instruct-2501-quantized.w8a8", True),
        ("RedHatAI/gpt-oss-20b", True),
        ("RedHatAI/DeepSeek-V2.5-1210-FP8", True),
        ("RedHatAI/Mixtral-8x7B-Instruct-v0.1", False),
    ],
)
def test_is_quantized(model_id, quantized):
    model_config = ModelConfig(model_id)
    # Just check that quantized field exists and is a boolean
    assert model_config.is_quantized == quantized


@pytest.mark.skipif(
    current_platform.is_rocm(), reason="Xformers backend is not supported on ROCm."
)
def test_get_bert_tokenization_sentence_transformer_config():
    model_id = "BAAI/bge-base-en-v1.5"
    bge_model_config = ModelConfig(model_id)

    bert_bge_model_config = bge_model_config._get_encoder_config()

    assert bert_bge_model_config["max_seq_length"] == 512
    assert bert_bge_model_config["do_lower_case"]


def test_rope_customization():
    TEST_ROPE_PARAMETERS = {
        "rope_theta": 16_000_000.0,
        "rope_type": "dynamic",
        "factor": 2.0,
    }
    LLAMA_ROPE_PARAMETERS = {"rope_theta": 500000.0, "rope_type": "default"}
    LONGCHAT_ROPE_PARAMETERS = {"rope_type": "linear", "factor": 8.0}

    llama_model_config = ModelConfig("meta-llama/Meta-Llama-3-8B-Instruct")
    assert (
        getattr(llama_model_config.hf_config, "rope_parameters", None)
        == LLAMA_ROPE_PARAMETERS
    )
    assert llama_model_config.max_model_len == 8192

    llama_model_config = ModelConfig(
        "meta-llama/Meta-Llama-3-8B-Instruct",
        hf_overrides={"rope_parameters": TEST_ROPE_PARAMETERS},
    )
    assert (
        getattr(llama_model_config.hf_config, "rope_parameters", None)
        == TEST_ROPE_PARAMETERS
    )
    assert llama_model_config.max_model_len == 16384

    longchat_model_config = ModelConfig("lmsys/longchat-13b-16k")
    # Check if LONGCHAT_ROPE_PARAMETERS entries are in longchat_model_config
    assert all(
        longchat_model_config.hf_config.rope_parameters.get(key) == value
        for key, value in LONGCHAT_ROPE_PARAMETERS.items()
    )
    assert longchat_model_config.max_model_len == 16384

    longchat_model_config = ModelConfig(
        "lmsys/longchat-13b-16k",
        hf_overrides={
            "rope_parameters": TEST_ROPE_PARAMETERS,
        },
    )
    assert (
        getattr(longchat_model_config.hf_config, "rope_parameters", None)
        == TEST_ROPE_PARAMETERS
    )
    assert longchat_model_config.max_model_len == 4096


def test_nested_hf_overrides():
    """Test that nested hf_overrides work correctly."""
    # Test with a model that has text_config
    model_config = ModelConfig(
        "Qwen/Qwen2-VL-2B-Instruct",
        hf_overrides={
            "text_config": {
                "hidden_size": 1024,
            },
        },
    )
    assert model_config.hf_config.text_config.hidden_size == 1024

    # Test with deeply nested overrides
    model_config = ModelConfig(
        "Qwen/Qwen2-VL-2B-Instruct",
        hf_overrides={
            "text_config": {
                "hidden_size": 2048,
                "num_attention_heads": 16,
            },
            "vision_config": {
                "hidden_size": 512,
            },
        },
    )
    assert model_config.hf_config.text_config.hidden_size == 2048
    assert model_config.hf_config.text_config.num_attention_heads == 16
    assert model_config.hf_config.vision_config.hidden_size == 512


def test_model_class_overrides_registers_target():
    """`model_class_overrides` redirects an architecture to a custom class."""
    from vllm.model_executor.models import ModelRegistry

    arch = "_TestModelClassOverrideArch"
    target = "vllm.model_executor.models.llama:LlamaForCausalLM"
    assert arch not in ModelRegistry.models

    model_config = ModelConfig(
        "facebook/opt-125m",
        model_class_overrides={arch: target},
    )
    try:
        # Accessing `.registry` is the chokepoint that applies the overrides;
        # it has already run during construction.
        registered = model_config.registry.models[arch]
        assert registered.module_name == "vllm.model_executor.models.llama"
        assert registered.class_name == "LlamaForCausalLM"
        # Idempotent: a second access does not re-register or error out.
        assert model_config.registry.models[arch] is registered
    finally:
        ModelRegistry.models.pop(arch, None)


@pytest.mark.skipif(
    current_platform.is_rocm(), reason="Encoder Decoder models not supported on ROCm."
)
@pytest.mark.parametrize(
    ("model_id", "is_encoder_decoder"),
    [
        ("facebook/opt-125m", False),
        ("openai/whisper-tiny", True),
        ("meta-llama/Llama-3.2-1B-Instruct", False),
    ],
)
def test_is_encoder_decoder(model_id, is_encoder_decoder):
    config = ModelConfig(model_id)

    assert config.is_encoder_decoder == is_encoder_decoder


@pytest.mark.parametrize(
    ("model_id", "uses_mrope"),
    [
        ("facebook/opt-125m", False),
        ("Qwen/Qwen2-VL-2B-Instruct", True),
    ],
)
def test_uses_mrope(model_id, uses_mrope):
    config = ModelConfig(model_id)

    assert config.uses_mrope == uses_mrope


def test_generation_config_loading():
    model_id = "Qwen/Qwen2.5-1.5B-Instruct"

    # When set generation_config to "vllm", the default generation config
    # will not be loaded.
    model_config = ModelConfig(model_id, generation_config="vllm")
    assert model_config.get_diff_sampling_param() == {}

    # When set generation_config to "auto", the default generation config
    # should be loaded.
    model_config = ModelConfig(model_id, generation_config="auto")

    correct_generation_config = {
        "repetition_penalty": 1.1,
        "temperature": 0.7,
        "top_p": 0.8,
        "top_k": 20,
    }

    assert model_config.get_diff_sampling_param() == correct_generation_config

    # The generation config could be overridden by the user.
    override_generation_config = {"temperature": 0.5, "top_k": 5}

    model_config = ModelConfig(
        model_id,
        generation_config="auto",
        override_generation_config=override_generation_config,
    )

    override_result = correct_generation_config.copy()
    override_result.update(override_generation_config)

    assert model_config.get_diff_sampling_param() == override_result

    # When generation_config is set to "vllm" and override_generation_config
    # is set, the override_generation_config should be used directly.
    model_config = ModelConfig(
        model_id,
        generation_config="vllm",
        override_generation_config=override_generation_config,
    )

    assert model_config.get_diff_sampling_param() == override_generation_config


@pytest.mark.parametrize(
    "pt_load_map_location",
    [
        DEVICE_TYPE,
        {"": DEVICE_TYPE},
    ],
)
def test_load_config_pt_load_map_location(pt_load_map_location):
    load_config = LoadConfig(pt_load_map_location=pt_load_map_location)
    config = VllmConfig(load_config=load_config)

    assert config.load_config.pt_load_map_location == pt_load_map_location


@pytest.mark.parametrize(
    ("model_id", "max_model_len", "expected_max_len", "should_raise"),
    [
        ("BAAI/bge-reranker-base", None, 512, False),
        ("BAAI/bge-reranker-base", 256, 256, False),
        ("BAAI/bge-reranker-base", 513, 512, True),
        ("deepseek-ai/DeepSeek-R1-Distill-Qwen-7B", None, 131072, False),
        ("deepseek-ai/DeepSeek-R1-Distill-Qwen-7B", 131073, 131072, True),
    ],
)
def test_get_and_verify_max_len(
    model_id, max_model_len, expected_max_len, should_raise
):
    """Test get_and_verify_max_len with different configurations."""
    model_config = ModelConfig(model_id)

    if should_raise:
        with pytest.raises(ValueError):
            model_config.get_and_verify_max_len(max_model_len)
    else:
        actual_max_len = model_config.get_and_verify_max_len(max_model_len)
        assert actual_max_len == expected_max_len


@pytest.mark.parametrize("max_model_len", [None, 1024])
@pytest.mark.parametrize(
    ("rope_parameters", "expected_max_len"),
    [
        ({"rope_type": "default"}, 4096),
        ({"rope_type": "linear", "factor": 2.0}, 8192),
        ({"rope_type": "longrope"}, 2048),
    ],
)
def test_get_and_verify_max_len_with_nope_layers(
    max_model_len, rope_parameters, expected_max_len
):
    """NoPE layers do not prevent deriving or scaling the context length."""
    from transformers import PreTrainedConfig

    from vllm.config.model import _get_and_verify_max_len
    from vllm.transformers_utils.model_arch_config_convertor import (
        ModelArchConfigConvertorBase,
    )

    hf_config = PreTrainedConfig(
        max_position_embeddings=4096,
        original_max_position_embeddings=2048,
    )
    hf_config.rope_parameters = {
        "full_attention": None,
        "sliding_attention": rope_parameters,
    }
    model_arch_config = ModelArchConfigConvertorBase(hf_config, hf_config).convert()

    actual_max_len = _get_and_verify_max_len(
        hf_config=hf_config,
        model_arch_config=model_arch_config,
        tokenizer_config=None,
        max_model_len=max_model_len,
        disable_sliding_window=False,
        sliding_window=None,
    )

    assert actual_max_len == (max_model_len or expected_max_len)
    assert hf_config.rope_parameters["full_attention"] is None


@pytest.mark.parametrize(
    ("rope_type", "factor", "expected_max_len"),
    [
        # TeleChat3-36B-Thinking: 32768 already scaled from 8192 by 4
        ("yarn", 4.0, 32768),
        # sarvam-105b: declares factor 40 but only serves 131072 of it
        ("deepseek_yarn", 40.0, 32768),
        ("deepseek_llama_scaling", 40.0, 32768),
        # Non-YaRN scaling still multiplies
        ("linear", 4.0, 131072),
    ],
)
def test_get_and_verify_max_len_yarn_is_already_scaled(
    rope_type, factor, expected_max_len
):
    """YaRN variants must not re-apply `factor` to max_position_embeddings.

    Transformers treats max_position_embeddings as the final context length
    for every YaRN variant, so scaling it again overstates the limit and lets
    requests past the end of the cos/sin cache.
    """
    from transformers import PreTrainedConfig

    from vllm.config.model import _get_and_verify_max_len
    from vllm.transformers_utils.model_arch_config_convertor import (
        ModelArchConfigConvertorBase,
    )

    hf_config = PreTrainedConfig(max_position_embeddings=32768)
    hf_config.rope_parameters = {
        "rope_type": rope_type,
        "factor": factor,
        "original_max_position_embeddings": 8192,
    }
    model_arch_config = ModelArchConfigConvertorBase(hf_config, hf_config).convert()

    actual_max_len = _get_and_verify_max_len(
        hf_config=hf_config,
        model_arch_config=model_arch_config,
        tokenizer_config=None,
        max_model_len=None,
        disable_sliding_window=False,
        sliding_window=None,
    )

    assert actual_max_len == expected_max_len


class MockConfig:
    """Simple mock object for testing maybe_pull_model_tokenizer_for_runai."""

    def __init__(self, model: str, tokenizer: str):
        self.model = model
        self.tokenizer = tokenizer
        self.model_weights = None


@pytest.mark.parametrize(
    "s3_url",
    [
        "s3://example-bucket-1/model/",
        "s3://example-bucket-2/model/",
    ],
)
@patch("vllm.transformers_utils.runai_utils.ObjectStorageModel.pull_files")
def test_s3_url_model_tokenizer_paths(mock_pull_files, s3_url):
    """Test that S3 URLs create deterministic local directories for model and
    tokenizer."""
    # Mock pull_files to avoid actually downloading files during tests
    mock_pull_files.return_value = None

    # Create first mock and run the method
    config1 = MockConfig(model=s3_url, tokenizer=s3_url)
    ModelConfig.maybe_pull_model_tokenizer_for_runai(config1, s3_url, s3_url)

    # Check that model and tokenizer point to existing directories
    assert os.path.exists(config1.model), (
        f"Model directory does not exist: {config1.model}"
    )
    assert os.path.isdir(config1.model), (
        f"Model path is not a directory: {config1.model}"
    )
    assert os.path.exists(config1.tokenizer), (
        f"Tokenizer directory does not exist: {config1.tokenizer}"
    )
    assert os.path.isdir(config1.tokenizer), (
        f"Tokenizer path is not a directory: {config1.tokenizer}"
    )

    # Verify that the paths are different from the original S3 URL
    assert config1.model != s3_url, "Model path should be converted to local directory"
    assert config1.tokenizer != s3_url, (
        "Tokenizer path should be converted to local directory"
    )

    # Store the original paths
    created_model_dir = config1.model
    create_tokenizer_dir = config1.tokenizer

    # Create a new mock and run the method with the same S3 URL
    config2 = MockConfig(model=s3_url, tokenizer=s3_url)
    ModelConfig.maybe_pull_model_tokenizer_for_runai(config2, s3_url, s3_url)

    # Check that the new directories exist
    assert os.path.exists(config2.model), (
        f"Model directory does not exist: {config2.model}"
    )
    assert os.path.isdir(config2.model), (
        f"Model path is not a directory: {config2.model}"
    )
    assert os.path.exists(config2.tokenizer), (
        f"Tokenizer directory does not exist: {config2.tokenizer}"
    )
    assert os.path.isdir(config2.tokenizer), (
        f"Tokenizer path is not a directory: {config2.tokenizer}"
    )

    # Verify that the paths are deterministic (same as before)
    assert config2.model == created_model_dir, (
        f"Model paths are not deterministic. "
        f"Original: {created_model_dir}, New: {config2.model}"
    )
    assert config2.tokenizer == create_tokenizer_dir, (
        f"Tokenizer paths are not deterministic. "
        f"Original: {create_tokenizer_dir}, New: {config2.tokenizer}"
    )


@patch("vllm.transformers_utils.runai_utils.ObjectStorageModel.pull_files")
def test_s3_url_different_models_create_different_directories(mock_pull_files):
    """Test that different S3 URLs create different local directories."""
    # Mock pull_files to avoid actually downloading files during tests
    mock_pull_files.return_value = None

    s3_url1 = "s3://example-bucket-1/model/"
    s3_url2 = "s3://example-bucket-2/model/"

    # Create mocks with different S3 URLs and run the method
    config1 = MockConfig(model=s3_url1, tokenizer=s3_url1)
    ModelConfig.maybe_pull_model_tokenizer_for_runai(config1, s3_url1, s3_url1)

    config2 = MockConfig(model=s3_url2, tokenizer=s3_url2)
    ModelConfig.maybe_pull_model_tokenizer_for_runai(config2, s3_url2, s3_url2)

    # Verify that different URLs produce different directories
    assert config1.model != config2.model, (
        f"Different S3 URLs should create different model directories. "
        f"URL1 model: {config1.model}, URL2 model: {config2.model}"
    )
    assert config1.tokenizer != config2.tokenizer, (
        f"Different S3 URLs should create different tokenizer directories. "
        f"URL1 tokenizer: {config1.tokenizer}, "
        f"URL2 tokenizer: {config2.tokenizer}"
    )

    # Verify that both sets of directories exist
    assert os.path.exists(config1.model) and os.path.isdir(config1.model)
    assert os.path.exists(config1.tokenizer) and os.path.isdir(config1.tokenizer)
    assert os.path.exists(config2.model) and os.path.isdir(config2.model)
    assert os.path.exists(config2.tokenizer) and os.path.isdir(config2.tokenizer)


@patch("vllm.transformers_utils.runai_utils.ObjectStorageModel.pull_files")
def test_s3_url_different_model_and_tokenizer(mock_pull_files):
    """Test that when model and tokenizer are different cloud URIs,
    pull_files receives the correct URI for each."""
    mock_pull_files.return_value = None

    model_url = "s3://bucket/model/"
    tokenizer_url = "s3://bucket/tokenizer/"

    config = MockConfig(model=model_url, tokenizer=tokenizer_url)
    ModelConfig.maybe_pull_model_tokenizer_for_runai(config, model_url, tokenizer_url)

    # pull_files should be called twice: once for model, once for tokenizer
    assert mock_pull_files.call_count == 2
    # First call: model URI with allow_pattern
    assert mock_pull_files.call_args_list[0][0][0] == model_url
    # Second call: tokenizer URI with ignore_pattern
    assert mock_pull_files.call_args_list[1][0][0] == tokenizer_url


@pytest.mark.parametrize(
    ("model_id", "expected_attn_type", "expected_result", "reason"),
    [
        # pooling models
        (
            "jason9693/Qwen2.5-1.5B-apeach",
            "decoder",
            True,
            "Pooling models with causal attn and LAST/ALL pooling support chunked prefill.",  # noqa: E501
        ),
        (
            "Qwen/Qwen3-Embedding-0.6B",
            "decoder",
            True,
            "Pooling models with causal attn and LAST/ALL pooling support chunked prefill.",  # noqa: E501
        ),
        (
            "Qwen/Qwen2.5-Math-PRM-7B",
            "decoder",
            False,
            "Pooling models with causal attn and LAST/STEP pooling do not support chunked prefill.",  # noqa: E501
        ),
        (
            "internlm/internlm2-1_8b-reward",
            "decoder",
            True,
            "Pooling models with causal attn and LAST/ALL pooling support chunked prefill.",  # noqa: E501
        ),
        (
            "BAAI/bge-base-en",
            "encoder_only",
            False,
            "Pooling models with bidirectional attn do not support chunked prefill.",  # noqa: E501
        ),
        (
            "boltuix/NeuroBERT-NER",
            "encoder_only",
            False,
            "Pooling models with bidirectional attn do not support chunked prefill.",  # noqa: E501
        ),
        (
            "papluca/xlm-roberta-base-language-detection",
            "encoder_only",
            False,
            "Pooling models with bidirectional attn do not support chunked prefill.",  # noqa: E501
        ),
        (
            "Alibaba-NLP/gte-Qwen2-1.5B-instruct",
            "encoder_only",
            False,
            "Pooling models with bidirectional attn do not support chunked prefill.",  # noqa: E501
        ),
        (
            "intfloat/e5-small",
            "encoder_only",
            False,
            "Pooling models with bidirectional attn do not support chunked prefill.",  # noqa: E501
        ),
        # multimodal models
        (
            "openai/clip-vit-base-patch32",
            "decoder",
            True,
            "Pooling models with causal attn and LAST/ALL pooling support chunked prefill.",  # noqa: E501
        ),
        (
            "google/siglip-base-patch16-224",
            "encoder_only",
            False,
            "Pooling models with bidirectional attn do not support chunked prefill.",  # noqa: E501
        ),
        # generate models
        (
            "Qwen/Qwen3-0.6B",
            "decoder",
            True,
            "Generative models support chunked prefill.",  # noqa: E501
        ),
        (
            "Qwen/Qwen3-Next-80B-A3B-Instruct",
            "hybrid",
            True,
            "Generative models support chunked prefill.",  # noqa: E501
        ),
        (
            "ibm-granite/granite-4.0-h-small",
            "hybrid",
            True,
            "Generative models support chunked prefill.",  # noqa: E501
        ),
        (
            "state-spaces/mamba-130m-hf",
            "attention_free",
            True,
            "Generative models support chunked prefill.",  # noqa: E501
        ),
        # encoder_decoder models
        (
            "openai/whisper-small",
            "encoder_decoder",
            False,
            "Encoder decoder models do not support chunked prefill.",  # noqa: E501
        ),
    ],
)
def test_is_chunked_prefill_supported(
    model_id: str,
    expected_attn_type: str,
    expected_result: bool,
    reason: str,
    caplog_vllm,
):
    model_config = ModelConfig(model_id, trust_remote_code=True)
    assert model_config.attn_type == expected_attn_type
    with caplog_vllm.at_level(level=logging.DEBUG, logger="vllm"):
        assert model_config.is_chunked_prefill_supported == expected_result
    assert reason in caplog_vllm.text


@pytest.mark.parametrize(
    ("model_id", "expected_attn_type", "expected_result", "reason"),
    [
        # pooling models
        (
            "jason9693/Qwen2.5-1.5B-apeach",
            "decoder",
            True,
            "Pooling models with causal attn and LAST/ALL pooling support prefix caching.",  # noqa: E501
        ),
        (
            "Qwen/Qwen3-Embedding-0.6B",
            "decoder",
            True,
            "Pooling models with causal attn and LAST/ALL pooling support prefix caching.",  # noqa: E501
        ),
        (
            "Qwen/Qwen2.5-Math-PRM-7B",
            "decoder",
            False,
            "Pooling models with causal attn and LAST/STEP pooling do not support prefix caching.",  # noqa: E501
        ),
        (
            "internlm/internlm2-1_8b-reward",
            "decoder",
            True,
            "Pooling models with causal attn and LAST/ALL pooling support prefix caching.",  # noqa: E501
        ),
        (
            "BAAI/bge-base-en",
            "encoder_only",
            False,
            "Pooling models with bidirectional attn do not support prefix caching.",  # noqa: E501
        ),
        (
            "boltuix/NeuroBERT-NER",
            "encoder_only",
            False,
            "Pooling models with bidirectional attn do not support prefix caching.",  # noqa: E501
        ),
        (
            "papluca/xlm-roberta-base-language-detection",
            "encoder_only",
            False,
            "Pooling models with bidirectional attn do not support prefix caching.",  # noqa: E501
        ),
        (
            "Alibaba-NLP/gte-Qwen2-1.5B-instruct",
            "encoder_only",
            False,
            "Pooling models with bidirectional attn do not support prefix caching.",  # noqa: E501
        ),
        (
            "intfloat/e5-small",
            "encoder_only",
            False,
            "Pooling models with bidirectional attn do not support prefix caching.",  # noqa: E501
        ),
        # multimodal models
        (
            "openai/clip-vit-base-patch32",
            "decoder",
            True,
            "Pooling models with causal attn and LAST/ALL pooling support prefix caching.",  # noqa: E501
        ),
        (
            "google/siglip-base-patch16-224",
            "encoder_only",
            False,
            "Pooling models with bidirectional attn do not support prefix caching.",  # noqa: E501
        ),
        # generate models
        (
            "Qwen/Qwen3-0.6B",
            "decoder",
            True,
            "Generative models support prefix caching.",  # noqa: E501
        ),
        (
            "Qwen/Qwen3-Next-80B-A3B-Instruct",
            "hybrid",
            True,
            "Generative hybrid models support prefix caching.",  # noqa: E501
        ),
        (
            "ibm-granite/granite-4.0-h-small",
            "hybrid",
            True,
            "Generative hybrid models support prefix caching.",  # noqa: E501
        ),
        (
            "state-spaces/mamba-130m-hf",
            "attention_free",
            False,
            "Attention free models do not support prefix caching since the feature is still experimental.",  # noqa: E501
        ),
        # encoder_decoder models
        (
            "openai/whisper-small",
            "encoder_decoder",
            False,
            "Encoder decoder models do not support prefix caching.",  # noqa: E501
        ),
    ],
)
def test_is_prefix_caching_supported(
    model_id: str,
    expected_attn_type: str,
    expected_result: bool,
    reason: str,
    caplog_vllm,
):
    model_config = ModelConfig(model_id, trust_remote_code=True)
    assert model_config.attn_type == expected_attn_type
    with caplog_vllm.at_level(level=logging.DEBUG, logger="vllm"):
        assert model_config.is_prefix_caching_supported == expected_result
    assert reason in caplog_vllm.text


@pytest.mark.parametrize(
    ("backend", "custom_ops", "expected"),
    [
        ("eager", [], True),
        ("eager", ["+fused_layernorm"], True),
        ("eager", ["all", "-fused_layernorm"], False),
        ("inductor", [], False),
        ("inductor", ["none", "+fused_layernorm"], True),
        ("inductor", ["none", "-fused_layernorm"], False),
    ],
)
def test_is_custom_op_enabled(backend: str, custom_ops: list[str], expected: bool):
    """Test that is_custom_op_enabled works correctly."""
    config = VllmConfig(
        compilation_config=CompilationConfig(backend=backend, custom_ops=custom_ops)
    )
    assert config.compilation_config.is_custom_op_enabled("fused_layernorm") is expected


def test_vllm_config_defaults_are_none():
    """Verify that optimization-level defaults are None when not set by user."""
    # Test all optimization levels to ensure defaults work correctly
    for opt_level in OptimizationLevel:
        config = object.__new__(VllmConfig)
        config.compilation_config = CompilationConfig()
        config.optimization_level = opt_level
        config.model_config = None

        # Use the global optimization level defaults
        default_config = OPTIMIZATION_LEVEL_TO_CONFIG[opt_level]

        # Verify that all pass_config values are None before defaults are applied
        for pass_k in default_config["compilation_config"]["pass_config"]:
            assert getattr(config.compilation_config.pass_config, pass_k) is None

        # Verify that other config values are None before defaults are applied
        for k in default_config["compilation_config"]:
            if k != "pass_config":
                assert getattr(config.compilation_config, k) is None


def test_validate_mamba_align_subblock_prefill():
    """Align mode permits configured prefill chunks smaller than a block."""
    config = SimpleNamespace(
        cache_config=SimpleNamespace(
            block_size=11392,
            mamba_cache_mode="align",
        ),
        parallel_config=SimpleNamespace(
            decode_context_parallel_size=1,
        ),
        scheduler_config=SimpleNamespace(
            max_num_batched_tokens=8192,
            long_prefill_token_threshold=4096,
            disable_chunked_mm_input=False,
        ),
        kv_transfer_config=None,
    )

    VllmConfig.validate_block_size(config)


@pytest.mark.parametrize(
    ("model_id", "compilation_config", "optimization_level"),
    [
        (
            None,
            CompilationConfig(backend="eager", custom_ops=["+quant_fp8"]),
            OptimizationLevel.O0,
        ),
        (None, CompilationConfig(), OptimizationLevel.O0),
        (None, CompilationConfig(), OptimizationLevel.O1),
        (None, CompilationConfig(), OptimizationLevel.O2),
        (None, CompilationConfig(), OptimizationLevel.O3),
        (
            "RedHatAI/Qwen3-8B-speculator.eagle3",
            CompilationConfig(backend="inductor", custom_ops=["+quant_fp8"]),
            OptimizationLevel.O2,
        ),
        (
            "RedHatAI/Qwen3-8B-speculator.eagle3",
            CompilationConfig(),
            OptimizationLevel.O0,
        ),
        (
            "RedHatAI/Qwen3-8B-speculator.eagle3",
            CompilationConfig(),
            OptimizationLevel.O1,
        ),
        (
            "RedHatAI/Qwen3-8B-speculator.eagle3",
            CompilationConfig(),
            OptimizationLevel.O2,
        ),
        (
            "RedHatAI/Qwen3-8B-speculator.eagle3",
            CompilationConfig(),
            OptimizationLevel.O3,
        ),
        ("RedHatAI/DeepSeek-V2.5-1210-FP8", CompilationConfig(), OptimizationLevel.O0),
        ("RedHatAI/DeepSeek-V2.5-1210-FP8", CompilationConfig(), OptimizationLevel.O1),
        ("RedHatAI/DeepSeek-V2.5-1210-FP8", CompilationConfig(), OptimizationLevel.O2),
        ("RedHatAI/DeepSeek-V2.5-1210-FP8", CompilationConfig(), OptimizationLevel.O3),
    ],
)
def test_vllm_config_defaults(model_id, compilation_config, optimization_level):
    """Test that optimization-level defaults are correctly applied."""
    model_config = None
    if model_id is not None:
        model_config = ModelConfig(model_id)
        vllm_config = VllmConfig(
            model_config=model_config,
            compilation_config=compilation_config,
            optimization_level=optimization_level,
        )
    else:
        vllm_config = VllmConfig(
            compilation_config=compilation_config,
            optimization_level=optimization_level,
        )
    # Use the global optimization level defaults
    default_config = OPTIMIZATION_LEVEL_TO_CONFIG[optimization_level]

    # Verify pass_config defaults (nested under compilation_config)
    pass_config_dict = default_config["compilation_config"]["pass_config"]
    for pass_k, pass_v in pass_config_dict.items():
        actual = getattr(vllm_config.compilation_config.pass_config, pass_k)
        expected = pass_v(vllm_config) if callable(pass_v) else pass_v
        assert actual == expected, (
            f"pass_config.{pass_k}: expected {expected}, got {actual}"
        )

    # Verify other compilation_config defaults
    compilation_config_dict = default_config["compilation_config"]
    for k, v in compilation_config_dict.items():
        if k == "pass_config":
            continue
        actual = getattr(vllm_config.compilation_config, k)
        expected = v(vllm_config) if callable(v) else v
        # On platforms without static graph support, __post_init__ forces
        # cudagraph_mode to NONE; expect that instead of the level default.
        if k == "cudagraph_mode" and not current_platform.support_static_graph_mode():
            expected = CUDAGraphMode.NONE
        assert actual == expected, (
            f"compilation_config.{k}: expected {expected}, got {actual}"
        )


def test_vllm_config_callable_defaults():
    """Test that callable defaults work in the config system.

    Verifies that lambdas in default configs can inspect VllmConfig properties
    (e.g., is_quantized, is_model_moe) to conditionally set optimization flags.
    """
    config_no_model = VllmConfig(optimization_level=OptimizationLevel.O2)

    # Callable that checks if model exists
    has_model = lambda cfg: cfg.model_config is not None
    assert has_model(config_no_model) is False

    # Test with quantized model
    quantized_model = ModelConfig("RedHatAI/Llama-3.2-1B-FP8")
    config_quantized = VllmConfig(
        model_config=quantized_model, optimization_level=OptimizationLevel.O2
    )
    enable_if_quantized = lambda cfg: (
        cfg.model_config is not None and cfg.model_config.is_quantized
    )
    assert enable_if_quantized(config_quantized) is True
    assert enable_if_quantized(config_no_model) is False

    # Test with MoE model
    moe_model = ModelConfig("deepseek-ai/DeepSeek-V2-Lite")
    config_moe = VllmConfig(
        model_config=moe_model, optimization_level=OptimizationLevel.O2
    )
    enable_if_sequential = lambda cfg: (
        cfg.model_config is not None and not cfg.model_config.is_moe
    )
    assert enable_if_sequential(config_moe) is False
    assert enable_if_sequential(config_quantized) is True


@pytest.mark.skipif(
    not current_platform.support_static_graph_mode(),
    reason="Explicit overrides may be force-overwritten without static graph support.",
)
def test_vllm_config_explicit_overrides():
    """Test that explicit property overrides work correctly with callable defaults.

    When users explicitly set configuration properties, those values
    take precedence over callable defaults, across different models and
    optimization levels.
    """
    from vllm.config.compilation import PassConfig

    quantized_model = ModelConfig("RedHatAI/Llama-3.2-1B-FP8")
    moe_model = ModelConfig("deepseek-ai/DeepSeek-V2-Lite")
    regular_model = ModelConfig("Qwen/Qwen1.5-7B")

    # Explicit compilation mode override on O0 (where default is NONE)
    compilation_config = CompilationConfig(mode=CompilationMode.VLLM_COMPILE)
    config = VllmConfig(
        optimization_level=OptimizationLevel.O0,
        compilation_config=compilation_config,
    )
    assert config.compilation_config.mode == CompilationMode.VLLM_COMPILE
    assert config.compilation_config.cudagraph_mode == CUDAGraphMode.NONE

    # Explicit pass config flags to override defaults
    pass_config = PassConfig(eliminate_noops=True, fuse_attn_quant=True)
    compilation_config = CompilationConfig(pass_config=pass_config)
    config = VllmConfig(
        optimization_level=OptimizationLevel.O0,
        compilation_config=compilation_config,
    )
    assert config.compilation_config.pass_config.eliminate_noops is True
    assert config.compilation_config.pass_config.fuse_attn_quant is True

    # Explicit cudagraph mode override on quantized model at O2
    pass_config = PassConfig(enable_qk_norm_rope_fusion=True)
    compilation_config = CompilationConfig(
        cudagraph_mode=CUDAGraphMode.NONE, pass_config=pass_config
    )
    config = VllmConfig(
        model_config=quantized_model,
        optimization_level=OptimizationLevel.O2,
        compilation_config=compilation_config,
    )
    assert config.compilation_config.cudagraph_mode == CUDAGraphMode.NONE
    assert config.compilation_config.pass_config.enable_qk_norm_rope_fusion is (
        current_platform.is_cuda_alike() or current_platform.is_xpu()
    )
    # Mode should still use default for O2
    assert config.compilation_config.mode == CompilationMode.VLLM_COMPILE

    # Different optimization levels with same model
    config_o0 = VllmConfig(
        model_config=regular_model, optimization_level=OptimizationLevel.O0
    )
    config_o2 = VllmConfig(
        model_config=regular_model, optimization_level=OptimizationLevel.O2
    )
    assert config_o0.compilation_config.mode == CompilationMode.NONE
    assert config_o2.compilation_config.mode == CompilationMode.VLLM_COMPILE
    assert config_o0.compilation_config.cudagraph_mode == CUDAGraphMode.NONE
    assert (
        config_o2.compilation_config.cudagraph_mode == CUDAGraphMode.FULL_AND_PIECEWISE
    )

    # Same optimization level across different model types
    config_moe_o2 = VllmConfig(
        model_config=moe_model, optimization_level=OptimizationLevel.O2
    )
    config_regular_o2 = VllmConfig(
        model_config=regular_model, optimization_level=OptimizationLevel.O2
    )
    config_quantized_o2 = VllmConfig(
        model_config=quantized_model, optimization_level=OptimizationLevel.O2
    )
    # All should have same base compilation settings at O2
    assert config_moe_o2.compilation_config.mode == CompilationMode.VLLM_COMPILE
    assert config_regular_o2.compilation_config.mode == CompilationMode.VLLM_COMPILE
    assert config_quantized_o2.compilation_config.mode == CompilationMode.VLLM_COMPILE
    assert (
        config_moe_o2.compilation_config.cudagraph_mode
        == CUDAGraphMode.FULL_AND_PIECEWISE
    )
    assert (
        config_regular_o2.compilation_config.cudagraph_mode
        == CUDAGraphMode.FULL_AND_PIECEWISE
    )

    # Override one field but not others
    pass_config = PassConfig(eliminate_noops=False)
    compilation_config = CompilationConfig(pass_config=pass_config)
    config = VllmConfig(
        model_config=regular_model,
        optimization_level=OptimizationLevel.O2,
        compilation_config=compilation_config,
    )
    # Explicit override should be respected
    assert config.compilation_config.pass_config.eliminate_noops is False
    # Other fields should still use defaults
    assert config.compilation_config.mode == CompilationMode.VLLM_COMPILE
    assert config.compilation_config.cudagraph_mode == CUDAGraphMode.FULL_AND_PIECEWISE


def test_fusion_pass_op_priority():
    """This test checks that custom op enablement & IR op priority
    correctly control default fusions"""
    # Default config, O2, rms_norm+quant fusion disabled
    cfg1 = VllmConfig()
    assert not cfg1.compilation_config.pass_config.fuse_norm_quant

    # rms_norm manually enabled, O1, rms_norm+quant fusion enabled
    cfg2 = VllmConfig(
        optimization_level=OptimizationLevel.O1,
        compilation_config=CompilationConfig(
            custom_ops=["+rms_norm"],
        ),
    )
    assert cfg2.compilation_config.pass_config.fuse_norm_quant

    # using custom kernel for RMSNorm via IR:
    # Note that vLLM IR only supports the non-residual rms_norm for now;
    # soon this will be resolved.
    cfg3 = VllmConfig(
        kernel_config=KernelConfig(
            ir_op_priority=IrOpPriorityConfig(rms_norm=["vllm_c"])
        )
    )
    assert cfg3.compilation_config.pass_config.fuse_norm_quant

    # block-fp8 model should enable quant_fp8 automatically
    cfg4 = VllmConfig(model_config=ModelConfig("Qwen/Qwen3-4B-FP8"))
    assert "+quant_fp8" in cfg4.compilation_config.custom_ops
    assert cfg4.compilation_config.pass_config.fuse_norm_quant


def test_scheduler_config_init():
    with pytest.raises(ValidationError):
        # Positional InitVars missing
        # (InitVars cannot have defaults otherwise they will become attributes)
        SchedulerConfig()

    with pytest.raises(AttributeError):
        # InitVar does not become an attribute
        print(SchedulerConfig.default_factory().max_model_len)


@pytest.mark.parametrize(
    (
        "model_id",
        "data_parallel_size",
        "external_lb",
        "expected_needs_coordinator",
    ),
    [
        # Non-MoE model with DP=1 should not need coordinator
        ("facebook/opt-125m", 1, False, False),
        # Non-MoE model with DP>1 internal LB should need coordinator
        ("facebook/opt-125m", 2, False, True),
        # MoE model with DP=1 should not need coordinator
        ("mistralai/Mixtral-8x7B-Instruct-v0.1", 1, False, False),
        # MoE model with DP>1 internal LB should need both coordinator
        # and wave coordination
        ("mistralai/Mixtral-8x7B-Instruct-v0.1", 2, False, True),
        # MoE model with DP>1 external LB needs coordinator for wave coordination
        # (wave coordination runs in coordinator process)
        ("mistralai/Mixtral-8x7B-Instruct-v0.1", 2, True, True),
    ],
)
def test_needs_dp_coordination(
    model_id,
    data_parallel_size,
    external_lb,
    expected_needs_coordinator,
):
    """Test that DP coordinator and wave coordination are configured correctly."""
    from vllm.config import ParallelConfig

    model_config = ModelConfig(model_id)
    parallel_config = ParallelConfig(
        data_parallel_size=data_parallel_size,
        data_parallel_external_lb=external_lb,
    )
    vllm_config = VllmConfig(model_config=model_config, parallel_config=parallel_config)

    assert vllm_config.needs_dp_coordinator == expected_needs_coordinator


def test_fault_tolerance_requires_single_api_server():
    """Fault tolerance assumes one AsyncMPClient manages all engines, so it
    is incompatible with API server scale-out (_api_process_count > 1)."""
    with pytest.raises(ValueError, match="single API server"):
        ParallelConfig(enable_fault_tolerance=True, _api_process_count=2)

    # Single API server (the FT-supported topology) is accepted.
    ParallelConfig(enable_fault_tolerance=True, _api_process_count=1)


def test_renderer_num_workers_with_mm_cache():
    """Disallow renderer_num_workers > 1 with the mm processor cache only for
    pooling models, whose preprocessing runs on the renderer workers."""
    mm_model = "Qwen/Qwen2-VL-2B-Instruct"

    # Should raise: pooling + multi-worker + cache enabled (default cache_gb=4)
    with pytest.raises(ValueError, match="renderer-num-workers"):
        ModelConfig(mm_model, runner="pooling", renderer_num_workers=4)

    # Should raise: pooling + multi-worker + explicit cache size
    with pytest.raises(ValueError, match="renderer-num-workers"):
        ModelConfig(
            mm_model,
            runner="pooling",
            renderer_num_workers=2,
            mm_processor_cache_gb=1.0,
        )

    # Should pass: pooling + multi-worker + cache disabled
    config = ModelConfig(
        mm_model, runner="pooling", renderer_num_workers=4, mm_processor_cache_gb=0
    )
    assert config.renderer_num_workers == 4

    # Should pass: generate models preprocess on the dedicated mm executor
    config = ModelConfig(mm_model, renderer_num_workers=4)
    assert config.renderer_num_workers == 4

    # Should pass: single worker + cache enabled (default)
    config = ModelConfig(mm_model, renderer_num_workers=1)
    assert config.renderer_num_workers == 1


def test_eagle_draft_model_config():
    """Test that EagleDraft model config is correctly set."""
    target_model_config = ModelConfig(
        "meta-llama/Meta-Llama-3-8B-Instruct", trust_remote_code=True
    )
    speculative_config = SpeculativeConfig(
        model="yuhuili/EAGLE-LLaMA3-Instruct-8B",
        num_speculative_tokens=1,
        target_model_config=target_model_config,
        target_parallel_config=ParallelConfig(),
    )
    draft_model_config = speculative_config.draft_model_config
    assert draft_model_config.hf_config.architectures == ["EagleLlamaForCausalLM"]
    assert draft_model_config.hf_text_config.architectures == ["EagleLlamaForCausalLM"]
    assert draft_model_config.hf_config.model_type == "eagle"
    assert draft_model_config.hf_text_config.model_type == "eagle"
    assert draft_model_config.architectures == ["EagleLlamaForCausalLM"]
    assert draft_model_config.architecture == "EagleLlamaForCausalLM"


def test_draft_sample_method_probabilistic_is_accepted():
    speculative_config = SpeculativeConfig(
        method="ngram",
        num_speculative_tokens=1,
        draft_sample_method="probabilistic",
    )
    assert speculative_config.draft_sample_method == "probabilistic"


@pytest.mark.parametrize("disable_eagle_block_drop", [False, True])
def test_eagle_block_drop_can_be_disabled_without_disabling_eagle(
    disable_eagle_block_drop: bool,
):
    # Start from an ngram config to avoid loading model metadata: these predicates
    # depend only on the speculative method and the new switch.
    speculative_config = SpeculativeConfig(
        method="ngram",
        num_speculative_tokens=3,
        disable_eagle_block_drop=disable_eagle_block_drop,
    )
    speculative_config.method = "eagle3"

    assert speculative_config.use_eagle()
    assert speculative_config.use_eagle_block_drop() is not disable_eagle_block_drop


def test_draft_sample_method_gumbel_is_rejected():
    with pytest.raises(ValidationError):
        SpeculativeConfig(
            method="ngram",
            num_speculative_tokens=1,
            draft_sample_method="gumbel",
        )


def _watermarked_vllm_config() -> VllmConfig:
    config = object.__new__(VllmConfig)
    config.watermark_config = WatermarkConfig(key=42)
    config.attention_config = AttentionConfig()
    config.speculative_config = None
    return config


def test_target_only_gumbel_allows_speculative_decoding(caplog_vllm, disable_log_dedup):
    config = _watermarked_vllm_config()
    config.watermark_config = WatermarkConfig(
        key=42, allow_target_only_watermarking=True
    )
    config.speculative_config = SimpleNamespace(
        method="mtp",
        draft_sample_method="probabilistic",
        rejection_sample_method="standard",
        parallel_drafting=False,
    )

    with caplog_vllm.at_level(logging.WARNING):
        config._check_watermarking_unsupported()

    assert "Target-only watermarking leaves accepted draft tokens" in caplog_vllm.text


def test_speculative_watermarking_without_context_dedup_does_not_warn(
    caplog_vllm, disable_log_dedup
):
    config = _watermarked_vllm_config()
    config.watermark_config = WatermarkConfig(
        algorithm="dual_key_gumbel", key=42, deduplicate_contexts="none"
    )
    config.speculative_config = SimpleNamespace(
        method="mtp",
        draft_sample_method="probabilistic",
        rejection_sample_method="standard",
        parallel_drafting=False,
    )

    caplog_vllm.clear()
    with caplog_vllm.at_level(logging.WARNING):
        config._check_watermarking_unsupported()

    assert "dedup" not in caplog_vllm.text.lower()


def test_speculative_context_dedup_logs_no_unsupported_warning(
    caplog_vllm, disable_log_dedup
):
    config = _watermarked_vllm_config()
    config.watermark_config = WatermarkConfig(
        algorithm="dual_key_gumbel", key=42, deduplicate_contexts="single_turn"
    )
    config.speculative_config = SimpleNamespace(
        method="mtp",
        draft_sample_method="probabilistic",
        rejection_sample_method="standard",
        parallel_drafting=False,
    )

    caplog_vllm.clear()
    with caplog_vllm.at_level(logging.WARNING):
        config._check_watermarking_unsupported()

    assert "dedup" not in caplog_vllm.text.lower()


def test_gumbel_rejects_speculative_decoding_without_target_only():
    config = _watermarked_vllm_config()
    config.watermark_config = WatermarkConfig(
        algorithm="gumbel", key=42, allow_target_only_watermarking=False
    )
    config.speculative_config = SimpleNamespace(
        method="mtp",
        draft_sample_method="probabilistic",
        rejection_sample_method="standard",
        parallel_drafting=False,
    )

    with pytest.raises(ValueError, match="'gumbel'.*allow_target_only_watermarking"):
        config._check_watermarking_unsupported()


def test_dual_key_gumbel_warns_that_configured_alpha_is_unused(
    caplog_vllm, disable_log_dedup
):
    config = _watermarked_vllm_config()
    config.watermark_config = WatermarkConfig(
        algorithm="dual_key_gumbel", key=42, alpha=0.25
    )
    config.speculative_config = SimpleNamespace(
        method="mtp",
        draft_sample_method="probabilistic",
        rejection_sample_method="standard",
        parallel_drafting=False,
    )

    with caplog_vllm.at_level(logging.WARNING):
        config._check_watermarking_unsupported()

    assert "The configured alpha=0.25 is not used" in caplog_vllm.text


@pytest.mark.parametrize(
    ("alpha", "speculative"),
    [(0.1, True), (0.25, False)],
    ids=["default-alpha", "no-specdec"],
)
def test_dual_key_gumbel_alpha_warning_is_not_emitted(
    caplog_vllm, disable_log_dedup, alpha, speculative
):
    config = _watermarked_vllm_config()
    config.watermark_config = WatermarkConfig(
        algorithm="dual_key_gumbel", key=42, alpha=alpha
    )
    if speculative:
        config.speculative_config = SimpleNamespace(
            method="mtp",
            draft_sample_method="probabilistic",
            rejection_sample_method="standard",
            parallel_drafting=False,
        )

    with caplog_vllm.at_level(logging.WARNING):
        config._check_watermarking_unsupported()

    assert "is not used" not in caplog_vllm.text


def test_dual_key_gumbel_requires_probabilistic_drafting():
    config = _watermarked_vllm_config()
    config.watermark_config = WatermarkConfig(algorithm="dual_key_gumbel", key=42)
    config.speculative_config = SimpleNamespace(
        method="mtp",
        draft_sample_method="greedy",
        rejection_sample_method="standard",
        parallel_drafting=False,
    )

    with pytest.raises(ValueError, match="draft_sample_method='probabilistic'"):
        config._check_watermarking_unsupported()


@pytest.mark.parametrize("method", ["eagle", "eagle3", "mtp"])
def test_dual_key_gumbel_supports_probabilistic_speculative_decoding(method):
    config = _watermarked_vllm_config()
    config.watermark_config = WatermarkConfig(algorithm="dual_key_gumbel", key=42)
    config.speculative_config = SimpleNamespace(
        method=method,
        draft_sample_method="probabilistic",
        rejection_sample_method="standard",
        parallel_drafting=False,
    )

    config._check_watermarking_unsupported()


def test_dual_key_gumbel_supports_dspark():
    config = _watermarked_vllm_config()
    config.watermark_config = WatermarkConfig(algorithm="dual_key_gumbel", key=42)
    config.speculative_config = SimpleNamespace(
        method="dspark",
        draft_sample_method="probabilistic",
        rejection_sample_method="standard",
        parallel_drafting=True,
    )

    config._check_watermarking_unsupported()


def test_dual_key_gumbel_rejects_non_autoregressive_speculation():
    config = _watermarked_vllm_config()
    config.watermark_config = WatermarkConfig(algorithm="dual_key_gumbel", key=42)
    config.speculative_config = SimpleNamespace(
        method="ngram",
        draft_sample_method="probabilistic",
        rejection_sample_method="standard",
        parallel_drafting=False,
    )

    with pytest.raises(ValueError, match="autoregressive model-based"):
        config._check_watermarking_unsupported()


@pytest.mark.parametrize(
    ("overrides", "match"),
    [
        ({"rejection_sample_method": "synthetic"}, "rejection_sample_method"),
        ({"rejection_sample_method": "block"}, "rejection_sample_method"),
        ({"parallel_drafting": True}, "Parallel speculative drafting"),
    ],
)
def test_dual_key_gumbel_rejects_incompatible_speculative_modes(overrides, match):
    config = _watermarked_vllm_config()
    config.watermark_config = WatermarkConfig(algorithm="dual_key_gumbel", key=42)
    values = {
        "method": "mtp",
        "draft_sample_method": "probabilistic",
        "rejection_sample_method": "standard",
        "parallel_drafting": False,
        **overrides,
    }
    config.speculative_config = SimpleNamespace(**values)

    with pytest.raises(ValueError, match=match):
        config._check_watermarking_unsupported()


def test_gumbel_watermark_rejects_beam_search():
    with pytest.raises(ValueError, match="Beam search is not supported"):
        _watermarked_vllm_config()._check_watermarking_unsupported(beam_search=True)


def test_gumbel_watermark_rejects_custom_sampler():
    with pytest.raises(ValueError, match="custom samplers are not supported"):
        _watermarked_vllm_config()._check_watermarking_unsupported(custom_sampler=True)


def test_watermark_key_must_fit_in_64_bits():
    with pytest.raises(ValueError, match="64 bits"):
        WatermarkConfig(key=2**64)


@pytest.mark.parametrize("alpha", [-0.1, 1.1])
def test_watermark_alpha_must_be_a_probability(alpha):
    with pytest.raises(ValidationError):
        WatermarkConfig(key=42, algorithm="dual_key_gumbel", alpha=alpha)


def test_unknown_watermark_prf_is_rejected():
    with pytest.raises(ValidationError):
        pydantic.TypeAdapter(WatermarkConfig).validate_python(
            {"key": 42, "prf": "unsupported"}
        )


def test_watermark_key_is_excluded_from_serialization():
    serialized = pydantic.TypeAdapter(WatermarkConfig).dump_python(
        WatermarkConfig(key=42), mode="json"
    )

    assert "key" not in serialized


def test_watermarking_forces_model_runner_v2(monkeypatch):
    monkeypatch.setenv("VLLM_USE_V2_MODEL_RUNNER", "0")

    with patch("vllm.config.vllm.logger.info_once") as info_once:
        assert _watermarked_vllm_config().use_v2_model_runner

    info_once.assert_called_once_with(
        "Watermarking requires Model Runner V2 and overrides "
        "VLLM_USE_V2_MODEL_RUNNER=0."
    )


@patch("vllm.config.speculative.ModelConfig")
def test_mtp_draft_uses_model_weights_not_local_cache(mock_model_config_cls):
    """Regression test: MTP + runai_streamer should use model_weights (original
    S3 URL) for the draft model, not model (local cache dir set by
    pull_runai_model_from_obj_storage)."""
    from unittest.mock import MagicMock

    s3_url = "s3://my-bucket/Qwen3-35B-A3B-FP8"
    local_cache = "/root/.cache/vllm/assets/model_streamer/abcd1234"

    mock_draft = MagicMock()
    mock_draft.model = local_cache
    mock_draft.hf_config.model_type = "deepseek_mtp"
    mock_draft.hf_config.n_predict = None
    mock_draft.max_model_len = 4096
    mock_model_config_cls.return_value = mock_draft

    target_config = MagicMock()
    target_config.model = local_cache
    target_config.model_weights = s3_url
    target_config.hf_text_config.model_type = "deepseek_v3"
    target_config.quantization = None
    target_config.max_model_len = 4096

    SpeculativeConfig(
        method="mtp",
        num_speculative_tokens=1,
        target_model_config=target_config,
        target_parallel_config=ParallelConfig(),
    )

    actual_model = mock_model_config_cls.call_args.kwargs["model"]
    assert actual_model == s3_url


def _make_qwen3_omni_dspark_configs():
    text_config = SimpleNamespace(
        hidden_size=2048,
        num_hidden_layers=48,
        num_attention_heads=32,
        num_key_value_heads=4,
        head_dim=128,
        vocab_size=152064,
    )
    target_model_config = SimpleNamespace(
        hf_config=SimpleNamespace(
            model_type="qwen3_omni_moe",
            architectures=["Qwen3OmniMoeForConditionalGeneration"],
            thinker_config=SimpleNamespace(text_config=text_config),
        ),
        hf_text_config=text_config,
        architectures=["Qwen3OmniMoeForConditionalGeneration"],
        get_hidden_size=lambda: 2048,
        get_total_num_hidden_layers=lambda: 48,
        get_vocab_size=lambda: 152064,
    )
    draft_hf_config = SimpleNamespace(
        model_type="qwen3",
        architectures=["Qwen3OmniDSparkModel"],
        hidden_size=2048,
        num_attention_heads=32,
        num_key_value_heads=4,
        head_dim=128,
        block_size=7,
        target_hidden_size=2048,
        target_layer_ids=[1, 9, 17, 25, 33],
        use_aux_hidden_state=True,
        markov_rank=64,
        markov_head_type="vanilla",
        sample_from_anchor=True,
        dspark_bonus_anchor=False,
        vocab_size=152064,
        draft_vocab_size=32000,
        mask_token_id=151669,
        rope_parameters={"rope_type": "default", "rope_theta": 1000000.0},
    )
    draft_model_config = SimpleNamespace(
        hf_config=draft_hf_config,
        architectures=["Qwen3OmniDSparkModel"],
    )
    return target_model_config, draft_model_config


@pytest.mark.skip_global_cleanup
def test_qwen3_omni_dspark_checkpoint_contract_is_accepted():
    target_config, draft_config = _make_qwen3_omni_dspark_configs()
    _validate_qwen3_omni_dspark(target_config, draft_config, 7)


@pytest.mark.skip_global_cleanup
def test_qwen3_omni_dspark_rejects_generic_qwen3_architecture():
    target_config, draft_config = _make_qwen3_omni_dspark_configs()
    draft_config.hf_config.architectures = ["Qwen3DSparkModel"]
    draft_config.architectures = ["Qwen3DSparkModel"]

    with pytest.raises(ValueError, match="must be converted first"):
        _validate_qwen3_omni_dspark(target_config, draft_config, 7)


@pytest.mark.parametrize(
    ("field", "value", "error"),
    [
        ("block_size", 5, "trained block_size"),
        ("target_hidden_size", 4096, "target_hidden_size"),
        ("hidden_size", 4096, "draft hidden_size"),
        ("num_attention_heads", 16, "num_attention_heads"),
        ("num_key_value_heads", 8, "num_key_value_heads"),
        ("head_dim", 64, "head_dim"),
        ("target_layer_ids", [7, 48], "zero-based text-layer"),
        ("target_layer_ids", [23, 7], "strictly increasing"),
        ("use_aux_hidden_state", False, "use_aux_hidden_state=true"),
        ("markov_rank", 0, "markov_rank"),
        ("markov_head_type", "gated", "markov_head_type='vanilla'"),
        ("sample_from_anchor", False, "sample_from_anchor=true"),
        ("dspark_bonus_anchor", True, "dspark_bonus_anchor=false"),
        ("vocab_size", 0, "input vocab_size must be a positive integer"),
    ],
)
@pytest.mark.skip_global_cleanup
def test_qwen3_omni_dspark_rejects_incompatible_checkpoint_fields(field, value, error):
    target_config, draft_config = _make_qwen3_omni_dspark_configs()
    setattr(draft_config.hf_config, field, value)

    with pytest.raises(ValueError, match=error):
        _validate_qwen3_omni_dspark(target_config, draft_config, 7)


@pytest.mark.skip_global_cleanup
def test_qwen3_omni_dspark_rejects_mrope_draft_positions():
    target_config, draft_config = _make_qwen3_omni_dspark_configs()
    draft_config.hf_config.rope_parameters["mrope_section"] = [24, 20, 20]

    with pytest.raises(ValueError, match="logical 1-D RoPE"):
        _validate_qwen3_omni_dspark(target_config, draft_config, 7)


@pytest.mark.skip_global_cleanup
def test_qwen3_omni_dspark_allows_draft_only_noise_token_row():
    target_config, draft_config = _make_qwen3_omni_dspark_configs()
    draft_config.hf_config.vocab_size = 152065
    draft_config.hf_config.draft_vocab_size = 152064
    draft_config.hf_config.mask_token_id = 152064

    _validate_qwen3_omni_dspark(target_config, draft_config, 7)


@pytest.mark.skip_global_cleanup
def test_qwen3_omni_dspark_allows_smaller_input_vocabulary():
    target_config, draft_config = _make_qwen3_omni_dspark_configs()
    draft_config.hf_config.vocab_size = 151936
    draft_config.hf_config.mask_token_id = 151669

    _validate_qwen3_omni_dspark(target_config, draft_config, 7)


def test_ir_op_priority_default():
    """Test that IR op priority defaults are set correctly."""
    from vllm.config.kernel import IrOpPriorityConfig

    # Assert default is applied to ops
    priority_config = IrOpPriorityConfig.with_default(["vllm_c", "native"])
    assert priority_config.rms_norm == ["vllm_c", "native"]
    assert priority_config.fused_add_rms_norm == ["vllm_c", "native"]

    # Assert single ops override the default
    priority_config = IrOpPriorityConfig.with_default(
        ["native"], rms_norm=["oink", "native"]
    )
    assert priority_config.rms_norm == ["oink", "native"]
    assert priority_config.fused_add_rms_norm == ["native"]


@pytest.mark.parametrize("mode", [CompilationMode.NONE, CompilationMode.VLLM_COMPILE])
@pytest.mark.parametrize("backend", ["inductor", "eager"])
def test_ir_op_platform_defaults_support_sparse_gelu(mode, backend):
    """Worker initialization must not select an unregistered sparse GELU provider."""
    from vllm import ir

    config = SimpleNamespace(
        compilation_config=CompilationConfig(mode=mode, backend=backend)
    )
    priority = current_platform.get_default_ir_op_priority(config)
    expected = ["triton", "native"] if current_platform.is_cuda() else ["native"]

    with priority.set_priority():
        assert ir.ops.gelu_and_mul_sparse.get_priority() == expected


def test_ir_op_priority_str():
    """Test that passing a comma-delimited string works."""
    from vllm.config.kernel import IrOpPriorityConfig

    priority_config = IrOpPriorityConfig(rms_norm="vllm_c")
    assert priority_config.rms_norm == ["vllm_c"]

    priority_config = IrOpPriorityConfig(rms_norm="vllm_c,native")
    assert priority_config.rms_norm == ["vllm_c", "native"]

    priority_config = IrOpPriorityConfig(rms_norm=" native, vllm_c ")
    assert priority_config.rms_norm == ["native", "vllm_c"]

    with pytest.raises(pydantic.ValidationError):
        # must be list of only strings
        priority_config = IrOpPriorityConfig(rms_norm=["vllm_c", 4, "native"])


def test_ir_op_priority_ctx():
    """Test that the priority-setting context sets priority correctly."""
    from vllm import ir
    from vllm.config.kernel import IrOpPriorityConfig

    priority = IrOpPriorityConfig.with_default(["native"], rms_norm=["vllm_c"])
    priority2 = IrOpPriorityConfig.with_default(
        ["native"], fused_add_rms_norm=["vllm_c"]
    )
    with priority.set_priority():
        assert ir.ops.rms_norm.get_priority() == ["vllm_c", "native"]
        assert ir.ops.fused_add_rms_norm.get_priority() == ["native"]
        with priority2.set_priority():
            assert ir.ops.rms_norm.get_priority() == ["native"]
            assert ir.ops.fused_add_rms_norm.get_priority() == ["vllm_c", "native"]

        # context restored
        assert ir.ops.rms_norm.get_priority() == ["vllm_c", "native"]
        assert ir.ops.fused_add_rms_norm.get_priority() == ["native"]

        with pytest.raises(ValueError), priority2.set_priority():
            assert ir.ops.rms_norm.get_priority() == ["native"]
            assert ir.ops.fused_add_rms_norm.get_priority() == ["vllm_c", "native"]

            raise ValueError

        # context restored even after exception
        assert ir.ops.rms_norm.get_priority() == ["vllm_c", "native"]
        assert ir.ops.fused_add_rms_norm.get_priority() == ["native"]


def test_load_config_rejects_invalid_safetensors_load_strategy():
    with pytest.raises(pydantic.ValidationError):
        LoadConfig(safetensors_load_strategy="not_a_real_strategy")


@pytest.mark.parametrize("bad_load_format", [None, 123])
def test_load_config_rejects_non_string_load_format(bad_load_format):
    with pytest.raises(pydantic.ValidationError):
        LoadConfig(load_format=bad_load_format)


# A real Qwen3-0.6B model revision that is used in the test below.
REVISION = "c1899de289a04d12100db370d81485cdf75e47ca"


@patch("vllm.config.model.resolve_revision", return_value=ResolvedRevision(REVISION))
def test_revision_resolved_for_model(mock_resolve):
    model = "Qwen/Qwen3-0.6B"
    config = ModelConfig(model)
    assert isinstance(config.revision, ResolvedRevision)
    assert config.revision.resolved == REVISION
    mock_resolve.assert_any_call(model, None, config.hf_token)


@pytest.mark.parametrize(
    ("layer_types", "expected_attention"),
    [
        # Qwen3-Next / Qwen3.5 spell their attention layers "full_attention".
        (["linear_attention", "full_attention"], 1),
        # GLM-5.3-Flash and Qwen4-Exp use sparse attention, which still caches
        # every token and so must count as attention.
        (["linear_attention", "deepseek_sparse_attention"], 1),
        (["linear_attention", "qwen_sparse_attention"], 1),
        (["linear_attention", "linear_attention"], 0),
    ],
)
def test_hybrid_layer_counts_sparse_attention_as_attention(
    layer_types, expected_attention
):
    """Sparse-attention layer types consume a full attention KV cache."""
    model_config = object.__new__(ModelConfig)
    model_config.hf_config = SimpleNamespace()
    model_config.hf_text_config = SimpleNamespace(layer_types=layer_types)
    model_config.model_arch_config = SimpleNamespace(text_model_type="glm5_next")

    parallel_config = SimpleNamespace()
    with (
        patch.object(ModelConfig, "is_hybrid", True),
        patch.object(ModelConfig, "has_noops", False),
        patch.object(ModelConfig, "is_attention_free", False),
        patch.object(
            ModelConfig,
            "get_layers_start_end_indices",
            return_value=(0, len(layer_types)),
        ),
    ):
        assert (
            model_config.get_num_layers_by_block_type(parallel_config, "attention")
            == expected_attention
        )
        assert (
            model_config.get_num_layers_by_block_type(
                parallel_config, "linear_attention"
            )
            == len(layer_types) - expected_attention
        )
