Prepare for OSS release v0.2.0

This commit prepares Headroom for public open source release with
comprehensive documentation, licensing, and community infrastructure.

License & Legal:
- Add Apache 2.0 LICENSE file
- Add NOTICE file with third-party attributions
- Add SECURITY.md for vulnerability reporting

Community:
- Add CONTRIBUTING.md with contribution guidelines
- Add CODE_OF_CONDUCT.md (Contributor Covenant)
- Add GitHub issue templates (bug report, feature request)
- Add pull request template

Documentation:
- Update README.md with compelling value proposition
- Add docs/getting-started.md
- Add docs/proxy.md for proxy server documentation
- Add docs/transforms.md for transform reference
- Add docs/api.md for API reference
- Add examples/README.md

Package Infrastructure:
- Add headroom/py.typed for PEP 561 compliance
- Add headroom/cli.py for CLI entry point
- Add .github/workflows/ci.yml for CI pipeline
- Add .github/workflows/publish.yml for PyPI publishing
- Update pyproject.toml with proper metadata

New Features:
- Add multi-provider support (Google, Cohere, LiteLLM, OpenAI-compatible)
- Add universal tokenizer registry with multiple backends
- Add model registry with pricing and context limits
- Add production proxy server with caching and rate limiting

Code Quality:
- Fix 83 lint issues via ruff auto-fix
- Fix version consistency (benchmarks 0.1.0 → 0.2.0)
- Add skip decorators for optional dependency tests
This commit is contained in:
chopratejas 2026-01-07 11:36:44 -08:00
parent 9c7d4512d6
commit 175746cc26
67 changed files with 9184 additions and 291 deletions

7
.github/FUNDING.yml vendored Normal file
View file

@ -0,0 +1,7 @@
# These are supported funding model platforms
github: [headroom-sdk]
# patreon: headroom
# open_collective: headroom
# ko_fi: headroom
# custom: ["https://headroom.dev/sponsor"]

53
.github/ISSUE_TEMPLATE/bug_report.md vendored Normal file
View file

@ -0,0 +1,53 @@
---
name: Bug Report
about: Report a bug to help us improve Headroom
title: '[BUG] '
labels: bug
assignees: ''
---
## Description
A clear and concise description of what the bug is.
## To Reproduce
Steps to reproduce the behavior:
1. Install headroom with '...'
2. Run this code '...'
3. See error
## Expected Behavior
What you expected to happen.
## Actual Behavior
What actually happened.
## Code Sample
```python
# Minimal code to reproduce the issue
from headroom import HeadroomClient
# Your code here
```
## Error Output
```
Paste any error messages or stack traces here
```
## Environment
- **Headroom version**: (run `python -c "import headroom; print(headroom.__version__)"`)
- **Python version**: (run `python --version`)
- **OS**: (e.g., macOS 14.0, Ubuntu 22.04, Windows 11)
- **LLM Provider**: (e.g., OpenAI, Anthropic)
## Additional Context
Add any other context about the problem here (logs, screenshots, etc.)

8
.github/ISSUE_TEMPLATE/config.yml vendored Normal file
View file

@ -0,0 +1,8 @@
blank_issues_enabled: true
contact_links:
- name: Questions & Discussions
url: https://github.com/headroom-sdk/headroom/discussions
about: Ask questions and discuss ideas in GitHub Discussions
- name: Documentation
url: https://headroom.dev/docs
about: Check out the documentation for guides and API reference

View file

@ -0,0 +1,44 @@
---
name: Feature Request
about: Suggest a new feature for Headroom
title: '[FEATURE] '
labels: enhancement
assignees: ''
---
## Problem Statement
A clear description of the problem you're trying to solve.
Ex: "I'm always frustrated when..."
## Proposed Solution
Describe the solution you'd like. Be as specific as possible.
## Use Case
Explain your use case and why this feature would be valuable:
- What type of application are you building?
- How would this feature help you?
- How many tokens/cost would this save?
## Alternatives Considered
Describe any alternative solutions or features you've considered.
## Example API (Optional)
If you have ideas about how the API should look:
```python
# How you'd like to use this feature
from headroom import SomeNewFeature
# Example usage
```
## Additional Context
- Are you willing to contribute this feature?
- Any relevant links, papers, or prior art?

56
.github/PULL_REQUEST_TEMPLATE.md vendored Normal file
View file

@ -0,0 +1,56 @@
## Description
Brief description of changes and motivation.
Fixes #(issue number)
## Type of Change
- [ ] 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
- Change 1
- Change 2
- Change 3
## Testing
Describe the tests you ran to verify your changes:
- [ ] Unit tests pass (`pytest`)
- [ ] Linting passes (`ruff check .`)
- [ ] Type checking passes (`mypy headroom`)
- [ ] New tests added for new functionality
- [ ] Manual testing performed
## Test Output
```
# Paste relevant test output here
pytest -v tests/test_your_feature.py
```
## Checklist
- [ ] My code follows the project's style guidelines
- [ ] I have performed a self-review of my code
- [ ] I have commented my code, particularly in hard-to-understand areas
- [ ] I have made corresponding changes to the documentation
- [ ] My changes generate no new warnings
- [ ] 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
- [ ] I have updated the CHANGELOG.md if applicable
## Screenshots (if applicable)
Add screenshots to help explain your changes.
## Additional Notes
Any additional information that reviewers should know.

108
.github/workflows/ci.yml vendored Normal file
View file

@ -0,0 +1,108 @@
name: CI
on:
push:
branches: [main]
pull_request:
branches: [main]
jobs:
test:
runs-on: ubuntu-latest
strategy:
fail-fast: false
matrix:
python-version: ["3.10", "3.11", "3.12"]
steps:
- uses: actions/checkout@v4
- name: Set up Python ${{ matrix.python-version }}
uses: actions/setup-python@v5
with:
python-version: ${{ matrix.python-version }}
- name: Cache pip packages
uses: actions/cache@v4
with:
path: ~/.cache/pip
key: ${{ runner.os }}-pip-${{ matrix.python-version }}-${{ hashFiles('pyproject.toml') }}
restore-keys: |
${{ runner.os }}-pip-${{ matrix.python-version }}-
- name: Install dependencies
run: |
python -m pip install --upgrade pip
pip install -e ".[dev]"
- name: Run linting
run: |
ruff check .
ruff format --check .
- name: Run type checking
run: |
mypy headroom --ignore-missing-imports
- name: Run tests
run: |
pytest -v --tb=short
- name: Run tests with coverage
if: matrix.python-version == '3.11'
run: |
pytest --cov=headroom --cov-report=xml --cov-report=term-missing
- name: Upload coverage to Codecov
if: matrix.python-version == '3.11'
uses: codecov/codecov-action@v4
with:
file: ./coverage.xml
fail_ci_if_error: false
test-extras:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.11"
- name: Install with relevance extras
run: |
python -m pip install --upgrade pip
pip install -e ".[dev,relevance]"
- name: Run relevance tests
run: |
pytest tests/test_relevance.py -v
build:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.11"
- name: Install build tools
run: |
python -m pip install --upgrade pip build twine
- name: Build package
run: |
python -m build
- name: Check package
run: |
twine check dist/*
- name: Upload artifacts
uses: actions/upload-artifact@v4
with:
name: dist
path: dist/

31
.github/workflows/publish.yml vendored Normal file
View file

@ -0,0 +1,31 @@
name: Publish to PyPI
on:
release:
types: [published]
jobs:
publish:
runs-on: ubuntu-latest
environment: pypi
permissions:
id-token: write # For trusted publishing
steps:
- uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.11"
- name: Install build tools
run: |
python -m pip install --upgrade pip build
- name: Build package
run: |
python -m build
- name: Publish to PyPI
uses: pypa/gh-action-pypi-publish@release/v1

68
.gitignore vendored
View file

@ -20,9 +20,11 @@ parts/
sdist/
var/
wheels/
share/python-wheels/
*.egg-info/
.installed.cfg
*.egg
MANIFEST
# PyInstaller
*.manifest
@ -45,6 +47,7 @@ coverage.xml
*.py,cover
.hypothesis/
.pytest_cache/
pytest_cache/
# Translations
*.mo
@ -59,16 +62,19 @@ venv/
ENV/
env.bak/
venv.bak/
.python-version
# Secrets and API keys
# Secrets and API keys - NEVER commit these
*.pem
*.key
secrets.json
credentials.json
.secrets
api_keys.txt
.anthropic
.openai
# IDE
# IDE and editors
.idea/
.vscode/
*.swp
@ -77,36 +83,86 @@ api_keys.txt
.project
.pydevproject
.settings/
*.sublime-project
*.sublime-workspace
.spyproject
.spyderproject
# Jupyter Notebook
.ipynb_checkpoints
*.ipynb
# macOS
.DS_Store
.AppleDouble
.LSOverride
._*
# Thumbnails
Icon?
._*
# Windows
Thumbs.db
ehthumbs.db
Desktop.ini
# Linux
*~
# Local configuration
local_settings.py
*.local.py
*.local.json
*.local.yaml
# Database
# Database files
*.db
*.sqlite
*.sqlite3
# Logs
# Log files
*.log
logs/
log/
# Temporary files
tmp/
temp/
*.tmp
*.bak
*.swp
# Benchmark results (keep the framework, not results)
/tmp/
# Benchmark results (keep framework, not results)
.benchmarks/
benchmark_results.json
benchmark_results/
# DeepEval cache
.deepeval/
# Headroom specific
headroom.db
headroom_*.db
*.jsonl
!tests/fixtures/*.jsonl
# Documentation build
docs/_build/
site/
# mypy
.mypy_cache/
.dmypy.json
dmypy.json
# Ruff
.ruff_cache/
# pyright
pyrightconfig.json
# Editor backup files
*~
\#*\#
.\#*

120
CHANGELOG.md Normal file
View file

@ -0,0 +1,120 @@
# Changelog
All notable changes to Headroom will be documented in this file.
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.1.0/),
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
## [Unreleased]
### Added
- Production-ready proxy server with caching, rate limiting, and metrics
- CLI command `headroom proxy` to start the proxy server
## [0.2.0] - 2025-01-07
### Added
- **SmartCrusher**: Statistical compression for tool outputs
- Keeps first/last K items, errors, anomalies, and relevance matches
- Variance-based change point detection
- Pattern detection (time series, logs, search results)
- **Relevance Scoring Engine**: ML-powered item relevance
- `BM25Scorer`: Fast keyword matching (zero dependencies)
- `EmbeddingScorer`: Semantic similarity with sentence-transformers
- `HybridScorer`: Adaptive combination of both methods
- **CacheAligner**: Prefix stabilization for better cache hits
- Dynamic date extraction
- Whitespace normalization
- Stable prefix hashing
- **RollingWindow**: Context management within token limits
- Drops oldest tool units first
- Never orphans tool results
- Preserves recent turns
- **Multi-Provider Support**:
- Anthropic with official `count_tokens` API
- Google with official `countTokens` API
- Cohere with official `tokenize` API
- Mistral with official tokenizer
- LiteLLM for unified interface
- **Integrations**:
- LangChain callback handler (`HeadroomOptimizer`)
- MCP (Model Context Protocol) utilities
- **Proxy Server** (`headroom.proxy`):
- Semantic caching with LRU eviction
- Token bucket rate limiting
- Retry with exponential backoff
- Cost tracking with budget enforcement
- Prometheus metrics endpoint
- Request logging (JSONL)
- **Pricing Registry**: Centralized model pricing with staleness tracking
- **Benchmarks**: Performance benchmarks for transforms and relevance scoring
### Changed
- Improved token counting accuracy across all providers
- Enhanced tool output compression with relevance-aware selection
### Fixed
- Mistral tokenizer API compatibility
- Google token counting for multi-turn conversations
## [0.1.0] - 2025-01-05
### Added
- Initial release
- `HeadroomClient`: OpenAI-compatible client wrapper
- `ToolCrusher`: Basic tool output compression
- Audit mode for observation without modification
- Optimize mode for applying transforms
- Simulate mode for previewing changes
- SQLite and JSONL storage backends
- HTML report generation
- Streaming support
### Safety Guarantees
- Never removes human content
- Never breaks tool ordering
- Parse failures are no-ops
- Preserves recency (last N turns)
---
## Migration Guide
### From 0.1.x to 0.2.x
The 0.2.0 release is backward compatible. New features are opt-in:
```python
# Old code still works
from headroom import HeadroomClient, OpenAIProvider
# New SmartCrusher (replaces ToolCrusher for better compression)
from headroom import SmartCrusher, SmartCrusherConfig
config = SmartCrusherConfig(
min_tokens_to_crush=200,
max_items_after_crush=50,
)
crusher = SmartCrusher(config)
# New relevance scoring
from headroom import create_scorer
scorer = create_scorer("hybrid") # or "bm25" for zero deps
```
### Using the Proxy
New in 0.2.0 - run Headroom as a proxy server:
```bash
# Start the proxy
python -m headroom.proxy.server --port 8787
# Use with Claude Code
ANTHROPIC_BASE_URL=http://localhost:8787 claude
```
[Unreleased]: https://github.com/headroom-sdk/headroom/compare/v0.2.0...HEAD
[0.2.0]: https://github.com/headroom-sdk/headroom/compare/v0.1.0...v0.2.0
[0.1.0]: https://github.com/headroom-sdk/headroom/releases/tag/v0.1.0

133
CODE_OF_CONDUCT.md Normal file
View file

@ -0,0 +1,133 @@
# Contributor Covenant Code of Conduct
## Our Pledge
We as members, contributors, and leaders pledge to make participation in our
community a harassment-free experience for everyone, regardless of age, body
size, visible or invisible disability, ethnicity, sex characteristics, gender
identity and expression, level of experience, education, socio-economic status,
nationality, personal appearance, race, caste, color, religion, or sexual
identity and orientation.
We pledge to act and interact in ways that contribute to an open, welcoming,
diverse, inclusive, and healthy community.
## Our Standards
Examples of behavior that contributes to a positive environment for our
community include:
* Demonstrating empathy and kindness toward other people
* Being respectful of differing opinions, viewpoints, and experiences
* Giving and gracefully accepting constructive feedback
* Accepting responsibility and apologizing to those affected by our mistakes,
and learning from the experience
* Focusing on what is best not just for us as individuals, but for the overall
community
Examples of unacceptable behavior include:
* The use of sexualized language or imagery, and sexual attention or advances of
any kind
* Trolling, insulting or derogatory comments, and personal or political attacks
* Public or private harassment
* Publishing others' private information, such as a physical or email address,
without their explicit permission
* Other conduct which could reasonably be considered inappropriate in a
professional setting
## Enforcement Responsibilities
Community leaders are responsible for clarifying and enforcing our standards of
acceptable behavior and will take appropriate and fair corrective action in
response to any behavior that they deem inappropriate, threatening, offensive,
or harmful.
Community leaders have the right and responsibility to remove, edit, or reject
comments, commits, code, wiki edits, issues, and other contributions that are
not aligned to this Code of Conduct, and will communicate reasons for moderation
decisions when appropriate.
## Scope
This Code of Conduct applies within all community spaces, and also applies when
an individual is officially representing the community in public spaces.
Examples of representing our community include using an official email address,
posting via an official social media account, or acting as an appointed
representative at an online or offline event.
## Enforcement
Instances of abusive, harassing, or otherwise unacceptable behavior may be
reported to the community leaders responsible for enforcement at
**conduct@headroom.dev**.
All complaints will be reviewed and investigated promptly and fairly.
All community leaders are obligated to respect the privacy and security of the
reporter of any incident.
## Enforcement Guidelines
Community leaders will follow these Community Impact Guidelines in determining
the consequences for any action they deem in violation of this Code of Conduct:
### 1. Correction
**Community Impact**: Use of inappropriate language or other behavior deemed
unprofessional or unwelcome in the community.
**Consequence**: A private, written warning from community leaders, providing
clarity around the nature of the violation and an explanation of why the
behavior was inappropriate. A public apology may be requested.
### 2. Warning
**Community Impact**: A violation through a single incident or series of
actions.
**Consequence**: A warning with consequences for continued behavior. No
interaction with the people involved, including unsolicited interaction with
those enforcing the Code of Conduct, for a specified period of time. This
includes avoiding interactions in community spaces as well as external channels
like social media. Violating these terms may lead to a temporary or permanent
ban.
### 3. Temporary Ban
**Community Impact**: A serious violation of community standards, including
sustained inappropriate behavior.
**Consequence**: A temporary ban from any sort of interaction or public
communication with the community for a specified period of time. No public or
private interaction with the people involved, including unsolicited interaction
with those enforcing the Code of Conduct, is allowed during this period.
Violating these terms may lead to a permanent ban.
### 4. Permanent Ban
**Community Impact**: Demonstrating a pattern of violation of community
standards, including sustained inappropriate behavior, harassment of an
individual, or aggression toward or disparagement of classes of individuals.
**Consequence**: A permanent ban from any sort of public interaction within the
community.
## Attribution
This Code of Conduct is adapted from the [Contributor Covenant][homepage],
version 2.1, available at
[https://www.contributor-covenant.org/version/2/1/code_of_conduct.html][v2.1].
Community Impact Guidelines were inspired by
[Mozilla's code of conduct enforcement ladder][Mozilla CoC].
For answers to common questions about this code of conduct, see the FAQ at
[https://www.contributor-covenant.org/faq][FAQ]. Translations are available at
[https://www.contributor-covenant.org/translations][translations].
[homepage]: https://www.contributor-covenant.org
[v2.1]: https://www.contributor-covenant.org/version/2/1/code_of_conduct.html
[Mozilla CoC]: https://github.com/mozilla/diversity
[FAQ]: https://www.contributor-covenant.org/faq
[translations]: https://www.contributor-covenant.org/translations

209
CONTRIBUTING.md Normal file
View file

@ -0,0 +1,209 @@
# Contributing to Headroom
Thank you for your interest in contributing to Headroom! This document provides guidelines and instructions for contributing.
## Code of Conduct
By participating in this project, you agree to abide by our [Code of Conduct](CODE_OF_CONDUCT.md).
## How to Contribute
### Reporting Bugs
Before creating a bug report, please check existing issues to avoid duplicates. When creating a bug report, include:
- **Clear title** describing the issue
- **Steps to reproduce** the behavior
- **Expected behavior** vs what actually happened
- **Environment details** (Python version, OS, Headroom version)
- **Code samples** or minimal reproduction if possible
### Suggesting Features
Feature requests are welcome! Please:
- Check existing issues/discussions first
- Clearly describe the use case and motivation
- Explain how it fits with Headroom's goals (context optimization, safety, determinism)
### Pull Requests
1. **Fork the repository** and create your branch from `main`
2. **Install development dependencies**:
```bash
pip install -e ".[dev]"
```
3. **Make your changes** following our coding standards
4. **Add tests** for new functionality
5. **Run the test suite**:
```bash
pytest
```
6. **Run linting**:
```bash
ruff check .
ruff format .
```
7. **Update documentation** if needed
8. **Submit your PR** with a clear description
## Development Setup
```bash
# Clone the repository
git clone https://github.com/headroom-sdk/headroom.git
cd headroom
# Create a virtual environment
python -m venv .venv
source .venv/bin/activate # or `.venv\Scripts\activate` on Windows
# Install in development mode with all dependencies
pip install -e ".[dev,relevance,proxy]"
# Run tests
pytest
# Run tests with coverage
pytest --cov=headroom --cov-report=html
```
## Coding Standards
### Style
- We use [Ruff](https://github.com/astral-sh/ruff) for linting and formatting
- Line length: 100 characters
- Use type hints for all public functions
- Follow PEP 8 naming conventions
### Code Organization
```
headroom/
├── __init__.py # Public API exports
├── client.py # HeadroomClient wrapper
├── config.py # Configuration dataclasses
├── transforms/ # Context transforms
│ ├── smart_crusher.py # Statistical compression
│ ├── cache_aligner.py # Cache optimization
│ └── rolling_window.py# Context windowing
├── relevance/ # Relevance scoring
├── providers/ # LLM provider adapters
├── proxy/ # Proxy server
└── storage/ # Metrics storage
```
### Testing
- Write tests for all new functionality
- Use pytest fixtures for common setup
- Test edge cases and error conditions
- Aim for >80% coverage on new code
Example test structure:
```python
class TestSmartCrusher:
"""Tests for SmartCrusher transform."""
def test_compresses_large_arrays(self):
"""Should compress arrays above token threshold."""
...
def test_preserves_errors(self):
"""Should never drop items containing errors."""
...
```
### Documentation
- Add docstrings to all public classes and functions
- Use Google-style docstrings
- Update README.md for user-facing changes
- Add examples for new features
```python
def compress_tool_output(
content: str,
max_items: int = 50,
) -> str:
"""Compress tool output while preserving important items.
Args:
content: The tool output content (usually JSON).
max_items: Maximum items to keep in arrays.
Returns:
Compressed content string.
Raises:
ValueError: If content is not valid JSON.
Example:
>>> compress_tool_output('[{"id": 1}, {"id": 2}]', max_items=1)
'[{"id": 1}]'
"""
```
## Pull Request Guidelines
### PR Title Format
Use conventional commit style:
- `feat: Add semantic caching to proxy`
- `fix: Handle empty tool outputs correctly`
- `docs: Update proxy documentation`
- `test: Add tests for CacheAligner`
- `refactor: Simplify rolling window logic`
### PR Description
Include:
- **What** changes were made
- **Why** the changes were needed
- **How** to test the changes
- **Breaking changes** if any
### Review Process
1. All PRs require at least one review
2. CI must pass (tests, linting, type checking)
3. Maintain or improve test coverage
4. Update CHANGELOG.md for notable changes
## Architecture Decisions
### Safety First
Headroom's core principle is **safety**. When in doubt:
- Never drop user/assistant content
- Never break tool call/response pairing
- Malformed content passes through unchanged
- Prefer false negatives over false positives
### Performance
- Transforms should add <50ms latency at P99
- Use lazy loading for optional dependencies
- Profile before optimizing
### Compatibility
- Support Python 3.10+
- Core functionality has minimal dependencies
- Optional features use extras (e.g., `pip install headroom[relevance]`)
## Getting Help
- **Questions**: Open a [Discussion](https://github.com/headroom-sdk/headroom/discussions)
- **Bugs**: Open an [Issue](https://github.com/headroom-sdk/headroom/issues)
- **Security**: Email security@headroom.dev (do not open public issues)
## Recognition
Contributors are recognized in:
- The CHANGELOG for their contributions
- The GitHub contributors page
- Release notes for significant features
Thank you for contributing to Headroom!

190
LICENSE Normal file
View file

@ -0,0 +1,190 @@
Apache License
Version 2.0, January 2004
http://www.apache.org/licenses/
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
1. Definitions.
"License" shall mean the terms and conditions for use, reproduction,
and distribution as defined by Sections 1 through 9 of this document.
"Licensor" shall mean the copyright owner or entity authorized by
the copyright owner that is granting the License.
"Legal Entity" shall mean the union of the acting entity and all
other entities that control, are controlled by, or are under common
control with that entity. For the purposes of this definition,
"control" means (i) the power, direct or indirect, to cause the
direction or management of such entity, whether by contract or
otherwise, or (ii) ownership of fifty percent (50%) or more of the
outstanding shares, or (iii) beneficial ownership of such entity.
"You" (or "Your") shall mean an individual or Legal Entity
exercising permissions granted by this License.
"Source" form shall mean the preferred form for making modifications,
including but not limited to software source code, documentation
source, and configuration files.
"Object" form shall mean any form resulting from mechanical
transformation or translation of a Source form, including but
not limited to compiled object code, generated documentation,
and conversions to other media types.
"Work" shall mean the work of authorship, whether in Source or
Object form, made available under the License, as indicated by a
copyright notice that is included in or attached to the work
(an example is provided in the Appendix below).
"Derivative Works" shall mean any work, whether in Source or Object
form, that is based on (or derived from) the Work and for which the
editorial revisions, annotations, elaborations, or other modifications
represent, as a whole, an original work of authorship. For the purposes
of this License, Derivative Works shall not include works that remain
separable from, or merely link (or bind by name) to the interfaces of,
the Work and Derivative Works thereof.
"Contribution" shall mean any work of authorship, including
the original version of the Work and any modifications or additions
to that Work or Derivative Works thereof, that is intentionally
submitted to the Licensor for inclusion in the Work by the copyright owner
or by an individual or Legal Entity authorized to submit on behalf of
the copyright owner. For the purposes of this definition, "submitted"
means any form of electronic, verbal, or written communication sent
to the Licensor or its representatives, including but not limited to
communication on electronic mailing lists, source code control systems,
and issue tracking systems that are managed by, or on behalf of, the
Licensor for the purpose of discussing and improving the Work, but
excluding communication that is conspicuously marked or otherwise
designated in writing by the copyright owner as "Not a Contribution."
"Contributor" shall mean Licensor and any individual or Legal Entity
on behalf of whom a Contribution has been received by Licensor and
subsequently incorporated within the Work.
2. Grant of Copyright License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
copyright license to reproduce, prepare Derivative Works of,
publicly display, publicly perform, sublicense, and distribute the
Work and such Derivative Works in Source or Object form.
3. Grant of Patent License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
(except as stated in this section) patent license to make, have made,
use, offer to sell, sell, import, and otherwise transfer the Work,
where such license applies only to those patent claims licensable
by such Contributor that are necessarily infringed by their
Contribution(s) alone or by combination of their Contribution(s)
with the Work to which such Contribution(s) was submitted. If You
institute patent litigation against any entity (including a
cross-claim or counterclaim in a lawsuit) alleging that the Work
or a Contribution incorporated within the Work constitutes direct
or contributory patent infringement, then any patent licenses
granted to You under this License for that Work shall terminate
as of the date such litigation is filed.
4. Redistribution. You may reproduce and distribute copies of the
Work or Derivative Works thereof in any medium, with or without
modifications, and in Source or Object form, provided that You
meet the following conditions:
(a) You must give any other recipients of the Work or
Derivative Works a copy of this License; and
(b) You must cause any modified files to carry prominent notices
stating that You changed the files; and
(c) You must retain, in the Source form of any Derivative Works
that You distribute, all copyright, patent, trademark, and
attribution notices from the Source form of the Work,
excluding those notices that do not pertain to any part of
the Derivative Works; and
(d) If the Work includes a "NOTICE" text file as part of its
distribution, then any Derivative Works that You distribute must
include a readable copy of the attribution notices contained
within such NOTICE file, excluding those notices that do not
pertain to any part of the Derivative Works, in at least one
of the following places: within a NOTICE text file distributed
as part of the Derivative Works; within the Source form or
documentation, if provided along with the Derivative Works; or,
within a display generated by the Derivative Works, if and
wherever such third-party notices normally appear. The contents
of the NOTICE file are for informational purposes only and
do not modify the License. You may add Your own attribution
notices within Derivative Works that You distribute, alongside
or as an addendum to the NOTICE text from the Work, provided
that such additional attribution notices cannot be construed
as modifying the License.
You may add Your own copyright statement to Your modifications and
may provide additional or different license terms and conditions
for use, reproduction, or distribution of Your modifications, or
for any such Derivative Works as a whole, provided Your use,
reproduction, and distribution of the Work otherwise complies with
the conditions stated in this License.
5. Submission of Contributions. Unless You explicitly state otherwise,
any Contribution intentionally submitted for inclusion in the Work
by You to the Licensor shall be under the terms and conditions of
this License, without any additional terms or conditions.
Notwithstanding the above, nothing herein shall supersede or modify
the terms of any separate license agreement you may have executed
with Licensor regarding such Contributions.
6. Trademarks. This License does not grant permission to use the trade
names, trademarks, service marks, or product names of the Licensor,
except as required for reasonable and customary use in describing the
origin of the Work and reproducing the content of the NOTICE file.
7. Disclaimer of Warranty. Unless required by applicable law or
agreed to in writing, Licensor provides the Work (and each
Contributor provides its Contributions) on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
implied, including, without limitation, any warranties or conditions
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
PARTICULAR PURPOSE. You are solely responsible for determining the
appropriateness of using or redistributing the Work and assume any
risks associated with Your exercise of permissions under this License.
8. Limitation of Liability. In no event and under no legal theory,
whether in tort (including negligence), contract, or otherwise,
unless required by applicable law (such as deliberate and grossly
negligent acts) or agreed to in writing, shall any Contributor be
liable to You for damages, including any direct, indirect, special,
incidental, or consequential damages of any character arising as a
result of this License or out of the use or inability to use the
Work (including but not limited to damages for loss of goodwill,
work stoppage, computer failure or malfunction, or any and all
other commercial damages or losses), even if such Contributor
has been advised of the possibility of such damages.
9. Accepting Warranty or Additional Liability. While redistributing
the Work or Derivative Works thereof, You may choose to offer,
and charge a fee for, acceptance of support, warranty, indemnity,
or other liability obligations and/or rights consistent with this
License. However, in accepting such obligations, You may act only
on Your own behalf and on Your sole responsibility, not on behalf
of any other Contributor, and only if You agree to indemnify,
defend, and hold each Contributor harmless for any liability
incurred by, or claims asserted against, such Contributor by reason
of your accepting any such warranty or additional liability.
END OF TERMS AND CONDITIONS
Copyright 2025 Headroom Contributors
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.

43
NOTICE Normal file
View file

@ -0,0 +1,43 @@
Headroom
Copyright 2025 Headroom Contributors
This product includes software developed by the Headroom Contributors.
Third-Party Licenses
====================
This software uses the following third-party libraries:
tiktoken
--------
Copyright (c) 2022 OpenAI, Shantanu Jain
Licensed under the MIT License
https://github.com/openai/tiktoken
Pydantic
--------
Copyright (c) 2017 to present Pydantic Services Inc. and individual contributors
Licensed under the MIT License
https://github.com/pydantic/pydantic
sentence-transformers (optional dependency)
-------------------------------------------
Copyright 2019 Nils Reimers
Licensed under the Apache License 2.0
https://github.com/UKPLab/sentence-transformers
Note: Some pretrained sentence-transformer models may have additional licensing
restrictions based on their training data. Please verify model-specific licenses
before commercial use.
FastAPI (optional dependency)
-----------------------------
Copyright (c) 2018 Sebastián Ramírez
Licensed under the MIT License
https://github.com/tiangolo/fastapi
NumPy (optional dependency)
---------------------------
Copyright (c) 2005-2024, NumPy Developers
Licensed under the BSD 3-Clause License
https://github.com/numpy/numpy

408
README.md
View file

@ -1,285 +1,263 @@
# Headroom
<p align="center">
<h1 align="center">Headroom</h1>
<p align="center">
<strong>The Context Optimization Layer for LLM Applications</strong>
</p>
<p align="center">
Cut your LLM costs by 50-90% without losing accuracy
</p>
</p>
A safe, deterministic Context Budget Controller for LLM APIs.
<p align="center">
<a href="https://github.com/headroom-sdk/headroom/actions/workflows/ci.yml">
<img src="https://github.com/headroom-sdk/headroom/actions/workflows/ci.yml/badge.svg" alt="CI">
</a>
<a href="https://pypi.org/project/headroom/">
<img src="https://img.shields.io/pypi/v/headroom.svg" alt="PyPI">
</a>
<a href="https://pypi.org/project/headroom/">
<img src="https://img.shields.io/pypi/pyversions/headroom.svg" alt="Python">
</a>
<a href="https://github.com/headroom-sdk/headroom/blob/main/LICENSE">
<img src="https://img.shields.io/badge/license-Apache%202.0-blue.svg" alt="License">
</a>
</p>
**Increase effective TPM headroom. Reduce latency. Never break correctness.**
---
## Features
## The Problem
- **Context MRI (Audit Mode)**: Analyze context waste without modifying requests
- **Tool Output Compression**: Safely compress large tool outputs
- **Cache-Aligned Prefixes**: Optimize for provider caching (OpenAI, etc.)
- **Rolling Window Management**: Keep context within token limits
- **Streaming Support**: Full pass-through streaming with metrics
- **Simulate Mode**: Preview optimizations before applying
AI coding agents and tool-using applications generate **massive contexts**:
## Installation
- Tool outputs with 1000s of search results, log entries, API responses
- Long conversation histories that hit token limits
- System prompts with dynamic dates that break provider caching
**Result**: You pay for tokens you don't need, and cache hits are rare.
## The Solution
Headroom is a **smart compression layer** that sits between your app and LLM providers. It applies three transforms:
| Transform | What It Does | Savings |
|-----------|--------------|---------|
| **SmartCrusher** | Compresses tool outputs statistically (keeps errors, anomalies, relevant items) | 70-90% |
| **CacheAligner** | Stabilizes prefixes so provider caching works | Up to 10x |
| **RollingWindow** | Manages context within limits without breaking tool calls | Prevents failures |
**Zero accuracy loss** - we keep what matters: errors, anomalies, relevant items.
## Quick Start
### Option 1: Proxy (Recommended)
Run Headroom as a proxy server - works with any client:
```bash
pip install headroom
# Start the proxy
headroom proxy --port 8787
# Use with Claude Code
ANTHROPIC_BASE_URL=http://localhost:8787 claude
# Use with any OpenAI-compatible client
OPENAI_BASE_URL=http://localhost:8787/v1 your-app
```
Or install from source:
### Option 2: Python SDK
```bash
git clone https://github.com/headroom-sdk/headroom
cd headroom
pip install -e ".[dev]"
```
## Quick Start
Wrap your existing client:
```python
from headroom import HeadroomClient
from openai import OpenAI
# Wrap any OpenAI-compatible client
base = OpenAI(api_key="...")
client = HeadroomClient(
original_client=base,
store_url="sqlite:///headroom.db",
default_mode="audit", # Start in observation mode
original_client=OpenAI(),
default_mode="optimize",
)
# Use exactly like the original client
response = client.chat.completions.create(
model="gpt-4o",
messages=[
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hello!"},
],
messages=[...],
)
print(response.choices[0].message.content)
```
### Option 3: LangChain Integration
```python
from langchain_openai import ChatOpenAI
from headroom.integrations import HeadroomOptimizer
llm = ChatOpenAI(model="gpt-4o", callbacks=[HeadroomOptimizer()])
```
## Features
### Smart Tool Output Compression
```python
# Before: 50KB tool response with 1000 items
{"results": [{"id": 1, ...}, {"id": 2, ...}, ... 1000 items ...]}
# After: ~2KB with important items preserved
# - First 3 items (context)
# - Last 2 items (recency)
# - All error items
# - Anomalous values (> 2 std dev)
# - Items matching user's query
```
### Cache-Aligned Prefixes
```python
# Before: Cache miss every day due to changing date
"You are helpful. Today is January 7, 2025."
# After: Stable prefix (cache hit!) + dynamic context
"You are helpful."
# [Dynamic context moved to end]
```
### Rolling Window
```python
# Automatically manages context within token limits
# - Drops oldest tool outputs first
# - Never orphans tool call/response pairs
# - Always preserves system prompt and recent turns
```
### Production Proxy Features
- **Semantic Caching**: LRU cache with TTL for repeated queries
- **Rate Limiting**: Token bucket (requests + tokens per minute)
- **Cost Tracking**: Budget enforcement (hourly/daily/monthly)
- **Prometheus Metrics**: `/metrics` endpoint for monitoring
- **Request Logging**: JSONL logs for debugging
## Installation
```bash
# Core (minimal dependencies)
pip install headroom
# With semantic relevance scoring
pip install headroom[relevance]
# With proxy server
pip install headroom[proxy]
# Everything
pip install headroom[all]
```
## Modes
### Audit Mode (Default)
Observe and log without making changes:
### Audit Mode (Observe Only)
```python
client = HeadroomClient(
original_client=base,
default_mode="audit",
)
# Logs metrics to SQLite but doesn't modify requests
response = client.chat.completions.create(...)
client = HeadroomClient(original_client=base, default_mode="audit")
# Logs metrics but doesn't modify requests
```
### Optimize Mode
Apply safe, deterministic transforms:
### Optimize Mode (Apply Transforms)
```python
response = client.chat.completions.create(
model="gpt-4o",
messages=[...],
headroom_mode="optimize", # Enable optimization
)
client = HeadroomClient(original_client=base, default_mode="optimize")
# Applies safe, deterministic transforms
```
### Simulate Mode
Preview what optimizations would do:
### Simulate Mode (Preview)
```python
plan = client.chat.completions.simulate(
model="gpt-4o",
messages=[...],
)
print(f"Tokens before: {plan.tokens_before}")
print(f"Tokens after: {plan.tokens_after}")
print(f"Tokens saved: {plan.tokens_saved}")
print(f"Transforms: {plan.transforms}")
print(f"Estimated savings: {plan.estimated_savings}")
plan = client.chat.completions.simulate(model="gpt-4o", messages=[...])
print(f"Would save {plan.tokens_saved} tokens ({plan.savings_percent:.1f}%)")
```
## Configuration
### Headroom Parameters
All headroom parameters are optional:
```python
response = client.chat.completions.create(
model="gpt-4o",
messages=[...],
from headroom import HeadroomClient, SmartCrusherConfig
# Headroom-specific parameters
headroom_mode="optimize", # "audit" | "optimize"
headroom_output_buffer_tokens=4000, # Reserve for output
headroom_keep_turns=2, # Never drop last N turns
headroom_tool_profiles={ # Per-tool compression
"search": {"max_array_items": 5},
},
# All other OpenAI parameters work normally
temperature=0.7,
max_tokens=1000,
)
```
### Model Context Limits
Override default context limits:
```python
client = HeadroomClient(
original_client=base,
model_context_limits={
"gpt-4o": 128000,
"my-custom-model": 32000,
},
default_mode="optimize",
smart_crusher_config=SmartCrusherConfig(
min_tokens_to_crush=200, # Only compress if > 200 tokens
max_items_after_crush=50, # Keep at most 50 items
keep_first=3, # Always keep first 3
keep_last=2, # Always keep last 2
relevance_threshold=0.3, # Keep items with relevance > 0.3
),
)
```
## Transforms
## Supported Providers
### 1. Tool Output Compression
Compresses large tool outputs while preserving structure:
- Truncates long arrays (keeps first N items)
- Truncates long strings with markers
- Limits nesting depth
- **Safe**: Malformed JSON is never modified
```python
# Before: 50KB tool response
{"results": [{"id": 1, ...}, {"id": 2, ...}, ... 1000 items ...]}
# After: ~2KB with marker
{"results": [{"id": 1, ...}, ..., {"__headroom_truncated": 995}]}
<headroom:tool_digest sha256="abc123">
```
### 2. Cache Alignment
Stabilizes prefixes for better cache hit rates:
- Extracts dynamic dates from system prompts
- Normalizes whitespace
- Computes stable prefix hash
```python
# Before: Cache miss every day due to date
"You are helpful. Current Date: 2024-01-15"
# After: Stable prefix, date moved to context
"You are helpful.
[Context: Current Date: 2024-01-15]"
```
### 3. Rolling Window
Keeps context within token limits:
- Drops oldest tool call units first
- Never orphans tool responses
- Preserves system prompt and recent turns
- Inserts dropped context markers
## Reporting
Generate HTML reports of context waste:
```python
from headroom import generate_report
generate_report(
store_url="sqlite:///headroom.db",
output_path="report.html",
)
```
Reports include:
- Waste histogram by category
- Top high-waste requests
- Cache alignment analysis
- Actionable recommendations
| Provider | Token Counting | Status |
|----------|----------------|--------|
| OpenAI | tiktoken | Full support |
| Anthropic | Official API | Full support |
| Google | Official API | Full support |
| Cohere | Official API | Full support |
| Mistral | Official tokenizer | Full support |
| LiteLLM | Via provider | Full support |
## Safety Guarantees
Headroom follows strict safety rules:
1. **Never removes human content**: User/assistant text is sacred
2. **Never breaks tool ordering**: Tool calls and responses stay paired
3. **Parse failures are no-ops**: Malformed content passes through unchanged
4. **Preserves recency**: Last N turns are always kept
1. **Never removes human content** - User/assistant text is sacred
2. **Never breaks tool ordering** - Tool calls and responses stay paired
3. **Parse failures are no-ops** - Malformed content passes through unchanged
4. **Preserves recency** - Last N turns are always kept
## Streaming
## Benchmarks
Full streaming support:
| Scenario | Before | After | Savings |
|----------|--------|-------|---------|
| Search results (1000 items) | 45,000 tokens | 4,500 tokens | 90% |
| Log analysis (500 entries) | 22,000 tokens | 3,300 tokens | 85% |
| API response (nested JSON) | 15,000 tokens | 2,250 tokens | 85% |
| Long conversation (50 turns) | 80,000 tokens | 32,000 tokens | 60% |
```python
stream = client.chat.completions.create(
model="gpt-4o",
messages=[...],
stream=True,
headroom_mode="optimize",
)
## Documentation
for chunk in stream:
print(chunk.choices[0].delta.content, end="")
```
- [Getting Started Guide](docs/getting-started.md)
- [Proxy Server Documentation](docs/proxy.md)
- [Transform Reference](docs/transforms.md)
- [API Reference](docs/api.md)
- [Examples](examples/)
## Storage Options
## Contributing
### SQLite (Default)
```python
client = HeadroomClient(
original_client=base,
store_url="sqlite:///headroom.db",
)
```
### JSONL
```python
client = HeadroomClient(
original_client=base,
store_url="jsonl:///var/log/headroom.jsonl",
)
```
## Metrics
Access stored metrics programmatically:
```python
# Get recent metrics
metrics = client.get_metrics(limit=100)
# Get summary stats
summary = client.get_summary()
print(f"Total tokens saved: {summary['total_tokens_saved']}")
```
## Development
We welcome contributions! Please see our [Contributing Guide](CONTRIBUTING.md) for details.
```bash
# Install dev dependencies
# Development setup
git clone https://github.com/headroom-sdk/headroom.git
cd headroom
pip install -e ".[dev]"
# Run tests
pytest
# Run linter
ruff check .
# Type check
mypy headroom
```
## License
MIT
Apache License 2.0 - see [LICENSE](LICENSE) for details.
## Contributing
## Links
Contributions welcome! Please read the contributing guidelines first.
- [GitHub](https://github.com/headroom-sdk/headroom)
- [PyPI](https://pypi.org/project/headroom/)
- [Documentation](https://headroom.dev/docs)
- [Discord](https://discord.gg/headroom)
---
<p align="center">
<sub>Built with care for the AI developer community</sub>
</p>

65
SECURITY.md Normal file
View file

@ -0,0 +1,65 @@
# Security Policy
## Supported Versions
| Version | Supported |
| ------- | ------------------ |
| 0.2.x | :white_check_mark: |
| 0.1.x | :x: |
## Reporting a Vulnerability
We take security vulnerabilities seriously. If you discover a security issue, please report it responsibly.
### How to Report
**Please DO NOT open a public GitHub issue for security vulnerabilities.**
Instead, please email us at: **security@headroom.dev**
Include the following information:
- Type of vulnerability (e.g., injection, data exposure, authentication bypass)
- Full path of the affected source file(s)
- Step-by-step instructions to reproduce the issue
- Proof-of-concept or exploit code (if possible)
- Impact assessment
### What to Expect
1. **Acknowledgment**: We will acknowledge receipt within 48 hours
2. **Assessment**: We will assess the vulnerability and determine its severity
3. **Updates**: We will keep you informed of our progress
4. **Resolution**: We aim to resolve critical issues within 7 days
5. **Credit**: With your permission, we will credit you in the security advisory
### Security Best Practices for Users
When using Headroom:
1. **API Keys**: Never commit API keys. Use environment variables.
2. **Proxy Exposure**: Don't expose the proxy server to the public internet without authentication
3. **Log Files**: Be aware that request logs may contain sensitive information
4. **Budget Limits**: Set budget limits to prevent unexpected costs
### Scope
The following are in scope for security reports:
- Headroom Python package (`pip install headroom`)
- Headroom proxy server
- Official integrations (LangChain, MCP)
The following are out of scope:
- Third-party integrations not maintained by us
- Issues in dependencies (report these to the upstream project)
- Social engineering attacks
## Security Features
Headroom includes several security features:
- **No credential storage**: We never store or log API keys
- **Passthrough mode**: Sensitive content passes through unchanged by default
- **Input validation**: All inputs are validated before processing
- **Safe defaults**: Security-conscious defaults out of the box
Thank you for helping keep Headroom and its users safe!

View file

@ -21,7 +21,7 @@ Performance Targets:
- HybridScorer: < 50ms for 100 items (with embeddings)
"""
__version__ = "0.1.0"
__version__ = "0.2.0"
from .scenarios.tool_outputs import (
generate_api_responses,

28
docs/README.md Normal file
View file

@ -0,0 +1,28 @@
# Headroom Documentation
Welcome to the Headroom documentation.
## Quick Links
- [Getting Started](getting-started.md)
- [Proxy Server](proxy.md)
- [Transforms](transforms.md)
- [API Reference](api.md)
- [Architecture](ARCHITECTURE.md)
## Overview
Headroom is the Context Optimization Layer for LLM applications. It reduces your LLM costs by 50-90% through intelligent context compression.
### Core Concepts
1. **Transforms**: Stateless functions that modify message arrays to reduce tokens
2. **Providers**: Adapters for different LLM providers (OpenAI, Anthropic, etc.)
3. **Pipeline**: Chains multiple transforms together
4. **Proxy**: HTTP server that applies transforms transparently
### Getting Help
- [GitHub Issues](https://github.com/headroom-sdk/headroom/issues) - Bug reports
- [GitHub Discussions](https://github.com/headroom-sdk/headroom/discussions) - Questions
- [Discord](https://discord.gg/headroom) - Community chat

346
docs/api.md Normal file
View file

@ -0,0 +1,346 @@
# API Reference
## HeadroomClient
The main entry point for Headroom SDK.
```python
from headroom import HeadroomClient
from openai import OpenAI
client = HeadroomClient(
original_client=OpenAI(),
default_mode="optimize",
)
```
### Constructor Parameters
| Parameter | Type | Default | Description |
|-----------|------|---------|-------------|
| `original_client` | `OpenAI \| Anthropic` | Required | The underlying LLM client |
| `provider` | `Provider` | Auto-detected | Token counting provider |
| `default_mode` | `str` | `"audit"` | Default mode: "audit", "optimize", "off" |
| `store_url` | `str` | `None` | Storage URL for metrics |
| `smart_crusher_config` | `SmartCrusherConfig` | Default | Compression settings |
| `cache_aligner_config` | `CacheAlignerConfig` | Default | Cache alignment settings |
| `rolling_window_config` | `RollingWindowConfig` | Default | Context window settings |
### Methods
#### `chat.completions.create(**kwargs)`
Create a chat completion with optional optimization.
```python
response = client.chat.completions.create(
model="gpt-4o",
messages=[...],
headroom_mode="optimize", # Override default mode
)
```
**Additional Parameters:**
| Parameter | Type | Description |
|-----------|------|-------------|
| `headroom_mode` | `str` | Override mode for this request |
| `headroom_query` | `str` | Query for relevance scoring |
#### `chat.completions.simulate(**kwargs)`
Preview optimization without making an API call.
```python
plan = client.chat.completions.simulate(
model="gpt-4o",
messages=[...],
)
print(f"Tokens before: {plan.tokens_before}")
print(f"Tokens after: {plan.tokens_after}")
print(f"Savings: {plan.savings_percent:.1f}%")
```
**Returns:** `SimulationResult`
---
## Configuration Classes
### SmartCrusherConfig
```python
from headroom import SmartCrusherConfig
config = SmartCrusherConfig(
min_tokens_to_crush=200,
max_items_after_crush=50,
keep_first=3,
keep_last=2,
relevance_threshold=0.3,
anomaly_std_threshold=2.0,
preserve_errors=True,
)
```
### CacheAlignerConfig
```python
from headroom import CacheAlignerConfig
config = CacheAlignerConfig(
extract_dates=True,
normalize_whitespace=True,
stable_prefix_min_tokens=100,
)
```
### RollingWindowConfig
```python
from headroom import RollingWindowConfig
config = RollingWindowConfig(
max_tokens=100000,
preserve_system=True,
preserve_recent_turns=5,
drop_oldest_first=True,
)
```
### RelevanceScorerConfig
```python
from headroom import RelevanceScorerConfig
config = RelevanceScorerConfig(
scorer_type="bm25", # "bm25", "embedding", or "hybrid"
embedding_model=None, # Model name for embedding scorer
hybrid_alpha=0.5, # Weight for hybrid scoring
)
```
---
## Data Models
### SimulationResult
Returned by `simulate()`.
```python
@dataclass
class SimulationResult:
tokens_before: int
tokens_after: int
tokens_saved: int
savings_percent: float
transforms_applied: list[str]
waste_signals: WasteSignals
```
### RequestMetrics
Metrics for a single request.
```python
@dataclass
class RequestMetrics:
request_id: str
timestamp: datetime
model: str
tokens_input_before: int
tokens_input_after: int
tokens_output: int
cost_before: float
cost_after: float
transforms_applied: list[str]
```
### WasteSignals
Detected waste in the request.
```python
@dataclass
class WasteSignals:
json_bloat_tokens: int
html_noise_tokens: int
whitespace_tokens: int
dynamic_date_tokens: int
repetition_tokens: int
```
---
## Providers
### OpenAIProvider
```python
from headroom import OpenAIProvider
provider = OpenAIProvider()
# Get token counter
counter = provider.get_token_counter("gpt-4o")
tokens = counter.count_text("Hello, world!")
# Get context limit
limit = provider.get_context_limit("gpt-4o") # 128000
# Estimate cost
cost = provider.estimate_cost(
input_tokens=1000,
output_tokens=500,
model="gpt-4o",
)
```
### AnthropicProvider
```python
from headroom import AnthropicProvider
from anthropic import Anthropic
provider = AnthropicProvider(client=Anthropic())
counter = provider.get_token_counter("claude-3-5-sonnet-latest")
tokens = counter.count_messages(messages) # Accurate count via API
```
---
## Relevance Scoring
### BM25Scorer
Fast keyword-based scoring (zero dependencies).
```python
from headroom import BM25Scorer
scorer = BM25Scorer()
scores = scorer.score_items(
items=["item 1", "item 2", ...],
query="search query",
)
```
### EmbeddingScorer
Semantic similarity scoring (requires `sentence-transformers`).
```python
from headroom import EmbeddingScorer, embedding_available
if embedding_available():
scorer = EmbeddingScorer(model="all-MiniLM-L6-v2")
scores = scorer.score_items(items, query)
```
### HybridScorer
Combines BM25 and embeddings.
```python
from headroom import HybridScorer
scorer = HybridScorer(alpha=0.5) # 50% BM25, 50% embedding
scores = scorer.score_items(items, query)
```
### create_scorer()
Factory function to create scorers.
```python
from headroom import create_scorer
# Auto-select best available scorer
scorer = create_scorer()
# Explicitly choose type
scorer = create_scorer(scorer_type="hybrid", alpha=0.7)
```
---
## Transforms (Direct Use)
### SmartCrusher
```python
from headroom import SmartCrusher
crusher = SmartCrusher()
result = crusher.crush(
data={"results": [...]},
query="user query",
)
```
### CacheAligner
```python
from headroom import CacheAligner
aligner = CacheAligner()
result = aligner.align(messages)
```
### RollingWindow
```python
from headroom import RollingWindow
window = RollingWindow(config)
result = window.apply(messages, max_tokens=100000)
```
### TransformPipeline
```python
from headroom import TransformPipeline
pipeline = TransformPipeline([
SmartCrusher(),
CacheAligner(),
RollingWindow(),
])
result = pipeline.transform(messages)
```
---
## Utilities
### Tokenizer
```python
from headroom import Tokenizer, count_tokens_text, count_tokens_messages
# Quick counting
tokens = count_tokens_text("Hello, world!", model="gpt-4o")
# With tokenizer instance
tokenizer = Tokenizer(model="gpt-4o")
tokens = tokenizer.count_text("Hello")
tokens = tokenizer.count_messages(messages)
```
### generate_report()
Generate HTML/Markdown reports from stored metrics.
```python
from headroom import generate_report
report = generate_report(
store_url="sqlite:///headroom.db",
format="html",
period="day",
)
```

109
docs/getting-started.md Normal file
View file

@ -0,0 +1,109 @@
# Getting Started with Headroom
This guide will help you get up and running with Headroom in under 5 minutes.
## Installation
```bash
# Core package (minimal dependencies)
pip install headroom
# With proxy server
pip install headroom[proxy]
# With semantic relevance (for smarter compression)
pip install headroom[relevance]
# Everything
pip install headroom[all]
```
## Quick Start: Proxy Mode (Recommended)
The easiest way to use Headroom is as a proxy server:
```bash
# Start the proxy
headroom proxy --port 8787
```
Then point your LLM client at it:
```bash
# Claude Code
ANTHROPIC_BASE_URL=http://localhost:8787 claude
# OpenAI-compatible clients
OPENAI_BASE_URL=http://localhost:8787/v1 your-app
```
That's it! All your requests now go through Headroom and get optimized automatically.
## Quick Start: Python SDK
If you want programmatic control:
```python
from headroom import HeadroomClient
from openai import OpenAI
# Create a wrapped client
client = HeadroomClient(
original_client=OpenAI(),
default_mode="optimize",
)
# Use exactly like the original
response = client.chat.completions.create(
model="gpt-4o",
messages=[
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hello!"},
],
)
```
## Modes
### Audit Mode
Observe without modifying:
```python
client = HeadroomClient(
original_client=OpenAI(),
default_mode="audit",
)
# Logs metrics but doesn't change requests
```
### Optimize Mode
Apply transforms to reduce tokens:
```python
client = HeadroomClient(
original_client=OpenAI(),
default_mode="optimize",
)
# Compresses tool outputs, aligns cache prefixes, etc.
```
### Simulate Mode
Preview what optimizations would do:
```python
plan = client.chat.completions.simulate(
model="gpt-4o",
messages=[...],
)
print(f"Would save {plan.tokens_saved} tokens")
print(f"Transforms: {plan.transforms_applied}")
```
## Next Steps
- [Proxy Server Documentation](proxy.md) - Configure the proxy
- [Transforms Reference](transforms.md) - Understand each transform
- [API Reference](api.md) - Full API documentation

173
docs/proxy.md Normal file
View file

@ -0,0 +1,173 @@
# Proxy Server Documentation
The Headroom proxy server is a production-ready HTTP server that applies context optimization to all requests passing through it.
## Starting the Proxy
```bash
# Basic usage
headroom proxy
# Custom port
headroom proxy --port 8080
# With all options
headroom proxy \
--host 0.0.0.0 \
--port 8787 \
--log-file /var/log/headroom.jsonl \
--budget 100.0
```
## Command Line Options
| Option | Default | Description |
|--------|---------|-------------|
| `--host` | `127.0.0.1` | Host to bind to |
| `--port` | `8787` | Port to bind to |
| `--no-optimize` | `false` | Disable optimization (passthrough mode) |
| `--no-cache` | `false` | Disable semantic caching |
| `--no-rate-limit` | `false` | Disable rate limiting |
| `--log-file` | None | Path to JSONL log file |
| `--budget` | None | Daily budget limit in USD |
## API Endpoints
### Health Check
```bash
curl http://localhost:8787/health
```
Response:
```json
{
"status": "healthy",
"optimize": true,
"stats": {
"total_requests": 42,
"tokens_saved": 15000,
"savings_percent": 45.2
}
}
```
### Detailed Statistics
```bash
curl http://localhost:8787/stats
```
### Prometheus Metrics
```bash
curl http://localhost:8787/metrics
```
### LLM APIs
The proxy supports both Anthropic and OpenAI API formats:
```bash
# Anthropic format
POST /v1/messages
# OpenAI format
POST /v1/chat/completions
```
## Using with Claude Code
```bash
# Start proxy
headroom proxy --port 8787
# In another terminal
ANTHROPIC_BASE_URL=http://localhost:8787 claude
```
## Using with Cursor
1. Start the proxy: `headroom proxy`
2. In Cursor settings, set the base URL to `http://localhost:8787`
## Using with OpenAI SDK
```python
from openai import OpenAI
client = OpenAI(
base_url="http://localhost:8787/v1",
api_key="your-api-key", # Still needed for upstream
)
```
## Features
### Semantic Caching
The proxy caches responses for repeated queries:
- LRU eviction with configurable max entries
- TTL-based expiration
- Cache key based on message content hash
### Rate Limiting
Token bucket rate limiting protects against runaway costs:
- Configurable requests per minute
- Configurable tokens per minute
- Per-API-key tracking
### Cost Tracking
Track spending and enforce budgets:
- Real-time cost estimation
- Budget periods: hourly, daily, monthly
- Automatic request rejection when over budget
### Prometheus Metrics
Export metrics for monitoring:
```
headroom_requests_total
headroom_tokens_saved_total
headroom_cost_usd_total
headroom_latency_ms_sum
```
## Configuration via Environment
```bash
export HEADROOM_HOST=0.0.0.0
export HEADROOM_PORT=8787
export HEADROOM_BUDGET=100.0
headroom proxy
```
## Running in Production
For production deployments:
```bash
# Use a process manager
pip install gunicorn
# Run with gunicorn
gunicorn headroom.proxy.server:app \
--workers 4 \
--bind 0.0.0.0:8787 \
--worker-class uvicorn.workers.UvicornWorker
```
Or with Docker:
```dockerfile
FROM python:3.11-slim
RUN pip install headroom[proxy]
EXPOSE 8787
CMD ["headroom", "proxy", "--host", "0.0.0.0"]
```

198
docs/transforms.md Normal file
View file

@ -0,0 +1,198 @@
# Transform Reference
Headroom provides three core transforms that work together to optimize LLM context.
## SmartCrusher
Statistical compression for JSON tool outputs.
### How It Works
SmartCrusher analyzes JSON arrays and selectively keeps important items:
1. **First/Last items** - Context for pagination and recency
2. **Error items** - 100% preservation of error states
3. **Anomalies** - Statistical outliers (> 2 std dev from mean)
4. **Relevant items** - Matches to user's query via BM25/embeddings
5. **Change points** - Significant transitions in data
### Configuration
```python
from headroom import SmartCrusherConfig
config = SmartCrusherConfig(
min_tokens_to_crush=200, # Only compress if > 200 tokens
max_items_after_crush=50, # Keep at most 50 items
keep_first=3, # Always keep first 3 items
keep_last=2, # Always keep last 2 items
relevance_threshold=0.3, # Keep items with relevance > 0.3
anomaly_std_threshold=2.0, # Keep items > 2 std dev from mean
preserve_errors=True, # Always keep error items
)
```
### Example
```python
from headroom import SmartCrusher
crusher = SmartCrusher(config)
# Before: 1000 search results (45,000 tokens)
tool_output = {"results": [...1000 items...]}
# After: ~50 important items (4,500 tokens) - 90% reduction
compressed = crusher.crush(tool_output, query="user's question")
```
### What Gets Preserved
| Category | Preserved | Why |
|----------|-----------|-----|
| Errors | 100% | Critical for debugging |
| First N | 100% | Context/pagination |
| Last N | 100% | Recency |
| Anomalies | All | Unusual values matter |
| Relevant | Top K | Match user's query |
| Others | Sampled | Statistical representation |
---
## CacheAligner
Prefix stabilization for improved cache hit rates.
### The Problem
LLM providers cache request prefixes. But dynamic content breaks caching:
```
"You are helpful. Today is January 7, 2025." # Changes daily = no cache
```
### The Solution
CacheAligner extracts dynamic content to stabilize the prefix:
```python
from headroom import CacheAligner
aligner = CacheAligner()
result = aligner.align(messages)
# Static prefix (cacheable):
# "You are helpful."
# Dynamic content moved to end:
# [Current date context]
```
### Configuration
```python
from headroom import CacheAlignerConfig
config = CacheAlignerConfig(
extract_dates=True, # Move dates to dynamic section
normalize_whitespace=True, # Consistent spacing
stable_prefix_min_tokens=100, # Min prefix size for alignment
)
```
### Cache Hit Improvement
| Scenario | Before | After |
|----------|--------|-------|
| Daily date in prompt | 0% hits | ~95% hits |
| Dynamic user context | ~10% hits | ~80% hits |
| Consistent prompts | ~90% hits | ~95% hits |
---
## RollingWindow
Context management within token limits.
### The Problem
Long conversations exceed context limits. Naive truncation breaks tool calls:
```
[tool_call: search] # Kept
[tool_result: ...] # Dropped = orphaned call!
```
### The Solution
RollingWindow drops complete tool units, preserving pairs:
```python
from headroom import RollingWindow
window = RollingWindow(config)
result = window.apply(messages, max_tokens=100000)
# Guarantees:
# 1. Tool calls paired with results
# 2. System prompt preserved
# 3. Recent turns kept
# 4. Oldest tool outputs dropped first
```
### Configuration
```python
from headroom import RollingWindowConfig
config = RollingWindowConfig(
max_tokens=100000, # Target token limit
preserve_system=True, # Always keep system prompt
preserve_recent_turns=5, # Keep last 5 user/assistant turns
drop_oldest_first=True, # Remove oldest tool outputs
)
```
### Drop Priority
1. **Oldest tool outputs** - First to go
2. **Old assistant messages** - Summary preserved
3. **Old user messages** - Only if necessary
4. **Never dropped**: System prompt, recent turns, active tool pairs
---
## TransformPipeline
Combine transforms for optimal results.
```python
from headroom import TransformPipeline, SmartCrusher, CacheAligner, RollingWindow
pipeline = TransformPipeline([
SmartCrusher(), # First: compress tool outputs
CacheAligner(), # Then: stabilize prefix
RollingWindow(), # Finally: fit in context
])
result = pipeline.transform(messages)
print(f"Saved {result.tokens_saved} tokens")
```
### Recommended Order
1. **SmartCrusher** - Reduce individual messages
2. **CacheAligner** - Optimize for caching
3. **RollingWindow** - Final size constraint
---
## Safety Guarantees
All transforms follow strict safety rules:
1. **Never remove human content** - User/assistant text is sacred
2. **Never break tool ordering** - Calls and results stay paired
3. **Parse failures are no-ops** - Malformed content passes through
4. **Preserves recency** - Last N turns always kept
5. **100% error preservation** - Error items never dropped

133
examples/README.md Normal file
View file

@ -0,0 +1,133 @@
# Headroom Examples
This directory contains examples demonstrating Headroom's capabilities.
## Quick Start Examples
### basic_usage.py
Basic integration with OpenAI client:
```bash
export OPENAI_API_KEY='your-key'
python examples/basic_usage.py
```
### anthropic_example.py
Integration with Anthropic Claude:
```bash
export ANTHROPIC_API_KEY='your-key'
python examples/anthropic_example.py
```
### streaming_example.py
Streaming responses with optimization:
```bash
export OPENAI_API_KEY='your-key'
python examples/streaming_example.py
```
## Evaluation Examples
### smart_vs_naive_eval.py
Compare SmartCrusher against naive truncation:
```bash
export OPENAI_API_KEY='your-key'
python examples/smart_vs_naive_eval.py
```
### real_world_eval.py
Comprehensive evaluation with Anthropic models:
```bash
export ANTHROPIC_API_KEY='your-key'
python examples/real_world_eval.py
```
### real_world_openai_eval.py
Comprehensive evaluation with OpenAI models:
```bash
export OPENAI_API_KEY='your-key'
python examples/real_world_openai_eval.py
```
## Demo Directories
### langchain_demo/
Full LangChain agent integration demo:
```bash
# No API key needed for compression demo
PYTHONPATH=. python -m examples.langchain_demo.show_compression
# Full comparison (requires API key)
export OPENAI_API_KEY='your-key'
PYTHONPATH=. python -m examples.langchain_demo.run_comparison
```
See [langchain_demo/README.md](langchain_demo/README.md) for details.
### mcp_demo/
MCP (Model Context Protocol) integration demo:
```bash
export OPENAI_API_KEY='your-key'
PYTHONPATH=. python -m examples.mcp_demo.run_agent_eval
```
## Running Examples
All examples can be run from the repository root:
```bash
# Install dependencies
pip install -e ".[dev]"
# Run any example
python examples/<example_name>.py
```
## Expected Results
| Example | Token Savings | Notes |
|---------|---------------|-------|
| basic_usage | 50-70% | Simple tool output compression |
| langchain_demo | 70-85% | Real agent with multiple tools |
| mcp_demo | 60-80% | MCP tool outputs |
| real_world_eval | 50-90% | Varies by scenario |
## Troubleshooting
**ModuleNotFoundError: No module named 'headroom'**
Run from the repository root with PYTHONPATH:
```bash
PYTHONPATH=. python examples/basic_usage.py
```
Or install in development mode:
```bash
pip install -e .
```
**API Key Errors**
Ensure your API keys are set:
```bash
export OPENAI_API_KEY='sk-...'
export ANTHROPIC_API_KEY='sk-ant-...'
```

View file

@ -57,15 +57,6 @@ from .config import (
WasteSignals,
)
from .providers import AnthropicProvider, OpenAIProvider, Provider, TokenCounter
from .reporting import generate_report
from .tokenizer import Tokenizer, count_tokens_messages, count_tokens_text
from .transforms import (
CacheAligner,
RollingWindow,
SmartCrusher,
ToolCrusher,
TransformPipeline,
)
from .relevance import (
BM25Scorer,
EmbeddingScorer,
@ -75,6 +66,15 @@ from .relevance import (
create_scorer,
embedding_available,
)
from .reporting import generate_report
from .tokenizer import Tokenizer, count_tokens_messages, count_tokens_text
from .transforms import (
CacheAligner,
RollingWindow,
SmartCrusher,
ToolCrusher,
TransformPipeline,
)
__version__ = "0.2.0"

185
headroom/cli.py Normal file
View file

@ -0,0 +1,185 @@
#!/usr/bin/env python3
"""Headroom CLI - The Context Optimization Layer for LLM Applications.
Usage:
headroom proxy [OPTIONS] Start the optimization proxy server
headroom --version Show version
headroom --help Show this help message
Examples:
# Start proxy on default port (8787)
headroom proxy
# Start proxy on custom port
headroom proxy --port 8080
# Start with optimization disabled (passthrough mode)
headroom proxy --no-optimize
# Use with Claude Code
ANTHROPIC_BASE_URL=http://localhost:8787 claude
"""
from __future__ import annotations
import argparse
import sys
def get_version() -> str:
"""Get the current version."""
try:
from headroom import __version__
return __version__
except ImportError:
return "unknown"
def cmd_proxy(args: argparse.Namespace) -> int:
"""Start the proxy server."""
try:
from headroom.proxy.server import ProxyConfig, run_server
except ImportError as e:
print("Error: Proxy dependencies not installed. Run: pip install headroom[proxy]")
print(f"Details: {e}")
return 1
config = ProxyConfig(
host=args.host,
port=args.port,
optimize=not args.no_optimize,
cache_enabled=not args.no_cache,
rate_limit_enabled=not args.no_rate_limit,
log_file=args.log_file,
budget_limit_usd=args.budget,
)
print(f"""
HEADROOM PROXY
The Context Optimization Layer for LLM Applications
Starting proxy server...
URL: http://{config.host}:{config.port}
Optimization: {'ENABLED' if config.optimize else 'DISABLED'}
Caching: {'ENABLED' if config.cache_enabled else 'DISABLED'}
Rate Limit: {'ENABLED' if config.rate_limit_enabled else 'DISABLED'}
Usage with Claude Code:
ANTHROPIC_BASE_URL=http://{config.host}:{config.port} claude
Usage with OpenAI-compatible clients:
OPENAI_BASE_URL=http://{config.host}:{config.port}/v1 your-app
Endpoints:
GET /health Health check
GET /stats Detailed statistics
GET /metrics Prometheus metrics
POST /v1/messages Anthropic API
POST /v1/chat/completions OpenAI API
Press Ctrl+C to stop.
""")
try:
run_server(config)
except KeyboardInterrupt:
print("\nShutting down...")
return 0
return 0
def cmd_version(args: argparse.Namespace) -> int:
"""Print version information."""
print(f"headroom {get_version()}")
return 0
def main(argv: list[str] | None = None) -> int:
"""Main CLI entry point."""
parser = argparse.ArgumentParser(
prog="headroom",
description="The Context Optimization Layer for LLM Applications",
formatter_class=argparse.RawDescriptionHelpFormatter,
epilog="""
Examples:
headroom proxy Start proxy on port 8787
headroom proxy --port 8080 Start proxy on port 8080
headroom proxy --no-optimize Passthrough mode (no optimization)
Environment Variables:
ANTHROPIC_API_KEY Your Anthropic API key (for proxying)
OPENAI_API_KEY Your OpenAI API key (for proxying)
Documentation: https://github.com/headroom-sdk/headroom
""",
)
parser.add_argument(
"--version", "-V",
action="store_true",
help="Show version and exit",
)
subparsers = parser.add_subparsers(dest="command", help="Commands")
# Proxy command
proxy_parser = subparsers.add_parser(
"proxy",
help="Start the optimization proxy server",
formatter_class=argparse.RawDescriptionHelpFormatter,
)
proxy_parser.add_argument(
"--host",
default="127.0.0.1",
help="Host to bind to (default: 127.0.0.1)",
)
proxy_parser.add_argument(
"--port", "-p",
type=int,
default=8787,
help="Port to bind to (default: 8787)",
)
proxy_parser.add_argument(
"--no-optimize",
action="store_true",
help="Disable optimization (passthrough mode)",
)
proxy_parser.add_argument(
"--no-cache",
action="store_true",
help="Disable semantic caching",
)
proxy_parser.add_argument(
"--no-rate-limit",
action="store_true",
help="Disable rate limiting",
)
proxy_parser.add_argument(
"--log-file",
help="Path to JSONL log file",
)
proxy_parser.add_argument(
"--budget",
type=float,
help="Daily budget limit in USD",
)
proxy_parser.set_defaults(func=cmd_proxy)
args = parser.parse_args(argv)
if args.version:
return cmd_version(args)
if args.command is None:
parser.print_help()
return 0
return args.func(args)
if __name__ == "__main__":
sys.exit(main())

View file

@ -2,8 +2,9 @@
from __future__ import annotations
from collections.abc import Iterator
from datetime import datetime
from typing import Any, Iterator
from typing import Any
from .config import (
HeadroomConfig,

View file

@ -301,8 +301,8 @@ class TransformResult:
transforms_applied: list[str]
markers_inserted: list[str] = field(default_factory=list)
warnings: list[str] = field(default_factory=list)
diff_artifact: "DiffArtifact | None" = None # Populated if generate_diff_artifact=True
cache_metrics: "CachePrefixMetrics | None" = None # Populated by CacheAligner
diff_artifact: DiffArtifact | None = None # Populated if generate_diff_artifact=True
cache_metrics: CachePrefixMetrics | None = None # Populated by CacheAligner
@dataclass

View file

@ -8,21 +8,20 @@ Install LangChain support: pip install headroom[langchain]
"""
from .langchain import (
HeadroomChatModel,
HeadroomCallbackHandler,
optimize_messages,
HeadroomChatModel,
HeadroomRunnable,
optimize_messages,
)
from .mcp import (
HeadroomMCPCompressor,
DEFAULT_MCP_PROFILES,
HeadroomMCPClientWrapper,
HeadroomMCPCompressor,
MCPCompressionResult,
MCPToolProfile,
compress_tool_result,
compress_tool_result_with_metrics,
create_headroom_mcp_proxy,
DEFAULT_MCP_PROFILES,
)
__all__ = [

View file

@ -29,9 +29,10 @@ from __future__ import annotations
import json
import logging
from collections.abc import Iterator, Sequence
from dataclasses import dataclass
from datetime import datetime
from typing import Any, Iterator, List, Optional, Sequence, Union
from typing import Any
from uuid import uuid4
# LangChain imports - these are optional dependencies
@ -378,7 +379,7 @@ class HeadroomChatModel(BaseChatModel):
**kwargs,
)
def bind_tools(self, tools: Sequence[Any], **kwargs) -> "HeadroomChatModel":
def bind_tools(self, tools: Sequence[Any], **kwargs) -> HeadroomChatModel:
"""Bind tools to the wrapped model."""
new_wrapped = self.wrapped_model.bind_tools(tools, **kwargs)
return HeadroomChatModel(

View file

@ -48,12 +48,13 @@ from __future__ import annotations
import json
import re
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import Any, Callable
from typing import Any
from headroom.config import HeadroomConfig, SmartCrusherConfig
from headroom.transforms import SmartCrusher
from headroom.providers import OpenAIProvider
from headroom.transforms import SmartCrusher
@dataclass

View file

@ -0,0 +1,39 @@
"""Model registry and capabilities database.
Provides a centralized registry of LLM models with their capabilities,
context limits, pricing, and provider information.
Usage:
from headroom.models import ModelRegistry, get_model_info
# Get info about a model
info = get_model_info("gpt-4o")
print(f"Context: {info.context_window}")
print(f"Provider: {info.provider}")
# List all models from a provider
models = ModelRegistry.list_models(provider="openai")
# Register a custom model
ModelRegistry.register(
"my-custom-model",
provider="custom",
context_window=32000,
)
"""
from .registry import (
ModelInfo,
ModelRegistry,
get_model_info,
list_models,
register_model,
)
__all__ = [
"ModelRegistry",
"ModelInfo",
"get_model_info",
"list_models",
"register_model",
]

749
headroom/models/registry.py Normal file
View file

@ -0,0 +1,749 @@
"""Model registry with capabilities database.
Centralized database of LLM models with their capabilities, context limits,
pricing, and provider information. Supports dynamic registration of custom
models and automatic provider detection.
"""
from __future__ import annotations
from dataclasses import dataclass
from datetime import date
from typing import Any
@dataclass(frozen=True)
class ModelInfo:
"""Information about an LLM model.
Attributes:
name: Model identifier.
provider: Provider name (openai, anthropic, etc.).
context_window: Maximum context window in tokens.
max_output_tokens: Maximum output tokens.
supports_tools: Whether model supports tool/function calling.
supports_vision: Whether model supports image inputs.
supports_streaming: Whether model supports streaming responses.
supports_json_mode: Whether model supports JSON output mode.
tokenizer_backend: Tokenizer backend to use.
input_cost_per_1m: Cost per 1M input tokens in USD.
output_cost_per_1m: Cost per 1M output tokens in USD.
cached_input_cost_per_1m: Cost per 1M cached input tokens.
pricing_date: Date pricing was last updated.
aliases: Alternative names for the model.
notes: Additional notes about the model.
"""
name: str
provider: str
context_window: int = 128000
max_output_tokens: int = 4096
supports_tools: bool = True
supports_vision: bool = False
supports_streaming: bool = True
supports_json_mode: bool = True
tokenizer_backend: str | None = None
input_cost_per_1m: float | None = None
output_cost_per_1m: float | None = None
cached_input_cost_per_1m: float | None = None
pricing_date: date | None = None
aliases: tuple[str, ...] = ()
notes: str = ""
# Built-in model database
# Pricing as of January 2025 - verify current rates
_MODELS: dict[str, ModelInfo] = {}
def _register_builtin_models() -> None:
"""Register built-in models."""
# ============================================================
# OpenAI Models
# ============================================================
# GPT-4o family
_MODELS["gpt-4o"] = ModelInfo(
name="gpt-4o",
provider="openai",
context_window=128000,
max_output_tokens=16384,
supports_tools=True,
supports_vision=True,
supports_streaming=True,
tokenizer_backend="tiktoken",
input_cost_per_1m=2.50,
output_cost_per_1m=10.00,
cached_input_cost_per_1m=1.25,
pricing_date=date(2025, 1, 6),
aliases=("gpt-4o-2024-11-20", "gpt-4o-2024-08-06", "gpt-4o-2024-05-13"),
notes="Latest GPT-4o with vision and tools",
)
_MODELS["gpt-4o-mini"] = ModelInfo(
name="gpt-4o-mini",
provider="openai",
context_window=128000,
max_output_tokens=16384,
supports_tools=True,
supports_vision=True,
supports_streaming=True,
tokenizer_backend="tiktoken",
input_cost_per_1m=0.15,
output_cost_per_1m=0.60,
cached_input_cost_per_1m=0.075,
pricing_date=date(2025, 1, 6),
aliases=("gpt-4o-mini-2024-07-18",),
notes="Cost-effective GPT-4o variant",
)
# o1 reasoning models
_MODELS["o1"] = ModelInfo(
name="o1",
provider="openai",
context_window=200000,
max_output_tokens=100000,
supports_tools=True,
supports_vision=True,
supports_streaming=True,
tokenizer_backend="tiktoken",
input_cost_per_1m=15.00,
output_cost_per_1m=60.00,
cached_input_cost_per_1m=7.50,
pricing_date=date(2025, 1, 6),
notes="Full reasoning model with extended thinking",
)
_MODELS["o1-mini"] = ModelInfo(
name="o1-mini",
provider="openai",
context_window=128000,
max_output_tokens=65536,
supports_tools=True,
supports_vision=False,
supports_streaming=True,
tokenizer_backend="tiktoken",
input_cost_per_1m=1.10,
output_cost_per_1m=4.40,
cached_input_cost_per_1m=0.55,
pricing_date=date(2025, 1, 6),
notes="Fast reasoning model",
)
_MODELS["o3-mini"] = ModelInfo(
name="o3-mini",
provider="openai",
context_window=200000,
max_output_tokens=100000,
supports_tools=True,
supports_vision=True,
supports_streaming=True,
tokenizer_backend="tiktoken",
input_cost_per_1m=1.10,
output_cost_per_1m=4.40,
cached_input_cost_per_1m=0.55,
pricing_date=date(2025, 1, 6),
notes="Latest reasoning model",
)
# GPT-4 Turbo
_MODELS["gpt-4-turbo"] = ModelInfo(
name="gpt-4-turbo",
provider="openai",
context_window=128000,
max_output_tokens=4096,
supports_tools=True,
supports_vision=True,
supports_streaming=True,
tokenizer_backend="tiktoken",
input_cost_per_1m=10.00,
output_cost_per_1m=30.00,
cached_input_cost_per_1m=5.00,
pricing_date=date(2025, 1, 6),
aliases=("gpt-4-turbo-preview", "gpt-4-turbo-2024-04-09"),
notes="GPT-4 Turbo with vision",
)
# GPT-4
_MODELS["gpt-4"] = ModelInfo(
name="gpt-4",
provider="openai",
context_window=8192,
max_output_tokens=4096,
supports_tools=True,
supports_vision=False,
supports_streaming=True,
tokenizer_backend="tiktoken",
input_cost_per_1m=30.00,
output_cost_per_1m=60.00,
pricing_date=date(2025, 1, 6),
aliases=("gpt-4-0613",),
notes="Original GPT-4",
)
_MODELS["gpt-4-32k"] = ModelInfo(
name="gpt-4-32k",
provider="openai",
context_window=32768,
max_output_tokens=4096,
supports_tools=True,
supports_vision=False,
supports_streaming=True,
tokenizer_backend="tiktoken",
input_cost_per_1m=60.00,
output_cost_per_1m=120.00,
pricing_date=date(2025, 1, 6),
notes="Extended context GPT-4",
)
# GPT-3.5
_MODELS["gpt-3.5-turbo"] = ModelInfo(
name="gpt-3.5-turbo",
provider="openai",
context_window=16385,
max_output_tokens=4096,
supports_tools=True,
supports_vision=False,
supports_streaming=True,
tokenizer_backend="tiktoken",
input_cost_per_1m=0.50,
output_cost_per_1m=1.50,
cached_input_cost_per_1m=0.25,
pricing_date=date(2025, 1, 6),
aliases=("gpt-3.5-turbo-0125", "gpt-3.5-turbo-1106"),
notes="Fast and cost-effective",
)
# ============================================================
# Anthropic Models
# ============================================================
_MODELS["claude-3-5-sonnet-20241022"] = ModelInfo(
name="claude-3-5-sonnet-20241022",
provider="anthropic",
context_window=200000,
max_output_tokens=8192,
supports_tools=True,
supports_vision=True,
supports_streaming=True,
tokenizer_backend="anthropic",
input_cost_per_1m=3.00,
output_cost_per_1m=15.00,
cached_input_cost_per_1m=0.30,
pricing_date=date(2025, 1, 6),
aliases=("claude-3-5-sonnet-latest", "claude-sonnet-4-20250514"),
notes="Claude 3.5 Sonnet - Best balance of speed and capability",
)
_MODELS["claude-3-5-haiku-20241022"] = ModelInfo(
name="claude-3-5-haiku-20241022",
provider="anthropic",
context_window=200000,
max_output_tokens=8192,
supports_tools=True,
supports_vision=True,
supports_streaming=True,
tokenizer_backend="anthropic",
input_cost_per_1m=0.80,
output_cost_per_1m=4.00,
cached_input_cost_per_1m=0.08,
pricing_date=date(2025, 1, 6),
aliases=("claude-3-5-haiku-latest",),
notes="Claude 3.5 Haiku - Fast and cost-effective",
)
_MODELS["claude-3-opus-20240229"] = ModelInfo(
name="claude-3-opus-20240229",
provider="anthropic",
context_window=200000,
max_output_tokens=4096,
supports_tools=True,
supports_vision=True,
supports_streaming=True,
tokenizer_backend="anthropic",
input_cost_per_1m=15.00,
output_cost_per_1m=75.00,
cached_input_cost_per_1m=1.50,
pricing_date=date(2025, 1, 6),
aliases=("claude-3-opus-latest",),
notes="Claude 3 Opus - Most capable",
)
_MODELS["claude-3-haiku-20240307"] = ModelInfo(
name="claude-3-haiku-20240307",
provider="anthropic",
context_window=200000,
max_output_tokens=4096,
supports_tools=True,
supports_vision=True,
supports_streaming=True,
tokenizer_backend="anthropic",
input_cost_per_1m=0.25,
output_cost_per_1m=1.25,
cached_input_cost_per_1m=0.03,
pricing_date=date(2025, 1, 6),
notes="Claude 3 Haiku - Legacy fast model",
)
# ============================================================
# Google Models
# ============================================================
_MODELS["gemini-2.0-flash"] = ModelInfo(
name="gemini-2.0-flash",
provider="google",
context_window=1000000,
max_output_tokens=8192,
supports_tools=True,
supports_vision=True,
supports_streaming=True,
tokenizer_backend="google",
input_cost_per_1m=0.10,
output_cost_per_1m=0.40,
pricing_date=date(2025, 1, 6),
aliases=("gemini-2.0-flash-exp",),
notes="Gemini 2.0 Flash - Fast multimodal",
)
_MODELS["gemini-1.5-pro"] = ModelInfo(
name="gemini-1.5-pro",
provider="google",
context_window=2000000,
max_output_tokens=8192,
supports_tools=True,
supports_vision=True,
supports_streaming=True,
tokenizer_backend="google",
input_cost_per_1m=1.25,
output_cost_per_1m=5.00,
pricing_date=date(2025, 1, 6),
aliases=("gemini-1.5-pro-latest",),
notes="Gemini 1.5 Pro - 2M context window",
)
_MODELS["gemini-1.5-flash"] = ModelInfo(
name="gemini-1.5-flash",
provider="google",
context_window=1000000,
max_output_tokens=8192,
supports_tools=True,
supports_vision=True,
supports_streaming=True,
tokenizer_backend="google",
input_cost_per_1m=0.075,
output_cost_per_1m=0.30,
pricing_date=date(2025, 1, 6),
aliases=("gemini-1.5-flash-latest",),
notes="Gemini 1.5 Flash - Cost-effective",
)
# ============================================================
# Meta Llama Models (open source)
# ============================================================
_MODELS["llama-3.3-70b"] = ModelInfo(
name="llama-3.3-70b",
provider="meta",
context_window=128000,
max_output_tokens=4096,
supports_tools=True,
supports_vision=False,
supports_streaming=True,
tokenizer_backend="huggingface",
aliases=("llama-3.3-70b-instruct", "meta-llama/Llama-3.3-70B-Instruct"),
notes="Llama 3.3 70B - Open source",
)
_MODELS["llama-3.1-405b"] = ModelInfo(
name="llama-3.1-405b",
provider="meta",
context_window=128000,
max_output_tokens=4096,
supports_tools=True,
supports_vision=False,
supports_streaming=True,
tokenizer_backend="huggingface",
aliases=("llama-3.1-405b-instruct", "meta-llama/Llama-3.1-405B-Instruct"),
notes="Llama 3.1 405B - Largest open source",
)
_MODELS["llama-3.1-70b"] = ModelInfo(
name="llama-3.1-70b",
provider="meta",
context_window=128000,
max_output_tokens=4096,
supports_tools=True,
supports_vision=False,
supports_streaming=True,
tokenizer_backend="huggingface",
aliases=("llama-3.1-70b-instruct", "meta-llama/Llama-3.1-70B-Instruct"),
notes="Llama 3.1 70B",
)
_MODELS["llama-3.1-8b"] = ModelInfo(
name="llama-3.1-8b",
provider="meta",
context_window=128000,
max_output_tokens=4096,
supports_tools=True,
supports_vision=False,
supports_streaming=True,
tokenizer_backend="huggingface",
aliases=("llama-3.1-8b-instruct", "meta-llama/Llama-3.1-8B-Instruct"),
notes="Llama 3.1 8B - Fast and efficient",
)
# ============================================================
# Mistral Models
# ============================================================
_MODELS["mistral-large"] = ModelInfo(
name="mistral-large",
provider="mistral",
context_window=128000,
max_output_tokens=4096,
supports_tools=True,
supports_vision=False,
supports_streaming=True,
tokenizer_backend="huggingface",
input_cost_per_1m=2.00,
output_cost_per_1m=6.00,
pricing_date=date(2025, 1, 6),
aliases=("mistral-large-latest",),
notes="Mistral Large - Best capability",
)
_MODELS["mistral-small"] = ModelInfo(
name="mistral-small",
provider="mistral",
context_window=32768,
max_output_tokens=4096,
supports_tools=True,
supports_vision=False,
supports_streaming=True,
tokenizer_backend="huggingface",
input_cost_per_1m=0.20,
output_cost_per_1m=0.60,
pricing_date=date(2025, 1, 6),
aliases=("mistral-small-latest",),
notes="Mistral Small - Cost-effective",
)
_MODELS["mixtral-8x7b"] = ModelInfo(
name="mixtral-8x7b",
provider="mistral",
context_window=32768,
max_output_tokens=4096,
supports_tools=True,
supports_vision=False,
supports_streaming=True,
tokenizer_backend="huggingface",
aliases=("mixtral-8x7b-instruct",),
notes="Mixtral 8x7B - MoE architecture",
)
_MODELS["mistral-7b"] = ModelInfo(
name="mistral-7b",
provider="mistral",
context_window=32768,
max_output_tokens=4096,
supports_tools=False,
supports_vision=False,
supports_streaming=True,
tokenizer_backend="huggingface",
aliases=("mistral-7b-instruct",),
notes="Mistral 7B - Open source",
)
# ============================================================
# DeepSeek Models
# ============================================================
_MODELS["deepseek-v3"] = ModelInfo(
name="deepseek-v3",
provider="deepseek",
context_window=128000,
max_output_tokens=8192,
supports_tools=True,
supports_vision=False,
supports_streaming=True,
tokenizer_backend="huggingface",
input_cost_per_1m=0.14,
output_cost_per_1m=0.28,
pricing_date=date(2025, 1, 6),
notes="DeepSeek V3 - High performance, low cost",
)
_MODELS["deepseek-coder"] = ModelInfo(
name="deepseek-coder",
provider="deepseek",
context_window=16384,
max_output_tokens=4096,
supports_tools=False,
supports_vision=False,
supports_streaming=True,
tokenizer_backend="huggingface",
notes="DeepSeek Coder - Specialized for code",
)
# ============================================================
# Qwen Models
# ============================================================
_MODELS["qwen2.5-72b"] = ModelInfo(
name="qwen2.5-72b",
provider="alibaba",
context_window=131072,
max_output_tokens=8192,
supports_tools=True,
supports_vision=False,
supports_streaming=True,
tokenizer_backend="huggingface",
aliases=("qwen2.5-72b-instruct",),
notes="Qwen 2.5 72B - Strong multilingual",
)
_MODELS["qwen2.5-7b"] = ModelInfo(
name="qwen2.5-7b",
provider="alibaba",
context_window=131072,
max_output_tokens=8192,
supports_tools=True,
supports_vision=False,
supports_streaming=True,
tokenizer_backend="huggingface",
aliases=("qwen2.5-7b-instruct",),
notes="Qwen 2.5 7B - Efficient",
)
# Initialize built-in models
_register_builtin_models()
# Build alias lookup
_ALIASES: dict[str, str] = {}
for model_name, info in _MODELS.items():
for alias in info.aliases:
_ALIASES[alias.lower()] = model_name
class ModelRegistry:
"""Registry of LLM models and their capabilities.
Singleton registry providing access to model information.
Supports built-in models and custom registration.
Example:
# Get model info
info = ModelRegistry.get("gpt-4o")
print(f"Context: {info.context_window}")
# Register custom model
ModelRegistry.register(
"my-model",
provider="custom",
context_window=32000,
)
# List models by provider
openai_models = ModelRegistry.list_models(provider="openai")
"""
@classmethod
def get(cls, model: str) -> ModelInfo | None:
"""Get model information.
Args:
model: Model name or alias.
Returns:
ModelInfo if found, None otherwise.
"""
model_lower = model.lower()
# Direct lookup
if model_lower in _MODELS:
return _MODELS[model_lower]
# Alias lookup
if model_lower in _ALIASES:
return _MODELS[_ALIASES[model_lower]]
# Prefix matching
for name, info in _MODELS.items():
if model_lower.startswith(name):
return info
return None
@classmethod
def register(
cls,
model: str,
provider: str,
context_window: int = 128000,
**kwargs: Any,
) -> ModelInfo:
"""Register a custom model.
Args:
model: Model name.
provider: Provider name.
context_window: Maximum context window.
**kwargs: Additional ModelInfo fields.
Returns:
Registered ModelInfo.
"""
info = ModelInfo(
name=model,
provider=provider,
context_window=context_window,
**kwargs,
)
_MODELS[model.lower()] = info
# Register aliases
for alias in info.aliases:
_ALIASES[alias.lower()] = model.lower()
return info
@classmethod
def list_models(
cls,
provider: str | None = None,
supports_tools: bool | None = None,
supports_vision: bool | None = None,
min_context: int | None = None,
) -> list[ModelInfo]:
"""List models matching criteria.
Args:
provider: Filter by provider.
supports_tools: Filter by tool support.
supports_vision: Filter by vision support.
min_context: Minimum context window.
Returns:
List of matching ModelInfo.
"""
results = []
for info in _MODELS.values():
if provider and info.provider != provider:
continue
if supports_tools is not None and info.supports_tools != supports_tools:
continue
if supports_vision is not None and info.supports_vision != supports_vision:
continue
if min_context and info.context_window < min_context:
continue
results.append(info)
return results
@classmethod
def list_providers(cls) -> list[str]:
"""List all known providers.
Returns:
List of provider names.
"""
return list(set(info.provider for info in _MODELS.values()))
@classmethod
def get_context_limit(cls, model: str, default: int = 128000) -> int:
"""Get context limit for a model.
Args:
model: Model name.
default: Default if model not found.
Returns:
Context window size.
"""
info = cls.get(model)
return info.context_window if info else default
@classmethod
def estimate_cost(
cls,
model: str,
input_tokens: int,
output_tokens: int,
cached_tokens: int = 0,
) -> float | None:
"""Estimate API cost for a model.
Args:
model: Model name.
input_tokens: Number of input tokens.
output_tokens: Number of output tokens.
cached_tokens: Number of cached input tokens.
Returns:
Estimated cost in USD, or None if pricing unknown.
"""
info = cls.get(model)
if not info or info.input_cost_per_1m is None:
return None
input_cost = (input_tokens / 1_000_000) * info.input_cost_per_1m
output_cost = (output_tokens / 1_000_000) * (info.output_cost_per_1m or 0)
if cached_tokens and info.cached_input_cost_per_1m:
# Adjust for cached tokens
regular_input = input_tokens - cached_tokens
cached_cost = (cached_tokens / 1_000_000) * info.cached_input_cost_per_1m
input_cost = (regular_input / 1_000_000) * info.input_cost_per_1m + cached_cost
return input_cost + output_cost
# Convenience functions
def get_model_info(model: str) -> ModelInfo | None:
"""Get information about a model.
Args:
model: Model name or alias.
Returns:
ModelInfo if found, None otherwise.
"""
return ModelRegistry.get(model)
def list_models(
provider: str | None = None,
**kwargs: Any,
) -> list[ModelInfo]:
"""List models matching criteria.
Args:
provider: Filter by provider.
**kwargs: Additional filter criteria.
Returns:
List of matching ModelInfo.
"""
return ModelRegistry.list_models(provider=provider, **kwargs)
def register_model(
model: str,
provider: str,
context_window: int = 128000,
**kwargs: Any,
) -> ModelInfo:
"""Register a custom model.
Args:
model: Model name.
provider: Provider name.
context_window: Maximum context window.
**kwargs: Additional ModelInfo fields.
Returns:
Registered ModelInfo.
"""
return ModelRegistry.register(model, provider, context_window, **kwargs)

View file

@ -4,7 +4,7 @@ from __future__ import annotations
import hashlib
import re
from typing import Any, TYPE_CHECKING
from typing import TYPE_CHECKING, Any
from .config import Block, WasteSignals

View file

@ -4,18 +4,21 @@ This module provides pricing information and cost estimation utilities
for various LLM providers including OpenAI and Anthropic.
"""
from .registry import CostEstimate, ModelPricing, PricingRegistry
from .openai_prices import (
LAST_UPDATED as OPENAI_LAST_UPDATED,
OPENAI_PRICES,
get_openai_registry,
)
from .anthropic_prices import (
LAST_UPDATED as ANTHROPIC_LAST_UPDATED,
ANTHROPIC_PRICES,
get_anthropic_registry,
)
from .anthropic_prices import (
LAST_UPDATED as ANTHROPIC_LAST_UPDATED,
)
from .openai_prices import (
LAST_UPDATED as OPENAI_LAST_UPDATED,
)
from .openai_prices import (
OPENAI_PRICES,
get_openai_registry,
)
from .registry import CostEstimate, ModelPricing, PricingRegistry
__all__ = [
# Core classes

View file

@ -4,7 +4,6 @@ from datetime import date
from .registry import ModelPricing, PricingRegistry
# Last verified date for pricing information
LAST_UPDATED = date(2025, 1, 6)

View file

@ -4,7 +4,6 @@ from datetime import date
from .registry import ModelPricing, PricingRegistry
# Last verified date for pricing information
LAST_UPDATED = date(2025, 1, 6)

View file

@ -2,7 +2,6 @@
from dataclasses import dataclass, field
from datetime import date, timedelta
from typing import Optional
@dataclass(frozen=True)
@ -15,11 +14,11 @@ class ModelPricing:
provider: str
input_per_1m: float
output_per_1m: float
cached_input_per_1m: Optional[float] = None
batch_input_per_1m: Optional[float] = None
batch_output_per_1m: Optional[float] = None
context_window: Optional[int] = None
notes: Optional[str] = None
cached_input_per_1m: float | None = None
batch_input_per_1m: float | None = None
batch_output_per_1m: float | None = None
context_window: int | None = None
notes: str | None = None
@dataclass
@ -27,9 +26,9 @@ class CostEstimate:
"""Result of a cost estimation calculation."""
cost_usd: float
breakdown: dict = field(default_factory=dict)
pricing_date: Optional[date] = None
pricing_date: date | None = None
is_stale: bool = False
warning: Optional[str] = None
warning: str | None = None
class PricingRegistry:
@ -41,8 +40,8 @@ class PricingRegistry:
def __init__(
self,
last_updated: date,
source_url: Optional[str] = None,
prices: Optional[dict[str, ModelPricing]] = None,
source_url: str | None = None,
prices: dict[str, ModelPricing] | None = None,
):
"""Initialize the pricing registry.
@ -55,7 +54,7 @@ class PricingRegistry:
self.source_url = source_url
self.prices: dict[str, ModelPricing] = prices or {}
def get_price(self, model: str) -> Optional[ModelPricing]:
def get_price(self, model: str) -> ModelPricing | None:
"""Get pricing for a specific model.
Args:
@ -75,7 +74,7 @@ class PricingRegistry:
age = date.today() - self.last_updated
return age > timedelta(days=self.STALENESS_THRESHOLD_DAYS)
def staleness_warning(self) -> Optional[str]:
def staleness_warning(self) -> str | None:
"""Get a warning message if pricing is stale.
Returns:

View file

@ -2,15 +2,60 @@
Providers encapsulate model-specific behavior like tokenization,
context limits, and cost estimation.
Supported Providers:
- OpenAIProvider: Native OpenAI models (GPT-4o, o1, etc.)
- AnthropicProvider: Claude models
- GoogleProvider: Google Gemini models
- CohereProvider: Cohere Command models
- OpenAICompatibleProvider: Universal provider for any OpenAI-compatible API
(Ollama, vLLM, Together, Groq, Fireworks, LM Studio, etc.)
- LiteLLMProvider: Universal provider via LiteLLM (100+ providers)
"""
from .anthropic import AnthropicProvider
from .base import Provider, TokenCounter
from .cohere import CohereProvider
from .google import GoogleProvider
from .litellm import (
LiteLLMProvider,
create_litellm_provider,
is_litellm_available,
)
from .openai import OpenAIProvider
from .openai_compatible import (
ModelCapabilities,
OpenAICompatibleProvider,
create_anyscale_provider,
create_fireworks_provider,
create_groq_provider,
create_lmstudio_provider,
create_ollama_provider,
create_together_provider,
create_vllm_provider,
)
__all__ = [
# Base
"Provider",
"TokenCounter",
# Native providers
"OpenAIProvider",
"AnthropicProvider",
"GoogleProvider",
"CohereProvider",
# Universal providers
"OpenAICompatibleProvider",
"ModelCapabilities",
"LiteLLMProvider",
"is_litellm_available",
# Factory functions
"create_ollama_provider",
"create_together_provider",
"create_groq_provider",
"create_fireworks_provider",
"create_anyscale_provider",
"create_vllm_provider",
"create_lmstudio_provider",
"create_litellm_provider",
]

View file

@ -0,0 +1,313 @@
"""Cohere provider for Headroom SDK.
Token counting uses Cohere's official tokenize API when a client
is provided. This gives accurate counts for all content types.
Usage:
import cohere
from headroom import CohereProvider
client = cohere.ClientV2() # Uses CO_API_KEY env var
provider = CohereProvider(client=client) # Accurate counting via API
# Or without client (uses estimation - less accurate)
provider = CohereProvider() # Warning: approximate counting
"""
from __future__ import annotations
import logging
import warnings
from datetime import date
from typing import Any
from headroom.tokenizers import EstimatingTokenCounter
from .base import Provider, TokenCounter
logger = logging.getLogger(__name__)
# Warning flags
_FALLBACK_WARNING_SHOWN = False
# Pricing metadata
_PRICING_LAST_UPDATED = date(2025, 1, 6)
# Cohere model context limits
_CONTEXT_LIMITS: dict[str, int] = {
# Command A (latest, 2025)
"command-a-03-2025": 256000,
"command-a": 256000,
# Command R+ (2024)
"command-r-plus-08-2024": 128000,
"command-r-plus": 128000,
# Command R (2024)
"command-r-08-2024": 128000,
"command-r": 128000,
# Command (legacy)
"command": 4096,
"command-light": 4096,
"command-nightly": 128000,
# Embed models
"embed-english-v3.0": 512,
"embed-multilingual-v3.0": 512,
"embed-english-light-v3.0": 512,
"embed-multilingual-light-v3.0": 512,
}
# Pricing per 1M tokens (input, output)
_PRICING: dict[str, tuple[float, float]] = {
"command-a-03-2025": (2.50, 10.00),
"command-a": (2.50, 10.00),
"command-r-plus-08-2024": (2.50, 10.00),
"command-r-plus": (2.50, 10.00),
"command-r-08-2024": (0.15, 0.60),
"command-r": (0.15, 0.60),
"command": (1.00, 2.00),
"command-light": (0.30, 0.60),
}
class CohereTokenCounter:
"""Token counter for Cohere models.
When a Cohere client is provided, uses the official tokenize API
for accurate counting. Falls back to estimation when no client
is available.
Usage:
import cohere
client = cohere.ClientV2()
# With API (accurate)
counter = CohereTokenCounter("command-r-plus", client=client)
# Without API (estimation)
counter = CohereTokenCounter("command-r-plus")
"""
def __init__(self, model: str, client: Any = None):
"""Initialize Cohere token counter.
Args:
model: Cohere model name.
client: Optional cohere.ClientV2 for API-based counting.
"""
global _FALLBACK_WARNING_SHOWN
self.model = model
self._client = client
self._use_api = client is not None
# Cohere uses ~4 chars per token
self._estimator = EstimatingTokenCounter(chars_per_token=4.0)
if not self._use_api and not _FALLBACK_WARNING_SHOWN:
warnings.warn(
"CohereProvider: No client provided, using estimation. "
"For accurate counting, pass a Cohere client: "
"CohereProvider(client=cohere.ClientV2())",
UserWarning,
stacklevel=4
)
_FALLBACK_WARNING_SHOWN = True
def count_text(self, text: str) -> int:
"""Count tokens in text.
Uses tokenize API if client available, otherwise estimates.
"""
if not text:
return 0
if self._use_api:
try:
response = self._client.tokenize(
text=text,
model=self.model,
)
return len(response.tokens)
except Exception as e:
logger.debug(f"Cohere tokenize API failed: {e}, using estimation")
return self._estimator.count_text(text)
def count_message(self, message: dict[str, Any]) -> int:
"""Count tokens in a message."""
content = self._extract_content(message)
tokens = self.count_text(content)
tokens += 4 # Message overhead (role tokens, etc.)
return tokens
def count_messages(self, messages: list[dict[str, Any]]) -> int:
"""Count tokens in messages."""
if not messages:
return 0
# For API-based counting, concatenate all content
if self._use_api:
try:
all_content = []
for msg in messages:
content = self._extract_content(msg)
role = msg.get("role", "user")
all_content.append(f"{role}: {content}")
full_text = "\n".join(all_content)
response = self._client.tokenize(
text=full_text,
model=self.model,
)
return len(response.tokens)
except Exception as e:
logger.debug(f"Cohere tokenize API failed: {e}, using estimation")
# Fallback to estimation
total = sum(self.count_message(msg) for msg in messages)
total += 3 # Priming tokens
return total
def _extract_content(self, message: dict[str, Any]) -> str:
"""Extract text content from message."""
content = message.get("content", "")
if isinstance(content, str):
return content
elif isinstance(content, list):
parts = []
for part in content:
if isinstance(part, dict) and part.get("type") == "text":
parts.append(part.get("text", ""))
elif isinstance(part, str):
parts.append(part)
return "\n".join(parts)
return str(content)
class CohereProvider(Provider):
"""Provider for Cohere Command models.
Supports Command R, Command R+, and Command A model families.
Example:
import cohere
client = cohere.ClientV2()
# With client (accurate token counting via API)
provider = CohereProvider(client=client)
# Without client (estimation-based counting)
provider = CohereProvider()
# Token counting
counter = provider.get_token_counter("command-r-plus")
tokens = counter.count_text("Hello, world!")
# Context limits
limit = provider.get_context_limit("command-a") # 256K tokens
# Cost estimation
cost = provider.estimate_cost(
input_tokens=100000,
output_tokens=10000,
model="command-r-plus",
)
"""
def __init__(self, client: Any = None):
"""Initialize Cohere provider.
Args:
client: Optional cohere.ClientV2 for API-based token counting.
If provided, uses tokenize API for accurate counts.
"""
self._client = client
@property
def name(self) -> str:
return "cohere"
def supports_model(self, model: str) -> bool:
"""Check if model is a known Cohere model."""
model_lower = model.lower()
if model_lower in _CONTEXT_LIMITS:
return True
# Check prefix match
for prefix in ["command-a", "command-r", "command", "embed-"]:
if model_lower.startswith(prefix):
return True
return False
def get_token_counter(self, model: str) -> TokenCounter:
"""Get token counter for a Cohere model.
Uses tokenize API if client was provided, otherwise estimates.
"""
if not self.supports_model(model):
raise ValueError(
f"Model '{model}' is not recognized as a Cohere model. "
f"Supported models: {list(_CONTEXT_LIMITS.keys())}"
)
return CohereTokenCounter(model, client=self._client)
def get_context_limit(self, model: str) -> int:
"""Get context limit for a Cohere model."""
model_lower = model.lower()
# Direct match
if model_lower in _CONTEXT_LIMITS:
return _CONTEXT_LIMITS[model_lower]
# Prefix match
for prefix, limit in [
("command-a", 256000),
("command-r-plus", 128000),
("command-r", 128000),
("command", 4096),
("embed-", 512),
]:
if model_lower.startswith(prefix):
return limit
raise ValueError(
f"Unknown context limit for model '{model}'. "
f"Known models: {list(_CONTEXT_LIMITS.keys())}"
)
def estimate_cost(
self,
input_tokens: int,
output_tokens: int,
model: str,
cached_tokens: int = 0,
) -> float | None:
"""Estimate cost for Cohere API call.
Args:
input_tokens: Number of input tokens.
output_tokens: Number of output tokens.
model: Model name.
cached_tokens: Not used by Cohere.
Returns:
Estimated cost in USD, or None if pricing unknown.
"""
model_lower = model.lower()
# Find pricing
input_price, output_price = None, None
for model_prefix, (inp, outp) in _PRICING.items():
if model_lower.startswith(model_prefix):
input_price, output_price = inp, outp
break
if input_price is None:
return None
input_cost = (input_tokens / 1_000_000) * input_price
output_cost = (output_tokens / 1_000_000) * output_price
return input_cost + output_cost
def get_output_buffer(self, model: str, default: int = 4000) -> int:
"""Get recommended output buffer."""
return default

View file

@ -0,0 +1,372 @@
"""Google Gemini provider for Headroom SDK.
Supports Google's Gemini models through two interfaces:
1. OpenAI-compatible endpoint (recommended for Headroom)
2. Native Google AI SDK (for advanced features)
Token counting uses Google's official countTokens API when a client
is provided. This gives accurate counts for all content types.
Usage:
import google.generativeai as genai
from headroom import GoogleProvider
genai.configure(api_key="your-api-key")
provider = GoogleProvider(client=genai) # Accurate counting via API
# Or without client (uses estimation - less accurate)
provider = GoogleProvider() # Warning: approximate counting
"""
from __future__ import annotations
import logging
import warnings
from datetime import date
from typing import Any
from headroom.tokenizers import EstimatingTokenCounter
from .base import Provider, TokenCounter
logger = logging.getLogger(__name__)
# Warning flags
_FALLBACK_WARNING_SHOWN = False
# Pricing metadata
_PRICING_LAST_UPDATED = date(2025, 1, 6)
# Google model context limits
_CONTEXT_LIMITS: dict[str, int] = {
# Gemini 2.0
"gemini-2.0-flash": 1000000,
"gemini-2.0-flash-exp": 1000000,
"gemini-2.0-flash-thinking": 1000000,
# Gemini 1.5
"gemini-1.5-pro": 2000000,
"gemini-1.5-pro-latest": 2000000,
"gemini-1.5-flash": 1000000,
"gemini-1.5-flash-latest": 1000000,
"gemini-1.5-flash-8b": 1000000,
# Gemini 1.0
"gemini-1.0-pro": 32768,
"gemini-pro": 32768,
}
# Pricing per 1M tokens (input, output)
# Note: Google has different pricing tiers based on context length
_PRICING: dict[str, tuple[float, float]] = {
"gemini-2.0-flash": (0.10, 0.40),
"gemini-2.0-flash-exp": (0.10, 0.40), # Experimental, may change
"gemini-1.5-pro": (1.25, 5.00), # Up to 128K context
"gemini-1.5-flash": (0.075, 0.30), # Up to 128K context
"gemini-1.5-flash-8b": (0.0375, 0.15),
"gemini-1.0-pro": (0.50, 1.50),
}
class GeminiTokenCounter:
"""Token counter for Gemini models.
When a google.generativeai client is provided, uses the official
countTokens API for accurate counting. Falls back to estimation
when no client is available.
Usage:
import google.generativeai as genai
genai.configure(api_key="...")
# With API (accurate)
counter = GeminiTokenCounter("gemini-2.0-flash", client=genai)
# Without API (estimation)
counter = GeminiTokenCounter("gemini-2.0-flash")
"""
def __init__(self, model: str, client: Any = None):
"""Initialize Gemini token counter.
Args:
model: Gemini model name.
client: Optional google.generativeai module for API-based counting.
"""
global _FALLBACK_WARNING_SHOWN
self.model = model
self._client = client
self._use_api = client is not None
self._genai_model = None
# Gemini uses ~4 chars per token (similar to GPT models)
self._estimator = EstimatingTokenCounter(chars_per_token=4.0)
if not self._use_api and not _FALLBACK_WARNING_SHOWN:
warnings.warn(
"GoogleProvider: No client provided, using estimation. "
"For accurate counting, pass google.generativeai: "
"GoogleProvider(client=genai)",
UserWarning,
stacklevel=4
)
_FALLBACK_WARNING_SHOWN = True
def _get_model(self):
"""Lazy-load the GenerativeModel for API calls."""
if self._genai_model is None and self._client is not None:
self._genai_model = self._client.GenerativeModel(self.model)
return self._genai_model
def count_text(self, text: str) -> int:
"""Count tokens in text.
Uses countTokens API if client available, otherwise estimates.
"""
if not text:
return 0
if self._use_api:
try:
model = self._get_model()
response = model.count_tokens(text)
return response.total_tokens
except Exception as e:
logger.debug(f"Google countTokens API failed: {e}, using estimation")
return self._estimator.count_text(text)
def count_message(self, message: dict[str, Any]) -> int:
"""Count tokens in a message."""
# For API-based counting, convert message to content and count
if self._use_api:
try:
content = self._message_to_content(message)
model = self._get_model()
response = model.count_tokens(content)
return response.total_tokens
except Exception as e:
logger.debug(f"Google countTokens API failed: {e}, using estimation")
# Fallback to estimation
return self._estimate_message(message)
def count_messages(self, messages: list[dict[str, Any]]) -> int:
"""Count tokens in messages.
Uses countTokens API with full conversation if available.
"""
if not messages:
return 0
if self._use_api:
try:
# Convert to Gemini content format
contents = [self._message_to_content(msg) for msg in messages]
model = self._get_model()
response = model.count_tokens(contents)
return response.total_tokens
except Exception as e:
logger.debug(f"Google countTokens API failed: {e}, using estimation")
# Fallback to estimation
total = sum(self._estimate_message(msg) for msg in messages)
total += 3 # Priming tokens
return total
def _message_to_content(self, message: dict[str, Any]) -> str:
"""Convert OpenAI-format message to text content for counting."""
content = message.get("content", "")
if isinstance(content, str):
return content
elif isinstance(content, list):
parts = []
for part in content:
if isinstance(part, dict) and part.get("type") == "text":
parts.append(part.get("text", ""))
elif isinstance(part, str):
parts.append(part)
return "\n".join(parts)
return str(content)
def _estimate_message(self, message: dict[str, Any]) -> int:
"""Estimate tokens in a message without API."""
tokens = 4 # Message overhead
role = message.get("role", "")
tokens += self._estimator.count_text(role)
content = message.get("content")
if content:
if isinstance(content, str):
tokens += self._estimator.count_text(content)
elif isinstance(content, list):
for part in content:
if isinstance(part, dict):
if part.get("type") == "text":
tokens += self._estimator.count_text(part.get("text", ""))
elif isinstance(part, str):
tokens += self._estimator.count_text(part)
return tokens
class GoogleProvider(Provider):
"""Provider for Google Gemini models.
Supports Gemini 1.5 and 2.0 model families through:
- OpenAI-compatible endpoint (generativelanguage.googleapis.com)
- Native Google AI SDK (for accurate token counting)
Example:
import google.generativeai as genai
genai.configure(api_key="...")
# With client (accurate token counting via API)
provider = GoogleProvider(client=genai)
# Without client (estimation-based counting)
provider = GoogleProvider()
# Token counting
counter = provider.get_token_counter("gemini-2.0-flash")
tokens = counter.count_text("Hello, world!")
# Context limits
limit = provider.get_context_limit("gemini-1.5-pro") # 2M tokens!
# Cost estimation
cost = provider.estimate_cost(
input_tokens=100000,
output_tokens=10000,
model="gemini-1.5-pro",
)
"""
# OpenAI-compatible endpoint for Gemini
OPENAI_COMPATIBLE_BASE_URL = "https://generativelanguage.googleapis.com/v1beta/openai"
def __init__(self, client: Any = None):
"""Initialize Google provider.
Args:
client: Optional google.generativeai module for API-based token counting.
If provided, uses countTokens API for accurate counts.
"""
self._client = client
@property
def name(self) -> str:
return "google"
def supports_model(self, model: str) -> bool:
"""Check if model is a known Gemini model."""
model_lower = model.lower()
if model_lower in _CONTEXT_LIMITS:
return True
# Check prefix match
for prefix in ["gemini-2", "gemini-1.5", "gemini-1.0", "gemini-pro"]:
if model_lower.startswith(prefix):
return True
return False
def get_token_counter(self, model: str) -> TokenCounter:
"""Get token counter for a Gemini model.
Uses countTokens API if client was provided, otherwise estimates.
"""
if not self.supports_model(model):
raise ValueError(
f"Model '{model}' is not recognized as a Google model. "
f"Supported models: {list(_CONTEXT_LIMITS.keys())}"
)
return GeminiTokenCounter(model, client=self._client)
def get_context_limit(self, model: str) -> int:
"""Get context limit for a Gemini model.
Note: Gemini 1.5 Pro has 2M token context!
"""
model_lower = model.lower()
# Direct match
if model_lower in _CONTEXT_LIMITS:
return _CONTEXT_LIMITS[model_lower]
# Prefix match
for prefix, limit in [
("gemini-2.0", 1000000),
("gemini-1.5-pro", 2000000),
("gemini-1.5-flash", 1000000),
("gemini-1.0", 32768),
("gemini-pro", 32768),
]:
if model_lower.startswith(prefix):
return limit
raise ValueError(
f"Unknown context limit for model '{model}'. "
f"Known models: {list(_CONTEXT_LIMITS.keys())}"
)
def estimate_cost(
self,
input_tokens: int,
output_tokens: int,
model: str,
cached_tokens: int = 0,
) -> float | None:
"""Estimate cost for Gemini API call.
Note: Google has tiered pricing based on context length.
This uses the standard pricing (up to 128K context).
For >128K context, actual costs may be higher.
Args:
input_tokens: Number of input tokens.
output_tokens: Number of output tokens.
model: Model name.
cached_tokens: Number of cached tokens (not used by Google).
Returns:
Estimated cost in USD, or None if pricing unknown.
"""
model_lower = model.lower()
# Find pricing
input_price, output_price = None, None
for model_prefix, (inp, outp) in _PRICING.items():
if model_lower.startswith(model_prefix):
input_price, output_price = inp, outp
break
if input_price is None:
return None
input_cost = (input_tokens / 1_000_000) * input_price
output_cost = (output_tokens / 1_000_000) * output_price
return input_cost + output_cost
def get_output_buffer(self, model: str, default: int = 4000) -> int:
"""Get recommended output buffer."""
# Gemini models can output up to 8K tokens
return min(8192, default)
@classmethod
def get_openai_compatible_url(cls, api_key: str) -> str:
"""Get OpenAI-compatible endpoint URL.
Use this with the OpenAI client:
from openai import OpenAI
client = OpenAI(
api_key=api_key,
base_url=GoogleProvider.get_openai_compatible_url(api_key),
)
Args:
api_key: Google AI API key.
Returns:
Base URL for OpenAI-compatible requests.
"""
return cls.OPENAI_COMPATIBLE_BASE_URL

View file

@ -0,0 +1,293 @@
"""LiteLLM provider for universal LLM support.
LiteLLM provides a unified interface to 100+ LLM providers:
- OpenAI, Azure OpenAI
- Anthropic
- Google (Vertex AI, AI Studio)
- AWS Bedrock
- Cohere
- Replicate
- Hugging Face
- Ollama
- Together AI
- Groq
- And many more...
This integration allows Headroom to work with any LiteLLM-supported
model without needing provider-specific implementations.
Requires: pip install litellm
"""
from __future__ import annotations
import logging
from typing import Any
from headroom.tokenizers import EstimatingTokenCounter
from .base import Provider, TokenCounter
logger = logging.getLogger(__name__)
# Check if litellm is available
try:
import litellm
from litellm import get_model_info as litellm_get_model_info
from litellm import model_cost as litellm_model_cost
from litellm import token_counter as litellm_token_counter
LITELLM_AVAILABLE = True
except ImportError:
LITELLM_AVAILABLE = False
litellm = None
litellm_token_counter = None
litellm_model_cost = None
litellm_get_model_info = None
def is_litellm_available() -> bool:
"""Check if LiteLLM is installed.
Returns:
True if litellm is available.
"""
return LITELLM_AVAILABLE
class LiteLLMTokenCounter:
"""Token counter using LiteLLM's token counting.
LiteLLM provides accurate token counting for most providers
by using the appropriate tokenizer for each model.
"""
def __init__(self, model: str):
"""Initialize LiteLLM token counter.
Args:
model: Model name in LiteLLM format (e.g., 'gpt-4o', 'claude-3-sonnet').
"""
if not LITELLM_AVAILABLE:
raise RuntimeError(
"LiteLLM is required for LiteLLMProvider. "
"Install with: pip install litellm"
)
self.model = model
# Fallback estimator for when litellm counting fails
self._fallback = EstimatingTokenCounter()
def count_text(self, text: str) -> int:
"""Count tokens in text using LiteLLM."""
if not text:
return 0
try:
# LiteLLM's token_counter expects messages format
# We wrap text in a simple message
return litellm_token_counter(
model=self.model,
messages=[{"role": "user", "content": text}],
)
except Exception as e:
logger.debug(f"LiteLLM token count failed for {self.model}: {e}")
return self._fallback.count_text(text)
def count_message(self, message: dict[str, Any]) -> int:
"""Count tokens in a single message."""
try:
return litellm_token_counter(
model=self.model,
messages=[message],
)
except Exception as e:
logger.debug(f"LiteLLM message count failed for {self.model}: {e}")
# Fallback to estimation
tokens = 4 # Base overhead
content = message.get("content", "")
if isinstance(content, str):
tokens += self._fallback.count_text(content)
return tokens
def count_messages(self, messages: list[dict[str, Any]]) -> int:
"""Count tokens in messages using LiteLLM."""
if not messages:
return 0
try:
return litellm_token_counter(
model=self.model,
messages=messages,
)
except Exception as e:
logger.debug(f"LiteLLM messages count failed for {self.model}: {e}")
# Fallback to estimation
total = sum(self.count_message(msg) for msg in messages)
total += 3 # Priming
return total
class LiteLLMProvider(Provider):
"""Provider using LiteLLM for universal model support.
LiteLLM supports 100+ LLM providers with a unified interface.
This provider leverages LiteLLM's:
- Token counting (accurate for most providers)
- Model info (context limits, capabilities)
- Cost estimation (from LiteLLM's model database)
Example:
from headroom.providers import LiteLLMProvider
provider = LiteLLMProvider()
# Works with any LiteLLM-supported model
counter = provider.get_token_counter("gpt-4o")
counter = provider.get_token_counter("claude-3-5-sonnet-20241022")
counter = provider.get_token_counter("gemini/gemini-1.5-pro")
counter = provider.get_token_counter("bedrock/anthropic.claude-v2")
counter = provider.get_token_counter("ollama/llama3")
Model Format:
LiteLLM uses a provider/model format for some providers:
- OpenAI: "gpt-4o" or "openai/gpt-4o"
- Anthropic: "claude-3-sonnet" or "anthropic/claude-3-sonnet"
- Google: "gemini/gemini-1.5-pro"
- Azure: "azure/gpt-4"
- Bedrock: "bedrock/anthropic.claude-v2"
- Ollama: "ollama/llama3"
See LiteLLM docs for full model list:
https://docs.litellm.ai/docs/providers
"""
def __init__(self):
"""Initialize LiteLLM provider."""
if not LITELLM_AVAILABLE:
raise RuntimeError(
"LiteLLM is required for LiteLLMProvider. "
"Install with: pip install litellm"
)
@property
def name(self) -> str:
return "litellm"
def supports_model(self, model: str) -> bool:
"""Check if LiteLLM supports this model.
LiteLLM supports most models, so this returns True
for any model. Actual support depends on credentials.
"""
return True # LiteLLM handles validation
def get_token_counter(self, model: str) -> TokenCounter:
"""Get token counter for a model."""
return LiteLLMTokenCounter(model)
def get_context_limit(self, model: str) -> int:
"""Get context limit using LiteLLM's model info."""
try:
info = litellm_get_model_info(model)
if info and "max_input_tokens" in info:
return info["max_input_tokens"]
if info and "max_tokens" in info:
return info["max_tokens"]
except Exception as e:
logger.debug(f"LiteLLM get_model_info failed for {model}: {e}")
# Fallback to reasonable default
return 128000
def get_output_buffer(self, model: str, default: int = 4000) -> int:
"""Get recommended output buffer."""
try:
info = litellm_get_model_info(model)
if info and "max_output_tokens" in info:
return min(info["max_output_tokens"], default)
except Exception:
pass
return default
def estimate_cost(
self,
input_tokens: int,
output_tokens: int,
model: str,
cached_tokens: int = 0,
) -> float | None:
"""Estimate cost using LiteLLM's cost database.
Args:
input_tokens: Number of input tokens.
output_tokens: Number of output tokens.
model: Model name.
cached_tokens: Cached tokens (may not be supported by all providers).
Returns:
Estimated cost in USD, or None if pricing unknown.
"""
try:
# LiteLLM's cost calculation
cost = litellm.completion_cost(
model=model,
prompt="", # We're using token counts directly
completion="",
prompt_tokens=input_tokens,
completion_tokens=output_tokens,
)
return cost
except Exception as e:
logger.debug(f"LiteLLM cost estimation failed for {model}: {e}")
return None
@classmethod
def list_supported_providers(cls) -> list[str]:
"""List providers supported by LiteLLM.
Returns:
List of provider names.
"""
if not LITELLM_AVAILABLE:
return []
# Major providers supported by LiteLLM
return [
"openai",
"anthropic",
"azure",
"google",
"vertex_ai",
"bedrock",
"cohere",
"replicate",
"huggingface",
"ollama",
"together_ai",
"groq",
"fireworks_ai",
"anyscale",
"deepinfra",
"perplexity",
"mistral",
"cloudflare",
"ai21",
"nlp_cloud",
"aleph_alpha",
"petals",
"baseten",
"openrouter",
"vllm",
"xinference",
"text-generation-inference",
]
def create_litellm_provider() -> LiteLLMProvider:
"""Create a LiteLLM provider.
Returns:
Configured LiteLLMProvider.
Raises:
RuntimeError: If LiteLLM is not installed.
"""
return LiteLLMProvider()

View file

@ -6,7 +6,6 @@ Cost estimates are APPROXIMATE - always verify against your actual billing.
from __future__ import annotations
import json
import warnings
from datetime import date
from functools import lru_cache

View file

@ -0,0 +1,521 @@
"""OpenAI-compatible provider for universal LLM support.
This provider supports any LLM service that implements the OpenAI API format:
- Ollama (local)
- vLLM (local/cloud)
- Together AI
- Groq
- Fireworks AI
- Anyscale
- LM Studio
- LocalAI
- Hugging Face Inference Endpoints
- Azure OpenAI
- And many more...
The key insight: 70%+ of LLM providers use OpenAI-compatible APIs,
so supporting this format gives near-universal coverage.
"""
from __future__ import annotations
import logging
from dataclasses import dataclass
from typing import Any
from headroom.tokenizers import get_tokenizer
from .base import Provider
logger = logging.getLogger(__name__)
@dataclass
class ModelCapabilities:
"""Model capability metadata.
Stores information about a model's capabilities and constraints
that the provider needs for token counting and cost estimation.
"""
model: str
context_window: int = 128000 # Default to 128K
max_output_tokens: int = 4096
supports_tools: bool = True
supports_vision: bool = False
supports_streaming: bool = True
tokenizer_backend: str | None = None # Force specific tokenizer
input_cost_per_1m: float | None = None # Cost per 1M input tokens
output_cost_per_1m: float | None = None # Cost per 1M output tokens
# Default context limits for common open models
# These are reasonable defaults; users can override
_DEFAULT_CONTEXT_LIMITS: dict[str, int] = {
# Llama 3 family
"llama-3": 8192,
"llama-3-8b": 8192,
"llama-3-70b": 8192,
"llama-3.1": 128000,
"llama-3.1-8b": 128000,
"llama-3.1-70b": 128000,
"llama-3.1-405b": 128000,
"llama-3.2": 128000,
"llama-3.3": 128000,
# Llama 2 family
"llama-2": 4096,
"llama-2-7b": 4096,
"llama-2-13b": 4096,
"llama-2-70b": 4096,
"codellama": 16384,
# Mistral family
"mistral": 32768,
"mistral-7b": 32768,
"mistral-nemo": 128000,
"mistral-small": 32768,
"mistral-large": 128000,
"mixtral": 32768,
"mixtral-8x7b": 32768,
"mixtral-8x22b": 65536,
# Qwen family
"qwen": 32768,
"qwen2": 32768,
"qwen2-7b": 32768,
"qwen2-72b": 32768,
"qwen2.5": 131072,
# DeepSeek
"deepseek": 32768,
"deepseek-coder": 16384,
"deepseek-v2": 128000,
"deepseek-v3": 128000,
# Yi
"yi": 32768,
"yi-34b": 32768,
# Phi
"phi-2": 2048,
"phi-3": 4096,
"phi-3-mini": 4096,
"phi-3-medium": 4096,
# Others
"falcon": 2048,
"falcon-40b": 2048,
"falcon-180b": 2048,
"gemma": 8192,
"gemma-2": 8192,
"starcoder": 8192,
"starcoder2": 16384,
}
class OpenAICompatibleTokenCounter:
"""Token counter for OpenAI-compatible providers.
Uses the TokenizerRegistry to get the appropriate tokenizer
for the model, falling back to estimation if needed.
"""
def __init__(
self,
model: str,
tokenizer_backend: str | None = None,
):
"""Initialize token counter.
Args:
model: Model name.
tokenizer_backend: Force specific tokenizer backend.
"""
self.model = model
self._tokenizer = get_tokenizer(model, backend=tokenizer_backend)
def count_text(self, text: str) -> int:
"""Count tokens in text."""
return self._tokenizer.count_text(text)
def count_message(self, message: dict[str, Any]) -> int:
"""Count tokens in a single message."""
# Use OpenAI-style message overhead
tokens = 4 # Base overhead
role = message.get("role", "")
tokens += self.count_text(role)
content = message.get("content")
if content:
if isinstance(content, str):
tokens += self.count_text(content)
elif isinstance(content, list):
for part in content:
if isinstance(part, dict):
if part.get("type") == "text":
tokens += self.count_text(part.get("text", ""))
elif isinstance(part, str):
tokens += self.count_text(part)
name = message.get("name")
if name:
tokens += self.count_text(name) + 1
tool_calls = message.get("tool_calls")
if tool_calls:
for tc in tool_calls:
func = tc.get("function", {})
tokens += self.count_text(func.get("name", ""))
tokens += self.count_text(func.get("arguments", ""))
tokens += 10
tool_call_id = message.get("tool_call_id")
if tool_call_id:
tokens += self.count_text(tool_call_id) + 2
return tokens
def count_messages(self, messages: list[dict[str, Any]]) -> int:
"""Count tokens in a list of messages."""
total = sum(self.count_message(msg) for msg in messages)
total += 3 # Priming tokens
return total
class OpenAICompatibleProvider(Provider):
"""Provider for OpenAI-compatible LLM services.
Works with any service implementing the OpenAI chat completions API:
- Ollama (local)
- vLLM (local/cloud)
- Together AI
- Groq
- Fireworks AI
- LM Studio
- LocalAI
- And many more...
Example:
# For Ollama
provider = OpenAICompatibleProvider(
name="ollama",
base_url="http://localhost:11434/v1",
default_model="llama3.1",
)
# For Together AI
provider = OpenAICompatibleProvider(
name="together",
base_url="https://api.together.xyz/v1",
)
# Get token counter for a specific model
counter = provider.get_token_counter("llama-3.1-8b")
"""
def __init__(
self,
name: str = "openai_compatible",
base_url: str | None = None,
api_key: str | None = None,
default_model: str | None = None,
models: dict[str, ModelCapabilities] | None = None,
):
"""Initialize OpenAI-compatible provider.
Args:
name: Provider name for identification.
base_url: API base URL (e.g., 'http://localhost:11434/v1').
api_key: API key (if required).
default_model: Default model for operations.
models: Custom model configurations.
"""
self._name = name
self.base_url = base_url
self.api_key = api_key
self.default_model = default_model
self._models: dict[str, ModelCapabilities] = models or {}
@property
def name(self) -> str:
return self._name
def register_model(
self,
model: str,
capabilities: ModelCapabilities | None = None,
**kwargs: Any,
) -> None:
"""Register a model with its capabilities.
Args:
model: Model name.
capabilities: Model capabilities object.
**kwargs: Alternative way to specify capabilities.
"""
if capabilities is not None:
self._models[model] = capabilities
else:
self._models[model] = ModelCapabilities(model=model, **kwargs)
def supports_model(self, model: str) -> bool:
"""Check if model is supported.
OpenAI-compatible providers support any model by default,
using estimation for token counting.
"""
return True # Always return True - we can estimate
def get_token_counter(self, model: str) -> OpenAICompatibleTokenCounter:
"""Get token counter for a model.
Uses the TokenizerRegistry to find the best tokenizer,
with fallback to estimation.
"""
tokenizer_backend = None
# Check for registered model with specific tokenizer
if model in self._models:
tokenizer_backend = self._models[model].tokenizer_backend
return OpenAICompatibleTokenCounter(model, tokenizer_backend)
def get_context_limit(self, model: str) -> int:
"""Get context limit for a model.
Priority:
1. Registered model capabilities
2. Default limits for known models
3. Prefix matching
4. Default 128K
"""
# Check registered models
if model in self._models:
return self._models[model].context_window
model_lower = model.lower()
# Check default limits
if model_lower in _DEFAULT_CONTEXT_LIMITS:
return _DEFAULT_CONTEXT_LIMITS[model_lower]
# Prefix match
for prefix, limit in _DEFAULT_CONTEXT_LIMITS.items():
if model_lower.startswith(prefix):
return limit
# Default to 128K for modern models
return 128000
def get_output_buffer(self, model: str, default: int = 4000) -> int:
"""Get recommended output buffer."""
if model in self._models:
return min(self._models[model].max_output_tokens, default)
return default
def estimate_cost(
self,
input_tokens: int,
output_tokens: int,
model: str,
cached_tokens: int = 0,
) -> float | None:
"""Estimate cost if pricing is configured.
Args:
input_tokens: Number of input tokens.
output_tokens: Number of output tokens.
model: Model name.
cached_tokens: Number of cached tokens.
Returns:
Estimated cost in USD, or None if pricing unknown.
"""
if model not in self._models:
return None
caps = self._models[model]
if caps.input_cost_per_1m is None or caps.output_cost_per_1m is None:
return None
input_cost = (input_tokens / 1_000_000) * caps.input_cost_per_1m
output_cost = (output_tokens / 1_000_000) * caps.output_cost_per_1m
return input_cost + output_cost
# Pre-configured provider factories for common services
def create_ollama_provider(
base_url: str = "http://localhost:11434/v1",
) -> OpenAICompatibleProvider:
"""Create provider for Ollama.
Ollama is a popular local LLM runner that supports many open models.
Args:
base_url: Ollama API URL (default: http://localhost:11434/v1).
Returns:
Configured provider.
"""
return OpenAICompatibleProvider(
name="ollama",
base_url=base_url,
)
def create_together_provider(
api_key: str | None = None,
) -> OpenAICompatibleProvider:
"""Create provider for Together AI.
Together AI offers high-performance inference for open models.
Args:
api_key: Together AI API key.
Returns:
Configured provider with Together AI pricing.
"""
provider = OpenAICompatibleProvider(
name="together",
base_url="https://api.together.xyz/v1",
api_key=api_key,
)
# Register common Together models with pricing
# Pricing as of Jan 2025 (verify current rates)
provider.register_model(
"meta-llama/Llama-3.1-8B-Instruct-Turbo",
context_window=128000,
input_cost_per_1m=0.18,
output_cost_per_1m=0.18,
)
provider.register_model(
"meta-llama/Llama-3.1-70B-Instruct-Turbo",
context_window=128000,
input_cost_per_1m=0.88,
output_cost_per_1m=0.88,
)
provider.register_model(
"meta-llama/Llama-3.1-405B-Instruct-Turbo",
context_window=128000,
input_cost_per_1m=3.50,
output_cost_per_1m=3.50,
)
return provider
def create_groq_provider(
api_key: str | None = None,
) -> OpenAICompatibleProvider:
"""Create provider for Groq.
Groq offers ultra-fast inference on custom hardware.
Args:
api_key: Groq API key.
Returns:
Configured provider with Groq pricing.
"""
provider = OpenAICompatibleProvider(
name="groq",
base_url="https://api.groq.com/openai/v1",
api_key=api_key,
)
# Register common Groq models with pricing
# Pricing as of Jan 2025 (verify current rates)
provider.register_model(
"llama-3.1-8b-instant",
context_window=128000,
input_cost_per_1m=0.05,
output_cost_per_1m=0.08,
)
provider.register_model(
"llama-3.1-70b-versatile",
context_window=128000,
input_cost_per_1m=0.59,
output_cost_per_1m=0.79,
)
provider.register_model(
"mixtral-8x7b-32768",
context_window=32768,
input_cost_per_1m=0.24,
output_cost_per_1m=0.24,
)
return provider
def create_fireworks_provider(
api_key: str | None = None,
) -> OpenAICompatibleProvider:
"""Create provider for Fireworks AI.
Args:
api_key: Fireworks API key.
Returns:
Configured provider.
"""
return OpenAICompatibleProvider(
name="fireworks",
base_url="https://api.fireworks.ai/inference/v1",
api_key=api_key,
)
def create_anyscale_provider(
api_key: str | None = None,
) -> OpenAICompatibleProvider:
"""Create provider for Anyscale Endpoints.
Args:
api_key: Anyscale API key.
Returns:
Configured provider.
"""
return OpenAICompatibleProvider(
name="anyscale",
base_url="https://api.endpoints.anyscale.com/v1",
api_key=api_key,
)
def create_vllm_provider(
base_url: str,
) -> OpenAICompatibleProvider:
"""Create provider for vLLM server.
vLLM is a high-performance inference engine.
Args:
base_url: vLLM server URL (e.g., 'http://localhost:8000/v1').
Returns:
Configured provider.
"""
return OpenAICompatibleProvider(
name="vllm",
base_url=base_url,
)
def create_lmstudio_provider(
base_url: str = "http://localhost:1234/v1",
) -> OpenAICompatibleProvider:
"""Create provider for LM Studio.
LM Studio is a desktop app for running local LLMs.
Args:
base_url: LM Studio API URL.
Returns:
Configured provider.
"""
return OpenAICompatibleProvider(
name="lmstudio",
base_url=base_url,
)

View file

@ -0,0 +1,19 @@
"""Headroom Proxy Server.
A transparent proxy that sits between LLM clients (Claude Code, Cursor, etc.)
and LLM APIs (Anthropic, OpenAI), applying Headroom optimizations.
Usage:
# Start the proxy
python -m headroom.proxy.server
# Use with Claude Code
ANTHROPIC_BASE_URL=http://localhost:8787 claude
# Use with Cursor (if using Anthropic)
Set base URL in Cursor settings to http://localhost:8787
"""
from .server import create_app, run_server
__all__ = ["create_app", "run_server"]

1399
headroom/proxy/server.py Normal file

File diff suppressed because it is too large Load diff

0
headroom/py.typed Normal file
View file

View file

@ -20,9 +20,8 @@ from __future__ import annotations
import math
import re
from collections import Counter
from typing import Any
from .base import RelevanceScore, RelevanceScorer, default_batch_score
from .base import RelevanceScore, RelevanceScorer
class BM25Scorer(RelevanceScorer):

View file

@ -20,7 +20,7 @@ Limitations:
from __future__ import annotations
import logging
from typing import TYPE_CHECKING, Any
from typing import TYPE_CHECKING
import numpy as np
@ -71,7 +71,7 @@ class EmbeddingScorer(RelevanceScorer):
Requires sentence-transformers: pip install headroom[relevance]
"""
_model_cache: dict[str, "SentenceTransformer"] = {}
_model_cache: dict[str, SentenceTransformer] = {}
def __init__(
self,
@ -93,7 +93,7 @@ class EmbeddingScorer(RelevanceScorer):
self.model_name = model_name
self.device = device
self.cache_model = cache_model
self._model: "SentenceTransformer | None" = None
self._model: SentenceTransformer | None = None
self._available: bool | None = None
@classmethod
@ -110,7 +110,7 @@ class EmbeddingScorer(RelevanceScorer):
except ImportError:
return False
def _get_model(self) -> "SentenceTransformer":
def _get_model(self) -> SentenceTransformer:
"""Get or load the sentence transformer model.
Returns:

View file

@ -11,7 +11,6 @@ from jinja2 import Template
from ..storage import create_storage
from ..utils import estimate_cost, format_cost
# HTML template embedded as string
REPORT_TEMPLATE = """
<!DOCTYPE html>

View file

@ -3,8 +3,9 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from collections.abc import Iterator
from datetime import datetime
from typing import Any, Iterator
from typing import Any
from ..config import RequestMetrics

View file

@ -3,9 +3,10 @@
from __future__ import annotations
import json
from collections.abc import Iterator
from datetime import datetime
from pathlib import Path
from typing import Any, Iterator
from typing import Any
from ..config import RequestMetrics
from ..utils import format_timestamp, parse_timestamp

View file

@ -4,9 +4,10 @@ from __future__ import annotations
import json
import sqlite3
from collections.abc import Iterator
from datetime import datetime
from pathlib import Path
from typing import Any, Iterator
from typing import Any
from ..config import RequestMetrics
from ..utils import format_timestamp, parse_timestamp

View file

@ -0,0 +1,72 @@
"""Pluggable tokenizer system for universal LLM support.
This module provides a registry-based tokenizer system that supports
multiple backends:
1. tiktoken - OpenAI models (GPT-3.5, GPT-4, GPT-4o)
2. HuggingFace - Open models (Llama, Mistral, Falcon, etc.)
3. Anthropic - Claude models (via SDK or estimation)
4. Estimation - Fallback for unknown models
Usage:
from headroom.tokenizers import TokenizerRegistry, get_tokenizer
# Auto-detect tokenizer from model name
tokenizer = get_tokenizer("gpt-4o")
tokens = tokenizer.count_text("Hello, world!")
# Get tokenizer for specific backend
tokenizer = get_tokenizer("llama-3-8b", backend="huggingface")
# Register custom tokenizer
TokenizerRegistry.register("my-model", my_tokenizer)
"""
from .base import BaseTokenizer, TokenCounter
from .estimator import CharacterCounter, EstimatingTokenCounter
from .registry import (
TokenizerRegistry,
get_tokenizer,
list_supported_models,
register_tokenizer,
)
from .tiktoken_counter import TiktokenCounter
# Lazy imports for optional dependencies
def get_huggingface_tokenizer():
"""Get HuggingFaceTokenizer class (requires transformers)."""
from .huggingface import HuggingFaceTokenizer
return HuggingFaceTokenizer
def get_mistral_tokenizer():
"""Get MistralTokenizer class (requires mistral-common)."""
from .mistral import MistralTokenizer
return MistralTokenizer
def is_mistral_tokenizer_available() -> bool:
"""Check if Mistral tokenizer is available."""
from .mistral import is_mistral_available
return is_mistral_available()
__all__ = [
# Registry
"TokenizerRegistry",
"get_tokenizer",
"register_tokenizer",
"list_supported_models",
# Base classes
"TokenCounter",
"BaseTokenizer",
# Implementations
"TiktokenCounter",
"EstimatingTokenCounter",
"CharacterCounter",
# Lazy loaders
"get_huggingface_tokenizer",
"get_mistral_tokenizer",
"is_mistral_tokenizer_available",
]

203
headroom/tokenizers/base.py Normal file
View file

@ -0,0 +1,203 @@
"""Base classes for tokenizer implementations.
Defines the TokenCounter protocol and BaseTokenizer class that all
tokenizer backends must implement.
"""
from __future__ import annotations
import json
from abc import ABC, abstractmethod
from typing import Any, Protocol, runtime_checkable
@runtime_checkable
class TokenCounter(Protocol):
"""Protocol for token counting implementations.
Any class implementing this protocol can be used with Headroom
for token counting. This allows integration with various
tokenizer backends (tiktoken, HuggingFace, custom, etc.).
"""
def count_text(self, text: str) -> int:
"""Count tokens in a text string.
Args:
text: The text to count tokens for.
Returns:
Number of tokens in the text.
"""
...
def count_messages(self, messages: list[dict[str, Any]]) -> int:
"""Count tokens in a list of chat messages.
Args:
messages: List of message dicts with 'role' and 'content'.
Returns:
Total token count including message overhead.
"""
...
class BaseTokenizer(ABC):
"""Abstract base class for tokenizer implementations.
Provides common functionality for counting messages while
requiring subclasses to implement text tokenization.
"""
# Token overhead per message (role, formatting, etc.)
# Override in subclasses for model-specific overhead
MESSAGE_OVERHEAD = 4
REPLY_OVERHEAD = 3 # Assistant reply start tokens
@abstractmethod
def count_text(self, text: str) -> int:
"""Count tokens in a text string. Must be implemented by subclasses."""
pass
def count_messages(self, messages: list[dict[str, Any]]) -> int:
"""Count tokens in a list of chat messages.
Uses OpenAI-style message counting as the baseline, which
works well for most models.
Args:
messages: List of message dicts.
Returns:
Total token count.
"""
total = 0
for message in messages:
# Base message overhead
total += self.MESSAGE_OVERHEAD
# Count role
role = message.get("role", "")
total += self.count_text(role)
# Count content
content = message.get("content")
if content is not None:
if isinstance(content, str):
total += self.count_text(content)
elif isinstance(content, list):
# Multi-part content (images, tool results, etc.)
total += self._count_content_parts(content)
# Count tool calls
tool_calls = message.get("tool_calls")
if tool_calls:
total += self._count_tool_calls(tool_calls)
# Count function call (legacy)
function_call = message.get("function_call")
if function_call:
total += self._count_function_call(function_call)
# Count name field
name = message.get("name")
if name:
total += self.count_text(name)
total += 1 # Name field overhead
# Reply start overhead
total += self.REPLY_OVERHEAD
return total
def _count_content_parts(self, parts: list[Any]) -> int:
"""Count tokens in multi-part content."""
total = 0
for part in parts:
if isinstance(part, dict):
part_type = part.get("type", "")
if part_type == "text":
total += self.count_text(part.get("text", ""))
elif part_type == "image_url":
# Images have fixed token cost (varies by model)
total += 85 # Base image token count
elif part_type == "tool_result":
content = part.get("content", "")
if isinstance(content, str):
total += self.count_text(content)
else:
total += self.count_text(json.dumps(content))
elif part_type == "tool_use":
total += self.count_text(part.get("name", ""))
total += self.count_text(json.dumps(part.get("input", {})))
else:
# Unknown type - estimate from JSON
total += self.count_text(json.dumps(part))
elif isinstance(part, str):
total += self.count_text(part)
return total
def _count_tool_calls(self, tool_calls: list[dict[str, Any]]) -> int:
"""Count tokens in tool calls."""
total = 0
for call in tool_calls:
total += 4 # Tool call overhead
if "function" in call:
func = call["function"]
total += self.count_text(func.get("name", ""))
total += self.count_text(func.get("arguments", ""))
if "id" in call:
total += self.count_text(call["id"])
return total
def _count_function_call(self, function_call: dict[str, Any]) -> int:
"""Count tokens in legacy function call."""
total = 4 # Function call overhead
total += self.count_text(function_call.get("name", ""))
total += self.count_text(function_call.get("arguments", ""))
return total
def encode(self, text: str) -> list[int]:
"""Encode text to token IDs.
Optional method - not all backends support encoding.
Default implementation raises NotImplementedError.
Args:
text: Text to encode.
Returns:
List of token IDs.
Raises:
NotImplementedError: If encoding is not supported.
"""
raise NotImplementedError(
f"{self.__class__.__name__} does not support encoding"
)
def decode(self, tokens: list[int]) -> str:
"""Decode token IDs to text.
Optional method - not all backends support decoding.
Default implementation raises NotImplementedError.
Args:
tokens: List of token IDs.
Returns:
Decoded text.
Raises:
NotImplementedError: If decoding is not supported.
"""
raise NotImplementedError(
f"{self.__class__.__name__} does not support decoding"
)

View file

@ -0,0 +1,199 @@
"""Estimation-based token counter for fallback scenarios.
When no exact tokenizer is available (e.g., unknown models, missing
dependencies), this provides a reasonable approximation based on
character/word heuristics calibrated against real tokenizers.
"""
from __future__ import annotations
import json
import re
from typing import Any
from .base import BaseTokenizer
class EstimatingTokenCounter(BaseTokenizer):
"""Token counter using estimation heuristics.
This is the fallback tokenizer used when:
- Model is unknown/unsupported
- Required tokenizer library not installed
- Speed is prioritized over accuracy
The estimation is calibrated against tiktoken cl100k_base and
provides ~90% accuracy for typical text. It tends to slightly
overestimate, which is safer for context window management.
Estimation Strategy:
- Base: ~4 characters per token (calibrated against GPT-4)
- Adjustments for code, URLs, numbers, whitespace
- Special handling for JSON structure
Example:
counter = EstimatingTokenCounter()
tokens = counter.count_text("Hello, world!")
print(f"Estimated tokens: {tokens}")
"""
# Calibration constants (derived from tiktoken analysis)
CHARS_PER_TOKEN = 4.0 # Average for English text
CHARS_PER_TOKEN_CODE = 3.5 # Code is denser
CHARS_PER_TOKEN_JSON = 3.2 # JSON has more structure
# Patterns for content type detection
CODE_PATTERN = re.compile(
r'(?:def |class |function |const |let |var |import |from |'
r'if \(|for \(|while \(|switch \(|try \{|catch \(|'
r'=>|->|\{\{|\}\}|;$)',
re.MULTILINE
)
JSON_PATTERN = re.compile(r'^\s*[\[\{]')
URL_PATTERN = re.compile(r'https?://\S+')
UUID_PATTERN = re.compile(
r'[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}',
re.IGNORECASE
)
def __init__(self, chars_per_token: float | None = None):
"""Initialize estimating counter.
Args:
chars_per_token: Override default chars per token ratio.
If None, auto-detects based on content type.
"""
self._fixed_ratio = chars_per_token
def count_text(self, text: str) -> int:
"""Estimate token count for text.
Args:
text: Text to count tokens for.
Returns:
Estimated number of tokens.
"""
if not text:
return 0
# Use fixed ratio if provided
if self._fixed_ratio is not None:
return max(1, int(len(text) / self._fixed_ratio + 0.5))
# Auto-detect content type and adjust ratio
ratio = self._detect_ratio(text)
# Apply ratio with minimum of 1 token
base_count = int(len(text) / ratio + 0.5)
# Add overhead for special patterns
overhead = self._count_special_overhead(text)
return max(1, base_count + overhead)
def _detect_ratio(self, text: str) -> float:
"""Detect optimal chars-per-token ratio based on content.
Args:
text: Text to analyze.
Returns:
Chars per token ratio.
"""
# Check for JSON
if self.JSON_PATTERN.match(text):
try:
json.loads(text)
return self.CHARS_PER_TOKEN_JSON
except (json.JSONDecodeError, ValueError):
pass
# Check for code
code_matches = len(self.CODE_PATTERN.findall(text))
if code_matches > len(text) / 500: # ~2 matches per KB
return self.CHARS_PER_TOKEN_CODE
return self.CHARS_PER_TOKEN
def _count_special_overhead(self, text: str) -> int:
"""Count additional tokens for special patterns.
URLs and UUIDs often tokenize into more tokens than
character count would suggest.
Args:
text: Text to analyze.
Returns:
Additional token overhead.
"""
overhead = 0
# URLs typically tokenize to more tokens
urls = self.URL_PATTERN.findall(text)
for url in urls:
# Each URL component adds overhead
overhead += url.count('/') + url.count('?') + url.count('&')
# UUIDs are typically 8-10 tokens despite being 36 chars
uuids = self.UUID_PATTERN.findall(text)
overhead += len(uuids) * 2 # Each UUID adds ~2 extra tokens
return overhead
def count_messages(self, messages: list[dict[str, Any]]) -> int:
"""Estimate tokens in chat messages.
Uses the base class implementation with estimation-based
text counting.
Args:
messages: List of chat messages.
Returns:
Estimated total token count.
"""
# Use base class implementation
return super().count_messages(messages)
def __repr__(self) -> str:
if self._fixed_ratio:
return f"EstimatingTokenCounter(chars_per_token={self._fixed_ratio})"
return "EstimatingTokenCounter(auto)"
class CharacterCounter(BaseTokenizer):
"""Simple character-based counter.
Uses a fixed character-to-token ratio. Useful for:
- Quick approximations
- Testing
- Models with unknown tokenization
This is less accurate than EstimatingTokenCounter but faster.
"""
def __init__(self, chars_per_token: float = 4.0):
"""Initialize character counter.
Args:
chars_per_token: Characters per token ratio.
"""
self.chars_per_token = chars_per_token
def count_text(self, text: str) -> int:
"""Count tokens based on character count.
Args:
text: Text to count.
Returns:
Estimated token count.
"""
if not text:
return 0
return max(1, int(len(text) / self.chars_per_token + 0.5))
def __repr__(self) -> str:
return f"CharacterCounter(chars_per_token={self.chars_per_token})"

View file

@ -0,0 +1,316 @@
"""HuggingFace tokenizer wrapper for open models.
Supports Llama, Mistral, Falcon, and other models with HuggingFace
tokenizers. Requires the `transformers` library.
"""
from __future__ import annotations
import logging
from functools import lru_cache
from typing import Any
from .base import BaseTokenizer
logger = logging.getLogger(__name__)
# Model name to HuggingFace tokenizer mapping
# Maps common model names to their HuggingFace tokenizer identifiers
MODEL_TO_TOKENIZER: dict[str, str] = {
# Llama 3 family
"llama-3": "meta-llama/Meta-Llama-3-8B",
"llama-3-8b": "meta-llama/Meta-Llama-3-8B",
"llama-3-70b": "meta-llama/Meta-Llama-3-70B",
"llama-3.1-8b": "meta-llama/Llama-3.1-8B",
"llama-3.1-70b": "meta-llama/Llama-3.1-70B",
"llama-3.1-405b": "meta-llama/Llama-3.1-405B",
"llama-3.2-1b": "meta-llama/Llama-3.2-1B",
"llama-3.2-3b": "meta-llama/Llama-3.2-3B",
"llama-3.3-70b": "meta-llama/Llama-3.3-70B-Instruct",
# Llama 2 family
"llama-2": "meta-llama/Llama-2-7b-hf",
"llama-2-7b": "meta-llama/Llama-2-7b-hf",
"llama-2-13b": "meta-llama/Llama-2-13b-hf",
"llama-2-70b": "meta-llama/Llama-2-70b-hf",
# CodeLlama
"codellama": "codellama/CodeLlama-7b-hf",
"codellama-7b": "codellama/CodeLlama-7b-hf",
"codellama-13b": "codellama/CodeLlama-13b-hf",
"codellama-34b": "codellama/CodeLlama-34b-hf",
# Mistral family
"mistral": "mistralai/Mistral-7B-v0.1",
"mistral-7b": "mistralai/Mistral-7B-v0.1",
"mistral-7b-v0.2": "mistralai/Mistral-7B-Instruct-v0.2",
"mistral-7b-v0.3": "mistralai/Mistral-7B-Instruct-v0.3",
"mistral-nemo": "mistralai/Mistral-Nemo-Base-2407",
"mistral-small": "mistralai/Mistral-Small-Instruct-2409",
"mistral-large": "mistralai/Mistral-Large-Instruct-2407",
# Mixtral
"mixtral": "mistralai/Mixtral-8x7B-v0.1",
"mixtral-8x7b": "mistralai/Mixtral-8x7B-v0.1",
"mixtral-8x22b": "mistralai/Mixtral-8x22B-v0.1",
# Qwen family
"qwen": "Qwen/Qwen-7B",
"qwen-7b": "Qwen/Qwen-7B",
"qwen-14b": "Qwen/Qwen-14B",
"qwen-72b": "Qwen/Qwen-72B",
"qwen2": "Qwen/Qwen2-7B",
"qwen2-7b": "Qwen/Qwen2-7B",
"qwen2-72b": "Qwen/Qwen2-72B",
"qwen2.5": "Qwen/Qwen2.5-7B",
"qwen2.5-7b": "Qwen/Qwen2.5-7B",
"qwen2.5-72b": "Qwen/Qwen2.5-72B",
# DeepSeek
"deepseek": "deepseek-ai/deepseek-llm-7b-base",
"deepseek-7b": "deepseek-ai/deepseek-llm-7b-base",
"deepseek-67b": "deepseek-ai/deepseek-llm-67b-base",
"deepseek-coder": "deepseek-ai/deepseek-coder-6.7b-base",
"deepseek-v2": "deepseek-ai/DeepSeek-V2",
"deepseek-v3": "deepseek-ai/DeepSeek-V3",
# Yi family
"yi": "01-ai/Yi-6B",
"yi-6b": "01-ai/Yi-6B",
"yi-34b": "01-ai/Yi-34B",
"yi-1.5": "01-ai/Yi-1.5-6B",
# Phi family
"phi-2": "microsoft/phi-2",
"phi-3": "microsoft/Phi-3-mini-4k-instruct",
"phi-3-mini": "microsoft/Phi-3-mini-4k-instruct",
"phi-3-small": "microsoft/Phi-3-small-8k-instruct",
"phi-3-medium": "microsoft/Phi-3-medium-4k-instruct",
# Falcon
"falcon": "tiiuae/falcon-7b",
"falcon-7b": "tiiuae/falcon-7b",
"falcon-40b": "tiiuae/falcon-40b",
"falcon-180b": "tiiuae/falcon-180B",
# StarCoder
"starcoder": "bigcode/starcoder",
"starcoder2": "bigcode/starcoder2-15b",
"starcoder2-3b": "bigcode/starcoder2-3b",
"starcoder2-7b": "bigcode/starcoder2-7b",
"starcoder2-15b": "bigcode/starcoder2-15b",
# MPT
"mpt-7b": "mosaicml/mpt-7b",
"mpt-30b": "mosaicml/mpt-30b",
# Gemma
"gemma": "google/gemma-7b",
"gemma-2b": "google/gemma-2b",
"gemma-7b": "google/gemma-7b",
"gemma-2": "google/gemma-2-9b",
"gemma-2-9b": "google/gemma-2-9b",
"gemma-2-27b": "google/gemma-2-27b",
}
@lru_cache(maxsize=16)
def _load_tokenizer(tokenizer_name: str):
"""Load and cache HuggingFace tokenizer.
Args:
tokenizer_name: HuggingFace model/tokenizer name.
Returns:
Loaded tokenizer, or None if unavailable.
"""
from transformers import AutoTokenizer
try:
return AutoTokenizer.from_pretrained(
tokenizer_name,
trust_remote_code=True,
)
except Exception as e:
logger.warning(f"Failed to load tokenizer {tokenizer_name}: {e}")
return None
def get_tokenizer_name(model: str) -> str:
"""Get HuggingFace tokenizer name for a model.
Args:
model: Model name.
Returns:
HuggingFace tokenizer identifier.
"""
model_lower = model.lower()
# Direct lookup
if model_lower in MODEL_TO_TOKENIZER:
return MODEL_TO_TOKENIZER[model_lower]
# Try prefix matching
for key, value in MODEL_TO_TOKENIZER.items():
if model_lower.startswith(key):
return value
# Assume model name is the tokenizer name
return model
class HuggingFaceTokenizer(BaseTokenizer):
"""Token counter using HuggingFace tokenizers.
Supports any model with a HuggingFace tokenizer, including:
- Llama family (Llama 2, Llama 3, CodeLlama)
- Mistral family (Mistral, Mixtral)
- Qwen family
- DeepSeek family
- Phi family
- Falcon, StarCoder, MPT, Gemma, etc.
Requires the `transformers` library:
pip install transformers
Some models may require authentication:
huggingface-cli login
Example:
counter = HuggingFaceTokenizer("llama-3-8b")
tokens = counter.count_text("Hello, world!")
"""
# Overhead per message (varies by model, this is a reasonable default)
MESSAGE_OVERHEAD = 4
REPLY_OVERHEAD = 3
def __init__(self, model: str):
"""Initialize HuggingFace tokenizer.
Args:
model: Model name (e.g., 'llama-3-8b', 'mistral-7b').
"""
self.model = model
self.tokenizer_name = get_tokenizer_name(model)
self._tokenizer = None # Lazy load
@property
def tokenizer(self):
"""Lazy-load the tokenizer."""
if self._tokenizer is None:
loaded = _load_tokenizer(self.tokenizer_name)
if loaded is not None:
self._tokenizer = loaded
else:
# Mark as unavailable
self._tokenizer = False
return self._tokenizer if self._tokenizer is not False else None
def _use_fallback(self) -> bool:
"""Check if we need to use fallback estimation."""
return self.tokenizer is None
def count_text(self, text: str) -> int:
"""Count tokens in text.
Falls back to estimation if tokenizer unavailable.
Args:
text: Text to tokenize.
Returns:
Number of tokens.
"""
if not text:
return 0
if self._use_fallback():
# Fall back to ~4 chars per token estimation
return max(1, int(len(text) / 4 + 0.5))
tokens = self.tokenizer.encode(text, add_special_tokens=False)
return len(tokens)
def count_messages(self, messages: list[dict[str, Any]]) -> int:
"""Count tokens in chat messages.
Uses the model's chat template if available, otherwise
falls back to base class implementation.
Args:
messages: List of chat messages.
Returns:
Total token count.
"""
if self._use_fallback():
# Use base class implementation with estimation
return super().count_messages(messages)
# Try to use chat template for accurate counting
if hasattr(self.tokenizer, "apply_chat_template"):
try:
# Apply chat template and count
formatted = self.tokenizer.apply_chat_template(
messages,
tokenize=True,
add_generation_prompt=True,
)
return len(formatted)
except Exception:
# Fall back to base implementation
pass
return super().count_messages(messages)
def encode(self, text: str) -> list[int]:
"""Encode text to token IDs.
Args:
text: Text to encode.
Returns:
List of token IDs.
Raises:
NotImplementedError: If tokenizer not available.
"""
if self._use_fallback():
raise NotImplementedError(
f"Encoding not available for {self.model} - "
f"tokenizer {self.tokenizer_name} could not be loaded"
)
return self.tokenizer.encode(text, add_special_tokens=False)
def decode(self, tokens: list[int]) -> str:
"""Decode token IDs to text.
Args:
tokens: List of token IDs.
Returns:
Decoded text.
Raises:
NotImplementedError: If tokenizer not available.
"""
if self._use_fallback():
raise NotImplementedError(
f"Decoding not available for {self.model} - "
f"tokenizer {self.tokenizer_name} could not be loaded"
)
return self.tokenizer.decode(tokens)
@classmethod
def is_available(cls) -> bool:
"""Check if HuggingFace tokenizers are available.
Returns:
True if transformers is installed.
"""
try:
import transformers
return True
except ImportError:
return False
@classmethod
def list_supported_models(cls) -> list[str]:
"""List models with known tokenizer mappings.
Returns:
List of supported model names.
"""
return list(MODEL_TO_TOKENIZER.keys())
def __repr__(self) -> str:
return f"HuggingFaceTokenizer(model={self.model!r}, tokenizer={self.tokenizer_name!r})"

View file

@ -0,0 +1,244 @@
"""Mistral tokenizer using the official mistral-common package.
Mistral AI released their tokenizer publicly, making accurate
token counting possible without API calls.
Requires: pip install mistral-common
"""
from __future__ import annotations
import logging
from functools import lru_cache
from typing import Any
from .base import BaseTokenizer
logger = logging.getLogger(__name__)
# Check if mistral-common is available
try:
from mistral_common.protocol.instruct.messages import (
AssistantMessage,
SystemMessage,
UserMessage,
)
from mistral_common.protocol.instruct.request import ChatCompletionRequest
from mistral_common.tokens.tokenizers.mistral import MistralTokenizer as _MistralTokenizer
MISTRAL_AVAILABLE = True
except ImportError:
MISTRAL_AVAILABLE = False
_MistralTokenizer = None
def is_mistral_available() -> bool:
"""Check if mistral-common is installed."""
return MISTRAL_AVAILABLE
# Model to tokenizer version mapping
MODEL_TO_VERSION = {
# Mistral models use v3 tokenizer (tekken)
"mistral-large": "v3",
"mistral-large-latest": "v3",
"mistral-small": "v3",
"mistral-small-latest": "v3",
"ministral-8b": "v3",
"ministral-3b": "v3",
"mistral-nemo": "v3",
"pixtral-12b": "v3",
"codestral": "v3",
"codestral-latest": "v3",
# Mixtral uses v1
"mixtral-8x7b": "v1",
"mixtral-8x22b": "v1",
"open-mixtral-8x7b": "v1",
"open-mixtral-8x22b": "v1",
# Mistral 7B uses v1
"mistral-7b": "v1",
"open-mistral-7b": "v1",
"mistral-7b-instruct": "v1",
}
@lru_cache(maxsize=4)
def _get_tokenizer(version: str):
"""Get and cache Mistral tokenizer by version."""
if not MISTRAL_AVAILABLE:
raise RuntimeError(
"mistral-common is required for MistralTokenizer. "
"Install with: pip install mistral-common"
)
if version == "v3":
return _MistralTokenizer.v3(is_tekken=True)
elif version == "v2":
return _MistralTokenizer.v2()
else: # v1
return _MistralTokenizer.v1()
def get_tokenizer_version(model: str) -> str:
"""Get tokenizer version for a model."""
model_lower = model.lower()
# Direct lookup
if model_lower in MODEL_TO_VERSION:
return MODEL_TO_VERSION[model_lower]
# Prefix matching
for prefix, version in [
("mistral-large", "v3"),
("mistral-small", "v3"),
("ministral", "v3"),
("codestral", "v3"),
("pixtral", "v3"),
("mistral-nemo", "v3"),
("mixtral", "v1"),
("mistral-7b", "v1"),
("open-mistral", "v1"),
]:
if model_lower.startswith(prefix):
return version
# Default to v3 for newer models
return "v3"
class MistralTokenizer(BaseTokenizer):
"""Token counter using Mistral's official tokenizer.
Uses mistral-common package for accurate token counting.
Requires: pip install mistral-common
Example:
counter = MistralTokenizer("mistral-large")
tokens = counter.count_text("Hello, world!")
"""
MESSAGE_OVERHEAD = 4
REPLY_OVERHEAD = 3
def __init__(self, model: str = "mistral-large"):
"""Initialize Mistral tokenizer.
Args:
model: Mistral model name.
"""
if not MISTRAL_AVAILABLE:
raise RuntimeError(
"mistral-common is required for MistralTokenizer. "
"Install with: pip install mistral-common"
)
self.model = model
self.version = get_tokenizer_version(model)
self._tokenizer = None # Lazy load
@property
def tokenizer(self):
"""Lazy-load the tokenizer (MistralTokenizer object)."""
if self._tokenizer is None:
self._tokenizer = _get_tokenizer(self.version)
return self._tokenizer
@property
def _text_tokenizer(self):
"""Get the underlying text tokenizer for encode/decode."""
return self.tokenizer.instruct_tokenizer.tokenizer
def count_text(self, text: str) -> int:
"""Count tokens in text.
Args:
text: Text to tokenize.
Returns:
Number of tokens.
"""
if not text:
return 0
tokens = self._text_tokenizer.encode(text, bos=False, eos=False)
return len(tokens)
def count_messages(self, messages: list[dict[str, Any]]) -> int:
"""Count tokens in chat messages.
Uses Mistral's chat template for accurate counting.
Args:
messages: List of chat messages.
Returns:
Total token count.
"""
if not messages:
return 0
try:
# Convert to Mistral message format
mistral_messages = []
for msg in messages:
role = msg.get("role", "user")
content = msg.get("content", "")
if isinstance(content, list):
# Multi-part content - extract text
text_parts = []
for part in content:
if isinstance(part, dict) and part.get("type") == "text":
text_parts.append(part.get("text", ""))
elif isinstance(part, str):
text_parts.append(part)
content = "\n".join(text_parts)
if role == "user":
mistral_messages.append(UserMessage(content=content))
elif role == "assistant":
mistral_messages.append(AssistantMessage(content=content))
elif role == "system":
mistral_messages.append(SystemMessage(content=content))
else:
# Tool messages etc - treat as user
mistral_messages.append(UserMessage(content=content))
# Encode with chat template
request = ChatCompletionRequest(messages=mistral_messages)
tokenized = self.tokenizer.encode_chat_completion(request)
return len(tokenized.tokens)
except Exception as e:
logger.debug(f"Mistral chat encoding failed: {e}, falling back to text counting")
# Fallback to base implementation
return super().count_messages(messages)
def encode(self, text: str) -> list[int]:
"""Encode text to token IDs.
Args:
text: Text to encode.
Returns:
List of token IDs.
"""
return self._text_tokenizer.encode(text, bos=False, eos=False)
def decode(self, tokens: list[int]) -> str:
"""Decode token IDs to text.
Args:
tokens: List of token IDs.
Returns:
Decoded text.
"""
return self._text_tokenizer.decode(tokens)
@classmethod
def is_available(cls) -> bool:
"""Check if Mistral tokenizer is available."""
return MISTRAL_AVAILABLE
def __repr__(self) -> str:
return f"MistralTokenizer(model={self.model!r}, version={self.version!r})"

View file

@ -0,0 +1,398 @@
"""Tokenizer registry for universal model support.
Provides automatic tokenizer selection based on model name with
support for multiple backends and custom tokenizers.
"""
from __future__ import annotations
import logging
import re
from collections.abc import Callable
from typing import TYPE_CHECKING
from .base import TokenCounter
from .estimator import EstimatingTokenCounter
if TYPE_CHECKING:
pass
logger = logging.getLogger(__name__)
# Model pattern matching for tokenizer selection
# Order matters - more specific patterns first
MODEL_PATTERNS: list[tuple[str, str]] = [
# OpenAI models -> tiktoken
(r"^gpt-4o", "tiktoken"),
(r"^gpt-4", "tiktoken"),
(r"^gpt-3\.5", "tiktoken"),
(r"^o1", "tiktoken"),
(r"^o3", "tiktoken"),
(r"^text-embedding", "tiktoken"),
(r"^text-davinci", "tiktoken"),
(r"^code-", "tiktoken"),
(r"^davinci", "tiktoken"),
(r"^curie", "tiktoken"),
(r"^babbage", "tiktoken"),
(r"^ada", "tiktoken"),
# Anthropic models -> estimation (Claude uses custom tokenizer)
(r"^claude-", "anthropic"),
# Llama family -> huggingface (when available)
(r"^llama", "huggingface"),
(r"^meta-llama", "huggingface"),
(r"^codellama", "huggingface"),
# Mistral family -> official mistral tokenizer
(r"^mistral", "mistral"),
(r"^mixtral", "mistral"),
(r"^codestral", "mistral"),
(r"^ministral", "mistral"),
(r"^pixtral", "mistral"),
# Google models -> estimation (Gemini uses SentencePiece)
(r"^gemini", "google"),
(r"^palm", "google"),
# Cohere models -> estimation
(r"^command", "cohere"),
# Open models commonly served via OpenAI-compatible APIs
(r"^phi-", "huggingface"),
(r"^qwen", "huggingface"),
(r"^deepseek", "huggingface"),
(r"^yi-", "huggingface"),
(r"^falcon", "huggingface"),
(r"^mpt-", "huggingface"),
(r"^starcoder", "huggingface"),
(r"^codegen", "huggingface"),
]
class TokenizerRegistry:
"""Registry for tokenizer instances and factories.
Supports:
- Automatic tokenizer selection based on model name
- Custom tokenizer registration
- Multiple backends (tiktoken, huggingface, estimation)
- Lazy loading of tokenizer dependencies
Example:
# Auto-detect tokenizer
tokenizer = TokenizerRegistry.get("gpt-4o")
# Register custom tokenizer
TokenizerRegistry.register("my-model", my_tokenizer)
# Use specific backend
tokenizer = TokenizerRegistry.get("llama-3", backend="huggingface")
"""
# Singleton registry instance
_instance: TokenizerRegistry | None = None
# Registered tokenizers (model -> tokenizer instance)
_tokenizers: dict[str, TokenCounter] = {}
# Registered factories (backend -> factory function)
_factories: dict[str, Callable[[str], TokenCounter]] = {}
# Cache for auto-detected tokenizers
_cache: dict[str, TokenCounter] = {}
def __new__(cls) -> TokenizerRegistry:
"""Singleton pattern."""
if cls._instance is None:
cls._instance = super().__new__(cls)
cls._instance._init_factories()
return cls._instance
def _init_factories(self) -> None:
"""Initialize default tokenizer factories."""
self._factories = {
"tiktoken": self._create_tiktoken,
"huggingface": self._create_huggingface,
"anthropic": self._create_anthropic,
"google": self._create_google,
"cohere": self._create_cohere,
"mistral": self._create_mistral,
"estimation": self._create_estimation,
}
@classmethod
def get(
cls,
model: str,
backend: str | None = None,
fallback: bool = True,
) -> TokenCounter:
"""Get tokenizer for a model.
Args:
model: Model name (e.g., 'gpt-4o', 'claude-3-sonnet').
backend: Force specific backend ('tiktoken', 'huggingface', etc.).
If None, auto-detects based on model name.
fallback: If True, fall back to estimation on errors.
Returns:
TokenCounter instance for the model.
Raises:
ValueError: If backend not found and fallback=False.
"""
registry = cls()
model_lower = model.lower()
# Check for explicitly registered tokenizer
if model_lower in registry._tokenizers:
return registry._tokenizers[model_lower]
# Check cache
cache_key = f"{model_lower}:{backend or 'auto'}"
if cache_key in registry._cache:
return registry._cache[cache_key]
# Create tokenizer
try:
tokenizer = registry._create_tokenizer(model, backend)
registry._cache[cache_key] = tokenizer
return tokenizer
except Exception as e:
if fallback:
logger.warning(
f"Failed to create tokenizer for {model}: {e}. "
"Falling back to estimation."
)
tokenizer = EstimatingTokenCounter()
registry._cache[cache_key] = tokenizer
return tokenizer
raise ValueError(f"No tokenizer available for {model}: {e}") from e
@classmethod
def register(
cls,
model: str,
tokenizer: TokenCounter | None = None,
factory: Callable[[str], TokenCounter] | None = None,
) -> None:
"""Register a tokenizer or factory for a model.
Args:
model: Model name to register.
tokenizer: Pre-instantiated tokenizer instance.
factory: Factory function that creates tokenizer for model.
Raises:
ValueError: If neither tokenizer nor factory provided.
"""
registry = cls()
model_lower = model.lower()
if tokenizer is not None:
registry._tokenizers[model_lower] = tokenizer
elif factory is not None:
registry._factories[model_lower] = factory
else:
raise ValueError("Must provide either tokenizer or factory")
# Clear cache for this model
keys_to_remove = [k for k in registry._cache if k.startswith(model_lower)]
for key in keys_to_remove:
del registry._cache[key]
@classmethod
def register_backend(
cls,
backend: str,
factory: Callable[[str], TokenCounter],
) -> None:
"""Register a backend factory.
Args:
backend: Backend name.
factory: Factory function (model: str) -> TokenCounter.
"""
registry = cls()
registry._factories[backend] = factory
@classmethod
def list_backends(cls) -> list[str]:
"""List available backends."""
registry = cls()
return list(registry._factories.keys())
@classmethod
def list_registered(cls) -> list[str]:
"""List explicitly registered models."""
registry = cls()
return list(registry._tokenizers.keys())
@classmethod
def clear_cache(cls) -> None:
"""Clear the tokenizer cache."""
registry = cls()
registry._cache.clear()
def _create_tokenizer(
self,
model: str,
backend: str | None,
) -> TokenCounter:
"""Create tokenizer for model.
Args:
model: Model name.
backend: Backend to use (or None for auto-detect).
Returns:
TokenCounter instance.
"""
if backend is None:
backend = self._detect_backend(model)
factory = self._factories.get(backend)
if factory is None:
raise ValueError(f"Unknown backend: {backend}")
return factory(model)
def _create_mistral(self, model: str) -> TokenCounter:
"""Create Mistral tokenizer using official mistral-common."""
try:
from .mistral import MistralTokenizer, is_mistral_available
if is_mistral_available():
return MistralTokenizer(model)
except ImportError:
pass
logger.warning(
"mistral-common not installed for Mistral tokenizer. "
"Install with: pip install mistral-common"
)
return EstimatingTokenCounter()
def _detect_backend(self, model: str) -> str:
"""Detect best backend for model.
Args:
model: Model name.
Returns:
Backend name.
"""
model_lower = model.lower()
for pattern, backend in MODEL_PATTERNS:
if re.match(pattern, model_lower):
return backend
# Default to estimation for unknown models
return "estimation"
def _create_tiktoken(self, model: str) -> TokenCounter:
"""Create tiktoken-based tokenizer."""
try:
from .tiktoken_counter import TiktokenCounter
return TiktokenCounter(model)
except ImportError:
logger.warning(
"tiktoken not installed. Install with: pip install tiktoken"
)
return EstimatingTokenCounter()
def _create_huggingface(self, model: str) -> TokenCounter:
"""Create HuggingFace-based tokenizer."""
try:
from .huggingface import HuggingFaceTokenizer
return HuggingFaceTokenizer(model)
except ImportError:
logger.warning(
"transformers not installed for HuggingFace tokenizer. "
"Install with: pip install transformers"
)
return EstimatingTokenCounter()
except Exception as e:
logger.warning(f"Failed to load HuggingFace tokenizer for {model}: {e}")
return EstimatingTokenCounter()
def _create_anthropic(self, model: str) -> TokenCounter:
"""Create Anthropic tokenizer.
Anthropic uses a custom tokenizer that's not publicly available.
We use estimation calibrated for Claude models.
"""
# Claude models use ~3.5 chars per token on average
return EstimatingTokenCounter(chars_per_token=3.5)
def _create_google(self, model: str) -> TokenCounter:
"""Create Google tokenizer.
Gemini uses SentencePiece which isn't easily accessible.
We use estimation calibrated for Gemini models.
"""
# Gemini models use ~4 chars per token
return EstimatingTokenCounter(chars_per_token=4.0)
def _create_cohere(self, model: str) -> TokenCounter:
"""Create Cohere tokenizer.
Cohere has its own tokenizer, we use estimation.
"""
return EstimatingTokenCounter(chars_per_token=4.0)
def _create_estimation(self, model: str) -> TokenCounter:
"""Create estimation-based tokenizer."""
return EstimatingTokenCounter()
# Convenience functions
def get_tokenizer(
model: str,
backend: str | None = None,
fallback: bool = True,
) -> TokenCounter:
"""Get tokenizer for a model.
This is the main entry point for getting tokenizers.
Args:
model: Model name (e.g., 'gpt-4o', 'claude-3-sonnet').
backend: Force specific backend ('tiktoken', 'huggingface', etc.).
fallback: If True, fall back to estimation on errors.
Returns:
TokenCounter instance.
Example:
tokenizer = get_tokenizer("gpt-4o")
tokens = tokenizer.count_text("Hello, world!")
"""
return TokenizerRegistry.get(model, backend, fallback)
def register_tokenizer(
model: str,
tokenizer: TokenCounter | None = None,
factory: Callable[[str], TokenCounter] | None = None,
) -> None:
"""Register a custom tokenizer for a model.
Args:
model: Model name.
tokenizer: Tokenizer instance.
factory: Factory function.
Example:
# Register instance
register_tokenizer("my-model", MyTokenizer())
# Register factory
register_tokenizer("my-model", factory=lambda m: MyTokenizer(m))
"""
TokenizerRegistry.register(model, tokenizer, factory)
def list_supported_models() -> dict[str, str]:
"""List models with known tokenizer mappings.
Returns:
Dict mapping model pattern to backend.
"""
return {pattern: backend for pattern, backend in MODEL_PATTERNS}

View file

@ -0,0 +1,247 @@
"""Tiktoken-based token counter for OpenAI models.
Tiktoken is OpenAI's fast BPE tokenizer used by GPT models.
It supports multiple encodings:
- cl100k_base: GPT-4, GPT-3.5-turbo, text-embedding-ada-002
- o200k_base: GPT-4o, GPT-4o-mini
- p50k_base: Codex models, text-davinci-002/003
- r50k_base: GPT-3 models (davinci, curie, etc.)
"""
from __future__ import annotations
from functools import lru_cache
from typing import Any
from .base import BaseTokenizer
# Model to encoding mapping
MODEL_TO_ENCODING = {
# GPT-4o family (o200k_base)
"gpt-4o": "o200k_base",
"gpt-4o-mini": "o200k_base",
"gpt-4o-2024-05-13": "o200k_base",
"gpt-4o-2024-08-06": "o200k_base",
"gpt-4o-2024-11-20": "o200k_base",
"gpt-4o-mini-2024-07-18": "o200k_base",
# o1 reasoning models (o200k_base)
"o1": "o200k_base",
"o1-mini": "o200k_base",
"o1-preview": "o200k_base",
"o3-mini": "o200k_base",
# GPT-4 family (cl100k_base)
"gpt-4": "cl100k_base",
"gpt-4-turbo": "cl100k_base",
"gpt-4-turbo-preview": "cl100k_base",
"gpt-4-0314": "cl100k_base",
"gpt-4-0613": "cl100k_base",
"gpt-4-32k": "cl100k_base",
"gpt-4-32k-0314": "cl100k_base",
"gpt-4-32k-0613": "cl100k_base",
"gpt-4-1106-preview": "cl100k_base",
"gpt-4-0125-preview": "cl100k_base",
"gpt-4-turbo-2024-04-09": "cl100k_base",
# GPT-3.5 family (cl100k_base)
"gpt-3.5-turbo": "cl100k_base",
"gpt-3.5-turbo-0301": "cl100k_base",
"gpt-3.5-turbo-0613": "cl100k_base",
"gpt-3.5-turbo-1106": "cl100k_base",
"gpt-3.5-turbo-0125": "cl100k_base",
"gpt-3.5-turbo-16k": "cl100k_base",
"gpt-3.5-turbo-16k-0613": "cl100k_base",
"gpt-3.5-turbo-instruct": "cl100k_base",
# Embeddings (cl100k_base)
"text-embedding-ada-002": "cl100k_base",
"text-embedding-3-small": "cl100k_base",
"text-embedding-3-large": "cl100k_base",
# Codex (p50k_base)
"code-davinci-002": "p50k_base",
"code-davinci-001": "p50k_base",
"code-cushman-002": "p50k_base",
"code-cushman-001": "p50k_base",
# Legacy GPT-3 (r50k_base)
"text-davinci-003": "p50k_base",
"text-davinci-002": "p50k_base",
"text-davinci-001": "r50k_base",
"text-curie-001": "r50k_base",
"text-babbage-001": "r50k_base",
"text-ada-001": "r50k_base",
"davinci": "r50k_base",
"curie": "r50k_base",
"babbage": "r50k_base",
"ada": "r50k_base",
}
# Default encoding for unknown models
DEFAULT_ENCODING = "cl100k_base"
@lru_cache(maxsize=8)
def _get_encoding(encoding_name: str):
"""Get tiktoken encoding, cached for performance."""
import tiktoken
return tiktoken.get_encoding(encoding_name)
def get_encoding_for_model(model: str) -> str:
"""Get the tiktoken encoding name for a model.
Args:
model: Model name (e.g., 'gpt-4o', 'gpt-3.5-turbo').
Returns:
Encoding name (e.g., 'o200k_base', 'cl100k_base').
"""
# Direct lookup
if model in MODEL_TO_ENCODING:
return MODEL_TO_ENCODING[model]
# Try prefix matching for versioned models
for prefix in ["gpt-4o", "gpt-4-turbo", "gpt-4", "gpt-3.5", "o1", "o3"]:
if model.startswith(prefix):
# Find any model with this prefix
for known_model, encoding in MODEL_TO_ENCODING.items():
if known_model.startswith(prefix):
return encoding
return DEFAULT_ENCODING
class TiktokenCounter(BaseTokenizer):
"""Token counter using tiktoken (OpenAI's tokenizer).
This is the most accurate tokenizer for OpenAI models and provides
a good approximation for many other models that use similar BPE
tokenization.
Example:
counter = TiktokenCounter("gpt-4o")
tokens = counter.count_text("Hello, world!")
print(f"Token count: {tokens}")
"""
# OpenAI-specific message overhead
MESSAGE_OVERHEAD = 3
REPLY_OVERHEAD = 3
def __init__(self, model: str = "gpt-4o"):
"""Initialize tiktoken counter.
Args:
model: Model name to determine encoding.
Defaults to 'gpt-4o' (o200k_base encoding).
"""
self.model = model
self.encoding_name = get_encoding_for_model(model)
self._encoding = None # Lazy load
@property
def encoding(self):
"""Lazy-load the encoding."""
if self._encoding is None:
self._encoding = _get_encoding(self.encoding_name)
return self._encoding
def count_text(self, text: str) -> int:
"""Count tokens in text using tiktoken.
Args:
text: Text to tokenize.
Returns:
Number of tokens.
"""
if not text:
return 0
return len(self.encoding.encode(text))
def count_messages(self, messages: list[dict[str, Any]]) -> int:
"""Count tokens in messages using OpenAI's exact formula.
This matches OpenAI's token counting for chat completions.
Args:
messages: List of chat messages.
Returns:
Total token count.
"""
total = 0
for message in messages:
# Every message has overhead for role and formatting
total += self.MESSAGE_OVERHEAD
for key, value in message.items():
if value is None:
continue
if key == "content":
if isinstance(value, str):
total += self.count_text(value)
elif isinstance(value, list):
# Multi-part content
for part in value:
if isinstance(part, dict):
if part.get("type") == "text":
total += self.count_text(part.get("text", ""))
elif part.get("type") == "image_url":
# Image tokens vary by detail level
detail = part.get("image_url", {}).get("detail", "auto")
if detail == "low":
total += 85
else:
total += 170 # Base for high detail
else:
total += self.count_text(str(part))
elif isinstance(part, str):
total += self.count_text(part)
elif key == "role":
total += self.count_text(value)
elif key == "name":
total += self.count_text(value)
total += 1 # Name adds 1 token
elif key == "tool_calls":
for tool_call in value:
total += 3 # Tool call overhead
if "function" in tool_call:
func = tool_call["function"]
total += self.count_text(func.get("name", ""))
total += self.count_text(func.get("arguments", ""))
if "id" in tool_call:
total += self.count_text(tool_call["id"])
elif key == "tool_call_id":
total += self.count_text(value)
elif key == "function_call":
total += self.count_text(value.get("name", ""))
total += self.count_text(value.get("arguments", ""))
# Every reply is primed with assistant
total += self.REPLY_OVERHEAD
return total
def encode(self, text: str) -> list[int]:
"""Encode text to token IDs.
Args:
text: Text to encode.
Returns:
List of token IDs.
"""
return self.encoding.encode(text)
def decode(self, tokens: list[int]) -> str:
"""Decode token IDs to text.
Args:
tokens: List of token IDs.
Returns:
Decoded text.
"""
return self.encoding.decode(tokens)
def __repr__(self) -> str:
return f"TiktokenCounter(model={self.model!r}, encoding={self.encoding_name!r})"

View file

@ -2,14 +2,13 @@
from __future__ import annotations
from typing import Any, TYPE_CHECKING
from typing import TYPE_CHECKING, Any
from ..config import (
CacheAlignerConfig,
DiffArtifact,
HeadroomConfig,
RollingWindowConfig,
SmartCrusherConfig,
ToolCrusherConfig,
TransformDiff,
TransformResult,

View file

@ -34,11 +34,10 @@ import statistics
from collections import Counter
from dataclasses import dataclass, field
from enum import Enum
from typing import TYPE_CHECKING, Any
from typing import Any
from ..config import RelevanceScorerConfig, TransformResult
from ..relevance import BM25Scorer, RelevanceScorer, create_scorer
from ..relevance import RelevanceScorer, create_scorer
# Legacy patterns for backwards compatibility (extract_query_anchors)
_UUID_PATTERN = re.compile(

View file

@ -2,7 +2,6 @@
from __future__ import annotations
import json
from typing import Any
from ..config import ToolCrusherConfig, TransformResult

View file

@ -4,39 +4,67 @@ build-backend = "hatchling.build"
[project]
name = "headroom"
version = "0.1.0"
description = "A safe, deterministic Context Budget Controller for LLM APIs"
version = "0.2.0"
description = "The Context Optimization Layer for LLM Applications - Cut costs by 50-90%"
readme = "README.md"
license = "MIT"
license = "Apache-2.0"
requires-python = ">=3.10"
authors = [
{ name = "Headroom Team" }
{ name = "Headroom Contributors" }
]
maintainers = [
{ name = "Headroom Contributors" }
]
keywords = [
"llm",
"openai",
"anthropic",
"claude",
"gpt",
"context",
"token",
"optimization",
"compression",
"caching",
"proxy",
"ai",
"machine-learning",
]
classifiers = [
"Development Status :: 3 - Alpha",
"Development Status :: 4 - Beta",
"Intended Audience :: Developers",
"License :: OSI Approved :: MIT License",
"License :: OSI Approved :: Apache Software License",
"Operating System :: OS Independent",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11",
"Programming Language :: Python :: 3.12",
"Topic :: Scientific/Engineering :: Artificial Intelligence",
"Topic :: Software Development :: Libraries :: Python Modules",
"Typing :: Typed",
]
dependencies = [
"tiktoken>=0.5.0",
"pydantic>=2.0.0",
"jinja2>=3.0.0",
]
[project.optional-dependencies]
# Semantic relevance scoring with embeddings
relevance = [
"sentence-transformers>=2.2.0",
"numpy>=1.24.0",
]
# Proxy server
proxy = [
"fastapi>=0.100.0",
"uvicorn>=0.23.0",
"httpx>=0.24.0",
]
# Report generation
reports = [
"jinja2>=3.0.0",
]
# Development dependencies
dev = [
"pytest>=7.0.0",
"pytest-cov>=4.0.0",
@ -44,16 +72,36 @@ dev = [
"ruff>=0.1.0",
"mypy>=1.0.0",
"openai>=1.0.0",
"anthropic>=0.18.0",
]
# All optional dependencies
all = [
"headroom[relevance,proxy,reports]",
]
[project.scripts]
headroom = "headroom.cli:main"
[project.urls]
Homepage = "https://github.com/headroom-sdk/headroom"
Documentation = "https://github.com/headroom-sdk/headroom#readme"
Documentation = "https://headroom.dev/docs"
Repository = "https://github.com/headroom-sdk/headroom"
Issues = "https://github.com/headroom-sdk/headroom/issues"
Changelog = "https://github.com/headroom-sdk/headroom/blob/main/CHANGELOG.md"
[tool.hatch.build.targets.wheel]
packages = ["headroom"]
[tool.hatch.build.targets.sdist]
include = [
"/headroom",
"/tests",
"/LICENSE",
"/NOTICE",
"/README.md",
"/CHANGELOG.md",
]
[tool.ruff]
target-version = "py310"
line-length = 100
@ -71,19 +119,43 @@ select = [
ignore = [
"E501", # line too long (handled by formatter)
"B008", # do not perform function calls in argument defaults
"B905", # zip without strict parameter
]
[tool.ruff.lint.isort]
known-first-party = ["headroom"]
[tool.ruff.format]
quote-style = "double"
indent-style = "space"
[tool.mypy]
python_version = "3.10"
warn_return_any = true
warn_unused_configs = true
disallow_untyped_defs = true
ignore_missing_imports = true
[tool.pytest.ini_options]
testpaths = ["tests"]
python_files = ["test_*.py"]
python_functions = ["test_*"]
addopts = "-v --tb=short"
asyncio_mode = "auto"
[tool.coverage.run]
source = ["headroom"]
branch = true
omit = [
"headroom/cli.py",
"*/tests/*",
]
[tool.coverage.report]
exclude_lines = [
"pragma: no cover",
"def __repr__",
"raise NotImplementedError",
"if TYPE_CHECKING:",
"if __name__ == .__main__.:",
]

260
tests/test_models.py Normal file
View file

@ -0,0 +1,260 @@
"""Tests for the model registry and capabilities database."""
from __future__ import annotations
import pytest
from datetime import date
from headroom.models import (
ModelRegistry,
ModelInfo,
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,
input_cost_per_1m=1.5,
output_cost_per_1m=3.0,
)
assert info.context_window == 32000
assert info.max_output_tokens == 8192
assert info.supports_tools is False
assert info.supports_vision is True
assert info.input_cost_per_1m == 1.5
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_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."""
cost = ModelRegistry.estimate_cost(
model="gpt-4o",
input_tokens=1000000,
output_tokens=0,
cached_tokens=500000, # Half cached
)
assert cost is not None
# 500K regular at $2.50/1M + 500K cached at $1.25/1M
# = $1.25 + $0.625 = $1.875
assert abs(cost - 1.875) < 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
assert info.input_cost_per_1m == 2.50
assert info.output_cost_per_1m == 10.00
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
assert info.cached_input_cost_per_1m == 0.30 # 90% cache discount
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

View file

@ -0,0 +1,124 @@
"""Tests for Cohere provider."""
from __future__ import annotations
import pytest
from headroom.providers import CohereProvider
class TestCohereProvider:
"""Tests for CohereProvider."""
@pytest.fixture
def provider(self):
"""Create Cohere provider without client (estimation mode)."""
return CohereProvider()
def test_name(self, provider):
"""Test provider name."""
assert provider.name == "cohere"
def test_supports_command_models(self, provider):
"""Test support for Command models."""
assert provider.supports_model("command-r-plus") is True
assert provider.supports_model("command-r") is True
assert provider.supports_model("command-a") is True
assert provider.supports_model("command") is True
def test_not_supports_other_models(self, provider):
"""Test non-support for other models."""
assert provider.supports_model("gpt-4o") is False
assert provider.supports_model("claude-3") is False
assert provider.supports_model("gemini-2.0") is False
def test_get_token_counter(self, provider):
"""Test getting token counter."""
counter = provider.get_token_counter("command-r-plus")
assert counter is not None
count = counter.count_text("Hello, world!")
assert count > 0
def test_get_context_limit_command_a(self, provider):
"""Test context limit for Command A (256K)."""
limit = provider.get_context_limit("command-a")
assert limit == 256000
def test_get_context_limit_command_r_plus(self, provider):
"""Test context limit for Command R+."""
limit = provider.get_context_limit("command-r-plus")
assert limit == 128000
def test_get_context_limit_command_r(self, provider):
"""Test context limit for Command R."""
limit = provider.get_context_limit("command-r")
assert limit == 128000
def test_get_context_limit_legacy_command(self, provider):
"""Test context limit for legacy Command."""
limit = provider.get_context_limit("command")
assert limit == 4096
def test_estimate_cost_command_r_plus(self, provider):
"""Test cost estimation for Command R+."""
cost = provider.estimate_cost(
input_tokens=1000000,
output_tokens=500000,
model="command-r-plus",
)
assert cost is not None
# 1M input * $2.50/1M + 0.5M output * $10.00/1M = $2.50 + $5.00 = $7.50
assert abs(cost - 7.50) < 0.01
def test_estimate_cost_command_r(self, provider):
"""Test cost estimation for Command R."""
cost = provider.estimate_cost(
input_tokens=1000000,
output_tokens=500000,
model="command-r",
)
assert cost is not None
# 1M input * $0.15/1M + 0.5M output * $0.60/1M = $0.15 + $0.30 = $0.45
assert abs(cost - 0.45) < 0.01
def test_estimate_cost_unknown_model(self, provider):
"""Test cost estimation returns None for unknown model."""
cost = provider.estimate_cost(
input_tokens=1000,
output_tokens=500,
model="unknown-model",
)
assert cost is None
class TestCohereTokenCounter:
"""Tests for CohereTokenCounter."""
@pytest.fixture
def counter(self):
"""Create token counter without client."""
provider = CohereProvider()
return provider.get_token_counter("command-r-plus")
def test_count_text_empty(self, counter):
"""Test counting empty text."""
assert counter.count_text("") == 0
def test_count_text_simple(self, counter):
"""Test counting simple text."""
count = counter.count_text("Hello, world!")
assert count > 0
assert count < 20 # Should be a few tokens
def test_count_messages(self, counter):
"""Test counting messages."""
messages = [
{"role": "user", "content": "Hello!"},
{"role": "assistant", "content": "Hi there!"},
]
count = counter.count_messages(messages)
assert count > 0
def test_count_messages_empty(self, counter):
"""Test counting empty messages."""
assert counter.count_messages([]) == 0

View file

@ -0,0 +1,293 @@
"""Tests for universal provider support.
Tests OpenAICompatibleProvider, GoogleProvider, and LiteLLMProvider.
"""
from __future__ import annotations
import pytest
def _transformers_available() -> bool:
"""Check if transformers is available."""
try:
import transformers # noqa: F401
return True
except ImportError:
return False
from headroom.providers import (
OpenAICompatibleProvider,
ModelCapabilities,
GoogleProvider,
create_ollama_provider,
create_together_provider,
create_groq_provider,
create_vllm_provider,
create_lmstudio_provider,
is_litellm_available,
)
class TestOpenAICompatibleProvider:
"""Tests for OpenAICompatibleProvider."""
def test_init_default(self):
"""Test initialization with defaults."""
provider = OpenAICompatibleProvider()
assert provider.name == "openai_compatible"
assert provider.base_url is None
def test_init_with_config(self):
"""Test initialization with configuration."""
provider = OpenAICompatibleProvider(
name="custom",
base_url="http://localhost:8080/v1",
api_key="test-key",
)
assert provider.name == "custom"
assert provider.base_url == "http://localhost:8080/v1"
assert provider.api_key == "test-key"
def test_supports_any_model(self):
"""Test that provider supports any model."""
provider = OpenAICompatibleProvider()
assert provider.supports_model("any-model") is True
assert provider.supports_model("llama-3") is True
assert provider.supports_model("custom-finetuned") is True
@pytest.mark.skipif(
not _transformers_available(),
reason="transformers not installed - needed for HuggingFace tokenizer"
)
def test_get_token_counter(self):
"""Test getting token counter."""
provider = OpenAICompatibleProvider()
counter = provider.get_token_counter("llama-3-8b")
assert counter is not None
# Should be able to count tokens
count = counter.count_text("Hello, world!")
assert count > 0
def test_get_context_limit_known_model(self):
"""Test context limit for known models."""
provider = OpenAICompatibleProvider()
# Llama 3.1 has 128K context
limit = provider.get_context_limit("llama-3.1-8b")
assert limit == 128000
def test_get_context_limit_unknown_model(self):
"""Test context limit for unknown models (defaults to 128K)."""
provider = OpenAICompatibleProvider()
limit = provider.get_context_limit("unknown-model")
assert limit == 128000
def test_register_model(self):
"""Test registering a custom model."""
provider = OpenAICompatibleProvider()
provider.register_model(
"my-model",
context_window=64000,
max_output_tokens=8192,
input_cost_per_1m=1.0,
output_cost_per_1m=2.0,
)
assert provider.get_context_limit("my-model") == 64000
def test_estimate_cost_registered_model(self):
"""Test cost estimation for registered model."""
provider = OpenAICompatibleProvider()
provider.register_model(
"priced-model",
input_cost_per_1m=1.0,
output_cost_per_1m=2.0,
)
cost = provider.estimate_cost(
input_tokens=1000000,
output_tokens=500000,
model="priced-model",
)
assert cost == 2.0 # 1.0 + 1.0
def test_estimate_cost_unknown_model(self):
"""Test cost estimation returns None for unknown model."""
provider = OpenAICompatibleProvider()
cost = provider.estimate_cost(
input_tokens=1000,
output_tokens=500,
model="unknown-model",
)
assert cost is None
class TestModelCapabilities:
"""Tests for ModelCapabilities dataclass."""
def test_default_values(self):
"""Test default capability values."""
caps = ModelCapabilities(model="test-model")
assert caps.context_window == 128000
assert caps.max_output_tokens == 4096
assert caps.supports_tools is True
assert caps.supports_vision is False
assert caps.supports_streaming is True
def test_custom_values(self):
"""Test custom capability values."""
caps = ModelCapabilities(
model="custom-model",
context_window=32000,
max_output_tokens=16384,
supports_tools=False,
supports_vision=True,
input_cost_per_1m=0.5,
output_cost_per_1m=1.5,
)
assert caps.context_window == 32000
assert caps.max_output_tokens == 16384
assert caps.supports_tools is False
assert caps.supports_vision is True
assert caps.input_cost_per_1m == 0.5
assert caps.output_cost_per_1m == 1.5
class TestGoogleProvider:
"""Tests for GoogleProvider."""
@pytest.fixture
def provider(self):
"""Create Google provider."""
return GoogleProvider()
def test_name(self, provider):
"""Test provider name."""
assert provider.name == "google"
def test_supports_gemini_models(self, provider):
"""Test support for Gemini models."""
assert provider.supports_model("gemini-2.0-flash") is True
assert provider.supports_model("gemini-1.5-pro") is True
assert provider.supports_model("gemini-1.5-flash") is True
def test_not_supports_other_models(self, provider):
"""Test non-support for other models."""
assert provider.supports_model("gpt-4o") is False
assert provider.supports_model("claude-3") is False
def test_get_token_counter(self, provider):
"""Test getting token counter."""
counter = provider.get_token_counter("gemini-2.0-flash")
assert counter is not None
count = counter.count_text("Hello, world!")
assert count > 0
def test_get_context_limit_gemini_2(self, provider):
"""Test context limit for Gemini 2.0."""
limit = provider.get_context_limit("gemini-2.0-flash")
assert limit == 1000000 # 1M tokens
def test_get_context_limit_gemini_1_5_pro(self, provider):
"""Test context limit for Gemini 1.5 Pro (2M!)."""
limit = provider.get_context_limit("gemini-1.5-pro")
assert limit == 2000000 # 2M tokens!
def test_estimate_cost(self, provider):
"""Test cost estimation."""
cost = provider.estimate_cost(
input_tokens=1000000,
output_tokens=500000,
model="gemini-2.0-flash",
)
assert cost is not None
# 1M input * $0.10 + 0.5M output * $0.40 = $0.10 + $0.20 = $0.30
assert abs(cost - 0.30) < 0.01
def test_openai_compatible_url(self):
"""Test OpenAI-compatible URL."""
url = GoogleProvider.get_openai_compatible_url("test-key")
assert "generativelanguage.googleapis.com" in url
class TestProviderFactoryFunctions:
"""Tests for provider factory functions."""
def test_create_ollama_provider(self):
"""Test creating Ollama provider."""
provider = create_ollama_provider()
assert provider.name == "ollama"
assert provider.base_url == "http://localhost:11434/v1"
def test_create_ollama_provider_custom_url(self):
"""Test creating Ollama provider with custom URL."""
provider = create_ollama_provider("http://192.168.1.100:11434/v1")
assert provider.base_url == "http://192.168.1.100:11434/v1"
def test_create_together_provider(self):
"""Test creating Together provider."""
provider = create_together_provider()
assert provider.name == "together"
assert "together.xyz" in provider.base_url
def test_create_groq_provider(self):
"""Test creating Groq provider."""
provider = create_groq_provider()
assert provider.name == "groq"
assert "groq.com" in provider.base_url
def test_create_vllm_provider(self):
"""Test creating vLLM provider."""
provider = create_vllm_provider("http://localhost:8000/v1")
assert provider.name == "vllm"
assert provider.base_url == "http://localhost:8000/v1"
def test_create_lmstudio_provider(self):
"""Test creating LM Studio provider."""
provider = create_lmstudio_provider()
assert provider.name == "lmstudio"
assert provider.base_url == "http://localhost:1234/v1"
class TestLiteLLMProvider:
"""Tests for LiteLLM provider."""
def test_is_litellm_available(self):
"""Test checking LiteLLM availability."""
result = is_litellm_available()
assert isinstance(result, bool)
@pytest.mark.skipif(
not is_litellm_available(),
reason="LiteLLM not installed",
)
def test_create_litellm_provider(self):
"""Test creating LiteLLM provider."""
from headroom.providers import create_litellm_provider
provider = create_litellm_provider()
assert provider.name == "litellm"
@pytest.mark.skipif(
not is_litellm_available(),
reason="LiteLLM not installed",
)
def test_litellm_supports_any_model(self):
"""Test LiteLLM supports any model."""
from headroom.providers import create_litellm_provider
provider = create_litellm_provider()
assert provider.supports_model("gpt-4o") is True
assert provider.supports_model("claude-3-sonnet") is True
assert provider.supports_model("any-model") is True
@pytest.mark.skipif(
not is_litellm_available(),
reason="LiteLLM not installed",
)
def test_litellm_list_providers(self):
"""Test listing LiteLLM providers."""
from headroom.providers import LiteLLMProvider
providers = LiteLLMProvider.list_supported_providers()
assert "openai" in providers
assert "anthropic" in providers
assert "ollama" in providers

443
tests/test_tokenizers.py Normal file
View file

@ -0,0 +1,443 @@
"""Tests for the pluggable tokenizer system."""
from __future__ import annotations
import pytest
from headroom.tokenizers import (
TokenizerRegistry,
get_tokenizer,
register_tokenizer,
list_supported_models,
TiktokenCounter,
EstimatingTokenCounter,
CharacterCounter,
TokenCounter,
BaseTokenizer,
is_mistral_tokenizer_available,
get_mistral_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)
@pytest.mark.skipif(
not is_mistral_tokenizer_available(),
reason="mistral-common not installed",
)
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")
@pytest.mark.skipif(
not is_mistral_tokenizer_available(),
reason="mistral-common not installed",
)
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"
@pytest.mark.skipif(
not is_mistral_tokenizer_available(),
reason="mistral-common not installed",
)
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"
@pytest.mark.skipif(
not is_mistral_tokenizer_available(),
reason="mistral-common not installed",
)
def test_count_text_empty(self):
"""Test counting empty text."""
MistralTokenizer = get_mistral_tokenizer()
counter = MistralTokenizer()
assert counter.count_text("") == 0
@pytest.mark.skipif(
not is_mistral_tokenizer_available(),
reason="mistral-common not installed",
)
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
@pytest.mark.skipif(
not is_mistral_tokenizer_available(),
reason="mistral-common not installed",
)
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
@pytest.mark.skipif(
not is_mistral_tokenizer_available(),
reason="mistral-common not installed",
)
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
@pytest.mark.skipif(
not is_mistral_tokenizer_available(),
reason="mistral-common not installed",
)
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
@pytest.mark.skipif(
not is_mistral_tokenizer_available(),
reason="mistral-common not installed",
)
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
@pytest.mark.skipif(
not is_mistral_tokenizer_available(),
reason="mistral-common not installed",
)
def test_implements_protocol(self):
"""Test MistralTokenizer implements TokenCounter protocol."""
MistralTokenizer = get_mistral_tokenizer()
counter = MistralTokenizer()
assert isinstance(counter, TokenCounter)
@pytest.mark.skipif(
not is_mistral_tokenizer_available(),
reason="mistral-common not installed",
)
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)
@pytest.mark.skipif(
not is_mistral_tokenizer_available(),
reason="mistral-common not installed",
)
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)
@pytest.mark.skipif(
not is_mistral_tokenizer_available(),
reason="mistral-common not installed",
)
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)
@pytest.mark.skipif(
not is_mistral_tokenizer_available(),
reason="mistral-common not installed",
)
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)