mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
Merge pull request #21085 from BerriAI/litellm_oss_staging_02_13_2026
Litellm oss staging 02 13 2026
This commit is contained in:
commit
e17c639fb1
21 changed files with 1406 additions and 365 deletions
|
|
@ -19,6 +19,7 @@
|
|||
"mcp-client-2025-11-20": "mcp-client-2025-11-20",
|
||||
"mcp-client-2025-04-04": "mcp-client-2025-04-04",
|
||||
"mcp-servers-2025-12-04": "mcp-servers-2025-12-04",
|
||||
"oauth-2025-04-20": "oauth-2025-04-20",
|
||||
"output-128k-2025-02-19": "output-128k-2025-02-19",
|
||||
"prompt-caching-scope-2026-01-05": "prompt-caching-scope-2026-01-05",
|
||||
"skills-2025-10-02": "skills-2025-10-02",
|
||||
|
|
|
|||
|
|
@ -29,6 +29,7 @@ from litellm.types.utils import (
|
|||
LLMResponseTypes,
|
||||
StandardLoggingGuardrailInformation,
|
||||
)
|
||||
from fastapi.exceptions import HTTPException
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
|
@ -650,6 +651,23 @@ class CustomGuardrail(CustomLogger):
|
|||
)
|
||||
return response
|
||||
|
||||
@staticmethod
|
||||
def _is_guardrail_intervention(e: Exception) -> bool:
|
||||
"""
|
||||
Returns True if the exception represents an intentional guardrail block
|
||||
(this was logged previously as an API failure - guardrail_failed_to_respond).
|
||||
|
||||
Guardrails signal intentional blocks by raising:
|
||||
- HTTPException with status 400 (content policy violation)
|
||||
- ModifyResponseException (passthrough mode violation)
|
||||
"""
|
||||
|
||||
if isinstance(e, ModifyResponseException):
|
||||
return True
|
||||
if isinstance(e, HTTPException) and e.status_code == 400:
|
||||
return True
|
||||
return False
|
||||
|
||||
def _process_error(
|
||||
self,
|
||||
e: Exception,
|
||||
|
|
@ -664,6 +682,11 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
This gets logged on downsteam Langfuse, DataDog, etc.
|
||||
"""
|
||||
guardrail_status: GuardrailStatus = (
|
||||
"guardrail_intervened"
|
||||
if self._is_guardrail_intervention(e)
|
||||
else "guardrail_failed_to_respond"
|
||||
)
|
||||
# For custom_code_guardrail scenario, log as "deny" instead of full exception
|
||||
# Check if this is from custom_code_guardrail by checking the class name
|
||||
guardrail_response: Union[Exception, str] = e
|
||||
|
|
@ -673,7 +696,7 @@ class CustomGuardrail(CustomLogger):
|
|||
self.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_json_response=guardrail_response,
|
||||
request_data=request_data,
|
||||
guardrail_status="guardrail_failed_to_respond",
|
||||
guardrail_status=guardrail_status,
|
||||
duration=duration,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
|
|
|
|||
|
|
@ -2331,7 +2331,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
result, LiteLLMBatch
|
||||
):
|
||||
litellm_params = self.litellm_params or {}
|
||||
litellm_metadata = litellm_params.get("litellm_metadata", {})
|
||||
litellm_metadata = litellm_params.get("litellm_metadata") or {}
|
||||
if (
|
||||
litellm_metadata.get("batch_ignore_default_logging", False) is True
|
||||
): # polling job will query these frequently, don't spam db logs
|
||||
|
|
|
|||
|
|
@ -38,9 +38,18 @@ def optionally_handle_anthropic_oauth(
|
|||
Returns:
|
||||
Tuple of (updated headers, api_key)
|
||||
"""
|
||||
# Check Authorization header (passthrough / forwarded requests)
|
||||
auth_header = headers.get("authorization", "")
|
||||
if auth_header and auth_header.startswith(f"Bearer {ANTHROPIC_OAUTH_TOKEN_PREFIX}"):
|
||||
api_key = auth_header.replace("Bearer ", "")
|
||||
headers.pop("x-api-key", None)
|
||||
headers["anthropic-beta"] = ANTHROPIC_OAUTH_BETA_HEADER
|
||||
headers["anthropic-dangerous-direct-browser-access"] = "true"
|
||||
return headers, api_key
|
||||
# Check api_key directly (standard chat/completion flow)
|
||||
if api_key and api_key.startswith(ANTHROPIC_OAUTH_TOKEN_PREFIX):
|
||||
headers.pop("x-api-key", None)
|
||||
headers["authorization"] = f"Bearer {api_key}"
|
||||
headers["anthropic-beta"] = ANTHROPIC_OAUTH_BETA_HEADER
|
||||
headers["anthropic-dangerous-direct-browser-access"] = "true"
|
||||
return headers, api_key
|
||||
|
|
@ -108,7 +117,9 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
if tools is None:
|
||||
return False
|
||||
for tool in tools:
|
||||
if "type" in tool and tool["type"].startswith(ANTHROPIC_HOSTED_TOOLS.WEB_SEARCH.value):
|
||||
if "type" in tool and tool["type"].startswith(
|
||||
ANTHROPIC_HOSTED_TOOLS.WEB_SEARCH.value
|
||||
):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
|
@ -134,111 +145,126 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
"""
|
||||
if not tools:
|
||||
return False
|
||||
|
||||
|
||||
for tool in tools:
|
||||
tool_type = tool.get("type", "")
|
||||
if tool_type in ["tool_search_tool_regex_20251119", "tool_search_tool_bm25_20251119"]:
|
||||
if tool_type in [
|
||||
"tool_search_tool_regex_20251119",
|
||||
"tool_search_tool_bm25_20251119",
|
||||
]:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def is_programmatic_tool_calling_used(self, tools: Optional[List]) -> bool:
|
||||
"""
|
||||
Check if programmatic tool calling is being used (tools with allowed_callers field).
|
||||
|
||||
|
||||
Returns True if any tool has allowed_callers containing 'code_execution_20250825'.
|
||||
"""
|
||||
if not tools:
|
||||
return False
|
||||
|
||||
|
||||
for tool in tools:
|
||||
# Check top-level allowed_callers
|
||||
allowed_callers = tool.get("allowed_callers", None)
|
||||
if allowed_callers and isinstance(allowed_callers, list):
|
||||
if "code_execution_20250825" in allowed_callers:
|
||||
return True
|
||||
|
||||
|
||||
# Check function.allowed_callers for OpenAI format tools
|
||||
function = tool.get("function", {})
|
||||
if isinstance(function, dict):
|
||||
function_allowed_callers = function.get("allowed_callers", None)
|
||||
if function_allowed_callers and isinstance(function_allowed_callers, list):
|
||||
if function_allowed_callers and isinstance(
|
||||
function_allowed_callers, list
|
||||
):
|
||||
if "code_execution_20250825" in function_allowed_callers:
|
||||
return True
|
||||
|
||||
|
||||
return False
|
||||
|
||||
|
||||
def is_input_examples_used(self, tools: Optional[List]) -> bool:
|
||||
"""
|
||||
Check if input_examples is being used in any tools.
|
||||
|
||||
|
||||
Returns True if any tool has input_examples field.
|
||||
"""
|
||||
if not tools:
|
||||
return False
|
||||
|
||||
|
||||
for tool in tools:
|
||||
# Check top-level input_examples
|
||||
input_examples = tool.get("input_examples", None)
|
||||
if input_examples and isinstance(input_examples, list) and len(input_examples) > 0:
|
||||
if (
|
||||
input_examples
|
||||
and isinstance(input_examples, list)
|
||||
and len(input_examples) > 0
|
||||
):
|
||||
return True
|
||||
|
||||
|
||||
# Check function.input_examples for OpenAI format tools
|
||||
function = tool.get("function", {})
|
||||
if isinstance(function, dict):
|
||||
function_input_examples = function.get("input_examples", None)
|
||||
if function_input_examples and isinstance(function_input_examples, list) and len(function_input_examples) > 0:
|
||||
if (
|
||||
function_input_examples
|
||||
and isinstance(function_input_examples, list)
|
||||
and len(function_input_examples) > 0
|
||||
):
|
||||
return True
|
||||
|
||||
|
||||
return False
|
||||
|
||||
def is_effort_used(self, optional_params: Optional[dict], model: Optional[str] = None) -> bool:
|
||||
|
||||
def is_effort_used(
|
||||
self, optional_params: Optional[dict], model: Optional[str] = None
|
||||
) -> bool:
|
||||
"""
|
||||
Check if effort parameter is being used.
|
||||
|
||||
|
||||
Returns True if effort-related parameters are present.
|
||||
"""
|
||||
if not optional_params:
|
||||
return False
|
||||
|
||||
|
||||
# Check if reasoning_effort is provided for Claude Opus 4.5
|
||||
if model and ("opus-4-5" in model.lower() or "opus_4_5" in model.lower()):
|
||||
reasoning_effort = optional_params.get("reasoning_effort")
|
||||
if reasoning_effort and isinstance(reasoning_effort, str):
|
||||
return True
|
||||
|
||||
|
||||
# Check if output_config is directly provided
|
||||
output_config = optional_params.get("output_config")
|
||||
if output_config and isinstance(output_config, dict):
|
||||
effort = output_config.get("effort")
|
||||
if effort and isinstance(effort, str):
|
||||
return True
|
||||
|
||||
|
||||
return False
|
||||
|
||||
def is_code_execution_tool_used(self, tools: Optional[List]) -> bool:
|
||||
"""
|
||||
Check if code execution tool is being used.
|
||||
|
||||
|
||||
Returns True if any tool has type "code_execution_20250825".
|
||||
"""
|
||||
if not tools:
|
||||
return False
|
||||
|
||||
|
||||
for tool in tools:
|
||||
tool_type = tool.get("type", "")
|
||||
if tool_type == "code_execution_20250825":
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def is_container_with_skills_used(self, optional_params: Optional[dict]) -> bool:
|
||||
"""
|
||||
Check if container with skills is being used.
|
||||
|
||||
|
||||
Returns True if optional_params contains container with skills.
|
||||
"""
|
||||
if not optional_params:
|
||||
return False
|
||||
|
||||
|
||||
container = optional_params.get("container")
|
||||
if container and isinstance(container, dict):
|
||||
skills = container.get("skills")
|
||||
|
|
@ -256,10 +282,10 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
def get_computer_tool_beta_header(self, computer_tool_version: str) -> str:
|
||||
"""
|
||||
Get the appropriate beta header for a given computer tool version.
|
||||
|
||||
|
||||
Args:
|
||||
computer_tool_version: The computer tool version (e.g., 'computer_20250124', 'computer_20241022')
|
||||
|
||||
|
||||
Returns:
|
||||
The corresponding beta header string
|
||||
"""
|
||||
|
|
@ -282,37 +308,37 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
) -> List[str]:
|
||||
"""
|
||||
Get list of common beta headers based on the features that are active.
|
||||
|
||||
|
||||
Returns:
|
||||
List of beta header strings
|
||||
"""
|
||||
from litellm.types.llms.anthropic import (
|
||||
ANTHROPIC_EFFORT_BETA_HEADER,
|
||||
)
|
||||
|
||||
|
||||
betas = []
|
||||
|
||||
|
||||
# Detect features
|
||||
effort_used = self.is_effort_used(optional_params, model)
|
||||
|
||||
|
||||
if effort_used:
|
||||
betas.append(ANTHROPIC_EFFORT_BETA_HEADER) # effort-2025-11-24
|
||||
|
||||
|
||||
if computer_tool_used:
|
||||
beta_header = self.get_computer_tool_beta_header(computer_tool_used)
|
||||
betas.append(beta_header)
|
||||
|
||||
|
||||
# Anthropic no longer requires the prompt-caching beta header
|
||||
# Prompt caching now works automatically when cache_control is used in messages
|
||||
# Reference: https://docs.anthropic.com/en/docs/build-with-claude/prompt-caching
|
||||
|
||||
|
||||
if file_id_used:
|
||||
betas.append("files-api-2025-04-14")
|
||||
betas.append("code-execution-2025-05-22")
|
||||
|
||||
|
||||
if mcp_server_used:
|
||||
betas.append("mcp-client-2025-04-04")
|
||||
|
||||
|
||||
return list(set(betas))
|
||||
|
||||
def get_anthropic_headers(
|
||||
|
|
@ -351,27 +377,35 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
# Tool search, programmatic tool calling, and input_examples all use the same beta header
|
||||
if tool_search_used or programmatic_tool_calling_used or input_examples_used:
|
||||
from litellm.types.llms.anthropic import ANTHROPIC_TOOL_SEARCH_BETA_HEADER
|
||||
|
||||
betas.add(ANTHROPIC_TOOL_SEARCH_BETA_HEADER)
|
||||
|
||||
|
||||
# Effort parameter uses a separate beta header
|
||||
if effort_used:
|
||||
from litellm.types.llms.anthropic import ANTHROPIC_EFFORT_BETA_HEADER
|
||||
|
||||
betas.add(ANTHROPIC_EFFORT_BETA_HEADER)
|
||||
|
||||
|
||||
# Code execution tool uses a separate beta header
|
||||
if code_execution_tool_used:
|
||||
betas.add("code-execution-2025-08-25")
|
||||
|
||||
|
||||
# Container with skills uses a separate beta header
|
||||
if container_with_skills_used:
|
||||
betas.add("skills-2025-10-02")
|
||||
|
||||
_is_oauth = api_key and api_key.startswith(ANTHROPIC_OAUTH_TOKEN_PREFIX)
|
||||
headers = {
|
||||
"anthropic-version": anthropic_version or "2023-06-01",
|
||||
"x-api-key": api_key,
|
||||
"accept": "application/json",
|
||||
"content-type": "application/json",
|
||||
}
|
||||
if _is_oauth:
|
||||
headers["authorization"] = f"Bearer {api_key}"
|
||||
headers["anthropic-dangerous-direct-browser-access"] = "true"
|
||||
betas.add(ANTHROPIC_OAUTH_BETA_HEADER)
|
||||
else:
|
||||
headers["x-api-key"] = api_key
|
||||
|
||||
if user_anthropic_beta_headers is not None:
|
||||
betas.update(user_anthropic_beta_headers)
|
||||
|
|
@ -381,7 +415,10 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
# Vertex AI requires web search beta header for web search to work
|
||||
if web_search_tool_used:
|
||||
from litellm.types.llms.anthropic import ANTHROPIC_BETA_HEADER_VALUES
|
||||
headers["anthropic-beta"] = ANTHROPIC_BETA_HEADER_VALUES.WEB_SEARCH_2025_03_05.value
|
||||
|
||||
headers[
|
||||
"anthropic-beta"
|
||||
] = ANTHROPIC_BETA_HEADER_VALUES.WEB_SEARCH_2025_03_05.value
|
||||
elif len(betas) > 0:
|
||||
headers["anthropic-beta"] = ",".join(betas)
|
||||
|
||||
|
|
@ -398,7 +435,9 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
api_base: Optional[str] = None,
|
||||
) -> Dict:
|
||||
# Check for Anthropic OAuth token in headers
|
||||
headers, api_key = optionally_handle_anthropic_oauth(headers=headers, api_key=api_key)
|
||||
headers, api_key = optionally_handle_anthropic_oauth(
|
||||
headers=headers, api_key=api_key
|
||||
)
|
||||
if api_key is None:
|
||||
raise litellm.AuthenticationError(
|
||||
message="Missing Anthropic API Key - A call is being made to anthropic but no key is set either in the environment variables or via params. Please set `ANTHROPIC_API_KEY` in your environment vars",
|
||||
|
|
@ -416,11 +455,15 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
file_id_used = self.is_file_id_used(messages=messages)
|
||||
web_search_tool_used = self.is_web_search_tool_used(tools=tools)
|
||||
tool_search_used = self.is_tool_search_used(tools=tools)
|
||||
programmatic_tool_calling_used = self.is_programmatic_tool_calling_used(tools=tools)
|
||||
programmatic_tool_calling_used = self.is_programmatic_tool_calling_used(
|
||||
tools=tools
|
||||
)
|
||||
input_examples_used = self.is_input_examples_used(tools=tools)
|
||||
effort_used = self.is_effort_used(optional_params=optional_params, model=model)
|
||||
code_execution_tool_used = self.is_code_execution_tool_used(tools=tools)
|
||||
container_with_skills_used = self.is_container_with_skills_used(optional_params=optional_params)
|
||||
container_with_skills_used = self.is_container_with_skills_used(
|
||||
optional_params=optional_params
|
||||
)
|
||||
user_anthropic_beta_headers = self._get_user_anthropic_beta_headers(
|
||||
anthropic_beta_header=headers.get("anthropic-beta")
|
||||
)
|
||||
|
|
@ -499,7 +542,7 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
def get_token_counter(self) -> Optional[BaseTokenCounter]:
|
||||
"""
|
||||
Factory method to create an Anthropic token counter.
|
||||
|
||||
|
||||
Returns:
|
||||
AnthropicTokenCounter instance for this provider.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -49,15 +49,15 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
# TODO: Add Anthropic `metadata` support
|
||||
# "metadata",
|
||||
]
|
||||
|
||||
|
||||
@staticmethod
|
||||
def _filter_billing_headers_from_system(system_param):
|
||||
"""
|
||||
Filter out x-anthropic-billing-header metadata from system parameter.
|
||||
|
||||
|
||||
Args:
|
||||
system_param: Can be a string or a list of system message content blocks
|
||||
|
||||
|
||||
Returns:
|
||||
Filtered system parameter (string or list), or None if all content was filtered
|
||||
"""
|
||||
|
|
@ -74,7 +74,9 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
text = content_block.get("text", "")
|
||||
content_type = content_block.get("type", "")
|
||||
# Skip text blocks that start with billing header
|
||||
if content_type == "text" and text.startswith("x-anthropic-billing-header:"):
|
||||
if content_type == "text" and text.startswith(
|
||||
"x-anthropic-billing-header:"
|
||||
):
|
||||
continue
|
||||
filtered_list.append(content_block)
|
||||
else:
|
||||
|
|
@ -111,11 +113,13 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
import os
|
||||
|
||||
# Check for Anthropic OAuth token in Authorization header
|
||||
headers, api_key = optionally_handle_anthropic_oauth(headers=headers, api_key=api_key)
|
||||
headers, api_key = optionally_handle_anthropic_oauth(
|
||||
headers=headers, api_key=api_key
|
||||
)
|
||||
if api_key is None:
|
||||
api_key = os.getenv("ANTHROPIC_API_KEY")
|
||||
|
||||
if "x-api-key" not in headers and api_key:
|
||||
if "x-api-key" not in headers and "authorization" not in headers and api_key:
|
||||
headers["x-api-key"] = api_key
|
||||
if "anthropic-version" not in headers:
|
||||
headers["anthropic-version"] = DEFAULT_ANTHROPIC_API_VERSION
|
||||
|
|
@ -149,7 +153,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
message="max_tokens is required for Anthropic /v1/messages API",
|
||||
status_code=400,
|
||||
)
|
||||
|
||||
|
||||
# Filter out x-anthropic-billing-header from system messages
|
||||
system_param = anthropic_messages_optional_request_params.get("system")
|
||||
if system_param is not None:
|
||||
|
|
@ -159,7 +163,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
else:
|
||||
# Remove system parameter if all content was filtered out
|
||||
anthropic_messages_optional_request_params.pop("system", None)
|
||||
|
||||
|
||||
####### get required params for all anthropic messages requests ######
|
||||
verbose_logger.debug(f"TRANSFORMATION DEBUG - Messages: {messages}")
|
||||
anthropic_messages_request: AnthropicMessagesRequest = AnthropicMessagesRequest(
|
||||
|
|
@ -244,25 +248,29 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
edits = context_management_param.get("edits", [])
|
||||
has_compact = False
|
||||
has_other = False
|
||||
|
||||
|
||||
for edit in edits:
|
||||
edit_type = edit.get("type", "")
|
||||
if edit_type == "compact_20260112":
|
||||
has_compact = True
|
||||
else:
|
||||
has_other = True
|
||||
|
||||
|
||||
# Add compact header if any compact edits exist
|
||||
if has_compact:
|
||||
beta_values.add(ANTHROPIC_BETA_HEADER_VALUES.COMPACT_2026_01_12.value)
|
||||
|
||||
|
||||
# Add context management header if any other edits exist
|
||||
if has_other:
|
||||
beta_values.add(ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value)
|
||||
beta_values.add(
|
||||
ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value
|
||||
)
|
||||
|
||||
# Check for structured outputs
|
||||
if optional_params.get("output_format") is not None:
|
||||
beta_values.add(ANTHROPIC_BETA_HEADER_VALUES.STRUCTURED_OUTPUT_2025_09_25.value)
|
||||
beta_values.add(
|
||||
ANTHROPIC_BETA_HEADER_VALUES.STRUCTURED_OUTPUT_2025_09_25.value
|
||||
)
|
||||
|
||||
# Check for fast mode
|
||||
if optional_params.get("speed") == "fast":
|
||||
|
|
|
|||
|
|
@ -480,7 +480,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
tool = {VertexToolName.COMPUTER_USE.value: computer_use_config}
|
||||
# Handle OpenAI-style web_search and web_search_preview tools
|
||||
# Transform them to Gemini's googleSearch tool
|
||||
elif "type" in tool and tool["type"] in ("web_search", "web_search_preview"):
|
||||
elif "type" in tool and tool["type"] in (
|
||||
"web_search",
|
||||
"web_search_preview",
|
||||
):
|
||||
verbose_logger.info(
|
||||
f"Gemini: Transforming OpenAI-style '{tool['type']}' tool to googleSearch"
|
||||
)
|
||||
|
|
@ -1196,6 +1199,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
"PROHIBITED_CONTENT": "The token generation was stopped as the response was flagged for the prohibited contents.",
|
||||
"SPII": "The token generation was stopped as the response was flagged for Sensitive Personally Identifiable Information (SPII) contents.",
|
||||
"IMAGE_SAFETY": "The token generation was stopped as the response was flagged for image safety reasons.",
|
||||
"IMAGE_PROHIBITED_CONTENT": "The token generation was stopped as the response was flagged for prohibited image content.",
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -1218,6 +1222,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
"SPII": "content_filter",
|
||||
"MALFORMED_FUNCTION_CALL": "malformed_function_call", # openai doesn't have a way of representing this
|
||||
"IMAGE_SAFETY": "content_filter",
|
||||
"IMAGE_PROHIBITED_CONTENT": "content_filter",
|
||||
}
|
||||
|
||||
def translate_exception_str(self, exception_string: str):
|
||||
|
|
@ -1630,7 +1635,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
completion_image_tokens = response_tokens_details.image_tokens or 0
|
||||
completion_audio_tokens = response_tokens_details.audio_tokens or 0
|
||||
calculated_text_tokens = (
|
||||
candidates_token_count - completion_image_tokens - completion_audio_tokens
|
||||
candidates_token_count
|
||||
- completion_image_tokens
|
||||
- completion_audio_tokens
|
||||
)
|
||||
response_tokens_details.text_tokens = calculated_text_tokens
|
||||
#########################################################
|
||||
|
|
@ -2248,6 +2255,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
citation_metadata # older approach - maintaining to prevent regressions
|
||||
)
|
||||
|
||||
## ADD TRAFFIC TYPE ##
|
||||
traffic_type = completion_response.get("usageMetadata", {}).get(
|
||||
"trafficType"
|
||||
)
|
||||
if traffic_type:
|
||||
model_response._hidden_params.setdefault("provider_specific_fields", {})["traffic_type"] = traffic_type
|
||||
|
||||
except Exception as e:
|
||||
raise VertexAIError(
|
||||
message="Received={}, Error converting to valid response block={}. File an issue if litellm error - https://github.com/BerriAI/litellm/issues".format(
|
||||
|
|
@ -2906,6 +2920,12 @@ class ModelResponseIterator:
|
|||
PromptTokensDetailsWrapper, usage.prompt_tokens_details
|
||||
).web_search_requests = web_search_requests
|
||||
|
||||
traffic_type = processed_chunk.get("usageMetadata", {}).get(
|
||||
"trafficType"
|
||||
)
|
||||
if traffic_type:
|
||||
model_response._hidden_params.setdefault("provider_specific_fields", {})["traffic_type"] = traffic_type
|
||||
|
||||
setattr(model_response, "usage", usage) # type: ignore
|
||||
|
||||
model_response._hidden_params["is_finished"] = False
|
||||
|
|
|
|||
|
|
@ -7383,6 +7383,16 @@ def stream_chunk_builder( # noqa: PLR0915
|
|||
|
||||
setattr(response, "usage", usage)
|
||||
|
||||
# Propagate provider_specific_fields from the last chunk (contains provider
|
||||
# metadata like traffic_type set during streaming)
|
||||
for chunk in reversed(chunks):
|
||||
hidden = getattr(chunk, "_hidden_params", None)
|
||||
if hidden and "provider_specific_fields" in hidden:
|
||||
response._hidden_params.setdefault(
|
||||
"provider_specific_fields", {}
|
||||
).update(hidden["provider_specific_fields"])
|
||||
break
|
||||
|
||||
# Add cost to usage object if include_cost_in_streaming_usage is True
|
||||
if litellm.include_cost_in_streaming_usage and logging_obj is not None:
|
||||
setattr(
|
||||
|
|
|
|||
|
|
@ -5,14 +5,9 @@ OpenAI Moderation Guardrail Integration for LiteLLM
|
|||
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
AsyncGenerator,
|
||||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
Optional,
|
||||
Type,
|
||||
Union,
|
||||
)
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -20,7 +15,7 @@ from fastapi import HTTPException
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
log_guardrail_information
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
|
|
@ -32,10 +27,8 @@ from litellm.types.utils import GenericGuardrailAPIInputs
|
|||
from .base import OpenAIGuardrailBase
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.llms.openai import OpenAIModerationResponse
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
from litellm.types.utils import ModelResponse, ModelResponseStream
|
||||
|
||||
|
||||
class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
|
||||
|
|
@ -236,108 +229,6 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
|
|||
# Moderation doesn't modify content, just blocks - return inputs unchanged
|
||||
return inputs
|
||||
|
||||
@log_guardrail_information
|
||||
async def async_post_call_streaming_iterator_hook(
|
||||
self,
|
||||
user_api_key_dict: "UserAPIKeyAuth",
|
||||
response: Any,
|
||||
request_data: Dict[str, Any],
|
||||
) -> AsyncGenerator["ModelResponseStream", None]:
|
||||
"""
|
||||
Process streaming response chunks for OpenAI moderation.
|
||||
|
||||
Collects all chunks from the stream, assembles them into a complete response,
|
||||
and applies moderation check. If content violates moderation policy, raises HTTPException.
|
||||
"""
|
||||
# Import here to avoid circular imports
|
||||
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
|
||||
from litellm.main import stream_chunk_builder
|
||||
from litellm.types.utils import TextCompletionResponse
|
||||
|
||||
verbose_proxy_logger.debug("OpenAI Moderation: Running streaming response scan")
|
||||
|
||||
# Collect all chunks to process them together
|
||||
all_chunks: List["ModelResponseStream"] = []
|
||||
async for chunk in response:
|
||||
all_chunks.append(chunk)
|
||||
|
||||
# Assemble the complete response from chunks
|
||||
assembled_model_response: Optional[
|
||||
Union["ModelResponse", TextCompletionResponse]
|
||||
] = stream_chunk_builder(
|
||||
chunks=all_chunks,
|
||||
)
|
||||
|
||||
if isinstance(assembled_model_response, (type(None), TextCompletionResponse)):
|
||||
# If we can't assemble a ModelResponse or it's a text completion,
|
||||
# just yield the original chunks without moderation
|
||||
verbose_proxy_logger.warning(
|
||||
"OpenAI Moderation: Could not assemble ModelResponse from chunks, skipping moderation"
|
||||
)
|
||||
for chunk in all_chunks:
|
||||
yield chunk
|
||||
return
|
||||
|
||||
# Extract response text for moderation
|
||||
response_text = self._extract_response_text(assembled_model_response)
|
||||
if response_text:
|
||||
verbose_proxy_logger.debug(
|
||||
f"OpenAI Moderation: Streaming response text: {response_text[:100]}..." # Log first 100 chars
|
||||
)
|
||||
|
||||
# Make moderation request - this will raise HTTPException if content is flagged
|
||||
moderation_response = await self.async_make_request(
|
||||
input_text=response_text,
|
||||
)
|
||||
|
||||
# Check if content is flagged and raise exception if needed
|
||||
self._check_moderation_result(moderation_response)
|
||||
|
||||
# If we reach here, content passed moderation - yield the original chunks
|
||||
mock_response = MockResponseIterator(model_response=assembled_model_response)
|
||||
|
||||
# Return the reconstructed stream
|
||||
async for chunk in mock_response:
|
||||
yield chunk
|
||||
|
||||
def _extract_response_text(self, response: "ModelResponse") -> Optional[str]:
|
||||
"""
|
||||
Extract text content from the model response for moderation.
|
||||
"""
|
||||
if not hasattr(response, "choices") or not response.choices:
|
||||
return None
|
||||
|
||||
response_texts = []
|
||||
for choice in response.choices:
|
||||
try:
|
||||
# Try to get content from message (chat completion)
|
||||
message = getattr(choice, "message", None)
|
||||
if message:
|
||||
content = getattr(message, "content", None)
|
||||
if content and isinstance(content, str):
|
||||
response_texts.append(content)
|
||||
continue
|
||||
|
||||
# Try to get text (text completion)
|
||||
text = getattr(choice, "text", None)
|
||||
if text and isinstance(text, str):
|
||||
response_texts.append(text)
|
||||
continue
|
||||
|
||||
# Try to get content from delta (streaming)
|
||||
delta = getattr(choice, "delta", None)
|
||||
if delta:
|
||||
content = getattr(delta, "content", None)
|
||||
if content and isinstance(content, str):
|
||||
response_texts.append(content)
|
||||
continue
|
||||
|
||||
except (AttributeError, TypeError):
|
||||
# Skip choices that don't have expected attributes
|
||||
continue
|
||||
|
||||
return "\n".join(response_texts) if response_texts else None
|
||||
|
||||
@staticmethod
|
||||
def get_config_model() -> Optional[Type["GuardrailConfigModel"]]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ from litellm.types.utils import GenericGuardrailAPIInputs
|
|||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
|
||||
GUARDRAIL_TIMEOUT = 5
|
||||
|
||||
|
|
@ -334,3 +335,11 @@ class ZscalerAIGuard(CustomGuardrail):
|
|||
user_facing_error = self._create_user_facing_error(f"{str(e)})")
|
||||
# This exception will be caught by the proxy and returned to the user
|
||||
raise HTTPException(status_code=500, detail=user_facing_error)
|
||||
|
||||
@staticmethod
|
||||
def get_config_model() -> Optional[type["GuardrailConfigModel"]]:
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.zscaler_ai_guard import (
|
||||
ZscalerAIGuardConfigModel,
|
||||
)
|
||||
|
||||
return ZscalerAIGuardConfigModel
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from litellm._uuid import uuid
|
|||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.proxy.types_utils.utils import get_instance_fn
|
||||
from litellm.secret_managers.main import get_secret
|
||||
from litellm.types.guardrails import (
|
||||
Guardrail,
|
||||
|
|
@ -489,7 +490,7 @@ class InMemoryGuardrailHandler:
|
|||
config_file_path: Optional[str] = None,
|
||||
) -> Optional[CustomGuardrail]:
|
||||
"""
|
||||
Initialize a Custom Guardrail from a python file
|
||||
Initialize a Custom Guardrail from a python file or module path
|
||||
|
||||
This initializes it by adding it to the litellm callback manager
|
||||
"""
|
||||
|
|
@ -498,26 +499,12 @@ class InMemoryGuardrailHandler:
|
|||
"GuardrailsAIException - Please pass the config_file_path to initialize_guardrails_v2"
|
||||
)
|
||||
|
||||
_file_name, _class_name = guardrail_type.split(".")
|
||||
verbose_proxy_logger.debug(
|
||||
"Initializing custom guardrail: %s, file_name: %s, class_name: %s",
|
||||
"Initializing custom guardrail: %s",
|
||||
guardrail_type,
|
||||
_file_name,
|
||||
_class_name,
|
||||
)
|
||||
|
||||
directory = os.path.dirname(config_file_path)
|
||||
module_file_path = os.path.join(directory, _file_name) + ".py"
|
||||
|
||||
spec = importlib.util.spec_from_file_location(_class_name, module_file_path) # type: ignore
|
||||
if not spec:
|
||||
raise ImportError(
|
||||
f"Could not find a module specification for {module_file_path}"
|
||||
)
|
||||
|
||||
module = importlib.util.module_from_spec(spec) # type: ignore
|
||||
spec.loader.exec_module(module) # type: ignore
|
||||
_guardrail_class = getattr(module, _class_name)
|
||||
_guardrail_class = get_instance_fn(guardrail_type, config_file_path=config_file_path)
|
||||
|
||||
mode = litellm_params.mode
|
||||
if mode is None:
|
||||
|
|
|
|||
|
|
@ -1512,6 +1512,12 @@ class LiteLLMCompletionResponsesConfig:
|
|||
user=getattr(chat_completion_response, "user", None),
|
||||
)
|
||||
responses_api_response._hidden_params = getattr(chat_completion_response, "_hidden_params", {})
|
||||
|
||||
# Surface provider-specific fields (generic passthrough from any provider)
|
||||
provider_fields = responses_api_response._hidden_params.get("provider_specific_fields")
|
||||
if provider_fields:
|
||||
responses_api_response.provider_specific_fields = provider_fields
|
||||
|
||||
return responses_api_response
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -41,6 +41,53 @@ from litellm.types.utils import GenericBudgetConfigType, StandardLoggingPayload
|
|||
DEFAULT_REDIS_SYNC_INTERVAL = 1
|
||||
|
||||
|
||||
class _LiteLLMParamsDictView:
|
||||
"""
|
||||
Lightweight attribute view over `litellm_params` dict.
|
||||
|
||||
This avoids pydantic construction in request hot-path while preserving
|
||||
attribute-style access used by `litellm.get_llm_provider(...)`.
|
||||
"""
|
||||
|
||||
__slots__ = ("_params",)
|
||||
|
||||
def __init__(self, params: Dict[str, Any]):
|
||||
self._params = params
|
||||
|
||||
def __getattr__(self, key: str) -> Any:
|
||||
return self._params.get(key)
|
||||
|
||||
def __getitem__(self, key: str) -> Any:
|
||||
return self._params.get(key)
|
||||
|
||||
def __contains__(self, key: str) -> bool:
|
||||
return key in self._params
|
||||
|
||||
def get(self, key: str, default: Any = None) -> Any:
|
||||
return self._params.get(key, default)
|
||||
|
||||
def keys(self):
|
||||
return self._params.keys()
|
||||
|
||||
def values(self):
|
||||
return self._params.values()
|
||||
|
||||
def items(self):
|
||||
return self._params.items()
|
||||
|
||||
def __iter__(self):
|
||||
return iter(self._params)
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self._params)
|
||||
|
||||
def dict(self) -> Dict[str, Any]:
|
||||
return dict(self._params)
|
||||
|
||||
def model_dump(self) -> Dict[str, Any]:
|
||||
return dict(self._params)
|
||||
|
||||
|
||||
class RouterBudgetLimiting(CustomLogger):
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -98,6 +145,7 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
cache_keys,
|
||||
provider_configs,
|
||||
deployment_configs,
|
||||
deployment_providers,
|
||||
) = await self._async_get_cache_keys_for_router_budget_limiting(
|
||||
healthy_deployments=healthy_deployments,
|
||||
request_kwargs=request_kwargs,
|
||||
|
|
@ -123,6 +171,7 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
healthy_deployments=healthy_deployments,
|
||||
provider_configs=provider_configs,
|
||||
deployment_configs=deployment_configs,
|
||||
deployment_providers=deployment_providers,
|
||||
spend_map=spend_map,
|
||||
potential_deployments=potential_deployments,
|
||||
request_tags=_get_tags_from_request_kwargs(
|
||||
|
|
@ -145,6 +194,7 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
healthy_deployments: List[Dict[str, Any]],
|
||||
provider_configs: Dict[str, GenericBudgetInfo],
|
||||
deployment_configs: Dict[str, GenericBudgetInfo],
|
||||
deployment_providers: List[Optional[str]],
|
||||
spend_map: Dict[str, float],
|
||||
request_tags: List[str],
|
||||
) -> Tuple[List[Dict[str, Any]], str]:
|
||||
|
|
@ -161,12 +211,15 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
"""
|
||||
# Filter deployments based on both provider and deployment budgets
|
||||
deployment_above_budget_info: str = ""
|
||||
for deployment in healthy_deployments:
|
||||
for idx, deployment in enumerate(healthy_deployments):
|
||||
is_within_budget = True
|
||||
|
||||
# Check provider budget
|
||||
if self.provider_budget_config:
|
||||
provider = self._get_llm_provider_for_deployment(deployment)
|
||||
if idx < len(deployment_providers):
|
||||
provider = deployment_providers[idx]
|
||||
else:
|
||||
provider = self._get_llm_provider_for_deployment(deployment)
|
||||
if provider in provider_configs:
|
||||
config = provider_configs[provider]
|
||||
if config.max_budget is None:
|
||||
|
|
@ -230,24 +283,32 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
self,
|
||||
healthy_deployments: List[Dict[str, Any]],
|
||||
request_kwargs: Optional[Dict] = None,
|
||||
) -> Tuple[List[str], Dict[str, GenericBudgetInfo], Dict[str, GenericBudgetInfo]]:
|
||||
) -> Tuple[
|
||||
List[str],
|
||||
Dict[str, GenericBudgetInfo],
|
||||
Dict[str, GenericBudgetInfo],
|
||||
List[Optional[str]],
|
||||
]:
|
||||
"""
|
||||
Returns list of cache keys to fetch from router cache for budget limiting and provider and deployment configs
|
||||
|
||||
Returns:
|
||||
Tuple[List[str], Dict[str, GenericBudgetInfo], Dict[str, GenericBudgetInfo]]:
|
||||
Tuple[List[str], Dict[str, GenericBudgetInfo], Dict[str, GenericBudgetInfo], List[Optional[str]]]:
|
||||
- List of cache keys to fetch from router cache for budget limiting
|
||||
- Dict of provider budget configs `provider_configs`
|
||||
- Dict of deployment budget configs `deployment_configs`
|
||||
- List of resolved providers aligned by deployment index `deployment_providers`
|
||||
"""
|
||||
cache_keys: List[str] = []
|
||||
provider_configs: Dict[str, GenericBudgetInfo] = {}
|
||||
deployment_configs: Dict[str, GenericBudgetInfo] = {}
|
||||
deployment_providers: List[Optional[str]] = []
|
||||
|
||||
for deployment in healthy_deployments:
|
||||
# Check provider budgets
|
||||
if self.provider_budget_config:
|
||||
provider = self._get_llm_provider_for_deployment(deployment)
|
||||
deployment_providers.append(provider)
|
||||
if provider is not None:
|
||||
budget_config = self._get_budget_config_for_provider(provider)
|
||||
if (
|
||||
|
|
@ -280,7 +341,12 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
cache_keys.append(
|
||||
f"tag_spend:{_tag}:{_tag_budget_config.budget_duration}"
|
||||
)
|
||||
return cache_keys, provider_configs, deployment_configs
|
||||
return (
|
||||
cache_keys,
|
||||
provider_configs,
|
||||
deployment_configs,
|
||||
deployment_providers,
|
||||
)
|
||||
|
||||
async def _get_or_set_budget_start_time(
|
||||
self, start_time_key: str, current_time: float, ttl_seconds: int
|
||||
|
|
@ -597,12 +663,23 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
|
||||
def _get_llm_provider_for_deployment(self, deployment: Dict) -> Optional[str]:
|
||||
try:
|
||||
_litellm_params: LiteLLM_Params = LiteLLM_Params(
|
||||
**deployment.get("litellm_params", {"model": ""})
|
||||
)
|
||||
deployment_litellm_params = deployment.get("litellm_params") or {}
|
||||
|
||||
if isinstance(deployment_litellm_params, LiteLLM_Params):
|
||||
model = deployment_litellm_params.model or ""
|
||||
provider_resolution_params: Any = deployment_litellm_params
|
||||
elif isinstance(deployment_litellm_params, dict):
|
||||
model = deployment_litellm_params.get("model") or ""
|
||||
provider_resolution_params = _LiteLLMParamsDictView(
|
||||
deployment_litellm_params
|
||||
)
|
||||
else:
|
||||
model = ""
|
||||
provider_resolution_params = _LiteLLMParamsDictView({})
|
||||
|
||||
_, custom_llm_provider, _, _ = litellm.get_llm_provider(
|
||||
model=_litellm_params.model,
|
||||
litellm_params=_litellm_params,
|
||||
model=str(model),
|
||||
litellm_params=provider_resolution_params,
|
||||
)
|
||||
except Exception:
|
||||
verbose_router_logger.error(
|
||||
|
|
|
|||
|
|
@ -0,0 +1,132 @@
|
|||
from typing import Optional
|
||||
|
||||
from pydantic import Field, model_validator
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.types.guardrails import GuardrailParamUITypes
|
||||
|
||||
from .base import GuardrailConfigModel
|
||||
|
||||
|
||||
class ZscalerAIGuardConfigModel(GuardrailConfigModel):
|
||||
api_key: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"API key for Zscaler AI Guard authentication. "
|
||||
"If not provided, falls back to ZSCALER_AI_GUARD_API_KEY environment variable."
|
||||
),
|
||||
)
|
||||
|
||||
api_base: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Zscaler AI Guard API endpoint. Determines policy resolution behavior:\n"
|
||||
"• /execute-policy (default) - Requires explicit policy_id in configuration\n"
|
||||
"• /resolve-and-execute-policy - Infers policy from user-api-key-alias header\n"
|
||||
"Default: https://api.us1.zseclipse.net/v1/detection/execute-policy\n"
|
||||
"Falls back to ZSCALER_AI_GUARD_URL environment variable."
|
||||
),
|
||||
json_schema_extra={
|
||||
"examples": [
|
||||
"https://api.us1.zseclipse.net/v1/detection/execute-policy",
|
||||
"https://api.us1.zseclipse.net/v1/detection/resolve-and-execute-policy",
|
||||
]
|
||||
},
|
||||
)
|
||||
|
||||
policy_id: Optional[int] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Global policy ID for Zscaler AI Guard. Required when using /execute-policy endpoint.\n\n"
|
||||
"Set to 0 or leave empty when using /resolve-and-execute-policy with dynamic policy resolution.\n"
|
||||
"Falls back to ZSCALER_AI_GUARD_POLICY_ID environment variable."
|
||||
),
|
||||
json_schema_extra={
|
||||
"ui_hint": "conditional_required",
|
||||
"condition": "Required when api_base ends with /execute-policy",
|
||||
},
|
||||
)
|
||||
|
||||
send_user_api_key_alias: Optional[bool] = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"Send user API key alias in request headers as 'user-api-key-alias'. "
|
||||
"CRITICAL when using /resolve-and-execute-policy endpoint - the policy is inferred from this value. "
|
||||
"Also useful for tracking/auditing with /execute-policy endpoint."
|
||||
),
|
||||
json_schema_extra={
|
||||
"ui_type": GuardrailParamUITypes.BOOL,
|
||||
"ui_hint": "recommended_when",
|
||||
"condition": "Recommended when api_base ends with /resolve-and-execute-policy",
|
||||
},
|
||||
)
|
||||
|
||||
send_user_api_key_user_id: Optional[bool] = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"Send user API key user_id in request headers as 'user-api-key-user-id'. "
|
||||
"Enables user-level tracking and analytics in Zscaler AI Guard."
|
||||
),
|
||||
json_schema_extra={"ui_type": GuardrailParamUITypes.BOOL},
|
||||
)
|
||||
|
||||
send_user_api_key_team_id: Optional[bool] = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"Send user API key team_id in request headers as 'user-api-key-team-id'. "
|
||||
"Enables team-level tracking and analytics in Zscaler AI Guard."
|
||||
),
|
||||
json_schema_extra={"ui_type": GuardrailParamUITypes.BOOL},
|
||||
)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_endpoint_configuration(self) -> "ZscalerAIGuardConfigModel":
|
||||
"""
|
||||
Validate configuration consistency between api_base and other fields.
|
||||
Provides warnings but doesn't block (since env vars might provide values).
|
||||
"""
|
||||
import os
|
||||
|
||||
# Resolve actual api_base value (including env fallback)
|
||||
api_base = self.api_base or os.getenv(
|
||||
"ZSCALER_AI_GUARD_URL",
|
||||
"https://api.us1.zseclipse.net/v1/detection/execute-policy",
|
||||
)
|
||||
|
||||
# Resolve actual policy_id value
|
||||
policy_id = self.policy_id
|
||||
if policy_id is None:
|
||||
env_policy = os.getenv("ZSCALER_AI_GUARD_POLICY_ID")
|
||||
if env_policy:
|
||||
try:
|
||||
policy_id = int(env_policy)
|
||||
except ValueError:
|
||||
verbose_proxy_logger.warning(
|
||||
f"ZSCALER_AI_GUARD_POLICY_ID env var is not a valid integer: {env_policy}"
|
||||
)
|
||||
|
||||
# Check for configuration issues
|
||||
is_resolve_policy = api_base.endswith("/resolve-and-execute-policy")
|
||||
is_execute_policy = api_base.endswith("/execute-policy") and not is_resolve_policy
|
||||
|
||||
# Scenario A: execute-policy without policy_id
|
||||
if is_execute_policy and (policy_id is None or policy_id < 1):
|
||||
verbose_proxy_logger.warning(
|
||||
"Using /execute-policy endpoint without a valid policy_id. "
|
||||
"Ensure ZSCALER_AI_GUARD_POLICY_ID environment variable is set, "
|
||||
"or provide policy_id via request/key/team metadata."
|
||||
)
|
||||
|
||||
# Scenario B: resolve-and-execute-policy without user_api_key_alias
|
||||
if is_resolve_policy and not self.send_user_api_key_alias:
|
||||
verbose_proxy_logger.warning(
|
||||
"Using /resolve-and-execute-policy endpoint without send_user_api_key_alias=true. "
|
||||
"The endpoint requires user-api-key-alias header to resolve the policy. "
|
||||
"Set send_user_api_key_alias to true or ensure the header is sent via other means."
|
||||
)
|
||||
|
||||
return self
|
||||
|
||||
@staticmethod
|
||||
def ui_friendly_name() -> str:
|
||||
return "Zscaler AI Guard"
|
||||
|
|
@ -1,84 +1,285 @@
|
|||
"""
|
||||
Tests for Anthropic OAuth token handling for Claude Code Max integration.
|
||||
Tests for Anthropic OAuth token handling in common_utils.
|
||||
|
||||
Verifies that OAuth tokens (sk-ant-oat*) are sent via Authorization: Bearer
|
||||
instead of x-api-key, per Anthropic's OAuth specification.
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
|
||||
# Add litellm to path
|
||||
sys.path.insert(0, os.path.abspath("../../../../.."))
|
||||
sys.path.insert(
|
||||
0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../.."))
|
||||
)
|
||||
|
||||
# Fake OAuth token for testing (not a real secret)
|
||||
FAKE_OAUTH_TOKEN = "sk-ant-oat01-fake-token-for-testing-123456789abcdef"
|
||||
FAKE_REGULAR_KEY = "sk-ant-api03-regular-key-for-testing-123456789"
|
||||
|
||||
|
||||
def test_oauth_detection_in_common_utils():
|
||||
"""Test 1: OAuth token detection in common_utils"""
|
||||
from litellm.llms.anthropic.common_utils import optionally_handle_anthropic_oauth
|
||||
class TestOptionallyHandleAnthropicOAuth:
|
||||
"""Tests for optionally_handle_anthropic_oauth function."""
|
||||
|
||||
headers = {"authorization": f"Bearer {FAKE_OAUTH_TOKEN}"}
|
||||
updated_headers, extracted_api_key = optionally_handle_anthropic_oauth(headers, None)
|
||||
def test_oauth_token_in_authorization_header(self):
|
||||
"""OAuth token in Authorization header should be detected and headers set correctly."""
|
||||
from litellm.llms.anthropic.common_utils import (
|
||||
optionally_handle_anthropic_oauth,
|
||||
)
|
||||
|
||||
assert extracted_api_key == FAKE_OAUTH_TOKEN
|
||||
assert updated_headers["anthropic-beta"] == "oauth-2025-04-20"
|
||||
assert updated_headers["anthropic-dangerous-direct-browser-access"] == "true"
|
||||
headers = {"authorization": f"Bearer {FAKE_OAUTH_TOKEN}"}
|
||||
updated_headers, extracted_api_key = optionally_handle_anthropic_oauth(
|
||||
headers, None
|
||||
)
|
||||
|
||||
assert extracted_api_key == FAKE_OAUTH_TOKEN
|
||||
assert updated_headers["anthropic-beta"] == "oauth-2025-04-20"
|
||||
assert updated_headers["anthropic-dangerous-direct-browser-access"] == "true"
|
||||
assert "x-api-key" not in updated_headers
|
||||
|
||||
def test_oauth_token_in_api_key_directly(self):
|
||||
"""OAuth token passed as api_key should set Authorization: Bearer header."""
|
||||
from litellm.llms.anthropic.common_utils import (
|
||||
optionally_handle_anthropic_oauth,
|
||||
)
|
||||
|
||||
headers = {}
|
||||
updated_headers, returned_api_key = optionally_handle_anthropic_oauth(
|
||||
headers, FAKE_OAUTH_TOKEN
|
||||
)
|
||||
|
||||
assert returned_api_key == FAKE_OAUTH_TOKEN
|
||||
assert updated_headers["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}"
|
||||
assert updated_headers["anthropic-beta"] == "oauth-2025-04-20"
|
||||
assert updated_headers["anthropic-dangerous-direct-browser-access"] == "true"
|
||||
assert "x-api-key" not in updated_headers
|
||||
|
||||
def test_oauth_removes_existing_x_api_key(self):
|
||||
"""When OAuth is detected, any existing x-api-key should be removed."""
|
||||
from litellm.llms.anthropic.common_utils import (
|
||||
optionally_handle_anthropic_oauth,
|
||||
)
|
||||
|
||||
headers = {"x-api-key": FAKE_OAUTH_TOKEN}
|
||||
updated_headers, _ = optionally_handle_anthropic_oauth(
|
||||
headers, FAKE_OAUTH_TOKEN
|
||||
)
|
||||
|
||||
assert "x-api-key" not in updated_headers
|
||||
assert updated_headers["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}"
|
||||
|
||||
def test_regular_api_key_unchanged(self):
|
||||
"""Regular API keys (non-OAuth) should pass through unmodified."""
|
||||
from litellm.llms.anthropic.common_utils import (
|
||||
optionally_handle_anthropic_oauth,
|
||||
)
|
||||
|
||||
headers = {}
|
||||
updated_headers, returned_api_key = optionally_handle_anthropic_oauth(
|
||||
headers, FAKE_REGULAR_KEY
|
||||
)
|
||||
|
||||
assert returned_api_key == FAKE_REGULAR_KEY
|
||||
assert "authorization" not in updated_headers
|
||||
assert "anthropic-dangerous-direct-browser-access" not in updated_headers
|
||||
assert "anthropic-beta" not in updated_headers
|
||||
|
||||
def test_regular_key_in_authorization_header(self):
|
||||
"""Non-OAuth token in Authorization header should not trigger OAuth handling."""
|
||||
from litellm.llms.anthropic.common_utils import (
|
||||
optionally_handle_anthropic_oauth,
|
||||
)
|
||||
|
||||
headers = {"authorization": f"Bearer {FAKE_REGULAR_KEY}"}
|
||||
updated_headers, returned_api_key = optionally_handle_anthropic_oauth(
|
||||
headers, FAKE_REGULAR_KEY
|
||||
)
|
||||
|
||||
assert returned_api_key == FAKE_REGULAR_KEY
|
||||
assert "anthropic-dangerous-direct-browser-access" not in updated_headers
|
||||
|
||||
def test_none_api_key_no_error(self):
|
||||
"""None api_key with empty headers should not raise errors."""
|
||||
from litellm.llms.anthropic.common_utils import (
|
||||
optionally_handle_anthropic_oauth,
|
||||
)
|
||||
|
||||
headers = {}
|
||||
updated_headers, returned_api_key = optionally_handle_anthropic_oauth(
|
||||
headers, None
|
||||
)
|
||||
|
||||
assert returned_api_key is None
|
||||
assert "authorization" not in updated_headers
|
||||
|
||||
|
||||
def test_oauth_integration_in_validate_environment():
|
||||
"""Test 2: OAuth integration in AnthropicConfig validate_environment"""
|
||||
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
|
||||
class TestGetAnthropicHeaders:
|
||||
"""Tests for get_anthropic_headers method with OAuth support."""
|
||||
|
||||
config = AnthropicModelInfo()
|
||||
headers = {"authorization": f"Bearer {FAKE_OAUTH_TOKEN}"}
|
||||
def test_oauth_token_uses_authorization_bearer(self):
|
||||
"""OAuth token should produce Authorization: Bearer header, not x-api-key."""
|
||||
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
|
||||
|
||||
updated_headers = config.validate_environment(
|
||||
headers=headers,
|
||||
model="claude-3-haiku-20240307",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
)
|
||||
config = AnthropicModelInfo()
|
||||
headers = config.get_anthropic_headers(
|
||||
api_key=FAKE_OAUTH_TOKEN,
|
||||
computer_tool_used=False,
|
||||
prompt_caching_set=False,
|
||||
pdf_used=False,
|
||||
is_vertex_request=False,
|
||||
)
|
||||
|
||||
assert updated_headers["x-api-key"] == FAKE_OAUTH_TOKEN
|
||||
assert updated_headers["anthropic-dangerous-direct-browser-access"] == "true"
|
||||
assert headers["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}"
|
||||
assert headers["anthropic-dangerous-direct-browser-access"] == "true"
|
||||
assert "oauth-2025-04-20" in headers.get("anthropic-beta", "")
|
||||
assert "x-api-key" not in headers
|
||||
|
||||
def test_regular_key_uses_x_api_key(self):
|
||||
"""Regular API key should produce x-api-key header, not Authorization."""
|
||||
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
|
||||
|
||||
config = AnthropicModelInfo()
|
||||
headers = config.get_anthropic_headers(
|
||||
api_key=FAKE_REGULAR_KEY,
|
||||
computer_tool_used=False,
|
||||
prompt_caching_set=False,
|
||||
pdf_used=False,
|
||||
is_vertex_request=False,
|
||||
)
|
||||
|
||||
assert headers["x-api-key"] == FAKE_REGULAR_KEY
|
||||
assert "authorization" not in headers
|
||||
assert "anthropic-dangerous-direct-browser-access" not in headers
|
||||
|
||||
def test_oauth_includes_standard_headers(self):
|
||||
"""OAuth path should still include standard Anthropic headers."""
|
||||
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
|
||||
|
||||
config = AnthropicModelInfo()
|
||||
headers = config.get_anthropic_headers(
|
||||
api_key=FAKE_OAUTH_TOKEN,
|
||||
computer_tool_used=False,
|
||||
prompt_caching_set=False,
|
||||
pdf_used=False,
|
||||
is_vertex_request=False,
|
||||
)
|
||||
|
||||
assert headers["anthropic-version"] == "2023-06-01"
|
||||
assert headers["accept"] == "application/json"
|
||||
assert headers["content-type"] == "application/json"
|
||||
|
||||
|
||||
def test_oauth_detection_in_messages_transformation():
|
||||
"""Test 3: OAuth detection in messages transformation"""
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
|
||||
AnthropicMessagesConfig,
|
||||
)
|
||||
class TestValidateEnvironmentOAuth:
|
||||
"""Tests for validate_environment with OAuth tokens."""
|
||||
|
||||
config = AnthropicMessagesConfig()
|
||||
headers = {"authorization": f"Bearer {FAKE_OAUTH_TOKEN}"}
|
||||
def test_oauth_via_authorization_header(self):
|
||||
"""validate_environment should produce correct headers for OAuth tokens."""
|
||||
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
|
||||
|
||||
updated_headers, _ = config.validate_anthropic_messages_environment(
|
||||
headers=headers,
|
||||
model="claude-3-haiku-20240307",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
)
|
||||
config = AnthropicModelInfo()
|
||||
headers = {"authorization": f"Bearer {FAKE_OAUTH_TOKEN}"}
|
||||
|
||||
assert updated_headers["x-api-key"] == FAKE_OAUTH_TOKEN
|
||||
assert "oauth-2025-04-20" in updated_headers["anthropic-beta"]
|
||||
assert updated_headers["anthropic-dangerous-direct-browser-access"] == "true"
|
||||
updated_headers = config.validate_environment(
|
||||
headers=headers,
|
||||
model="claude-sonnet-4-5-20250929",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
)
|
||||
|
||||
assert updated_headers["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}"
|
||||
assert updated_headers["anthropic-dangerous-direct-browser-access"] == "true"
|
||||
assert "oauth-2025-04-20" in updated_headers.get("anthropic-beta", "")
|
||||
assert "x-api-key" not in updated_headers
|
||||
|
||||
def test_oauth_via_api_key_param(self):
|
||||
"""validate_environment with OAuth token as api_key should use Bearer auth."""
|
||||
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
|
||||
|
||||
config = AnthropicModelInfo()
|
||||
headers = {}
|
||||
|
||||
updated_headers = config.validate_environment(
|
||||
headers=headers,
|
||||
model="claude-sonnet-4-5-20250929",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key=FAKE_OAUTH_TOKEN,
|
||||
api_base=None,
|
||||
)
|
||||
|
||||
assert updated_headers["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}"
|
||||
assert updated_headers["anthropic-dangerous-direct-browser-access"] == "true"
|
||||
assert "x-api-key" not in updated_headers
|
||||
|
||||
def test_regular_key_via_api_key_param(self):
|
||||
"""validate_environment with regular API key should use x-api-key."""
|
||||
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
|
||||
|
||||
config = AnthropicModelInfo()
|
||||
headers = {}
|
||||
|
||||
updated_headers = config.validate_environment(
|
||||
headers=headers,
|
||||
model="claude-sonnet-4-5-20250929",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key=FAKE_REGULAR_KEY,
|
||||
api_base=None,
|
||||
)
|
||||
|
||||
assert updated_headers["x-api-key"] == FAKE_REGULAR_KEY
|
||||
assert "authorization" not in updated_headers
|
||||
assert "anthropic-dangerous-direct-browser-access" not in updated_headers
|
||||
|
||||
|
||||
def test_regular_api_keys_still_work():
|
||||
"""Test 4: Regular API keys still work (regression test)"""
|
||||
from litellm.llms.anthropic.common_utils import optionally_handle_anthropic_oauth
|
||||
class TestPassthroughOAuth:
|
||||
"""Tests for passthrough messages endpoint with OAuth tokens."""
|
||||
|
||||
regular_key = "sk-ant-api03-regular-key-123"
|
||||
headers = {"authorization": f"Bearer {regular_key}"}
|
||||
def test_passthrough_oauth_no_x_api_key(self):
|
||||
"""Passthrough endpoint should not add x-api-key for OAuth tokens."""
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
|
||||
AnthropicMessagesConfig,
|
||||
)
|
||||
|
||||
updated_headers, extracted_api_key = optionally_handle_anthropic_oauth(headers, regular_key)
|
||||
config = AnthropicMessagesConfig()
|
||||
headers = {"authorization": f"Bearer {FAKE_OAUTH_TOKEN}"}
|
||||
|
||||
# Regular key should be unchanged
|
||||
assert extracted_api_key == regular_key
|
||||
# OAuth headers should NOT be added
|
||||
assert "anthropic-dangerous-direct-browser-access" not in updated_headers
|
||||
updated_headers, _ = config.validate_anthropic_messages_environment(
|
||||
headers=headers,
|
||||
model="claude-sonnet-4-5-20250929",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
)
|
||||
|
||||
assert "oauth-2025-04-20" in updated_headers.get("anthropic-beta", "")
|
||||
assert updated_headers["anthropic-dangerous-direct-browser-access"] == "true"
|
||||
assert "x-api-key" not in updated_headers
|
||||
|
||||
def test_passthrough_regular_key_uses_x_api_key(self):
|
||||
"""Passthrough endpoint should still use x-api-key for regular API keys."""
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
|
||||
AnthropicMessagesConfig,
|
||||
)
|
||||
|
||||
config = AnthropicMessagesConfig()
|
||||
headers = {}
|
||||
|
||||
updated_headers, _ = config.validate_anthropic_messages_environment(
|
||||
headers=headers,
|
||||
model="claude-sonnet-4-5-20250929",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key=FAKE_REGULAR_KEY,
|
||||
api_base=None,
|
||||
)
|
||||
|
||||
assert updated_headers["x-api-key"] == FAKE_REGULAR_KEY
|
||||
assert "authorization" not in updated_headers
|
||||
|
|
|
|||
|
|
@ -3338,3 +3338,94 @@ def test_chunk_parser_handles_prompt_feedback_block_with_usage():
|
|||
assert result.usage.completion_tokens == 0, f"completion_tokens should be 0, got {result.usage.completion_tokens}"
|
||||
assert result.usage.total_tokens == 8175, f"total_tokens should be 8175, got {result.usage.total_tokens}"
|
||||
|
||||
|
||||
def test_vertex_ai_traffic_type_preserved_in_hidden_params_streaming():
|
||||
"""Test trafficType is preserved in _hidden_params for streaming."""
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
ModelResponseIterator,
|
||||
)
|
||||
|
||||
chunk = {
|
||||
"candidates": [{"content": {"parts": [{"text": "Hello"}]}}],
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 100,
|
||||
"candidatesTokenCount": 200,
|
||||
"totalTokenCount": 300,
|
||||
"trafficType": "ON_DEMAND",
|
||||
},
|
||||
}
|
||||
|
||||
iterator = ModelResponseIterator(
|
||||
streaming_response=[], sync_stream=True, logging_obj=MagicMock()
|
||||
)
|
||||
result = iterator.chunk_parser(chunk)
|
||||
|
||||
assert result._hidden_params["provider_specific_fields"]["traffic_type"] == "ON_DEMAND"
|
||||
|
||||
|
||||
def test_vertex_ai_traffic_type_preserved_in_hidden_params_non_streaming():
|
||||
"""Test trafficType is preserved in _hidden_params for non-streaming."""
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
VertexGeminiConfig,
|
||||
)
|
||||
|
||||
completion_response = {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {"parts": [{"text": "Hello"}], "role": "model"},
|
||||
"finishReason": "STOP",
|
||||
}
|
||||
],
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 50,
|
||||
"candidatesTokenCount": 100,
|
||||
"totalTokenCount": 150,
|
||||
"trafficType": "PROVISIONED_THROUGHPUT",
|
||||
},
|
||||
}
|
||||
|
||||
raw_response = MagicMock()
|
||||
raw_response.json.return_value = completion_response
|
||||
|
||||
result = VertexGeminiConfig().transform_response(
|
||||
model="gemini-pro",
|
||||
raw_response=raw_response,
|
||||
model_response=ModelResponse(),
|
||||
logging_obj=MagicMock(),
|
||||
request_data={},
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
|
||||
assert result._hidden_params["provider_specific_fields"]["traffic_type"] == "PROVISIONED_THROUGHPUT"
|
||||
|
||||
|
||||
def test_vertex_ai_traffic_type_surfaced_in_responses_api():
|
||||
"""Test trafficType is surfaced as provider_specific_fields in ResponsesAPIResponse."""
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
LiteLLMCompletionResponsesConfig,
|
||||
)
|
||||
|
||||
# Create a ModelResponse with provider_specific_fields in _hidden_params
|
||||
from litellm.types.utils import Choices, Message
|
||||
|
||||
model_response = ModelResponse()
|
||||
model_response._hidden_params["provider_specific_fields"] = {"traffic_type": "ON_DEMAND"}
|
||||
model_response.choices = [
|
||||
Choices(
|
||||
message=Message(content="Hello", role="assistant"),
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
)
|
||||
]
|
||||
|
||||
responses_api_response = LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response(
|
||||
request_input="test",
|
||||
chat_completion_response=model_response,
|
||||
responses_api_request={},
|
||||
)
|
||||
|
||||
assert responses_api_response.provider_specific_fields["traffic_type"] == "ON_DEMAND"
|
||||
|
||||
|
|
|
|||
|
|
@ -7,7 +7,6 @@ import sys
|
|||
|
||||
sys.path.insert(0, os.path.abspath("../../../../../.."))
|
||||
|
||||
import asyncio
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -26,7 +25,7 @@ async def test_openai_moderation_guardrail_init():
|
|||
guardrail = OpenAIModerationGuardrail(
|
||||
guardrail_name="test-openai-moderation",
|
||||
)
|
||||
|
||||
|
||||
assert guardrail.guardrail_name == "test-openai-moderation"
|
||||
assert guardrail.api_key == "test-key"
|
||||
assert guardrail.model == "omni-moderation-latest"
|
||||
|
|
@ -49,27 +48,27 @@ async def test_openai_moderation_guardrail_adds_to_litellm_callbacks():
|
|||
# Clear existing callbacks for clean test
|
||||
original_callbacks = litellm.callbacks.copy()
|
||||
litellm.logging_callback_manager._reset_all_callbacks()
|
||||
|
||||
|
||||
try:
|
||||
with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}):
|
||||
guardrail_litellm_params = LitellmParams(
|
||||
guardrail=SupportedGuardrailIntegrations.OPENAI_MODERATION,
|
||||
api_key="test-key",
|
||||
model="omni-moderation-latest",
|
||||
mode="pre_call"
|
||||
mode="pre_call",
|
||||
)
|
||||
guardrail = openai_initialize_guardrail(
|
||||
litellm_params=guardrail_litellm_params,
|
||||
guardrail=Guardrail(
|
||||
guardrail_name="test-openai-moderation",
|
||||
litellm_params=guardrail_litellm_params
|
||||
)
|
||||
litellm_params=guardrail_litellm_params,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
# Check that the guardrail was added to litellm callbacks
|
||||
assert guardrail in litellm.callbacks
|
||||
assert len(litellm.callbacks) == 1
|
||||
|
||||
|
||||
# Verify it's the correct guardrail
|
||||
callback = litellm.callbacks[0]
|
||||
assert isinstance(callback, OpenAIModerationGuardrail)
|
||||
|
|
@ -85,12 +84,12 @@ async def test_openai_moderation_guardrail_adds_to_litellm_callbacks():
|
|||
async def test_openai_moderation_guardrail_safe_content():
|
||||
"""Test OpenAI moderation guardrail with safe content via apply_guardrail"""
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
|
||||
with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}):
|
||||
guardrail = OpenAIModerationGuardrail(
|
||||
guardrail_name="test-openai-moderation",
|
||||
)
|
||||
|
||||
|
||||
# Mock safe moderation response
|
||||
mock_response = OpenAIModerationResponse(
|
||||
id="modr-123",
|
||||
|
|
@ -118,25 +117,29 @@ async def test_openai_moderation_guardrail_safe_content():
|
|||
"harassment": [],
|
||||
"self-harm": [],
|
||||
"violence": [],
|
||||
}
|
||||
},
|
||||
)
|
||||
]
|
||||
],
|
||||
)
|
||||
|
||||
with patch.object(guardrail, 'async_make_request', return_value=mock_response):
|
||||
|
||||
with patch.object(guardrail, "async_make_request", return_value=mock_response):
|
||||
# Test apply_guardrail with safe content using structured_messages
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
structured_messages=[
|
||||
{"role": "user", "content": "Hello, how are you today?"}
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data={"messages": [{"role": "user", "content": "Hello, how are you today?"}]},
|
||||
input_type="request"
|
||||
request_data={
|
||||
"messages": [
|
||||
{"role": "user", "content": "Hello, how are you today?"}
|
||||
]
|
||||
},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
|
||||
# Should return the original inputs unchanged
|
||||
assert result == inputs
|
||||
|
||||
|
|
@ -145,12 +148,12 @@ async def test_openai_moderation_guardrail_safe_content():
|
|||
async def test_openai_moderation_guardrail_apply_guardrail():
|
||||
"""Test OpenAI moderation guardrail apply_guardrail method (unified guardrail interface)"""
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
|
||||
with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}):
|
||||
guardrail = OpenAIModerationGuardrail(
|
||||
guardrail_name="test-openai-moderation",
|
||||
)
|
||||
|
||||
|
||||
# Mock safe moderation response
|
||||
mock_response = OpenAIModerationResponse(
|
||||
id="modr-123",
|
||||
|
|
@ -178,37 +181,37 @@ async def test_openai_moderation_guardrail_apply_guardrail():
|
|||
"harassment": [],
|
||||
"self-harm": [],
|
||||
"violence": [],
|
||||
}
|
||||
},
|
||||
)
|
||||
]
|
||||
],
|
||||
)
|
||||
|
||||
with patch.object(guardrail, 'async_make_request', return_value=mock_response):
|
||||
|
||||
with patch.object(guardrail, "async_make_request", return_value=mock_response):
|
||||
# Test apply_guardrail with texts (embeddings-style input)
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
texts=["Hello, how are you?", "What is the weather?"]
|
||||
)
|
||||
|
||||
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data={},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
|
||||
# Should return inputs unchanged (moderation doesn't modify, only blocks)
|
||||
assert result == inputs
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_moderation_guardrail_harmful_content():
|
||||
"""Test OpenAI moderation guardrail with harmful content via apply_guardrail"""
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
|
||||
with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}):
|
||||
guardrail = OpenAIModerationGuardrail(
|
||||
guardrail_name="test-openai-moderation",
|
||||
)
|
||||
|
||||
|
||||
# Mock harmful moderation response
|
||||
mock_response = OpenAIModerationResponse(
|
||||
id="modr-123",
|
||||
|
|
@ -236,40 +239,51 @@ async def test_openai_moderation_guardrail_harmful_content():
|
|||
"harassment": [],
|
||||
"self-harm": [],
|
||||
"violence": [],
|
||||
}
|
||||
},
|
||||
)
|
||||
]
|
||||
],
|
||||
)
|
||||
|
||||
with patch.object(guardrail, 'async_make_request', return_value=mock_response):
|
||||
|
||||
with patch.object(guardrail, "async_make_request", return_value=mock_response):
|
||||
# Test apply_guardrail with harmful content using structured_messages
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
structured_messages=[
|
||||
{"role": "user", "content": "This is hateful content"}
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
# Should raise HTTPException
|
||||
from fastapi import HTTPException
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data={"messages": [{"role": "user", "content": "This is hateful content"}]},
|
||||
input_type="request"
|
||||
request_data={
|
||||
"messages": [
|
||||
{"role": "user", "content": "This is hateful content"}
|
||||
]
|
||||
},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "Violated OpenAI moderation policy" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_moderation_guardrail_streaming_safe_content():
|
||||
"""Test OpenAI moderation guardrail with streaming safe content"""
|
||||
"""Test OpenAI moderation guardrail with streaming safe content via UnifiedLLMGuardrails"""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
UnifiedLLMGuardrails,
|
||||
)
|
||||
|
||||
with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}):
|
||||
guardrail = OpenAIModerationGuardrail(
|
||||
guardrail_name="test-openai-moderation",
|
||||
event_hook="post_call",
|
||||
)
|
||||
|
||||
unified_guardrail = UnifiedLLMGuardrails()
|
||||
|
||||
# Mock safe moderation response
|
||||
mock_response = OpenAIModerationResponse(
|
||||
id="modr-123",
|
||||
|
|
@ -297,72 +311,85 @@ async def test_openai_moderation_guardrail_streaming_safe_content():
|
|||
"harassment": [],
|
||||
"self-harm": [],
|
||||
"violence": [],
|
||||
}
|
||||
},
|
||||
)
|
||||
]
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
# Mock streaming chunks
|
||||
async def mock_stream():
|
||||
# Simulate streaming chunks with safe content
|
||||
chunks = [
|
||||
MagicMock(choices=[MagicMock(delta=MagicMock(content="Hello "))]),
|
||||
MagicMock(choices=[MagicMock(delta=MagicMock(content="world"))]),
|
||||
MagicMock(choices=[MagicMock(delta=MagicMock(content="!"))])
|
||||
]
|
||||
for chunk in chunks:
|
||||
chunk1 = MagicMock()
|
||||
chunk1.model = "gpt-4"
|
||||
chunk1.choices = [MagicMock()]
|
||||
chunk1.choices[0].delta = MagicMock()
|
||||
chunk1.choices[0].delta.content = "Hello "
|
||||
chunk1.choices[0].finish_reason = None
|
||||
|
||||
chunk2 = MagicMock()
|
||||
chunk2.model = "gpt-4"
|
||||
chunk2.choices = [MagicMock()]
|
||||
chunk2.choices[0].delta = MagicMock()
|
||||
chunk2.choices[0].delta.content = "world"
|
||||
chunk2.choices[0].finish_reason = None
|
||||
|
||||
# Last chunk with finish_reason
|
||||
chunk3 = MagicMock()
|
||||
chunk3.model = "gpt-4"
|
||||
chunk3.choices = [MagicMock()]
|
||||
chunk3.choices[0].delta = MagicMock()
|
||||
chunk3.choices[0].delta.content = "!"
|
||||
chunk3.choices[0].finish_reason = "stop"
|
||||
|
||||
for chunk in [chunk1, chunk2, chunk3]:
|
||||
yield chunk
|
||||
|
||||
# Mock the stream_chunk_builder to return a proper ModelResponse
|
||||
|
||||
# Mock for stream_chunk_builder
|
||||
mock_model_response = MagicMock()
|
||||
mock_model_response.choices = [
|
||||
MagicMock(message=MagicMock(content="Hello world!"))
|
||||
]
|
||||
|
||||
with patch.object(guardrail, 'async_make_request', return_value=mock_response), \
|
||||
patch('litellm.main.stream_chunk_builder', return_value=mock_model_response), \
|
||||
patch('litellm.llms.base_llm.base_model_iterator.MockResponseIterator') as mock_iterator:
|
||||
|
||||
# Mock the iterator to yield the original chunks
|
||||
async def mock_yield_chunks():
|
||||
chunks = [
|
||||
MagicMock(choices=[MagicMock(delta=MagicMock(content="Hello "))]),
|
||||
MagicMock(choices=[MagicMock(delta=MagicMock(content="world"))]),
|
||||
MagicMock(choices=[MagicMock(delta=MagicMock(content="!"))])
|
||||
]
|
||||
for chunk in chunks:
|
||||
yield chunk
|
||||
|
||||
mock_iterator.return_value.__aiter__ = lambda self: mock_yield_chunks()
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="test")
|
||||
mock_model_response.choices = [MagicMock()]
|
||||
mock_model_response.choices[0].message = MagicMock()
|
||||
mock_model_response.choices[0].message.content = "Hello world!"
|
||||
|
||||
with patch.object(guardrail, "async_make_request", return_value=mock_response), patch(
|
||||
"litellm.llms.openai.chat.guardrail_translation.handler.stream_chunk_builder",
|
||||
return_value=mock_model_response,
|
||||
):
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test", request_route="/chat/completions"
|
||||
)
|
||||
request_data = {
|
||||
"messages": [
|
||||
{"role": "user", "content": "Hello, how are you today?"}
|
||||
]
|
||||
"messages": [{"role": "user", "content": "Hello, how are you today?"}],
|
||||
"guardrail_to_apply": guardrail,
|
||||
"metadata": {"guardrails": ["test-openai-moderation"]},
|
||||
}
|
||||
|
||||
# Test streaming hook with safe content
|
||||
|
||||
# Test streaming hook with safe content via UnifiedLLMGuardrails
|
||||
result_chunks = []
|
||||
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
|
||||
async for chunk in unified_guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=mock_stream(),
|
||||
request_data=request_data
|
||||
request_data=request_data,
|
||||
):
|
||||
result_chunks.append(chunk)
|
||||
|
||||
|
||||
# Should return all chunks without blocking
|
||||
assert len(result_chunks) == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_moderation_guardrail_streaming_harmful_content():
|
||||
"""Test OpenAI moderation guardrail with streaming harmful content"""
|
||||
"""Test OpenAI moderation guardrail with streaming harmful content via UnifiedLLMGuardrails"""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
UnifiedLLMGuardrails,
|
||||
)
|
||||
|
||||
with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}):
|
||||
guardrail = OpenAIModerationGuardrail(
|
||||
guardrail_name="test-openai-moderation",
|
||||
event_hook="post_call",
|
||||
)
|
||||
|
||||
unified_guardrail = UnifiedLLMGuardrails()
|
||||
|
||||
# Mock harmful moderation response
|
||||
mock_response = OpenAIModerationResponse(
|
||||
id="modr-123",
|
||||
|
|
@ -390,46 +417,74 @@ async def test_openai_moderation_guardrail_streaming_harmful_content():
|
|||
"harassment": [],
|
||||
"self-harm": [],
|
||||
"violence": [],
|
||||
}
|
||||
},
|
||||
)
|
||||
]
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
# Mock streaming chunks with harmful content
|
||||
async def mock_stream():
|
||||
chunks = [
|
||||
MagicMock(choices=[MagicMock(delta=MagicMock(content="This is "))]),
|
||||
MagicMock(choices=[MagicMock(delta=MagicMock(content="harmful content"))])
|
||||
]
|
||||
for chunk in chunks:
|
||||
# First chunk - no finish_reason
|
||||
chunk1 = MagicMock()
|
||||
chunk1.model = "gpt-4"
|
||||
chunk1.choices = [MagicMock()]
|
||||
chunk1.choices[0].delta = MagicMock()
|
||||
chunk1.choices[0].delta.content = "This is "
|
||||
chunk1.choices[0].finish_reason = None
|
||||
|
||||
# Last chunk - with finish_reason to signal end of stream
|
||||
chunk2 = MagicMock()
|
||||
chunk2.model = "gpt-4"
|
||||
chunk2.choices = [MagicMock()]
|
||||
chunk2.choices[0].delta = MagicMock()
|
||||
chunk2.choices[0].delta.content = "harmful content"
|
||||
chunk2.choices[0].finish_reason = "stop"
|
||||
|
||||
for chunk in [chunk1, chunk2]:
|
||||
yield chunk
|
||||
|
||||
# Mock the stream_chunk_builder to return a ModelResponse with harmful content
|
||||
mock_model_response = MagicMock()
|
||||
mock_model_response.choices = [
|
||||
MagicMock(message=MagicMock(content="This is harmful content"))
|
||||
]
|
||||
|
||||
with patch.object(guardrail, 'async_make_request', return_value=mock_response), \
|
||||
patch('litellm.main.stream_chunk_builder', return_value=mock_model_response):
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="test")
|
||||
|
||||
# Mock for stream_chunk_builder - use real litellm types so isinstance checks pass
|
||||
from litellm.types.utils import ModelResponse
|
||||
import litellm
|
||||
mock_model_response = ModelResponse(
|
||||
id="mock-response",
|
||||
model="gpt-4",
|
||||
choices=[
|
||||
litellm.Choices(
|
||||
index=0,
|
||||
message=litellm.Message(
|
||||
role="assistant",
|
||||
content="This is harmful content",
|
||||
),
|
||||
finish_reason="stop",
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
with patch.object(guardrail, "async_make_request", return_value=mock_response), patch(
|
||||
"litellm.llms.openai.chat.guardrail_translation.handler.stream_chunk_builder",
|
||||
return_value=mock_model_response,
|
||||
):
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test", request_route="/chat/completions"
|
||||
)
|
||||
request_data = {
|
||||
"messages": [
|
||||
{"role": "user", "content": "Generate harmful content"}
|
||||
]
|
||||
"messages": [{"role": "user", "content": "Generate harmful content"}],
|
||||
"guardrail_to_apply": guardrail,
|
||||
"metadata": {"guardrails": ["test-openai-moderation"]},
|
||||
}
|
||||
|
||||
|
||||
# Should raise HTTPException when processing streaming harmful content
|
||||
from fastapi import HTTPException
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
result_chunks = []
|
||||
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
|
||||
async for chunk in unified_guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=mock_stream(),
|
||||
request_data=request_data
|
||||
request_data=request_data,
|
||||
):
|
||||
result_chunks.append(chunk)
|
||||
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "Violated OpenAI moderation policy" in str(exc_info.value.detail)
|
||||
assert "Violated OpenAI moderation policy" in str(exc_info.value.detail)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,172 @@
|
|||
import pytest
|
||||
from unittest.mock import MagicMock, patch
|
||||
import os
|
||||
from litellm.proxy.guardrails.guardrail_hooks.openai.moderations import (
|
||||
OpenAIModerationGuardrail,
|
||||
)
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
UnifiedLLMGuardrails,
|
||||
)
|
||||
from litellm.types.utils import ModelResponseStream, ModelResponse
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_moderation_guardrail_streaming_latency():
|
||||
"""
|
||||
Test that the OpenAI Moderation guardrail, when run via UnifiedLLMGuardrails,
|
||||
supports streaming (fast time-to-first-token) instead of buffering.
|
||||
"""
|
||||
with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}):
|
||||
# 1. Initialize the specific guardrail with proper event_hook
|
||||
openai_guardrail = OpenAIModerationGuardrail(
|
||||
guardrail_name="test-openai-moderation",
|
||||
event_hook="post_call",
|
||||
)
|
||||
|
||||
# 2. Initialize the Unified Guardrail system (which invokes the specific guardrail)
|
||||
unified_guardrail = UnifiedLLMGuardrails()
|
||||
|
||||
# Mock safe moderation response
|
||||
mock_mod_response = MagicMock()
|
||||
mock_mod_response.results = []
|
||||
|
||||
# Mock streaming chunks (no artificial delay - test deterministically)
|
||||
async def mock_stream():
|
||||
chunks_data = ["Hello", " ", "world", "!", " Goodbye"]
|
||||
for i, content in enumerate(chunks_data):
|
||||
chunk = MagicMock(spec=ModelResponseStream)
|
||||
chunk.model = "gpt-4"
|
||||
choice = MagicMock()
|
||||
choice.delta = MagicMock()
|
||||
choice.delta.content = content
|
||||
# Last chunk gets finish_reason
|
||||
choice.finish_reason = "stop" if i == len(chunks_data) - 1 else None
|
||||
chunk.choices = [choice]
|
||||
yield chunk
|
||||
|
||||
# Mock for stream_chunk_builder to return a simple ModelResponse
|
||||
mock_model_response = MagicMock(spec=ModelResponse)
|
||||
mock_model_response.choices = [MagicMock()]
|
||||
mock_model_response.choices[0].message = MagicMock()
|
||||
mock_model_response.choices[0].message.content = "Hello world! Goodbye"
|
||||
|
||||
# Patch the network call in the specific guardrail
|
||||
with patch.object(
|
||||
openai_guardrail, "async_make_request", return_value=mock_mod_response
|
||||
), patch(
|
||||
"litellm.llms.openai.chat.guardrail_translation.handler.stream_chunk_builder",
|
||||
return_value=mock_model_response,
|
||||
):
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test", request_route="/chat/completions"
|
||||
)
|
||||
request_data = {
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"guardrail_to_apply": openai_guardrail,
|
||||
"metadata": {
|
||||
"guardrails": ["test-openai-moderation"],
|
||||
"guardrail_config": {"streaming_sampling_rate": 1},
|
||||
}, # Check every chunk for test
|
||||
}
|
||||
|
||||
chunks_received = 0
|
||||
first_chunk_yielded = False
|
||||
|
||||
# Call the hook on UnifiedLLMGuardrails
|
||||
async for chunk in unified_guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=mock_stream(),
|
||||
request_data=request_data,
|
||||
):
|
||||
if not first_chunk_yielded:
|
||||
first_chunk_yielded = True
|
||||
chunks_received += 1
|
||||
|
||||
# Deterministic assertions (no flaky timing checks)
|
||||
assert first_chunk_yielded, "Expected at least one chunk to be yielded"
|
||||
assert chunks_received == 5, f"Expected 5 chunks, got {chunks_received}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_moderation_guardrail_streaming_harmful_content():
|
||||
"""
|
||||
Test that harmful content is caught during streaming via UnifiedLLMGuardrails
|
||||
"""
|
||||
from fastapi import HTTPException
|
||||
|
||||
with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}):
|
||||
openai_guardrail = OpenAIModerationGuardrail(
|
||||
guardrail_name="test-openai-moderation",
|
||||
event_hook="post_call",
|
||||
)
|
||||
unified_guardrail = UnifiedLLMGuardrails()
|
||||
|
||||
# Mock harmful moderation response
|
||||
mock_mod_response = MagicMock()
|
||||
mock_mod_response.results = [
|
||||
MagicMock(
|
||||
flagged=True, categories={"hate": True}, category_scores={"hate": 0.99}
|
||||
)
|
||||
]
|
||||
|
||||
async def mock_stream():
|
||||
chunks_data = ["This ", "is ", "harmful ", "content"]
|
||||
for i, content in enumerate(chunks_data):
|
||||
chunk = MagicMock(spec=ModelResponseStream)
|
||||
chunk.model = "gpt-4"
|
||||
choice = MagicMock()
|
||||
choice.delta = MagicMock()
|
||||
choice.delta.content = content
|
||||
# Last chunk gets finish_reason
|
||||
choice.finish_reason = "stop" if i == len(chunks_data) - 1 else None
|
||||
chunk.choices = [choice]
|
||||
yield chunk
|
||||
|
||||
# Mock for stream_chunk_builder - use real litellm types so isinstance checks pass
|
||||
import litellm
|
||||
|
||||
mock_model_response = ModelResponse(
|
||||
id="mock-response",
|
||||
model="gpt-4",
|
||||
choices=[
|
||||
litellm.Choices(
|
||||
index=0,
|
||||
message=litellm.Message(
|
||||
role="assistant",
|
||||
content="This is harmful content",
|
||||
),
|
||||
finish_reason="stop",
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
openai_guardrail, "async_make_request", return_value=mock_mod_response
|
||||
), patch(
|
||||
"litellm.llms.openai.chat.guardrail_translation.handler.stream_chunk_builder",
|
||||
return_value=mock_model_response,
|
||||
):
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test", request_route="/chat/completions"
|
||||
)
|
||||
request_data = {
|
||||
"messages": [{"role": "user", "content": "generate hate"}],
|
||||
"guardrail_to_apply": openai_guardrail,
|
||||
"metadata": {
|
||||
"guardrails": ["test-openai-moderation"],
|
||||
"guardrail_config": {"streaming_sampling_rate": 1},
|
||||
},
|
||||
}
|
||||
|
||||
# Should raise HTTPException
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=mock_stream(),
|
||||
request_data=request_data,
|
||||
):
|
||||
pass
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "Violated OpenAI moderation policy" in str(exc_info.value.detail)
|
||||
|
|
@ -1122,6 +1122,83 @@ async def test_model_armor_non_model_response():
|
|||
assert not guardrail.async_handler.post.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_armor_guardrail_status_intervened_vs_failed():
|
||||
"""
|
||||
regression test for bug where _process_error always set 'guardrail_failed_to_respond'
|
||||
even for intentional blocks (error 400).
|
||||
"""
|
||||
mock_user_api_key_dict = UserAPIKeyAuth()
|
||||
mock_cache = MagicMock(spec=DualCache)
|
||||
|
||||
#1: Blocked content should raise exception and show guardrail status: guardrail_intervened"
|
||||
guardrail = ModelArmorGuardrail(
|
||||
template_id="test-template",
|
||||
project_id="test-project",
|
||||
location="us-central1",
|
||||
guardrail_name="model-armor-test",
|
||||
)
|
||||
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json = AsyncMock(return_value={
|
||||
"sanitizationResult": {
|
||||
"filterMatchState": "MATCH_FOUND",
|
||||
"filterResults": {
|
||||
"rai": {
|
||||
"raiFilterResult": {
|
||||
"matchState": "MATCH_FOUND",
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
guardrail._ensure_access_token_async = AsyncMock(return_value=("token", "test-project"))
|
||||
with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=mock_response)):
|
||||
request_data = {
|
||||
"model": "gpt-4",
|
||||
"messages": [{"role": "user", "content": "bad content"}],
|
||||
"metadata": {"guardrails": ["model-armor-test"]},
|
||||
}
|
||||
with pytest.raises(HTTPException):
|
||||
await guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
cache=mock_cache,
|
||||
data=request_data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
info = request_data["metadata"]["standard_logging_guardrail_information"]
|
||||
assert info[0]["guardrail_status"] == "guardrail_intervened"
|
||||
|
||||
#2: if an API error - guardrail status should be guardrail_failed_to_respond"
|
||||
guardrail2 = ModelArmorGuardrail(
|
||||
template_id="test-template",
|
||||
project_id="test-project",
|
||||
location="us-central1",
|
||||
guardrail_name="model-armor-test2",
|
||||
fail_on_error=True,
|
||||
)
|
||||
|
||||
guardrail2._ensure_access_token_async = AsyncMock(side_effect=ConnectionError("timeout"))
|
||||
request_data2 = {
|
||||
"model": "gpt-4",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"metadata": {"guardrails": ["model-armor-test2"]},
|
||||
}
|
||||
with pytest.raises(ConnectionError):
|
||||
await guardrail2.async_pre_call_hook(
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
cache=mock_cache,
|
||||
data=request_data2,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
info2 = request_data2["metadata"]["standard_logging_guardrail_information"]
|
||||
assert info2[0]["guardrail_status"] == "guardrail_failed_to_respond"
|
||||
|
||||
|
||||
def mock_open(read_data=''):
|
||||
"""Helper to create a mock file object"""
|
||||
import io
|
||||
|
|
|
|||
|
|
@ -0,0 +1,232 @@
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.router_strategy.budget_limiter import RouterBudgetLimiting
|
||||
from litellm.types.router import LiteLLM_Params
|
||||
from litellm.types.utils import BudgetConfig
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def disable_budget_sync(monkeypatch):
|
||||
async def noop(*args, **kwargs):
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.router_strategy.budget_limiter.RouterBudgetLimiting.periodic_sync_in_memory_spend_with_redis",
|
||||
noop,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_llm_provider_for_deployment_dict_does_not_require_litellm_params_instantiation(
|
||||
disable_budget_sync, monkeypatch
|
||||
):
|
||||
class RaiseOnInit:
|
||||
def __init__(self, *args, **kwargs):
|
||||
raise AssertionError("LiteLLM_Params should not be instantiated in hot path")
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.router_strategy.budget_limiter.LiteLLM_Params",
|
||||
RaiseOnInit,
|
||||
)
|
||||
|
||||
provider_budget = RouterBudgetLimiting(
|
||||
dual_cache=DualCache(),
|
||||
provider_budget_config={},
|
||||
)
|
||||
|
||||
deployment = {"litellm_params": {"model": "openai/gpt-4o-mini"}}
|
||||
provider = provider_budget._get_llm_provider_for_deployment(deployment)
|
||||
|
||||
assert provider == "openai"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_llm_provider_for_deployment_dict_view_supports_mapping_and_attr_access(
|
||||
disable_budget_sync, monkeypatch
|
||||
):
|
||||
observed = {}
|
||||
|
||||
def _future_style_get_llm_provider(
|
||||
model,
|
||||
custom_llm_provider=None,
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
litellm_params=None,
|
||||
):
|
||||
assert litellm_params is not None
|
||||
observed["model_attr"] = litellm_params.model
|
||||
observed["provider_get"] = litellm_params.get("custom_llm_provider")
|
||||
observed["api_base_item"] = litellm_params["api_base"]
|
||||
observed["has_api_key"] = "api_key" in litellm_params
|
||||
observed["model_dump"] = litellm_params.model_dump()
|
||||
return model, "openai", None, None
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.router_strategy.budget_limiter.litellm.get_llm_provider",
|
||||
_future_style_get_llm_provider,
|
||||
)
|
||||
|
||||
provider_budget = RouterBudgetLimiting(
|
||||
dual_cache=DualCache(),
|
||||
provider_budget_config={},
|
||||
)
|
||||
|
||||
deployment = {
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"custom_llm_provider": "openai",
|
||||
"api_base": "https://api.openai.com/v1",
|
||||
}
|
||||
}
|
||||
provider = provider_budget._get_llm_provider_for_deployment(deployment)
|
||||
|
||||
assert provider == "openai"
|
||||
assert observed["model_attr"] == "openai/gpt-4o-mini"
|
||||
assert observed["provider_get"] == "openai"
|
||||
assert observed["api_base_item"] == "https://api.openai.com/v1"
|
||||
assert observed["has_api_key"] is False
|
||||
assert observed["model_dump"]["model"] == "openai/gpt-4o-mini"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_filter_deployments_resolves_provider_once_per_deployment(
|
||||
disable_budget_sync, monkeypatch
|
||||
):
|
||||
provider_budget = RouterBudgetLimiting(
|
||||
dual_cache=DualCache(),
|
||||
provider_budget_config={
|
||||
"openai": BudgetConfig(budget_duration="1d", max_budget=100.0),
|
||||
},
|
||||
)
|
||||
|
||||
healthy_deployments = [
|
||||
{
|
||||
"model_name": "gpt-4o-mini",
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini"},
|
||||
"model_info": {"id": "deployment-1"},
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-4o-mini",
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini"},
|
||||
"model_info": {"id": "deployment-2"},
|
||||
},
|
||||
]
|
||||
|
||||
provider_resolution_calls = 0
|
||||
|
||||
def _count_provider_calls(deployment):
|
||||
nonlocal provider_resolution_calls
|
||||
provider_resolution_calls += 1
|
||||
return "openai"
|
||||
|
||||
monkeypatch.setattr(
|
||||
provider_budget,
|
||||
"_get_llm_provider_for_deployment",
|
||||
_count_provider_calls,
|
||||
)
|
||||
|
||||
filtered_deployments = await provider_budget.async_filter_deployments(
|
||||
model="gpt-4o-mini",
|
||||
healthy_deployments=healthy_deployments,
|
||||
messages=[],
|
||||
request_kwargs={},
|
||||
parent_otel_span=None,
|
||||
)
|
||||
|
||||
assert len(filtered_deployments) == len(healthy_deployments)
|
||||
assert provider_resolution_calls == len(healthy_deployments)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_filter_deployments_does_not_recompute_provider_when_resolved_none(
|
||||
disable_budget_sync, monkeypatch
|
||||
):
|
||||
provider_budget = RouterBudgetLimiting(
|
||||
dual_cache=DualCache(),
|
||||
provider_budget_config={
|
||||
"openai": BudgetConfig(budget_duration="1d", max_budget=100.0),
|
||||
},
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-4o-mini",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"max_budget": 100.0,
|
||||
"budget_duration": "1d",
|
||||
},
|
||||
"model_info": {"id": "deployment-1"},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
healthy_deployments = [
|
||||
{
|
||||
"model_name": "gpt-4o-mini",
|
||||
"litellm_params": {"model": "unknown-provider/model"},
|
||||
"model_info": {"id": "deployment-1"},
|
||||
}
|
||||
]
|
||||
|
||||
provider_resolution_calls = 0
|
||||
|
||||
def _provider_returns_none(deployment):
|
||||
nonlocal provider_resolution_calls
|
||||
provider_resolution_calls += 1
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(
|
||||
provider_budget,
|
||||
"_get_llm_provider_for_deployment",
|
||||
_provider_returns_none,
|
||||
)
|
||||
|
||||
filtered_deployments = await provider_budget.async_filter_deployments(
|
||||
model="gpt-4o-mini",
|
||||
healthy_deployments=healthy_deployments,
|
||||
messages=[],
|
||||
request_kwargs={},
|
||||
parent_otel_span=None,
|
||||
)
|
||||
|
||||
assert len(filtered_deployments) == len(healthy_deployments)
|
||||
assert provider_resolution_calls == len(healthy_deployments)
|
||||
|
||||
|
||||
def _legacy_provider_resolution(deployment):
|
||||
"""
|
||||
Reference implementation used before hot-path optimization.
|
||||
"""
|
||||
try:
|
||||
_litellm_params = LiteLLM_Params(**deployment.get("litellm_params", {"model": ""}))
|
||||
_, custom_llm_provider, _, _ = litellm.get_llm_provider(
|
||||
model=_litellm_params.model,
|
||||
litellm_params=_litellm_params,
|
||||
)
|
||||
except Exception:
|
||||
return None
|
||||
return custom_llm_provider
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"deployment",
|
||||
[
|
||||
{"litellm_params": {"model": "openai/gpt-4o-mini"}},
|
||||
{"litellm_params": {"model": "gpt-4o-mini", "custom_llm_provider": "openai"}},
|
||||
{"litellm_params": {"model": "unknown-provider/model"}},
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_llm_provider_for_deployment_matches_legacy_behavior(
|
||||
disable_budget_sync, deployment
|
||||
):
|
||||
provider_budget = RouterBudgetLimiting(
|
||||
dual_cache=DualCache(),
|
||||
provider_budget_config={},
|
||||
)
|
||||
|
||||
current_provider = provider_budget._get_llm_provider_for_deployment(deployment)
|
||||
legacy_provider = _legacy_provider_resolution(deployment)
|
||||
|
||||
assert current_provider == legacy_provider
|
||||
5
ui/litellm-dashboard/public/assets/logos/zscaler.svg
Normal file
5
ui/litellm-dashboard/public/assets/logos/zscaler.svg
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
<svg width="50" height="41" viewBox="0 0 35 22" fill="none" xmlns="http://www.w3.org/2000/svg">
|
||||
<path d="M34.9305 11.1952C35.4985 14.7139 32.4864 16.6962 29.4313 16.9252C27.4867 20.976 20.6324 23.7107 14.6971 20.762C12.1552 21.1535 10.764 20.5092 9.6896 19.2752C11.861 16.4425 18.4366 11.3459 26.47 13.9755C30.7569 15.3849 31.8715 12.0843 30.7961 10.8513C26.7505 6.19284 17.6134 10.3865 17.2879 10.7353C20.8729 5.16699 33.7065 3.60491 34.9305 11.1952ZM22.6298 4.80521C22.6522 4.79728 19.7125 3.74467 15.6034 5.54462C15.4367 5.4832 15.2735 5.41239 15.1146 5.33251C19.0701 2.66529 22.5359 1.53437 25.5001 1.97445C23.7042 -0.11988 15.0833 -1.65817 10.1432 3.36208C4.03683 2.14096 -0.233535 7.38224 0.00989967 12.1765C0.253334 16.9708 5.41043 19.9502 8.34533 19.1424C8.41578 19.1335 8.48704 19.1335 8.55748 19.1424C9.21348 15.9766 11.8658 8.61029 22.6298 4.80521Z" fill="#2160E1"/>
|
||||
</svg>
|
||||
|
||||
<!--65 41, 50 32-->
|
||||
|
After Width: | Height: | Size: 906 B |
|
|
@ -104,6 +104,7 @@ export const shouldRenderContentFilterConfigSettings = (provider: string | null)
|
|||
const asset_logos_folder = "../ui/assets/logos/";
|
||||
|
||||
export const guardrailLogoMap: Record<string, string> = {
|
||||
"Zscaler AI Guard": `${asset_logos_folder}zscaler.svg`,
|
||||
"Presidio PII": `${asset_logos_folder}presidio.png`,
|
||||
"Bedrock Guardrail": `${asset_logos_folder}bedrock.svg`,
|
||||
Lakera: `${asset_logos_folder}lakeraai.jpeg`,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue