Skip to content
33 changes: 30 additions & 3 deletions pyrit/converter/converter.py
Original file line number Diff line number Diff line change
Expand Up @@ -198,8 +198,21 @@ def _get_random_seed_override(self) -> int | None:
"""
return self._seed

@staticmethod
def _mark_converted(text: str, *, start_token: str, end_token: str) -> str:
# A nested selective converter already marks its own output, so wrapping it again would nest the tokens.
if start_token in text:
return text
return f"{start_token}{text}{end_token}"

async def convert_tokens_async(
self, *, prompt: str, input_type: PromptDataType = "text", start_token: str = "⟪", end_token: str = "⟫"
self,
*,
prompt: str,
input_type: PromptDataType = "text",
start_token: str = "⟪",
end_token: str = "⟫",
keep_tokens: bool = False,
) -> ConverterResult:
"""
Convert marked text regions, consuming their delimiters and preserving all unmarked text.
Expand All @@ -216,6 +229,9 @@ async def convert_tokens_async(
relatively distinct.
end_token (str): The token indicating the end of a substring to be converted. Defaults to "⟫" which is
relatively distinct.
keep_tokens (bool): When True, each converted region stays wrapped in the start and end tokens so
a later stage can find it again. Without delimiters, the whole converted prompt is wrapped.
Defaults to False.

Returns:
ConverterResult: The prompt with specified substrings converted.
Expand All @@ -231,7 +247,13 @@ async def convert_tokens_async(

spans = self._get_token_spans(prompt=prompt, start_token=start_token, end_token=end_token)
if not spans:
return await self.convert_async(prompt=prompt, input_type=input_type)
result = await self.convert_async(prompt=prompt, input_type=input_type)
if keep_tokens and result.output_type == "text":
result = ConverterResult(
output_text=self._mark_converted(result.output_text, start_token=start_token, end_token=end_token),
output_type="text",
)
return result

if not self.input_supported("text") or not self.output_supported("text"):
raise ValueError("Selected-region conversion requires a converter supporting text input and text output.")
Expand All @@ -245,7 +267,12 @@ async def convert_tokens_async(
parts: list[str] = []
previous_end = 0
for (start, end), converted in zip(spans, converted_parts, strict=True):
parts.extend((prompt[previous_end:start], converted.output_text))
converted_text = (
self._mark_converted(converted.output_text, start_token=start_token, end_token=end_token)
if keep_tokens
else converted.output_text
)
parts.extend((prompt[previous_end:start], converted_text))
previous_end = end
parts.append(prompt[previous_end:])
return ConverterResult(output_text="".join(parts), output_type="text")
Expand Down
19 changes: 15 additions & 4 deletions pyrit/converter/selective_text_converter.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT license.

import inspect
from typing import Any

from pyrit.converter.converter import Converter, ConverterResult
from pyrit.converter.text_selection_strategy import (
Expand Down Expand Up @@ -162,16 +164,21 @@ async def convert_async(self, *, prompt: str, input_type: PromptDataType = "text

# If using TokenSelectionStrategy, delegate to convert_tokens_async
if self._is_token_based:
# With preserve_tokens, each converted region keeps its own tokens so later stages still see only
# the originally selected text as marked. keep_tokens is only passed when needed and supported, so
# custom overrides with the older convert_tokens_async signature keep working.
kwargs: dict[str, Any] = {}
if self._preserve_tokens and self._accepts_keep_tokens():
kwargs["keep_tokens"] = True
result = await self._sub_converter.convert_tokens_async(
prompt=prompt,
input_type="text",
start_token=self._start_token,
end_token=self._end_token,
**kwargs,
)
# If preserve_tokens is True, the tokens are already in the result
# If False, convert_tokens_async removes them
if self._preserve_tokens and self._start_token not in result.output_text:
# Wrap the result with tokens if they were removed
if self._preserve_tokens and not kwargs and self._start_token not in result.output_text:
# Older override without keep_tokens: fall back to wrapping the whole result.
result = ConverterResult(
output_text=f"{self._start_token}{result.output_text}{self._end_token}", output_type="text"
)
Expand All @@ -181,6 +188,10 @@ async def convert_async(self, *, prompt: str, input_type: PromptDataType = "text
return await self._convert_word_level_async(prompt=prompt)
return await self._convert_char_level_async(prompt=prompt)

def _accepts_keep_tokens(self) -> bool:
parameters = inspect.signature(self._sub_converter.convert_tokens_async).parameters.values()
return any(p.name == "keep_tokens" or p.kind is inspect.Parameter.VAR_KEYWORD for p in parameters)

async def _convert_word_level_async(self, *, prompt: str) -> ConverterResult:
"""
Convert selected words using word-level selection strategy.
Expand Down
16 changes: 16 additions & 0 deletions tests/unit/converter/test_converter.py
Original file line number Diff line number Diff line change
Expand Up @@ -714,3 +714,19 @@ def test_llm_based_converters_validate_target_requirements(setup_memory, convert
with patch("pyrit.prompt_target.common.target_requirements.TargetRequirements.validate") as mock_validate:
converter_class(**converter_args)
mock_validate.assert_called_once_with(target=setup_memory)


async def test_convert_tokens_keep_tokens_rewraps_each_region_async() -> None:
converter = Base64Converter()

result = await converter.convert_tokens_async(prompt="a ⟪b⟫ c ⟪d⟫", keep_tokens=True)

assert result.output_text == "a ⟪Yg==⟫ c ⟪ZA==⟫"


async def test_convert_tokens_keep_tokens_without_markers_wraps_prompt_async() -> None:
converter = Base64Converter()

result = await converter.convert_tokens_async(prompt="b", keep_tokens=True)

assert result.output_text == "⟪Yg==⟫"
93 changes: 93 additions & 0 deletions tests/unit/converter/test_selective_text_converter.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,16 +9,20 @@
ROT13Converter,
SelectiveTextConverter,
)
from pyrit.converter.converter import ConverterResult
from pyrit.converter.text_selection_strategy import (
IndexSelectionStrategy,
KeywordSelectionStrategy,
PositionSelectionStrategy,
ProportionSelectionStrategy,
RangeSelectionStrategy,
RegexSelectionStrategy,
TokenSelectionStrategy,
WordIndexSelectionStrategy,
WordPositionSelectionStrategy,
WordProportionSelectionStrategy,
)
from pyrit.models import PromptDataType


class TestSelectiveTextConverter:
Expand Down Expand Up @@ -297,3 +301,92 @@ def test_identifier_uses_empty_strategy_params_for_char_level_strategies(self):
params = converter.get_identifier().params
assert params["selection_strategy"] == "IndexSelectionStrategy"
assert params["selection_strategy_params"] == {}


class TestTokenSelectionChaining:
"""Token-based stages must only convert the regions an earlier stage marked."""

async def test_token_stages_keep_each_region_marked(self):
first = SelectiveTextConverter(
sub_converter=Base64Converter(),
selection_strategy=WordPositionSelectionStrategy(start_proportion=0.5, end_proportion=1.0),
preserve_tokens=True,
)
second = SelectiveTextConverter(
sub_converter=ROT13Converter(), selection_strategy=TokenSelectionStrategy(), preserve_tokens=True
)
third = SelectiveTextConverter(
sub_converter=Base64Converter(), selection_strategy=TokenSelectionStrategy(), preserve_tokens=True
)

stage1 = (await first.convert_async(prompt="tell me how to do it")).output_text
stage2 = (await second.convert_async(prompt=stage1)).output_text
stage3 = (await third.convert_async(prompt=stage2)).output_text

assert stage1 == "tell me how ⟪dG8=⟫ ⟪ZG8=⟫ ⟪aXQ=⟫"
assert stage2 == "tell me how ⟪qT8=⟫ ⟪MT8=⟫ ⟪nKD=⟫"
assert stage3 == "tell me how ⟪cVQ4PQ==⟫ ⟪TVQ4PQ==⟫ ⟪bktEPQ==⟫"

async def test_token_stage_without_preserve_tokens_drops_markers(self):
converter = SelectiveTextConverter(sub_converter=ROT13Converter(), selection_strategy=TokenSelectionStrategy())

result = await converter.convert_async(prompt="keep ⟪this⟫ and ⟪that⟫")

assert result.output_text == "keep guvf and gung"

async def test_token_stage_without_markers_wraps_whole_prompt(self):
converter = SelectiveTextConverter(
sub_converter=ROT13Converter(), selection_strategy=TokenSelectionStrategy(), preserve_tokens=True
)

result = await converter.convert_async(prompt="hello")

assert result.output_text == "⟪uryyb⟫"

async def test_nested_selective_converter_keeps_one_pair_of_markers(self):
inner = SelectiveTextConverter(
sub_converter=ROT13Converter(), selection_strategy=TokenSelectionStrategy(), preserve_tokens=True
)
outer = SelectiveTextConverter(
sub_converter=inner, selection_strategy=TokenSelectionStrategy(), preserve_tokens=True
)
later = SelectiveTextConverter(
sub_converter=Base64Converter(), selection_strategy=TokenSelectionStrategy(), preserve_tokens=True
)

nested = (await outer.convert_async(prompt="prefix ⟪word⟫ suffix")).output_text
assert nested == "prefix ⟪jbeq⟫ suffix"

result = await later.convert_async(prompt=nested)
assert result.output_text == "prefix ⟪amJlcQ==⟫ suffix"


class _LegacyTokenConverter(ROT13Converter):
"""A custom converter overriding convert_tokens_async with the older signature (no keep_tokens)."""

async def convert_tokens_async(
self, *, prompt: str, input_type: PromptDataType = "text", start_token: str = "⟪", end_token: str = "⟫"
) -> ConverterResult:
return ConverterResult(
output_text=prompt.replace(start_token, "").replace(end_token, "").upper(), output_type="text"
)


class TestLegacyConvertTokensOverride:
async def test_override_without_keep_tokens_works_by_default(self):
converter = SelectiveTextConverter(
sub_converter=_LegacyTokenConverter(), selection_strategy=TokenSelectionStrategy()
)

result = await converter.convert_async(prompt="a ⟪b⟫ c")

assert result.output_text == "A B C"

async def test_override_without_keep_tokens_falls_back_to_wrapping_with_preserve_tokens(self):
converter = SelectiveTextConverter(
sub_converter=_LegacyTokenConverter(), selection_strategy=TokenSelectionStrategy(), preserve_tokens=True
)

result = await converter.convert_async(prompt="a ⟪b⟫ c")

assert result.output_text == "⟪A B C⟫"