From c530e5085b58b1814249871a803bcb8462604429 Mon Sep 17 00:00:00 2001 From: moe-berri Date: Tue, 15 Sep 2026 17:10:08 -0700 Subject: [PATCH] fix(memory): replay complete saved Responses conversations --- litellm/proxy/memory/continuation.py | 3 +- litellm/proxy/memory/gateway.py | 34 +++++------- .../proxy/memory/test_memory_v2_boundaries.py | 55 +++++++++++++++++-- 3 files changed, 66 insertions(+), 26 deletions(-) diff --git a/litellm/proxy/memory/continuation.py b/litellm/proxy/memory/continuation.py index be0efbfb04a..293247969a5 100644 --- a/litellm/proxy/memory/continuation.py +++ b/litellm/proxy/memory/continuation.py @@ -48,7 +48,8 @@ class MemoryContinuation(BaseModel): response: Mapping[str, object] | None = None upstream_ids: tuple[str, ...] = () - pending_results: tuple[Mapping[str, object], ...] = () + input: tuple[Mapping[str, object], ...] | None = None + previous_response_id: str | None = None permission_revision: str | None = None diff --git a/litellm/proxy/memory/gateway.py b/litellm/proxy/memory/gateway.py index 77a089d0330..4db3677d429 100644 --- a/litellm/proxy/memory/gateway.py +++ b/litellm/proxy/memory/gateway.py @@ -89,7 +89,6 @@ class GatewayMemoryLoop: self.last_response: Mapping[str, object] | None = None self.headers: Mapping[str, str] = MappingProxyType({}) self.costs: tuple[float | None, ...] = () - self.pending_results: tuple[Mapping[str, object], ...] = () async def prepare(self) -> None: previous: Final = self.original.get("previous_response_id") @@ -102,20 +101,26 @@ class GatewayMemoryLoop: ) if isinstance(previous, str) and previous.startswith("resp_litellm_memory_") and previous_patch is None: raise HTTPException(status_code=404, detail="Memory response not found or expired") + if previous_patch is not None and previous_patch.input is None: + raise HTTPException(status_code=409, detail="Memory response history unavailable; start a new conversation") field: Final = "input" if self.route == "aresponses" else "messages" functions: Final = memory_functions(self.store.access) injected: Final = inject_server_tools( { # mutable-ok: Native provider JSON containers. - **self.original, + **{ + key: value + for key, value in self.original.items() + if previous_patch is None or key not in ("previous_response_id", "_litellm_addressed_response_id") + }, field: [ # mutable-ok: Native provider JSON containers. - *(previous_patch.pending_results if previous_patch else ()), + *((previous_patch.input or ()) if previous_patch else ()), *self.visible_input, ], **( { # mutable-ok: Native provider JSON containers. - "previous_response_id": previous_patch.upstream_ids[-1] + "previous_response_id": previous_patch.previous_response_id } - if previous_patch + if previous_patch and previous_patch.previous_response_id else { # mutable-ok: Native provider JSON containers. } ), @@ -258,11 +263,15 @@ class GatewayMemoryLoop: if self.continuations is None or self.original.get("store") is False: return response: Final = self.stream.response() + root_response: Final = self.data.get("previous_response_id") try: await self.continuations.save( str(response["id"]), MemoryContinuation( - response=response, upstream_ids=self.upstream_ids, pending_results=self.pending_results + response=response, + upstream_ids=self.upstream_ids, + input=transcript_items(self.data, self.route), + previous_response_id=root_response if isinstance(root_response, str) else None, ), ) except Exception: @@ -308,23 +317,10 @@ class GatewayMemoryLoop: [await execute_memory_tool(self.store, call, self.visible_input) for call in memory_calls] ) if memory_calls: - self.pending_results = ( - tuple( - { # mutable-ok: Native provider JSON containers. - "type": "function_call_output", - "call_id": call["id"], - "output": json.dumps(dict(result)), - } - for call, result in zip(memory_calls, results) - ) - if self.route == "aresponses" - else () - ) self.data = continue_server_tools( self.data, self.route, response, memory_calls, tuple(dict(result) for result in results) ) else: - self.pending_results = () field: Final = "input" if self.route == "aresponses" else "messages" self.data = { # mutable-ok: Native provider JSON containers. **self.data, 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 7847b647187..c58eb08bc40 100644 --- a/tests/test_litellm/proxy/memory/test_memory_v2_boundaries.py +++ b/tests/test_litellm/proxy/memory/test_memory_v2_boundaries.py @@ -106,6 +106,7 @@ async def test_saved_response_reads_are_scoped_and_never_return_internal_input( permission_revision=access_for().continuation_revision, response={"id": "resp_litellm_memory_test", "output": [{"type": "message", "content": []}]}, upstream_ids=("native=one",), + input=({"type": "function_call_output", "call_id": "private", "output": "Private memory content"},), ) prisma_edge.db.litellm_memorycontinuation.find_first.return_value = ( None if operation == "missing" else SimpleNamespace(payload=patch.model_dump()) @@ -1110,13 +1111,23 @@ async def test_rounds_share_trace_and_keep_live_auth_objects(prisma_edge: MagicM @pytest.mark.asyncio -async def test_previous_response_uses_owned_upstream_and_pending_tool_outputs(prisma_edge: MagicMock) -> None: +@pytest.mark.parametrize("root_response", (None, "native-before-memory")) +async def test_previous_response_replays_complete_history_and_pending_client_tools( + prisma_edge: MagicMock, root_response: str | None +) -> None: pending = {"type": "function_call_output", "call_id": "memory-call", "output": "Memory saved"} + history = ( + {"role": "user", "content": "Read the file and recall its port"}, + {"type": "function_call", "name": "litellm_memory_search", "call_id": "memory-call", "arguments": "{}"}, + {"type": "function_call", "name": "read_file", "call_id": "client-call", "arguments": "{}"}, + pending, + ) patch = MemoryContinuation( permission_revision=access_for().continuation_revision, response={"id": "resp_litellm_memory_owned"}, upstream_ids=("native-first", "native-last"), - pending_results=(pending,), + input=history, + previous_response_id=root_response, ) prisma_edge.db.litellm_memorycontinuation.find_first.return_value = SimpleNamespace(payload=patch.model_dump()) execute = AsyncMock(return_value=JSONResponse(provider_response("aresponses", "answer"))) @@ -1126,7 +1137,7 @@ async def test_previous_response_uses_owned_upstream_and_pending_tool_outputs(pr { "previous_response_id": "resp_litellm_memory_owned", "input": [{"type": "function_call_output", "call_id": "client-call", "output": "File contents"}], - "store": False, + "_litellm_addressed_response_id": "resp_litellm_memory_owned", }, "aresponses", store(prisma_edge), @@ -1135,9 +1146,41 @@ async def test_previous_response_uses_owned_upstream_and_pending_tool_outputs(pr async for _ in loop.run(): pass body = execute.call_args.args[1] - assert body["previous_response_id"] == "native-last" - assert pending in body["input"] - assert any(item.get("call_id") == "client-call" for item in body["input"]) + assert body.get("previous_response_id") == root_response + assert "_litellm_addressed_response_id" not in body + assert body["input"] == [ + *history, + {"type": "function_call_output", "call_id": "client-call", "output": "File contents"}, + ] + saved = json.loads(prisma_edge.db.litellm_memorycontinuation.upsert.call_args.kwargs["data"]["create"]["payload"]) + assert saved["input"] == [*body["input"], *provider_response("aresponses", "answer")["output"]] + assert saved["previous_response_id"] == root_response + assert loop.visible_input == ( + {"type": "function_call_output", "call_id": "client-call", "output": "File contents"}, + ) + + +@pytest.mark.asyncio +async def test_legacy_response_without_history_fails_before_calling_provider(prisma_edge: MagicMock) -> None: + prisma_edge.db.litellm_memorycontinuation.find_first.return_value = SimpleNamespace( + payload=MemoryContinuation( + permission_revision=access_for().continuation_revision, upstream_ids=("native-last",) + ).model_dump() + ) + execute = AsyncMock() + loop = GatewayMemoryLoop( + execute, + request(), + {"previous_response_id": "resp_litellm_memory_old", "input": "Continue"}, + "aresponses", + store(prisma_edge), + UserAPIKeyAuth(), + ) + with pytest.raises(HTTPException) as error: + await loop.prepare() + assert error.value.status_code == 409 + assert "start a new conversation" in error.value.detail + execute.assert_not_awaited() @pytest.mark.asyncio