headroom/tests/test_models.py
julienguarino 17ecad9d89
fix(gemini): resolve Google model capabilities through ModelRegistry (#1276)
## Description

Google model capability lookup was still tied to static provider tables
for support checks and context limits. That made plausible future Gemini
model ids fail token counting or context lookup even when they clearly
belonged to the Google provider family.

This change adds a tolerant `ModelRegistry.resolve()` runtime lookup
path and routes the Google provider through it. Exact built-in registry
matches still win first, LiteLLM pricing metadata can supply live limits
when available, and provider-scoped family fallbacks cover future Gemini
ids without letting Google claim unrelated models.

## Type of Change

- [x] Bug fix (non-breaking change that fixes an issue)
- [ ] New feature (non-breaking change that adds functionality)
- [ ] Breaking change (fix or feature that would cause existing
functionality to change)
- [ ] Documentation update
- [ ] Performance improvement
- [ ] Code refactoring (no functional changes)

## Changes Made

- Added `ModelRegistry.resolve()` as a tolerant runtime capability
resolver.
- Added provider-scoped Google/Gemini family fallbacks for plausible
future model ids.
- Added support for LiteLLM-style `gemini/gemini-...` model ids in
provider inference and family fallback matching.
- Updated `GoogleProvider.supports_model()` and
`GoogleProvider.get_context_limit()` to use the shared model registry
path.
- Added regression tests for future Gemini ids, legacy Gemini context
limits, and unrelated model rejection.

## Testing

- [x] Unit tests pass (`pytest`)
- [x] Linting passes (`ruff check .`)
- [ ] Type checking passes (`mypy headroom`)
- [x] New tests added for new functionality
- [x] Manual testing performed

### Test Output

```text
uv run --no-project --with pytest --with opentelemetry-api --with pydantic --with tiktoken --with litellm --with click --with rich python -B -m pytest tests/test_provider_model_fallback.py tests/test_models.py

65 passed

uv run --no-project --with ruff ruff check headroom/models/registry.py headroom/providers/google.py tests/test_provider_model_fallback.py tests/test_models.py

All checks passed!

uv run --no-project --with ruff ruff format --check headroom/models/registry.py headroom/providers/google.py tests/test_provider_model_fallback.py tests/test_models.py

4 files already formatted
```

## Real Behavior Proof

- Environment: macOS arm64 local checkout, Python 3.13 virtualenv for
editable install; deployed smoke test in a Cloud Run staging service
using an earlier commit from this fork branch before the review
follow-up.
- Exact command / steps: installed `headroom-ai[langchain]` from the
fork branch in the staging service, triggered long-context requests that
activate Headroom's LangChain compression path, then checked Cloud Run
logs after 2026-06-22 12:20 Europe/Paris.
- Observed result: Headroom initialized successfully, compressed
conversation memory (`23255 -> 5618 chars`), and no logs matched the
previous model-resolution failure signatures (`not recognized as a
Google model`, `Unknown context limit`).
- Not tested: staging was not rerun after the `gemini/gemini-...` review
follow-up; that prefix path is covered by local regression tests. Full
repository `uv run pytest` on local macOS is currently blocked by a
native `maturin`/`esaxx-rs` compile failure (`fatal error: 'cstdint'
file not found`). Type checking was not run.

## 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
- [x] I have updated the CHANGELOG.md if applicable

## Screenshots (if applicable)

N/A

## Additional Notes

- Documentation changes are not included because this is a runtime
compatibility fix with no public API or user-facing configuration
change.
- Full local test execution should be retried in CI or a Linux
environment where the native Rust extension build is healthy.

Co-authored-by: Julien Guarino <julien.guarino@fashiondata.io>
2026-06-26 23:31:56 -05:00

296 lines
10 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_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