# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

from types import SimpleNamespace
from unittest.mock import MagicMock

import pytest
import torch

import lerobot.policies.factory as policy_factory
from lerobot.configs import FeatureType
from lerobot.utils.constants import ACTION


def test_make_policy_keeps_peft_adapter_and_base_revisions_separate(monkeypatch):
    cfg = SimpleNamespace(
        type="mock",
        device="cpu",
        pretrained_path="user/adapter",
        pretrained_revision="adapter-sha",
        use_peft=True,
        input_features={},
        output_features={},
    )
    dataset_meta = SimpleNamespace(features={}, stats={})

    base_policy = torch.nn.Linear(1, 1)
    policy_from_pretrained = MagicMock(return_value=base_policy)
    policy_class = SimpleNamespace(from_pretrained=policy_from_pretrained)
    monkeypatch.setattr(policy_factory, "get_policy_class", lambda _: policy_class)
    monkeypatch.setattr(policy_factory, "dataset_to_policy_features", lambda _: {})
    monkeypatch.setattr(policy_factory, "validate_visual_features_consistency", lambda *args: None)

    peft_config = SimpleNamespace(
        base_model_name_or_path="user/base-policy",
        revision="base-sha",
    )
    peft_config_from_pretrained = MagicMock(return_value=peft_config)
    adapted_policy = torch.nn.Linear(1, 1)
    peft_model_from_pretrained = MagicMock(return_value=adapted_policy)
    require_package = MagicMock()
    monkeypatch.setattr(policy_factory, "require_package", require_package)
    monkeypatch.setattr(
        policy_factory,
        "PeftConfig",
        SimpleNamespace(from_pretrained=peft_config_from_pretrained),
    )
    monkeypatch.setattr(
        policy_factory,
        "PeftModel",
        SimpleNamespace(from_pretrained=peft_model_from_pretrained),
    )

    policy = policy_factory.make_policy(cfg, ds_meta=dataset_meta)

    assert policy is adapted_policy
    require_package.assert_called_once_with("peft", extra="peft")
    peft_config_from_pretrained.assert_called_once_with(
        "user/adapter",
        revision="adapter-sha",
    )
    policy_from_pretrained.assert_called_once_with(
        config=cfg,
        dataset_stats=dataset_meta.stats,
        dataset_meta=dataset_meta,
        pretrained_name_or_path="user/base-policy",
        revision="base-sha",
    )
    peft_model_from_pretrained.assert_called_once_with(
        base_policy,
        "user/adapter",
        config=peft_config,
        revision="adapter-sha",
        is_trainable=True,
    )


@pytest.mark.parametrize("action_key", [ACTION, "actions"])
@pytest.mark.parametrize(
    ("raw_names", "expected_names"),
    [
        (["shoulder", "elbow", "gripper"], ["shoulder", "elbow", "gripper"]),
        (("shoulder", "elbow", "gripper"), ["shoulder", "elbow", "gripper"]),
        ({"motors": ["shoulder", "elbow", "gripper"]}, ["shoulder", "elbow", "gripper"]),
        ({"right": ["shoulder", "elbow"], "left": ["gripper"]}, ["shoulder", "elbow", "gripper"]),
        ({"right": ("shoulder", "elbow"), "left": ("gripper",)}, ["shoulder", "elbow", "gripper"]),
        ({"unused": [], "motors": ["shoulder", "elbow", "gripper"]}, ["shoulder", "elbow", "gripper"]),
        ({"shoulder": 0, "elbow": 1, "gripper": 2}, ["shoulder", "elbow", "gripper"]),
        ([], []),
        ({}, []),
        (None, None),
    ],
)
def test_make_policy_reads_action_names(monkeypatch, action_key, raw_names, expected_names):
    cfg = SimpleNamespace(
        type="mock",
        device="cpu",
        pretrained_path=None,
        use_peft=False,
        input_features={},
        output_features={},
        action_feature_names=None,
    )
    dataset_meta = SimpleNamespace(
        features={
            action_key: {
                "dtype": "float32",
                "shape": (3,),
                "names": raw_names,
            }
        },
        stats={},
    )

    policy = torch.nn.Linear(1, 1)
    policy_class = MagicMock(return_value=policy)
    monkeypatch.setattr(policy_factory, "get_policy_class", lambda _: policy_class)

    result = policy_factory.make_policy(
        cfg,
        ds_meta=dataset_meta,
        rename_map={action_key: ACTION} if action_key != ACTION else None,
    )

    assert result is policy
    assert cfg.action_feature_names == expected_names
    assert dataset_meta.features[action_key]["names"] == raw_names
    assert cfg.output_features[ACTION].type is FeatureType.ACTION


def test_make_policy_loads_resume_weights_and_keeps_the_parent(monkeypatch):
    """A resume loads the checkpoint, while `pretrained_path` keeps naming the fine-tuned-from model."""
    cfg = SimpleNamespace(
        type="mock",
        device="cpu",
        pretrained_path="user/base-policy",
        pretrained_revision=None,
        use_peft=False,
        input_features={},
        output_features={},
    )
    seen_while_building = []

    def from_pretrained(**kwargs):
        seen_while_building.append(cfg.pretrained_path)
        policy = torch.nn.Linear(1, 1)
        policy.config = cfg
        # Like FLUX3, which records its load source on the config.
        cfg.pretrained_path = str(kwargs["pretrained_name_or_path"])
        return policy

    monkeypatch.setattr(
        policy_factory, "get_policy_class", lambda _: SimpleNamespace(from_pretrained=from_pretrained)
    )
    monkeypatch.setattr(policy_factory, "dataset_to_policy_features", lambda _: {})
    monkeypatch.setattr(policy_factory, "validate_visual_features_consistency", lambda *args: None)

    policy = policy_factory.make_policy(
        cfg,
        ds_meta=SimpleNamespace(features={}, stats={}),
        pretrained_path="run/checkpoints/000002/pretrained_model",
    )

    assert [str(path) for path in seen_while_building] == ["run/checkpoints/000002/pretrained_model"]
    assert cfg.pretrained_path == "user/base-policy"
    assert policy.config.pretrained_path == "user/base-policy"
