fix(memory): keep structured streams alive during preparation

This commit is contained in:
moe-berri 2026-09-15 00:14:03 -07:00
parent 4450bfdf79
commit eb72042a0a
7 changed files with 113 additions and 75 deletions

View file

@ -165,20 +165,7 @@ class ServerToolStream:
self.indices = MappingProxyType({**self.indices, index: None if hidden else self.content_count})
if not hidden:
self.content_count += 1
if isinstance(index, int):
mapped: Final = self.indices.get(index)
if mapped is None:
return ()
return (
self._emit(
{ # mutable-ok: Native provider JSON containers.
**data,
"index": mapped,
},
str(kind),
),
)
return (self._emit(data, str(kind)),)
return self._indexed_event(data, "index", str(kind))
def _responses(self, data: Mapping[str, object]) -> tuple[bytes, ...]:
kind: Final = str(data.get("type", ""))
@ -212,20 +199,24 @@ class ServerToolStream:
self.indices = MappingProxyType({**self.indices, index: None if hidden else self.content_count})
if not hidden:
self.content_count += 1
if isinstance(index, int):
mapped: Final = self.indices.get(index)
if mapped is None:
return ()
return (
self._emit(
{ # mutable-ok: Native provider JSON containers.
**data,
"output_index": mapped,
},
kind,
),
)
return (self._emit(data, kind),)
return self._indexed_event(data, "output_index", kind)
def _indexed_event(self, data: Mapping[str, object], index_field: str, event: str) -> tuple[bytes, ...]:
index: Final = data.get(index_field)
if not isinstance(index, int):
return (self._emit(data, event),)
mapped: Final = self.indices.get(index)
if mapped is None:
return ()
return (
self._emit(
{ # mutable-ok: Native provider JSON containers.
**data,
index_field: mapped,
},
event,
),
)
def _chat(self, data: Mapping[str, object]) -> tuple[bytes, ...]:
if self.response_id is None:

View file

@ -1015,7 +1015,9 @@ _TTFT_KEEPALIVE_HEADERS: Final[Mapping[str, str]] = MappingProxyType(
)
def ttft_keepalive_interval(request_data: Mapping[str, object], llm_router: Router | None = None) -> float | None:
def ttft_keepalive_interval(
request_data: Mapping[str, object], llm_router: Router | None = None, *, default_interval: float | None = None
) -> float | None:
"""The operator's keepalive interval, but only for a request that asked to stream.
Resolved through the deployments the request could land on, so a deployment's
@ -1030,7 +1032,8 @@ def ttft_keepalive_interval(request_data: Mapping[str, object], llm_router: Rout
if llm_router is not None and isinstance(requested_model, str)
else ()
)
return resolve_ttft_keepalive_interval(deployments, litellm.sse_keepalive_ping_interval_seconds)
configured: Final = litellm.sse_keepalive_ping_interval_seconds
return resolve_ttft_keepalive_interval(deployments, default_interval if configured is None else configured)
async def _aclose_late_response(produced: Response) -> None:

View file

@ -216,9 +216,6 @@ class MemoryContinuations:
)
return self.validate_patch(row.payload) if row is not None else None
async def save(self, anchor: str, patch: MemoryContinuation) -> None:
await self.save_many(((anchor, patch),))
async def save_many(self, patches: tuple[tuple[str, MemoryContinuation], ...]) -> None:
namespace: Final = await self.store.authorize_namespace()
payloads: Final = tuple(

View file

@ -414,10 +414,18 @@ async def process_gateway_memory(
raise HTTPException(status_code=404, detail="Memory response not found or expired")
return None
validate_memory_request(data, request)
from litellm.proxy.common_request_processing import (
_UpstreamClosingStreamingResponse, # pyright: ignore[reportPrivateUsage] # Reuse disconnect cleanup for the prefetched stream.
ttft_keepalive_interval,
)
from litellm.proxy.proxy_server import app, llm_router
loop: Final = GatewayMemoryLoop(app, request, data, route, store)
iterator: Final = loop.run()
iterator: Final = wrap_passthrough_sse_bytes_with_keepalive_pings(
loop.run(),
ping_interval_seconds=ttft_keepalive_interval(data, llm_router, default_interval=5.0),
upstream_headers=MappingProxyType({"content-type": "text/event-stream"}),
)
if not loop.streaming:
async for _ in iterator:
pass
@ -449,17 +457,8 @@ async def process_gateway_memory(
finally:
await iterator.aclose()
from litellm.proxy.common_request_processing import (
_UpstreamClosingStreamingResponse, # pyright: ignore[reportPrivateUsage] # Reuse cleanup when a client disconnects before consuming the prefetched stream.
ttft_keepalive_interval,
)
return _UpstreamClosingStreamingResponse(
wrap_passthrough_sse_bytes_with_keepalive_pings(
stream(),
ping_interval_seconds=ttft_keepalive_interval(data, llm_router),
upstream_headers=MappingProxyType({"content-type": "text/event-stream"}),
),
stream(),
media_type="text/event-stream",
headers={ # mutable-ok: Native ASGI response headers.
**loop.response_headers(),

View file

@ -1,6 +1,5 @@
import json
from datetime import datetime, timezone
from typing import Final
import pytest
from collections.abc import Mapping
@ -54,26 +53,6 @@ def _local_model_cost_map(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
@pytest.mark.parametrize(
"duration,second_rate,expected",
((2.0, 0.00016, 0.00032), (2.0, 0.0, 0.0), (2.0, None, 0.000416), (0.0, 0.00016, 0.000416)),
)
def test_audio_duration_and_tokens_bill_only_once(duration: float, second_rate: float | None, expected: float) -> None:
info: Final[ModelInfo] = {
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
"input_cost_per_audio_token": 6.5e-6,
"input_cost_per_audio_per_second": second_rate,
}
usage: Final = Usage(
prompt_tokens=64,
completion_tokens=0,
prompt_tokens_details=PromptTokensDetailsWrapper(audio_tokens=64, audio_length_seconds=duration, text_tokens=0),
)
cost, _ = generic_cost_per_token("audio-billing-fixture", usage, "vertex_ai", model_info=info)
assert cost == pytest.approx(expected)
@pytest.mark.parametrize("prompt_tokens", [100, 200000, 200001])
@pytest.mark.parametrize("read_rate", [None, 0.0, 0.25e-6])
@pytest.mark.parametrize("service_tier", [None, "priority"])

View file

@ -1073,7 +1073,9 @@ async def test_continuation_quota_rejects_excess_without_writing(
) -> None:
prisma_edge.db.query_raw.return_value = [{"key_count": key_count, "bytes": used_bytes}]
with pytest.raises(HTTPException) as exc:
await MemoryContinuations(store(prisma_edge), "aresponses").save("response", MemoryContinuation(replaces=1))
await MemoryContinuations(store(prisma_edge), "aresponses").save_many(
(("response", MemoryContinuation(replaces=1)),)
)
assert exc.value.status_code == 429
prisma_edge.db.litellm_memorycontinuation.upsert.assert_not_awaited()
@ -1086,7 +1088,7 @@ async def test_continuation_quota_shares_namespace_lock_across_keys_and_allows_r
other_key = MemoryIdentity("b" * 64, "owner", "team", "project", "org", False)
for identity in (_IDENTITY, other_key):
continuations = MemoryContinuations(MemoryStore(prisma_edge, access_for(identity)), "aresponses")
await continuations.save("response", MemoryContinuation(replaces=1, response={"text": "é漢字"}))
await continuations.save_many((("response", MemoryContinuation(replaces=1, response={"text": "é漢字"})),))
query = prisma_edge.db.query_raw.call_args.args
assert query[1:4] == (identity.namespace, identity.key_id, [continuations.identifier("response")])
assert json.loads(query[4])[0]["response"]["text"] == "é漢字"
@ -1313,3 +1315,65 @@ async def test_structured_output_hides_preparation_and_restores_final_constraint
elif route == "anthropic_messages":
assert wire.count(b'"type": "message_start"') == 1
assert wire.count(b'"type": "message_stop"') == 1
@pytest.mark.asyncio
@pytest.mark.parametrize("configured_interval", (None, 0, 0.01))
async def test_structured_preparation_pings_before_headers_and_honors_explicit_disable(
prisma_edge: MagicMock, configured_interval: float | None
) -> None:
import asyncio
from unittest.mock import patch
from starlette.responses import StreamingResponse
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.memory.gateway import process_gateway_memory
waiting = asyncio.Event()
release = asyncio.Event()
closed = asyncio.Event()
async def provider(scope: Scope, receive: Receive, send: Send) -> None:
body = json.loads((await receive())["body"])
assert body["stream"] is False and "response_format" not in body
waiting.set()
try:
await release.wait()
await JSONResponse({"error": "private provider details"}, status_code=429)(scope, receive, send)
finally:
closed.set()
with (
patch( # test-quality-ok: Exercise default and explicit operator configuration through the real keepalive selector.
"litellm.sse_keepalive_ping_interval_seconds", configured_interval
),
patch.multiple( # test-quality-ok: Inject the external provider HTTP boundary and preserve real gateway dispatch.
"litellm.proxy.proxy_server", app=provider, llm_router=None
),
patch( # test-quality-ok: Inject authorized persistence; execute the actual memory loop and keepalive wrapper.
"litellm.proxy.memory.gateway.gateway_memory_store", new=AsyncMock(return_value=store(prisma_edge))
),
):
pending = asyncio.create_task(
process_gateway_memory(
{"model": "test", "stream": True, "messages": [], "response_format": {"type": "json_object"}},
request(),
UserAPIKeyAuth(),
"acompletion",
)
)
await asyncio.wait_for(waiting.wait(), timeout=1)
if configured_interval == 0:
with pytest.raises(TimeoutError):
await asyncio.wait_for(pending, timeout=0.05)
else:
response = await asyncio.wait_for(pending, timeout=6)
assert isinstance(response, StreamingResponse)
assert not release.is_set() and not closed.is_set()
assert await anext(response.body_iterator) == b": ping\n\n"
release.set()
remaining = b"".join([chunk async for chunk in response.body_iterator])
assert b'"code": "429"' in remaining and b"private provider details" not in remaining
assert closed.is_set()
prisma_edge.db.litellm_memorycontinuation.upsert.assert_not_awaited()

View file

@ -7793,18 +7793,23 @@ async def test_ttft_keepalive_cancels_the_in_flight_call_when_the_client_gives_u
@pytest.mark.parametrize(
"request_data, global_interval, expected",
"request_data, global_interval, default_interval, expected",
[
({"stream": True}, 30.0, 30.0),
({"stream": True}, None, None),
({"stream": False}, 30.0, None),
({}, 30.0, None),
({"stream": "true"}, 30.0, None),
({"stream": True}, 30.0, None, 30.0),
({"stream": True}, None, None, None),
({"stream": True}, None, 5.0, 5.0),
({"stream": True}, 0, 5.0, None),
({"stream": True}, 30.0, 5.0, 30.0),
({"stream": False}, 30.0, 5.0, None),
({}, 30.0, 5.0, None),
({"stream": "true"}, 30.0, 5.0, None),
],
)
def test_ttft_keepalive_interval_only_arms_for_a_streaming_request(request_data, global_interval, expected):
def test_ttft_keepalive_interval_only_arms_for_a_streaming_request(
request_data, global_interval, default_interval, expected
):
with patch.object(litellm, "sse_keepalive_ping_interval_seconds", global_interval):
assert ttft_keepalive_interval(request_data) == expected
assert ttft_keepalive_interval(request_data, default_interval=default_interval) == expected
@pytest.mark.asyncio