mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-27 14:17:10 -04:00
## 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.
606 lines
21 KiB
Python
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"}]
|