fix(memory): replay complete saved Responses conversations

This commit is contained in:
moe-berri 2026-09-15 17:10:08 -07:00
parent 396b18b74b
commit c530e5085b
3 changed files with 66 additions and 26 deletions

View file

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

View file

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

View file

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