mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-27 14:17:10 -04:00
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:
parent
9c7d4512d6
commit
175746cc26
67 changed files with 9184 additions and 291 deletions
7
.github/FUNDING.yml
vendored
Normal file
7
.github/FUNDING.yml
vendored
Normal 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
53
.github/ISSUE_TEMPLATE/bug_report.md
vendored
Normal 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
8
.github/ISSUE_TEMPLATE/config.yml
vendored
Normal 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
|
||||
44
.github/ISSUE_TEMPLATE/feature_request.md
vendored
Normal file
44
.github/ISSUE_TEMPLATE/feature_request.md
vendored
Normal 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
56
.github/PULL_REQUEST_TEMPLATE.md
vendored
Normal 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
108
.github/workflows/ci.yml
vendored
Normal 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
31
.github/workflows/publish.yml
vendored
Normal 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
68
.gitignore
vendored
|
|
@ -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
120
CHANGELOG.md
Normal 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
133
CODE_OF_CONDUCT.md
Normal 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
209
CONTRIBUTING.md
Normal 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
190
LICENSE
Normal 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
43
NOTICE
Normal 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
408
README.md
|
|
@ -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
65
SECURITY.md
Normal 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!
|
||||
|
|
@ -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
28
docs/README.md
Normal 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
346
docs/api.md
Normal 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
109
docs/getting-started.md
Normal 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
173
docs/proxy.md
Normal 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
198
docs/transforms.md
Normal 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
133
examples/README.md
Normal 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-...'
|
||||
```
|
||||
|
|
@ -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
185
headroom/cli.py
Normal 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())
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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__ = [
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
39
headroom/models/__init__.py
Normal file
39
headroom/models/__init__.py
Normal 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
749
headroom/models/registry.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
313
headroom/providers/cohere.py
Normal file
313
headroom/providers/cohere.py
Normal 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
|
||||
372
headroom/providers/google.py
Normal file
372
headroom/providers/google.py
Normal 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
|
||||
293
headroom/providers/litellm.py
Normal file
293
headroom/providers/litellm.py
Normal 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()
|
||||
|
|
@ -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
|
||||
|
|
|
|||
521
headroom/providers/openai_compatible.py
Normal file
521
headroom/providers/openai_compatible.py
Normal 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,
|
||||
)
|
||||
19
headroom/proxy/__init__.py
Normal file
19
headroom/proxy/__init__.py
Normal 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
1399
headroom/proxy/server.py
Normal file
File diff suppressed because it is too large
Load diff
0
headroom/py.typed
Normal file
0
headroom/py.typed
Normal 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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
72
headroom/tokenizers/__init__.py
Normal file
72
headroom/tokenizers/__init__.py
Normal 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
203
headroom/tokenizers/base.py
Normal 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"
|
||||
)
|
||||
199
headroom/tokenizers/estimator.py
Normal file
199
headroom/tokenizers/estimator.py
Normal 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})"
|
||||
316
headroom/tokenizers/huggingface.py
Normal file
316
headroom/tokenizers/huggingface.py
Normal 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})"
|
||||
244
headroom/tokenizers/mistral.py
Normal file
244
headroom/tokenizers/mistral.py
Normal 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})"
|
||||
398
headroom/tokenizers/registry.py
Normal file
398
headroom/tokenizers/registry.py
Normal 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}
|
||||
247
headroom/tokenizers/tiktoken_counter.py
Normal file
247
headroom/tokenizers/tiktoken_counter.py
Normal 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})"
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
from ..config import ToolCrusherConfig, TransformResult
|
||||
|
|
|
|||
|
|
@ -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
260
tests/test_models.py
Normal 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
|
||||
124
tests/test_providers/test_cohere.py
Normal file
124
tests/test_providers/test_cohere.py
Normal 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
|
||||
293
tests/test_providers/test_universal.py
Normal file
293
tests/test_providers/test_universal.py
Normal 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
443
tests/test_tokenizers.py
Normal 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)
|
||||
Loading…
Add table
Add a link
Reference in a new issue