mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(memory): replay complete saved Responses conversations
This commit is contained in:
parent
396b18b74b
commit
c530e5085b
3 changed files with 66 additions and 26 deletions
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue