mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(memory): isolate private tool output budgets
This commit is contained in:
parent
c530e5085b
commit
06860d0311
5 changed files with 142 additions and 21 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue