From 278c331f2a69fdde4d97fbe3c6fcc67101e9a1f0 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Wed, 24 Jun 2026 15:05:20 +0000 Subject: [PATCH] fix: handle streaming and SSO edge cases --- .../adapters/streaming_iterator.py | 29 +++++++- .../adapters/transformation.py | 24 +++++-- litellm/proxy/common_request_processing.py | 23 +++++- .../guardrail_hooks/bedrock_guardrails.py | 4 +- litellm/proxy/management_endpoints/ui_sso.py | 52 +++++++++++--- .../test_bedrock_guardrails.py | 6 ++ .../test_streaming_iterator_combined_chunk.py | 10 +++ .../proxy/management_endpoints/test_ui_sso.py | 34 +++++++++ .../proxy/test_common_request_processing.py | 71 +++++++++++++++++++ 9 files changed, 229 insertions(+), 24 deletions(-) diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py index ddb6e7b7021..598350701c3 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py @@ -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 diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index 37b8813d499..3912df8fd6e 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -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: diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 61ef049bda4..1df0afbded2 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index ca46880b0d9..fb790a44779 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -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 diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index aa70ef99f1a..76f00846c21 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -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( diff --git a/tests/guardrails_tests/test_bedrock_guardrails.py b/tests/guardrails_tests/test_bedrock_guardrails.py index a23e89e576c..9bdb6dd159d 100644 --- a/tests/guardrails_tests/test_bedrock_guardrails.py +++ b/tests/guardrails_tests/test_bedrock_guardrails.py @@ -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 diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_combined_chunk.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_combined_chunk.py index d67de0dcaf8..c763883ab94 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_combined_chunk.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_combined_chunk.py @@ -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)]) diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 9b1523e1d73..aa0ced50e02 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -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(): """ diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 4e56fe68856..547f27bf5e2 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -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