Spaces:
Build error
Build error
| """Tests for the pluggable tokenizer system.""" | |
| from __future__ import annotations | |
| import pytest | |
| from headroom.tokenizers import ( | |
| BaseTokenizer, | |
| CharacterCounter, | |
| EstimatingTokenCounter, | |
| TiktokenCounter, | |
| TokenCounter, | |
| TokenizerRegistry, | |
| get_mistral_tokenizer, | |
| get_tokenizer, | |
| is_mistral_tokenizer_available, | |
| list_supported_models, | |
| register_tokenizer, | |
| ) | |
| class TestTiktokenCounter: | |
| """Tests for TiktokenCounter.""" | |
| def test_init_default_model(self): | |
| """Test initialization with default model.""" | |
| counter = TiktokenCounter() | |
| assert counter.model == "gpt-4o" | |
| assert counter.encoding_name == "o200k_base" | |
| def test_init_gpt4_model(self): | |
| """Test initialization with GPT-4.""" | |
| counter = TiktokenCounter("gpt-4") | |
| assert counter.model == "gpt-4" | |
| assert counter.encoding_name == "cl100k_base" | |
| def test_count_text_empty(self): | |
| """Test counting empty text.""" | |
| counter = TiktokenCounter() | |
| assert counter.count_text("") == 0 | |
| def test_count_text_simple(self): | |
| """Test counting simple text.""" | |
| counter = TiktokenCounter() | |
| count = counter.count_text("Hello, world!") | |
| assert count > 0 | |
| assert count < 10 # Should be a few tokens | |
| def test_count_text_unicode(self): | |
| """Test counting text with unicode.""" | |
| counter = TiktokenCounter() | |
| count = counter.count_text("Hello, 世界!") | |
| assert count > 0 | |
| def test_count_messages_single(self): | |
| """Test counting single message.""" | |
| counter = TiktokenCounter() | |
| messages = [{"role": "user", "content": "Hello!"}] | |
| count = counter.count_messages(messages) | |
| assert count > 0 | |
| def test_count_messages_with_tool_calls(self): | |
| """Test counting messages with tool calls.""" | |
| counter = TiktokenCounter() | |
| messages = [ | |
| {"role": "user", "content": "Search for Python"}, | |
| { | |
| "role": "assistant", | |
| "tool_calls": [ | |
| { | |
| "id": "call_123", | |
| "type": "function", | |
| "function": { | |
| "name": "search", | |
| "arguments": '{"query": "Python"}', | |
| }, | |
| } | |
| ], | |
| }, | |
| { | |
| "role": "tool", | |
| "tool_call_id": "call_123", | |
| "content": "Results...", | |
| }, | |
| ] | |
| count = counter.count_messages(messages) | |
| assert count > 0 | |
| def test_encode_decode_roundtrip(self): | |
| """Test encode/decode roundtrip.""" | |
| counter = TiktokenCounter() | |
| text = "Hello, world!" | |
| tokens = counter.encode(text) | |
| decoded = counter.decode(tokens) | |
| assert decoded == text | |
| def test_repr(self): | |
| """Test string representation.""" | |
| counter = TiktokenCounter("gpt-4o") | |
| assert "TiktokenCounter" in repr(counter) | |
| assert "gpt-4o" in repr(counter) | |
| class TestEstimatingTokenCounter: | |
| """Tests for EstimatingTokenCounter.""" | |
| def test_init_default(self): | |
| """Test initialization with defaults.""" | |
| counter = EstimatingTokenCounter() | |
| assert counter._fixed_ratio is None | |
| def test_init_fixed_ratio(self): | |
| """Test initialization with fixed ratio.""" | |
| counter = EstimatingTokenCounter(chars_per_token=3.5) | |
| assert counter._fixed_ratio == 3.5 | |
| def test_count_text_empty(self): | |
| """Test counting empty text.""" | |
| counter = EstimatingTokenCounter() | |
| assert counter.count_text("") == 0 | |
| def test_count_text_simple(self): | |
| """Test counting simple text.""" | |
| counter = EstimatingTokenCounter() | |
| text = "Hello, world!" | |
| count = counter.count_text(text) | |
| assert count > 0 | |
| # Rough estimate: 13 chars / 4 chars per token ≈ 3-4 tokens | |
| assert 2 <= count <= 6 | |
| def test_count_text_fixed_ratio(self): | |
| """Test counting with fixed ratio.""" | |
| counter = EstimatingTokenCounter(chars_per_token=5.0) | |
| text = "x" * 50 # 50 chars | |
| count = counter.count_text(text) | |
| assert count == 10 # 50 / 5 = 10 | |
| def test_count_text_minimum_one(self): | |
| """Test minimum of 1 token.""" | |
| counter = EstimatingTokenCounter() | |
| assert counter.count_text("x") >= 1 | |
| def test_count_messages(self): | |
| """Test counting messages.""" | |
| counter = EstimatingTokenCounter() | |
| messages = [ | |
| {"role": "user", "content": "Hello!"}, | |
| {"role": "assistant", "content": "Hi there!"}, | |
| ] | |
| count = counter.count_messages(messages) | |
| assert count > 0 | |
| def test_json_detection(self): | |
| """Test JSON content detection.""" | |
| counter = EstimatingTokenCounter() | |
| json_text = '{"name": "test", "value": 123}' | |
| # Should use JSON ratio | |
| count = counter.count_text(json_text) | |
| assert count > 0 | |
| def test_code_detection(self): | |
| """Test code content detection.""" | |
| counter = EstimatingTokenCounter() | |
| code_text = """ | |
| def hello(): | |
| return "Hello, world!" | |
| """ | |
| count = counter.count_text(code_text) | |
| assert count > 0 | |
| def test_repr(self): | |
| """Test string representation.""" | |
| counter = EstimatingTokenCounter() | |
| assert "EstimatingTokenCounter" in repr(counter) | |
| class TestCharacterCounter: | |
| """Tests for CharacterCounter.""" | |
| def test_init_default(self): | |
| """Test initialization with default ratio.""" | |
| counter = CharacterCounter() | |
| assert counter.chars_per_token == 4.0 | |
| def test_init_custom_ratio(self): | |
| """Test initialization with custom ratio.""" | |
| counter = CharacterCounter(chars_per_token=3.5) | |
| assert counter.chars_per_token == 3.5 | |
| def test_count_text(self): | |
| """Test counting text.""" | |
| counter = CharacterCounter(chars_per_token=4.0) | |
| text = "x" * 40 # 40 chars | |
| count = counter.count_text(text) | |
| assert count == 10 # 40 / 4 = 10 | |
| def test_count_text_empty(self): | |
| """Test counting empty text.""" | |
| counter = CharacterCounter() | |
| assert counter.count_text("") == 0 | |
| class TestTokenizerRegistry: | |
| """Tests for TokenizerRegistry.""" | |
| def test_get_openai_model(self): | |
| """Test getting tokenizer for OpenAI model.""" | |
| tokenizer = get_tokenizer("gpt-4o") | |
| assert isinstance(tokenizer, TiktokenCounter) | |
| def test_get_anthropic_model(self): | |
| """Test getting tokenizer for Anthropic model.""" | |
| tokenizer = get_tokenizer("claude-3-sonnet") | |
| assert isinstance(tokenizer, EstimatingTokenCounter) | |
| def test_get_unknown_model_fallback(self): | |
| """Test fallback for unknown model.""" | |
| tokenizer = get_tokenizer("unknown-model-xyz") | |
| assert isinstance(tokenizer, EstimatingTokenCounter) | |
| def test_get_with_specific_backend(self): | |
| """Test forcing specific backend.""" | |
| tokenizer = get_tokenizer("any-model", backend="estimation") | |
| assert isinstance(tokenizer, EstimatingTokenCounter) | |
| def test_register_custom_tokenizer(self): | |
| """Test registering custom tokenizer.""" | |
| custom = EstimatingTokenCounter(chars_per_token=3.0) | |
| register_tokenizer("my-custom-model", tokenizer=custom) | |
| retrieved = get_tokenizer("my-custom-model") | |
| assert retrieved is custom | |
| def test_list_supported_models(self): | |
| """Test listing supported models.""" | |
| models = list_supported_models() | |
| assert isinstance(models, dict) | |
| assert "gpt-4o" in str(models) or "^gpt-4o" in str(models) | |
| def test_clear_cache(self): | |
| """Test clearing tokenizer cache.""" | |
| # Get a tokenizer to populate cache | |
| get_tokenizer("gpt-4o") | |
| # Clear cache | |
| TokenizerRegistry.clear_cache() | |
| # Should still work after clearing | |
| tokenizer = get_tokenizer("gpt-4o") | |
| assert tokenizer is not None | |
| class TestTokenCounterProtocol: | |
| """Tests for TokenCounter protocol.""" | |
| def test_tiktoken_implements_protocol(self): | |
| """Test TiktokenCounter implements protocol.""" | |
| counter = TiktokenCounter() | |
| assert isinstance(counter, TokenCounter) | |
| def test_estimating_implements_protocol(self): | |
| """Test EstimatingTokenCounter implements protocol.""" | |
| counter = EstimatingTokenCounter() | |
| assert isinstance(counter, TokenCounter) | |
| def test_character_implements_protocol(self): | |
| """Test CharacterCounter implements protocol.""" | |
| counter = CharacterCounter() | |
| assert isinstance(counter, TokenCounter) | |
| class TestBaseTokenizer: | |
| """Tests for BaseTokenizer base class.""" | |
| def test_message_overhead_constant(self): | |
| """Test message overhead constant.""" | |
| assert BaseTokenizer.MESSAGE_OVERHEAD == 4 | |
| def test_reply_overhead_constant(self): | |
| """Test reply overhead constant.""" | |
| assert BaseTokenizer.REPLY_OVERHEAD == 3 | |
| class TestMistralTokenizer: | |
| """Tests for Mistral tokenizer using official mistral-common.""" | |
| def test_is_available(self): | |
| """Test availability check.""" | |
| result = is_mistral_tokenizer_available() | |
| assert isinstance(result, bool) | |
| def test_get_mistral_tokenizer_class(self): | |
| """Test getting MistralTokenizer class.""" | |
| MistralTokenizer = get_mistral_tokenizer() | |
| assert MistralTokenizer is not None | |
| assert hasattr(MistralTokenizer, "count_text") | |
| def test_init_default_model(self): | |
| """Test initialization with default model.""" | |
| MistralTokenizer = get_mistral_tokenizer() | |
| counter = MistralTokenizer() | |
| assert counter.model == "mistral-large" | |
| assert counter.version == "v3" | |
| def test_init_mixtral_model(self): | |
| """Test initialization with Mixtral model (uses v1).""" | |
| MistralTokenizer = get_mistral_tokenizer() | |
| counter = MistralTokenizer("mixtral-8x7b") | |
| assert counter.version == "v1" | |
| def test_count_text_empty(self): | |
| """Test counting empty text.""" | |
| MistralTokenizer = get_mistral_tokenizer() | |
| counter = MistralTokenizer() | |
| assert counter.count_text("") == 0 | |
| def test_count_text_simple(self): | |
| """Test counting simple text.""" | |
| MistralTokenizer = get_mistral_tokenizer() | |
| counter = MistralTokenizer() | |
| count = counter.count_text("Hello, world!") | |
| assert count > 0 | |
| assert count < 10 | |
| def test_count_text_unicode(self): | |
| """Test counting text with unicode.""" | |
| MistralTokenizer = get_mistral_tokenizer() | |
| counter = MistralTokenizer() | |
| count = counter.count_text("Hello, 世界!") | |
| assert count > 0 | |
| def test_count_messages(self): | |
| """Test counting messages.""" | |
| MistralTokenizer = get_mistral_tokenizer() | |
| counter = MistralTokenizer() | |
| messages = [ | |
| {"role": "user", "content": "Hello!"}, | |
| {"role": "assistant", "content": "Hi there!"}, | |
| ] | |
| count = counter.count_messages(messages) | |
| assert count > 0 | |
| def test_count_messages_with_system(self): | |
| """Test counting messages with system prompt.""" | |
| MistralTokenizer = get_mistral_tokenizer() | |
| counter = MistralTokenizer() | |
| messages = [ | |
| {"role": "system", "content": "You are a helpful assistant."}, | |
| {"role": "user", "content": "Hello!"}, | |
| ] | |
| count = counter.count_messages(messages) | |
| assert count > 0 | |
| def test_encode_decode_roundtrip(self): | |
| """Test encode/decode roundtrip.""" | |
| MistralTokenizer = get_mistral_tokenizer() | |
| counter = MistralTokenizer() | |
| text = "Hello, world!" | |
| tokens = counter.encode(text) | |
| decoded = counter.decode(tokens) | |
| assert decoded == text | |
| def test_implements_protocol(self): | |
| """Test MistralTokenizer implements TokenCounter protocol.""" | |
| MistralTokenizer = get_mistral_tokenizer() | |
| counter = MistralTokenizer() | |
| assert isinstance(counter, TokenCounter) | |
| def test_repr(self): | |
| """Test string representation.""" | |
| MistralTokenizer = get_mistral_tokenizer() | |
| counter = MistralTokenizer("mistral-large") | |
| assert "MistralTokenizer" in repr(counter) | |
| assert "mistral-large" in repr(counter) | |
| def test_registry_returns_mistral_for_mistral_models(self): | |
| """Test registry returns Mistral tokenizer for Mistral models.""" | |
| tokenizer = get_tokenizer("mistral-large") | |
| MistralTokenizer = get_mistral_tokenizer() | |
| assert isinstance(tokenizer, MistralTokenizer) | |
| def test_registry_returns_mistral_for_mixtral(self): | |
| """Test registry returns Mistral tokenizer for Mixtral models.""" | |
| tokenizer = get_tokenizer("mixtral-8x7b") | |
| MistralTokenizer = get_mistral_tokenizer() | |
| assert isinstance(tokenizer, MistralTokenizer) | |
| def test_registry_returns_mistral_for_codestral(self): | |
| """Test registry returns Mistral tokenizer for Codestral models.""" | |
| tokenizer = get_tokenizer("codestral") | |
| MistralTokenizer = get_mistral_tokenizer() | |
| assert isinstance(tokenizer, MistralTokenizer) | |