import pytest

from deepteam.vulnerabilities import CustomVulnerability


class TestCustomVulnerability:

    def test_custom_vulnerability_initialize(self):
        custom_vulnerability = CustomVulnerability(
            criteria="Criteria",
            name="Name",
            custom_prompt="Custom prompt",
            types=["type1", "type2"],
        )
        assert custom_vulnerability.criteria == "Criteria"
        assert custom_vulnerability.name == "Name"
        assert custom_vulnerability.custom_prompt == "Custom prompt"
        assert custom_vulnerability._get_metric(None) is not None
        assert sorted(
            type.value for type in custom_vulnerability.types
        ) == sorted(["type1", "type2"])

    def test_custom_vulnerability_no_types(self):
        """CustomVulnerability should not raise when types is omitted."""
        custom_vulnerability = CustomVulnerability(
            criteria="Criteria",
            name="Data Leakage",
        )
        assert custom_vulnerability.name == "Data Leakage"
        assert custom_vulnerability.criteria == "Criteria"
        # Should have a default type derived from the name
        type_values = [t.value for t in custom_vulnerability.types]
        assert type_values == ["Data Leakage"]

    def test_custom_vulnerability_no_types_get_name(self):
        """get_name works when types is omitted."""
        custom_vulnerability = CustomVulnerability(
            criteria="Test criteria",
            name="Test Vuln",
        )
        assert custom_vulnerability.get_name() == "Test Vuln"

    def test_custom_vulnerability_single_type(self):
        """CustomVulnerability works with a single type."""
        custom_vulnerability = CustomVulnerability(
            criteria="Criteria",
            name="Name",
            types=["only_type"],
        )
        type_values = [t.value for t in custom_vulnerability.types]
        assert type_values == ["only_type"]
