diff --git a/.circleci/config.yml b/.circleci/config.yml index 38fdaf3609d..cf69ff68da6 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -2541,7 +2541,6 @@ jobs: paths: - litellm-docker-database.tar.zst - test_bad_database_url: machine: image: ubuntu-2204:2024.04.1 diff --git a/litellm/constants.py b/litellm/constants.py index e36746326cc..fb765c0226c 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1443,6 +1443,12 @@ CLI_JWT_EXPIRATION_HOURS = int( or os.getenv("LITELLM_CLI_JWT_EXPIRATION_HOURS") or 24 ) +# Comma-separated allowlisted OIDC claim map for CLI SSO polling, e.g. +# "employment_type->acme_employment_type,org_info.department->department" +CLI_SSO_CLAIM_MAP = ( + os.getenv("CLI_SSO_CLAIM_MAP") or os.getenv("LITELLM_CLI_SSO_CLAIM_MAP") or "" +) +CLI_SSO_CLAIM_MAX_SCALAR_LENGTH = 1024 ########################### UI SESSION DURATION ########################### # Duration for UI login session (username/password, SSO, invitation links). Format: "30s", "30m", "24h", "7d" diff --git a/litellm/litellm_core_utils/get_supported_openai_params.py b/litellm/litellm_core_utils/get_supported_openai_params.py index 9d8bd7523db..b8cdc8210fc 100644 --- a/litellm/litellm_core_utils/get_supported_openai_params.py +++ b/litellm/litellm_core_utils/get_supported_openai_params.py @@ -11,6 +11,7 @@ def get_supported_openai_params( # noqa: PLR0915 request_type: Literal[ "chat_completion", "embeddings", "transcription" ] = "chat_completion", + base_model: Optional[str] = None, ) -> Optional[list]: """ Returns the supported openai params for a given model + provider @@ -20,6 +21,11 @@ def get_supported_openai_params( # noqa: PLR0915 get_supported_openai_params(model="anthropic.claude-3", custom_llm_provider="bedrock") ``` + Args: + base_model: For Azure, the true underlying model (e.g. ``"azure/gpt-5.2"``) + when the deployment name differs. Used for model-type detection so that + non-standard deployment names route to the correct config. + Returns: - List if custom_llm_provider is mapped - None if unmapped @@ -32,17 +38,21 @@ def get_supported_openai_params( # noqa: PLR0915 if custom_llm_provider in LlmProvidersSet: provider_config = litellm.ProviderConfigManager.get_provider_chat_config( - model=model, provider=LlmProviders(custom_llm_provider) + model=model, + provider=LlmProviders(custom_llm_provider), + base_model=base_model, ) elif custom_llm_provider.split("/")[0] in LlmProvidersSet: provider_config = litellm.ProviderConfigManager.get_provider_chat_config( - model=model, provider=LlmProviders(custom_llm_provider.split("/")[0]) + model=model, + provider=LlmProviders(custom_llm_provider.split("/")[0]), + base_model=base_model, ) else: provider_config = None if provider_config and request_type == "chat_completion": - return provider_config.get_supported_openai_params(model=model) + return provider_config.get_supported_openai_params(model=base_model or model) if custom_llm_provider == "bedrock": return litellm.AmazonConverseConfig().get_supported_openai_params(model=model) @@ -130,16 +140,23 @@ def get_supported_openai_params( # noqa: PLR0915 model=model ) elif custom_llm_provider == "azure": - if litellm.AzureOpenAIO1Config().is_o_series_model(model=model): + _azure_detection_model = base_model or model + if litellm.AzureOpenAIO1Config().is_o_series_model( + model=_azure_detection_model + ): return litellm.AzureOpenAIO1Config().get_supported_openai_params( - model=model + model=_azure_detection_model ) - elif litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(model=model): + elif litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model( + model=_azure_detection_model + ): return litellm.AzureOpenAIGPT5Config().get_supported_openai_params( - model=model + model=_azure_detection_model ) else: - return litellm.AzureOpenAIConfig().get_supported_openai_params(model=model) + return litellm.AzureOpenAIConfig().get_supported_openai_params( + model=_azure_detection_model + ) elif custom_llm_provider == "openrouter": return litellm.OpenrouterConfig().get_supported_openai_params(model=model) elif custom_llm_provider == "vercel_ai_gateway": diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 3ee56dfc5ca..32ae61d7f58 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -1204,12 +1204,8 @@ def get_last_user_message(messages: List[AllMessageValues]) -> Optional[str]: {"role": "assistant", "content": "I'm good, thank you!"}, {"role": "user", "content": "What is the weather in Tokyo?"}, ] - get_user_prompt(messages) -> "What is the weather in Tokyo?" + get_last_user_message(messages) -> "What is the weather in Tokyo?" """ - from litellm.litellm_core_utils.prompt_templates.common_utils import ( - convert_content_list_to_str, - ) - if not messages: return None diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index 0e198daf089..51a1e739a0f 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -1476,7 +1476,7 @@ class LiteLLMAnthropicMessagesAdapter: for choice in choices: if choice.delta.content is not None and len(choice.delta.content) > 0: text += choice.delta.content - if choice.delta.tool_calls is not None: + if choice.delta.tool_calls: partial_json = "" for tool in choice.delta.tool_calls: if ( diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index 9291269d153..734b8ecef16 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -239,7 +239,9 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): ) data = {"model": None, "messages": messages, **optional_params} - elif litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(model=model): + elif litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model( + model=litellm_params.get("base_model") or model + ): data = litellm.AzureOpenAIGPT5Config().transform_request( model=model, messages=messages, diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index 9dd2b055a12..b659c1b0a0a 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -44,6 +44,12 @@ else: # (e.g. "us-east-1", "eu-west-2", "us-gov-west-1", "cn-north-1"). _VALID_AWS_REGION_PATTERN = re.compile(r"\A[a-z0-9-]+\Z") +# Regional STS hostnames, e.g. sts.eu-west-1.amazonaws.com or +# vpce-xxx.sts.eu-west-1.vpce.amazonaws.com +_STS_REGION_FROM_ENDPOINT_PATTERN = re.compile( + r"(?:^|\.)sts(?:-fips)?\.([a-z0-9-]+)\.(?:amazonaws\.com(?:\.cn)?|vpce\.amazonaws\.com)" +) + class Boto3CredentialsInfo(BaseModel): credentials: Credentials @@ -450,6 +456,24 @@ class BaseAWSLLM: model_id = BaseAWSLLM.encode_model_id(model_id=model_id) else: model_id = model + # Strip LiteLLM routing prefixes (e.g. "bedrock/", "invoke/", + # "bedrock/invoke/", "bedrock/converse/") that are not part of the + # actual Bedrock model ID. The converse path already does this; the + # invoke path must do the same so that ARN models such as + # bedrock/arn:aws:bedrock:…:inference-profile/global.anthropic.… + # are not forwarded verbatim to the Bedrock API, which would produce + # a malformed URL and cause botocore's EventStreamBuffer to receive + # a JSON error body instead of a binary event-stream — surfaced as a + # misleading ChecksumMismatch (0x223a7b22 == ':{"'). + # Use strip_bedrock_routing_prefix (no break) so compound prefixes + # like "bedrock/invoke/arn:..." are fully stripped in one call. + from litellm.llms.bedrock.common_utils import strip_bedrock_routing_prefix + + model_id = strip_bedrock_routing_prefix(model_id) + # URL-encode ARNs so colons and slashes are safe in the URL path. + if model_id.startswith("arn:"): + model_id = BaseAWSLLM.encode_model_id(model_id=model_id) + return model_id model_id = model_id.replace("invoke/", "", 1) if provider == "llama" and "llama/" in model_id: @@ -633,6 +657,40 @@ class BaseAWSLLM: "Region names must contain only lowercase letters, digits, and hyphens." ) + @staticmethod + def _parse_sts_region_from_endpoint( + aws_sts_endpoint: Optional[str], + ) -> Optional[str]: + """Extract region from sts.{region}.amazonaws.com or vpce-x.sts.{region}.vpce.amazonaws.com.""" + if not aws_sts_endpoint: + return None + host = urllib.parse.urlparse(aws_sts_endpoint).hostname or "" + match = _STS_REGION_FROM_ENDPOINT_PATTERN.search(host) + return match.group(1) if match else None + + @staticmethod + def _resolve_sts_region(aws_sts_endpoint: Optional[str] = None) -> Optional[str]: + """STS signing region: parsed from aws_sts_endpoint else AWS_REGION / AWS_DEFAULT_REGION.""" + return ( + BaseAWSLLM._parse_sts_region_from_endpoint(aws_sts_endpoint) + or os.getenv("AWS_REGION") + or os.getenv("AWS_DEFAULT_REGION") + ) + + def _build_sts_client_kwargs( + self, + aws_sts_endpoint: Optional[str] = None, + ssl_verify: Optional[Union[bool, str]] = None, + ) -> dict: + """STS client kwargs with aligned endpoint_url and region_name (SigV4).""" + kwargs: dict = {"verify": self._get_ssl_verify(ssl_verify)} + if aws_sts_endpoint is not None: + kwargs["endpoint_url"] = aws_sts_endpoint + sts_region = self._resolve_sts_region(aws_sts_endpoint) + if sts_region is not None: + kwargs["region_name"] = sts_region + return kwargs + def get_aws_region_name_for_non_llm_api_calls( self, aws_region_name: Optional[str] = None, @@ -787,11 +845,6 @@ class BaseAWSLLM: f"IN Web Identity Token: {aws_web_identity_token} | Role Name: {aws_role_name} | Session Name: {aws_session_name}" ) - if aws_sts_endpoint is None: - sts_endpoint = f"https://sts.{aws_region_name}.amazonaws.com" - else: - sts_endpoint = aws_sts_endpoint - oidc_token = get_secret(aws_web_identity_token) if oidc_token is None: @@ -800,13 +853,13 @@ class BaseAWSLLM: status_code=401, ) + sts_client_kwargs = self._build_sts_client_kwargs( + aws_sts_endpoint=aws_sts_endpoint, + ssl_verify=ssl_verify, + ) + with tracer.trace("boto3.client(sts)"): - sts_client = boto3.client( - "sts", - region_name=aws_region_name, - endpoint_url=sts_endpoint, - verify=self._get_ssl_verify(ssl_verify), - ) + sts_client = boto3.client("sts", **sts_client_kwargs) # https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRoleWithWebIdentity.html # https://boto3.amazonaws.com/v1/documentation/api/latest/reference/services/sts/client/assume_role_with_web_identity.html @@ -847,7 +900,6 @@ class BaseAWSLLM: irsa_role_arn: str, aws_role_name: str, aws_session_name: str, - region: str, web_identity_token_file: str, aws_external_id: Optional[str] = None, aws_sts_endpoint: Optional[str] = None, @@ -862,12 +914,10 @@ class BaseAWSLLM: with open(web_identity_token_file, "r") as f: web_identity_token = f.read().strip() - irsa_sts_kwargs: dict = { - "region_name": region, - "verify": self._get_ssl_verify(ssl_verify), - } - if aws_sts_endpoint is not None: - irsa_sts_kwargs["endpoint_url"] = aws_sts_endpoint + irsa_sts_kwargs = self._build_sts_client_kwargs( + aws_sts_endpoint=aws_sts_endpoint, + ssl_verify=ssl_verify, + ) # Create an STS client without credentials with tracer.trace("boto3.client(sts) for manual IRSA"): @@ -924,7 +974,6 @@ class BaseAWSLLM: self, aws_role_name: str, aws_session_name: str, - region: str, aws_external_id: Optional[str] = None, aws_sts_endpoint: Optional[str] = None, ssl_verify: Optional[Union[bool, str]] = None, @@ -932,12 +981,10 @@ class BaseAWSLLM: """Handle same-account role assumption for IRSA.""" import boto3 - irsa_sts_kwargs: dict = { - "region_name": region, - "verify": self._get_ssl_verify(ssl_verify), - } - if aws_sts_endpoint is not None: - irsa_sts_kwargs["endpoint_url"] = aws_sts_endpoint + irsa_sts_kwargs = self._build_sts_client_kwargs( + aws_sts_endpoint=aws_sts_endpoint, + ssl_verify=ssl_verify, + ) verbose_logger.debug("Same account role assumption, using automatic IRSA") with tracer.trace("boto3.client(sts) with automatic IRSA"): @@ -1010,12 +1057,6 @@ class BaseAWSLLM: web_identity_token_file = os.getenv("AWS_WEB_IDENTITY_TOKEN_FILE") irsa_role_arn = os.getenv("AWS_ROLE_ARN") - region = ( - aws_region_name - or os.getenv("AWS_REGION") - or os.getenv("AWS_DEFAULT_REGION") - ) - # If we have IRSA environment variables and no explicit credentials, # we need to use the web identity token flow if ( @@ -1031,16 +1072,12 @@ class BaseAWSLLM: ) try: - # Use passed-in region when set, else env, else default (align with AssumeRole path) - region = region or "us-east-1" - # Check if we need to do cross-account role assumption if aws_role_name != irsa_role_arn: sts_response = self._handle_irsa_cross_account( irsa_role_arn, aws_role_name, aws_session_name, - region, web_identity_token_file, aws_external_id, aws_sts_endpoint=aws_sts_endpoint, @@ -1050,7 +1087,6 @@ class BaseAWSLLM: sts_response = self._handle_irsa_same_account( aws_role_name, aws_session_name, - region, aws_external_id, aws_sts_endpoint=aws_sts_endpoint, ssl_verify=ssl_verify, @@ -1074,11 +1110,10 @@ class BaseAWSLLM: # In EKS/IRSA environments, use ambient credentials (no explicit keys needed) # This allows the web identity token to work automatically - sts_client_kwargs: dict = {"verify": self._get_ssl_verify(ssl_verify)} - if region is not None: - sts_client_kwargs["region_name"] = region - if aws_sts_endpoint is not None: - sts_client_kwargs["endpoint_url"] = aws_sts_endpoint + sts_client_kwargs = self._build_sts_client_kwargs( + aws_sts_endpoint=aws_sts_endpoint, + ssl_verify=ssl_verify, + ) if aws_access_key_id is None and aws_secret_access_key is None: with tracer.trace("boto3.client(sts)"): sts_client = boto3.client("sts", **sts_client_kwargs) diff --git a/litellm/llms/cohere/embed/v1_transformation.py b/litellm/llms/cohere/embed/v1_transformation.py index feca9cb5b88..82c901e7eca 100644 --- a/litellm/llms/cohere/embed/v1_transformation.py +++ b/litellm/llms/cohere/embed/v1_transformation.py @@ -110,15 +110,35 @@ class CohereEmbeddingConfig: additional_args={"complete_input_dict": data}, original_response=response_json, ) + return self._populate_embedding_response( + response_json=response_json, + model_response=model_response, + model=model, + encoding=encoding, + input=input, + ) + + def _populate_embedding_response( + self, + response_json: dict, + model_response: EmbeddingResponse, + model: str, + encoding: Any, + input: list, + ) -> EmbeddingResponse: """ - response + Parse a Cohere embed response body into an OpenAI-style EmbeddingResponse. + + Split out from `_transform_response` so callers that already log + `post_call` themselves (e.g. SageMaker's embedding handler) can reuse + the parsing without triggering a second `post_call`. + + Response shape: { 'object': "list", - 'data': [ - - ] - 'model', - 'usage' + 'data': [...], + 'model', + 'usage', } """ embeddings = response_json["embeddings"] @@ -149,9 +169,6 @@ class CohereEmbeddingConfig: model_response.object = "list" model_response.data = output_data model_response.model = model - input_tokens = 0 - for text in input: - input_tokens += len(encoding.encode(text)) setattr( model_response, diff --git a/litellm/llms/openai/responses/transformation.py b/litellm/llms/openai/responses/transformation.py index b7d5340d8d4..5043d25ee37 100644 --- a/litellm/llms/openai/responses/transformation.py +++ b/litellm/llms/openai/responses/transformation.py @@ -126,9 +126,21 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): litellm_params: GenericLiteLLMParams, headers: dict, ) -> Dict: - """No transform applied since inputs are in OpenAI spec already""" + """Strip Anthropic-only `cache_control` markers before sending to OpenAI. + + OpenAI's Responses API rejects unknown fields on input content blocks + with HTTP 400 ("Unknown parameter: 'input[0].content[0].cache_control'"). + Chat Completions strips these in + `remove_cache_control_flag_from_messages_and_tools`; mirror that here. + """ input = self._validate_input_param(input) + tools = response_api_optional_request_params.get("tools") + input, tools = self.remove_cache_control_flag_from_input_and_tools( + model=model, input=input, tools=tools + ) + if tools is not None: + response_api_optional_request_params["tools"] = tools final_request_params = dict( ResponsesAPIRequestParams( model=model, input=input, **response_api_optional_request_params @@ -137,6 +149,38 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): return final_request_params + def remove_cache_control_flag_from_input_and_tools( + self, + model: str, # allows overrides to selectively run this + input: Union[str, ResponseInputParam], + tools: Optional[List[ALL_RESPONSES_API_TOOL_PARAMS]] = None, + ) -> Tuple[ + Union[str, ResponseInputParam], + Optional[List[ALL_RESPONSES_API_TOOL_PARAMS]], + ]: + """Sibling of `remove_cache_control_flag_from_messages_and_tools` on + the chat path. Strips Anthropic-only `cache_control` markers from + Responses API input content blocks and tools. + + `filter_value_from_dict` mutates each dict in place, so the same + objects are returned. + """ + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + filter_value_from_dict, + ) + + if isinstance(input, list): + for item in input: + if isinstance(item, dict): + filter_value_from_dict(cast(dict, item), "cache_control") + + if tools is not None: + for tool in tools: + if isinstance(tool, dict): + filter_value_from_dict(cast(dict, tool), "cache_control") + + return input, tools + def _validate_input_param( self, input: Union[str, ResponseInputParam] ) -> Union[str, ResponseInputParam]: @@ -604,6 +648,12 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): url = str(parsed_url.copy_with(path=compact_path)) input = self._validate_input_param(input) + tools = response_api_optional_request_params.get("tools") + input, tools = self.remove_cache_control_flag_from_input_and_tools( + model=model, input=input, tools=tools + ) + if tools is not None: + response_api_optional_request_params["tools"] = tools data = dict( ResponsesAPIRequestParams( model=model, input=input, **response_api_optional_request_params diff --git a/litellm/llms/sagemaker/completion/handler.py b/litellm/llms/sagemaker/completion/handler.py index efbb218f575..de7be18e8ba 100644 --- a/litellm/llms/sagemaker/completion/handler.py +++ b/litellm/llms/sagemaker/completion/handler.py @@ -578,7 +578,7 @@ class SagemakerLLM(BaseAWSLLM): logger_fn=None, ): """ - Supports both Huggingface Jumpstart embeddings and Voyage models + Supports Hugging Face (TGI), Voyage, and Cohere embedding endpoints """ ### BOTO3 INIT import boto3 diff --git a/litellm/llms/sagemaker/embedding/cohere_transformation.py b/litellm/llms/sagemaker/embedding/cohere_transformation.py new file mode 100644 index 00000000000..fdb67202ebb --- /dev/null +++ b/litellm/llms/sagemaker/embedding/cohere_transformation.py @@ -0,0 +1,141 @@ +""" +Translate from OpenAI's `/v1/embeddings` to Sagemaker's `/invoke` + +In the native Cohere embed format for self-hosted Cohere endpoints +(AWS Marketplace / JumpStart). Cohere containers expect +`{"texts": [...], "input_type": "..."}` and reject the HuggingFace TGI shape +`{"inputs": [...]}` with `422 EmbedReqV2.inputs is of type string but should +be of type Object`. + +Reference: https://docs.cohere.com/v2/reference/embed +""" + +from typing import TYPE_CHECKING, Any, List, Optional, Union, cast + +if TYPE_CHECKING: + from litellm.types.llms.openai import AllEmbeddingInputValues + +from httpx._models import Headers, Response + +import litellm +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig +from litellm.llms.bedrock.embed.cohere_transformation import ( + BedrockCohereEmbeddingConfig, +) +from litellm.llms.cohere.embed.v1_transformation import CohereEmbeddingConfig +from litellm.types.utils import EmbeddingResponse + +from ..common_utils import SagemakerError + + +class SagemakerCohereEmbeddingConfig(BaseEmbeddingConfig): + """ + SageMaker invoke payload for self-hosted Cohere embed models. + """ + + def __init__(self) -> None: + pass + + def get_supported_openai_params(self, model: str) -> List[str]: + return ["encoding_format", "dimensions", "input_type"] + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + optional_params = BedrockCohereEmbeddingConfig().map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + ) + if "input_type" in non_default_params: + optional_params["input_type"] = non_default_params["input_type"] + return optional_params + + def get_error_class( + self, error_message: str, status_code: int, headers: Union[dict, Headers] + ) -> BaseLLMException: + return SagemakerError( + message=error_message, status_code=status_code, headers=headers + ) + + def transform_embedding_request( + self, + model: str, + input: "AllEmbeddingInputValues", + optional_params: dict, + headers: dict, + ) -> dict: + """ + Transform embedding request for Cohere models on SageMaker + """ + if isinstance(input, str): + input_list: List[str] = [input] + elif isinstance(input, list): + if input and (isinstance(input[0], list) or isinstance(input[0], int)): + raise ValueError("Input must be a list of strings") + input_list = cast(List[str], input) + else: + input_list = [str(input)] + + return dict( + BedrockCohereEmbeddingConfig()._transform_request( + model=model, + input=input_list, + inference_params=optional_params, + ) + ) + + def transform_embedding_response( + self, + model: str, + raw_response: Response, + model_response: "EmbeddingResponse", + logging_obj: Any, + api_key: Optional[str] = None, + request_data: dict = {}, + optional_params: dict = {}, + litellm_params: dict = {}, + ) -> "EmbeddingResponse": + """ + Transform embedding response for Cohere models on SageMaker. + + Uses `CohereEmbeddingConfig._populate_embedding_response` (not + `_transform_response`) so we do not log `post_call` a second time + — the SageMaker embedding handler already logs `post_call` before + invoking this transform. + """ + input_value = ( + logging_obj.model_call_details.get("input") + or request_data.get("texts") + or request_data.get("images") + or [] + ) + if isinstance(input_value, str): + input_value = [input_value] + + return CohereEmbeddingConfig()._populate_embedding_response( + response_json=raw_response.json(), + model_response=model_response, + model=model, + encoding=litellm.encoding, + input=input_value, + ) + + def validate_environment( + self, + headers: dict, + model: str, + messages: List[Any], + optional_params: dict, + litellm_params: dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> dict: + """ + Validate environment for SageMaker Cohere embeddings + """ + return {"Content-Type": "application/json"} diff --git a/litellm/llms/sagemaker/embedding/transformation.py b/litellm/llms/sagemaker/embedding/transformation.py index 09bdb9295e7..5e2aa99534f 100644 --- a/litellm/llms/sagemaker/embedding/transformation.py +++ b/litellm/llms/sagemaker/embedding/transformation.py @@ -11,12 +11,13 @@ if TYPE_CHECKING: from httpx._models import Headers, Response -from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig from litellm.llms.base_llm.chat.transformation import BaseLLMException -from litellm.types.utils import Usage, EmbeddingResponse +from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig from litellm.llms.voyage.embedding.transformation import VoyageEmbeddingConfig +from litellm.types.utils import EmbeddingResponse, Usage from ..common_utils import SagemakerError +from .cohere_transformation import SagemakerCohereEmbeddingConfig class SagemakerEmbeddingConfig(BaseEmbeddingConfig): @@ -38,17 +39,20 @@ class SagemakerEmbeddingConfig(BaseEmbeddingConfig): Returns: Appropriate embedding config instance """ - if "voyage" in model.lower(): + model_lower = model.lower() + if "voyage" in model_lower: return VoyageEmbeddingConfig() - else: - return cls() + if "cohere" in model_lower: + return SagemakerCohereEmbeddingConfig() + return cls() def get_supported_openai_params(self, model: str) -> List[str]: - # Check if this is an embedding model - if "voyage" in model.lower(): + model_lower = model.lower() + if "voyage" in model_lower: return VoyageEmbeddingConfig().get_supported_openai_params(model) - else: - return [] + if "cohere" in model_lower: + return SagemakerCohereEmbeddingConfig().get_supported_openai_params(model) + return [] def map_openai_params( self, diff --git a/litellm/main.py b/litellm/main.py index b5364f8ba17..e17a5ad9a48 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -1491,7 +1491,9 @@ def completion( # type: ignore # noqa: PLR0915 provider.value for provider in LlmProviders ]: provider_config = ProviderConfigManager.get_provider_chat_config( - model=model, provider=LlmProviders(custom_llm_provider) + model=model, + provider=LlmProviders(custom_llm_provider), + base_model=base_model, ) if provider_config is not None: @@ -1550,6 +1552,7 @@ def completion( # type: ignore # noqa: PLR0915 "safety_identifier": safety_identifier, "service_tier": service_tier, "allowed_openai_params": kwargs.get("allowed_openai_params"), + "base_model": base_model, } optional_params = get_optional_params( **optional_param_args, **non_default_params @@ -1670,6 +1673,10 @@ def completion( # type: ignore # noqa: PLR0915 reasoning_summary=_reasoning_summary_for_bridge, ) + # Use base_model (the true underlying model) for Azure model-type + # detection when the deployment name differs from the model name. + _azure_detection_model = base_model or model + if responses_api_model_info.get("mode") == "responses": from litellm.completion_extras import responses_api_bridge @@ -1713,7 +1720,9 @@ def completion( # type: ignore # noqa: PLR0915 and OpenAIGPT5Config.is_model_gpt_5_model(model) ) or ( custom_llm_provider == "azure" - and litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(model) + and litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model( + _azure_detection_model + ) ): optional_params, _ = strip_reasoning_summary_aliases_from_optional_params( optional_params @@ -1766,7 +1775,9 @@ def completion( # type: ignore # noqa: PLR0915 if max_retries is not None: optional_params["max_retries"] = max_retries - if litellm.AzureOpenAIO1Config().is_o_series_model(model=model): + if litellm.AzureOpenAIO1Config().is_o_series_model( + model=_azure_detection_model + ): ## LOAD CONFIG - if set config = litellm.AzureOpenAIO1Config.get_config() for k, v in config.items(): diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index a72e8e34a49..d0e9ad7b2a4 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -12,6 +12,7 @@ import hashlib import json import os import re +import time from typing import Any, Callable, Dict, List, Literal, Optional, Set, Tuple, Union, cast from urllib.parse import urlparse @@ -250,6 +251,10 @@ class MCPServerManager: } """ self._upstream_initialize_instructions_by_server_id: Dict[str, str] = {} + # Per-server monotonic timestamp of last upstream prefetch attempt (success, + # empty result, or failure). Used to throttle re-probes for servers that do + # not return instructions, and to apply a short cooldown after failures. + self._upstream_initialize_instructions_probed_at: Dict[str, float] = {} def _remember_upstream_initialize_instructions( self, server: MCPServer, client: MCPClient @@ -260,6 +265,80 @@ class MCPServerManager: raw ).strip() + async def _ensure_upstream_initialize_instructions_cached( + self, server: MCPServer + ) -> None: + """ + Open one upstream session and cache InitializeResult.instructions if missing. + + No-op when: + - YAML/DB instructions are set on the server record, + - server is OpenAPI (spec_path), + - non-empty upstream instructions are already cached, + - auth preconditions match health_check_server's skip rules + (per-user auth / missing static auth token), + - a prior probe attempt for this server is within + MCP_HEALTH_CHECK_TIMEOUT seconds (the probe is a health-check-shaped + op and already uses this knob for its inner call timeout; reusing it + as the cooldown avoids reconnecting on every gateway initialize when + upstream returns empty or fails). + """ + if server.spec_path: + return + if server.instructions and server.instructions.strip(): + return + if self._upstream_initialize_instructions_by_server_id.get(server.server_id): + return + if server.requires_per_user_auth: + return + if ( + server.auth_type + and server.auth_type != MCPAuth.none + and server.auth_type != MCPAuth.aws_sigv4 + and not server.authentication_token + ): + return + + last_probed_at = self._upstream_initialize_instructions_probed_at.get( + server.server_id + ) + if ( + last_probed_at is not None + and (time.monotonic() - last_probed_at) < MCP_HEALTH_CHECK_TIMEOUT + ): + return + + # Record the attempt up-front so that a failure / empty response does not + # cause every subsequent initialize request to re-open the upstream session. + self._upstream_initialize_instructions_probed_at[server.server_id] = ( + time.monotonic() + ) + + try: + extra_headers: Optional[Dict[str, str]] = ( + dict(server.static_headers) if server.static_headers else None + ) + client = await self._create_mcp_client( + server=server, + mcp_auth_header=None, + extra_headers=extra_headers, + stdio_env=None, + ) + + async def _noop(_session): + return "ok" + + await asyncio.wait_for( + client.run_with_session(_noop), timeout=MCP_HEALTH_CHECK_TIMEOUT + ) + self._remember_upstream_initialize_instructions(server, client) + except Exception as e: + verbose_logger.debug( + "Upstream initialize instructions prefetch failed for %s: %s", + server.name, + e, + ) + def get_registry(self) -> Dict[str, MCPServer]: """ Get the registered MCP Servers from the registry and union with the config MCP Servers @@ -280,6 +359,7 @@ class MCPServerManager: """ verbose_logger.debug("Loading MCP Servers from config-----") self._upstream_initialize_instructions_by_server_id.clear() + self._upstream_initialize_instructions_probed_at.clear() # Track which aliases have been used to ensure only first occurrence is used used_aliases = set() @@ -3141,6 +3221,7 @@ class MCPServerManager: verbose_logger.debug("Loading MCP servers from database into registry...") self._upstream_initialize_instructions_by_server_id.clear() + self._upstream_initialize_instructions_probed_at.clear() # perform authz check to filter the mcp servers user has access to prisma_client = get_prisma_client_or_throw( diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 5205426edf3..f31005be0cb 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1165,7 +1165,7 @@ if MCP_AVAILABLE: def _merge_gateway_initialize_instructions( allowed_mcp_servers: List[MCPServer], ) -> Optional[str]: - """YAML/DB override, else in-memory upstream text from list_tools / health_check / call_tool.""" + """YAML/DB override, else upstream text (prefetch on init, or list_tools / health_check / call_tool cache).""" if not allowed_mcp_servers: return None @@ -1206,6 +1206,20 @@ if MCP_AVAILABLE: mcp_servers=mcp_servers, client_ip=client_ip, ) + if allowed: + # return_exceptions=True: a per-server probe failure (incl. CancelledError + # bubbled from anyio task group teardown on connection refused) must not + # cancel sibling probes or 500 the gateway initialize request. + await asyncio.gather( + *[ + global_mcp_server_manager._ensure_upstream_initialize_instructions_cached( + s + ) + for s in allowed + if s is not None + ], + return_exceptions=True, + ) merged = _merge_gateway_initialize_instructions(allowed_mcp_servers=allowed) tok = _mcp_gateway_initialize_instructions.set(merged) try: diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 09bb8057203..14f198e0f12 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -2353,7 +2353,9 @@ class ExperimentalUIJWTToken: @staticmethod def get_cli_jwt_auth_token( - user_info: LiteLLM_UserTable, team_id: Optional[str] = None + user_info: LiteLLM_UserTable, + team_id: Optional[str] = None, + team_alias: Optional[str] = None, ) -> str: """ Generate a JWT token for CLI authentication with configurable expiration. @@ -2364,6 +2366,7 @@ class ExperimentalUIJWTToken: Args: user_info: User information from the database team_id: Team ID for the user (optional, uses user's team if available) + team_alias: Team alias for the selected team, if available Returns: Encrypted JWT token string @@ -2397,6 +2400,7 @@ class ExperimentalUIJWTToken: expires=expires, user_id=user_info.user_id, team_id=_team_id, + team_alias=team_alias, models=user_info.models, max_parallel_requests=None, user_role=LitellmUserRoles(user_info.user_role), diff --git a/litellm/proxy/client/README.md b/litellm/proxy/client/README.md index adf562d69c5..9fbc6f2197d 100644 --- a/litellm/proxy/client/README.md +++ b/litellm/proxy/client/README.md @@ -350,7 +350,7 @@ The CLI provides three authentication commands: 4. **User Authentication**: User completes SSO authentication in browser 5. **Callback Processing**: SSO provider redirects back to proxy with state parameter 6. **User Code Verification**: Browser confirms the verification code shown in the CLI -7. **Polling**: CLI polls `/sso/cli/poll/{login_id}` with the polling secret header until the JWT is ready +7. **Polling**: CLI polls `/sso/cli/poll/{login_id}` with the polling secret header until the JWT is ready. When `CLI_SSO_CLAIM_MAP` is configured on the proxy, the poll response may include `attribution_metadata` (allowlisted scalar OIDC claims for client attribution). 8. **Token Storage**: CLI saves the authentication token to `~/.litellm/token.json` ### Benefits of This Approach diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/base.py b/litellm/proxy/guardrails/guardrail_hooks/azure/base.py index fab65884a9f..7d2dfce0711 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/base.py @@ -2,6 +2,9 @@ import re from typing import TYPE_CHECKING, Any, Dict, List, Optional from litellm._logging import verbose_proxy_logger +from litellm.litellm_core_utils.prompt_templates.common_utils import ( + get_last_user_message, +) from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, @@ -134,32 +137,4 @@ class AzureGuardrailBase: ] get_user_prompt(messages) -> "What is the weather in Tokyo?" """ - from litellm.litellm_core_utils.prompt_templates.common_utils import ( - convert_content_list_to_str, - ) - - if not messages: - return None - - # Iterate from the end to find the last consecutive block of user messages - user_messages = [] - for message in reversed(messages): - if message.get("role") == "user": - user_messages.append(message) - else: - # Stop when we hit a non-user message - break - - if not user_messages: - return None - - # Reverse to get the messages in chronological order - user_messages.reverse() - - user_prompt = "" - for message in user_messages: - text_content = convert_content_list_to_str(message) - user_prompt += text_content + "\n" - - result = user_prompt.strip() - return result if result else None + return get_last_user_message(messages) diff --git a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/__init__.py new file mode 100644 index 00000000000..9ead2a63b60 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/__init__.py @@ -0,0 +1,57 @@ +from typing import TYPE_CHECKING + +from litellm.types.guardrails import SupportedGuardrailIntegrations + +from .purview_dlp import MicrosoftPurviewDLPGuardrail + +if TYPE_CHECKING: + from litellm.types.guardrails import Guardrail, LitellmParams + + +def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"): + import litellm + + tenant_id = getattr(litellm_params, "tenant_id", None) + client_id = getattr(litellm_params, "client_id", None) + + # client_secret can be passed via the standard api_key field or as + # a dedicated client_secret parameter. + client_secret = litellm_params.api_key or getattr( + litellm_params, "client_secret", None + ) + + if not tenant_id: + raise ValueError("Microsoft Purview: tenant_id is required") + if not client_id: + raise ValueError("Microsoft Purview: client_id is required") + if not client_secret: + raise ValueError("Microsoft Purview: client_secret (or api_key) is required") + + guardrail_name = guardrail.get("guardrail_name") + if not guardrail_name: + raise ValueError("Microsoft Purview: guardrail_name is required") + + purview_guardrail = MicrosoftPurviewDLPGuardrail( + guardrail_name=guardrail_name, + tenant_id=str(tenant_id), + client_id=str(client_id), + client_secret=str(client_secret), + purview_app_name=str( + getattr(litellm_params, "purview_app_name", None) or "LiteLLM" + ), + user_id_field=str(getattr(litellm_params, "user_id_field", None) or "user_id"), + event_hook=litellm_params.mode, + default_on=litellm_params.default_on, + ) + + litellm.logging_callback_manager.add_litellm_callback(purview_guardrail) + return purview_guardrail + + +guardrail_initializer_registry = { + SupportedGuardrailIntegrations.MICROSOFT_PURVIEW.value: initialize_guardrail, +} + +guardrail_class_registry = { + SupportedGuardrailIntegrations.MICROSOFT_PURVIEW.value: MicrosoftPurviewDLPGuardrail, +} diff --git a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py new file mode 100644 index 00000000000..a7ed1d40913 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py @@ -0,0 +1,515 @@ +import threading +import time +import uuid +from collections import OrderedDict +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple + +from litellm._logging import verbose_proxy_logger +from litellm.litellm_core_utils.url_utils import encode_url_path_segment +from litellm.litellm_core_utils.prompt_templates.common_utils import ( + convert_content_list_to_str, +) +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + httpxSpecialProvider, +) + +if TYPE_CHECKING: + from litellm.types.llms.openai import AllMessageValues + +GRAPH_API_BASE = "https://graph.microsoft.com/v1.0" +TOKEN_ENDPOINT_TEMPLATE = ( + "https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/token" +) +GRAPH_SCOPE = "https://graph.microsoft.com/.default" + +# Protection scope cache TTL in seconds (1 hour, per Microsoft recommendation). +SCOPE_CACHE_TTL_SECONDS = 3600.0 + + +class PurviewGuardrailBase: + """ + Base class for Microsoft Purview guardrails. + + Manages OAuth2 client-credentials token acquisition, protection scope + computation with ETag caching, and authenticated POST calls to the + Microsoft Graph API. + """ + + def __init__( + self, + tenant_id: str, + client_id: str, + client_secret: str, + purview_app_name: str = "LiteLLM", + user_id_field: str = "user_id", + **kwargs: Any, + ): + # Forward remaining kwargs to the next class in the MRO + # (typically CustomGuardrail). + super().__init__(**kwargs) + + self.async_handler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.GuardrailCallback + ) + self.tenant_id = tenant_id + self.client_id = client_id + self.client_secret = client_secret + self.purview_app_name = purview_app_name + self.user_id_field = user_id_field + + # Token cache: (access_token, expires_at_epoch) + self._token_cache: Optional[Tuple[str, float]] = None + + # Protection scope cache: user_id -> (etag, scope_response, fetched_at) + # Capped at 1000 entries (LRU eviction) to avoid unbounded growth. + self._scope_cache: OrderedDict[str, Tuple[str, Dict[str, Any], float]] = ( + OrderedDict() + ) + self._scope_cache_maxsize = 1000 + # Use a threading.Lock (not asyncio.Lock) because this lock is acquired + # from both the proxy's main asyncio event loop and from short-lived + # event loops created by the logging_hook thread fallback. In Python + # 3.10+ an asyncio.Lock is bound to the first event loop that acquires + # it and raises RuntimeError from any other loop, which would silently + # break audit logging via the thread fallback. All critical sections + # below are pure in-memory dict ops with no awaits, so a synchronous + # lock is both correct and sufficient. + self._cache_lock = threading.Lock() + + @staticmethod + def _encode_graph_user_id(user_id: str) -> str: + """Percent-encode Entra user id for Graph ``/users/{id}/...`` path segments.""" + return encode_url_path_segment(user_id, field_name="user_id") + + # ------------------------------------------------------------------ + # OAuth2 token management + # ------------------------------------------------------------------ + + async def _get_access_token(self) -> str: + """Acquire or return cached OAuth2 token via client_credentials grant.""" + now = time.time() + with self._cache_lock: + if self._token_cache and self._token_cache[1] > now + 60: + return self._token_cache[0] + + url = TOKEN_ENDPOINT_TEMPLATE.format(tenant_id=self.tenant_id) + data = { + "grant_type": "client_credentials", + "client_id": self.client_id, + "client_secret": self.client_secret, + "scope": GRAPH_SCOPE, + } + response = await self.async_handler.post( + url=url, + data=data, + headers={"Content-Type": "application/x-www-form-urlencoded"}, + ) + response.raise_for_status() + token_data = response.json() + access_token = token_data["access_token"] + expires_in = int(token_data.get("expires_in", 3599)) + # Recompute ``now`` after the await so the expiry reflects when the + # token was actually received, not when the request started. + with self._cache_lock: + self._token_cache = (access_token, time.time() + expires_in) + verbose_proxy_logger.debug( + "Purview: acquired new OAuth2 token (expires_in=%ds)", expires_in + ) + return access_token + + # ------------------------------------------------------------------ + # Graph API helpers + # ------------------------------------------------------------------ + + async def _graph_post( + self, + url: str, + json_body: Dict[str, Any], + extra_headers: Optional[Dict[str, str]] = None, + ) -> Tuple[Dict[str, Any], Dict[str, str]]: + """POST to Graph API with bearer auth. + + Returns: + Tuple of (response_json, response_headers). + """ + token = await self._get_access_token() + headers = { + "Authorization": f"Bearer {token}", + "Content-Type": "application/json", + } + if extra_headers: + headers.update(extra_headers) + + verbose_proxy_logger.debug("Purview Graph POST %s", url) + response = await self.async_handler.post( + url=url, headers=headers, json=json_body + ) + response.raise_for_status() + response_json: Dict[str, Any] = response.json() + response_headers = dict(response.headers) + verbose_proxy_logger.debug("Purview Graph response: %s", response_json) + return response_json, response_headers + + # ------------------------------------------------------------------ + # Protection scopes + # ------------------------------------------------------------------ + + async def _compute_protection_scopes( + self, user_id: str + ) -> Tuple[str, Dict[str, Any]]: + """Call protectionScopes/compute and cache with ETag. + + Returns: + Tuple of (etag, scope_response). + """ + encoded_user_id = self._encode_graph_user_id(user_id) + now = time.time() + + with self._cache_lock: + cached = self._scope_cache.get(user_id) + if cached and (now - cached[2]) < SCOPE_CACHE_TTL_SECONDS: + self._scope_cache.move_to_end(user_id) + return cached[0], cached[1] + + url = ( + f"{GRAPH_API_BASE}/users/{encoded_user_id}" + "/dataSecurityAndGovernance/protectionScopes/compute" + ) + body: Dict[str, Any] = { + "activities": "uploadText,downloadText", + "locations": [ + { + "@odata.type": "microsoft.graph.policyLocationApplication", + "value": self.client_id, + } + ], + } + + response_json, response_headers = await self._graph_post(url, body) + etag = response_headers.get("etag", response_headers.get("ETag", "")) + + # Recompute ``now`` after the await so the TTL reflects when the + # scope response was actually received, not when the request started. + fetched_at = time.time() + with self._cache_lock: + self._scope_cache[user_id] = (etag, response_json, fetched_at) + # Move refreshed entry to the end so it is treated as most-recently-used. + # OrderedDict.__setitem__ preserves existing insertion order for known + # keys, so an explicit move_to_end() call is required. + self._scope_cache.move_to_end(user_id) + # Evict least-recently-used entry when cache exceeds max size. + while len(self._scope_cache) > self._scope_cache_maxsize: + self._scope_cache.popitem(last=False) + return etag, response_json + + # ------------------------------------------------------------------ + # Process content + # ------------------------------------------------------------------ + + async def _process_content( + self, + user_id: str, + text: str, + activity: str, + etag: str, + correlation_id: Optional[str] = None, + ) -> Dict[str, Any]: + """Call processContent for DLP policy evaluation. + + Args: + user_id: Entra object ID of the user. + text: The content to evaluate. + activity: ``"uploadText"`` for prompts, ``"downloadText"`` for responses. + etag: Cached ETag from protectionScopes/compute. + correlation_id: Optional conversation/thread ID. + """ + encoded_user_id = self._encode_graph_user_id(user_id) + url = ( + f"{GRAPH_API_BASE}/users/{encoded_user_id}" + "/dataSecurityAndGovernance/processContent" + ) + body: Dict[str, Any] = { + "contentToProcess": { + "contentEntries": [ + { + "@odata.type": "microsoft.graph.processConversationMetadata", + "identifier": str(uuid.uuid4()), + "content": { + "@odata.type": "microsoft.graph.textContent", + "data": text, + }, + "name": f"{self.purview_app_name} message", + "correlationId": correlation_id or str(uuid.uuid4()), + "sequenceNumber": 0, + "isTruncated": False, + } + ], + "activityMetadata": {"activity": activity}, + "deviceMetadata": {}, + "protectedAppMetadata": { + "name": self.purview_app_name, + "version": "1.0", + "applicationLocation": { + "@odata.type": "microsoft.graph.policyLocationApplication", + "value": self.client_id, + }, + }, + "integratedAppMetadata": { + "name": self.purview_app_name, + "version": "1.0", + }, + } + } + + extra_headers: Dict[str, str] = {} + if etag: + extra_headers["If-None-Match"] = etag + + response_json, _ = await self._graph_post(url, body, extra_headers) + + # If policies changed, invalidate scope cache so next call re-fetches. + if response_json.get("protectionScopeState") == "modified": + with self._cache_lock: + self._scope_cache.pop(user_id, None) + + return response_json + + # ------------------------------------------------------------------ + # User ID resolution + # ------------------------------------------------------------------ + + def _resolve_user_id( + self, data: Dict[str, Any], user_api_key_dict: Any + ) -> Optional[str]: + """Resolve the Entra user object ID from request data or auth context. + + Returns the strongest available identity walking down four sources, in + decreasing trust order: + + 1. ``user_api_key_dict.user_id`` — LiteLLM key / JWT-bound user + 2. ``user_api_key_dict.end_user_id`` — request-derived + 3. ``metadata["user_api_key_user_id"]`` — proxy-injected from the key + 4. ``metadata[user_id_field]`` — caller-supplied + + Used only by blocking-mode resolution to disambiguate "no identity at + all" from "caller supplied an untrusted identity" for the error + message. Neither blocking nor audit DLP feeds the untrusted + fallbacks (2, 4) into Purview itself. + """ + trusted = self._resolve_trusted_user_id(data, user_api_key_dict) + if trusted: + return trusted + + if hasattr(user_api_key_dict, "end_user_id") and user_api_key_dict.end_user_id: + return str(user_api_key_dict.end_user_id) + + metadata = data.get("metadata") or data.get("litellm_metadata") or {} + uid = metadata.get("user_api_key_user_id") + if uid: + return str(uid) + + uid = metadata.get(self.user_id_field) + if uid: + return str(uid) + + return None + + @staticmethod + def _logging_kwargs_metadata(kwargs: Dict[str, Any]) -> Dict[str, Any]: + """Metadata dict from ``model_call_details`` / logging kwargs.""" + litellm_params = kwargs.get("litellm_params") or {} + if not isinstance(litellm_params, dict): + return {} + md = litellm_params.get("metadata") + return md if isinstance(md, dict) else {} + + def _resolve_trusted_user_id( + self, data: Dict[str, Any], user_api_key_dict: Any + ) -> Optional[str]: + """Resolve user ID from API-key/JWT-bound identity for blocking DLP. + + Uses only ``UserAPIKeyAuth.user_id`` (bound on the LiteLLM key or JWT). + Intentionally omits ``UserAPIKeyAuth.end_user_id`` because the proxy sets + it from caller-controlled request fields (``user``, ``metadata.user_id``, + ``safety_identifier``, custom headers, etc.) via + ``get_end_user_id_from_request_body``. + + Also omits ``metadata[user_id_field]`` and + ``metadata["user_api_key_user_id"]`` for the same impersonation risk when + the key has no bound user. + + Returns ``None`` when no authenticated identity is available. Blocking + hooks must fail closed rather than skip the DLP check. + """ + if hasattr(user_api_key_dict, "user_id") and user_api_key_dict.user_id: + return str(user_api_key_dict.user_id) + + return None + + def _resolve_user_id_from_logging_kwargs( + self, kwargs: Dict[str, Any] + ) -> Optional[str]: + """Trusted-identity-only resolver for logging-only hooks. + + Uses only the proxy-injected ``user_api_key_user_id`` (populated from + the API-key/JWT-bound ``UserAPIKeyAuth.user_id`` after the proxy + strips every caller-supplied ``user_api_key_*`` key from the request + metadata). Caller-influenceable sources (``user_api_key_end_user_id``, + ``metadata[user_id_field]``) are not used here so a caller cannot + cause Purview audit records to be written under a victim's identity. + Returns ``None`` when no trusted identity is available so the audit + is skipped rather than misattributed. + """ + md = self._logging_kwargs_metadata(kwargs) + uid = md.get("user_api_key_user_id") or kwargs.get("user_api_key_user_id") + if uid: + return str(uid) + return None + + # ------------------------------------------------------------------ + # Policy action evaluation + # ------------------------------------------------------------------ + + @staticmethod + def _should_block(response: Dict[str, Any]) -> bool: + """Return True if any policyAction requires blocking.""" + for action in response.get("policyActions", []): + odata_type = action.get("@odata.type", "") + action_field = action.get("action", "") + + if "restrictAccessAction" in odata_type or action_field == "restrictAccess": + restriction = action.get("restrictionAction", "") + if restriction == "block": + return True + return False + + # ------------------------------------------------------------------ + # Prompt text for DLP + # ------------------------------------------------------------------ + + @staticmethod + def is_token_id_prompt(prompt: Any) -> bool: + """Return True if ``prompt`` carries OpenAI completions token ids. + + Covers every list shape that ``completion_prompt_to_str`` cannot decode + for Purview, including flat ``list[int]`` (single token-id prompt), + ``list[list[int]]`` (multi-prompt token-id batches), and mixed lists + that include any token-id sub-array. + """ + if not isinstance(prompt, list) or not prompt: + return False + for x in prompt: + if isinstance(x, int): + return True + if isinstance(x, list) and x and any(isinstance(y, int) for y in x): + return True + return False + + @staticmethod + def completion_prompt_to_str(prompt: Any) -> Optional[str]: + """Normalize OpenAI ``/v1/completions`` ``prompt`` for text DLP. + + Supports string prompts and list-of-string prompts. List-of-token-id prompts + are skipped (no plaintext for Purview to evaluate). + """ + if prompt is None: + return None + if isinstance(prompt, str): + stripped = prompt.strip() + return stripped or None + if isinstance(prompt, list) and prompt: + if all(isinstance(x, str) for x in prompt): + joined = "\n".join(s.strip() for s in prompt if isinstance(s, str)) + return joined.strip() or None + if all(isinstance(x, int) for x in prompt): + verbose_proxy_logger.debug( + "Purview DLP: completions prompt is token ids only; skipping text scan" + ) + return None + str_parts = [x for x in prompt if isinstance(x, str)] + if str_parts: + joined = "\n".join(s.strip() for s in str_parts) + return joined.strip() or None + return None + + @staticmethod + def _extract_tool_call_args_from_message(message: Any) -> List[str]: + """Return plaintext arguments strings from tool_calls and function_call fields. + + Covers both the request path (assistant messages in chat histories that + carry tool_calls / function_call) and the response path (model-generated + tool calls returned in a ModelResponse). Both dict-style and object-style + representations are handled. + """ + args: List[str] = [] + + # tool_calls: [{"function": {"arguments": "..."}}] + tool_calls = ( + message.get("tool_calls") + if isinstance(message, dict) + else getattr(message, "tool_calls", None) + ) + if tool_calls: + for tc in tool_calls: + fn = ( + tc.get("function") + if isinstance(tc, dict) + else getattr(tc, "function", None) + ) + if fn is None: + continue + arguments = ( + fn.get("arguments") + if isinstance(fn, dict) + else getattr(fn, "arguments", None) + ) + if isinstance(arguments, str) and arguments.strip(): + args.append(arguments) + + # Legacy function_call: {"arguments": "..."} + function_call = ( + message.get("function_call") + if isinstance(message, dict) + else getattr(message, "function_call", None) + ) + if function_call is not None: + arguments = ( + function_call.get("arguments") + if isinstance(function_call, dict) + else getattr(function_call, "arguments", None) + ) + if isinstance(arguments, str) and arguments.strip(): + args.append(arguments) + + return args + + def get_prompt_text_for_dlp( + self, messages: List["AllMessageValues"] + ) -> Optional[str]: + """Concatenate text from every chat message (all roles) for pre-call DLP. + + Evaluates the same payload the model receives, not only the trailing user + turn. Each message is separated by ``\\n\\n`` so that tokens at message + boundaries are not merged (e.g., ``"end of msg1\\n\\nstart of msg2"`` + rather than ``"end of msg1start of msg2"``), which preserves DLP pattern + detection accuracy across message boundaries. + + Tool-call arguments (``tool_calls[].function.arguments`` and + ``function_call.arguments``) are included alongside message content so + that sensitive data hidden in function arguments is not bypassed. + """ + if not messages: + return None + parts: List[str] = [] + for msg in messages: + segments: List[str] = [] + content = convert_content_list_to_str(message=msg).strip() + if content: + segments.append(content) + segments.extend(self._extract_tool_call_args_from_message(msg)) + combined = "\n".join(segments) + if combined.strip(): + parts.append(combined.strip()) + text = "\n\n".join(parts) + return text or None diff --git a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py new file mode 100644 index 00000000000..ee0bac64d4f --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py @@ -0,0 +1,734 @@ +""" +Microsoft Purview DLP Guardrail for LiteLLM. + +Supports three modes: +- pre_call: Block sensitive data in prompts before they reach the LLM. +- post_call: Block sensitive data in LLM responses. +- logging_only: Log interactions to Purview for audit/compliance without blocking. +""" + +import asyncio +import threading +import uuid +from datetime import datetime +from typing import ( + TYPE_CHECKING, + Any, + AsyncGenerator, + Dict, + List, + Optional, + Tuple, + Type, + Union, + cast, +) + +import httpx +from fastapi import HTTPException + +from litellm._logging import verbose_proxy_logger +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, +) +from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.utils import ( + Choices, + GuardrailStatus, + ModelResponse, + ModelResponseStream, + ResponsesAPIResponse, + TextChoices, + TextCompletionResponse, +) + +from .base import PurviewGuardrailBase + +if TYPE_CHECKING: + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.proxy.guardrails.guardrail_hooks.base import ( + GuardrailConfigModel, + ) + from litellm.types.utils import ( + CallTypesLiteral, + EmbeddingResponse, + ImageResponse, + ) + + +class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): + """ + Microsoft Purview DLP guardrail. + + Evaluates prompts and responses against Microsoft Purview DLP policies + via the Microsoft Graph ``processContent`` API. + """ + + def __init__( + self, + guardrail_name: str, + tenant_id: str, + client_id: str, + client_secret: str, + purview_app_name: str = "LiteLLM", + user_id_field: str = "user_id", + **kwargs: Any, + ): + supported_event_hooks = [ + GuardrailEventHooks.pre_call, + GuardrailEventHooks.post_call, + GuardrailEventHooks.logging_only, + ] + + super().__init__( + tenant_id=tenant_id, + client_id=client_id, + client_secret=client_secret, + purview_app_name=purview_app_name, + user_id_field=user_id_field, + guardrail_name=guardrail_name, + supported_event_hooks=supported_event_hooks, + **kwargs, + ) + self.guardrail_provider = "microsoft_purview" + verbose_proxy_logger.info( + "Initialized Microsoft Purview DLP Guardrail: %s", + guardrail_name, + ) + + @staticmethod + def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + return None # Config model can be added later for UI support + + # ------------------------------------------------------------------ + # Core DLP check + # ------------------------------------------------------------------ + + async def _check_content( + self, + user_id: str, + text: str, + activity: str, + request_data: Dict[str, Any], + block_on_violation: bool = True, + ) -> Dict[str, Any]: + """Evaluate content against Purview DLP policies. + + Args: + user_id: Entra object ID. + text: Content to evaluate. + activity: ``"uploadText"`` or ``"downloadText"``. + request_data: Original request dict (used for logging metadata). + block_on_violation: If False, log only — do not raise. + + Returns: + The processContent response dict. + """ + start_time = datetime.now() + status: GuardrailStatus = "success" + response: Dict[str, Any] = {} + + try: + etag, _ = await self._compute_protection_scopes(user_id) + correlation_id = request_data.get("litellm_call_id") or str(uuid.uuid4()) + response = await self._process_content( + user_id=user_id, + text=text, + activity=activity, + etag=etag, + correlation_id=correlation_id, + ) + + if self._should_block(response): + status = "guardrail_intervened" + except HTTPException: + status = "guardrail_failed_to_respond" + raise + except httpx.HTTPStatusError as exc: + # Preserve the upstream Graph API status code (e.g. 429, 503) so + # callers can distinguish a transient infrastructure error from a + # DLP policy block (signaled separately as HTTP 400 below) and can + # implement retry-after handling on rate limits. 401/403 upstream + # responses indicate a proxy-side credential / consent problem the + # caller can do nothing about, so they are mapped to 502. + status = "guardrail_failed_to_respond" + if block_on_violation: + upstream_status = exc.response.status_code + client_status = ( + 502 if upstream_status in (401, 403) else upstream_status + ) + headers: Optional[Dict[str, str]] = None + retry_after = exc.response.headers.get("retry-after") + if retry_after: + headers = {"Retry-After": retry_after} + raise HTTPException( + status_code=client_status, + detail={ + "error": "Microsoft Purview DLP: upstream policy evaluation failed", + "activity": activity, + "upstream_status": upstream_status, + "exception": str(exc), + }, + headers=headers, + ) from exc + verbose_proxy_logger.warning( + "Purview DLP: API/network error in logging-only mode (not re-raised): %s", + exc, + ) + except Exception as exc: + status = "guardrail_failed_to_respond" + if block_on_violation: + raise HTTPException( + status_code=400, + detail={ + "error": "Microsoft Purview DLP: upstream policy evaluation failed", + "activity": activity, + "exception": str(exc), + }, + ) from exc + verbose_proxy_logger.warning( + "Purview DLP: API/network error in logging-only mode (not re-raised): %s", + exc, + ) + finally: + end_time = datetime.now() + self.add_standard_logging_guardrail_information_to_request_data( + guardrail_provider=self.guardrail_provider, + guardrail_json_response=response, + request_data=request_data, + guardrail_status=status, + start_time=start_time.timestamp(), + end_time=end_time.timestamp(), + duration=(end_time - start_time).total_seconds(), + ) + + if block_on_violation and status == "guardrail_intervened": + raise HTTPException( + status_code=400, + detail={ + "error": "Microsoft Purview DLP: Content blocked by policy", + "activity": activity, + }, + ) + + return response + + @staticmethod + def _extract_responses_api_function_call_args(result: Any) -> List[str]: + """Return tool-call argument strings from a ``ResponsesAPIResponse.output``. + + ``ResponsesAPIResponse.output_text`` only aggregates ``output_text`` + content blocks and ignores ``function_call`` items. Model-generated + tool-call arguments can themselves contain sensitive data, so we + extract them explicitly to keep DLP coverage consistent with the + chat (``ModelResponse``) path. + """ + args: List[str] = [] + output = getattr(result, "output", None) + if not output: + return args + for item in output: + if isinstance(item, dict): + item_type = item.get("type") + arguments = item.get("arguments") + else: + item_type = getattr(item, "type", None) + arguments = getattr(item, "arguments", None) + if item_type == "function_call" and isinstance(arguments, str): + if arguments.strip(): + args.append(arguments) + return args + + def _completion_response_text_parts(self, result: Any) -> List[str]: + """Collect non-empty text segments from chat, text completions, or responses API. + + Includes assistant message content *and* model-generated tool-call + arguments so that sensitive data returned inside function calls is not + missed by the DLP scan. + """ + parts: List[str] = [] + if isinstance(result, TextCompletionResponse) and result.choices: + for text_choice in result.choices: + if not isinstance(text_choice, TextChoices): + continue + raw = text_choice.get("text") + if isinstance(raw, str) and raw.strip(): + parts.append(raw) + elif isinstance(result, ResponsesAPIResponse): + text = result.output_text + if text and text.strip(): + parts.append(text) + # Include tool-call arguments from ``function_call`` output items + # (``output_text`` ignores them). + parts.extend(self._extract_responses_api_function_call_args(result)) + elif isinstance(result, ModelResponse) and result.choices: + for chat_choice in result.choices: + if not isinstance(chat_choice, Choices): + continue + msg = chat_choice.message + if msg is None: + continue + raw = ( + msg.get("content") + if isinstance(msg, dict) + else getattr(msg, "content", None) + ) + if isinstance(raw, str) and raw.strip(): + parts.append(raw) + # Include tool-call arguments returned by the model + parts.extend(self._extract_tool_call_args_from_message(msg)) + return parts + + def _assemble_responses_api_from_chunks( + self, chunks: List[Any] + ) -> Tuple[bool, Optional[ResponsesAPIResponse]]: + """Extract the final ``ResponsesAPIResponse`` from a buffered Responses API stream. + + Returns a ``(is_responses_api_stream, assembled)`` tuple so the caller + can distinguish "not a Responses API stream" (fall through to + ``stream_chunk_builder``) from "Responses API stream but no final + response event was received" (fail closed with an accurate error). + When the stream is a Responses API stream the latest event carrying a + ``ResponsesAPIResponse`` body is returned (``response.completed``, or + ``response.failed`` / ``response.incomplete`` as fallbacks). + """ + looks_like_responses_api = False + final: Optional[ResponsesAPIResponse] = None + for chunk in chunks: + event_type = getattr(chunk, "type", None) + if isinstance(event_type, str) and event_type.startswith("response."): + looks_like_responses_api = True + candidate = getattr(chunk, "response", None) + if isinstance(candidate, ResponsesAPIResponse): + final = candidate + return looks_like_responses_api, final + + def _responses_api_input_to_str( + self, data: Dict[str, Any], raise_on_failure: bool = False + ) -> Optional[str]: + """Extract DLP-scannable text from a Responses API request ``input`` field. + + ``input`` may be a plain string or a list of input items (messages). In + the latter case the items are converted to chat messages via the standard + LiteLLM transformation and then concatenated by ``get_prompt_text_for_dlp``. + + When ``raise_on_failure`` is True (blocking mode), a transformation error + raises ``HTTPException`` so the request is fail-closed. In logging-only + mode the error is swallowed and ``None`` is returned so audit attempts on + the response side can still run. + """ + from litellm.responses.litellm_completion_transformation.transformation import ( + LiteLLMCompletionResponsesConfig, + ) + + input_data = data.get("input") + if input_data is None and not data.get("instructions"): + return None + try: + # Always transform via messages so ``instructions`` become a system message + # (string ``input`` alone would skip instructions and bypass DLP). + messages = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( + input=input_data if input_data is not None else "", + responses_api_request=data, + ) + return self.get_prompt_text_for_dlp(cast(List[Any], messages)) + except Exception: + verbose_proxy_logger.warning( + "Purview DLP: failed to transform responses API input", + exc_info=True, + ) + if raise_on_failure: + raise HTTPException( + status_code=400, + detail={ + "error": ( + "Microsoft Purview DLP: Responses API input could " + "not be transformed for DLP scanning in blocking mode" + ), + }, + ) + return None + + # ------------------------------------------------------------------ + # Identity resolution for blocking modes + # ------------------------------------------------------------------ + + def _resolve_user_id_for_blocking( + self, + data: Dict[str, Any], + user_api_key_dict: Any, + ) -> str: + """Resolve user ID for blocking (pre_call / post_call) DLP hooks. + + Uses only trusted proxy-authenticated sources (``_resolve_trusted_user_id``). + Caller-supplied ``UserAPIKeyAuth.end_user_id`` (from request ``user``, + ``metadata.user_id``, ``safety_identifier``, etc.) and + ``metadata[user_id_field]`` are rejected (fail closed) because they can + impersonate another Entra user's Purview policy. + + Raises ``HTTPException`` when no API-key-bound ``user_id`` exists or when + only caller-influenceable identity fields are available (fail closed). + """ + trusted_id = self._resolve_trusted_user_id(data, user_api_key_dict) + if trusted_id: + return trusted_id + + if self._resolve_user_id(data, user_api_key_dict): + raise HTTPException( + status_code=400, + detail={ + "error": ( + "Microsoft Purview DLP: No proxy-authenticated user identity; " + "bind user_id to the API key (caller-supplied metadata cannot " + "be used for blocking DLP)" + ), + }, + ) + + raise HTTPException( + status_code=400, + detail={ + "error": ( + "Microsoft Purview DLP: No proxy-authenticated user identity; " + "bind user_id to the API key for blocking DLP" + ), + }, + ) + + # ------------------------------------------------------------------ + # Pre-call hook — DLP on prompts + # ------------------------------------------------------------------ + + @log_guardrail_information + async def async_pre_call_hook( + self, + user_api_key_dict: "UserAPIKeyAuth", + cache: Any, + data: Dict[str, Any], + call_type: "CallTypesLiteral", + ) -> Optional[Dict[str, Any]]: + """Check user prompt against Purview DLP policies before LLM call.""" + user_id = self._resolve_user_id_for_blocking(data, user_api_key_dict) + + prompt_text: Optional[str] = None + if call_type in ("responses", "aresponses"): + # Route Responses API calls to the responses-specific extractor + # before the generic ``messages`` branch. This mirrors + # ``async_logging_hook`` and ensures ``instructions`` (system + # prompt) content is included in the DLP scan, and prevents a + # crafted ``messages`` key in the request from being scanned in + # place of the actual ``input``. + prompt_text = self._responses_api_input_to_str(data, raise_on_failure=True) + elif call_type in ("text_completion", "atext_completion"): + raw_prompt = data.get("prompt") + # Reject every token-id prompt shape Purview cannot evaluate — + # flat ``list[int]`` (single prompt), ``list[list[int]]`` (multi-prompt + # batches), and mixed lists that include any token-id sub-array. + # Empty/whitespace-only strings also yield ``prompt_text is None`` but + # contain no sensitive data and pass through harmlessly below. + if self.is_token_id_prompt(raw_prompt): + raise HTTPException( + status_code=400, + detail={ + "error": ( + "Microsoft Purview DLP: Token-id completion prompts " + "cannot be scanned for DLP in blocking mode" + ), + }, + ) + prompt_text = self.completion_prompt_to_str(raw_prompt) + else: + messages: Optional[List] = data.get("messages") + if messages: + prompt_text = self.get_prompt_text_for_dlp(cast(List[Any], messages)) + + if not prompt_text: + return data + + await self._check_content( + user_id=user_id, + text=prompt_text, + activity="uploadText", + request_data=data, + block_on_violation=True, + ) + return data + + # ------------------------------------------------------------------ + # Post-call hook — DLP on responses + # ------------------------------------------------------------------ + + @log_guardrail_information + async def async_post_call_success_hook( + self, + data: dict, + user_api_key_dict: "UserAPIKeyAuth", + response: Union[Any, ModelResponse, "EmbeddingResponse", "ImageResponse"], + ) -> Any: + """Check LLM response against Purview DLP policies (non-streaming only). + + Streaming responses are handled by ``async_post_call_streaming_iterator_hook`` + which buffers all chunks before scanning. The proxy automatically skips + this hook for requests that have a streaming iterator hook defined. + """ + user_id = self._resolve_user_id_for_blocking(data, user_api_key_dict) + + parts = self._completion_response_text_parts(response) + + if parts: + combined = "\n\n---\n\n".join(parts) + await self._check_content( + user_id=user_id, + text=combined, + activity="downloadText", + request_data=data, + block_on_violation=True, + ) + return response + + async def async_post_call_streaming_iterator_hook( + self, + user_api_key_dict: "UserAPIKeyAuth", + response: Any, + request_data: dict, + ) -> AsyncGenerator[ModelResponseStream, None]: + """Check streaming LLM responses against Purview DLP policies. + + All chunks are buffered before the DLP scan so that no content is + delivered to the client if a policy violation is detected. After a + clean scan the assembled response is re-yielded chunk-by-chunk via a + ``MockResponseIterator`` so the caller receives normal streaming output. + + The proxy automatically skips ``async_post_call_success_hook`` for + guardrails that define this method, preventing duplicate scans. + """ + from litellm.llms.base_llm.base_model_iterator import MockResponseIterator + from litellm.main import stream_chunk_builder + + # Resolve user ID up-front so identity failures don't waste work + # buffering and assembling the stream. + user_id = self._resolve_user_id_for_blocking(request_data, user_api_key_dict) + + # Buffer the entire stream before any DLP scan. + all_chunks: List[ModelResponseStream] = [] + async for chunk in response: + all_chunks.append(chunk) + + # Responses API streams emit typed events (e.g. ``response.completed``) + # whose final event carries the full ``ResponsesAPIResponse`` — these + # are not understood by ``stream_chunk_builder`` (which is built for + # chat/text-completion deltas). Detect and scan them via the same + # ``_completion_response_text_parts`` path used by non-streaming. + ( + is_responses_api_stream, + responses_api_assembled, + ) = self._assemble_responses_api_from_chunks(all_chunks) + if is_responses_api_stream: + if responses_api_assembled is None: + # Fail closed: Responses API events were seen but no final + # ``response.completed`` / ``response.failed`` / + # ``response.incomplete`` event carrying a ``ResponsesAPIResponse`` + # body was received, so we cannot scan the content. + raise HTTPException( + status_code=400, + detail={ + "error": ( + "Microsoft Purview DLP: Incomplete Responses API " + "stream — no final response event received for " + "DLP scanning; blocking response." + ), + }, + ) + parts = self._completion_response_text_parts(responses_api_assembled) + if parts: + combined = "\n\n---\n\n".join(parts) + await self._check_content( + user_id=user_id, + text=combined, + activity="downloadText", + request_data=request_data, + block_on_violation=True, + ) + for chunk in all_chunks: + yield chunk + return + + assembled_response = stream_chunk_builder(chunks=all_chunks) + + if assembled_response is None and all_chunks: + # Fail closed: stream_chunk_builder dropped all chunks, so we cannot + # scan the content. Refuse to release the buffered chunks. + raise HTTPException( + status_code=400, + detail={ + "error": ( + "Microsoft Purview DLP: Unable to assemble streamed " + "response for scanning; blocking response." + ), + }, + ) + + if isinstance( + assembled_response, (TextCompletionResponse, ResponsesAPIResponse) + ): + parts = self._completion_response_text_parts(assembled_response) + if parts: + combined = "\n\n---\n\n".join(parts) + await self._check_content( + user_id=user_id, + text=combined, + activity="downloadText", + request_data=request_data, + block_on_violation=True, + ) + for chunk in all_chunks: + yield chunk + return + + if not isinstance(assembled_response, ModelResponse): + # Non-content response (e.g. embeddings) — pass through unchanged. + for chunk in all_chunks: + yield chunk + return + + parts = self._completion_response_text_parts(assembled_response) + if parts: + combined = "\n\n---\n\n".join(parts) + # Raises HTTPException(400) on violation — no chunks are yielded. + await self._check_content( + user_id=user_id, + text=combined, + activity="downloadText", + request_data=request_data, + block_on_violation=True, + ) + + # DLP passed — re-yield chunks from the assembled chat response. + mock_response = MockResponseIterator(model_response=assembled_response) + async for chunk in mock_response: + yield chunk + + # ------------------------------------------------------------------ + # Logging-only hook — audit without blocking + # ------------------------------------------------------------------ + + def logging_hook( + self, kwargs: dict, result: Any, call_type: str + ) -> Tuple[dict, Any]: + """Fire-and-forget async audit logging; returns original (kwargs, result) immediately. + + In the proxy's async success path, litellm independently calls both + ``logging_hook`` (sync) and ``async_logging_hook`` (async) for every + ``CustomGuardrail`` callback. To avoid making two complete sets of + Purview API calls per request, this sync hook is a no-op whenever an + event loop is running — the framework's async path will invoke + ``async_logging_hook`` directly. + + For genuine sync-only call paths (no running event loop, so the async + success handler will not fire either), schedule ``async_logging_hook`` + on a short-lived background daemon thread so audit logging still runs + without blocking the caller on two Graph API round-trips. + """ + + try: + asyncio.get_running_loop() + # Async context — let the framework's async success handler invoke + # async_logging_hook to avoid duplicate Purview API calls. Log so + # the deferral is observable if the framework ever stops dispatching + # async_logging_hook on a given code path (otherwise audit silently + # drops). + verbose_proxy_logger.debug( + "Purview audit: deferring to async_logging_hook (running event loop detected)" + ) + return kwargs, result + except RuntimeError: + pass + + async def _log_safe() -> None: + try: + await self.async_logging_hook( + kwargs=kwargs, result=result, call_type=call_type + ) + except Exception as exc: + verbose_proxy_logger.error( + "Purview audit background logging error: %s", exc + ) + + def _run_in_new_loop() -> None: + new_loop = asyncio.new_event_loop() + try: + asyncio.set_event_loop(new_loop) + new_loop.run_until_complete(_log_safe()) + finally: + new_loop.close() + asyncio.set_event_loop(None) + + thread = threading.Thread(target=_run_in_new_loop, daemon=True) + thread.start() + + return kwargs, result + + async def async_logging_hook( + self, kwargs: dict, result: Any, call_type: str + ) -> Tuple[dict, Any]: + """Send both prompt and response to Purview for audit logging. + + Errors are logged but never raised — this mode is non-blocking. + Each audit call (prompt and response) is wrapped in its own try/except + so a failure on the first does not prevent the second from running. + """ + user_id = self._resolve_user_id_from_logging_kwargs(kwargs) + if not user_id: + verbose_proxy_logger.debug("Purview audit: no user_id, skipping") + return kwargs, result + + # Log prompt (uploadText) + try: + prompt_text: Optional[str] = None + if call_type in ("responses", "aresponses"): + # Responses API: route to the responses-specific extractor + # before the generic ``messages`` branch. litellm's logging + # pipeline stores the raw responses ``input`` (a string or a + # list of input items) under ``model_call_details["messages"]`` + # via ``function_setup``, which is NOT the chat message format + # ``get_prompt_text_for_dlp`` expects. Use the original + # ``input`` / ``instructions`` keys that ``pre_call`` and + # ``update_environment_variables`` persist on the call details. + prompt_text = self._responses_api_input_to_str(kwargs) + elif call_type in ("text_completion", "atext_completion"): + prompt_text = self.completion_prompt_to_str(kwargs.get("prompt")) + else: + messages = kwargs.get("messages") + if messages: + prompt_text = self.get_prompt_text_for_dlp( + cast(List[Any], messages) + ) + + if prompt_text: + await self._check_content( + user_id=user_id, + text=prompt_text, + activity="uploadText", + request_data=kwargs, + block_on_violation=False, + ) + except Exception as e: + verbose_proxy_logger.error("Purview audit logging error (prompt): %s", e) + + # Log response (downloadText) — runs regardless of prompt audit outcome + try: + parts = self._completion_response_text_parts(result) + if parts: + combined = "\n\n---\n\n".join(parts) + await self._check_content( + user_id=user_id, + text=combined, + activity="downloadText", + request_data=kwargs, + block_on_violation=False, + ) + except Exception as e: + verbose_proxy_logger.error("Purview audit logging error (response): %s", e) + + return kwargs, result diff --git a/litellm/proxy/guardrails/guardrail_hooks/openai/base.py b/litellm/proxy/guardrails/guardrail_hooks/openai/base.py index 872d09cd886..281afacd5c4 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/openai/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/openai/base.py @@ -1,5 +1,9 @@ from typing import TYPE_CHECKING, List, Optional +from litellm.litellm_core_utils.prompt_templates.common_utils import ( + get_last_user_message, +) + if TYPE_CHECKING: from litellm.types.llms.openai import AllMessageValues @@ -21,32 +25,4 @@ class OpenAIGuardrailBase: ] get_user_prompt(messages) -> "What is the weather in Tokyo?" """ - from litellm.litellm_core_utils.prompt_templates.common_utils import ( - convert_content_list_to_str, - ) - - if not messages: - return None - - # Iterate from the end to find the last consecutive block of user messages - user_messages = [] - for message in reversed(messages): - if message.get("role") == "user": - user_messages.append(message) - else: - # Stop when we hit a non-user message - break - - if not user_messages: - return None - - # Reverse to get the messages in chronological order - user_messages.reverse() - - user_prompt = "" - for message in user_messages: - text_content = convert_content_list_to_str(message) - user_prompt += text_content + "\n" - - result = user_prompt.strip() - return result if result else None + return get_last_user_message(messages) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index ff3bbf47389..d6082899c02 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -43,6 +43,8 @@ from litellm.caching.dual_cache import DualCache from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid from litellm.constants import ( + CLI_SSO_CLAIM_MAP, + CLI_SSO_CLAIM_MAX_SCALAR_LENGTH, CLI_SSO_SESSION_CACHE_KEY_PREFIX, CLI_SSO_SESSION_TTL_SECONDS, LITELLM_CLI_SOURCE_IDENTIFIER, @@ -140,6 +142,20 @@ _CLI_SSO_START_RATE_LIMIT_WINDOW_SECONDS = 60 _CLI_SSO_START_RATE_LIMIT_MAX_ATTEMPTS = 30 _CLI_SSO_USER_CODE_ALPHABET = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789" _CLI_SSO_LOGIN_ID_RE = re.compile(r"^cli-[A-Za-z0-9_-]{12,124}$") +_CLI_SSO_SCALAR_TYPES = (str, int, float, bool) +_CLI_SSO_DEST_KEY_RE = re.compile(r"^[A-Za-z0-9_.-]+$") +_CLI_SSO_SECRET_KEY_FRAGMENTS = frozenset( + { + "access_token", + "api_key", + "client_secret", + "id_token", + "password", + "private_key", + "refresh_token", + "secret", + } +) def _hash_cli_sso_secret(secret: str) -> str: @@ -225,6 +241,239 @@ def _verify_cli_sso_poll_secret(flow: dict, poll_secret: Optional[str]) -> bool: return secrets.compare_digest(supplied_poll_secret_hash, expected_poll_secret_hash) +def _parse_cli_sso_claim_map() -> List[Tuple[str, str]]: + """ + Parse CLI_SSO_CLAIM_MAP / LITELLM_CLI_SSO_CLAIM_MAP. + + Format: comma-separated ``source_claim->metadata_key`` pairs, e.g. + ``employment_type->acme_employment_type,org_info.department->department``. + Destination keys may use an optional ``metadata.`` prefix; values are stored + on the LiteLLM user's ``metadata`` JSON column. + """ + claim_map_raw = CLI_SSO_CLAIM_MAP.strip() + if not claim_map_raw: + return [] + + parsed: List[Tuple[str, str]] = [] + for entry in claim_map_raw.split(","): + entry = entry.strip() + if not entry or "->" not in entry: + continue + source_claim, dest_key = entry.split("->", 1) + source_claim = source_claim.strip() + dest_key = dest_key.strip() + if dest_key.startswith("metadata."): + dest_key = dest_key[len("metadata.") :] + if source_claim and dest_key: + parsed.append((source_claim, dest_key)) + return parsed + + +def _is_safe_cli_sso_metadata_dest_key(dest_key: str) -> bool: + if not dest_key or not _CLI_SSO_DEST_KEY_RE.fullmatch(dest_key): + return False + lowered = dest_key.lower() + return not any(fragment in lowered for fragment in _CLI_SSO_SECRET_KEY_FRAGMENTS) + + +def _is_safe_cli_sso_scalar_claim_value(value: Any) -> bool: + if not isinstance(value, _CLI_SSO_SCALAR_TYPES): + return False + if isinstance(value, str): + if len(value) > CLI_SSO_CLAIM_MAX_SCALAR_LENGTH: + return False + if value.startswith("eyJ") and value.count(".") >= 2: + return False + return True + + +def _sso_result_to_dict(result: Union[CustomOpenID, OpenID, dict]) -> Dict[str, Any]: + if isinstance(result, dict): + return result + if hasattr(result, "model_dump"): + dumped = result.model_dump() + if isinstance(dumped, dict): + return cast(Dict[str, Any], dumped) + return {} + + +def _get_nested_claim_value(data: Dict[str, Any], claim_path: str) -> Any: + """Resolve a dot-notation claim path against an SSO result dict. + + Unlike ``get_nested_value``, this does not strip a leading ``metadata.`` + prefix, since OIDC claims may legitimately use ``metadata`` as a top-level + key. + """ + if not claim_path: + return None + if claim_path in data: + return data[claim_path] + placeholder = "\x00" + parts = claim_path.replace("\\.", placeholder).split(".") + parts = [p.replace(placeholder, ".") for p in parts] + current: Any = data + for part in parts: + if isinstance(current, dict) and part in current: + current = current[part] + else: + return None + return current + + +def _extract_sso_claim_value( + result: Union[CustomOpenID, OpenID, dict], claim_path: str +) -> Any: + extra_fields = getattr(result, "extra_fields", None) + if isinstance(extra_fields, dict): + if claim_path in extra_fields: + return extra_fields[claim_path] + nested = _get_nested_claim_value(extra_fields, claim_path) + if nested is not None: + return nested + + if isinstance(result, dict): + return _get_nested_claim_value(result, claim_path) + + result_dict = _sso_result_to_dict(result) + return _get_nested_claim_value(result_dict, claim_path) + + +def _set_nested_metadata_value( + metadata: Dict[str, Any], key_path: str, value: Any +) -> None: + placeholder = "\x00" + parts = key_path.replace("\\.", placeholder).split(".") + parts = [p.replace(placeholder, ".") for p in parts] + current: Any = metadata + for part in parts[:-1]: + existing = current.get(part) + if not isinstance(existing, dict): + existing = {} + current[part] = existing + current = existing + current[parts[-1]] = value + + +def _flatten_cli_sso_metadata_for_poll( + metadata: Dict[str, Any], +) -> Dict[str, Union[str, int, float, bool]]: + """Expose scalar attribution metadata as a flat dict for CLI poll responses.""" + flattened: Dict[str, Union[str, int, float, bool]] = {} + stack: List[Tuple[str, Any]] = [("", metadata)] + while stack: + prefix, value = stack.pop() + if isinstance(value, dict): + for key, nested in value.items(): + nested_prefix = f"{prefix}.{key}" if prefix else key + stack.append((nested_prefix, nested)) + elif _is_safe_cli_sso_scalar_claim_value(value): + flattened[prefix] = value + return flattened + + +def build_cli_sso_attribution_metadata( + result: Union[CustomOpenID, OpenID, dict], +) -> Dict[str, Any]: + """ + Build allowlisted, non-secret scalar attribution metadata from an SSO result. + + Sources are configured via CLI_SSO_CLAIM_MAP / LITELLM_CLI_SSO_CLAIM_MAP and + may include claims captured by GENERIC_USER_EXTRA_ATTRIBUTES on CustomOpenID. + """ + claim_map = _parse_cli_sso_claim_map() + if not claim_map: + return {} + + metadata: Dict[str, Any] = {} + for source_claim, dest_key in claim_map: + if not _is_safe_cli_sso_metadata_dest_key(dest_key): + verbose_proxy_logger.debug( + f"Skipping unsafe CLI SSO metadata destination key: {dest_key}" + ) + continue + + raw_value = _extract_sso_claim_value(result=result, claim_path=source_claim) + if not _is_safe_cli_sso_scalar_claim_value(raw_value): + continue + + _set_nested_metadata_value( + metadata=metadata, key_path=dest_key, value=raw_value + ) + + return metadata + + +def _merge_cli_sso_attribution_metadata( + existing_metadata: Dict[str, Any], attribution_metadata: Dict[str, Any] +) -> Dict[str, Any]: + """Merge attribution metadata into existing user metadata in-place. + + Preserves original value types (in particular, string claim values that + happen to look numeric are NOT coerced to ``int``/``float``). Nested dicts + are merged iteratively so attribution claims do not clobber unrelated keys + under the same parent. + """ + pending: List[Tuple[Dict[str, Any], Dict[str, Any]]] = [ + (existing_metadata, attribution_metadata) + ] + while pending: + target, source = pending.pop() + for key, value in source.items(): + if value is None: + continue + existing_value = target.get(key) + if isinstance(value, dict) and isinstance(existing_value, dict): + pending.append((existing_value, value)) + else: + target[key] = value + return existing_metadata + + +async def _persist_cli_sso_user_metadata( + prisma_client: PrismaClient, + user_id: str, + attribution_metadata: Dict[str, Any], +) -> None: + if not attribution_metadata: + return + + try: + user_row = await prisma_client.db.litellm_usertable.find_unique( + where={"user_id": user_id} + ) + existing_metadata: Dict[str, Any] = {} + if user_row is not None: + row_metadata = user_row.metadata + if isinstance(row_metadata, dict): + existing_metadata = deepcopy(row_metadata) + + merged_metadata = _merge_cli_sso_attribution_metadata( + existing_metadata=existing_metadata, + attribution_metadata=attribution_metadata, + ) + await prisma_client.db.litellm_usertable.update_many( + where={"user_id": user_id}, + data={"metadata": merged_metadata}, + ) + verbose_proxy_logger.info( + f"Persisted CLI SSO attribution metadata for user {user_id}: " + f"{list(_flatten_cli_sso_metadata_for_poll(attribution_metadata).keys())}" + ) + except Exception as e: + verbose_proxy_logger.error( + f"Failed to persist CLI SSO attribution metadata for user {user_id}: {e}" + ) + + +def _cli_poll_attribution_metadata_from_session( + session_data: Dict[str, Any], +) -> Dict[str, Union[str, int, float, bool]]: + stored = session_data.get("attribution_metadata") + if isinstance(stored, dict): + return _flatten_cli_sso_metadata_for_poll(stored) + return {} + + def _render_cli_sso_verification_page( verify_url: str, browser_complete_token: str ) -> str: @@ -1674,7 +1923,12 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa: key_id = state_parts[1] if len(state_parts) > 1 else None verbose_proxy_logger.info("CLI SSO callback detected") - return await cli_sso_callback(request=request, key=key_id, result=result) + return await cli_sso_callback( + request=request, + key=key_id, + result=result, + received_response=received_response, + ) # Control-plane cross-origin: read return_to from cookie. # Starlette's cookie_parser already handles RFC 2109 unquoting. @@ -1692,15 +1946,144 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa: ) +async def _build_cli_sso_user_defined_values( + result: Union[OpenID, dict], + parsed_openid_result: ParsedOpenIDResult, +) -> Optional[SSOUserDefinedValues]: + from litellm.proxy.proxy_server import user_custom_sso + + user_id = parsed_openid_result.get("user_id") + if user_custom_sso is not None: + if inspect.iscoroutinefunction(user_custom_sso): + return await user_custom_sso(result) # type: ignore + raise ValueError("user_custom_sso must be a coroutine function") + if user_id is None: + return None + return SSOUserDefinedValues( + models=[], + user_id=user_id, + user_email=parsed_openid_result.get("user_email"), + max_budget=litellm.max_internal_user_budget, + user_role=parsed_openid_result.get("user_role"), + budget_duration=litellm.internal_user_budget_duration, + ) + + +async def _fetch_cli_sso_team_details( + prisma_client: PrismaClient, + teams: List[str], +) -> List[Dict[str, Any]]: + team_details: List[Dict[str, Any]] = [] + try: + if teams: + prisma_teams = await prisma_client.db.litellm_teamtable.find_many( + where={"team_id": {"in": teams}} + ) + for team_row in prisma_teams: + team_dict = team_row.model_dump() + team_details.append( + { + "team_id": team_dict.get("team_id"), + "team_alias": team_dict.get("team_alias"), + } + ) + except Exception as e: + verbose_proxy_logger.error( + f"Error fetching team details for CLI SSO session: {e}" + ) + return team_details + + +async def _complete_cli_sso_callback_session( + *, + request: Request, + key: str, + flow: dict, + result: Union[OpenID, dict], + parsed_openid_result: ParsedOpenIDResult, + user_defined_values: Optional[SSOUserDefinedValues], + prisma_client: PrismaClient, + user_api_key_cache: UserApiKeyCache, + proxy_logging_obj: ProxyLogging, +): + from fastapi.responses import HTMLResponse + + user_id = parsed_openid_result.get("user_id") + user_email = parsed_openid_result.get("user_email") + user_info = await get_user_info_from_db( + result=result, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + user_email=user_email, + user_defined_values=user_defined_values, + alternate_user_id=user_id, + ) + if user_info is None: + raise HTTPException( + status_code=500, detail="Failed to retrieve user information from SSO" + ) + if not user_info.user_id: + raise HTTPException( + status_code=500, detail="Failed to retrieve user information from SSO" + ) + + teams: List[str] = [] + if hasattr(user_info, "teams") and user_info.teams: + teams = user_info.teams if isinstance(user_info.teams, list) else [] + + team_details = await _fetch_cli_sso_team_details( + prisma_client=prisma_client, teams=teams + ) + attribution_metadata = build_cli_sso_attribution_metadata(result=result) + if attribution_metadata: + await _persist_cli_sso_user_metadata( + prisma_client=prisma_client, + user_id=cast(str, user_info.user_id), + attribution_metadata=attribution_metadata, + ) + + flow["session_data"] = { + "user_id": cast(str, user_info.user_id), + "user_role": user_info.user_role, + "models": user_info.models if hasattr(user_info, "models") else [], + "user_email": user_email, + "teams": teams, + "team_details": team_details, + "attribution_metadata": attribution_metadata, + } + flow["sso_complete"] = True + browser_complete_token = secrets.token_urlsafe(32) + flow["browser_complete_token_hash"] = _hash_cli_sso_secret(browser_complete_token) + _set_cli_sso_flow(login_id=key, cache=user_api_key_cache, flow=flow) + + verbose_proxy_logger.info( + f"Stored CLI SSO session for user: {user_info.user_id}, teams: {teams}, num_teams: {len(teams)}" + ) + verify_url = get_custom_url( + request_base_url=str(request.base_url), + route=f"sso/cli/complete/{key}", + ) + return HTMLResponse( + content=_render_cli_sso_verification_page( + verify_url=verify_url, + browser_complete_token=browser_complete_token, + ), + status_code=200, + ) + + async def cli_sso_callback( request: Request, key: Optional[str] = None, result: Optional[Union[OpenID, dict]] = None, + received_response: Optional[dict] = None, ): """CLI SSO callback - stores session info for JWT generation on polling""" verbose_proxy_logger.info("CLI SSO callback") from litellm.proxy.proxy_server import ( + general_settings, prisma_client, proxy_logging_obj, user_api_key_cache, @@ -1722,92 +2105,40 @@ async def cli_sso_callback( # After None check, cast to non-None type for type checker result_non_none: Union[OpenID, dict] = cast(Union[OpenID, dict], result) - parsed_openid_result = SSOAuthenticationHandler._get_user_email_and_id_from_result( - result=result_non_none - ) - verbose_proxy_logger.debug(f"parsed_openid_result: {parsed_openid_result}") - try: - # Get full user info from DB - user_info = await get_user_info_from_db( + parsed_openid_result = ( + SSOAuthenticationHandler._get_user_email_and_id_from_result( + result=result_non_none, + generic_client_id=os.getenv("GENERIC_CLIENT_ID", None), + ) + ) + verbose_proxy_logger.debug(f"parsed_openid_result: {parsed_openid_result}") + user_defined_values = await _build_cli_sso_user_defined_values( result=result_non_none, + parsed_openid_result=parsed_openid_result, + ) + + SSOAuthenticationHandler.verify_user_in_restricted_sso_group( + general_settings=general_settings, + result=result_non_none, + received_response=received_response, + ) + + return await _complete_cli_sso_callback_session( + request=request, + key=cast(str, key), + flow=flow, + result=result_non_none, + parsed_openid_result=parsed_openid_result, + user_defined_values=user_defined_values, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, - user_email=parsed_openid_result.get("user_email"), - user_defined_values=None, - alternate_user_id=parsed_openid_result.get("user_id"), ) - - if user_info is None: - raise HTTPException( - status_code=500, detail="Failed to retrieve user information from SSO" - ) - - # Get all teams from user_info - CLI will let user select which one - teams: List[str] = [] - if hasattr(user_info, "teams") and user_info.teams: - teams = user_info.teams if isinstance(user_info.teams, list) else [] - - # Also fetch team aliases for a better CLI UX. We keep the original - # "teams" list of IDs for backwards compatibility and add an - # optional "team_details" field containing objects with both - # team_id and team_alias. - team_details: List[Dict[str, Any]] = [] - try: - if teams: - prisma_teams = await prisma_client.db.litellm_teamtable.find_many( - where={"team_id": {"in": teams}} - ) - for team_row in prisma_teams: - team_dict = team_row.model_dump() - team_details.append( - { - "team_id": team_dict.get("team_id"), - "team_alias": team_dict.get("team_alias"), - } - ) - except Exception as e: - # If anything goes wrong here, fall back gracefully without - # impacting the SSO flow. - verbose_proxy_logger.error( - f"Error fetching team details for CLI SSO session: {e}" - ) - - session_data = { - "user_id": user_info.user_id, - "user_role": user_info.user_role, - "models": user_info.models if hasattr(user_info, "models") else [], - "user_email": parsed_openid_result.get("user_email"), - "teams": teams, - # Optional rich metadata for clients that want nicer display - "team_details": team_details, - } - - flow["session_data"] = session_data - flow["sso_complete"] = True - browser_complete_token = secrets.token_urlsafe(32) - flow["browser_complete_token_hash"] = _hash_cli_sso_secret( - browser_complete_token - ) - _set_cli_sso_flow(login_id=cast(str, key), cache=user_api_key_cache, flow=flow) - - verbose_proxy_logger.info( - f"Stored CLI SSO session for user: {user_info.user_id}, teams: {teams}, num_teams: {len(teams)}" - ) - - from fastapi.responses import HTMLResponse - - verify_url = get_custom_url( - request_base_url=str(request.base_url), - route=f"sso/cli/complete/{key}", - ) - html_content = _render_cli_sso_verification_page( - verify_url=verify_url, - browser_complete_token=browser_complete_token, - ) - return HTMLResponse(content=html_content, status_code=200) - + except ProxyException: + raise + except HTTPException: + raise except Exception as e: verbose_proxy_logger.error(f"Error with CLI SSO callback: {e}") raise HTTPException( @@ -1873,13 +2204,19 @@ async def cli_poll_key( team_details_response = [ {"team_id": t, "team_alias": None} for t in user_teams ] - return { + poll_response: Dict[str, Any] = { "status": "ready", "user_id": user_id, "teams": user_teams, "team_details": team_details_response, "requires_team_selection": True, } + attribution_metadata = _cli_poll_attribution_metadata_from_session( + session_data + ) + if attribution_metadata: + poll_response["attribution_metadata"] = attribution_metadata + return poll_response # Validate team_id if provided if team_id is not None: @@ -1892,6 +2229,17 @@ async def cli_poll_key( # If no team_id provided and user has 0 or 1 team, use first team (or None) team_id = user_teams[0] if len(user_teams) > 0 else None + team_alias = None + if team_id and isinstance(user_team_details, list): + team_alias = next( + ( + team.get("team_alias") + for team in user_team_details + if team.get("team_id") == team_id + ), + None, + ) + # Create user object for JWT generation user_info = LiteLLM_UserTable( user_id=user_id, @@ -1903,7 +2251,7 @@ async def cli_poll_key( # Generate CLI JWT on-demand (expiration configurable via LITELLM_CLI_JWT_EXPIRATION_HOURS) # Pass selected team_id to ensure JWT has correct team jwt_token = ExperimentalUIJWTToken.get_cli_jwt_auth_token( - user_info=user_info, team_id=team_id + user_info=user_info, team_id=team_id, team_alias=team_alias ) # Delete cache entry (single-use) @@ -1912,7 +2260,7 @@ async def cli_poll_key( verbose_proxy_logger.info( f"CLI JWT generated for user: {user_id}, team: {team_id}" ) - return { + poll_response = { "status": "ready", "key": jwt_token, "user_id": user_id, @@ -1922,6 +2270,12 @@ async def cli_poll_key( # present nicer information if needed. "team_details": user_team_details, } + attribution_metadata = _cli_poll_attribution_metadata_from_session( + session_data + ) + if attribution_metadata: + poll_response["attribution_metadata"] = attribution_metadata + return poll_response else: return {"status": "pending"} diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index da8da1b486f..c4e72cb7dc5 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -9,6 +9,7 @@ from functools import lru_cache from typing import Any, Dict, List, Literal, Optional import httpx +from openai._streaming import SSEDecoder import litellm from litellm.constants import ( @@ -27,7 +28,7 @@ from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfi from litellm.responses.utils import ResponsesAPIRequestUtils from litellm.types.llms.openai import ResponsesAPIStreamEvents from litellm.types.utils import CallTypes -from litellm.utils import CustomStreamWrapper, async_post_call_success_deployment_hook +from litellm.utils import async_post_call_success_deployment_hook @lru_cache(maxsize=1) @@ -120,10 +121,10 @@ class BaseResponsesAPIStreamingIterator: if not chunk: return None - # Handle SSE format (data: {...}) - chunk = CustomStreamWrapper._strip_sse_data_from_chunk(chunk) - if chunk is None: - return None + # NOTE: ``SSEDecoder`` already strips the SSE ``data:`` field prefix, so + # the value passed in here is the raw field content. Do not re-run + # ``_strip_sse_data_from_chunk`` on it — doing so would incorrectly mangle + # payloads whose actual JSON value happens to start with ``data:``. # Handle "[DONE]" marker if chunk == STREAM_SSE_DONE_STRING: @@ -634,7 +635,7 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): request_data, call_type, ) - self.stream_iterator = response.aiter_lines() + self.stream_iterator = SSEDecoder().aiter_bytes(response.aiter_bytes()) def __aiter__(self): return self @@ -645,13 +646,13 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): while True: # Get the next chunk from the stream try: - chunk = await self.stream_iterator.__anext__() + sse = await self.stream_iterator.__anext__() except StopAsyncIteration: self.finished = True raise StopAsyncIteration self._check_max_streaming_duration() - result = self._process_chunk(chunk) + result = self._process_chunk(sse.data) if self.finished: raise StopAsyncIteration @@ -708,7 +709,7 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): request_data, call_type, ) - self.stream_iterator = response.iter_lines() + self.stream_iterator = SSEDecoder().iter_bytes(response.iter_bytes()) def __iter__(self): return self @@ -719,13 +720,13 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): while True: # Get the next chunk from the stream try: - chunk = next(self.stream_iterator) + sse = next(self.stream_iterator) except StopIteration: self.finished = True raise StopIteration self._check_max_streaming_duration() - result = self._process_chunk(chunk) + result = self._process_chunk(sse.data) if self.finished: raise StopIteration diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 0a51ce3d456..0430c570e14 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -5,6 +5,9 @@ from typing import Any, Dict, List, Literal, Optional, Union from pydantic import BaseModel, ConfigDict, Field, field_validator from typing_extensions import Required, TypedDict +from litellm.types.proxy.guardrails.guardrail_hooks.akto import ( + AktoConfigModel, +) from litellm.types.proxy.guardrails.guardrail_hooks.block_code_execution import ( BlockCodeExecutionGuardrailConfigModel, ) @@ -17,9 +20,6 @@ from litellm.types.proxy.guardrails.guardrail_hooks.grayswan import ( from litellm.types.proxy.guardrails.guardrail_hooks.ibm import ( IBMGuardrailsBaseConfigModel, ) -from litellm.types.proxy.guardrails.guardrail_hooks.akto import ( - AktoConfigModel, -) from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import ( ContentFilterCategoryConfig, ) @@ -93,6 +93,7 @@ class SupportedGuardrailIntegrations(Enum): GENERIC_GUARDRAIL_API = "generic_guardrail_api" QUALIFIRE = "qualifire" CUSTOM_CODE = "custom_code" + MICROSOFT_PURVIEW = "microsoft_purview" SEMANTIC_GUARD = "semantic_guard" MCP_END_USER_PERMISSION = "mcp_end_user_permission" BLOCK_CODE_EXECUTION = "block_code_execution" diff --git a/litellm/utils.py b/litellm/utils.py index 2487d39bd0d..c28a88e0f1c 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -2933,6 +2933,13 @@ def register_model(model_cost: Union[str, dict]): # noqa: PLR0915 except Exception: existing_model = {} model_cost_key = key + # ``get_model_info`` returns ``litellm_provider: None`` when the + # provider is unknown (e.g. custom deployments registered via + # ``Router.add_deployment``). Persisting that None into + # ``litellm.model_cost`` causes ``_check_provider_match`` to drop + # custom pricing on subsequent cost lookups. + if existing_model.get("litellm_provider") is None: + existing_model.pop("litellm_provider", None) ## override / add new keys to the existing model cost dictionary updated_dictionary = _update_dictionary(existing_model, value) litellm.model_cost.setdefault(model_cost_key, {}).update(updated_dictionary) @@ -3343,6 +3350,21 @@ def get_optional_params_embeddings( # noqa: PLR0915 model=model, drop_params=drop_params if drop_params is not None else False, ) + # Provider-only params (e.g. Cohere input_type) are not in + # OPENAI_EMBEDDING_PARAMS, so embedding_pre_process drops them from + # non_default_params before map_openai_params. Restore only those extras + # from passed_params — skip OPENAI_EMBEDDING_PARAMS to avoid duplicating + # values already mapped (e.g. dimensions -> output_dimension). + if supported_params: + for param in supported_params: + if param in OPENAI_EMBEDDING_PARAMS: + continue + if ( + param in passed_params + and passed_params[param] is not None + and param not in optional_params + ): + optional_params[param] = passed_params[param] ## raise exception if non-default value passed for non-openai/azure embedding calls elif custom_llm_provider == "openai": # 'dimensions` is only supported in `text-embedding-3` and later models @@ -4019,16 +4041,23 @@ def get_optional_params( # noqa: PLR0915 thinking: Optional[AnthropicThinkingParam] = None, web_search_options: Optional[OpenAIWebSearchOptions] = None, safety_identifier: Optional[str] = None, + base_model: Optional[str] = None, **kwargs, ): passed_params = locals().copy() special_params = passed_params.pop("kwargs") + # Remove base_model from passed_params so it doesn't interfere with + # non_default_params / _check_valid_arg — it's a routing hint, not an + # OpenAI param. + passed_params.pop("base_model", None) provider_config: Optional[BaseConfig] = None if custom_llm_provider is not None and custom_llm_provider in [ provider.value for provider in LlmProviders ]: provider_config = ProviderConfigManager.get_provider_chat_config( - model=model, provider=LlmProviders(custom_llm_provider) + model=model, + provider=LlmProviders(custom_llm_provider), + base_model=base_model, ) non_default_params = pre_process_non_default_params( passed_params=passed_params, @@ -4091,7 +4120,7 @@ def get_optional_params( # noqa: PLR0915 sys.modules[__name__], "get_supported_openai_params" ) supported_params = get_supported_openai_params( - model=model, custom_llm_provider=custom_llm_provider + model=model, custom_llm_provider=custom_llm_provider, base_model=base_model ) if supported_params is None: supported_params = get_supported_openai_params( @@ -4702,22 +4731,27 @@ def get_optional_params( # noqa: PLR0915 ), ) elif custom_llm_provider == "azure": - if litellm.AzureOpenAIO1Config().is_o_series_model(model=model): + _azure_detection_model = base_model or model + if litellm.AzureOpenAIO1Config().is_o_series_model( + model=_azure_detection_model + ): optional_params = litellm.AzureOpenAIO1Config().map_openai_params( non_default_params=non_default_params, optional_params=optional_params, - model=model, + model=_azure_detection_model, drop_params=( drop_params if drop_params is not None and isinstance(drop_params, bool) else False ), ) - elif litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(model=model): + elif litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model( + model=_azure_detection_model + ): optional_params = litellm.AzureOpenAIGPT5Config().map_openai_params( non_default_params=non_default_params, optional_params=optional_params, - model=model, + model=_azure_detection_model, drop_params=( drop_params if drop_params is not None and isinstance(drop_params, bool) @@ -4739,7 +4773,7 @@ def get_optional_params( # noqa: PLR0915 optional_params = litellm.AzureOpenAIConfig().map_openai_params( non_default_params=non_default_params, optional_params=optional_params, - model=model, + model=_azure_detection_model, api_version=api_version, # type: ignore drop_params=( drop_params @@ -5510,9 +5544,15 @@ def _get_model_info_from_model_cost(key: str) -> dict: def _check_provider_match(model_info: dict, custom_llm_provider: Optional[str]) -> bool: """ Check if the model info provider matches the custom provider. + + A missing ``litellm_provider`` key and a ``litellm_provider`` set to + ``None`` both mean "no specific provider constraint" and are treated + as a wildcard match. ``register_model`` may persist ``None`` here via + ``get_model_info`` when a deployment is registered without a provider, + so normalising the two cases keeps custom pricing applied consistently. """ if custom_llm_provider and ( - "litellm_provider" in model_info + model_info.get("litellm_provider") is not None and model_info["litellm_provider"] != custom_llm_provider ): if custom_llm_provider == "vertex_ai" and model_info[ @@ -8124,10 +8164,8 @@ class ProviderConfigManager: # Format: (factory_function, needs_model_parameter: bool) LlmProviders.OPENAI: (lambda: litellm.OpenAIGPTConfig(), False), LlmProviders.ANTHROPIC: (lambda: litellm.AnthropicConfig(), False), - LlmProviders.AZURE: ( - lambda model: ProviderConfigManager._get_azure_config(model), - True, - ), + # AZURE is handled as a special case in get_provider_chat_config() + # so that base_model can be threaded through for model-type detection. LlmProviders.AZURE_AI: ( lambda model: ProviderConfigManager._get_azure_ai_config(model), True, @@ -8267,11 +8305,19 @@ class ProviderConfigManager: } @staticmethod - def _get_azure_config(model: str) -> BaseConfig: - """Get Azure config based on model type.""" - if litellm.AzureOpenAIO1Config().is_o_series_model(model=model): + def _get_azure_config(model: str, base_model: Optional[str] = None) -> BaseConfig: + """Get Azure config based on model type. + + When *base_model* is provided (e.g. ``"azure/gpt-5.2"``), it is used + for model-type detection instead of *model* (the deployment name). + This allows non-standard deployment names like ``"azure/foo"`` to be + routed through the correct config when the user specifies the true + underlying model via ``base_model``. + """ + detection_model = base_model or model + if litellm.AzureOpenAIO1Config().is_o_series_model(model=detection_model): return litellm.AzureOpenAIO1Config() - if litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(model=model): + if litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(model=detection_model): return litellm.AzureOpenAIGPT5Config() return litellm.AzureOpenAIConfig() @@ -8329,13 +8375,18 @@ class ProviderConfigManager: @staticmethod def get_provider_chat_config( # noqa: PLR0915 - model: str, provider: LlmProviders + model: str, + provider: LlmProviders, + base_model: Optional[str] = None, ) -> Optional[BaseConfig]: """ Returns the provider config for a given provider. Uses O(1) dictionary lookup for fast provider resolution. Python classes take priority over JSON (they have custom overrides). + + For Azure, *base_model* (when set) drives model-type detection so that + non-standard deployment names still route to the correct config. """ # Handle OpenAI special cases (O-series and GPT-5 models) if provider == LlmProviders.OPENAI: @@ -8344,6 +8395,12 @@ class ProviderConfigManager: if litellm.OpenAIGPT5Config.is_model_gpt_5_model(model=model): return litellm.OpenAIGPT5Config() + # Handle Azure before the generic map so base_model can be threaded through + if provider == LlmProviders.AZURE: + return ProviderConfigManager._get_azure_config( + model=model, base_model=base_model + ) + # Initialize provider config map lazily (avoids circular imports) if ProviderConfigManager._PROVIDER_CONFIG_MAP is None: ProviderConfigManager._PROVIDER_CONFIG_MAP = ( diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 31a5993a240..2140493ec4a 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -15006,10 +15006,16 @@ "supports_service_tier": true }, "gemini-3.1-flash-lite": { - "cache_read_input_token_cost": 4.5e-08, - "cache_read_input_token_cost_per_audio_token": 9e-08, - "input_cost_per_audio_token": 9e-07, - "input_cost_per_token": 4.5e-07, + "cache_read_input_token_cost": 2.5e-08, + "cache_read_input_token_cost_batches": 1.25e-08, + "cache_read_input_token_cost_flex": 1.25e-08, + "cache_read_input_token_cost_per_audio_token": 5e-08, + "cache_read_input_token_cost_priority": 4.5e-08, + "input_cost_per_audio_token": 5e-07, + "input_cost_per_token": 2.5e-07, + "input_cost_per_token_batches": 1.25e-07, + "input_cost_per_token_flex": 1.25e-07, + "input_cost_per_token_priority": 4.5e-07, "litellm_provider": "vertex_ai-language-models", "max_audio_length_hours": 8.4, "max_audio_per_prompt": 1, @@ -15021,9 +15027,12 @@ "max_video_length": 1, "max_videos_per_prompt": 10, "mode": "chat", - "output_cost_per_reasoning_token": 2.7e-06, - "output_cost_per_token": 2.7e-06, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", + "output_cost_per_reasoning_token": 1.5e-06, + "output_cost_per_token": 1.5e-06, + "output_cost_per_token_batches": 7.5e-07, + "output_cost_per_token_flex": 7.5e-07, + "output_cost_per_token_priority": 2.7e-06, + "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-3.1-flash-lite", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -17128,10 +17137,16 @@ "supports_service_tier": true }, "gemini/gemini-3.1-flash-lite": { - "cache_read_input_token_cost": 4.5e-08, - "cache_read_input_token_cost_per_audio_token": 9e-08, - "input_cost_per_audio_token": 9e-07, - "input_cost_per_token": 4.5e-07, + "cache_read_input_token_cost": 2.5e-08, + "cache_read_input_token_cost_batches": 1.25e-08, + "cache_read_input_token_cost_flex": 1.25e-08, + "cache_read_input_token_cost_per_audio_token": 5e-08, + "cache_read_input_token_cost_priority": 4.5e-08, + "input_cost_per_audio_token": 5e-07, + "input_cost_per_token": 2.5e-07, + "input_cost_per_token_batches": 1.25e-07, + "input_cost_per_token_flex": 1.25e-07, + "input_cost_per_token_priority": 4.5e-07, "litellm_provider": "gemini", "max_audio_length_hours": 8.4, "max_audio_per_prompt": 1, @@ -17143,10 +17158,13 @@ "max_video_length": 1, "max_videos_per_prompt": 10, "mode": "chat", - "output_cost_per_reasoning_token": 2.7e-06, - "output_cost_per_token": 2.7e-06, + "output_cost_per_reasoning_token": 1.5e-06, + "output_cost_per_token": 1.5e-06, + "output_cost_per_token_batches": 7.5e-07, + "output_cost_per_token_flex": 7.5e-07, + "output_cost_per_token_priority": 2.7e-06, "rpm": 15, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-3.1-flash-lite", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -33932,10 +33950,16 @@ "web_search_billing_unit": "per_query" }, "vertex_ai/gemini-3.1-flash-lite": { - "cache_read_input_token_cost": 4.5e-08, - "cache_read_input_token_cost_per_audio_token": 9e-08, - "input_cost_per_audio_token": 9e-07, - "input_cost_per_token": 4.5e-07, + "cache_read_input_token_cost": 2.5e-08, + "cache_read_input_token_cost_batches": 1.25e-08, + "cache_read_input_token_cost_flex": 1.25e-08, + "cache_read_input_token_cost_per_audio_token": 5e-08, + "cache_read_input_token_cost_priority": 4.5e-08, + "input_cost_per_audio_token": 5e-07, + "input_cost_per_token": 2.5e-07, + "input_cost_per_token_batches": 1.25e-07, + "input_cost_per_token_flex": 1.25e-07, + "input_cost_per_token_priority": 4.5e-07, "litellm_provider": "vertex_ai-language-models", "max_audio_length_hours": 8.4, "max_audio_per_prompt": 1, @@ -33947,8 +33971,11 @@ "max_video_length": 1, "max_videos_per_prompt": 10, "mode": "chat", - "output_cost_per_reasoning_token": 2.7e-06, - "output_cost_per_token": 2.7e-06, + "output_cost_per_reasoning_token": 1.5e-06, + "output_cost_per_token": 1.5e-06, + "output_cost_per_token_batches": 7.5e-07, + "output_cost_per_token_flex": 7.5e-07, + "output_cost_per_token_priority": 2.7e-06, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", "supported_endpoints": [ "/v1/chat/completions", diff --git a/tests/code_coverage_tests/check_licenses.py b/tests/code_coverage_tests/check_licenses.py index 668aefa8024..5fb2b495c24 100644 --- a/tests/code_coverage_tests/check_licenses.py +++ b/tests/code_coverage_tests/check_licenses.py @@ -31,6 +31,13 @@ DEFAULT_TRANSITIVE_PIN_PACKAGES = ( "wheel", ) +# SPDX license expressions (PEP 639 "License-Expression") join identifiers with +# the uppercase operators OR / AND / WITH. The split is case-sensitive: the +# lowercase "-or-later" inside an identifier such as "GPL-2.0-or-later" is part +# of the identifier, not an operator. +_SPDX_OPERATOR_SPLIT = re.compile(r"\s+(?:OR|AND)\s+") +_SPDX_WITH_SUFFIX = re.compile(r"\s+WITH\s+.*", re.DOTALL) + @dataclass class PackageLicense: @@ -109,21 +116,86 @@ class LicenseChecker: def get_package_license_from_pypi( self, package_name: str, version: str ) -> Optional[str]: - """Fetch license information for a package from PyPI.""" + """Fetch license information for a package from PyPI. + + Prefers the PEP 639 SPDX expression (``info.license_expression``), + falls back to the legacy free-text ``info.license`` field, and as a + last resort derives the license from the ``License :: OSI Approved :: + ...`` trove classifiers. + """ try: url = f"https://pypi.org/pypi/{package_name}/{version}/json" response = requests.get(url, timeout=10) response.raise_for_status() - data = response.json() - return data.get("info", {}).get("license") + info = response.json().get("info", {}) or {} + return ( + info.get("license_expression") + or info.get("license") + or self._license_from_classifiers(info.get("classifiers") or []) + ) except Exception as e: print( f"Warning: Failed to fetch license for {package_name} {version}: {str(e)}" ) return None - def is_license_acceptable(self, license_str: str) -> Tuple[bool, str]: - """Check if a license is acceptable based on configured lists.""" + @staticmethod + def _license_from_classifiers(classifiers: List[str]) -> Optional[str]: + """Derive a license name from the ``License :: OSI Approved :: ...`` trove classifiers.""" + prefix = "License :: OSI Approved :: " + for classifier in classifiers: + if classifier.startswith(prefix): + license_name = classifier[len(prefix) :].strip() + if license_name: + return license_name + return None + + @staticmethod + def _split_spdx_expression(license_str: str) -> Optional[List[str]]: + """Split an SPDX license expression into its component identifiers. + + Returns ``None`` when the string is not a recognizable SPDX expression + (for example a free-text license blob), so callers fall back to + whole-string matching. + """ + if "OR" not in license_str and "AND" not in license_str: + return None + + components: List[str] = [] + normalized = license_str.replace("(", " ").replace(")", " ") + for part in _SPDX_OPERATOR_SPLIT.split(normalized): + # Drop any "WITH " suffix: the exception qualifies the + # preceding license, it is not itself a license to authorize. + identifier = _SPDX_WITH_SUFFIX.sub("", part).strip() + if not identifier: + continue + # SPDX short-form identifiers are single whitespace-free tokens; a + # component with internal whitespace means this is free text. + if any(char.isspace() for char in identifier): + return None + components.append(identifier) + + return components if len(components) > 1 else None + + def is_license_acceptable(self, license_str: Optional[str]) -> Tuple[bool, str]: + """Check if a license (or compound SPDX expression) is acceptable.""" + if not license_str: + return False, "Unknown license" + + components = self._split_spdx_expression(license_str) + if components is None: + return self._is_single_license_acceptable(license_str) + + # Compound SPDX expression: conservatively require every component to + # be acceptable on its own (the safe direction for a CI gate). + for component in components: + is_acceptable, reason = self._is_single_license_acceptable(component) + if not is_acceptable: + return False, f"{reason} (in SPDX expression '{license_str}')" + return True, f"All SPDX components authorized: {', '.join(components)}" + + def _is_single_license_acceptable(self, license_str: str) -> Tuple[bool, str]: + """Check if a single license identifier is acceptable based on configured lists.""" if not license_str: return False, "Unknown license" diff --git a/tests/code_coverage_tests/liccheck.ini b/tests/code_coverage_tests/liccheck.ini index 0d1a6f0b045..5a09403c570 100644 --- a/tests/code_coverage_tests/liccheck.ini +++ b/tests/code_coverage_tests/liccheck.ini @@ -90,7 +90,6 @@ jinja2: >=3.1.4 # BSD 3-Clause License litellm-proxy-extras: >=0.1.1 # MIT License litellm-enterprise: >=0.1.1 # LiteLLM Enterprise License a2a-sdk: >=0.3.22 # Apache 2.0 license -pydantic-settings: >=2.14.1 # MIT License manually verified (uses PEP 639 License-Expression: MIT, not the legacy License field, so liccheck reports it as unknown) anyio: >=4.5.0 # Unknown license httpx-aiohttp: >=0.1.4 # Unknown license backoff: >=2.2.1 # Unknown license @@ -156,7 +155,6 @@ pytest: >=9.0.3 # MIT license pytest-postgresql: >=7.0.2 # LGPLv3+ license pytest-xdist: >=3.8.0 # MIT License ruff: >=0.15.3 # MIT License -black: >=26.3.1 # MIT License manually verified (uses PEP 639 License-Expression: MIT, not the legacy License field, so liccheck reports it as unknown) types-requests: >=2.32.4.20260107 # Apache 2.0 license (typeshed) types-pyyaml: >=6.0.12.20250915 # Apache 2.0 license (typeshed) fakeredis: >=2.34.1 # BSD license diff --git a/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py b/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py index e2c50810cc2..37fcc602d37 100644 --- a/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py +++ b/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py @@ -41,6 +41,62 @@ from litellm.types.llms.openai import ( class TestBaseResponsesAPIStreamingIterator: """Test cases for BaseResponsesAPIStreamingIterator""" + @pytest.mark.asyncio + async def test_responses_streaming_iterator_parses_u2028_in_sse_json(self): + """ + U+2028 inside JSON must not split the SSE event. httpx aiter_lines uses + str.splitlines() and drops response.completed; OpenAI SSEDecoder does not. + """ + from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator + + u2028 = "\u2028" + payload = json.dumps( + { + "type": "response.completed", + "response": {"instructions": f"eligible{u2028}promo"}, + } + ) + sse_bytes = f"data: {payload}\n\n".encode("utf-8") + + async def mock_aiter_bytes(): + yield sse_bytes + + mock_response = Mock() + mock_response.headers = {} + mock_response.aiter_bytes = mock_aiter_bytes + + mock_logging_obj = Mock(spec=LiteLLMLoggingObj) + mock_logging_obj.model_call_details = {"litellm_params": {}} + mock_config = Mock(spec=BaseResponsesAPIConfig) + + mock_responses_api_response = Mock(spec=ResponsesAPIResponse) + mock_responses_api_response.id = "resp_u2028" + mock_completed_event = Mock(spec=ResponseCompletedEvent) + mock_completed_event.type = ResponsesAPIStreamEvents.RESPONSE_COMPLETED + mock_completed_event.response = mock_responses_api_response + mock_config.transform_streaming_response.return_value = mock_completed_event + + iterator = ResponsesAPIStreamingIterator( + response=mock_response, + model="gpt-5.5", + responses_api_provider_config=mock_config, + logging_obj=mock_logging_obj, + litellm_metadata={"model_info": {"id": "model_123"}}, + custom_llm_provider="openai", + ) + + chunks = [] + with ( + patch("asyncio.create_task"), + patch("litellm.responses.streaming_iterator.executor"), + ): + async for chunk in iterator: + chunks.append(chunk) + + assert len(chunks) == 1 + assert chunks[0].type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + assert iterator.completed_response is not None + def test_process_chunk_with_response_completed_event(self): """ Test that _process_chunk correctly processes a ResponseCompletedEvent @@ -270,7 +326,7 @@ class TestBaseResponsesAPIStreamingIterator: # Mock dependencies mock_response = Mock() mock_response.headers = {} - mock_response.aiter_lines = Mock() + mock_response.aiter_bytes = Mock() mock_logging_obj = Mock(spec=LiteLLMLoggingObj) mock_logging_obj.model_call_details = {"litellm_params": {}} mock_logging_obj.async_success_handler = Mock() @@ -334,12 +390,10 @@ class TestBaseResponsesAPIStreamingIterator: mock_response = Mock() mock_response.headers = {} - # Create an async iterator that raises StopAsyncIteration after yielding one chunk - async def mock_aiter_lines(): - yield 'data: {"type": "response.output_text.delta", "delta": "test"}' - # Normal end of stream - raise StopAsyncIteration + async def mock_aiter_bytes(): + yield b'data: {"type": "response.output_text.delta", "delta": "test"}\n\n' - mock_response.aiter_lines = mock_aiter_lines + mock_response.aiter_bytes = mock_aiter_bytes mock_logging_obj = Mock(spec=LiteLLMLoggingObj) mock_logging_obj.model_call_details = {"litellm_params": {}} @@ -396,12 +450,10 @@ class TestBaseResponsesAPIStreamingIterator: mock_response = Mock() mock_response.headers = {} - # Create a sync iterator that raises StopIteration after yielding one chunk - def mock_iter_lines(): - yield 'data: {"type": "response.output_text.delta", "delta": "test"}' - # Normal end of stream - raise StopIteration + def mock_iter_bytes(): + yield b'data: {"type": "response.output_text.delta", "delta": "test"}\n\n' - mock_response.iter_lines = mock_iter_lines + mock_response.iter_bytes = mock_iter_bytes mock_logging_obj = Mock(spec=LiteLLMLoggingObj) mock_logging_obj.model_call_details = {"litellm_params": {}} @@ -450,7 +502,7 @@ class TestBaseResponsesAPIStreamingIterator: mock_response = Mock() mock_response.headers = {} - mock_response.aiter_lines = Mock() + mock_response.aiter_bytes = Mock() mock_logging_obj = Mock(spec=LiteLLMLoggingObj) mock_logging_obj.model_call_details = {"litellm_params": {}} mock_logging_obj.async_failure_handler = Mock() @@ -532,7 +584,7 @@ class TestBaseResponsesAPIStreamingIterator: mock_response = Mock() mock_response.headers = {} - mock_response.aiter_lines = Mock() + mock_response.aiter_bytes = Mock() mock_logging_obj = Mock(spec=LiteLLMLoggingObj) mock_logging_obj.model_call_details = {"litellm_params": {}} mock_logging_obj.async_failure_handler = Mock() diff --git a/tests/local_testing/test_streaming.py b/tests/local_testing/test_streaming.py index b1a93c380b2..10f351714e1 100644 --- a/tests/local_testing/test_streaming.py +++ b/tests/local_testing/test_streaming.py @@ -993,6 +993,11 @@ def test_vertex_ai_stream(provider): except litellm.RateLimitError as e: pass + except litellm.exceptions.MidStreamFallbackError as e: + # Streaming 429s are wrapped in MidStreamFallbackError so the + # Router can fall back; treat as a transient rate-limit pass. + if not isinstance(e.original_exception, litellm.RateLimitError): + pytest.fail(f"Error occurred: {e}") except Exception as e: pytest.fail(f"Error occurred: {e}") diff --git a/tests/proxy_admin_ui_tests/e2e_ui_tests/login_to_ui.spec.ts b/tests/proxy_admin_ui_tests/e2e_ui_tests/login_to_ui.spec.ts deleted file mode 100644 index e5a397a6a66..00000000000 --- a/tests/proxy_admin_ui_tests/e2e_ui_tests/login_to_ui.spec.ts +++ /dev/null @@ -1,51 +0,0 @@ -/* - -Login to Admin UI -Basic UI Test - -Click on all the tabs ensure nothing is broken -*/ - -import { test, expect } from "@playwright/test"; - -test("admin login test", async ({ page }) => { - // Go to the specified URL - await page.goto("http://localhost:4000/ui"); - await page.waitForLoadState("networkidle"); - - await page.screenshot({ path: "test-results/login_before.png" }); - - // Enter "admin" in the username input field - await page.fill('input[placeholder="Enter your username"]', "admin"); - - // Enter "gm" in the password input field - await page.fill('input[placeholder="Enter your password"]', "gm"); - - page.screenshot({ path: "test-results/login_after_inputs.png" }); - - // Optionally, you can add an assertion to verify the login button is enabled - const loginButton = page.getByRole("button", { name: "Login" }); - await expect(loginButton).toBeEnabled(); - - // Optionally, you can click the login button to submit the form - await loginButton.click(); - const tabs = [ - "Virtual Keys", - "Playground", - "Models", - "Usage", - "Teams", - "Internal User", - "Settings", - "Experimental", - "API Reference", - "AI Hub", - ]; - - for (const tab of tabs) { - const tabElement = page.locator("span.ant-menu-title-content", { - hasText: tab, - }); - await tabElement.click(); - } -}); diff --git a/tests/proxy_admin_ui_tests/e2e_ui_tests/redirect-fail-screenshot.png b/tests/proxy_admin_ui_tests/e2e_ui_tests/redirect-fail-screenshot.png deleted file mode 100644 index b2e33251260..00000000000 Binary files a/tests/proxy_admin_ui_tests/e2e_ui_tests/redirect-fail-screenshot.png and /dev/null differ diff --git a/tests/proxy_admin_ui_tests/e2e_ui_tests/require_auth_for_dashboard.spec.ts b/tests/proxy_admin_ui_tests/e2e_ui_tests/require_auth_for_dashboard.spec.ts deleted file mode 100644 index 4e4bd2fcd93..00000000000 --- a/tests/proxy_admin_ui_tests/e2e_ui_tests/require_auth_for_dashboard.spec.ts +++ /dev/null @@ -1,37 +0,0 @@ -// tests/auth.spec.ts -import { test, expect } from "@playwright/test"; - -test.describe("Authentication Checks", () => { - test("should redirect unauthenticated user from a protected page", async ({ - page, - }) => { - test.setTimeout(30000); - - page.on("console", (msg) => console.log("PAGE LOG:", msg.text())); - - const protectedPageUrl = "http://localhost:4000/ui?page=llm-playground"; - const expectedRedirectUrl = "http://localhost:4000/ui/login/"; - - console.log( - `Attempting to navigate to protected page: ${protectedPageUrl}` - ); - - await page.goto(protectedPageUrl); - - console.log(`Navigation initiated. Current URL: ${page.url()}`); - - try { - await page.waitForURL(expectedRedirectUrl, { timeout: 10000 }); - console.log(`Waited for URL. Current URL is now: ${page.url()}`); - } catch (error) { - console.error( - `Timeout waiting for URL: ${expectedRedirectUrl}. Current URL: ${page.url()}` - ); - await page.screenshot({ path: "redirect-fail-screenshot.png" }); - throw error; - } - - await expect(page).toHaveURL(expectedRedirectUrl); - console.log(`Assertion passed: Page URL is ${expectedRedirectUrl}`); - }); -}); diff --git a/tests/proxy_admin_ui_tests/e2e_ui_tests/search_users.spec.ts b/tests/proxy_admin_ui_tests/e2e_ui_tests/search_users.spec.ts deleted file mode 100644 index d72c44ab8cc..00000000000 --- a/tests/proxy_admin_ui_tests/e2e_ui_tests/search_users.spec.ts +++ /dev/null @@ -1,222 +0,0 @@ -/* -Search Users in Admin UI -E2E Test for user search functionality - -Tests: -1. Navigate to Internal Users tab -2. Verify search input exists -3. Test search functionality -4. Verify results update -5. Test filtering by email, user ID, and SSO user ID -*/ - -import { test, expect } from "@playwright/test"; - -test("user search test", async ({ page }) => { - // Set a longer timeout for the entire test - test.setTimeout(60000); - - // Enable console logging - page.on("console", (msg) => console.log("PAGE LOG:", msg.text())); - - // Login first - await page.goto("http://localhost:4000/ui"); - await page.waitForLoadState("networkidle"); - console.log("Navigated to login page"); - - page.screenshot({ path: "test-results/search_users_before_login.png" }); - - // Wait for login form to be visible - await page.waitForSelector('input[placeholder="Enter your username"]', { - timeout: 10000, - }); - console.log("Login form is visible"); - - await page.fill('input[placeholder="Enter your username"]', "admin"); - await page.fill('input[placeholder="Enter your password"]', "gm"); - console.log("Filled login credentials"); - - const loginButton = page.getByRole("button", { name: "Login" }); - await expect(loginButton).toBeEnabled(); - await loginButton.click(); - console.log("Clicked login button"); - - // Wait for navigation to complete and dashboard to load - await page.waitForLoadState("networkidle"); - console.log("Page loaded after login"); - - // Take a screenshot for debugging - await page.screenshot({ path: "after-login.png" }); - console.log("Took screenshot after login"); - - // Try to find the Internal User tab with more debugging - console.log("Looking for Internal User tab..."); - const internalUserTab = page.locator("span.ant-menu-title-content", { - hasText: "Internal User", - }); - - // Wait for the tab to be visible - await internalUserTab.waitFor({ state: "visible", timeout: 10000 }); - console.log("Internal User tab is visible"); - - // Take another screenshot before clicking - await page.screenshot({ path: "before-tab-click.png" }); - console.log("Took screenshot before tab click"); - - await internalUserTab.click(); - console.log("Clicked Internal User tab"); - - // Wait for the page to load and table to be visible - await page.waitForSelector("tbody tr", { timeout: 30000 }); - await page.waitForTimeout(2000); // Additional wait for table to stabilize - console.log("Table is visible"); - - // Take a final screenshot - await page.screenshot({ path: "after-tab-click.png" }); - console.log("Took screenshot after tab click"); - - // Verify search input exists - const searchInput = page.locator('input[placeholder="Search by email..."]'); - await expect(searchInput).toBeVisible(); - console.log("Search input is visible"); - - // Test search functionality - const initialUserCount = await page.locator("tbody tr").count(); - console.log(`Initial user count: ${initialUserCount}`); - - // Perform a search - const testEmail = "test@"; - await searchInput.fill(testEmail); - console.log("Filled search input"); - - // Wait for the debounced search to complete - await page.waitForTimeout(500); - console.log("Waited for debounce"); - - // Wait for the results count to update - await page.waitForFunction((initialCount) => { - const currentCount = document.querySelectorAll("tbody tr").length; - return currentCount !== initialCount; - }, initialUserCount); - console.log("Results updated"); - - const filteredUserCount = await page.locator("tbody tr").count(); - console.log(`Filtered user count: ${filteredUserCount}`); - - expect(filteredUserCount).toBeDefined(); - - // Clear the search - await searchInput.clear(); - console.log("Cleared search"); - - await page.waitForTimeout(500); - console.log("Waited for debounce after clear"); - - await page.waitForFunction((initialCount) => { - const currentCount = document.querySelectorAll("tbody tr").length; - return currentCount === initialCount; - }, initialUserCount); - console.log("Results reset"); - - const resetUserCount = await page.locator("tbody tr").count(); - console.log(`Reset user count: ${resetUserCount}`); - - expect(resetUserCount).toBe(initialUserCount); -}); - -test("user filter test", async ({ page }) => { - // Set a longer timeout for the entire test - test.setTimeout(60000); - - // Enable console logging - page.on("console", (msg) => console.log("PAGE LOG:", msg.text())); - - // Login first - await page.goto("http://localhost:4000/ui"); - await page.waitForLoadState("networkidle"); - console.log("Navigated to login page"); - - // Wait for login form to be visible - await page.waitForSelector('input[placeholder="Enter your username"]', { - timeout: 10000, - }); - console.log("Login form is visible"); - - await page.fill('input[placeholder="Enter your username"]', "admin"); - await page.fill('input[placeholder="Enter your password"]', "gm"); - console.log("Filled login credentials"); - - const loginButton = page.getByRole("button", { name: "Login" }); - await expect(loginButton).toBeEnabled(); - await loginButton.click(); - console.log("Clicked login button"); - - // Wait for navigation to complete and dashboard to load - await page.waitForLoadState("networkidle"); - console.log("Page loaded after login"); - - // Navigate to Internal Users tab - const internalUserTab = page.locator("span.ant-menu-title-content", { - hasText: "Internal User", - }); - await internalUserTab.waitFor({ state: "visible", timeout: 10000 }); - await internalUserTab.click(); - console.log("Clicked Internal User tab"); - - // Wait for the page to load and table to be visible - await page.waitForSelector("tbody tr", { timeout: 30000 }); - await page.waitForTimeout(2000); // Additional wait for table to stabilize - console.log("Table is visible"); - - // Get initial user count - const initialUserCount = await page.locator("tbody tr").count(); - console.log(`Initial user count: ${initialUserCount}`); - - // Click the filter button to show additional filters - const filterButton = page.getByRole("button", { - name: "Filters", - exact: true, - }); - await filterButton.click(); - console.log("Clicked filter button"); - await page.waitForTimeout(500); // Wait for filters to appear - - // Test user ID filter - const userIdInput = page.locator('input[placeholder="Filter by User ID"]'); - await expect(userIdInput).toBeVisible(); - console.log("User ID filter is visible"); - - await userIdInput.fill("user"); - console.log("Filled user ID filter"); - await page.waitForTimeout(1000); - const userIdFilteredCount = await page.locator("tbody tr").count(); - console.log(`User ID filtered count: ${userIdFilteredCount}`); - expect(userIdFilteredCount).toBeLessThan(initialUserCount); - - // Clear user ID filter - await userIdInput.clear(); - await page.waitForTimeout(1000); - console.log("Cleared user ID filter"); - - // Test SSO user ID filter - const ssoUserIdInput = page.locator('input[placeholder="Filter by SSO ID"]'); - await expect(ssoUserIdInput).toBeVisible(); - console.log("SSO user ID filter is visible"); - - await ssoUserIdInput.fill("sso"); - console.log("Filled SSO user ID filter"); - await page.waitForTimeout(1000); - const ssoUserIdFilteredCount = await page.locator("tbody tr").count(); - console.log(`SSO user ID filtered count: ${ssoUserIdFilteredCount}`); - expect(ssoUserIdFilteredCount).toBeLessThan(initialUserCount); - - // Clear SSO user ID filter - await ssoUserIdInput.clear(); - await page.waitForTimeout(5000); - console.log("Cleared SSO user ID filter"); - - // Verify count returns to initial after clearing all filters - const finalUserCount = await page.locator("tbody tr").count(); - console.log(`Final user count: ${finalUserCount}`); - expect(finalUserCount).toBe(initialUserCount); -}); diff --git a/tests/proxy_admin_ui_tests/e2e_ui_tests/team_admin.spec.ts b/tests/proxy_admin_ui_tests/e2e_ui_tests/team_admin.spec.ts deleted file mode 100644 index a753c724b37..00000000000 --- a/tests/proxy_admin_ui_tests/e2e_ui_tests/team_admin.spec.ts +++ /dev/null @@ -1,250 +0,0 @@ -import { test, expect } from "@playwright/test"; -import { loginToUI } from "../utils/login"; - -// test.describe("Invite User, Set Password, and Login", () => { -// let testEmail: string; -// const testPassword = "Password123!"; // Define a password -// const teamName1 = `team-invite-test-1-${Date.now()}`; -// const teamName2 = `team-invite-test-2-${Date.now()}`; -// const keyName1 = `key-${teamName1}`; -// const keyName2 = `key-${teamName2}`; - -// test.beforeEach(async ({ page }) => { -// await loginToUI(page); // Login as admin first -// await page.goto("http://localhost:4000/ui?page=teams"); - -// // --- Create Team 1 --- -// await page.getByRole("button", { name: "+ Create New Team" }).click(); -// await page -// .getByLabel("Team Name") -// .waitFor({ state: "visible", timeout: 5000 }); // Wait for label -// await page.getByLabel("Team Name").click(); -// await page.getByLabel("Team Name").fill(teamName1); -// await page.getByRole("button", { name: "Create Team" }).click(); -// // Wait for the modal to close or for a success message if applicable -// await expect( -// page.locator(".ant-modal-wrap").filter({ hasText: "Create New Team" }) -// ).not.toBeVisible({ timeout: 10000 }); -// console.log(`Created Team 1: ${teamName1}`); - -// // --- Create Team 2 --- -// await page.getByRole("button", { name: "+ Create New Team" }).click(); -// await page -// .getByLabel("Team Name") -// .waitFor({ state: "visible", timeout: 5000 }); // Wait for label -// await page.getByLabel("Team Name").click(); -// await page.getByLabel("Team Name").fill(teamName2); -// await page.getByRole("button", { name: "Create Team" }).click(); -// // Wait for the modal to close or for a success message if applicable -// await expect( -// page.locator(".ant-modal-wrap").filter({ hasText: "Create New Team" }) -// ).not.toBeVisible({ timeout: 10000 }); -// console.log(`Created Team 2: ${teamName2}`); - -// // // Verify both teams are listed -// // await page.goto("http://localhost:4000/ui?page=teams"); // Refresh or ensure on teams page -// // await page.waitForTimeout(3000); -// await expect(page.getByText(teamName1)).toBeVisible({ timeout: 10000 }); -// await expect(page.getByText(teamName2)).toBeVisible({ timeout: 10000 }); - -// // --- Navigate to Keys Page --- -// await page.goto("http://localhost:4000/ui?page=api-keys"); -// await page.waitForTimeout(3000); -// await expect( -// page.getByRole("button", { name: "+ Create New Key" }) -// ).toBeVisible(); // Wait for page load - -// // --- Create Key for Team 1 --- -// await page.getByRole("button", { name: "+ Create New Key" }).click(); -// const createKeyModal1 = page -// .locator(".ant-modal-wrap") -// .filter({ hasText: "Key Ownership" }); -// await expect(createKeyModal1).toBeVisible(); - -// // Select Team 1 -// await createKeyModal1 -// .locator(".ant-select-selector >> input") -// .first() -// .click(); // Click to open team dropdown -// await createKeyModal1 -// .locator(".ant-select-selector >> input") -// .first() -// .fill(teamName1); - -// await page -// .locator(".ant-select-item-option") -// .filter({ hasText: teamName1 }) -// .first() -// .click(); // Click specific team name - -// // Enter Key Name 1 -// await page.fill('input[id="key_alias"]', keyName1); - -// // Click on models dropdown -// await page.locator("input#models").click(); -// await page.waitForSelector( -// '.ant-select-item-option[title="All Team Models"]' -// ); -// await page -// .locator('.ant-select-item-option[title="All Team Models"]') -// .click(); - -// // Click Create Key -// await createKeyModal1.getByRole("button", { name: "Create Key" }).click(); - -// // Close the Key Generated modal (which appears after successful creation) -// const keyGeneratedModal1 = page -// .locator(".ant-modal-wrap") -// .filter({ hasText: "Save your Key" }); -// await expect(keyGeneratedModal1).toBeVisible({ timeout: 10000 }); -// await keyGeneratedModal1.locator('button[aria-label="Close"]').click(); -// await expect(keyGeneratedModal1).not.toBeVisible(); // Wait for close -// console.log(`Created Key 1: ${keyName1} for Team: ${teamName1}`); - -// // --- Create Key for Team 2 --- -// await page.getByRole("button", { name: "+ Create New Key" }).click(); -// const createKeyModal2 = page -// .locator(".ant-modal-wrap") -// .filter({ hasText: "Key Ownership" }); -// await expect(createKeyModal2).toBeVisible(); - -// // Select Team 2 -// await createKeyModal2 -// .locator(".ant-select-selector >> input") -// .first() -// .click(); // Click to open team dropdown -// await page -// .locator(".ant-select-item-option") -// .filter({ hasText: teamName2 }) -// .click(); // Click specific team name - -// // Enter Key Name 2 -// await page.fill('input[id="key_alias"]', keyName2); - -// // Click on models dropdown -// await page.locator("input#models").click(); -// await page.waitForSelector( -// '.ant-select-item-option[title="All Team Models"]' -// ); -// await page -// .locator('.ant-select-item-option[title="All Team Models"]') -// .click(); - -// // Click Create Key -// await createKeyModal2.getByRole("button", { name: "Create Key" }).click(); - -// // Close the Key Generated modal -// const keyGeneratedModal2 = page -// .locator(".ant-modal-wrap") -// .filter({ hasText: "Save your Key" }); -// await expect(keyGeneratedModal2).toBeVisible({ timeout: 10000 }); -// await keyGeneratedModal2.locator('button[aria-label="Close"]').click(); -// await expect(keyGeneratedModal2).not.toBeVisible(); // Wait for close -// console.log(`Created Key 2: ${keyName2} for Team: ${teamName2}`); -// }); - -// test("Invite user, set password via link, and login", async ({ page }) => { -// // Navigate to Users page -// await page.goto("http://localhost:4000/ui?page=users"); - -// // Go to Internal User tab -// const internalUserTab = page.locator("span.ant-menu-title-content", { -// hasText: "Internal User", -// }); -// await internalUserTab.waitFor({ state: "visible", timeout: 10000 }); -// await internalUserTab.click(); - -// // --- Invite User Flow --- -// await page.getByRole("button", { name: "+ Invite User" }).click(); - -// // Wait for the invite user modal to be visible -// const inviteModal = page -// .locator(".ant-modal-wrap") -// .filter({ hasText: "Invite User" }); -// await expect(inviteModal).toBeVisible(); - -// testEmail = `test-${Date.now()}@litellm.ai`; // Use a unique email -// // Assuming the email input is the first one with 'base-input' test id inside the modal -// await inviteModal.getByTestId("base-input").first().fill(testEmail); - -// // Select Global Admin Role (or another appropriate role) -// const globalRoleLabel = inviteModal.getByLabel("Global Proxy Role"); -// await globalRoleLabel.click(); -// // Wait for the dropdown option to be visible before clicking -// const adminRoleOption = page.getByTitle("Admin (All Permissions)", { -// exact: true, -// }); -// await adminRoleOption.waitFor({ state: "visible", timeout: 5000 }); -// await adminRoleOption.click(); - -// // Select Team - Add explicit wait before clicking -// const teamIdLabel = inviteModal.getByLabel("Team ID"); -// // Wait for the label associated with the Team ID select to be visible -// await teamIdLabel.waitFor({ state: "visible", timeout: 10000 }); // Increased timeout for safety -// await teamIdLabel.click(); - -// // Wait for the team name option to be visible in the dropdown -// const teamNameOption = page.getByText(teamName1, { exact: true }); -// await teamNameOption.waitFor({ state: "visible", timeout: 5000 }); -// await teamNameOption.click(); - -// // Create User -// await inviteModal.getByRole("button", { name: "Create User" }).click(); - -// // --- Capture Invitation Link --- -// const invitationModal = page -// .locator(".ant-modal-wrap") -// .filter({ hasText: "Invitation Link" }); -// await expect(invitationModal).toBeVisible({ timeout: 15000 }); // Wait longer for modal - -// // Locate the text element containing the URL more reliably -// const invitationUrl = await page -// .locator("div.flex.justify-between.pt-5.pb-2") // find the correct div -// .filter({ hasText: "Invitation Link" }) // find the div that has text "Invitation Link" -// .locator("p") // find all

inside that div -// .nth(1) // pick the second

(index 1) -// .innerText(); - -// // Close Invitation Link Modal -// await page -// .locator(".ant-modal-wrap") -// .filter({ hasText: "Invitation Link" }) -// .locator('button[aria-label="Close"]') -// .click(); - -// // Close Invite User Modal -// await page -// .locator(".ant-modal-wrap") -// .filter({ hasText: "Invite User" }) -// .locator('button[aria-label="Close"]') -// .click(); - -// // Open invite link as new page (simulate invited user) -// const context = await page.context()?.browser()?.newContext(); -// const invitedUserPage = await context?.newPage(); -// if (!invitedUserPage) { -// throw new Error("invitedUserPage is undefined"); -// } -// await invitedUserPage?.goto(invitationUrl || ""); - -// //Insert new password -// await invitedUserPage?.fill("input#password", testPassword); - -// //Click on submit -// await invitedUserPage?.getByRole("button", { name: "Sign Up" }).click(); - -// // // --- Verify Keys Created --- -// // await invitedUserPage?.waitForSelector("table"); - -// // // Verify keyName1 (associated with user's team) IS visible in the table -// // const keyTable = invitedUserPage.locator('table'); // Locate the table element -// // await expect(keyTable).toBeVisible({ timeout: 10000 }); // Ensure table exists -// // // Use getByText within the table scope to find the key name -// // await expect(keyTable.getByText(keyName1, { exact: true })).toBeVisible({ timeout: 10000 }); -// // console.log(`Verified key ${keyName1} is visible for user ${testEmail}`); - -// // // Verify keyName2 (associated with the *other* team) IS NOT visible -// // await expect(keyTable.getByText(keyName2, { exact: true })).not.toBeVisible(); -// // console.log(`Verified key ${keyName2} is NOT visible for user ${testEmail}`); -// }); -// }); diff --git a/tests/proxy_admin_ui_tests/e2e_ui_tests/view_internal_user.spec.ts b/tests/proxy_admin_ui_tests/e2e_ui_tests/view_internal_user.spec.ts deleted file mode 100644 index 832832d8ae8..00000000000 --- a/tests/proxy_admin_ui_tests/e2e_ui_tests/view_internal_user.spec.ts +++ /dev/null @@ -1,72 +0,0 @@ -/* -Test view internal user page -*/ - -import { test, expect } from "@playwright/test"; - -test("view internal user page", async ({ page }) => { - // Go to the specified URL - await page.goto("http://localhost:4000/ui"); - await page.waitForLoadState("networkidle"); - - page.screenshot({ path: "test-results/view_internal_user_before_login.png" }); - - // Enter "admin" in the username input field - await page.fill('input[placeholder="Enter your username"]', "admin"); - - // Enter "gm" in the password input field - await page.fill('input[placeholder="Enter your password"]', "gm"); - - // Click the login button - const loginButton = page.getByRole("button", { name: "Login" }); - await expect(loginButton).toBeEnabled(); - await loginButton.click(); - - // Wait for the Internal User tab and click it - const tabElement = page.locator("span.ant-menu-title-content", { - hasText: "Internal User", - }); - await tabElement.click(); - - // Wait for the table to load - await page.waitForSelector("tbody tr", { timeout: 10000 }); - await page.waitForTimeout(2000); // Additional wait for table to stabilize - await page.waitForLoadState("networkidle"); - - // Test all expected fields are present - // Verify that the API Keys column is rendered for all users - // The UI renders badges in each row - we just verify the column structure exists - const rowCount = await page.locator("tbody tr").count(); - expect(rowCount).toBeGreaterThan(0); - - const userIdHeader = await page.locator("th", { hasText: "User ID" }); - await expect(userIdHeader).toBeVisible({ timeout: 10000 }); - - // test pagination - // Wait for pagination controls to be visible - await page.waitForSelector(".flex.justify-between.items-center", { - timeout: 5000, - }); - - // Check if we're on the first page by looking at the results count - const resultsText = - (await page.locator(".text-sm.text-gray-700").textContent()) || ""; - const isFirstPage = resultsText.includes("1 -"); - - if (isFirstPage) { - // On first page, previous button should be disabled - const prevButton = page.locator("button", { hasText: "Previous" }); - await expect(prevButton).toBeDisabled(); - } - - // Next button should be enabled if there are more pages - const nextButton = page.locator("button", { hasText: "Next" }); - const totalResults = - (await page.locator(".text-sm.text-gray-700").textContent()) || ""; - const hasMorePages = - totalResults.includes("of") && !totalResults.includes("1 - 25 of 25"); - - if (hasMorePages) { - await expect(nextButton).toBeEnabled(); - } -}); diff --git a/tests/proxy_admin_ui_tests/e2e_ui_tests/view_user_info.spec.ts b/tests/proxy_admin_ui_tests/e2e_ui_tests/view_user_info.spec.ts deleted file mode 100644 index adda3088f12..00000000000 --- a/tests/proxy_admin_ui_tests/e2e_ui_tests/view_user_info.spec.ts +++ /dev/null @@ -1,124 +0,0 @@ -import { test, expect } from "@playwright/test"; -import { loginToUI } from "../utils/login"; - -test.describe("User Info View", () => { - test("should display user info when clicking on user ID", async ({ - page, - }) => { - await page.goto("http://localhost:4000/ui"); - await page.waitForLoadState("networkidle"); - - page.screenshot({ - path: "test-results/view_user_info_before_login.png", - }); - - // Enter "admin" in the username input field - await page.fill('input[placeholder="Enter your username"]', "admin"); - page.screenshot({ - path: "test-results/view_user_info_after_username_input.png", - }); - - // Enter "gm" in the password input field - await page.fill('input[placeholder="Enter your password"]', "gm"); - page.screenshot({ - path: "test-results/view_user_info_after_password_input.png", - }); - - // Click the login button - const loginButton = page.getByRole("button", { name: "Login" }); - await expect(loginButton).toBeEnabled(); - await loginButton.click(); - page.screenshot({ - path: "test-results/view_user_info_after_login_button_click.png", - }); - - // Wait for navigation to complete and dashboard to load - await page.waitForLoadState("networkidle"); - const tabElement = page.locator("span.ant-menu-title-content", { - hasText: "Internal User", - }); - await tabElement.click(); - page.screenshot({ - path: "test-results/view_user_info_after_internal_user_tab_click.png", - }); - // Wait for loading state to disappear - await page.waitForSelector('text="🚅 Loading users..."', { - state: "hidden", - timeout: 10000, - }); - page.screenshot({ path: "test-results/view_user_info_after_loading.png" }); - // Wait for users table to load - await page.waitForSelector("table"); - page.screenshot({ - path: "test-results/view_user_info_after_table_load.png", - }); - // Get the first user ID cell - const firstUserIdCell = page.locator( - "table tbody tr:first-child td:first-child" - ); - const userId = await firstUserIdCell.textContent(); - console.log("Found user ID:", userId); - - // Click on the user ID - await firstUserIdCell.click(); - await page.waitForLoadState("networkidle"); - - // Check for tabs - await expect(page.locator('button:has-text("Overview")')).toBeVisible({ - timeout: 10000, - }); - await expect(page.locator('button:has-text("Details")')).toBeVisible({ - timeout: 10000, - }); - - // Switch to details tab - await page.locator('button:has-text("Details")').click(); - - // Check details section - await expect(page.locator("text=User ID")).toBeVisible(); - await expect(page.locator("text=Email")).toBeVisible(); - - // Go back to users list - await page.locator('button:has-text("Back to Users")').click(); - - // Verify we're back on the users page - await expect(page.locator("table")).toBeVisible(); - await expect( - page.locator('input[placeholder="Search by email..."]') - ).toBeVisible(); - }); - - // test("should handle user deletion", async ({ page }) => { - // // Wait for users table to load - // await page.waitForSelector("table"); - - // // Get the first user ID cell - // const firstUserIdCell = page.locator( - // "table tbody tr:first-child td:first-child" - // ); - // const userId = await firstUserIdCell.textContent(); - - // // Click on the user ID - // await firstUserIdCell.click(); - - // // Wait for user info view to load - // await page.waitForSelector('h1:has-text("User")'); - - // // Click delete button - // await page.locator('button:has-text("Delete User")').click(); - - // // Confirm deletion in modal - // await page.locator('button:has-text("Delete")').click(); - - // // Verify success message - // await expect(page.locator("text=User deleted successfully")).toBeVisible(); - - // // Verify we're back on the users page - // await expect(page.locator('h1:has-text("Users")')).toBeVisible(); - - // // Verify user is no longer in the table - // if (userId) { - // await expect(page.locator(`text=${userId}`)).not.toBeVisible(); - // } - // }); -}); diff --git a/tests/proxy_admin_ui_tests/package-lock.json b/tests/proxy_admin_ui_tests/package-lock.json deleted file mode 100644 index 8c79edf9ad1..00000000000 --- a/tests/proxy_admin_ui_tests/package-lock.json +++ /dev/null @@ -1,97 +0,0 @@ -{ - "name": "proxy_admin_ui_tests", - "version": "1.0.0", - "lockfileVersion": 3, - "requires": true, - "packages": { - "": { - "name": "proxy_admin_ui_tests", - "version": "1.0.0", - "license": "ISC", - "devDependencies": { - "@playwright/test": "^1.47.2", - "@types/node": "^22.5.5" - } - }, - "node_modules/@playwright/test": { - "version": "1.56.1", - "resolved": "https://registry.npmjs.org/@playwright/test/-/test-1.56.1.tgz", - "integrity": "sha512-vSMYtL/zOcFpvJCW71Q/OEGQb7KYBPAdKh35WNSkaZA75JlAO8ED8UN6GUNTm3drWomcbcqRPFqQbLae8yBTdg==", - "dev": true, - "license": "Apache-2.0", - "dependencies": { - "playwright": "1.56.1" - }, - "bin": { - "playwright": "cli.js" - }, - "engines": { - "node": ">=18" - } - }, - "node_modules/@types/node": { - "version": "22.19.1", - "resolved": "https://registry.npmjs.org/@types/node/-/node-22.19.1.tgz", - "integrity": "sha512-LCCV0HdSZZZb34qifBsyWlUmok6W7ouER+oQIGBScS8EsZsQbrtFTUrDX4hOl+CS6p7cnNC4td+qrSVGSCTUfQ==", - "dev": true, - "license": "MIT", - "dependencies": { - "undici-types": "~6.21.0" - } - }, - "node_modules/fsevents": { - "version": "2.3.2", - "resolved": "https://registry.npmjs.org/fsevents/-/fsevents-2.3.2.tgz", - "integrity": "sha512-xiqMQR4xAeHTuB9uWm+fFRcIOgKBMiOBP+eXiyT7jsgVCq1bkVygt00oASowB7EdtpOHaaPgKt812P9ab+DDKA==", - "dev": true, - "hasInstallScript": true, - "license": "MIT", - "optional": true, - "os": [ - "darwin" - ], - "engines": { - "node": "^8.16.0 || ^10.6.0 || >=11.0.0" - } - }, - "node_modules/playwright": { - "version": "1.56.1", - "resolved": "https://registry.npmjs.org/playwright/-/playwright-1.56.1.tgz", - "integrity": "sha512-aFi5B0WovBHTEvpM3DzXTUaeN6eN0qWnTkKx4NQaH4Wvcmc153PdaY2UBdSYKaGYw+UyWXSVyxDUg5DoPEttjw==", - "dev": true, - "license": "Apache-2.0", - "dependencies": { - "playwright-core": "1.56.1" - }, - "bin": { - "playwright": "cli.js" - }, - "engines": { - "node": ">=18" - }, - "optionalDependencies": { - "fsevents": "2.3.2" - } - }, - "node_modules/playwright-core": { - "version": "1.56.1", - "resolved": "https://registry.npmjs.org/playwright-core/-/playwright-core-1.56.1.tgz", - "integrity": "sha512-hutraynyn31F+Bifme+Ps9Vq59hKuUCz7H1kDOcBs+2oGguKkWTU50bBWrtz34OUWmIwpBTWDxaRPXrIXkgvmQ==", - "dev": true, - "license": "Apache-2.0", - "bin": { - "playwright-core": "cli.js" - }, - "engines": { - "node": ">=18" - } - }, - "node_modules/undici-types": { - "version": "6.21.0", - "resolved": "https://registry.npmjs.org/undici-types/-/undici-types-6.21.0.tgz", - "integrity": "sha512-iwDZqg0QAGrg9Rav5H4n0M64c3mkR59cJ6wQp+7C4nI0gsmExaedaYLNO44eT4AtBBwjbTiGPMlt2Md0T9H9JQ==", - "dev": true, - "license": "MIT" - } - } -} diff --git a/tests/proxy_admin_ui_tests/package.json b/tests/proxy_admin_ui_tests/package.json deleted file mode 100644 index 5933490fb1d..00000000000 --- a/tests/proxy_admin_ui_tests/package.json +++ /dev/null @@ -1,14 +0,0 @@ -{ - "name": "proxy_admin_ui_tests", - "version": "1.0.0", - "description": "", - "main": "index.js", - "scripts": {}, - "keywords": [], - "author": "", - "license": "ISC", - "devDependencies": { - "@playwright/test": "1.56.1", - "@types/node": "22.19.1" - } -} diff --git a/tests/proxy_admin_ui_tests/playwright.config.ts b/tests/proxy_admin_ui_tests/playwright.config.ts deleted file mode 100644 index 8b66c47394a..00000000000 --- a/tests/proxy_admin_ui_tests/playwright.config.ts +++ /dev/null @@ -1,84 +0,0 @@ -import { defineConfig, devices } from '@playwright/test'; - -/** - * Read environment variables from file. - * https://github.com/motdotla/dotenv - */ -// import dotenv from 'dotenv'; -// import path from 'path'; -// dotenv.config({ path: path.resolve(__dirname, '.env') }); - -/** - * See https://playwright.dev/docs/test-configuration. - */ -export default defineConfig({ - testDir: './e2e_ui_tests', - testIgnore: ['**/tests/pass_through_tests/**', '../pass_through_tests/**/*'], - testMatch: '**/*.spec.ts', // Only run files ending in .spec.ts - /* Run tests in files in parallel */ - fullyParallel: true, - /* Fail the build on CI if you accidentally left test.only in the source code. */ - forbidOnly: !!process.env.CI, - /* Retry on CI only */ - retries: process.env.CI ? 2 : 0, - /* Opt out of parallel tests on CI. */ - workers: process.env.CI ? 1 : undefined, - /* Reporter to use. See https://playwright.dev/docs/test-reporters */ - reporter: 'html', - /* Shared settings for all the projects below. See https://playwright.dev/docs/api/class-testoptions. */ - use: { - /* Base URL to use in actions like `await page.goto('/')`. */ - // baseURL: 'http://127.0.0.1:3000', - - /* Collect trace when retrying the failed test. See https://playwright.dev/docs/trace-viewer */ - trace: 'on-first-retry', - }, - - /* Configure projects for major browsers */ - projects: [ - { - name: 'chromium', - use: { ...devices['Desktop Chrome'] }, - }, - - { - name: 'firefox', - use: { ...devices['Desktop Firefox'] }, - }, - - { - name: 'webkit', - use: { ...devices['Desktop Safari'] }, - }, - - /* Test against mobile viewports. */ - // { - // name: 'Mobile Chrome', - // use: { ...devices['Pixel 5'] }, - // }, - // { - // name: 'Mobile Safari', - // use: { ...devices['iPhone 12'] }, - // }, - - /* Test against branded browsers. */ - // { - // name: 'Microsoft Edge', - // use: { ...devices['Desktop Edge'], channel: 'msedge' }, - // }, - // { - // name: 'Google Chrome', - // use: { ...devices['Desktop Chrome'], channel: 'chrome' }, - // }, - ], - timeout: 4*60*1000, - expect: { - timeout: 10 * 1000 - } - /* Run your local dev server before starting the tests */ - // webServer: { - // command: 'npm run start', - // url: 'http://127.0.0.1:3000', - // reuseExistingServer: !process.env.CI, - // }, -}); diff --git a/tests/proxy_admin_ui_tests/utils/login.ts b/tests/proxy_admin_ui_tests/utils/login.ts deleted file mode 100644 index 25858d9f570..00000000000 --- a/tests/proxy_admin_ui_tests/utils/login.ts +++ /dev/null @@ -1,27 +0,0 @@ -import { Page, expect } from "@playwright/test"; - -export async function loginToUI(page: Page) { - // Login first - await page.goto("http://localhost:4000/ui"); - await page.waitForLoadState("networkidle"); - console.log("Navigated to login page"); - - page.screenshot({ path: "test-results/login_utils_before.png" }); - // Wait for login form to be visible - await page.waitForSelector('input[placeholder="Enter your username"]', { - timeout: 10000, - }); - console.log("Login form is visible"); - - await page.fill('input[placeholder="Enter your username"]', "admin"); - await page.fill('input[placeholder="Enter your password"]', "gm"); - console.log("Filled login credentials"); - - const loginButton = page.getByRole("button", { name: "Login" }); - await expect(loginButton).toBeEnabled(); - await loginButton.click(); - console.log("Clicked login button"); - - // Wait for navigation to complete - await page.waitForURL("**/*"); -} diff --git a/tests/proxy_behavior/management/conftest.py b/tests/proxy_behavior/management/conftest.py index 3432f4ad6cf..fa0bef86280 100644 --- a/tests/proxy_behavior/management/conftest.py +++ b/tests/proxy_behavior/management/conftest.py @@ -11,6 +11,7 @@ import pytest_asyncio import yaml from prisma import Json +from litellm.proxy.utils import hash_token MASTER_KEY = "sk-1234" SCRATCH_PREFIX = "scratch-" @@ -106,12 +107,19 @@ async def create_scratch_key( user_id: str, team_id: Optional[str] = None, organization_id: Optional[str] = None, + key_alias: Optional[str] = None, ) -> str: """Seed a scratch-tagged key via /key/generate; returns its cleartext. Shared by the write-scenario matrices (key update/regenerate/delete). + key_alias defaults to scratch_prefix; pass a distinct scratch-prefixed + alias when a single scenario needs more than one key (/key/generate + enforces unique aliases). """ - body: Dict[str, Any] = {"key_alias": scratch_prefix, "user_id": user_id} + body: Dict[str, Any] = { + "key_alias": key_alias or scratch_prefix, + "user_id": user_id, + } if team_id is not None: body["team_id"] = team_id if organization_id is not None: @@ -132,6 +140,8 @@ async def create_scratch_team( organization_id: Optional[str] = None, admin_user_ids: Optional[list] = None, member_user_ids: Optional[list] = None, + team_member_permissions: Optional[list] = None, + models: Optional[list] = None, ) -> str: """Raw-seed a scratch-tagged team row; returns its team_id. @@ -142,6 +152,9 @@ async def create_scratch_team( members_with_roles JSON, so a raw-seeded team exercises them exactly as a /team/new-created team would. team_id must start with the scratch prefix so the `scratch` fixture reclaims the row. + + team_member_permissions / models seed the matching raw columns — needed + by the team-key-permission and team-model matrices. """ admin_user_ids = list(admin_user_ids or []) member_user_ids = list(member_user_ids or []) @@ -157,10 +170,72 @@ async def create_scratch_team( } if organization_id is not None: data["organization_id"] = organization_id + if team_member_permissions is not None: + data["team_member_permissions"] = team_member_permissions + if models is not None: + data["models"] = models await prisma.db.litellm_teamtable.create(data=data) return team_id +@dataclass(frozen=True) +class SeededActor: + user_id: str + cleartext: str + hashed: str + + +async def create_scratch_actor( + prisma, + scratch_prefix: str, + *, + user_role: str, + org_admin_of: tuple = (), + organization_id: Optional[str] = None, + suffix: str = "actor", +) -> SeededActor: + """Mint a scratch-prefixed user + verification token (+ org memberships). + + Reclaimed by the existing `scratch` teardown, which sweeps + litellm_usertable, litellm_verificationtoken, and + litellm_organizationmembership by scratch prefix — no bespoke cleanup + needed. Does NOT write litellm_teammembership against world teams: the + teardown reclaims that table only by team_id prefix, so a scratch actor + needing team membership must join a scratch team instead. The cleartext + is hashed with the real hash_token so the key authenticates end-to-end; + models=[] satisfies LiteLLM_VerificationTokenView. + """ + user_id = f"{scratch_prefix}-{suffix}" + cleartext = "sk-" + uuid.uuid4().hex + hashed = hash_token(cleartext) + await prisma.db.litellm_usertable.create( + data={ + "user_id": user_id, + "user_role": user_role, + "organization_id": organization_id, + } + ) + token_data: Dict[str, Any] = { + "token": hashed, + "key_name": f"{scratch_prefix}-{suffix}-key", + "key_alias": f"{scratch_prefix}-{suffix}-alias", + "user_id": user_id, + "models": [], + } + if organization_id is not None: + token_data["organization_id"] = organization_id + await prisma.db.litellm_verificationtoken.create(data=token_data) + for org_id in org_admin_of: + await prisma.db.litellm_organizationmembership.create( + data={ + "user_id": user_id, + "organization_id": org_id, + "user_role": "org_admin", + } + ) + return SeededActor(user_id=user_id, cleartext=cleartext, hashed=hashed) + + @pytest_asyncio.fixture async def scratch(prisma): handle = Scratch(prefix=f"{SCRATCH_PREFIX}{uuid.uuid4().hex[:12]}") diff --git a/tests/proxy_behavior/management/test_key_aliases.py b/tests/proxy_behavior/management/test_key_aliases.py new file mode 100644 index 00000000000..38ce5cdfaf3 --- /dev/null +++ b/tests/proxy_behavior/management/test_key_aliases.py @@ -0,0 +1,119 @@ +import uuid +from typing import FrozenSet + +import pytest + +from litellm.proxy.utils import hash_token + +from .actors import TEAM_ALPHA, TEAM_BETA, Actor + +pytestmark = pytest.mark.asyncio(loop_scope="session") + + +# GET /key/aliases scopes non-admins via _apply_non_admin_alias_scope: a +# non-admin sees an alias only if it owns the key (user_id match) or the key +# belongs to one of its teams. PROXY_ADMIN sees every alias. The seeded keys: +# own — owned by INTERNAL_USER, no team -> user_id scope only +# alpha — owned by OWNER, team TEAM_ALPHA -> team scope for alpha members +# beta — owned by CROSS_ORG_USER, TEAM_BETA +async def _seed_alias_keys(prisma, prefix: str, world) -> dict: + spec = { + "own": (Actor.INTERNAL_USER, None), + "alpha": (Actor.OWNER, TEAM_ALPHA), + "beta": (Actor.CROSS_ORG_USER, TEAM_BETA), + } + out = {} + for tag, (owner, team_id) in spec.items(): + alias = f"{prefix}-{tag}" + data = { + "token": hash_token("sk-" + uuid.uuid4().hex), + "key_name": f"{prefix}-{tag}-key", + "key_alias": alias, + "user_id": world.keys[owner].user_id, + "models": [], + } + if team_id is not None: + data["team_id"] = team_id + await prisma.db.litellm_verificationtoken.create(data=data) + out[tag] = alias + return out + + +async def _fetch_aliases(proxy_client, caller_cleartext: str, query: str) -> set: + resp = await proxy_client.get( + f"/key/aliases?{query}&size=100", + headers={"Authorization": f"Bearer {caller_cleartext}"}, + ) + assert resp.status_code == 200, resp.text + return set(resp.json()["aliases"]) + + +# ORG_ADMIN-role callers are stopped 401 by the management-route gate before +# the handler runs — /key/aliases carries no org context. Every other actor +# reaches the handler and is scoped by _apply_non_admin_alias_scope. +_VISIBILITY = { + Actor.PROXY_ADMIN: (200, frozenset({"own", "alpha", "beta"})), + Actor.ORG_ADMIN: (401, None), + Actor.TEAM_ADMIN: (200, frozenset({"alpha"})), + Actor.INTERNAL_USER: (200, frozenset({"own", "alpha"})), + Actor.OWNER: (200, frozenset({"alpha"})), + Actor.UNRELATED_SAME_ORG: (200, frozenset({"alpha"})), + Actor.CROSS_ORG_USER: (200, frozenset({"beta"})), + Actor.SERVICE_ACCOUNT: (200, frozenset({"alpha"})), + Actor.ORG_B_ADMIN: (401, None), +} + + +@pytest.mark.parametrize( + "actor,expected_status,expected_tags", + [(a, s, t) for a, (s, t) in _VISIBILITY.items()], + ids=[a.value for a in _VISIBILITY], +) +async def test_key_aliases_visibility( + actor: Actor, + expected_status: int, + expected_tags: FrozenSet[str], + proxy_client, + prisma, + scratch, + world, +): + aliases = await _seed_alias_keys(prisma, scratch.prefix, world) + known = {v: k for k, v in aliases.items()} + + resp = await proxy_client.get( + f"/key/aliases?search={scratch.prefix}&size=100", + headers={"Authorization": f"Bearer {world.keys[actor].cleartext}"}, + ) + assert ( + resp.status_code == expected_status + ), f"{actor.value}: {resp.status_code} {resp.text}" + if expected_status != 200: + return + + visible = {known[a] for a in resp.json()["aliases"] if a in known} + assert visible == set( + expected_tags + ), f"{actor.value}: expected {sorted(expected_tags)}, got {sorted(visible)}" + + +async def test_key_aliases_team_id_filter(proxy_client, prisma, scratch, world): + """team_id filter narrows the result to keys of that team.""" + aliases = await _seed_alias_keys(prisma, scratch.prefix, world) + returned = await _fetch_aliases( + proxy_client, + world.keys[Actor.PROXY_ADMIN].cleartext, + f"search={scratch.prefix}&team_id={TEAM_ALPHA}", + ) + assert returned & set(aliases.values()) == {aliases["alpha"]} + + +async def test_key_aliases_search_filter(proxy_client, prisma, scratch, world): + """search is a case-insensitive substring match on key_alias.""" + aliases = await _seed_alias_keys(prisma, scratch.prefix, world) + returned = await _fetch_aliases( + proxy_client, + world.keys[Actor.PROXY_ADMIN].cleartext, + f"search={aliases['beta']}", + ) + assert returned & set(aliases.values()) == {aliases["beta"]} diff --git a/tests/proxy_behavior/management/test_key_block_unblock.py b/tests/proxy_behavior/management/test_key_block_unblock.py new file mode 100644 index 00000000000..37aa0c0219a --- /dev/null +++ b/tests/proxy_behavior/management/test_key_block_unblock.py @@ -0,0 +1,159 @@ +import uuid + +import pytest + +from litellm.proxy.utils import hash_token + +from .actors import TEAM_ALPHA, TEAM_BETA, Actor +from .conftest import create_scratch_key + +pytestmark = pytest.mark.asyncio(loop_scope="session") + + +# POST /key/block + /key/unblock. PROXY_ADMIN bypasses. ORG_ADMIN-role callers +# are stopped 401 by the management-route gate BEFORE the handler runs — the +# body carries no organization_id, so the gate has no org context and falls +# back to proxy-admin-only. The handler's own _check_key_admin_access org-admin +# branch is therefore unreachable via these routes. INTERNAL_USER-role callers +# do reach _check_key_admin_access: a team admin of the key's team passes (200); +# everyone else (incl. a teamless "self" key with no team to admin) is 403. +_SCENARIOS = [ + ("self/proxy_admin", Actor.PROXY_ADMIN, "self", 200), + ("self/org_admin", Actor.ORG_ADMIN, "self", 401), + ("self/team_admin", Actor.TEAM_ADMIN, "self", 403), + ("self/internal_user", Actor.INTERNAL_USER, "self", 403), + ("self/cross_org_user", Actor.CROSS_ORG_USER, "self", 403), + ("owner/proxy_admin", Actor.PROXY_ADMIN, "owner", 200), + ("owner/org_admin", Actor.ORG_ADMIN, "owner", 401), + ("owner/team_admin", Actor.TEAM_ADMIN, "owner", 200), + ("owner/internal_user", Actor.INTERNAL_USER, "owner", 403), + ("owner/owner", Actor.OWNER, "owner", 403), + ("owner/unrelated_same_org", Actor.UNRELATED_SAME_ORG, "owner", 403), + ("owner/cross_org_user", Actor.CROSS_ORG_USER, "owner", 403), + ("owner/service_account", Actor.SERVICE_ACCOUNT, "owner", 403), + ("owner/org_b_admin", Actor.ORG_B_ADMIN, "owner", 401), + ("cross_org/proxy_admin", Actor.PROXY_ADMIN, "cross_org", 200), + ("cross_org/org_admin", Actor.ORG_ADMIN, "cross_org", 401), + ("cross_org/team_admin", Actor.TEAM_ADMIN, "cross_org", 403), + ("cross_org/cross_org_user", Actor.CROSS_ORG_USER, "cross_org", 403), + ("cross_org/org_b_admin", Actor.ORG_B_ADMIN, "cross_org", 401), +] + + +async def _seed_target(proxy_client, seeder, scratch_prefix, world, shape, caller): + if shape == "self": + return await create_scratch_key( + proxy_client, seeder, scratch_prefix, user_id=caller.user_id + ) + if shape == "owner": + return await create_scratch_key( + proxy_client, + seeder, + scratch_prefix, + user_id=world.keys[Actor.OWNER].user_id, + team_id=TEAM_ALPHA, + ) + if shape == "cross_org": + return await create_scratch_key( + proxy_client, + seeder, + scratch_prefix, + user_id=world.keys[Actor.CROSS_ORG_USER].user_id, + team_id=TEAM_BETA, + ) + pytest.fail(f"unknown shape={shape}") # pragma: no cover + + +@pytest.mark.parametrize("route", ["block", "unblock"]) +@pytest.mark.parametrize( + "actor,shape,expected_status", + [(a, sh, s) for (_id, a, sh, s) in _SCENARIOS], + ids=[s[0] for s in _SCENARIOS], +) +async def test_key_block_unblock_authz_matrix( + route: str, + actor: Actor, + shape: str, + expected_status: int, + proxy_client, + prisma, + scratch, + world, +): + caller = world.keys[actor] + seeder = world.keys[Actor.PROXY_ADMIN].cleartext + target_cleartext = await _seed_target( + proxy_client, seeder, scratch.prefix, world, shape, caller + ) + target_hashed = hash_token(target_cleartext) + + # /unblock starts from a blocked row so a 200 is observable as True->False. + if route == "unblock": + await prisma.db.litellm_verificationtoken.update( + where={"token": target_hashed}, data={"blocked": True} + ) + + resp = await proxy_client.post( + f"/key/{route}", + headers={"Authorization": f"Bearer {caller.cleartext}"}, + json={"key": target_cleartext}, + ) + assert ( + resp.status_code == expected_status + ), f"{route} {actor.value} {shape}: {resp.status_code} {resp.text}" + + row = await prisma.db.litellm_verificationtoken.find_unique( + where={"token": target_hashed} + ) + assert row is not None + # A never-blocked key reads back blocked=None; treat that as not-blocked. + if expected_status == 200: + assert bool(row.blocked) is (route == "block") + else: + # A denial leaves the blocked column at its pre-request value. + assert bool(row.blocked) is (route == "unblock"), "denied but blocked mutated" + + +async def test_key_block_unblock_round_trip(proxy_client, prisma, scratch, world): + """PROXY_ADMIN block then unblock flips the blocked column True then False.""" + admin = world.keys[Actor.PROXY_ADMIN] + target = await create_scratch_key( + proxy_client, admin.cleartext, scratch.prefix, user_id=admin.user_id + ) + hashed = hash_token(target) + headers = {"Authorization": f"Bearer {admin.cleartext}"} + + blocked = await proxy_client.post( + "/key/block", headers=headers, json={"key": target} + ) + assert blocked.status_code == 200, blocked.text + row = await prisma.db.litellm_verificationtoken.find_unique(where={"token": hashed}) + assert row is not None and row.blocked is True + + unblocked = await proxy_client.post( + "/key/unblock", headers=headers, json={"key": target} + ) + assert unblocked.status_code == 200, unblocked.text + row = await prisma.db.litellm_verificationtoken.find_unique(where={"token": hashed}) + assert row is not None and row.blocked is False + + +@pytest.mark.parametrize("route", ["block", "unblock"]) +@pytest.mark.parametrize( + "actor", [Actor.PROXY_ADMIN, Actor.TEAM_ADMIN], ids=["proxy_admin", "team_admin"] +) +async def test_key_block_unblock_missing_key_returns_404( + route: str, actor: Actor, proxy_client, world +): + """A well-formed but unseeded key is 404 — not 401/403 — for both the + PROXY_ADMIN existence check and the non-admin _check_key_admin_access path.""" + caller = world.keys[actor] + missing = "sk-" + uuid.uuid4().hex + resp = await proxy_client.post( + f"/key/{route}", + headers={"Authorization": f"Bearer {caller.cleartext}"}, + json={"key": missing}, + ) + assert ( + resp.status_code == 404 + ), f"{route} {actor.value}: {resp.status_code} {resp.text}" diff --git a/tests/proxy_behavior/management/test_key_bulk_update.py b/tests/proxy_behavior/management/test_key_bulk_update.py new file mode 100644 index 00000000000..1a57998cece --- /dev/null +++ b/tests/proxy_behavior/management/test_key_bulk_update.py @@ -0,0 +1,123 @@ +import uuid + +import pytest + +from litellm.proxy.utils import hash_token + +from .actors import Actor +from .conftest import create_scratch_key + +pytestmark = pytest.mark.asyncio(loop_scope="session") + +_MARKER_BUDGET = 42.0 + + +# POST /key/bulk_update is PROXY_ADMIN-only. The handler's own gate is +# user_role != PROXY_ADMIN -> 403, but ORG_ADMIN-role callers never reach it: +# the management-route gate 401s them first (the body carries no org context, +# and /key/bulk_update is an internal_user route, not an org-admin one). +# INTERNAL_USER-role callers clear the route gate and hit the handler's 403. +_MATRIX = [ + ("proxy_admin", Actor.PROXY_ADMIN, 200), + ("org_admin", Actor.ORG_ADMIN, 401), + ("team_admin", Actor.TEAM_ADMIN, 403), + ("internal_user", Actor.INTERNAL_USER, 403), + ("owner", Actor.OWNER, 403), + ("unrelated_same_org", Actor.UNRELATED_SAME_ORG, 403), + ("cross_org_user", Actor.CROSS_ORG_USER, 403), + ("service_account", Actor.SERVICE_ACCOUNT, 403), +] + + +@pytest.mark.parametrize( + "actor,expected_status", + [(a, s) for (_id, a, s) in _MATRIX], + ids=[s[0] for s in _MATRIX], +) +async def test_key_bulk_update_authz_matrix( + actor: Actor, + expected_status: int, + proxy_client, + prisma, + scratch, + world, +): + caller = world.keys[actor] + seeder = world.keys[Actor.PROXY_ADMIN].cleartext + target = await create_scratch_key( + proxy_client, seeder, scratch.prefix, user_id=caller.user_id + ) + hashed = hash_token(target) + + resp = await proxy_client.post( + "/key/bulk_update", + headers={"Authorization": f"Bearer {caller.cleartext}"}, + json={"keys": [{"key": target, "max_budget": _MARKER_BUDGET}]}, + ) + assert ( + resp.status_code == expected_status + ), f"{actor.value}: {resp.status_code} {resp.text}" + + row = await prisma.db.litellm_verificationtoken.find_unique(where={"token": hashed}) + assert row is not None + if expected_status == 200: + body = resp.json() + assert len(body["successful_updates"]) == 1 + assert body["failed_updates"] == [] + assert row.max_budget == _MARKER_BUDGET + else: + assert row.max_budget != _MARKER_BUDGET, "denied but key mutated" + + +async def test_key_bulk_update_empty_keys_is_400(proxy_client, world): + """An empty batch is rejected 400 before any per-key processing.""" + resp = await proxy_client.post( + "/key/bulk_update", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={"keys": []}, + ) + assert resp.status_code == 400, resp.text + + +async def test_key_bulk_update_over_max_batch_is_400(proxy_client, world): + """A batch larger than the 500-key cap is rejected 400.""" + items = [{"key": "sk-" + uuid.uuid4().hex} for _ in range(501)] + resp = await proxy_client.post( + "/key/bulk_update", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={"keys": items}, + ) + assert resp.status_code == 400, resp.text + + +async def test_key_bulk_update_per_key_failure_is_isolated( + proxy_client, prisma, scratch, world +): + """One bad key in the batch does not abort the others — it lands in + failed_updates while the valid key is still updated.""" + admin = world.keys[Actor.PROXY_ADMIN] + valid = await create_scratch_key( + proxy_client, admin.cleartext, scratch.prefix, user_id=admin.user_id + ) + missing = "sk-" + uuid.uuid4().hex + + resp = await proxy_client.post( + "/key/bulk_update", + headers={"Authorization": f"Bearer {admin.cleartext}"}, + json={ + "keys": [ + {"key": valid, "max_budget": _MARKER_BUDGET}, + {"key": missing, "max_budget": _MARKER_BUDGET}, + ] + }, + ) + assert resp.status_code == 200, resp.text + body = resp.json() + assert body["total_requested"] == 2 + assert len(body["successful_updates"]) == 1 + assert len(body["failed_updates"]) == 1 + + row = await prisma.db.litellm_verificationtoken.find_unique( + where={"token": hash_token(valid)} + ) + assert row is not None and row.max_budget == _MARKER_BUDGET diff --git a/tests/proxy_behavior/management/test_key_delete.py b/tests/proxy_behavior/management/test_key_delete.py index 05844ac0031..0b483edc056 100644 --- a/tests/proxy_behavior/management/test_key_delete.py +++ b/tests/proxy_behavior/management/test_key_delete.py @@ -1,3 +1,5 @@ +import uuid + import pytest from litellm.proxy.utils import hash_token @@ -99,3 +101,13 @@ async def test_key_delete_authz_matrix( else: assert row is not None, f"{actor.value}: denied but row vanished" assert auth_check.status_code == 200 + + +async def test_key_delete_missing_key_is_404(proxy_client, world): + """Deleting a key absent from the DB is a 404 — not 401/403.""" + resp = await proxy_client.post( + "/key/delete", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={"keys": ["sk-" + uuid.uuid4().hex]}, + ) + assert resp.status_code == 404, resp.text diff --git a/tests/proxy_behavior/management/test_key_health.py b/tests/proxy_behavior/management/test_key_health.py new file mode 100644 index 00000000000..62147e7fa13 --- /dev/null +++ b/tests/proxy_behavior/management/test_key_health.py @@ -0,0 +1,24 @@ +import pytest + +from .actors import Actor + +pytestmark = pytest.mark.asyncio(loop_scope="session") + + +# POST /key/health has no role gate — it reflects the caller's OWN key logging +# metadata. The world keys carry no "logging" metadata, so every authenticated +# actor gets 200 with key="healthy". This pins auth-required + route coverage. +@pytest.mark.parametrize("actor", list(Actor), ids=[a.value for a in Actor]) +async def test_key_health_each_actor_is_healthy(actor: Actor, proxy_client, world): + caller = world.keys[actor] + resp = await proxy_client.post( + "/key/health", + headers={"Authorization": f"Bearer {caller.cleartext}"}, + ) + assert resp.status_code == 200, f"{actor.value}: {resp.status_code} {resp.text}" + assert resp.json()["key"] == "healthy" + + +async def test_key_health_requires_auth(proxy_client): + resp = await proxy_client.post("/key/health") + assert resp.status_code == 401, resp.text diff --git a/tests/proxy_behavior/management/test_key_info_v2.py b/tests/proxy_behavior/management/test_key_info_v2.py new file mode 100644 index 00000000000..b0fb27a19fa --- /dev/null +++ b/tests/proxy_behavior/management/test_key_info_v2.py @@ -0,0 +1,82 @@ +import uuid + +import pytest + +from .actors import Actor + +pytestmark = pytest.mark.asyncio(loop_scope="session") + + +# POST /v2/key/info resolves the posted keys, then drops any key the caller +# cannot see via _can_user_query_key_info — silently, no 403. A non-admin sees +# a key it owns (user_id match) or a key whose team it belongs to. The world's +# TEAM_ALPHA members all see each other's keys; CROSS_ORG_USER and the org +# admins see only their own. The request is posted with every world key, and +# the returned info set is asserted to equal the visible subset. +_ALPHA_KEYS = frozenset( + { + Actor.TEAM_ADMIN, + Actor.INTERNAL_USER, + Actor.OWNER, + Actor.UNRELATED_SAME_ORG, + Actor.SERVICE_ACCOUNT, + } +) +_VISIBILITY = { + Actor.PROXY_ADMIN: frozenset(Actor), + Actor.ORG_ADMIN: frozenset({Actor.ORG_ADMIN}), + Actor.TEAM_ADMIN: _ALPHA_KEYS, + Actor.INTERNAL_USER: _ALPHA_KEYS, + Actor.OWNER: _ALPHA_KEYS, + Actor.UNRELATED_SAME_ORG: _ALPHA_KEYS, + Actor.SERVICE_ACCOUNT: _ALPHA_KEYS, + Actor.CROSS_ORG_USER: frozenset({Actor.CROSS_ORG_USER}), + Actor.ORG_B_ADMIN: frozenset({Actor.ORG_B_ADMIN}), +} + + +@pytest.mark.parametrize( + "actor,expected_visible", + list(_VISIBILITY.items()), + ids=[a.value for a in _VISIBILITY], +) +async def test_key_info_v2_visibility(actor, expected_visible, proxy_client, world): + caller = world.keys[actor] + user_id_to_actor = {world.keys[a].user_id: a for a in Actor} + + resp = await proxy_client.post( + "/v2/key/info", + headers={"Authorization": f"Bearer {caller.cleartext}"}, + json={"keys": [world.keys[a].cleartext for a in Actor]}, + ) + assert resp.status_code == 200, f"{actor.value}: {resp.status_code} {resp.text}" + + visible = { + user_id_to_actor[entry["user_id"]] + for entry in resp.json()["info"] + if entry.get("user_id") in user_id_to_actor + } + assert visible == set(expected_visible), ( + f"{actor.value}: expected {sorted(a.value for a in expected_visible)}, " + f"got {sorted(a.value for a in visible)}" + ) + + +async def test_key_info_v2_no_body_is_422(proxy_client, world): + """A request with no body is a 422 — the handler has no keys to resolve.""" + resp = await proxy_client.post( + "/v2/key/info", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + ) + assert resp.status_code == 422, resp.text + + +async def test_key_info_v2_unknown_key_returns_empty_info(proxy_client, world): + """Keys that resolve to no rows yield an empty info list, not an error.""" + resp = await proxy_client.post( + "/v2/key/info", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={"keys": ["sk-" + uuid.uuid4().hex]}, + ) + assert resp.status_code == 200, resp.text + assert resp.json()["info"] == [] diff --git a/tests/proxy_behavior/management/test_key_list.py b/tests/proxy_behavior/management/test_key_list.py index bda8788c9a7..0ed101d5868 100644 --- a/tests/proxy_behavior/management/test_key_list.py +++ b/tests/proxy_behavior/management/test_key_list.py @@ -2,7 +2,10 @@ from typing import FrozenSet import pytest -from .actors import Actor +from litellm.proxy.utils import hash_token + +from .actors import TEAM_ALPHA, Actor +from .conftest import create_scratch_key pytestmark = pytest.mark.asyncio(loop_scope="session") @@ -61,3 +64,108 @@ async def test_key_list_visibility( f"{actor.value}: expected {sorted(a.value for a in expected_visible)}, " f"got {sorted(a.value for a in visible_seeded)}" ) + + +async def _list_hashes(proxy_client, caller_cleartext: str, query: str) -> set: + resp = await proxy_client.get( + f"/key/list?{query}&size=100", + headers={"Authorization": f"Bearer {caller_cleartext}"}, + ) + assert resp.status_code == 200, resp.text + hashes: set = set() + for entry in resp.json().get("keys", []): + tok = entry.get("token") if isinstance(entry, dict) else entry + if tok: + hashes.add(tok) + return hashes + + +async def test_key_list_admin_key_alias_substring_match(proxy_client, scratch, world): + """A PROXY_ADMIN's key_alias filter is a case-insensitive substring match; + a narrower fragment selects the subset whose alias contains it.""" + admin = world.keys[Actor.PROXY_ADMIN] + a = await create_scratch_key( + proxy_client, + admin.cleartext, + scratch.prefix, + user_id=admin.user_id, + key_alias=f"{scratch.prefix}-sub-a", + ) + b = await create_scratch_key( + proxy_client, + admin.cleartext, + scratch.prefix, + user_id=admin.user_id, + key_alias=f"{scratch.prefix}-sub-b", + ) + seeded = {hash_token(a), hash_token(b)} + + broad = await _list_hashes( + proxy_client, admin.cleartext, f"key_alias={scratch.prefix}-sub" + ) + assert broad & seeded == seeded + + narrow = await _list_hashes( + proxy_client, admin.cleartext, f"key_alias={scratch.prefix}-sub-a" + ) + assert narrow & seeded == {hash_token(a)} + + +async def test_key_list_non_admin_key_alias_is_exact_match( + proxy_client, scratch, world +): + """A non-admin's key_alias filter is exact-match only — substring filtering + is restricted to admins. The full alias matches; a fragment does not.""" + caller = world.keys[Actor.INTERNAL_USER] + alias = f"{scratch.prefix}-exact" + key = await create_scratch_key( + proxy_client, + world.keys[Actor.PROXY_ADMIN].cleartext, + scratch.prefix, + user_id=caller.user_id, + key_alias=alias, + ) + key_hash = hash_token(key) + + exact = await _list_hashes(proxy_client, caller.cleartext, f"key_alias={alias}") + assert key_hash in exact + + fragment = await _list_hashes( + proxy_client, caller.cleartext, f"key_alias={scratch.prefix}-exac" + ) + assert key_hash not in fragment + + +async def test_key_list_team_id_filter(proxy_client, scratch, world): + """A team_id filter narrows the listing to keys of that team.""" + admin = world.keys[Actor.PROXY_ADMIN] + team_key = await create_scratch_key( + proxy_client, + admin.cleartext, + scratch.prefix, + user_id=world.keys[Actor.OWNER].user_id, + team_id=TEAM_ALPHA, + key_alias=f"{scratch.prefix}-team", + ) + no_team_key = await create_scratch_key( + proxy_client, + admin.cleartext, + scratch.prefix, + user_id=admin.user_id, + key_alias=f"{scratch.prefix}-noteam", + ) + + hashes = await _list_hashes(proxy_client, admin.cleartext, f"team_id={TEAM_ALPHA}") + assert hash_token(team_key) in hashes + assert hash_token(no_team_key) not in hashes + + +async def test_key_list_non_admin_cannot_filter_other_team(proxy_client, world): + """A non-admin filtering by a team it does not belong to is rejected 403.""" + resp = await proxy_client.get( + f"/key/list?team_id={world.team_beta_id}", + headers={ + "Authorization": f"Bearer {world.keys[Actor.INTERNAL_USER].cleartext}" + }, + ) + assert resp.status_code == 403, resp.text diff --git a/tests/proxy_behavior/management/test_key_regenerate.py b/tests/proxy_behavior/management/test_key_regenerate.py index a3289144eef..724b8b6d65b 100644 --- a/tests/proxy_behavior/management/test_key_regenerate.py +++ b/tests/proxy_behavior/management/test_key_regenerate.py @@ -1,5 +1,10 @@ +import litellm import pytest +from litellm.types.proxy.management_endpoints.ui_sso import ( + LiteLLM_UpperboundKeyGenerateParams, +) + from .actors import TEAM_ALPHA, TEAM_BETA, Actor from .conftest import create_scratch_key @@ -115,3 +120,46 @@ async def test_key_path_regenerate_smoke(proxy_client, scratch, world): assert new_cleartext.startswith("sk-") and new_cleartext != target_cleartext assert (await _info(proxy_client, target_cleartext)).status_code == 401 assert (await _info(proxy_client, new_cleartext)).status_code == 200 + + +async def test_key_regenerate_enforces_upperbound_key_params( + proxy_client, scratch, world, monkeypatch +): + """Regenerate runs _enforce_upperbound_key_params: a max_budget above + litellm.upperbound_key_generate_params is rejected 400, a value within the + bound is accepted. Pins #26340 (db8ef44323) — regenerate previously + bypassed the upperbound. upperbound_key_generate_params is module-level + litellm.* state, so monkeypatch save/restores it.""" + admin = world.keys[Actor.PROXY_ADMIN] + over_key = await create_scratch_key( + proxy_client, + admin.cleartext, + scratch.prefix, + user_id=admin.user_id, + key_alias=f"{scratch.prefix}-over", + ) + within_key = await create_scratch_key( + proxy_client, + admin.cleartext, + scratch.prefix, + user_id=admin.user_id, + key_alias=f"{scratch.prefix}-within", + ) + monkeypatch.setattr( + litellm, + "upperbound_key_generate_params", + LiteLLM_UpperboundKeyGenerateParams(max_budget=100.0), + ) + headers = {"Authorization": f"Bearer {admin.cleartext}"} + + over = await proxy_client.post( + "/key/regenerate", headers=headers, json={"key": over_key, "max_budget": 500.0} + ) + assert over.status_code == 400, over.text + + within = await proxy_client.post( + "/key/regenerate", + headers=headers, + json={"key": within_key, "max_budget": 50.0}, + ) + assert within.status_code == 200, within.text diff --git a/tests/proxy_behavior/management/test_key_reset_spend.py b/tests/proxy_behavior/management/test_key_reset_spend.py new file mode 100644 index 00000000000..fb1c266f655 --- /dev/null +++ b/tests/proxy_behavior/management/test_key_reset_spend.py @@ -0,0 +1,136 @@ +import uuid + +import pytest + +from litellm.proxy.utils import hash_token + +from .actors import TEAM_ALPHA, TEAM_BETA, Actor +from .conftest import create_scratch_key + +pytestmark = pytest.mark.asyncio(loop_scope="session") + +_SEED_SPEND = 5.0 +_RESET_TO = 2.0 + + +# POST /key/{key}/reset_spend. The target key is pre-seeded with spend=5.0 so +# reset_to=2.0 always clears _validate_reset_spend_value (which runs before +# authz). _check_proxy_or_team_admin_for_key then allows only PROXY_ADMIN or a +# team admin of the key's team — there is no org-admin branch, and a teamless +# "self" key has no team to admin. ORG_ADMIN-role callers are stopped 401 at +# the management-route gate before the handler runs. +_SCENARIOS = [ + ("self/proxy_admin", Actor.PROXY_ADMIN, "self", 200), + ("self/org_admin", Actor.ORG_ADMIN, "self", 401), + ("self/team_admin", Actor.TEAM_ADMIN, "self", 403), + ("self/internal_user", Actor.INTERNAL_USER, "self", 403), + ("self/cross_org_user", Actor.CROSS_ORG_USER, "self", 403), + ("team_alpha/proxy_admin", Actor.PROXY_ADMIN, "team_alpha", 200), + ("team_alpha/org_admin", Actor.ORG_ADMIN, "team_alpha", 401), + ("team_alpha/team_admin", Actor.TEAM_ADMIN, "team_alpha", 200), + ("team_alpha/internal_user", Actor.INTERNAL_USER, "team_alpha", 403), + ("team_alpha/owner", Actor.OWNER, "team_alpha", 403), + ("team_alpha/unrelated_same_org", Actor.UNRELATED_SAME_ORG, "team_alpha", 403), + ("team_alpha/cross_org_user", Actor.CROSS_ORG_USER, "team_alpha", 403), + ("team_alpha/service_account", Actor.SERVICE_ACCOUNT, "team_alpha", 403), + ("team_alpha/org_b_admin", Actor.ORG_B_ADMIN, "team_alpha", 401), + ("team_beta/proxy_admin", Actor.PROXY_ADMIN, "team_beta", 200), + ("team_beta/org_admin", Actor.ORG_ADMIN, "team_beta", 401), + ("team_beta/team_admin", Actor.TEAM_ADMIN, "team_beta", 403), + ("team_beta/cross_org_user", Actor.CROSS_ORG_USER, "team_beta", 403), + ("team_beta/org_b_admin", Actor.ORG_B_ADMIN, "team_beta", 401), +] + + +async def _seed_target(proxy_client, seeder, prefix, world, shape, caller) -> str: + if shape == "self": + return await create_scratch_key( + proxy_client, seeder, prefix, user_id=caller.user_id + ) + if shape == "team_alpha": + return await create_scratch_key( + proxy_client, + seeder, + prefix, + user_id=world.keys[Actor.OWNER].user_id, + team_id=TEAM_ALPHA, + ) + if shape == "team_beta": + return await create_scratch_key( + proxy_client, + seeder, + prefix, + user_id=world.keys[Actor.CROSS_ORG_USER].user_id, + team_id=TEAM_BETA, + ) + pytest.fail(f"unknown shape={shape}") # pragma: no cover + + +@pytest.mark.parametrize( + "actor,shape,expected_status", + [(a, sh, s) for (_id, a, sh, s) in _SCENARIOS], + ids=[s[0] for s in _SCENARIOS], +) +async def test_key_reset_spend_authz_matrix( + actor: Actor, + shape: str, + expected_status: int, + proxy_client, + prisma, + scratch, + world, +): + caller = world.keys[actor] + seeder = world.keys[Actor.PROXY_ADMIN].cleartext + target = await _seed_target( + proxy_client, seeder, scratch.prefix, world, shape, caller + ) + hashed = hash_token(target) + await prisma.db.litellm_verificationtoken.update( + where={"token": hashed}, data={"spend": _SEED_SPEND} + ) + + resp = await proxy_client.post( + f"/key/{target}/reset_spend", + headers={"Authorization": f"Bearer {caller.cleartext}"}, + json={"reset_to": _RESET_TO}, + ) + assert ( + resp.status_code == expected_status + ), f"{actor.value} {shape}: {resp.status_code} {resp.text}" + + row = await prisma.db.litellm_verificationtoken.find_unique(where={"token": hashed}) + assert row is not None + if expected_status == 200: + assert row.spend == _RESET_TO + else: + assert row.spend == _SEED_SPEND, "denied but spend reset" + + +@pytest.mark.parametrize( + "actor", [Actor.PROXY_ADMIN, Actor.TEAM_ADMIN], ids=["proxy_admin", "team_admin"] +) +async def test_key_reset_spend_missing_key_is_404(actor: Actor, proxy_client, world): + """A well-formed but unseeded key is 404 before any spend validation.""" + resp = await proxy_client.post( + f"/key/sk-{uuid.uuid4().hex}/reset_spend", + headers={"Authorization": f"Bearer {world.keys[actor].cleartext}"}, + json={"reset_to": 0.0}, + ) + assert resp.status_code == 404, resp.text + + +async def test_key_reset_spend_above_current_spend_is_400( + proxy_client, prisma, scratch, world +): + """reset_to above the key's current spend is rejected 400.""" + admin = world.keys[Actor.PROXY_ADMIN] + target = await create_scratch_key( + proxy_client, admin.cleartext, scratch.prefix, user_id=admin.user_id + ) + resp = await proxy_client.post( + f"/key/{target}/reset_spend", + headers={"Authorization": f"Bearer {admin.cleartext}"}, + json={"reset_to": 1.0}, + ) + assert resp.status_code == 400, resp.text diff --git a/tests/proxy_behavior/management/test_key_service_account_generate.py b/tests/proxy_behavior/management/test_key_service_account_generate.py new file mode 100644 index 00000000000..3b5bbe39754 --- /dev/null +++ b/tests/proxy_behavior/management/test_key_service_account_generate.py @@ -0,0 +1,98 @@ +import pytest + +from .actors import Actor + +pytestmark = pytest.mark.asyncio(loop_scope="session") + + +# POST /key/service-account/generate. PROXY_ADMIN always passes. ORG_ADMIN-role +# callers are stopped 401 by the management-route gate (the body carries a +# team_id but no organization_id, so the org-admin route branch never matches). +# INTERNAL_USER-role callers reach the handler: a team admin of the target team +# passes (200); a "user"-role member is 401 (no service-account-generate +# permission); a non-member is 400 ("not assigned to team"). A request with no +# team_id is 400 ("team_id is required") for every actor that reaches the handler. +_SCENARIOS = [ + ("own/proxy_admin", Actor.PROXY_ADMIN, "own", 200), + ("own/org_admin", Actor.ORG_ADMIN, "own", 401), + ("own/team_admin", Actor.TEAM_ADMIN, "own", 200), + ("own/internal_user", Actor.INTERNAL_USER, "own", 401), + ("own/owner", Actor.OWNER, "own", 401), + ("own/unrelated_same_org", Actor.UNRELATED_SAME_ORG, "own", 401), + ("own/cross_org_user", Actor.CROSS_ORG_USER, "own", 400), + ("own/service_account", Actor.SERVICE_ACCOUNT, "own", 401), + ("own/org_b_admin", Actor.ORG_B_ADMIN, "own", 401), + ("cross_org/proxy_admin", Actor.PROXY_ADMIN, "cross_org", 200), + ("cross_org/org_admin", Actor.ORG_ADMIN, "cross_org", 401), + ("cross_org/team_admin", Actor.TEAM_ADMIN, "cross_org", 400), + ("cross_org/internal_user", Actor.INTERNAL_USER, "cross_org", 400), + ("cross_org/cross_org_user", Actor.CROSS_ORG_USER, "cross_org", 401), + ("cross_org/org_b_admin", Actor.ORG_B_ADMIN, "cross_org", 401), + ("none/proxy_admin", Actor.PROXY_ADMIN, "none", 400), + ("none/org_admin", Actor.ORG_ADMIN, "none", 401), + ("none/team_admin", Actor.TEAM_ADMIN, "none", 400), + ("none/internal_user", Actor.INTERNAL_USER, "none", 400), + ("none/cross_org_user", Actor.CROSS_ORG_USER, "none", 400), +] + + +@pytest.mark.parametrize( + "actor,team_target,expected_status", + [(a, t, s) for (_id, a, t, s) in _SCENARIOS], + ids=[s[0] for s in _SCENARIOS], +) +async def test_key_service_account_generate_authz_matrix( + actor: Actor, + team_target: str, + expected_status: int, + proxy_client, + prisma, + scratch, + world, +): + caller = world.keys[actor] + team_id = { + "own": world.team_alpha_id, + "cross_org": world.team_beta_id, + "none": None, + }[team_target] + + body = {"key_alias": scratch.prefix} + if team_id is not None: + body["team_id"] = team_id + + resp = await proxy_client.post( + "/key/service-account/generate", + headers={"Authorization": f"Bearer {caller.cleartext}"}, + json=body, + ) + assert ( + resp.status_code == expected_status + ), f"{actor.value} {team_target}: {resp.status_code} {resp.text}" + + rows = await prisma.db.litellm_verificationtoken.find_many( + where={"key_alias": scratch.prefix} + ) + if expected_status == 200: + assert len(rows) == 1 + # A service-account key belongs to the team, not a user. + assert rows[0].user_id is None + assert rows[0].team_id == team_id + else: + assert rows == [], f"{actor.value}: denied but key row leaked" + + +async def test_key_service_account_generate_unknown_team_is_400( + proxy_client, prisma, scratch, world +): + """A team_id absent from the database is rejected 400.""" + resp = await proxy_client.post( + "/key/service-account/generate", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={"key_alias": scratch.prefix, "team_id": scratch.tag("no-such-team")}, + ) + assert resp.status_code == 400, resp.text + rows = await prisma.db.litellm_verificationtoken.find_many( + where={"key_alias": scratch.prefix} + ) + assert rows == [] diff --git a/tests/proxy_behavior/management/test_key_update.py b/tests/proxy_behavior/management/test_key_update.py index 36ddefa5750..7b7f6f5558b 100644 --- a/tests/proxy_behavior/management/test_key_update.py +++ b/tests/proxy_behavior/management/test_key_update.py @@ -1,3 +1,5 @@ +import uuid + import pytest from litellm.proxy.utils import hash_token @@ -98,3 +100,85 @@ async def test_key_update_authz_matrix( assert row.models == [MARKER_MODEL] else: assert row.models != [MARKER_MODEL], "denied but row mutated" + + +async def _seed_shape(proxy_client, seeder, prefix, world, shape, caller) -> str: + if shape == "self": + return await create_scratch_key( + proxy_client, seeder, prefix, user_id=caller.user_id + ) + if shape == "owner": + return await create_scratch_key( + proxy_client, + seeder, + prefix, + user_id=world.keys[Actor.OWNER].user_id, + team_id=TEAM_ALPHA, + ) + if shape == "cross_org": + return await create_scratch_key( + proxy_client, + seeder, + prefix, + user_id=world.keys[Actor.CROSS_ORG_USER].user_id, + team_id=TEAM_BETA, + ) + pytest.fail(f"unknown shape={shape}") # pragma: no cover + + +async def test_key_update_missing_key_is_404(proxy_client, world): + """An update targeting a key absent from the DB is a 404 — not 401/403.""" + resp = await proxy_client.post( + "/key/update", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={"key": "sk-" + uuid.uuid4().hex, "models": [MARKER_MODEL]}, + ) + assert resp.status_code == 404, resp.text + + +# A denied /key/update must not partially apply: the budget/limit columns are +# left untouched. Each scenario is a denial cell from the matrix above. +_DENIED_BUDGET = [ + ("team_admin/self", Actor.TEAM_ADMIN, "self", 403), + ("internal_user/owner", Actor.INTERNAL_USER, "owner", 403), + ("cross_org_user/cross_org", Actor.CROSS_ORG_USER, "cross_org", 401), +] + + +@pytest.mark.parametrize( + "actor,target_shape,expected_status", + [(a, t, s) for (_id, a, t, s) in _DENIED_BUDGET], + ids=[s[0] for s in _DENIED_BUDGET], +) +async def test_key_update_denied_does_not_touch_budget_counters( + actor: Actor, + target_shape: str, + expected_status: int, + proxy_client, + prisma, + scratch, + world, +): + caller = world.keys[actor] + seeder = world.keys[Actor.PROXY_ADMIN].cleartext + target = await _seed_shape( + proxy_client, seeder, scratch.prefix, world, target_shape, caller + ) + target_hashed = hash_token(target) + + resp = await proxy_client.post( + "/key/update", + headers={"Authorization": f"Bearer {caller.cleartext}"}, + json={"key": target, "max_budget": 999.0, "tpm_limit": 888, "rpm_limit": 777}, + ) + assert ( + resp.status_code == expected_status + ), f"{actor.value} {target_shape}: {resp.status_code} {resp.text}" + + row = await prisma.db.litellm_verificationtoken.find_unique( + where={"token": target_hashed} + ) + assert row is not None + assert row.max_budget is None, "denied but max_budget applied" + assert row.tpm_limit is None, "denied but tpm_limit applied" + assert row.rpm_limit is None, "denied but rpm_limit applied" diff --git a/tests/proxy_behavior/management/test_route_coverage.py b/tests/proxy_behavior/management/test_route_coverage.py new file mode 100644 index 00000000000..1139e251a59 --- /dev/null +++ b/tests/proxy_behavior/management/test_route_coverage.py @@ -0,0 +1,91 @@ +"""PR3.M1 — codified route coverage. + +Every route declared in the two management-endpoint source files must be +exercised by at least one behavior-suite scenario. This is a permanent +regression guard: a future route added without a behavior test fails CI here, +the same way test_no_management_imports.py codifies the G3 import grep. +""" + +import ast +import pathlib +import re + +REPO_ROOT = pathlib.Path(__file__).resolve().parents[3] +SOURCE_FILES = [ + REPO_ROOT / "litellm/proxy/management_endpoints/key_management_endpoints.py", + REPO_ROOT / "litellm/proxy/management_endpoints/team_endpoints.py", +] +TEST_DIR = pathlib.Path(__file__).resolve().parent +SELF = pathlib.Path(__file__).resolve() + +# Captures the route literal from `@router.(""` — `\s*` spans +# newlines so multi-line decorators are matched too. +_ROUTE_DECORATOR = re.compile( + r"@router\.(?:get|post|put|delete|patch)\(\s*[\"']([^\"']+)[\"']" +) + + +def _source_routes() -> set: + routes: set = set() + for path in SOURCE_FILES: + routes.update(_ROUTE_DECORATOR.findall(path.read_text())) + return routes + + +def _route_to_regex(route: str) -> re.Pattern: + # A plain path param ({team_id}) matches a single path segment; a Starlette + # ':path' param ({key:path}) matches across '/'. Keeping plain params + # slash-bounded stops a loose regex from falsely reporting a future + # multi-segment route as already covered. + pattern = ["^"] + pos = 0 + for match in re.finditer(r"\{([^}]+)\}", route): + pattern.append(re.escape(route[pos : match.start()])) + pattern.append("[^?]+" if match.group(1).endswith(":path") else "[^/?]+") + pos = match.end() + pattern.append(re.escape(route[pos:]) + "$") + return re.compile("".join(pattern)) + + +def _test_urls() -> set: + """Every request-URL string literal across the behavior test suite. + + f-strings are reconstructed with each interpolation collapsed to a single + placeholder char, so f"/key/{target}/regenerate" becomes /key/X/regenerate. + Query strings are dropped — coverage is a path-level property. + """ + urls: set = set() + for path in sorted(TEST_DIR.glob("test_*.py")): + if path.resolve() == SELF: + continue + tree = ast.parse(path.read_text()) + for node in ast.walk(tree): + literal = None + if isinstance(node, ast.Constant) and isinstance(node.value, str): + literal = node.value + elif isinstance(node, ast.JoinedStr): + chunks = [] + for value in node.values: + if isinstance(value, ast.Constant) and isinstance(value.value, str): + chunks.append(value.value) + else: + chunks.append("X") # interpolated path / query segment + literal = "".join(chunks) + if literal and literal.startswith("/"): + urls.add(literal.split("?", 1)[0]) + return urls + + +def test_every_management_route_has_a_behavior_scenario(): + routes = _source_routes() + assert routes, "no @router routes parsed — the decorator regex is stale" + + urls = _test_urls() + uncovered = sorted( + route + for route in routes + if not any(_route_to_regex(route).match(url) for url in urls) + ) + assert ( + not uncovered + ), "management routes with no behavior-suite scenario:\n " + "\n ".join(uncovered) diff --git a/tests/proxy_behavior/management/test_scratch_teardown.py b/tests/proxy_behavior/management/test_scratch_teardown.py index 689c60fc78a..bcb53935558 100644 --- a/tests/proxy_behavior/management/test_scratch_teardown.py +++ b/tests/proxy_behavior/management/test_scratch_teardown.py @@ -1,13 +1,16 @@ import pytest -from .conftest import MASTER_KEY, SCRATCH_PREFIX +from litellm.proxy._types import LitellmUserRoles + +from .actors import ORG_A, ORG_B +from .conftest import MASTER_KEY, SCRATCH_PREFIX, create_scratch_actor pytestmark = pytest.mark.asyncio(loop_scope="session") -# The two tests run in file order: _a writes a scratch-tagged key and asserts -# it lands; _b runs after _a's fixture teardown and asserts no scratch row -# survived. A leak in either direction fails _b on the next collection. +# The minting tests run in file order, then _b runs after their fixture +# teardown and asserts no scratch row survived in any reclaimed table. A leak +# in either direction fails _b on the next collection. async def test_a_scratch_key_lands_in_db(proxy_client, prisma, scratch): @@ -24,8 +27,35 @@ async def test_a_scratch_key_lands_in_db(proxy_client, prisma, scratch): assert len(rows) == 1 +async def test_a2_scratch_actor_lands_in_db(proxy_client, prisma, scratch): + actor = await create_scratch_actor( + prisma, + scratch.prefix, + user_role=LitellmUserRoles.ORG_ADMIN.value, + org_admin_of=(ORG_A, ORG_B), + ) + user_row = await prisma.db.litellm_usertable.find_unique( + where={"user_id": actor.user_id} + ) + assert user_row is not None + info = await proxy_client.get( + "/key/info", headers={"Authorization": f"Bearer {actor.cleartext}"} + ) + assert info.status_code == 200, info.text + memberships = await prisma.db.litellm_organizationmembership.find_many( + where={"user_id": actor.user_id} + ) + assert {m.organization_id for m in memberships} == {ORG_A, ORG_B} + + async def test_b_scratch_namespace_is_clean(prisma): - rows = await prisma.db.litellm_verificationtoken.find_many( + tokens = await prisma.db.litellm_verificationtoken.find_many( where={"key_alias": {"startswith": SCRATCH_PREFIX}} ) - assert rows == [] + users = await prisma.db.litellm_usertable.find_many( + where={"user_id": {"startswith": SCRATCH_PREFIX}} + ) + memberships = await prisma.db.litellm_organizationmembership.find_many( + where={"user_id": {"startswith": SCRATCH_PREFIX}} + ) + assert tokens == [] and users == [] and memberships == [] diff --git a/tests/proxy_behavior/management/test_team_available.py b/tests/proxy_behavior/management/test_team_available.py new file mode 100644 index 00000000000..874c8dd4df7 --- /dev/null +++ b/tests/proxy_behavior/management/test_team_available.py @@ -0,0 +1,21 @@ +import pytest + +from .actors import Actor + +pytestmark = pytest.mark.asyncio(loop_scope="session") + + +# GET /team/available lists teams from +# litellm.default_internal_user_params["available_teams"]. The behavior world +# configures no available_teams, so the handler returns [] for every actor +# before it even reads the caller — this is the route-coverage + default-path +# pin. /team/available is an info route, so every authenticated actor reaches +# the handler. +@pytest.mark.parametrize("actor", list(Actor), ids=[a.value for a in Actor]) +async def test_team_available_default_is_empty(actor: Actor, proxy_client, world): + resp = await proxy_client.get( + "/team/available", + headers={"Authorization": f"Bearer {world.keys[actor].cleartext}"}, + ) + assert resp.status_code == 200, f"{actor.value}: {resp.status_code} {resp.text}" + assert resp.json() == [] diff --git a/tests/proxy_behavior/management/test_team_block_unblock.py b/tests/proxy_behavior/management/test_team_block_unblock.py new file mode 100644 index 00000000000..9412e51b909 --- /dev/null +++ b/tests/proxy_behavior/management/test_team_block_unblock.py @@ -0,0 +1,114 @@ +import pytest + +from .actors import Actor +from .conftest import create_scratch_team + +pytestmark = pytest.mark.asyncio(loop_scope="session") + + +# POST /team/block + /team/unblock. The handler gate is _verify_team_access +# (proxy admin / team admin / org admin), but the management-route gate fronts +# it: the request carries the team's organization_id so an org admin of that +# org clears the gate's org-scoped branch. A team admin is an INTERNAL_USER +# and these are not internal_user routes, so a team admin can never reach the +# handler — only PROXY_ADMIN and an org admin of the team's own org pass. +_MATRIX = [ + ("alpha/proxy_admin", Actor.PROXY_ADMIN, "alpha", 200), + ("alpha/org_admin", Actor.ORG_ADMIN, "alpha", 200), + ("alpha/team_admin", Actor.TEAM_ADMIN, "alpha", 401), + ("alpha/internal_user", Actor.INTERNAL_USER, "alpha", 401), + ("alpha/owner", Actor.OWNER, "alpha", 401), + ("alpha/unrelated_same_org", Actor.UNRELATED_SAME_ORG, "alpha", 401), + ("alpha/cross_org_user", Actor.CROSS_ORG_USER, "alpha", 401), + ("alpha/service_account", Actor.SERVICE_ACCOUNT, "alpha", 401), + ("alpha/org_b_admin", Actor.ORG_B_ADMIN, "alpha", 401), + ("beta/proxy_admin", Actor.PROXY_ADMIN, "beta", 200), + ("beta/org_admin", Actor.ORG_ADMIN, "beta", 401), + ("beta/org_b_admin", Actor.ORG_B_ADMIN, "beta", 200), +] + + +async def _seed_target(prisma, world, shape: str, team_id: str) -> str: + """Raw-seed the scratch target team; returns its organization_id.""" + org_id = world.org_a_id if shape == "alpha" else world.org_b_id + await create_scratch_team(prisma, team_id, organization_id=org_id) + return org_id + + +@pytest.mark.parametrize("route", ["block", "unblock"]) +@pytest.mark.parametrize( + "actor,shape,expected_status", + [(a, sh, s) for (_id, a, sh, s) in _MATRIX], + ids=[s[0] for s in _MATRIX], +) +async def test_team_block_unblock_authz_matrix( + route: str, + actor: Actor, + shape: str, + expected_status: int, + proxy_client, + prisma, + scratch, + world, +): + org_id = await _seed_target(prisma, world, shape, scratch.prefix) + caller = world.keys[actor] + + # /unblock starts from a blocked row so a 200 is observable as True->False. + if route == "unblock": + await prisma.db.litellm_teamtable.update( + where={"team_id": scratch.prefix}, data={"blocked": True} + ) + + resp = await proxy_client.post( + f"/team/{route}", + headers={"Authorization": f"Bearer {caller.cleartext}"}, + json={"team_id": scratch.prefix, "organization_id": org_id}, + ) + assert ( + resp.status_code == expected_status + ), f"{route} {actor.value} {shape}: {resp.status_code} {resp.text}" + + row = await prisma.db.litellm_teamtable.find_unique( + where={"team_id": scratch.prefix} + ) + assert row is not None + if expected_status == 200: + assert bool(row.blocked) is (route == "block") + else: + assert bool(row.blocked) is (route == "unblock"), "denied but blocked mutated" + + +async def test_team_block_unblock_round_trip(proxy_client, prisma, scratch, world): + """PROXY_ADMIN block then unblock flips the blocked column True then False.""" + await create_scratch_team(prisma, scratch.prefix, organization_id=world.org_a_id) + headers = {"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"} + + blocked = await proxy_client.post( + "/team/block", headers=headers, json={"team_id": scratch.prefix} + ) + assert blocked.status_code == 200, blocked.text + row = await prisma.db.litellm_teamtable.find_unique( + where={"team_id": scratch.prefix} + ) + assert row is not None and row.blocked is True + + unblocked = await proxy_client.post( + "/team/unblock", headers=headers, json={"team_id": scratch.prefix} + ) + assert unblocked.status_code == 200, unblocked.text + row = await prisma.db.litellm_teamtable.find_unique( + where={"team_id": scratch.prefix} + ) + assert row is not None and row.blocked is False + + +@pytest.mark.parametrize("route", ["block", "unblock"]) +async def test_team_block_unblock_missing_team_is_404(route: str, proxy_client, world): + """A team_id absent from the DB is 404 — the existence check precedes authz.""" + resp = await proxy_client.post( + f"/team/{route}", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={"team_id": "behavior-pin-no-such-team"}, + ) + assert resp.status_code == 404, resp.text diff --git a/tests/proxy_behavior/management/test_team_bulk_member_add.py b/tests/proxy_behavior/management/test_team_bulk_member_add.py new file mode 100644 index 00000000000..fc83cd414e5 --- /dev/null +++ b/tests/proxy_behavior/management/test_team_bulk_member_add.py @@ -0,0 +1,105 @@ +import pytest + +from .actors import Actor +from .conftest import create_scratch_team + +pytestmark = pytest.mark.asyncio(loop_scope="session") + + +def _member_ids(row) -> list: + return [m["user_id"] for m in (row.members_with_roles or [])] + + +async def test_team_bulk_member_add_proxy_admin_adds_explicit_members( + proxy_client, prisma, scratch, world +): + """PROXY_ADMIN bulk-adds an explicit member list to a scratch team.""" + await create_scratch_team(prisma, scratch.prefix, organization_id=world.org_a_id) + new_member = scratch.tag("m1") + resp = await proxy_client.post( + "/team/bulk_member_add", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={ + "team_id": scratch.prefix, + "members": [{"user_id": new_member, "role": "user"}], + }, + ) + assert resp.status_code == 200, resp.text + row = await prisma.db.litellm_teamtable.find_unique( + where={"team_id": scratch.prefix} + ) + assert row is not None and new_member in _member_ids(row) + + +async def test_team_bulk_member_add_empty_members_is_400( + proxy_client, prisma, scratch, world +): + """An empty member list (with all_users unset) is rejected 400.""" + await create_scratch_team(prisma, scratch.prefix, organization_id=world.org_a_id) + resp = await proxy_client.post( + "/team/bulk_member_add", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={"team_id": scratch.prefix, "members": []}, + ) + assert resp.status_code == 400, resp.text + + +async def test_team_bulk_member_add_over_max_batch_is_400( + proxy_client, prisma, scratch, world +): + """A member list larger than the 500-member cap is rejected 400.""" + await create_scratch_team(prisma, scratch.prefix, organization_id=world.org_a_id) + members = [ + {"user_id": f"{scratch.prefix}-u{i}", "role": "user"} for i in range(501) + ] + resp = await proxy_client.post( + "/team/bulk_member_add", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={"team_id": scratch.prefix, "members": members}, + ) + assert resp.status_code == 400, resp.text + + +@pytest.mark.parametrize( + "actor", + [Actor.TEAM_ADMIN, Actor.INTERNAL_USER], + ids=["team_admin", "internal_user"], +) +async def test_team_bulk_member_add_non_admin_is_401( + actor: Actor, proxy_client, prisma, scratch, world +): + """/team/bulk_member_add is neither an internal_user nor a self-managed + route — a non-proxy-admin with no org context is 401 at the route gate.""" + await create_scratch_team(prisma, scratch.prefix, organization_id=world.org_a_id) + resp = await proxy_client.post( + "/team/bulk_member_add", + headers={"Authorization": f"Bearer {world.keys[actor].cleartext}"}, + json={ + "team_id": scratch.prefix, + "members": [{"user_id": scratch.tag("m"), "role": "user"}], + }, + ) + assert resp.status_code == 401, f"{actor.value}: {resp.status_code} {resp.text}" + + +async def test_team_bulk_member_add_all_users_proxy_admin( + proxy_client, prisma, scratch, world +): + """all_users=True pulls every user in the DB into the team. The route is + reachable only by PROXY_ADMIN (the route gate 401s every other actor — even + an org admin with organization_id in the body), so the handler's own + all_users PROXY_ADMIN gate is never the deciding check at the boundary.""" + await create_scratch_team(prisma, scratch.prefix, organization_id=world.org_a_id) + resp = await proxy_client.post( + "/team/bulk_member_add", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={"team_id": scratch.prefix, "all_users": True}, + ) + assert resp.status_code == 200, resp.text + row = await prisma.db.litellm_teamtable.find_unique( + where={"team_id": scratch.prefix} + ) + assert row is not None + member_ids = _member_ids(row) + # every world actor is a user in the DB, so all are now team members + assert world.keys[Actor.INTERNAL_USER].user_id in member_ids diff --git a/tests/proxy_behavior/management/test_team_daily_activity.py b/tests/proxy_behavior/management/test_team_daily_activity.py new file mode 100644 index 00000000000..7a1e70b91fc --- /dev/null +++ b/tests/proxy_behavior/management/test_team_daily_activity.py @@ -0,0 +1,63 @@ +import pytest + +from .actors import Actor + +pytestmark = pytest.mark.asyncio(loop_scope="session") + + +# GET /team/daily/activity. A proxy admin (admin view) sees activity for any +# team. A non-admin is scoped to user_info.teams: a bare query defaults to its +# own teams (200), and an explicit team_ids filter naming a team it does not +# belong to is 404 (the VERIA-43 fix). Org admins have no team memberships, so +# they behave like a non-member for any specific team. +_MEMBERS = { + "alpha": { + Actor.TEAM_ADMIN, + Actor.INTERNAL_USER, + Actor.OWNER, + Actor.UNRELATED_SAME_ORG, + Actor.SERVICE_ACCOUNT, + }, + "beta": {Actor.CROSS_ORG_USER}, +} + + +def _expected(actor: Actor, team: str) -> int: + if team == "none" or actor == Actor.PROXY_ADMIN: + return 200 + return 200 if actor in _MEMBERS.get(team, set()) else 404 + + +_CASES = [ + (f"{team}/{actor.value}", actor, team, _expected(actor, team)) + for team in ("none", "alpha", "beta") + for actor in Actor +] + + +# start_date / end_date are required by the handler — pin only the team-scope +# authz, not the date validation. +_DATES = "start_date=2024-01-01&end_date=2024-12-31" + + +@pytest.mark.parametrize( + "actor,team,expected_status", + [(a, t, s) for (_id, a, t, s) in _CASES], + ids=[c[0] for c in _CASES], +) +async def test_team_daily_activity_matrix( + actor: Actor, team: str, expected_status: int, proxy_client, world +): + query = _DATES + if team == "alpha": + query += f"&team_ids={world.team_alpha_id}" + elif team == "beta": + query += f"&team_ids={world.team_beta_id}" + + resp = await proxy_client.get( + f"/team/daily/activity?{query}", + headers={"Authorization": f"Bearer {world.keys[actor].cleartext}"}, + ) + assert ( + resp.status_code == expected_status + ), f"{actor.value} -> {team}: {resp.status_code} {resp.text}" diff --git a/tests/proxy_behavior/management/test_team_delete.py b/tests/proxy_behavior/management/test_team_delete.py new file mode 100644 index 00000000000..bbf0a6563f3 --- /dev/null +++ b/tests/proxy_behavior/management/test_team_delete.py @@ -0,0 +1,78 @@ +import pytest + +from .actors import Actor +from .conftest import create_scratch_team + +pytestmark = pytest.mark.asyncio(loop_scope="session") + + +# POST /team/delete runs per-team _verify_team_access. The request carries the +# team's organization_id so an org admin of that org clears the management- +# route gate; a team admin is an INTERNAL_USER on a non-internal_user route, +# so a team admin never reaches the handler. Only PROXY_ADMIN and an org admin +# of the team's own org can delete it. +_MATRIX = [ + ("alpha/proxy_admin", Actor.PROXY_ADMIN, "alpha", 200), + ("alpha/org_admin", Actor.ORG_ADMIN, "alpha", 200), + ("alpha/team_admin", Actor.TEAM_ADMIN, "alpha", 401), + ("alpha/internal_user", Actor.INTERNAL_USER, "alpha", 401), + ("alpha/cross_org_user", Actor.CROSS_ORG_USER, "alpha", 401), + ("alpha/org_b_admin", Actor.ORG_B_ADMIN, "alpha", 401), + ("beta/proxy_admin", Actor.PROXY_ADMIN, "beta", 200), + ("beta/org_admin", Actor.ORG_ADMIN, "beta", 401), + ("beta/org_b_admin", Actor.ORG_B_ADMIN, "beta", 200), +] + + +@pytest.mark.parametrize( + "actor,shape,expected_status", + [(a, sh, s) for (_id, a, sh, s) in _MATRIX], + ids=[s[0] for s in _MATRIX], +) +async def test_team_delete_authz_matrix( + actor: Actor, + shape: str, + expected_status: int, + proxy_client, + prisma, + scratch, + world, +): + org_id = world.org_a_id if shape == "alpha" else world.org_b_id + await create_scratch_team(prisma, scratch.prefix, organization_id=org_id) + caller = world.keys[actor] + + resp = await proxy_client.post( + "/team/delete", + headers={"Authorization": f"Bearer {caller.cleartext}"}, + json={"team_ids": [scratch.prefix], "organization_id": org_id}, + ) + assert ( + resp.status_code == expected_status + ), f"{actor.value} {shape}: {resp.status_code} {resp.text}" + + row = await prisma.db.litellm_teamtable.find_unique( + where={"team_id": scratch.prefix} + ) + if expected_status == 200: + assert row is None, "deleted but team row survives" + else: + assert row is not None, "denied but team row vanished" + + +async def test_team_delete_batch_with_missing_id_deletes_nothing( + proxy_client, prisma, scratch, world +): + """A batch is validated whole before any deletion: one missing team_id + fails the request 404 and the accessible team in the batch survives.""" + await create_scratch_team(prisma, scratch.prefix, organization_id=world.org_a_id) + resp = await proxy_client.post( + "/team/delete", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={"team_ids": [scratch.prefix, "behavior-pin-no-such-team"]}, + ) + assert resp.status_code == 404, resp.text + row = await prisma.db.litellm_teamtable.find_unique( + where={"team_id": scratch.prefix} + ) + assert row is not None, "batch aborted but the accessible team was deleted" diff --git a/tests/proxy_behavior/management/test_team_filter_ui.py b/tests/proxy_behavior/management/test_team_filter_ui.py new file mode 100644 index 00000000000..69cbabf72a1 --- /dev/null +++ b/tests/proxy_behavior/management/test_team_filter_ui.py @@ -0,0 +1,39 @@ +import pytest + +from .actors import Actor + +pytestmark = pytest.mark.asyncio(loop_scope="session") + + +# GET /team/filter/ui (ui_view_teams) — include_in_schema=False. The handler +# body has no role/org check and never reads user_api_key_dict, but the +# endpoint is still effectively PROXY-ADMIN-only as its docstring claims: the +# management-route gate fronts it (not an internal_user / info / org-admin +# route) and 401s every non-proxy-admin before the handler runs. PROXY_ADMIN +# reaches the unscoped find_many and sees teams across every org. +@pytest.mark.parametrize("actor", list(Actor), ids=[a.value for a in Actor]) +async def test_team_filter_ui_is_proxy_admin_only(actor: Actor, proxy_client, world): + resp = await proxy_client.get( + "/team/filter/ui", + headers={"Authorization": f"Bearer {world.keys[actor].cleartext}"}, + ) + expected = 200 if actor == Actor.PROXY_ADMIN else 401 + assert ( + resp.status_code == expected + ), f"{actor.value}: {resp.status_code} {resp.text}" + + +async def test_team_filter_ui_proxy_admin_sees_cross_org_teams(proxy_client, world): + """The handler runs an unscoped query — PROXY_ADMIN sees teams from every + org, including the three seeded world teams.""" + resp = await proxy_client.get( + "/team/filter/ui", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + ) + assert resp.status_code == 200, resp.text + team_ids = {t.get("team_id") for t in resp.json() if isinstance(t, dict)} + assert { + world.team_alpha_id, + world.team_beta_id, + world.team_gamma_id, + } <= team_ids diff --git a/tests/proxy_behavior/management/test_team_key_bulk_update.py b/tests/proxy_behavior/management/test_team_key_bulk_update.py new file mode 100644 index 00000000000..5acf0c8185c --- /dev/null +++ b/tests/proxy_behavior/management/test_team_key_bulk_update.py @@ -0,0 +1,217 @@ +import uuid + +import pytest + +from litellm.proxy._types import KeyManagementRoutes +from litellm.proxy.utils import hash_token + +from .actors import Actor +from .conftest import create_scratch_key, create_scratch_team + +pytestmark = pytest.mark.asyncio(loop_scope="session") + +_MARKER_BUDGET = 42.0 +_KEY_UPDATE = KeyManagementRoutes.KEY_UPDATE.value + + +# POST /team/key/bulk_update — PROXY_ADMIN bypasses; otherwise +# can_team_member_execute_key_management_endpoint runs with route=KEY_UPDATE. +# A team admin always passes; a "user"-role member passes only when the team's +# team_member_permissions grants /key/update; a non-member is 401. ORG_ADMIN is +# stopped 401 at the management-route gate before the handler (the body has a +# team_id but no organization_id, so the org-admin route branch never matches). +_MATRIX = [ + ("admin/proxy_admin", Actor.PROXY_ADMIN, "admin", 200), + ("admin/internal_user", Actor.INTERNAL_USER, "admin", 200), + ("member_allowed/internal_user", Actor.INTERNAL_USER, "member_allowed", 200), + ("member_denied/internal_user", Actor.INTERNAL_USER, "member_denied", 401), + ("nonmember/internal_user", Actor.INTERNAL_USER, "nonmember", 401), + ("nonmember/org_admin", Actor.ORG_ADMIN, "nonmember", 401), + ("nonmember/proxy_admin", Actor.PROXY_ADMIN, "nonmember", 200), +] + + +async def _seed_team_key(prisma, proxy_client, prefix: str, world, shape: str) -> str: + """Raw-seed the scratch team for `shape`, return a team key's cleartext.""" + internal = world.keys[Actor.INTERNAL_USER].user_id + owner = world.keys[Actor.OWNER].user_id + if shape == "admin": + await create_scratch_team( + prisma, prefix, organization_id=world.org_a_id, admin_user_ids=[internal] + ) + key_owner = internal + elif shape == "member_allowed": + await create_scratch_team( + prisma, + prefix, + organization_id=world.org_a_id, + admin_user_ids=[owner], + member_user_ids=[internal], + team_member_permissions=[_KEY_UPDATE], + ) + key_owner = owner + elif shape == "member_denied": + await create_scratch_team( + prisma, + prefix, + organization_id=world.org_a_id, + admin_user_ids=[owner], + member_user_ids=[internal], + team_member_permissions=[], + ) + key_owner = owner + elif shape == "nonmember": + await create_scratch_team( + prisma, prefix, organization_id=world.org_a_id, admin_user_ids=[owner] + ) + key_owner = owner + else: + pytest.fail(f"unknown shape={shape}") # pragma: no cover + return await create_scratch_key( + proxy_client, + world.keys[Actor.PROXY_ADMIN].cleartext, + prefix, + user_id=key_owner, + team_id=prefix, + ) + + +@pytest.mark.parametrize( + "actor,shape,expected_status", + [(a, sh, s) for (_id, a, sh, s) in _MATRIX], + ids=[s[0] for s in _MATRIX], +) +async def test_team_key_bulk_update_authz_matrix( + actor: Actor, + shape: str, + expected_status: int, + proxy_client, + prisma, + scratch, + world, +): + key = await _seed_team_key(prisma, proxy_client, scratch.prefix, world, shape) + hashed = hash_token(key) + caller = world.keys[actor] + + resp = await proxy_client.post( + "/team/key/bulk_update", + headers={"Authorization": f"Bearer {caller.cleartext}"}, + json={ + "team_id": scratch.prefix, + "key_ids": [key], + "update_fields": {"max_budget": _MARKER_BUDGET}, + }, + ) + assert ( + resp.status_code == expected_status + ), f"{actor.value} {shape}: {resp.status_code} {resp.text}" + + row = await prisma.db.litellm_verificationtoken.find_unique(where={"token": hashed}) + assert row is not None + if expected_status == 200: + assert len(resp.json()["successful_updates"]) == 1 + assert row.max_budget == _MARKER_BUDGET + else: + assert row.max_budget != _MARKER_BUDGET, "denied but key mutated" + + +async def test_team_key_bulk_update_requires_team_id( + proxy_client, prisma, scratch, world +): + """An empty team_id is rejected 400.""" + resp = await proxy_client.post( + "/team/key/bulk_update", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={ + "team_id": "", + "key_ids": ["sk-" + uuid.uuid4().hex], + "update_fields": {"max_budget": _MARKER_BUDGET}, + }, + ) + assert resp.status_code == 400, resp.text + + +async def test_team_key_bulk_update_all_keys_in_team( + proxy_client, prisma, scratch, world +): + """all_keys_in_team=True broadcasts the update to every key in the team.""" + admin = world.keys[Actor.PROXY_ADMIN].cleartext + await create_scratch_team(prisma, scratch.prefix, organization_id=world.org_a_id) + keys = [ + await create_scratch_key( + proxy_client, + admin, + scratch.prefix, + user_id=world.keys[Actor.OWNER].user_id, + team_id=scratch.prefix, + key_alias=f"{scratch.prefix}-k{i}", + ) + for i in range(2) + ] + + resp = await proxy_client.post( + "/team/key/bulk_update", + headers={"Authorization": f"Bearer {admin}"}, + json={ + "team_id": scratch.prefix, + "all_keys_in_team": True, + "update_fields": {"max_budget": _MARKER_BUDGET}, + }, + ) + assert resp.status_code == 200, resp.text + assert len(resp.json()["successful_updates"]) == 2 + for key in keys: + row = await prisma.db.litellm_verificationtoken.find_unique( + where={"token": hash_token(key)} + ) + assert row is not None and row.max_budget == _MARKER_BUDGET + + +async def test_team_key_bulk_update_no_keys_found_is_404( + proxy_client, prisma, scratch, world +): + """all_keys_in_team=True on a team with no keys is a top-level 404.""" + await create_scratch_team(prisma, scratch.prefix, organization_id=world.org_a_id) + resp = await proxy_client.post( + "/team/key/bulk_update", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={ + "team_id": scratch.prefix, + "all_keys_in_team": True, + "update_fields": {"max_budget": _MARKER_BUDGET}, + }, + ) + assert resp.status_code == 404, resp.text + + +async def test_team_key_bulk_update_missing_key_is_isolated( + proxy_client, prisma, scratch, world +): + """A key_id absent from the team lands in failed_updates; the batch still + returns 200 and the real key is updated.""" + admin = world.keys[Actor.PROXY_ADMIN].cleartext + real = await _seed_team_key( + prisma, proxy_client, scratch.prefix, world, "nonmember" + ) + missing = "sk-" + uuid.uuid4().hex + + resp = await proxy_client.post( + "/team/key/bulk_update", + headers={"Authorization": f"Bearer {admin}"}, + json={ + "team_id": scratch.prefix, + "key_ids": [real, missing], + "update_fields": {"max_budget": _MARKER_BUDGET}, + }, + ) + assert resp.status_code == 200, resp.text + body = resp.json() + assert body["total_requested"] == 2 + assert len(body["successful_updates"]) == 1 + assert len(body["failed_updates"]) == 1 + + row = await prisma.db.litellm_verificationtoken.find_unique( + where={"token": hash_token(real)} + ) + assert row is not None and row.max_budget == _MARKER_BUDGET diff --git a/tests/proxy_behavior/management/test_team_list_v2.py b/tests/proxy_behavior/management/test_team_list_v2.py new file mode 100644 index 00000000000..81178ad73c0 --- /dev/null +++ b/tests/proxy_behavior/management/test_team_list_v2.py @@ -0,0 +1,141 @@ +from typing import FrozenSet, Optional + +import pytest + +from .actors import Actor + +pytestmark = pytest.mark.asyncio(loop_scope="session") + + +def _seeded(team_ids: set, world) -> set: + known = { + world.team_alpha_id: "alpha", + world.team_beta_id: "beta", + world.team_gamma_id: "gamma", + } + return {known[t] for t in team_ids if t in known} + + +async def _v2_team_ids(proxy_client, caller_cleartext: str, extra: str = "") -> set: + """Walk every /v2/team/list page and collect the returned team_ids.""" + ids: set = set() + page = 1 + while True: + resp = await proxy_client.get( + f"/v2/team/list?page={page}&page_size=100{extra}", + headers={"Authorization": f"Bearer {caller_cleartext}"}, + ) + assert resp.status_code == 200, resp.text + body = resp.json() + teams = body.get("teams", []) or [] + for t in teams: + tid = t.get("team_id") if isinstance(t, dict) else None + if tid: + ids.add(tid) + if page * 100 >= (body.get("total") or 0) or not teams: + return ids + page += 1 + + +# GET /v2/team/list is an info route reachable by every actor, but +# _enforce_list_team_v2_access still gates a BARE query: a proxy admin sees +# all teams, an org admin sees its orgs' teams, and a regular user — who has +# passed no user_id filter — is rejected 401 ("only admins can query all +# teams"). A regular user must scope the query to its own user_id. +_BARE = [ + ("proxy_admin", Actor.PROXY_ADMIN, 200, frozenset({"alpha", "beta", "gamma"})), + ("org_admin", Actor.ORG_ADMIN, 200, frozenset({"alpha", "gamma"})), + ("org_b_admin", Actor.ORG_B_ADMIN, 200, frozenset({"beta"})), + ("team_admin", Actor.TEAM_ADMIN, 401, None), + ("internal_user", Actor.INTERNAL_USER, 401, None), + ("owner", Actor.OWNER, 401, None), + ("unrelated_same_org", Actor.UNRELATED_SAME_ORG, 401, None), + ("cross_org_user", Actor.CROSS_ORG_USER, 401, None), + ("service_account", Actor.SERVICE_ACCOUNT, 401, None), +] + + +@pytest.mark.parametrize( + "actor,expected_status,expected_visible", + [(a, s, v) for (_id, a, s, v) in _BARE], + ids=[s[0] for s in _BARE], +) +async def test_team_list_v2_bare( + actor: Actor, + expected_status: int, + expected_visible: Optional[FrozenSet[str]], + proxy_client, + world, +): + caller = world.keys[actor] + if expected_status != 200: + resp = await proxy_client.get( + "/v2/team/list", + headers={"Authorization": f"Bearer {caller.cleartext}"}, + ) + assert resp.status_code == expected_status, resp.text + return + + visible = _seeded(await _v2_team_ids(proxy_client, caller.cleartext), world) + assert visible == set( + expected_visible + ), f"{actor.value}: expected {sorted(expected_visible)}, got {sorted(visible)}" + + +# A regular user scoping the query to its own user_id is allowed, and sees +# exactly the teams it belongs to. +_OWN = { + Actor.TEAM_ADMIN: frozenset({"alpha"}), + Actor.INTERNAL_USER: frozenset({"alpha"}), + Actor.OWNER: frozenset({"alpha"}), + Actor.UNRELATED_SAME_ORG: frozenset({"alpha"}), + Actor.CROSS_ORG_USER: frozenset({"beta"}), + Actor.SERVICE_ACCOUNT: frozenset({"alpha"}), +} + + +@pytest.mark.parametrize( + "actor,expected_visible", list(_OWN.items()), ids=[a.value for a in _OWN] +) +async def test_team_list_v2_own_user_id_query( + actor: Actor, expected_visible: FrozenSet[str], proxy_client, world +): + caller = world.keys[actor] + visible = _seeded( + await _v2_team_ids( + proxy_client, caller.cleartext, f"&user_id={caller.user_id}" + ), + world, + ) + assert visible == set( + expected_visible + ), f"{actor.value}: expected {sorted(expected_visible)}, got {sorted(visible)}" + + +async def test_team_list_v2_user_id_filter_other_user_is_401(proxy_client, world): + """A regular user filtering by another user's user_id is rejected 401.""" + resp = await proxy_client.get( + f"/v2/team/list?user_id={world.keys[Actor.OWNER].user_id}", + headers={ + "Authorization": f"Bearer {world.keys[Actor.INTERNAL_USER].cleartext}" + }, + ) + assert resp.status_code == 401, resp.text + + +async def test_team_list_v2_org_filter_foreign_org_is_403(proxy_client, world): + """An org admin filtering by an organization it does not administer is 403.""" + resp = await proxy_client.get( + f"/v2/team/list?organization_id={world.org_b_id}", + headers={"Authorization": f"Bearer {world.keys[Actor.ORG_ADMIN].cleartext}"}, + ) + assert resp.status_code == 403, resp.text + + +async def test_team_list_v2_invalid_status_is_400(proxy_client, world): + """status accepts only 'deleted' — any other value is 400.""" + resp = await proxy_client.get( + "/v2/team/list?status=bogus", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + ) + assert resp.status_code == 400, resp.text diff --git a/tests/proxy_behavior/management/test_team_member_me.py b/tests/proxy_behavior/management/test_team_member_me.py new file mode 100644 index 00000000000..bfbbe0504ae --- /dev/null +++ b/tests/proxy_behavior/management/test_team_member_me.py @@ -0,0 +1,83 @@ +import uuid + +import pytest + +from litellm.proxy.utils import hash_token + +from .actors import Actor + +pytestmark = pytest.mark.asyncio(loop_scope="session") + + +# GET /team/{team_id}/members/me resolves the CALLER's own membership row. +# A caller that is not a member of the team is 404 — even PROXY_ADMIN, which +# is not in any seeded team. The route is self-managed, so every actor reaches +# the handler. TEAM_GAMMA has no members, so every actor is 404 there. +_MEMBERS = { + "alpha": { + Actor.TEAM_ADMIN, + Actor.INTERNAL_USER, + Actor.OWNER, + Actor.UNRELATED_SAME_ORG, + Actor.SERVICE_ACCOUNT, + }, + "beta": {Actor.CROSS_ORG_USER}, + "gamma": set(), +} + +_CASES = [ + (f"{team}/{actor.value}", actor, team, 200 if actor in members else 404) + for team, members in _MEMBERS.items() + for actor in Actor +] + + +@pytest.mark.parametrize( + "actor,team,expected_status", + [(a, t, s) for (_id, a, t, s) in _CASES], + ids=[c[0] for c in _CASES], +) +async def test_team_member_me_matrix( + actor: Actor, team: str, expected_status: int, proxy_client, world +): + team_id = { + "alpha": world.team_alpha_id, + "beta": world.team_beta_id, + "gamma": world.team_gamma_id, + }[team] + caller = world.keys[actor] + + resp = await proxy_client.get( + f"/team/{team_id}/members/me", + headers={"Authorization": f"Bearer {caller.cleartext}"}, + ) + assert ( + resp.status_code == expected_status + ), f"{actor.value} -> {team}: {resp.status_code} {resp.text}" + + if expected_status == 200: + body = resp.json() + assert body["user_id"] == caller.user_id + assert body["team_id"] == team_id + + +async def test_team_member_me_team_key_without_user_id_is_400( + proxy_client, prisma, scratch, world +): + """A key with no associated user_id (a team / service-account key) cannot + resolve 'me' — the caller has no identity to look up — so it is 400.""" + cleartext = "sk-" + uuid.uuid4().hex + await prisma.db.litellm_verificationtoken.create( + data={ + "token": hash_token(cleartext), + "key_name": f"{scratch.prefix}-teamkey", + "key_alias": f"{scratch.prefix}-teamkey", + "team_id": world.team_alpha_id, + "models": [], + } + ) + resp = await proxy_client.get( + f"/team/{world.team_alpha_id}/members/me", + headers={"Authorization": f"Bearer {cleartext}"}, + ) + assert resp.status_code == 400, resp.text diff --git a/tests/proxy_behavior/management/test_team_model.py b/tests/proxy_behavior/management/test_team_model.py new file mode 100644 index 00000000000..3564e8df83a --- /dev/null +++ b/tests/proxy_behavior/management/test_team_model.py @@ -0,0 +1,78 @@ +import pytest + +from .actors import Actor +from .conftest import create_scratch_team + +pytestmark = pytest.mark.asyncio(loop_scope="session") + +_MARKER_MODEL = "behavior-pin-team-model-marker" +_ROUTE_URL = {"add": "/team/model/add", "delete": "/team/model/delete"} + + +# POST /team/model/add + /team/model/delete. The handler gate is PROXY_ADMIN +# or team admin or org admin, but the management-route gate fronts it — these +# are neither internal_user nor org-admin nor info routes, so every +# non-proxy-admin is 401 before the handler runs. Only PROXY_ADMIN reaches the +# handler, making the team-admin / org-admin handler branches unreachable here. +_MATRIX = [ + ("proxy_admin", Actor.PROXY_ADMIN, 200), + ("org_admin", Actor.ORG_ADMIN, 401), + ("team_admin", Actor.TEAM_ADMIN, 401), + ("internal_user", Actor.INTERNAL_USER, 401), + ("owner", Actor.OWNER, 401), + ("unrelated_same_org", Actor.UNRELATED_SAME_ORG, 401), + ("cross_org_user", Actor.CROSS_ORG_USER, 401), + ("service_account", Actor.SERVICE_ACCOUNT, 401), + ("org_b_admin", Actor.ORG_B_ADMIN, 401), +] + + +@pytest.mark.parametrize("route", ["add", "delete"]) +@pytest.mark.parametrize( + "actor,expected_status", + [(a, s) for (_id, a, s) in _MATRIX], + ids=[s[0] for s in _MATRIX], +) +async def test_team_model_authz_matrix( + route: str, + actor: Actor, + expected_status: int, + proxy_client, + prisma, + scratch, + world, +): + initial = [] if route == "add" else [_MARKER_MODEL] + await create_scratch_team( + prisma, scratch.prefix, organization_id=world.org_a_id, models=initial + ) + caller = world.keys[actor] + + resp = await proxy_client.post( + _ROUTE_URL[route], + headers={"Authorization": f"Bearer {caller.cleartext}"}, + json={"team_id": scratch.prefix, "models": [_MARKER_MODEL]}, + ) + assert ( + resp.status_code == expected_status + ), f"{route} {actor.value}: {resp.status_code} {resp.text}" + + row = await prisma.db.litellm_teamtable.find_unique( + where={"team_id": scratch.prefix} + ) + assert row is not None + if expected_status == 200: + assert (_MARKER_MODEL in row.models) is (route == "add") + else: + assert list(row.models) == initial, "denied but models mutated" + + +@pytest.mark.parametrize("route", ["add", "delete"]) +async def test_team_model_missing_team_is_404(route: str, proxy_client, world): + """A team_id absent from the DB is 404 — the existence check precedes authz.""" + resp = await proxy_client.post( + _ROUTE_URL[route], + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={"team_id": "behavior-pin-no-such-team", "models": [_MARKER_MODEL]}, + ) + assert resp.status_code == 404, resp.text diff --git a/tests/proxy_behavior/management/test_team_permissions.py b/tests/proxy_behavior/management/test_team_permissions.py new file mode 100644 index 00000000000..5d16702fe6c --- /dev/null +++ b/tests/proxy_behavior/management/test_team_permissions.py @@ -0,0 +1,170 @@ +import litellm +import pytest + +from litellm.proxy._types import KeyManagementRoutes + +from .actors import Actor +from .conftest import create_scratch_team + +pytestmark = pytest.mark.asyncio(loop_scope="session") + +_PERM = KeyManagementRoutes.KEY_INFO.value + + +# GET /team/permissions_list and POST /team/permissions_update are self-managed +# routes, so every actor reaches the handler. Both grant access to PROXY_ADMIN, +# the team admin, or an org admin of the team's org. The scratch team is in +# ORG_A with TEAM_ADMIN as its team admin. +_MATRIX = [ + ("proxy_admin", Actor.PROXY_ADMIN, 200), + ("org_admin", Actor.ORG_ADMIN, 200), + ("team_admin", Actor.TEAM_ADMIN, 200), + ("internal_user", Actor.INTERNAL_USER, 403), + ("owner", Actor.OWNER, 403), + ("unrelated_same_org", Actor.UNRELATED_SAME_ORG, 403), + ("cross_org_user", Actor.CROSS_ORG_USER, 403), + ("service_account", Actor.SERVICE_ACCOUNT, 403), + ("org_b_admin", Actor.ORG_B_ADMIN, 403), +] + + +async def _seed_team(prisma, scratch_prefix, world) -> None: + await create_scratch_team( + prisma, + scratch_prefix, + organization_id=world.org_a_id, + admin_user_ids=[world.keys[Actor.TEAM_ADMIN].user_id], + member_user_ids=[ + world.keys[Actor.INTERNAL_USER].user_id, + world.keys[Actor.OWNER].user_id, + world.keys[Actor.UNRELATED_SAME_ORG].user_id, + world.keys[Actor.SERVICE_ACCOUNT].user_id, + ], + ) + + +@pytest.mark.parametrize( + "actor,expected_status", + [(a, s) for (_id, a, s) in _MATRIX], + ids=[s[0] for s in _MATRIX], +) +async def test_team_permissions_list_authz_matrix( + actor: Actor, expected_status: int, proxy_client, prisma, scratch, world +): + await _seed_team(prisma, scratch.prefix, world) + resp = await proxy_client.get( + f"/team/permissions_list?team_id={scratch.prefix}", + headers={"Authorization": f"Bearer {world.keys[actor].cleartext}"}, + ) + assert ( + resp.status_code == expected_status + ), f"{actor.value}: {resp.status_code} {resp.text}" + if expected_status == 200: + assert resp.json()["team_id"] == scratch.prefix + + +@pytest.mark.parametrize( + "actor,expected_status", + [(a, s) for (_id, a, s) in _MATRIX], + ids=[s[0] for s in _MATRIX], +) +async def test_team_permissions_update_authz_matrix( + actor: Actor, expected_status: int, proxy_client, prisma, scratch, world +): + await _seed_team(prisma, scratch.prefix, world) + resp = await proxy_client.post( + "/team/permissions_update", + headers={"Authorization": f"Bearer {world.keys[actor].cleartext}"}, + json={"team_id": scratch.prefix, "team_member_permissions": [_PERM]}, + ) + assert ( + resp.status_code == expected_status + ), f"{actor.value}: {resp.status_code} {resp.text}" + + row = await prisma.db.litellm_teamtable.find_unique( + where={"team_id": scratch.prefix} + ) + assert row is not None + if expected_status == 200: + assert _PERM in (row.team_member_permissions or []) + else: + assert _PERM not in (row.team_member_permissions or []), "denied but mutated" + + +async def test_team_permissions_available_team_self_join_divergence( + proxy_client, prisma, scratch, world, monkeypatch +): + """permissions_list honours the available-team self-join — a non-admin can + READ an available team's permissions — but permissions_update deliberately + does not: the same caller is 403 on update. default_internal_user_params is + module-level litellm.* state, so monkeypatch save/restores it.""" + await create_scratch_team(prisma, scratch.prefix, organization_id=world.org_a_id) + monkeypatch.setattr( + litellm, "default_internal_user_params", {"available_teams": [scratch.prefix]} + ) + caller = world.keys[Actor.CROSS_ORG_USER] # non-admin, unrelated to the team + + listed = await proxy_client.get( + f"/team/permissions_list?team_id={scratch.prefix}", + headers={"Authorization": f"Bearer {caller.cleartext}"}, + ) + assert listed.status_code == 200, listed.text + + updated = await proxy_client.post( + "/team/permissions_update", + headers={"Authorization": f"Bearer {caller.cleartext}"}, + json={"team_id": scratch.prefix, "team_member_permissions": [_PERM]}, + ) + assert updated.status_code == 403, updated.text + + +# POST /team/permissions_bulk_update is PROXY_ADMIN-only. ORG_ADMIN-role +# callers are stopped 401 by the management-route gate; INTERNAL_USER-role +# callers, on a route that is neither internal_user nor self-managed, are 401 +# there too — only PROXY_ADMIN reaches the handler's own admin gate. +_BULK_MATRIX = [ + ("proxy_admin", Actor.PROXY_ADMIN, 200), + ("org_admin", Actor.ORG_ADMIN, 401), + ("team_admin", Actor.TEAM_ADMIN, 401), + ("internal_user", Actor.INTERNAL_USER, 401), + ("cross_org_user", Actor.CROSS_ORG_USER, 401), + ("org_b_admin", Actor.ORG_B_ADMIN, 401), +] + + +@pytest.mark.parametrize( + "actor,expected_status", + [(a, s) for (_id, a, s) in _BULK_MATRIX], + ids=[s[0] for s in _BULK_MATRIX], +) +async def test_team_permissions_bulk_update_authz_matrix( + actor: Actor, expected_status: int, proxy_client, prisma, scratch, world +): + await create_scratch_team(prisma, scratch.prefix, organization_id=world.org_a_id) + resp = await proxy_client.post( + "/team/permissions_bulk_update", + headers={"Authorization": f"Bearer {world.keys[actor].cleartext}"}, + json={"team_ids": [scratch.prefix], "permissions": [_PERM]}, + ) + assert ( + resp.status_code == expected_status + ), f"{actor.value}: {resp.status_code} {resp.text}" + + row = await prisma.db.litellm_teamtable.find_unique( + where={"team_id": scratch.prefix} + ) + assert row is not None + if expected_status == 200: + assert _PERM in (row.team_member_permissions or []) + else: + assert _PERM not in (row.team_member_permissions or []), "denied but mutated" + + +async def test_team_permissions_bulk_update_no_selector_is_400(proxy_client, world): + """Neither team_ids nor apply_to_all_teams is a 400.""" + resp = await proxy_client.post( + "/team/permissions_bulk_update", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={"permissions": [_PERM]}, + ) + assert resp.status_code == 400, resp.text diff --git a/tests/proxy_behavior/management/test_team_update.py b/tests/proxy_behavior/management/test_team_update.py index 3baf2b2148f..9b21911cef2 100644 --- a/tests/proxy_behavior/management/test_team_update.py +++ b/tests/proxy_behavior/management/test_team_update.py @@ -1,7 +1,9 @@ import pytest +from litellm.proxy._types import LitellmUserRoles + from .actors import Actor -from .conftest import create_scratch_team +from .conftest import create_scratch_actor, create_scratch_team pytestmark = pytest.mark.asyncio(loop_scope="session") @@ -130,8 +132,8 @@ async def test_team_update_requires_proxy_admin_without_org_context( # in ORG_A; each scenario relocates it to ORG_B. PROXY_ADMIN bypasses; # ORG_B_ADMIN clears the route gate (dest-org admin) but fails # _verify_team_access on the source team (403); the rest fail the route gate -# (401). The relocation-allowed branch needs a caller who is org admin of both -# orgs — no seeded actor is, so it is left to a later slice. +# (401). The relocation-*allowed* branch (caller is org admin of both orgs) is +# covered by test_team_update_org_relocation_allowed_for_dual_org_admin below. _RELOCATION = [ ("proxy_admin", Actor.PROXY_ADMIN, 200), ("org_b_admin", Actor.ORG_B_ADMIN, 403), @@ -174,3 +176,34 @@ async def test_team_update_org_relocation_gate( assert row.organization_id == world.org_b_id else: assert row.organization_id == world.org_a_id, "denied but team relocated" + + +async def test_team_update_org_relocation_allowed_for_dual_org_admin( + proxy_client, prisma, scratch, world +): + """Relocation-allowed branch: a caller who is org admin of BOTH the source + and destination org may relocate a team between them. Completes the + _RELOCATION matrix, whose allowed branch PR2 left open — no seeded actor is + a dual-org admin, so one is minted with create_scratch_actor.""" + actor = await create_scratch_actor( + prisma, + scratch.prefix, + user_role=LitellmUserRoles.ORG_ADMIN.value, + org_admin_of=(world.org_a_id, world.org_b_id), + ) + team_id = await create_scratch_team( + prisma, scratch.tag("team"), organization_id=world.org_a_id + ) + + resp = await proxy_client.post( + "/team/update", + headers={"Authorization": f"Bearer {actor.cleartext}"}, + json={"team_id": team_id, "organization_id": world.org_b_id}, + ) + assert resp.status_code == 200, resp.text + + row = await prisma.db.litellm_teamtable.find_unique(where={"team_id": team_id}) + assert row is not None + assert ( + row.organization_id == world.org_b_id + ), "dual-org admin relocation not applied" diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py index 11465e6f718..44530fecebd 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py @@ -1203,6 +1203,51 @@ def test_streaming_chunk_with_both_text_and_tool_calls_issue_18238(): assert content_block_start["id"] == "toolu_bdrk_013xRVejhv3ybmLEGCoZib2b" +def test_streaming_chunk_with_text_and_empty_tool_calls_returns_text_delta(): + """ + Some OpenAI-compatible providers emit `tool_calls: []` on regular text chunks. + + Empty tool_calls should be treated as no tool call so the Anthropic adapter + does not shadow text with an empty input_json_delta. + """ + choices = [ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta( + provider_specific_fields=None, + content="Hello from vLLM", + role="assistant", + function_call=None, + tool_calls=[], + audio=None, + ), + logprobs=None, + ) + ] + + adapter = LiteLLMAnthropicMessagesAdapter() + + ( + type_of_content, + content_block_delta, + ) = adapter._translate_streaming_openai_chunk_to_anthropic(choices=choices) + + assert type_of_content == "text_delta" + assert content_block_delta["type"] == "text_delta" + assert content_block_delta["text"] == "Hello from vLLM" + + ( + block_type, + content_block_start, + ) = adapter._translate_streaming_openai_chunk_to_anthropic_content_block( + choices=choices + ) + + assert block_type == "text" + assert content_block_start == {"type": "text", "text": ""} + + # ============================================================================ # Cache Control Transformation Tests # ============================================================================ diff --git a/tests/test_litellm/llms/azure/chat/test_azure_base_model_routing.py b/tests/test_litellm/llms/azure/chat/test_azure_base_model_routing.py new file mode 100644 index 00000000000..1e8e23c38ca --- /dev/null +++ b/tests/test_litellm/llms/azure/chat/test_azure_base_model_routing.py @@ -0,0 +1,274 @@ +"""Tests for decoupling Azure deployment IDs from underlying model names. + +When users name their Azure deployment something non-standard (e.g. "my-deployment-id"), +setting ``base_model`` should drive model-type detection (o-series, gpt-5, +etc.) so the correct config, supported params, and param mapping are used. +""" + +import pytest + +import litellm +from litellm.llms.azure.chat.gpt_5_transformation import AzureOpenAIGPT5Config +from litellm.llms.azure.chat.o_series_transformation import AzureOpenAIO1Config +from litellm.utils import ProviderConfigManager, get_optional_params + + +# --------------------------------------------------------------------------- +# _get_azure_config — routes to the correct config based on base_model +# --------------------------------------------------------------------------- +class TestGetAzureConfigWithBaseModel: + """ProviderConfigManager._get_azure_config should use base_model for detection.""" + + def test_should_return_gpt5_config_when_base_model_is_gpt5(self): + config = ProviderConfigManager._get_azure_config( + model="my-deployment-id", base_model="azure/gpt-5.2" + ) + assert isinstance(config, AzureOpenAIGPT5Config) + + def test_should_return_o_series_config_when_base_model_is_o_series(self): + config = ProviderConfigManager._get_azure_config( + model="my-deployment-id", base_model="azure/o4-mini" + ) + assert isinstance(config, AzureOpenAIO1Config) + + def test_should_return_default_config_when_base_model_is_regular(self): + config = ProviderConfigManager._get_azure_config( + model="my-deployment-id", base_model="azure/gpt-4o" + ) + assert type(config).__name__ == "AzureOpenAIConfig" + + def test_should_fallback_to_model_when_base_model_is_none(self): + config = ProviderConfigManager._get_azure_config( + model="gpt-5.2", base_model=None + ) + assert isinstance(config, AzureOpenAIGPT5Config) + + def test_should_return_default_config_when_both_are_non_standard(self): + config = ProviderConfigManager._get_azure_config( + model="my-deployment-id", base_model=None + ) + assert type(config).__name__ == "AzureOpenAIConfig" + + +# --------------------------------------------------------------------------- +# get_provider_chat_config — threads base_model through for Azure +# --------------------------------------------------------------------------- +class TestGetProviderChatConfigWithBaseModel: + """get_provider_chat_config should pass base_model to Azure config selection.""" + + def test_should_return_gpt5_config_for_custom_deployment_with_base_model(self): + from litellm.types.utils import LlmProviders + + config = ProviderConfigManager.get_provider_chat_config( + model="my-deployment-id", + provider=LlmProviders.AZURE, + base_model="azure/gpt-5", + ) + assert isinstance(config, AzureOpenAIGPT5Config) + + def test_should_return_o_series_config_for_custom_deployment_with_base_model(self): + from litellm.types.utils import LlmProviders + + config = ProviderConfigManager.get_provider_chat_config( + model="my-other-deployment", + provider=LlmProviders.AZURE, + base_model="azure/o3-mini", + ) + assert isinstance(config, AzureOpenAIO1Config) + + +# --------------------------------------------------------------------------- +# get_supported_openai_params — base_model drives Azure param detection +# --------------------------------------------------------------------------- +class TestGetSupportedOpenAIParamsWithBaseModel: + """get_supported_openai_params should use base_model for Azure detection.""" + + def test_should_return_gpt5_params_for_custom_deployment_with_gpt5_base_model( + self, + ): + params = litellm.get_supported_openai_params( + model="my-deployment-id", + custom_llm_provider="azure", + base_model="azure/gpt-5", + ) + assert params is not None + assert "reasoning_effort" in params + # gpt-5 maps max_tokens -> max_completion_tokens, verifying we got GPT-5 config + assert "max_completion_tokens" in params + + def test_should_return_o_series_params_for_custom_deployment_with_o_series_base_model( + self, + ): + params = litellm.get_supported_openai_params( + model="my-other-deployment", + custom_llm_provider="azure", + base_model="azure/o4-mini", + ) + assert params is not None + assert "reasoning_effort" in params + + def test_should_return_regular_params_when_no_base_model(self): + """When base_model is not set and model is non-standard, default Azure config.""" + params = litellm.get_supported_openai_params( + model="my-deployment-id", + custom_llm_provider="azure", + ) + assert params is not None + # Default Azure config supports temperature + assert "temperature" in params + + +# --------------------------------------------------------------------------- +# get_optional_params — base_model drives Azure param mapping +# --------------------------------------------------------------------------- +class TestGetOptionalParamsWithBaseModel: + """get_optional_params should use base_model for Azure model-type detection.""" + + def test_should_map_max_tokens_for_custom_deployment_with_gpt5_base_model(self): + """A non-standard deployment name + gpt-5 base_model should map max_tokens -> max_completion_tokens.""" + params = get_optional_params( + model="my-deployment-id", + custom_llm_provider="azure", + max_tokens=100, + base_model="azure/gpt-5", + ) + assert params.get("max_completion_tokens") == 100 + assert "max_tokens" not in params + + def test_should_keep_max_tokens_for_custom_deployment_without_base_model(self): + """A non-standard deployment name without base_model should use default Azure config.""" + params = get_optional_params( + model="my-deployment-id", + custom_llm_provider="azure", + max_tokens=100, + api_version="2024-05-01-preview", + ) + # Default AzureOpenAIConfig keeps max_tokens as-is (or maps based on api_version) + assert "max_tokens" in params or "max_completion_tokens" in params + + def test_should_support_reasoning_effort_for_custom_deployment_with_o_series_base_model( + self, + ): + """A non-standard deployment name + o-series base_model should accept reasoning_effort.""" + params = get_optional_params( + model="my-other-deployment", + custom_llm_provider="azure", + reasoning_effort="low", + base_model="azure/o4-mini", + ) + assert params.get("reasoning_effort") == "low" + + def test_should_reject_temperature_for_custom_deployment_with_gpt5_base_model( + self, + ): + """A non-standard deployment + gpt-5 base_model should reject temperature.""" + with pytest.raises(litellm.UnsupportedParamsError): + get_optional_params( + model="my-deployment-id", + custom_llm_provider="azure", + temperature=0.5, + base_model="azure/gpt-5", + ) + + +# --------------------------------------------------------------------------- +# Backward compatibility — existing patterns still work +# --------------------------------------------------------------------------- +class TestBackwardCompatibility: + """Existing model-name-based and prefix-based patterns must keep working.""" + + def test_should_detect_gpt5_from_model_name(self): + config = ProviderConfigManager._get_azure_config(model="gpt-5.2") + assert isinstance(config, AzureOpenAIGPT5Config) + + def test_should_detect_gpt5_from_gpt5_series_prefix(self): + config = ProviderConfigManager._get_azure_config( + model="gpt5_series/my-deployment" + ) + assert isinstance(config, AzureOpenAIGPT5Config) + + def test_should_detect_o_series_from_model_name(self): + config = ProviderConfigManager._get_azure_config(model="o4-mini") + assert isinstance(config, AzureOpenAIO1Config) + + def test_should_detect_o_series_from_o_series_prefix(self): + config = ProviderConfigManager._get_azure_config(model="o_series/my-deployment") + assert isinstance(config, AzureOpenAIO1Config) + + def test_should_handle_gpt5_chat_model_correctly(self): + """gpt-5-chat models should NOT be routed to GPT-5 config.""" + config = ProviderConfigManager._get_azure_config(model="gpt-5-chat") + assert type(config).__name__ == "AzureOpenAIConfig" + + def test_base_model_overrides_model_detection(self): + """base_model should take priority over model for type detection.""" + # model looks like o-series, but base_model says gpt-5 + config = ProviderConfigManager._get_azure_config( + model="o3-mini", base_model="azure/gpt-5.2" + ) + assert isinstance(config, AzureOpenAIGPT5Config) + + +# --------------------------------------------------------------------------- +# Deep config method awareness — base_model flows into config internals +# --------------------------------------------------------------------------- +class TestBaseModelFlowsIntoConfigInternals: + """base_model should be used by config internal methods (e.g. is_model_gpt_5_2_model).""" + + def test_should_support_logprobs_for_prefixed_deployment_with_gpt52_base_model( + self, + ): + """Deployment 'my-gpt-5.2' with base_model='azure/gpt-5.2' should support logprobs.""" + params = litellm.get_supported_openai_params( + model="gpt5_series/my-gpt-5.2", + custom_llm_provider="azure", + base_model="azure/gpt-5.2", + ) + assert params is not None + assert "logprobs" in params + assert "top_logprobs" in params + + def test_should_support_logprobs_for_plain_deployment_with_gpt52_base_model(self): + """Deployment 'my-deployment-id' with base_model='azure/gpt-5.2' should support logprobs.""" + params = litellm.get_supported_openai_params( + model="my-deployment-id", + custom_llm_provider="azure", + base_model="azure/gpt-5.2", + ) + assert params is not None + assert "logprobs" in params + assert "top_logprobs" in params + + def test_should_not_support_logprobs_for_gpt5_base_model(self): + """Deployment with base_model='azure/gpt-5' (not 5.2) should NOT support logprobs.""" + params = litellm.get_supported_openai_params( + model="my-deployment-id", + custom_llm_provider="azure", + base_model="azure/gpt-5", + ) + assert params is not None + assert "logprobs" not in params + assert "top_logprobs" not in params + + def test_should_pass_logprobs_through_get_optional_params(self): + """logprobs should pass validation in get_optional_params when base_model is gpt-5.2.""" + params = get_optional_params( + model="gpt5_series/my-gpt-5.2", + custom_llm_provider="azure", + logprobs=True, + top_logprobs=5, + base_model="azure/gpt-5.2", + ) + assert params.get("logprobs") is True + assert params.get("top_logprobs") == 5 + + def test_should_map_max_tokens_for_prefixed_deployment_with_gpt5_base_model(self): + """my-gpt-5.2 with base_model should correctly map max_tokens -> max_completion_tokens.""" + params = get_optional_params( + model="gpt5_series/my-gpt-5.2", + custom_llm_provider="azure", + max_tokens=200, + base_model="azure/gpt-5.2", + ) + assert params.get("max_completion_tokens") == 200 + assert "max_tokens" not in params diff --git a/tests/test_litellm/llms/bedrock/test_base_aws_llm.py b/tests/test_litellm/llms/bedrock/test_base_aws_llm.py index a4969e5dacc..3f91f6ac26e 100644 --- a/tests/test_litellm/llms/bedrock/test_base_aws_llm.py +++ b/tests/test_litellm/llms/bedrock/test_base_aws_llm.py @@ -869,14 +869,18 @@ def test_different_roles_without_session_names_should_not_share_cache(): ({}, {"verify": True}), ( {"aws_region_name": "us-east-1"}, - {"region_name": "us-east-1", "verify": True}, + {"verify": True}, ), ( {"aws_sts_endpoint": "https://sts.eu-west-1.amazonaws.com"}, - {"endpoint_url": "https://sts.eu-west-1.amazonaws.com", "verify": True}, + { + "endpoint_url": "https://sts.eu-west-1.amazonaws.com", + "region_name": "eu-west-1", + "verify": True, + }, ), ], - ids=["no_region_or_endpoint", "regional_sts", "explicit_sts_endpoint"], + ids=["no_region_or_endpoint", "bedrock_region_ignored_for_sts", "explicit_sts_endpoint"], ) def test_eks_irsa_ambient_credentials_used(role_kwargs, expected_client_kwargs): """ @@ -925,6 +929,316 @@ def test_eks_irsa_ambient_credentials_used(role_kwargs, expected_client_kwargs): assert ttl is not None +@pytest.mark.parametrize( + "endpoint,expected_region", + [ + ("https://sts.eu-west-1.amazonaws.com", "eu-west-1"), + ("https://sts.us-east-1.amazonaws.com", "us-east-1"), + ("https://sts-fips.us-east-1.amazonaws.com", "us-east-1"), + ("https://sts-fips.us-gov-west-1.amazonaws.com", "us-gov-west-1"), + ("https://sts.us-gov-west-1.amazonaws.com", "us-gov-west-1"), + ("https://sts.cn-north-1.amazonaws.com.cn", "cn-north-1"), + ( + "https://vpce-abc123.sts.eu-west-1.vpce.amazonaws.com", + "eu-west-1", + ), + ("https://sts.amazonaws.com", None), + ("https://invalid.example.com", None), + ], +) +def test_parse_sts_region_from_endpoint(endpoint, expected_region): + assert BaseAWSLLM._parse_sts_region_from_endpoint(endpoint) == expected_region + + +@pytest.mark.parametrize( + "env,aws_sts_endpoint,expected_region", + [ + ({}, None, None), + ({"AWS_REGION": "us-east-1"}, None, "us-east-1"), + ({"AWS_DEFAULT_REGION": "ap-southeast-1"}, None, "ap-southeast-1"), + ({}, "https://sts.eu-west-1.amazonaws.com", "eu-west-1"), + ( + {"AWS_REGION": "us-east-1"}, + "https://sts.eu-west-1.amazonaws.com", + "eu-west-1", + ), + ({}, "https://sts.amazonaws.com", None), + ( + {}, + "https://vpce-abc.sts.eu-central-1.vpce.amazonaws.com", + "eu-central-1", + ), + ], + ids=[ + "no_env_no_endpoint", + "env_region", + "env_default_region", + "parsed_from_endpoint", + "parsed_endpoint_over_env", + "global_endpoint", + "vpce_endpoint", + ], +) +def test_resolve_sts_region(env, aws_sts_endpoint, expected_region): + with patch.dict(os.environ, env, clear=True): + assert ( + BaseAWSLLM._resolve_sts_region(aws_sts_endpoint=aws_sts_endpoint) + == expected_region + ) + + +@pytest.mark.parametrize( + "env,aws_sts_endpoint,ssl_verify,expected", + [ + ({}, None, None, {"verify": True}), + ( + {"AWS_REGION": "us-east-1"}, + None, + None, + {"verify": True, "region_name": "us-east-1"}, + ), + ( + {}, + "https://sts.eu-west-1.amazonaws.com", + None, + { + "verify": True, + "endpoint_url": "https://sts.eu-west-1.amazonaws.com", + "region_name": "eu-west-1", + }, + ), + ( + {"AWS_REGION": "us-east-1"}, + "https://sts.eu-west-1.amazonaws.com", + None, + { + "verify": True, + "endpoint_url": "https://sts.eu-west-1.amazonaws.com", + "region_name": "eu-west-1", + }, + ), + ( + {}, + "https://sts.amazonaws.com", + None, + {"verify": True, "endpoint_url": "https://sts.amazonaws.com"}, + ), + ( + {}, + "https://vpce-abc.sts.eu-central-1.vpce.amazonaws.com", + None, + { + "verify": True, + "endpoint_url": "https://vpce-abc.sts.eu-central-1.vpce.amazonaws.com", + "region_name": "eu-central-1", + }, + ), + ({}, None, False, {"verify": False}), + ( + {"AWS_DEFAULT_REGION": "ap-southeast-1"}, + None, + None, + {"verify": True, "region_name": "ap-southeast-1"}, + ), + ], + ids=[ + "default_verify_only", + "env_region", + "endpoint_with_parsed_region", + "endpoint_parsed_over_env", + "global_endpoint_no_region", + "vpce_endpoint", + "ssl_verify_false", + "env_default_region", + ], +) +def test_build_sts_client_kwargs(env, aws_sts_endpoint, ssl_verify, expected): + base_aws_llm = BaseAWSLLM() + with patch.dict(os.environ, env, clear=True): + assert ( + base_aws_llm._build_sts_client_kwargs( + aws_sts_endpoint=aws_sts_endpoint, + ssl_verify=ssl_verify, + ) + == expected + ) + + +def test_irsa_cross_account_sts_client_uses_resolved_region(): + """IRSA cross-account path must use _build_sts_client_kwargs (env region, not Bedrock).""" + base_aws_llm = BaseAWSLLM() + import tempfile + + with tempfile.NamedTemporaryFile(mode="w", delete=False) as f: + f.write("test-web-identity-token") + token_file = f.name + + try: + with patch.dict( + os.environ, + { + "AWS_WEB_IDENTITY_TOKEN_FILE": token_file, + "AWS_ROLE_ARN": "arn:aws:iam::111111111111:role/eks-service-account-role", + "AWS_REGION": "eu-west-1", + }, + clear=True, + ): + mock_sts_client = MagicMock() + mock_sts_client.assume_role_with_web_identity.return_value = { + "Credentials": { + "AccessKeyId": "temp-key", + "SecretAccessKey": "temp-secret", + "SessionToken": "temp-token", + "Expiration": datetime.now(timezone.utc) + timedelta(hours=1), + } + } + mock_sts_client.assume_role.return_value = { + "Credentials": { + "AccessKeyId": "assumed-key", + "SecretAccessKey": "assumed-secret", + "SessionToken": "assumed-token", + "Expiration": datetime.now(timezone.utc) + timedelta(hours=1), + } + } + + with patch( + "boto3.client", return_value=mock_sts_client + ) as mock_boto3_client: + base_aws_llm._auth_with_aws_role( + aws_access_key_id=None, + aws_secret_access_key=None, + aws_session_token=None, + aws_role_name="arn:aws:iam::222222222222:role/target-role", + aws_session_name="test-session", + aws_region_name="eu-central-1", + ) + + for call in mock_boto3_client.call_args_list: + assert call.args == ("sts",) + assert call.kwargs["region_name"] == "eu-west-1" + assert call.kwargs["verify"] is True + finally: + os.unlink(token_file) + + +def test_web_identity_token_sts_client_uses_build_sts_client_kwargs(): + base_aws_llm = BaseAWSLLM() + mock_sts_client = MagicMock() + mock_sts_client.assume_role_with_web_identity.return_value = { + "Credentials": { + "AccessKeyId": "key", + "SecretAccessKey": "secret", + "SessionToken": "token", + "Expiration": datetime.now(timezone.utc) + timedelta(hours=1), + }, + "PackedPolicySize": 0, + } + + with patch.dict(os.environ, {"AWS_REGION": "eu-west-1"}, clear=True): + with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client: + with patch( + "litellm.llms.bedrock.base_aws_llm.get_secret", + return_value="oidc-token", + ): + base_aws_llm._auth_with_web_identity_token( + aws_web_identity_token="my-token", + aws_role_name="arn:aws:iam::111111111111:role/target", + aws_session_name="test-session", + aws_region_name="eu-central-1", + aws_sts_endpoint="https://sts.eu-west-1.amazonaws.com", + ) + + mock_boto3_client.assert_called_once_with( + "sts", + verify=True, + endpoint_url="https://sts.eu-west-1.amazonaws.com", + region_name="eu-west-1", + ) + + +def test_sts_uses_workload_region_not_bedrock_region(): + """Air-gapped: Bedrock in eu-central-1, STS VPC endpoint in eu-west-1 via AWS_REGION.""" + base_aws_llm = BaseAWSLLM() + mock_expiry = MagicMock() + mock_expiry.tzinfo = timezone.utc + time_diff = MagicMock() + time_diff.total_seconds.return_value = 3600 + mock_expiry.__sub__ = MagicMock(return_value=time_diff) + mock_sts_client = MagicMock() + mock_sts_client.assume_role.return_value = { + "Credentials": { + "AccessKeyId": "assumed-access-key", + "SecretAccessKey": "assumed-secret-key", + "SessionToken": "assumed-session-token", + "Expiration": mock_expiry, + } + } + + with patch.dict(os.environ, {"AWS_REGION": "eu-west-1"}, clear=True): + with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client: + base_aws_llm._auth_with_aws_role( + aws_access_key_id=None, + aws_secret_access_key=None, + aws_session_token=None, + aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole", + aws_session_name="test-session", + aws_region_name="eu-central-1", + ) + mock_boto3_client.assert_called_with( + "sts", + region_name="eu-west-1", + verify=True, + ) + + +def test_sts_endpoint_region_matches_bedrock_region_param(): + """aws_sts_endpoint signing region must not follow aws_region_name when they differ.""" + base_aws_llm = BaseAWSLLM() + mock_expiry = MagicMock() + mock_expiry.tzinfo = timezone.utc + time_diff = MagicMock() + time_diff.total_seconds.return_value = 3600 + mock_expiry.__sub__ = MagicMock(return_value=time_diff) + mock_sts_client = MagicMock() + mock_sts_client.assume_role.return_value = { + "Credentials": { + "AccessKeyId": "assumed-access-key", + "SecretAccessKey": "assumed-secret-key", + "SessionToken": "assumed-session-token", + "Expiration": mock_expiry, + } + } + + env_without_irsa = { + k: v + for k, v in os.environ.items() + if k + not in ( + "AWS_ROLE_ARN", + "AWS_WEB_IDENTITY_TOKEN_FILE", + "AWS_REGION", + "AWS_DEFAULT_REGION", + ) + } + with patch.dict(env_without_irsa, clear=True): + with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client: + base_aws_llm._auth_with_aws_role( + aws_access_key_id=None, + aws_secret_access_key=None, + aws_session_token=None, + aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole", + aws_session_name="test-session", + aws_region_name="eu-central-1", + aws_sts_endpoint="https://sts.eu-west-1.amazonaws.com", + ) + mock_boto3_client.assert_called_with( + "sts", + endpoint_url="https://sts.eu-west-1.amazonaws.com", + region_name="eu-west-1", + verify=True, + ) + + @pytest.mark.parametrize( "role_kwargs,expected_client_kwargs", [ @@ -940,7 +1254,6 @@ def test_eks_irsa_ambient_credentials_used(role_kwargs, expected_client_kwargs): ( {"aws_region_name": "us-east-1"}, { - "region_name": "us-east-1", "aws_access_key_id": "explicit-access-key", "aws_secret_access_key": "explicit-secret-key", "aws_session_token": "assumed-session-token", @@ -951,6 +1264,7 @@ def test_eks_irsa_ambient_credentials_used(role_kwargs, expected_client_kwargs): {"aws_sts_endpoint": "https://sts.eu-west-1.amazonaws.com"}, { "endpoint_url": "https://sts.eu-west-1.amazonaws.com", + "region_name": "eu-west-1", "aws_access_key_id": "explicit-access-key", "aws_secret_access_key": "explicit-secret-key", "aws_session_token": "assumed-session-token", @@ -958,7 +1272,7 @@ def test_eks_irsa_ambient_credentials_used(role_kwargs, expected_client_kwargs): }, ), ], - ids=["no_region_or_endpoint", "regional_sts", "explicit_sts_endpoint"], + ids=["no_region_or_endpoint", "bedrock_region_ignored_for_sts", "explicit_sts_endpoint"], ) def test_explicit_credentials_used_when_provided(role_kwargs, expected_client_kwargs): """ @@ -2112,3 +2426,102 @@ def test_is_already_running_as_role_ssl_verify_passed(): mock_boto3_client.assert_called_once_with( "sts", verify="/path/to/ca-bundle.crt" ) + + +# --------------------------------------------------------------------------- +# LIT-3274: get_bedrock_model_id must strip "bedrock/" prefix and URL-encode +# ARNs for the invoke path (invoke-with-response-stream). Without this fix +# the Bedrock API receives a malformed URL, returns a JSON error body, and +# botocore's EventStreamBuffer raises ChecksumMismatch instead of the real +# error. 0x223a7b22 == ':{\"' — the start of a JSON object. +# --------------------------------------------------------------------------- + + +class TestGetBedrockModelIdArnHandling: + """Unit tests for get_bedrock_model_id with inference-profile ARNs.""" + + ARN = "arn:aws:bedrock:us-east-1:086734376398:inference-profile/global.anthropic.claude-sonnet-4-5-20250929-v1:0" + + def _call(self, model: str, optional_params: dict | None = None) -> str: + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + provider = BaseAWSLLM.get_bedrock_invoke_provider(model) + return BaseAWSLLM.get_bedrock_model_id( + model=model, + provider=provider, + optional_params=optional_params or {}, + ) + + def test_arn_with_bedrock_prefix_is_stripped_and_encoded(self): + """bedrock/arn:... must not appear verbatim in the model_id.""" + model_id = self._call(f"bedrock/{self.ARN}") + assert ( + "bedrock/arn" not in model_id + ), f"'bedrock/' prefix not stripped; got: {model_id}" + # Must be URL-encoded (colons → %3A) + assert "%3A" in model_id, f"ARN not URL-encoded; got: {model_id}" + assert "%2F" in model_id, f"ARN slashes not URL-encoded; got: {model_id}" + + def test_arn_with_compound_bedrock_invoke_prefix_is_fully_stripped_and_encoded( + self, + ): + """bedrock/invoke/arn:... — compound prefix — must be fully stripped. + + The old fix used ``break`` after the first matched prefix, so + ``bedrock/invoke/arn:...`` would only strip ``bedrock/``, leaving + ``invoke/arn:...``. The subsequent ``.replace('invoke/', '')`` call + then returned the bare unencoded ARN, reproducing the same + malformed-URL bug the fix aimed to prevent. + + strip_bedrock_routing_prefix() has no break and handles this correctly. + """ + model_id = self._call(f"bedrock/invoke/{self.ARN}") + assert ( + "invoke/" not in model_id + ), f"'invoke/' prefix not stripped; got: {model_id}" + assert ( + "bedrock/" not in model_id + ), f"'bedrock/' prefix not stripped; got: {model_id}" + assert "%3A" in model_id, f"ARN not URL-encoded; got: {model_id}" + assert "%2F" in model_id, f"ARN slashes not URL-encoded; got: {model_id}" + + def test_bare_arn_is_encoded(self): + """Direct ARN without routing prefix must also be URL-encoded.""" + model_id = self._call(self.ARN) + assert "%3A" in model_id, f"ARN not URL-encoded; got: {model_id}" + assert "%2F" in model_id, f"ARN slashes not URL-encoded; got: {model_id}" + + def test_arn_url_matches_expected(self): + """Full URL built from messages config must match expected encoded form.""" + import urllib.parse + from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import ( + AmazonAnthropicClaudeMessagesConfig, + ) + + config = AmazonAnthropicClaudeMessagesConfig() + url = config.get_complete_url( + api_base=None, + api_key=None, + model=f"bedrock/{self.ARN}", + optional_params={"aws_region_name": "us-east-1"}, + litellm_params={}, + stream=True, + ) + encoded_arn = urllib.parse.quote(self.ARN, safe="") + expected = ( + f"https://bedrock-runtime.us-east-1.amazonaws.com" + f"/model/{encoded_arn}/invoke-with-response-stream" + ) + assert ( + url == expected + ), f"URL mismatch:\n got: {url}\n expected: {expected}" + + def test_regular_model_id_unaffected(self): + """Non-ARN model IDs must continue to work as before.""" + model_id = self._call("anthropic.claude-3-sonnet-20240229-v1:0") + assert model_id == "anthropic.claude-3-sonnet-20240229-v1:0" + + def test_invoke_prefixed_model_unaffected(self): + """invoke/ prefix stripping still works after the fix.""" + model_id = self._call("invoke/anthropic.claude-3-sonnet-20240229-v1:0") + assert model_id == "anthropic.claude-3-sonnet-20240229-v1:0" diff --git a/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py b/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py index acb9fa9b64c..4b2e9471fb7 100644 --- a/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py +++ b/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py @@ -86,6 +86,90 @@ class TestOpenAIResponsesAPIConfig: self.validate_responses_api_request_params(result, expected_fields) + def test_transform_strips_cache_control_from_input_content_blocks(self): + """`cache_control` markers (Anthropic-only) must be stripped from + Responses API input content blocks before sending to OpenAI. + + OpenAI rejects unknown params on input content blocks with HTTP 400: + "Unknown parameter: 'input[0].content[0].cache_control'" + Chat Completions strips these via + `remove_cache_control_flag_from_messages_and_tools`; the Responses + path must do the same. + """ + input_with_cache_control = [ + { + "role": "user", + "content": [ + { + "type": "input_text", + "text": "Hello", + "cache_control": {"type": "ephemeral"}, + } + ], + } + ] + + result = self.config.transform_responses_api_request( + model=self.model, + input=input_with_cache_control, + response_api_optional_request_params={}, + litellm_params={}, + headers={}, + ) + + assert "cache_control" not in result["input"][0]["content"][0] + assert result["input"][0]["content"][0]["type"] == "input_text" + assert result["input"][0]["content"][0]["text"] == "Hello" + + def test_transform_strips_cache_control_from_tools(self): + """`cache_control` markers must also be stripped from tools for + symmetry with the Chat Completions path. OpenAI currently accepts + cache_control on tools silently but stripping keeps the wire payload + clean and matches `remove_cache_control_flag_from_messages_and_tools`. + """ + tools_with_cache_control = [ + { + "type": "function", + "name": "get_weather", + "description": "Get the weather", + "parameters": { + "type": "object", + "properties": {"city": {"type": "string"}}, + }, + "cache_control": {"type": "ephemeral"}, + } + ] + + result = self.config.transform_responses_api_request( + model=self.model, + input="hi", + response_api_optional_request_params={"tools": tools_with_cache_control}, + litellm_params={}, + headers={}, + ) + + assert "cache_control" not in result["tools"][0] + assert result["tools"][0]["name"] == "get_weather" + + def test_transform_preserves_input_without_cache_control(self): + """Inputs without cache_control must pass through unmodified.""" + input_clean = [ + { + "role": "user", + "content": [{"type": "input_text", "text": "Hello"}], + } + ] + + result = self.config.transform_responses_api_request( + model=self.model, + input=input_clean, + response_api_optional_request_params={}, + litellm_params={}, + headers={}, + ) + + assert result["input"] == input_clean + def test_transform_streaming_response(self): """Test streaming response transformation""" # Test with a text delta event diff --git a/tests/test_litellm/llms/sagemaker/test_sagemaker_embedding_voyage.py b/tests/test_litellm/llms/sagemaker/test_sagemaker_embedding_voyage.py index a36aec32d13..943a3160bb7 100644 --- a/tests/test_litellm/llms/sagemaker/test_sagemaker_embedding_voyage.py +++ b/tests/test_litellm/llms/sagemaker/test_sagemaker_embedding_voyage.py @@ -17,6 +17,9 @@ import pytest sys.path.insert(0, os.path.abspath("../../../../..")) from litellm import embedding +from litellm.llms.sagemaker.embedding.cohere_transformation import ( + SagemakerCohereEmbeddingConfig, +) from litellm.llms.sagemaker.embedding.transformation import SagemakerEmbeddingConfig from litellm.llms.voyage.embedding.transformation import VoyageEmbeddingConfig from litellm.types.utils import EmbeddingResponse, Usage @@ -54,6 +57,172 @@ class TestSagemakerEmbeddingFactory: assert isinstance(config2, VoyageEmbeddingConfig) assert isinstance(config3, VoyageEmbeddingConfig) + def test_get_model_config_cohere_model(self): + """Cohere SageMaker endpoints route to SagemakerCohereEmbeddingConfig""" + for endpoint_name in ( + "cohere.embed-multilingual-v3", + "cohere-embed-english-v3-prod", + "my-cohere-marketplace-endpoint", + "COHERE-EMBED-V4", + ): + config = SagemakerEmbeddingConfig.get_model_config(endpoint_name) + assert isinstance(config, SagemakerCohereEmbeddingConfig), endpoint_name + + +class TestSagemakerCohereEmbeddingConfig: + """Cohere-specific SageMaker embedding request/response transforms""" + + def setup_method(self): + self.config = SagemakerCohereEmbeddingConfig() + + MODEL = "cohere.embed-multilingual-v3" + + def test_transform_request_uses_cohere_payload(self): + """Bug repro: request must use `texts` + `input_type`, not HF `inputs`""" + result = self.config.transform_embedding_request( + model=self.MODEL, + input=["hello"], + optional_params={"input_type": "search_query"}, + headers={}, + ) + assert "inputs" not in result + assert result["texts"] == ["hello"] + assert result["input_type"] == "search_query" + + def test_transform_request_default_input_type(self): + result = self.config.transform_embedding_request( + model=self.MODEL, + input=["hello"], + optional_params={}, + headers={}, + ) + assert result["texts"] == ["hello"] + assert result["input_type"] == "search_document" + + def test_transform_request_normalizes_string_input(self): + result = self.config.transform_embedding_request( + model=self.MODEL, + input="hello", + optional_params={}, + headers={}, + ) + assert result["texts"] == ["hello"] + + def test_map_openai_params_dimensions_to_output_dimension(self): + params = self.config.map_openai_params( + non_default_params={"dimensions": 512, "encoding_format": "float"}, + optional_params={}, + model=self.MODEL, + drop_params=False, + ) + assert params["output_dimension"] == 512 + assert params["embedding_types"] == ["float"] + + def test_map_openai_params_input_type_from_non_default_params(self): + params = self.config.map_openai_params( + non_default_params={"input_type": "search_query"}, + optional_params={}, + model=self.MODEL, + drop_params=False, + ) + assert params["input_type"] == "search_query" + + def test_get_optional_params_embeddings_preserves_input_type(self): + """Exercises get_optional_params_embeddings, not transform in isolation.""" + from litellm.utils import get_optional_params_embeddings + + optional_params = get_optional_params_embeddings( + model=self.MODEL, + custom_llm_provider="sagemaker", + input_type="search_query", + ) + assert optional_params.get("input_type") == "search_query" + + body = self.config.transform_embedding_request( + model=self.MODEL, + input=["hello"], + optional_params=optional_params, + headers={}, + ) + assert body["texts"] == ["hello"] + assert body["input_type"] == "search_query" + + def test_get_optional_params_embeddings_maps_dimensions_without_duplicate(self): + """dimensions must map to output_dimension only, not also stay as dimensions.""" + from litellm.utils import get_optional_params_embeddings + + optional_params = get_optional_params_embeddings( + model=self.MODEL, + custom_llm_provider="sagemaker", + dimensions=512, + input_type="search_query", + ) + assert optional_params.get("output_dimension") == 512 + assert "dimensions" not in optional_params + assert optional_params.get("input_type") == "search_query" + + def test_transform_response_parses_cohere_payload(self): + cohere_response = { + "embeddings": [[0.1, 0.2, 0.3]], + "meta": {"billed_units": {"input_tokens": 2}}, + } + mock_response = httpx.Response( + status_code=200, + content=json.dumps(cohere_response).encode("utf-8"), + headers={"content-type": "application/json"}, + ) + logging_obj = MagicMock() + logging_obj.model_call_details = {"input": ["hello"]} + + result = self.config.transform_embedding_response( + model=self.MODEL, + raw_response=mock_response, + model_response=EmbeddingResponse(), + logging_obj=logging_obj, + api_key=None, + request_data={"texts": ["hello"], "input_type": "search_query"}, + optional_params={}, + litellm_params={}, + ) + + assert result.object == "list" + assert len(result.data) == 1 + assert result.data[0]["embedding"] == [0.1, 0.2, 0.3] + assert result.usage.prompt_tokens == 2 + + def test_transform_response_does_not_double_call_post_call(self): + """ + Greptile review fix: SageMaker handler already calls + `logging_obj.post_call` once before invoking + `transform_embedding_response`. The transform must NOT call it again, + otherwise callbacks, cost calculators, and log handlers double-fire + for every Cohere SageMaker embedding call. + """ + cohere_response = { + "embeddings": [[0.1, 0.2, 0.3]], + "meta": {"billed_units": {"input_tokens": 2}}, + } + mock_response = httpx.Response( + status_code=200, + content=json.dumps(cohere_response).encode("utf-8"), + headers={"content-type": "application/json"}, + ) + logging_obj = MagicMock() + logging_obj.model_call_details = {"input": ["hello"]} + + self.config.transform_embedding_response( + model=self.MODEL, + raw_response=mock_response, + model_response=EmbeddingResponse(), + logging_obj=logging_obj, + api_key=None, + request_data={"texts": ["hello"], "input_type": "search_query"}, + optional_params={}, + litellm_params={}, + ) + + logging_obj.post_call.assert_not_called() + class TestVoyageEmbeddingConfig: """Test Voyage-specific embedding configuration""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index f2fd73f3f22..d62720ed36f 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -3036,6 +3036,206 @@ class TestMergeGatewayInitializeInstructions: ) +class TestEnsureUpstreamInitializeInstructionsCached: + @pytest.mark.asyncio + async def test_skips_when_yaml_instructions_set(self): + from unittest.mock import AsyncMock, patch + + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + server = _make_instruction_server( + server_id="yaml-only", instructions="from yaml" + ) + with patch.object( + global_mcp_server_manager, "_create_mcp_client", AsyncMock() + ) as mock_create: + await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached( + server + ) + mock_create.assert_not_awaited() + + @pytest.mark.asyncio + async def test_skips_when_already_cached(self): + from unittest.mock import AsyncMock, patch + + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + server = _make_instruction_server(server_id="cached-only", instructions=None) + global_mcp_server_manager._upstream_initialize_instructions_by_server_id[ + "cached-only" + ] = "warm" + try: + with patch.object( + global_mcp_server_manager, "_create_mcp_client", AsyncMock() + ) as mock_create: + await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached( + server + ) + mock_create.assert_not_awaited() + finally: + global_mcp_server_manager._upstream_initialize_instructions_by_server_id.pop( + "cached-only", None + ) + + @pytest.mark.asyncio + async def test_skips_when_spec_path_set(self): + from unittest.mock import AsyncMock, patch + + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + server = _make_instruction_server( + server_id="openapi-spec", spec_path="/openapi.json", url=None + ) + with patch.object( + global_mcp_server_manager, "_create_mcp_client", AsyncMock() + ) as mock_create: + await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached( + server + ) + mock_create.assert_not_awaited() + + @pytest.mark.asyncio + async def test_runs_upstream_session_and_caches(self): + from unittest.mock import AsyncMock, MagicMock, patch + + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + server = _make_instruction_server(server_id="cold-server", instructions=None) + fake_client = MagicMock() + fake_client.run_with_session = AsyncMock(return_value="ok") + fake_client._last_initialize_instructions = " upstream says hi " + + with patch.object( + global_mcp_server_manager, + "_create_mcp_client", + AsyncMock(return_value=fake_client), + ): + try: + await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached( + server + ) + assert ( + global_mcp_server_manager._upstream_initialize_instructions_by_server_id[ + "cold-server" + ] + == "upstream says hi" + ) + finally: + global_mcp_server_manager._upstream_initialize_instructions_by_server_id.pop( + "cold-server", None + ) + global_mcp_server_manager._upstream_initialize_instructions_probed_at.pop( + "cold-server", None + ) + + @pytest.mark.asyncio + async def test_cooldown_after_empty_upstream_response(self): + """Upstream returns no instructions → next call within cooldown must not reconnect.""" + from unittest.mock import AsyncMock, MagicMock, patch + + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + server = _make_instruction_server(server_id="empty-server", instructions=None) + fake_client = MagicMock() + fake_client.run_with_session = AsyncMock(return_value="ok") + fake_client._last_initialize_instructions = None # upstream sent nothing + + create = AsyncMock(return_value=fake_client) + with patch.object(global_mcp_server_manager, "_create_mcp_client", create): + try: + await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached( + server + ) + await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached( + server + ) + assert create.await_count == 1, ( + "Second probe within cooldown must not reconnect to upstream" + ) + assert ( + "empty-server" + not in global_mcp_server_manager._upstream_initialize_instructions_by_server_id + ) + assert ( + "empty-server" + in global_mcp_server_manager._upstream_initialize_instructions_probed_at + ) + finally: + global_mcp_server_manager._upstream_initialize_instructions_probed_at.pop( + "empty-server", None + ) + + @pytest.mark.asyncio + async def test_cooldown_after_upstream_failure(self): + """run_with_session raises → cooldown applies, no immediate retry.""" + from unittest.mock import AsyncMock, MagicMock, patch + + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + server = _make_instruction_server(server_id="boom-server", instructions=None) + fake_client = MagicMock() + fake_client.run_with_session = AsyncMock(side_effect=RuntimeError("upstream down")) + fake_client._last_initialize_instructions = None + + create = AsyncMock(return_value=fake_client) + with patch.object(global_mcp_server_manager, "_create_mcp_client", create): + try: + await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached( + server + ) + await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached( + server + ) + assert create.await_count == 1, ( + "Second probe within cooldown must not reconnect after failure" + ) + assert ( + "boom-server" + not in global_mcp_server_manager._upstream_initialize_instructions_by_server_id + ) + assert ( + "boom-server" + in global_mcp_server_manager._upstream_initialize_instructions_probed_at + ) + finally: + global_mcp_server_manager._upstream_initialize_instructions_probed_at.pop( + "boom-server", None + ) + + @pytest.mark.asyncio + async def test_reload_resets_probe_cooldown(self): + """load_servers_from_config clears the negative-cache map so reloads re-probe.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + global_mcp_server_manager._upstream_initialize_instructions_probed_at[ + "reload-target" + ] = 1.0 + try: + await global_mcp_server_manager.load_servers_from_config({}) + assert ( + "reload-target" + not in global_mcp_server_manager._upstream_initialize_instructions_probed_at + ) + finally: + global_mcp_server_manager._upstream_initialize_instructions_probed_at.pop( + "reload-target", None + ) + + class TestGatewayCreateInitializationOptions: """Tests for the patched server.create_initialization_options via ContextVar.""" diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 35a3bd7f657..116ba83f42e 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -127,6 +127,23 @@ def test_get_experimental_ui_login_jwt_auth_token_valid(valid_sso_user_defined_v assert expires <= now + timedelta(minutes=10, seconds=2) +def test_get_cli_jwt_auth_token_includes_team_alias(valid_sso_user_defined_values): + token = ExperimentalUIJWTToken.get_cli_jwt_auth_token( + valid_sso_user_defined_values, + team_id="team-123", + team_alias="test-team", + ) + + decrypted_token = decrypt_value_helper( + token, key="ui_hash_key", exception_type="debug" + ) + assert decrypted_token is not None + token_data = json.loads(decrypted_token) + + assert token_data["team_id"] == "team-123" + assert token_data["team_alias"] == "test-team" + + def test_get_experimental_ui_login_jwt_auth_token_uses_10_min_expiry( valid_sso_user_defined_values, ): diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_microsoft_purview.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_microsoft_purview.py new file mode 100644 index 00000000000..cc89cea58d2 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_microsoft_purview.py @@ -0,0 +1,2659 @@ +"""Unit tests for the Microsoft Purview DLP guardrail.""" + +import asyncio +import time +from unittest.mock import AsyncMock, Mock, patch + +import httpx +import pytest +from fastapi import HTTPException + +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_hooks.microsoft_purview.base import ( + PurviewGuardrailBase, +) +from litellm.proxy.guardrails.guardrail_hooks.microsoft_purview.purview_dlp import ( + MicrosoftPurviewDLPGuardrail, +) + + +def _make_guardrail(**kwargs) -> MicrosoftPurviewDLPGuardrail: + """Helper to construct a guardrail with test defaults.""" + defaults = { + "guardrail_name": "test-purview", + "tenant_id": "test-tenant-id", + "client_id": "test-client-id", + "client_secret": "test-client-secret", + } + defaults.update(kwargs) + return MicrosoftPurviewDLPGuardrail(**defaults) + + +def _mock_token_response(): + """Mock a successful OAuth2 token response.""" + resp = Mock() + resp.json.return_value = { + "access_token": "mock-access-token", + "expires_in": 3600, + } + return resp + + +def _mock_graph_response(policy_actions=None, protection_scope_state="unchanged"): + """Mock a processContent Graph API response.""" + resp = Mock() + body = { + "protectionScopeState": protection_scope_state, + "policyActions": policy_actions or [], + "processingErrors": [], + } + resp.json.return_value = body + resp.headers = {"ETag": "test-etag-123"} + return resp + + +def _mock_scope_response(): + """Mock a protectionScopes/compute Graph API response.""" + resp = Mock() + resp.json.return_value = { + "value": [ + { + "activities": "uploadText,downloadText", + "executionMode": "evaluateInline", + "policyActions": [], + } + ] + } + resp.headers = {"ETag": "scope-etag-123"} + return resp + + +# --------------------------------------------------------------- +# _should_block +# --------------------------------------------------------------- + + +class TestShouldBlock: + def test_empty_policy_actions(self): + assert PurviewGuardrailBase._should_block({"policyActions": []}) is False + + def test_no_policy_actions_key(self): + assert PurviewGuardrailBase._should_block({}) is False + + def test_restrict_access_block(self): + response = { + "policyActions": [ + { + "@odata.type": "#microsoft.graph.restrictAccessAction", + "action": "restrictAccess", + "restrictionAction": "block", + } + ] + } + assert PurviewGuardrailBase._should_block(response) is True + + def test_restrict_access_non_block(self): + response = { + "policyActions": [ + { + "@odata.type": "#microsoft.graph.restrictAccessAction", + "action": "restrictAccess", + "restrictionAction": "warn", + } + ] + } + assert PurviewGuardrailBase._should_block(response) is False + + def test_non_restrict_action(self): + response = { + "policyActions": [ + { + "@odata.type": "#microsoft.graph.auditAction", + "action": "audit", + } + ] + } + assert PurviewGuardrailBase._should_block(response) is False + + def test_multiple_actions_one_blocks(self): + response = { + "policyActions": [ + {"action": "audit"}, + { + "@odata.type": "#microsoft.graph.restrictAccessAction", + "action": "restrictAccess", + "restrictionAction": "block", + }, + ] + } + assert PurviewGuardrailBase._should_block(response) is True + + +# --------------------------------------------------------------- +# completion prompt normalization (text completions API) +# --------------------------------------------------------------- + + +class TestCompletionPromptToStr: + def test_string_prompt(self): + assert PurviewGuardrailBase.completion_prompt_to_str(" hi ") == "hi" + + def test_list_of_strings(self): + assert PurviewGuardrailBase.completion_prompt_to_str(["a", "b"]) == "a\nb" + + def test_token_ids_returns_none(self): + assert PurviewGuardrailBase.completion_prompt_to_str([1, 2, 3]) is None + + def test_empty(self): + assert PurviewGuardrailBase.completion_prompt_to_str("") is None + assert PurviewGuardrailBase.completion_prompt_to_str([]) is None + + +# --------------------------------------------------------------- +# User ID resolution +# --------------------------------------------------------------- + + +class TestResolveUserId: + def test_from_metadata_when_no_auth_identity(self): + guardrail = _make_guardrail() + data = {"metadata": {"user_id": "entra-user-123"}} + auth = UserAPIKeyAuth(api_key="test-key-no-user") + assert guardrail._resolve_user_id(data, auth) == "entra-user-123" + + def test_authenticated_user_id_overrides_metadata(self): + """Key user_id must win over spoofed metadata[user_id_field].""" + guardrail = _make_guardrail() + data = {"metadata": {"user_id": "spoofed-entra-id"}} + auth = UserAPIKeyAuth(api_key="test", user_id="real-entra-id") + assert guardrail._resolve_user_id(data, auth) == "real-entra-id" + + def test_user_api_key_metadata_before_custom_field(self): + """Proxy-injected user_api_key_user_id wins over arbitrary metadata field.""" + guardrail = _make_guardrail(user_id_field="entra_id") + data = { + "metadata": { + "user_api_key_user_id": "from-proxy-111", + "entra_id": "metadata-222", + } + } + auth = UserAPIKeyAuth(api_key="test") + assert guardrail._resolve_user_id(data, auth) == "from-proxy-111" + + def test_custom_field_when_no_stronger_source(self): + guardrail = _make_guardrail(user_id_field="entra_id") + data = {"metadata": {"entra_id": "custom-user-456"}} + auth = UserAPIKeyAuth(api_key="test") + assert guardrail._resolve_user_id(data, auth) == "custom-user-456" + + def test_from_user_api_key_dict_user_id(self): + guardrail = _make_guardrail() + auth = UserAPIKeyAuth(api_key="test", user_id="key-user-789") + assert guardrail._resolve_user_id({}, auth) == "key-user-789" + + def test_from_end_user_id(self): + guardrail = _make_guardrail() + auth = UserAPIKeyAuth(api_key="test", end_user_id="end-user-101") + assert guardrail._resolve_user_id({}, auth) == "end-user-101" + + def test_end_user_id_after_key_user_id(self): + """When both key user_id and end_user_id exist, key user_id is used first.""" + guardrail = _make_guardrail() + auth = UserAPIKeyAuth( + api_key="test", user_id="key-owner", end_user_id="end-user-101" + ) + assert guardrail._resolve_user_id({}, auth) == "key-owner" + + def test_none_when_missing(self): + guardrail = _make_guardrail() + auth = UserAPIKeyAuth(api_key="test") + assert guardrail._resolve_user_id({}, auth) is None + + +# --------------------------------------------------------------- +# Pre-call hook +# --------------------------------------------------------------- + + +class TestPreCallHook: + @pytest.mark.asyncio + async def test_pre_call_allow(self): + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + cache=None, + data={"messages": [{"role": "user", "content": "Hello, how are you?"}]}, + call_type="completion", + ) + + mock_check.assert_called_once() + assert mock_check.call_args.kwargs["activity"] == "uploadText" + assert mock_check.call_args.kwargs["block_on_violation"] is True + + @pytest.mark.asyncio + async def test_pre_call_success_returns_request_data(self): + """After a successful DLP pass, the hook must return the same data dict (not None).""" + guardrail = _make_guardrail() + payload = { + "messages": [{"role": "user", "content": "Hello, how are you?"}], + "litellm_call_id": "call-abc", + } + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + out = await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + cache=None, + data=payload, + call_type="completion", + ) + + assert out is payload + + @pytest.mark.asyncio + async def test_pre_call_block(self): + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.side_effect = HTTPException( + status_code=400, + detail={"error": "Microsoft Purview DLP: Content blocked by policy"}, + ) + + with pytest.raises(HTTPException) as exc_info: + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth( + api_key="test", user_id="user-123" + ), + cache=None, + data={ + "messages": [ + { + "role": "user", + "content": "SSN: 123-45-6789", + } + ] + }, + call_type="completion", + ) + + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_pre_call_no_user_id_raises(self): + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + with pytest.raises(HTTPException) as exc_info: + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + cache=None, + data={"messages": [{"role": "user", "content": "Hello"}]}, + call_type="completion", + ) + + mock_check.assert_not_called() + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_pre_call_no_messages_skips(self): + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + cache=None, + data={}, + call_type="completion", + ) + + mock_check.assert_not_called() + + +class TestPreCallFullTranscript: + @pytest.mark.asyncio + async def test_pre_call_sends_all_message_roles_to_dlp(self): + """DLP text must include system / prior turns, not only the last user block.""" + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + cache=None, + data={ + "messages": [ + {"role": "system", "content": "SYSTEM_SENSITIVE"}, + {"role": "user", "content": "EARLIER_USER"}, + {"role": "assistant", "content": "reply"}, + {"role": "user", "content": "final benign"}, + ] + }, + call_type="completion", + ) + + mock_check.assert_called_once() + sent = mock_check.call_args.kwargs["text"] + assert "SYSTEM_SENSITIVE" in sent + assert "EARLIER_USER" in sent + assert "final benign" in sent + + +# --------------------------------------------------------------- +# Post-call hook +# --------------------------------------------------------------- + + +class TestPostCallHook: + @pytest.mark.asyncio + async def test_post_call_allow(self): + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _make_guardrail() + response = ModelResponse( + choices=[ + Choices( + index=0, message=Message(content="Safe response", role="assistant") + ) + ], + ) + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + result = await guardrail.async_post_call_success_hook( + data={}, + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + response=response, + ) + + mock_check.assert_called_once() + assert mock_check.call_args.kwargs["activity"] == "downloadText" + assert result is response + + @pytest.mark.asyncio + async def test_post_call_block(self): + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _make_guardrail() + response = ModelResponse( + choices=[ + Choices( + index=0, + message=Message( + content="Credit card: 4532-6677-8521-3500", + role="assistant", + ), + ) + ], + ) + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.side_effect = HTTPException( + status_code=400, + detail={"error": "Microsoft Purview DLP: Content blocked by policy"}, + ) + + with pytest.raises(HTTPException) as exc_info: + await guardrail.async_post_call_success_hook( + data={}, + user_api_key_dict=UserAPIKeyAuth( + api_key="test", user_id="user-123" + ), + response=response, + ) + + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_post_call_no_user_id_raises(self): + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _make_guardrail() + response = ModelResponse( + choices=[ + Choices(index=0, message=Message(content="Response", role="assistant")) + ], + ) + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + with pytest.raises(HTTPException) as exc_info: + await guardrail.async_post_call_success_hook( + data={}, + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=response, + ) + + mock_check.assert_not_called() + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_post_call_scans_all_choices(self): + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _make_guardrail() + response = ModelResponse( + choices=[ + Choices( + index=0, + message=Message(content="First completion", role="assistant"), + ), + Choices( + index=1, + message=Message(content="Second completion body", role="assistant"), + ), + ], + ) + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + await guardrail.async_post_call_success_hook( + data={}, + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + response=response, + ) + + mock_check.assert_called_once() + combined = mock_check.call_args.kwargs["text"] + assert "First completion" in combined + assert "Second completion body" in combined + + +class TestTextCompletionHooks: + @pytest.mark.asyncio + async def test_pre_call_text_completion_uses_prompt(self): + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + cache=None, + data={"prompt": "Completions API prompt body"}, + call_type="text_completion", + ) + + mock_check.assert_called_once() + assert mock_check.call_args.kwargs["text"] == "Completions API prompt body" + assert mock_check.call_args.kwargs["activity"] == "uploadText" + + @pytest.mark.asyncio + async def test_post_call_text_completion_all_choices(self): + from litellm.types.utils import TextChoices, TextCompletionResponse + + guardrail = _make_guardrail() + response = TextCompletionResponse( + model="gpt-3.5-turbo-instruct", + choices=[ + TextChoices(text="alpha", index=0), + TextChoices(text="beta", index=1), + ], + ) + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + await guardrail.async_post_call_success_hook( + data={}, + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + response=response, + ) + + mock_check.assert_called_once() + combined = mock_check.call_args.kwargs["text"] + assert "alpha" in combined + assert "beta" in combined + + +# --------------------------------------------------------------- +# Responses API hooks +# --------------------------------------------------------------- + + +class TestResponsesAPIHooks: + @pytest.mark.asyncio + async def test_pre_call_responses_api_string_input(self): + """Pre-call hook must scan plain-string ``input`` on responses call type.""" + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + cache=None, + data={"input": "SSN: 123-45-6789"}, + call_type="responses", + ) + + mock_check.assert_called_once() + assert mock_check.call_args.kwargs["activity"] == "uploadText" + assert "SSN: 123-45-6789" in mock_check.call_args.kwargs["text"] + + @pytest.mark.asyncio + async def test_pre_call_aresponses_string_input(self): + """Pre-call hook must scan ``input`` on ``aresponses`` call type too.""" + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + cache=None, + data={"input": "sensitive content"}, + call_type="aresponses", + ) + + mock_check.assert_called_once() + assert "sensitive content" in mock_check.call_args.kwargs["text"] + + @pytest.mark.asyncio + async def test_pre_call_responses_api_list_input(self): + """Pre-call hook must extract text from structured list ``input``.""" + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + cache=None, + data={ + "input": [{"role": "user", "content": "Secret phrase: alpha bravo"}] + }, + call_type="responses", + ) + + mock_check.assert_called_once() + assert "Secret phrase: alpha bravo" in mock_check.call_args.kwargs["text"] + + @pytest.mark.asyncio + async def test_pre_call_responses_api_no_input_skips(self): + """Pre-call hook must not call _check_content when ``input`` is absent.""" + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + cache=None, + data={}, + call_type="responses", + ) + + mock_check.assert_not_called() + + @pytest.mark.asyncio + async def test_pre_call_responses_string_input_includes_instructions(self): + """Benign string ``input`` must still scan ``instructions`` (system message).""" + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + cache=None, + data={ + "input": "benign user text", + "instructions": "SYSTEM_SENSITIVE in instructions", + }, + call_type="responses", + ) + + mock_check.assert_called_once() + sent = mock_check.call_args.kwargs["text"] + assert "benign user text" in sent + assert "SYSTEM_SENSITIVE in instructions" in sent + + @pytest.mark.asyncio + async def test_pre_call_responses_instructions_only(self): + """Requests with only ``instructions`` (no ``input``) must still be scanned.""" + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + cache=None, + data={"instructions": "policy text in instructions only"}, + call_type="responses", + ) + + mock_check.assert_called_once() + assert ( + "policy text in instructions only" + in mock_check.call_args.kwargs["text"] + ) + + @pytest.mark.asyncio + async def test_post_call_responses_api_output_text(self): + """Post-call hook must scan text from ``ResponsesAPIResponse.output``.""" + from litellm.types.llms.openai import ResponsesAPIResponse + + guardrail = _make_guardrail() + response = ResponsesAPIResponse( + id="resp-1", + created_at=0, + output=[ + { + "type": "message", + "id": "msg-1", + "status": "completed", + "role": "assistant", + "content": [ + {"type": "output_text", "text": "card 4111-1111-1111-1111"} + ], + } + ], + ) + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + result = await guardrail.async_post_call_success_hook( + data={}, + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + response=response, + ) + + mock_check.assert_called_once() + assert mock_check.call_args.kwargs["activity"] == "downloadText" + assert "card 4111-1111-1111-1111" in mock_check.call_args.kwargs["text"] + assert result is response + + @pytest.mark.asyncio + async def test_post_call_responses_api_empty_output_skips(self): + """Post-call hook must not call _check_content when output has no text.""" + from litellm.types.llms.openai import ResponsesAPIResponse + + guardrail = _make_guardrail() + response = ResponsesAPIResponse( + id="resp-2", + created_at=0, + output=[], + ) + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + await guardrail.async_post_call_success_hook( + data={}, + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + response=response, + ) + + mock_check.assert_not_called() + + @pytest.mark.asyncio + async def test_logging_hook_responses_api_input_and_output(self): + """Logging hook must scan both ``input`` and ``ResponsesAPIResponse.output``.""" + from litellm.types.llms.openai import ResponsesAPIResponse + + guardrail = _make_guardrail() + result_response = ResponsesAPIResponse( + id="resp-3", + created_at=0, + output=[ + { + "type": "message", + "id": "msg-2", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "response body"}], + } + ], + ) + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + await guardrail.async_logging_hook( + kwargs={ + "input": "prompt body", + "litellm_params": { + "metadata": {"user_api_key_user_id": "user-123"} + }, + }, + result=result_response, + call_type="responses", + ) + + assert mock_check.call_count == 2 + activities = {c.kwargs["activity"] for c in mock_check.call_args_list} + assert activities == {"uploadText", "downloadText"} + texts = {c.kwargs["text"] for c in mock_check.call_args_list} + assert any("prompt body" in t for t in texts) + assert any("response body" in t for t in texts) + + @pytest.mark.asyncio + async def test_logging_hook_responses_api_with_messages_key_set(self): + """Responses-API prompt audit must fire even when ``kwargs["messages"]`` is + also set to the raw responses input. + + litellm's logging pipeline (``function_setup`` + + ``update_environment_variables``) stores the raw responses ``input`` + under ``model_call_details["messages"]``. The audit must still extract + the prompt via the responses-specific path, not silently fall through + the generic ``messages`` branch with the wrong format. + """ + from litellm.types.llms.openai import ResponsesAPIResponse + + guardrail = _make_guardrail() + result_response = ResponsesAPIResponse( + id="resp-msgkey", + created_at=0, + output=[ + { + "type": "message", + "id": "msg-3", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "response body"}], + } + ], + ) + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + await guardrail.async_logging_hook( + kwargs={ + "input": "prompt body", + "instructions": "system instructions", + # Simulate litellm's logging path which mirrors the raw + # responses input under "messages". + "messages": "prompt body", + "litellm_params": { + "metadata": {"user_api_key_user_id": "user-123"} + }, + }, + result=result_response, + call_type="aresponses", + ) + + assert mock_check.call_count == 2 + activities = {c.kwargs["activity"] for c in mock_check.call_args_list} + assert activities == {"uploadText", "downloadText"} + upload_calls = [ + c + for c in mock_check.call_args_list + if c.kwargs["activity"] == "uploadText" + ] + assert len(upload_calls) == 1 + upload_text = upload_calls[0].kwargs["text"] + assert "prompt body" in upload_text + assert "system instructions" in upload_text + + +# --------------------------------------------------------------- +# Logging hook user resolution +# --------------------------------------------------------------- + + +class TestLoggingResolveUserId: + def test_logging_prefers_user_api_key_user_id_in_metadata(self): + guardrail = _make_guardrail() + kwargs = { + "litellm_params": { + "metadata": { + "user_api_key_user_id": "trusted-from-proxy", + "user_id": "metadata-spoof", + } + } + } + assert ( + guardrail._resolve_user_id_from_logging_kwargs(kwargs) + == "trusted-from-proxy" + ) + + def test_logging_ignores_caller_supplied_user_id_field(self): + """Caller-controlled ``metadata[user_id_field]`` must not drive Purview audit attribution.""" + guardrail = _make_guardrail() + kwargs = {"litellm_params": {"metadata": {"user_id": "only-metadata-user"}}} + assert guardrail._resolve_user_id_from_logging_kwargs(kwargs) is None + + def test_logging_kwargs_level_user_api_key_user_id(self): + """Top-level ``kwargs["user_api_key_user_id"]`` is also a proxy-injected source.""" + guardrail = _make_guardrail() + kwargs = { + "user_api_key_user_id": "from-top-level", + "litellm_params": {"metadata": {}}, + } + assert ( + guardrail._resolve_user_id_from_logging_kwargs(kwargs) == "from-top-level" + ) + + def test_logging_returns_none_when_no_trusted_identity(self): + guardrail = _make_guardrail() + kwargs = {"litellm_params": {"metadata": {}}} + assert guardrail._resolve_user_id_from_logging_kwargs(kwargs) is None + + +# --------------------------------------------------------------- +# _check_content — integration-level +# --------------------------------------------------------------- + + +class TestCheckContent: + @pytest.mark.asyncio + async def test_check_content_allow(self): + guardrail = _make_guardrail() + + with ( + patch.object( + guardrail, + "_compute_protection_scopes", + new_callable=AsyncMock, + return_value=("etag-1", {}), + ), + patch.object( + guardrail, + "_process_content", + new_callable=AsyncMock, + return_value={ + "protectionScopeState": "unchanged", + "policyActions": [], + }, + ), + ): + result = await guardrail._check_content( + user_id="user-1", + text="Hello world", + activity="uploadText", + request_data={}, + block_on_violation=True, + ) + + assert result["policyActions"] == [] + + @pytest.mark.asyncio + async def test_check_content_block(self): + guardrail = _make_guardrail() + + with ( + patch.object( + guardrail, + "_compute_protection_scopes", + new_callable=AsyncMock, + return_value=("etag-1", {}), + ), + patch.object( + guardrail, + "_process_content", + new_callable=AsyncMock, + return_value={ + "protectionScopeState": "unchanged", + "policyActions": [ + { + "@odata.type": "#microsoft.graph.restrictAccessAction", + "action": "restrictAccess", + "restrictionAction": "block", + } + ], + }, + ), + ): + with pytest.raises(HTTPException) as exc_info: + await guardrail._check_content( + user_id="user-1", + text="SSN: 123-45-6789", + activity="uploadText", + request_data={}, + block_on_violation=True, + ) + + assert exc_info.value.status_code == 400 + assert "blocked by policy" in str(exc_info.value.detail) + + @pytest.mark.asyncio + async def test_check_content_logging_only_no_block(self): + """In logging_only mode, violations should NOT raise.""" + guardrail = _make_guardrail() + + with ( + patch.object( + guardrail, + "_compute_protection_scopes", + new_callable=AsyncMock, + return_value=("etag-1", {}), + ), + patch.object( + guardrail, + "_process_content", + new_callable=AsyncMock, + return_value={ + "protectionScopeState": "unchanged", + "policyActions": [ + { + "@odata.type": "#microsoft.graph.restrictAccessAction", + "action": "restrictAccess", + "restrictionAction": "block", + } + ], + }, + ), + ): + # Should NOT raise even though violation detected + result = await guardrail._check_content( + user_id="user-1", + text="SSN: 123-45-6789", + activity="uploadText", + request_data={}, + block_on_violation=False, + ) + + assert len(result["policyActions"]) == 1 + + +# --------------------------------------------------------------- +# Token caching +# --------------------------------------------------------------- + + +class TestTokenCaching: + @pytest.mark.asyncio + async def test_token_cached(self): + guardrail = _make_guardrail() + + with patch.object( + guardrail.async_handler, "post", return_value=_mock_token_response() + ) as mock_post: + token1 = await guardrail._get_access_token() + token2 = await guardrail._get_access_token() + + assert token1 == "mock-access-token" + assert token2 == "mock-access-token" + # Should only call the token endpoint once (cached) + assert mock_post.call_count == 1 + + @pytest.mark.asyncio + async def test_token_refreshed_on_expiry(self): + guardrail = _make_guardrail() + + with patch.object( + guardrail.async_handler, "post", return_value=_mock_token_response() + ) as mock_post: + await guardrail._get_access_token() + + # Expire the token + guardrail._token_cache = ("old-token", time.time() - 10) + + await guardrail._get_access_token() + + # Should have called token endpoint twice + assert mock_post.call_count == 2 + + @pytest.mark.asyncio + async def test_token_http_error_propagates(self): + """OAuth2 4xx/5xx responses must surface as HTTPStatusError, not KeyError.""" + guardrail = _make_guardrail() + + error_resp = Mock() + error_resp.json.return_value = { + "error": "invalid_client", + "error_description": "client secret is wrong", + } + error_resp.raise_for_status = Mock( + side_effect=httpx.HTTPStatusError( + "401 Unauthorized", + request=httpx.Request("POST", "https://login.microsoftonline.com/"), + response=httpx.Response(401), + ) + ) + + with patch.object(guardrail.async_handler, "post", return_value=error_resp): + with pytest.raises(httpx.HTTPStatusError): + await guardrail._get_access_token() + + # Failure must not poison the cache. + assert guardrail._token_cache is None + + +# --------------------------------------------------------------- +# Graph POST HTTP error propagation +# --------------------------------------------------------------- + + +class TestGraphPostHttpError: + @pytest.mark.asyncio + async def test_graph_post_http_error_propagates(self): + """Non-2xx Graph API responses must raise rather than return error JSON.""" + guardrail = _make_guardrail() + + error_resp = Mock() + error_resp.json.return_value = { + "error": {"code": "Forbidden", "message": "no access"} + } + error_resp.headers = {} + error_resp.raise_for_status = Mock( + side_effect=httpx.HTTPStatusError( + "403 Forbidden", + request=httpx.Request("POST", "https://graph.microsoft.com/"), + response=httpx.Response(403), + ) + ) + + with ( + patch.object( + guardrail, "_get_access_token", new_callable=AsyncMock + ) as mock_token, + patch.object(guardrail.async_handler, "post", return_value=error_resp), + ): + mock_token.return_value = "mock-token" + + with pytest.raises(httpx.HTTPStatusError): + await guardrail._graph_post( + "https://graph.microsoft.com/v1.0/users/u/example", + {"foo": "bar"}, + ) + + @pytest.mark.asyncio + async def test_compute_protection_scopes_http_error_propagates(self): + """A Graph error on protectionScopes/compute must not be cached as success.""" + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_graph_post", new_callable=AsyncMock + ) as mock_post: + mock_post.side_effect = httpx.HTTPStatusError( + "429 Too Many Requests", + request=httpx.Request("POST", "https://graph.microsoft.com/"), + response=httpx.Response(429), + ) + + with pytest.raises(httpx.HTTPStatusError): + await guardrail._compute_protection_scopes("user-err") + + # The failed compute must not populate the scope cache. + assert "user-err" not in guardrail._scope_cache + + +# --------------------------------------------------------------- +# Protection scope caching +# --------------------------------------------------------------- + + +class TestScopeCaching: + @pytest.mark.asyncio + async def test_scope_cached(self): + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_graph_post", new_callable=AsyncMock + ) as mock_post: + mock_post.return_value = ( + { + "value": [ + {"activities": "uploadText", "executionMode": "evaluateInline"} + ] + }, + {"ETag": "scope-etag"}, + ) + + etag1, _ = await guardrail._compute_protection_scopes("user-1") + etag2, _ = await guardrail._compute_protection_scopes("user-1") + + assert etag1 == "scope-etag" + assert etag2 == "scope-etag" + assert mock_post.call_count == 1 + + @pytest.mark.asyncio + async def test_scope_cache_lru_keeps_hot_user_on_eviction(self): + """Frequently accessed users should not be evicted before cold entries.""" + guardrail = _make_guardrail() + guardrail._scope_cache_maxsize = 3 + + scope_payload = ( + { + "value": [ + {"activities": "uploadText", "executionMode": "evaluateInline"} + ] + }, + {"ETag": "scope-etag"}, + ) + + with patch.object( + guardrail, "_graph_post", new_callable=AsyncMock + ) as mock_post: + mock_post.return_value = scope_payload + + await guardrail._compute_protection_scopes("user-a") + await guardrail._compute_protection_scopes("user-b") + await guardrail._compute_protection_scopes("user-c") + assert mock_post.call_count == 3 + + await guardrail._compute_protection_scopes("user-a") + assert mock_post.call_count == 3 + + await guardrail._compute_protection_scopes("user-d") + assert mock_post.call_count == 4 + + await guardrail._compute_protection_scopes("user-a") + assert mock_post.call_count == 4 + assert "user-a" in guardrail._scope_cache + assert "user-b" not in guardrail._scope_cache + + @pytest.mark.asyncio + async def test_scope_cache_refresh_moves_to_end_of_lru(self): + """Refreshing a stale entry must move it to the MRU end of the OrderedDict. + + Before the fix, OrderedDict.__setitem__ preserved the original insertion + position for existing keys, causing the just-refreshed entry to be the + next candidate for LRU eviction. + """ + guardrail = _make_guardrail() + guardrail._scope_cache_maxsize = 2 + + scope_payload = ( + {"value": []}, + {"ETag": "scope-etag"}, + ) + + with patch.object( + guardrail, "_graph_post", new_callable=AsyncMock + ) as mock_post: + mock_post.return_value = scope_payload + + # Populate cache: user-a (older), user-b (newer) + await guardrail._compute_protection_scopes("user-a") + await guardrail._compute_protection_scopes("user-b") + assert mock_post.call_count == 2 + + # Expire user-a's entry so it is re-fetched on the next access. + old_etag, old_scope, _ = guardrail._scope_cache["user-a"] + guardrail._scope_cache["user-a"] = (old_etag, old_scope, 0.0) + + # Re-fetch user-a — should move it to the MRU end. + await guardrail._compute_protection_scopes("user-a") + assert mock_post.call_count == 3 + + # Adding a third user must evict user-b (the true LRU), not user-a. + await guardrail._compute_protection_scopes("user-c") + assert mock_post.call_count == 4 + + assert "user-a" in guardrail._scope_cache, "user-a was wrongly evicted" + assert ( + "user-b" not in guardrail._scope_cache + ), "user-b should have been evicted" + assert "user-c" in guardrail._scope_cache + + @pytest.mark.asyncio + async def test_scope_invalidated_on_modified(self): + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_graph_post", new_callable=AsyncMock + ) as mock_post: + # First call: compute scopes + mock_post.return_value = ( + {"value": []}, + {"ETag": "etag-1"}, + ) + await guardrail._compute_protection_scopes("user-1") + + # processContent returns modified + mock_post.return_value = ( + {"protectionScopeState": "modified", "policyActions": []}, + {}, + ) + await guardrail._process_content("user-1", "text", "uploadText", "etag-1") + + # Scope cache should be invalidated + assert "user-1" not in guardrail._scope_cache + + +# --------------------------------------------------------------- +# get_prompt_text_for_dlp — message separator +# --------------------------------------------------------------- + + +class TestGetPromptTextForDlp: + def test_single_message_no_extra_separator(self): + """A single message is returned as-is (no leading/trailing separator).""" + guardrail = _make_guardrail() + result = guardrail.get_prompt_text_for_dlp( + [{"role": "user", "content": "Hello"}] + ) + assert result == "Hello" + + def test_messages_separated_by_double_newline(self): + """Adjacent messages must NOT be concatenated without a separator. + + Before the fix, "end of msg1" + "start of msg2" became + "end of msg1start of msg2", mangling DLP pattern detection. + """ + guardrail = _make_guardrail() + result = guardrail.get_prompt_text_for_dlp( + [ + {"role": "system", "content": "end of msg1"}, + {"role": "user", "content": "start of msg2"}, + ] + ) + assert result is not None + assert "end of msg1" in result + assert "start of msg2" in result + # Separator must be present between messages + assert "end of msg1start of msg2" not in result + assert "end of msg1\n\nstart of msg2" in result + + def test_empty_messages_returns_none(self): + guardrail = _make_guardrail() + assert guardrail.get_prompt_text_for_dlp([]) is None + + def test_whitespace_only_messages_skipped(self): + guardrail = _make_guardrail() + result = guardrail.get_prompt_text_for_dlp( + [ + {"role": "system", "content": " "}, + {"role": "user", "content": "real content"}, + ] + ) + assert result == "real content" + + def test_multi_role_conversation_preserves_all_content(self): + guardrail = _make_guardrail() + result = guardrail.get_prompt_text_for_dlp( + [ + {"role": "system", "content": "SYSTEM"}, + {"role": "user", "content": "USER1"}, + {"role": "assistant", "content": "ASSISTANT"}, + {"role": "user", "content": "USER2"}, + ] + ) + assert result is not None + for token in ("SYSTEM", "USER1", "ASSISTANT", "USER2"): + assert token in result + + +# --------------------------------------------------------------- +# logging_hook — non-blocking fire-and-forget +# --------------------------------------------------------------- + + +class TestLoggingHookNonBlocking: + @pytest.mark.asyncio + async def test_logging_hook_does_not_block_running_loop(self): + """logging_hook must return immediately without blocking the event loop. + + Before the fix, logging_hook called future.result() which blocked the + event loop thread for the full round-trip of the two Graph API calls. + """ + guardrail = _make_guardrail() + call_count = 0 + + async def slow_async_hook(**_kwargs): + nonlocal call_count + await asyncio.sleep(0.05) + call_count += 1 + return _kwargs.get("kwargs", {}), _kwargs.get("result") + + with patch.object(guardrail, "async_logging_hook", side_effect=slow_async_hook): + # Call logging_hook from within a running event loop + result = guardrail.logging_hook( + kwargs={"messages": [{"role": "user", "content": "test"}]}, + result=None, + call_type="completion", + ) + + # Must return (kwargs, result) unchanged without waiting for async work + assert result[0]["messages"][0]["content"] == "test" + assert result[1] is None + + def test_logging_hook_returns_original_kwargs_and_result(self): + """Return value must be the original (kwargs, result) tuple unchanged.""" + guardrail = _make_guardrail() + kwargs = {"messages": [{"role": "user", "content": "hello"}]} + result_obj = {"some": "result"} + + with patch.object( + guardrail, + "async_logging_hook", + new_callable=AsyncMock, + return_value=(kwargs, result_obj), + ): + out = guardrail.logging_hook( + kwargs=kwargs, + result=result_obj, + call_type="completion", + ) + + assert out == (kwargs, result_obj) + + +# --------------------------------------------------------------- +# Initializer validation +# --------------------------------------------------------------- + + +class TestInitializerValidation: + def test_missing_tenant_id(self): + from litellm.proxy.guardrails.guardrail_hooks.microsoft_purview import ( + initialize_guardrail, + ) + + litellm_params = Mock( + spec=[ + "tenant_id", + "client_id", + "client_secret", + "purview_app_name", + "user_id_field", + "api_key", + "mode", + "default_on", + ] + ) + litellm_params.tenant_id = None + litellm_params.client_id = None + litellm_params.client_secret = None + litellm_params.api_key = "secret" + litellm_params.mode = "pre_call" + + with pytest.raises(ValueError, match="tenant_id is required"): + initialize_guardrail(litellm_params, {"guardrail_name": "test"}) + + def test_missing_client_id(self): + from litellm.proxy.guardrails.guardrail_hooks.microsoft_purview import ( + initialize_guardrail, + ) + + litellm_params = Mock( + spec=[ + "tenant_id", + "client_id", + "client_secret", + "purview_app_name", + "user_id_field", + "api_key", + "mode", + "default_on", + ] + ) + litellm_params.tenant_id = "test-tenant" + litellm_params.client_id = None + litellm_params.client_secret = None + litellm_params.api_key = "secret" + litellm_params.mode = "pre_call" + + with pytest.raises(ValueError, match="client_id is required"): + initialize_guardrail(litellm_params, {"guardrail_name": "test"}) + + def test_missing_client_secret(self): + from litellm.proxy.guardrails.guardrail_hooks.microsoft_purview import ( + initialize_guardrail, + ) + + litellm_params = Mock( + spec=[ + "tenant_id", + "client_id", + "client_secret", + "purview_app_name", + "user_id_field", + "api_key", + "mode", + "default_on", + ] + ) + litellm_params.tenant_id = "test-tenant" + litellm_params.client_id = "test-client" + litellm_params.client_secret = None + litellm_params.api_key = None + litellm_params.mode = "pre_call" + + with pytest.raises(ValueError, match="client_secret"): + initialize_guardrail(litellm_params, {"guardrail_name": "test"}) + + +# --------------------------------------------------------------- +# _check_content — API error handling with block_on_violation=False +# --------------------------------------------------------------- + + +class TestCheckContentApiErrorHandling: + @pytest.mark.asyncio + async def test_api_error_reraises_when_block_on_violation_true(self): + """API/network errors must surface as HTTPException(400) when block_on_violation=True.""" + guardrail = _make_guardrail() + + with patch.object( + guardrail, + "_compute_protection_scopes", + new_callable=AsyncMock, + side_effect=RuntimeError("network failure"), + ): + with pytest.raises(HTTPException) as exc_info: + await guardrail._check_content( + user_id="user-1", + text="some content", + activity="uploadText", + request_data={}, + block_on_violation=True, + ) + + assert exc_info.value.status_code == 400 + assert isinstance(exc_info.value.detail, dict) + assert "upstream policy evaluation failed" in exc_info.value.detail.get( + "error", "" + ) + assert "network failure" in exc_info.value.detail.get("exception", "") + assert isinstance(exc_info.value.__cause__, RuntimeError) + + @pytest.mark.asyncio + async def test_http_exception_passes_through_unchanged(self): + """HTTPException from upstream layers must propagate as-is (not wrapped).""" + guardrail = _make_guardrail() + inner = HTTPException(status_code=403, detail="forbidden") + + with patch.object( + guardrail, + "_compute_protection_scopes", + new_callable=AsyncMock, + side_effect=inner, + ): + with pytest.raises(HTTPException) as exc_info: + await guardrail._check_content( + user_id="user-1", + text="some content", + activity="uploadText", + request_data={}, + block_on_violation=True, + ) + + assert exc_info.value is inner + + @pytest.mark.asyncio + async def test_api_error_not_reraised_when_block_on_violation_false(self): + """API/network errors must be swallowed (logged only) when block_on_violation=False.""" + guardrail = _make_guardrail() + + with patch.object( + guardrail, + "_compute_protection_scopes", + new_callable=AsyncMock, + side_effect=RuntimeError("network failure"), + ): + # Must NOT raise — should return empty dict + result = await guardrail._check_content( + user_id="user-1", + text="some content", + activity="uploadText", + request_data={}, + block_on_violation=False, + ) + + assert isinstance(result, dict) + + @pytest.mark.asyncio + async def test_process_content_error_not_reraised_when_block_on_violation_false( + self, + ): + """Errors from _process_content itself must also be suppressed in logging-only mode.""" + guardrail = _make_guardrail() + + with ( + patch.object( + guardrail, + "_compute_protection_scopes", + new_callable=AsyncMock, + return_value=("etag-1", {}), + ), + patch.object( + guardrail, + "_process_content", + new_callable=AsyncMock, + side_effect=ConnectionError("timeout"), + ), + ): + result = await guardrail._check_content( + user_id="user-1", + text="some content", + activity="uploadText", + request_data={}, + block_on_violation=False, + ) + + assert isinstance(result, dict) + + @pytest.mark.asyncio + async def test_http_status_error_preserves_upstream_status_code(self): + """Upstream Graph 429 must surface as 429 with Retry-After (not a generic 400).""" + guardrail = _make_guardrail() + upstream_resp = httpx.Response( + status_code=429, + headers={"Retry-After": "30"}, + request=httpx.Request("POST", "https://graph.microsoft.com/v1.0/x"), + ) + upstream_err = httpx.HTTPStatusError( + "rate limited", request=upstream_resp.request, response=upstream_resp + ) + + with patch.object( + guardrail, + "_compute_protection_scopes", + new_callable=AsyncMock, + side_effect=upstream_err, + ): + with pytest.raises(HTTPException) as exc_info: + await guardrail._check_content( + user_id="user-1", + text="some content", + activity="uploadText", + request_data={}, + block_on_violation=True, + ) + + assert exc_info.value.status_code == 429 + assert exc_info.value.headers == {"Retry-After": "30"} + assert isinstance(exc_info.value.detail, dict) + assert exc_info.value.detail.get("upstream_status") == 429 + assert isinstance(exc_info.value.__cause__, httpx.HTTPStatusError) + + @pytest.mark.asyncio + async def test_http_status_error_401_maps_to_502(self): + """Upstream 401/403 (proxy creds problem) should be exposed as 502, not 401/403.""" + guardrail = _make_guardrail() + upstream_resp = httpx.Response( + status_code=401, + request=httpx.Request("POST", "https://graph.microsoft.com/v1.0/x"), + ) + upstream_err = httpx.HTTPStatusError( + "unauthorized", request=upstream_resp.request, response=upstream_resp + ) + + with patch.object( + guardrail, + "_compute_protection_scopes", + new_callable=AsyncMock, + side_effect=upstream_err, + ): + with pytest.raises(HTTPException) as exc_info: + await guardrail._check_content( + user_id="user-1", + text="some content", + activity="uploadText", + request_data={}, + block_on_violation=True, + ) + + assert exc_info.value.status_code == 502 + assert exc_info.value.detail.get("upstream_status") == 401 + + +# --------------------------------------------------------------- +# async_logging_hook — independent prompt/response audit calls +# --------------------------------------------------------------- + + +class TestAsyncLoggingHookIndependence: + @pytest.mark.asyncio + async def test_response_audit_runs_even_if_prompt_audit_fails(self): + """A failure in the prompt audit must not prevent the response audit from running.""" + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _make_guardrail() + response = ModelResponse( + choices=[ + Choices( + index=0, + message=Message(content="response text", role="assistant"), + ) + ], + ) + + call_activities: list = [] + + async def fake_check_content(**kwargs): + activity = kwargs.get("activity") + if activity == "uploadText": + raise RuntimeError("simulated prompt API failure") + call_activities.append(activity) + return {"policyActions": []} + + with patch.object(guardrail, "_check_content", side_effect=fake_check_content): + await guardrail.async_logging_hook( + kwargs={ + "messages": [{"role": "user", "content": "prompt"}], + "litellm_params": { + "metadata": {"user_api_key_user_id": "user-123"} + }, + }, + result=response, + call_type="completion", + ) + + # The response audit must still have been attempted + assert "downloadText" in call_activities + + @pytest.mark.asyncio + async def test_prompt_audit_runs_even_if_response_audit_fails(self): + """A failure in the response audit must not affect the prompt audit result.""" + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _make_guardrail() + response = ModelResponse( + choices=[ + Choices( + index=0, + message=Message(content="response text", role="assistant"), + ) + ], + ) + + call_activities: list = [] + + async def fake_check_content(**kwargs): + activity = kwargs.get("activity") + if activity == "downloadText": + raise RuntimeError("simulated response API failure") + call_activities.append(activity) + return {"policyActions": []} + + with patch.object(guardrail, "_check_content", side_effect=fake_check_content): + await guardrail.async_logging_hook( + kwargs={ + "messages": [{"role": "user", "content": "prompt"}], + "litellm_params": { + "metadata": {"user_api_key_user_id": "user-123"} + }, + }, + result=response, + call_type="completion", + ) + + assert "uploadText" in call_activities + + @pytest.mark.asyncio + async def test_logging_hook_returns_original_when_both_audits_fail(self): + """async_logging_hook must always return (kwargs, result) even if both audits fail.""" + guardrail = _make_guardrail() + + with patch.object( + guardrail, + "_check_content", + new_callable=AsyncMock, + side_effect=RuntimeError("total failure"), + ): + kwargs = { + "messages": [{"role": "user", "content": "prompt"}], + "litellm_params": {"metadata": {"user_api_key_user_id": "user-123"}}, + } + result_obj = {"some": "result"} + out_kwargs, out_result = await guardrail.async_logging_hook( + kwargs=kwargs, + result=result_obj, + call_type="completion", + ) + + assert out_kwargs is kwargs + assert out_result is result_obj + + +# --------------------------------------------------------------- +# Tool-call argument extraction +# --------------------------------------------------------------- + + +class TestExtractToolCallArgs: + def test_dict_message_with_tool_calls(self): + msg = { + "role": "assistant", + "content": None, + "tool_calls": [ + {"function": {"arguments": '{"ssn": "123-45-6789"}'}}, + {"function": {"arguments": '{"card": "4111-1111-1111-1111"}'}}, + ], + } + args = MicrosoftPurviewDLPGuardrail._extract_tool_call_args_from_message(msg) + assert '{"ssn": "123-45-6789"}' in args + assert '{"card": "4111-1111-1111-1111"}' in args + + def test_dict_message_with_function_call(self): + msg = { + "role": "assistant", + "content": None, + "function_call": {"name": "lookup", "arguments": '{"query": "secret"}'}, + } + args = MicrosoftPurviewDLPGuardrail._extract_tool_call_args_from_message(msg) + assert '{"query": "secret"}' in args + + def test_object_message_with_tool_calls(self): + from litellm.types.utils import Message + + msg = Message( + role="assistant", + content=None, + tool_calls=[ + { + "id": "tc1", + "type": "function", + "function": {"name": "fn", "arguments": '{"x": 1}'}, + }, + ], + ) + args = MicrosoftPurviewDLPGuardrail._extract_tool_call_args_from_message(msg) + assert '{"x": 1}' in args + + def test_message_with_no_tool_calls(self): + msg = {"role": "user", "content": "hello"} + args = MicrosoftPurviewDLPGuardrail._extract_tool_call_args_from_message(msg) + assert args == [] + + def test_empty_arguments_skipped(self): + msg = { + "role": "assistant", + "content": None, + "tool_calls": [{"function": {"arguments": " "}}], + } + args = MicrosoftPurviewDLPGuardrail._extract_tool_call_args_from_message(msg) + assert args == [] + + +# --------------------------------------------------------------- +# Tool-call arguments included in DLP text extraction (prompt) +# --------------------------------------------------------------- + + +class TestGetPromptTextToolCalls: + def test_tool_call_args_included_in_prompt_scan(self): + """Sensitive data in tool_calls[].function.arguments must appear in DLP text.""" + guardrail = _make_guardrail() + messages = [ + {"role": "user", "content": "benign query"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "tc1", + "type": "function", + "function": { + "name": "lookup", + "arguments": '{"ssn": "123-45-6789"}', + }, + } + ], + }, + ] + text = guardrail.get_prompt_text_for_dlp(messages) + assert text is not None + assert "benign query" in text + assert '{"ssn": "123-45-6789"}' in text + + def test_function_call_args_included_in_prompt_scan(self): + """Legacy function_call.arguments must also appear in DLP text.""" + guardrail = _make_guardrail() + messages = [ + { + "role": "assistant", + "content": "Calling function", + "function_call": { + "name": "search", + "arguments": '{"credit_card": "4111-1111-1111-1111"}', + }, + } + ] + text = guardrail.get_prompt_text_for_dlp(messages) + assert text is not None + assert "Calling function" in text + assert '{"credit_card": "4111-1111-1111-1111"}' in text + + def test_content_only_message_unchanged(self): + """Messages without tool calls must still produce the same output.""" + guardrail = _make_guardrail() + messages = [ + {"role": "system", "content": "You are helpful."}, + {"role": "user", "content": "Tell me a joke."}, + ] + text = guardrail.get_prompt_text_for_dlp(messages) + assert text is not None + assert "You are helpful." in text + assert "Tell me a joke." in text + + @pytest.mark.asyncio + async def test_pre_call_hook_scans_tool_call_args(self): + """async_pre_call_hook must include tool_call arguments in the text sent to Purview.""" + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + await guardrail.async_pre_call_hook( + user_api_key_dict=__import__( + "litellm.proxy._types", fromlist=["UserAPIKeyAuth"] + ).UserAPIKeyAuth(api_key="test", user_id="user-123"), + cache=None, + data={ + "messages": [ + {"role": "user", "content": "benign"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "tc1", + "type": "function", + "function": { + "name": "do_thing", + "arguments": '{"password": "hunter2"}', + }, + } + ], + }, + ] + }, + call_type="completion", + ) + + mock_check.assert_called_once() + sent_text = mock_check.call_args.kwargs["text"] + assert '{"password": "hunter2"}' in sent_text + + +# --------------------------------------------------------------- +# Tool-call arguments included in DLP text extraction (response) +# --------------------------------------------------------------- + + +class TestCompletionResponseTextPartsToolCalls: + def test_response_tool_call_args_included(self): + """Model-generated tool_call arguments must appear in the DLP scan text.""" + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _make_guardrail() + response = ModelResponse( + choices=[ + Choices( + index=0, + message=Message( + role="assistant", + content=None, + tool_calls=[ + { + "id": "tc1", + "type": "function", + "function": { + "name": "exfil", + "arguments": '{"data": "secret-value"}', + }, + } + ], + ), + ) + ], + ) + parts = guardrail._completion_response_text_parts(response) + assert any("secret-value" in p for p in parts) + + def test_response_with_content_and_tool_calls(self): + """Both message content and tool_call arguments must be included.""" + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _make_guardrail() + response = ModelResponse( + choices=[ + Choices( + index=0, + message=Message( + role="assistant", + content="Here is the result", + tool_calls=[ + { + "id": "tc2", + "type": "function", + "function": { + "name": "fn", + "arguments": '{"ssn": "123-45-6789"}', + }, + } + ], + ), + ) + ], + ) + parts = guardrail._completion_response_text_parts(response) + combined = " ".join(parts) + assert "Here is the result" in combined + assert '{"ssn": "123-45-6789"}' in combined + + @pytest.mark.asyncio + async def test_post_call_hook_scans_response_tool_call_args(self): + """async_post_call_success_hook must send tool_call arguments to Purview.""" + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _make_guardrail() + response = ModelResponse( + choices=[ + Choices( + index=0, + message=Message( + role="assistant", + content=None, + tool_calls=[ + { + "id": "tc3", + "type": "function", + "function": { + "name": "retrieve", + "arguments": '{"credit_card": "4111-1111-1111-1111"}', + }, + } + ], + ), + ) + ], + ) + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + await guardrail.async_post_call_success_hook( + data={}, + user_api_key_dict=__import__( + "litellm.proxy._types", fromlist=["UserAPIKeyAuth"] + ).UserAPIKeyAuth(api_key="test", user_id="user-123"), + response=response, + ) + + mock_check.assert_called_once() + sent_text = mock_check.call_args.kwargs["text"] + assert '{"credit_card": "4111-1111-1111-1111"}' in sent_text + + def test_responses_api_function_call_args_included(self): + """Function-call arguments in ``ResponsesAPIResponse.output`` must be DLP-scanned.""" + from litellm.types.llms.openai import ResponsesAPIResponse + + guardrail = _make_guardrail() + response = ResponsesAPIResponse( + id="resp-tc-1", + created_at=0, + output=[ + { + "type": "message", + "id": "msg-tc-1", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "calling tool"}], + }, + { + "type": "function_call", + "id": "fc-1", + "call_id": "call-1", + "name": "exfil", + "arguments": '{"ssn": "123-45-6789"}', + }, + ], + ) + parts = guardrail._completion_response_text_parts(response) + combined = " ".join(parts) + assert "calling tool" in combined + assert '{"ssn": "123-45-6789"}' in combined + + def test_responses_api_function_call_args_only(self): + """Function-call args must be scanned even when no ``output_text`` blocks exist.""" + from litellm.types.llms.openai import ResponsesAPIResponse + + guardrail = _make_guardrail() + response = ResponsesAPIResponse( + id="resp-tc-2", + created_at=0, + output=[ + { + "type": "function_call", + "id": "fc-2", + "call_id": "call-2", + "name": "exfil", + "arguments": '{"secret": "hunter2"}', + } + ], + ) + parts = guardrail._completion_response_text_parts(response) + assert any('{"secret": "hunter2"}' in p for p in parts) + + +# --------------------------------------------------------------- +# Graph user id path encoding +# --------------------------------------------------------------- + + +class TestGraphUserIdEncoding: + def test_encode_graph_user_id_percent_encodes_special_chars(self): + from urllib.parse import quote + + raw = "user/with%special" + encoded = PurviewGuardrailBase._encode_graph_user_id(raw) + assert encoded == quote(raw, safe="") + assert "/" not in encoded + + @pytest.mark.asyncio + async def test_compute_protection_scopes_uses_encoded_path(self): + guardrail = _make_guardrail() + guardrail._scope_cache.clear() + + mock_resp = _mock_scope_response() + + async def _capture_post(url, **kwargs): + assert "/users/" in url + assert "user%2Fwith%25special" in url + return mock_resp + + guardrail.async_handler.post = AsyncMock(side_effect=_capture_post) + + with patch.object( + guardrail, "_get_access_token", new_callable=AsyncMock + ) as mock_token: + mock_token.return_value = "tok" + await guardrail._compute_protection_scopes("user/with%special") + + guardrail.async_handler.post.assert_called_once() + + +# --------------------------------------------------------------- +# _resolve_trusted_user_id +# --------------------------------------------------------------- + + +class TestResolveTrustedUserId: + def test_trusted_user_id_from_api_key_dict(self): + guardrail = _make_guardrail() + auth = UserAPIKeyAuth(api_key="test", user_id="auth-user-111") + assert guardrail._resolve_trusted_user_id({}, auth) == "auth-user-111" + + def test_end_user_id_not_trusted_for_blocking(self): + """end_user_id is request-derived; must not be used for blocking DLP.""" + guardrail = _make_guardrail() + auth = UserAPIKeyAuth(api_key="test", end_user_id="end-user-222") + assert guardrail._resolve_trusted_user_id({}, auth) is None + + def test_metadata_user_api_key_user_id_not_trusted_without_auth(self): + """Metadata user_api_key_user_id is not trusted when the key has no user_id.""" + guardrail = _make_guardrail() + auth = UserAPIKeyAuth(api_key="test") + data = {"metadata": {"user_api_key_user_id": "proxy-user-333"}} + assert guardrail._resolve_trusted_user_id(data, auth) is None + + def test_trusted_user_id_returns_none_for_caller_supplied_only(self): + """Caller-supplied metadata must NOT be returned by _resolve_trusted_user_id.""" + guardrail = _make_guardrail() + auth = UserAPIKeyAuth(api_key="test") + data = {"metadata": {"user_id": "caller-supplied-444"}} + assert guardrail._resolve_trusted_user_id(data, auth) is None + + def test_trusted_prefers_key_user_id_over_end_user_id(self): + guardrail = _make_guardrail() + auth = UserAPIKeyAuth( + api_key="test", user_id="key-owner", end_user_id="end-user" + ) + assert guardrail._resolve_trusted_user_id({}, auth) == "key-owner" + + +# --------------------------------------------------------------- +# _resolve_user_id_from_logging_kwargs — caller-influenceable identity rejected +# --------------------------------------------------------------- + + +class TestLoggingRejectsCallerInfluenceableIdentity: + """``end_user_id`` is derived from caller-controllable request fields + (``user``, ``metadata.user_id``, ``safety_identifier``) so it must not + drive Purview audit attribution either. + """ + + def test_end_user_id_in_metadata_is_ignored(self): + guardrail = _make_guardrail() + kwargs = { + "litellm_params": { + "metadata": { + "user_api_key_end_user_id": "end-user-from-metadata", + } + } + } + assert guardrail._resolve_user_id_from_logging_kwargs(kwargs) is None + + def test_end_user_id_at_top_level_kwargs_is_ignored(self): + guardrail = _make_guardrail() + kwargs = { + "user_api_key_end_user_id": "end-user-from-kwargs", + "litellm_params": {"metadata": {}}, + } + assert guardrail._resolve_user_id_from_logging_kwargs(kwargs) is None + + +# --------------------------------------------------------------- +# _resolve_user_id_for_blocking — security warning path +# --------------------------------------------------------------- + + +class TestResolveUserIdForBlocking: + def test_trusted_id_returned_without_warning(self, caplog): + import logging + + guardrail = _make_guardrail() + auth = UserAPIKeyAuth(api_key="test", user_id="trusted-111") + with caplog.at_level(logging.WARNING): + result = guardrail._resolve_user_id_for_blocking({}, auth) + assert result == "trusted-111" + assert "SECURITY" not in caplog.text + + def test_caller_supplied_id_raises_http_exception(self): + guardrail = _make_guardrail() + auth = UserAPIKeyAuth(api_key="test") + data = {"metadata": {"user_id": "caller-supplied-999"}} + with pytest.raises(HTTPException) as exc_info: + guardrail._resolve_user_id_for_blocking(data, auth) + assert exc_info.value.status_code == 400 + assert "proxy-authenticated" in str(exc_info.value.detail) + + def test_no_id_raises_http_exception(self): + guardrail = _make_guardrail() + auth = UserAPIKeyAuth(api_key="test") + with pytest.raises(HTTPException) as exc_info: + guardrail._resolve_user_id_for_blocking({}, auth) + assert exc_info.value.status_code == 400 + assert "bind user_id" in str(exc_info.value.detail) + + def test_end_user_id_only_raises_for_blocking(self): + """Request-derived end_user_id cannot drive blocking Purview checks.""" + guardrail = _make_guardrail() + auth = UserAPIKeyAuth(api_key="test", end_user_id="caller-end-user") + with pytest.raises(HTTPException) as exc_info: + guardrail._resolve_user_id_for_blocking({}, auth) + assert exc_info.value.status_code == 400 + assert "proxy-authenticated" in str(exc_info.value.detail) + + +# --------------------------------------------------------------- +# Token-id prompt handling in pre_call blocking mode +# --------------------------------------------------------------- + + +class TestTokenIdPromptHandling: + @pytest.mark.asyncio + async def test_token_id_prompt_raises_in_blocking_mode(self): + """Pure token-id prompts must be rejected in blocking pre_call mode.""" + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + with pytest.raises(HTTPException) as exc_info: + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="u1"), + cache=None, + data={"prompt": [1, 2, 3, 100, 200]}, + call_type="text_completion", + ) + + mock_check.assert_not_called() + assert exc_info.value.status_code == 400 + assert "Token-id" in str(exc_info.value.detail) + + @pytest.mark.asyncio + async def test_missing_prompt_skips_without_warning(self, caplog): + """No prompt at all → silently skip (not a token-id bypass case).""" + import logging + + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + with caplog.at_level(logging.WARNING): + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="u1"), + cache=None, + data={}, + call_type="text_completion", + ) + + mock_check.assert_not_called() + assert "token-id" not in caplog.text.lower() + + @pytest.mark.asyncio + async def test_string_prompt_still_scanned(self): + """Normal string prompts must still be sent to Purview.""" + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="u1"), + cache=None, + data={"prompt": "sensitive text"}, + call_type="text_completion", + ) + + mock_check.assert_called_once() + assert mock_check.call_args.kwargs["text"] == "sensitive text" + + @pytest.mark.asyncio + @pytest.mark.parametrize("empty_prompt", ["", " ", "\n\t "]) + async def test_empty_or_whitespace_prompt_passes_through(self, empty_prompt): + """Empty/whitespace-only string prompts must not be flagged as token-id prompts.""" + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + data = {"prompt": empty_prompt} + result = await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="u1"), + cache=None, + data=data, + call_type="text_completion", + ) + + mock_check.assert_not_called() + assert result is data + assert result["prompt"] == empty_prompt + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "raw_prompt", + [ + [[1, 2, 3]], + [[1, 2], [3, 4]], + ["benign text", [99, 100]], + ], + ) + async def test_nested_token_id_prompt_raises_in_blocking_mode(self, raw_prompt): + """Nested/mixed token-id prompts must also be rejected in blocking pre_call mode.""" + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + with pytest.raises(HTTPException) as exc_info: + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="u1"), + cache=None, + data={"prompt": raw_prompt}, + call_type="text_completion", + ) + + mock_check.assert_not_called() + assert exc_info.value.status_code == 400 + assert "Token-id" in str(exc_info.value.detail) + + +class TestIsTokenIdPrompt: + @pytest.mark.parametrize( + "prompt,expected", + [ + ([1, 2, 3], True), + ([[1, 2, 3]], True), + ([[1, 2], [3, 4]], True), + (["hi", [1, 2]], True), + (["a", "b"], False), + ([], False), + ("hello", False), + (None, False), + ], + ) + def test_is_token_id_prompt(self, prompt, expected): + assert PurviewGuardrailBase.is_token_id_prompt(prompt) is expected + + +# --------------------------------------------------------------- +# Streaming iterator hook +# --------------------------------------------------------------- + + +class TestStreamingIteratorHook: + @pytest.mark.asyncio + async def test_streaming_clean_response_yields_all_chunks(self): + """Clean stream: all chunks must be re-yielded after DLP passes.""" + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _make_guardrail() + + assembled_response = ModelResponse( + choices=[ + Choices( + index=0, + message=Message(content="safe response", role="assistant"), + ) + ] + ) + + async def fake_response_stream(): + yield assembled_response + + with ( + patch("litellm.main.stream_chunk_builder", return_value=assembled_response), + patch( + "litellm.llms.base_llm.base_model_iterator.MockResponseIterator" + ) as mock_iterator_cls, + patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check, + ): + mock_check.return_value = {"policyActions": []} + + async def _iter_chunks(): + yield assembled_response + + mock_iterator_cls.return_value.__aiter__ = lambda s: _iter_chunks() + + chunks = [] + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + response=fake_response_stream(), + request_data={"metadata": {"user_id": "user-123"}}, + ): + chunks.append(chunk) + + mock_check.assert_called_once() + assert mock_check.call_args.kwargs["activity"] == "downloadText" + assert len(chunks) > 0 + + @pytest.mark.asyncio + async def test_streaming_violation_raises_before_any_chunk(self): + """A policy violation must raise HTTPException before yielding any chunk.""" + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _make_guardrail() + + assembled_response = ModelResponse( + choices=[ + Choices( + index=0, + message=Message( + content="SSN: 123-45-6789", + role="assistant", + ), + ) + ] + ) + + async def fake_response_stream(): + yield assembled_response + + with ( + patch("litellm.main.stream_chunk_builder", return_value=assembled_response), + patch.object( + guardrail, + "_check_content", + new_callable=AsyncMock, + side_effect=HTTPException( + status_code=400, + detail={ + "error": "Microsoft Purview DLP: Content blocked by policy" + }, + ), + ), + ): + chunks = [] + with pytest.raises(HTTPException) as exc_info: + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth( + api_key="test", user_id="user-123" + ), + response=fake_response_stream(), + request_data={"metadata": {"user_id": "user-123"}}, + ): + chunks.append(chunk) + + assert exc_info.value.status_code == 400 + assert len(chunks) == 0 # No chunks yielded before the block + + @pytest.mark.asyncio + async def test_streaming_no_user_id_raises_before_yield(self): + """No resolvable user_id → fail closed before any chunk is yielded.""" + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _make_guardrail() + + assembled_response = ModelResponse( + choices=[ + Choices( + index=0, + message=Message(content="some content", role="assistant"), + ) + ] + ) + + async def fake_response_stream(): + yield assembled_response + + with patch( + "litellm.main.stream_chunk_builder", return_value=assembled_response + ): + chunks = [] + with pytest.raises(HTTPException) as exc_info: + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), # no user_id + response=fake_response_stream(), + request_data={}, + ): + chunks.append(chunk) + + assert exc_info.value.status_code == 400 + assert len(chunks) == 0 + + @pytest.mark.asyncio + async def test_streaming_text_completion_scans_before_yield(self): + """Streamed /v1/completions must be DLP-scanned via TextCompletionResponse.""" + from litellm.types.utils import TextChoices, TextCompletionResponse + + guardrail = _make_guardrail() + + assembled_response = TextCompletionResponse( + model="gpt-3.5-turbo-instruct", + choices=[TextChoices(text="completion body", index=0)], + ) + + async def fake_response_stream(): + yield assembled_response + + with ( + patch("litellm.main.stream_chunk_builder", return_value=assembled_response), + patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check, + ): + mock_check.return_value = {"policyActions": []} + + chunks = [] + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + response=fake_response_stream(), + request_data={}, + ): + chunks.append(chunk) + + mock_check.assert_called_once() + assert mock_check.call_args.kwargs["text"] == "completion body" + assert len(chunks) > 0 + + @pytest.mark.asyncio + async def test_streaming_responses_api_scans_completed_event(self): + """Streamed Responses API: assembled ResponsesAPIResponse must be DLP-scanned.""" + from litellm.types.llms.openai import ( + ResponseCompletedEvent, + ResponseCreatedEvent, + ResponsesAPIResponse, + ResponsesAPIStreamEvents, + ) + + guardrail = _make_guardrail() + + completed_response = ResponsesAPIResponse( + id="resp-stream", + created_at=0, + output=[ + { + "type": "message", + "id": "msg-stream", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "streamed output"}], + } + ], + ) + created_event = ResponseCreatedEvent( + type=ResponsesAPIStreamEvents.RESPONSE_CREATED, + response=completed_response, + ) + completed_event = ResponseCompletedEvent( + type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + response=completed_response, + ) + + async def fake_response_stream(): + yield created_event + yield completed_event + + with ( + patch("litellm.main.stream_chunk_builder") as mock_stream_builder, + patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check, + ): + mock_check.return_value = {"policyActions": []} + + chunks = [] + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + response=fake_response_stream(), + request_data={}, + ): + chunks.append(chunk) + + mock_stream_builder.assert_not_called() + mock_check.assert_called_once() + assert mock_check.call_args.kwargs["activity"] == "downloadText" + assert mock_check.call_args.kwargs["text"] == "streamed output" + assert chunks == [created_event, completed_event] + + @pytest.mark.asyncio + async def test_streaming_responses_api_violation_blocks_before_yield(self): + """Responses API stream with a DLP violation must raise before any chunk is yielded.""" + from litellm.types.llms.openai import ( + ResponseCompletedEvent, + ResponsesAPIResponse, + ResponsesAPIStreamEvents, + ) + + guardrail = _make_guardrail() + + completed_response = ResponsesAPIResponse( + id="resp-stream-block", + created_at=0, + output=[ + { + "type": "message", + "id": "msg-stream-block", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "SSN: 123-45-6789"}], + } + ], + ) + completed_event = ResponseCompletedEvent( + type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + response=completed_response, + ) + + async def fake_response_stream(): + yield completed_event + + with patch.object( + guardrail, + "_check_content", + new_callable=AsyncMock, + side_effect=HTTPException( + status_code=400, + detail={"error": "Microsoft Purview DLP: Content blocked by policy"}, + ), + ): + chunks = [] + with pytest.raises(HTTPException) as exc_info: + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth( + api_key="test", user_id="user-123" + ), + response=fake_response_stream(), + request_data={}, + ): + chunks.append(chunk) + + assert exc_info.value.status_code == 400 + assert len(chunks) == 0 + + +# --------------------------------------------------------------- +# Auto-discovery registration +# --------------------------------------------------------------- + + +class TestRegistration: + def test_registry_contains_microsoft_purview(self): + from litellm.proxy.guardrails.guardrail_hooks.microsoft_purview import ( + guardrail_class_registry, + guardrail_initializer_registry, + ) + + assert "microsoft_purview" in guardrail_initializer_registry + assert "microsoft_purview" in guardrail_class_registry + assert ( + guardrail_class_registry["microsoft_purview"] + is MicrosoftPurviewDLPGuardrail + ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 23216542f35..c763e9c0e98 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -2438,6 +2438,7 @@ class TestCLIKeyRegenerationFlow: request=mock_request, key="cli-new-session-key-456", result=mock_result, + received_response=None, ) def test_get_redirect_url_does_not_include_existing_key_in_url(self): @@ -2496,6 +2497,11 @@ class TestCLIKeyRegenerationFlow: "user_id": "test-user-789", "user_role": "internal_user", "teams": ["team-a", "team-b", "team-c"], + "team_details": [ + {"team_id": "team-a", "team_alias": "Team A"}, + {"team_id": "team-b", "team_alias": "Team B"}, + {"team_id": "team-c", "team_alias": "Team C"}, + ], "models": ["gpt-4"], "user_email": "test@example.com", } @@ -2550,6 +2556,7 @@ class TestCLIKeyRegenerationFlow: mock_get_jwt.assert_called_once() jwt_call_args = mock_get_jwt.call_args assert jwt_call_args.kwargs["team_id"] == selected_team + assert jwt_call_args.kwargs["team_alias"] == "Team B" # Verify session was deleted after JWT generation mock_cache.delete_cache.assert_called_once() @@ -5552,6 +5559,289 @@ def test_generic_response_convertor_extra_attributes_missing_field(monkeypatch): assert result.extra_fields["another_missing"] is None +class TestCliSsoAttributionMetadata: + """CLI SSO allowlisted OIDC claim persistence and poll exposure.""" + + def test_parse_cli_sso_claim_map(self, monkeypatch): + from litellm.proxy.management_endpoints import ui_sso + + monkeypatch.setattr( + ui_sso, + "CLI_SSO_CLAIM_MAP", + "employment_type->metadata.acme_employment_type, org_info.department -> department", + ) + assert ui_sso._parse_cli_sso_claim_map() == [ + ("employment_type", "acme_employment_type"), + ("org_info.department", "department"), + ] + + def test_build_cli_sso_attribution_metadata_filters_non_scalars(self, monkeypatch): + from litellm.proxy.management_endpoints import ui_sso + from litellm.proxy.management_endpoints.types import CustomOpenID + + monkeypatch.setattr( + ui_sso, + "CLI_SSO_CLAIM_MAP", + "employment_type->acme_employment_type,access_token->should_drop,group->groups", + ) + + result = CustomOpenID( + id="user-1", + email="user@example.com", + display_name="User", + provider="generic", + team_ids=[], + extra_fields={ + "employment_type": "full_time", + "access_token": "eyJhbGciOiJIUzI1NiJ9.payload.signature", + "group": ["team-a", "team-b"], + }, + ) + + metadata = ui_sso.build_cli_sso_attribution_metadata(result=result) + assert metadata == {"acme_employment_type": "full_time"} + + def test_build_cli_sso_attribution_metadata_from_oidc_dict(self, monkeypatch): + from litellm.proxy.management_endpoints import ui_sso + + monkeypatch.setattr( + ui_sso, + "CLI_SSO_CLAIM_MAP", + "org_info.department->department", + ) + + metadata = ui_sso.build_cli_sso_attribution_metadata( + result={ + "sub": "user-1", + "email": "user@example.com", + "org_info": {"department": "Engineering"}, + } + ) + assert metadata == {"department": "Engineering"} + + @pytest.mark.asyncio + async def test_cli_sso_callback_passes_user_defined_values_for_new_users(self): + """First CLI SSO login must supply SSOUserDefinedValues so upsert can create the user.""" + from litellm.proxy._types import LiteLLM_UserTable + from litellm.proxy.management_endpoints import ui_sso + from litellm.proxy.management_endpoints.types import CustomOpenID + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://internal-proxy.local/" + session_key = "cli-session-new-user" + mock_user_info = LiteLLM_UserTable( + user_id="cli-test-user", + user_role="internal_user", + teams=[], + models=[], + ) + mock_sso_result = CustomOpenID( + id="cli-test-user", + email="cli-test@example.com", + display_name="cli-test-user", + provider="generic", + team_ids=[], + ) + mock_cache = MagicMock() + mock_cache.get_cache.return_value = { + "poll_secret_hash": "poll-secret-hash", + "user_code_hash": "user-code-hash", + "sso_complete": False, + "user_code_verified": False, + "session_data": None, + } + get_user_info_mock = AsyncMock(return_value=mock_user_info) + + with ( + patch( + "litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db", + get_user_info_mock, + ), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.user_custom_sso", None), + ): + await ui_sso.cli_sso_callback( + request=mock_request, + key=session_key, + result=mock_sso_result, + ) + + get_user_info_mock.assert_awaited_once() + assert get_user_info_mock.call_args.kwargs["user_defined_values"] is not None + assert ( + get_user_info_mock.call_args.kwargs["user_defined_values"]["user_id"] + == "cli-test-user" + ) + + @pytest.mark.asyncio + async def test_cli_sso_callback_rejects_restricted_sso_group(self): + """CLI SSO must enforce restricted_sso_group before upserting the user.""" + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints import ui_sso + from litellm.proxy.management_endpoints.types import CustomOpenID + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://internal-proxy.local/" + mock_cache = MagicMock() + mock_cache.get_cache.return_value = { + "poll_secret_hash": "poll-secret-hash", + "user_code_hash": "user-code-hash", + "sso_complete": False, + "user_code_verified": False, + "session_data": None, + } + mock_sso_result = CustomOpenID( + id="cli-test-user", + email="cli-test@example.com", + display_name="cli-test-user", + provider="generic", + team_ids=["other-group"], + ) + + with ( + patch( + "litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db", + new=AsyncMock(), + ) as get_user_info_mock, + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.user_custom_sso", None), + patch( + "litellm.proxy.proxy_server.general_settings", + { + "ui_access_mode": { + "type": "restricted_sso_group", + "restricted_sso_group": "required-group", + } + }, + ), + ): + with pytest.raises(ProxyException): + await ui_sso.cli_sso_callback( + request=mock_request, + key="cli-session-restricted", + result=mock_sso_result, + received_response={"groups": ["other-group"]}, + ) + + get_user_info_mock.assert_not_awaited() + + @pytest.mark.asyncio + async def test_cli_sso_callback_persists_attribution_metadata(self, monkeypatch): + from litellm.proxy._types import LiteLLM_UserTable + from litellm.proxy.management_endpoints import ui_sso + + monkeypatch.setattr( + ui_sso, + "CLI_SSO_CLAIM_MAP", + "employment_type->acme_employment_type", + ) + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://internal-proxy.local/" + session_key = "cli-session-4567890" + mock_user_info = LiteLLM_UserTable( + user_id="test-user-123", + user_role="internal_user", + teams=["team1"], + models=["gpt-4"], + ) + mock_sso_result = { + "user_email": "test@example.com", + "user_id": "test-user-123", + "employment_type": "contractor", + } + mock_cache = MagicMock() + mock_cache.get_cache.return_value = { + "poll_secret_hash": "poll-secret-hash", + "user_code_hash": "user-code-hash", + "sso_complete": False, + "user_code_verified": False, + "session_data": None, + } + mock_prisma = MagicMock() + mock_prisma.db.litellm_usertable.find_unique = AsyncMock( + return_value=MagicMock(metadata={"auth_provider": "generic"}) + ) + mock_prisma.db.litellm_usertable.update_many = AsyncMock() + + with ( + patch.dict( + os.environ, + { + "PROXY_BASE_URL": "https://test.example.com", + "SERVER_ROOT_PATH": "", + }, + ), + patch( + "litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db", + return_value=mock_user_info, + ), + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.user_custom_sso", None), + patch( + "litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page", + return_value="Success", + ), + ): + await ui_sso.cli_sso_callback( + request=mock_request, + key=session_key, + result=mock_sso_result, + ) + + flow_data = mock_cache.set_cache.call_args.kwargs["value"] + assert flow_data["session_data"]["attribution_metadata"] == { + "acme_employment_type": "contractor" + } + mock_prisma.db.litellm_usertable.update_many.assert_awaited_once() + update_data = mock_prisma.db.litellm_usertable.update_many.call_args.kwargs[ + "data" + ] + assert update_data["metadata"]["acme_employment_type"] == "contractor" + assert update_data["metadata"]["auth_provider"] == "generic" + + @pytest.mark.asyncio + async def test_cli_poll_key_returns_attribution_metadata(self, monkeypatch): + from litellm.proxy.management_endpoints.ui_sso import ( + _hash_cli_sso_secret, + cli_poll_key, + ) + + session_key = "cli-session-789123" + session_data = { + "user_id": "test-user-456", + "user_role": "internal_user", + "teams": ["team-a", "team-b"], + "models": ["gpt-4"], + "attribution_metadata": { + "acme_employment_type": "full_time", + "org": {"cost_center": "CC-42"}, + }, + } + mock_cache = MagicMock() + mock_cache.get_cache.return_value = { + "poll_secret_hash": _hash_cli_sso_secret("poll-secret"), + "sso_complete": True, + "user_code_verified": True, + "session_data": session_data, + } + + with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache): + result = await cli_poll_key( + key_id=session_key, + team_id=None, + x_litellm_cli_poll_secret="poll-secret", + ) + + assert result["attribution_metadata"] == { + "acme_employment_type": "full_time", + "org.cost_center": "CC-42", + } + + class TestValidateReturnTo: """Tests for SSOAuthenticationHandler._validate_return_to""" diff --git a/tests/test_litellm/test_check_licenses.py b/tests/test_litellm/test_check_licenses.py new file mode 100644 index 00000000000..4d72f185a25 --- /dev/null +++ b/tests/test_litellm/test_check_licenses.py @@ -0,0 +1,211 @@ +"""Tests for the dependency license checker at tests/code_coverage_tests/check_licenses.py. + +Focus: PEP 639 license metadata. Packages that adopt PEP 639 publish their +license as an SPDX expression in ``info.license_expression`` and often leave the +legacy ``info.license`` field null, so the checker must read the new field (and +fall back to trove classifiers) instead of reporting "Unknown license". + +PyPI HTTP responses are mocked — these tests never hit the network. +""" + +import os +import sys +from pathlib import Path + +_CODE_COVERAGE_DIR = os.path.join( + os.path.dirname(os.path.abspath(__file__)), "..", "code_coverage_tests" +) +sys.path.insert(0, _CODE_COVERAGE_DIR) + +import check_licenses # noqa: E402 + +_LICCHECK_INI = Path(_CODE_COVERAGE_DIR) / "liccheck.ini" + + +class _FakeResponse: + def __init__(self, payload): + self._payload = payload + + def raise_for_status(self): + return None + + def json(self): + return self._payload + + +def _make_checker(): + return check_licenses.LicenseChecker(config_file=_LICCHECK_INI) + + +def _patch_pypi(monkeypatch, info): + """Make PyPI return a JSON response with the given ``info`` block.""" + + def _fake_get(url, timeout=None): + return _FakeResponse({"info": info}) + + monkeypatch.setattr(check_licenses.requests, "get", _fake_get) + + +# -------------------------------------------------------------------------- +# get_package_license_from_pypi: license metadata resolution +# -------------------------------------------------------------------------- + + +def test_get_license_prefers_license_expression(monkeypatch): + """(a) PEP 639 packages publish the SPDX expression in license_expression.""" + _patch_pypi( + monkeypatch, + {"license_expression": "MIT", "license": None, "classifiers": []}, + ) + checker = _make_checker() + assert checker.get_package_license_from_pypi("black", "26.3.1") == "MIT" + + +def test_license_expression_wins_when_both_present(monkeypatch): + """license_expression takes precedence over the legacy license field.""" + _patch_pypi( + monkeypatch, + {"license_expression": "Apache-2.0", "license": "stale free text"}, + ) + checker = _make_checker() + assert checker.get_package_license_from_pypi("pkg", "1.0.0") == "Apache-2.0" + + +def test_get_license_falls_back_to_legacy_license(monkeypatch): + """(b) Pre-PEP-639 packages only set the legacy free-text license field.""" + _patch_pypi( + monkeypatch, + {"license_expression": None, "license": "MIT License", "classifiers": []}, + ) + checker = _make_checker() + assert checker.get_package_license_from_pypi("pkg", "1.0.0") == "MIT License" + + +def test_get_license_falls_back_to_classifiers(monkeypatch): + """(c) Some packages express the license only through trove classifiers.""" + _patch_pypi( + monkeypatch, + { + "license_expression": None, + "license": None, + "classifiers": [ + "Programming Language :: Python :: 3", + "License :: OSI Approved :: Apache Software License", + ], + }, + ) + checker = _make_checker() + assert ( + checker.get_package_license_from_pypi("pkg", "1.0.0") + == "Apache Software License" + ) + + +def test_get_license_returns_none_when_unset(monkeypatch): + """(d) With no license metadata at all the license stays unknown.""" + _patch_pypi( + monkeypatch, + {"license_expression": None, "license": None, "classifiers": []}, + ) + checker = _make_checker() + assert checker.get_package_license_from_pypi("pkg", "1.0.0") is None + + +def test_get_license_returns_none_on_request_failure(monkeypatch): + """Network/HTTP failures are swallowed and reported as unknown.""" + + def _boom(url, timeout=None): + raise RuntimeError("network down") + + monkeypatch.setattr(check_licenses.requests, "get", _boom) + checker = _make_checker() + assert checker.get_package_license_from_pypi("pkg", "1.0.0") is None + + +# -------------------------------------------------------------------------- +# is_license_acceptable: SPDX identifiers and compound expressions +# -------------------------------------------------------------------------- + + +def test_spdx_identifiers_are_authorized(): + """Plain SPDX identifiers match the legacy-spelled authorized list as-is.""" + checker = _make_checker() + for identifier in ("MIT", "Apache-2.0", "BSD-3-Clause"): + is_ok, reason = checker.is_license_acceptable(identifier) + assert is_ok is True, f"{identifier}: {reason}" + + +def test_spdx_compound_or_expression_is_authorized(): + checker = _make_checker() + is_ok, reason = checker.is_license_acceptable("MIT OR Apache-2.0") + assert is_ok is True, reason + + +def test_spdx_with_exception_in_compound_is_authorized(): + """The 'WITH ' suffix is stripped; the base license is checked.""" + checker = _make_checker() + is_ok, reason = checker.is_license_acceptable( + "Apache-2.0 WITH LLVM-exception OR MIT" + ) + assert is_ok is True, reason + + +def test_spdx_gpl3_is_rejected(): + """GPL-3.0 spellings must fail — they match no authorized license.""" + checker = _make_checker() + for expr in ("GPL-3.0-only", "GPL-3.0-or-later"): + is_ok, reason = checker.is_license_acceptable(expr) + assert is_ok is False, f"{expr} unexpectedly accepted: {reason}" + + +def test_spdx_compound_with_copyleft_component_is_rejected(): + """A permissive-OR-copyleft expression is conservatively rejected.""" + checker = _make_checker() + is_ok, _ = checker.is_license_acceptable("MIT OR GPL-3.0-only") + assert is_ok is False + + +def test_or_later_identifier_is_not_split_as_operator(): + """The lowercase '-or-later' inside an identifier is not the SPDX OR operator.""" + assert ( + check_licenses.LicenseChecker._split_spdx_expression("GPL-2.0-or-later") is None + ) + + +def test_free_text_license_is_not_treated_as_spdx(): + """Free-text license blobs fall back to whole-string substring matching.""" + free_text = "MIT License AND additional redistribution permissions" + assert check_licenses.LicenseChecker._split_spdx_expression(free_text) is None + checker = _make_checker() + assert checker.is_license_acceptable(free_text)[0] is True + + +def test_unknown_license_is_reported(): + checker = _make_checker() + is_ok, reason = checker.is_license_acceptable(None) + assert is_ok is False + assert reason == "Unknown license" + + +# -------------------------------------------------------------------------- +# check_package: end-to-end resolution + acceptability +# -------------------------------------------------------------------------- + + +def test_check_package_accepts_pep639_package(monkeypatch): + """A PEP 639 package whose license lives only in license_expression passes.""" + _patch_pypi( + monkeypatch, + {"license_expression": "MIT", "license": None, "classifiers": []}, + ) + checker = _make_checker() + assert checker.check_package("some-pep639-pkg", "1.0.0") is True + + +def test_check_package_rejects_package_without_license(monkeypatch): + _patch_pypi( + monkeypatch, + {"license_expression": None, "license": None, "classifiers": []}, + ) + checker = _make_checker() + assert checker.check_package("mystery-pkg", "1.0.0") is False diff --git a/tests/test_litellm/test_register_model_custom_pricing.py b/tests/test_litellm/test_register_model_custom_pricing.py index 1efd698fb64..719cb8eecd2 100644 --- a/tests/test_litellm/test_register_model_custom_pricing.py +++ b/tests/test_litellm/test_register_model_custom_pricing.py @@ -190,3 +190,164 @@ def test_build_custom_pricing_entry_time_based(): assert entry["litellm_provider"] == "openai" assert entry["input_cost_per_second"] == 0.01 assert entry["output_cost_per_second"] == 0.02 + + +def test_register_model_strips_none_litellm_provider(): + """``get_model_info`` returns ``litellm_provider: None`` for deployments + registered without a provider (e.g. ``Router.add_deployment`` flows). + ``register_model`` must not persist that None into ``model_cost``, + otherwise ``_check_provider_match`` will drop custom pricing on + subsequent cost lookups. + + Regression test for https://github.com/BerriAI/litellm/issues/28336. + """ + from litellm.utils import _check_provider_match + + model_key = "test-custom-pricing-no-provider-28336" + litellm.model_cost.pop(model_key, None) + + try: + litellm.register_model( + { + model_key: { + "input_cost_per_token": 0.001, + "output_cost_per_token": 0.002, + } + } + ) + + registered = litellm.model_cost.get(model_key) + assert registered is not None, f"{model_key} should be in model_cost" + # The key may be absent entirely, but if present it must not be None. + assert ( + "litellm_provider" not in registered + or registered["litellm_provider"] is not None + ) + # Downstream consumers must accept this entry for any provider, + # mirroring what the cost calculator does. + assert _check_provider_match(registered, "openai") is True + assert _check_provider_match(registered, "anthropic") is True + finally: + litellm.model_cost.pop(model_key, None) + + +def test_register_model_strips_none_litellm_provider_from_get_model_info(monkeypatch): + """Directly exercise the strip in ``register_model``. + + The companion test above hits the ``except Exception`` branch where + ``existing_model`` is an empty dict, so the ``pop`` is a no-op. This + test patches ``get_model_info`` to return the failure mode the strip + was added to handle, namely a populated dict whose ``litellm_provider`` + is ``None``. Without the strip, the merged entry in + ``litellm.model_cost`` would carry ``litellm_provider: None`` and + ``_check_provider_match`` would drop custom pricing. + + Regression test for https://github.com/BerriAI/litellm/issues/28336. + """ + from litellm import utils as litellm_utils + from litellm.utils import _check_provider_match + + model_key = "test-strip-none-provider-from-get-model-info-28336" + litellm.model_cost.pop(model_key, None) + + def _fake_get_model_info(model, *args, **kwargs): + assert model == model_key + return { + "key": model_key, + "litellm_provider": None, + "mode": "chat", + "max_tokens": 4096, + } + + # ``register_model`` calls ``get_model_info.cache_clear`` via + # ``_invalidate_model_cost_lowercase_map``, so the replacement must + # expose a no-op ``cache_clear`` attribute. + _fake_get_model_info.cache_clear = lambda: None + monkeypatch.setattr(litellm_utils, "get_model_info", _fake_get_model_info) + + try: + litellm.register_model( + { + model_key: { + "input_cost_per_token": 0.001, + "output_cost_per_token": 0.002, + } + } + ) + + registered = litellm.model_cost.get(model_key) + assert registered is not None, f"{model_key} should be in model_cost" + # The strip must have removed the None-valued provider that + # ``get_model_info`` returned. The key may be absent entirely, but + # it must never be present with value ``None``. + assert "litellm_provider" not in registered or ( + registered["litellm_provider"] is not None + ), ( + "register_model failed to strip litellm_provider=None returned " + f"by get_model_info, got {registered.get('litellm_provider')!r}" + ) + # Metadata from the patched ``get_model_info`` must still flow + # through, so we know the strip did not nuke the rest of the entry. + assert registered.get("mode") == "chat" + assert registered.get("max_tokens") == 4096 + # And custom pricing from the registration call must be preserved. + assert registered.get("input_cost_per_token") == 0.001 + assert registered.get("output_cost_per_token") == 0.002 + # Downstream _check_provider_match must accept any provider for + # this entry, mirroring the cost calculator path. + assert _check_provider_match(registered, "openai") is True + assert _check_provider_match(registered, "anthropic") is True + finally: + litellm.model_cost.pop(model_key, None) + + +def test_register_model_router_add_deployment_custom_pricing_applies(): + """End-to-end regression for https://github.com/BerriAI/litellm/issues/28336. + + ``Router.add_deployment`` registers custom pricing without passing + ``litellm_provider``. Cost calculation must still pick up the custom + pricing instead of falling back to the default provider price. + """ + from litellm import Router + + model_key = "router-add-deployment-custom-pricing-28336" + deployment_model = f"openai/{model_key}" + litellm.model_cost.pop(model_key, None) + litellm.model_cost.pop(deployment_model, None) + + router = Router( + model_list=[ + { + "model_name": model_key, + "litellm_params": { + "model": deployment_model, + "api_key": "fake-key-for-registration", + "input_cost_per_token": 0.00042, + "output_cost_per_token": 0.00084, + }, + "model_info": {"id": "deployment-28336"}, + } + ] + ) + + try: + # ``add_deployment`` runs as part of ``Router.__init__``; the + # registered entry must not block ``_check_provider_match`` for + # the deployment's provider. + from litellm.utils import _check_provider_match + + registered_keys = [ + k for k in (deployment_model, model_key) if k in litellm.model_cost + ] + assert registered_keys, ( + "Router.add_deployment did not register custom pricing for " + f"{model_key} / {deployment_model}" + ) + for k in registered_keys: + assert _check_provider_match(litellm.model_cost[k], "openai") is True, ( + f"custom pricing for {k} was dropped by _check_provider_match" + ) + finally: + litellm.model_cost.pop(model_key, None) + litellm.model_cost.pop(deployment_model, None) + del router diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index de286aede93..0efb3083139 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -1140,6 +1140,34 @@ def test_check_provider_match(): assert litellm.utils._check_provider_match(model_info, "openai") is False +def test_check_provider_match_none_value_matches_any_provider(): + """ + A ``litellm_provider`` of None must be treated the same as a missing + key: both mean "no provider constraint" and should match any + ``custom_llm_provider``. + + Regression test for https://github.com/BerriAI/litellm/issues/28336. + Before the fix, ``register_model`` persisted ``litellm_provider: None`` + via ``get_model_info`` for deployments registered without a provider + (e.g. ``Router.add_deployment``), which caused ``_check_provider_match`` + to drop custom pricing intermittently. + """ + # Missing key already returned True; None must behave identically. + assert litellm.utils._check_provider_match({}, "openai") is True + assert ( + litellm.utils._check_provider_match({"litellm_provider": None}, "openai") + is True + ) + assert ( + litellm.utils._check_provider_match({"litellm_provider": None}, "anthropic") + is True + ) + # When custom_llm_provider is also None nothing constrains the match. + assert ( + litellm.utils._check_provider_match({"litellm_provider": None}, None) is True + ) + + def test_get_provider_rerank_config(): """ Test the get_provider_rerank_config function for various providers