fix(memory): isolate private tool output budgets

This commit is contained in:
moe-berri 2026-09-15 18:38:43 -07:00
parent c530e5085b
commit 06860d0311
5 changed files with 142 additions and 21 deletions

View file

@ -10,11 +10,16 @@ ServerToolRoute: TypeAlias = Literal["acompletion", "aresponses", "anthropic_mes
_LIST: Final = TypeAdapter(tuple[object, ...])
_OBJECT: Final = TypeAdapter(dict[str, object])
_OUTPUT_FIELDS: Final = frozenset(("response_format", "text", "output_format", "output_config"))
_FINAL_FIELDS: Final = _OUTPUT_FIELDS | frozenset(("tools", "tool_choice", "stream_options"))
_TOKEN_LIMITS: Final = frozenset(("max_tokens", "max_completion_tokens", "max_output_tokens"))
_TOOL_OUTPUT_BUDGET: Final = 4096
_FINAL_FIELDS: Final = _OUTPUT_FIELDS | _TOKEN_LIMITS | frozenset(("tools", "tool_choice", "stream_options"))
def has_server_output_constraint(data: Mapping[str, object]) -> bool:
return any(
isinstance(limit := data.get(field), int) and not isinstance(limit, bool) and 0 < limit < _TOOL_OUTPUT_BUDGET
for field in _TOKEN_LIMITS
) or any(
isinstance(value := data.get(field), dict)
and isinstance(
nested := _OBJECT.validate_python(value).get("format") if field in ("text", "output_config") else value,
@ -25,9 +30,24 @@ def has_server_output_constraint(data: Mapping[str, object]) -> bool:
)
def prepare_server_tool_context(data: Mapping[str, object], server_names: frozenset[str]) -> Mapping[str, object]:
def prepare_server_tool_context(
data: Mapping[str, object], server_names: frozenset[str], route: ServerToolRoute
) -> Mapping[str, object]:
limit_field: Final = (
"max_output_tokens"
if route == "aresponses"
else "max_completion_tokens"
if "max_completion_tokens" in data
else "max_tokens"
)
client_limit: Final = data.get(limit_field)
return { # mutable-ok: Provider wire format requires native JSON containers.
**{key: value for key, value in data.items() if key not in _OUTPUT_FIELDS and key != "stream_options"},
**{
key: value
for key, value in data.items()
if key not in _OUTPUT_FIELDS | _TOKEN_LIMITS and key != "stream_options"
},
limit_field: max(_TOOL_OUTPUT_BUDGET, client_limit if isinstance(client_limit, int) else 0),
**{
key: remainder
for key in ("text", "output_config")

View file

@ -7,6 +7,7 @@ from typing import Final
from fastapi import HTTPException
from pydantic import BaseModel, ConfigDict
from litellm.litellm_core_utils.prompt_templates.server_tool_responses import response_has_client_tools
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
from litellm.proxy.memory.policy import memory_digest, memory_primary_client
from litellm.proxy.memory.store import MemoryStore
@ -52,6 +53,25 @@ class MemoryContinuation(BaseModel):
previous_response_id: str | None = None
permission_revision: str | None = None
def can_resume(self, server_names: frozenset[str]) -> bool:
if self.input is None:
return False
if (
self.response
and self.response.get("status") == "incomplete"
and response_has_client_tools(self.response, "aresponses", frozenset())
):
return False
completed: Final = frozenset(
item.get("call_id") for item in self.input if item.get("type") == "function_call_output"
)
return not any(
item.get("type") == "function_call"
and item.get("name") in server_names
and item.get("call_id") not in completed
for item in self.input
)
class MemoryContinuations:
def __init__(self, store: MemoryStore) -> None:

View file

@ -101,7 +101,7 @@ 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:
if previous_patch is not None and not previous_patch.can_resume(MEMORY_TOOL_NAMES):
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)
@ -134,11 +134,11 @@ class GatewayMemoryLoop:
self.data = injected
if self.preparing_output:
self.data = append_server_reference(
prepare_server_tool_context(self.data, MEMORY_TOOL_NAMES),
prepare_server_tool_context(self.data, MEMORY_TOOL_NAMES, self.route),
self.route,
"If needed, use the available memory tools for this request. "
"The final response will be generated separately with the client's output "
"format and application tools. Do not call application tools during this preparation.",
"requirements and application tools. Do not call application tools during this preparation.",
)
self.prepared = True
@ -264,21 +264,19 @@ class GatewayMemoryLoop:
return
response: Final = self.stream.response()
root_response: Final = self.data.get("previous_response_id")
continuation: Final = MemoryContinuation(
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,
)
try:
await self.continuations.save(
str(response["id"]),
MemoryContinuation(
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,
),
)
if continuation.can_resume(MEMORY_TOOL_NAMES):
await self.continuations.save(str(response["id"]), continuation)
return
except Exception:
verbose_proxy_logger.warning("Memory response retention unavailable; returning the completed answer")
self.stream.client_response_fields = MappingProxyType(
{**self.stream.client_response_fields, "store": False}
)
self.stream.client_response_fields = MappingProxyType({**self.stream.client_response_fields, "store": False})
def response_headers(self) -> Mapping[str, str]:
cost_header: Final = (
@ -331,7 +329,7 @@ class GatewayMemoryLoop:
}
if client_calls or not memory_calls:
return True
if round_index + 2 >= _MAX_ROUNDS:
if not self.preparing_output and round_index + 2 >= _MAX_ROUNDS:
self.data = restore_client_output(self.data, self.original)
return False

View file

@ -1161,10 +1161,19 @@ async def test_previous_response_replays_complete_history_and_pending_client_too
@pytest.mark.asyncio
async def test_legacy_response_without_history_fails_before_calling_provider(prisma_edge: MagicMock) -> None:
@pytest.mark.parametrize(
"history",
(
None,
({"type": "function_call", "name": "litellm_memory_capture", "call_id": "partial", "arguments": '{"key":'},),
),
)
async def test_unusable_response_history_fails_before_calling_provider(prisma_edge: MagicMock, history: object) -> None:
prisma_edge.db.litellm_memorycontinuation.find_first.return_value = SimpleNamespace(
payload=MemoryContinuation(
permission_revision=access_for().continuation_revision, upstream_ids=("native-last",)
permission_revision=access_for().continuation_revision,
upstream_ids=("native-last",),
input=history,
).model_dump()
)
execute = AsyncMock()
@ -1183,6 +1192,79 @@ async def test_legacy_response_without_history_fails_before_calling_provider(pri
execute.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"route,limit_field",
(
("acompletion", "max_tokens"),
("acompletion", "max_completion_tokens"),
("anthropic_messages", "max_tokens"),
("aresponses", "max_output_tokens"),
),
)
@pytest.mark.parametrize("streaming", (False, True))
async def test_private_tools_have_room_but_final_answer_keeps_client_token_cap(
prisma_edge: MagicMock, route: ServerToolRoute, limit_field: str, streaming: bool
) -> None:
observed = []
prisma_edge.db.litellm_memorytable.find_first.return_value = row()
call = {"id": "read-entry", "name": "litellm_memory_read", "arguments": {"id": "entry"}}
async def execute(inner: Request, body: dict[str, object], auth: UserAPIKeyAuth) -> Response:
observed.append(body)
reply = provider_response(route, "ok" if len(observed) == 3 else "", (call,) if len(observed) == 1 else ())
return wire_response(reply, route, body.get("stream") is True)
loop = GatewayMemoryLoop(
execute,
request(),
{
"input": "Recall the port and reply ok",
"messages": [{"role": "user", "content": "Recall the port and reply ok"}],
limit_field: 40,
"stream": streaming,
},
route,
store(prisma_edge),
UserAPIKeyAuth(),
)
chunks = b"".join([chunk async for chunk in loop.run()])
assert [body[limit_field] for body in observed] == [4096, 4096, 40]
assert "litellm_memory_read" not in json.dumps(observed[-1].get("tools", []))
assert "ok" in json.dumps(loop.stream.response())
if streaming:
assert b"ok" in chunks and b"litellm_memory_read" not in chunks
@pytest.mark.asyncio
@pytest.mark.parametrize("streaming", (False, True))
@pytest.mark.parametrize("tool_name", ("litellm_memory_capture", "read_file"))
async def test_truncated_tool_response_is_not_saved_as_a_resumable_conversation(
prisma_edge: MagicMock, streaming: bool, tool_name: str
) -> None:
call = {"id": "partial", "name": tool_name, "arguments": {"key": "unfinished"}}
execute = AsyncMock(
return_value=wire_response(
provider_response("aresponses", "", (call,), truncated=True), "aresponses", streaming
)
)
loop = GatewayMemoryLoop(
execute,
request(),
{"input": "Remember this", "stream": streaming},
"aresponses",
store(prisma_edge),
UserAPIKeyAuth(),
)
chunks = b"".join([chunk async for chunk in loop.run()])
assert loop.stream.response()["status"] == "incomplete"
assert loop.stream.response()["store"] is False
prisma_edge.db.litellm_memorycontinuation.upsert.assert_not_awaited()
prisma_edge.db.litellm_memorytable.create.assert_not_awaited()
if streaming:
assert b'"store": false' in chunks
@pytest.mark.asyncio
@pytest.mark.parametrize(
"body",

View file

@ -36,6 +36,7 @@ def test_structured_output_restores_provider_json_enforcement_after_memory_prepa
"Search memory before answering",
),
frozenset(("memory_search",)),
"acompletion",
)
provider: Final = get_optional_params(
model="claude-sonnet-5",