fix(memory): retain duplicate directive cache positions
Some checks are pending
ai-gateway image / ai-gateway release image (push) Waiting to run
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run

This commit is contained in:
moe-berri 2026-09-14 14:25:31 -07:00
parent ec294d1b91
commit 91ba66db9e
2 changed files with 32 additions and 8 deletions

View file

@ -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)

View file

@ -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(