mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-27 14:17:10 -04:00
737 lines
28 KiB
Python
737 lines
28 KiB
Python
"""Tests for TOIN Field-Level Learning.
|
|
|
|
These tests PROVE that field-level learning actually works:
|
|
1. FieldSemantics correctly infers types from retrieval patterns
|
|
2. TOIN populates field_semantics from retrievals
|
|
3. SmartCrusher uses learned semantics to detect important items
|
|
4. End-to-end: important items are preserved based on learned behavior
|
|
|
|
No hardcoded patterns - all learning is behavior-based.
|
|
"""
|
|
|
|
import hashlib
|
|
|
|
import pytest
|
|
|
|
from headroom.telemetry import (
|
|
ToolIntelligenceNetwork,
|
|
ToolPattern,
|
|
ToolSignature,
|
|
reset_toin,
|
|
)
|
|
from headroom.telemetry.models import FieldSemantics
|
|
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def reset_globals():
|
|
"""Reset global state before each test."""
|
|
reset_toin()
|
|
yield
|
|
reset_toin()
|
|
|
|
|
|
def _hash_value(value) -> str:
|
|
"""Hash a value the same way TOIN does."""
|
|
value_str = str(value)
|
|
return hashlib.sha256(value_str.encode()).hexdigest()[:8]
|
|
|
|
|
|
def _hash_field(field_name: str) -> str:
|
|
"""Hash a field name the same way TOIN does."""
|
|
return hashlib.sha256(field_name.encode()).hexdigest()[:8]
|
|
|
|
|
|
class TestFieldSemanticsLearning:
|
|
"""Test that FieldSemantics correctly learns from retrieval patterns."""
|
|
|
|
def test_identifier_type_inference(self):
|
|
"""PROVES: Field used with exact-match queries + high uniqueness = identifier."""
|
|
fs = FieldSemantics(field_hash="test123")
|
|
|
|
# Simulate: user retrieves items by unique IDs (exact match queries)
|
|
# Each ID is different - high uniqueness
|
|
for i in range(10):
|
|
fs.record_retrieval_value(_hash_value(f"id_{i}"), operator="=")
|
|
|
|
# Simulate compression stats: all values are unique
|
|
fs.record_compression_stats(
|
|
unique_values=100,
|
|
total_values=100, # uniqueness ratio = 1.0
|
|
most_common_value_hash=_hash_value("id_0"),
|
|
most_common_frequency=0.01, # No dominant value
|
|
)
|
|
fs.record_compression_stats(
|
|
unique_values=100,
|
|
total_values=100,
|
|
most_common_value_hash=_hash_value("id_1"),
|
|
most_common_frequency=0.01,
|
|
)
|
|
|
|
# Now infer the type
|
|
fs.infer_type()
|
|
|
|
# VERIFY: Should be classified as identifier
|
|
assert fs.inferred_type == "identifier", (
|
|
f"Expected 'identifier' but got '{fs.inferred_type}'. "
|
|
"High uniqueness + exact match queries should = identifier"
|
|
)
|
|
assert fs.confidence > 0.5, "Should have reasonable confidence"
|
|
|
|
def test_error_indicator_type_inference(self):
|
|
"""PROVES: Field with dominant default + retrievals for non-default = error_indicator."""
|
|
fs = FieldSemantics(field_hash="status_field")
|
|
|
|
# Simulate: most items have status="success" (the default)
|
|
# But user only retrieves items with status="error" or "failed"
|
|
error_hash = _hash_value("error")
|
|
failed_hash = _hash_value("failed")
|
|
success_hash = _hash_value("success")
|
|
|
|
# User retrieves "error" and "failed" values (non-default)
|
|
for _ in range(5):
|
|
fs.record_retrieval_value(error_hash, operator="=")
|
|
fs.record_retrieval_value(failed_hash, operator="=")
|
|
|
|
# Compression stats: 90% have "success" (the default)
|
|
fs.record_compression_stats(
|
|
unique_values=3, # "success", "error", "failed"
|
|
total_values=100,
|
|
most_common_value_hash=success_hash,
|
|
most_common_frequency=0.9, # 90% are "success"
|
|
)
|
|
fs.record_compression_stats(
|
|
unique_values=3,
|
|
total_values=100,
|
|
most_common_value_hash=success_hash,
|
|
most_common_frequency=0.9,
|
|
)
|
|
|
|
fs.infer_type()
|
|
|
|
# VERIFY: Should be error_indicator
|
|
assert fs.inferred_type == "error_indicator", (
|
|
f"Expected 'error_indicator' but got '{fs.inferred_type}'. "
|
|
"Dominant default + retrieval of non-default = error indicator"
|
|
)
|
|
assert fs.default_value_hash == success_hash, "Default should be 'success'"
|
|
assert fs.confidence > 0.5, "Should have reasonable confidence"
|
|
|
|
def test_status_type_inference(self):
|
|
"""PROVES: Low cardinality + specific values retrieved = status."""
|
|
fs = FieldSemantics(field_hash="state_field")
|
|
|
|
# Simulate: field has few unique values (low cardinality)
|
|
# User retrieves the same few values repeatedly
|
|
pending_hash = _hash_value("pending")
|
|
active_hash = _hash_value("active")
|
|
|
|
for _ in range(6):
|
|
fs.record_retrieval_value(pending_hash, operator="=")
|
|
for _ in range(4):
|
|
fs.record_retrieval_value(active_hash, operator="=")
|
|
|
|
# Compression stats: only 5 unique values across 100 items
|
|
fs.record_compression_stats(
|
|
unique_values=5,
|
|
total_values=100, # uniqueness ratio = 0.05 (very low)
|
|
most_common_value_hash=None,
|
|
most_common_frequency=0.3, # No overwhelming default
|
|
)
|
|
fs.record_compression_stats(
|
|
unique_values=5,
|
|
total_values=100,
|
|
most_common_value_hash=None,
|
|
most_common_frequency=0.3,
|
|
)
|
|
|
|
fs.infer_type()
|
|
|
|
# VERIFY: Should be status
|
|
assert fs.inferred_type == "status", (
|
|
f"Expected 'status' but got '{fs.inferred_type}'. "
|
|
"Low cardinality + specific values retrieved = status"
|
|
)
|
|
|
|
def test_score_type_inference(self):
|
|
"""PROVES: Range queries = score type."""
|
|
fs = FieldSemantics(field_hash="relevance_field")
|
|
|
|
# Simulate: user queries with range operators (top-N behavior)
|
|
for _ in range(8):
|
|
fs.record_retrieval_value(_hash_value("0.95"), operator=">")
|
|
for _ in range(4):
|
|
fs.record_retrieval_value(_hash_value("0.90"), operator=">=")
|
|
|
|
# Compression stats
|
|
fs.record_compression_stats(
|
|
unique_values=50,
|
|
total_values=100,
|
|
most_common_value_hash=None,
|
|
most_common_frequency=0.1,
|
|
)
|
|
fs.record_compression_stats(
|
|
unique_values=50,
|
|
total_values=100,
|
|
most_common_value_hash=None,
|
|
most_common_frequency=0.1,
|
|
)
|
|
|
|
fs.infer_type()
|
|
|
|
# VERIFY: Should be score
|
|
assert fs.inferred_type == "score", (
|
|
f"Expected 'score' but got '{fs.inferred_type}'. "
|
|
"Range queries (>, >=) should = score type"
|
|
)
|
|
|
|
def test_content_type_inference(self):
|
|
"""PROVES: Contains/text search queries = content type."""
|
|
fs = FieldSemantics(field_hash="description_field")
|
|
|
|
# Simulate: user does text search on this field
|
|
for i in range(10):
|
|
fs.record_retrieval_value(_hash_value(f"search_term_{i}"), operator="contains")
|
|
|
|
# Compression stats: high uniqueness (different descriptions)
|
|
fs.record_compression_stats(
|
|
unique_values=90,
|
|
total_values=100,
|
|
most_common_value_hash=None,
|
|
most_common_frequency=0.05,
|
|
)
|
|
fs.record_compression_stats(
|
|
unique_values=90,
|
|
total_values=100,
|
|
most_common_value_hash=None,
|
|
most_common_frequency=0.05,
|
|
)
|
|
|
|
fs.infer_type()
|
|
|
|
# VERIFY: Should be content
|
|
assert fs.inferred_type == "content", (
|
|
f"Expected 'content' but got '{fs.inferred_type}'. "
|
|
"Contains queries should = content type"
|
|
)
|
|
|
|
def test_is_value_important_for_error_indicator(self):
|
|
"""PROVES: For error_indicator, non-default values are important."""
|
|
fs = FieldSemantics(field_hash="status")
|
|
|
|
error_hash = _hash_value("error")
|
|
success_hash = _hash_value("success")
|
|
|
|
# Set up as error_indicator
|
|
fs.inferred_type = "error_indicator"
|
|
fs.confidence = 0.8
|
|
fs.default_value_hash = success_hash
|
|
fs.important_value_hashes = [error_hash]
|
|
|
|
# VERIFY
|
|
assert fs.is_value_important(error_hash), "Error value should be important"
|
|
assert not fs.is_value_important(success_hash), "Default value should NOT be important"
|
|
|
|
def test_is_value_important_for_status(self):
|
|
"""PROVES: For status fields, retrieved values are important."""
|
|
fs = FieldSemantics(field_hash="state")
|
|
|
|
pending_hash = _hash_value("pending")
|
|
unknown_hash = _hash_value("never_retrieved")
|
|
|
|
# Set up as status
|
|
fs.inferred_type = "status"
|
|
fs.confidence = 0.7
|
|
fs.value_retrieval_frequency = {pending_hash: 5}
|
|
|
|
# VERIFY
|
|
assert fs.is_value_important(pending_hash), "Retrieved value should be important"
|
|
assert not fs.is_value_important(unknown_hash), (
|
|
"Never-retrieved value should NOT be important"
|
|
)
|
|
|
|
|
|
class TestTOINFieldLearningIntegration:
|
|
"""Test that TOIN correctly integrates field-level learning."""
|
|
|
|
def test_record_retrieval_populates_field_semantics(self):
|
|
"""PROVES: record_retrieval with items actually populates field_semantics."""
|
|
toin = ToolIntelligenceNetwork()
|
|
|
|
# Create a tool signature
|
|
items = [
|
|
{"id": "123", "status": "ok", "value": 100},
|
|
{"id": "456", "status": "error", "value": 200},
|
|
]
|
|
sig = ToolSignature.from_items(items)
|
|
|
|
# Record retrieval with items - THIS IS WHERE LEARNING HAPPENS
|
|
toin.record_retrieval(
|
|
tool_signature_hash=sig.structure_hash,
|
|
retrieval_type="full",
|
|
query="status=error",
|
|
query_fields=["status"],
|
|
retrieved_items=items,
|
|
)
|
|
|
|
# VERIFY: pattern should have field_semantics populated
|
|
pattern = toin._patterns.get(sig.structure_hash)
|
|
assert pattern is not None, "Pattern should exist"
|
|
assert len(pattern.field_semantics) > 0, (
|
|
f"field_semantics should be populated after retrieval. Got: {pattern.field_semantics}"
|
|
)
|
|
|
|
# Check that field hashes match expected fields
|
|
expected_field_hashes = {_hash_field(f) for f in ["id", "status", "value"]}
|
|
actual_field_hashes = set(pattern.field_semantics.keys())
|
|
assert expected_field_hashes == actual_field_hashes, (
|
|
f"Expected field hashes {expected_field_hashes}, got {actual_field_hashes}"
|
|
)
|
|
|
|
def test_repeated_retrievals_trigger_type_inference(self):
|
|
"""PROVES: After multiple retrievals, TOIN infers field types."""
|
|
toin = ToolIntelligenceNetwork()
|
|
|
|
# Simulate items with status field
|
|
items = [
|
|
{"status": "success"},
|
|
{"status": "success"},
|
|
{"status": "success"},
|
|
{"status": "error"},
|
|
]
|
|
sig = ToolSignature.from_items(items)
|
|
|
|
# Simulate 6 retrievals (enough to trigger inference at retrieval 5)
|
|
for _ in range(6):
|
|
# User always retrieves items with error status
|
|
toin.record_retrieval(
|
|
tool_signature_hash=sig.structure_hash,
|
|
retrieval_type="full",
|
|
query="status=error",
|
|
query_fields=["status"],
|
|
retrieved_items=[{"status": "error"}],
|
|
)
|
|
# Also record compression to get enough data for inference
|
|
toin.record_compression(
|
|
tool_signature=sig,
|
|
original_count=100,
|
|
compressed_count=10,
|
|
original_tokens=1000,
|
|
compressed_tokens=100,
|
|
strategy="top_n",
|
|
items=items,
|
|
)
|
|
|
|
# VERIFY: After enough retrievals, type should be inferred
|
|
pattern = toin._patterns.get(sig.structure_hash)
|
|
status_hash = _hash_field("status")
|
|
assert status_hash in pattern.field_semantics, "Status field should be tracked"
|
|
|
|
status_sem = pattern.field_semantics[status_hash]
|
|
# The type should be inferred (not unknown) after enough data
|
|
assert status_sem.retrieval_count >= 6, (
|
|
f"Should have 6+ retrievals, got {status_sem.retrieval_count}"
|
|
)
|
|
|
|
def test_get_recommendation_includes_field_semantics(self):
|
|
"""PROVES: get_recommendation returns learned field_semantics."""
|
|
toin = ToolIntelligenceNetwork()
|
|
|
|
# Set up pattern with learned field_semantics
|
|
items = [{"id": "123", "status": "ok"}]
|
|
sig = ToolSignature.from_items(items)
|
|
|
|
# Record enough data
|
|
for i in range(10):
|
|
toin.record_retrieval(
|
|
tool_signature_hash=sig.structure_hash,
|
|
retrieval_type="full",
|
|
query=f"id={i}",
|
|
query_fields=["id"],
|
|
retrieved_items=[{"id": str(i), "status": "ok"}],
|
|
)
|
|
toin.record_compression(
|
|
tool_signature=sig,
|
|
original_count=100,
|
|
compressed_count=10,
|
|
original_tokens=1000,
|
|
compressed_tokens=100,
|
|
strategy="top_n",
|
|
items=items,
|
|
)
|
|
|
|
# Get recommendation
|
|
hint = toin.get_recommendation(sig)
|
|
|
|
# VERIFY: hint should include field_semantics
|
|
assert hint is not None, "Should get a recommendation"
|
|
# field_semantics might be empty if confidence is too low,
|
|
# but the attribute should exist
|
|
assert hasattr(hint, "field_semantics"), "Hint should have field_semantics attribute"
|
|
|
|
def test_field_semantics_persisted_correctly(self):
|
|
"""PROVES: field_semantics survives to_dict/from_dict round-trip."""
|
|
# ToolPattern is already imported at module level from headroom.telemetry
|
|
|
|
# Create pattern with field_semantics
|
|
pattern = ToolPattern(tool_signature_hash="test123")
|
|
|
|
# Add field semantics
|
|
fs = FieldSemantics(field_hash="field123")
|
|
fs.inferred_type = "error_indicator"
|
|
fs.confidence = 0.8
|
|
fs.important_value_hashes = ["value1", "value2"]
|
|
fs.default_value_hash = "default"
|
|
pattern.field_semantics["field123"] = fs
|
|
|
|
# Round-trip through dict
|
|
d = pattern.to_dict()
|
|
pattern2 = ToolPattern.from_dict(d)
|
|
|
|
# VERIFY: field_semantics preserved
|
|
assert "field123" in pattern2.field_semantics, "field_semantics should be preserved"
|
|
fs2 = pattern2.field_semantics["field123"]
|
|
assert fs2.inferred_type == "error_indicator"
|
|
assert fs2.confidence == 0.8
|
|
assert fs2.important_value_hashes == ["value1", "value2"]
|
|
|
|
|
|
class TestSmartCrusherUsesLearnedSemantics:
|
|
"""Test that SmartCrusher actually uses learned field semantics."""
|
|
|
|
def test_detect_items_by_learned_semantics_finds_important_items(self):
|
|
"""PROVES: _detect_items_by_learned_semantics correctly identifies items."""
|
|
from headroom.transforms.smart_crusher import _detect_items_by_learned_semantics
|
|
|
|
# Create field semantics that knows "error" is important
|
|
status_hash = _hash_field("status")
|
|
error_hash = _hash_value("error")
|
|
success_hash = _hash_value("success")
|
|
|
|
fs = FieldSemantics(field_hash=status_hash)
|
|
fs.inferred_type = "error_indicator"
|
|
fs.confidence = 0.8
|
|
fs.default_value_hash = success_hash
|
|
fs.important_value_hashes = [error_hash]
|
|
fs.value_retrieval_frequency = {error_hash: 10}
|
|
|
|
field_semantics = {status_hash: fs}
|
|
|
|
# Test items - index 1 has error status
|
|
items = [
|
|
{"status": "success", "message": "all good"},
|
|
{"status": "error", "message": "something failed"}, # <-- This should be detected
|
|
{"status": "success", "message": "also good"},
|
|
]
|
|
|
|
# VERIFY
|
|
important_indices = _detect_items_by_learned_semantics(items, field_semantics)
|
|
assert 1 in important_indices, (
|
|
f"Index 1 (error status) should be detected as important. "
|
|
f"Got indices: {important_indices}"
|
|
)
|
|
assert 0 not in important_indices, "Index 0 (success) should not be important"
|
|
assert 2 not in important_indices, "Index 2 (success) should not be important"
|
|
|
|
def test_detect_items_handles_empty_semantics(self):
|
|
"""PROVES: Function handles edge cases gracefully."""
|
|
from headroom.transforms.smart_crusher import _detect_items_by_learned_semantics
|
|
|
|
items = [{"status": "ok"}]
|
|
|
|
# Empty semantics
|
|
assert _detect_items_by_learned_semantics(items, {}) == []
|
|
assert _detect_items_by_learned_semantics(items, None) == []
|
|
assert _detect_items_by_learned_semantics([], {"x": FieldSemantics(field_hash="x")}) == []
|
|
|
|
def test_detect_items_requires_confidence(self):
|
|
"""PROVES: Low confidence semantics are ignored."""
|
|
from headroom.transforms.smart_crusher import _detect_items_by_learned_semantics
|
|
|
|
status_hash = _hash_field("status")
|
|
error_hash = _hash_value("error")
|
|
|
|
fs = FieldSemantics(field_hash=status_hash)
|
|
fs.inferred_type = "error_indicator"
|
|
fs.confidence = 0.1 # TOO LOW
|
|
fs.important_value_hashes = [error_hash]
|
|
|
|
items = [{"status": "error"}]
|
|
|
|
# VERIFY: Low confidence = ignored
|
|
result = _detect_items_by_learned_semantics(items, {status_hash: fs})
|
|
assert result == [], "Low confidence semantics should be ignored"
|
|
|
|
|
|
class TestEndToEndFieldLearning:
|
|
"""End-to-end tests proving the full learning pipeline works."""
|
|
|
|
def test_full_pipeline_learns_and_applies(self):
|
|
"""PROVES: End-to-end learning from retrieval to compression."""
|
|
toin = ToolIntelligenceNetwork()
|
|
|
|
# PHASE 1: Learning - User retrieves items with error status
|
|
# Simulating: "show me all failed items"
|
|
items_with_errors = [
|
|
{"id": "1", "status": "success", "data": "..."},
|
|
{"id": "2", "status": "error", "data": "..."}, # Retrieved
|
|
{"id": "3", "status": "success", "data": "..."},
|
|
{"id": "4", "status": "failed", "data": "..."}, # Retrieved
|
|
]
|
|
sig = ToolSignature.from_items(items_with_errors)
|
|
|
|
# User keeps retrieving error/failed items (learning behavior)
|
|
for _ in range(5):
|
|
toin.record_retrieval(
|
|
tool_signature_hash=sig.structure_hash,
|
|
retrieval_type="full",
|
|
query="status!=success",
|
|
query_fields=["status"],
|
|
retrieved_items=[
|
|
{"id": "2", "status": "error", "data": "..."},
|
|
{"id": "4", "status": "failed", "data": "..."},
|
|
],
|
|
)
|
|
toin.record_compression(
|
|
tool_signature=sig,
|
|
original_count=100,
|
|
compressed_count=10,
|
|
original_tokens=1000,
|
|
compressed_tokens=100,
|
|
strategy="top_n",
|
|
items=[
|
|
{"id": str(i), "status": "success" if i % 5 != 0 else "error", "data": "..."}
|
|
for i in range(100)
|
|
],
|
|
)
|
|
|
|
# PHASE 2: Verify learning occurred
|
|
pattern = toin._patterns.get(sig.structure_hash)
|
|
assert pattern is not None
|
|
assert len(pattern.field_semantics) > 0, "Should have learned field semantics"
|
|
|
|
# Check status field was learned
|
|
status_hash = _hash_field("status")
|
|
if status_hash in pattern.field_semantics:
|
|
status_sem = pattern.field_semantics[status_hash]
|
|
# Error value should be tracked
|
|
error_hash = _hash_value("error")
|
|
failed_hash = _hash_value("failed")
|
|
assert (
|
|
error_hash in status_sem.important_value_hashes
|
|
or failed_hash in status_sem.important_value_hashes
|
|
), "Error/failed values should be marked as important"
|
|
|
|
def test_recommendation_hint_includes_learned_semantics(self):
|
|
"""PROVES: TOIN recommendation includes learned field semantics for SmartCrusher."""
|
|
toin = ToolIntelligenceNetwork()
|
|
|
|
# Set up sufficient learning
|
|
items = [{"status": "ok", "id": "123"}]
|
|
sig = ToolSignature.from_items(items)
|
|
|
|
# Create pattern with confident field semantics by directly manipulating internal state
|
|
# (This is a test - in real code, patterns are created via record_* methods)
|
|
pattern = ToolPattern(tool_signature_hash=sig.structure_hash)
|
|
toin._patterns[sig.structure_hash] = pattern
|
|
|
|
status_hash = _hash_field("status")
|
|
fs = FieldSemantics(field_hash=status_hash)
|
|
fs.inferred_type = "error_indicator"
|
|
fs.confidence = 0.8 # High confidence
|
|
fs.retrieval_count = 10
|
|
fs.important_value_hashes = [_hash_value("error")]
|
|
pattern.field_semantics[status_hash] = fs
|
|
|
|
# Ensure pattern has enough data for recommendation
|
|
pattern.total_compressions = 10
|
|
pattern.sample_size = 10 # Required by min_samples_for_recommendation
|
|
pattern.confidence = 0.5
|
|
|
|
# Get recommendation
|
|
hint = toin.get_recommendation(sig)
|
|
|
|
# VERIFY
|
|
assert hint is not None
|
|
assert len(hint.field_semantics) > 0, (
|
|
f"Recommendation should include field_semantics. Got: {hint.field_semantics}"
|
|
)
|
|
assert status_hash in hint.field_semantics
|
|
|
|
|
|
class TestFieldSemanticsMemoryBounds:
|
|
"""Test that memory bounds are enforced."""
|
|
|
|
def test_important_values_bounded(self):
|
|
"""PROVES: important_value_hashes stays within bounds."""
|
|
fs = FieldSemantics(field_hash="test")
|
|
|
|
# Add more values than MAX_IMPORTANT_VALUES
|
|
for i in range(fs.MAX_IMPORTANT_VALUES + 20):
|
|
fs.record_retrieval_value(_hash_value(f"value_{i}"))
|
|
|
|
# VERIFY: bounded
|
|
assert len(fs.important_value_hashes) <= fs.MAX_IMPORTANT_VALUES
|
|
|
|
def test_value_frequency_bounded(self):
|
|
"""PROVES: value_retrieval_frequency stays within bounds."""
|
|
fs = FieldSemantics(field_hash="test")
|
|
|
|
# Add more values than MAX_VALUE_FREQUENCY_ENTRIES
|
|
for i in range(fs.MAX_VALUE_FREQUENCY_ENTRIES + 20):
|
|
fs.record_retrieval_value(_hash_value(f"value_{i}"))
|
|
|
|
# VERIFY: bounded
|
|
assert len(fs.value_retrieval_frequency) <= fs.MAX_VALUE_FREQUENCY_ENTRIES
|
|
|
|
|
|
class TestProductionCodePath:
|
|
"""Integration tests for the ACTUAL production code path.
|
|
|
|
These tests verify that CompressionStore -> TOIN integration works,
|
|
not just TOIN in isolation. This is critical because the unit tests
|
|
can pass while the production integration is broken.
|
|
"""
|
|
|
|
def test_compression_store_passes_items_to_toin(self):
|
|
"""PROVES: CompressionStore.process_pending_feedback passes retrieved_items to TOIN.
|
|
|
|
This is the integration test that would have caught the original bug
|
|
where retrieved_items was never passed to TOIN in production.
|
|
"""
|
|
import json
|
|
|
|
from headroom.cache.compression_store import CompressionStore
|
|
from headroom.telemetry.toin import get_toin, reset_toin
|
|
|
|
reset_toin()
|
|
toin = get_toin()
|
|
|
|
# Create a store with feedback enabled
|
|
store = CompressionStore(max_entries=100, default_ttl=300, enable_feedback=True)
|
|
|
|
# Store some compressed content with items that have distinct field values
|
|
items = [
|
|
{"id": "123", "status": "success", "value": 100},
|
|
{"id": "456", "status": "error", "value": 200},
|
|
{"id": "789", "status": "success", "value": 300},
|
|
]
|
|
compressed_json = json.dumps(items)
|
|
original_json = json.dumps(items * 10) # Original was bigger
|
|
|
|
# Create a tool signature hash for this structure
|
|
sig = ToolSignature.from_items(items)
|
|
|
|
hash_key = store.store(
|
|
original=original_json,
|
|
compressed=compressed_json,
|
|
original_tokens=1000,
|
|
compressed_tokens=100,
|
|
original_item_count=30,
|
|
compressed_item_count=3,
|
|
tool_name="test_api",
|
|
tool_call_id="call_123",
|
|
tool_signature_hash=sig.structure_hash,
|
|
compression_strategy="top_n",
|
|
)
|
|
|
|
# Simulate a retrieval (this triggers the feedback loop)
|
|
store.retrieve(hash_key, query="status=error")
|
|
|
|
# Process pending feedback - THIS IS WHERE THE BUG WAS
|
|
store.process_pending_feedback()
|
|
|
|
# VERIFY: TOIN should have received the items and learned from them
|
|
pattern = toin._patterns.get(sig.structure_hash)
|
|
assert pattern is not None, "TOIN should have a pattern for this tool"
|
|
|
|
# The key assertion: field_semantics should be populated
|
|
# This would have FAILED before the fix because retrieved_items wasn't passed
|
|
assert len(pattern.field_semantics) > 0, (
|
|
"TOIN should have learned field semantics from the retrieved items. "
|
|
"If this fails, CompressionStore is not passing retrieved_items to TOIN."
|
|
)
|
|
|
|
# Verify specific fields were learned
|
|
id_hash = _hash_field("id")
|
|
status_hash = _hash_field("status")
|
|
value_hash = _hash_field("value")
|
|
|
|
learned_fields = set(pattern.field_semantics.keys())
|
|
expected_fields = {id_hash, status_hash, value_hash}
|
|
assert expected_fields == learned_fields, (
|
|
f"Expected fields {expected_fields}, got {learned_fields}"
|
|
)
|
|
|
|
def test_compression_store_handles_wrapped_arrays(self):
|
|
"""PROVES: CompressionStore correctly extracts items from wrapped arrays."""
|
|
import json
|
|
|
|
from headroom.cache.compression_store import CompressionStore
|
|
from headroom.telemetry.toin import get_toin, reset_toin
|
|
|
|
reset_toin()
|
|
toin = get_toin()
|
|
|
|
store = CompressionStore(max_entries=100, default_ttl=300, enable_feedback=True)
|
|
|
|
# Content wrapped in {"results": [...]} pattern
|
|
items = [{"name": "test", "score": 0.95}]
|
|
wrapped_content = json.dumps({"results": items, "total": 1})
|
|
|
|
sig = ToolSignature.from_items(items)
|
|
|
|
hash_key = store.store(
|
|
original=wrapped_content,
|
|
compressed=wrapped_content,
|
|
original_tokens=100,
|
|
compressed_tokens=100,
|
|
original_item_count=1,
|
|
compressed_item_count=1,
|
|
tool_name="search_api",
|
|
tool_call_id="call_456",
|
|
tool_signature_hash=sig.structure_hash,
|
|
compression_strategy="top_n",
|
|
)
|
|
|
|
store.retrieve(hash_key)
|
|
store.process_pending_feedback()
|
|
|
|
# VERIFY: Items were extracted from wrapped structure
|
|
pattern = toin._patterns.get(sig.structure_hash)
|
|
assert pattern is not None
|
|
assert len(pattern.field_semantics) > 0, (
|
|
"TOIN should extract items from wrapped arrays like {'results': [...]}"
|
|
)
|
|
|
|
def test_compression_store_handles_invalid_json(self):
|
|
"""PROVES: CompressionStore gracefully handles invalid JSON."""
|
|
from headroom.cache.compression_store import CompressionStore
|
|
from headroom.telemetry.toin import get_toin, reset_toin
|
|
|
|
reset_toin()
|
|
get_toin() # Initialize TOIN for feedback loop
|
|
|
|
store = CompressionStore(max_entries=100, default_ttl=300, enable_feedback=True)
|
|
|
|
# Store invalid JSON content
|
|
invalid_json = "not valid json {"
|
|
|
|
hash_key = store.store(
|
|
original=invalid_json,
|
|
compressed=invalid_json,
|
|
original_tokens=10,
|
|
compressed_tokens=10,
|
|
original_item_count=0,
|
|
compressed_item_count=0,
|
|
tool_name="broken_api",
|
|
tool_call_id="call_789",
|
|
tool_signature_hash="invalid123",
|
|
compression_strategy="none",
|
|
)
|
|
|
|
# This should not crash
|
|
store.retrieve(hash_key)
|
|
store.process_pending_feedback()
|
|
|
|
# VERIFY: No crash, pattern may or may not exist but no exception
|
|
# The main assertion is that we got here without exception
|