diff --git a/headroom/cache/compression_feedback.py b/headroom/cache/compression_feedback.py index 20413ade4..e4a320444 100644 --- a/headroom/cache/compression_feedback.py +++ b/headroom/cache/compression_feedback.py @@ -33,6 +33,8 @@ import time from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any +from .compression_strategy_outcomes import CompressionStrategyOutcomes + if TYPE_CHECKING: from .compression_store import CompressionStore, RetrievalEvent @@ -94,28 +96,33 @@ class LocalToolPattern: def strategy_retrieval_rate(self, strategy: str) -> float: """Get retrieval rate for a specific compression strategy.""" - compressions = self.strategy_compressions.get(strategy, 0) - if compressions == 0: - return 0.0 - retrievals = self.strategy_retrievals.get(strategy, 0) - return retrievals / compressions + return self.strategy_outcomes.retrieval_rate(strategy) def best_strategy(self) -> str | None: """Find the strategy with lowest retrieval rate (most successful).""" - if not self.strategy_compressions: - return None + return self.strategy_outcomes.best_strategy() - best = None - best_rate = 1.0 + @property + def strategy_outcomes(self) -> CompressionStrategyOutcomes: + """Strategy outcome view backed by this pattern's public counters.""" + return CompressionStrategyOutcomes( + compressions=self.strategy_compressions, + retrievals=self.strategy_retrievals, + ) - for strategy in self.strategy_compressions: - rate = self.strategy_retrieval_rate(strategy) - # Only consider strategies with enough samples - if self.strategy_compressions[strategy] >= 3 and rate < best_rate: - best_rate = rate - best = strategy + def record_strategy_compression(self, strategy: str) -> None: + """Record strategy compression outcome.""" + outcomes = self.strategy_outcomes + outcomes.record_compression(strategy) + self.strategy_compressions = outcomes.compressions + self.strategy_retrievals = outcomes.retrievals - return best + def record_strategy_retrieval(self, strategy: str) -> None: + """Record strategy retrieval outcome.""" + outcomes = self.strategy_outcomes + outcomes.record_retrieval(strategy) + self.strategy_compressions = outcomes.compressions + self.strategy_retrievals = outcomes.retrievals @dataclass @@ -235,15 +242,7 @@ class CompressionFeedback: # Track strategy usage if strategy: - pattern.strategy_compressions[strategy] = ( - pattern.strategy_compressions.get(strategy, 0) + 1 - ) - - # CRITICAL FIX: When truncating strategy dicts, keep them in sync - # to prevent desync between compressions and retrievals. - # Both dicts must have the same keys for accurate retrieval rate calculation. - if len(pattern.strategy_compressions) > 50: - self._truncate_strategy_dicts(pattern) + pattern.record_strategy_compression(strategy) # Track signature hash for TOIN correlation if tool_signature_hash: @@ -291,14 +290,7 @@ class CompressionFeedback: # Track strategy retrievals (for success rate calculation) if strategy: - pattern.strategy_retrievals[strategy] = ( - pattern.strategy_retrievals.get(strategy, 0) + 1 - ) - - # CRITICAL FIX: When truncating strategy dicts, keep them in sync - # to prevent desync between compressions and retrievals. - if len(pattern.strategy_retrievals) > 50: - self._truncate_strategy_dicts(pattern) + pattern.record_strategy_retrieval(strategy) # Track query patterns if event.query: @@ -318,40 +310,11 @@ class CompressionFeedback: self._extract_field_hints(pattern, event.query) def _truncate_strategy_dicts(self, pattern: LocalToolPattern) -> None: - """Truncate strategy_compressions and strategy_retrievals in sync. - - CRITICAL FIX: Both dicts must have the same keys for accurate retrieval - rate calculation. When truncating, we keep the union of top strategies - from both dicts, then truncate both to the same key set. - """ - # Get top 40 strategies from each dict (using 40 to allow union to stay under 50) - top_compressions = { - k - for k, _ in sorted( - pattern.strategy_compressions.items(), - key=lambda x: x[1], - reverse=True, - )[:40] - } - top_retrievals = { - k - for k, _ in sorted( - pattern.strategy_retrievals.items(), - key=lambda x: x[1], - reverse=True, - )[:40] - } - - # Keep union of top strategies from both - keys_to_keep = top_compressions | top_retrievals - - # Truncate both dicts to same keys - pattern.strategy_compressions = { - k: v for k, v in pattern.strategy_compressions.items() if k in keys_to_keep - } - pattern.strategy_retrievals = { - k: v for k, v in pattern.strategy_retrievals.items() if k in keys_to_keep - } + """Truncate strategy counters using the shared strategy outcome domain.""" + outcomes = pattern.strategy_outcomes + outcomes.prune() + pattern.strategy_compressions = outcomes.compressions + pattern.strategy_retrievals = outcomes.retrievals def _extract_field_hints(self, pattern: LocalToolPattern, query: str) -> None: """Extract potential field names from search queries. diff --git a/headroom/cache/compression_strategy_outcomes.py b/headroom/cache/compression_strategy_outcomes.py new file mode 100644 index 000000000..04cfdf578 --- /dev/null +++ b/headroom/cache/compression_strategy_outcomes.py @@ -0,0 +1,99 @@ +"""Strategy outcome accounting for local compression feedback.""" + +from __future__ import annotations + +from dataclasses import dataclass, field + + +@dataclass +class CompressionStrategyOutcomes: + """Track compression and retrieval outcomes by compression strategy.""" + + compressions: dict[str, int] = field(default_factory=dict) + retrievals: dict[str, int] = field(default_factory=dict) + max_strategies: int = 50 + top_strategies_per_counter: int = 40 + minimum_samples_for_recommendation: int = 3 + + def record_compression(self, strategy: str) -> None: + """Record one compression for a strategy.""" + self.compressions[strategy] = self.compressions.get(strategy, 0) + 1 + self.prune() + + def record_retrieval(self, strategy: str) -> None: + """Record one retrieval for a strategy.""" + self.retrievals[strategy] = self.retrievals.get(strategy, 0) + 1 + self.prune() + + def retrieval_rate(self, strategy: str) -> float: + """Return the retrievals-per-compression rate for one strategy.""" + compressions = self.compressions.get(strategy, 0) + if compressions == 0: + return 0.0 + return self.retrievals.get(strategy, 0) / compressions + + def best_strategy(self) -> str | None: + """Return the sampled strategy with the lowest retrieval rate.""" + best = None + best_rate = 1.0 + + for strategy, compression_count in self.compressions.items(): + if compression_count < self.minimum_samples_for_recommendation: + continue + + rate = self.retrieval_rate(strategy) + if rate < best_rate: + best = strategy + best_rate = rate + + return best + + def prune(self) -> None: + """Bound counters while preserving the highest-signal strategies.""" + if ( + len(self.compressions) <= self.max_strategies + and len(self.retrievals) <= self.max_strategies + ): + return + + keys_to_keep = self._keys_to_keep() + self.compressions = { + strategy: count + for strategy, count in self.compressions.items() + if strategy in keys_to_keep + } + self.retrievals = { + strategy: count + for strategy, count in self.retrievals.items() + if strategy in keys_to_keep + } + + def _keys_to_keep(self) -> set[str]: + top_compressions = self._top_keys(self.compressions) + top_retrievals = self._top_keys(self.retrievals) + candidate_keys = top_compressions | top_retrievals + + if len(candidate_keys) <= self.max_strategies: + return candidate_keys + + ranked_keys = sorted( + candidate_keys, + key=lambda strategy: ( + self.compressions.get(strategy, 0) + self.retrievals.get(strategy, 0), + self.compressions.get(strategy, 0), + self.retrievals.get(strategy, 0), + strategy, + ), + reverse=True, + ) + return set(ranked_keys[: self.max_strategies]) + + def _top_keys(self, counts: dict[str, int]) -> set[str]: + return { + strategy + for strategy, _ in sorted( + counts.items(), + key=lambda item: (item[1], item[0]), + reverse=True, + )[: self.top_strategies_per_counter] + } diff --git a/tests/test_compression_strategy_outcomes.py b/tests/test_compression_strategy_outcomes.py new file mode 100644 index 000000000..d872bde5a --- /dev/null +++ b/tests/test_compression_strategy_outcomes.py @@ -0,0 +1,45 @@ +from headroom.cache.compression_strategy_outcomes import CompressionStrategyOutcomes + + +def test_retrieval_rate_is_zero_without_strategy_compressions(): + outcomes = CompressionStrategyOutcomes(retrievals={"sample": 2}) + + assert outcomes.retrieval_rate("sample") == 0.0 + + +def test_best_strategy_requires_minimum_samples(): + outcomes = CompressionStrategyOutcomes( + compressions={"under_sampled": 2, "sampled": 3}, + retrievals={"under_sampled": 0, "sampled": 1}, + ) + + assert outcomes.best_strategy() == "sampled" + + +def test_best_strategy_uses_lowest_retrieval_rate(): + outcomes = CompressionStrategyOutcomes( + compressions={"top_n": 10, "smart_sample": 10}, + retrievals={"top_n": 7, "smart_sample": 2}, + ) + + assert outcomes.retrieval_rate("smart_sample") == 0.2 + assert outcomes.best_strategy() == "smart_sample" + + +def test_recording_prunes_strategy_counters_to_bounded_high_signal_set(): + outcomes = CompressionStrategyOutcomes(max_strategies=10, top_strategies_per_counter=8) + + for index in range(30): + strategy = f"strategy_{index:02d}" + for _ in range(index + 1): + outcomes.record_compression(strategy) + + for index in range(30): + strategy = f"strategy_{index:02d}" + for _ in range(30 - index): + outcomes.record_retrieval(strategy) + + assert len(outcomes.compressions) <= 10 + assert len(outcomes.retrievals) <= 10 + assert "strategy_29" in outcomes.compressions + assert "strategy_00" in outcomes.retrievals