mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
Merge 6a92487cec into e6c4580a31
This commit is contained in:
commit
d59080100b
2 changed files with 89 additions and 2 deletions
|
|
@ -6,7 +6,7 @@ import math
|
|||
import traceback
|
||||
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from functools import lru_cache
|
||||
from functools import lru_cache, partial
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, Protocol, TypeAlias, TypeVar, overload
|
||||
|
||||
|
|
@ -170,6 +170,8 @@ StreamChunkSerializer = Callable[[Any], str]
|
|||
# Type alias for streaming error serializer (ProxyException -> wire format)
|
||||
StreamErrorSerializer = Callable[[ProxyException], str]
|
||||
|
||||
_disconnect_stream_assembly_limiter: Final = anyio.CapacityLimiter(1)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.proxy_server import ProxyConfig as _ProxyConfig
|
||||
|
||||
|
|
@ -364,11 +366,16 @@ async def _bill_partial_streamed_spend_on_disconnect(request_data: dict, respons
|
|||
)
|
||||
messages: Final[object] = getattr(response, "messages", None)
|
||||
try:
|
||||
partial_response: Final = litellm.stream_chunk_builder(
|
||||
build_partial_response: Final = partial(
|
||||
litellm.stream_chunk_builder,
|
||||
chunks=chunks,
|
||||
messages=messages if isinstance(messages, list) else None,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
partial_response: Final = await anyio.to_thread.run_sync(
|
||||
build_partial_response,
|
||||
limiter=_disconnect_stream_assembly_limiter,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # partial billing is best-effort; never break stream teardown
|
||||
verbose_proxy_logger.debug("Failed to assemble partial streamed response for disconnect billing: %s", e)
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -2,6 +2,8 @@ import asyncio
|
|||
import copy
|
||||
import datetime
|
||||
import json
|
||||
import threading
|
||||
import time
|
||||
from types import SimpleNamespace
|
||||
from typing import AsyncGenerator, Callable, Final, Optional
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
|
@ -5649,6 +5651,84 @@ class TestStreamingClientDisconnectBilling:
|
|||
assert standard_logging_object["total_tokens"] > 0
|
||||
assert standard_logging_object["response_cost"] >= 0.002
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disconnect_partial_response_assembly_does_not_block_event_loop(self):
|
||||
response = await self._start_partial_stream()
|
||||
builder_started = threading.Event()
|
||||
builder_finished = threading.Event()
|
||||
|
||||
def blocking_builder(*args, **kwargs):
|
||||
builder_started.set()
|
||||
time.sleep(0.2)
|
||||
builder_finished.set()
|
||||
|
||||
async def event_loop_progressed_during_assembly():
|
||||
while not builder_started.is_set():
|
||||
await asyncio.sleep(0)
|
||||
await asyncio.sleep(0.01)
|
||||
return not builder_finished.is_set()
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.common_request_processing.litellm.stream_chunk_builder",
|
||||
side_effect=blocking_builder,
|
||||
):
|
||||
billing_task = asyncio.create_task(
|
||||
_bill_partial_streamed_spend_on_disconnect(
|
||||
{"litellm_logging_obj": response.logging_obj}, response
|
||||
)
|
||||
)
|
||||
event_loop_progressed = await event_loop_progressed_during_assembly()
|
||||
billed = await billing_task
|
||||
|
||||
assert event_loop_progressed is True
|
||||
assert billed is False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disconnect_partial_response_assembly_limits_cpu_concurrency(self):
|
||||
response = await self._start_partial_stream()
|
||||
first_builder_started = threading.Event()
|
||||
release_first_builder = threading.Event()
|
||||
second_builder_started = threading.Event()
|
||||
builder_claim_lock = threading.Lock()
|
||||
|
||||
def blocking_builder(*args, **kwargs):
|
||||
with builder_claim_lock:
|
||||
if first_builder_started.is_set():
|
||||
second_builder_started.set()
|
||||
return
|
||||
first_builder_started.set()
|
||||
release_first_builder.wait(timeout=5)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.common_request_processing.litellm.stream_chunk_builder",
|
||||
side_effect=blocking_builder,
|
||||
):
|
||||
billing_tasks = (
|
||||
asyncio.create_task(
|
||||
_bill_partial_streamed_spend_on_disconnect(
|
||||
{"litellm_logging_obj": response.logging_obj}, response
|
||||
)
|
||||
),
|
||||
asyncio.create_task(
|
||||
_bill_partial_streamed_spend_on_disconnect(
|
||||
{"litellm_logging_obj": response.logging_obj}, response
|
||||
)
|
||||
),
|
||||
)
|
||||
try:
|
||||
for _ in range(100):
|
||||
if first_builder_started.is_set():
|
||||
break
|
||||
await asyncio.sleep(0.01)
|
||||
assert first_builder_started.is_set()
|
||||
await asyncio.sleep(0.05)
|
||||
assert second_builder_started.is_set() is False
|
||||
finally:
|
||||
release_first_builder.set()
|
||||
await asyncio.gather(*billing_tasks)
|
||||
|
||||
assert second_builder_started.is_set() is True
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_completed_stream_does_not_double_bill_on_late_disconnect(self):
|
||||
recorder = _RecordingSuccessLogger()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue