diff --git a/pyrit/converter/token_smuggling/variation_selector_smuggler_converter.py b/pyrit/converter/token_smuggling/variation_selector_smuggler_converter.py index ad2ecd1a57..11b68eed32 100644 --- a/pyrit/converter/token_smuggling/variation_selector_smuggler_converter.py +++ b/pyrit/converter/token_smuggling/variation_selector_smuggler_converter.py @@ -47,10 +47,13 @@ def __init__( Default is True. Raises: - ValueError: If an unsupported action or ``encoding_mode`` is provided. + ValueError: If an unsupported action is provided or ``base_char_utf8`` is not exactly one character. """ super().__init__(action=action) - self.utf8_base_char = base_char_utf8 if base_char_utf8 is not None else "😊" + base_char = base_char_utf8 if base_char_utf8 is not None else "😊" + if len(base_char) != 1: + raise ValueError("base_char_utf8 must be exactly one character.") + self.utf8_base_char = base_char self.embed_in_base = embed_in_base def _build_identifier(self) -> ComponentIdentifier: diff --git a/tests/unit/converter/test_variation_selector_smuggler_converter.py b/tests/unit/converter/test_variation_selector_smuggler_converter.py index 59e242093f..4bd7cb0f9d 100644 --- a/tests/unit/converter/test_variation_selector_smuggler_converter.py +++ b/tests/unit/converter/test_variation_selector_smuggler_converter.py @@ -43,6 +43,23 @@ def test_variation_selector_invalid_action(): VariationSelectorSmugglerConverter(action="invalid") +@pytest.mark.parametrize("base_char", ["", "ab", "😊x"]) +def test_variation_selector_invalid_base_char(base_char): + with pytest.raises(ValueError, match="base_char_utf8 must be exactly one character"): + VariationSelectorSmugglerConverter(base_char_utf8=base_char) + + +async def test_variation_selector_custom_base_char_roundtrip(): + encoder = VariationSelectorSmugglerConverter(action="encode", base_char_utf8="A") + encoded = await encoder.convert_async(prompt="test", input_type="text") + + decoder = VariationSelectorSmugglerConverter(action="decode", base_char_utf8="A") + decoded = await decoder.convert_async(prompt=encoded.output_text, input_type="text") + + assert encoded.output_text.startswith("A") + assert decoded.output_text == "test" + + async def test_variation_selector_input_not_supported(): converter = VariationSelectorSmugglerConverter(action="encode") with pytest.raises(ValueError, match="Input type not supported"):