mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
219 lines
7.5 KiB
Python
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
|