Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 30 additions & 3 deletions pyrit/converter/bijection_converter.py
Original file line number Diff line number Diff line change
Expand Up @@ -256,14 +256,19 @@ def get_teaching_instructions(self) -> str:
str: Instructions for using this digit mapping and its markers.
"""
mapping_str = ", ".join(f"{key}={value}" for key, value in self._mapping.items())
examples = "; ".join(f'"{text}" encodes to "{self.encode(prompt=text)}"' for text in ("it's", "I'm"))
examples = "; ".join(f'"{text}" encodes to "{self.encode(prompt=text)}"' for text in ("it's", "I'm", "top 10"))
return (
"Use this substitution notation for the rest of the conversation. Replace each lowercase English "
f"letter using this table of {self._num_digits}-digit tokens: {mapping_str}. "
"For uppercase letters, prefix the lowercase letter's token with one apostrophe "
f"({self._CASE_MARKER}). Encode each literal apostrophe as two apostrophes "
f"({self._CASE_MARKER * 2}). Preserve spaces and all other punctuation. "
f"({self._CASE_MARKER * 2}). Encode each literal digit as one tilde "
f"({self._LITERAL_MARKER}) followed by that digit, and each literal tilde as two tildes "
f"({self._LITERAL_MARKER * 2}), so that every unescaped run of digits is made only of table tokens. "
"Preserve spaces and all other punctuation. "
"Join adjacent digit tokens without separators. To decode, scan from left to right. "
"Consume a tilde together with the single character after it first: two tildes are one literal tilde, "
"and a tilde before a digit is that literal digit. "
"Consume doubled apostrophes as one literal apostrophe before checking for a single uppercase marker. "
"Reverse the table for each digit token, making its letter uppercase only when preceded by that marker. "
f"Examples: {examples}. When a user message is in this notation, read it by reversing these rules, "
Expand Down Expand Up @@ -299,9 +304,19 @@ def _build_identifier(self) -> ComponentIdentifier:
# doubled on encode and collapsed back on decode.
_CASE_MARKER = "'"

# A digit in the plaintext collides with the digit tokens the same way a literal
# apostrophe collides with _CASE_MARKER. Tokens are bare digit runs joined without
# separators, so a passed-through "123" is indistinguishable from encoded letters:
# decode() reads num_digits characters at a time and consumes whatever prefix of the
# literal number happens to be in the mapping ("abc 123 xyz" round-trips to
# "abc w3 xyz" whenever "12" is some letter's token). Every literal digit is therefore
# escaped with a marker that can never appear inside a token, which leaves every bare
# digit run in the encoded text a pure concatenation of tokens.
_LITERAL_MARKER = "~"

def encode(self, *, prompt: str) -> str:
"""
Encode text using digit tokens and uppercase markers.
Encode text using digit tokens, uppercase markers, and literal-digit escapes.

Args:
prompt (str): The prompt to encode.
Expand All @@ -316,6 +331,10 @@ def encode(self, *, prompt: str) -> str:
encoded += (self._CASE_MARKER + token) if char.isupper() else token
elif char == self._CASE_MARKER:
encoded += self._CASE_MARKER * 2
elif char == self._LITERAL_MARKER:
encoded += self._LITERAL_MARKER * 2
elif char in string.digits:
encoded += self._LITERAL_MARKER + char
else:
encoded += char
return encoded
Expand Down Expand Up @@ -352,6 +371,14 @@ def decode(self, encoded_text: str) -> str:
decoded = ""
i = 0
while i < len(encoded_text):
# Literal escapes bind to exactly one following character, so they are resolved
# before any digit run is considered for token lookup.
if encoded_text[i] == self._LITERAL_MARKER:
escaped = encoded_text[i + 1 : i + 2]
if escaped == self._LITERAL_MARKER or escaped in string.digits:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Could we check that escaped is non-empty before consuming it? If a model response ends with a lone ~, the slice is "", and Python evaluates "" in string.digits as True. The decoder then silently drops that marker instead of preserving it through the existing fallback:

converter = DigitBijectionConverter(seed=42)
converter.decode("~")  # Returns "" instead of "~"
converter.decode(converter.encode(prompt="top") + "~")  # Returns "top" instead of "top~"

This can happen with malformed or truncated responses. A guard would let an incomplete escape reach the fallback:

if escaped and (escaped == self._LITERAL_MARKER or escaped in string.digits):

Please also add direct decoder cases for "~" and "~~~". The round-trip tests cannot catch this because the encoder always doubles literal tildes.

decoded += escaped
i += 2
continue
if encoded_text[i] == self._CASE_MARKER and encoded_text[i + 1 : i + 2] == self._CASE_MARKER:
decoded += self._CASE_MARKER
i += 2
Expand Down
43 changes: 43 additions & 0 deletions tests/unit/converter/test_bijection_converter.py
Original file line number Diff line number Diff line change
Expand Up @@ -128,6 +128,49 @@ async def test_digit_converter_literal_apostrophe_round_trip():
assert converter.decode(encoded.output_text) == "it's"


async def test_digit_converter_literal_digit_round_trip():
"""A literal number survives encoding instead of being read back as letter tokens."""
custom_mapping = {letter: str(index + 10) for index, letter in enumerate(string.ascii_lowercase)}
converter = DigitBijectionConverter(mapping=custom_mapping)

encoded = await converter.convert_async(prompt="abc 123 xyz")

# Unescaped, "12" is c's token and the literal number would decode as "w3"/"c3".
assert encoded.output_text == "101112 ~1~2~3 333435"
assert converter.decode(encoded.output_text) == "abc 123 xyz"


async def test_digit_converter_literal_marker_round_trip():
"""The escape character itself is doubled so it can still be sent literally."""
custom_mapping = {letter: str(index + 10) for index, letter in enumerate(string.ascii_lowercase)}
converter = DigitBijectionConverter(mapping=custom_mapping)

encoded = await converter.convert_async(prompt="a~b")

assert encoded.output_text == "10~~11"
assert converter.decode(encoded.output_text) == "a~b"


@pytest.mark.parametrize("num_digits", [2, 3, 4])
@pytest.mark.parametrize(
"prompt",
["abc 123 xyz", "CVE-2021-44228", "pi is 3.14159", "BOB's 7 cats~", "~12", "0000000000"],
)
def test_digit_converter_round_trips_digit_bearing_prompts(num_digits: int, prompt: str):
"""Digit-bearing prompts round-trip for every mapping width, not just lucky mappings."""
for seed in range(12):
converter = DigitBijectionConverter(num_digits=num_digits, seed=seed)
assert converter.decode(converter.encode(prompt=prompt)) == prompt


def test_digit_converter_teaching_instructions_cover_literal_digits():
"""The target is told the escape rule, since it has to apply the same one."""
instructions = DigitBijectionConverter(seed=42).get_teaching_instructions()

assert "literal digit" in instructions
assert "~" in instructions


async def test_digit_converter_uppercase_letter_after_apostrophe_round_trip():
custom_mapping = {letter: str(index + 10) for index, letter in enumerate(string.ascii_lowercase)}
converter = DigitBijectionConverter(mapping=custom_mapping)
Expand Down