diff --git a/litellm/litellm_core_utils/prompt_templates/server_tools.py b/litellm/litellm_core_utils/prompt_templates/server_tools.py index 81fde70543b..b6c689af8de 100644 --- a/litellm/litellm_core_utils/prompt_templates/server_tools.py +++ b/litellm/litellm_core_utils/prompt_templates/server_tools.py @@ -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") diff --git a/litellm/proxy/memory/continuation.py b/litellm/proxy/memory/continuation.py index 293247969a5..ce9d06c3bd7 100644 --- a/litellm/proxy/memory/continuation.py +++ b/litellm/proxy/memory/continuation.py @@ -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: diff --git a/litellm/proxy/memory/gateway.py b/litellm/proxy/memory/gateway.py index 4db3677d429..7c4da315b67 100644 --- a/litellm/proxy/memory/gateway.py +++ b/litellm/proxy/memory/gateway.py @@ -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 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 c58eb08bc40..53a77cb6251 100644 --- a/tests/test_litellm/proxy/memory/test_memory_v2_boundaries.py +++ b/tests/test_litellm/proxy/memory/test_memory_v2_boundaries.py @@ -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", diff --git a/tests/test_litellm/proxy/memory/test_memory_v2_protocols.py b/tests/test_litellm/proxy/memory/test_memory_v2_protocols.py index 5173d19d460..5c416372a47 100644 --- a/tests/test_litellm/proxy/memory/test_memory_v2_protocols.py +++ b/tests/test_litellm/proxy/memory/test_memory_v2_protocols.py @@ -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",