mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix: handle streaming and SSO edge cases
This commit is contained in:
parent
ac679561f5
commit
278c331f2a
9 changed files with 229 additions and 24 deletions
|
|
@ -97,6 +97,14 @@ class _CombinedChunkSplitter:
|
|||
finish_delta.thinking_blocks = None
|
||||
return [content_chunk, finish_chunk]
|
||||
|
||||
@property
|
||||
def chunks(self) -> Any:
|
||||
return getattr(self._stream, "chunks", None)
|
||||
|
||||
@property
|
||||
def messages(self) -> Any:
|
||||
return getattr(self._stream, "messages", None)
|
||||
|
||||
def __iter__(self) -> "Iterator[Any]":
|
||||
return self
|
||||
|
||||
|
|
@ -189,6 +197,23 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
text="",
|
||||
)
|
||||
|
||||
@property
|
||||
def chunks(self) -> Any:
|
||||
return getattr(self.completion_stream, "chunks", None)
|
||||
|
||||
@property
|
||||
def messages(self) -> Any:
|
||||
return getattr(self.completion_stream, "messages", None)
|
||||
|
||||
@staticmethod
|
||||
def _is_empty_choices_without_usage(chunk: Any) -> bool:
|
||||
if getattr(chunk, "choices", None):
|
||||
return False
|
||||
if getattr(chunk, "usage", None) is not None:
|
||||
return False
|
||||
hidden_params = getattr(chunk, "_hidden_params", None)
|
||||
return not (isinstance(hidden_params, dict) and hidden_params.get("usage"))
|
||||
|
||||
def _merge_usage_into_held_stop_reason_chunk(self, chunk: Any) -> Dict[str, Any]:
|
||||
"""Merge usage data from ``chunk`` into the held ``message_delta`` chunk.
|
||||
|
||||
|
|
@ -388,7 +413,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
if chunk == "None" or chunk is None:
|
||||
raise Exception
|
||||
|
||||
if not chunk.choices and getattr(chunk, "usage", None) is None:
|
||||
if self._is_empty_choices_without_usage(chunk):
|
||||
continue
|
||||
|
||||
should_start_new_block = self._should_start_new_content_block(chunk)
|
||||
|
|
@ -612,7 +637,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
if chunk == "None" or chunk is None:
|
||||
raise Exception
|
||||
|
||||
if not chunk.choices and getattr(chunk, "usage", None) is None:
|
||||
if self._is_empty_choices_without_usage(chunk):
|
||||
continue
|
||||
|
||||
# Check if we need to start a new content block
|
||||
|
|
|
|||
|
|
@ -1480,9 +1480,25 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
current_content_block_index: int,
|
||||
applied_edits: Optional[List[AppliedEdit]] = None,
|
||||
) -> Union[ContentBlockDelta, MessageBlockDelta]:
|
||||
if getattr(response, "usage", None) is not None:
|
||||
litellm_usage_chunk: Optional[Usage] = response.usage # type: ignore
|
||||
elif (
|
||||
hasattr(response, "_hidden_params")
|
||||
and "usage" in response._hidden_params
|
||||
):
|
||||
litellm_usage_chunk = response._hidden_params["usage"]
|
||||
else:
|
||||
litellm_usage_chunk = None
|
||||
|
||||
## base case - final chunk w/ finish reason, or a usage-only chunk
|
||||
## (choices=[]) that carries trailing usage. See #30761.
|
||||
if not response.choices or response.choices[0].finish_reason is not None:
|
||||
has_finish_reason = (
|
||||
bool(response.choices) and response.choices[0].finish_reason is not None
|
||||
)
|
||||
has_usage_only_chunk = (
|
||||
not response.choices and litellm_usage_chunk is not None
|
||||
)
|
||||
if has_finish_reason or has_usage_only_chunk:
|
||||
stop_reason = (
|
||||
self._translate_openai_finish_reason_to_anthropic(
|
||||
response.choices[0].finish_reason
|
||||
|
|
@ -1491,12 +1507,6 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
else None
|
||||
)
|
||||
delta = MessageDelta(stop_reason=stop_reason)
|
||||
if getattr(response, "usage", None) is not None:
|
||||
litellm_usage_chunk: Optional[Usage] = response.usage # type: ignore
|
||||
elif hasattr(response, "_hidden_params") and "usage" in response._hidden_params:
|
||||
litellm_usage_chunk = response._hidden_params["usage"]
|
||||
else:
|
||||
litellm_usage_chunk = None
|
||||
if litellm_usage_chunk is not None:
|
||||
usage_delta = self._translate_openai_usage_to_anthropic_usage_delta(litellm_usage_chunk)
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -2407,6 +2407,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
response: Any,
|
||||
stream_completed: bool = False,
|
||||
client_disconnected: bool = False,
|
||||
streamed_chunks: list[Any] | None = None,
|
||||
) -> None:
|
||||
with anyio.CancelScope(shield=True):
|
||||
should_record_client_disconnect = client_disconnected or (not stream_completed)
|
||||
|
|
@ -2433,11 +2434,12 @@ class ProxyBaseLLMRequestProcessing:
|
|||
await ProxyBaseLLMRequestProcessing._bill_partial_stream_on_disconnect(
|
||||
response,
|
||||
request_data,
|
||||
streamed_chunks,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _bill_partial_stream_on_disconnect(
|
||||
response: object, request_data: dict
|
||||
response: object, request_data: dict, streamed_chunks: list[Any] | None = None
|
||||
) -> None:
|
||||
"""Record SpendLogs for tokens already produced when a stream is cut off.
|
||||
|
||||
|
|
@ -2451,9 +2453,22 @@ class ProxyBaseLLMRequestProcessing:
|
|||
normal completion already logged.
|
||||
"""
|
||||
logging_obj = request_data.get("litellm_logging_obj")
|
||||
chunks = getattr(response, "chunks", None)
|
||||
response_chunks = getattr(response, "chunks", None)
|
||||
chunks = (
|
||||
response_chunks
|
||||
if isinstance(response_chunks, list) and response_chunks
|
||||
else streamed_chunks
|
||||
)
|
||||
if logging_obj is None or not chunks:
|
||||
return
|
||||
first_chunk = chunks[0]
|
||||
if isinstance(first_chunk, (bytes, bytearray)):
|
||||
return
|
||||
if isinstance(first_chunk, dict):
|
||||
if "choices" not in first_chunk:
|
||||
return
|
||||
elif not isinstance(first_chunk, str) and not hasattr(first_chunk, "choices"):
|
||||
return
|
||||
# Optimization, not a correctness guard: dispatch_success_handlers is the
|
||||
# authoritative de-dup via has_dispatched_final_stream_success. This just
|
||||
# skips the stream_chunk_builder assembly when completion already logged.
|
||||
|
|
@ -2530,6 +2545,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
stream_completed = False
|
||||
client_disconnected = False
|
||||
delivered_chunk = False
|
||||
streamed_chunks: list[Any] = []
|
||||
try:
|
||||
str_so_far = ""
|
||||
async for chunk in proxy_logging_obj.async_post_call_streaming_iterator_hook(
|
||||
|
|
@ -2572,6 +2588,8 @@ class ProxyBaseLLMRequestProcessing:
|
|||
# awaited above, so a cancellation during it still leaves this
|
||||
# False and refunds.
|
||||
delivered_chunk = True
|
||||
if not isinstance(chunk, (bytes, bytearray)):
|
||||
streamed_chunks.append(chunk)
|
||||
yield serialize_chunk(chunk)
|
||||
stream_completed = True
|
||||
except (asyncio.CancelledError, GeneratorExit):
|
||||
|
|
@ -2625,6 +2643,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
response=response,
|
||||
stream_completed=stream_completed,
|
||||
client_disconnected=client_disconnected,
|
||||
streamed_chunks=streamed_chunks,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -256,7 +256,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
def _normalize_checks(checks: object | None) -> dict[str, Any] | None:
|
||||
"""Normalize the configured `checks` into a plain dict for the API body.
|
||||
|
||||
Accepts a pydantic ``BedrockChecksConfigModel`` or a raw dict; drops empty /
|
||||
Accepts a pydantic ``BedrockChecksConfigModel`` or a raw dict; drops None /
|
||||
unknown keys. Returns None when no usable check is configured (=> ApplyGuardrail).
|
||||
"""
|
||||
if checks is None:
|
||||
|
|
@ -268,7 +268,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
cleaned = {
|
||||
key: value
|
||||
for key, value in checks.items()
|
||||
if key in _BEDROCK_CHECKS_KNOWN_KEYS and value
|
||||
if key in _BEDROCK_CHECKS_KNOWN_KEYS and value is not None
|
||||
}
|
||||
return cleaned or None
|
||||
|
||||
|
|
|
|||
|
|
@ -863,6 +863,11 @@ async def google_login(
|
|||
param="premium_user",
|
||||
code=status.HTTP_403_FORBIDDEN,
|
||||
)
|
||||
await _enforce_free_sso_user_limit(
|
||||
prisma_client=prisma_client,
|
||||
premium_user=premium_user,
|
||||
block_at_limit=False,
|
||||
)
|
||||
|
||||
####### Detect DB + MASTER KEY in .env #######
|
||||
missing_env_vars = show_missing_vars_in_env()
|
||||
|
|
@ -1495,6 +1500,29 @@ def get_disabled_non_admin_personal_key_creation():
|
|||
return bool("proxy_admin" in allowed_user_roles)
|
||||
|
||||
|
||||
def _free_tier_sso_user_limit_error() -> ProxyException:
|
||||
return ProxyException(
|
||||
message="You must be a LiteLLM Enterprise user to use SSO for more than 5 users. If you have a license please set `LITELLM_LICENSE` in your env. If you want to obtain a license meet with us here: https://enterprise.litellm.ai/demo You are seeing this error message because You set one of `MICROSOFT_CLIENT_ID`, `GOOGLE_CLIENT_ID`, or `GENERIC_CLIENT_ID` in your env. Please unset this",
|
||||
type=ProxyErrorTypes.auth_error,
|
||||
param="premium_user",
|
||||
code=status.HTTP_403_FORBIDDEN,
|
||||
)
|
||||
|
||||
|
||||
async def _enforce_free_sso_user_limit(
|
||||
prisma_client: PrismaClient | None,
|
||||
premium_user: bool,
|
||||
block_at_limit: bool,
|
||||
) -> None:
|
||||
if premium_user or prisma_client is None:
|
||||
return
|
||||
total_users = await prisma_client.db.litellm_usertable.count()
|
||||
if total_users is None:
|
||||
return
|
||||
if total_users > 5 or (block_at_limit and total_users >= 5):
|
||||
raise _free_tier_sso_user_limit_error()
|
||||
|
||||
|
||||
async def get_existing_user_info_from_db(
|
||||
user_id: Optional[str],
|
||||
user_email: Optional[str],
|
||||
|
|
@ -2220,16 +2248,11 @@ async def insert_sso_user(
|
|||
if user_defined_values is None:
|
||||
raise ValueError("user_defined_values is None")
|
||||
|
||||
if not premium_user and prisma_client is not None:
|
||||
# Check if under 'free SSO user' limit
|
||||
total_users = await prisma_client.db.litellm_usertable.count()
|
||||
if total_users is not None and total_users >= 5:
|
||||
raise ProxyException(
|
||||
message="You must be a LiteLLM Enterprise user to use SSO for more than 5 users. If you have a license please set `LITELLM_LICENSE` in your env. If you want to obtain a license meet with us here: https://enterprise.litellm.ai/demo You are seeing this error message because You set one of `MICROSOFT_CLIENT_ID`, `GOOGLE_CLIENT_ID`, or `GENERIC_CLIENT_ID` in your env. Please unset this",
|
||||
type=ProxyErrorTypes.auth_error,
|
||||
param="premium_user",
|
||||
code=status.HTTP_403_FORBIDDEN,
|
||||
)
|
||||
await _enforce_free_sso_user_limit(
|
||||
prisma_client=prisma_client,
|
||||
premium_user=premium_user,
|
||||
block_at_limit=True,
|
||||
)
|
||||
# Apply default_internal_user_params
|
||||
if litellm.default_internal_user_params:
|
||||
# Preserve the SSO-extracted role if it's a valid LiteLLM role,
|
||||
|
|
@ -3046,7 +3069,14 @@ class SSOAuthenticationHandler:
|
|||
from litellm.proxy.utils import get_prisma_client_or_throw
|
||||
from litellm.types.proxy.ui_sso import ReturnedUITokenObject
|
||||
|
||||
prisma_client = get_prisma_client_or_throw("Prisma client is None, connect a database to your proxy")
|
||||
prisma_client = get_prisma_client_or_throw(
|
||||
"Prisma client is None, connect a database to your proxy"
|
||||
)
|
||||
await _enforce_free_sso_user_limit(
|
||||
prisma_client=prisma_client,
|
||||
premium_user=premium_user,
|
||||
block_at_limit=False,
|
||||
)
|
||||
|
||||
# User is Authe'd in - generate key for the UI to access Proxy
|
||||
parsed_openid_result = SSOAuthenticationHandler._get_user_email_and_id_from_result(
|
||||
|
|
|
|||
|
|
@ -14,6 +14,12 @@ from litellm.caching import DualCache
|
|||
from unittest.mock import MagicMock, AsyncMock, patch
|
||||
|
||||
|
||||
def test_bedrock_normalize_checks_keeps_empty_known_check_config():
|
||||
assert BedrockGuardrail._normalize_checks({"contentFilter": {}}) == {
|
||||
"contentFilter": {}
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_guardrails_pii_masking():
|
||||
# Create proper mock objects
|
||||
|
|
|
|||
|
|
@ -177,6 +177,16 @@ def test_is_combined_false_when_choices_empty():
|
|||
assert _CombinedChunkSplitter._is_combined(SimpleNamespace(choices=[])) is False
|
||||
|
||||
|
||||
def test_wrapper_exposes_underlying_chunks_for_disconnect_billing():
|
||||
chunks = [SimpleNamespace(choices=[])]
|
||||
messages = [{"role": "user", "content": "hi"}]
|
||||
upstream = SimpleNamespace(chunks=chunks, messages=messages)
|
||||
wrapper = AnthropicStreamWrapper(completion_stream=upstream, model="claude-x")
|
||||
|
||||
assert wrapper.chunks is chunks
|
||||
assert wrapper.messages is messages
|
||||
|
||||
|
||||
def test_is_combined_false_when_delta_missing():
|
||||
"""A finish chunk whose choice has no delta is not combined."""
|
||||
chunk = SimpleNamespace(choices=[SimpleNamespace(finish_reason="stop", delta=None)])
|
||||
|
|
|
|||
|
|
@ -769,6 +769,40 @@ def test_generic_response_convertor_normalizes_email():
|
|||
assert result.display_name == "Test User"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_free_sso_login_blocks_existing_users_when_over_limit():
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.ui_sso import (
|
||||
_enforce_free_sso_user_limit,
|
||||
)
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_usertable.count = AsyncMock(return_value=6)
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await _enforce_free_sso_user_limit(
|
||||
prisma_client=mock_prisma,
|
||||
premium_user=False,
|
||||
block_at_limit=False,
|
||||
)
|
||||
|
||||
assert str(exc_info.value.code) == "403"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_free_sso_login_allows_existing_users_at_limit():
|
||||
from litellm.proxy.management_endpoints.ui_sso import _enforce_free_sso_user_limit
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_usertable.count = AsyncMock(return_value=5)
|
||||
|
||||
await _enforce_free_sso_user_limit(
|
||||
prisma_client=mock_prisma,
|
||||
premium_user=False,
|
||||
block_at_limit=False,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_insert_sso_user_blocks_when_at_user_limit():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -3507,6 +3507,77 @@ class TestStreamingClientDisconnectLogging:
|
|||
logging_obj.dispatch_success_handlers.assert_awaited_once()
|
||||
assert order == ["aclose", "bill"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_streaming_data_generator_bills_partial_chunks_without_response_chunks(
|
||||
self, monkeypatch
|
||||
):
|
||||
from litellm.proxy.common_request_processing import (
|
||||
ProxyBaseLLMRequestProcessing,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.utils.ProxyLogging._fire_deferred_stream_logging",
|
||||
MagicMock(),
|
||||
)
|
||||
|
||||
chunk = {
|
||||
"id": "chatcmpl-test",
|
||||
"model": "gpt-4",
|
||||
"choices": [{"delta": {"content": "partial"}}],
|
||||
}
|
||||
partial_response = MagicMock()
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {"metadata": {}, "litellm_params": {}}
|
||||
logging_obj._on_deferred_stream_complete = None
|
||||
logging_obj.dispatch_success_handlers = AsyncMock()
|
||||
|
||||
class StreamWithoutChunks:
|
||||
aclose = AsyncMock()
|
||||
|
||||
async def mock_streaming_iterator(*_args, **_kwargs):
|
||||
yield chunk
|
||||
|
||||
mock_proxy_logging = MagicMock(spec=ProxyLogging)
|
||||
mock_proxy_logging.async_post_call_streaming_iterator_hook = (
|
||||
mock_streaming_iterator
|
||||
)
|
||||
mock_proxy_logging._release_max_parallel_requests_on_disconnect = MagicMock()
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.is_disconnected = AsyncMock(return_value=True)
|
||||
request_data = {
|
||||
"model": "gpt-4",
|
||||
"metadata": {},
|
||||
"litellm_params": {"metadata": {}},
|
||||
"litellm_logging_obj": logging_obj,
|
||||
}
|
||||
|
||||
with patch.object(
|
||||
litellm, "stream_chunk_builder", return_value=partial_response
|
||||
) as builder:
|
||||
gen = ProxyBaseLLMRequestProcessing.async_streaming_data_generator(
|
||||
response=StreamWithoutChunks(),
|
||||
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
|
||||
request_data=request_data,
|
||||
proxy_logging_obj=mock_proxy_logging,
|
||||
serialize_chunk=lambda stream_chunk: stream_chunk,
|
||||
serialize_error=lambda proxy_exc: proxy_exc,
|
||||
request=mock_request,
|
||||
)
|
||||
assert await gen.__anext__() is chunk
|
||||
await gen.aclose()
|
||||
|
||||
builder.assert_called_once()
|
||||
assert builder.call_args.kwargs["chunks"] == [chunk]
|
||||
logging_obj.dispatch_success_handlers.assert_awaited_once_with(
|
||||
partial_response,
|
||||
start_time=None,
|
||||
end_time=None,
|
||||
cache_hit=False,
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_streaming_data_generator_records_499_on_early_aclose(
|
||||
self, monkeypatch
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue