fix: handle streaming and SSO edge cases

This commit is contained in:
Cursor Agent 2026-06-24 15:05:20 +00:00 • committed by Sameer Kankute
parent ac679561f5
commit 278c331f2a
No known key found for this signature in database
9 changed files with 229 additions and 24 deletions

View file

@ -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

View file

@ -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:

View file

@ -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

View file

@ -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

View file

@ -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(

View file

@ -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

View file

@ -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)])

View file

@ -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():
"""

View file

@ -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