Merge branch 'litellm_internal_staging' of https://github.com/BerriAI/litellm into litellm_govcloud_profiles_lit6421

This commit is contained in:
mateo-berri 2026-09-03 18:41:00 -07:00
commit 1763052ff5
33 changed files with 2776 additions and 359 deletions

View file

@ -79,6 +79,11 @@ test_paths:
- tests/load_tests/test_otel_load_test.py
- tests/load_tests/test_vertex_embeddings_load_test.py
- tests/load_tests/test_vertex_load_tests.py
- reason: >-
Env-gated saturation benchmark requires a live proxy and provider credentials, so it is run
locally rather than in pull-request jobs
paths:
- tests/load_tests/test_granian_admission_saturation.py
- reason: >-
A local-only agent rig: test_a2a_completion_bridge.py needs a LangGraph server on
localhost:2024 and test_a2a.py drives a live A2A endpoint, so neither can run in a

View file

@ -36,7 +36,7 @@ AZURE_STORAGE_TOKEN_SCOPE: Final = "https://storage.azure.com/.default"
def _cached_credential_chain_token_provider() -> Callable[[], str]:
return get_azure_ad_token_provider(
azure_scope=AZURE_STORAGE_TOKEN_SCOPE,
azure_credential=AzureCredentialType.DefaultAzureCredential,
azure_credential=AzureCredentialType.DeploymentIdentityCredential,
)

View file

@ -8,23 +8,31 @@ Routes to native Cortex REST API endpoints based on model:
Ref: https://docs.snowflake.com/en/user-guide/snowflake-cortex/cortex-rest-api
"""
import copy
import json
import re
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Protocol, TypedDict
import httpx
from typing_extensions import ReadOnly
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolCallChunk
from litellm.litellm_core_utils.prompt_templates.factory import (
anthropic_process_openai_file_message,
convert_to_anthropic_tool_result,
create_anthropic_image_param,
select_anthropic_content_block_type_for_file,
)
from litellm.llms.anthropic.chat.handler import ModelResponseIterator as AnthropicStreamParser
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
from litellm.llms.anthropic.common_utils import normalize_cache_control_in_anthropic_payload
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolCallChunk, ChatCompletionToolMessage
from litellm.types.utils import (
ChatCompletionMessageToolCall,
ChatCompletionUsageBlock,
Choices,
Function,
GenericStreamingChunk,
Message,
ModelResponse,
Usage,
ModelResponseStream,
)
from ...base_llm.base_model_iterator import BaseModelResponseIterator
@ -93,6 +101,103 @@ def _is_claude_model(model: str) -> bool:
return any(name.startswith(p) for p in _CLAUDE_MODEL_PREFIXES)
def _convert_image_url_to_anthropic(block: Mapping[str, object]) -> object:
"""One OpenAI ``image_url`` block in the native shape Cortex accepts.
Cortex documents base64 sources only, so remote URLs are inlined the way every
other base64-only Anthropic dialect (Bedrock invoke, Vertex) inlines them, and
pdf/text data URIs become document blocks rather than malformed image blocks.
"""
image_url: Final = block.get("image_url")
url: Final = image_url if isinstance(image_url, str) else _image_url_field(image_url, "url")
if not url:
return block
converted: Final = (
anthropic_process_openai_file_message({"type": "file", "file": {"file_data": url}})
if select_anthropic_content_block_type_for_file(_data_uri_media_type(url)) == "document"
else create_anthropic_image_param(
image_url if isinstance(image_url, dict) else url, # mutable-ok: caller's JSON block
format=_image_url_field(image_url, "format"),
is_bedrock_invoke=True,
)
)
cache_control: Final = block.get("cache_control")
if cache_control is None:
return converted
return {**converted, "cache_control": cache_control} # mutable-ok: JSON wire block
def _image_url_field(image_url: object, key: str) -> str | None:
value: Final = image_url.get(key) if isinstance(image_url, dict) else None
return value if isinstance(value, str) else None
def _data_uri_media_type(url: str) -> str:
match: Final = re.match(r"data:([^;,]+)", url)
return match.group(1) if match else ""
def _convert_image_url_blocks_to_anthropic(content: object) -> object:
if not isinstance(content, list):
return content
return [ # mutable-ok: JSON wire blocks
_convert_image_url_to_anthropic(block)
if isinstance(block, Mapping) and block.get("type") == "image_url"
else block
for block in content
]
def _convert_tool_result_to_anthropic(
content: object, tool_call_id: str, cache_control: object
) -> Mapping[str, object]:
"""The Anthropic ``tool_result`` block for one OpenAI tool message.
Delegating to the shared converter keeps image, document and per-block cache
breakpoints identical to every other Anthropic dialect; only the plain-string
and non-list shapes it does not model are handled here.
"""
if not isinstance(content, list):
plain: Final[dict[str, object]] = { # mutable-ok: JSON wire block
"type": "tool_result",
"tool_use_id": tool_call_id,
"content": content if isinstance(content, str) else json.dumps(content),
}
return {**plain, "cache_control": cache_control} if cache_control is not None else plain
converted: Final = convert_to_anthropic_tool_result(
ChatCompletionToolMessage(role="tool", tool_call_id=tool_call_id, content=content),
force_base64=True,
)
if cache_control is None:
return converted
return {**converted, "cache_control": cache_control} # mutable-ok: JSON wire block
def _signed_thinking_blocks(msg: object) -> list[dict[str, object]]: # mutable-ok: JSON wire blocks
"""The assistant turn's thinking blocks that can legally be echoed back.
Only signed blocks round-trip: Cortex rejects a thinking block whose signature is
missing, which is what an unsigned block from a non-thinking turn would produce.
"""
blocks: Final = msg.get("thinking_blocks") if isinstance(msg, dict) else getattr(msg, "thinking_blocks", None)
if not isinstance(blocks, list):
return [] # mutable-ok: JSON wire blocks
return [ # mutable-ok: JSON wire blocks
dict(block)
for block in blocks
if isinstance(block, Mapping) and (block.get("signature") or block.get("type") == "redacted_thinking")
]
def _clean_input_schema(schema: object) -> object: # mutable-ok: JSON schema copy
return (
{key: value for key, value in schema.items() if key != "$schema"}
if isinstance(schema, Mapping)
else schema # mutable-ok: JSON schema copy
) # mutable-ok: JSON schema copy
class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
"""
Snowflake Cortex REST API — unified provider.
@ -178,7 +283,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
if "description" in func:
anthropic_tool["description"] = func["description"]
if "parameters" in func:
anthropic_tool["input_schema"] = func["parameters"]
anthropic_tool["input_schema"] = _clean_input_schema(func["parameters"])
else:
anthropic_tool["input_schema"] = {
"type": "object",
@ -186,10 +291,16 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
}
anthropic_tools.append(anthropic_tool)
else:
anthropic_tools.append(tool)
anthropic_tools.append(
{**tool, "input_schema": _clean_input_schema(tool["input_schema"])} # mutable-ok: JSON wire tool
if "input_schema" in tool
else tool
)
return anthropic_tools
def _extract_system_and_messages(self, messages: list[AllMessageValues]) -> tuple[str | None, list[dict]]:
def _extract_system_and_messages( # mutable-ok: JSON wire messages
self, messages: list[AllMessageValues]
) -> tuple[list[dict] | None, list[dict]]:
"""
Split messages into system prompt and conversation turns for Anthropic format.
@ -197,26 +308,39 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
- assistant messages with tool_calls → tool_use content blocks
- tool role messages → user role with tool_result content blocks
"""
system_parts: Final[list[str]] = []
conversation: Final[list[dict]] = []
system_parts: Final[list[dict]] = [] # mutable-ok: JSON wire messages
conversation: Final[list[dict]] = [] # mutable-ok: JSON wire messages
for msg in messages:
if isinstance(msg, dict):
role = msg.get("role", "")
content: Any = msg.get("content", "")
msg_cache_control: object = msg.get("cache_control")
else:
role = getattr(msg, "role", "")
content = getattr(msg, "content", "")
msg_cache_control = getattr(msg, "cache_control", None)
if role == "system":
if isinstance(content, str) and content:
system_parts.append(content)
system_parts.append({"type": "text", "text": content}) # mutable-ok: JSON wire system block
elif isinstance(content, list):
system_parts.append("\n".join(b.get("text", "") for b in content if b.get("type") == "text"))
system_parts.extend(
{ # mutable-ok: JSON wire system block
"type": "text",
"text": block.get("text", ""),
**(
{"cache_control": block["cache_control"]} if "cache_control" in block else {}
), # mutable-ok: JSON wire block
}
for block in content
if isinstance(block, Mapping) and block.get("type") == "text"
)
elif role == "assistant":
tool_calls = msg.get("tool_calls") if isinstance(msg, dict) else getattr(msg, "tool_calls", None)
thinking_blocks = _signed_thinking_blocks(msg)
if tool_calls:
content_blocks: list[dict[str, object]] = []
content_blocks: list[dict[str, object]] = list(thinking_blocks) # mutable-ok: JSON wire blocks
if content:
content_blocks.append({"type": "text", "text": content})
for tc in tool_calls:
@ -239,18 +363,26 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
}
)
conversation.append({"role": "assistant", "content": content_blocks})
elif thinking_blocks:
thinking_content = (
[
*thinking_blocks,
*copy.deepcopy(content),
]
if isinstance(content, list)
else [*thinking_blocks, *([{"type": "text", "text": content}] if content else [])]
) # rebind-ok: loop-local normalized content
conversation.append({"role": "assistant", "content": thinking_content})
else:
conversation.append({"role": "assistant", "content": content})
elif role == "tool":
tool_call_id = (
tool_call_id_value = (
msg.get("tool_call_id", "") if isinstance(msg, dict) else getattr(msg, "tool_call_id", "")
)
tool_content = content if isinstance(content, str) else json.dumps(content)
tool_result_block = {
"type": "tool_result",
"tool_use_id": tool_call_id,
"content": tool_content,
}
tool_call_id = (
tool_call_id_value if isinstance(tool_call_id_value, str) else ""
) # rebind-ok: normalized loop value
tool_result_block = _convert_tool_result_to_anthropic(content, tool_call_id, msg_cache_control)
if (
conversation
and conversation[-1]["role"] == "user"
@ -260,11 +392,18 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
):
conversation[-1]["content"].append(tool_result_block)
else:
conversation.append({"role": "user", "content": [tool_result_block]})
conversation.append(
{"role": "user", "content": [tool_result_block]} # mutable-ok: JSON wire message
) # mutable-ok: JSON wire message
else:
conversation.append({"role": role, "content": content})
conversation.append( # mutable-ok: JSON wire message
{ # mutable-ok: JSON wire message
"role": role,
"content": _convert_image_url_blocks_to_anthropic(content),
} # mutable-ok: JSON wire message
)
system: Final[str | None] = "\n\n".join(system_parts) if system_parts else None
system: Final[list[dict] | None] = system_parts if system_parts else None # mutable-ok: JSON wire messages
return system, conversation
def transform_request(
@ -339,7 +478,9 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
extra_body: dict,
) -> dict:
"""Anthropic Messages format for /messages endpoint."""
system, conversation = self._extract_system_and_messages(messages)
passthrough_system: Final = optional_params.pop("system", None)
extracted_system, conversation = self._extract_system_and_messages(messages)
system: Final = passthrough_system if passthrough_system is not None else extracted_system
if "tools" in optional_params:
optional_params["tools"] = self._transform_tools_to_anthropic(optional_params["tools"])
@ -353,16 +494,19 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
model_name: Final = model.removeprefix("snowflake/")
body: Final[dict[str, object]] = {
"model": model_name,
"messages": conversation,
"stream": stream,
**optional_params,
**extra_body,
}
body: Final[dict[str, object]] = normalize_cache_control_in_anthropic_payload( # mutable-ok: JSON wire body
{ # mutable-ok: JSON wire body
"model": model_name,
"messages": conversation,
"stream": stream,
**optional_params,
**extra_body, # mutable-ok: JSON wire body
}
)
if system is not None:
body["system"] = system
body["system"] = normalize_cache_control_in_anthropic_payload( # mutable-ok: JSON wire payload
{"system": system} # mutable-ok: JSON wire payload
)["system"]
if "max_tokens" not in body:
body["max_tokens"] = 4096 # reasonable default; Anthropic API max varies by model
@ -435,23 +579,10 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
additional_args={"complete_input_dict": request_data},
)
text_content = ""
tool_calls: Final = []
for block in response_json.get("content", []):
if block.get("type") == "text":
text_content += block.get("text", "")
elif block.get("type") == "tool_use":
tool_calls.append(
ChatCompletionMessageToolCall(
id=block.get("id", ""),
type="function",
function=Function(
name=block.get("name", ""),
arguments=json.dumps(block.get("input", {})),
),
)
)
anthropic_config: Final = AnthropicConfig()
text_content, _, thinking_blocks, reasoning_content, tool_calls, _, _, _ = (
anthropic_config.extract_response_content(completion_response=dict(response_json))
)
_stop_reason_map: Final = {
"end_turn": "stop",
@ -461,9 +592,13 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
}
finish_reason: Final = _stop_reason_map.get(response_json.get("stop_reason", "end_turn"), "stop")
message: Final = Message(content=text_content or None, role="assistant")
if tool_calls:
message.tool_calls = tool_calls
message: Final = Message(
content=text_content or None,
role="assistant",
tool_calls=tool_calls or None,
thinking_blocks=thinking_blocks,
reasoning_content=reasoning_content,
)
choice: Final = Choices(
finish_reason=finish_reason,
@ -471,11 +606,13 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
message=message,
)
usage_data: Final = response_json.get("usage", {})
usage: Final = Usage(
prompt_tokens=usage_data.get("input_tokens", 0),
completion_tokens=usage_data.get("output_tokens", 0),
total_tokens=usage_data.get("input_tokens", 0) + usage_data.get("output_tokens", 0),
# Cortex reports prompt-cache creation/read counts alongside input_tokens; the
# shared calculator folds them into prompt_tokens_details so cached input is
# visible and billed at its own rate.
usage: Final = anthropic_config.calculate_usage(
usage_object=response_json.get("usage", {}),
reasoning_content=reasoning_content,
completion_response=dict(response_json),
)
model_response.choices = [choice]
@ -516,15 +653,19 @@ class SnowflakeStreamingHandler(BaseModelResponseIterator):
json_mode: bool | None = False,
):
super().__init__(streaming_response=streaming_response, sync_stream=sync_stream)
self._tool_index = 0
self._tool_id = ""
self._tool_name = ""
self._input_tokens = 0
# Cortex streams the Anthropic SSE dialect on /messages, so its events are parsed
# by Anthropic's own parser: thinking deltas, signatures and prompt-cache usage
# all arrive the way they do on every other Anthropic-dialect provider.
self._anthropic_parser: Final = AnthropicStreamParser(
streaming_response=streaming_response,
sync_stream=sync_stream,
json_mode=json_mode,
)
def chunk_parser(self, chunk: dict) -> GenericStreamingChunk:
def chunk_parser(self, chunk: dict) -> GenericStreamingChunk | ModelResponseStream:
if "choices" in chunk:
return self._parse_openai_chunk(chunk)
return self._parse_anthropic_chunk(chunk)
return self._anthropic_parser.chunk_parser(chunk)
def _parse_openai_chunk(self, chunk: dict) -> GenericStreamingChunk:
choices: Final = chunk.get("choices", [])
@ -566,117 +707,3 @@ class SnowflakeStreamingHandler(BaseModelResponseIterator):
index=choice.get("index", 0),
tool_use=tool_use,
)
def _parse_anthropic_chunk(self, chunk: dict) -> GenericStreamingChunk:
event_type: Final = chunk.get("type", "")
if event_type == "message_start":
message: Final = chunk.get("message", {})
usage_data = message.get("usage", {})
self._input_tokens = usage_data.get("input_tokens", 0)
return GenericStreamingChunk(
text="",
is_finished=False,
finish_reason="",
usage=None,
index=0,
tool_use=None,
)
elif event_type == "content_block_delta":
delta = chunk.get("delta", {})
delta_type: Final = delta.get("type", "")
if delta_type == "text_delta":
return GenericStreamingChunk(
text=delta.get("text", ""),
is_finished=False,
finish_reason="",
usage=None,
index=chunk.get("index", 0),
tool_use=None,
)
elif delta_type == "input_json_delta":
return GenericStreamingChunk(
text="",
is_finished=False,
finish_reason="",
usage=None,
index=chunk.get("index", 0),
tool_use=ChatCompletionToolCallChunk(
id=self._tool_id,
type="function",
function={
"name": self._tool_name,
"arguments": delta.get("partial_json", ""),
},
index=self._tool_index,
),
)
elif event_type == "content_block_start":
content_block: Final = chunk.get("content_block", {})
if content_block.get("type") == "tool_use":
self._tool_id = content_block.get("id", "")
self._tool_name = content_block.get("name", "")
self._tool_index = chunk.get("index", 0)
return GenericStreamingChunk(
text="",
is_finished=False,
finish_reason="",
usage=None,
index=chunk.get("index", 0),
tool_use=ChatCompletionToolCallChunk(
id=self._tool_id,
type="function",
function={"name": self._tool_name, "arguments": ""},
index=self._tool_index,
),
)
elif event_type == "message_delta":
delta = chunk.get("delta", {})
stop_reason: Final = delta.get("stop_reason", "")
usage_data = chunk.get("usage", {})
_stop_map: Final = {
"end_turn": "stop",
"max_tokens": "length",
"tool_use": "tool_calls",
"stop_sequence": "stop",
}
usage = None
if usage_data or self._input_tokens:
output_t: Final = usage_data.get("output_tokens", 0)
input_t: Final = self._input_tokens or usage_data.get("input_tokens", 0)
usage = ChatCompletionUsageBlock(
prompt_tokens=input_t,
completion_tokens=output_t,
total_tokens=input_t + output_t,
)
return GenericStreamingChunk(
text="",
is_finished=True,
finish_reason=_stop_map.get(stop_reason, "stop"),
usage=usage,
index=0,
tool_use=None,
)
elif event_type == "message_stop":
return GenericStreamingChunk(
text="",
is_finished=True,
finish_reason="stop",
usage=None,
index=0,
tool_use=None,
)
return GenericStreamingChunk(
text="",
is_finished=False,
finish_reason="",
usage=None,
index=0,
tool_use=None,
)

View file

@ -3638,30 +3638,48 @@ class MCPServerManager:
user_api_key_auth: UserAPIKeyAuth | None,
raw_headers: Mapping[str, str] | None = None,
) -> None:
"""Run the OBO exchange for a caller-supplied subject at the transport edge.
"""Mint an exchange-backed server's upstream credential at the transport edge.
Single-server routes call this before the MCP session opens, where an HTTP status and
``WWW-Authenticate`` still reach the client. A rejected subject raises the RFC 9728
challenge and any other ``CredError`` maps onto its public HTTP status, so an exchange
failure surfaces as a failure instead of the session continuing into an empty tool list.
A successful exchange is cached by the exchanger, so the session's list/call reuses it.
Each mode pre-flights only where it would resolve the subject the session goes on to use,
which is what keeps the pre-flight from reaching a verdict the session would contradict.
``oauth2_token_exchange`` mints from the caller's inbound bearer, so without one there is
nothing to exchange and the missing-subject case stays the preemptive challenge's job.
``oauth2_id_jag`` is the mirror image: tool listing resolves it from the identity assertion
captured for this user at SSO login and never from the inbound bearer, so the pre-flight is
faithful exactly when no identity bearer was sent (a LiteLLM key in ``Authorization`` is not one),
and a caller that did send one is passed through
untouched rather than judged against a subject the listing will not use. That store-sourced
case is the one whose missing-assertion 412 and store-outage 503 the session cannot report.
Only OBO has a discovery challenge to raise; ID-JAG's failures are plain statuses whose body
already names what the user has to do, so they map through ``raise_public`` as at egress.
"""
if server.auth_type != MCPAuth.oauth2_token_exchange:
return
if not self._extract_bearer_token(oauth2_headers, None):
return
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
spec: Final = to_server_spec(resolved_server)
if spec is None or not isinstance(spec.config, TokenExchangeConfig):
return
subject_token: Final = self._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth)
if subject_token is None:
match server.auth_type:
case MCPAuth.oauth2_token_exchange:
if not self._extract_bearer_token(oauth2_headers, None):
return
case MCPAuth.oauth2_id_jag:
if subject_token is not None:
return
case _:
return
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
spec: Final = _to_server_spec_fail_closed(resolved_server)
if spec is None or not isinstance(spec.config, (TokenExchangeConfig, IdJagConfig)):
return
if subject_token is None and isinstance(spec.config, TokenExchangeConfig):
raise_token_exchange_challenge(resolved_server, root_path=get_server_root_path())
match await self._cred_provider.resolve_credentials(to_subject(user_api_key_auth, subject_token), spec):
case Ok(_):
return
case Error(err):
if err.tag == "unauthorized":
if err.tag == "unauthorized" and isinstance(spec.config, TokenExchangeConfig):
raise_token_exchange_challenge(
resolved_server,
root_path=get_server_root_path(),

View file

@ -3851,15 +3851,15 @@ if MCP_AVAILABLE:
raise_token_exchange_challenge(server, root_path=get_server_root_path())
# token_exchange (OBO) with a subject present: run the exchange here at the transport
# edge, so a rejected subject raises the RFC 9728 challenge (and a gateway fault its
# public status) instead of the session opening and list_tools masking the failure as
# an empty tool list. Gated to single-server routes; the multi-server aggregate keeps
# absorbing per-server auth failures so one bad server cannot 401 the whole connect.
# Exchange-backed modes (token_exchange's OBO mint, id_jag's stored-assertion mint): run
# the exchange here at the transport edge, so a rejected subject raises the RFC 9728
# challenge and any other failure its public status, instead of the session opening and
# list_tools masking it as an empty tool list. The manager owns which modes pre-flight
# and what each mints from. Gated to single-server routes the key may reach; the
# multi-server aggregate keeps absorbing per-server auth failures so one bad server
# cannot 401 the whole connect.
if (
server
and server.auth_type == MCPAuth.oauth2_token_exchange
and oauth2_headers
and len(mcp_servers or []) == 1
and server.server_id
in frozenset(

View file

@ -2404,6 +2404,15 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
"""
completion_model: str | None = Field(None, description="proxy level default model for all chat completion calls")
max_in_flight_requests_per_worker: int | None = Field(
None, gt=0, description="maximum concurrent requests handled by each worker"
)
max_queued_requests_per_worker: int | None = Field(
None, ge=0, description="maximum requests waiting for a worker slot"
)
admission_queue_timeout_seconds: float = Field(
1.0, gt=0, description="maximum time a request waits for a worker slot"
)
plugins: list[PluginConfig] | None = Field(
None, description="external services registered as embeddable UI plugins"
)

View file

@ -387,33 +387,22 @@ def _get_wildcard_models(
all_wildcard_models: Final = []
for model in unique_models:
if _check_wildcard_routing(model=model):
if return_wildcard_routes: # will add the wildcard route to the list eg: anthropic/*.
if return_wildcard_routes:
all_wildcard_models.append(model)
## get litellm params from model
if llm_router is not None:
model_list = llm_router.get_model_list(model_name=model, team_id=team_id)
if model_list:
for router_model in model_list:
wildcard_models = get_known_models_from_wildcard(
models_to_remove.add(model)
model_list = llm_router.get_model_list(model_name=model, team_id=team_id) if llm_router else None
if model_list:
for router_model in model_list:
all_wildcard_models.extend(
get_known_models_from_wildcard(
wildcard_model=model,
litellm_params=LiteLLM_Params(**router_model["litellm_params"]),
)
all_wildcard_models.extend(wildcard_models)
else:
# Router has no deployment for this wildcard (e.g., BYOK team models)
# Fall back to expanding from known provider models
wildcard_models = get_known_models_from_wildcard(wildcard_model=model, litellm_params=None)
if wildcard_models:
models_to_remove.add(model)
all_wildcard_models.extend(wildcard_models)
)
else:
# get all known provider models
wildcard_models = get_known_models_from_wildcard(wildcard_model=model, litellm_params=None)
if wildcard_models:
models_to_remove.add(model)
all_wildcard_models.extend(wildcard_models)
all_wildcard_models.extend(get_known_models_from_wildcard(wildcard_model=model, litellm_params=None))
for model in models_to_remove:
unique_models.remove(model)

View file

@ -50,6 +50,9 @@ from litellm.proxy.health_check import (
perform_health_check,
run_with_timeout,
)
from litellm.proxy.middleware.admission_control_middleware import (
get_admission_control_stats,
)
from litellm.proxy.middleware.in_flight_requests_middleware import (
get_in_flight_requests,
)
@ -63,6 +66,13 @@ from litellm.secret_managers.main import get_secret_bool
#### Health ENDPOINTS ####
class _HealthBacklogResponse(TypedDict):
in_flight_requests: ReadOnly[int]
admitted_requests: ReadOnly[int]
queued_requests: ReadOnly[int]
rejected_requests: ReadOnly[int]
def _reject_os_environ_references(params: dict) -> None:
"""
Validate that the provided params do not contain any ``os.environ/``
@ -1759,7 +1769,14 @@ async def health_backlog():
for the event loop to get to them, adding latency before LiteLLM even starts
its own timer.
"""
return {"in_flight_requests": get_in_flight_requests()}
stats: Final = get_admission_control_stats()
response: Final[_HealthBacklogResponse] = {
"in_flight_requests": get_in_flight_requests(),
"admitted_requests": stats.admitted,
"queued_requests": stats.queued,
"rejected_requests": stats.rejected_total,
}
return response
@router.get(

View file

@ -12,7 +12,7 @@ All /policy management endpoints
import copy
import json
import os
from collections.abc import AsyncIterator
from collections.abc import AsyncGenerator, AsyncIterator
from typing import TYPE_CHECKING, Final, Literal, cast
from fastapi import APIRouter, Depends, HTTPException, Request
@ -20,6 +20,7 @@ from fastapi.responses import Response, StreamingResponse
from pydantic import BaseModel, Field
from typing_extensions import TypedDict
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import (
COMPETITOR_LLM_TEMPERATURE,
@ -32,6 +33,10 @@ from litellm.llms.openai.chat.guardrail_translation.handler import (
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.sse_keepalive import (
SSE_COMMENT_PING,
wrap_sse_stream_with_keepalive_pings,
)
from litellm.proxy.guardrails.guardrail_hooks.custom_code import (
RESPONSE_REJECTION_GUARDRAIL_CODE,
CustomCodeGuardrail,
@ -811,7 +816,7 @@ async def _stream_competitor_events(
llm_enrichment: dict,
brand_name: str,
model: str,
) -> AsyncIterator[str]:
) -> AsyncGenerator[str, None]:
"""Stream competitor names as SSE events, then emit a final 'done' event."""
competitors: Final[list[str]] = list(data.competitors or [])
@ -883,7 +888,11 @@ async def enrich_policy_template_stream(
model: Final = data.model or DEFAULT_COMPETITOR_DISCOVERY_MODEL
return StreamingResponse(
_stream_competitor_events(data, template, llm_enrichment, brand_name, model),
wrap_sse_stream_with_keepalive_pings(
_stream_competitor_events(data, template, llm_enrichment, brand_name, model),
ping_interval_seconds=litellm.sse_keepalive_ping_interval_seconds,
ping_chunk=SSE_COMMENT_PING,
),
media_type="text/event-stream",
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},
)

View file

@ -4,7 +4,7 @@ usage/spend data by querying the aggregated daily activity endpoints.
"""
import json
from collections.abc import AsyncIterator, Awaitable, Callable, Mapping, Sequence
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Mapping, Sequence
from datetime import date
from typing import Any, Final, Literal, Protocol, cast, overload
@ -543,7 +543,7 @@ async def stream_usage_ai_chat(
model: str | None = None,
user_id: str | None = None,
is_admin: bool = False,
) -> AsyncIterator[str]:
) -> AsyncGenerator[str, None]:
"""Stream SSE events: status → tool_call → chunk → done."""
resolved_model: Final = (model or "").strip() or DEFAULT_COMPETITOR_DISCOVERY_MODEL
truncated: Final = messages[-MAX_CHAT_MESSAGES:] if len(messages) > MAX_CHAT_MESSAGES else messages

View file

@ -10,8 +10,13 @@ from fastapi import APIRouter, Depends, Request
from fastapi.responses import StreamingResponse
from pydantic import BaseModel, Field
import litellm
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.sse_keepalive import (
SSE_COMMENT_PING,
wrap_sse_stream_with_keepalive_pings,
)
router: Final = APIRouter()
@ -56,11 +61,15 @@ async def usage_ai_chat(
messages: Final = [{"role": m.role, "content": m.content} for m in data.messages]
return StreamingResponse(
stream_usage_ai_chat(
messages=messages,
model=data.model,
user_id=user_id,
is_admin=is_admin,
wrap_sse_stream_with_keepalive_pings(
stream_usage_ai_chat(
messages=messages,
model=data.model,
user_id=user_id,
is_admin=is_admin,
),
ping_interval_seconds=litellm.sse_keepalive_ping_interval_seconds,
ping_chunk=SSE_COMMENT_PING,
),
media_type="text/event-stream",
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},

View file

@ -0,0 +1,315 @@
import asyncio
import os
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from functools import lru_cache
from typing import Annotated, Final, Protocol, TypeAlias, runtime_checkable
from pydantic import Field, TypeAdapter, ValidationError
from starlette.responses import JSONResponse
from starlette.types import ASGIApp, Receive, Scope, Send
from litellm._logging import verbose_proxy_logger
_EXEMPT_PATHS: Final[frozenset[str]] = frozenset(
{
"/health/liveliness",
"/health/liveness",
"/health/readiness",
"/health/readiness/details",
"/health/backlog",
"/health/drain",
"/metrics",
"/metrics/",
}
)
@dataclass(frozen=True, slots=True)
class AdmissionControlSettings:
max_in_flight_requests: int
max_queued_requests: int
queue_timeout_seconds: float
AdmissionControlSettingsGetter: TypeAlias = Callable[[], AdmissionControlSettings | None] # mutable-ok: Callable params
@dataclass(frozen=True, slots=True)
class AdmissionControlStats:
admitted: int
queued: int
rejected_total: int
@runtime_checkable
class _Gauge(Protocol):
def inc(self, amount: float = 1) -> None: ...
def dec(self, amount: float = 1) -> None: ...
@runtime_checkable
class _CounterChild(Protocol):
def inc(self, amount: float = 1) -> None: ...
@runtime_checkable
class _Counter(Protocol):
def labels(self, reason: str) -> _CounterChild: ...
@dataclass(frozen=True, slots=True)
class AdmissionControlMetrics:
admitted_gauge: _Gauge
queued_gauge: _Gauge
rejected_counter: _Counter
AdmissionControlMetricsFactory: TypeAlias = Callable[[], AdmissionControlMetrics | None] # mutable-ok: Callable params
class AdmissionControlState:
"""Per-process admission counters and the in-flight semaphore shared by one worker's requests."""
def __init__(self, metrics_factory: AdmissionControlMetricsFactory) -> None:
self._metrics_factory = metrics_factory
self._metrics: AdmissionControlMetrics | None = None
self._metrics_init_attempted = False
self._admitted = 0
self._queued = 0
self._rejected_total = 0
self._semaphore: asyncio.Semaphore | None = None
self._semaphore_loop: asyncio.AbstractEventLoop | None = None
def get_stats(self) -> AdmissionControlStats:
return AdmissionControlStats(
admitted=self._admitted,
queued=self._queued,
rejected_total=self._rejected_total,
)
def get_semaphore(self, max_in_flight_requests: int) -> asyncio.Semaphore:
loop: Final = asyncio.get_running_loop()
if self._semaphore_loop is not loop:
self._semaphore = asyncio.Semaphore(max_in_flight_requests)
self._semaphore_loop = loop
semaphore: Final = self._semaphore
if semaphore is None:
raise RuntimeError("Admission control semaphore was not initialized")
return semaphore
def record_admission(self) -> None:
self._admitted += 1
metrics: Final = self._get_metrics()
if metrics is not None:
metrics.admitted_gauge.inc()
def record_release(self) -> None:
self._admitted -= 1
metrics: Final = self._get_metrics()
if metrics is not None:
metrics.admitted_gauge.dec()
def record_queue(self) -> None:
self._queued += 1
metrics: Final = self._get_metrics()
if metrics is not None:
metrics.queued_gauge.inc()
def record_dequeue(self) -> None:
self._queued -= 1
metrics: Final = self._get_metrics()
if metrics is not None:
metrics.queued_gauge.dec()
def record_rejection(self, reason: str) -> None:
self._rejected_total += 1
metrics: Final = self._get_metrics()
if metrics is not None:
metrics.rejected_counter.labels(reason=reason).inc()
def _get_metrics(self) -> AdmissionControlMetrics | None:
if not self._metrics_init_attempted:
self._metrics_init_attempted = True
self._metrics = self._metrics_factory()
return self._metrics
class AdmissionControlMiddleware:
def __init__(
self,
app: ASGIApp,
get_settings: AdmissionControlSettingsGetter,
state: AdmissionControlState,
) -> None:
self.app = app
self.get_settings = get_settings
self.state = state
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
if scope["type"] != "http":
await self.app(scope, receive, send)
return
settings: Final = self.get_settings()
if settings is None or _get_route_path(scope) in _EXEMPT_PATHS:
await self.app(scope, receive, send)
return
state: Final = self.state
semaphore: Final = state.get_semaphore(settings.max_in_flight_requests)
if not semaphore.locked():
await semaphore.acquire()
state.record_admission()
elif state.get_stats().queued >= settings.max_queued_requests:
state.record_rejection("queue_full")
await _overloaded_response(state)(scope, receive, send)
return
else:
state.record_queue()
try:
await asyncio.wait_for(
semaphore.acquire(),
timeout=settings.queue_timeout_seconds,
)
except asyncio.TimeoutError:
state.record_dequeue()
state.record_rejection("queue_timeout")
await _overloaded_response(state)(scope, receive, send)
return
except asyncio.CancelledError:
state.record_dequeue()
raise
state.record_dequeue()
state.record_admission()
try:
await self.app(scope, receive, send)
finally:
semaphore.release()
state.record_release()
def _get_route_path(scope: Scope) -> str:
"""Strip the ASGI root_path (SERVER_ROOT_PATH) the same way Starlette does before route matching."""
path: Final[str] = scope["path"]
root_path: Final[str] = scope.get("root_path", "")
if not root_path or not path.startswith(root_path):
return path
if path == root_path:
return ""
if path[len(root_path)] == "/":
return path[len(root_path) :]
return path
def _create_gauge(gauge_type: Callable[..., object], name: str, description: str) -> _Gauge:
metric: Final = (
gauge_type(name, description, multiprocess_mode="livesum")
if "PROMETHEUS_MULTIPROC_DIR" in os.environ
else gauge_type(name, description)
)
if not isinstance(metric, _Gauge):
raise TypeError("Admission gauge has an unexpected type")
return metric
def create_prometheus_admission_metrics() -> AdmissionControlMetrics | None:
try:
from prometheus_client import Counter, Gauge
return AdmissionControlMetrics(
admitted_gauge=_create_gauge(
Gauge,
"litellm_admission_admitted_requests",
"Number of requests admitted by this worker",
),
queued_gauge=_create_gauge(
Gauge,
"litellm_admission_queued_requests",
"Number of requests queued by this worker",
),
rejected_counter=Counter( # mutable-ok: Prometheus requires runtime Counter construction
"litellm_admission_rejected_requests_total",
"Number of requests rejected by this worker",
labelnames=("reason",),
),
)
except (ImportError, ValueError):
return None
admission_control_state: Final = AdmissionControlState(create_prometheus_admission_metrics)
def get_admission_control_stats() -> AdmissionControlStats:
return admission_control_state.get_stats()
_PositiveInt: TypeAlias = Annotated[int, Field(gt=0)]
_NonNegativeInt: TypeAlias = Annotated[int, Field(ge=0)]
_PositiveFloat: TypeAlias = Annotated[float, Field(gt=0)]
_AdmissionControlRaw: TypeAlias = int | float | str | None
def _hashable(value: object) -> _AdmissionControlRaw:
return value if value is None or isinstance(value, (int, float, str)) else repr(value)
_POSITIVE_INT_ADAPTER: Final[TypeAdapter[int]] = TypeAdapter(_PositiveInt)
_NON_NEGATIVE_INT_ADAPTER: Final[TypeAdapter[int]] = TypeAdapter(_NonNegativeInt)
_POSITIVE_FLOAT_ADAPTER: Final[TypeAdapter[float]] = TypeAdapter(_PositiveFloat)
@lru_cache(maxsize=16)
def _parse_admission_control_settings(
max_in_flight_raw: _AdmissionControlRaw,
max_queued_raw: _AdmissionControlRaw,
queue_timeout_raw: _AdmissionControlRaw,
) -> AdmissionControlSettings | None:
try:
max_in_flight: Final = _POSITIVE_INT_ADAPTER.validate_python(max_in_flight_raw)
max_queued: Final = (
max_in_flight if max_queued_raw is None else _NON_NEGATIVE_INT_ADAPTER.validate_python(max_queued_raw)
)
queue_timeout: Final = _POSITIVE_FLOAT_ADAPTER.validate_python(queue_timeout_raw)
except ValidationError as exc:
verbose_proxy_logger.error(
"Ignoring invalid admission control settings, per-worker admission control is disabled: %s",
exc,
)
return None
return AdmissionControlSettings(
max_in_flight_requests=max_in_flight,
max_queued_requests=max_queued,
queue_timeout_seconds=queue_timeout,
)
def get_admission_control_settings(settings: Mapping[str, object]) -> AdmissionControlSettings | None:
max_in_flight_raw: Final = settings.get("max_in_flight_requests_per_worker")
if max_in_flight_raw is None:
return None
return _parse_admission_control_settings(
_hashable(max_in_flight_raw),
_hashable(settings.get("max_queued_requests_per_worker")),
_hashable(settings.get("admission_queue_timeout_seconds", 1.0)),
)
def _overloaded_response(state: AdmissionControlState) -> JSONResponse:
stats: Final = state.get_stats()
return JSONResponse(
status_code=503,
headers={"retry-after": "1"}, # mutable-ok: Starlette expects a plain headers mapping
content={ # mutable-ok: Starlette serializes a plain response mapping
"error": { # mutable-ok: nested response mapping
"message": (
f"Worker at capacity: {stats.admitted} in-flight, {stats.queued} queued requests. Retry later."
),
"type": "overloaded_error",
"code": "503",
}
},
)

View file

@ -9,10 +9,11 @@ Use litellm with Anthropic SDK, Vertex AI SDK, Cohere SDK, etc.
from __future__ import annotations
import hmac
import inspect
import json
import os
import re
from collections.abc import Callable, Mapping
from collections.abc import AsyncGenerator, Callable, Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Annotated, Final, cast
@ -32,6 +33,7 @@ from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
from litellm.proxy._types import *
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.auth.route_checks import RouteChecks
@ -40,6 +42,7 @@ from litellm.proxy.auth.user_api_key_auth import (
user_api_key_auth,
user_api_key_auth_websocket,
)
from litellm.proxy.common_request_processing import open_sse_before_first_byte
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_headers,
@ -47,6 +50,9 @@ from litellm.proxy.common_utils.http_parsing_utils import (
get_form_data,
get_request_body,
)
from litellm.proxy.common_utils.sse_keepalive import (
wrap_passthrough_sse_bytes_with_keepalive_pings,
)
from litellm.proxy.pass_through_endpoints.common_utils import get_litellm_virtual_key
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
HttpPassThroughEndpointHelpers,
@ -1478,6 +1484,74 @@ def is_azure_ai_search_service_level_index_create(method: str, endpoint: str) ->
return path == "indexes" or path.endswith("/indexes")
async def _relay_upstream_bytes(upstream: AsyncGenerator[bytes, bytes]) -> AsyncGenerator[bytes, None]:
try:
async for chunk in upstream:
yield chunk
finally:
await upstream.aclose()
async def _relay_azure_router_model(
llm_router: litellm.Router,
model: str,
endpoint: str,
request: Request,
request_body: Mapping[str, object],
is_streaming_request: bool,
user_api_key_dict: UserAPIKeyAuth,
) -> Response:
result: Final = await llm_router.allm_passthrough_route(
model=model,
method=request.method,
endpoint=endpoint,
request_query_params=request.query_params,
request_headers=_safe_get_request_headers(request),
stream=is_streaming_request,
content=None,
data=None,
files=None,
json=(request_body if request.headers.get("content-type") == "application/json" else None),
params=None,
headers=None,
cookies=None,
litellm_metadata=get_passthrough_router_request_metadata(user_api_key_dict),
)
if not is_streaming_request:
upstream: Final = cast(httpx.Response, result)
return Response(
content=await upstream.aread(),
status_code=upstream.status_code,
headers=HttpPassThroughEndpointHelpers.get_response_headers(headers=upstream.headers, custom_headers=None),
)
if inspect.isasyncgen(result):
sse_headers: Final = {"content-type": "text/event-stream"}
return StreamingResponse(
content=wrap_passthrough_sse_bytes_with_keepalive_pings(
stream=_relay_upstream_bytes(result),
ping_interval_seconds=litellm.sse_keepalive_ping_interval_seconds,
upstream_headers=sse_headers,
),
status_code=200,
headers=sse_headers,
)
upstream_stream: Final = cast(AsyncPassthroughStreamingResponse, result)
return StreamingResponse(
content=wrap_passthrough_sse_bytes_with_keepalive_pings(
stream=_relay_upstream_bytes(upstream_stream),
ping_interval_seconds=litellm.sse_keepalive_ping_interval_seconds,
upstream_headers=upstream_stream.headers,
),
status_code=upstream_stream.status_code,
headers=HttpPassThroughEndpointHelpers.get_response_headers(
headers=upstream_stream.headers, custom_headers=None
),
)
@router.api_route(
"/azure_ai/{endpoint:path}",
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],
@ -1528,55 +1602,18 @@ async def azure_proxy_route(
if is_router_model:
request_body = await get_request_body(request)
is_streaming_request = is_passthrough_request_streaming(request_body)
result = await llm_router.allm_passthrough_route(
model=part,
method=request.method,
endpoint=endpoint,
request_query_params=request.query_params,
request_headers=_safe_get_request_headers(request),
stream=is_streaming_request,
content=None,
data=None,
files=None,
json=(request_body if request.headers.get("content-type") == "application/json" else None),
params=None,
headers=None,
cookies=None,
litellm_metadata=get_passthrough_router_request_metadata(user_api_key_dict),
)
if is_streaming_request:
# Check if result is an async generator (from _async_streaming)
import inspect
if inspect.isasyncgen(result):
# Result is already an async generator, use it directly
return StreamingResponse(
content=result,
status_code=200,
headers={"content-type": "text/event-stream"},
)
else:
# Result is an httpx.Response, use aiter_bytes()
result = cast(httpx.Response, result)
return StreamingResponse(
content=result.aiter_bytes(),
status_code=result.status_code,
headers=HttpPassThroughEndpointHelpers.get_response_headers(
headers=result.headers,
custom_headers=None,
),
)
# Non-streaming response
result = cast(httpx.Response, result)
content = await result.aread()
return Response(
content=content,
status_code=result.status_code,
headers=HttpPassThroughEndpointHelpers.get_response_headers(
headers=result.headers,
custom_headers=None,
return await open_sse_before_first_byte(
_relay_azure_router_model(
llm_router=llm_router,
model=part,
endpoint=endpoint,
request=request,
request_body=request_body,
is_streaming_request=is_streaming_request,
user_api_key_dict=user_api_key_dict,
),
ping_interval_seconds=(
litellm.sse_keepalive_ping_interval_seconds if is_streaming_request else None
),
)
elif is_vector_store_index:

View file

@ -583,6 +583,11 @@ try:
except ImportError:
build_billing_metrics_recorder = None
shutdown_billing_metrics_recorder = None
from litellm.proxy.middleware.admission_control_middleware import (
AdmissionControlMiddleware,
admission_control_state,
get_admission_control_settings,
)
from litellm.proxy.middleware.in_flight_requests_middleware import (
InFlightRequestsMiddleware,
)
@ -15233,20 +15238,33 @@ async def async_queue_request(
if llm_router is None:
raise HTTPException(status_code=500, detail={"error": CommonProxyErrors.no_llm_router.value})
response: Final = await llm_router.schedule_acompletion(**data)
router: Final = llm_router
if "stream" in data and data["stream"] is True: # use generate_responses to stream responses
return StreamingResponse(
async_data_generator(
user_api_key_dict=user_api_key_dict,
response=response,
request_data=data,
request=request,
),
media_type="text/event-stream",
async def produce_queue_stream() -> StreamingResponse:
return StreamingResponse(
async_data_generator(
user_api_key_dict=user_api_key_dict,
response=await router.schedule_acompletion(**data),
request_data=data,
request=request,
),
media_type="text/event-stream",
)
async def audit_late_failure(exc: Exception) -> HTTPException | None:
return await proxy_logging_obj.post_call_failure_hook(
user_api_key_dict=user_api_key_dict, original_exception=exc, request_data=data
)
return await open_sse_before_first_byte(
produce_queue_stream(),
ping_interval_seconds=ttft_keepalive_interval(data, router),
on_late_failure=audit_late_failure,
)
response: Final = await router.schedule_acompletion(**data)
fastapi_response.headers.update({"x-litellm-priority": str(data["priority"])})
return response
except Exception as e:
@ -16502,6 +16520,9 @@ _GENERAL_SETTINGS_CONFIG_LIST_FIELD_TYPES: Final[Mapping[str, str]] = MappingPro
{
"max_parallel_requests": "Integer",
"global_max_parallel_requests": "Integer",
"max_in_flight_requests_per_worker": "Integer",
"max_queued_requests_per_worker": "Integer",
"admission_queue_timeout_seconds": "Float",
"max_request_size_mb": "Integer",
"max_batch_file_size_mb": "Integer",
"max_file_size_mb": "Integer",
@ -18177,6 +18198,11 @@ app.add_middleware(
get_max_request_size_mb=lambda: general_settings.get("max_request_size_mb"),
is_request_size_limit_enabled=lambda: premium_user is True,
)
app.add_middleware(
AdmissionControlMiddleware,
get_settings=lambda: get_admission_control_settings(general_settings),
state=admission_control_state,
)
async def _stream_mcp_asgi_response(handle_fn, scope: dict, receive) -> "StreamingResponse":

View file

@ -23,11 +23,14 @@ from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
LiteLLM_ManagedVectorStore,
)
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.proxy._types import *
from litellm.proxy.auth.auth_utils import is_request_body_safe
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
open_sse_before_first_byte,
ttft_keepalive_interval,
)
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_headers,
@ -48,6 +51,7 @@ from litellm.proxy.vector_store_endpoints.utils import (
assert_user_can_access_vector_store_id,
)
from litellm.repositories.table_repositories import ManagedVectorStoresRepository
from litellm.types.utils import ModelResponse
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
@ -756,43 +760,53 @@ async def rag_query(
merged_retrieval_config.get("custom_llm_provider"),
)
# Call query
response: Final = await litellm.aquery(
model=model,
messages=messages,
retrieval_config=merged_retrieval_config,
vector_store_params=store_data,
rerank=rerank,
stream=stream,
router=llm_router,
**request_data,
)
hidden_params: Final = getattr(response, "_hidden_params", {}) or {}
custom_headers: Final = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=user_api_key_dict,
call_id=hidden_params.get("litellm_call_id", None) or "",
model_id=hidden_params.get("model_id", None) or "",
cache_key=hidden_params.get("cache_key", None) or "",
api_base=hidden_params.get("api_base", None) or "",
version=version,
response_cost=hidden_params.get("response_cost", None),
request_data=request_data,
)
if isinstance(response, CustomStreamWrapper):
return StreamingResponse(
select_data_generator(
response=response,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
request=request,
),
media_type="text/event-stream",
headers=custom_headers,
async def query() -> ModelResponse:
return await litellm.aquery(
model=model,
messages=messages,
retrieval_config=merged_retrieval_config,
vector_store_params=store_data,
rerank=rerank,
stream=stream,
router=llm_router,
**request_data,
)
fastapi_response.headers.update(custom_headers)
def custom_headers_for(response: ModelResponse) -> Mapping[str, str]:
hidden_params: Final = getattr(response, "_hidden_params", {}) or {}
return ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=user_api_key_dict,
call_id=hidden_params.get("litellm_call_id", None) or "",
model_id=hidden_params.get("model_id", None) or "",
cache_key=hidden_params.get("cache_key", None) or "",
api_base=hidden_params.get("api_base", None) or "",
version=version,
response_cost=hidden_params.get("response_cost", None),
request_data=request_data,
)
if stream:
async def produce_stream() -> StreamingResponse:
response: Final = await query()
return StreamingResponse(
select_data_generator(
response=response,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
request=request,
),
media_type="text/event-stream",
headers=custom_headers_for(response),
)
return await open_sse_before_first_byte(
produce_stream(),
ping_interval_seconds=ttft_keepalive_interval(data, llm_router),
)
response: Final = await query()
fastapi_response.headers.update(custom_headers_for(response))
return response
except HTTPException:

View file

@ -57,9 +57,11 @@ def get_azure_ad_token_provider(
from azure import identity
from azure.identity import (
CertificateCredential,
ChainedTokenCredential,
ClientSecretCredential,
DefaultAzureCredential,
ManagedIdentityCredential,
WorkloadIdentityCredential,
get_bearer_token_provider,
)
@ -101,6 +103,28 @@ def get_azure_ad_token_provider(
# DefaultAzureCredential doesn't require explicit environment variables
# It automatically discovers credentials from the environment (managed identity, CLI, etc.)
credential = DefaultAzureCredential()
elif cred == AzureCredentialType.DeploymentIdentityCredential:
# DefaultAzureCredential cannot express this: excluding its developer credentials still
# leaves one managed identity link, which AZURE_CLIENT_ID pins to a user assigned identity,
# so a host running as a system assigned identity never gets asked
workload_client_id: Final = os.environ.get("AZURE_CLIENT_ID")
workload_tenant_id: Final = os.environ.get("AZURE_TENANT_ID")
workload_token_file: Final = os.environ.get("AZURE_FEDERATED_TOKEN_FILE")
credential = ChainedTokenCredential(
*(
(
WorkloadIdentityCredential(
client_id=workload_client_id,
tenant_id=workload_tenant_id,
token_file_path=workload_token_file,
),
)
if workload_client_id and workload_tenant_id and workload_token_file
else ()
),
*((ManagedIdentityCredential(client_id=workload_client_id),) if workload_client_id else ()),
ManagedIdentityCredential(),
)
else:
cred_cls: Final = getattr(identity, cred)
credential = cred_cls()

View file

@ -6,3 +6,4 @@ class AzureCredentialType(str, Enum):
ManagedIdentityCredential = "ManagedIdentityCredential"
CertificateCredential = "CertificateCredential"
DefaultAzureCredential = "DefaultAzureCredential"
DeploymentIdentityCredential = "DeploymentIdentityCredential"

View file

@ -0,0 +1,150 @@
import asyncio
import os
import socket
import subprocess
import sys
import time
from pathlib import Path
from typing import Final
import httpx
import pytest
pytestmark = pytest.mark.skipif(
os.environ.get("LITELLM_RUN_SATURATION_BENCHMARK") != "1",
reason="set LITELLM_RUN_SATURATION_BENCHMARK=1 to run the saturation benchmark",
)
def _free_port() -> int:
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as listener:
listener.bind(("127.0.0.1", 0))
return int(listener.getsockname()[1])
def _percentile(values: list[float], percentile: float) -> float:
return sorted(values)[min(int(len(values) * percentile), len(values) - 1)]
@pytest.mark.asyncio
async def test_granian_admission_control_saturation(tmp_path: Path) -> None:
fake_port: Final = _free_port()
proxy_port: Final = _free_port()
fake_script: Final = Path(__file__).parents[1] / "_fake_openai_endpoint_server.py"
config_path: Final = tmp_path / "saturation_config.yaml"
config_path.write_text(
f"""model_list:
- model_name: slow-endpoint
litellm_params:
model: openai/slow-endpoint
api_base: http://127.0.0.1:{fake_port}/v1
general_settings:
master_key: sk-saturation
max_in_flight_requests_per_worker: 8
max_queued_requests_per_worker: 8
admission_queue_timeout_seconds: 0.5
"""
)
fake_process: Final = subprocess.Popen(
[sys.executable, str(fake_script), "--host", "127.0.0.1", "--port", str(fake_port)],
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
)
try:
proxy_process: Final = subprocess.Popen(
[
sys.executable,
"-m",
"litellm.proxy.proxy_cli",
"--config",
str(config_path),
"--run_granian",
"--num_workers",
"1",
"--port",
str(proxy_port),
],
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
)
try:
async with httpx.AsyncClient(base_url=f"http://127.0.0.1:{proxy_port}") as client:
deadline: Final = time.monotonic() + 60
while time.monotonic() < deadline:
try:
response: Final = await client.get("/health/liveliness", timeout=2)
if response.status_code == 200:
break
except httpx.HTTPError:
pass
await asyncio.sleep(0.25)
else:
raise AssertionError("Granian proxy did not become healthy")
liveness_latencies: Final[list[float]] = []
stop_sampling: Final = asyncio.Event()
async def sample_liveness() -> None:
while not stop_sampling.is_set():
start: Final = time.perf_counter()
try:
response = await client.get("/health/liveliness", timeout=2)
response.raise_for_status()
liveness_latencies.append(time.perf_counter() - start)
except httpx.HTTPError:
pass
await asyncio.sleep(0.05)
async def send_completion() -> tuple[int, float, bool]:
start: Final = time.perf_counter()
response = await client.post(
"/chat/completions",
headers={"Authorization": "Bearer sk-saturation"},
json={
"model": "slow-endpoint",
"messages": [{"role": "user", "content": "hello"}],
},
timeout=10,
)
return response.status_code, time.perf_counter() - start, "retry-after" in response.headers
sampler: Final = asyncio.create_task(sample_liveness())
results: Final = await asyncio.gather(*(send_completion() for _ in range(200)))
stop_sampling.set()
await sampler
statuses: Final = [result[0] for result in results]
latencies: Final = [result[1] for result in results]
rejected: Final = [result for result in results if result[0] == 503]
assert set(statuses) <= {200, 503}
assert rejected
assert all(result[2] for result in rejected)
assert _percentile(latencies, 0.99) < 5
assert liveness_latencies
assert _percentile(liveness_latencies, 0.95) < 0.5
duration: Final = max(latencies)
print(
"\nmetric value\n"
f"rps {len(results) / duration:.2f}\n"
f"200 count {statuses.count(200)}\n"
f"503 count {statuses.count(503)}\n"
f"p50 {_percentile(latencies, 0.50):.3f}s\n"
f"p95 {_percentile(latencies, 0.95):.3f}s\n"
f"p99 {_percentile(latencies, 0.99):.3f}s\n"
f"liveness p95 {_percentile(liveness_latencies, 0.95):.3f}s"
)
finally:
proxy_process.terminate()
try:
proxy_process.wait(timeout=10)
except subprocess.TimeoutExpired:
proxy_process.kill()
proxy_process.wait()
finally:
fake_process.terminate()
try:
fake_process.wait(timeout=10)
except subprocess.TimeoutExpired:
fake_process.kill()
fake_process.wait()

View file

@ -42,6 +42,7 @@ def workload_identity_env_vars(monkeypatch):
"AZURE_STORAGE_ENDPOINT_SUFFIX",
"AZURE_CLIENT_SECRET",
"AZURE_CREDENTIAL",
"AZURE_TOKEN_CREDENTIALS",
"AZURE_SCOPE",
):
monkeypatch.delenv(unset, raising=False)
@ -206,10 +207,28 @@ def test_default_chain_provider_is_storage_scoped_and_built_once_per_process():
assert first() == "chain-token"
mock_builder.assert_called_once_with(
azure_scope="https://storage.azure.com/.default",
azure_credential=AzureCredentialType.DefaultAzureCredential,
azure_credential=AzureCredentialType.DeploymentIdentityCredential,
)
def test_storage_chain_reaches_only_the_identities_a_deployment_carries(workload_identity_env_vars):
"""
The chain runs on a server, where a developer sign-in is a person and not the deployment, so
the storage token must come from workload identity or managed identity or from nothing
"""
_cached_credential_chain_token_provider.cache_clear()
with patch("azure.identity.get_bearer_token_provider", return_value=lambda: "chain-token") as bearer:
_cached_credential_chain_token_provider()
_cached_credential_chain_token_provider.cache_clear()
bearer.assert_called_once()
with bearer.call_args.args[0] as chain:
assert {type(link).__name__ for link in chain.credentials} == {
"WorkloadIdentityCredential",
"ManagedIdentityCredential",
}
@pytest.mark.asyncio
async def test_chain_tokens_are_read_from_the_provider_on_every_refresh(
workload_identity_env_vars,

View file

@ -17,7 +17,7 @@ import pytest
import litellm
from litellm import completion, acompletion
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.snowflake.chat.transformation import SnowflakeConfig
from litellm.llms.snowflake.chat.transformation import SnowflakeConfig, SnowflakeStreamingHandler
from litellm.types.utils import ModelResponse
@ -114,8 +114,7 @@ class TestSnowflakeToolTransformation:
)
assert transformed_request["tool_choice"] == value, (
f"tool_choice='{value}' should pass through unchanged, "
f"got {transformed_request['tool_choice']}"
f"tool_choice='{value}' should pass through unchanged, got {transformed_request['tool_choice']}"
)
def test_transform_response_with_tool_calls(self):
@ -159,9 +158,7 @@ class TestSnowflakeToolTransformation:
headers={"Content-Type": "application/json"},
)
model_response = ModelResponse(
choices=[litellm.Choices(index=0, message=litellm.Message())]
)
model_response = ModelResponse(choices=[litellm.Choices(index=0, message=litellm.Message())])
logging_obj = MagicMock()
@ -232,9 +229,7 @@ class TestSnowflakeToolTransformation:
headers={"Content-Type": "application/json"},
)
model_response = ModelResponse(
choices=[litellm.Choices(index=0, message=litellm.Message())]
)
model_response = ModelResponse(choices=[litellm.Choices(index=0, message=litellm.Message())])
logging_obj = MagicMock()
@ -280,9 +275,7 @@ class TestSnowflakeToolTransformation:
headers={"Content-Type": "application/json"},
)
model_response = ModelResponse(
choices=[litellm.Choices(index=0, message=litellm.Message())]
)
model_response = ModelResponse(choices=[litellm.Choices(index=0, message=litellm.Message())])
logging_obj = MagicMock()
@ -300,10 +293,7 @@ class TestSnowflakeToolTransformation:
# Verify standard response works
assert isinstance(result, ModelResponse)
assert (
result.choices[0].message.content
== "Hello! I'm doing well, thank you for asking."
)
assert result.choices[0].message.content == "Hello! I'm doing well, thank you for asking."
def test_get_supported_openai_params_includes_tools(self):
"""
@ -318,6 +308,385 @@ class TestSnowflakeToolTransformation:
assert "max_tokens" in supported_params
class TestSnowflakeCortexClaudeFixes:
def setup_method(self):
self.config = SnowflakeConfig()
@staticmethod
def _transform(messages, optional_params=None):
return SnowflakeConfig().transform_request(
model="snowflake/claude-sonnet-4-6",
messages=messages,
optional_params=optional_params or {},
litellm_params={},
headers={},
)
def test_thinking_is_offered_on_every_claude_model(self):
"""Cortex documents extended thinking (budget_tokens) for Claude generally, so a
4.6-only gate would silently drop it on the models that do support it."""
for model in (
"snowflake/claude-sonnet-4-6",
"snowflake/claude-sonnet-4-5",
"snowflake/claude-3-7-sonnet",
"snowflake/claude-4-opus",
):
assert "thinking" in self.config.get_supported_openai_params(model), model
assert "thinking" not in self.config.get_supported_openai_params("snowflake/llama3.1-70b")
def test_system_blocks_preserve_cache_control_and_strip_ttl(self):
body = self._transform(
[
{
"role": "system",
"content": [
{
"type": "text",
"text": "You are helpful",
"cache_control": {"type": "ephemeral", "ttl": "1h"},
}
],
},
{"role": "user", "content": "hi"},
]
)
assert body["system"] == [{"type": "text", "text": "You are helpful", "cache_control": {"type": "ephemeral"}}]
def test_direct_system_param_is_normalized(self):
body = self._transform(
[{"role": "user", "content": "hi"}],
{"system": [{"type": "text", "text": "direct", "cache_control": {"type": "ephemeral", "ttl": "1h"}}]},
)
assert body["system"] == [{"type": "text", "text": "direct", "cache_control": {"type": "ephemeral"}}]
def test_message_and_tool_cache_control_are_normalized(self):
body = self._transform(
[
{
"role": "user",
"content": [{"type": "text", "text": "hi", "cache_control": {"type": "ephemeral", "ttl": "1h"}}],
}
],
{
"tools": [
{
"name": "f",
"input_schema": {"type": "object", "properties": {}},
"cache_control": {"type": "ephemeral", "ttl": "1h"},
}
]
},
)
assert body["messages"][0]["content"][0]["cache_control"] == {"type": "ephemeral"}
assert body["tools"][0]["cache_control"] == {"type": "ephemeral"}
def test_extra_body_message_override_is_normalized(self):
body = self._transform(
[{"role": "user", "content": "original"}],
{
"extra_body": {
"messages": [
{
"role": "user",
"content": [
{
"type": "text",
"text": "override",
"cache_control": {"type": "ephemeral", "ttl": "1h"},
}
],
}
]
}
},
)
assert body["messages"][0]["content"][0]["cache_control"] == {"type": "ephemeral"}
def test_image_blocks_are_converted_to_anthropic_source(self):
body = self._transform(
[
{
"role": "user",
"content": [
{
"type": "image_url",
"image_url": {"url": "data:image/png;base64,ZmFrZQ==", "format": "image/jpeg"},
}
],
}
]
)
assert body["messages"][0]["content"] == [
{"type": "image", "source": {"type": "base64", "media_type": "image/jpeg", "data": "ZmFrZQ=="}}
]
def test_tool_result_image_list_is_converted(self):
body = self._transform(
[
{"role": "user", "content": "look"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{"id": "call_1", "type": "function", "function": {"name": "read", "arguments": "{}"}}
],
},
{
"role": "tool",
"tool_call_id": "call_1",
"content": [{"type": "image_url", "image_url": {"url": "data:image/png;base64,ZmFrZQ=="}}],
},
]
)
assert body["messages"][2]["content"][0]["content"] == [
{"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "ZmFrZQ=="}}
]
def test_tool_result_preserves_cache_control(self):
"""A cache breakpoint the bridge puts on a tool message must survive onto the tool_result."""
for tool_content in ("done", [{"type": "text", "text": "done"}]):
body = self._transform(
[
{"role": "user", "content": "look"},
{
"role": "tool",
"tool_call_id": "call_1",
"content": tool_content,
"cache_control": {"type": "ephemeral", "ttl": "1h"},
},
]
)
tool_result = body["messages"][1]["content"][0]
assert tool_result["cache_control"] == {"type": "ephemeral"}, tool_content
def test_pdf_data_uri_becomes_a_document_block(self):
"""A bridged pdf data URI is a document block; forwarding it as an image is malformed."""
body = self._transform(
[
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url": "data:application/pdf;base64,ZmFrZQ=="}},
],
}
]
)
assert body["messages"][0]["content"] == [
{
"type": "document",
"source": {"type": "base64", "media_type": "application/pdf", "data": "ZmFrZQ=="},
}
]
def test_multipart_tool_result_preserves_text_and_converts_image(self):
body = self._transform(
[
{"role": "user", "content": "look"},
{
"role": "tool",
"tool_call_id": "call_1",
"content": [
{"type": "text", "text": "first"},
{"type": "image_url", "image_url": {"url": "data:image/png;base64,ZmFrZQ=="}},
{"type": "text", "text": "last"},
],
},
]
)
assert body["messages"][1]["content"][0]["content"] == [
{"type": "text", "text": "first"},
{"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "ZmFrZQ=="}},
{"type": "text", "text": "last"},
]
def test_plain_text_tool_result_remains_string(self):
body = self._transform(
[{"role": "user", "content": "look"}, {"role": "tool", "tool_call_id": "call_1", "content": "done"}]
)
assert body["messages"][1]["content"][0]["content"] == "done"
def test_anthropic_tool_schema_strips_only_top_level_schema_key(self):
tools = [
{
"name": "f",
"input_schema": {"$schema": "schema", "type": "object", "properties": {"$schema": {"type": "string"}}},
}
]
body = self._transform([{"role": "user", "content": "hi"}], {"tools": tools})
schema = body["tools"][0]["input_schema"]
assert "$schema" not in schema
assert "$schema" in schema["properties"]
def test_tool_schema_strips_only_top_level_schema_key(self):
tools = [
{
"type": "function",
"function": {
"name": "f",
"parameters": {
"$schema": "https://json-schema.org/draft/2020-12/schema",
"type": "object",
"properties": {"$schema": {"type": "string"}},
},
},
}
]
body = self._transform([{"role": "user", "content": "hi"}], {"tools": tools})
schema = body["tools"][0]["input_schema"]
assert "$schema" not in schema
assert "$schema" in schema["properties"]
def test_streaming_tool_identity_is_emitted_only_on_start(self):
handler = SnowflakeStreamingHandler(streaming_response=[], sync_stream=True)
start = handler.chunk_parser(
{
"type": "content_block_start",
"index": 0,
"content_block": {"type": "tool_use", "id": "tool_1", "name": "read"},
}
)
first_delta = handler.chunk_parser(
{
"type": "content_block_delta",
"index": 0,
"delta": {"type": "input_json_delta", "partial_json": '{"path":'},
}
)
second_delta = handler.chunk_parser(
{
"type": "content_block_delta",
"index": 0,
"delta": {"type": "input_json_delta", "partial_json": '"/tmp"}'},
}
)
def _tool_call(chunk):
return chunk.choices[0].delta.tool_calls[0]
assert _tool_call(start).id == "tool_1"
assert _tool_call(start).function.name == "read"
assert _tool_call(first_delta).id is None
assert _tool_call(first_delta).function.name is None
assert _tool_call(second_delta).id is None
assert _tool_call(second_delta).function.name is None
assert _tool_call(first_delta).function.arguments == '{"path":'
assert _tool_call(second_delta).function.arguments == '"/tmp"}'
def test_signed_thinking_blocks_lead_the_assistant_turn(self):
"""Multi-turn tool use with thinking only works if the signed block is echoed back first."""
body = self._transform(
[
{"role": "user", "content": "hi"},
{
"role": "assistant",
"content": None,
"thinking_blocks": [
{"type": "thinking", "thinking": "391", "signature": "Eto"},
{"type": "thinking", "thinking": "unsigned"},
],
"tool_calls": [
{"id": "call_1", "type": "function", "function": {"name": "read", "arguments": "{}"}}
],
},
]
)
blocks = body["messages"][1]["content"]
assert blocks[0] == {"type": "thinking", "thinking": "391", "signature": "Eto"}
assert [b["type"] for b in blocks] == ["thinking", "tool_use"]
def test_signed_thinking_blocks_lead_a_plain_text_assistant_turn(self):
"""A thinking response without a tool call must also round-trip on the next request."""
body = self._transform(
[
{"role": "user", "content": "hi"},
{
"role": "assistant",
"content": "391",
"thinking_blocks": [{"type": "thinking", "thinking": "391", "signature": "Eto"}],
},
{"role": "user", "content": "continue"},
]
)
assert body["messages"][1] == {
"role": "assistant",
"content": [
{"type": "thinking", "thinking": "391", "signature": "Eto"},
{"type": "text", "text": "391"},
],
}
def test_signed_thinking_blocks_preserve_list_content(self):
"""Cached assistant text reaches this transform as a content list, not a string."""
body = self._transform(
[
{"role": "user", "content": "hi"},
{
"role": "assistant",
"content": [{"type": "text", "text": "391", "cache_control": {"type": "ephemeral"}}],
"thinking_blocks": [{"type": "thinking", "thinking": "391", "signature": "Eto"}],
},
{"role": "user", "content": "continue"},
]
)
assert body["messages"][1] == {
"role": "assistant",
"content": [
{"type": "thinking", "thinking": "391", "signature": "Eto"},
{"type": "text", "text": "391", "cache_control": {"type": "ephemeral"}},
],
}
def test_thinking_only_assistant_turn_sends_no_empty_text_block(self):
"""Anthropic-shaped APIs reject empty text blocks, so a content-less thinking turn is thinking only."""
body = self._transform(
[
{"role": "user", "content": "hi"},
{
"role": "assistant",
"content": None,
"thinking_blocks": [{"type": "thinking", "thinking": "391", "signature": "Eto"}],
},
{"role": "user", "content": "continue"},
]
)
assert body["messages"][1]["content"] == [{"type": "thinking", "thinking": "391", "signature": "Eto"}]
def test_streaming_surfaces_thinking_and_prompt_cache_usage(self):
"""Cortex streams thinking deltas, signatures and cache counts; all must reach the caller."""
handler = SnowflakeStreamingHandler(streaming_response=[], sync_stream=True)
handler.chunk_parser(
{
"type": "message_start",
"message": {"usage": {"input_tokens": 18, "cache_creation_input_tokens": 1323}},
}
)
thinking = handler.chunk_parser(
{
"type": "content_block_delta",
"index": 0,
"delta": {"type": "thinking_delta", "thinking": "391"},
}
)
signature = handler.chunk_parser(
{
"type": "content_block_delta",
"index": 0,
"delta": {"type": "signature_delta", "signature": "Eto"},
}
)
final = handler.chunk_parser(
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
"usage": {"output_tokens": 8, "cache_read_input_tokens": 1323},
}
)
assert thinking.choices[0].delta.reasoning_content == "391"
assert signature.choices[0].delta.thinking_blocks[0]["signature"] == "Eto"
assert final.usage.prompt_tokens_details.cached_tokens == 1323
class TestSnowFlakeCompletion:
model_name = "mistral"
@ -380,10 +749,7 @@ class TestSnowFlakeCompletion:
# PAT key was used
post_kwargs = mock_post.call_args_list[-1][1]
assert "xxxxx" in post_kwargs["headers"]["Authorization"]
assert (
post_kwargs["headers"]["X-Snowflake-Authorization-Token-Type"]
== "PROGRAMMATIC_ACCESS_TOKEN"
)
assert post_kwargs["headers"]["X-Snowflake-Authorization-Token-Type"] == "PROGRAMMATIC_ACCESS_TOKEN"
# account id was used
assert "AAAA-BBBB" in post_kwargs["url"]
@ -495,9 +861,7 @@ class TestSnowflakeChatCompletion:
)
mock_post.assert_called_once()
else:
with patch.object(
AsyncHTTPHandler, "post", new_callable=AsyncMock, return_value=mock_resp
) as mock_post:
with patch.object(AsyncHTTPHandler, "post", new_callable=AsyncMock, return_value=mock_resp) as mock_post:
response = asyncio.run(
acompletion(
model="snowflake/mistral-7b",
@ -580,8 +944,4 @@ class TestSnowflakeChatCompletion:
chunks_received = asyncio.run(_run())
assert len(chunks_received) > 0
content = "".join(
c.choices[0].delta.content
for c in chunks_received
if c.choices[0].delta.content
)
content = "".join(c.choices[0].delta.content for c in chunks_received if c.choices[0].delta.content)

View file

@ -338,7 +338,7 @@ class TestAnthropicConfigRequest:
litellm_params={},
headers={},
)
assert body["system"] == "You are helpful."
assert body["system"] == [{"type": "text", "text": "You are helpful."}]
assert all(m["role"] != "system" for m in body["messages"])
assert body["messages"][0] == {"role": "user", "content": "Hello"}
@ -422,6 +422,64 @@ class TestAnthropicConfigResponse:
assert result.usage.completion_tokens == 5
assert result.usage.total_tokens == 15
def test_prompt_cache_usage_is_surfaced(self):
"""Cortex reports cache creation/read counts; dropping them hides caching and bills cached input at full price."""
raw = httpx.Response(
200,
json={
"id": "msg_1",
"model": "claude-sonnet-4-6",
"content": [{"type": "text", "text": "hi"}],
"stop_reason": "end_turn",
"usage": {"input_tokens": 18, "cache_creation_input_tokens": 1323, "cache_read_input_tokens": 0},
},
)
result = self.cfg.transform_response(
model="snowflake/claude-sonnet-4-6",
raw_response=raw,
model_response=ModelResponse(),
logging_obj=_mock_logging(),
request_data={},
messages=[],
optional_params={},
litellm_params={},
encoding=None,
)
assert result.usage.prompt_tokens == 1341
assert result.usage.prompt_tokens_details.cache_creation_tokens == 1323
assert result.usage.prompt_tokens_details.cached_tokens == 0
def test_thinking_block_and_signature_are_preserved(self):
"""The signature must survive so a client can echo the thinking block on the next turn."""
raw = httpx.Response(
200,
json={
"id": "msg_1",
"model": "claude-sonnet-4-6",
"content": [
{"type": "thinking", "thinking": "391", "signature": "Eto"},
{"type": "text", "text": "391"},
],
"stop_reason": "end_turn",
"usage": {"input_tokens": 10, "output_tokens": 5},
},
)
result = self.cfg.transform_response(
model="snowflake/claude-sonnet-4-6",
raw_response=raw,
model_response=ModelResponse(),
logging_obj=_mock_logging(),
request_data={},
messages=[],
optional_params={},
litellm_params={},
encoding=None,
)
message = result.choices[0].message
assert message.content == "391"
assert message.reasoning_content == "391"
assert message.thinking_blocks[0]["signature"] == "Eto"
def test_stop_reason_end_turn_maps_to_stop(self):
raw = _make_anthropic_response()
result = self.cfg.transform_response(

View file

@ -8201,6 +8201,129 @@ class TestPreemptive401ModeAware:
await self._run(delegate, self.LITELLM_KEY_HEADERS, has_stored_token=False)
class TestSingleServerPreflightReachesIdJag:
"""The connect-time preflight is what turns a credential failure into an HTTP status the client
can read. An oauth2_id_jag server has to reach it: its subject comes from the assertion stored at
SSO login, so the failure is decided before any IdP call and there is nothing later in the session
that can report it (tools/list degrades to an empty list, tools/call to 'tool not found')."""
def _id_jag_server(self) -> MCPServer:
return MCPServer(
server_id="id-idjag",
name="idjag",
alias="idjag",
server_name="idjag",
url="https://idjag.test/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_id_jag,
client_id="gateway-client",
client_secret="gateway-secret",
token_exchange_endpoint="https://org-idp.test/oauth2/token",
id_jag_resource_token_endpoint="https://resource-as.test/oauth2/token",
mcp_info={"server_name": "idjag"},
)
async def _run(self, server: MCPServer, mcp_servers: list[str], preflight: AsyncMock) -> None:
from litellm.proxy._experimental.mcp_server import server as server_module
with (
patch.object( # test-quality-ok: route wiring must use the manager's configured server
server_module.global_mcp_server_manager,
"get_mcp_server_by_name",
return_value=server,
),
patch.object( # test-quality-ok: route wiring must invoke the manager preflight
server_module.global_mcp_server_manager,
"preflight_token_exchange",
preflight,
),
patch.object( # test-quality-ok: allowed-set resolution needs the DB; the test controls its answer
server_module, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])
),
):
await server_module._raise_preemptive_401_for_unauthenticated_servers(
scope={"type": "http", "method": "POST", "path": "/mcp/idjag", "headers": []},
mcp_servers=mcp_servers,
oauth2_headers={"Authorization": "Bearer sk-litellm-virtual-key"},
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
client_ip=None,
)
@pytest.mark.asyncio
async def test_id_jag_single_server_route_surfaces_the_preflight_status(self):
"""The 412 the preflight raises must propagate out of connect, not be swallowed."""
server = self._id_jag_server()
preflight = AsyncMock(side_effect=HTTPException(status_code=412, detail="no stored assertion"))
with pytest.raises(HTTPException) as exc:
await self._run(server, ["idjag"], preflight)
assert exc.value.status_code == 412
assert preflight.await_args.kwargs["server"] is server
@pytest.mark.asyncio
async def test_token_exchange_without_a_bearer_still_challenges_and_never_pre_flights(self):
"""The already-shipped OBO path must be untouched by the call site dropping its mode test.
A token_exchange server with no inbound bearer has nothing to exchange, so it still gets the
RFC 9728 discovery challenge from the block above and the preflight is never reached; pushing
a subject-less exchange through the resolver would turn that challenge into some other status
and strand a client that only had to SSO and retry."""
from litellm.proxy._experimental.mcp_server import server as server_module
token_exchange = MCPServer(
server_id="id-obo",
name="obo",
alias="obo",
server_name="obo",
url="https://obo.test/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_token_exchange,
token_exchange_endpoint="https://idp.test/oauth2/token",
client_id="cid",
client_secret="csec",
mcp_info={"server_name": "obo"},
)
preflight = AsyncMock()
with (
patch.object( # test-quality-ok: route wiring must use the manager's configured server
server_module.global_mcp_server_manager,
"get_mcp_server_by_name",
return_value=token_exchange,
),
patch.object( # test-quality-ok: route wiring must invoke the manager preflight
server_module.global_mcp_server_manager,
"preflight_token_exchange",
preflight,
),
pytest.raises(HTTPException) as exc,
):
await server_module._raise_preemptive_401_for_unauthenticated_servers(
scope={"type": "http", "method": "POST", "path": "/mcp/obo", "headers": []},
mcp_servers=["obo"],
oauth2_headers=None,
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
client_ip=None,
)
assert exc.value.status_code == 401
headers = exc.value.headers or {}
assert "resource_metadata" in (headers.get("WWW-Authenticate") or headers.get("www-authenticate") or "")
preflight.assert_not_awaited()
@pytest.mark.asyncio
async def test_id_jag_multi_server_route_still_absorbs_the_failure(self):
"""The aggregate contract is unchanged: with more than one target the preflight does not run,
so one server with no stored assertion cannot fail the whole connect."""
preflight = AsyncMock(side_effect=HTTPException(status_code=412, detail="no stored assertion"))
await self._run(self._id_jag_server(), ["idjag", "other"], preflight)
preflight.assert_not_awaited()
def _make_obo_server(alias: str) -> MCPServer:
return MCPServer(
server_id=f"id-{alias}",

View file

@ -2616,6 +2616,165 @@ class TestMCPServerManager:
await manager.preflight_token_exchange(server=server, oauth2_headers=None, user_api_key_auth=None)
assert resolved == ["good-subject"]
def _id_jag_server(self, server_id: str) -> "MCPServer":
return MCPServer(
server_id=server_id,
name=f"{server_id}-server",
url="https://up.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_id_jag,
client_id="gateway-client",
client_secret="gateway-secret",
token_exchange_endpoint="https://org-idp.example/oauth2/token",
id_jag_resource_token_endpoint="https://resource-as.example/oauth2/token",
)
@pytest.mark.asyncio
async def test_preflight_id_jag_surfaces_missing_assertion_as_a_plain_412(self):
"""ID-JAG's missing/expired-assertion precondition must reach the client as a 412 whose body
names the fix, at the transport edge. Without the preflight the session opens and the caller
gets a 200 with an empty tool list and then 'tool not found', which is not what happened.
412 is a precondition, not an RFC 9728 discovery challenge, so it carries no
WWW-Authenticate: there is nothing for the client to discover and retry against."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import CredError
summary = (
"ID-JAG requires an IdP identity assertion for this user and none is stored. "
"Sign in through LiteLLM SSO so the gateway captures one."
)
class _FakeProvider:
async def resolve_credentials(self, subject, server):
return Error(CredError.of_precondition_required(summary))
manager = MCPServerManager(cred_provider=_FakeProvider())
with pytest.raises(HTTPException) as exc_info:
await manager.preflight_token_exchange(
server=self._id_jag_server("id-jag-preflight-412"),
oauth2_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
)
assert exc_info.value.status_code == 412
assert summary in exc_info.value.detail
assert not (exc_info.value.headers or {})
@pytest.mark.asyncio
async def test_preflight_id_jag_surfaces_assertion_store_outage_as_503(self):
"""A store outage is the other failure the session would swallow, and it is a different
answer than 412: the user has nothing to fix by signing in again."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import CredError
class _FakeProvider:
async def resolve_credentials(self, subject, server):
return Error(CredError.of_upstream_unavailable("assertion store unreachable"))
manager = MCPServerManager(cred_provider=_FakeProvider())
with pytest.raises(HTTPException) as exc_info:
await manager.preflight_token_exchange(
server=self._id_jag_server("id-jag-preflight-503"),
oauth2_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
)
assert exc_info.value.status_code == 503
@pytest.mark.asyncio
async def test_preflight_id_jag_preflights_litellm_key_and_skips_identity_bearer(self):
"""ID-JAG preflights when Authorization carries a LiteLLM key, but skips a caller identity
bearer that the session passes through unchanged."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
StaticHeaderAuth,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
subjects = []
class _FakeProvider:
async def resolve_credentials(self, subject, server):
subjects.append(
(
subject.subject_id,
subject.inbound_token.get_secret_value() if subject.inbound_token else None,
)
)
return Ok(StaticHeaderAuth("Bearer minted-id-jag", header_name="Authorization"))
manager = MCPServerManager(cred_provider=_FakeProvider())
caller = UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1")
await manager.preflight_token_exchange(
server=self._id_jag_server("id-jag-preflight-key"),
oauth2_headers={"Authorization": "Bearer sk-litellm-virtual-key"},
raw_headers={"authorization": "Bearer sk-litellm-virtual-key"},
user_api_key_auth=caller,
)
assert subjects == [("u-1", None)]
await manager.preflight_token_exchange(
server=self._id_jag_server("id-jag-preflight-identity"),
oauth2_headers={"Authorization": "Bearer caller-idp-id-token"},
raw_headers={
"x-litellm-api-key": "Bearer sk-admission-key",
"authorization": "Bearer caller-idp-id-token",
},
user_api_key_auth=UserAPIKeyAuth(api_key="hashed-key", user_id="u-1"),
)
assert subjects == [("u-1", None)]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"server_fields",
[
{"auth_type": MCPAuth.none},
{"auth_type": MCPAuth.api_key, "authentication_token": "static-upstream-key"},
{"auth_type": MCPAuth.bearer_token, "authentication_token": "static-upstream-key"},
{
"auth_type": MCPAuth.oauth2,
"oauth2_flow": "client_credentials",
"client_id": "cid",
"client_secret": "csec",
"token_url": "https://idp.example.com/token",
},
{"auth_type": MCPAuth.true_passthrough},
],
)
async def test_preflight_resolves_nothing_for_a_mode_that_does_not_pre_flight(self, server_fields):
"""The manager is the only thing deciding which modes pre-flight, so it has to reject every
other mode itself. The single-server call site no longer tests the mode before calling, so a
mode that falls through here would start resolving its credential a second time, at connect,
for flows that never had a connect-time resolution at all."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import CredError
calls = []
class _FakeProvider:
async def resolve_credentials(self, subject, server):
calls.append(server.server_id)
return Error(CredError.of_misconfigured("the preflight must never get here"))
manager = MCPServerManager(cred_provider=_FakeProvider())
server = MCPServer(
server_id="not-pre-flighted",
name="not-pre-flighted-server",
url="https://up.example.com/mcp",
transport=MCPTransport.http,
**server_fields,
)
assert (
await manager.preflight_token_exchange(
server=server,
oauth2_headers={"Authorization": "Bearer sk-litellm-virtual-key"},
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
)
is None
)
assert calls == []
@pytest.mark.asyncio
@pytest.mark.parametrize(
"authorization",
@ -2638,7 +2797,9 @@ class TestMCPServerManager:
resolved: Final[list[str | None]] = []
class _FakeProvider:
async def resolve_credentials(self, subject: Subject, server: ServerSpec) -> Ok[StaticHeaderAuth, CredError]:
async def resolve_credentials(
self, subject: Subject, server: ServerSpec
) -> Ok[StaticHeaderAuth, CredError]:
resolved.append(subject.inbound_token.get_secret_value() if subject.inbound_token else None)
return Ok(StaticHeaderAuth("Bearer MINTED", header_name="Authorization"))
@ -2664,7 +2825,9 @@ class TestMCPServerManager:
resolved: Final[list[str | None]] = []
class _FakeProvider:
async def resolve_credentials(self, subject: Subject, server: ServerSpec) -> Ok[StaticHeaderAuth, CredError]:
async def resolve_credentials(
self, subject: Subject, server: ServerSpec
) -> Ok[StaticHeaderAuth, CredError]:
resolved.append(subject.inbound_token.get_secret_value() if subject.inbound_token else None)
return Ok(StaticHeaderAuth("Bearer MINTED", header_name="Authorization"))

View file

@ -754,6 +754,84 @@ def test_expand_wildcard_invalid_litellm_params_passthrough():
assert result == [deployment]
def test_get_complete_model_list_excludes_wildcard_routes_by_default():
"""Regression (LIT-4108): a wildcard with a matching router deployment leaked into /v1/models."""
from litellm import Router
from litellm.proxy.auth.model_checks import get_complete_model_list
router = Router(
model_list=[
{
"model_name": "bedrock/*",
"litellm_params": {"model": "bedrock/*"},
},
{
"model_name": "gpt-4",
"litellm_params": {"model": "openai/gpt-4"},
},
]
)
result = get_complete_model_list(
key_models=[],
team_models=[],
proxy_model_list=["bedrock/*", "gpt-4"],
user_model=None,
infer_model_from_keys=False,
return_wildcard_routes=False,
llm_router=router,
)
assert "bedrock/*" not in result
assert "gpt-4" in result
assert any(m.startswith("bedrock/") for m in result)
def test_get_complete_model_list_excludes_wildcard_routes_without_router():
from litellm.proxy.auth.model_checks import get_complete_model_list
result = get_complete_model_list(
key_models=[],
team_models=[],
proxy_model_list=["bedrock/*", "gpt-4"],
user_model=None,
infer_model_from_keys=False,
return_wildcard_routes=False,
llm_router=None,
)
assert "bedrock/*" not in result
assert "gpt-4" in result
assert any(m.startswith("bedrock/") for m in result)
def test_get_complete_model_list_includes_wildcard_routes_when_requested():
from litellm import Router
from litellm.proxy.auth.model_checks import get_complete_model_list
router = Router(
model_list=[
{
"model_name": "bedrock/*",
"litellm_params": {"model": "bedrock/*"},
},
]
)
result = get_complete_model_list(
key_models=[],
team_models=[],
proxy_model_list=["bedrock/*"],
user_model=None,
infer_model_from_keys=False,
return_wildcard_routes=True,
llm_router=router,
)
assert result.count("bedrock/*") == 1
assert any(m.startswith("bedrock/") and m != "bedrock/*" for m in result)
def test_add_known_models_refreshes_models_by_provider_for_wildcard_expansion():
"""models_by_provider was a frozen import-time snapshot of set unions, so cost map
reloads (which call add_known_models) never reached wildcard expansion until a

View file

@ -1238,6 +1238,18 @@ def test_health_liveness_endpoint(proxy_client):
print(f"\n/health/liveness response time: {duration_ms:.2f}ms")
def test_health_backlog_includes_admission_control_stats(proxy_client):
response = proxy_client.get("/health/backlog")
assert response.status_code == 200, response.text
assert set(response.json()) == {
"in_flight_requests",
"admitted_requests",
"queued_requests",
"rejected_requests",
}
def test_health_readiness(proxy_client):
"""
Test /health/readiness endpoint.

View file

@ -225,3 +225,71 @@ def test_compute_overall_action_all_passed():
def test_compute_overall_action_empty():
assert _compute_overall_action([]) == "passed"
class TestEnrichPolicyTemplateStreamKeepalive:
async def _collect_endpoint_body(self, monkeypatch, interval, delay=0.3) -> tuple[list[bytes], dict]:
import asyncio
from unittest.mock import MagicMock
import litellm
import litellm.proxy.management_endpoints.policy_endpoints.endpoints as policy_endpoints
import litellm.proxy.proxy_server as proxy_server
from fastapi.responses import StreamingResponse
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.management_endpoints.policy_endpoints.endpoints import (
EnrichTemplateRequest,
enrich_policy_template_stream,
)
monkeypatch.setattr(litellm, "sse_keepalive_ping_interval_seconds", interval)
async def _name_chunks():
await asyncio.sleep(delay)
chunk = MagicMock()
chunk.choices = [MagicMock()]
chunk.choices[0].delta.content = "Rival Air\n"
yield chunk
class SlowRouter:
async def acompletion(self, **kwargs):
return _name_chunks()
async def _no_variations(competitors, model):
return {}
monkeypatch.setattr(proxy_server, "llm_router", SlowRouter())
monkeypatch.setattr(policy_endpoints, "_generate_competitor_variations", _no_variations)
response = await enrich_policy_template_stream(
data=EnrichTemplateRequest(
template_id="competitor-mention-detection",
parameters={"brand_name": "Acme"},
model="gpt-5.4-mini",
),
request=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN),
)
assert isinstance(response, StreamingResponse)
chunks = [chunk if isinstance(chunk, bytes) else chunk.encode() async for chunk in response.body_iterator]
return chunks, dict(response.headers)
@pytest.mark.asyncio
async def test_endpoint_pings_while_competitor_discovery_is_still_running(self, monkeypatch):
chunks, headers = await self._collect_endpoint_body(monkeypatch, interval=0.05)
assert headers["content-type"].startswith("text/event-stream")
assert headers["cache-control"] == "no-cache"
assert headers["x-accel-buffering"] == "no"
assert chunks[0] == b": ping\n\n"
assert chunks.count(b": ping\n\n") >= 3
assert b'data: {"type": "competitor", "name": "Rival Air"}\n\n' in chunks
assert chunks[-1].startswith(b'data: {"type": "done"')
@pytest.mark.asyncio
async def test_endpoint_stream_is_untouched_while_keepalives_are_unconfigured(self, monkeypatch):
chunks, _ = await self._collect_endpoint_body(monkeypatch, interval=None, delay=0.15)
assert b": ping\n\n" not in chunks
assert chunks[0] == b'data: {"type": "competitor", "name": "Rival Air"}\n\n'
assert chunks[-1].startswith(b'data: {"type": "done"')

View file

@ -466,3 +466,61 @@ class TestUsageAiChatServiceAccountGuard:
is_admin=False,
)
assert "Endpoint-level guard missing" in str(exc_info.value)
class TestUsageAiChatKeepalive:
async def _collect_endpoint_body(self, monkeypatch, interval, delay=0.3) -> tuple[list[bytes], dict]:
import asyncio
import litellm
from fastapi.responses import StreamingResponse
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.management_endpoints.usage_endpoints.endpoints import (
ChatMessage,
UsageAIChatRequest,
usage_ai_chat,
)
monkeypatch.setattr(litellm, "sse_keepalive_ping_interval_seconds", interval)
async def slow_acompletion(**kwargs):
await asyncio.sleep(delay)
response = MagicMock()
response.choices = [MagicMock()]
response.choices[0].message.tool_calls = None
response.choices[0].message.content = "Total spend is $50.25"
return response
with patch( # test-quality-ok: the stream calls the module-level litellm.acompletion directly; no injection seam
"litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm.acompletion",
new=AsyncMock(side_effect=slow_acompletion),
):
response = await usage_ai_chat(
data=UsageAIChatRequest(messages=[ChatMessage(role="user", content="hi")], model="gpt-4o-mini"),
request=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN),
)
assert isinstance(response, StreamingResponse)
chunks = [chunk if isinstance(chunk, bytes) else chunk.encode() async for chunk in response.body_iterator]
return chunks, dict(response.headers)
@pytest.mark.asyncio
async def test_endpoint_pings_while_the_planning_completion_is_still_running(self, monkeypatch):
chunks, headers = await self._collect_endpoint_body(monkeypatch, interval=0.05)
assert headers["content-type"].startswith("text/event-stream")
assert headers["cache-control"] == "no-cache"
assert headers["x-accel-buffering"] == "no"
assert chunks[0].startswith(b'data: {"type": "status"')
assert chunks[1] == b": ping\n\n"
assert chunks.count(b": ping\n\n") >= 3
assert b'"content": "Total spend is $50.25"' in b"".join(chunks)
assert chunks[-1] == b'data: {"type": "done"}\n\n'
@pytest.mark.asyncio
async def test_endpoint_stream_is_untouched_while_keepalives_are_unconfigured(self, monkeypatch):
chunks, _ = await self._collect_endpoint_body(monkeypatch, interval=None, delay=0.15)
assert b": ping\n\n" not in chunks
assert chunks[0].startswith(b'data: {"type": "status"')
assert chunks[-1] == b'data: {"type": "done"}\n\n'

View file

@ -0,0 +1,402 @@
import asyncio
import json
from typing import Final
import pytest
from starlette.middleware.base import BaseHTTPMiddleware
from starlette.types import ASGIApp, Message, Receive, Scope, Send
from litellm.proxy.middleware.admission_control_middleware import (
AdmissionControlMetrics,
AdmissionControlMiddleware,
AdmissionControlSettings,
AdmissionControlState,
AdmissionControlStats,
_parse_admission_control_settings,
create_prometheus_admission_metrics,
get_admission_control_settings,
)
@pytest.fixture
def state() -> AdmissionControlState:
return AdmissionControlState(lambda: None)
async def _call(
middleware: AdmissionControlMiddleware,
path: str = "/",
root_path: str = "",
) -> tuple[Message, ...]:
messages: Final[list[Message]] = []
async def receive() -> Message:
return {"type": "http.request", "body": b"", "more_body": False}
async def send(message: Message) -> None:
messages.append(message)
scope: Final[Scope] = {
"type": "http",
"path": path,
"root_path": root_path,
"method": "GET",
"headers": [],
}
await middleware(scope, receive, send)
return tuple(messages)
def _handler_with_release(
started: asyncio.Event,
release: asyncio.Event,
) -> ASGIApp:
async def handler(scope: Scope, receive: Receive, send: Send) -> None:
started.set()
await release.wait()
await send({"type": "http.response.start", "status": 200, "headers": []})
await send({"type": "http.response.body", "body": b"ok", "more_body": False})
return handler
def test_is_not_base_http_middleware() -> None:
assert not issubclass(AdmissionControlMiddleware, BaseHTTPMiddleware)
@pytest.mark.asyncio
async def test_capacity_rejects_excess_and_releases_queued_request(state: AdmissionControlState) -> None:
started: Final = asyncio.Event()
release: Final = asyncio.Event()
middleware: Final = AdmissionControlMiddleware(
_handler_with_release(started, release),
lambda: AdmissionControlSettings(1, 1, 1.0),
state,
)
first: Final = asyncio.create_task(_call(middleware))
await started.wait()
second: Final = asyncio.create_task(_call(middleware))
await asyncio.sleep(0)
assert state.get_stats().queued == 1
third: Final = await _call(middleware)
assert third[0]["status"] == 503
headers: Final = dict(third[0]["headers"])
assert headers[b"retry-after"] == b"1"
assert headers[b"content-type"] == b"application/json"
assert json.loads(third[1]["body"])["error"] == {
"message": "Worker at capacity: 1 in-flight, 1 queued requests. Retry later.",
"type": "overloaded_error",
"code": "503",
}
assert state.get_stats().rejected_total == 1
release.set()
assert (await first)[0]["status"] == 200
assert (await second)[0]["status"] == 200
assert state.get_stats() == AdmissionControlStats(0, 0, 1)
@pytest.mark.asyncio
async def test_pending_waiter_is_not_skipped_after_admission_is_released(state: AdmissionControlState) -> None:
started: Final = asyncio.Event()
release: Final = asyncio.Event()
third_trigger: Final = asyncio.Event()
middleware: Final = AdmissionControlMiddleware(
_handler_with_release(started, release),
lambda: AdmissionControlSettings(1, 2, 1.0),
state,
)
first: Final = asyncio.create_task(_call(middleware))
await started.wait()
second: Final = asyncio.create_task(_call(middleware))
await asyncio.sleep(0)
async def call_third() -> tuple[Message, ...]:
await third_trigger.wait()
return await _call(middleware)
third: Final = asyncio.create_task(call_third())
await asyncio.sleep(0)
release.set()
third_trigger.set()
await asyncio.sleep(0)
assert state.get_stats().queued == 2
await asyncio.gather(first, second, third)
@pytest.mark.asyncio
async def test_queue_timeout_rejects_and_decrements_queue(state: AdmissionControlState) -> None:
started: Final = asyncio.Event()
release: Final = asyncio.Event()
middleware: Final = AdmissionControlMiddleware(
_handler_with_release(started, release),
lambda: AdmissionControlSettings(1, 1, 0.05),
state,
)
first: Final = asyncio.create_task(_call(middleware))
await started.wait()
start_time: Final = asyncio.get_running_loop().time()
second: Final = await _call(middleware)
elapsed: Final = asyncio.get_running_loop().time() - start_time
assert second[0]["status"] == 503
assert elapsed < 0.5
assert state.get_stats().queued == 0
assert state.get_stats().rejected_total == 1
release.set()
await first
@pytest.mark.asyncio
@pytest.mark.parametrize(
("root_path", "probe_path"),
(
("", "/health/liveliness"),
("/proxy", "/proxy/health/liveliness"),
("/proxy", "/proxy/metrics"),
),
)
async def test_exempt_path_passes_through_when_saturated(
state: AdmissionControlState,
root_path: str,
probe_path: str,
) -> None:
started: Final = asyncio.Event()
release: Final = asyncio.Event()
async def handler(scope: Scope, receive: Receive, send: Send) -> None:
if scope["path"] == "/":
started.set()
await release.wait()
await send({"type": "http.response.start", "status": 200, "headers": []})
await send({"type": "http.response.body", "body": b"ok", "more_body": False})
middleware: Final = AdmissionControlMiddleware(handler, lambda: AdmissionControlSettings(1, 0, 1.0), state)
first: Final = asyncio.create_task(_call(middleware))
await started.wait()
health: Final = await _call(middleware, probe_path, root_path)
assert health[0]["status"] == 200
blocked: Final = await _call(middleware, "/proxy/v1/chat/completions", root_path)
assert blocked[0]["status"] == 503
lookalike: Final = await _call(middleware, "/proxyhealth/liveliness", "/proxy")
assert lookalike[0]["status"] == 503
release.set()
await first
@pytest.mark.asyncio
async def test_non_http_scope_passes_through_when_saturated(state: AdmissionControlState) -> None:
seen: Final[list[str]] = []
async def handler(scope: Scope, receive: Receive, send: Send) -> None:
seen.append(scope["type"])
middleware: Final = AdmissionControlMiddleware(handler, lambda: AdmissionControlSettings(1, 0, 1.0), state)
state.record_admission()
async def receive() -> Message:
return {"type": "lifespan.startup"}
async def send(message: Message) -> None:
return None
await middleware({"type": "lifespan"}, receive, send)
assert seen == ["lifespan"]
@pytest.mark.asyncio
async def test_none_settings_does_not_limit_concurrency() -> None:
active: Final = [0]
peak: Final = [0]
all_started: Final = asyncio.Event()
release: Final = asyncio.Event()
async def handler(scope: Scope, receive: Receive, send: Send) -> None:
active[0] += 1
peak[0] = max(peak[0], active[0])
if active[0] == 3:
all_started.set()
await release.wait()
active[0] -= 1
await send({"type": "http.response.start", "status": 200, "headers": []})
await send({"type": "http.response.body", "body": b"ok", "more_body": False})
middleware: Final = AdmissionControlMiddleware(handler, lambda: None, AdmissionControlState(lambda: None))
requests: Final = tuple(asyncio.create_task(_call(middleware)) for _ in range(3))
await all_started.wait()
assert peak[0] == 3
release.set()
results: Final = await asyncio.gather(*requests)
assert tuple(result[0]["status"] for result in results) == (200, 200, 200)
@pytest.mark.asyncio
async def test_cancelling_queued_request_does_not_leak_counter(state: AdmissionControlState) -> None:
started: Final = asyncio.Event()
release: Final = asyncio.Event()
middleware: Final = AdmissionControlMiddleware(
_handler_with_release(started, release),
lambda: AdmissionControlSettings(1, 1, 1.0),
state,
)
first: Final = asyncio.create_task(_call(middleware))
await started.wait()
queued: Final = asyncio.create_task(_call(middleware))
await asyncio.sleep(0)
queued.cancel()
with pytest.raises(asyncio.CancelledError):
await queued
assert state.get_stats().queued == 0
release.set()
await first
@pytest.mark.asyncio
async def test_streaming_response_holds_admission_until_final_body(state: AdmissionControlState) -> None:
first_chunk_sent: Final = asyncio.Event()
finish_stream: Final = asyncio.Event()
async def handler(scope: Scope, receive: Receive, send: Send) -> None:
await send({"type": "http.response.start", "status": 200, "headers": []})
await send({"type": "http.response.body", "body": b"first", "more_body": True})
first_chunk_sent.set()
await finish_stream.wait()
await send({"type": "http.response.body", "body": b"last", "more_body": False})
middleware: Final = AdmissionControlMiddleware(
handler,
lambda: AdmissionControlSettings(1, 1, 1.0),
state,
)
first: Final = asyncio.create_task(_call(middleware))
await first_chunk_sent.wait()
second: Final = asyncio.create_task(_call(middleware))
await asyncio.sleep(0)
assert not second.done()
assert state.get_stats().queued == 1
finish_stream.set()
assert (await first)[0]["status"] == 200
assert (await second)[0]["status"] == 200
assert state.get_stats().admitted == 0
assert state.get_stats().queued == 0
class _FakeGauge:
def __init__(self) -> None:
self.value = 0.0
def inc(self, amount: float = 1) -> None:
self.value += amount
def dec(self, amount: float = 1) -> None:
self.value -= amount
class _FakeCounter:
def __init__(self) -> None:
self.by_reason: Final[dict[str, _FakeGauge]] = {}
def labels(self, reason: str) -> _FakeGauge:
return self.by_reason.setdefault(reason, _FakeGauge())
@pytest.mark.asyncio
async def test_metrics_track_admitted_queued_and_rejected() -> None:
admitted: Final = _FakeGauge()
queued: Final = _FakeGauge()
rejected: Final = _FakeCounter()
state: Final = AdmissionControlState(
lambda: AdmissionControlMetrics(admitted_gauge=admitted, queued_gauge=queued, rejected_counter=rejected)
)
started: Final = asyncio.Event()
release: Final = asyncio.Event()
middleware: Final = AdmissionControlMiddleware(
_handler_with_release(started, release),
lambda: AdmissionControlSettings(1, 1, 0.05),
state,
)
first: Final = asyncio.create_task(_call(middleware))
await started.wait()
second: Final = asyncio.create_task(_call(middleware))
await asyncio.sleep(0)
assert (admitted.value, queued.value) == (1.0, 1.0)
await _call(middleware)
assert rejected.by_reason["queue_full"].value == 1.0
await second
assert rejected.by_reason["queue_timeout"].value == 1.0
release.set()
await first
assert (admitted.value, queued.value) == (0.0, 0.0)
def test_create_prometheus_admission_metrics_registers_named_metrics() -> None:
from prometheus_client import REGISTRY
metrics: Final = create_prometheus_admission_metrics()
if metrics is not None:
metrics.admitted_gauge.inc()
metrics.queued_gauge.inc()
metrics.rejected_counter.labels(reason="queue_full").inc()
assert REGISTRY.get_sample_value("litellm_admission_admitted_requests") == 1.0
assert REGISTRY.get_sample_value("litellm_admission_queued_requests") == 1.0
assert REGISTRY.get_sample_value("litellm_admission_rejected_requests_total", {"reason": "queue_full"}) is not None
assert create_prometheus_admission_metrics() is None
@pytest.mark.parametrize(
("settings", "expected"),
(
({}, None),
({"max_in_flight_requests_per_worker": None}, None),
({"max_in_flight_requests_per_worker": 0}, None),
({"max_in_flight_requests_per_worker": "many"}, None),
({"max_in_flight_requests_per_worker": 3, "max_queued_requests_per_worker": -1}, None),
({"max_in_flight_requests_per_worker": 3, "admission_queue_timeout_seconds": 0}, None),
({"max_in_flight_requests_per_worker": 3, "admission_queue_timeout_seconds": -0.5}, None),
(
{"max_in_flight_requests_per_worker": 3, "max_queued_requests_per_worker": 0},
AdmissionControlSettings(3, 0, 1.0),
),
(
{"max_in_flight_requests_per_worker": 3},
AdmissionControlSettings(3, 3, 1.0),
),
(
{
"max_in_flight_requests_per_worker": 3,
"max_queued_requests_per_worker": 5,
"admission_queue_timeout_seconds": 0.25,
},
AdmissionControlSettings(3, 5, 0.25),
),
),
)
def test_get_admission_control_settings(
settings: dict[str, object],
expected: AdmissionControlSettings | None,
) -> None:
assert get_admission_control_settings(settings) == expected
def test_invalid_admission_control_settings_logs_once(caplog: pytest.LogCaptureFixture) -> None:
_parse_admission_control_settings.cache_clear()
caplog.set_level("ERROR")
settings: Final = {"max_in_flight_requests_per_worker": [1]}
assert get_admission_control_settings(settings) is None
assert get_admission_control_settings(settings) is None
messages: Final = tuple(
record.message
for record in caplog.records
if record.message.startswith("Ignoring invalid admission control settings")
)
assert len(messages) == 1

View file

@ -5116,3 +5116,93 @@ class TestAzureRouterModelStreamingDispatch:
assert result.status_code == 200
body = b"".join([chunk async for chunk in result.body_iterator])
assert body == upstream_body
class TestAzureRouterModelStreamingKeepalive:
async def _dispatch(self, monkeypatch, interval, headers_delay=0.0, body_delay=0.0) -> StreamingResponse:
import asyncio
import litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints as ep
import litellm.proxy.proxy_server as proxy_server
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
monkeypatch.setattr(litellm, "sse_keepalive_ping_interval_seconds", interval)
class _StallingBody(httpx.AsyncByteStream):
async def __aiter__(self):
await asyncio.sleep(body_delay)
yield b"data: hello\n\n"
async def _upstream_response() -> httpx.Response:
await asyncio.sleep(headers_delay)
return httpx.Response(
200,
headers={"content-type": "text/event-stream", "x-upstream": "kept"},
stream=_StallingBody(),
request=httpx.Request("POST", "https://my-azure.openai.azure.com/openai/deployments/gpt-5/x"),
)
logging_obj = MagicMock()
logging_obj.async_flush_passthrough_collected_chunks = AsyncMock()
class StreamingRouter:
async def allm_passthrough_route(self, **kwargs):
return await AsyncPassthroughStreamingResponse(
response=_upstream_response(),
litellm_logging_obj=logging_obj,
provider_config=MagicMock(),
)
async def fake_get_request_body(_request):
return {"model": "gpt-5", "stream": True}
monkeypatch.setattr(proxy_server, "llm_router", StreamingRouter())
monkeypatch.setattr(ep, "get_request_body", fake_get_request_body)
monkeypatch.setattr(ep, "is_passthrough_request_using_router_model", lambda *a, **k: True)
request = MagicMock(spec=Request)
request.method = "POST"
request.headers = {"content-type": "application/json"}
request.query_params = {}
result = await azure_proxy_route(
endpoint="openai/deployments/gpt-5/chat/completions",
request=request,
fastapi_response=MagicMock(spec=Response),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-token"),
)
assert isinstance(result, StreamingResponse)
return result
@pytest.mark.asyncio
async def test_pings_while_upstream_headers_are_still_pending(self, monkeypatch):
result = await self._dispatch(monkeypatch, interval=0.05, headers_delay=0.3)
chunks = [chunk async for chunk in result.body_iterator]
assert result.status_code == 200
assert result.headers["x-accel-buffering"] == "no"
assert chunks[0] == b": ping\n\n"
assert chunks.count(b": ping\n\n") >= 3
assert b"".join(chunks).endswith(b"data: hello\n\n")
@pytest.mark.asyncio
async def test_pings_while_upstream_body_is_still_pending(self, monkeypatch):
result = await self._dispatch(monkeypatch, interval=0.05, body_delay=0.3)
chunks = [chunk async for chunk in result.body_iterator]
assert result.status_code == 200
assert result.headers["x-upstream"] == "kept"
assert chunks[0] == b": ping\n\n"
assert chunks.count(b": ping\n\n") >= 3
assert chunks[-1] == b"data: hello\n\n"
@pytest.mark.asyncio
async def test_relays_upstream_bytes_untouched_while_keepalives_are_unconfigured(self, monkeypatch):
result = await self._dispatch(monkeypatch, interval=None, headers_delay=0.15, body_delay=0.15)
chunks = [chunk async for chunk in result.body_iterator]
assert result.headers["x-upstream"] == "kept"
assert chunks == [b"data: hello\n\n"]

View file

@ -1819,3 +1819,92 @@ async def test_run_thread_stream_is_untouched_while_keepalives_are_unconfigured(
assert not any(chunk.startswith(": ping") for chunk in chunks)
assert chunks[-1] == "data: [DONE]\n\n"
# ---------------------------------------------------------------------------
# async_queue_request: SSE keepalives during the time-to-first-token
# ---------------------------------------------------------------------------
async def _queue_streaming(monkeypatch, interval, delay=0.3, fails_with=None):
_patch_logging_flags(monkeypatch)
monkeypatch.setattr(litellm, "sse_keepalive_ping_interval_seconds", interval)
router = MagicMock()
router.get_model_list.return_value = []
async def _schedule_after_the_scheduler_queue_drains(**kwargs):
await asyncio.sleep(delay)
if fails_with is not None:
raise fails_with
return _async_iter([_simple_chunk(content="queued reply")])
router.schedule_acompletion = _schedule_after_the_scheduler_queue_drains
monkeypatch.setattr(ps, "llm_router", router)
request = MagicMock()
request.url = "http://testserver/queue/chat/completions"
request.method = "POST"
request.headers = {}
request.json = AsyncMock(
return_value={
"model": "gpt-4",
"messages": [{"role": "user", "content": "hi"}],
"priority": 0,
"stream": True,
}
)
request.is_disconnected = AsyncMock(return_value=False)
return await ps.async_queue_request(
request=request,
fastapi_response=Response(),
user_api_key_dict=_user_auth(),
)
@pytest.mark.asyncio
async def test_queue_request_pings_while_the_scheduler_is_still_waiting(monkeypatch):
response = await _queue_streaming(monkeypatch, interval=0.05)
assert isinstance(response, StreamingResponse)
assert response.headers["x-accel-buffering"] == "no"
chunks = [chunk async for chunk in response.body_iterator]
assert chunks[0] == b": ping\n\n"
assert chunks.count(b": ping\n\n") >= 3
assert b'"content":"queued reply"' in chunks[-2]
assert chunks[-1] == b"data: [DONE]\n\n"
@pytest.mark.asyncio
async def test_queue_request_audits_a_failure_that_arrives_after_the_first_ping(monkeypatch):
audited = []
async def _record_failure(*, user_api_key_dict, original_exception, request_data, **kwargs):
audited.append(original_exception)
return None
monkeypatch.setattr(ps.proxy_logging_obj, "post_call_failure_hook", _record_failure)
boom = RuntimeError("scheduler died after the wire was already open")
response = await _queue_streaming(monkeypatch, interval=0.05, fails_with=boom)
assert isinstance(response, StreamingResponse)
chunks = [chunk async for chunk in response.body_iterator]
assert chunks[0] == b": ping\n\n"
assert audited == [boom]
assert json.loads(chunks[-2].removeprefix(b"data: "))["error"]["code"] == "500"
assert chunks[-1] == b"data: [DONE]\n\n"
@pytest.mark.asyncio
async def test_queue_request_stream_is_untouched_while_keepalives_are_unconfigured(monkeypatch):
response = await _queue_streaming(monkeypatch, interval=None, delay=0.15)
assert isinstance(response, StreamingResponse)
chunks = [chunk if isinstance(chunk, bytes) else chunk.encode() async for chunk in response.body_iterator]
assert not any(chunk.startswith(b": ping") for chunk in chunks)
assert chunks[-1] == b"data: [DONE]\n\n"

View file

@ -11,7 +11,6 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi.testclient import TestClient
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.proxy_server import app
@ -324,6 +323,100 @@ def test_rag_query_stream_returns_event_stream(client_internal_user):
assert "data: [DONE]" in response.text
def test_rag_query_stream_pings_while_retrieval_is_still_running(client_internal_user, monkeypatch):
import asyncio
import litellm as litellm_module
monkeypatch.setattr(litellm_module, "sse_keepalive_ping_interval_seconds", 0.05)
async def slow_aquery(**kwargs):
await asyncio.sleep(0.3)
return await litellm_module.acompletion(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "What is the codename?"}],
mock_response="The codename is AZURE-FALCON-42.",
stream=True,
api_key="test-key",
)
with (
patch( # test-quality-ok: the handler calls the module-level litellm.aquery directly; no injection seam
"litellm.proxy.rag_endpoints.endpoints.litellm.aquery",
new=AsyncMock(side_effect=slow_aquery),
),
patch("litellm.vector_store_registry", None), # test-quality-ok: proxy module global, no injection seam
patch("litellm.proxy.proxy_server.prisma_client", None), # test-quality-ok: proxy module global, no injection seam
):
response = client_internal_user.post(
"/v1/rag/query",
json={
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "What is the codename?"}],
"retrieval_config": {
"vector_store_id": "vs_test_123",
"custom_llm_provider": "openai",
},
"stream": True,
},
)
assert response.status_code == 200, response.text
assert response.headers.get("content-type", "").startswith("text/event-stream")
assert response.headers["x-accel-buffering"] == "no"
assert response.text.startswith(": ping\n\n")
assert response.text.count(": ping\n\n") >= 3
assert '"object":"chat.completion.chunk"' in response.text
assert response.text.endswith("data: [DONE]\n\n")
def test_rag_query_stream_keeps_response_headers_when_retrieval_beats_the_keepalive(
client_internal_user, monkeypatch
):
import litellm as litellm_module
monkeypatch.setattr(litellm_module, "sse_keepalive_ping_interval_seconds", 5)
async def fast_aquery(**kwargs):
response = await litellm_module.acompletion(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "What is the codename?"}],
mock_response="The codename is AZURE-FALCON-42.",
stream=True,
api_key="test-key",
)
response._hidden_params["response_cost"] = 3.45e-06
return response
with (
patch( # test-quality-ok: the handler calls the module-level litellm.aquery directly; no injection seam
"litellm.proxy.rag_endpoints.endpoints.litellm.aquery",
new=AsyncMock(side_effect=fast_aquery),
),
patch("litellm.vector_store_registry", None), # test-quality-ok: proxy module global, no injection seam
patch("litellm.proxy.proxy_server.prisma_client", None), # test-quality-ok: proxy module global, no injection seam
):
response = client_internal_user.post(
"/v1/rag/query",
json={
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "What is the codename?"}],
"retrieval_config": {
"vector_store_id": "vs_test_123",
"custom_llm_provider": "openai",
},
"stream": True,
},
)
assert response.status_code == 200, response.text
assert response.headers.get("content-type", "").startswith("text/event-stream")
assert response.headers.get("x-litellm-response-cost") == "3.45e-06"
assert not response.text.startswith(": ping")
assert '"object":"chat.completion.chunk"' in response.text
assert response.text.endswith("data: [DONE]\n\n")
def test_rag_query_merges_managed_store_params(client_internal_user):
"""
Regression: /v1/rag/query must consult the managed vector store registry

View file

@ -6,6 +6,7 @@ from unittest.mock import MagicMock, patch
# Adds the grandparent directory to sys.path to allow importing project modules
import pytest
from azure.core.exceptions import ClientAuthenticationError
from litellm.secret_managers.get_azure_ad_token_provider import (
get_azure_ad_token_provider,
@ -16,6 +17,143 @@ from litellm.types.secret_managers.get_azure_ad_token_provider import (
)
class TestDeploymentIdentityCredential:
@staticmethod
def _chain_for(credential_type):
with patch("azure.identity.get_bearer_token_provider", return_value=lambda: "token") as bearer:
get_azure_ad_token_provider(
azure_scope="https://storage.azure.com/.default",
azure_credential=credential_type,
)
bearer.assert_called_once()
with bearer.call_args.args[0] as chain:
return {type(link).__name__ for link in chain.credentials}
@staticmethod
def _managed_identity_client_ids(credential_type):
with patch("azure.identity.get_bearer_token_provider", return_value=lambda: "token") as bearer:
get_azure_ad_token_provider(
azure_scope="https://storage.azure.com/.default",
azure_credential=credential_type,
)
with bearer.call_args.args[0] as chain:
return [
(link._credential._settings or {}).get("client_id")
for link in chain.credentials
if type(link).__name__ == "ManagedIdentityCredential"
]
@patch.dict(
os.environ,
{
"AZURE_CLIENT_ID": "workload-identity-client-id",
"AZURE_TENANT_ID": "workload-identity-tenant-id",
"AZURE_FEDERATED_TOKEN_FILE": "/var/run/secrets/azure/tokens/azure-identity-token",
},
clear=True,
)
def test_deployment_identity_reaches_workload_and_managed_identity_only(self):
assert self._chain_for(AzureCredentialType.DeploymentIdentityCredential) == {
"WorkloadIdentityCredential",
"ManagedIdentityCredential",
}
@patch.dict(
os.environ,
{
"AZURE_CLIENT_ID": "workload-identity-client-id",
"AZURE_TENANT_ID": "workload-identity-tenant-id",
"AZURE_FEDERATED_TOKEN_FILE": "/var/run/secrets/azure/tokens/azure-identity-token",
"AZURE_TOKEN_CREDENTIALS": "dev",
},
clear=True,
)
def test_deployment_identity_survives_a_developer_only_token_credentials_setting(self):
"""AZURE_TOKEN_CREDENTIALS=dev asks the SDK for developer credentials only, which is every
credential this chain drops, so the deployment's own identity has to win over it"""
assert self._chain_for(AzureCredentialType.DeploymentIdentityCredential) == {
"WorkloadIdentityCredential",
"ManagedIdentityCredential",
}
@patch.dict(
os.environ,
{
"AZURE_CLIENT_ID": "azure-openai-client-id",
"AZURE_CLIENT_SECRET": "azure-openai-client-secret",
"AZURE_TENANT_ID": "azure-openai-tenant-id",
},
clear=True,
)
def test_default_azure_credential_keeps_its_full_chain(self):
"""Azure OpenAI callers pass DefaultAzureCredential and must be unaffected by the
narrowing that the storage callback asks for"""
full_chain = self._chain_for(AzureCredentialType.DefaultAzureCredential)
assert "EnvironmentCredential" in full_chain
assert "AzureCliCredential" in full_chain
assert "EnvironmentCredential" not in self._chain_for(
AzureCredentialType.DeploymentIdentityCredential
)
@patch.dict(
os.environ,
{
"AZURE_CLIENT_ID": "azure-openai-client-id",
"AZURE_CLIENT_SECRET": "azure-openai-client-secret",
"AZURE_TENANT_ID": "azure-openai-tenant-id",
},
clear=True,
)
def test_deployment_identity_refuses_to_mint_a_token_for_a_configured_service_principal(self):
"""A host carrying only an Azure OpenAI client secret must get no token at all, and the
refusal must name the identities that were actually tried"""
provider = get_azure_ad_token_provider(
azure_scope="https://storage.azure.com/.default",
azure_credential=AzureCredentialType.DeploymentIdentityCredential,
)
with pytest.raises(ClientAuthenticationError) as refusal:
provider()
assert "ManagedIdentityCredential" in str(refusal.value)
assert "EnvironmentCredential" not in str(refusal.value)
assert "AzureCliCredential" not in str(refusal.value)
assert "azure-openai-client-secret" not in str(refusal.value)
@patch.dict(
os.environ,
{
"AZURE_CLIENT_ID": "azure-openai-client-id",
"AZURE_CLIENT_SECRET": "azure-openai-client-secret",
"AZURE_TENANT_ID": "azure-openai-tenant-id",
},
clear=True,
)
def test_deployment_identity_still_reaches_a_system_assigned_managed_identity(self):
"""AZURE_CLIENT_ID names one identity for the whole proxy, and pointing it at Azure OpenAI
must not hide the system assigned identity the host runs as"""
client_ids = self._managed_identity_client_ids(AzureCredentialType.DeploymentIdentityCredential)
assert "azure-openai-client-id" in client_ids
assert None in client_ids
@patch.dict(
os.environ,
{
"AZURE_CLIENT_ID": "user-assigned-identity-client-id",
"AZURE_TOKEN_CREDENTIALS": "dev",
},
clear=True,
)
def test_deployment_identity_keeps_the_user_assigned_identity_under_a_dev_only_setting(self):
"""AZURE_TOKEN_CREDENTIALS=dev asks the SDK for developer credentials only, and the
identity a host actually runs as has to survive that"""
assert "user-assigned-identity-client-id" in self._managed_identity_client_ids(
AzureCredentialType.DeploymentIdentityCredential
)
class TestGetAzureAdTokenProvider:
@patch.dict(
os.environ,

View file

@ -25501,6 +25501,12 @@ export interface components {
* @description Documents all the fields supported by `general_settings` in config.yaml
*/
ConfigGeneralSettings: {
/**
* Admission Queue Timeout Seconds
* @description maximum time a request waits for a worker slot
* @default 1
*/
admission_queue_timeout_seconds: number;
/**
* Alert To Webhook Url
* @description Mapping of alert type to webhook url. e.g. `alert_to_webhook_url: {'budget_alerts': 'https://nothooks.slack.com/services/T00000000/B00000000/XXXXXXXXXXXXXXXXXXXXXXXX'}`
@ -25709,11 +25715,21 @@ export interface components {
* @description max file size in MB for /v1/files uploads, for any purpose, if a file is larger than this size it will be rejected before being forwarded to the provider
*/
max_file_size_mb?: number | null;
/**
* Max In Flight Requests Per Worker
* @description maximum concurrent requests handled by each worker
*/
max_in_flight_requests_per_worker?: number | null;
/**
* Max Parallel Requests
* @description maximum parallel requests for each api key
*/
max_parallel_requests?: number | null;
/**
* Max Queued Requests Per Worker
* @description maximum requests waiting for a worker slot
*/
max_queued_requests_per_worker?: number | null;
/**
* Max Request Size Mb
* @description max request size in MB, if a request is larger than this size it will be rejected