Skip to content
Closed
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
2 changes: 1 addition & 1 deletion haystack/token_counters/tiktoken_counter.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,7 +74,7 @@ def count(self, messages: list[ChatMessage], tools: ToolsType | None = None) ->
if not messages and not tools:
return 0
self.warm_up()
text_tokens = len(self._encoder.encode(_rendered_conversation(messages) + _rendered_tools(tools)))
text_tokens = len(self._encoder.encode_ordinary(_rendered_conversation(messages) + _rendered_tools(tools)))
return text_tokens + _non_text_tokens(
messages, tokens_per_image=self.tokens_per_image, tokens_per_file=self.tokens_per_file
)
Expand Down
5 changes: 5 additions & 0 deletions releasenotes/notes/fix-tiktoken-literal-special-tokens.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
fixes:
- |
Fix ``TiktokenCounter`` raising ``ValueError`` when user messages or tool schemas contain a literal tiktoken
special-token marker such as ``<|endoftext|>``. These markers are now counted as ordinary text, matching the
behavior of the document splitters.
21 changes: 20 additions & 1 deletion test/token_counters/test_tiktoken_counter.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,11 +26,16 @@ class _FakeEncoder:

def __init__(self) -> None:
self.encoded: list[str] = []
self.ordinary_encoded: list[str] = []

def encode(self, text: str) -> list[int]:
self.encoded.append(text)
return list(range(len(text.split())))

def encode_ordinary(self, text: str) -> list[int]:
self.ordinary_encoded.append(text)
return list(range(len(text.split())))


@pytest.fixture
def fake_encoder(monkeypatch: pytest.MonkeyPatch) -> _FakeEncoder:
Expand Down Expand Up @@ -90,7 +95,7 @@ def test_encodes_the_rendered_conversation(self, fake_encoder):

TiktokenCounter().count(messages)

rendered = fake_encoder.encoded[0]
rendered = fake_encoder.ordinary_encoded[0]
assert "[assistant] looking" in rendered
assert '[assistant -> tool_call] search({"q": "x"})' in rendered
assert "[tool:search] found it" in rendered
Expand All @@ -116,6 +121,12 @@ def search(query: Annotated[str, "the search query"]) -> str:
return "result"


@tool
def search_with_literal_marker(query: Annotated[str, "the search query"]) -> str:
"""Search documentation that may contain the literal marker <|endoftext|>."""
return "result"


class TestTiktokenCounterTools:
def test_tool_schemas_add_to_the_count(self, fake_encoder):
# A provider is sent the schemas alongside the messages, so they consume tokens too.
Expand All @@ -127,6 +138,14 @@ def test_tool_schemas_add_to_the_count(self, fake_encoder):
def test_tools_can_be_counted_without_messages(self, fake_encoder):
assert TiktokenCounter().count([], tools=[search]) > 0

def test_literal_special_token_in_message_is_counted_as_text(self, fake_encoder):
assert TiktokenCounter().count([ChatMessage.from_user("literal <|endoftext|> marker")]) > 0
assert "<|endoftext|>" in fake_encoder.ordinary_encoded[0]

def test_literal_special_token_in_tool_schema_is_counted_as_text(self, fake_encoder):
assert TiktokenCounter().count([], tools=[search_with_literal_marker]) > 0
assert "<|endoftext|>" in fake_encoder.ordinary_encoded[0]

def test_nothing_to_measure_is_zero(self):
assert TiktokenCounter().count([]) == 0
assert TiktokenCounter().count([], tools=None) == 0
Expand Down