mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge branch 'litellm_internal_staging' of https://github.com/BerriAI/litellm into litellm_govcloud_profiles_lit6421
This commit is contained in:
commit
1763052ff5
33 changed files with 2776 additions and 359 deletions
5
.github/ci-coverage-allowlist.yml
vendored
5
.github/ci-coverage-allowlist.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
|
|
|
|||
315
litellm/proxy/middleware/admission_control_middleware.py
Normal file
315
litellm/proxy/middleware/admission_control_middleware.py
Normal 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",
|
||||
}
|
||||
},
|
||||
)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -6,3 +6,4 @@ class AzureCredentialType(str, Enum):
|
|||
ManagedIdentityCredential = "ManagedIdentityCredential"
|
||||
CertificateCredential = "CertificateCredential"
|
||||
DefaultAzureCredential = "DefaultAzureCredential"
|
||||
DeploymentIdentityCredential = "DeploymentIdentityCredential"
|
||||
|
|
|
|||
150
tests/load_tests/test_granian_admission_saturation.py
Normal file
150
tests/load_tests/test_granian_admission_saturation.py
Normal 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()
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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"')
|
||||
|
|
|
|||
|
|
@ -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'
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
16
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
16
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue