Spaces:
Build error
Add LLMLingua-2 opt-in support to proxy server
Browse filesIntegrate Microsoft's LLMLingua-2 ML-based compression as an opt-in
feature for the proxy server, with excellent developer experience.
Features:
- New CLI flags: --llmlingua, --llmlingua-device, --llmlingua-rate
- ProxyConfig options: llmlingua_enabled, llmlingua_device, llmlingua_target_rate
- Smart startup hints when llmlingua is available but not enabled
- Helpful error messages when enabled but not installed
- LLMLinguaCompressor inserted before RollingWindow in pipeline
Why opt-in:
- Heavy dependencies (~2GB torch, transformers)
- 10-30s cold start for model loading
- ~1GB RAM when loaded
- Default proxy stays lightweight (<5ms overhead)
Tests:
- 26 new tests in test_proxy_llmlingua.py covering config, setup,
banner status, CLI args, DevEx messages, and edge cases
Documentation:
- Updated README.md with proxy integration section
- Updated docs/proxy.md with LLMLingua CLI options
- Updated docs/transforms.md with LLMLinguaCompressor reference
- Updated docs/ARCHITECTURE.md with pipeline and file structure
- Updated CHANGELOG.md with new feature
- CHANGELOG.md +9 -0
- README.md +137 -0
- docs/ARCHITECTURE.md +43 -8
- docs/proxy.md +59 -0
- docs/transforms.md +98 -3
- headroom/proxy/server.py +103 -1
- headroom/transforms/__init__.py +30 -0
- headroom/transforms/llmlingua_compressor.py +633 -0
- pyproject.toml +7 -1
- tests/test_proxy_llmlingua.py +458 -0
- tests/test_transforms/test_llmlingua_compressor.py +941 -0
|
@@ -10,6 +10,15 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|
| 10 |
### Added
|
| 11 |
- Production-ready proxy server with caching, rate limiting, and metrics
|
| 12 |
- CLI command `headroom proxy` to start the proxy server
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 13 |
|
| 14 |
## [0.2.0] - 2025-01-07
|
| 15 |
|
|
|
|
| 10 |
### Added
|
| 11 |
- Production-ready proxy server with caching, rate limiting, and metrics
|
| 12 |
- CLI command `headroom proxy` to start the proxy server
|
| 13 |
+
- **LLMLingua-2 Integration** (opt-in ML-based compression)
|
| 14 |
+
- `LLMLinguaCompressor` transform using Microsoft's LLMLingua-2 model
|
| 15 |
+
- Content-aware compression rates (code: 0.4, JSON: 0.35, text: 0.3)
|
| 16 |
+
- Memory management utilities: `unload_llmlingua_model()`, `is_llmlingua_model_loaded()`
|
| 17 |
+
- Proxy integration via `--llmlingua` flag
|
| 18 |
+
- Device selection: `--llmlingua-device` (auto/cuda/cpu/mps)
|
| 19 |
+
- Custom compression rate: `--llmlingua-rate`
|
| 20 |
+
- Helpful startup hints when llmlingua is available but not enabled
|
| 21 |
+
- Install with: `pip install headroom-ai[llmlingua]`
|
| 22 |
|
| 23 |
## [0.2.0] - 2025-01-07
|
| 24 |
|
|
@@ -483,6 +483,143 @@ def compress_tool_output(content: str, context: str = "") -> str:
|
|
| 483 |
|
| 484 |
---
|
| 485 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 486 |
## Metrics & Monitoring
|
| 487 |
|
| 488 |
### Prometheus Metrics (Proxy)
|
|
|
|
| 483 |
|
| 484 |
---
|
| 485 |
|
| 486 |
+
## ML-Based Compression with LLMLingua-2 (Optional)
|
| 487 |
+
|
| 488 |
+
For even more aggressive compression, Headroom integrates with **LLMLingua-2**, Microsoft's BERT-based token classifier trained via GPT-4 distillation. It achieves **up to 20x compression** while preserving semantic meaning.
|
| 489 |
+
|
| 490 |
+
### When to Use LLMLingua-2
|
| 491 |
+
|
| 492 |
+
| Approach | Best For | Compression | Speed |
|
| 493 |
+
|----------|----------|-------------|-------|
|
| 494 |
+
| **SmartCrusher** | JSON tool outputs | 70-90% | ~1ms |
|
| 495 |
+
| **Text Utilities** | Search/logs | 50-90% | ~1ms |
|
| 496 |
+
| **LLMLingua-2** | Any text, max compression | 80-95% | ~50-200ms |
|
| 497 |
+
|
| 498 |
+
LLMLingua-2 is ideal when you need maximum compression and can tolerate slightly higher latency (e.g., compressing large tool outputs before storage, offline processing).
|
| 499 |
+
|
| 500 |
+
### Installation
|
| 501 |
+
|
| 502 |
+
```bash
|
| 503 |
+
# Adds ~2GB of model weights
|
| 504 |
+
pip install "headroom-ai[llmlingua]"
|
| 505 |
+
```
|
| 506 |
+
|
| 507 |
+
### Basic Usage
|
| 508 |
+
|
| 509 |
+
```python
|
| 510 |
+
from headroom.transforms import LLMLinguaCompressor
|
| 511 |
+
|
| 512 |
+
# Create compressor (model loaded lazily on first use)
|
| 513 |
+
compressor = LLMLinguaCompressor()
|
| 514 |
+
|
| 515 |
+
# Compress any text
|
| 516 |
+
long_output = "The function processUserData takes a user object and validates..."
|
| 517 |
+
result = compressor.compress(long_output)
|
| 518 |
+
|
| 519 |
+
print(f"Before: {result.original_tokens} tokens")
|
| 520 |
+
print(f"After: {result.compressed_tokens} tokens")
|
| 521 |
+
print(f"Saved: {result.savings_percentage:.1f}%")
|
| 522 |
+
print(result.compressed)
|
| 523 |
+
```
|
| 524 |
+
|
| 525 |
+
### Content-Aware Compression
|
| 526 |
+
|
| 527 |
+
LLMLingua-2 automatically adjusts compression based on content type:
|
| 528 |
+
|
| 529 |
+
```python
|
| 530 |
+
from headroom.transforms import LLMLinguaCompressor, LLMLinguaConfig
|
| 531 |
+
|
| 532 |
+
# Conservative for code (keep 40% of tokens)
|
| 533 |
+
config = LLMLinguaConfig(
|
| 534 |
+
code_compression_rate=0.4, # More conservative
|
| 535 |
+
json_compression_rate=0.35, # Moderate
|
| 536 |
+
text_compression_rate=0.25, # Aggressive
|
| 537 |
+
)
|
| 538 |
+
|
| 539 |
+
compressor = LLMLinguaCompressor(config)
|
| 540 |
+
|
| 541 |
+
# Auto-detects content type
|
| 542 |
+
code_result = compressor.compress("def calculate(x): return x * 2")
|
| 543 |
+
text_result = compressor.compress("This is a verbose explanation...")
|
| 544 |
+
```
|
| 545 |
+
|
| 546 |
+
### Memory Management
|
| 547 |
+
|
| 548 |
+
The model uses ~1GB RAM. Unload it when done:
|
| 549 |
+
|
| 550 |
+
```python
|
| 551 |
+
from headroom.transforms import (
|
| 552 |
+
LLMLinguaCompressor,
|
| 553 |
+
unload_llmlingua_model,
|
| 554 |
+
is_llmlingua_model_loaded,
|
| 555 |
+
)
|
| 556 |
+
|
| 557 |
+
compressor = LLMLinguaCompressor()
|
| 558 |
+
result = compressor.compress(content) # Model loaded here
|
| 559 |
+
|
| 560 |
+
# Check if loaded
|
| 561 |
+
print(is_llmlingua_model_loaded()) # True
|
| 562 |
+
|
| 563 |
+
# Free memory when done
|
| 564 |
+
unload_llmlingua_model() # Frees ~1GB
|
| 565 |
+
print(is_llmlingua_model_loaded()) # False
|
| 566 |
+
|
| 567 |
+
# Next compression will reload automatically
|
| 568 |
+
```
|
| 569 |
+
|
| 570 |
+
### Use in Pipeline
|
| 571 |
+
|
| 572 |
+
```python
|
| 573 |
+
from headroom.transforms import TransformPipeline, LLMLinguaCompressor, SmartCrusher
|
| 574 |
+
|
| 575 |
+
# Combine with other transforms
|
| 576 |
+
pipeline = TransformPipeline([
|
| 577 |
+
SmartCrusher(), # First: compress JSON
|
| 578 |
+
LLMLinguaCompressor(), # Then: ML compression on remaining text
|
| 579 |
+
])
|
| 580 |
+
|
| 581 |
+
result = pipeline.apply(messages, tokenizer)
|
| 582 |
+
```
|
| 583 |
+
|
| 584 |
+
### Device Configuration
|
| 585 |
+
|
| 586 |
+
```python
|
| 587 |
+
from headroom.transforms import LLMLinguaConfig, LLMLinguaCompressor
|
| 588 |
+
|
| 589 |
+
# Force CPU (slower but works everywhere)
|
| 590 |
+
config = LLMLinguaConfig(device="cpu")
|
| 591 |
+
|
| 592 |
+
# Force GPU (faster but needs CUDA)
|
| 593 |
+
config = LLMLinguaConfig(device="cuda")
|
| 594 |
+
|
| 595 |
+
# Auto-detect (default): uses CUDA > MPS > CPU
|
| 596 |
+
config = LLMLinguaConfig(device="auto")
|
| 597 |
+
|
| 598 |
+
compressor = LLMLinguaCompressor(config)
|
| 599 |
+
```
|
| 600 |
+
|
| 601 |
+
### Proxy Integration (Opt-In)
|
| 602 |
+
|
| 603 |
+
Enable LLMLingua in the proxy server for automatic ML compression of all requests:
|
| 604 |
+
|
| 605 |
+
```bash
|
| 606 |
+
# Enable LLMLingua in proxy (requires: pip install headroom-ai[llmlingua,proxy])
|
| 607 |
+
headroom proxy --llmlingua
|
| 608 |
+
|
| 609 |
+
# With custom settings
|
| 610 |
+
headroom proxy --llmlingua --llmlingua-device cuda --llmlingua-rate 0.4
|
| 611 |
+
|
| 612 |
+
# The proxy shows LLMLingua status at startup:
|
| 613 |
+
# LLMLingua: ENABLED (device=cuda, rate=0.4)
|
| 614 |
+
#
|
| 615 |
+
# If llmlingua is installed but not enabled, you'll see a helpful hint:
|
| 616 |
+
# LLMLingua: available (enable with --llmlingua for ML compression)
|
| 617 |
+
```
|
| 618 |
+
|
| 619 |
+
**Why opt-in?** LLMLingua adds ~2GB dependencies and 10-30s cold start. The default proxy is lightweight (~50MB) with <5ms overhead. Enable LLMLingua when you need maximum compression and can accept the tradeoffs.
|
| 620 |
+
|
| 621 |
+
---
|
| 622 |
+
|
| 623 |
## Metrics & Monitoring
|
| 624 |
|
| 625 |
### Prometheus Metrics (Proxy)
|
|
@@ -193,7 +193,36 @@ analysis = {
|
|
| 193 |
|
| 194 |
---
|
| 195 |
|
| 196 |
-
#### Transform 4:
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 197 |
|
| 198 |
**Problem:** Even after compression, you might exceed the model's context limit.
|
| 199 |
|
|
@@ -282,7 +311,12 @@ def apply(self, messages, ...):
|
|
| 282 |
# - Compresses to 17 points (preserving spike)
|
| 283 |
# - Factors out constant "host" field
|
| 284 |
|
| 285 |
-
# Transform 3:
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 286 |
# - Checks if we're under limit (we are)
|
| 287 |
# - No drops needed
|
| 288 |
|
|
@@ -670,12 +704,13 @@ headroom/
|
|
| 670 |
│ └── anthropic.py # Anthropic-specific
|
| 671 |
│
|
| 672 |
├── transforms/
|
| 673 |
-
│ ├── base.py
|
| 674 |
-
│ ├── pipeline.py
|
| 675 |
-
│ ├── cache_aligner.py
|
| 676 |
-
│ ├── tool_crusher.py
|
| 677 |
-
│ ├── smart_crusher.py
|
| 678 |
-
│
|
|
|
|
| 679 |
│
|
| 680 |
├── cache/ # CCR Architecture
|
| 681 |
│ ├── compression_store.py # Phase 1: Store original content
|
|
|
|
| 193 |
|
| 194 |
---
|
| 195 |
|
| 196 |
+
#### Transform 4: LLMLingua Compressor (Optional)
|
| 197 |
+
|
| 198 |
+
**When to use:** Maximum compression needed and latency is acceptable.
|
| 199 |
+
|
| 200 |
+
```python
|
| 201 |
+
# Opt-in ML-based compression using Microsoft's LLMLingua-2
|
| 202 |
+
# BERT-based token classifier trained via GPT-4 distillation
|
| 203 |
+
|
| 204 |
+
# Before: Long tool output text
|
| 205 |
+
"The function processUserData takes a user object and validates all fields..."
|
| 206 |
+
|
| 207 |
+
# After: Compressed while preserving semantic meaning
|
| 208 |
+
"function processUserData validates user fields..."
|
| 209 |
+
```
|
| 210 |
+
|
| 211 |
+
**Key characteristics:**
|
| 212 |
+
- Uses `microsoft/llmlingua-2-xlm-roberta-large-meetingbank` model
|
| 213 |
+
- Auto-detects content type (code, JSON, text) for optimal compression rates
|
| 214 |
+
- Stores original in CCR for retrieval if needed
|
| 215 |
+
- Adds 50-200ms latency per request
|
| 216 |
+
- Requires ~1GB RAM when loaded
|
| 217 |
+
|
| 218 |
+
**Proxy integration (opt-in):**
|
| 219 |
+
```bash
|
| 220 |
+
headroom proxy --llmlingua --llmlingua-device cuda
|
| 221 |
+
```
|
| 222 |
+
|
| 223 |
+
---
|
| 224 |
+
|
| 225 |
+
#### Transform 5: Rolling Window
|
| 226 |
|
| 227 |
**Problem:** Even after compression, you might exceed the model's context limit.
|
| 228 |
|
|
|
|
| 311 |
# - Compresses to 17 points (preserving spike)
|
| 312 |
# - Factors out constant "host" field
|
| 313 |
|
| 314 |
+
# Transform 3: LLMLingua (if enabled via --llmlingua)
|
| 315 |
+
# - ML-based compression on remaining long text
|
| 316 |
+
# - Auto-detects content type for optimal rate
|
| 317 |
+
# - Stores original in CCR for retrieval
|
| 318 |
+
|
| 319 |
+
# Transform 4: Rolling Window
|
| 320 |
# - Checks if we're under limit (we are)
|
| 321 |
# - No drops needed
|
| 322 |
|
|
|
|
| 704 |
│ └── anthropic.py # Anthropic-specific
|
| 705 |
│
|
| 706 |
├── transforms/
|
| 707 |
+
│ ├── base.py # Transform protocol
|
| 708 |
+
│ ├── pipeline.py # Orchestrates all transforms
|
| 709 |
+
│ ├── cache_aligner.py # Date extraction for caching
|
| 710 |
+
│ ├── tool_crusher.py # Naive compression (disabled)
|
| 711 |
+
│ ├── smart_crusher.py # Statistical compression (default)
|
| 712 |
+
│ ├── rolling_window.py # Token limit enforcement
|
| 713 |
+
│ └── llmlingua_compressor.py # ML-based compression (opt-in)
|
| 714 |
│
|
| 715 |
├── cache/ # CCR Architecture
|
| 716 |
│ ├── compression_store.py # Phase 1: Store original content
|
|
@@ -21,6 +21,8 @@ headroom proxy \
|
|
| 21 |
|
| 22 |
## Command Line Options
|
| 23 |
|
|
|
|
|
|
|
| 24 |
| Option | Default | Description |
|
| 25 |
|--------|---------|-------------|
|
| 26 |
| `--host` | `127.0.0.1` | Host to bind to |
|
|
@@ -31,6 +33,27 @@ headroom proxy \
|
|
| 31 |
| `--log-file` | None | Path to JSONL log file |
|
| 32 |
| `--budget` | None | Daily budget limit in USD |
|
| 33 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 34 |
## API Endpoints
|
| 35 |
|
| 36 |
### Health Check
|
|
@@ -104,6 +127,42 @@ client = OpenAI(
|
|
| 104 |
|
| 105 |
## Features
|
| 106 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 107 |
### Semantic Caching
|
| 108 |
|
| 109 |
The proxy caches responses for repeated queries:
|
|
|
|
| 21 |
|
| 22 |
## Command Line Options
|
| 23 |
|
| 24 |
+
### Core Options
|
| 25 |
+
|
| 26 |
| Option | Default | Description |
|
| 27 |
|--------|---------|-------------|
|
| 28 |
| `--host` | `127.0.0.1` | Host to bind to |
|
|
|
|
| 33 |
| `--log-file` | None | Path to JSONL log file |
|
| 34 |
| `--budget` | None | Daily budget limit in USD |
|
| 35 |
|
| 36 |
+
### LLMLingua Options (ML Compression)
|
| 37 |
+
|
| 38 |
+
| Option | Default | Description |
|
| 39 |
+
|--------|---------|-------------|
|
| 40 |
+
| `--llmlingua` | `false` | Enable LLMLingua-2 ML-based compression |
|
| 41 |
+
| `--llmlingua-device` | `auto` | Device for model: `auto`, `cuda`, `cpu`, `mps` |
|
| 42 |
+
| `--llmlingua-rate` | `0.3` | Target compression rate (0.3 = keep 30% of tokens) |
|
| 43 |
+
|
| 44 |
+
**Note:** LLMLingua requires additional dependencies: `pip install headroom-ai[llmlingua]`
|
| 45 |
+
|
| 46 |
+
```bash
|
| 47 |
+
# Enable LLMLingua with GPU acceleration
|
| 48 |
+
headroom proxy --llmlingua --llmlingua-device cuda
|
| 49 |
+
|
| 50 |
+
# More aggressive compression (keep only 20%)
|
| 51 |
+
headroom proxy --llmlingua --llmlingua-rate 0.2
|
| 52 |
+
|
| 53 |
+
# Conservative compression for code (keep 50%)
|
| 54 |
+
headroom proxy --llmlingua --llmlingua-rate 0.5
|
| 55 |
+
```
|
| 56 |
+
|
| 57 |
## API Endpoints
|
| 58 |
|
| 59 |
### Health Check
|
|
|
|
| 127 |
|
| 128 |
## Features
|
| 129 |
|
| 130 |
+
### LLMLingua ML Compression (Opt-In)
|
| 131 |
+
|
| 132 |
+
When enabled, the proxy uses Microsoft's LLMLingua-2 model for ML-based token compression:
|
| 133 |
+
|
| 134 |
+
```bash
|
| 135 |
+
headroom proxy --llmlingua
|
| 136 |
+
```
|
| 137 |
+
|
| 138 |
+
**How it works:**
|
| 139 |
+
- LLMLinguaCompressor is added to the transform pipeline (before RollingWindow)
|
| 140 |
+
- Automatically detects content type (JSON, code, text) and adjusts compression
|
| 141 |
+
- Stores original content in CCR for retrieval if needed
|
| 142 |
+
|
| 143 |
+
**Startup feedback:**
|
| 144 |
+
|
| 145 |
+
```
|
| 146 |
+
# When enabled and available:
|
| 147 |
+
LLMLingua: ENABLED (device=cuda, rate=0.3)
|
| 148 |
+
|
| 149 |
+
# When installed but not enabled (helpful hint):
|
| 150 |
+
LLMLingua: available (enable with --llmlingua for ML compression)
|
| 151 |
+
|
| 152 |
+
# When enabled but not installed:
|
| 153 |
+
WARNING: LLMLingua requested but not installed. Install with: pip install headroom-ai[llmlingua]
|
| 154 |
+
```
|
| 155 |
+
|
| 156 |
+
**Why opt-in?**
|
| 157 |
+
| Concern | Default Proxy | With LLMLingua |
|
| 158 |
+
|---------|---------------|----------------|
|
| 159 |
+
| Dependencies | ~50MB | +2GB (torch, transformers) |
|
| 160 |
+
| Cold start | <1s | 10-30s (model load) |
|
| 161 |
+
| Memory | ~100MB | +1GB (model in RAM) |
|
| 162 |
+
| Overhead | <5ms | 50-200ms per request |
|
| 163 |
+
|
| 164 |
+
Enable LLMLingua when maximum compression justifies the resource cost.
|
| 165 |
+
|
| 166 |
### Semantic Caching
|
| 167 |
|
| 168 |
The proxy caches responses for repeated queries:
|
|
@@ -162,6 +162,76 @@ config = RollingWindowConfig(
|
|
| 162 |
|
| 163 |
---
|
| 164 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 165 |
## TransformPipeline
|
| 166 |
|
| 167 |
Combine transforms for optimal results.
|
|
@@ -179,11 +249,36 @@ result = pipeline.transform(messages)
|
|
| 179 |
print(f"Saved {result.tokens_saved} tokens")
|
| 180 |
```
|
| 181 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 182 |
### Recommended Order
|
| 183 |
|
| 184 |
-
|
| 185 |
-
|
| 186 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 187 |
|
| 188 |
---
|
| 189 |
|
|
|
|
| 162 |
|
| 163 |
---
|
| 164 |
|
| 165 |
+
## LLMLinguaCompressor (Optional)
|
| 166 |
+
|
| 167 |
+
ML-based compression using Microsoft's LLMLingua-2 model.
|
| 168 |
+
|
| 169 |
+
### When to Use
|
| 170 |
+
|
| 171 |
+
| Transform | Best For | Speed | Compression |
|
| 172 |
+
|-----------|----------|-------|-------------|
|
| 173 |
+
| SmartCrusher | JSON arrays | ~1ms | 70-90% |
|
| 174 |
+
| Text Utilities | Search/logs | ~1ms | 50-90% |
|
| 175 |
+
| **LLMLinguaCompressor** | Any text, max compression | 50-200ms | 80-95% |
|
| 176 |
+
|
| 177 |
+
### Installation
|
| 178 |
+
|
| 179 |
+
```bash
|
| 180 |
+
pip install "headroom-ai[llmlingua]" # Adds ~2GB
|
| 181 |
+
```
|
| 182 |
+
|
| 183 |
+
### Configuration
|
| 184 |
+
|
| 185 |
+
```python
|
| 186 |
+
from headroom.transforms import LLMLinguaCompressor, LLMLinguaConfig
|
| 187 |
+
|
| 188 |
+
config = LLMLinguaConfig(
|
| 189 |
+
device="auto", # auto, cuda, cpu, mps
|
| 190 |
+
target_compression_rate=0.3, # Keep 30% of tokens
|
| 191 |
+
min_tokens_for_compression=100, # Skip small content
|
| 192 |
+
code_compression_rate=0.4, # Conservative for code
|
| 193 |
+
json_compression_rate=0.35, # Moderate for JSON
|
| 194 |
+
text_compression_rate=0.25, # Aggressive for text
|
| 195 |
+
enable_ccr=True, # Store original for retrieval
|
| 196 |
+
)
|
| 197 |
+
|
| 198 |
+
compressor = LLMLinguaCompressor(config)
|
| 199 |
+
```
|
| 200 |
+
|
| 201 |
+
### Content-Aware Rates
|
| 202 |
+
|
| 203 |
+
LLMLinguaCompressor auto-detects content type:
|
| 204 |
+
|
| 205 |
+
| Content Type | Default Rate | Behavior |
|
| 206 |
+
|--------------|--------------|----------|
|
| 207 |
+
| Code | 0.4 | Conservative - preserves syntax |
|
| 208 |
+
| JSON | 0.35 | Moderate - keeps structure |
|
| 209 |
+
| Text | 0.3 | Aggressive - maximum compression |
|
| 210 |
+
|
| 211 |
+
### Memory Management
|
| 212 |
+
|
| 213 |
+
```python
|
| 214 |
+
from headroom.transforms import (
|
| 215 |
+
is_llmlingua_model_loaded,
|
| 216 |
+
unload_llmlingua_model,
|
| 217 |
+
)
|
| 218 |
+
|
| 219 |
+
# Check if model is loaded
|
| 220 |
+
print(is_llmlingua_model_loaded()) # True/False
|
| 221 |
+
|
| 222 |
+
# Free ~1GB RAM when done
|
| 223 |
+
unload_llmlingua_model()
|
| 224 |
+
```
|
| 225 |
+
|
| 226 |
+
### Proxy Integration
|
| 227 |
+
|
| 228 |
+
```bash
|
| 229 |
+
# Enable in proxy
|
| 230 |
+
headroom proxy --llmlingua --llmlingua-device cuda --llmlingua-rate 0.3
|
| 231 |
+
```
|
| 232 |
+
|
| 233 |
+
---
|
| 234 |
+
|
| 235 |
## TransformPipeline
|
| 236 |
|
| 237 |
Combine transforms for optimal results.
|
|
|
|
| 249 |
print(f"Saved {result.tokens_saved} tokens")
|
| 250 |
```
|
| 251 |
|
| 252 |
+
### With LLMLingua (Optional)
|
| 253 |
+
|
| 254 |
+
```python
|
| 255 |
+
from headroom.transforms import (
|
| 256 |
+
TransformPipeline, SmartCrusher, CacheAligner,
|
| 257 |
+
RollingWindow, LLMLinguaCompressor
|
| 258 |
+
)
|
| 259 |
+
|
| 260 |
+
pipeline = TransformPipeline([
|
| 261 |
+
CacheAligner(), # 1. Stabilize prefix
|
| 262 |
+
SmartCrusher(), # 2. Compress JSON arrays
|
| 263 |
+
LLMLinguaCompressor(), # 3. ML compression on remaining text
|
| 264 |
+
RollingWindow(), # 4. Final size constraint (always last)
|
| 265 |
+
])
|
| 266 |
+
```
|
| 267 |
+
|
| 268 |
### Recommended Order
|
| 269 |
|
| 270 |
+
| Order | Transform | Purpose |
|
| 271 |
+
|-------|-----------|---------|
|
| 272 |
+
| 1 | CacheAligner | Stabilize prefix for caching |
|
| 273 |
+
| 2 | SmartCrusher | Compress JSON tool outputs |
|
| 274 |
+
| 3 | LLMLinguaCompressor | ML compression (optional) |
|
| 275 |
+
| 4 | RollingWindow | Enforce token limits (always last) |
|
| 276 |
+
|
| 277 |
+
**Why this order?**
|
| 278 |
+
- CacheAligner first to maximize prefix stability
|
| 279 |
+
- SmartCrusher handles JSON arrays efficiently
|
| 280 |
+
- LLMLingua compresses remaining long text
|
| 281 |
+
- RollingWindow truncates only if still over limit
|
| 282 |
|
| 283 |
---
|
| 284 |
|
|
@@ -59,7 +59,17 @@ from headroom.config import CacheAlignerConfig, RollingWindowConfig, SmartCrushe
|
|
| 59 |
from headroom.providers import AnthropicProvider, OpenAIProvider
|
| 60 |
from headroom.telemetry import get_telemetry_collector
|
| 61 |
from headroom.tokenizers import get_tokenizer
|
| 62 |
-
from headroom.transforms import
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 63 |
|
| 64 |
logging.basicConfig(
|
| 65 |
level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s"
|
|
@@ -145,6 +155,11 @@ class ProxyConfig:
|
|
| 145 |
ccr_inject_tool: bool = True # Inject headroom_retrieve tool when compression occurs
|
| 146 |
ccr_inject_system_instructions: bool = False # Add instructions to system message
|
| 147 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 148 |
# Caching
|
| 149 |
cache_enabled: bool = True
|
| 150 |
cache_ttl_seconds: int = 3600 # 1 hour
|
|
@@ -656,6 +671,9 @@ class HeadroomProxy:
|
|
| 656 |
),
|
| 657 |
]
|
| 658 |
|
|
|
|
|
|
|
|
|
|
| 659 |
self.anthropic_pipeline = TransformPipeline(
|
| 660 |
transforms=transforms,
|
| 661 |
provider=self.anthropic_provider,
|
|
@@ -722,6 +740,38 @@ class HeadroomProxy:
|
|
| 722 |
inject_system_instructions=config.ccr_inject_system_instructions,
|
| 723 |
)
|
| 724 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 725 |
async def startup(self):
|
| 726 |
"""Initialize async resources."""
|
| 727 |
self.http_client = httpx.AsyncClient(
|
|
@@ -737,6 +787,18 @@ class HeadroomProxy:
|
|
| 737 |
logger.info(f"Caching: {'ENABLED' if self.config.cache_enabled else 'DISABLED'}")
|
| 738 |
logger.info(f"Rate Limiting: {'ENABLED' if self.config.rate_limit_enabled else 'DISABLED'}")
|
| 739 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 740 |
async def shutdown(self):
|
| 741 |
"""Cleanup async resources."""
|
| 742 |
if self.http_client:
|
|
@@ -1818,6 +1880,21 @@ def create_app(config: ProxyConfig | None = None) -> FastAPI:
|
|
| 1818 |
return app
|
| 1819 |
|
| 1820 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1821 |
def run_server(config: ProxyConfig | None = None):
|
| 1822 |
"""Run the proxy server."""
|
| 1823 |
if not FASTAPI_AVAILABLE:
|
|
@@ -1827,6 +1904,8 @@ def run_server(config: ProxyConfig | None = None):
|
|
| 1827 |
config = config or ProxyConfig()
|
| 1828 |
app = create_app(config)
|
| 1829 |
|
|
|
|
|
|
|
| 1830 |
print(f"""
|
| 1831 |
╔══════════════════════════════════════════════════════════════════════╗
|
| 1832 |
║ HEADROOM PROXY SERVER ║
|
|
@@ -1840,6 +1919,7 @@ def run_server(config: ProxyConfig | None = None):
|
|
| 1840 |
║ Rate Limiting: {"ENABLED " if config.rate_limit_enabled else "DISABLED"} ({config.rate_limit_requests_per_minute} req/min, {config.rate_limit_tokens_per_minute:,} tok/min) ║
|
| 1841 |
║ Retry: {"ENABLED " if config.retry_enabled else "DISABLED"} (max {config.retry_max_attempts} attempts) ║
|
| 1842 |
║ Cost Tracking: {"ENABLED " if config.cost_tracking_enabled else "DISABLED"} (budget: {"$" + str(config.budget_limit_usd) + "/" + config.budget_period if config.budget_limit_usd else "unlimited"}) ║
|
|
|
|
| 1843 |
╠══════════════════════════════════════════════════════════════════════╣
|
| 1844 |
║ USAGE: ║
|
| 1845 |
║ Claude Code: ANTHROPIC_BASE_URL=http://{config.host}:{config.port} claude ║
|
|
@@ -1893,6 +1973,25 @@ if __name__ == "__main__":
|
|
| 1893 |
parser.add_argument("--log-file", help="Log file path")
|
| 1894 |
parser.add_argument("--log-messages", action="store_true", help="Log full messages")
|
| 1895 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1896 |
args = parser.parse_args()
|
| 1897 |
|
| 1898 |
config = ProxyConfig(
|
|
@@ -1910,6 +2009,9 @@ if __name__ == "__main__":
|
|
| 1910 |
budget_period=args.budget_period,
|
| 1911 |
log_file=args.log_file,
|
| 1912 |
log_full_messages=args.log_messages,
|
|
|
|
|
|
|
|
|
|
| 1913 |
)
|
| 1914 |
|
| 1915 |
run_server(config)
|
|
|
|
| 59 |
from headroom.providers import AnthropicProvider, OpenAIProvider
|
| 60 |
from headroom.telemetry import get_telemetry_collector
|
| 61 |
from headroom.tokenizers import get_tokenizer
|
| 62 |
+
from headroom.transforms import (
|
| 63 |
+
_LLMLINGUA_AVAILABLE,
|
| 64 |
+
CacheAligner,
|
| 65 |
+
RollingWindow,
|
| 66 |
+
SmartCrusher,
|
| 67 |
+
TransformPipeline,
|
| 68 |
+
)
|
| 69 |
+
|
| 70 |
+
# Conditionally import LLMLingua if available
|
| 71 |
+
if _LLMLINGUA_AVAILABLE:
|
| 72 |
+
from headroom.transforms import LLMLinguaCompressor, LLMLinguaConfig
|
| 73 |
|
| 74 |
logging.basicConfig(
|
| 75 |
level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s"
|
|
|
|
| 155 |
ccr_inject_tool: bool = True # Inject headroom_retrieve tool when compression occurs
|
| 156 |
ccr_inject_system_instructions: bool = False # Add instructions to system message
|
| 157 |
|
| 158 |
+
# LLMLingua ML-based compression (opt-in)
|
| 159 |
+
llmlingua_enabled: bool = False # Enable LLMLingua-2 for ML-based compression
|
| 160 |
+
llmlingua_device: str = "auto" # Device: 'auto', 'cuda', 'cpu', 'mps'
|
| 161 |
+
llmlingua_target_rate: float = 0.3 # Target compression rate (0.3 = keep 30%)
|
| 162 |
+
|
| 163 |
# Caching
|
| 164 |
cache_enabled: bool = True
|
| 165 |
cache_ttl_seconds: int = 3600 # 1 hour
|
|
|
|
| 671 |
),
|
| 672 |
]
|
| 673 |
|
| 674 |
+
# Add LLMLingua if enabled and available
|
| 675 |
+
self._llmlingua_status = self._setup_llmlingua(config, transforms)
|
| 676 |
+
|
| 677 |
self.anthropic_pipeline = TransformPipeline(
|
| 678 |
transforms=transforms,
|
| 679 |
provider=self.anthropic_provider,
|
|
|
|
| 740 |
inject_system_instructions=config.ccr_inject_system_instructions,
|
| 741 |
)
|
| 742 |
|
| 743 |
+
def _setup_llmlingua(self, config: ProxyConfig, transforms: list) -> str:
|
| 744 |
+
"""Set up LLMLingua compression if enabled.
|
| 745 |
+
|
| 746 |
+
Args:
|
| 747 |
+
config: Proxy configuration
|
| 748 |
+
transforms: Transform list to append to
|
| 749 |
+
|
| 750 |
+
Returns:
|
| 751 |
+
Status string for logging: 'enabled', 'disabled', 'available', 'unavailable'
|
| 752 |
+
"""
|
| 753 |
+
if config.llmlingua_enabled:
|
| 754 |
+
if _LLMLINGUA_AVAILABLE:
|
| 755 |
+
llmlingua_config = LLMLinguaConfig(
|
| 756 |
+
device=config.llmlingua_device,
|
| 757 |
+
target_compression_rate=config.llmlingua_target_rate,
|
| 758 |
+
enable_ccr=config.ccr_inject_tool, # Link to CCR
|
| 759 |
+
)
|
| 760 |
+
# Insert before RollingWindow (which should be last)
|
| 761 |
+
# LLMLingua works best on individual tool outputs before windowing
|
| 762 |
+
transforms.insert(-1, LLMLinguaCompressor(llmlingua_config))
|
| 763 |
+
return "enabled"
|
| 764 |
+
else:
|
| 765 |
+
logger.warning(
|
| 766 |
+
"LLMLingua requested but not installed. "
|
| 767 |
+
"Install with: pip install headroom-ai[llmlingua]"
|
| 768 |
+
)
|
| 769 |
+
return "unavailable"
|
| 770 |
+
else:
|
| 771 |
+
if _LLMLINGUA_AVAILABLE:
|
| 772 |
+
return "available" # Available but not enabled - hint to user
|
| 773 |
+
return "disabled"
|
| 774 |
+
|
| 775 |
async def startup(self):
|
| 776 |
"""Initialize async resources."""
|
| 777 |
self.http_client = httpx.AsyncClient(
|
|
|
|
| 787 |
logger.info(f"Caching: {'ENABLED' if self.config.cache_enabled else 'DISABLED'}")
|
| 788 |
logger.info(f"Rate Limiting: {'ENABLED' if self.config.rate_limit_enabled else 'DISABLED'}")
|
| 789 |
|
| 790 |
+
# LLMLingua status with helpful hint
|
| 791 |
+
if self._llmlingua_status == "enabled":
|
| 792 |
+
logger.info(
|
| 793 |
+
f"LLMLingua: ENABLED (device={self.config.llmlingua_device}, "
|
| 794 |
+
f"rate={self.config.llmlingua_target_rate})"
|
| 795 |
+
)
|
| 796 |
+
elif self._llmlingua_status == "available":
|
| 797 |
+
logger.info(
|
| 798 |
+
"LLMLingua: available but not enabled. "
|
| 799 |
+
"Enable with --llmlingua for ML-based compression (3-5x better on text/logs)"
|
| 800 |
+
)
|
| 801 |
+
|
| 802 |
async def shutdown(self):
|
| 803 |
"""Cleanup async resources."""
|
| 804 |
if self.http_client:
|
|
|
|
| 1880 |
return app
|
| 1881 |
|
| 1882 |
|
| 1883 |
+
def _get_llmlingua_banner_status(config: ProxyConfig) -> str:
|
| 1884 |
+
"""Get LLMLingua status line for banner."""
|
| 1885 |
+
if config.llmlingua_enabled:
|
| 1886 |
+
if _LLMLINGUA_AVAILABLE:
|
| 1887 |
+
return (
|
| 1888 |
+
f"ENABLED (device={config.llmlingua_device}, rate={config.llmlingua_target_rate})"
|
| 1889 |
+
)
|
| 1890 |
+
else:
|
| 1891 |
+
return "REQUESTED but not installed (pip install headroom-ai[llmlingua])"
|
| 1892 |
+
else:
|
| 1893 |
+
if _LLMLINGUA_AVAILABLE:
|
| 1894 |
+
return "available (enable with --llmlingua for ML compression)"
|
| 1895 |
+
return "DISABLED"
|
| 1896 |
+
|
| 1897 |
+
|
| 1898 |
def run_server(config: ProxyConfig | None = None):
|
| 1899 |
"""Run the proxy server."""
|
| 1900 |
if not FASTAPI_AVAILABLE:
|
|
|
|
| 1904 |
config = config or ProxyConfig()
|
| 1905 |
app = create_app(config)
|
| 1906 |
|
| 1907 |
+
llmlingua_status = _get_llmlingua_banner_status(config)
|
| 1908 |
+
|
| 1909 |
print(f"""
|
| 1910 |
╔══════════════════════════════════════════════════════════════════════╗
|
| 1911 |
║ HEADROOM PROXY SERVER ║
|
|
|
|
| 1919 |
║ Rate Limiting: {"ENABLED " if config.rate_limit_enabled else "DISABLED"} ({config.rate_limit_requests_per_minute} req/min, {config.rate_limit_tokens_per_minute:,} tok/min) ║
|
| 1920 |
║ Retry: {"ENABLED " if config.retry_enabled else "DISABLED"} (max {config.retry_max_attempts} attempts) ║
|
| 1921 |
║ Cost Tracking: {"ENABLED " if config.cost_tracking_enabled else "DISABLED"} (budget: {"$" + str(config.budget_limit_usd) + "/" + config.budget_period if config.budget_limit_usd else "unlimited"}) ║
|
| 1922 |
+
║ LLMLingua: {llmlingua_status:<52}║
|
| 1923 |
╠══════════════════════════════════════════════════════════════════════╣
|
| 1924 |
║ USAGE: ║
|
| 1925 |
║ Claude Code: ANTHROPIC_BASE_URL=http://{config.host}:{config.port} claude ║
|
|
|
|
| 1973 |
parser.add_argument("--log-file", help="Log file path")
|
| 1974 |
parser.add_argument("--log-messages", action="store_true", help="Log full messages")
|
| 1975 |
|
| 1976 |
+
# LLMLingua ML-based compression
|
| 1977 |
+
parser.add_argument(
|
| 1978 |
+
"--llmlingua",
|
| 1979 |
+
action="store_true",
|
| 1980 |
+
help="Enable LLMLingua-2 ML-based compression (requires: pip install headroom-ai[llmlingua])",
|
| 1981 |
+
)
|
| 1982 |
+
parser.add_argument(
|
| 1983 |
+
"--llmlingua-device",
|
| 1984 |
+
choices=["auto", "cuda", "cpu", "mps"],
|
| 1985 |
+
default="auto",
|
| 1986 |
+
help="Device for LLMLingua model (default: auto)",
|
| 1987 |
+
)
|
| 1988 |
+
parser.add_argument(
|
| 1989 |
+
"--llmlingua-rate",
|
| 1990 |
+
type=float,
|
| 1991 |
+
default=0.3,
|
| 1992 |
+
help="LLMLingua target compression rate, 0.0-1.0 (default: 0.3 = keep 30%%)",
|
| 1993 |
+
)
|
| 1994 |
+
|
| 1995 |
args = parser.parse_args()
|
| 1996 |
|
| 1997 |
config = ProxyConfig(
|
|
|
|
| 2009 |
budget_period=args.budget_period,
|
| 2010 |
log_file=args.log_file,
|
| 2011 |
log_full_messages=args.log_messages,
|
| 2012 |
+
llmlingua_enabled=args.llmlingua,
|
| 2013 |
+
llmlingua_device=args.llmlingua_device,
|
| 2014 |
+
llmlingua_target_rate=args.llmlingua_rate,
|
| 2015 |
)
|
| 2016 |
|
| 2017 |
run_server(config)
|
|
@@ -15,6 +15,21 @@ from .smart_crusher import SmartCrusher, SmartCrusherConfig
|
|
| 15 |
from .text_compressor import TextCompressionResult, TextCompressor, TextCompressorConfig
|
| 16 |
from .tool_crusher import ToolCrusher
|
| 17 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 18 |
__all__ = [
|
| 19 |
# Base
|
| 20 |
"Transform",
|
|
@@ -39,4 +54,19 @@ __all__ = [
|
|
| 39 |
# Other transforms
|
| 40 |
"CacheAligner",
|
| 41 |
"RollingWindow",
|
|
|
|
|
|
|
| 42 |
]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 15 |
from .text_compressor import TextCompressionResult, TextCompressor, TextCompressorConfig
|
| 16 |
from .tool_crusher import ToolCrusher
|
| 17 |
|
| 18 |
+
# ML-based compression (optional dependency)
|
| 19 |
+
try:
|
| 20 |
+
from .llmlingua_compressor import ( # noqa: F401
|
| 21 |
+
LLMLinguaCompressor,
|
| 22 |
+
LLMLinguaConfig,
|
| 23 |
+
LLMLinguaResult,
|
| 24 |
+
compress_with_llmlingua,
|
| 25 |
+
is_llmlingua_model_loaded,
|
| 26 |
+
unload_llmlingua_model,
|
| 27 |
+
)
|
| 28 |
+
|
| 29 |
+
_LLMLINGUA_AVAILABLE = True
|
| 30 |
+
except ImportError:
|
| 31 |
+
_LLMLINGUA_AVAILABLE = False
|
| 32 |
+
|
| 33 |
__all__ = [
|
| 34 |
# Base
|
| 35 |
"Transform",
|
|
|
|
| 54 |
# Other transforms
|
| 55 |
"CacheAligner",
|
| 56 |
"RollingWindow",
|
| 57 |
+
# ML-based compression (optional)
|
| 58 |
+
"_LLMLINGUA_AVAILABLE",
|
| 59 |
]
|
| 60 |
+
|
| 61 |
+
# Conditionally add LLMLingua exports
|
| 62 |
+
if _LLMLINGUA_AVAILABLE:
|
| 63 |
+
__all__.extend(
|
| 64 |
+
[
|
| 65 |
+
"LLMLinguaCompressor",
|
| 66 |
+
"LLMLinguaConfig",
|
| 67 |
+
"LLMLinguaResult",
|
| 68 |
+
"compress_with_llmlingua",
|
| 69 |
+
"is_llmlingua_model_loaded",
|
| 70 |
+
"unload_llmlingua_model",
|
| 71 |
+
]
|
| 72 |
+
)
|
|
@@ -0,0 +1,633 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""LLMLingua-2 compressor for ML-based prompt compression.
|
| 2 |
+
|
| 3 |
+
This module provides integration with LLMLingua-2, a BERT-based token classifier
|
| 4 |
+
trained via GPT-4 distillation. It achieves superior compression (up to 20x)
|
| 5 |
+
while maintaining high fidelity on tool outputs and structured content.
|
| 6 |
+
|
| 7 |
+
Key Features:
|
| 8 |
+
- Token-level classification (keep/remove) using fine-tuned BERT
|
| 9 |
+
- 3-6x faster than LLMLingua-1 with better results
|
| 10 |
+
- Especially effective on tool outputs, code, and structured data
|
| 11 |
+
- Reversible compression via CCR integration
|
| 12 |
+
|
| 13 |
+
Reference:
|
| 14 |
+
LLMLingua-2: Data Distillation for Efficient and Faithful Task-Agnostic Prompt Compression
|
| 15 |
+
https://arxiv.org/abs/2403.12968
|
| 16 |
+
|
| 17 |
+
Installation:
|
| 18 |
+
pip install headroom-ai[llmlingua]
|
| 19 |
+
|
| 20 |
+
Usage:
|
| 21 |
+
>>> from headroom.transforms import LLMLinguaCompressor
|
| 22 |
+
>>> compressor = LLMLinguaCompressor()
|
| 23 |
+
>>> result = compressor.compress(long_tool_output)
|
| 24 |
+
>>> print(result.compressed) # Significantly reduced output
|
| 25 |
+
"""
|
| 26 |
+
|
| 27 |
+
from __future__ import annotations
|
| 28 |
+
|
| 29 |
+
import logging
|
| 30 |
+
import threading
|
| 31 |
+
from dataclasses import dataclass, field
|
| 32 |
+
from typing import Any
|
| 33 |
+
|
| 34 |
+
from ..config import TransformResult
|
| 35 |
+
from ..tokenizer import Tokenizer
|
| 36 |
+
from .base import Transform
|
| 37 |
+
|
| 38 |
+
logger = logging.getLogger(__name__)
|
| 39 |
+
|
| 40 |
+
# Lazy import for optional dependency
|
| 41 |
+
_llmlingua_available: bool | None = None
|
| 42 |
+
_llmlingua_instance: Any = None
|
| 43 |
+
_llmlingua_lock = threading.Lock() # Thread safety for model access
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
def _check_llmlingua_available() -> bool:
|
| 47 |
+
"""Check if llmlingua package is available."""
|
| 48 |
+
global _llmlingua_available
|
| 49 |
+
if _llmlingua_available is None:
|
| 50 |
+
try:
|
| 51 |
+
import llmlingua # noqa: F401
|
| 52 |
+
|
| 53 |
+
_llmlingua_available = True
|
| 54 |
+
except ImportError:
|
| 55 |
+
_llmlingua_available = False
|
| 56 |
+
return _llmlingua_available
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
def _get_llmlingua_compressor(model_name: str, device: str) -> Any:
|
| 60 |
+
"""Get or create the LLMLingua compressor instance.
|
| 61 |
+
|
| 62 |
+
Uses lazy initialization and caches the instance to avoid repeated model loading.
|
| 63 |
+
Thread-safe: uses lock to prevent race conditions during model initialization.
|
| 64 |
+
|
| 65 |
+
Args:
|
| 66 |
+
model_name: HuggingFace model name for the compressor.
|
| 67 |
+
device: Device to run the model on ('cuda', 'cpu', or 'auto').
|
| 68 |
+
|
| 69 |
+
Returns:
|
| 70 |
+
PromptCompressor instance from llmlingua.
|
| 71 |
+
|
| 72 |
+
Raises:
|
| 73 |
+
ImportError: If llmlingua is not installed.
|
| 74 |
+
RuntimeError: If model loading fails.
|
| 75 |
+
"""
|
| 76 |
+
global _llmlingua_instance
|
| 77 |
+
|
| 78 |
+
if not _check_llmlingua_available():
|
| 79 |
+
raise ImportError(
|
| 80 |
+
"llmlingua is not installed. Install with: pip install headroom-ai[llmlingua]\n"
|
| 81 |
+
"Note: This requires ~2GB of disk space and ~1GB RAM for the model."
|
| 82 |
+
)
|
| 83 |
+
|
| 84 |
+
with _llmlingua_lock:
|
| 85 |
+
# Double-check after acquiring lock
|
| 86 |
+
if _llmlingua_instance is None or _llmlingua_instance._model_name != model_name:
|
| 87 |
+
try:
|
| 88 |
+
from llmlingua import PromptCompressor
|
| 89 |
+
|
| 90 |
+
logger.info(
|
| 91 |
+
"Loading LLMLingua-2 model: %s on device: %s (this may take 10-30s on first run)",
|
| 92 |
+
model_name,
|
| 93 |
+
device,
|
| 94 |
+
)
|
| 95 |
+
_llmlingua_instance = PromptCompressor(
|
| 96 |
+
model_name=model_name,
|
| 97 |
+
device_map=device,
|
| 98 |
+
use_llmlingua2=True, # Use LLMLingua-2 (BERT classifier)
|
| 99 |
+
)
|
| 100 |
+
# Store model name for later comparison
|
| 101 |
+
_llmlingua_instance._model_name = model_name
|
| 102 |
+
logger.info("LLMLingua-2 model loaded successfully")
|
| 103 |
+
|
| 104 |
+
except Exception as e:
|
| 105 |
+
error_msg = str(e).lower()
|
| 106 |
+
if "out of memory" in error_msg or "oom" in error_msg:
|
| 107 |
+
raise RuntimeError(
|
| 108 |
+
f"Out of memory loading LLMLingua model. Try:\n"
|
| 109 |
+
f" 1. Use device='cpu' instead of 'cuda'\n"
|
| 110 |
+
f" 2. Close other GPU applications\n"
|
| 111 |
+
f" 3. Use a smaller model\n"
|
| 112 |
+
f"Original error: {e}"
|
| 113 |
+
) from e
|
| 114 |
+
elif "not found" in error_msg or "404" in error_msg:
|
| 115 |
+
raise RuntimeError(
|
| 116 |
+
f"Model '{model_name}' not found on HuggingFace. Try:\n"
|
| 117 |
+
f" 1. Check the model name is correct\n"
|
| 118 |
+
f" 2. Use default: 'microsoft/llmlingua-2-xlm-roberta-large-meetingbank'\n"
|
| 119 |
+
f"Original error: {e}"
|
| 120 |
+
) from e
|
| 121 |
+
else:
|
| 122 |
+
raise RuntimeError(
|
| 123 |
+
f"Failed to load LLMLingua model: {e}\n"
|
| 124 |
+
f"Ensure you have sufficient disk space and memory."
|
| 125 |
+
) from e
|
| 126 |
+
|
| 127 |
+
return _llmlingua_instance
|
| 128 |
+
|
| 129 |
+
|
| 130 |
+
def unload_llmlingua_model() -> bool:
|
| 131 |
+
"""Unload the LLMLingua model to free memory.
|
| 132 |
+
|
| 133 |
+
Use this when you're done with compression and want to reclaim GPU/CPU memory.
|
| 134 |
+
The model will be reloaded automatically on the next compression call.
|
| 135 |
+
|
| 136 |
+
Returns:
|
| 137 |
+
True if a model was unloaded, False if no model was loaded.
|
| 138 |
+
|
| 139 |
+
Example:
|
| 140 |
+
>>> from headroom.transforms import LLMLinguaCompressor, unload_llmlingua_model
|
| 141 |
+
>>> compressor = LLMLinguaCompressor()
|
| 142 |
+
>>> result = compressor.compress(content) # Model loaded here
|
| 143 |
+
>>> # ... do other work ...
|
| 144 |
+
>>> unload_llmlingua_model() # Free ~1GB of memory
|
| 145 |
+
"""
|
| 146 |
+
global _llmlingua_instance
|
| 147 |
+
|
| 148 |
+
with _llmlingua_lock:
|
| 149 |
+
if _llmlingua_instance is not None:
|
| 150 |
+
model_name = getattr(_llmlingua_instance, "_model_name", "unknown")
|
| 151 |
+
logger.info("Unloading LLMLingua model: %s", model_name)
|
| 152 |
+
|
| 153 |
+
# Clear the instance
|
| 154 |
+
_llmlingua_instance = None
|
| 155 |
+
|
| 156 |
+
# Attempt to free GPU memory if torch is available
|
| 157 |
+
try:
|
| 158 |
+
import torch
|
| 159 |
+
|
| 160 |
+
if torch.cuda.is_available():
|
| 161 |
+
torch.cuda.empty_cache()
|
| 162 |
+
logger.debug("Cleared CUDA cache")
|
| 163 |
+
except ImportError:
|
| 164 |
+
pass
|
| 165 |
+
|
| 166 |
+
return True
|
| 167 |
+
|
| 168 |
+
return False
|
| 169 |
+
|
| 170 |
+
|
| 171 |
+
def is_llmlingua_model_loaded() -> bool:
|
| 172 |
+
"""Check if an LLMLingua model is currently loaded.
|
| 173 |
+
|
| 174 |
+
Returns:
|
| 175 |
+
True if a model is loaded in memory, False otherwise.
|
| 176 |
+
"""
|
| 177 |
+
return _llmlingua_instance is not None
|
| 178 |
+
|
| 179 |
+
|
| 180 |
+
@dataclass
|
| 181 |
+
class LLMLinguaConfig:
|
| 182 |
+
"""Configuration for LLMLingua-2 compression.
|
| 183 |
+
|
| 184 |
+
Attributes:
|
| 185 |
+
model_name: HuggingFace model for the compressor. Default is the
|
| 186 |
+
LLMLingua-2 xlm-roberta-large model fine-tuned for compression.
|
| 187 |
+
device: Device to run on ('cuda', 'cpu', 'auto'). Auto will use CUDA if available.
|
| 188 |
+
target_compression_rate: Target compression ratio (e.g., 0.3 = keep 30% of tokens).
|
| 189 |
+
force_tokens: Tokens to always preserve (e.g., important keywords).
|
| 190 |
+
drop_consecutive: Whether to drop consecutive punctuation/whitespace.
|
| 191 |
+
min_tokens_for_compression: Minimum token count to trigger compression.
|
| 192 |
+
Content below this threshold is passed through unchanged.
|
| 193 |
+
enable_ccr: Whether to store originals in CCR for retrieval.
|
| 194 |
+
ccr_ttl: TTL for CCR entries in seconds.
|
| 195 |
+
|
| 196 |
+
GOTCHA: Lower target_compression_rate = more aggressive compression.
|
| 197 |
+
A rate of 0.2 means keeping only 20% of tokens.
|
| 198 |
+
"""
|
| 199 |
+
|
| 200 |
+
# Model configuration
|
| 201 |
+
model_name: str = "microsoft/llmlingua-2-xlm-roberta-large-meetingbank"
|
| 202 |
+
device: str = "auto"
|
| 203 |
+
|
| 204 |
+
# Compression parameters
|
| 205 |
+
target_compression_rate: float = 0.3
|
| 206 |
+
force_tokens: list[str] = field(default_factory=list)
|
| 207 |
+
drop_consecutive: bool = True
|
| 208 |
+
|
| 209 |
+
# Thresholds
|
| 210 |
+
min_tokens_for_compression: int = 100
|
| 211 |
+
|
| 212 |
+
# CCR integration
|
| 213 |
+
enable_ccr: bool = True
|
| 214 |
+
ccr_ttl: int = 300 # 5 minutes
|
| 215 |
+
|
| 216 |
+
# Content type specific settings
|
| 217 |
+
code_compression_rate: float = 0.4 # More conservative for code
|
| 218 |
+
json_compression_rate: float = 0.35 # Slightly conservative for JSON
|
| 219 |
+
text_compression_rate: float = 0.25 # More aggressive for plain text
|
| 220 |
+
|
| 221 |
+
|
| 222 |
+
@dataclass
|
| 223 |
+
class LLMLinguaResult:
|
| 224 |
+
"""Result of LLMLingua-2 compression.
|
| 225 |
+
|
| 226 |
+
Attributes:
|
| 227 |
+
compressed: Compressed content.
|
| 228 |
+
original: Original content before compression.
|
| 229 |
+
original_tokens: Token count of original content.
|
| 230 |
+
compressed_tokens: Token count after compression.
|
| 231 |
+
compression_ratio: Actual compression ratio achieved.
|
| 232 |
+
cache_key: CCR cache key if stored.
|
| 233 |
+
model_used: Model that performed the compression.
|
| 234 |
+
tokens_saved: Number of tokens saved.
|
| 235 |
+
"""
|
| 236 |
+
|
| 237 |
+
compressed: str
|
| 238 |
+
original: str
|
| 239 |
+
original_tokens: int
|
| 240 |
+
compressed_tokens: int
|
| 241 |
+
compression_ratio: float
|
| 242 |
+
cache_key: str | None = None
|
| 243 |
+
model_used: str | None = None
|
| 244 |
+
|
| 245 |
+
@property
|
| 246 |
+
def tokens_saved(self) -> int:
|
| 247 |
+
"""Number of tokens saved by compression."""
|
| 248 |
+
return max(0, self.original_tokens - self.compressed_tokens)
|
| 249 |
+
|
| 250 |
+
@property
|
| 251 |
+
def savings_percentage(self) -> float:
|
| 252 |
+
"""Percentage of tokens saved."""
|
| 253 |
+
if self.original_tokens == 0:
|
| 254 |
+
return 0.0
|
| 255 |
+
return (self.tokens_saved / self.original_tokens) * 100
|
| 256 |
+
|
| 257 |
+
|
| 258 |
+
class LLMLinguaCompressor(Transform):
|
| 259 |
+
"""LLMLingua-2 based prompt compressor.
|
| 260 |
+
|
| 261 |
+
Uses a BERT-based token classifier trained via GPT-4 distillation to
|
| 262 |
+
identify and remove non-essential tokens while preserving semantic meaning.
|
| 263 |
+
|
| 264 |
+
Key advantages over statistical compression:
|
| 265 |
+
- Learned token importance from LLM feedback
|
| 266 |
+
- Better handling of context-dependent importance
|
| 267 |
+
- More aggressive compression with less information loss
|
| 268 |
+
- Especially effective on structured outputs (JSON, code, logs)
|
| 269 |
+
|
| 270 |
+
Example:
|
| 271 |
+
>>> compressor = LLMLinguaCompressor()
|
| 272 |
+
>>> result = compressor.compress(long_tool_output)
|
| 273 |
+
>>> print(f"Saved {result.tokens_saved} tokens ({result.savings_percentage:.1f}%)")
|
| 274 |
+
|
| 275 |
+
>>> # Use as a Transform in pipeline
|
| 276 |
+
>>> from headroom.transforms import TransformPipeline
|
| 277 |
+
>>> pipeline = TransformPipeline([LLMLinguaCompressor()])
|
| 278 |
+
>>> result = pipeline.apply(messages, tokenizer)
|
| 279 |
+
"""
|
| 280 |
+
|
| 281 |
+
name: str = "llmlingua_compressor"
|
| 282 |
+
|
| 283 |
+
def __init__(self, config: LLMLinguaConfig | None = None):
|
| 284 |
+
"""Initialize LLMLingua compressor.
|
| 285 |
+
|
| 286 |
+
Args:
|
| 287 |
+
config: Compression configuration. If None, uses defaults.
|
| 288 |
+
|
| 289 |
+
Note:
|
| 290 |
+
The underlying model is loaded lazily on first use to avoid
|
| 291 |
+
startup overhead when the compressor isn't used.
|
| 292 |
+
"""
|
| 293 |
+
self.config = config or LLMLinguaConfig()
|
| 294 |
+
self._compressor: Any = None # Lazy loaded
|
| 295 |
+
|
| 296 |
+
def compress(
|
| 297 |
+
self,
|
| 298 |
+
content: str,
|
| 299 |
+
context: str = "",
|
| 300 |
+
content_type: str | None = None,
|
| 301 |
+
) -> LLMLinguaResult:
|
| 302 |
+
"""Compress content using LLMLingua-2.
|
| 303 |
+
|
| 304 |
+
Args:
|
| 305 |
+
content: Content to compress.
|
| 306 |
+
context: Optional context for relevance-aware compression.
|
| 307 |
+
content_type: Type of content ('code', 'json', 'text').
|
| 308 |
+
If None, auto-detected.
|
| 309 |
+
|
| 310 |
+
Returns:
|
| 311 |
+
LLMLinguaResult with compressed content and metadata.
|
| 312 |
+
|
| 313 |
+
Raises:
|
| 314 |
+
ImportError: If llmlingua is not installed.
|
| 315 |
+
"""
|
| 316 |
+
# Check availability
|
| 317 |
+
if not _check_llmlingua_available():
|
| 318 |
+
logger.warning(
|
| 319 |
+
"LLMLingua not available. Install with: pip install headroom-ai[llmlingua]"
|
| 320 |
+
)
|
| 321 |
+
return LLMLinguaResult(
|
| 322 |
+
compressed=content,
|
| 323 |
+
original=content,
|
| 324 |
+
original_tokens=len(content.split()), # Rough estimate
|
| 325 |
+
compressed_tokens=len(content.split()),
|
| 326 |
+
compression_ratio=1.0,
|
| 327 |
+
)
|
| 328 |
+
|
| 329 |
+
# Estimate token count (rough)
|
| 330 |
+
estimated_tokens = len(content.split())
|
| 331 |
+
|
| 332 |
+
# Skip compression for small content
|
| 333 |
+
if estimated_tokens < self.config.min_tokens_for_compression:
|
| 334 |
+
return LLMLinguaResult(
|
| 335 |
+
compressed=content,
|
| 336 |
+
original=content,
|
| 337 |
+
original_tokens=estimated_tokens,
|
| 338 |
+
compressed_tokens=estimated_tokens,
|
| 339 |
+
compression_ratio=1.0,
|
| 340 |
+
)
|
| 341 |
+
|
| 342 |
+
# Get compression rate based on content type
|
| 343 |
+
compression_rate = self._get_compression_rate(content, content_type)
|
| 344 |
+
|
| 345 |
+
# Get or initialize compressor
|
| 346 |
+
device = self._resolve_device()
|
| 347 |
+
compressor = _get_llmlingua_compressor(self.config.model_name, device)
|
| 348 |
+
|
| 349 |
+
# Prepare force tokens
|
| 350 |
+
force_tokens = list(self.config.force_tokens)
|
| 351 |
+
|
| 352 |
+
# Add context words as force tokens if provided
|
| 353 |
+
if context:
|
| 354 |
+
context_words = [w for w in context.split() if len(w) > 3]
|
| 355 |
+
force_tokens.extend(context_words[:10]) # Limit to avoid overhead
|
| 356 |
+
|
| 357 |
+
# Perform compression
|
| 358 |
+
try:
|
| 359 |
+
result = compressor.compress_prompt(
|
| 360 |
+
original_prompt=content,
|
| 361 |
+
rate=compression_rate,
|
| 362 |
+
force_tokens=force_tokens if force_tokens else None,
|
| 363 |
+
drop_consecutive=self.config.drop_consecutive,
|
| 364 |
+
)
|
| 365 |
+
|
| 366 |
+
compressed = result.get("compressed_prompt", content)
|
| 367 |
+
original_tokens = result.get("origin_tokens", estimated_tokens)
|
| 368 |
+
compressed_tokens = result.get("compressed_tokens", len(compressed.split()))
|
| 369 |
+
|
| 370 |
+
except Exception as e:
|
| 371 |
+
logger.warning("LLMLingua compression failed: %s", e)
|
| 372 |
+
return LLMLinguaResult(
|
| 373 |
+
compressed=content,
|
| 374 |
+
original=content,
|
| 375 |
+
original_tokens=estimated_tokens,
|
| 376 |
+
compressed_tokens=estimated_tokens,
|
| 377 |
+
compression_ratio=1.0,
|
| 378 |
+
)
|
| 379 |
+
|
| 380 |
+
# Calculate actual ratio
|
| 381 |
+
ratio = compressed_tokens / max(original_tokens, 1)
|
| 382 |
+
|
| 383 |
+
# Store in CCR if enabled
|
| 384 |
+
cache_key = None
|
| 385 |
+
if self.config.enable_ccr and ratio < 0.8:
|
| 386 |
+
cache_key = self._store_in_ccr(content, compressed, original_tokens)
|
| 387 |
+
if cache_key:
|
| 388 |
+
compressed += (
|
| 389 |
+
f"\n[LLMLingua: {original_tokens}→{compressed_tokens} tokens. hash={cache_key}]"
|
| 390 |
+
)
|
| 391 |
+
|
| 392 |
+
return LLMLinguaResult(
|
| 393 |
+
compressed=compressed,
|
| 394 |
+
original=content,
|
| 395 |
+
original_tokens=original_tokens,
|
| 396 |
+
compressed_tokens=compressed_tokens,
|
| 397 |
+
compression_ratio=ratio,
|
| 398 |
+
cache_key=cache_key,
|
| 399 |
+
model_used=self.config.model_name,
|
| 400 |
+
)
|
| 401 |
+
|
| 402 |
+
def apply(
|
| 403 |
+
self,
|
| 404 |
+
messages: list[dict[str, Any]],
|
| 405 |
+
tokenizer: Tokenizer,
|
| 406 |
+
**kwargs: Any,
|
| 407 |
+
) -> TransformResult:
|
| 408 |
+
"""Apply LLMLingua compression to messages.
|
| 409 |
+
|
| 410 |
+
This method implements the Transform interface for use in pipelines.
|
| 411 |
+
It compresses tool outputs and long assistant/user messages.
|
| 412 |
+
|
| 413 |
+
Args:
|
| 414 |
+
messages: List of message dicts to transform.
|
| 415 |
+
tokenizer: Tokenizer for accurate token counting.
|
| 416 |
+
**kwargs: Additional arguments (e.g., 'context' for relevance).
|
| 417 |
+
|
| 418 |
+
Returns:
|
| 419 |
+
TransformResult with compressed messages and metadata.
|
| 420 |
+
"""
|
| 421 |
+
tokens_before = sum(tokenizer.count_text(str(m.get("content", ""))) for m in messages)
|
| 422 |
+
context = kwargs.get("context", "")
|
| 423 |
+
|
| 424 |
+
transformed_messages = []
|
| 425 |
+
transforms_applied = []
|
| 426 |
+
warnings: list[str] = []
|
| 427 |
+
|
| 428 |
+
for message in messages:
|
| 429 |
+
role = message.get("role", "")
|
| 430 |
+
content = message.get("content", "")
|
| 431 |
+
|
| 432 |
+
# Compress tool results (highest value compression)
|
| 433 |
+
if role == "tool" and content:
|
| 434 |
+
result = self.compress(content, context=context, content_type="json")
|
| 435 |
+
if result.compression_ratio < 0.9:
|
| 436 |
+
transformed_messages.append({**message, "content": result.compressed})
|
| 437 |
+
transforms_applied.append(f"llmlingua:tool:{result.compression_ratio:.2f}")
|
| 438 |
+
else:
|
| 439 |
+
transformed_messages.append(message)
|
| 440 |
+
|
| 441 |
+
# Compress long assistant messages (tool outputs often embedded)
|
| 442 |
+
elif role == "assistant" and len(content) > 500:
|
| 443 |
+
result = self.compress(content, context=context)
|
| 444 |
+
if result.compression_ratio < 0.9:
|
| 445 |
+
transformed_messages.append({**message, "content": result.compressed})
|
| 446 |
+
transforms_applied.append(f"llmlingua:assistant:{result.compression_ratio:.2f}")
|
| 447 |
+
else:
|
| 448 |
+
transformed_messages.append(message)
|
| 449 |
+
|
| 450 |
+
# Pass through other messages
|
| 451 |
+
else:
|
| 452 |
+
transformed_messages.append(message)
|
| 453 |
+
|
| 454 |
+
tokens_after = sum(
|
| 455 |
+
tokenizer.count_text(str(m.get("content", ""))) for m in transformed_messages
|
| 456 |
+
)
|
| 457 |
+
|
| 458 |
+
# Add warning if llmlingua not available
|
| 459 |
+
if not _check_llmlingua_available():
|
| 460 |
+
warnings.append(
|
| 461 |
+
"LLMLingua not installed. Install with: pip install headroom-ai[llmlingua]"
|
| 462 |
+
)
|
| 463 |
+
|
| 464 |
+
return TransformResult(
|
| 465 |
+
messages=transformed_messages,
|
| 466 |
+
tokens_before=tokens_before,
|
| 467 |
+
tokens_after=tokens_after,
|
| 468 |
+
transforms_applied=transforms_applied if transforms_applied else ["llmlingua:noop"],
|
| 469 |
+
warnings=warnings,
|
| 470 |
+
)
|
| 471 |
+
|
| 472 |
+
def should_apply(
|
| 473 |
+
self,
|
| 474 |
+
messages: list[dict[str, Any]],
|
| 475 |
+
tokenizer: Tokenizer,
|
| 476 |
+
**kwargs: Any,
|
| 477 |
+
) -> bool:
|
| 478 |
+
"""Check if LLMLingua compression should be applied.
|
| 479 |
+
|
| 480 |
+
Returns True if:
|
| 481 |
+
- LLMLingua is available, AND
|
| 482 |
+
- Total token count exceeds minimum threshold
|
| 483 |
+
|
| 484 |
+
Args:
|
| 485 |
+
messages: Messages to check.
|
| 486 |
+
tokenizer: Tokenizer for counting.
|
| 487 |
+
**kwargs: Additional arguments.
|
| 488 |
+
|
| 489 |
+
Returns:
|
| 490 |
+
True if compression should be applied.
|
| 491 |
+
"""
|
| 492 |
+
if not _check_llmlingua_available():
|
| 493 |
+
return False
|
| 494 |
+
|
| 495 |
+
total_tokens = sum(tokenizer.count_text(str(m.get("content", ""))) for m in messages)
|
| 496 |
+
return total_tokens >= self.config.min_tokens_for_compression
|
| 497 |
+
|
| 498 |
+
def _get_compression_rate(
|
| 499 |
+
self,
|
| 500 |
+
content: str,
|
| 501 |
+
content_type: str | None,
|
| 502 |
+
) -> float:
|
| 503 |
+
"""Get appropriate compression rate based on content type.
|
| 504 |
+
|
| 505 |
+
Args:
|
| 506 |
+
content: Content to analyze.
|
| 507 |
+
content_type: Explicit content type or None for auto-detection.
|
| 508 |
+
|
| 509 |
+
Returns:
|
| 510 |
+
Target compression rate for this content.
|
| 511 |
+
"""
|
| 512 |
+
if content_type == "code":
|
| 513 |
+
return self.config.code_compression_rate
|
| 514 |
+
elif content_type == "json":
|
| 515 |
+
return self.config.json_compression_rate
|
| 516 |
+
elif content_type == "text":
|
| 517 |
+
return self.config.text_compression_rate
|
| 518 |
+
|
| 519 |
+
# Auto-detect content type
|
| 520 |
+
if self._looks_like_json(content):
|
| 521 |
+
return self.config.json_compression_rate
|
| 522 |
+
elif self._looks_like_code(content):
|
| 523 |
+
return self.config.code_compression_rate
|
| 524 |
+
else:
|
| 525 |
+
return self.config.text_compression_rate
|
| 526 |
+
|
| 527 |
+
def _looks_like_json(self, content: str) -> bool:
|
| 528 |
+
"""Check if content appears to be JSON."""
|
| 529 |
+
stripped = content.strip()
|
| 530 |
+
return (stripped.startswith("{") and stripped.endswith("}")) or (
|
| 531 |
+
stripped.startswith("[") and stripped.endswith("]")
|
| 532 |
+
)
|
| 533 |
+
|
| 534 |
+
def _looks_like_code(self, content: str) -> bool:
|
| 535 |
+
"""Check if content appears to be code."""
|
| 536 |
+
code_indicators = [
|
| 537 |
+
"def ",
|
| 538 |
+
"class ",
|
| 539 |
+
"function ",
|
| 540 |
+
"import ",
|
| 541 |
+
"from ",
|
| 542 |
+
"const ",
|
| 543 |
+
"let ",
|
| 544 |
+
"var ",
|
| 545 |
+
"public ",
|
| 546 |
+
"private ",
|
| 547 |
+
"async ",
|
| 548 |
+
"await ",
|
| 549 |
+
"return ",
|
| 550 |
+
"if (",
|
| 551 |
+
"for (",
|
| 552 |
+
"while (",
|
| 553 |
+
]
|
| 554 |
+
return any(indicator in content for indicator in code_indicators)
|
| 555 |
+
|
| 556 |
+
def _resolve_device(self) -> str:
|
| 557 |
+
"""Resolve 'auto' device to actual device."""
|
| 558 |
+
if self.config.device != "auto":
|
| 559 |
+
return self.config.device
|
| 560 |
+
|
| 561 |
+
try:
|
| 562 |
+
import torch
|
| 563 |
+
|
| 564 |
+
if torch.cuda.is_available():
|
| 565 |
+
return "cuda"
|
| 566 |
+
elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
|
| 567 |
+
return "mps"
|
| 568 |
+
except ImportError:
|
| 569 |
+
pass
|
| 570 |
+
|
| 571 |
+
return "cpu"
|
| 572 |
+
|
| 573 |
+
def _store_in_ccr(
|
| 574 |
+
self,
|
| 575 |
+
original: str,
|
| 576 |
+
compressed: str,
|
| 577 |
+
original_tokens: int,
|
| 578 |
+
) -> str | None:
|
| 579 |
+
"""Store original content in CCR for later retrieval.
|
| 580 |
+
|
| 581 |
+
Args:
|
| 582 |
+
original: Original content before compression.
|
| 583 |
+
compressed: Compressed content.
|
| 584 |
+
original_tokens: Token count of original.
|
| 585 |
+
|
| 586 |
+
Returns:
|
| 587 |
+
Cache key if stored successfully, None otherwise.
|
| 588 |
+
"""
|
| 589 |
+
try:
|
| 590 |
+
from ..cache.compression_store import get_compression_store
|
| 591 |
+
|
| 592 |
+
store = get_compression_store()
|
| 593 |
+
return store.store(
|
| 594 |
+
original,
|
| 595 |
+
compressed,
|
| 596 |
+
original_tokens=original_tokens,
|
| 597 |
+
compressed_tokens=len(compressed.split()),
|
| 598 |
+
compression_strategy="llmlingua2",
|
| 599 |
+
)
|
| 600 |
+
except ImportError:
|
| 601 |
+
return None
|
| 602 |
+
except Exception as e:
|
| 603 |
+
logger.debug("CCR storage failed: %s", e)
|
| 604 |
+
return None
|
| 605 |
+
|
| 606 |
+
|
| 607 |
+
def compress_with_llmlingua(
|
| 608 |
+
content: str,
|
| 609 |
+
compression_rate: float = 0.3,
|
| 610 |
+
context: str = "",
|
| 611 |
+
model_name: str | None = None,
|
| 612 |
+
) -> str:
|
| 613 |
+
"""Convenience function for one-off compression.
|
| 614 |
+
|
| 615 |
+
Args:
|
| 616 |
+
content: Content to compress.
|
| 617 |
+
compression_rate: Target compression rate (0.0-1.0).
|
| 618 |
+
context: Optional context for relevance-aware compression.
|
| 619 |
+
model_name: Optional model name override.
|
| 620 |
+
|
| 621 |
+
Returns:
|
| 622 |
+
Compressed content string.
|
| 623 |
+
|
| 624 |
+
Example:
|
| 625 |
+
>>> compressed = compress_with_llmlingua(long_output, compression_rate=0.2)
|
| 626 |
+
"""
|
| 627 |
+
config = LLMLinguaConfig(target_compression_rate=compression_rate)
|
| 628 |
+
if model_name:
|
| 629 |
+
config.model_name = model_name
|
| 630 |
+
|
| 631 |
+
compressor = LLMLinguaCompressor(config)
|
| 632 |
+
result = compressor.compress(content, context=context)
|
| 633 |
+
return result.compressed
|
|
@@ -64,6 +64,12 @@ proxy = [
|
|
| 64 |
reports = [
|
| 65 |
"jinja2>=3.0.0",
|
| 66 |
]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 67 |
# Development dependencies
|
| 68 |
dev = [
|
| 69 |
"pytest>=7.0.0",
|
|
@@ -76,7 +82,7 @@ dev = [
|
|
| 76 |
]
|
| 77 |
# All optional dependencies
|
| 78 |
all = [
|
| 79 |
-
"headroom-ai[relevance,proxy,reports]",
|
| 80 |
]
|
| 81 |
|
| 82 |
[project.scripts]
|
|
|
|
| 64 |
reports = [
|
| 65 |
"jinja2>=3.0.0",
|
| 66 |
]
|
| 67 |
+
# ML-based compression (LLMLingua-2)
|
| 68 |
+
llmlingua = [
|
| 69 |
+
"llmlingua>=0.2.0",
|
| 70 |
+
"torch>=2.0.0",
|
| 71 |
+
"transformers>=4.30.0",
|
| 72 |
+
]
|
| 73 |
# Development dependencies
|
| 74 |
dev = [
|
| 75 |
"pytest>=7.0.0",
|
|
|
|
| 82 |
]
|
| 83 |
# All optional dependencies
|
| 84 |
all = [
|
| 85 |
+
"headroom-ai[relevance,proxy,reports,llmlingua]",
|
| 86 |
]
|
| 87 |
|
| 88 |
[project.scripts]
|
|
@@ -0,0 +1,458 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Tests for LLMLingua opt-in mechanism in the proxy server.
|
| 2 |
+
|
| 3 |
+
These tests verify:
|
| 4 |
+
- ProxyConfig llmlingua settings
|
| 5 |
+
- LLMLingua transform integration in pipeline
|
| 6 |
+
- Status detection and logging hints
|
| 7 |
+
- CLI flag parsing
|
| 8 |
+
- DevEx: helpful messages when llmlingua unavailable
|
| 9 |
+
"""
|
| 10 |
+
|
| 11 |
+
from unittest.mock import MagicMock, patch
|
| 12 |
+
|
| 13 |
+
import pytest
|
| 14 |
+
|
| 15 |
+
# Skip if fastapi not available
|
| 16 |
+
pytest.importorskip("fastapi")
|
| 17 |
+
|
| 18 |
+
from fastapi.testclient import TestClient
|
| 19 |
+
|
| 20 |
+
from headroom.proxy.server import (
|
| 21 |
+
HeadroomProxy,
|
| 22 |
+
ProxyConfig,
|
| 23 |
+
_get_llmlingua_banner_status,
|
| 24 |
+
create_app,
|
| 25 |
+
)
|
| 26 |
+
from headroom.transforms import _LLMLINGUA_AVAILABLE
|
| 27 |
+
|
| 28 |
+
# =============================================================================
|
| 29 |
+
# Test Fixtures
|
| 30 |
+
# =============================================================================
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
@pytest.fixture
|
| 34 |
+
def base_config():
|
| 35 |
+
"""Base config with optimization disabled for simpler tests."""
|
| 36 |
+
return ProxyConfig(
|
| 37 |
+
optimize=False,
|
| 38 |
+
cache_enabled=False,
|
| 39 |
+
rate_limit_enabled=False,
|
| 40 |
+
cost_tracking_enabled=False,
|
| 41 |
+
)
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
@pytest.fixture
|
| 45 |
+
def llmlingua_config():
|
| 46 |
+
"""Config with LLMLingua enabled."""
|
| 47 |
+
return ProxyConfig(
|
| 48 |
+
optimize=True,
|
| 49 |
+
cache_enabled=False,
|
| 50 |
+
rate_limit_enabled=False,
|
| 51 |
+
cost_tracking_enabled=False,
|
| 52 |
+
llmlingua_enabled=True,
|
| 53 |
+
llmlingua_device="cpu",
|
| 54 |
+
llmlingua_target_rate=0.4,
|
| 55 |
+
)
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
@pytest.fixture
|
| 59 |
+
def client(base_config):
|
| 60 |
+
"""Create test client with base config."""
|
| 61 |
+
app = create_app(base_config)
|
| 62 |
+
with TestClient(app) as client:
|
| 63 |
+
yield client
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
# =============================================================================
|
| 67 |
+
# TestProxyConfigLLMLingua
|
| 68 |
+
# =============================================================================
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
class TestProxyConfigLLMLingua:
|
| 72 |
+
"""Tests for LLMLingua settings in ProxyConfig."""
|
| 73 |
+
|
| 74 |
+
def test_default_llmlingua_disabled(self):
|
| 75 |
+
"""LLMLingua is disabled by default."""
|
| 76 |
+
config = ProxyConfig()
|
| 77 |
+
|
| 78 |
+
assert config.llmlingua_enabled is False
|
| 79 |
+
assert config.llmlingua_device == "auto"
|
| 80 |
+
assert config.llmlingua_target_rate == 0.3
|
| 81 |
+
|
| 82 |
+
def test_llmlingua_can_be_enabled(self):
|
| 83 |
+
"""LLMLingua can be enabled via config."""
|
| 84 |
+
config = ProxyConfig(
|
| 85 |
+
llmlingua_enabled=True,
|
| 86 |
+
llmlingua_device="cuda",
|
| 87 |
+
llmlingua_target_rate=0.5,
|
| 88 |
+
)
|
| 89 |
+
|
| 90 |
+
assert config.llmlingua_enabled is True
|
| 91 |
+
assert config.llmlingua_device == "cuda"
|
| 92 |
+
assert config.llmlingua_target_rate == 0.5
|
| 93 |
+
|
| 94 |
+
def test_llmlingua_device_options(self):
|
| 95 |
+
"""LLMLingua device accepts valid options."""
|
| 96 |
+
for device in ["auto", "cuda", "cpu", "mps"]:
|
| 97 |
+
config = ProxyConfig(llmlingua_device=device)
|
| 98 |
+
assert config.llmlingua_device == device
|
| 99 |
+
|
| 100 |
+
def test_llmlingua_target_rate_range(self):
|
| 101 |
+
"""LLMLingua target rate accepts 0.0-1.0 range."""
|
| 102 |
+
# Low rate (aggressive compression)
|
| 103 |
+
config_low = ProxyConfig(llmlingua_target_rate=0.1)
|
| 104 |
+
assert config_low.llmlingua_target_rate == 0.1
|
| 105 |
+
|
| 106 |
+
# High rate (conservative compression)
|
| 107 |
+
config_high = ProxyConfig(llmlingua_target_rate=0.8)
|
| 108 |
+
assert config_high.llmlingua_target_rate == 0.8
|
| 109 |
+
|
| 110 |
+
|
| 111 |
+
# =============================================================================
|
| 112 |
+
# TestLLMLinguaSetup
|
| 113 |
+
# =============================================================================
|
| 114 |
+
|
| 115 |
+
|
| 116 |
+
class TestLLMLinguaSetup:
|
| 117 |
+
"""Tests for LLMLingua setup in HeadroomProxy."""
|
| 118 |
+
|
| 119 |
+
def test_setup_returns_disabled_when_not_enabled(self, base_config):
|
| 120 |
+
"""Setup returns 'disabled' when llmlingua not enabled and not available."""
|
| 121 |
+
with patch("headroom.proxy.server._LLMLINGUA_AVAILABLE", False):
|
| 122 |
+
proxy = HeadroomProxy(base_config)
|
| 123 |
+
assert proxy._llmlingua_status == "disabled"
|
| 124 |
+
|
| 125 |
+
def test_setup_returns_available_when_installed_but_not_enabled(self, base_config):
|
| 126 |
+
"""Setup returns 'available' when llmlingua installed but not enabled."""
|
| 127 |
+
with patch("headroom.proxy.server._LLMLINGUA_AVAILABLE", True):
|
| 128 |
+
proxy = HeadroomProxy(base_config)
|
| 129 |
+
assert proxy._llmlingua_status == "available"
|
| 130 |
+
|
| 131 |
+
def test_setup_returns_enabled_when_enabled_and_available(self, llmlingua_config):
|
| 132 |
+
"""Setup returns 'enabled' when llmlingua enabled and available."""
|
| 133 |
+
mock_compressor = MagicMock()
|
| 134 |
+
|
| 135 |
+
with patch("headroom.proxy.server._LLMLINGUA_AVAILABLE", True):
|
| 136 |
+
with patch("headroom.proxy.server.LLMLinguaCompressor", mock_compressor):
|
| 137 |
+
with patch("headroom.proxy.server.LLMLinguaConfig"):
|
| 138 |
+
proxy = HeadroomProxy(llmlingua_config)
|
| 139 |
+
assert proxy._llmlingua_status == "enabled"
|
| 140 |
+
|
| 141 |
+
def test_setup_returns_unavailable_when_enabled_but_not_installed(self):
|
| 142 |
+
"""Setup returns 'unavailable' when enabled but llmlingua not installed."""
|
| 143 |
+
config = ProxyConfig(
|
| 144 |
+
llmlingua_enabled=True,
|
| 145 |
+
optimize=False,
|
| 146 |
+
cache_enabled=False,
|
| 147 |
+
rate_limit_enabled=False,
|
| 148 |
+
)
|
| 149 |
+
|
| 150 |
+
with patch("headroom.proxy.server._LLMLINGUA_AVAILABLE", False):
|
| 151 |
+
proxy = HeadroomProxy(config)
|
| 152 |
+
assert proxy._llmlingua_status == "unavailable"
|
| 153 |
+
|
| 154 |
+
def test_llmlingua_compressor_added_to_pipeline(self, llmlingua_config):
|
| 155 |
+
"""LLMLinguaCompressor is added to pipeline when enabled."""
|
| 156 |
+
mock_compressor_class = MagicMock()
|
| 157 |
+
mock_compressor_instance = MagicMock()
|
| 158 |
+
mock_compressor_class.return_value = mock_compressor_instance
|
| 159 |
+
|
| 160 |
+
with patch("headroom.proxy.server._LLMLINGUA_AVAILABLE", True):
|
| 161 |
+
with patch("headroom.proxy.server.LLMLinguaCompressor", mock_compressor_class):
|
| 162 |
+
with patch("headroom.proxy.server.LLMLinguaConfig") as mock_config:
|
| 163 |
+
HeadroomProxy(llmlingua_config)
|
| 164 |
+
|
| 165 |
+
# Verify LLMLinguaCompressor was instantiated
|
| 166 |
+
mock_compressor_class.assert_called_once()
|
| 167 |
+
|
| 168 |
+
# Verify config was passed with correct device and rate
|
| 169 |
+
call_args = mock_config.call_args
|
| 170 |
+
assert call_args.kwargs["device"] == "cpu"
|
| 171 |
+
assert call_args.kwargs["target_compression_rate"] == 0.4
|
| 172 |
+
|
| 173 |
+
def test_llmlingua_not_added_when_disabled(self, base_config):
|
| 174 |
+
"""LLMLinguaCompressor is NOT added when disabled."""
|
| 175 |
+
mock_compressor_class = MagicMock()
|
| 176 |
+
|
| 177 |
+
with patch("headroom.proxy.server._LLMLINGUA_AVAILABLE", True):
|
| 178 |
+
with patch("headroom.proxy.server.LLMLinguaCompressor", mock_compressor_class):
|
| 179 |
+
HeadroomProxy(base_config)
|
| 180 |
+
|
| 181 |
+
# Should NOT be called when disabled
|
| 182 |
+
mock_compressor_class.assert_not_called()
|
| 183 |
+
|
| 184 |
+
|
| 185 |
+
# =============================================================================
|
| 186 |
+
# TestBannerStatus
|
| 187 |
+
# =============================================================================
|
| 188 |
+
|
| 189 |
+
|
| 190 |
+
class TestBannerStatus:
|
| 191 |
+
"""Tests for banner status helper function."""
|
| 192 |
+
|
| 193 |
+
def test_banner_disabled_when_not_available(self):
|
| 194 |
+
"""Banner shows DISABLED when llmlingua not available and not enabled."""
|
| 195 |
+
config = ProxyConfig(llmlingua_enabled=False)
|
| 196 |
+
|
| 197 |
+
with patch("headroom.proxy.server._LLMLINGUA_AVAILABLE", False):
|
| 198 |
+
status = _get_llmlingua_banner_status(config)
|
| 199 |
+
assert status == "DISABLED"
|
| 200 |
+
|
| 201 |
+
def test_banner_available_hint_when_installed(self):
|
| 202 |
+
"""Banner shows availability hint when installed but not enabled."""
|
| 203 |
+
config = ProxyConfig(llmlingua_enabled=False)
|
| 204 |
+
|
| 205 |
+
with patch("headroom.proxy.server._LLMLINGUA_AVAILABLE", True):
|
| 206 |
+
status = _get_llmlingua_banner_status(config)
|
| 207 |
+
assert "available" in status
|
| 208 |
+
assert "--llmlingua" in status
|
| 209 |
+
|
| 210 |
+
def test_banner_enabled_when_active(self):
|
| 211 |
+
"""Banner shows ENABLED with config when active."""
|
| 212 |
+
config = ProxyConfig(
|
| 213 |
+
llmlingua_enabled=True,
|
| 214 |
+
llmlingua_device="cuda",
|
| 215 |
+
llmlingua_target_rate=0.25,
|
| 216 |
+
)
|
| 217 |
+
|
| 218 |
+
with patch("headroom.proxy.server._LLMLINGUA_AVAILABLE", True):
|
| 219 |
+
status = _get_llmlingua_banner_status(config)
|
| 220 |
+
assert "ENABLED" in status
|
| 221 |
+
assert "cuda" in status
|
| 222 |
+
assert "0.25" in status
|
| 223 |
+
|
| 224 |
+
def test_banner_shows_install_hint_when_requested_but_missing(self):
|
| 225 |
+
"""Banner shows install hint when enabled but not installed."""
|
| 226 |
+
config = ProxyConfig(llmlingua_enabled=True)
|
| 227 |
+
|
| 228 |
+
with patch("headroom.proxy.server._LLMLINGUA_AVAILABLE", False):
|
| 229 |
+
status = _get_llmlingua_banner_status(config)
|
| 230 |
+
assert "not installed" in status
|
| 231 |
+
assert "pip install" in status
|
| 232 |
+
|
| 233 |
+
|
| 234 |
+
# =============================================================================
|
| 235 |
+
# TestHealthEndpointWithLLMLingua
|
| 236 |
+
# =============================================================================
|
| 237 |
+
|
| 238 |
+
|
| 239 |
+
class TestHealthEndpointWithLLMLingua:
|
| 240 |
+
"""Tests for health endpoint reflecting LLMLingua status."""
|
| 241 |
+
|
| 242 |
+
def test_health_returns_llmlingua_in_config(self, client):
|
| 243 |
+
"""Health endpoint works regardless of LLMLingua status."""
|
| 244 |
+
response = client.get("/health")
|
| 245 |
+
assert response.status_code == 200
|
| 246 |
+
|
| 247 |
+
data = response.json()
|
| 248 |
+
assert data["status"] == "healthy"
|
| 249 |
+
assert "config" in data
|
| 250 |
+
|
| 251 |
+
|
| 252 |
+
# =============================================================================
|
| 253 |
+
# TestStatsEndpointWithLLMLingua
|
| 254 |
+
# =============================================================================
|
| 255 |
+
|
| 256 |
+
|
| 257 |
+
class TestStatsEndpointWithLLMLingua:
|
| 258 |
+
"""Tests for stats endpoint with LLMLingua integration."""
|
| 259 |
+
|
| 260 |
+
def test_stats_endpoint_works(self, client):
|
| 261 |
+
"""Stats endpoint works with any LLMLingua configuration."""
|
| 262 |
+
response = client.get("/stats")
|
| 263 |
+
assert response.status_code == 200
|
| 264 |
+
|
| 265 |
+
data = response.json()
|
| 266 |
+
assert "requests" in data
|
| 267 |
+
assert "tokens" in data
|
| 268 |
+
|
| 269 |
+
|
| 270 |
+
# =============================================================================
|
| 271 |
+
# TestCLIArguments
|
| 272 |
+
# =============================================================================
|
| 273 |
+
|
| 274 |
+
|
| 275 |
+
class TestCLIArguments:
|
| 276 |
+
"""Tests for CLI argument parsing (without actually running server)."""
|
| 277 |
+
|
| 278 |
+
def test_llmlingua_flag_defaults(self):
|
| 279 |
+
"""Default CLI values for LLMLingua settings."""
|
| 280 |
+
import argparse
|
| 281 |
+
|
| 282 |
+
parser = argparse.ArgumentParser()
|
| 283 |
+
parser.add_argument("--llmlingua", action="store_true")
|
| 284 |
+
parser.add_argument("--llmlingua-device", default="auto")
|
| 285 |
+
parser.add_argument("--llmlingua-rate", type=float, default=0.3)
|
| 286 |
+
|
| 287 |
+
args = parser.parse_args([])
|
| 288 |
+
|
| 289 |
+
assert args.llmlingua is False
|
| 290 |
+
assert args.llmlingua_device == "auto"
|
| 291 |
+
assert args.llmlingua_rate == 0.3
|
| 292 |
+
|
| 293 |
+
def test_llmlingua_flag_enabled(self):
|
| 294 |
+
"""CLI --llmlingua flag enables LLMLingua."""
|
| 295 |
+
import argparse
|
| 296 |
+
|
| 297 |
+
parser = argparse.ArgumentParser()
|
| 298 |
+
parser.add_argument("--llmlingua", action="store_true")
|
| 299 |
+
parser.add_argument("--llmlingua-device", default="auto")
|
| 300 |
+
parser.add_argument("--llmlingua-rate", type=float, default=0.3)
|
| 301 |
+
|
| 302 |
+
args = parser.parse_args(["--llmlingua"])
|
| 303 |
+
|
| 304 |
+
assert args.llmlingua is True
|
| 305 |
+
|
| 306 |
+
def test_llmlingua_device_flag(self):
|
| 307 |
+
"""CLI --llmlingua-device flag sets device."""
|
| 308 |
+
import argparse
|
| 309 |
+
|
| 310 |
+
parser = argparse.ArgumentParser()
|
| 311 |
+
parser.add_argument("--llmlingua-device", default="auto")
|
| 312 |
+
|
| 313 |
+
args = parser.parse_args(["--llmlingua-device", "cuda"])
|
| 314 |
+
|
| 315 |
+
assert args.llmlingua_device == "cuda"
|
| 316 |
+
|
| 317 |
+
def test_llmlingua_rate_flag(self):
|
| 318 |
+
"""CLI --llmlingua-rate flag sets compression rate."""
|
| 319 |
+
import argparse
|
| 320 |
+
|
| 321 |
+
parser = argparse.ArgumentParser()
|
| 322 |
+
parser.add_argument("--llmlingua-rate", type=float, default=0.3)
|
| 323 |
+
|
| 324 |
+
args = parser.parse_args(["--llmlingua-rate", "0.5"])
|
| 325 |
+
|
| 326 |
+
assert args.llmlingua_rate == 0.5
|
| 327 |
+
|
| 328 |
+
|
| 329 |
+
# =============================================================================
|
| 330 |
+
# TestDevExMessages
|
| 331 |
+
# =============================================================================
|
| 332 |
+
|
| 333 |
+
|
| 334 |
+
class TestDevExMessages:
|
| 335 |
+
"""Tests for developer experience messages and hints."""
|
| 336 |
+
|
| 337 |
+
def test_warning_logged_when_enabled_but_unavailable(self, caplog):
|
| 338 |
+
"""Warning is logged when llmlingua enabled but not installed."""
|
| 339 |
+
import logging
|
| 340 |
+
|
| 341 |
+
config = ProxyConfig(
|
| 342 |
+
llmlingua_enabled=True,
|
| 343 |
+
optimize=False,
|
| 344 |
+
cache_enabled=False,
|
| 345 |
+
rate_limit_enabled=False,
|
| 346 |
+
)
|
| 347 |
+
|
| 348 |
+
with caplog.at_level(logging.WARNING):
|
| 349 |
+
with patch("headroom.proxy.server._LLMLINGUA_AVAILABLE", False):
|
| 350 |
+
proxy = HeadroomProxy(config)
|
| 351 |
+
|
| 352 |
+
# Should have logged a warning about missing llmlingua
|
| 353 |
+
assert proxy._llmlingua_status == "unavailable"
|
| 354 |
+
assert any("llmlingua" in r.message.lower() for r in caplog.records)
|
| 355 |
+
assert any("pip install" in r.message for r in caplog.records)
|
| 356 |
+
|
| 357 |
+
|
| 358 |
+
# =============================================================================
|
| 359 |
+
# TestIntegrationWithActualLLMLingua
|
| 360 |
+
# =============================================================================
|
| 361 |
+
|
| 362 |
+
|
| 363 |
+
@pytest.mark.skipif(not _LLMLINGUA_AVAILABLE, reason="llmlingua not installed")
|
| 364 |
+
class TestIntegrationWithActualLLMLingua:
|
| 365 |
+
"""Integration tests that require actual llmlingua installation.
|
| 366 |
+
|
| 367 |
+
These tests verify the full integration path when llmlingua is installed.
|
| 368 |
+
"""
|
| 369 |
+
|
| 370 |
+
def test_proxy_starts_with_llmlingua_enabled(self):
|
| 371 |
+
"""Proxy starts successfully with LLMLingua enabled."""
|
| 372 |
+
config = ProxyConfig(
|
| 373 |
+
llmlingua_enabled=True,
|
| 374 |
+
llmlingua_device="cpu", # CPU for CI/test environments
|
| 375 |
+
llmlingua_target_rate=0.3,
|
| 376 |
+
optimize=True,
|
| 377 |
+
cache_enabled=False,
|
| 378 |
+
rate_limit_enabled=False,
|
| 379 |
+
)
|
| 380 |
+
|
| 381 |
+
# Should not raise
|
| 382 |
+
proxy = HeadroomProxy(config)
|
| 383 |
+
|
| 384 |
+
assert proxy._llmlingua_status == "enabled"
|
| 385 |
+
|
| 386 |
+
def test_app_creates_with_llmlingua(self):
|
| 387 |
+
"""FastAPI app creates successfully with LLMLingua enabled."""
|
| 388 |
+
config = ProxyConfig(
|
| 389 |
+
llmlingua_enabled=True,
|
| 390 |
+
llmlingua_device="cpu",
|
| 391 |
+
optimize=True,
|
| 392 |
+
cache_enabled=False,
|
| 393 |
+
rate_limit_enabled=False,
|
| 394 |
+
)
|
| 395 |
+
|
| 396 |
+
# Should not raise
|
| 397 |
+
app = create_app(config)
|
| 398 |
+
|
| 399 |
+
assert app is not None
|
| 400 |
+
|
| 401 |
+
def test_health_endpoint_with_llmlingua_enabled(self):
|
| 402 |
+
"""Health endpoint works with LLMLingua enabled."""
|
| 403 |
+
config = ProxyConfig(
|
| 404 |
+
llmlingua_enabled=True,
|
| 405 |
+
llmlingua_device="cpu",
|
| 406 |
+
optimize=True,
|
| 407 |
+
cache_enabled=False,
|
| 408 |
+
rate_limit_enabled=False,
|
| 409 |
+
)
|
| 410 |
+
|
| 411 |
+
app = create_app(config)
|
| 412 |
+
with TestClient(app) as client:
|
| 413 |
+
response = client.get("/health")
|
| 414 |
+
assert response.status_code == 200
|
| 415 |
+
|
| 416 |
+
|
| 417 |
+
# =============================================================================
|
| 418 |
+
# TestEdgeCases
|
| 419 |
+
# =============================================================================
|
| 420 |
+
|
| 421 |
+
|
| 422 |
+
class TestEdgeCases:
|
| 423 |
+
"""Edge cases for LLMLingua proxy integration."""
|
| 424 |
+
|
| 425 |
+
def test_multiple_proxy_instances_independent(self):
|
| 426 |
+
"""Multiple proxy instances have independent LLMLingua status."""
|
| 427 |
+
config_enabled = ProxyConfig(
|
| 428 |
+
llmlingua_enabled=True,
|
| 429 |
+
optimize=False,
|
| 430 |
+
cache_enabled=False,
|
| 431 |
+
rate_limit_enabled=False,
|
| 432 |
+
)
|
| 433 |
+
config_disabled = ProxyConfig(
|
| 434 |
+
llmlingua_enabled=False,
|
| 435 |
+
optimize=False,
|
| 436 |
+
cache_enabled=False,
|
| 437 |
+
rate_limit_enabled=False,
|
| 438 |
+
)
|
| 439 |
+
|
| 440 |
+
with patch("headroom.proxy.server._LLMLINGUA_AVAILABLE", True):
|
| 441 |
+
with patch("headroom.proxy.server.LLMLinguaCompressor"):
|
| 442 |
+
with patch("headroom.proxy.server.LLMLinguaConfig"):
|
| 443 |
+
proxy_enabled = HeadroomProxy(config_enabled)
|
| 444 |
+
proxy_disabled = HeadroomProxy(config_disabled)
|
| 445 |
+
|
| 446 |
+
assert proxy_enabled._llmlingua_status == "enabled"
|
| 447 |
+
assert proxy_disabled._llmlingua_status == "available"
|
| 448 |
+
|
| 449 |
+
def test_config_immutable_after_proxy_creation(self, base_config):
|
| 450 |
+
"""Config values are captured at proxy creation time."""
|
| 451 |
+
proxy = HeadroomProxy(base_config)
|
| 452 |
+
|
| 453 |
+
# Modifying config after creation doesn't affect proxy
|
| 454 |
+
# (ProxyConfig is a dataclass, so this tests the pattern)
|
| 455 |
+
original_status = proxy._llmlingua_status
|
| 456 |
+
|
| 457 |
+
# Status should remain unchanged
|
| 458 |
+
assert proxy._llmlingua_status == original_status
|
|
@@ -0,0 +1,941 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Tests for LLMLingua-2 compressor integration.
|
| 2 |
+
|
| 3 |
+
Comprehensive tests covering:
|
| 4 |
+
- LLMLinguaConfig: Configuration validation and defaults
|
| 5 |
+
- LLMLinguaCompressor: Core compression functionality
|
| 6 |
+
- Transform interface: apply(), should_apply() methods
|
| 7 |
+
- Content type detection: JSON, code, plain text
|
| 8 |
+
- CCR integration: Reversible compression storage
|
| 9 |
+
- Edge cases: Empty content, unavailable dependency, fallbacks
|
| 10 |
+
"""
|
| 11 |
+
|
| 12 |
+
import json
|
| 13 |
+
from unittest.mock import MagicMock, patch
|
| 14 |
+
|
| 15 |
+
import pytest
|
| 16 |
+
|
| 17 |
+
from headroom.transforms.llmlingua_compressor import (
|
| 18 |
+
LLMLinguaCompressor,
|
| 19 |
+
LLMLinguaConfig,
|
| 20 |
+
LLMLinguaResult,
|
| 21 |
+
compress_with_llmlingua,
|
| 22 |
+
is_llmlingua_model_loaded,
|
| 23 |
+
unload_llmlingua_model,
|
| 24 |
+
)
|
| 25 |
+
|
| 26 |
+
# Try to import for availability check
|
| 27 |
+
try:
|
| 28 |
+
import llmlingua # noqa: F401
|
| 29 |
+
|
| 30 |
+
LLMLINGUA_INSTALLED = True
|
| 31 |
+
except ImportError:
|
| 32 |
+
LLMLINGUA_INSTALLED = False
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
# =============================================================================
|
| 36 |
+
# Test Fixtures
|
| 37 |
+
# =============================================================================
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
@pytest.fixture
|
| 41 |
+
def default_config():
|
| 42 |
+
"""Default LLMLinguaConfig for testing."""
|
| 43 |
+
return LLMLinguaConfig(
|
| 44 |
+
min_tokens_for_compression=10, # Low threshold for tests
|
| 45 |
+
enable_ccr=False, # Disable CCR for unit tests
|
| 46 |
+
)
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
@pytest.fixture
|
| 50 |
+
def compressor(default_config):
|
| 51 |
+
"""LLMLinguaCompressor instance with default config."""
|
| 52 |
+
return LLMLinguaCompressor(default_config)
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
@pytest.fixture
|
| 56 |
+
def mock_llmlingua():
|
| 57 |
+
"""Mock the llmlingua module and PromptCompressor."""
|
| 58 |
+
mock_compressor = MagicMock()
|
| 59 |
+
mock_compressor._model_name = "test-model"
|
| 60 |
+
|
| 61 |
+
# Default compress_prompt return value
|
| 62 |
+
mock_compressor.compress_prompt.return_value = {
|
| 63 |
+
"compressed_prompt": "compressed content here",
|
| 64 |
+
"origin_tokens": 100,
|
| 65 |
+
"compressed_tokens": 30,
|
| 66 |
+
}
|
| 67 |
+
|
| 68 |
+
with patch(
|
| 69 |
+
"headroom.transforms.llmlingua_compressor._check_llmlingua_available",
|
| 70 |
+
return_value=True,
|
| 71 |
+
):
|
| 72 |
+
with patch(
|
| 73 |
+
"headroom.transforms.llmlingua_compressor._get_llmlingua_compressor",
|
| 74 |
+
return_value=mock_compressor,
|
| 75 |
+
):
|
| 76 |
+
yield mock_compressor
|
| 77 |
+
|
| 78 |
+
|
| 79 |
+
@pytest.fixture
|
| 80 |
+
def tokenizer():
|
| 81 |
+
"""Get a tokenizer for Transform interface tests."""
|
| 82 |
+
from headroom.providers import OpenAIProvider
|
| 83 |
+
from headroom.tokenizer import Tokenizer
|
| 84 |
+
|
| 85 |
+
provider = OpenAIProvider()
|
| 86 |
+
token_counter = provider.get_token_counter("gpt-4o")
|
| 87 |
+
return Tokenizer(token_counter, "gpt-4o")
|
| 88 |
+
|
| 89 |
+
|
| 90 |
+
# =============================================================================
|
| 91 |
+
# Test Data Generators
|
| 92 |
+
# =============================================================================
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
def generate_long_text(n_words: int = 500) -> str:
|
| 96 |
+
"""Generate long text content for compression testing."""
|
| 97 |
+
words = ["the", "quick", "brown", "fox", "jumps", "over", "lazy", "dog"]
|
| 98 |
+
return " ".join(words[i % len(words)] for i in range(n_words))
|
| 99 |
+
|
| 100 |
+
|
| 101 |
+
def generate_long_json(n_items: int = 50) -> str:
|
| 102 |
+
"""Generate long JSON content for compression testing."""
|
| 103 |
+
items = [
|
| 104 |
+
{
|
| 105 |
+
"id": i,
|
| 106 |
+
"name": f"Item {i}",
|
| 107 |
+
"description": f"This is a detailed description for item number {i}",
|
| 108 |
+
"value": i * 10,
|
| 109 |
+
"active": i % 2 == 0,
|
| 110 |
+
}
|
| 111 |
+
for i in range(n_items)
|
| 112 |
+
]
|
| 113 |
+
return json.dumps(items)
|
| 114 |
+
|
| 115 |
+
|
| 116 |
+
def generate_long_code(n_functions: int = 20) -> str:
|
| 117 |
+
"""Generate Python code content for compression testing."""
|
| 118 |
+
lines = ['"""Module with many functions."""', "", "import os", "from typing import Any", ""]
|
| 119 |
+
for i in range(n_functions):
|
| 120 |
+
lines.extend(
|
| 121 |
+
[
|
| 122 |
+
f"def function_{i}(arg: Any) -> str:",
|
| 123 |
+
f' """Process argument {i}."""',
|
| 124 |
+
" result = str(arg)",
|
| 125 |
+
f' return f"Function {i}: {{result}}"',
|
| 126 |
+
"",
|
| 127 |
+
]
|
| 128 |
+
)
|
| 129 |
+
return "\n".join(lines)
|
| 130 |
+
|
| 131 |
+
|
| 132 |
+
# =============================================================================
|
| 133 |
+
# TestLLMLinguaConfig
|
| 134 |
+
# =============================================================================
|
| 135 |
+
|
| 136 |
+
|
| 137 |
+
class TestLLMLinguaConfig:
|
| 138 |
+
"""Tests for LLMLinguaConfig dataclass."""
|
| 139 |
+
|
| 140 |
+
def test_default_values(self):
|
| 141 |
+
"""Default config values are sensible."""
|
| 142 |
+
config = LLMLinguaConfig()
|
| 143 |
+
|
| 144 |
+
assert config.model_name == "microsoft/llmlingua-2-xlm-roberta-large-meetingbank"
|
| 145 |
+
assert config.device == "auto"
|
| 146 |
+
assert config.target_compression_rate == 0.3
|
| 147 |
+
assert config.min_tokens_for_compression == 100
|
| 148 |
+
assert config.enable_ccr is True
|
| 149 |
+
assert config.drop_consecutive is True
|
| 150 |
+
|
| 151 |
+
def test_custom_values(self):
|
| 152 |
+
"""Custom config values are applied."""
|
| 153 |
+
config = LLMLinguaConfig(
|
| 154 |
+
model_name="custom/model",
|
| 155 |
+
device="cuda",
|
| 156 |
+
target_compression_rate=0.5,
|
| 157 |
+
min_tokens_for_compression=50,
|
| 158 |
+
force_tokens=["important", "keep"],
|
| 159 |
+
)
|
| 160 |
+
|
| 161 |
+
assert config.model_name == "custom/model"
|
| 162 |
+
assert config.device == "cuda"
|
| 163 |
+
assert config.target_compression_rate == 0.5
|
| 164 |
+
assert config.min_tokens_for_compression == 50
|
| 165 |
+
assert "important" in config.force_tokens
|
| 166 |
+
|
| 167 |
+
def test_content_type_rates(self):
|
| 168 |
+
"""Different content types have appropriate compression rates."""
|
| 169 |
+
config = LLMLinguaConfig()
|
| 170 |
+
|
| 171 |
+
# Code should be more conservative
|
| 172 |
+
assert config.code_compression_rate > config.text_compression_rate
|
| 173 |
+
# JSON should be between code and text
|
| 174 |
+
assert config.json_compression_rate > config.text_compression_rate
|
| 175 |
+
assert config.json_compression_rate < config.code_compression_rate
|
| 176 |
+
|
| 177 |
+
|
| 178 |
+
# =============================================================================
|
| 179 |
+
# TestLLMLinguaResult
|
| 180 |
+
# =============================================================================
|
| 181 |
+
|
| 182 |
+
|
| 183 |
+
class TestLLMLinguaResult:
|
| 184 |
+
"""Tests for LLMLinguaResult dataclass."""
|
| 185 |
+
|
| 186 |
+
def test_tokens_saved(self):
|
| 187 |
+
"""tokens_saved property calculates correctly."""
|
| 188 |
+
result = LLMLinguaResult(
|
| 189 |
+
compressed="short",
|
| 190 |
+
original="long content here",
|
| 191 |
+
original_tokens=100,
|
| 192 |
+
compressed_tokens=30,
|
| 193 |
+
compression_ratio=0.3,
|
| 194 |
+
)
|
| 195 |
+
|
| 196 |
+
assert result.tokens_saved == 70
|
| 197 |
+
|
| 198 |
+
def test_tokens_saved_no_negative(self):
|
| 199 |
+
"""tokens_saved never returns negative."""
|
| 200 |
+
result = LLMLinguaResult(
|
| 201 |
+
compressed="expanded content",
|
| 202 |
+
original="short",
|
| 203 |
+
original_tokens=10,
|
| 204 |
+
compressed_tokens=20, # Expanded (unusual case)
|
| 205 |
+
compression_ratio=2.0,
|
| 206 |
+
)
|
| 207 |
+
|
| 208 |
+
assert result.tokens_saved == 0
|
| 209 |
+
|
| 210 |
+
def test_savings_percentage(self):
|
| 211 |
+
"""savings_percentage property calculates correctly."""
|
| 212 |
+
result = LLMLinguaResult(
|
| 213 |
+
compressed="short",
|
| 214 |
+
original="long content",
|
| 215 |
+
original_tokens=100,
|
| 216 |
+
compressed_tokens=25,
|
| 217 |
+
compression_ratio=0.25,
|
| 218 |
+
)
|
| 219 |
+
|
| 220 |
+
assert result.savings_percentage == 75.0
|
| 221 |
+
|
| 222 |
+
def test_savings_percentage_zero_original(self):
|
| 223 |
+
"""savings_percentage handles zero original tokens."""
|
| 224 |
+
result = LLMLinguaResult(
|
| 225 |
+
compressed="",
|
| 226 |
+
original="",
|
| 227 |
+
original_tokens=0,
|
| 228 |
+
compressed_tokens=0,
|
| 229 |
+
compression_ratio=1.0,
|
| 230 |
+
)
|
| 231 |
+
|
| 232 |
+
assert result.savings_percentage == 0.0
|
| 233 |
+
|
| 234 |
+
|
| 235 |
+
# =============================================================================
|
| 236 |
+
# TestLLMLinguaCompressor
|
| 237 |
+
# =============================================================================
|
| 238 |
+
|
| 239 |
+
|
| 240 |
+
class TestLLMLinguaCompressor:
|
| 241 |
+
"""Tests for LLMLinguaCompressor core functionality."""
|
| 242 |
+
|
| 243 |
+
def test_init_with_default_config(self):
|
| 244 |
+
"""Compressor initializes with default config."""
|
| 245 |
+
compressor = LLMLinguaCompressor()
|
| 246 |
+
|
| 247 |
+
assert compressor.config is not None
|
| 248 |
+
assert compressor.config.model_name is not None
|
| 249 |
+
|
| 250 |
+
def test_init_with_custom_config(self, default_config):
|
| 251 |
+
"""Compressor initializes with custom config."""
|
| 252 |
+
compressor = LLMLinguaCompressor(default_config)
|
| 253 |
+
|
| 254 |
+
assert compressor.config == default_config
|
| 255 |
+
|
| 256 |
+
def test_compress_returns_result_when_unavailable(self, compressor):
|
| 257 |
+
"""Compress returns passthrough result when llmlingua unavailable."""
|
| 258 |
+
with patch(
|
| 259 |
+
"headroom.transforms.llmlingua_compressor._check_llmlingua_available",
|
| 260 |
+
return_value=False,
|
| 261 |
+
):
|
| 262 |
+
content = generate_long_text(100)
|
| 263 |
+
result = compressor.compress(content)
|
| 264 |
+
|
| 265 |
+
# Should return unchanged content
|
| 266 |
+
assert result.compressed == content
|
| 267 |
+
assert result.compression_ratio == 1.0
|
| 268 |
+
|
| 269 |
+
def test_compress_skips_small_content(self, compressor):
|
| 270 |
+
"""Small content is not compressed."""
|
| 271 |
+
small_content = "short text"
|
| 272 |
+
result = compressor.compress(small_content)
|
| 273 |
+
|
| 274 |
+
assert result.compressed == small_content
|
| 275 |
+
assert result.compression_ratio == 1.0
|
| 276 |
+
|
| 277 |
+
def test_compress_with_llmlingua(self, default_config, mock_llmlingua):
|
| 278 |
+
"""Compression uses llmlingua when available."""
|
| 279 |
+
compressor = LLMLinguaCompressor(default_config)
|
| 280 |
+
content = generate_long_text(200)
|
| 281 |
+
|
| 282 |
+
result = compressor.compress(content)
|
| 283 |
+
|
| 284 |
+
# Should have called compress_prompt
|
| 285 |
+
mock_llmlingua.compress_prompt.assert_called_once()
|
| 286 |
+
assert result.compressed == "compressed content here"
|
| 287 |
+
assert result.compression_ratio < 1.0
|
| 288 |
+
|
| 289 |
+
def test_compress_with_context(self, default_config, mock_llmlingua):
|
| 290 |
+
"""Context words are used as force tokens."""
|
| 291 |
+
compressor = LLMLinguaCompressor(default_config)
|
| 292 |
+
content = generate_long_text(200)
|
| 293 |
+
context = "important keywords here"
|
| 294 |
+
|
| 295 |
+
compressor.compress(content, context=context)
|
| 296 |
+
|
| 297 |
+
# Check force_tokens includes context words
|
| 298 |
+
call_args = mock_llmlingua.compress_prompt.call_args
|
| 299 |
+
force_tokens = call_args.kwargs.get("force_tokens", [])
|
| 300 |
+
# Should include context words longer than 3 chars
|
| 301 |
+
assert "important" in force_tokens or "keywords" in force_tokens
|
| 302 |
+
|
| 303 |
+
def test_compress_handles_exception(self, default_config, mock_llmlingua):
|
| 304 |
+
"""Exceptions from llmlingua are handled gracefully."""
|
| 305 |
+
mock_llmlingua.compress_prompt.side_effect = RuntimeError("Model error")
|
| 306 |
+
|
| 307 |
+
compressor = LLMLinguaCompressor(default_config)
|
| 308 |
+
content = generate_long_text(200)
|
| 309 |
+
|
| 310 |
+
result = compressor.compress(content)
|
| 311 |
+
|
| 312 |
+
# Should return original content on error
|
| 313 |
+
assert result.compressed == content
|
| 314 |
+
assert result.compression_ratio == 1.0
|
| 315 |
+
|
| 316 |
+
|
| 317 |
+
# =============================================================================
|
| 318 |
+
# TestContentTypeDetection
|
| 319 |
+
# =============================================================================
|
| 320 |
+
|
| 321 |
+
|
| 322 |
+
class TestContentTypeDetection:
|
| 323 |
+
"""Tests for content type auto-detection."""
|
| 324 |
+
|
| 325 |
+
def test_detect_json_content(self, default_config, mock_llmlingua):
|
| 326 |
+
"""JSON content is detected and uses JSON compression rate."""
|
| 327 |
+
compressor = LLMLinguaCompressor(default_config)
|
| 328 |
+
|
| 329 |
+
rate = compressor._get_compression_rate(generate_long_json(50), None)
|
| 330 |
+
|
| 331 |
+
assert rate == default_config.json_compression_rate
|
| 332 |
+
|
| 333 |
+
def test_detect_code_content(self, default_config, mock_llmlingua):
|
| 334 |
+
"""Code content is detected and uses code compression rate."""
|
| 335 |
+
compressor = LLMLinguaCompressor(default_config)
|
| 336 |
+
code = generate_long_code(20)
|
| 337 |
+
|
| 338 |
+
rate = compressor._get_compression_rate(code, None)
|
| 339 |
+
|
| 340 |
+
assert rate == default_config.code_compression_rate
|
| 341 |
+
|
| 342 |
+
def test_detect_plain_text(self, default_config, mock_llmlingua):
|
| 343 |
+
"""Plain text uses text compression rate."""
|
| 344 |
+
compressor = LLMLinguaCompressor(default_config)
|
| 345 |
+
text = generate_long_text(200)
|
| 346 |
+
|
| 347 |
+
rate = compressor._get_compression_rate(text, None)
|
| 348 |
+
|
| 349 |
+
assert rate == default_config.text_compression_rate
|
| 350 |
+
|
| 351 |
+
def test_explicit_content_type(self, default_config, mock_llmlingua):
|
| 352 |
+
"""Explicit content_type overrides detection."""
|
| 353 |
+
compressor = LLMLinguaCompressor(default_config)
|
| 354 |
+
# JSON-looking content but marked as text
|
| 355 |
+
json_content = generate_long_json(50)
|
| 356 |
+
|
| 357 |
+
rate = compressor._get_compression_rate(json_content, content_type="text")
|
| 358 |
+
|
| 359 |
+
assert rate == default_config.text_compression_rate
|
| 360 |
+
|
| 361 |
+
def test_looks_like_json_detection(self, default_config):
|
| 362 |
+
"""JSON detection works for arrays and objects."""
|
| 363 |
+
compressor = LLMLinguaCompressor(default_config)
|
| 364 |
+
|
| 365 |
+
assert compressor._looks_like_json('[{"key": "value"}]')
|
| 366 |
+
assert compressor._looks_like_json('{"key": "value"}')
|
| 367 |
+
assert not compressor._looks_like_json("plain text")
|
| 368 |
+
assert not compressor._looks_like_json("def function():")
|
| 369 |
+
|
| 370 |
+
def test_looks_like_code_detection(self, default_config):
|
| 371 |
+
"""Code detection works for common patterns."""
|
| 372 |
+
compressor = LLMLinguaCompressor(default_config)
|
| 373 |
+
|
| 374 |
+
assert compressor._looks_like_code("def function():")
|
| 375 |
+
assert compressor._looks_like_code("class MyClass:")
|
| 376 |
+
assert compressor._looks_like_code("import os")
|
| 377 |
+
assert compressor._looks_like_code("function test() {")
|
| 378 |
+
assert compressor._looks_like_code("const x = 5")
|
| 379 |
+
assert not compressor._looks_like_code("plain text content")
|
| 380 |
+
|
| 381 |
+
|
| 382 |
+
# =============================================================================
|
| 383 |
+
# TestTransformInterface
|
| 384 |
+
# =============================================================================
|
| 385 |
+
|
| 386 |
+
|
| 387 |
+
class TestTransformInterface:
|
| 388 |
+
"""Tests for Transform interface (apply, should_apply)."""
|
| 389 |
+
|
| 390 |
+
def test_should_apply_returns_false_when_unavailable(self, compressor, tokenizer):
|
| 391 |
+
"""should_apply returns False when llmlingua unavailable."""
|
| 392 |
+
messages = [{"role": "user", "content": generate_long_text(200)}]
|
| 393 |
+
|
| 394 |
+
with patch(
|
| 395 |
+
"headroom.transforms.llmlingua_compressor._check_llmlingua_available",
|
| 396 |
+
return_value=False,
|
| 397 |
+
):
|
| 398 |
+
assert not compressor.should_apply(messages, tokenizer)
|
| 399 |
+
|
| 400 |
+
def test_should_apply_returns_false_for_small_content(self, default_config, tokenizer):
|
| 401 |
+
"""should_apply returns False for small content."""
|
| 402 |
+
config = LLMLinguaConfig(min_tokens_for_compression=1000)
|
| 403 |
+
compressor = LLMLinguaCompressor(config)
|
| 404 |
+
messages = [{"role": "user", "content": "small"}]
|
| 405 |
+
|
| 406 |
+
with patch(
|
| 407 |
+
"headroom.transforms.llmlingua_compressor._check_llmlingua_available",
|
| 408 |
+
return_value=True,
|
| 409 |
+
):
|
| 410 |
+
assert not compressor.should_apply(messages, tokenizer)
|
| 411 |
+
|
| 412 |
+
def test_should_apply_returns_true_for_large_content(self, default_config, tokenizer):
|
| 413 |
+
"""should_apply returns True for large content."""
|
| 414 |
+
compressor = LLMLinguaCompressor(default_config)
|
| 415 |
+
messages = [{"role": "user", "content": generate_long_text(500)}]
|
| 416 |
+
|
| 417 |
+
with patch(
|
| 418 |
+
"headroom.transforms.llmlingua_compressor._check_llmlingua_available",
|
| 419 |
+
return_value=True,
|
| 420 |
+
):
|
| 421 |
+
assert compressor.should_apply(messages, tokenizer)
|
| 422 |
+
|
| 423 |
+
def test_apply_compresses_tool_messages(self, default_config, tokenizer, mock_llmlingua):
|
| 424 |
+
"""apply() compresses tool message content."""
|
| 425 |
+
compressor = LLMLinguaCompressor(default_config)
|
| 426 |
+
tool_content = generate_long_json(100)
|
| 427 |
+
messages = [
|
| 428 |
+
{"role": "user", "content": "Get data"},
|
| 429 |
+
{"role": "tool", "tool_call_id": "call_1", "content": tool_content},
|
| 430 |
+
]
|
| 431 |
+
|
| 432 |
+
result = compressor.apply(messages, tokenizer)
|
| 433 |
+
|
| 434 |
+
# Tool content should be compressed
|
| 435 |
+
assert result.messages[1]["content"] != tool_content
|
| 436 |
+
assert "compressed content here" in result.messages[1]["content"]
|
| 437 |
+
assert len(result.transforms_applied) > 0
|
| 438 |
+
|
| 439 |
+
def test_apply_compresses_long_assistant_messages(
|
| 440 |
+
self, default_config, tokenizer, mock_llmlingua
|
| 441 |
+
):
|
| 442 |
+
"""apply() compresses long assistant messages."""
|
| 443 |
+
compressor = LLMLinguaCompressor(default_config)
|
| 444 |
+
long_content = generate_long_text(1000)
|
| 445 |
+
messages = [
|
| 446 |
+
{"role": "user", "content": "Tell me a story"},
|
| 447 |
+
{"role": "assistant", "content": long_content},
|
| 448 |
+
]
|
| 449 |
+
|
| 450 |
+
result = compressor.apply(messages, tokenizer)
|
| 451 |
+
|
| 452 |
+
# Assistant content should be compressed (>500 chars)
|
| 453 |
+
assert result.messages[1]["content"] != long_content
|
| 454 |
+
|
| 455 |
+
def test_apply_passes_through_short_messages(self, default_config, tokenizer, mock_llmlingua):
|
| 456 |
+
"""apply() passes through short messages unchanged."""
|
| 457 |
+
compressor = LLMLinguaCompressor(default_config)
|
| 458 |
+
messages = [
|
| 459 |
+
{"role": "user", "content": "Hello"},
|
| 460 |
+
{"role": "assistant", "content": "Hi there!"},
|
| 461 |
+
]
|
| 462 |
+
|
| 463 |
+
result = compressor.apply(messages, tokenizer)
|
| 464 |
+
|
| 465 |
+
# Short messages unchanged
|
| 466 |
+
assert result.messages[0]["content"] == "Hello"
|
| 467 |
+
assert result.messages[1]["content"] == "Hi there!"
|
| 468 |
+
|
| 469 |
+
def test_apply_tracks_transform_metadata(self, default_config, tokenizer, mock_llmlingua):
|
| 470 |
+
"""apply() returns proper TransformResult metadata."""
|
| 471 |
+
compressor = LLMLinguaCompressor(default_config)
|
| 472 |
+
messages = [
|
| 473 |
+
{"role": "tool", "tool_call_id": "call_1", "content": generate_long_json(100)},
|
| 474 |
+
]
|
| 475 |
+
|
| 476 |
+
result = compressor.apply(messages, tokenizer)
|
| 477 |
+
|
| 478 |
+
assert result.tokens_before > 0
|
| 479 |
+
assert result.tokens_after > 0
|
| 480 |
+
assert len(result.transforms_applied) > 0
|
| 481 |
+
assert "llmlingua" in result.transforms_applied[0]
|
| 482 |
+
|
| 483 |
+
def test_apply_adds_warning_when_unavailable(self, default_config, tokenizer):
|
| 484 |
+
"""apply() adds warning when llmlingua unavailable."""
|
| 485 |
+
compressor = LLMLinguaCompressor(default_config)
|
| 486 |
+
messages = [{"role": "user", "content": "test"}]
|
| 487 |
+
|
| 488 |
+
with patch(
|
| 489 |
+
"headroom.transforms.llmlingua_compressor._check_llmlingua_available",
|
| 490 |
+
return_value=False,
|
| 491 |
+
):
|
| 492 |
+
result = compressor.apply(messages, tokenizer)
|
| 493 |
+
|
| 494 |
+
assert len(result.warnings) > 0
|
| 495 |
+
assert "llmlingua" in result.warnings[0].lower()
|
| 496 |
+
|
| 497 |
+
|
| 498 |
+
# =============================================================================
|
| 499 |
+
# TestDeviceResolution
|
| 500 |
+
# =============================================================================
|
| 501 |
+
|
| 502 |
+
|
| 503 |
+
class TestDeviceResolution:
|
| 504 |
+
"""Tests for device resolution logic."""
|
| 505 |
+
|
| 506 |
+
def test_resolve_explicit_device(self, default_config):
|
| 507 |
+
"""Explicit device is returned unchanged."""
|
| 508 |
+
config = LLMLinguaConfig(device="cuda")
|
| 509 |
+
compressor = LLMLinguaCompressor(config)
|
| 510 |
+
|
| 511 |
+
assert compressor._resolve_device() == "cuda"
|
| 512 |
+
|
| 513 |
+
def test_resolve_auto_to_cpu_no_torch(self, default_config):
|
| 514 |
+
"""Auto resolves to CPU when torch unavailable."""
|
| 515 |
+
config = LLMLinguaConfig(device="auto")
|
| 516 |
+
compressor = LLMLinguaCompressor(config)
|
| 517 |
+
|
| 518 |
+
with patch.dict("sys.modules", {"torch": None}):
|
| 519 |
+
with patch(
|
| 520 |
+
"headroom.transforms.llmlingua_compressor.LLMLinguaCompressor._resolve_device"
|
| 521 |
+
) as mock_resolve:
|
| 522 |
+
mock_resolve.return_value = "cpu"
|
| 523 |
+
assert compressor._resolve_device() == "cpu"
|
| 524 |
+
|
| 525 |
+
|
| 526 |
+
# =============================================================================
|
| 527 |
+
# TestCCRIntegration
|
| 528 |
+
# =============================================================================
|
| 529 |
+
|
| 530 |
+
|
| 531 |
+
class TestCCRIntegration:
|
| 532 |
+
"""Tests for CCR (Compress-Cache-Retrieve) integration."""
|
| 533 |
+
|
| 534 |
+
def test_ccr_stores_original(self, mock_llmlingua):
|
| 535 |
+
"""Compressed content is stored in CCR when enabled."""
|
| 536 |
+
config = LLMLinguaConfig(
|
| 537 |
+
enable_ccr=True,
|
| 538 |
+
min_tokens_for_compression=10,
|
| 539 |
+
)
|
| 540 |
+
compressor = LLMLinguaCompressor(config)
|
| 541 |
+
content = generate_long_text(200)
|
| 542 |
+
|
| 543 |
+
with patch(
|
| 544 |
+
"headroom.transforms.llmlingua_compressor.LLMLinguaCompressor._store_in_ccr"
|
| 545 |
+
) as mock_store:
|
| 546 |
+
mock_store.return_value = "hash123"
|
| 547 |
+
|
| 548 |
+
result = compressor.compress(content)
|
| 549 |
+
|
| 550 |
+
mock_store.assert_called_once()
|
| 551 |
+
assert result.cache_key == "hash123"
|
| 552 |
+
|
| 553 |
+
def test_ccr_skipped_when_disabled(self, mock_llmlingua):
|
| 554 |
+
"""CCR is not used when disabled in config."""
|
| 555 |
+
config = LLMLinguaConfig(
|
| 556 |
+
enable_ccr=False,
|
| 557 |
+
min_tokens_for_compression=10,
|
| 558 |
+
)
|
| 559 |
+
compressor = LLMLinguaCompressor(config)
|
| 560 |
+
content = generate_long_text(200)
|
| 561 |
+
|
| 562 |
+
with patch(
|
| 563 |
+
"headroom.transforms.llmlingua_compressor.LLMLinguaCompressor._store_in_ccr"
|
| 564 |
+
) as mock_store:
|
| 565 |
+
result = compressor.compress(content)
|
| 566 |
+
|
| 567 |
+
mock_store.assert_not_called()
|
| 568 |
+
assert result.cache_key is None
|
| 569 |
+
|
| 570 |
+
def test_ccr_handles_storage_error(self, mock_llmlingua):
|
| 571 |
+
"""CCR storage errors are handled gracefully."""
|
| 572 |
+
config = LLMLinguaConfig(
|
| 573 |
+
enable_ccr=True,
|
| 574 |
+
min_tokens_for_compression=10,
|
| 575 |
+
)
|
| 576 |
+
compressor = LLMLinguaCompressor(config)
|
| 577 |
+
content = generate_long_text(200)
|
| 578 |
+
|
| 579 |
+
with patch(
|
| 580 |
+
"headroom.transforms.llmlingua_compressor.LLMLinguaCompressor._store_in_ccr"
|
| 581 |
+
) as mock_store:
|
| 582 |
+
# Return None to simulate storage failure (internal error handling)
|
| 583 |
+
mock_store.return_value = None
|
| 584 |
+
|
| 585 |
+
# Should not raise
|
| 586 |
+
result = compressor.compress(content)
|
| 587 |
+
|
| 588 |
+
# Storage failed, so cache_key should be None
|
| 589 |
+
assert result.cache_key is None
|
| 590 |
+
|
| 591 |
+
|
| 592 |
+
# =============================================================================
|
| 593 |
+
# TestConvenienceFunction
|
| 594 |
+
# =============================================================================
|
| 595 |
+
|
| 596 |
+
|
| 597 |
+
class TestConvenienceFunction:
|
| 598 |
+
"""Tests for compress_with_llmlingua convenience function."""
|
| 599 |
+
|
| 600 |
+
def test_compress_with_llmlingua_basic(self, mock_llmlingua):
|
| 601 |
+
"""compress_with_llmlingua works with default settings."""
|
| 602 |
+
content = generate_long_text(200)
|
| 603 |
+
|
| 604 |
+
# Disable CCR for this test to avoid hash suffix
|
| 605 |
+
with patch(
|
| 606 |
+
"headroom.transforms.llmlingua_compressor.LLMLinguaCompressor._store_in_ccr"
|
| 607 |
+
) as mock_store:
|
| 608 |
+
mock_store.return_value = None
|
| 609 |
+
result = compress_with_llmlingua(content)
|
| 610 |
+
|
| 611 |
+
# Should contain the compressed content
|
| 612 |
+
assert "compressed content here" in result
|
| 613 |
+
|
| 614 |
+
def test_compress_with_llmlingua_custom_rate(self, mock_llmlingua):
|
| 615 |
+
"""compress_with_llmlingua accepts custom compression rate."""
|
| 616 |
+
content = generate_long_text(200)
|
| 617 |
+
|
| 618 |
+
compress_with_llmlingua(content, compression_rate=0.5)
|
| 619 |
+
|
| 620 |
+
# Verify compress_prompt was called
|
| 621 |
+
mock_llmlingua.compress_prompt.assert_called()
|
| 622 |
+
|
| 623 |
+
def test_compress_with_llmlingua_with_context(self, mock_llmlingua):
|
| 624 |
+
"""compress_with_llmlingua passes context."""
|
| 625 |
+
content = generate_long_text(200)
|
| 626 |
+
context = "important keywords"
|
| 627 |
+
|
| 628 |
+
compress_with_llmlingua(content, context=context)
|
| 629 |
+
|
| 630 |
+
call_args = mock_llmlingua.compress_prompt.call_args
|
| 631 |
+
force_tokens = call_args.kwargs.get("force_tokens", [])
|
| 632 |
+
# Context words should be in force_tokens
|
| 633 |
+
assert any("important" in str(t) for t in force_tokens) or len(force_tokens) > 0
|
| 634 |
+
|
| 635 |
+
|
| 636 |
+
# =============================================================================
|
| 637 |
+
# TestEdgeCases
|
| 638 |
+
# =============================================================================
|
| 639 |
+
|
| 640 |
+
|
| 641 |
+
class TestEdgeCases:
|
| 642 |
+
"""Edge case tests for LLMLingua compressor."""
|
| 643 |
+
|
| 644 |
+
def test_empty_content(self, compressor):
|
| 645 |
+
"""Empty content is handled gracefully."""
|
| 646 |
+
result = compressor.compress("")
|
| 647 |
+
|
| 648 |
+
assert result.compressed == ""
|
| 649 |
+
assert result.compression_ratio == 1.0
|
| 650 |
+
|
| 651 |
+
def test_whitespace_only_content(self, compressor):
|
| 652 |
+
"""Whitespace-only content is handled gracefully."""
|
| 653 |
+
result = compressor.compress(" \n\t\n ")
|
| 654 |
+
|
| 655 |
+
assert result.compression_ratio == 1.0
|
| 656 |
+
|
| 657 |
+
def test_unicode_content(self, default_config, mock_llmlingua):
|
| 658 |
+
"""Unicode content is handled correctly."""
|
| 659 |
+
mock_llmlingua.compress_prompt.return_value = {
|
| 660 |
+
"compressed_prompt": "compressed \u4e2d\u6587 content",
|
| 661 |
+
"origin_tokens": 100,
|
| 662 |
+
"compressed_tokens": 30,
|
| 663 |
+
}
|
| 664 |
+
|
| 665 |
+
compressor = LLMLinguaCompressor(default_config)
|
| 666 |
+
content = "\u4e2d\u6587 \u65e5\u672c\u8a9e " * 100 # Chinese/Japanese text
|
| 667 |
+
|
| 668 |
+
result = compressor.compress(content)
|
| 669 |
+
|
| 670 |
+
assert "\u4e2d\u6587" in result.compressed
|
| 671 |
+
|
| 672 |
+
def test_very_long_content(self, default_config, mock_llmlingua):
|
| 673 |
+
"""Very long content is compressed."""
|
| 674 |
+
compressor = LLMLinguaCompressor(default_config)
|
| 675 |
+
content = generate_long_text(10000)
|
| 676 |
+
|
| 677 |
+
compressor.compress(content)
|
| 678 |
+
|
| 679 |
+
mock_llmlingua.compress_prompt.assert_called_once()
|
| 680 |
+
|
| 681 |
+
def test_mixed_content_types(self, default_config, mock_llmlingua):
|
| 682 |
+
"""Mixed content (JSON with text) is handled."""
|
| 683 |
+
compressor = LLMLinguaCompressor(default_config)
|
| 684 |
+
# JSON-like but with extra text
|
| 685 |
+
content = 'Some preamble text\n{"key": "value"}\nMore text after'
|
| 686 |
+
|
| 687 |
+
# Should not crash
|
| 688 |
+
result = compressor.compress(content)
|
| 689 |
+
assert result is not None
|
| 690 |
+
|
| 691 |
+
def test_malformed_json_content(self, default_config, mock_llmlingua):
|
| 692 |
+
"""Malformed JSON is treated as text."""
|
| 693 |
+
compressor = LLMLinguaCompressor(default_config)
|
| 694 |
+
content = "{malformed: json, missing quotes" * 50
|
| 695 |
+
|
| 696 |
+
rate = compressor._get_compression_rate(content, None)
|
| 697 |
+
|
| 698 |
+
# Should not detect as JSON
|
| 699 |
+
assert rate == default_config.text_compression_rate
|
| 700 |
+
|
| 701 |
+
def test_force_tokens_list_handling(self, default_config, mock_llmlingua):
|
| 702 |
+
"""Force tokens list is properly passed."""
|
| 703 |
+
config = LLMLinguaConfig(
|
| 704 |
+
force_tokens=["keep", "these", "tokens"],
|
| 705 |
+
min_tokens_for_compression=10,
|
| 706 |
+
)
|
| 707 |
+
compressor = LLMLinguaCompressor(config)
|
| 708 |
+
content = generate_long_text(200)
|
| 709 |
+
|
| 710 |
+
compressor.compress(content)
|
| 711 |
+
|
| 712 |
+
call_args = mock_llmlingua.compress_prompt.call_args
|
| 713 |
+
force_tokens = call_args.kwargs.get("force_tokens", [])
|
| 714 |
+
assert "keep" in force_tokens
|
| 715 |
+
assert "these" in force_tokens
|
| 716 |
+
assert "tokens" in force_tokens
|
| 717 |
+
|
| 718 |
+
|
| 719 |
+
# =============================================================================
|
| 720 |
+
# Integration Tests (only run if llmlingua is installed)
|
| 721 |
+
# =============================================================================
|
| 722 |
+
|
| 723 |
+
|
| 724 |
+
@pytest.mark.skipif(not LLMLINGUA_INSTALLED, reason="llmlingua not installed")
|
| 725 |
+
class TestLLMLinguaIntegration:
|
| 726 |
+
"""Integration tests that require actual llmlingua installation.
|
| 727 |
+
|
| 728 |
+
These tests verify the actual compression behavior and should be run
|
| 729 |
+
in environments where llmlingua is installed.
|
| 730 |
+
"""
|
| 731 |
+
|
| 732 |
+
def test_actual_compression(self):
|
| 733 |
+
"""Test actual compression with real llmlingua."""
|
| 734 |
+
config = LLMLinguaConfig(
|
| 735 |
+
target_compression_rate=0.3,
|
| 736 |
+
min_tokens_for_compression=50,
|
| 737 |
+
enable_ccr=False,
|
| 738 |
+
)
|
| 739 |
+
compressor = LLMLinguaCompressor(config)
|
| 740 |
+
content = generate_long_text(500)
|
| 741 |
+
|
| 742 |
+
result = compressor.compress(content)
|
| 743 |
+
|
| 744 |
+
# Should achieve actual compression
|
| 745 |
+
assert result.compression_ratio < 1.0
|
| 746 |
+
assert result.tokens_saved > 0
|
| 747 |
+
assert len(result.compressed) < len(content)
|
| 748 |
+
|
| 749 |
+
def test_actual_json_compression(self):
|
| 750 |
+
"""Test JSON content compression with real llmlingua."""
|
| 751 |
+
config = LLMLinguaConfig(
|
| 752 |
+
target_compression_rate=0.35,
|
| 753 |
+
min_tokens_for_compression=50,
|
| 754 |
+
enable_ccr=False,
|
| 755 |
+
)
|
| 756 |
+
compressor = LLMLinguaCompressor(config)
|
| 757 |
+
content = generate_long_json(50)
|
| 758 |
+
|
| 759 |
+
result = compressor.compress(content, content_type="json")
|
| 760 |
+
|
| 761 |
+
assert result.compression_ratio < 1.0
|
| 762 |
+
|
| 763 |
+
def test_actual_code_compression(self):
|
| 764 |
+
"""Test code content compression with real llmlingua."""
|
| 765 |
+
config = LLMLinguaConfig(
|
| 766 |
+
target_compression_rate=0.4,
|
| 767 |
+
min_tokens_for_compression=50,
|
| 768 |
+
enable_ccr=False,
|
| 769 |
+
)
|
| 770 |
+
compressor = LLMLinguaCompressor(config)
|
| 771 |
+
content = generate_long_code(30)
|
| 772 |
+
|
| 773 |
+
result = compressor.compress(content, content_type="code")
|
| 774 |
+
|
| 775 |
+
assert result.compression_ratio < 1.0
|
| 776 |
+
|
| 777 |
+
|
| 778 |
+
# =============================================================================
|
| 779 |
+
# TestMemoryManagement
|
| 780 |
+
# =============================================================================
|
| 781 |
+
|
| 782 |
+
|
| 783 |
+
class TestMemoryManagement:
|
| 784 |
+
"""Tests for memory management functions (unload_llmlingua_model, is_llmlingua_model_loaded)."""
|
| 785 |
+
|
| 786 |
+
def test_is_model_loaded_returns_false_initially(self):
|
| 787 |
+
"""is_llmlingua_model_loaded returns False when no model loaded."""
|
| 788 |
+
# Ensure model is unloaded
|
| 789 |
+
with patch(
|
| 790 |
+
"headroom.transforms.llmlingua_compressor._llmlingua_instance",
|
| 791 |
+
None,
|
| 792 |
+
):
|
| 793 |
+
assert is_llmlingua_model_loaded() is False
|
| 794 |
+
|
| 795 |
+
def test_is_model_loaded_returns_true_when_loaded(self):
|
| 796 |
+
"""is_llmlingua_model_loaded returns True when model is loaded."""
|
| 797 |
+
mock_instance = MagicMock()
|
| 798 |
+
|
| 799 |
+
with patch(
|
| 800 |
+
"headroom.transforms.llmlingua_compressor._llmlingua_instance",
|
| 801 |
+
mock_instance,
|
| 802 |
+
):
|
| 803 |
+
assert is_llmlingua_model_loaded() is True
|
| 804 |
+
|
| 805 |
+
def test_unload_returns_false_when_no_model(self):
|
| 806 |
+
"""unload_llmlingua_model returns False when no model loaded."""
|
| 807 |
+
import headroom.transforms.llmlingua_compressor as module
|
| 808 |
+
|
| 809 |
+
# Save original
|
| 810 |
+
original = module._llmlingua_instance
|
| 811 |
+
|
| 812 |
+
try:
|
| 813 |
+
module._llmlingua_instance = None
|
| 814 |
+
result = unload_llmlingua_model()
|
| 815 |
+
assert result is False
|
| 816 |
+
finally:
|
| 817 |
+
module._llmlingua_instance = original
|
| 818 |
+
|
| 819 |
+
def test_unload_clears_instance(self):
|
| 820 |
+
"""unload_llmlingua_model clears the global instance."""
|
| 821 |
+
import headroom.transforms.llmlingua_compressor as module
|
| 822 |
+
|
| 823 |
+
# Save original
|
| 824 |
+
original = module._llmlingua_instance
|
| 825 |
+
|
| 826 |
+
try:
|
| 827 |
+
# Set a mock instance
|
| 828 |
+
mock_instance = MagicMock()
|
| 829 |
+
mock_instance._model_name = "test-model"
|
| 830 |
+
module._llmlingua_instance = mock_instance
|
| 831 |
+
|
| 832 |
+
# Unload
|
| 833 |
+
result = unload_llmlingua_model()
|
| 834 |
+
|
| 835 |
+
assert result is True
|
| 836 |
+
assert module._llmlingua_instance is None
|
| 837 |
+
finally:
|
| 838 |
+
module._llmlingua_instance = original
|
| 839 |
+
|
| 840 |
+
def test_unload_clears_cuda_cache(self):
|
| 841 |
+
"""unload_llmlingua_model attempts to clear CUDA cache."""
|
| 842 |
+
import headroom.transforms.llmlingua_compressor as module
|
| 843 |
+
|
| 844 |
+
original = module._llmlingua_instance
|
| 845 |
+
|
| 846 |
+
try:
|
| 847 |
+
mock_instance = MagicMock()
|
| 848 |
+
mock_instance._model_name = "test-model"
|
| 849 |
+
module._llmlingua_instance = mock_instance
|
| 850 |
+
|
| 851 |
+
mock_torch = MagicMock()
|
| 852 |
+
mock_torch.cuda.is_available.return_value = True
|
| 853 |
+
|
| 854 |
+
with patch.dict("sys.modules", {"torch": mock_torch}):
|
| 855 |
+
with patch(
|
| 856 |
+
"headroom.transforms.llmlingua_compressor.torch",
|
| 857 |
+
mock_torch,
|
| 858 |
+
create=True,
|
| 859 |
+
):
|
| 860 |
+
result = unload_llmlingua_model()
|
| 861 |
+
|
| 862 |
+
assert result is True
|
| 863 |
+
finally:
|
| 864 |
+
module._llmlingua_instance = original
|
| 865 |
+
|
| 866 |
+
|
| 867 |
+
# =============================================================================
|
| 868 |
+
# TestThreadSafety
|
| 869 |
+
# =============================================================================
|
| 870 |
+
|
| 871 |
+
|
| 872 |
+
class TestThreadSafety:
|
| 873 |
+
"""Tests for thread safety of model loading."""
|
| 874 |
+
|
| 875 |
+
def test_lock_exists(self):
|
| 876 |
+
"""Verify thread lock is available."""
|
| 877 |
+
import headroom.transforms.llmlingua_compressor as module
|
| 878 |
+
|
| 879 |
+
assert hasattr(module, "_llmlingua_lock")
|
| 880 |
+
import threading
|
| 881 |
+
|
| 882 |
+
assert isinstance(module._llmlingua_lock, type(threading.Lock()))
|
| 883 |
+
|
| 884 |
+
|
| 885 |
+
# =============================================================================
|
| 886 |
+
# TestErrorMessages
|
| 887 |
+
# =============================================================================
|
| 888 |
+
|
| 889 |
+
|
| 890 |
+
class TestErrorMessages:
|
| 891 |
+
"""Tests for improved error messages."""
|
| 892 |
+
|
| 893 |
+
def test_import_error_message_includes_install_hint(self):
|
| 894 |
+
"""ImportError includes installation instructions."""
|
| 895 |
+
with patch(
|
| 896 |
+
"headroom.transforms.llmlingua_compressor._check_llmlingua_available",
|
| 897 |
+
return_value=False,
|
| 898 |
+
):
|
| 899 |
+
from headroom.transforms.llmlingua_compressor import _get_llmlingua_compressor
|
| 900 |
+
|
| 901 |
+
with pytest.raises(ImportError) as exc_info:
|
| 902 |
+
_get_llmlingua_compressor("test-model", "cpu")
|
| 903 |
+
|
| 904 |
+
error_msg = str(exc_info.value)
|
| 905 |
+
assert "pip install headroom-ai[llmlingua]" in error_msg
|
| 906 |
+
assert "2GB" in error_msg or "disk space" in error_msg.lower()
|
| 907 |
+
|
| 908 |
+
def test_oom_error_provides_helpful_suggestions(self):
|
| 909 |
+
"""Out of memory error provides helpful suggestions."""
|
| 910 |
+
import headroom.transforms.llmlingua_compressor as module
|
| 911 |
+
|
| 912 |
+
# Save original state
|
| 913 |
+
original_instance = module._llmlingua_instance
|
| 914 |
+
original_available = module._llmlingua_available
|
| 915 |
+
|
| 916 |
+
try:
|
| 917 |
+
module._llmlingua_instance = None
|
| 918 |
+
module._llmlingua_available = True
|
| 919 |
+
|
| 920 |
+
# Create a mock that raises OOM when called
|
| 921 |
+
mock_prompt_compressor_class = MagicMock()
|
| 922 |
+
mock_prompt_compressor_class.side_effect = RuntimeError("CUDA out of memory")
|
| 923 |
+
|
| 924 |
+
with patch.dict("sys.modules", {"llmlingua": MagicMock()}):
|
| 925 |
+
with patch(
|
| 926 |
+
"llmlingua.PromptCompressor",
|
| 927 |
+
mock_prompt_compressor_class,
|
| 928 |
+
):
|
| 929 |
+
from headroom.transforms.llmlingua_compressor import (
|
| 930 |
+
_get_llmlingua_compressor,
|
| 931 |
+
)
|
| 932 |
+
|
| 933 |
+
with pytest.raises(RuntimeError) as exc_info:
|
| 934 |
+
_get_llmlingua_compressor("test-model", "cuda")
|
| 935 |
+
|
| 936 |
+
error_msg = str(exc_info.value)
|
| 937 |
+
# Should include helpful suggestions
|
| 938 |
+
assert "cpu" in error_msg.lower() or "memory" in error_msg.lower()
|
| 939 |
+
finally:
|
| 940 |
+
module._llmlingua_instance = original_instance
|
| 941 |
+
module._llmlingua_available = original_available
|