mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge pull request #18105 from BerriAI/litellm_staging_12_17_2025
Staging - 12/17/2025
This commit is contained in:
commit
95b2e72dfc
19 changed files with 1047 additions and 384 deletions
|
|
@ -34,8 +34,8 @@ RUN pip wheel --no-cache-dir --wheel-dir=/wheels/ -r requirements.txt
|
|||
# Runtime stage
|
||||
FROM $LITELLM_RUNTIME_IMAGE AS runtime
|
||||
|
||||
# Update dependencies and clean up
|
||||
RUN apk upgrade --no-cache
|
||||
# Update dependencies and clean up, install libsndfile for audio processing
|
||||
RUN apk upgrade --no-cache && apk add --no-cache libsndfile
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
|
|
|
|||
|
|
@ -795,6 +795,8 @@ def image_edit(
|
|||
model=model,
|
||||
image_edit_provider_config=image_edit_provider_config,
|
||||
image_edit_optional_params=image_edit_optional_params,
|
||||
drop_params=kwargs.get("drop_params"),
|
||||
additional_drop_params=kwargs.get("additional_drop_params"),
|
||||
)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
from io import BufferedReader, BytesIO
|
||||
from typing import Any, Dict, cast, get_type_hints
|
||||
from typing import Any, Dict, List, Optional, cast, get_type_hints
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.token_counter import get_image_type
|
||||
|
|
@ -14,41 +14,53 @@ class ImageEditRequestUtils:
|
|||
model: str,
|
||||
image_edit_provider_config: BaseImageEditConfig,
|
||||
image_edit_optional_params: ImageEditOptionalRequestParams,
|
||||
drop_params: Optional[bool] = None,
|
||||
additional_drop_params: Optional[List[str]] = None,
|
||||
) -> Dict:
|
||||
"""
|
||||
Get optional parameters for the image edit API.
|
||||
|
||||
Args:
|
||||
params: Dictionary of all parameters
|
||||
model: The model name
|
||||
image_edit_provider_config: The provider configuration for image edit API
|
||||
image_edit_optional_params: The optional parameters for the image edit API
|
||||
drop_params: If True, silently drop unsupported parameters instead of raising
|
||||
additional_drop_params: List of additional parameter names to drop
|
||||
|
||||
Returns:
|
||||
A dictionary of supported parameters for the image edit API
|
||||
"""
|
||||
# Remove None values and internal parameters
|
||||
|
||||
# Get supported parameters for the model
|
||||
supported_params = image_edit_provider_config.get_supported_openai_params(model)
|
||||
|
||||
# Check for unsupported parameters
|
||||
should_drop = litellm.drop_params is True or drop_params is True
|
||||
|
||||
filtered_optional_params = dict(image_edit_optional_params)
|
||||
if additional_drop_params:
|
||||
for param in additional_drop_params:
|
||||
filtered_optional_params.pop(param, None)
|
||||
|
||||
unsupported_params = [
|
||||
param
|
||||
for param in image_edit_optional_params
|
||||
for param in filtered_optional_params
|
||||
if param not in supported_params
|
||||
]
|
||||
|
||||
if unsupported_params:
|
||||
raise litellm.UnsupportedParamsError(
|
||||
model=model,
|
||||
message=f"The following parameters are not supported for model {model}: {', '.join(unsupported_params)}",
|
||||
)
|
||||
if should_drop:
|
||||
for param in unsupported_params:
|
||||
filtered_optional_params.pop(param, None)
|
||||
else:
|
||||
raise litellm.UnsupportedParamsError(
|
||||
model=model,
|
||||
message=f"The following parameters are not supported for model {model}: {', '.join(unsupported_params)}",
|
||||
)
|
||||
|
||||
# Map parameters to provider-specific format
|
||||
mapped_params = image_edit_provider_config.map_openai_params(
|
||||
image_edit_optional_params=image_edit_optional_params,
|
||||
image_edit_optional_params=cast(
|
||||
ImageEditOptionalRequestParams, filtered_optional_params
|
||||
),
|
||||
model=model,
|
||||
drop_params=litellm.drop_params,
|
||||
drop_params=should_drop,
|
||||
)
|
||||
|
||||
return mapped_params
|
||||
|
|
|
|||
|
|
@ -815,7 +815,20 @@ class PrometheusLogger(CustomLogger):
|
|||
user_api_key_auth_metadata: Optional[dict] = standard_logging_payload[
|
||||
"metadata"
|
||||
].get("user_api_key_auth_metadata")
|
||||
|
||||
# Include top-level metadata fields (excluding nested dictionaries)
|
||||
# This allows accessing fields like requester_ip_address from top-level metadata
|
||||
top_level_metadata = standard_logging_payload.get("metadata", {})
|
||||
top_level_fields: Dict[str, Any] = {}
|
||||
if isinstance(top_level_metadata, dict):
|
||||
top_level_fields = {
|
||||
k: v
|
||||
for k, v in top_level_metadata.items()
|
||||
if not isinstance(v, dict) # Exclude nested dicts to avoid conflicts
|
||||
}
|
||||
|
||||
combined_metadata: Dict[str, Any] = {
|
||||
**top_level_fields, # Include top-level fields first
|
||||
**(_requester_metadata if _requester_metadata else {}),
|
||||
**(user_api_key_auth_metadata if user_api_key_auth_metadata else {}),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -253,20 +253,39 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
task_mappings: List[Tuple[int, Optional[int]]] = []
|
||||
# Track (content_index, None) for each text
|
||||
|
||||
response_content = response.get("content", [])
|
||||
# Handle both dict and object responses
|
||||
response_content: List[Any] = []
|
||||
if isinstance(response, dict):
|
||||
response_content = response.get("content", []) or []
|
||||
elif hasattr(response, "content"):
|
||||
content = getattr(response, "content", None)
|
||||
response_content = content or []
|
||||
else:
|
||||
response_content = []
|
||||
|
||||
if not response_content:
|
||||
return response
|
||||
|
||||
# Step 1: Extract all text content and tool calls from response
|
||||
for content_idx, content_block in enumerate(response_content):
|
||||
# Check if this is a text or tool_use block by checking the 'type' field
|
||||
if isinstance(content_block, dict) and content_block.get("type") in [
|
||||
"text",
|
||||
"tool_use",
|
||||
]:
|
||||
# Cast to dict to handle the union type properly
|
||||
# Handle both dict and Pydantic object content blocks
|
||||
block_dict: Dict[str, Any] = {}
|
||||
if isinstance(content_block, dict):
|
||||
block_type = content_block.get("type")
|
||||
block_dict = cast(Dict[str, Any], content_block)
|
||||
elif hasattr(content_block, "type"):
|
||||
block_type = getattr(content_block, "type", None)
|
||||
# Convert Pydantic object to dict for processing
|
||||
if hasattr(content_block, "model_dump"):
|
||||
block_dict = content_block.model_dump()
|
||||
else:
|
||||
block_dict = {"type": block_type, "text": getattr(content_block, "text", None)}
|
||||
else:
|
||||
continue
|
||||
|
||||
if block_type in ["text", "tool_use"]:
|
||||
self._extract_output_text_and_images(
|
||||
content_block=cast(Dict[str, Any], content_block),
|
||||
content_block=block_dict,
|
||||
content_idx=content_idx,
|
||||
texts_to_check=texts_to_check,
|
||||
images_to_check=images_to_check,
|
||||
|
|
@ -530,7 +549,11 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
|
||||
Override this method to customize text content detection.
|
||||
"""
|
||||
response_content = response.get("content", [])
|
||||
if isinstance(response, dict):
|
||||
response_content = response.get("content", [])
|
||||
else:
|
||||
response_content = getattr(response, "content", None) or []
|
||||
|
||||
if not response_content:
|
||||
return False
|
||||
for content_block in response_content:
|
||||
|
|
@ -590,7 +613,16 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
mapping = task_mappings[task_idx]
|
||||
content_idx = cast(int, mapping[0])
|
||||
|
||||
response_content = response.get("content", [])
|
||||
# Handle both dict and object responses
|
||||
response_content: List[Any] = []
|
||||
if isinstance(response, dict):
|
||||
response_content = response.get("content", []) or []
|
||||
elif hasattr(response, "content"):
|
||||
content = getattr(response, "content", None)
|
||||
response_content = content or []
|
||||
else:
|
||||
continue
|
||||
|
||||
if not response_content:
|
||||
continue
|
||||
|
||||
|
|
@ -601,7 +633,11 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
content_block = response_content[content_idx]
|
||||
|
||||
# Verify it's a text block and update the text field
|
||||
if isinstance(content_block, dict) and content_block.get("type") == "text":
|
||||
# Cast to dict to handle the union type properly for assignment
|
||||
content_block = cast("AnthropicResponseTextBlock", content_block)
|
||||
content_block["text"] = guardrail_response
|
||||
# Handle both dict and Pydantic object content blocks
|
||||
if isinstance(content_block, dict):
|
||||
if content_block.get("type") == "text":
|
||||
cast(Dict[str, Any], content_block)["text"] = guardrail_response
|
||||
elif hasattr(content_block, "type") and getattr(content_block, "type", None) == "text":
|
||||
# Update Pydantic object's text attribute
|
||||
if hasattr(content_block, "text"):
|
||||
content_block.text = guardrail_response
|
||||
|
|
|
|||
|
|
@ -30,7 +30,7 @@ Output: response.output is List[GenericResponseOutputItem] where each has:
|
|||
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
|
||||
|
||||
from openai.types.responses import ResponseFunctionToolCall
|
||||
from openai.types.responses.response_function_tool_call import ResponseFunctionToolCall
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -299,8 +299,25 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
task_mappings: List[Tuple[int, int]] = []
|
||||
# Track (output_item_index, content_index) for each text
|
||||
|
||||
# Handle both dict and Pydantic object responses
|
||||
if isinstance(response, dict):
|
||||
response_output = response.get("output", [])
|
||||
elif hasattr(response, "output"):
|
||||
response_output = response.output or []
|
||||
else:
|
||||
verbose_proxy_logger.debug(
|
||||
"OpenAI Responses API: No output found in response"
|
||||
)
|
||||
return response
|
||||
|
||||
if not response_output:
|
||||
verbose_proxy_logger.debug(
|
||||
"OpenAI Responses API: Empty output in response"
|
||||
)
|
||||
return response
|
||||
|
||||
# Step 1: Extract all text content and tool calls from response output
|
||||
for output_idx, output_item in enumerate(response.output):
|
||||
for output_idx, output_item in enumerate(response_output):
|
||||
self._extract_output_text_and_images(
|
||||
output_item=output_item,
|
||||
output_idx=output_idx,
|
||||
|
|
@ -538,13 +555,18 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
content: Optional[Union[List[OutputText], List[dict]]] = None
|
||||
if isinstance(output_item, BaseModel):
|
||||
try:
|
||||
output_item_dump = output_item.model_dump()
|
||||
generic_response_output_item = GenericResponseOutputItem.model_validate(
|
||||
output_item.model_dump()
|
||||
output_item_dump
|
||||
)
|
||||
if generic_response_output_item.content:
|
||||
content = generic_response_output_item.content
|
||||
except Exception:
|
||||
return
|
||||
# Try to extract content directly from output_item if validation fails
|
||||
if hasattr(output_item, "content") and output_item.content:
|
||||
content = output_item.content
|
||||
else:
|
||||
return
|
||||
elif isinstance(output_item, dict):
|
||||
content = output_item.get("content", [])
|
||||
else:
|
||||
|
|
@ -582,22 +604,53 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
|
||||
Override this method to customize how responses are applied.
|
||||
"""
|
||||
# Handle both dict and Pydantic object responses
|
||||
if isinstance(response, dict):
|
||||
response_output = response.get("output", [])
|
||||
elif hasattr(response, "output"):
|
||||
response_output = response.output or []
|
||||
else:
|
||||
return
|
||||
|
||||
for task_idx, guardrail_response in enumerate(responses):
|
||||
mapping = task_mappings[task_idx]
|
||||
output_idx = cast(int, mapping[0])
|
||||
content_idx = cast(int, mapping[1])
|
||||
|
||||
output_item = response.output[output_idx]
|
||||
if output_idx >= len(response_output):
|
||||
continue
|
||||
|
||||
# Handle both GenericResponseOutputItem and dict
|
||||
output_item = response_output[output_idx]
|
||||
|
||||
# Handle both GenericResponseOutputItem, BaseModel, and dict
|
||||
if isinstance(output_item, GenericResponseOutputItem):
|
||||
content_item = output_item.content[content_idx]
|
||||
if isinstance(content_item, OutputText):
|
||||
content_item.text = guardrail_response
|
||||
elif isinstance(content_item, dict):
|
||||
content_item["text"] = guardrail_response
|
||||
if output_item.content and content_idx < len(output_item.content):
|
||||
content_item = output_item.content[content_idx]
|
||||
if isinstance(content_item, OutputText):
|
||||
content_item.text = guardrail_response
|
||||
elif isinstance(content_item, dict):
|
||||
content_item["text"] = guardrail_response
|
||||
elif isinstance(output_item, BaseModel):
|
||||
# Handle other Pydantic models by converting to GenericResponseOutputItem
|
||||
try:
|
||||
generic_item = GenericResponseOutputItem.model_validate(
|
||||
output_item.model_dump()
|
||||
)
|
||||
if generic_item.content and content_idx < len(generic_item.content):
|
||||
content_item = generic_item.content[content_idx]
|
||||
if isinstance(content_item, OutputText):
|
||||
content_item.text = guardrail_response
|
||||
# Update the original response output
|
||||
if hasattr(output_item, "content") and output_item.content:
|
||||
original_content = output_item.content[content_idx]
|
||||
if hasattr(original_content, "text"):
|
||||
original_content.text = guardrail_response
|
||||
except Exception:
|
||||
pass
|
||||
elif isinstance(output_item, dict):
|
||||
content = output_item.get("content", [])
|
||||
if content and content_idx < len(content):
|
||||
if isinstance(content[content_idx], dict):
|
||||
content[content_idx]["text"] = guardrail_response
|
||||
elif hasattr(content[content_idx], "text"):
|
||||
content[content_idx].text = guardrail_response
|
||||
|
|
|
|||
|
|
@ -114,19 +114,24 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM):
|
|||
"""
|
||||
Get the complete URL for Vertex AI Gemini generateContent API
|
||||
"""
|
||||
vertex_project = self._resolve_vertex_project()
|
||||
vertex_location = self._resolve_vertex_location()
|
||||
vertex_project = (
|
||||
litellm_params.get("vertex_project") or self._resolve_vertex_project()
|
||||
)
|
||||
vertex_location = (
|
||||
litellm_params.get("vertex_location") or self._resolve_vertex_location()
|
||||
)
|
||||
|
||||
if not vertex_project or not vertex_location:
|
||||
raise ValueError("vertex_project and vertex_location are required for Vertex AI")
|
||||
|
||||
# Use the model name as provided, handling vertex_ai prefix
|
||||
model_name = model
|
||||
if model.startswith("vertex_ai/"):
|
||||
model_name = model.replace("vertex_ai/", "")
|
||||
|
||||
if api_base:
|
||||
base_url = api_base.rstrip("/")
|
||||
elif vertex_location == "global":
|
||||
base_url = "https://aiplatform.googleapis.com"
|
||||
else:
|
||||
base_url = f"https://{vertex_location}-aiplatform.googleapis.com"
|
||||
|
||||
|
|
|
|||
|
|
@ -616,6 +616,14 @@ def get_model_from_request(
|
|||
if match:
|
||||
model = match.group(1)
|
||||
|
||||
# If still not found, extract from Vertex AI passthrough route
|
||||
# Pattern: /vertex_ai/.../models/{model_id}:*
|
||||
# Example: /vertex_ai/v1/.../models/gemini-1.5-pro:generateContent
|
||||
if model is None and "/vertex" in route.lower():
|
||||
vertex_match = re.search(r"/models/([^/:]+)", route)
|
||||
if vertex_match:
|
||||
model = vertex_match.group(1)
|
||||
|
||||
return model
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -40,6 +40,12 @@ def initialize_guardrail(
|
|||
),
|
||||
categories=_get_config_value(litellm_params, optional_params, "categories"),
|
||||
policy_id=_get_config_value(litellm_params, optional_params, "policy_id"),
|
||||
streaming_end_of_stream_only=_get_config_value(
|
||||
litellm_params, optional_params, "streaming_end_of_stream_only"
|
||||
) or False,
|
||||
streaming_sampling_rate=_get_config_value(
|
||||
litellm_params, optional_params, "streaming_sampling_rate"
|
||||
) or 5,
|
||||
event_hook=litellm_params.mode,
|
||||
default_on=litellm_params.default_on,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,26 +1,25 @@
|
|||
"""Gray Swan Cygnal guardrail integration."""
|
||||
|
||||
import os
|
||||
from typing import Any, Dict, Literal, Optional, Union
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
ModifyResponseException,
|
||||
)
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_guardrail_to_applied_guardrails_header,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.utils import Choices, LLMResponseTypes, ModelResponse
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
|
||||
class GraySwanGuardrailMissingSecrets(Exception):
|
||||
|
|
@ -35,6 +34,15 @@ class GraySwanGuardrail(CustomGuardrail):
|
|||
"""
|
||||
Guardrail that calls Gray Swan's Cygnal monitoring endpoint.
|
||||
|
||||
Uses the unified guardrail system via `apply_guardrail` method,
|
||||
which automatically works with all LiteLLM endpoints:
|
||||
- OpenAI Chat Completions
|
||||
- OpenAI Responses API
|
||||
- OpenAI Text Completions
|
||||
- Anthropic Messages
|
||||
- Image Generation
|
||||
- And more...
|
||||
|
||||
see: https://docs.grayswan.ai/cygnal/monitor-requests
|
||||
"""
|
||||
|
||||
|
|
@ -54,6 +62,8 @@ class GraySwanGuardrail(CustomGuardrail):
|
|||
reasoning_mode: Optional[str] = None,
|
||||
categories: Optional[Dict[str, str]] = None,
|
||||
policy_id: Optional[str] = None,
|
||||
streaming_end_of_stream_only: bool = False,
|
||||
streaming_sampling_rate: int = 5,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
self.async_handler = get_async_httpx_client(
|
||||
|
|
@ -88,6 +98,16 @@ class GraySwanGuardrail(CustomGuardrail):
|
|||
self.categories = categories
|
||||
self.policy_id = policy_id
|
||||
|
||||
# Streaming configuration
|
||||
self.streaming_end_of_stream_only = streaming_end_of_stream_only
|
||||
self.streaming_sampling_rate = streaming_sampling_rate
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"GraySwan __init__: streaming_end_of_stream_only=%s, streaming_sampling_rate=%s",
|
||||
streaming_end_of_stream_only,
|
||||
streaming_sampling_rate,
|
||||
)
|
||||
|
||||
supported_event_hooks = [
|
||||
GuardrailEventHooks.pre_call,
|
||||
GuardrailEventHooks.during_call,
|
||||
|
|
@ -101,217 +121,227 @@ class GraySwanGuardrail(CustomGuardrail):
|
|||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Guardrail hook entry points
|
||||
# Debug override to trace post_call issues
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@log_guardrail_information
|
||||
async def async_pre_call_hook(
|
||||
def should_run_guardrail(self, data, event_type) -> bool:
|
||||
"""Override to add debug logging."""
|
||||
result = super().should_run_guardrail(data, event_type)
|
||||
# Check if apply_guardrail is in __dict__
|
||||
has_apply_guardrail = "apply_guardrail" in type(self).__dict__
|
||||
verbose_proxy_logger.debug(
|
||||
"GraySwan DEBUG: should_run_guardrail event_type=%s, result=%s, event_hook=%s, has_apply_guardrail=%s, class=%s",
|
||||
event_type,
|
||||
result,
|
||||
self.event_hook,
|
||||
has_apply_guardrail,
|
||||
type(self).__name__,
|
||||
)
|
||||
return result
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Unified Guardrail Interface (works with ALL endpoints automatically)
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
cache,
|
||||
data: dict,
|
||||
call_type: Literal[
|
||||
"completion",
|
||||
"text_completion",
|
||||
"embeddings",
|
||||
"image_generation",
|
||||
"moderation",
|
||||
"audio_transcription",
|
||||
"pass_through_endpoint",
|
||||
"rerank",
|
||||
"mcp_call",
|
||||
"anthropic_messages",
|
||||
],
|
||||
) -> Optional[Union[Exception, str, dict]]:
|
||||
if (
|
||||
self.should_run_guardrail(
|
||||
data=data, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
is not True
|
||||
):
|
||||
return data
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
"""
|
||||
Apply Gray Swan guardrail to extracted text content.
|
||||
|
||||
verbose_proxy_logger.debug("Gray Swan Guardrail: pre-call hook triggered")
|
||||
This method is called by the unified guardrail system which handles
|
||||
extracting text from any request format (OpenAI, Anthropic, etc.).
|
||||
|
||||
messages = data.get("messages")
|
||||
if not messages:
|
||||
verbose_proxy_logger.debug("Gray Swan Guardrail: No messages in data")
|
||||
return data
|
||||
Args:
|
||||
inputs: Dictionary containing:
|
||||
- texts: List of texts to scan
|
||||
- images: Optional list of images (not currently used by GraySwan)
|
||||
- tool_calls: Optional list of tool calls (not currently used)
|
||||
request_data: The original request data
|
||||
input_type: "request" for pre-call, "response" for post-call
|
||||
logging_obj: Optional logging object
|
||||
|
||||
dynamic_body = self.get_guardrail_dynamic_request_body_params(data) or {}
|
||||
Returns:
|
||||
GenericGuardrailAPIInputs - texts may be replaced with violation message in passthrough mode
|
||||
|
||||
Raises:
|
||||
HTTPException: If content is blocked (block mode)
|
||||
Exception: If guardrail check fails
|
||||
"""
|
||||
# DEBUG: Log when apply_guardrail is called
|
||||
verbose_proxy_logger.debug(
|
||||
"GraySwan DEBUG: apply_guardrail called with input_type=%s, texts=%s",
|
||||
input_type,
|
||||
inputs.get("texts", [])[:100] if inputs.get("texts") else "NONE",
|
||||
)
|
||||
|
||||
texts = inputs.get("texts", [])
|
||||
if not texts:
|
||||
verbose_proxy_logger.debug("Gray Swan Guardrail: No texts to scan")
|
||||
return inputs
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"Gray Swan Guardrail: Scanning %d text(s) for %s",
|
||||
len(texts),
|
||||
input_type,
|
||||
)
|
||||
|
||||
# Convert texts to messages format for GraySwan API
|
||||
# Use "user" role for request content, "assistant" for response content
|
||||
role = "assistant" if input_type == "response" else "user"
|
||||
messages = [{"role": role, "content": text} for text in texts]
|
||||
|
||||
# Get dynamic params from request metadata
|
||||
dynamic_body = self.get_guardrail_dynamic_request_body_params(request_data) or {}
|
||||
|
||||
# Prepare and send payload
|
||||
payload = self._prepare_payload(messages, dynamic_body)
|
||||
if payload is None:
|
||||
verbose_proxy_logger.debug(
|
||||
"Gray Swan Guardrail: no content to scan; skipping request"
|
||||
)
|
||||
return data
|
||||
return inputs
|
||||
|
||||
await self.run_grayswan_guardrail(payload, data, GuardrailEventHooks.pre_call)
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=data, guardrail_name=self.guardrail_name
|
||||
# Call GraySwan API
|
||||
response_json = await self._call_grayswan_api(payload)
|
||||
# Process response
|
||||
is_output = input_type == "response"
|
||||
result = self._process_response_internal(
|
||||
response_json=response_json,
|
||||
request_data=request_data,
|
||||
inputs=inputs,
|
||||
is_output=is_output,
|
||||
)
|
||||
return data
|
||||
|
||||
@log_guardrail_information
|
||||
async def async_moderation_hook(
|
||||
self,
|
||||
data: dict,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
call_type: Literal[
|
||||
"completion",
|
||||
"embeddings",
|
||||
"image_generation",
|
||||
"moderation",
|
||||
"audio_transcription",
|
||||
"responses",
|
||||
"mcp_call",
|
||||
"anthropic_messages",
|
||||
],
|
||||
) -> Optional[Union[Exception, str, dict]]:
|
||||
if (
|
||||
self.should_run_guardrail(
|
||||
data=data, event_type=GuardrailEventHooks.during_call
|
||||
)
|
||||
is not True
|
||||
):
|
||||
return data
|
||||
|
||||
verbose_proxy_logger.debug("GraySwan Guardrail: during-call hook triggered")
|
||||
|
||||
messages = data.get("messages")
|
||||
if not messages:
|
||||
verbose_proxy_logger.debug("Gray Swan Guardrail: No messages in data")
|
||||
return data
|
||||
|
||||
dynamic_body = self.get_guardrail_dynamic_request_body_params(data) or {}
|
||||
|
||||
payload = self._prepare_payload(messages, dynamic_body)
|
||||
if payload is None:
|
||||
verbose_proxy_logger.debug(
|
||||
"Gray Swan Guardrail: no content to scan; skipping request"
|
||||
)
|
||||
return data
|
||||
|
||||
await self.run_grayswan_guardrail(
|
||||
payload, data, GuardrailEventHooks.during_call
|
||||
)
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=data, guardrail_name=self.guardrail_name
|
||||
)
|
||||
return data
|
||||
|
||||
@log_guardrail_information
|
||||
async def async_post_call_success_hook(
|
||||
self,
|
||||
data: dict,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: LLMResponseTypes,
|
||||
) -> LLMResponseTypes:
|
||||
if (
|
||||
self.should_run_guardrail(
|
||||
data=data, event_type=GuardrailEventHooks.post_call
|
||||
)
|
||||
is not True
|
||||
):
|
||||
return response
|
||||
|
||||
verbose_proxy_logger.debug("GraySwan Guardrail: post-call hook triggered")
|
||||
|
||||
response_dict = response.model_dump() if hasattr(response, "model_dump") else {} # type: ignore[union-attr]
|
||||
response_messages = [
|
||||
msg if isinstance(msg, dict) else msg.model_dump()
|
||||
for choice in response_dict.get("choices", [])
|
||||
if isinstance(choice, dict)
|
||||
for msg in [choice.get("message")]
|
||||
if msg is not None
|
||||
]
|
||||
|
||||
if not response_messages:
|
||||
verbose_proxy_logger.debug(
|
||||
"Gray Swan Guardrail: no response messages detected; skipping post-call scan"
|
||||
)
|
||||
return response
|
||||
|
||||
dynamic_body = self.get_guardrail_dynamic_request_body_params(data) or {}
|
||||
|
||||
payload = self._prepare_payload(response_messages, dynamic_body)
|
||||
if payload is None:
|
||||
verbose_proxy_logger.debug(
|
||||
"Gray Swan Guardrail: no content to scan; skipping request"
|
||||
)
|
||||
return response
|
||||
|
||||
await self.run_grayswan_guardrail(payload, data, GuardrailEventHooks.post_call)
|
||||
|
||||
# If passthrough mode and detection info exists, replace response content with violation message
|
||||
if self.on_flagged_action == "passthrough" and "metadata" in data:
|
||||
guardrail_detections = data.get("metadata", {}).get(
|
||||
"guardrail_detections", []
|
||||
)
|
||||
if guardrail_detections:
|
||||
# Replace the model response content with guardrail violation message
|
||||
violation_message = self._format_violation_message(
|
||||
guardrail_detections, is_output=True
|
||||
)
|
||||
|
||||
# Handle ModelResponse (OpenAI-style chat/text completions)
|
||||
# Use isinstance to narrow the type for mypy
|
||||
if isinstance(response, ModelResponse) and response.choices:
|
||||
verbose_proxy_logger.debug(
|
||||
"Gray Swan Guardrail: Replacing response content in ModelResponse format"
|
||||
)
|
||||
for choice in response.choices:
|
||||
# Handle chat completion format (message.content)
|
||||
# Choices has message attribute, StreamingChoices has delta
|
||||
if isinstance(choice, Choices) and hasattr(choice, "message") and hasattr(
|
||||
choice.message, "content"
|
||||
):
|
||||
choice.message.content = violation_message
|
||||
# Handle text completion format (text)
|
||||
# Text attribute might be set dynamically, use setattr
|
||||
elif hasattr(choice, "text"):
|
||||
setattr(choice, "text", violation_message)
|
||||
|
||||
# Update finish_reason to indicate content filtering
|
||||
if hasattr(choice, "finish_reason"):
|
||||
choice.finish_reason = "content_filter"
|
||||
|
||||
# Handle AnthropicMessagesResponse format
|
||||
elif hasattr(response, "content") and isinstance(response.content, list): # type: ignore
|
||||
verbose_proxy_logger.debug(
|
||||
"Gray Swan Guardrail: Replacing response content in Anthropic Messages format"
|
||||
)
|
||||
# Replace content blocks with text block containing violation message
|
||||
response.content = [ # type: ignore
|
||||
{"type": "text", "text": violation_message}
|
||||
]
|
||||
# Update stop_reason if present
|
||||
if hasattr(response, "stop_reason"):
|
||||
response.stop_reason = "end_turn" # type: ignore
|
||||
|
||||
else:
|
||||
verbose_proxy_logger.warning(
|
||||
"Gray Swan Guardrail: Passthrough mode enabled but response format not recognized. "
|
||||
"Cannot replace content. Response type: %s",
|
||||
type(response).__name__,
|
||||
)
|
||||
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=data, guardrail_name=self.guardrail_name
|
||||
)
|
||||
return response
|
||||
return result
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Core GraySwan interaction
|
||||
# Legacy Test Interface (for backward compatibility)
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def run_grayswan_guardrail(
|
||||
async def run_grayswan_guardrail(self, payload: dict) -> Dict[str, Any]:
|
||||
"""
|
||||
Run the GraySwan guardrail on a payload.
|
||||
|
||||
This is a legacy method for testing purposes.
|
||||
|
||||
Args:
|
||||
payload: The payload to scan
|
||||
|
||||
Returns:
|
||||
Dict containing the GraySwan API response
|
||||
"""
|
||||
response_json = await self._call_grayswan_api(payload)
|
||||
# Call the legacy response processor (for test compatibility)
|
||||
self._process_grayswan_response(response_json)
|
||||
return response_json
|
||||
|
||||
def _process_grayswan_response(
|
||||
self,
|
||||
payload: dict,
|
||||
response_json: dict,
|
||||
data: Optional[dict] = None,
|
||||
hook_type: Optional[GuardrailEventHooks] = None,
|
||||
):
|
||||
) -> None:
|
||||
"""
|
||||
Legacy method for processing GraySwan API responses.
|
||||
|
||||
This method is maintained for backward compatibility with existing tests.
|
||||
It handles the test scenarios where responses need to be processed with
|
||||
knowledge of the request context (pre/during/post call hooks).
|
||||
|
||||
Args:
|
||||
response_json: Response from GraySwan API
|
||||
data: Optional request data (for passthrough exceptions)
|
||||
hook_type: Optional GuardrailEventHooks for determining behavior
|
||||
"""
|
||||
violation_score = float(response_json.get("violation", 0.0) or 0.0)
|
||||
violated_rules = response_json.get("violated_rules", [])
|
||||
mutation_detected = response_json.get("mutation")
|
||||
ipi_detected = response_json.get("ipi")
|
||||
|
||||
flagged = violation_score >= self.violation_threshold
|
||||
if not flagged:
|
||||
verbose_proxy_logger.debug(
|
||||
"Gray Swan Guardrail: content passed (score=%s, threshold=%s)",
|
||||
violation_score,
|
||||
self.violation_threshold,
|
||||
)
|
||||
return
|
||||
|
||||
verbose_proxy_logger.warning(
|
||||
"Gray Swan Guardrail: violation score %.3f exceeds threshold %.3f",
|
||||
violation_score,
|
||||
self.violation_threshold,
|
||||
)
|
||||
|
||||
detection_info = {
|
||||
"guardrail": "grayswan",
|
||||
"flagged": True,
|
||||
"violation_score": violation_score,
|
||||
"violated_rules": violated_rules,
|
||||
"mutation": mutation_detected,
|
||||
"ipi": ipi_detected,
|
||||
}
|
||||
|
||||
# Determine if this is input (pre-call/during-call) or output (post-call)
|
||||
if hook_type is not None:
|
||||
is_input = hook_type in [
|
||||
GuardrailEventHooks.pre_call,
|
||||
GuardrailEventHooks.during_call,
|
||||
]
|
||||
else:
|
||||
is_input = True
|
||||
|
||||
if self.on_flagged_action == "block":
|
||||
violation_location = "output" if (not is_input) else "input"
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Blocked by Gray Swan Guardrail",
|
||||
"violation_location": violation_location,
|
||||
"violation": violation_score,
|
||||
"violated_rules": violated_rules,
|
||||
"mutation": mutation_detected,
|
||||
"ipi": ipi_detected,
|
||||
},
|
||||
)
|
||||
elif self.on_flagged_action == "passthrough":
|
||||
# For passthrough mode, we need to handle violations
|
||||
detections = [detection_info]
|
||||
violation_message = self._format_violation_message(
|
||||
detections, is_output=not is_input
|
||||
)
|
||||
verbose_proxy_logger.info(
|
||||
"Gray Swan Guardrail: Passthrough mode - handling violation"
|
||||
)
|
||||
|
||||
# If hook_type is provided and in pre/during call, raise exception
|
||||
if hook_type in [GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call]:
|
||||
# Raise ModifyResponseException to short-circuit LLM call
|
||||
if data is None:
|
||||
data = {}
|
||||
self.raise_passthrough_exception(
|
||||
violation_message=violation_message,
|
||||
request_data=data,
|
||||
detection_info=detection_info,
|
||||
)
|
||||
elif hook_type == GuardrailEventHooks.post_call:
|
||||
# For post-call, store detection info in metadata
|
||||
if data is None:
|
||||
data = {}
|
||||
if "metadata" not in data:
|
||||
data["metadata"] = {}
|
||||
if "guardrail_detections" not in data["metadata"]:
|
||||
data["metadata"]["guardrail_detections"] = []
|
||||
data["metadata"]["guardrail_detections"].append(detection_info)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Core GraySwan API interaction
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _call_grayswan_api(self, payload: dict) -> Dict[str, Any]:
|
||||
"""Call the GraySwan monitoring API."""
|
||||
headers = self._prepare_headers()
|
||||
|
||||
try:
|
||||
|
|
@ -326,15 +356,107 @@ class GraySwanGuardrail(CustomGuardrail):
|
|||
verbose_proxy_logger.debug(
|
||||
"Gray Swan Guardrail: monitor response %s", safe_dumps(result)
|
||||
)
|
||||
return result
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc: # pragma: no cover - depends on HTTP client behaviour
|
||||
except Exception as exc:
|
||||
verbose_proxy_logger.exception(
|
||||
"Gray Swan Guardrail: API request failed: %s", exc
|
||||
)
|
||||
raise GraySwanGuardrailAPIError(str(exc)) from exc
|
||||
|
||||
self._process_grayswan_response(result, data, hook_type)
|
||||
def _process_response_internal(
|
||||
self,
|
||||
response_json: Dict[str, Any],
|
||||
request_data: dict,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
is_output: bool,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
"""
|
||||
Process GraySwan API response and handle violations.
|
||||
|
||||
Args:
|
||||
response_json: Response from GraySwan API
|
||||
request_data: Original request data
|
||||
inputs: The inputs being scanned
|
||||
is_output: True if scanning model output, False for input
|
||||
|
||||
Returns:
|
||||
GenericGuardrailAPIInputs - possibly modified with violation message
|
||||
|
||||
Raises:
|
||||
HTTPException: If content is blocked (block mode)
|
||||
"""
|
||||
violation_score = float(response_json.get("violation", 0.0) or 0.0)
|
||||
violated_rules = response_json.get("violated_rule_descriptions", [])
|
||||
mutation_detected = response_json.get("mutation")
|
||||
ipi_detected = response_json.get("ipi")
|
||||
|
||||
flagged = violation_score >= self.violation_threshold
|
||||
if not flagged:
|
||||
verbose_proxy_logger.debug(
|
||||
"Gray Swan Guardrail: content passed (score=%s, threshold=%s)",
|
||||
violation_score,
|
||||
self.violation_threshold,
|
||||
)
|
||||
return inputs
|
||||
|
||||
verbose_proxy_logger.warning(
|
||||
"Gray Swan Guardrail: violation score %.3f exceeds threshold %.3f",
|
||||
violation_score,
|
||||
self.violation_threshold,
|
||||
)
|
||||
|
||||
detection_info = {
|
||||
"guardrail": "grayswan",
|
||||
"flagged": True,
|
||||
"violation_score": violation_score,
|
||||
"violated_rules": violated_rules,
|
||||
"mutation": mutation_detected,
|
||||
"ipi": ipi_detected,
|
||||
}
|
||||
|
||||
if self.on_flagged_action == "block":
|
||||
violation_location = "output" if is_output else "input"
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Blocked by Gray Swan Guardrail",
|
||||
"violation_location": violation_location,
|
||||
"violation": violation_score,
|
||||
"violated_rules": violated_rules,
|
||||
"mutation": mutation_detected,
|
||||
"ipi": ipi_detected,
|
||||
},
|
||||
)
|
||||
elif self.on_flagged_action == "monitor":
|
||||
verbose_proxy_logger.info(
|
||||
"Gray Swan Guardrail: Monitoring mode - allowing flagged content"
|
||||
)
|
||||
return inputs
|
||||
elif self.on_flagged_action == "passthrough":
|
||||
# Replace content with violation message
|
||||
violation_message = self._format_violation_message(
|
||||
detection_info, is_output=is_output
|
||||
)
|
||||
verbose_proxy_logger.info(
|
||||
"Gray Swan Guardrail: Passthrough mode - replacing content with violation message"
|
||||
)
|
||||
|
||||
if not is_output:
|
||||
# For pre-call (request), raise exception to short-circuit LLM call
|
||||
# and return synthetic response with violation message
|
||||
self.raise_passthrough_exception(
|
||||
violation_message=violation_message,
|
||||
request_data=request_data,
|
||||
detection_info=detection_info,
|
||||
)
|
||||
|
||||
# For post-call (response), replace texts and let unified system apply them
|
||||
inputs["texts"] = [violation_message]
|
||||
return inputs
|
||||
|
||||
return inputs
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Helpers
|
||||
|
|
@ -348,10 +470,9 @@ class GraySwanGuardrail(CustomGuardrail):
|
|||
}
|
||||
|
||||
def _prepare_payload(
|
||||
self, messages: list[dict], dynamic_body: dict
|
||||
self, messages: List[Dict[str, str]], dynamic_body: dict
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
payload: Dict[str, Any] = {}
|
||||
payload["messages"] = messages
|
||||
payload: Dict[str, Any] = {"messages": messages}
|
||||
|
||||
categories = dynamic_body.get("categories") or self.categories
|
||||
if categories:
|
||||
|
|
@ -367,128 +488,41 @@ class GraySwanGuardrail(CustomGuardrail):
|
|||
|
||||
return payload
|
||||
|
||||
def _process_grayswan_response(
|
||||
self,
|
||||
response_json: Dict[str, Any],
|
||||
data: Optional[dict] = None,
|
||||
hook_type: Optional[GuardrailEventHooks] = None,
|
||||
) -> None:
|
||||
violation_score = float(response_json.get("violation", 0.0) or 0.0)
|
||||
violated_rules = response_json.get("violated_rules", [])
|
||||
mutation_detected = response_json.get("mutation")
|
||||
ipi_detected = response_json.get("ipi")
|
||||
|
||||
flagged = violation_score >= self.violation_threshold
|
||||
if not flagged:
|
||||
verbose_proxy_logger.debug(
|
||||
"Gray Swan Guardrail: request passed (score=%s, rules=%s)",
|
||||
violation_score,
|
||||
violated_rules,
|
||||
)
|
||||
return
|
||||
|
||||
verbose_proxy_logger.warning(
|
||||
"Gray Swan Guardrail: violation score %.3f exceeds threshold %.3f",
|
||||
violation_score,
|
||||
self.violation_threshold,
|
||||
)
|
||||
|
||||
if self.on_flagged_action == "block":
|
||||
# Determine if violation was in input or output
|
||||
violation_location = (
|
||||
"output"
|
||||
if hook_type == GuardrailEventHooks.post_call
|
||||
else "input"
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Blocked by Gray Swan Guardrail",
|
||||
"violation_location": violation_location,
|
||||
"violation": violation_score,
|
||||
"violated_rules": violated_rules,
|
||||
"mutation": mutation_detected,
|
||||
"ipi": ipi_detected,
|
||||
},
|
||||
)
|
||||
elif self.on_flagged_action == "monitor":
|
||||
verbose_proxy_logger.info(
|
||||
"Gray Swan Guardrail: Monitoring mode - allowing flagged content to proceed"
|
||||
)
|
||||
elif self.on_flagged_action == "passthrough":
|
||||
# Store detection info
|
||||
detection_info = {
|
||||
"guardrail": "grayswan",
|
||||
"flagged": True,
|
||||
"violation_score": violation_score,
|
||||
"violated_rules": violated_rules,
|
||||
"mutation": mutation_detected,
|
||||
"ipi": ipi_detected,
|
||||
}
|
||||
|
||||
# For pre_call and during_call, raise exception to short-circuit LLM call
|
||||
if hook_type in (
|
||||
GuardrailEventHooks.pre_call,
|
||||
GuardrailEventHooks.during_call,
|
||||
):
|
||||
verbose_proxy_logger.info(
|
||||
"Gray Swan Guardrail: Passthrough mode - raising exception to short-circuit LLM call"
|
||||
)
|
||||
violation_message = self._format_violation_message(
|
||||
[detection_info], is_output=False
|
||||
)
|
||||
self.raise_passthrough_exception(
|
||||
violation_message=violation_message,
|
||||
request_data=data or {},
|
||||
detection_info=detection_info,
|
||||
)
|
||||
|
||||
# For post_call, store in metadata to replace response later
|
||||
verbose_proxy_logger.info(
|
||||
"Gray Swan Guardrail: Passthrough mode - storing detection info in metadata"
|
||||
)
|
||||
if data is not None:
|
||||
if "metadata" not in data:
|
||||
data["metadata"] = {}
|
||||
if "guardrail_detections" not in data["metadata"]:
|
||||
data["metadata"]["guardrail_detections"] = []
|
||||
data["metadata"]["guardrail_detections"].append(detection_info)
|
||||
|
||||
def _format_violation_message(
|
||||
self, guardrail_detections: list, is_output: bool = False
|
||||
self, detection_info: Any, is_output: bool = False
|
||||
) -> str:
|
||||
"""
|
||||
Format guardrail detections into a user-friendly violation message.
|
||||
Format detection info into a user-friendly violation message.
|
||||
|
||||
Args:
|
||||
guardrail_detections: List of detection info dictionaries
|
||||
is_output: True if violation is in model output (post_call), False if in input (pre_call/during_call)
|
||||
detection_info: Can be either:
|
||||
- A single dict with violation_score, violated_rules, mutation, ipi keys
|
||||
- A list of such dicts (legacy format)
|
||||
is_output: True if violation is in model output, False if in input
|
||||
|
||||
Returns:
|
||||
Formatted violation message string
|
||||
"""
|
||||
if not guardrail_detections:
|
||||
return "Content was flagged by guardrail"
|
||||
# Handle legacy format where detection_info is a list
|
||||
if isinstance(detection_info, list) and len(detection_info) > 0:
|
||||
detection_info = detection_info[0]
|
||||
|
||||
violation_score = detection_info.get("violation_score", 0.0)
|
||||
violated_rules = detection_info.get("violated_rules", [])
|
||||
mutation = detection_info.get("mutation", False)
|
||||
ipi = detection_info.get("ipi", False)
|
||||
|
||||
# Get the most recent detection (should be from this guardrail)
|
||||
detection = guardrail_detections[-1]
|
||||
|
||||
violation_score = detection.get("violation_score", 0.0)
|
||||
violated_rules = detection.get("violated_rules", [])
|
||||
mutation = detection.get("mutation", False)
|
||||
ipi = detection.get("ipi", False)
|
||||
|
||||
# Indicate whether violation was in input or output
|
||||
violation_location = "the model response" if is_output else "input query"
|
||||
|
||||
message_parts = [
|
||||
f"Sorry I can't help with that. According to the Gray Swan Cygnal Guardrail, the {violation_location} has a violation score of {violation_score:.2f}.",
|
||||
f"Sorry I can't help with that. According to the Gray Swan Cygnal Guardrail, "
|
||||
f"the {violation_location} has a violation score of {violation_score:.2f}.",
|
||||
]
|
||||
|
||||
if violated_rules:
|
||||
message_parts.append(
|
||||
f"It was violating the rule(s): {', '.join(map(str, violated_rules))}."
|
||||
)
|
||||
formatted_rules = self._format_violated_rules(violated_rules)
|
||||
if formatted_rules:
|
||||
message_parts.append(f"It was violating the rule(s): {formatted_rules}.")
|
||||
|
||||
if mutation:
|
||||
message_parts.append(
|
||||
|
|
@ -496,31 +530,51 @@ class GraySwanGuardrail(CustomGuardrail):
|
|||
)
|
||||
|
||||
if ipi:
|
||||
message_parts.append("Indirect Prompt Injection was DETECTED.")
|
||||
message_parts.append(
|
||||
"Indirect Prompt Injection was DETECTED."
|
||||
)
|
||||
|
||||
return "\n".join(message_parts)
|
||||
|
||||
def _resolve_threshold(self, threshold: Optional[float]) -> float:
|
||||
if threshold is not None:
|
||||
return min(max(threshold, 0.0), 1.0)
|
||||
def _format_violated_rules(self, violated_rules: List) -> str:
|
||||
"""Format violated rules list into a readable string."""
|
||||
formatted: List[str] = []
|
||||
for rule in violated_rules:
|
||||
if isinstance(rule, dict):
|
||||
# New format: {'rule': 6, 'name': 'Illegal Activities...', 'description': '...'}
|
||||
rule_num = rule.get("rule", "")
|
||||
rule_name = rule.get("name", "")
|
||||
rule_desc = rule.get("description", "")
|
||||
if rule_num and rule_name:
|
||||
if rule_desc:
|
||||
formatted.append(f"#{rule_num} {rule_name}: {rule_desc}")
|
||||
else:
|
||||
formatted.append(f"#{rule_num} {rule_name}")
|
||||
elif rule_name:
|
||||
formatted.append(rule_name)
|
||||
else:
|
||||
formatted.append(str(rule))
|
||||
else:
|
||||
# Legacy format: simple value
|
||||
formatted.append(str(rule))
|
||||
|
||||
return ", ".join(formatted)
|
||||
|
||||
def _resolve_threshold(self, value: Optional[float]) -> float:
|
||||
if value is not None:
|
||||
return float(value)
|
||||
env_val = os.getenv("GRAYSWAN_VIOLATION_THRESHOLD")
|
||||
if env_val:
|
||||
try:
|
||||
return float(env_val)
|
||||
except ValueError:
|
||||
pass
|
||||
return 0.5
|
||||
|
||||
def _resolve_reasoning_mode(self, candidate: Optional[str]) -> Optional[str]:
|
||||
if candidate is None:
|
||||
return None
|
||||
normalised = candidate.strip().lower()
|
||||
if normalised in self.SUPPORTED_REASONING_MODES:
|
||||
return normalised
|
||||
verbose_proxy_logger.warning(
|
||||
"Gray Swan Guardrail: ignoring unsupported reasoning_mode '%s'",
|
||||
candidate,
|
||||
)
|
||||
def _resolve_reasoning_mode(self, value: Optional[str]) -> Optional[str]:
|
||||
if value and value.lower() in self.SUPPORTED_REASONING_MODES:
|
||||
return value.lower()
|
||||
env_val = os.getenv("GRAYSWAN_REASONING_MODE")
|
||||
if env_val and env_val.lower() in self.SUPPORTED_REASONING_MODES:
|
||||
return env_val.lower()
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def get_config_model():
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.grayswan import (
|
||||
GraySwanGuardrailConfigModel,
|
||||
)
|
||||
|
||||
return GraySwanGuardrailConfigModel
|
||||
|
|
|
|||
|
|
@ -180,7 +180,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
call_type: Optional[CallTypesLiteral] = None
|
||||
if user_api_key_dict.request_route is not None:
|
||||
call_types = get_call_types_for_route(user_api_key_dict.request_route)
|
||||
if call_types is not None:
|
||||
if call_types is not None and len(call_types) > 0:
|
||||
call_type = call_types[0]
|
||||
if call_type is None:
|
||||
call_type = _infer_call_type(call_type=None, completion_response=response)
|
||||
|
|
@ -213,7 +213,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
|
||||
return response
|
||||
|
||||
async def async_post_call_streaming_iterator_hook(
|
||||
async def async_post_call_streaming_iterator_hook( # noqa: PLR0915
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: Any,
|
||||
|
|
@ -238,19 +238,36 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
"guardrail_to_apply", None
|
||||
)
|
||||
|
||||
# Get sampling rate from guardrail config or optional_params, default to 5
|
||||
# Get streaming configuration from guardrail or optional_params
|
||||
sampling_rate = 5
|
||||
end_of_stream_only = False # If True, only apply guardrail at end of stream
|
||||
|
||||
if guardrail_to_apply is not None:
|
||||
# Check guardrail config first
|
||||
guardrail_config = getattr(guardrail_to_apply, "guardrail_config", {})
|
||||
sampling_rate = guardrail_config.get(
|
||||
"streaming_sampling_rate", sampling_rate
|
||||
# Check direct attributes on guardrail first
|
||||
sampling_rate = getattr(
|
||||
guardrail_to_apply, "streaming_sampling_rate", sampling_rate
|
||||
)
|
||||
end_of_stream_only = getattr(
|
||||
guardrail_to_apply, "streaming_end_of_stream_only", end_of_stream_only
|
||||
)
|
||||
|
||||
# Also check guardrail_config dict if present
|
||||
guardrail_config = getattr(guardrail_to_apply, "guardrail_config", {})
|
||||
if isinstance(guardrail_config, dict):
|
||||
sampling_rate = guardrail_config.get(
|
||||
"streaming_sampling_rate", sampling_rate
|
||||
)
|
||||
end_of_stream_only = guardrail_config.get(
|
||||
"streaming_end_of_stream_only", end_of_stream_only
|
||||
)
|
||||
|
||||
# Also check optional_params as fallback
|
||||
sampling_rate = self.optional_params.get(
|
||||
"streaming_sampling_rate", sampling_rate
|
||||
)
|
||||
end_of_stream_only = self.optional_params.get(
|
||||
"streaming_end_of_stream_only", end_of_stream_only
|
||||
)
|
||||
|
||||
if guardrail_to_apply is None:
|
||||
async for item in response:
|
||||
|
|
@ -306,6 +323,11 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
yield remaining_item
|
||||
return
|
||||
|
||||
# If end_of_stream_only mode, yield chunks without processing
|
||||
if end_of_stream_only:
|
||||
yield item
|
||||
continue
|
||||
|
||||
# Process chunk based on sampling rate
|
||||
if chunk_counter % sampling_rate == 0:
|
||||
|
||||
|
|
|
|||
|
|
@ -19,6 +19,10 @@ from litellm.types.guardrails import (
|
|||
LitellmParams,
|
||||
SupportedGuardrailIntegrations,
|
||||
)
|
||||
from litellm.proxy.guardrails.guardrail_hooks.grayswan import (
|
||||
GraySwanGuardrail,
|
||||
initialize_guardrail as initialize_grayswan,
|
||||
)
|
||||
|
||||
from .guardrail_initializers import (
|
||||
initialize_bedrock,
|
||||
|
|
@ -36,9 +40,12 @@ guardrail_initializer_registry = {
|
|||
SupportedGuardrailIntegrations.PRESIDIO.value: initialize_presidio,
|
||||
SupportedGuardrailIntegrations.HIDE_SECRETS.value: initialize_hide_secrets,
|
||||
SupportedGuardrailIntegrations.TOOL_PERMISSION.value: initialize_tool_permission,
|
||||
SupportedGuardrailIntegrations.GRAYSWAN.value: initialize_grayswan,
|
||||
}
|
||||
|
||||
guardrail_class_registry: Dict[str, Type[CustomGuardrail]] = {}
|
||||
guardrail_class_registry: Dict[str, Type[CustomGuardrail]] = {
|
||||
SupportedGuardrailIntegrations.GRAYSWAN.value: GraySwanGuardrail
|
||||
}
|
||||
|
||||
|
||||
def get_guardrail_initializer_from_hooks():
|
||||
|
|
|
|||
|
|
@ -1,12 +1,16 @@
|
|||
import asyncio
|
||||
import time
|
||||
from typing import Any, AsyncIterator, cast
|
||||
from uuid import uuid4
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import ModifyResponseException
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse
|
||||
from litellm.types.responses.main import DeleteResponseResult
|
||||
|
||||
router = APIRouter()
|
||||
|
|
@ -169,6 +173,28 @@ async def responses_api(
|
|||
user_api_base=user_api_base,
|
||||
version=version,
|
||||
)
|
||||
except ModifyResponseException as e:
|
||||
# Guardrail passthrough: return violation message in Responses API format (200)
|
||||
_data = e.request_data
|
||||
await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
original_exception=e,
|
||||
request_data=_data,
|
||||
)
|
||||
|
||||
violation_text = e.message
|
||||
response_obj = ResponsesAPIResponse(
|
||||
id=f"resp_{uuid4()}",
|
||||
object="response",
|
||||
created_at=int(time.time()),
|
||||
model=e.model or data.get("model"),
|
||||
output=cast(Any, [{"content": [{"type": "text", "text": violation_text}]}]),
|
||||
status="completed",
|
||||
usage=ResponseAPIUsage(
|
||||
input_tokens=0, output_tokens=0, total_tokens=0
|
||||
),
|
||||
)
|
||||
return response_obj
|
||||
except Exception as e:
|
||||
raise await processor._handle_llm_api_exception(
|
||||
e=e,
|
||||
|
|
|
|||
|
|
@ -7258,6 +7258,8 @@ class ProviderConfigManager:
|
|||
return litellm.AzureOpenAIGPT5Config()
|
||||
return litellm.AzureOpenAIConfig()
|
||||
elif litellm.LlmProviders.AZURE_AI == provider:
|
||||
if "claude" in model.lower():
|
||||
return litellm.AzureAnthropicConfig()
|
||||
return litellm.AzureAIStudioConfig()
|
||||
elif litellm.LlmProviders.AZURE_TEXT == provider:
|
||||
return litellm.AzureOpenAITextConfig()
|
||||
|
|
|
|||
|
|
@ -1124,6 +1124,124 @@ def test_get_custom_labels_from_metadata_tags(monkeypatch):
|
|||
assert get_custom_labels_from_metadata(metadata) == {}
|
||||
|
||||
|
||||
def test_get_custom_labels_from_top_level_metadata(monkeypatch):
|
||||
"""
|
||||
Test that get_custom_labels_from_metadata can extract fields from top-level metadata,
|
||||
such as requester_ip_address, not just from nested dictionaries like requester_metadata.
|
||||
"""
|
||||
monkeypatch.setattr(
|
||||
"litellm.custom_prometheus_metadata_labels",
|
||||
["requester_ip_address", "user_api_key_alias"],
|
||||
)
|
||||
# Simulate metadata structure with top-level fields
|
||||
metadata = {
|
||||
"requester_ip_address": "10.48.203.20", # Top-level field
|
||||
"user_api_key_alias": "TestAlias", # Top-level field
|
||||
"requester_metadata": {"nested_field": "nested_value"}, # Nested dict (excluded)
|
||||
"user_api_key_auth_metadata": {"another_nested": "value"}, # Nested dict (excluded)
|
||||
}
|
||||
result = get_custom_labels_from_metadata(metadata)
|
||||
assert result == {
|
||||
"requester_ip_address": "10.48.203.20",
|
||||
"user_api_key_alias": "TestAlias",
|
||||
}
|
||||
|
||||
|
||||
def test_get_custom_labels_from_top_level_and_nested_metadata(monkeypatch):
|
||||
"""
|
||||
Test that get_custom_labels_from_metadata can extract fields from both top-level
|
||||
and nested metadata (requester_metadata, user_api_key_auth_metadata).
|
||||
"""
|
||||
monkeypatch.setattr(
|
||||
"litellm.custom_prometheus_metadata_labels",
|
||||
[
|
||||
"requester_ip_address", # Top-level
|
||||
"metadata.foo", # From requester_metadata
|
||||
"metadata.bar", # From user_api_key_auth_metadata
|
||||
],
|
||||
)
|
||||
# Simulate combined_metadata structure as it would appear after merging
|
||||
# This is what gets passed to get_custom_labels_from_metadata
|
||||
combined_metadata = {
|
||||
"requester_ip_address": "10.48.203.20", # Top-level field
|
||||
"foo": "bar_value", # From requester_metadata (spread)
|
||||
"bar": "baz_value", # From user_api_key_auth_metadata (spread)
|
||||
}
|
||||
result = get_custom_labels_from_metadata(combined_metadata)
|
||||
assert result == {
|
||||
"requester_ip_address": "10.48.203.20",
|
||||
"metadata_foo": "bar_value",
|
||||
"metadata_bar": "baz_value",
|
||||
}
|
||||
|
||||
|
||||
async def test_async_log_success_event_with_top_level_metadata(prometheus_logger, monkeypatch):
|
||||
"""
|
||||
Test that async_log_success_event correctly extracts custom labels from top-level metadata
|
||||
fields like requester_ip_address, not just from nested dictionaries.
|
||||
"""
|
||||
# Configure custom metadata labels to extract requester_ip_address
|
||||
monkeypatch.setattr(
|
||||
"litellm.custom_prometheus_metadata_labels", ["requester_ip_address"]
|
||||
)
|
||||
|
||||
# Create standard logging payload with requester_ip_address at top-level metadata
|
||||
standard_logging_object = create_standard_logging_payload()
|
||||
standard_logging_object["metadata"]["requester_ip_address"] = "10.48.203.20"
|
||||
standard_logging_object["metadata"]["requester_metadata"] = {} # Empty nested dict
|
||||
standard_logging_object["metadata"]["user_api_key_auth_metadata"] = {} # Empty nested dict
|
||||
|
||||
kwargs = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"stream": True,
|
||||
"litellm_params": {
|
||||
"metadata": {
|
||||
"user_api_key": "test_key",
|
||||
"user_api_key_user_id": "test_user",
|
||||
"user_api_key_team_id": "test_team",
|
||||
"user_api_key_end_user_id": "test_end_user",
|
||||
}
|
||||
},
|
||||
"start_time": datetime.now(),
|
||||
"completion_start_time": datetime.now(),
|
||||
"api_call_start_time": datetime.now(),
|
||||
"end_time": datetime.now() + timedelta(seconds=1),
|
||||
"standard_logging_object": standard_logging_object,
|
||||
}
|
||||
response_obj = MagicMock()
|
||||
|
||||
# Mock the prometheus client methods
|
||||
prometheus_logger.litellm_requests_metric = MagicMock()
|
||||
prometheus_logger.litellm_spend_metric = MagicMock()
|
||||
prometheus_logger.litellm_tokens_metric = MagicMock()
|
||||
prometheus_logger.litellm_input_tokens_metric = MagicMock()
|
||||
prometheus_logger.litellm_output_tokens_metric = MagicMock()
|
||||
prometheus_logger.litellm_remaining_team_budget_metric = MagicMock()
|
||||
prometheus_logger.litellm_remaining_api_key_budget_metric = MagicMock()
|
||||
prometheus_logger.litellm_remaining_api_key_requests_for_model = MagicMock()
|
||||
prometheus_logger.litellm_remaining_api_key_tokens_for_model = MagicMock()
|
||||
prometheus_logger.litellm_llm_api_time_to_first_token_metric = MagicMock()
|
||||
prometheus_logger.litellm_llm_api_latency_metric = MagicMock()
|
||||
prometheus_logger.litellm_request_total_latency_metric = MagicMock()
|
||||
|
||||
await prometheus_logger.async_log_success_event(
|
||||
kwargs, response_obj, kwargs["start_time"], kwargs["end_time"]
|
||||
)
|
||||
|
||||
# Verify that the metrics were called with labels including requester_ip_address
|
||||
# Check that labels() was called - the actual labels dict should include requester_ip_address
|
||||
assert prometheus_logger.litellm_requests_metric.labels.called
|
||||
assert prometheus_logger.litellm_spend_metric.labels.called
|
||||
|
||||
# Get the actual call arguments to verify requester_ip_address is included
|
||||
# The custom labels should be extracted and included in the label factory
|
||||
call_args = prometheus_logger.litellm_requests_metric.labels.call_args
|
||||
assert call_args is not None
|
||||
# The labels() method receives a dict with label names and values
|
||||
# We can't easily assert the exact values without checking the internal implementation,
|
||||
# but we've verified the function is called, which means the extraction happened
|
||||
|
||||
|
||||
def test_get_custom_labels_from_tags(monkeypatch):
|
||||
from litellm.integrations.prometheus import get_custom_labels_from_tags
|
||||
|
||||
|
|
|
|||
|
|
@ -311,3 +311,56 @@ def test_get_internal_user_header_from_mapping_no_internal_returns_none():
|
|||
single_mapping = {"header_name": "X-Only-Customer", "litellm_user_role": "customer"}
|
||||
result = LiteLLMProxyRequestSetup.get_internal_user_header_from_mapping(single_mapping)
|
||||
assert result is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"request_data, route, expected_model",
|
||||
[
|
||||
# Vertex AI passthrough URL patterns
|
||||
(
|
||||
{},
|
||||
"/vertex_ai/v1/projects/my-project/locations/us-central1/publishers/google/models/gemini-1.5-pro:generateContent",
|
||||
"gemini-1.5-pro"
|
||||
),
|
||||
(
|
||||
{},
|
||||
"/vertex_ai/v1beta1/projects/my-project/locations/us-central1/publishers/google/models/gemini-1.0-pro:streamGenerateContent",
|
||||
"gemini-1.0-pro"
|
||||
),
|
||||
(
|
||||
{},
|
||||
"/vertex_ai/v1/projects/my-project/locations/asia-southeast1/publishers/google/models/gemini-2.0-flash:generateContent",
|
||||
"gemini-2.0-flash"
|
||||
),
|
||||
# Model without method suffix (no colon) - should still extract
|
||||
(
|
||||
{},
|
||||
"/vertex_ai/v1/projects/my-project/locations/us-central1/publishers/google/models/gemini-pro",
|
||||
"gemini-pro" # Should match even without colon
|
||||
),
|
||||
# Request body model takes precedence over URL
|
||||
(
|
||||
{"model": "gpt-4o"},
|
||||
"/vertex_ai/v1/projects/my-project/locations/us-central1/publishers/google/models/gemini-1.5-pro:generateContent",
|
||||
"gpt-4o"
|
||||
),
|
||||
# Non-vertex route should not extract from vertex pattern
|
||||
(
|
||||
{},
|
||||
"/openai/v1/chat/completions",
|
||||
None
|
||||
),
|
||||
# Azure deployment pattern should still work
|
||||
(
|
||||
{},
|
||||
"/openai/deployments/my-deployment/chat/completions",
|
||||
"my-deployment"
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_get_model_from_request_vertex_ai_passthrough(request_data, route, expected_model):
|
||||
"""Test that get_model_from_request correctly extracts Vertex AI model from URL"""
|
||||
from litellm.proxy.auth.auth_utils import get_model_from_request
|
||||
|
||||
model = get_model_from_request(request_data, route)
|
||||
assert model == expected_model
|
||||
|
|
|
|||
170
tests/test_litellm/images/test_image_edit_utils.py
Normal file
170
tests/test_litellm/images/test_image_edit_utils.py
Normal file
|
|
@ -0,0 +1,170 @@
|
|||
from typing import Any, Dict, List
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.images.utils import ImageEditRequestUtils
|
||||
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
|
||||
from litellm.types.images.main import ImageEditOptionalRequestParams
|
||||
|
||||
|
||||
class MockImageEditConfig(BaseImageEditConfig):
|
||||
def get_supported_openai_params(self, model: str) -> List[str]:
|
||||
return ["size", "quality"]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
image_edit_optional_params: ImageEditOptionalRequestParams,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> Dict[str, Any]:
|
||||
return dict(image_edit_optional_params)
|
||||
|
||||
def get_complete_url(
|
||||
self, model: str, api_base: str, litellm_params: dict
|
||||
) -> str:
|
||||
return "https://example.com/api"
|
||||
|
||||
def validate_environment(
|
||||
self, headers: dict, model: str, api_key: str = None
|
||||
) -> dict:
|
||||
return headers
|
||||
|
||||
def transform_image_edit_request(self, *args, **kwargs):
|
||||
return {}, []
|
||||
|
||||
def transform_image_edit_response(self, *args, **kwargs):
|
||||
return MagicMock()
|
||||
|
||||
|
||||
class TestImageEditRequestUtilsDropParams:
|
||||
def setup_method(self):
|
||||
self.config = MockImageEditConfig()
|
||||
self.model = "test-model"
|
||||
self._original_drop_params = getattr(litellm, "drop_params", None)
|
||||
|
||||
def teardown_method(self):
|
||||
if self._original_drop_params is None:
|
||||
if hasattr(litellm, "drop_params"):
|
||||
delattr(litellm, "drop_params")
|
||||
else:
|
||||
litellm.drop_params = self._original_drop_params
|
||||
|
||||
def test_unsupported_params_raises_without_drop(self):
|
||||
litellm.drop_params = False
|
||||
optional_params: ImageEditOptionalRequestParams = {
|
||||
"size": "1024x1024",
|
||||
"unsupported_param": "value",
|
||||
}
|
||||
|
||||
with pytest.raises(litellm.UnsupportedParamsError) as exc_info:
|
||||
ImageEditRequestUtils.get_optional_params_image_edit(
|
||||
model=self.model,
|
||||
image_edit_provider_config=self.config,
|
||||
image_edit_optional_params=optional_params,
|
||||
)
|
||||
|
||||
assert "unsupported_param" in str(exc_info.value)
|
||||
|
||||
def test_drop_params_global_setting(self):
|
||||
litellm.drop_params = True
|
||||
optional_params: ImageEditOptionalRequestParams = {
|
||||
"size": "1024x1024",
|
||||
"unsupported_param": "value",
|
||||
}
|
||||
|
||||
result = ImageEditRequestUtils.get_optional_params_image_edit(
|
||||
model=self.model,
|
||||
image_edit_provider_config=self.config,
|
||||
image_edit_optional_params=optional_params,
|
||||
)
|
||||
|
||||
assert "size" in result
|
||||
assert "unsupported_param" not in result
|
||||
|
||||
def test_drop_params_explicit_parameter(self):
|
||||
litellm.drop_params = False
|
||||
optional_params: ImageEditOptionalRequestParams = {
|
||||
"size": "1024x1024",
|
||||
"unsupported_param": "value",
|
||||
}
|
||||
|
||||
result = ImageEditRequestUtils.get_optional_params_image_edit(
|
||||
model=self.model,
|
||||
image_edit_provider_config=self.config,
|
||||
image_edit_optional_params=optional_params,
|
||||
drop_params=True,
|
||||
)
|
||||
|
||||
assert "size" in result
|
||||
assert "unsupported_param" not in result
|
||||
|
||||
def test_additional_drop_params(self):
|
||||
litellm.drop_params = False
|
||||
optional_params: ImageEditOptionalRequestParams = {
|
||||
"size": "1024x1024",
|
||||
"quality": "high",
|
||||
}
|
||||
|
||||
result = ImageEditRequestUtils.get_optional_params_image_edit(
|
||||
model=self.model,
|
||||
image_edit_provider_config=self.config,
|
||||
image_edit_optional_params=optional_params,
|
||||
additional_drop_params=["quality"],
|
||||
)
|
||||
|
||||
assert "size" in result
|
||||
assert "quality" not in result
|
||||
|
||||
def test_drop_params_false_with_global_true(self):
|
||||
litellm.drop_params = True
|
||||
optional_params: ImageEditOptionalRequestParams = {
|
||||
"size": "1024x1024",
|
||||
"unsupported_param": "value",
|
||||
}
|
||||
|
||||
result = ImageEditRequestUtils.get_optional_params_image_edit(
|
||||
model=self.model,
|
||||
image_edit_provider_config=self.config,
|
||||
image_edit_optional_params=optional_params,
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert "size" in result
|
||||
assert "unsupported_param" not in result
|
||||
|
||||
def test_supported_params_pass_through(self):
|
||||
litellm.drop_params = False
|
||||
optional_params: ImageEditOptionalRequestParams = {
|
||||
"size": "1024x1024",
|
||||
"quality": "high",
|
||||
}
|
||||
|
||||
result = ImageEditRequestUtils.get_optional_params_image_edit(
|
||||
model=self.model,
|
||||
image_edit_provider_config=self.config,
|
||||
image_edit_optional_params=optional_params,
|
||||
)
|
||||
|
||||
assert result["size"] == "1024x1024"
|
||||
assert result["quality"] == "high"
|
||||
|
||||
def test_additional_drop_params_with_unsupported_and_drop_true(self):
|
||||
litellm.drop_params = True
|
||||
optional_params: ImageEditOptionalRequestParams = {
|
||||
"size": "1024x1024",
|
||||
"quality": "high",
|
||||
"unsupported_param": "value",
|
||||
}
|
||||
|
||||
result = ImageEditRequestUtils.get_optional_params_image_edit(
|
||||
model=self.model,
|
||||
image_edit_provider_config=self.config,
|
||||
image_edit_optional_params=optional_params,
|
||||
additional_drop_params=["quality"],
|
||||
)
|
||||
|
||||
assert "size" in result
|
||||
assert "quality" not in result
|
||||
assert "unsupported_param" not in result
|
||||
|
|
@ -140,6 +140,55 @@ class TestVertexAIGeminiImageEditTransformation:
|
|||
headers={},
|
||||
)
|
||||
|
||||
def test_get_complete_url_from_litellm_params(self) -> None:
|
||||
"""Test vertex_project/vertex_location read from litellm_params first"""
|
||||
url = self.config.get_complete_url(
|
||||
model="gemini-2.5-flash",
|
||||
api_base=None,
|
||||
litellm_params={
|
||||
"vertex_project": "params-project",
|
||||
"vertex_location": "us-east1",
|
||||
},
|
||||
)
|
||||
assert "params-project" in url
|
||||
assert "us-east1" in url
|
||||
|
||||
def test_get_complete_url_global_location(self) -> None:
|
||||
"""Test global location uses correct base URL without region prefix"""
|
||||
url = self.config.get_complete_url(
|
||||
model="gemini-2.5-flash",
|
||||
api_base=None,
|
||||
litellm_params={
|
||||
"vertex_project": "test-project",
|
||||
"vertex_location": "global",
|
||||
},
|
||||
)
|
||||
assert "aiplatform.googleapis.com" in url
|
||||
assert "global-aiplatform.googleapis.com" not in url
|
||||
assert "/locations/global/" in url
|
||||
|
||||
def test_get_complete_url_litellm_params_overrides_env(self) -> None:
|
||||
"""Test litellm_params takes precedence over environment variables"""
|
||||
with patch.dict(
|
||||
os.environ,
|
||||
{
|
||||
"VERTEXAI_PROJECT": "env-project",
|
||||
"VERTEXAI_LOCATION": "us-central1",
|
||||
},
|
||||
):
|
||||
url = self.config.get_complete_url(
|
||||
model="gemini-2.5-flash",
|
||||
api_base=None,
|
||||
litellm_params={
|
||||
"vertex_project": "params-project",
|
||||
"vertex_location": "eu-west1",
|
||||
},
|
||||
)
|
||||
assert "params-project" in url
|
||||
assert "eu-west1" in url
|
||||
assert "env-project" not in url
|
||||
assert "us-central1" not in url
|
||||
|
||||
|
||||
class TestVertexAIImagenImageEditTransformation:
|
||||
def setup_method(self) -> None:
|
||||
|
|
|
|||
|
|
@ -2602,3 +2602,30 @@ class TestIsCachedMessage:
|
|||
"""Empty list content should return False."""
|
||||
message = {"role": "user", "content": []}
|
||||
assert is_cached_message(message) is False
|
||||
|
||||
|
||||
def test_azure_ai_claude_provider_config():
|
||||
"""Test that Azure AI Claude models return AzureAnthropicConfig for proper tool transformation."""
|
||||
from litellm import AzureAnthropicConfig, AzureAIStudioConfig
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
# Claude models should return AzureAnthropicConfig
|
||||
config = ProviderConfigManager.get_provider_chat_config(
|
||||
model="claude-sonnet-4-5",
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
)
|
||||
assert isinstance(config, AzureAnthropicConfig)
|
||||
|
||||
# Test case-insensitive matching
|
||||
config = ProviderConfigManager.get_provider_chat_config(
|
||||
model="Claude-Opus-4",
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
)
|
||||
assert isinstance(config, AzureAnthropicConfig)
|
||||
|
||||
# Non-Claude models should return AzureAIStudioConfig
|
||||
config = ProviderConfigManager.get_provider_chat_config(
|
||||
model="mistral-large",
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
)
|
||||
assert isinstance(config, AzureAIStudioConfig)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue