litellm/tests/test_litellm/compression/test_compress.py
jesus 71186e4ec1 merge: main into litellm_headroom_protect_cached_prefix
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
2026-09-15 04:12:28 +00:00

219 lines
7.5 KiB
Python

"""
Unit tests for litellm.compression.compress helpers.
get_protected_indices is the shared policy for which messages a compressor may
never rewrite. It is consumed by compress() and by the Headroom guardrail, so
the two agree on what "never compress this" means.
"""
from litellm.compression.compress import compress, get_protected_indices
from litellm.types.utils import CallTypes
def test_protects_system_last_user_and_last_assistant():
messages = [
{"role": "system", "content": "sys"},
{"role": "user", "content": "old question"},
{"role": "assistant", "content": "old answer"},
{"role": "user", "content": "newer question"},
{"role": "assistant", "content": "newer answer"},
{"role": "user", "content": "live instruction"},
]
assert sorted(get_protected_indices(messages)) == [0, 4, 5]
def test_history_is_not_protected():
messages = [
{"role": "user", "content": "old question"},
{"role": "assistant", "content": "old answer"},
{"role": "tool", "tool_call_id": "t1", "content": "old tool output"},
{"role": "user", "content": "live instruction"},
]
protected = sorted(get_protected_indices(messages))
assert protected == [1, 3]
# The tool row and the older user turn stay compressible; protection that
# covered everything would make compression a no-op.
assert 0 not in protected
assert 2 not in protected
def test_every_system_row_is_protected():
messages = [
{"role": "system", "content": "first"},
{"role": "user", "content": "q"},
{"role": "system", "content": "second, injected mid conversation"},
{"role": "user", "content": "live"},
]
assert sorted(get_protected_indices(messages)) == [0, 2, 3]
def test_no_user_or_assistant_rows():
assert sorted(get_protected_indices([{"role": "system", "content": "sys"}])) == [0]
assert get_protected_indices([]) == ()
def test_rows_before_last_cache_control_breakpoint_are_protected():
messages = [
{"role": "system", "content": "sys"},
{"role": "user", "content": "old question"},
{
"role": "assistant",
"content": "old answer",
"tool_calls": [{"id": "t1", "type": "function", "function": {"name": "Read", "arguments": "{}"}}],
},
{"role": "tool", "tool_call_id": "t1", "content": "large file body"},
{
"role": "user",
"content": [{"type": "text", "text": "cached turn", "cache_control": {"type": "ephemeral"}}],
},
{
"role": "assistant",
"content": "ack",
"tool_calls": [{"id": "t2", "type": "function", "function": {"name": "Bash", "arguments": "{}"}}],
},
{"role": "tool", "tool_call_id": "t2", "content": "later tool output"},
{"role": "user", "content": "live instruction"},
]
protected = sorted(get_protected_indices(messages))
assert protected == [0, 1, 2, 3, 4, 5, 7]
assert 6 not in protected
def test_cache_control_directly_on_message_protects_prefix():
messages = [
{"role": "system", "content": "sys"},
{"role": "tool", "tool_call_id": "before", "content": "large file body"},
{"role": "user", "content": "old question"},
{
"role": "tool",
"tool_call_id": "marked",
"content": "cached tool",
"cache_control": {"type": "ephemeral"},
},
{"role": "tool", "tool_call_id": "after", "content": "later tool output"},
{"role": "assistant", "content": "ack"},
{"role": "user", "content": "live instruction"},
]
protected = sorted(get_protected_indices(messages))
assert 1 in protected
assert 3 in protected
assert 4 not in protected
def test_no_cache_control_leaves_history_compressible():
messages = [
{"role": "system", "content": "sys"},
{"role": "user", "content": "old question"},
{"role": "assistant", "content": "old answer"},
{"role": "tool", "tool_call_id": "t1", "content": "large file body"},
{"role": "user", "content": "live instruction"},
]
assert sorted(get_protected_indices(messages)) == [0, 2, 4]
def test_non_mapping_content_parts_are_not_cache_control():
messages = [
{"role": "system", "content": "sys"},
{"role": "user", "content": ["not", "a", "dict"]},
{"role": "assistant", "content": "old answer"},
{"role": "tool", "tool_call_id": "t1", "content": "plain string"},
{"role": "user", "content": "live instruction"},
]
protected = sorted(get_protected_indices(messages))
assert protected == [0, 2, 4]
assert 1 not in protected
assert 3 not in protected
def test_mid_history_cache_control_part_is_protected():
messages = [
{"role": "user", "content": "old question"},
{"role": "assistant", "content": "old answer"},
{
"role": "user",
"content": [
{"type": "text", "text": "a large cached tool result", "cache_control": {"type": "ephemeral"}},
],
},
{"role": "assistant", "content": "ack"},
{"role": "user", "content": "live instruction"},
]
assert sorted(get_protected_indices(messages)) == [0, 1, 2, 3, 4]
def test_cache_control_directly_on_message_is_protected():
messages = [
{"role": "user", "content": "old question", "cache_control": {"type": "ephemeral"}},
{"role": "assistant", "content": "old answer"},
{"role": "user", "content": "live instruction"},
]
assert sorted(get_protected_indices(messages)) == [0, 1, 2]
def test_cache_control_protection_does_not_duplicate_already_protected_rows():
# The last user row is already protected by role; marking it too must not
# produce a duplicate index.
messages = [
{"role": "system", "content": "sys"},
{"role": "user", "content": "live", "cache_control": {"type": "ephemeral"}},
]
protected = get_protected_indices(messages)
assert sorted(protected) == [0, 1]
assert len(protected) == len(set(protected))
def test_content_that_is_not_a_list_of_mappings_is_not_treated_as_cache_control():
# Defensive: a plain string content, or a list of non-dict items, must not
# raise or be misread as carrying a breakpoint.
messages = [
{"role": "assistant", "content": "plain string content"},
{"role": "user", "content": ["not", "a", "dict", "list"]},
{"role": "user", "content": "live instruction"},
]
assert sorted(get_protected_indices(messages)) == [0, 2]
def test_compress_keeps_part_level_cache_control_row_verbatim():
stale_log = {"role": "user", "content": [{"type": "text", "text": "stale log line " * 2000}]}
pinned = {
"role": "user",
"content": [
{"type": "text", "text": "cached tool result " * 2000, "cache_control": {"type": "ephemeral"}},
],
}
messages = [
pinned,
{"role": "assistant", "content": "old answer"},
stale_log,
{"role": "assistant", "content": "ack"},
{"role": "user", "content": "live instruction"},
]
result = compress(
messages,
model="gpt-4o",
call_type=CallTypes.anthropic_messages,
compression_trigger=1000,
compression_target=500,
)
assert len(result["messages"]) == len(messages)
assert result["messages"][0] == pinned
assert result["messages"][2] != stale_log
assert len(result["cache"]) >= 1