headroom/tests/test_models.py
Abhay Singh b699bedf95
fix(models): version-boundary longest-prefix match in ModelRegistry.get (#1658)
## Description

`ModelRegistry.get()` has a prefix fallback for versioned model ids. It
accepted
**any** registered name as a bare `str.startswith` prefix and returned
the
**first** match in dict-insertion order:

```python
for name, info in _MODELS.items():
    if model_lower.startswith(name):
        return info
```

Two concrete failures fall out of that:

- `gpt-4` is registered before `gpt-4-32k`, so `get("gpt-4-32k-0613")`
matches
`gpt-4` first and returns an **8192**-token window instead of
`gpt-4-32k`'s
  **32768**.
- `gpt-4.1` / `gpt-4.5-preview` aren't registered, so they also match
`gpt-4`
and inherit its **8192**-token window — even though they're much larger,
  distinct models.

`get_context_limit()` reads straight from `get()` (no LiteLLM fallback),
so both
cases make the proxy believe a nearly-empty context is almost full and
compress
far too aggressively — or reject — on requests that are actually small.
This is
silent: no error, just a wrong number driving every downstream
compression
decision for those models.

## Fix

The fallback now:

1. Only matches when the registered name ends at a **version boundary**
in the
query — the next character must be a separator (`-`, `/`, `:`, `@`, `_`)
— so
`gpt-4.1`'s `.` no longer matches `gpt-4` (it falls through to the
caller's
   default instead of a wrong 8192).
2. Picks the **longest** qualifying name, so `gpt-4-32k-0613` →
`gpt-4-32k`.

Exact and alias lookups are unchanged, and boundary-separated variants
like
`gpt-4o-new-version` still resolve to `gpt-4o`.

## Type of Change

- [x] Bug fix (non-breaking change that fixes an issue)

## Changes Made

- `headroom/models/registry.py`: replace the first-match `startswith`
prefix loop in `ModelRegistry.get` with a
longest-prefix-at-a-version-boundary match.
- `tests/test_models.py`: add regression tests — `gpt-4-32k-0613` →
`gpt-4-32k` (32768), and `gpt-4.1`/`gpt-4.5-preview` no longer resolve
to gpt-4's 8192 window.
- `CHANGELOG.md`: Bug Fixes entry under Unreleased.

## Testing

- [x] New tests added for the fixed behavior (`tests/test_models.py`)
- [x] Linting passes (`ruff check`) and formatting is clean (`ruff
format --check`)
- [ ] Full `pytest` run deferred to CI — see Real Behavior Proof for why
I verify the logic with a dependency-free script locally.

```text
$ uv run ruff check headroom/models/registry.py tests/test_models.py
All checks passed!
$ uv run ruff format --check headroom/models/registry.py tests/test_models.py
2 files already formatted
```

## Real Behavior Proof

- Environment: Windows 11, Python 3.12.11, headroom built from this
branch (`uv sync --extra dev`). Importing `headroom` pulls in the
torch/transformers stack; a full `pytest` run exhausts memory and gets
OOM-killed on this box, so I verify the matching logic with a
dependency-free script (only stdlib) and leave the full pytest to CI.
- Exact command / steps: replicated the relevant `_MODELS` registration
order (`gpt-4o`, `gpt-4-turbo`, `gpt-4`, `gpt-4-32k`) and the new
longest-prefix-with-boundary loop in a standalone script (no `headroom`
import), then asserted the resolved context windows.
- Observed result: `gpt-4-32k-0613` resolves to 32768 (was 8192 under
first-match), `gpt-4.1`/`gpt-4.5-preview` fall through to the caller
default (no longer 8192), and `gpt-4o-new-version` / `gpt-4` /
`gpt-4-0613` resolve exactly as before:

```text
OK: gpt-4-32k-0613 -> 32768 (was 8192 under old first-prefix-wins)
OK: gpt-4.1 / gpt-4.5-preview -> default (not 8192)
OK: gpt-4o-new-version, gpt-4, gpt-4-0613 still resolve as before
REGISTRY LOGIC VERIFIED
```

- Not tested: I did not add explicit registry entries for
`gpt-4.1`/`gpt-4.5` (their real windows) — that's a data addition,
separate from this matching-logic fix; today they fall back to the
caller's default, which is honest for an unregistered model and strictly
better than the previous wrong 8192. Full local `pytest` deferred to CI
(OOM, per above).

## Review Readiness

- [x] I have performed a self-review
- [x] This PR is ready for human review

## Checklist

- [x] My code follows the project's style guidelines
- [x] I have performed a self-review of my code
- [x] I have commented my code, particularly in hard-to-understand areas
- [ ] I have made corresponding changes to the documentation
- [x] My changes generate no new warnings
- [x] I have added tests that prove my fix is effective or that my
feature works
- [ ] New and existing unit tests pass locally with my changes — ran
lint + a standalone logic check; full pytest deferred to CI (local OOM,
disclosed above)
- [x] I have updated the CHANGELOG.md if applicable

## Additional Notes

- No new dependencies; pure logic change in one function plus tests.
- Found via a read-through of the registry while looking at how context
limits drive compression decisions.
2026-07-10 23:07:31 -05:00

314 lines
12 KiB
Python

"""Tests for the model registry and capabilities database."""
from __future__ import annotations
from unittest.mock import patch
import pytest
from headroom.models import (
ModelInfo,
ModelRegistry,
get_model_info,
list_models,
register_model,
)
class TestModelInfo:
"""Tests for ModelInfo dataclass."""
def test_default_values(self):
"""Test default values."""
info = ModelInfo(name="test", provider="test-provider")
assert info.context_window == 128000
assert info.max_output_tokens == 4096
assert info.supports_tools is True
assert info.supports_vision is False
assert info.supports_streaming is True
def test_custom_values(self):
"""Test custom values."""
info = ModelInfo(
name="custom-model",
provider="custom",
context_window=32000,
max_output_tokens=8192,
supports_tools=False,
supports_vision=True,
)
assert info.context_window == 32000
assert info.max_output_tokens == 8192
assert info.supports_tools is False
assert info.supports_vision is True
def test_frozen(self):
"""Test that ModelInfo is frozen (immutable)."""
info = ModelInfo(name="test", provider="test")
with pytest.raises(AttributeError):
info.name = "changed"
class TestModelRegistry:
"""Tests for ModelRegistry."""
def test_get_openai_model(self):
"""Test getting OpenAI model info."""
info = ModelRegistry.get("gpt-4o")
assert info is not None
assert info.provider == "openai"
assert info.context_window == 128000
def test_get_anthropic_model(self):
"""Test getting Anthropic model info."""
info = ModelRegistry.get("claude-3-5-sonnet-20241022")
assert info is not None
assert info.provider == "anthropic"
assert info.context_window == 200000
def test_get_google_model(self):
"""Test getting Google model info."""
info = ModelRegistry.get("gemini-1.5-pro")
assert info is not None
assert info.provider == "google"
assert info.context_window == 2000000 # 2M!
def test_get_by_alias(self):
"""Test getting model by alias."""
info = ModelRegistry.get("gpt-4o-2024-11-20")
assert info is not None
assert info.name == "gpt-4o"
def test_get_unknown_model(self):
"""Test getting unknown model returns None."""
info = ModelRegistry.get("unknown-model-xyz")
assert info is None
def test_get_prefix_matching(self):
"""Test prefix matching for versioned models."""
info = ModelRegistry.get("gpt-4o-new-version")
assert info is not None
assert info.name == "gpt-4o"
def test_get_prefix_matching_prefers_longest_registered_name(self):
"""`gpt-4-32k-0613` must resolve to `gpt-4-32k` (32768), not the
shorter `gpt-4` (8192) that is registered first."""
info = ModelRegistry.get("gpt-4-32k-0613")
assert info is not None
assert info.name == "gpt-4-32k"
assert ModelRegistry.get_context_limit("gpt-4-32k-0613") == 32768
def test_get_prefix_matching_requires_version_boundary(self):
"""`gpt-4.1`/`gpt-4.5` are distinct models, not variants of `gpt-4`.
A `.`-separated suffix must not match `gpt-4`, so they no longer
inherit gpt-4's 8192-token window (they fall back to the default)."""
assert ModelRegistry.get("gpt-4.1") is None
assert ModelRegistry.get("gpt-4.5-preview") is None
# Not silently reported as an 8192-token model:
assert ModelRegistry.get_context_limit("gpt-4.1") != 8192
assert ModelRegistry.get_context_limit("gpt-4.1", default=100) == 100
def test_resolve_future_google_family_fallback(self):
"""Resolve should return provider-scoped fallbacks for plausible future models."""
with patch("headroom.models.registry.get_model_pricing", return_value=None):
info = ModelRegistry.resolve("gemini-3-pro-preview", provider="google")
assert info is not None
assert info.provider == "google"
assert info.context_window == 1000000
assert info.tokenizer_backend == "google"
def test_resolve_google_litellm_prefixed_family_fallback(self):
"""Resolve should support LiteLLM-style Gemini provider prefixes."""
with patch("headroom.models.registry.get_model_pricing", return_value=None):
info = ModelRegistry.resolve("gemini/gemini-3-pro-preview", provider="google")
assert info is not None
assert info.provider == "google"
assert info.context_window == 1000000
assert info.tokenizer_backend == "google"
def test_resolve_does_not_claim_unrelated_models_for_google(self):
"""Provider-scoped resolution should not mask unrelated model catalogs."""
assert ModelRegistry.resolve("not-a-google-model", provider="google") is None
assert ModelRegistry.resolve("gpt-4o", provider="google") is None
def test_register_custom_model(self):
"""Test registering custom model."""
info = ModelRegistry.register(
"my-custom-model",
provider="custom",
context_window=64000,
supports_vision=True,
)
assert info.name == "my-custom-model"
assert info.provider == "custom"
assert info.context_window == 64000
# Should be retrievable
retrieved = ModelRegistry.get("my-custom-model")
assert retrieved is not None
assert retrieved.context_window == 64000
def test_list_models_all(self):
"""Test listing all models."""
models = ModelRegistry.list_models()
assert len(models) > 0
def test_list_models_by_provider(self):
"""Test listing models by provider."""
openai_models = ModelRegistry.list_models(provider="openai")
assert len(openai_models) > 0
assert all(m.provider == "openai" for m in openai_models)
def test_list_models_with_tools(self):
"""Test listing models with tool support."""
models = ModelRegistry.list_models(supports_tools=True)
assert len(models) > 0
assert all(m.supports_tools for m in models)
def test_list_models_with_vision(self):
"""Test listing models with vision support."""
models = ModelRegistry.list_models(supports_vision=True)
assert len(models) > 0
assert all(m.supports_vision for m in models)
def test_list_models_min_context(self):
"""Test listing models with minimum context."""
models = ModelRegistry.list_models(min_context=1000000)
assert len(models) > 0
assert all(m.context_window >= 1000000 for m in models)
def test_list_providers(self):
"""Test listing all providers."""
providers = ModelRegistry.list_providers()
assert "openai" in providers
assert "anthropic" in providers
assert "google" in providers
def test_get_context_limit(self):
"""Test getting context limit."""
limit = ModelRegistry.get_context_limit("gpt-4o")
assert limit == 128000
def test_get_context_limit_unknown(self):
"""Test getting context limit for unknown model."""
limit = ModelRegistry.get_context_limit("unknown", default=32000)
assert limit == 32000
def test_estimate_cost(self):
"""Test cost estimation."""
cost = ModelRegistry.estimate_cost(
model="gpt-4o",
input_tokens=1000000,
output_tokens=500000,
)
assert cost is not None
# GPT-4o: $2.50/1M input + $10.00/1M output * 0.5 = $2.50 + $5.00 = $7.50
assert abs(cost - 7.50) < 0.01
def test_estimate_cost_with_cache(self):
"""Test cost estimation with cached tokens.
Note: LiteLLM's basic cost estimation doesn't support cached token pricing.
The cached_tokens parameter is accepted but not currently factored into cost.
"""
cost = ModelRegistry.estimate_cost(
model="gpt-4o",
input_tokens=1000000,
output_tokens=0,
cached_tokens=500000, # Not currently used by LiteLLM
)
assert cost is not None
# With LiteLLM, all 1M tokens are charged at input rate: $2.50
assert abs(cost - 2.50) < 0.01
def test_estimate_cost_unknown_model(self):
"""Test cost estimation for unknown model."""
cost = ModelRegistry.estimate_cost(
model="unknown-model",
input_tokens=1000,
output_tokens=500,
)
assert cost is None
class TestConvenienceFunctions:
"""Tests for convenience functions."""
def test_get_model_info(self):
"""Test get_model_info function."""
info = get_model_info("gpt-4o")
assert info is not None
assert info.name == "gpt-4o"
def test_list_models(self):
"""Test list_models function."""
models = list_models(provider="anthropic")
assert len(models) > 0
def test_register_model(self):
"""Test register_model function."""
info = register_model(
"test-function-model",
provider="test",
context_window=16000,
)
assert info.name == "test-function-model"
class TestBuiltInModels:
"""Tests for built-in model data."""
def test_gpt4o_info(self):
"""Test GPT-4o model info."""
info = get_model_info("gpt-4o")
assert info.provider == "openai"
assert info.context_window == 128000
assert info.supports_tools is True
assert info.supports_vision is True
# Pricing is now fetched from LiteLLM, not stored in ModelInfo
pricing = ModelRegistry.get_pricing("gpt-4o")
assert pricing is not None
assert pricing[0] == 2.50 # input cost per 1M
assert pricing[1] == 10.00 # output cost per 1M
def test_o1_info(self):
"""Test o1 model info."""
info = get_model_info("o1")
assert info.provider == "openai"
assert info.context_window == 200000 # 200K context
assert info.max_output_tokens == 100000 # 100K output
def test_claude_info(self):
"""Test Claude model info."""
info = get_model_info("claude-3-5-sonnet-20241022")
assert info.provider == "anthropic"
assert info.context_window == 200000
# Pricing fetched from LiteLLM (falls back to alias for retired models)
pricing = ModelRegistry.get_pricing("claude-sonnet-4-20250514")
assert pricing is not None
assert pricing[0] == 3.00 # input cost per 1M
assert pricing[1] == 15.00 # output cost per 1M
# Retired model alias should also resolve
alias_pricing = ModelRegistry.get_pricing("claude-3-5-sonnet-20241022")
assert alias_pricing is not None
def test_gemini_info(self):
"""Test Gemini model info."""
info = get_model_info("gemini-1.5-pro")
assert info.provider == "google"
assert info.context_window == 2000000 # 2M tokens!
def test_llama_info(self):
"""Test Llama model info."""
info = get_model_info("llama-3.1-8b")
assert info.provider == "meta"
assert info.context_window == 128000
assert info.tokenizer_backend == "huggingface"
def test_mistral_info(self):
"""Test Mistral model info."""
info = get_model_info("mistral-large")
assert info.provider == "mistral"
assert info.supports_tools is True