mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-19 00:01:29 +00:00
fix(memory): retain duplicate directive cache positions
This commit is contained in:
parent
ec294d1b91
commit
91ba66db9e
2 changed files with 32 additions and 8 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue