From 91ba66db9ee7d4aea0dde5dff125d1a524293934 Mon Sep 17 00:00:00 2001 From: moe-berri Date: Mon, 14 Sep 2026 14:25:31 -0700 Subject: [PATCH] fix(memory): retain duplicate directive cache positions --- litellm/proxy/memory/continuation.py | 17 +++++++------- .../proxy/memory/test_memory_v2_boundaries.py | 23 +++++++++++++++++++ 2 files changed, 32 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/memory/continuation.py b/litellm/proxy/memory/continuation.py index c916b76c532..8cd6e0ea49e 100644 --- a/litellm/proxy/memory/continuation.py +++ b/litellm/proxy/memory/continuation.py @@ -179,16 +179,17 @@ class MemoryContinuations: prefix: Final = result[: -(patch.replaces - 1)] if patch.replaces > 1 else result # Clients move cache breakpoints between turns. Reuse their current # directives rather than restoring an obsolete cached copy. - directives: Final = MappingProxyType( - { - prefix_hashes((item,), self.route)[0]: item - for item in items[index + 1 - patch.replaces : index + 1] - if item.get("role") == "system" - } + current_directives: Final = tuple( + item for item in items[index + 1 - patch.replaces : index + 1] if item.get("role") == "system" ) + positions: Final = tuple( + position for position, item in enumerate(patch.replacement) if item.get("role") == "system" + ) + if len(positions) != len(current_directives): + raise HTTPException(status_code=409, detail="Invalid memory continuation directives") + directives: Final = MappingProxyType(dict(zip(positions, current_directives))) replacement: Final = tuple( - directives.get(prefix_hashes((item,), self.route)[0], item) if item.get("role") == "system" else item - for item in patch.replacement + directives.get(position, item) for position, item in enumerate(patch.replacement) ) return _append_items(prefix, replacement, self.route) diff --git a/tests/test_litellm/proxy/memory/test_memory_v2_boundaries.py b/tests/test_litellm/proxy/memory/test_memory_v2_boundaries.py index 650c71939e5..64d0a2be176 100644 --- a/tests/test_litellm/proxy/memory/test_memory_v2_boundaries.py +++ b/tests/test_litellm/proxy/memory/test_memory_v2_boundaries.py @@ -655,6 +655,29 @@ async def test_claude_output_directives_reach_search_answer_and_reflection_round assert sum(message["role"] == "system" for message in object_items(loop.data.get("messages"))) == 1 +@pytest.mark.asyncio +async def test_duplicate_directives_preserve_each_current_cache_breakpoint(prisma_edge: MagicMock) -> None: + first = {"role": "system", "content": [{"type": "text", "text": "Same directive"}]} + second = { + "role": "system", + "content": [{"type": "text", "text": "Same directive", "cache_control": {"type": "ephemeral"}}], + } + assistant = {"role": "assistant", "content": [{"type": "text", "text": "Reply"}]} + items = (first, second, assistant) + continuations = MemoryContinuations(store(prisma_edge), "anthropic_messages") + patch = MemoryContinuation( + replaces=3, replacement=({"role": "user", "content": "Memory reference"}, second, first, assistant) + ) + prisma_edge.db.litellm_memorycontinuation.find_many.return_value = [ + SimpleNamespace( + id=continuations.identifier(prefix_hashes(items, "anthropic_messages")[-1]), payload=patch.model_dump() + ) + ] + restored = await continuations.restore(items) + assert restored[1:3] == (first, second) + assert restored[-1] == assistant + + @pytest.mark.asyncio @pytest.mark.parametrize("bad_id,count", [(True, 1), (False, 17)]) async def test_invalid_model_calls_are_rejected_before_storage(