headroom/tests/integrations/test_strands/test_hooks_unit.py
Rod Boev dcb674b5e4
fix(compression): honor qualified CCR names across integrations (#2698)
## Description

Three compression consumers compare tool names against the bare literal
`headroom_retrieve`, so the qualified forms MCP clients actually send
(`mcp__Headroom__headroom_retrieve`, `mcp_Headroom_headroom_retrieve`)
slip past the guard and get recompressed. `SmartCrusher.apply` has the
bare comparison at both its OpenAI `role=tool` site and its Anthropic
`tool_result` block site; the LangGraph compressor and the Strands hook
have no tool-name check at all. Recompressing already-retrieved CCR
content mints a new `<<ccr:hash>>` marker the agent cannot redeem.

`headroom.config.is_tool_excluded` already owns alias resolution,
including the MCP wrapper forms. This routes all three consumers through
it instead of adding a second name matcher. Closes #2656.

## Type of Change

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

## Changes Made

- `SmartCrusher.apply` routes both its `role=tool` and its Anthropic
`tool_result` guards through `is_tool_excluded`
- `_should_skip` in the LangGraph compressor takes the tool name and
skips excluded tools; tool-call names are indexed by id so a
`ToolMessage` without a copied `name` is still classifiable
- `_should_skip_compression` in the Strands hook takes the tool name and
skips excluded tools, recording `tool_excluded`
- regressions for the qualified and bare names across all three
consumers, the Anthropic block shape, the MCP wrapper entry point, and a
near-match name that must still compress
- a LangGraph regression for incomplete tool-call metadata that
continues to a later qualified call

## Testing

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

### Test Output

`pytest tests/test_smart_crusher.py tests/integrations/test_langgraph.py
tests/integrations/test_strands
tests/test_transforms/test_smart_crusher_ccr_retrieve_exemption.py -q`

```text
tests\test_smart_crusher.py ............                                     [ 10%]
tests\integrations\test_langgraph.py .....                                   [ 15%]
tests\integrations\test_strands\test_ccr_exclusion.py .....                  [ 19%]
tests\integrations\test_strands\test_hooks.py sssssssss                      [ 27%]
tests\integrations\test_strands\test_hooks_unit.py ssssssssssssssssssssssssssssssssss [ 57%]
tests\integrations\test_strands\test_model.py ssssssssssssssss               [ 71%]
tests\integrations\test_strands\test_model_unit.py sssssssssssssssssssssssssss [ 95%]
tests\test_transforms\test_smart_crusher_ccr_retrieve_exemption.py .....      [100%]

28 passed, 86 skipped in the focused invariant suite
```

## Real Behavior Proof

- Environment: Windows 11, Python 3.12.13, `headroom._core` built
- Exact command / steps: `uv run pytest tests/test_smart_crusher.py
tests/integrations/test_langgraph.py tests/integrations/test_strands
-q`, and the same suite run against the pre-change implementation with
the new tests in place
- Observed result: before the change, five regressions fail.
`SmartCrusher` returns non-byte-identical content for a
`mcp__Headroom__headroom_retrieve` result, the LangGraph compressor
replaces the message content, and the Strands hook returns
`"compressed"` in place of the tool output. After the change all three
preserve the content byte-for-byte, incomplete LangGraph tool-call
metadata is ignored while the later qualified call remains indexed, the
Strands hook records `tool_excluded` and never calls the crusher, and
`HeadroomMCPCompressor.compress` returns the payload unchanged.
`mcp__Headroom__headroom_retrieve_extra` still compresses in all three,
and the Kompress and ContentRouter suites are unchanged.
- Not tested: the optional Strands package, so the additions to
`tests/integrations/test_strands/test_hooks_unit.py` skip locally

## Review Readiness

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

## Checklist

- [x] My code follows the project's style guidelines
- [x] I have performed a self-review of my code
- [x] I have commented my code, particularly in hard-to-understand areas
- [ ] I have made corresponding changes to the documentation
- [x] My changes generate no new warnings
- [x] I have added tests that prove my fix is effective or that my
feature works
- [x] New and existing unit tests pass locally with my changes
- [x] I did **not** edit `CHANGELOG.md` — it is generated by
release-please from my Conventional Commit PR title (a CI guard enforces
this)

## Additional Notes

One deliberate divergence from the issue: the suggested snippet passes
`DEFAULT_VERBATIM_EXCLUDE_TOOLS` to `is_tool_excluded`, but that
constant holds only `WebSearch`, `WebFetch`, `web_search`, `web_fetch`.
Applied literally it would drop `headroom_retrieve` from the comparison
entirely and delete the #1077 guard these two SmartCrusher sites exist
to enforce. This passes `(CCR_TOOL_NAME,)` so each guard keeps doing the
one thing it documents. If you'd rather these paths also honor the
verbatim-exclude set, the tuple can become `(CCR_TOOL_NAME,
*DEFAULT_VERBATIM_EXCLUDE_TOOLS)` — the CCR name has to stay in it
either way.

Adjacent work: PR #2654 covers `ContentRouter` only.
2026-08-03 20:17:39 -07:00

606 lines
21 KiB
Python

"""Unit tests for Strands HeadroomHookProvider.
These tests use mocks and do NOT require AWS credentials or strands-agents.
They test the internal logic of HeadroomHookProvider in isolation.
For real integration tests, see test_hooks.py.
"""
from __future__ import annotations
import json
import threading
from datetime import datetime, timezone
from unittest.mock import MagicMock
import pytest
# Check if strands-agents is installed for proper skip handling
try:
import strands # noqa: F401
STRANDS_AVAILABLE = True
except ImportError:
STRANDS_AVAILABLE = False
# Skip all tests if Strands not installed
pytestmark = pytest.mark.skipif(not STRANDS_AVAILABLE, reason="strands-agents not installed")
class TestHeadroomHookProviderInit:
"""Tests for HeadroomHookProvider initialization."""
def test_init_with_defaults(self):
"""Initialize with default settings."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider()
assert hook.compress_tool_outputs is True
assert hook.min_tokens_to_compress == 100
assert hook.preserve_errors is True
assert hook.total_tokens_saved == 0
assert hook.metrics_history == []
def test_init_with_custom_config(self):
"""Initialize with custom configuration."""
from headroom import HeadroomConfig
from headroom.integrations.strands import HeadroomHookProvider
config = HeadroomConfig()
config.smart_crusher.min_tokens_to_crush = 200
config.smart_crusher.max_items_after_crush = 20
hook = HeadroomHookProvider(
compress_tool_outputs=False,
min_tokens_to_compress=500,
config=config,
preserve_errors=False,
)
assert hook.compress_tool_outputs is False
assert hook.min_tokens_to_compress == 500
assert hook.config is config
assert hook.preserve_errors is False
def test_init_creates_default_config_if_none(self):
"""Initialize creates a default HeadroomConfig if none provided."""
from headroom import HeadroomConfig
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider()
assert hook.config is not None
assert isinstance(hook.config, HeadroomConfig)
class TestRegisterHooks:
"""Tests for HeadroomHookProvider.register_hooks method."""
def test_register_hooks_adds_callback_to_registry(self):
"""register_hooks adds AfterToolCallEvent callback to registry."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider(compress_tool_outputs=True)
mock_registry = MagicMock()
hook.register_hooks(mock_registry)
# Should have registered exactly one callback for AfterToolCallEvent
assert mock_registry.add_callback.call_count == 1
def test_register_hooks_skips_when_compression_disabled(self):
"""register_hooks does not register callbacks when compression is disabled."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider(compress_tool_outputs=False)
mock_registry = MagicMock()
hook.register_hooks(mock_registry)
# Should not have registered any callbacks
assert mock_registry.add_callback.call_count == 0
class TestCrusherLazyInit:
"""Tests for SmartCrusher lazy initialization."""
def test_crusher_is_lazily_initialized(self):
"""SmartCrusher is not created until first access."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider()
# Directly check internal state - crusher should be None initially
assert hook._crusher is None
# Access the crusher property
crusher = hook.crusher
# Now it should be initialized
assert crusher is not None
assert hook._crusher is crusher
def test_crusher_uses_configured_min_tokens(self):
"""SmartCrusher uses min_tokens_to_compress from hook config."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider(min_tokens_to_compress=250)
crusher = hook.crusher
# The crusher config should have our min_tokens setting
assert crusher.config.min_tokens_to_crush == 250
class TestTokenEstimation:
"""Tests for _estimate_tokens helper method."""
def test_estimate_tokens_empty_string(self):
"""Estimate returns 0 for empty string."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider()
assert hook._estimate_tokens("") == 0
def test_estimate_tokens_short_string(self):
"""Estimate uses ~4 chars per token heuristic."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider()
# 12 chars = 3 tokens (12 // 4)
assert hook._estimate_tokens("hello world!") == 3
# 20 chars = 5 tokens
assert hook._estimate_tokens("a" * 20) == 5
class TestExtractTextContent:
"""Tests for _extract_text_content helper method."""
def test_extract_from_text_content(self):
"""Extract text from content with text field."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider()
result = {"content": [{"text": "Hello world"}]}
extracted = hook._extract_text_content(result)
assert extracted == "Hello world"
def test_extract_from_json_content(self):
"""Extract and serialize JSON content."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider()
result = {"content": [{"json": {"key": "value"}}]}
extracted = hook._extract_text_content(result)
assert extracted == '{"key": "value"}'
def test_extract_empty_content(self):
"""Return empty string for empty content."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider()
result = {"content": []}
extracted = hook._extract_text_content(result)
assert extracted == ""
def test_extract_missing_content(self):
"""Return empty string for missing content key."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider()
result = {}
extracted = hook._extract_text_content(result)
assert extracted == ""
class TestShouldSkipCompression:
"""Tests for _should_skip_compression helper method."""
def test_skip_when_compression_disabled(self):
"""Skip compression when compress_tool_outputs is False."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider(compress_tool_outputs=False)
result = {"content": [{"text": "data"}]}
skip_reason = hook._should_skip_compression(result)
assert skip_reason == "compression_disabled"
def test_skip_error_results_when_preserve_errors_true(self):
"""Skip error results when preserve_errors is True."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider(preserve_errors=True)
result = {"status": "error", "content": [{"text": "Error message"}]}
skip_reason = hook._should_skip_compression(result)
assert skip_reason == "error_result_preserved"
def test_allow_error_results_when_preserve_errors_false(self):
"""Allow error results when preserve_errors is False."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider(preserve_errors=False)
result = {"status": "error", "content": [{"text": "Error message"}]}
skip_reason = hook._should_skip_compression(result)
assert skip_reason is None
def test_skip_empty_content(self):
"""Skip results with empty content."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider()
result = {"content": []}
skip_reason = hook._should_skip_compression(result)
assert skip_reason == "empty_content"
def test_allow_valid_content(self):
"""Allow results with valid content."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider()
result = {"content": [{"text": "some data"}]}
skip_reason = hook._should_skip_compression(result)
assert skip_reason is None
@pytest.mark.parametrize(
"tool_name",
["mcp__Headroom__headroom_retrieve", "mcp_Headroom_headroom_retrieve"],
)
def test_skip_qualified_ccr_retrieval_results(self, tool_name):
"""Qualified CCR retrieval names use the shared exclusion authority."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider()
result = {"content": [{"text": "retrieved content"}]}
assert hook._should_skip_compression(result, tool_name) == "tool_excluded"
def test_near_match_ccr_tool_name_is_not_excluded(self):
"""A similar qualified name remains eligible for compression."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider()
result = {"content": [{"text": "retrieved content"}]}
assert (
hook._should_skip_compression(result, "mcp__Headroom__headroom_retrieve_extra") is None
)
class TestCompressToolResult:
"""Tests for _compress_tool_result hook handler."""
def test_compress_large_tool_output(self):
"""Compresses large tool output and tracks metrics."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider(
compress_tool_outputs=True,
min_tokens_to_compress=10, # Low threshold for testing
)
# Create large JSON output (50 items)
large_data = [{"id": i, "value": f"item-{i}", "data": "x" * 50} for i in range(50)]
large_json = json.dumps(large_data)
mock_event = MagicMock()
mock_event.tool_use = {"name": "get_items", "toolUseId": "tool-123"}
mock_event.result = {"content": [{"text": large_json}]}
hook._compress_tool_result(mock_event)
# Verify metrics were recorded
assert len(hook.metrics_history) == 1
metrics = hook.metrics_history[0]
assert metrics.tool_name == "get_items"
assert metrics.tool_use_id == "tool-123"
assert metrics.tokens_before > 0
def test_skip_compression_below_threshold(self):
"""Does not compress output below token threshold."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider(
compress_tool_outputs=True,
min_tokens_to_compress=10000, # High threshold
)
mock_event = MagicMock()
mock_event.tool_use = {"name": "small_tool", "toolUseId": "tool-456"}
mock_event.result = {"content": [{"text": '{"status": "ok"}'}]}
hook._compress_tool_result(mock_event)
# Metrics should show skipped compression
assert len(hook.metrics_history) == 1
metrics = hook.metrics_history[0]
assert metrics.was_compressed is False
assert "below_threshold" in metrics.skip_reason
def test_skip_compression_when_disabled(self):
"""Does not compress when compression is disabled."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider(compress_tool_outputs=False)
mock_event = MagicMock()
mock_event.tool_use = {"name": "test_tool", "toolUseId": "tool-789"}
mock_event.result = {"content": [{"text": '{"data": "value"}'}]}
hook._compress_tool_result(mock_event)
# Metrics should show compression disabled
assert len(hook.metrics_history) == 1
metrics = hook.metrics_history[0]
assert metrics.was_compressed is False
assert metrics.skip_reason == "compression_disabled"
def test_qualified_ccr_retrieval_result_is_preserved(self):
"""The production hook skips qualified CCR retrieval results."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider(min_tokens_to_compress=1)
original = json.dumps([{"id": i, "data": "x" * 50} for i in range(50)])
mock_event = MagicMock()
mock_event.tool_use = {
"name": "mcp__Headroom__headroom_retrieve",
"toolUseId": "tool-ccr",
}
mock_event.result = {"content": [{"text": original}]}
hook._compress_tool_result(mock_event)
assert mock_event.result["content"][0]["text"] == original
assert hook.metrics_history[-1].skip_reason == "tool_excluded"
class TestMetricsTracking:
"""Tests for metrics tracking and aggregation."""
def test_total_tokens_saved_accumulates(self):
"""total_tokens_saved accumulates across compressions."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider(
compress_tool_outputs=True,
min_tokens_to_compress=10,
)
# Simulate two compressions with savings
for i in range(2):
large_data = [{"id": j, "data": "x" * 100} for j in range(50)]
mock_event = MagicMock()
mock_event.tool_use = {"name": f"tool_{i}", "toolUseId": f"id_{i}"}
mock_event.result = {"content": [{"text": json.dumps(large_data)}]}
hook._compress_tool_result(mock_event)
# Should have accumulated some savings
compressed_count = sum(1 for m in hook.metrics_history if m.was_compressed)
if compressed_count > 0:
assert hook.total_tokens_saved >= 0
def test_metrics_history_bounded_to_100(self):
"""metrics_history keeps only last 100 entries."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider(
compress_tool_outputs=True,
min_tokens_to_compress=10,
)
# Directly add 150 metrics
for i in range(150):
hook._record_metrics(
request_id=f"req_{i}",
tool_name=f"tool_{i}",
tool_use_id=f"id_{i}",
tokens_before=100,
tokens_after=50,
was_compressed=True,
skip_reason=None,
)
# Should be bounded at 100
assert len(hook.metrics_history) == 100
# Should contain the most recent entries
last_metric = hook.metrics_history[-1]
assert last_metric.request_id == "req_149"
class TestGetSavingsSummary:
"""Tests for get_savings_summary method."""
def test_empty_summary(self):
"""Returns zero values when no metrics recorded."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider()
summary = hook.get_savings_summary()
assert summary["total_requests"] == 0
assert summary["compressed_requests"] == 0
assert summary["total_tokens_saved"] == 0
assert summary["average_savings_percent"] == 0.0
def test_summary_with_compressions(self):
"""Returns correct summary with recorded compressions."""
from headroom.integrations.strands import HeadroomHookProvider
from headroom.integrations.strands.hooks import CompressionMetrics
hook = HeadroomHookProvider()
# Add metrics manually
hook._metrics_history = [
CompressionMetrics(
request_id="1",
timestamp=datetime.now(timezone.utc),
tool_name="tool_a",
tool_use_id="id_1",
tokens_before=100,
tokens_after=60,
tokens_saved=40,
savings_percent=40.0,
was_compressed=True,
skip_reason=None,
),
CompressionMetrics(
request_id="2",
timestamp=datetime.now(timezone.utc),
tool_name="tool_b",
tool_use_id="id_2",
tokens_before=200,
tokens_after=100,
tokens_saved=100,
savings_percent=50.0,
was_compressed=True,
skip_reason=None,
),
CompressionMetrics(
request_id="3",
timestamp=datetime.now(timezone.utc),
tool_name="tool_c",
tool_use_id="id_3",
tokens_before=50,
tokens_after=50,
tokens_saved=0,
savings_percent=0.0,
was_compressed=False,
skip_reason="below_threshold",
),
]
hook._total_tokens_saved = 140
summary = hook.get_savings_summary()
assert summary["total_requests"] == 3
assert summary["compressed_requests"] == 2
assert summary["total_tokens_saved"] == 140
assert summary["average_savings_percent"] == 45.0 # (40 + 50) / 2
assert summary["total_tokens_before"] == 350
assert summary["total_tokens_after"] == 210
class TestReset:
"""Tests for reset method."""
def test_reset_clears_all_state(self):
"""reset() clears all tracked state."""
from headroom.integrations.strands import HeadroomHookProvider
from headroom.integrations.strands.hooks import CompressionMetrics
hook = HeadroomHookProvider()
# Add some state
hook._metrics_history = [
CompressionMetrics(
request_id="1",
timestamp=datetime.now(timezone.utc),
tool_name="test",
tool_use_id="id_1",
tokens_before=100,
tokens_after=50,
tokens_saved=50,
savings_percent=50.0,
was_compressed=True,
)
]
hook._total_tokens_saved = 50
# Reset
hook.reset()
# Verify all state cleared
assert hook._metrics_history == []
assert hook._total_tokens_saved == 0
assert hook.total_tokens_saved == 0
assert len(hook.metrics_history) == 0
class TestThreadSafety:
"""Tests for thread-safety of metrics tracking."""
def test_concurrent_metric_recording(self):
"""Metrics recording is thread-safe."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider()
def record_metrics(thread_id):
for i in range(10):
hook._record_metrics(
request_id=f"thread_{thread_id}_req_{i}",
tool_name=f"tool_{thread_id}_{i}",
tool_use_id=f"id_{thread_id}_{i}",
tokens_before=100,
tokens_after=50,
was_compressed=True,
skip_reason=None,
)
threads = []
for t_id in range(5):
t = threading.Thread(target=record_metrics, args=(t_id,))
threads.append(t)
t.start()
for t in threads:
t.join()
# Should have recorded 50 metrics (5 threads * 10 each)
# But bounded to 100, so if we had more it would be truncated
assert len(hook.metrics_history) == 50
assert hook.total_tokens_saved == 50 * 50 # 50 metrics * 50 tokens each
class TestUpdateResultContent:
"""Tests for _update_result_content helper method."""
def test_update_preserves_json_structure(self):
"""Updates preserve JSON structure when possible."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider()
result = {"content": [{"json": {"original": "data"}}]}
compressed = '{"compressed": "data"}'
hook._update_result_content(result, compressed)
# Should update with parsed JSON
assert result["content"] == [{"json": {"compressed": "data"}}]
def test_update_uses_text_for_non_json(self):
"""Updates use text format for non-JSON content."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider()
result = {"content": [{"text": "original text"}]}
compressed = "compressed text"
hook._update_result_content(result, compressed)
assert result["content"] == [{"text": "compressed text"}]
def test_update_creates_content_if_empty(self):
"""Creates content list if missing."""
from headroom.integrations.strands import HeadroomHookProvider
hook = HeadroomHookProvider()
result = {"content": []}
hook._update_result_content(result, "new content")
assert result["content"] == [{"text": "new content"}]