diff --git a/.github/workflows/mutation-test.yml b/.github/workflows/mutation-test.yml new file mode 100644 index 00000000000..8094ca57467 --- /dev/null +++ b/.github/workflows/mutation-test.yml @@ -0,0 +1,131 @@ +name: "Mutation Test (manual)" + +# Manually-triggered mutation testing. Runs mutmut against the scope +# configured in [tool.mutmut] in pyproject.toml (currently the +# litellm/proxy/management_endpoints/ folder). Intended cadence is roughly +# weekly — clicked from the Actions tab when someone wants a fresh report. +# +# Uploads a structured `mutation-report.md` (Meta ACH-style: original + +# mutated function with `# MUTANT START`/`# MUTANT END` delimiters + the +# existing tests + a task instruction) as a workflow artifact. Failures +# do not block anything because nothing depends on this workflow. + +on: + workflow_dispatch: + +permissions: + contents: read + +concurrency: + group: mutation-test-${{ github.ref }} + cancel-in-progress: true + +jobs: + mutation: + name: Run mutmut + runs-on: ubuntu-latest + # Whole-folder mutation against ~15 files / ~7.5k LOC can take hours. + # 350 minutes is just under the GitHub-hosted job cap of 360 minutes. + timeout-minutes: 350 + + steps: + - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + with: + persist-credentials: false + + - name: Set up Python + uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 + with: + python-version: "3.12" + + - name: Set up uv + uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 + with: + version: "0.10.9" + + - name: Cache uv dependencies + uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0 + with: + path: | + ~/.cache/uv + .venv + key: ${{ runner.os }}-uv-${{ hashFiles('uv.lock') }} + restore-keys: | + ${{ runner.os }}-uv- + + - name: Install dependencies + run: | + uv sync --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router + + - name: Generate Prisma client + env: + PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache + run: | + uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma + + # mutmut 3.x runs tests inside a `mutants/` sandbox where it injects + # mutation trampolines. uv installs the project as editable by default, + # which puts the original source dir on sys.path via a .pth file and + # shadows the sandbox copy — so tests would never exercise the mutated + # code. Reinstalling non-editable removes the .pth shadow. + - name: Reinstall litellm non-editable (so mutants/ is not shadowed) + run: | + uv pip uninstall litellm + uv pip install . --no-deps + + # pytest-retry's pytest_configure hook crashes with + # `INTERNALERROR: no option named 'filtered_exceptions'` when invoked + # via mutmut's in-process pytest.main() call. The entry-point name + # doesn't normalize cleanly with `-p no:`, so just remove the + # package outright. Reruns are wrong for mutation testing anyway — + # rerunning a "failed" mutant test would mask which mutants are killed. + - name: Remove pytest plugins that conflict with mutmut + run: | + uv pip uninstall pytest-retry || true + + - name: Run mutmut + env: + # Make the mutants/ sandbox win over site-packages on sys.path so the + # trampolined files are imported instead of the installed copy. + PYTHONPATH: ${{ github.workspace }}/mutants + run: | + set -o pipefail + mkdir -p mutants + uv run --no-sync --with mutmut==3.5.0 mutmut run 2>&1 | tee mutmut-run.log + + # Generate the structured report. The script embeds the enclosing + # function source for each survivor (via Python AST) and includes the + # existing test files, so an LLM agent has enough context to write + # killing tests without further file lookups. Modeled on Meta's ACH + # prompt template (arXiv 2501.12862). + - name: Generate detailed mutation report + if: always() + run: | + set +e + uv run --no-sync --with mutmut==3.5.0 mutmut export-cicd-stats > /dev/null 2>&1 + uv run --no-sync --with mutmut==3.5.0 mutmut results > mutmut-results.txt 2>&1 + uv run --no-sync python scripts/mutation_report.py + # The full report can be very long for big test files; the run-page + # summary cuts off at 1 MB. Append the head of the report (summary + # + survivor list) and link out to the artifact for the full body. + { + head -c 900000 mutation-report.md + echo "" + echo "" + echo "_Full report (with embedded function bodies and test files) is in the workflow artifact._" + } >> "$GITHUB_STEP_SUMMARY" + + - name: Upload mutmut artifacts + if: always() + uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1 + with: + name: mutmut-${{ github.run_id }}-${{ github.run_attempt }} + path: | + mutation-report.md + mutmut-results.txt + mutmut-run.log + mutants/mutmut-stats.json + mutants/mutmut-cicd-stats.json + mutants/litellm/proxy/management_endpoints/**/*.py + if-no-files-found: warn + retention-days: 14 diff --git a/litellm/__init__.py b/litellm/__init__.py index fd3d47ec154..e1b367fb234 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1426,6 +1426,12 @@ if TYPE_CHECKING: ) from .llms.datarobot.chat.transformation import DataRobotConfig as DataRobotConfig from .llms.anthropic.chat.transformation import AnthropicConfig as AnthropicConfig + from .llms.bedrock.claude_platform.transformation import ( + BedrockClaudePlatformConfig as BedrockClaudePlatformConfig, + ) + from .llms.bedrock.claude_platform.messages_transformation import ( + BedrockClaudePlatformMessagesConfig as BedrockClaudePlatformMessagesConfig, + ) from .llms.anthropic.completion.transformation import ( AnthropicTextConfig as AnthropicTextConfig, ) diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index 119e62a5b38..3531e8d96b9 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -131,6 +131,7 @@ LLM_CONFIG_NAMES = ( "OpenrouterConfig", "DataRobotConfig", "AnthropicConfig", + "BedrockClaudePlatformConfig", "AnthropicTextConfig", "GroqSTTConfig", "TritonConfig", @@ -170,6 +171,7 @@ LLM_CONFIG_NAMES = ( "SagemakerNovaConfig", "CohereChatConfig", "AnthropicMessagesConfig", + "BedrockClaudePlatformMessagesConfig", "AmazonAnthropicClaudeMessagesConfig", "AmazonMantleMessagesConfig", "TogetherAIConfig", @@ -610,6 +612,10 @@ _LLM_CONFIGS_IMPORT_MAP = { "OpenrouterConfig": (".llms.openrouter.chat.transformation", "OpenrouterConfig"), "DataRobotConfig": (".llms.datarobot.chat.transformation", "DataRobotConfig"), "AnthropicConfig": (".llms.anthropic.chat.transformation", "AnthropicConfig"), + "BedrockClaudePlatformConfig": ( + ".llms.bedrock.claude_platform.transformation", + "BedrockClaudePlatformConfig", + ), "AnthropicTextConfig": ( ".llms.anthropic.completion.transformation", "AnthropicTextConfig", @@ -712,6 +718,10 @@ _LLM_CONFIGS_IMPORT_MAP = { ".llms.anthropic.experimental_pass_through.messages.transformation", "AnthropicMessagesConfig", ), + "BedrockClaudePlatformMessagesConfig": ( + ".llms.bedrock.claude_platform.messages_transformation", + "BedrockClaudePlatformMessagesConfig", + ), "AmazonAnthropicClaudeMessagesConfig": ( ".llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation", "AmazonAnthropicClaudeMessagesConfig", diff --git a/litellm/_redis.py b/litellm/_redis.py index f12afbac297..1c11ea829ba 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -19,6 +19,7 @@ import redis.asyncio as async_redis # type: ignore from litellm import get_secret, get_secret_str from litellm._redis_credential_provider import ( + AzureADCredentialProvider, GCPIAMCredentialProvider, _generate_gcp_iam_access_token, ) @@ -27,6 +28,8 @@ from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker from ._logging import verbose_logger +AZURE_REDIS_SCOPE = "https://redis.azure.com/.default" + def _get_redis_kwargs(): arg_spec = inspect.getfullargspec(redis.Redis) @@ -43,6 +46,10 @@ def _get_redis_kwargs(): "redis_connect_func", "gcp_service_account", "gcp_ssl_ca_certs", + "azure_redis_ad_token", + "azure_client_id", + "azure_tenant_id", + "azure_client_secret", ] available_args = [x for x in arg_spec.args if x not in exclude_args] + include_args @@ -89,6 +96,10 @@ def _get_redis_cluster_kwargs(client=None): ) # Needed for sync clusters and IAM detection available_args.append("gcp_service_account") available_args.append("gcp_ssl_ca_certs") + available_args.append("azure_redis_ad_token") + available_args.append("azure_client_id") + available_args.append("azure_tenant_id") + available_args.append("azure_client_secret") available_args.append("max_connections") return available_args @@ -155,6 +166,125 @@ def create_gcp_iam_redis_connect_func( return iam_connect +def _build_azure_credential( + azure_client_id: Optional[str] = None, + azure_tenant_id: Optional[str] = None, + azure_client_secret: Optional[str] = None, +): + """ + Build a long-lived Azure credential object. + + Azure SDK credentials cache tokens internally and handle expiry/refresh + transparently, so this should be called once and the result reused. + """ + try: + from azure.identity import ( + ClientSecretCredential, + DefaultAzureCredential, + ManagedIdentityCredential, + ) + except ImportError: + raise ImportError( + "azure-identity is required for Azure AD Redis authentication. " + "Install it with: pip install azure-identity" + ) + + _client_id = azure_client_id or os.environ.get("AZURE_CLIENT_ID") + _tenant_id = azure_tenant_id or os.environ.get("AZURE_TENANT_ID") + _client_secret = azure_client_secret or os.environ.get("AZURE_CLIENT_SECRET") + + if _client_id and _tenant_id and _client_secret: + return ClientSecretCredential( + client_id=_client_id, + tenant_id=_tenant_id, + client_secret=_client_secret, + ) + elif _client_id: + return ManagedIdentityCredential(client_id=_client_id) + else: + return DefaultAzureCredential() + + +def _generate_azure_ad_redis_token( + azure_client_id: Optional[str] = None, + azure_tenant_id: Optional[str] = None, + azure_client_secret: Optional[str] = None, +) -> str: + """ + One-shot helper that builds a credential and fetches a single Azure AD + access token for Redis. Each call rebuilds the credential and performs a + network round-trip, so it should not be used in steady-state Redis flows + — the sync (``create_azure_ad_redis_connect_func``) and async paths + (``AzureADCredentialProvider``) keep the credential alive across + connections so the Azure SDK's internal cache + silent refresh apply. + """ + credential = _build_azure_credential( + azure_client_id=azure_client_id, + azure_tenant_id=azure_tenant_id, + azure_client_secret=azure_client_secret, + ) + token = credential.get_token(AZURE_REDIS_SCOPE) + return token.token + + +def create_azure_ad_redis_connect_func( + azure_client_id: Optional[str] = None, + azure_tenant_id: Optional[str] = None, + azure_client_secret: Optional[str] = None, +) -> Callable: + """ + Creates a custom Redis connection function for Azure AD authentication. + + Used for sync Redis clients. The credential is created once (captured by the + closure) and reused across connections — the Azure SDK handles token caching + and silent renewal internally. Only ``get_token`` is called per connection. + """ + credential = _build_azure_credential( + azure_client_id=azure_client_id, + azure_tenant_id=azure_tenant_id, + azure_client_secret=azure_client_secret, + ) + + def ad_connect(self): + """Initialize the connection and authenticate using Azure AD""" + from redis.exceptions import ( + AuthenticationError, + AuthenticationWrongNumberOfArgsError, + ) + from redis.utils import str_if_bytes + + self._parser.on_connect(self) + + access_token = credential.get_token(AZURE_REDIS_SCOPE).token + + # Only include username when explicitly set — sending AUTH "" + # is invalid for most ACL-configured Azure Redis instances. + username = os.environ.get("REDIS_USERNAME", "") + if username: + auth_args = (username, access_token) + else: + auth_args = (access_token,) + + self.send_command("AUTH", *auth_args, check_health=False) + + try: + auth_response = self.read_response() + except AuthenticationWrongNumberOfArgsError: + # Fallback: try with just the token (Redis < 6 / no ACL) + self.send_command("AUTH", access_token, check_health=False) + auth_response = self.read_response() + + if str_if_bytes(auth_response) != "OK": + raise AuthenticationError("Azure AD authentication failed for Redis") + + # Attach the live credential object so async paths can wrap it in + # AzureADCredentialProvider for refresh-aware token retrieval. The raw + # client_id/tenant_id/secret are intentionally NOT exposed here — the + # credential closure already holds them. + ad_connect._azure_credential = credential # type: ignore[attr-defined] + return ad_connect + + def get_redis_url_from_environment(): if "REDIS_URL" in os.environ: return os.environ["REDIS_URL"] @@ -179,7 +309,7 @@ def get_redis_url_from_environment(): return f"{redis_protocol}://{auth_part}{os.environ['REDIS_HOST']}:{os.environ['REDIS_PORT']}" -def _get_redis_client_logic(**env_overrides): +def _get_redis_client_logic(**env_overrides): # noqa: PLR0915 """ Common functionality across sync + async redis client implementations """ @@ -253,6 +383,52 @@ def _get_redis_client_logic(**env_overrides): if _gcp_ssl_ca_certs and redis_kwargs.get("ssl", False): redis_kwargs["ssl_ca_certs"] = _gcp_ssl_ca_certs + # Handle Azure AD authentication (after GCP IAM block) + _azure_redis_ad_token = redis_kwargs.get("azure_redis_ad_token") or get_secret( + "REDIS_AZURE_AD_TOKEN" + ) + + _azure_ad_enabled = ( + _azure_redis_ad_token is not None + and str(_azure_redis_ad_token).lower() == "true" + ) + + if _azure_ad_enabled and _gcp_service_account is not None: + verbose_logger.warning( + "Both GCP IAM (gcp_service_account) and Azure AD (azure_redis_ad_token) are configured for Redis. " + "Using GCP IAM. Remove one to avoid misconfiguration." + ) + + if _azure_ad_enabled and _gcp_service_account is None: + _azure_client_id = redis_kwargs.get("azure_client_id") or get_secret_str( + "AZURE_CLIENT_ID" + ) + _azure_tenant_id = redis_kwargs.get("azure_tenant_id") or get_secret_str( + "AZURE_TENANT_ID" + ) + _azure_client_secret = redis_kwargs.get( + "azure_client_secret" + ) or get_secret_str("AZURE_CLIENT_SECRET") + + verbose_logger.debug("Setting up Azure AD authentication for Redis.") + redis_kwargs["redis_connect_func"] = create_azure_ad_redis_connect_func( + azure_client_id=_azure_client_id, + azure_tenant_id=_azure_tenant_id, + azure_client_secret=_azure_client_secret, + ) + # Marker for async paths to detect Azure AD auth. The live credential + # object is attached separately as `_azure_credential` by + # `create_azure_ad_redis_connect_func`; the raw client_id/tenant_id/secret + # are intentionally NOT exposed on the function to avoid leaking + # credentials via inspection or logging. + redis_kwargs["redis_connect_func"]._azure_redis_ad_token = True # type: ignore[attr-defined] + + # Always remove Azure-specific kwargs that shouldn't be passed to Redis client + redis_kwargs.pop("azure_redis_ad_token", None) + redis_kwargs.pop("azure_client_id", None) + redis_kwargs.pop("azure_tenant_id", None) + redis_kwargs.pop("azure_client_secret", None) + if "url" in redis_kwargs and redis_kwargs["url"] is not None: # Only strip host/port/db/password when not routing to a cluster. # When startup_nodes is also present the cluster path takes priority and @@ -373,7 +549,7 @@ def get_redis_client(**env_overrides): return redis.Redis(**redis_kwargs) -def get_redis_async_client( +def get_redis_async_client( # noqa: PLR0915 connection_pool: Optional[async_redis.BlockingConnectionPool] = None, **env_overrides, ) -> Union[async_redis.Redis, async_redis.RedisCluster]: @@ -398,6 +574,14 @@ def get_redis_async_client( cluster_kwargs["credential_provider"] = GCPIAMCredentialProvider( redis_connect_func._gcp_service_account ) + # Handle Azure AD authentication for async clusters via CredentialProvider + # so the credential's internal cache + silent refresh runs per connection + # (mirrors GCP IAM above; avoids static-token-baked-in-pool expiry). + elif redis_connect_func and hasattr(redis_connect_func, "_azure_credential"): + cluster_kwargs["credential_provider"] = AzureADCredentialProvider( + redis_connect_func._azure_credential, + username=os.environ.get("REDIS_USERNAME") or None, + ) new_startup_nodes: List[ClusterNode] = [] @@ -431,6 +615,22 @@ def get_redis_async_client( # Check for Redis Sentinel if "sentinel_nodes" in redis_kwargs and "service_name" in redis_kwargs: return _init_async_redis_sentinel(redis_kwargs) + + # Wrap GCP / Azure AD auth in a CredentialProvider for the standard async + # Redis client. The async client doesn't support redis_connect_func, but it + # does honour credential_provider — which is called per connection, so the + # underlying SDK can refresh tokens silently before they expire. + redis_connect_func = redis_kwargs.pop("redis_connect_func", None) + if redis_connect_func and hasattr(redis_connect_func, "_azure_credential"): + redis_kwargs["credential_provider"] = AzureADCredentialProvider( + redis_connect_func._azure_credential, + username=os.environ.get("REDIS_USERNAME") or None, + ) + elif redis_connect_func and hasattr(redis_connect_func, "_gcp_service_account"): + redis_kwargs["credential_provider"] = GCPIAMCredentialProvider( + redis_connect_func._gcp_service_account + ) + _pretty_print_redis_config(redis_kwargs=redis_kwargs) if connection_pool is not None: @@ -464,6 +664,21 @@ def get_redis_connection_pool( redis_kwargs["max_connections"], ) return async_redis.BlockingConnectionPool.from_url(**pool_kwargs) + + # Wrap GCP / Azure AD auth in a CredentialProvider so pool-managed + # connections re-fetch tokens via the SDK's internal cache + silent refresh + # rather than reusing a single token captured at pool creation. + redis_connect_func = redis_kwargs.pop("redis_connect_func", None) + if redis_connect_func and hasattr(redis_connect_func, "_azure_credential"): + redis_kwargs["credential_provider"] = AzureADCredentialProvider( + redis_connect_func._azure_credential, + username=os.environ.get("REDIS_USERNAME") or None, + ) + elif redis_connect_func and hasattr(redis_connect_func, "_gcp_service_account"): + redis_kwargs["credential_provider"] = GCPIAMCredentialProvider( + redis_connect_func._gcp_service_account + ) + connection_class = async_redis.Connection if "ssl" in redis_kwargs: connection_class = async_redis.SSLConnection diff --git a/litellm/_redis_credential_provider.py b/litellm/_redis_credential_provider.py index 70725fe12c4..586b1c7716c 100644 --- a/litellm/_redis_credential_provider.py +++ b/litellm/_redis_credential_provider.py @@ -1,10 +1,13 @@ import asyncio import threading import time -from typing import Dict, Tuple +from typing import Any, Dict, Optional, Tuple, Union from redis.credentials import CredentialProvider # type: ignore[attr-defined] +# Azure AD scope for Redis Cache for Azure. +AZURE_REDIS_SCOPE = "https://redis.azure.com/.default" + # GCP IAM tokens are valid for 1 hour. Cache for 55 minutes to refresh before expiry. _GCP_IAM_TOKEN_TTL_SECONDS = 3300 @@ -101,3 +104,33 @@ class GCPIAMCredentialProvider(CredentialProvider): _get_cached_gcp_iam_token, self._gcp_service_account ) return (token,) + + +class AzureADCredentialProvider(CredentialProvider): + """ + redis.credentials.CredentialProvider implementation that supplies Azure AD + tokens for Redis authentication. + + Wraps an azure-identity credential object so the Azure SDK's internal token + cache and silent refresh are honoured on every Redis connection. This avoids + the static-token-baked-in-pool issue where pool-managed connections would + fail authentication after the initial token expired (~1 hour TTL). + """ + + def __init__(self, credential: Any, username: Optional[str] = None) -> None: + self._credential = credential + self._username = username + + def get_credentials(self) -> Union[Tuple[str], Tuple[str, str]]: + token = self._credential.get_token(AZURE_REDIS_SCOPE).token + if self._username: + return (self._username, token) + return (token,) + + async def get_credentials_async(self) -> Union[Tuple[str], Tuple[str, str]]: + token_obj = await asyncio.to_thread( + self._credential.get_token, AZURE_REDIS_SCOPE + ) + if self._username: + return (self._username, token_obj.token) + return (token_obj.token,) diff --git a/litellm/batches/main.py b/litellm/batches/main.py index 259439d4d09..15ee9303969 100644 --- a/litellm/batches/main.py +++ b/litellm/batches/main.py @@ -617,24 +617,35 @@ def retrieve_batch( _is_async = kwargs.pop("aretrieve_batch", False) is True client = kwargs.get("client", None) - # Check if this is an async invoke ARN (different from regular batch ARN) - # Async invoke ARNs have format: arn:aws(-[^:]+)?:bedrock:[a-z0-9-]{1,20}:[0-9]{12}:async-invoke/[a-z0-9]{12} - if ( - batch_id.startswith("arn:aws") - and ":bedrock:" in batch_id - and ":async-invoke/" in batch_id - ): - # Handle async invoke status check - # Remove aws_region_name from kwargs to avoid duplicate parameter - async_kwargs = kwargs.copy() - async_kwargs.pop("aws_region_name", None) + # Bedrock has two distinct ARN families that need different APIs: + # * async-invoke ARNs (Twelve Labs Marengo embeddings) -> bedrock-runtime data plane + # * model-invocation-job ARNs (CreateModelInvocationJob batch) -> bedrock control plane + # They live on different AWS service endpoints and can't share a handler. + # ARN shapes: + # arn:aws(-[^:]+)?:bedrock:::async-invoke/ + # arn:aws(-[^:]+)?:bedrock:::model-invocation-job/ + if batch_id.startswith("arn:aws") and ":bedrock:" in batch_id: + if ":async-invoke/" in batch_id: + # Remove aws_region_name from kwargs to avoid duplicate parameter + async_kwargs = kwargs.copy() + async_kwargs.pop("aws_region_name", None) - return BedrockBatchesHandler._handle_async_invoke_status( - batch_id=batch_id, - aws_region_name=kwargs.get("aws_region_name", "us-east-1"), - logging_obj=litellm_logging_obj, - **async_kwargs, - ) + return BedrockBatchesHandler._handle_async_invoke_status( + batch_id=batch_id, + aws_region_name=kwargs.get("aws_region_name", "us-east-1"), + logging_obj=litellm_logging_obj, + **async_kwargs, + ) + if ":model-invocation-job/" in batch_id: + mij_kwargs = kwargs.copy() + mij_kwargs.pop("aws_region_name", None) + + return BedrockBatchesHandler._handle_model_invocation_job_status( + batch_id=batch_id, + aws_region_name=kwargs.get("aws_region_name"), + logging_obj=litellm_logging_obj, + **mij_kwargs, + ) # Try to use provider config first (for providers like bedrock) model: Optional[str] = kwargs.get("model", None) diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index da3b9184edb..32423f23314 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -119,6 +119,20 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): def __init__(self): pass + def _normalize_tool_choice_for_responses_api(self, tool_choice: Any) -> Any: + """Chat tool_choice uses function.name; Responses API expects top-level name.""" + if not isinstance(tool_choice, dict) or tool_choice.get("type") != "function": + return tool_choice + if isinstance(tool_choice.get("name"), str) and tool_choice.get("name"): + # Return only Responses shape so stray chat ``function`` key is not sent upstream. + return {"type": "function", "name": tool_choice["name"]} + fn = tool_choice.get("function") + if isinstance(fn, dict): + fn_name = fn.get("name") + if isinstance(fn_name, str) and fn_name: + return {"type": "function", "name": fn_name} + return tool_choice + def _handle_raw_dict_response_item( self, item: Dict[str, Any], index: int ) -> Tuple[Optional[Any], int]: @@ -309,6 +323,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): text_format = self._transform_response_format_to_text_format(value) if text_format: responses_api_request["text"] = text_format # type: ignore + elif key == "tool_choice": + responses_api_request["tool_choice"] = ( # type: ignore[assignment] + self._normalize_tool_choice_for_responses_api(value) + ) elif key in ResponsesAPIOptionalRequestParams.__annotations__.keys(): responses_api_request[key] = value # type: ignore elif key == "previous_response_id": diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index 0af016feb26..abce3c3e1c7 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -1771,17 +1771,41 @@ class OpenTelemetry(CustomLogger): value=safe_dumps(transformed_messages), ) - if kwargs.get("system_instructions"): - transformed_system_instructions = ( - self._transform_messages_to_otel_semantic_conventions( - kwargs.get("system_instructions") + # Coalesce the different kwarg names that carry the system + # prompt depending on the call path: + # - "system_instructions" — Vertex AI Gemini chat-completion + # - "instructions" — OpenAI Responses API + # - "system" — Anthropic Messages API + # Use `is not None` rather than truthiness to avoid falsy + # values (e.g. []) falling through to the wrong kwarg. + system_instructions = ( + kwargs.get("system_instructions") + if kwargs.get("system_instructions") is not None + else ( + kwargs.get("instructions") + if kwargs.get("instructions") is not None + else kwargs.get("system") + ) + ) + if system_instructions: + if isinstance(system_instructions, str): + # Plain text system prompt — no transformation needed + self.safe_set_attribute( + span=span, + key=SpanAttributes.GEN_AI_SYSTEM_INSTRUCTIONS.value, + value=system_instructions, + ) + else: + transformed_system_instructions = ( + self._transform_messages_to_otel_semantic_conventions( + system_instructions + ) + ) + self.safe_set_attribute( + span=span, + key=SpanAttributes.GEN_AI_SYSTEM_INSTRUCTIONS.value, + value=safe_dumps(transformed_system_instructions), ) - ) - self.safe_set_attribute( - span=span, - key=SpanAttributes.GEN_AI_SYSTEM_INSTRUCTIONS.value, - value=safe_dumps(transformed_system_instructions), - ) self.safe_set_attribute( span=span, @@ -1840,6 +1864,57 @@ class OpenTelemetry(CustomLogger): value=value, ) + elif response_obj.get("output"): + # Responses API: ResponsesAPIResponse has an "output" + # list instead of "choices". Each item with + # type="message" contains a "content" list of + # OutputText objects (type="output_text"). + output_items = response_obj.get("output") + output_messages = self._transform_responses_api_output_to_otel( + output_items + ) + if output_messages: + self.safe_set_attribute( + span=span, + key=SpanAttributes.GEN_AI_OUTPUT_MESSAGES.value, + value=safe_dumps(output_messages), + ) + + # Emit per-tool-call span attributes (parity with + # the choices branch that calls _tool_calls_kv_pair). + # Convert Responses API function_call items to the + # ChatCompletionMessageToolCall format expected by + # _tool_calls_kv_pair. + tool_calls = [] + for out_item in output_items: + item_d = self._to_dict(out_item) + if item_d and item_d.get("type") == "function_call": + tool_calls.append( + { + "function": { + "name": item_d.get("name", ""), + "arguments": item_d.get("arguments", ""), + } + } + ) + if tool_calls: + kv_pairs = OpenTelemetry._tool_calls_kv_pair(tool_calls) # type: ignore + for key, value in kv_pairs.items(): + self.safe_set_attribute( + span=span, + key=key, + value=value, + ) + + # Extract finish reason from ResponsesAPIResponse.status + status = response_obj.get("status") + if status: + self.safe_set_attribute( + span=span, + key=SpanAttributes.GEN_AI_RESPONSE_FINISH_REASONS.value, + value=safe_dumps([status]), + ) + except Exception as e: self.handle_callback_failure( callback_name=self.callback_name or "opentelemetry" @@ -1935,6 +2010,78 @@ class OpenTelemetry(CustomLogger): transformed.append(transformed_msg) return transformed + @staticmethod + def _to_dict(obj) -> Optional[dict]: + """Normalize an object to a plain dict. + + Handles three forms that appear in practice: + + 1. Plain ``dict`` — returned as-is. + 2. LiteLLM's ``BaseLiteLLMOpenAIResponseObject`` — exposes a + ``.get()`` method that delegates to ``__dict__``. + 3. Raw Pydantic v2 models from the ``openai`` SDK (e.g. + ``ResponseOutputMessage``, ``ResponseOutputText``) — these do + **not** have ``.get()`` but do have ``.model_dump()``. + + Returns ``None`` for anything else so callers can skip it. + """ + if isinstance(obj, dict): + return obj + if hasattr(obj, "get"): + # BaseLiteLLMOpenAIResponseObject duck-type + return obj # type: ignore[return-value] + if hasattr(obj, "model_dump"): + # Raw Pydantic v2 model (e.g. openai SDK types) + return obj.model_dump() # type: ignore[union-attr] + return None + + def _transform_responses_api_output_to_otel(self, output: List) -> List[dict]: + """ + Transform Responses API output items into OTEL GenAI 1.38 format. + + The Responses API returns output as a list of items, each with a + ``type`` field. Message items (``type="message"``) contain a + ``content`` list of ``OutputText`` objects with ``type="output_text"`` + and ``text`` fields. + + Items may be plain dicts, LiteLLM wrapper objects (with ``.get()``), + or raw Pydantic v2 models from the ``openai`` SDK (with + ``.model_dump()``). We normalize each item to a dict via + ``_to_dict`` before processing. + + This method converts them to the same ``{"role": ..., "parts": [...]}`` + format used by ``_transform_choices_to_otel_semantic_conventions``. + """ + transformed = [] + for raw_item in output: + item = self._to_dict(raw_item) + if item is None: + continue + if item.get("type") == "message": + role = item.get("role", "assistant") + parts = [] + for raw_content in item.get("content", []): + content = self._to_dict(raw_content) + if content is None: + continue + if content.get("type") == "output_text": + text = content.get("text", "") + if text: + parts.append({"type": "text", "content": text}) + if parts: + transformed.append({"role": role, "parts": parts}) + elif item.get("type") == "function_call": + # Surface tool calls from Responses API output + part: dict = { + "type": "tool_call", + "name": item.get("name", ""), + "arguments": item.get("arguments", ""), + } + if item.get("call_id"): + part["id"] = item["call_id"] + transformed.append({"role": "assistant", "parts": [part]}) + return transformed + def set_raw_request_attributes(self, span: Span, kwargs, response_obj): try: # Only set provider-specific raw payload attributes on this span. diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 2f11a3fccb5..1ce80207552 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -1809,9 +1809,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): Translate messages to anthropic format. """ ## VALIDATE REQUEST - """ - Anthropic doesn't support tool calling without `tools=` param specified. - """ + """Anthropic requires ``tools`` when messages include tool blocks; LiteLLM injects a dummy tool if omitted (no ``modify_params`` needed).""" from litellm.litellm_core_utils.prompt_templates.factory import ( anthropic_messages_pt, ) @@ -1821,16 +1819,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): and messages is not None and has_tool_call_blocks(messages) ): - if litellm.modify_params: - optional_params["tools"], _ = self._map_tools( - add_dummy_tool(custom_llm_provider="anthropic") - ) - else: - raise litellm.UnsupportedParamsError( - message="Anthropic doesn't support tool calling without `tools=` param specified. Pass `tools=` param OR set `litellm.modify_params = True` // `litellm_settings::modify_params: True` to add dummy tool to the request.", - model="", - llm_provider="anthropic", - ) + optional_params["tools"], _ = self._map_tools( + add_dummy_tool(custom_llm_provider="anthropic") + ) # Drop thinking param if thinking is enabled but thinking_blocks are missing # This prevents the error: "Expected thinking or redacted_thinking, but found tool_use" diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py index 0c59e812e0b..a515e7f82b8 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py @@ -454,6 +454,16 @@ def anthropic_messages_handler( "display": "summarized", } + # Mirror Router._get_timeout: prefer `stream_timeout` when streaming, + # fall back to `timeout`. Coerce string form to float for httpx. + _resolved_timeout = ( + litellm_params.stream_timeout + if stream and litellm_params.stream_timeout is not None + else litellm_params.timeout + ) + if isinstance(_resolved_timeout, str): + _resolved_timeout = float(_resolved_timeout) + return base_llm_http_handler.anthropic_messages_handler( model=model, messages=messages, @@ -469,5 +479,6 @@ def anthropic_messages_handler( api_key=api_key, api_base=api_base, stream=stream, + timeout=_resolved_timeout, kwargs=kwargs, ) diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index 0885775932c..9dd2b055a12 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -1428,7 +1428,13 @@ class BaseAWSLLM: def _sign_request( self, - service_name: Literal["bedrock", "sagemaker", "bedrock-agentcore", "s3vectors"], + service_name: Literal[ + "bedrock", + "sagemaker", + "bedrock-agentcore", + "s3vectors", + "aws-external-anthropic", + ], headers: dict, optional_params: dict, request_data: dict, diff --git a/litellm/llms/bedrock/batches/handler.py b/litellm/llms/bedrock/batches/handler.py index f141bbd9ab4..c071f331337 100644 --- a/litellm/llms/bedrock/batches/handler.py +++ b/litellm/llms/bedrock/batches/handler.py @@ -1,8 +1,79 @@ +from datetime import datetime +from typing import Any, Optional, cast + from openai.types.batch import BatchRequestCounts from openai.types.batch import Metadata as OpenAIBatchMetadata from litellm.types.utils import LiteLLMBatch +# AWS Bedrock model-invocation-job statuses → OpenAI Batch statuses. +# Mirrors the mapping used by `BedrockBatchesConfig.transform_create_batch_response` +# so create / retrieve return consistent statuses. +_BEDROCK_MIJ_STATUS_TO_OPENAI = { + "Submitted": "validating", + "Validating": "validating", + "Scheduled": "validating", + "InProgress": "in_progress", + "Stopping": "cancelling", + "Stopped": "cancelled", + "Completed": "completed", + "PartiallyCompleted": "completed", + "Failed": "failed", + "Expired": "expired", +} + + +def _extract_region_from_bedrock_arn(arn: str) -> Optional[str]: + """ARN shape: ``arn:aws:bedrock:::/``""" + try: + parts = arn.split(":") + if len(parts) >= 4 and parts[2] == "bedrock": + return parts[3] or None + except Exception: + pass + return None + + +def _extract_job_id_from_arn(arn: str) -> Optional[str]: + """``arn:aws:bedrock:::model-invocation-job/`` -> ````.""" + if ":model-invocation-job/" not in arn: + return None + return arn.rsplit("/", 1)[-1] or None + + +def _predict_output_file_uri( + output_prefix: str, input_uri: str, job_id: Optional[str] +) -> Optional[str]: + """ + Compute the deterministic per-job result file URI Bedrock writes to. + + Bedrock lays results out as:: + + //.out + + We compute it client-side so OpenAI-style ``client.files.content(output_file_id)`` + works without an extra S3 ``ListObjectsV2`` round-trip. Returns ``None`` if we + don't have enough info; callers should fall back to the bare prefix. + """ + if not output_prefix or not input_uri or not job_id: + return None + if not output_prefix.endswith("/"): + output_prefix = output_prefix + "/" + input_basename = input_uri.rsplit("/", 1)[-1] + if not input_basename: + return None + return f"{output_prefix}{job_id}/{input_basename}.out" + + +def _to_epoch(value: Any) -> Optional[int]: + if value is None: + return None + if isinstance(value, (int, float)): + return int(value) + if isinstance(value, datetime): + return int(value.timestamp()) + return None + class BedrockBatchesHandler: """ @@ -97,3 +168,173 @@ class BedrockBatchesHandler: with concurrent.futures.ThreadPoolExecutor() as executor: future = executor.submit(run_in_thread) return future.result() + + @staticmethod + def _handle_model_invocation_job_status( + batch_id: str, + aws_region_name: Optional[str] = None, + logging_obj=None, + **kwargs, + ) -> "LiteLLMBatch": + """ + Handle ``GetModelInvocationJob`` status check for AWS Bedrock bulk batch + inference jobs (the ARN type returned by ``CreateModelInvocationJob``). + + ``CreateModelInvocationJob`` lives on the Bedrock **control plane** + (``bedrock..amazonaws.com``), distinct from the data-plane + ``bedrock-runtime`` endpoint that serves Twelve Labs async-invoke ARNs. + The two ARN families therefore can't share a handler — see + ``litellm/batches/main.py`` for the dispatch. + + Args: + batch_id: A ``arn:aws:bedrock:::model-invocation-job/`` + ARN (or just the trailing job id; both are accepted by + ``GetModelInvocationJob``). + aws_region_name: Region for the boto3 ``bedrock`` client. If omitted, + we fall back to parsing the region out of ``batch_id`` itself. + logging_obj: Optional litellm logging object. + **kwargs: Optional AWS credential overrides + (``aws_access_key_id``, ``aws_secret_access_key``, + ``aws_session_token``, ``aws_profile_name``, + ``aws_role_name``, ``aws_session_name``, + ``aws_web_identity_token``, ``aws_sts_endpoint``, + ``aws_external_id``). Unknown keys are ignored. + + Returns: + ``LiteLLMBatch`` shaped like an OpenAI Batch resource. Note that + ``request_counts`` is always ``(0, 0, 0)`` because + ``GetModelInvocationJob`` does not surface per-record counts; + callers that need accurate counts should parse + ``manifest.json.out`` from the output S3 prefix. + """ + try: + import boto3 + except ImportError as exc: + raise ImportError( + "Missing boto3 to call bedrock. Run 'pip install boto3'." + ) from exc + + # Resolve region: explicit > parsed-from-ARN > us-east-1 (boto3 default). + region = ( + aws_region_name or _extract_region_from_bedrock_arn(batch_id) or "us-east-1" + ) + + # Resolve credentials through the same path the rest of the bedrock + # provider uses, so model_list / env / role-assumption configs are + # honored. We instantiate BedrockBatchesConfig (which extends + # BaseAWSLLM) lazily to avoid a circular import at module load. + from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig + + creds = BedrockBatchesConfig().get_credentials( + aws_access_key_id=kwargs.get("aws_access_key_id"), + aws_secret_access_key=kwargs.get("aws_secret_access_key"), + aws_session_token=kwargs.get("aws_session_token"), + aws_region_name=region, + aws_session_name=kwargs.get("aws_session_name"), + aws_profile_name=kwargs.get("aws_profile_name"), + aws_role_name=kwargs.get("aws_role_name"), + aws_web_identity_token=kwargs.get("aws_web_identity_token"), + aws_sts_endpoint=kwargs.get("aws_sts_endpoint"), + aws_external_id=kwargs.get("aws_external_id"), + ) + + client = boto3.client( + "bedrock", + region_name=region, + aws_access_key_id=creds.access_key, + aws_secret_access_key=creds.secret_key, + aws_session_token=creds.token, + ) + + if logging_obj is not None: + # Use the bare job id in the logged URL so we don't double up the + # `model-invocation-job/` segment when `batch_id` is a full ARN. + # `GetModelInvocationJob` accepts either form, but only the bare id + # produces a sensible-looking URL in logs. + url_path_id = _extract_job_id_from_arn(batch_id) or batch_id + logging_obj.pre_call( + input=batch_id, + api_key="", + additional_args={ + "complete_input_dict": {"jobIdentifier": batch_id}, + "api_base": ( + f"https://bedrock.{region}.amazonaws.com/" + f"model-invocation-job/{url_path_id}" + ), + }, + ) + + response = client.get_model_invocation_job(jobIdentifier=batch_id) + + if logging_obj is not None: + logging_obj.post_call( + input=batch_id, + api_key="", + original_response=response, + additional_args={"complete_input_dict": {"jobIdentifier": batch_id}}, + ) + + bedrock_status = str(response.get("status", "")) + openai_status = cast( + Any, + _BEDROCK_MIJ_STATUS_TO_OPENAI.get(bedrock_status, "in_progress"), + ) + + input_uri = ( + response.get("inputDataConfig", {}) + .get("s3InputDataConfig", {}) + .get("s3Uri", "") + ) + output_prefix = ( + response.get("outputDataConfig", {}) + .get("s3OutputDataConfig", {}) + .get("s3Uri", "") + ) + + # Bedrock returns the output *prefix* the user supplied at job creation. + # Actual results land at //.out — we + # surface that single-file URI as `output_file_id` so the OpenAI-style + # download flow works without an extra S3 listing call. We deliberately + # do NOT fall back to the bare prefix when prediction fails: a prefix + # is not a downloadable object, so handing it back as `output_file_id` + # would reproduce the very NoSuchKey bug this handler exists to fix. + # The bare prefix is preserved in metadata for callers that want the + # `manifest.json.out` or want to do their own listing. + job_arn = response.get("jobArn", batch_id) + job_id = _extract_job_id_from_arn(job_arn) + output_file_uri = _predict_output_file_uri(output_prefix, input_uri, job_id) + + completed_at = _to_epoch(response.get("endTime")) + + # Note: metadata uses "" (not None) for unknown URIs to satisfy the + # OpenAI Batch metadata schema, which is `dict[str, str]`. The + # `output_file_id` field on the LiteLLMBatch itself does carry None + # correctly (see below), so callers should branch on that, not on + # `metadata["output_file_uri"]`. + openai_batch_metadata: OpenAIBatchMetadata = { + "model_arn": response.get("modelId", ""), + "job_arn": job_arn, + "job_name": response.get("jobName", ""), + "failure_message": response.get("message") or "", + "input_s3_uri": input_uri, + "output_s3_uri": output_prefix, + "output_file_uri": output_file_uri or "", + } + + return LiteLLMBatch( + id=job_arn, + object="batch", + status=openai_status, + created_at=_to_epoch(response.get("submitTime")) or 0, + in_progress_at=_to_epoch(response.get("lastModifiedTime")), + completed_at=completed_at if openai_status == "completed" else None, + failed_at=completed_at if openai_status == "failed" else None, + cancelled_at=completed_at if openai_status == "cancelled" else None, + expired_at=completed_at if openai_status == "expired" else None, + request_counts=BatchRequestCounts(total=0, completed=0, failed=0), + metadata=openai_batch_metadata, + completion_window="24h", + endpoint="/v1/chat/completions", + input_file_id=input_uri, + output_file_id=output_file_uri if openai_status == "completed" else None, + ) diff --git a/litellm/llms/bedrock/claude_platform/__init__.py b/litellm/llms/bedrock/claude_platform/__init__.py new file mode 100644 index 00000000000..88d4e9783c7 --- /dev/null +++ b/litellm/llms/bedrock/claude_platform/__init__.py @@ -0,0 +1,8 @@ +from .transformation import ( + BedrockClaudePlatformConfig, +) +from .messages_transformation import ( + BedrockClaudePlatformMessagesConfig, +) + +__all__ = ["BedrockClaudePlatformConfig", "BedrockClaudePlatformMessagesConfig"] diff --git a/litellm/llms/bedrock/claude_platform/common_utils.py b/litellm/llms/bedrock/claude_platform/common_utils.py new file mode 100644 index 00000000000..121221518c8 --- /dev/null +++ b/litellm/llms/bedrock/claude_platform/common_utils.py @@ -0,0 +1,107 @@ +from typing import Literal, Optional, Tuple + +import litellm +from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM +from litellm.secret_managers.main import get_secret_str + + +CLAUDE_PLATFORM_SERVICE_NAME: Literal["aws-external-anthropic"] = ( + "aws-external-anthropic" +) +CLAUDE_PLATFORM_BEDROCK_ROUTE = "claude_platform/" + + +def strip_claude_platform_route(model: str) -> str: + if model.startswith(CLAUDE_PLATFORM_BEDROCK_ROUTE): + return model.replace(CLAUDE_PLATFORM_BEDROCK_ROUTE, "", 1) + return model + + +class BedrockClaudePlatformMixin(BaseAWSLLM): + @staticmethod + def _get_workspace_id(optional_params: dict, litellm_params: dict) -> Optional[str]: + workspace_id = ( + optional_params.get("workspace_id") + or litellm_params.get("workspace_id") + or optional_params.get("aws_workspace_id") + or litellm_params.get("aws_workspace_id") + or optional_params.get("anthropic-workspace-id") + or litellm_params.get("anthropic-workspace-id") + ) + if workspace_id is None: + workspace_id = optional_params.get( + "anthropic_workspace_id" + ) or litellm_params.get("anthropic_workspace_id") + if workspace_id is not None: + return str(workspace_id) + return get_secret_str("ANTHROPIC_AWS_WORKSPACE_ID") or get_secret_str( + "ANTHROPIC_WORKSPACE_ID" + ) + + def _get_required_aws_region_name(self, optional_params: dict) -> str: + aws_region_name = ( + optional_params.get("aws_region_name") + or get_secret_str("AWS_REGION_NAME") + or get_secret_str("AWS_REGION") + or get_secret_str("AWS_DEFAULT_REGION") + ) + if aws_region_name is None: + raise litellm.AuthenticationError( + message=( + "Missing AWS region for Claude Platform on AWS. Pass " + "`aws_region_name` or set a standard AWS region environment value." + ), + llm_provider="bedrock", + model="", + ) + self._validate_aws_region_name(str(aws_region_name)) + return str(aws_region_name) + + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + optional_params: dict, + litellm_params: dict, + stream: Optional[bool] = None, + ) -> str: + api_base = ( + api_base + or litellm.api_base + or get_secret_str("ANTHROPIC_AWS_BASE_URL") + or get_secret_str("ANTHROPIC_AWS_API_BASE") + ) + if api_base is None: + aws_region_name = self._get_required_aws_region_name(optional_params) + api_base = ( + f"https://{CLAUDE_PLATFORM_SERVICE_NAME}.{aws_region_name}.api.aws" + ) + if not api_base.endswith("/v1/messages"): + api_base = f"{api_base.rstrip('/')}/v1/messages" + return api_base + + def sign_request( + self, + headers: dict, + optional_params: dict, + request_data: dict, + api_base: str, + api_key: Optional[str] = None, + model: Optional[str] = None, + stream: Optional[bool] = None, + fake_stream: Optional[bool] = None, + ) -> Tuple[dict, Optional[bytes]]: + if api_key or get_secret_str("ANTHROPIC_AWS_API_KEY"): + return headers, None + + return self._sign_request( + service_name=CLAUDE_PLATFORM_SERVICE_NAME, + headers=headers, + optional_params=optional_params, + request_data=request_data, + api_base=api_base, + model=model, + stream=stream, + fake_stream=fake_stream, + ) diff --git a/litellm/llms/bedrock/claude_platform/messages_transformation.py b/litellm/llms/bedrock/claude_platform/messages_transformation.py new file mode 100644 index 00000000000..66158196322 --- /dev/null +++ b/litellm/llms/bedrock/claude_platform/messages_transformation.py @@ -0,0 +1,71 @@ +from typing import Any, Dict, List, Optional, Tuple + +import litellm +from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( + DEFAULT_ANTHROPIC_API_VERSION, + AnthropicMessagesConfig, +) +from litellm.secret_managers.main import get_secret_str +from litellm.types.router import GenericLiteLLMParams + +from .common_utils import BedrockClaudePlatformMixin, strip_claude_platform_route + + +class BedrockClaudePlatformMessagesConfig( + BedrockClaudePlatformMixin, AnthropicMessagesConfig +): + def validate_anthropic_messages_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, + ) -> Tuple[dict, Optional[str]]: + workspace_id = self._get_workspace_id(optional_params, litellm_params) + if workspace_id is None: + raise litellm.AuthenticationError( + message=( + "Missing workspace ID for Claude Platform on AWS. Pass " + "`workspace_id` or configure the provider workspace setting." + ), + llm_provider="bedrock", + model=model, + ) + + resolved_api_key = api_key or get_secret_str("ANTHROPIC_AWS_API_KEY") + headers = { + **headers, + "anthropic-version": headers.get( + "anthropic-version", DEFAULT_ANTHROPIC_API_VERSION + ), + "content-type": headers.get("content-type", "application/json"), + "anthropic-workspace-id": workspace_id, + } + if resolved_api_key and "x-api-key" not in headers: + headers["x-api-key"] = resolved_api_key + + headers = self._update_headers_with_anthropic_beta( + headers=headers, + optional_params=optional_params, + ) + + return headers, api_base + + def transform_anthropic_messages_request( + self, + model: str, + messages: List[Dict], + anthropic_messages_optional_request_params: Dict, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Dict: + return super().transform_anthropic_messages_request( + model=strip_claude_platform_route(model), + messages=messages, + anthropic_messages_optional_request_params=anthropic_messages_optional_request_params, + litellm_params=litellm_params, + headers=headers, + ) diff --git a/litellm/llms/bedrock/claude_platform/transformation.py b/litellm/llms/bedrock/claude_platform/transformation.py new file mode 100644 index 00000000000..0167c457c96 --- /dev/null +++ b/litellm/llms/bedrock/claude_platform/transformation.py @@ -0,0 +1,94 @@ +from typing import Any, Dict, List, Optional + +import litellm +from litellm.llms.anthropic.chat.transformation import AnthropicConfig +from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.openai import AllMessageValues + +from .common_utils import BedrockClaudePlatformMixin + + +class BedrockClaudePlatformConfig(BedrockClaudePlatformMixin, AnthropicConfig): + """ + Bedrock Claude Platform uses Anthropic's Messages API with AWS gateway auth. + """ + + @property + def custom_llm_provider(self) -> Optional[str]: + return "bedrock" + + def validate_environment( + self, + headers: dict, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> Dict: + workspace_id = self._get_workspace_id(optional_params, litellm_params) + if workspace_id is None: + raise litellm.AuthenticationError( + message=( + "Missing workspace ID for Claude Platform on AWS. Pass " + "`workspace_id` or configure the provider workspace setting." + ), + llm_provider="bedrock", + model=model, + ) + + api_key = api_key or get_secret_str("ANTHROPIC_AWS_API_KEY") + anthropic_headers = self.get_anthropic_headers( + api_key=api_key, + auth_token=None, + computer_tool_used=self.is_computer_tool_used( + tools=optional_params.get("tools") + ), + prompt_caching_set=self.is_cache_control_set(messages=messages), + pdf_used=self.is_pdf_used(messages=messages), + file_id_used=self.is_file_id_used(messages=messages), + mcp_server_used=self.is_mcp_server_used( + mcp_servers=optional_params.get("mcp_servers") + ), + web_search_tool_used=self.is_web_search_tool_used( + tools=optional_params.get("tools") + ), + tool_search_used=self.is_tool_search_used( + tools=optional_params.get("tools") + ), + programmatic_tool_calling_used=self.is_programmatic_tool_calling_used( + tools=optional_params.get("tools") + ), + input_examples_used=self.is_input_examples_used( + tools=optional_params.get("tools") + ), + effort_used=self.is_effort_used( + optional_params=optional_params, model=model + ), + user_anthropic_beta_headers=self._get_user_anthropic_beta_headers( + anthropic_beta_header=headers.get("anthropic-beta") + ), + code_execution_tool_used=self.is_code_execution_tool_used( + tools=optional_params.get("tools") + ), + container_with_skills_used=self.is_container_with_skills_used( + optional_params=optional_params + ), + ) + anthropic_headers["anthropic-workspace-id"] = workspace_id + return {**headers, **anthropic_headers} + + def get_model_response_iterator( + self, + streaming_response: Any, + sync_stream: bool, + json_mode: Optional[bool] = False, + ) -> Any: + from litellm.llms.anthropic.chat.handler import ModelResponseIterator + + return ModelResponseIterator( + streaming_response=streaming_response, + sync_stream=sync_stream, + json_mode=bool(json_mode), + ) diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 856a525f773..0256d5d4b95 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -692,6 +692,7 @@ class BedrockModelInfo(BaseLLMModelInfo): ) -> Literal[ "converse", "invoke", + "claude_platform", "converse_like", "agent", "agentcore", @@ -706,6 +707,7 @@ class BedrockModelInfo(BaseLLMModelInfo): str, Literal[ "invoke", + "claude_platform", "converse_like", "converse", "agent", @@ -716,6 +718,7 @@ class BedrockModelInfo(BaseLLMModelInfo): ], ] = { "invoke/": "invoke", + "claude_platform/": "claude_platform", "converse_like/": "converse_like", "converse/": "converse", "agent/": "agent", @@ -753,6 +756,36 @@ class BedrockModelInfo(BaseLLMModelInfo): """ return "converse/" in model + @staticmethod + def _explicit_claude_platform_route(model: str) -> bool: + """ + Check if the model is an explicit Claude Platform on AWS route. + """ + return "claude_platform/" in model + + @staticmethod + def get_claude_platform_model(model: str) -> str: + """ + Strip the Claude Platform route prefix from a Bedrock model name. + """ + return model.replace("claude_platform/", "", 1) + + @staticmethod + def map_claude_platform_auth_params( + passed_params: dict, optional_params: dict + ) -> dict: + """ + Map Claude Platform route auth params that are not OpenAI request params. + """ + for key in ( + "workspace_id", + "aws_workspace_id", + "anthropic_workspace_id", + ): + if key in passed_params: + optional_params[key] = passed_params[key] + return optional_params + @staticmethod def _explicit_invoke_route(model: str) -> bool: """ @@ -815,6 +848,12 @@ class BedrockModelInfo(BaseLLMModelInfo): All other routes should return None since they will go through litellm.completion """ + ######################################################### + # Claude Platform route uses Anthropic Messages API via the AWS gateway. + ######################################################### + if BedrockModelInfo._explicit_claude_platform_route(model): + return litellm.BedrockClaudePlatformMessagesConfig() + ######################################################### # Converse routes should go through litellm.completion() if BedrockModelInfo._explicit_converse_route(model): @@ -860,7 +899,9 @@ def get_bedrock_chat_config(model: str): base_model = BedrockModelInfo.get_base_model(model) # Handle explicit routes first - if bedrock_route == "converse" or bedrock_route == "converse_like": + if bedrock_route == "claude_platform": + return litellm.BedrockClaudePlatformConfig() + elif bedrock_route == "converse" or bedrock_route == "converse_like": return litellm.AmazonConverseConfig() elif bedrock_route == "openai": return litellm.AmazonBedrockOpenAIConfig() diff --git a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py index aae2bc5e289..151e0e404a0 100644 --- a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py @@ -408,6 +408,47 @@ class AmazonAnthropicClaudeMessagesConfig( if self._supports_tool_search_on_bedrock(model): beta_set.add("tool-search-tool-2025-10-19") + @staticmethod + def _filter_context_management_for_bedrock_invoke( + anthropic_messages_request: Dict, + beta_set: set, + ) -> None: + """ + Bedrock InvokeModel accepts ``context_management`` only when it carries + ``compact_20260112`` edits paired with the ``compact-2026-01-12`` + anthropic-beta header. Other edit types (notably ``clear_thinking_20251015``, + which Claude Code sends on every request) are LiteLLM-internal and would + cause Bedrock to 400 with ``"context_management: Extra inputs are not + permitted"``. + + Filter the edits list to the supported subset, add the beta header when + compact edits remain, and drop ``context_management`` entirely when no + supported edits are left so the safety-net allowlist can pass it through. + + Ref: https://github.com/BerriAI/litellm/issues/27532 + """ + cm = anthropic_messages_request.get("context_management") + if not isinstance(cm, dict): + return + edits = cm.get("edits") + if not isinstance(edits, list): + anthropic_messages_request.pop("context_management", None) + return + + compact_edits = [ + e + for e in edits + if isinstance(e, dict) and e.get("type") == "compact_20260112" + ] + if compact_edits: + beta_set.add("compact-2026-01-12") + anthropic_messages_request["context_management"] = { + **cm, + "edits": compact_edits, + } + else: + anthropic_messages_request.pop("context_management", None) + def _convert_output_format_to_inline_schema( self, output_format: Dict, @@ -551,6 +592,11 @@ class AmazonAnthropicClaudeMessagesConfig( if injected_thinking_for_clear_thinking: beta_set.add("interleaved-thinking-2025-05-14") + self._filter_context_management_for_bedrock_invoke( + anthropic_messages_request=anthropic_messages_request, + beta_set=beta_set, + ) + self._get_tool_search_beta_header_for_bedrock( model=model, tool_search_used=tool_search_used, @@ -597,8 +643,9 @@ class AmazonAnthropicClaudeMessagesConfig( anthropic_messages_request.pop("output_config", None) # 7. Final safety net: filter top-level fields to the Bedrock Invoke allowlist. - # Catches Anthropic-only extensions (context_management, output_config, speed, - # mcp_servers, ...) and any future additions Claude Code may start sending. + # Catches Anthropic-only extensions (output_config, speed, mcp_servers, ...) + # and any future additions Claude Code may start sending. ``context_management`` + # has already been pre-filtered to its Bedrock-supported subset above. allowed = self.BEDROCK_INVOKE_ALLOWED_TOP_LEVEL_FIELDS stripped = sorted(k for k in anthropic_messages_request if k not in allowed) if stripped: diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index fa1253d9005..32f0623d1bc 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -1860,6 +1860,7 @@ class BaseLLMHTTPHandler: litellm_params: GenericLiteLLMParams, api_key: Optional[str], model: str, + timeout: Optional[Union[float, httpx.Timeout]] = None, ) -> httpx.Response: max_attempts = max( provider_config.max_retry_on_anthropic_messages_http_error, 1 @@ -1874,6 +1875,7 @@ class BaseLLMHTTPHandler: data=signed_json_body or json.dumps(request_body), stream=stream or False, logging_obj=logging_obj, + timeout=timeout, ) response.raise_for_status() return response @@ -1928,6 +1930,7 @@ class BaseLLMHTTPHandler: api_key: Optional[str] = None, api_base: Optional[str] = None, stream: Optional[bool] = False, + timeout: Optional[Union[float, httpx.Timeout]] = None, kwargs: Optional[Dict[str, Any]] = None, ) -> Union[AnthropicMessagesResponse, AsyncIterator]: from litellm.litellm_core_utils.get_provider_specific_headers import ( @@ -2065,6 +2068,7 @@ class BaseLLMHTTPHandler: litellm_params=litellm_params, api_key=api_key, model=model, + timeout=timeout, ) # used for logging + cost tracking @@ -2131,6 +2135,7 @@ class BaseLLMHTTPHandler: api_key: Optional[str] = None, api_base: Optional[str] = None, stream: Optional[bool] = False, + timeout: Optional[Union[float, httpx.Timeout]] = None, kwargs: Optional[Dict[str, Any]] = None, ) -> Union[ AnthropicMessagesResponse, @@ -2153,6 +2158,7 @@ class BaseLLMHTTPHandler: api_key=api_key, api_base=api_base, stream=stream, + timeout=timeout, kwargs=kwargs, ) raise ValueError("anthropic_messages_handler is not implemented for sync calls") diff --git a/litellm/llms/openai/chat/gpt_5_transformation.py b/litellm/llms/openai/chat/gpt_5_transformation.py index 4e34d10b187..9ccb2e1c267 100644 --- a/litellm/llms/openai/chat/gpt_5_transformation.py +++ b/litellm/llms/openai/chat/gpt_5_transformation.py @@ -3,7 +3,10 @@ from typing import Optional, Union import litellm -from litellm.utils import _is_explicitly_disabled_factory, _supports_factory +from litellm.utils import ( + _is_explicitly_disabled_factory, + _supports_factory, +) from .gpt_transformation import OpenAIGPTConfig diff --git a/litellm/llms/ovhcloud/audio_transcription/transformation.py b/litellm/llms/ovhcloud/audio_transcription/transformation.py index 7ff6dc986be..f49f31d7ecd 100644 --- a/litellm/llms/ovhcloud/audio_transcription/transformation.py +++ b/litellm/llms/ovhcloud/audio_transcription/transformation.py @@ -156,5 +156,17 @@ class OVHCloudAudioTranscriptionConfig(BaseAudioTranscriptionConfig): text = response_json.get("text") or response_json.get("transcript") or "" response = TranscriptionResponse(text=text) + # OVHCloud field migration (deadline: 2026-05-11): + # `duration` is replaced by `seconds` in STT responses. + # Prefer `seconds`, fall back to `duration`, normalize to `duration` + # so downstream consumers see a consistent key. + duration = ( + response_json["seconds"] + if "seconds" in response_json and response_json["seconds"] is not None + else response_json.get("duration") + ) + if duration is not None: + response_json["duration"] = duration + response._hidden_params = response_json return response diff --git a/litellm/llms/ovhcloud/chat/transformation.py b/litellm/llms/ovhcloud/chat/transformation.py index ae9271ddb16..62f51f1e9da 100644 --- a/litellm/llms/ovhcloud/chat/transformation.py +++ b/litellm/llms/ovhcloud/chat/transformation.py @@ -13,6 +13,7 @@ from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig from litellm.llms.ovhcloud.utils import OVHCloudException from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator from litellm.llms.base_llm.chat.transformation import BaseLLMException + from litellm.types.llms.openai import AllMessageValues @@ -98,10 +99,16 @@ class OVHCloudChatCompletionStreamingHandler(BaseModelResponseIterator): new_choices = [] for choice in chunk["choices"]: - if "delta" in choice and "reasoning" in choice["delta"]: - choice["delta"]["reasoning_content"] = choice["delta"].get( - "reasoning" - ) + if "delta" in choice: + delta = choice["delta"] + # OVHCloud field migration (deadline: 2026-05-11): + # `reasoning_content` is replaced by `reasoning`. + # Normalise to `reasoning_content` so downstream consumers + # see a consistent key during the transition window. + reasoning_new = delta.get("reasoning") + reasoning_legacy = delta.get("reasoning_content") + if reasoning_new is not None and reasoning_legacy is None: + delta["reasoning_content"] = reasoning_new new_choices.append(choice) return ModelResponseStream( diff --git a/litellm/main.py b/litellm/main.py index 051a82fdd19..52a256fdd05 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -1,3 +1,5 @@ +# LiteLLM main module: public completion, embedding, streaming, and moderation entrypoints. +# # +-----------------------------------------------+ # | | # | Give Feedback / Get Help | @@ -59,7 +61,13 @@ import litellm from litellm import client # Other utils are imported directly to avoid circular imports -from litellm.utils import exception_type, get_litellm_params, get_optional_params +from litellm.utils import ( + exception_type, + get_litellm_params, + get_optional_params, + peek_reasoning_summary_aliases, + strip_reasoning_summary_aliases_from_optional_params, +) # Logging is imported lazily when needed to avoid loading litellm_logging at import time if TYPE_CHECKING: @@ -946,6 +954,7 @@ def responses_api_bridge_check( web_search_options: Optional[OpenAIWebSearchOptions] = None, tools: Optional[List[Any]] = None, reasoning_effort: Optional[Any] = None, + reasoning_summary: Optional[Any] = None, ) -> Tuple[dict, str]: model_info: Dict[str, Any] = {} @@ -982,14 +991,23 @@ def responses_api_bridge_check( mode = "responses" model_info["mode"] = mode - # OpenAI/Azure gpt-5.4+ chat-completions calls with both tools + reasoning_effort - # must be bridged to Responses API. + # OpenAI/Azure GPT-5 chat-completions that need Responses-only fields (e.g. + # ``reasoningSummary`` in ``extra_body``) must be bridged; Chat Completions rejects + # those keys. + # + # - gpt-5.4+: tools + reasoning_effort (original) or any reasoning-summary alias. + # - Older GPT-5 names (e.g. ``gpt-5``, ``gpt-5.1``): bridge only when a reasoning + # summary alias is present with ``reasoning_effort`` (tools alone stay on chat). if ( custom_llm_provider in ("openai", "azure") - and OpenAIGPT5Config.is_model_gpt_5_4_plus_model(model) - and tools - and reasoning_effort is not None and model_info.get("mode") != "responses" + and OpenAIGPT5Config.is_model_gpt_5_model(model) + and not OpenAIGPT5Config.is_model_gpt_5_search_model(model) + and reasoning_effort is not None + and ( + reasoning_summary is not None + or (OpenAIGPT5Config.is_model_gpt_5_4_plus_model(model) and tools) + ) ): model_info["mode"] = "responses" model = model.replace("responses/", "") @@ -1634,8 +1652,10 @@ def completion( # type: ignore # noqa: PLR0915 ## RESPONSES API BRIDGE LOGIC ## - check if model has 'mode: responses' in litellm.model_cost map # Only run the second bridge check if the first one didn't already # detect responses mode (e.g. via the "responses/" prefix). The second - # check handles cases like gpt-5.4+ with tools+reasoning_effort that - # the first (early) check doesn't cover. + # check handles cases like gpt-5.4+ with tools+reasoning_effort or + # reasoningSummary/reasoning_summary without tools (AI SDK) that the first + # (early) check doesn't cover. + _reasoning_summary_for_bridge = peek_reasoning_summary_aliases(optional_params) if responses_api_model_info.get("mode") != "responses": responses_api_model_info, model = responses_api_bridge_check( model=model, @@ -1643,14 +1663,29 @@ def completion( # type: ignore # noqa: PLR0915 web_search_options=web_search_options, tools=tools, reasoning_effort=reasoning_effort, + reasoning_summary=_reasoning_summary_for_bridge, ) if responses_api_model_info.get("mode") == "responses": from litellm.completion_extras import responses_api_bridge + optional_params, rs_val = ( + strip_reasoning_summary_aliases_from_optional_params(optional_params) + ) + if isinstance(reasoning_effort, dict) and "summary" in reasoning_effort: - optional_params = dict(optional_params) optional_params["reasoning_effort"] = reasoning_effort + elif rs_val is not None: + eff = optional_params.get("reasoning_effort", reasoning_effort) + if isinstance(eff, dict): + optional_params["reasoning_effort"] = {**eff, "summary": rs_val} + elif eff is not None: + optional_params["reasoning_effort"] = { + "effort": eff, + "summary": rs_val, + } + else: + optional_params["reasoning_effort"] = {"summary": rs_val} return responses_api_bridge.completion( model=model, @@ -1669,6 +1704,16 @@ def completion( # type: ignore # noqa: PLR0915 encoding=_get_encoding(), stream=stream, ) + elif ( + custom_llm_provider == "openai" + and OpenAIGPT5Config.is_model_gpt_5_model(model) + ) or ( + custom_llm_provider == "azure" + and litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(model) + ): + optional_params, _ = strip_reasoning_summary_aliases_from_optional_params( + optional_params + ) if custom_llm_provider == "azure": # azure configs @@ -3813,7 +3858,33 @@ def completion( # type: ignore # noqa: PLR0915 ) bedrock_route = BedrockModelInfo.get_bedrock_route(model) - if bedrock_route == "converse": + if bedrock_route == "claude_platform": + provider_config = ProviderConfigManager.get_provider_chat_config( + model=model, + provider=LlmProviders.BEDROCK, + ) + model = BedrockModelInfo.get_claude_platform_model(model) + response = base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + custom_llm_provider="bedrock", + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, + client=client, + provider_config=provider_config, + ) + return response + elif bedrock_route == "converse": model = model.replace("converse/", "") response = bedrock_converse_chat_completion.completion( model=model, diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 1794cd14381..62691641234 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -12,7 +12,8 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.proxy._experimental.mcp_server.oauth_utils import ( TOKEN_NO_CACHE_HEADERS, - validate_loopback_redirect_uri, + get_request_base_url, + validate_trusted_redirect_uri, ) from litellm.proxy.auth.ip_address_utils import IPAddressUtils from litellm.proxy.common_utils.encrypt_decrypt_utils import ( @@ -29,51 +30,6 @@ router = APIRouter( ) -def get_request_base_url(request: Request) -> str: - """ - Get the base URL for the request, considering X-Forwarded-* headers. - - X-Forwarded-Proto / X-Forwarded-Host / X-Forwarded-Port are only honoured - when the request comes from a configured trusted proxy - (``use_x_forwarded_for`` enabled AND caller in ``mcp_trusted_proxy_ranges``). - Otherwise the request's literal ``base_url`` is returned, so an - untrusted caller cannot poison OAuth-discovery / redirect_uri values - by injecting headers. - - Args: - request: FastAPI Request object - - Returns: - The reconstructed base URL (e.g., "https://proxy.example.com") - """ - base_url = str(request.base_url).rstrip("/") - parsed = urlparse(base_url) - - if not IPAddressUtils.is_request_from_trusted_proxy(request): - return base_url - - x_forwarded_proto = request.headers.get("X-Forwarded-Proto") - x_forwarded_host = request.headers.get("X-Forwarded-Host") - x_forwarded_port = request.headers.get("X-Forwarded-Port") - - scheme = x_forwarded_proto if x_forwarded_proto else parsed.scheme - - if x_forwarded_host: - # X-Forwarded-Host may already include port (e.g., "example.com:8080") - if ":" in x_forwarded_host and not x_forwarded_host.startswith("["): - netloc = x_forwarded_host - elif x_forwarded_port: - netloc = f"{x_forwarded_host}:{x_forwarded_port}" - else: - netloc = x_forwarded_host - else: - netloc = parsed.netloc - if x_forwarded_port and ":" not in netloc: - netloc = f"{netloc}:{x_forwarded_port}" - - return urlunparse((scheme, netloc, parsed.path, "", "", "")) - - def encode_state_with_base_url( base_url: str, original_state: str, @@ -127,12 +83,14 @@ def decode_state_hash(encrypted_state: str) -> dict: return state_data -def _get_validated_client_redirect_uri(state_data: Dict[str, Any]) -> str: - """Return a loopback client redirect URI from OAuth state.""" +def _get_validated_client_redirect_uri( + request: Request, state_data: Dict[str, Any] +) -> str: + """Return a trusted (same-origin or loopback) client redirect URI from OAuth state.""" redirect_uri = state_data.get("client_redirect_uri") or state_data.get("base_url") if not redirect_uri or not isinstance(redirect_uri, str): raise HTTPException(status_code=400, detail="Invalid redirect URI") - validate_loopback_redirect_uri(redirect_uri) + validate_trusted_redirect_uri(request, redirect_uri) return redirect_uri @@ -338,12 +296,12 @@ async def authorize_with_server( status_code=400, detail="MCP server authorization url is not set" ) - # Loopback-only redirect_uri. The URI is encrypted into the OAuth - # state and decoded on /callback to redirect the user back; a non- - # loopback URI would be an open-redirect + code-theft primitive - # (VERIA-57 root cause B). MCP clients are native apps — loopback is - # the spec-compliant callback pattern. - validate_loopback_redirect_uri(redirect_uri) + # Loopback OR same-origin redirect_uri. The URI is encrypted into the + # OAuth state and decoded on /callback to redirect the user back; + # restricting to trusted origins blocks the open-redirect + + # code-theft primitive (VERIA-57 root cause B). Loopback supports + # native MCP clients; same-origin supports the proxy's own UI callback. + validate_trusted_redirect_uri(request, redirect_uri) parsed = urlparse(redirect_uri) base_url = urlunparse(parsed._replace(query="")) request_base_url = get_request_base_url(request) @@ -660,17 +618,18 @@ async def token_endpoint( @router.get("/callback") -async def callback(code: str, state: str): +async def callback(request: Request, code: str, state: str): try: state_data = decode_state_hash(state) original_state = state_data["original_state"] - # Re-validate loopback at the sink. /authorize rejects non-loopback + # Re-validate at the sink. /authorize rejects untrusted # redirect_uri before encoding into state, but encrypted states # minted before that check was added have no expiry and remain - # valid indefinitely. Validating here blocks the open-redirect + - # code-theft primitive even for pre-fix states. - redirect_uri = _get_validated_client_redirect_uri(state_data) + # valid indefinitely. Validating here (same-origin OR loopback) + # blocks the open-redirect + code-theft primitive even for pre-fix + # states while allowing the UI's same-origin callback to work. + redirect_uri = _get_validated_client_redirect_uri(request, state_data) params = {"code": code, "state": original_state} complete_returned_url = _append_query_params(redirect_uri, params) diff --git a/litellm/proxy/_experimental/mcp_server/oauth_utils.py b/litellm/proxy/_experimental/mcp_server/oauth_utils.py index b13cf83058c..343d1bee613 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_utils.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_utils.py @@ -2,15 +2,63 @@ (BYOK + discoverable / pass-through OAuth proxy).""" from ipaddress import ip_address -from urllib.parse import urlparse +from urllib.parse import urlparse, urlunparse -from fastapi import HTTPException +from fastapi import HTTPException, Request + +from litellm._logging import verbose_logger +from litellm.proxy.auth.ip_address_utils import IPAddressUtils # RFC 6749 §5.1 / OAuth 2.1 draft-15 §4.1.3: token-endpoint responses # must not be cached — both success and error bodies may reveal secrets. TOKEN_NO_CACHE_HEADERS = {"Cache-Control": "no-store", "Pragma": "no-cache"} +def get_request_base_url(request: Request) -> str: + """ + Get the base URL for the request, considering X-Forwarded-* headers. + + X-Forwarded-Proto / X-Forwarded-Host / X-Forwarded-Port are only honoured + when the request comes from a configured trusted proxy + (``use_x_forwarded_for`` enabled AND caller in ``mcp_trusted_proxy_ranges``). + Otherwise the request's literal ``base_url`` is returned, so an + untrusted caller cannot poison OAuth-discovery / redirect_uri values + by injecting headers. + + Args: + request: FastAPI Request object + + Returns: + The reconstructed base URL (e.g., "https://proxy.example.com") + """ + base_url = str(request.base_url).rstrip("/") + parsed = urlparse(base_url) + + if not IPAddressUtils.is_request_from_trusted_proxy(request): + return base_url + + x_forwarded_proto = request.headers.get("X-Forwarded-Proto") + x_forwarded_host = request.headers.get("X-Forwarded-Host") + x_forwarded_port = request.headers.get("X-Forwarded-Port") + + scheme = x_forwarded_proto if x_forwarded_proto else parsed.scheme + + if x_forwarded_host: + # X-Forwarded-Host may already include port (e.g., "example.com:8080") + if ":" in x_forwarded_host and not x_forwarded_host.startswith("["): + netloc = x_forwarded_host + elif x_forwarded_port: + netloc = f"{x_forwarded_host}:{x_forwarded_port}" + else: + netloc = x_forwarded_host + else: + netloc = parsed.netloc + if x_forwarded_port and ":" not in netloc: + netloc = f"{netloc}:{x_forwarded_port}" + + return urlunparse((scheme, netloc, parsed.path, "", "", "")) + + def validate_loopback_redirect_uri(redirect_uri: str) -> None: """Require a loopback ``redirect_uri`` (OAuth 2.1 §4.1.2.1 + RFC 8252 §7.3 native-app pattern). MCP clients are native apps that listen on @@ -46,3 +94,60 @@ def validate_loopback_redirect_uri(redirect_uri: str) -> None: # don't let it bubble up as a 500. pass raise HTTPException(status_code=400, detail="invalid_request") + + +def validate_trusted_redirect_uri(request: Request, redirect_uri: str) -> None: + """Accept same-origin (proxy's own origin) OR loopback ``redirect_uri``. + + Same-origin is required for the LiteLLM UI's OAuth flow: the UI + redirects to ``/ui/mcp/oauth/callback`` which is not loopback + but is on the proxy's own trusted HTTPS origin. An attacker cannot + host content on the proxy's own origin without already owning the + proxy, so the open-redirect / code-theft primitive that motivated + :func:`validate_loopback_redirect_uri` does not apply here. + + Loopback continues to be accepted for native MCP clients (per + OAuth 2.1 §4.1.2.1 + RFC 8252 §7.3). + + Use this in the discoverable OAuth proxy endpoints that serve both + native clients and the proxy's own UI. BYOK endpoints that only + support native clients should keep + :func:`validate_loopback_redirect_uri`. + """ + try: + parsed = urlparse(redirect_uri) + except ValueError: + raise HTTPException(status_code=400, detail="invalid_request") + if parsed.scheme not in ("http", "https"): + raise HTTPException(status_code=400, detail="invalid_request") + if parsed.fragment: + raise HTTPException(status_code=400, detail="invalid_request") + + # Same-origin: scheme + netloc (host[:port]) must match the proxy's + # own base URL at this request (honouring trusted X-Forwarded-*). + try: + proxy_base = urlparse(get_request_base_url(request)) + if ( + parsed.netloc + and parsed.scheme == proxy_base.scheme + and parsed.netloc.lower() == proxy_base.netloc.lower() + ): + return + except Exception as exc: + # If we can't determine the proxy's origin, fall through to + # loopback. Log so the failure is diagnosable in production. + verbose_logger.warning( + "validate_trusted_redirect_uri: could not determine proxy origin, " + "falling back to loopback-only check. error=%s", + exc, + ) + + host = (parsed.hostname or "").lower() + if host == "localhost": + return + try: + if ip_address(host).is_loopback: + return + except ValueError: + pass + raise HTTPException(status_code=400, detail="invalid_request") diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index d58612d5054..6128f485748 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -239,6 +239,7 @@ class KeyManagementRoutes(str, enum.Enum): KEY_BLOCK = "/key/block" KEY_UNBLOCK = "/key/unblock" KEY_BULK_UPDATE = "/key/bulk_update" + TEAM_KEY_BULK_UPDATE = "/team/key/bulk_update" KEY_RESET_SPEND = "/key/{key_id}/reset_spend" # info and health routes @@ -540,6 +541,7 @@ class LiteLLMRoutes(enum.Enum): KeyManagementRoutes.KEY_BLOCK.value, KeyManagementRoutes.KEY_UNBLOCK.value, KeyManagementRoutes.KEY_BULK_UPDATE.value, + KeyManagementRoutes.TEAM_KEY_BULK_UPDATE.value, KeyManagementRoutes.TEAM_DAILY_ACTIVITY.value, KeyManagementRoutes.SPEND_LOGS.value, KeyManagementRoutes.KEY_RESET_SPEND.value, @@ -675,6 +677,10 @@ class LiteLLMRoutes(enum.Enum): "/global/activity", "/global/activity/model", "/global/activity/cache_hits", + # Tag usage endpoints scope internal users to tags produced by + # their own keys in tag_management_endpoints.py. + "/tag/daily/activity", + "/tag/list", "/v1/models/{model_id}", "/models/{model_id}", "/guardrails/list", @@ -687,7 +693,16 @@ class LiteLLMRoutes(enum.Enum): + compliance_check_routes ) - internal_user_view_only_routes = spend_tracking_routes + internal_user_view_only_routes = ( + spend_tracking_routes + + compliance_check_routes + + [ + # Tag usage endpoints scope internal viewers to tags produced by + # their own keys in tag_management_endpoints.py. + "/tag/daily/activity", + "/tag/list", + ] + ) self_managed_routes = [ "/team/member_add", @@ -4346,10 +4361,16 @@ class JWTRoutingOverride(BaseModel): A rule matches when all provided selectors match token claims. If matched, request is routed to the configured auth path. + + Wildcard selectors use shell-style patterns (* and ?) and are matched with + case-sensitive semantics; use the same casing your IdP emits in JWT claims. + Space-delimited tokenization applies only to the ``scope`` claim (OAuth/OIDC + scope strings), not to ``iss``, ``aud``, or ``client_id``. """ iss: Union[str, List[str]] client_id: Optional[Union[str, List[str]]] = None + scope: Optional[Union[str, List[str]]] = None aud: Optional[Union[str, List[str]]] = None path: Literal["oauth2"] = "oauth2" diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index d1fd5818f35..5b116c142b2 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -224,6 +224,41 @@ class JWTHandler: return [] + def get_all_jwt_team_ids(self, token: dict) -> List[str]: + """ + Return team IDs from both the plural ``team_ids_jwt_field`` and the + singular ``team_id_jwt_field`` claim (string or list of strings), as a + deduplicated list preserving plural-first order. + + Membership-reconciliation paths (SSO callback, JWT-bearer sync) need + to consider both claim shapes. Reading only the plural field — as + callers historically did — silently dropped users whose IdP populates + the singular field, which is what Okta and Auth0 default to when a + user has a single primary team. + + This intentionally does NOT consult ``team_id_default``: that fallback + is a property of how the JWT-bearer auth flow resolves a single + request-bound team, not of the token's claims. Callers that want the + default-team behavior should still go through ``get_team_id``. + """ + team_ids: List[str] = list(self.get_team_ids_from_jwt(token)) + if self.litellm_jwtauth.team_id_jwt_field is not None: + singular = get_nested_value( + data=token, + key_path=self.litellm_jwtauth.team_id_jwt_field, + default=None, + ) + if isinstance(singular, list): + for item in singular: + if item is None: + continue + sid = str(item) + if sid and sid not in team_ids: + team_ids.append(sid) + elif singular and str(singular) not in team_ids: + team_ids.append(str(singular)) + return team_ids + def get_end_user_id( self, token: dict, default_value: Optional[str] ) -> Optional[str]: diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py index 8bcfbb67539..a9c36dfc512 100644 --- a/litellm/proxy/auth/route_checks.py +++ b/litellm/proxy/auth/route_checks.py @@ -50,6 +50,7 @@ _PROXY_ADMIN_VIEW_ONLY_BLOCKED_ROUTES = frozenset( KeyManagementRoutes.KEY_BLOCK.value, KeyManagementRoutes.KEY_UNBLOCK.value, KeyManagementRoutes.KEY_BULK_UPDATE.value, + KeyManagementRoutes.TEAM_KEY_BULK_UPDATE.value, ] ) @@ -671,6 +672,7 @@ class RouteChecks: "/key/service-account/generate", "/key/block", "/key/unblock", + "/team/key/bulk_update", ] ) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 4778549befc..39f0af19484 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -8,6 +8,7 @@ Returns a UserAPIKeyAuth object if the API key is valid """ import asyncio +import fnmatch import re import secrets from datetime import datetime, timezone @@ -183,22 +184,54 @@ def _get_bearer_token_or_received_api_key(api_key: str) -> str: def _routing_selector_matches_claim( - selector_value: Optional[Any], claim_value: Optional[Any] + selector_value: Optional[Any], + claim_value: Optional[Any], + *, + split_space_delimited: bool = False, ) -> bool: if selector_value is None: return True - selector_list = ( + selector_list: List[str] = ( [str(v) for v in selector_value] if isinstance(selector_value, list) else [str(selector_value)] ) + if claim_value is None: + return False + if isinstance(claim_value, list): claim_list = [str(v) for v in claim_value] - return any(v in claim_list for v in selector_list) + elif ( + split_space_delimited + and isinstance(claim_value, str) + and " " in claim_value.strip() + ): + # OAuth/OIDC often sends scope as a single space-delimited string. Only split + # for the scope selector: iss/aud/client_id must stay exact full-string match + # on unverified claims (see routing override security review). The elif guard + # (`" " in claim_value.strip()`) ensures at least two non-empty tokens survive. + claim_list = [v for v in claim_value.strip().split(" ") if v] + else: + claim_list = [str(claim_value)] - return str(claim_value) in selector_list if claim_value is not None else False + def _selector_matches_claim(selector: str, claim: str) -> bool: + # NOTE: wildcard matching is case-sensitive (fnmatch.fnmatchcase). + if "*" in selector or "?" in selector: + # Without scope splitting, do not let `*` span whitespace: a malformed + # iss like "trusted.example.com evil.com" must not match "trusted.*". + # Scope uses split_space_delimited so each claim token is checked separately. + if not split_space_delimited and any(ch.isspace() for ch in claim): + return False + return fnmatch.fnmatchcase(claim, selector) + return selector == claim + + return any( + _selector_matches_claim(selector=s, claim=c) + for s in selector_list + for c in claim_list + ) def _matches_routing_override( @@ -209,6 +242,11 @@ def _matches_routing_override( and _routing_selector_matches_claim( override.client_id, token_claims.get("client_id") ) + and _routing_selector_matches_claim( + override.scope, + token_claims.get("scope"), + split_space_delimited=True, + ) and _routing_selector_matches_claim(override.aud, token_claims.get("aud")) ) diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index 0928ce914da..71537cc62e6 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -2,7 +2,7 @@ import asyncio import json import time from datetime import datetime, timezone -from typing import Any, List, Literal, Optional, Union +from typing import Any, Callable, List, Literal, Optional, Union import litellm from litellm._logging import verbose_proxy_logger @@ -83,93 +83,139 @@ class ResetBudgetJob: "Failed to reset spend counter %s: %s", counter_key, e ) + @staticmethod + async def _invalidate_user_api_key_cache_entry(cache_key: str) -> None: + """Drop a stale management-cache entry so the next read fetches from DB. + + Some entity types (notably tags and end-users) are not handled by + SpendCounterReseed.from_db, so when a spend counter expires the + budget check falls back to ``cached_obj.spend``. If that cached + object lingers in ``user_api_key_cache`` past a budget reset, the + stale ``.spend`` keeps the entity blocked indefinitely. Deleting + the cache entry forces the next auth-time fetch to reload the + zeroed row from Postgres. + """ + try: + from litellm.proxy.proxy_server import user_api_key_cache + + await user_api_key_cache.async_delete_cache(key=cache_key) + except Exception as e: + verbose_proxy_logger.warning( + "Failed to invalidate user_api_key_cache entry %s: %s", + cache_key, + e, + ) + + async def _cascade_reset_spend_for_budget_link( + self, + budgets_to_reset: List[LiteLLM_BudgetTableFull], + table: Any, + counter_key_fn: Callable[[Any], str], + log_subject: str, + extra_where: Optional[dict] = None, + cache_key_fn: Optional[Callable[[Any], str]] = None, + ): + """ + Generic cascade: zero spend on rows whose budget_id is in the reset set. + + ``cache_key_fn`` is optional: when provided, after the DB update each + matching row's entry in ``user_api_key_cache`` is also dropped. This + is required for entities whose spend counter is read with the cached + object's ``.spend`` as fallback (tags, end-users) — otherwise the + stale cached object pins enforcement to the pre-reset spend until + its TTL expires. + """ + budget_ids = [b.budget_id for b in budgets_to_reset if b.budget_id is not None] + if not budget_ids: + return + + where: dict = {"budget_id": {"in": budget_ids}} + if extra_where: + where.update(extra_where) + + try: + rows = await table.find_many(where=where) + except Exception as e: + rows = [] + verbose_proxy_logger.warning( + "Failed to fetch %s for counter invalidation: %s", log_subject, e + ) + + update_result = await table.update_many(where=where, data={"spend": 0}) + + for row in rows: + await self._invalidate_spend_counter(counter_key_fn(row)) + if cache_key_fn is not None: + await self._invalidate_user_api_key_cache_entry(cache_key_fn(row)) + + return update_result + async def reset_budget_for_litellm_team_members( self, budgets_to_reset: List[LiteLLM_BudgetTableFull] ): """ Resets the budget for all LiteLLM Team Members if their budget has expired """ - budget_ids = [ - budget.budget_id - for budget in budgets_to_reset - if budget.budget_id is not None - ] - - try: - memberships = await self.prisma_client.db.litellm_teammembership.find_many( - where={"budget_id": {"in": budget_ids}} - ) - except Exception as e: - memberships = [] - verbose_proxy_logger.warning( - "Failed to fetch team memberships for counter invalidation: %s", e - ) - - update_result = await self.prisma_client.db.litellm_teammembership.update_many( - where={"budget_id": {"in": budget_ids}}, - data={ - "spend": 0, - }, + return await self._cascade_reset_spend_for_budget_link( + budgets_to_reset=budgets_to_reset, + table=self.prisma_client.db.litellm_teammembership, + counter_key_fn=lambda m: f"spend:team_member:{m.user_id}:{m.team_id}", + log_subject="team memberships", ) - for m in memberships: - await self._invalidate_spend_counter( - f"spend:team_member:{m.user_id}:{m.team_id}" - ) - - return update_result - async def reset_budget_for_keys_linked_to_budgets( self, budgets_to_reset: List[LiteLLM_BudgetTableFull] ): """ Resets the spend for keys linked to budget tiers that are being reset. - This handles keys that have budget_id but no budget_duration set on the key - itself. Keys with budget_id rely on their linked budget tier's reset schedule - rather than having their own budget_duration. - - Keys that have their own budget_duration are already handled by - reset_budget_for_litellm_keys() and are excluded here to avoid - double-resetting. + Excludes keys with their own budget_duration; those are reset by + reset_budget_for_litellm_keys() to avoid double-resetting. """ - budget_ids = [ - budget.budget_id - for budget in budgets_to_reset - if budget.budget_id is not None - ] - if not budget_ids: - return - - where_clause: dict = { - "budget_id": {"in": budget_ids}, - "budget_duration": None, # only keys without their own reset schedule - "spend": {"gt": 0}, # only reset keys that have accumulated spend - } - - try: - keys = await self.prisma_client.db.litellm_verificationtoken.find_many( - where=where_clause - ) - except Exception as e: - keys = [] - verbose_proxy_logger.warning( - "Failed to fetch keys for counter invalidation: %s", e - ) - - update_result = ( - await self.prisma_client.db.litellm_verificationtoken.update_many( - where=where_clause, - data={ - "spend": 0, - }, - ) + return await self._cascade_reset_spend_for_budget_link( + budgets_to_reset=budgets_to_reset, + table=self.prisma_client.db.litellm_verificationtoken, + counter_key_fn=lambda k: f"spend:key:{k.token}", + log_subject="keys", + extra_where={"budget_duration": None, "spend": {"gt": 0}}, ) - for k in keys: - await self._invalidate_spend_counter(f"spend:key:{k.token}") + async def reset_budget_for_orgs_linked_to_budgets( + self, budgets_to_reset: List[LiteLLM_BudgetTableFull] + ): + """ + Resets the spend for orgs linked to budget tiers that are being reset. + """ + return await self._cascade_reset_spend_for_budget_link( + budgets_to_reset=budgets_to_reset, + table=self.prisma_client.db.litellm_organizationtable, + counter_key_fn=lambda o: f"spend:org:{o.organization_id}", + log_subject="orgs", + extra_where={"spend": {"gt": 0}}, + ) - return update_result + async def reset_budget_for_tags_linked_to_budgets( + self, budgets_to_reset: List[LiteLLM_BudgetTableFull] + ): + """ + Resets the spend for tags linked to budget tiers that are being reset. + + Also drops each tag's ``user_api_key_cache`` entry so the next + ``_tag_max_budget_check`` reloads the zeroed row from the DB. + ``SpendCounterReseed.from_db`` intentionally returns ``None`` for + tags, so the budget check falls back to the cached + ``LiteLLM_TagTable.spend`` once the spend counter expires; without + this invalidation, that stale ``.spend`` keeps the tag over-budget + indefinitely. + """ + return await self._cascade_reset_spend_for_budget_link( + budgets_to_reset=budgets_to_reset, + table=self.prisma_client.db.litellm_tagtable, + counter_key_fn=lambda t: f"spend:tag:{t.tag_name}", + log_subject="tags", + extra_where={"spend": {"gt": 0}}, + cache_key_fn=lambda t: f"tag:{t.tag_name}", + ) async def reset_budget_for_litellm_budget_table(self): """ @@ -237,6 +283,14 @@ class ResetBudgetJob: budgets_to_reset=budgets_to_reset ) + await self.reset_budget_for_orgs_linked_to_budgets( + budgets_to_reset=budgets_to_reset + ) + + await self.reset_budget_for_tags_linked_to_budgets( + budgets_to_reset=budgets_to_reset + ) + if endusers_to_reset is not None and len(endusers_to_reset) > 0: for enduser in endusers_to_reset: try: diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 3697e498bb4..7611c9c9692 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -1131,10 +1131,12 @@ class DBSpendUpdateWriter: timeout=timedelta(seconds=60) ) as transaction: async with transaction.batch_() as batcher: - for ( - user_id, - response_cost, - ) in user_list_transactions.items(): + # Sort by ID for consistent lock ordering across pods to prevent deadlocks. + # batch_() issues statements sequentially within the tx, so iteration + # order = lock acquisition order. + for user_id, response_cost in sorted( + user_list_transactions.items() + ): batcher.litellm_usertable.update_many( where={"user_id": user_id}, data={"spend": {"increment": response_cost}}, @@ -1186,10 +1188,10 @@ class DBSpendUpdateWriter: timeout=timedelta(seconds=60) ) as transaction: async with transaction.batch_() as batcher: - for ( - token, - response_cost, - ) in key_list_transactions.items(): + # Sort by token for consistent lock ordering across pods to prevent deadlocks. + for token, response_cost in sorted( + key_list_transactions.items() + ): batcher.litellm_verificationtoken.update_many( # 'update_many' prevents error from being raised if no row exists where={"token": token}, data={ @@ -1230,10 +1232,10 @@ class DBSpendUpdateWriter: timeout=timedelta(seconds=60) ) as transaction: async with transaction.batch_() as batcher: - for ( - team_id, - response_cost, - ) in team_list_transactions.items(): + # Sort by team_id for consistent lock ordering across pods to prevent deadlocks. + for team_id, response_cost in sorted( + team_list_transactions.items() + ): verbose_proxy_logger.debug( "Updating spend for team id={} by {}".format( team_id, response_cost @@ -1288,10 +1290,11 @@ class DBSpendUpdateWriter: timeout=timedelta(seconds=60) ) as transaction: async with transaction.batch_() as batcher: - for ( - key, - response_cost, - ) in team_member_list_transactions.items(): + # Sort by composite key for consistent lock ordering across pods to prevent deadlocks. + # Key format "team_id::::user_id::" makes the string sort equivalent to sorting by (team_id, user_id). + for key, response_cost in sorted( + team_member_list_transactions.items() + ): # key is "team_id::::user_id::" team_id = key.split("::")[1] user_id = key.split("::")[3] @@ -1348,10 +1351,10 @@ class DBSpendUpdateWriter: timeout=timedelta(seconds=60) ) as transaction: async with transaction.batch_() as batcher: - for ( - org_id, - response_cost, - ) in org_list_transactions.items(): + # Sort by org_id for consistent lock ordering across pods to prevent deadlocks. + for org_id, response_cost in sorted( + org_list_transactions.items() + ): batcher.litellm_organizationtable.update_many( # 'update_many' prevents error from being raised if no row exists where={"organization_id": org_id}, data={"spend": {"increment": response_cost}}, @@ -1439,7 +1442,10 @@ class DBSpendUpdateWriter: timeout=timedelta(seconds=60) ) as transaction: async with transaction.batch_() as batcher: - for entity_id, response_cost in transactions.items(): + # Sort by entity_id for consistent lock ordering across pods to prevent deadlocks. + for entity_id, response_cost in sorted( + transactions.items() + ): verbose_proxy_logger.debug( f"Updating spend for {entity_name} {where_field}={entity_id} by {response_cost}" ) diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index 5351391e5e1..e55f3b6e16b 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -242,6 +242,10 @@ async def list_guardrails_v2( gid = guardrail.get("guardrail_id") if gid in seen_guardrail_ids: continue + # Skip stale DB-backed entries — the DB row was deleted (likely by + # another pod) and reconciliation hasn't fired yet on this pod. + if gid is not None and IN_MEMORY_GUARDRAIL_HANDLER.get_source(gid) == "db": + continue if not is_admin: g_team_id = guardrail.get("team_id") if g_team_id is not None and g_team_id not in caller_team_ids: @@ -360,7 +364,7 @@ async def create_guardrail( try: IN_MEMORY_GUARDRAIL_HANDLER.initialize_guardrail( - guardrail=cast(Guardrail, result) + guardrail=cast(Guardrail, result), source="db" ) verbose_proxy_logger.info( f"Immediate sync: Successfully initialized guardrail '{guardrail_name}' (ID: {guardrail_id})" @@ -1017,7 +1021,7 @@ async def approve_guardrail_submission( } try: IN_MEMORY_GUARDRAIL_HANDLER.initialize_guardrail( - guardrail=cast(Guardrail, guardrail_dict) + guardrail=cast(Guardrail, guardrail_dict), source="db" ) verbose_proxy_logger.info( "Approved guardrail %s (ID: %s) and initialized in memory", @@ -1295,10 +1299,18 @@ async def get_guardrail_info(guardrail_id: str): guardrail_id=guardrail_id, prisma_client=prisma_client ) if result is None: - result = IN_MEMORY_GUARDRAIL_HANDLER.get_guardrail_by_id( + in_memory = IN_MEMORY_GUARDRAIL_HANDLER.get_guardrail_by_id( guardrail_id=guardrail_id ) - guardrail_definition_location = GUARDRAIL_DEFINITION_LOCATION.CONFIG + # Only return config-loaded entries here. A DB-backed entry that's + # missing from the DB is stale (deleted on another pod, awaiting + # reconciliation on this one) and must surface as 404. + if ( + in_memory is not None + and IN_MEMORY_GUARDRAIL_HANDLER.get_source(guardrail_id) == "config" + ): + result = in_memory + guardrail_definition_location = GUARDRAIL_DEFINITION_LOCATION.CONFIG if result is None: raise HTTPException( diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 838fb2e01ad..aafcc5f1819 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -3,7 +3,7 @@ import importlib import os from datetime import datetime, timezone -from typing import Any, Dict, List, Optional, Type, cast +from typing import Any, Dict, List, Literal, Optional, Set, Type, cast import litellm from litellm import Router @@ -403,11 +403,19 @@ class InMemoryGuardrailHandler: Guardrail id to CustomGuardrail object mapping """ + self._sources: Dict[str, Literal["db", "config"]] = {} + """ + Guardrail id to provenance marker. "db" entries are reconciled against + the DB on each polling tick; "config" entries are owned by proxy_config.yaml + and never deleted by reconciliation. + """ + def initialize_guardrail( self, guardrail: Guardrail, config_file_path: Optional[str] = None, llm_router: Optional["Router"] = None, + source: Literal["db", "config"] = "config", ) -> Optional[Guardrail]: """ Initialize a guardrail from a dictionary and add it to the litellm callback manager @@ -420,6 +428,10 @@ class InMemoryGuardrailHandler: verbose_proxy_logger.debug( "guardrail_id already exists in IN_MEMORY_GUARDRAILS" ) + # Honor the caller's source even on the early-return path so a + # racing polling tick or a hot-reload of config can correct an + # entry's provenance. + self._sources[guardrail_id] = source return self.IN_MEMORY_GUARDRAILS[guardrail_id] custom_guardrail_callback: Optional[CustomGuardrail] = None @@ -497,6 +509,7 @@ class InMemoryGuardrailHandler: # store references to the guardrail in memory self.IN_MEMORY_GUARDRAILS[guardrail_id] = parsed_guardrail self.guardrail_id_to_custom_guardrail[guardrail_id] = custom_guardrail_callback + self._sources[guardrail_id] = source return parsed_guardrail @@ -557,7 +570,10 @@ class InMemoryGuardrailHandler: return _guardrail_callback def update_in_memory_guardrail( - self, guardrail_id: str, guardrail: Guardrail + self, + guardrail_id: str, + guardrail: Guardrail, + source: Literal["db", "config"] = "db", ) -> None: """ Update a guardrail in memory @@ -566,6 +582,7 @@ class InMemoryGuardrailHandler: - updates the guardrail params in litellm.callback_manager """ self.IN_MEMORY_GUARDRAILS[guardrail_id] = guardrail + self._sources[guardrail_id] = source custom_guardrail_callback = self.guardrail_id_to_custom_guardrail.get( guardrail_id @@ -584,6 +601,7 @@ class InMemoryGuardrailHandler: """ # Remove from in-memory storage self.IN_MEMORY_GUARDRAILS.pop(guardrail_id, None) + self._sources.pop(guardrail_id, None) # Remove the callback from litellm.callbacks custom_guardrail_callback = self.guardrail_id_to_custom_guardrail.pop( @@ -608,6 +626,34 @@ class InMemoryGuardrailHandler: """ return self.IN_MEMORY_GUARDRAILS.get(guardrail_id) + def get_source(self, guardrail_id: str) -> Optional[Literal["db", "config"]]: + """ + Return the provenance of an in-memory guardrail. + """ + return self._sources.get(guardrail_id) + + def reconcile_db_guardrails(self, db_guardrail_ids: Set[str]) -> List[str]: + """ + Drop in-memory entries that originated from the DB but are no longer + present in db_guardrail_ids. Config-loaded guardrails are never touched. + + Called by the periodic DB polling tick so that a guardrail deleted + on another pod is eventually purged from this pod's memory + callbacks. + """ + stale_ids = [ + guardrail_id + for guardrail_id, source in self._sources.items() + if source == "db" and guardrail_id not in db_guardrail_ids + ] + for guardrail_id in stale_ids: + verbose_proxy_logger.info( + "Reconcile: removing stale DB-backed guardrail '%s' from memory " + "(deleted in DB by another pod)", + guardrail_id, + ) + self.delete_in_memory_guardrail(guardrail_id) + return stale_ids + def _has_guardrail_params_changed( self, guardrail_id: str, new_guardrail: Guardrail ) -> bool: @@ -661,7 +707,10 @@ class InMemoryGuardrailHandler: return len(changed_fields) > 0 def reinitialize_guardrail( - self, guardrail: Guardrail, config_file_path: Optional[str] = None + self, + guardrail: Guardrail, + config_file_path: Optional[str] = None, + source: Literal["db", "config"] = "config", ) -> Optional[Guardrail]: """ Force re-initialization of a guardrail even if it exists in memory. @@ -680,7 +729,7 @@ class InMemoryGuardrailHandler: # Initialize fresh (will add new callback to litellm.callbacks) return self.initialize_guardrail( - guardrail=guardrail, config_file_path=config_file_path + guardrail=guardrail, config_file_path=config_file_path, source=source ) def sync_guardrail_from_db( @@ -701,9 +750,15 @@ class InMemoryGuardrailHandler: f"Guardrail '{guardrail_name}' (ID: {guardrail_id}) params changed, re-initializing..." ) return self.reinitialize_guardrail( - guardrail=guardrail, config_file_path=config_file_path + guardrail=guardrail, + config_file_path=config_file_path, + source="db", ) + # Params unchanged but the entry is still DB-backed; make sure the + # source marker reflects that even if it was previously set differently + # (e.g. a config entry whose UUID later collided with a DB row). + self._sources[guardrail_id] = "db" return self.IN_MEMORY_GUARDRAILS.get(guardrail_id) diff --git a/litellm/proxy/guardrails/init_guardrails.py b/litellm/proxy/guardrails/init_guardrails.py index d742cc223b4..83f1281dc02 100644 --- a/litellm/proxy/guardrails/init_guardrails.py +++ b/litellm/proxy/guardrails/init_guardrails.py @@ -30,6 +30,7 @@ def init_guardrails_v2( guardrail=cast(Guardrail, guardrail), config_file_path=config_file_path, llm_router=llm_router, + source="config", ) if initialized_guardrail: guardrail_list.append(initialized_guardrail) diff --git a/litellm/proxy/health_check_utils/shared_health_check_manager.py b/litellm/proxy/health_check_utils/shared_health_check_manager.py index 5b8370fece8..ad58bc7e286 100644 --- a/litellm/proxy/health_check_utils/shared_health_check_manager.py +++ b/litellm/proxy/health_check_utils/shared_health_check_manager.py @@ -253,27 +253,63 @@ class SharedHealthCheckManager: # Always release the lock await self.release_health_check_lock() else: - # Lock not acquired, wait briefly and try to get cached results + # If Redis is not configured, skip polling — there is no cache + # to wait for. + if self.redis_cache is None: + return await perform_health_check( + model_list=model_list, + details=details, + max_concurrency=max_concurrency, + ) + + # Lock not acquired — poll for cached results until the lock + # holder finishes or the lock expires, rather than falling back + # to a redundant local health check after only 2 seconds. verbose_proxy_logger.debug( "Pod %s waiting for other pod to complete health check", self.pod_id ) - # Wait a bit for the other pod to complete - await asyncio.sleep(2) + poll_interval = 5 # seconds between cache checks + max_wait = self.lock_ttl # wait at most as long as the lock can live + elapsed = 0 - # Try to get cached results again - cached_results = await self.get_cached_health_check_results() - if cached_results is not None: - return ( - cached_results.get("healthy_endpoints", []), - cached_results.get("unhealthy_endpoints", []), - {}, - ) + while elapsed < max_wait: + await asyncio.sleep(poll_interval) + elapsed += poll_interval - # Still no cache, fall back to local health check + cached_results = await self.get_cached_health_check_results() + if cached_results is not None: + verbose_proxy_logger.info( + "Pod %s using cached health check results after waiting %ds", + self.pod_id, + elapsed, + ) + return ( + cached_results.get("healthy_endpoints", []), + cached_results.get("unhealthy_endpoints", []), + {}, + ) + + # Check if the lock is still held — if it was released without + # caching (e.g. the holder crashed), stop waiting early. + try: + lock_key = self.get_health_check_lock_key() + current_owner = await self.redis_cache.async_get_cache(lock_key) + if current_owner is None: + verbose_proxy_logger.debug( + "Pod %s detected lock released without cache, stopping wait", + self.pod_id, + ) + break + except Exception: + # Redis hiccup — continue polling rather than crashing out + pass + + # Exhausted wait — fall back to local health check verbose_proxy_logger.warning( - "Pod %s falling back to local health check (no cache available)", + "Pod %s falling back to local health check after waiting %ds (no cache available)", self.pod_id, + elapsed, ) return await perform_health_check( diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 096e23e673d..2c857c0fde8 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -1742,29 +1742,54 @@ async def test_model_connection( # Look up model configuration from router if model name is provided # This gets the litellm_params from proxy config (with resolved env vars) config_litellm_params: dict = {} - if model_name and llm_router is not None: + if llm_router is not None: + # Prefer disambiguation by deployment id (`model_info.id`) when + # the caller supplies it. This is required when multiple + # deployments share a `model_name` (e.g. wildcard `openai/*` + # with multiple `api_base` values for failover): the UI's + # "Test Connection" button targets a specific row, and that + # row's id is the only thing that uniquely identifies which + # deployment to probe. Without this, all duplicates collapse + # onto `deployments[0]`. + request_model_info = model_info or {} + request_model_id = request_model_info.get("id") try: - # First try to find by proxy model_name (e.g., "gpt-4o") - deployments = llm_router.get_model_list(model_name=model_name) - - # If not found, try to find by litellm model name (e.g., "azure/gpt-4o") - if not deployments or len(deployments) == 0: - all_deployments = llm_router.get_model_list(model_name=None) - if all_deployments: - for deployment in all_deployments: - if ( - deployment.get("litellm_params", {}).get("model") - == model_name - ): - deployments = [deployment] - break - - if deployments and len(deployments) > 0: - # Use the first deployment's litellm_params as base config - # These already have resolved environment variables from proxy config - config_litellm_params = dict( - deployments[0].get("litellm_params", {}) + deployment_by_id = None + if request_model_id: + deployment_by_id = llm_router.get_deployment( + model_id=request_model_id ) + + if deployment_by_id is not None: + config_litellm_params = deployment_by_id.litellm_params.model_dump( + exclude_none=True + ) + elif model_name: + # Fall back to model_name lookup for callers (e.g. the + # "Add Model" wizard, or curl) that don't supply an id. + # First try to find by proxy model_name (e.g., "gpt-4o") + deployments = llm_router.get_model_list(model_name=model_name) + + # If not found, try to find by litellm model name + # (e.g., "azure/gpt-4o") + if not deployments or len(deployments) == 0: + all_deployments = llm_router.get_model_list(model_name=None) + if all_deployments: + for deployment in all_deployments: + if ( + deployment.get("litellm_params", {}).get("model") + == model_name + ): + deployments = [deployment] + break + + if deployments and len(deployments) > 0: + # Use the first deployment's litellm_params as base + # config. These already have resolved environment + # variables from proxy config. + config_litellm_params = dict( + deployments[0].get("litellm_params", {}) + ) except Exception as e: verbose_proxy_logger.debug( f"Could not find model {model_name} in router: {e}. " diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index a63613c5836..b97e7c5e693 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -1,5 +1,6 @@ import asyncio import copy +import json import re import time from collections import OrderedDict @@ -794,8 +795,17 @@ class LiteLLMProxyRequestSetup: ) ) for k, v in litellm_logging_metadata_headers.items(): - if v is not None: + if v is None: + continue + # httpx requires header values to be str or bytes; coerce numbers/bools + # to str and JSON-encode dict/list (e.g. user_api_key_spend is float, + # user_api_key_auth_metadata is dict). See #27458. + if isinstance(v, (dict, list)): + returned_headers["x-litellm-{}".format(k)] = json.dumps(v) + elif isinstance(v, (str, bytes)): returned_headers["x-litellm-{}".format(k)] = v + else: + returned_headers["x-litellm-{}".format(k)] = str(v) return returned_headers @@ -1731,6 +1741,7 @@ async def add_litellm_data_to_request( # noqa: PLR0915 data=data, user_api_key_dict=user_api_key_dict, pre_alias_model_name=_pre_alias_model, + llm_router=llm_router, ) ## ENFORCED PARAMS CHECK @@ -1864,6 +1875,7 @@ def _apply_credential_overrides_from_model_config( data: dict, user_api_key_dict: UserAPIKeyAuth, pre_alias_model_name: Optional[str] = None, + llm_router: Optional[Router] = None, ) -> None: """ Walk the model_config precedence chain in team/project metadata. @@ -1899,10 +1911,19 @@ def _apply_credential_overrides_from_model_config( if not project_model_config and not team_model_config: return - # Extract provider hint from model name (e.g. "azure/gpt-4" -> "azure") + # Extract provider hint from model name (e.g. "azure/gpt-4" -> "azure"). + # When the user-facing name has no provider prefix, fall back to the + # deployment's litellm_params so multi-provider defaultconfig entries + # don't silently match the first dict key (#27516). provider: Optional[str] = None if "/" in model_name: provider = model_name.split("/", 1)[0] + elif llm_router is not None: + provider = _resolve_provider_from_deployment( + llm_router=llm_router, + model_name=model_name, + pre_alias_model_name=pre_alias_model_name, + ) credential_name = _resolve_credential_from_model_config( model_name=model_name, @@ -1938,6 +1959,48 @@ def _apply_credential_overrides_from_model_config( ) +def _resolve_provider_from_deployment( + llm_router: Router, + model_name: str, + pre_alias_model_name: Optional[str] = None, +) -> Optional[str]: + """ + Resolve a provider hint from the deployment's litellm_params when the + user-facing model name has no provider prefix. + + Tries the post-alias name first (the resolved model group), then the + pre-alias name. Returns None if no deployment is found or the deployment + has no usable provider info. + """ + candidates = [model_name] + if pre_alias_model_name and pre_alias_model_name != model_name: + candidates.append(pre_alias_model_name) + + for name in candidates: + try: + deployment = llm_router.get_deployment_by_model_group_name( + model_group_name=name + ) + except Exception: + deployment = None + if deployment is None: + continue + + litellm_params = getattr(deployment, "litellm_params", None) + if litellm_params is None: + continue + + custom_provider = getattr(litellm_params, "custom_llm_provider", None) + if custom_provider: + return custom_provider + + deployment_model = getattr(litellm_params, "model", "") or "" + if "/" in deployment_model: + return deployment_model.split("/", 1)[0] + + return None + + def _resolve_credential_from_model_config( model_name: str, project_model_config: Optional[dict], diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index b112af1fe20..4439a55c1c6 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -88,8 +88,8 @@ from litellm.router import Router from litellm.secret_managers.main import get_secret from litellm.types.proxy.management_endpoints.key_management_endpoints import ( BulkUpdateKeyRequest, - BulkUpdateKeyRequestItem, BulkUpdateKeyResponse, + BulkUpdateTeamKeysRequest, FailedKeyUpdate, SuccessfulKeyUpdate, ) @@ -1881,7 +1881,7 @@ async def _get_and_validate_existing_key( async def _process_single_key_update( - key_update_item: BulkUpdateKeyRequestItem, + update_key_request: UpdateKeyRequest, user_api_key_dict: UserAPIKeyAuth, litellm_changed_by: Optional[str], prisma_client: Optional[PrismaClient], @@ -1889,6 +1889,7 @@ async def _process_single_key_update( proxy_logging_obj: Any, llm_router: Optional[Router], user_custom_key_update: Optional[Callable] = None, + existing_key_row: Optional[LiteLLM_VerificationToken] = None, ) -> Dict[str, Any]: """ Process a single key update with all validations and checks. @@ -1897,13 +1898,14 @@ async def _process_single_key_update( including validation, permission checks, team checks, and database updates. Args: - key_update_item: The key update request item + update_key_request: Fully-constructed UpdateKeyRequest for the target key user_api_key_dict: The authenticated user's API key info litellm_changed_by: Optional header for tracking who made the change prisma_client: Prisma client instance user_api_key_cache: User API key cache proxy_logging_obj: Proxy logging object llm_router: LLM router instance + existing_key_row: Optional pre-fetched key row to avoid redundant lookups Returns: Dict containing the updated key information @@ -1912,13 +1914,14 @@ async def _process_single_key_update( HTTPException: For various validation and permission errors """ # Validate max_budget - _validate_max_budget(key_update_item.max_budget) + _validate_max_budget(update_key_request.max_budget) # Get and validate existing key - existing_key_row = await _get_and_validate_existing_key( - token=key_update_item.key, - prisma_client=prisma_client, - ) + if existing_key_row is None: + existing_key_row = await _get_and_validate_existing_key( + token=update_key_request.key, + prisma_client=prisma_client, + ) # Check team member permissions if prisma_client is not None: @@ -1930,15 +1933,6 @@ async def _process_single_key_update( user_api_key_cache=user_api_key_cache, ) - # Create UpdateKeyRequest from BulkUpdateKeyRequestItem - update_key_request = UpdateKeyRequest( - key=key_update_item.key, - budget_id=key_update_item.budget_id, - max_budget=key_update_item.max_budget, - team_id=key_update_item.team_id, - tags=key_update_item.tags, - ) - # Custom key update hook if user_custom_key_update is not None: if inspect.iscoroutinefunction(user_custom_key_update): @@ -2003,12 +1997,12 @@ async def _process_single_key_update( detail={"error": "Database not connected"}, ) - _data = {**non_default_values, "token": key_update_item.key} - response = await prisma_client.update_data(token=key_update_item.key, data=_data) + _data = {**non_default_values, "token": update_key_request.key} + response = await prisma_client.update_data(token=update_key_request.key, data=_data) # Delete cache await _delete_cache_key_object( - hashed_token=_hash_token_if_needed(key_update_item.key), + hashed_token=_hash_token_if_needed(update_key_request.key), user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, ) @@ -2598,9 +2592,15 @@ async def bulk_update_keys( for key_update_item in data.keys: try: - # Process single key update using reusable function + update_key_request = UpdateKeyRequest( + key=key_update_item.key, + budget_id=key_update_item.budget_id, + max_budget=key_update_item.max_budget, + team_id=key_update_item.team_id, + tags=key_update_item.tags, + ) updated_key_info = await _process_single_key_update( - key_update_item=key_update_item, + update_key_request=update_key_request, user_api_key_dict=user_api_key_dict, litellm_changed_by=litellm_changed_by, prisma_client=prisma_client, @@ -2665,6 +2665,223 @@ async def bulk_update_keys( ) +def _build_failed_team_key_update( + token: str, + exception: Exception, + existing_key_row: Optional[LiteLLM_VerificationToken], +) -> FailedKeyUpdate: + """Normalize an exception from the per-key update loop into a FailedKeyUpdate.""" + if isinstance(exception, HTTPException): + detail = exception.detail + if isinstance(detail, dict): + error_message = detail.get("error", str(exception)) + else: + error_message = str(detail) + elif isinstance(exception, ProxyException): + error_message = exception.message + else: + error_message = str(exception) + + key_info: Optional[Dict[str, Any]] = None + if existing_key_row is not None: + if hasattr(existing_key_row, "model_dump"): + key_info = existing_key_row.model_dump() + elif hasattr(existing_key_row, "dict"): + key_info = existing_key_row.dict() + if key_info: + key_info.pop("token", None) + + return FailedKeyUpdate(key=token, key_info=key_info, failed_reason=error_message) + + +@router.post( + "/team/key/bulk_update", + tags=["key management"], + dependencies=[Depends(user_api_key_auth)], + response_model=BulkUpdateKeyResponse, +) +@management_endpoint_wrapper +async def bulk_update_team_keys( + data: BulkUpdateTeamKeysRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + litellm_changed_by: Optional[str] = Header( + None, + description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", + ), +): + """ + Apply one update payload to many keys inside a single team. + + Pass `team_id` plus either `key_ids` or `all_keys_in_team=True`. The + `update_fields` payload is broadcast to every selected key. Per-key + failures are returned in `failed_updates` rather than aborting the batch. + + Callable by proxy admins, or by team admins with `KEY_UPDATE` permission. + """ + from litellm.proxy.proxy_server import ( + llm_router, + prisma_client, + proxy_logging_obj, + user_api_key_cache, + user_custom_key_update, + ) + + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": "Database not connected"}, + ) + + if not data.team_id: + raise HTTPException( + status_code=400, + detail={"error": "team_id is required"}, + ) + + MAX_BATCH_SIZE = 500 + if data.key_ids is not None and len(data.key_ids) > MAX_BATCH_SIZE: + raise HTTPException( + status_code=400, + detail={ + "error": f"Maximum {MAX_BATCH_SIZE} keys can be updated at once. Found {len(data.key_ids)} key_ids." + }, + ) + + if data.all_keys_in_team: + # "all" excludes blocked/expired — bulk refresh shouldn't revive a key an admin disabled. + # `blocked` is Boolean? with no default; `/key/generate` writes NULL. Prisma's `NOT` + # excludes NULLs, so explicitly OR `false` with `null` to include them. + now = datetime.now(timezone.utc) + existing_keys = await prisma_client.db.litellm_verificationtoken.find_many( + where={ + "team_id": data.team_id, + "AND": [ + {"OR": [{"blocked": False}, {"blocked": None}]}, + {"OR": [{"expires": None}, {"expires": {"gt": now}}]}, + ], + }, + order={"token": "asc"}, + take=MAX_BATCH_SIZE + 1, + ) + if len(existing_keys) > MAX_BATCH_SIZE: + raise HTTPException( + status_code=400, + detail={ + "error": f"Team {data.team_id} has more than {MAX_BATCH_SIZE} keys. Use `key_ids` to update in batches of {MAX_BATCH_SIZE}." + }, + ) + requested_tokens = [row.token for row in existing_keys] + else: + if data.key_ids is None or len(data.key_ids) == 0: + raise HTTPException( + status_code=400, + detail={ + "error": "key_ids must be provided when all_keys_in_team is False" + }, + ) + # Dedupe by hashed form — duplicates collapse to one update. + requested_tokens = [] + hashed_key_ids = [] + seen_hashes = set() + for k in data.key_ids: + h = _hash_token_if_needed(k) + if h in seen_hashes: + continue + seen_hashes.add(h) + requested_tokens.append(k) + hashed_key_ids.append(h) + existing_keys = await prisma_client.db.litellm_verificationtoken.find_many( + where={"team_id": data.team_id, "token": {"in": hashed_key_ids}} + ) + + # Anchor membership check on data.team_id (not existing_keys[0]); empty result must still gate non-admins. + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value: + auth_anchor = ( + existing_keys[0] + if existing_keys + else LiteLLM_VerificationToken( + token="__team_scope_auth_check__", + team_id=data.team_id, + models=[], + ) + ) + await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint( + user_api_key_dict=user_api_key_dict, + route=KeyManagementRoutes.KEY_UPDATE, + prisma_client=prisma_client, + existing_key_row=auth_anchor, + user_api_key_cache=user_api_key_cache, + ) + + # Block metadata.allowed_passthrough_routes for non-admins — the runtime + # route checker reads it from key/team metadata to grant passthrough. + _check_passthrough_routes_caller_permission( + data=data.update_fields, user_api_key_dict=user_api_key_dict + ) + + if not requested_tokens: + raise HTTPException( + status_code=404, + detail={"error": f"No keys found for team {data.team_id}"}, + ) + + existing_by_token = {row.token: row for row in existing_keys} + update_field_dict = data.update_fields.model_dump(exclude_unset=True) + + successful_updates: List[SuccessfulKeyUpdate] = [] + failed_updates: List[FailedKeyUpdate] = [] + + for token in requested_tokens: + db_token = _hash_token_if_needed(token) + try: + if db_token not in existing_by_token: + raise HTTPException( + status_code=404, + detail={"error": f"Key not found in team {data.team_id}"}, + ) + + # team_id from validated scope, never user payload — drives _check_team_key_limits. + update_key_request = UpdateKeyRequest( + key=token, + team_id=data.team_id, + **update_field_dict, + ) + updated_key_info = await _process_single_key_update( + update_key_request=update_key_request, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=litellm_changed_by, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + llm_router=llm_router, + user_custom_key_update=user_custom_key_update, + existing_key_row=existing_by_token[db_token], + ) + + successful_updates.append( + SuccessfulKeyUpdate(key=token, key_info=updated_key_info) + ) + + except Exception as e: + # Log the hashed prefix — `token` may be a raw sk-... and ERROR logs persist. + verbose_proxy_logger.exception( + f"Failed to update key {db_token[:12]}... in team {data.team_id}: {e}" + ) + failed_updates.append( + _build_failed_team_key_update( + token=token, + exception=e, + existing_key_row=existing_by_token.get(db_token), + ) + ) + + return BulkUpdateKeyResponse( + total_requested=len(requested_tokens), + successful_updates=successful_updates, + failed_updates=failed_updates, + ) + + async def validate_key_team_change( key: LiteLLM_VerificationToken, team: LiteLLM_TeamTable, diff --git a/litellm/proxy/management_endpoints/tag_management_endpoints.py b/litellm/proxy/management_endpoints/tag_management_endpoints.py index 0e60820aab1..49d9b67a28a 100644 --- a/litellm/proxy/management_endpoints/tag_management_endpoints.py +++ b/litellm/proxy/management_endpoints/tag_management_endpoints.py @@ -12,12 +12,13 @@ All /tag management endpoints import asyncio import json -from typing import TYPE_CHECKING, Dict, List, Optional +from datetime import datetime +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union -from fastapi import APIRouter, Depends, HTTPException +from fastapi import APIRouter, Depends, HTTPException, Query from litellm._logging import verbose_proxy_logger -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import UserAPIKeyAuth, user_api_key_has_admin_view from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.management_endpoints.common_daily_activity import ( SpendAnalyticsPaginatedResponse, @@ -39,6 +40,72 @@ if TYPE_CHECKING: router = APIRouter() +async def _get_internal_user_api_keys( + prisma_client, + user_api_key_dict: UserAPIKeyAuth, +) -> List[str]: + user_role = user_api_key_dict.user_role + if user_role is None or not user_role.is_internal_user_role: + return [] + + user_api_keys = set() + if user_api_key_dict.api_key: + user_api_keys.add(user_api_key_dict.api_key) + + user_id = user_api_key_dict.user_id + if user_id is None: + return sorted(user_api_keys) + + key_records = await prisma_client.db.litellm_verificationtoken.find_many( + where={"user_id": user_id}, + select={"token": True}, + ) + user_api_keys.update( + key_record.token + for key_record in key_records + if getattr(key_record, "token", None) + ) + + return sorted(user_api_keys) + + +async def _get_tag_list_scope( + prisma_client, + user_api_key_dict: UserAPIKeyAuth, +) -> Optional[Dict[str, dict]]: + user_role = user_api_key_dict.user_role + if user_api_key_has_admin_view(user_api_key_dict) or ( + user_role is None or not user_role.is_internal_user_role + ): + return None + + scoped_api_keys = await _get_internal_user_api_keys( + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + ) + return {"api_key": {"in": scoped_api_keys}} + + +async def _get_tag_daily_activity_api_key_filter( + prisma_client, + user_api_key_dict: UserAPIKeyAuth, + requested_api_key: Optional[str], +) -> Optional[Union[str, List[str]]]: + user_role = user_api_key_dict.user_role + if user_api_key_has_admin_view(user_api_key_dict) or ( + user_role is None or not user_role.is_internal_user_role + ): + return requested_api_key + + scoped_api_keys = await _get_internal_user_api_keys( + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + ) + if requested_api_key is not None: + return requested_api_key if requested_api_key in scoped_api_keys else [] + return scoped_api_keys + + async def _get_model_names(prisma_client, model_ids: list) -> Dict[str, str]: """Helper function to get model names from model IDs""" try: @@ -395,6 +462,32 @@ async def info_tag( raise HTTPException(status_code=500, detail=str(e)) +def _validate_tag_list_date_range( + start_date: Optional[str], end_date: Optional[str] +) -> None: + """Require both dates together, and enforce YYYY-MM-DD format with start <= end.""" + if (start_date is None) != (end_date is None): + raise HTTPException( + status_code=400, + detail="start_date and end_date must be provided together", + ) + if start_date is None: + return + try: + start = datetime.strptime(start_date, "%Y-%m-%d") + end = datetime.strptime(end_date, "%Y-%m-%d") # type: ignore[arg-type] + except ValueError as e: + raise HTTPException( + status_code=400, + detail=f"Invalid date format, expected YYYY-MM-DD: {e}", + ) + if start > end: + raise HTTPException( + status_code=400, + detail="start_date must be on or before end_date", + ) + + @router.get( "/tag/list", tags=["tag management"], @@ -402,6 +495,18 @@ async def info_tag( ) async def list_tags( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + start_date: Optional[str] = Query( + None, + description=( + "Optional start date (YYYY-MM-DD). When provided together with " + "end_date, dynamic tags are limited to those active in the window. " + "Stored tags are always returned." + ), + ), + end_date: Optional[str] = Query( + None, + description="Optional end date (YYYY-MM-DD). Must be given with start_date.", + ), ): """ List all available tags with their budget information. @@ -411,10 +516,44 @@ async def list_tags( if prisma_client is None: raise HTTPException(status_code=500, detail="Database not connected") + _validate_tag_list_date_range(start_date, end_date) + try: + tag_scope = await _get_tag_list_scope( + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + ) + + ## QUERY DYNAMIC TAGS ## + # Use group_by instead of find_many(distinct=["tag"]). + # Prisma's distinct fetches all columns for all rows and deduplicates + # in application code, which is extremely slow on large tables. + # See: https://www.prisma.io/docs/orm/prisma-client/queries/aggregation-grouping-summarizing#distinct-under-the-hood + dynamic_tag_where: Dict[str, Any] = {"tag": {"not": None}} + if tag_scope: + dynamic_tag_where = {**dynamic_tag_where, **tag_scope} + if start_date is not None and end_date is not None: + dynamic_tag_where["date"] = {"gte": start_date, "lte": end_date} + + dynamic_tag_rows = await prisma_client.db.litellm_dailytagspend.group_by( + by=["tag"], + where=dynamic_tag_where, + min={"created_at": True}, + max={"updated_at": True}, + ) + + used_tag_names = [row["tag"] for row in dynamic_tag_rows if row["tag"]] + if tag_scope is not None and not used_tag_names: + return [] + + stored_tag_where = ( + {"tag_name": {"in": used_tag_names}} if tag_scope is not None else None + ) + ## QUERY STORED TAGS ## tag_records = await prisma_client.db.litellm_tagtable.find_many( - include={"litellm_budget_table": True} + where=stored_tag_where, + include={"litellm_budget_table": True}, ) stored_tag_names = set() @@ -448,18 +587,6 @@ async def list_tags( list_of_tags.append(tag_dict) - ## QUERY DYNAMIC TAGS ## - # Use group_by instead of find_many(distinct=["tag"]). - # Prisma's distinct fetches all columns for all rows and deduplicates - # in application code, which is extremely slow on large tables. - # See: https://www.prisma.io/docs/orm/prisma-client/queries/aggregation-grouping-summarizing#distinct-under-the-hood - dynamic_tag_rows = await prisma_client.db.litellm_dailytagspend.group_by( - by=["tag"], - where={"tag": {"not": None}}, - min={"created_at": True}, - max={"updated_at": True}, - ) - dynamic_tag_config = [ { "name": row["tag"], @@ -527,6 +654,7 @@ async def get_tag_daily_activity( api_key: Optional[str] = None, page: int = 1, page_size: int = 10, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ Get daily activity for specific tags or all tags. @@ -545,8 +673,18 @@ async def get_tag_daily_activity( """ from litellm.proxy.proxy_server import prisma_client + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + # Convert comma-separated tags string to list if provided tag_list = tags.split(",") if tags else None + scoped_api_key_filter = await _get_tag_daily_activity_api_key_filter( + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + requested_api_key=api_key, + ) + if scoped_api_key_filter == []: + return SpendAnalyticsPaginatedResponse(results=[]) return await get_daily_activity( prisma_client=prisma_client, @@ -557,7 +695,7 @@ async def get_tag_daily_activity( start_date=start_date, end_date=end_date, model=model, - api_key=api_key, + api_key=scoped_api_key_filter, page=page, page_size=page_size, # metadata_metrics_func=None because litellm_dailytagspend rows are diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index b13118db6fe..6e2e2bedac1 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -740,7 +740,7 @@ def generic_response_convertor( all_teams = [] if sso_jwt_handler is not None: - team_ids = sso_jwt_handler.get_team_ids_from_jwt(cast(dict, response)) + team_ids = sso_jwt_handler.get_all_jwt_team_ids(cast(dict, response)) all_teams.extend(team_ids) if team_mappings is not None and team_mappings.team_ids_jwt_field is not None: @@ -755,7 +755,7 @@ def generic_response_convertor( f"Loaded team_ids from DB team_mappings.team_ids_jwt_field='{team_mappings.team_ids_jwt_field}': {team_ids_from_db_mapping}" ) else: - team_ids = jwt_handler.get_team_ids_from_jwt(cast(dict, response)) + team_ids = jwt_handler.get_all_jwt_team_ids(cast(dict, response)) all_teams.extend(team_ids) # Determine user role based on role_mappings if available diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 6359b48654b..5fc8c44b2d8 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -805,7 +805,8 @@ def run_server( # noqa: PLR0915 ) db_connection_pool_limit = 100 - db_connection_timeout = 60 + # Starts optional due to config fallback checks; guaranteed non-None before use. + db_connection_timeout: Optional[Union[int, float]] = 60 general_settings = {} ### GET DB TOKEN FOR IAM AUTH ### @@ -914,10 +915,15 @@ def run_server( # noqa: PLR0915 "database_connection_pool_limit", LiteLLMDatabaseConnectionPool.database_connection_pool_limit.value, ) - db_connection_timeout = general_settings.get( - "database_connection_pool_timeout", - LiteLLMDatabaseConnectionPool.database_connection_pool_timeout.value, - ) + db_connection_timeout = general_settings.get("database_connection_timeout") + if db_connection_timeout is None: + db_connection_timeout = general_settings.get( + "database_connection_pool_timeout" + ) + if db_connection_timeout is None: + db_connection_timeout = ( + LiteLLMDatabaseConnectionPool.database_connection_pool_timeout.value + ) if database_url and database_url.startswith("os.environ/"): original_dir = os.getcwd() # set the working directory to where this script is diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 493519f2328..4a538a28e03 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1061,6 +1061,52 @@ vertex_live_passthrough_vertex_base = VertexBase() from fastapi.routing import APIWebSocketRoute +def _inject_websocket_stubs_into_openapi_schema( + openapi_schema: dict, websocket_routes: list +) -> dict: + """ + Add a synthetic GET stub for each WebSocket route so it appears in Swagger UI. + + Merges into any existing path entry rather than replacing it — a WebSocket route + that shares its path with an HTTP route must not erase the HTTP operation. If + a "get" operation is already documented on the path, the WebSocket stub is + skipped to preserve the real GET. + """ + for route in websocket_routes: + base_path = route.path.split("{")[0].rstrip("?") + + parameters = [] + try: + if hasattr(route, "dependant") and route.dependant is not None: + # Handle both FastAPI <0.120 and >=0.120 + query_params = getattr(route.dependant, "query_params", []) + if query_params: + for param in query_params: + parameters.append( + { + "name": param.name, + "in": "query", + "required": param.required, + "schema": {"type": "string"}, + } + ) + except (AttributeError, TypeError): + pass + + path_entry = openapi_schema["paths"].setdefault(base_path, {}) + if "get" not in path_entry: + path_entry["get"] = { + "summary": f"WebSocket: {route.name or base_path}", + "description": "WebSocket connection endpoint", + "operationId": f"websocket_{route.name or base_path.replace('/', '_')}", + "parameters": parameters, + "responses": {"101": {"description": "WebSocket Protocol Switched"}}, + "tags": ["WebSocket"], + } + + return openapi_schema + + def get_openapi_schema(): if app.openapi_schema: return app.openapi_schema @@ -1083,43 +1129,11 @@ def get_openapi_schema(): route for route in app.routes if isinstance(route, APIWebSocketRoute) ] - # Add each WebSocket route to the schema - for route in websocket_routes: - # Get the base path without query parameters - base_path = route.path.split("{")[0].rstrip("?") - - # Extract parameters from the route - parameters = [] - try: - if hasattr(route, "dependant") and route.dependant is not None: - # Handle both FastAPI <0.120 and >=0.120 - query_params = getattr(route.dependant, "query_params", []) - if query_params: - for param in query_params: - parameters.append( - { - "name": param.name, - "in": "query", - "required": param.required, - "schema": { - "type": "string" - }, # You can make this more specific if needed - } - ) - except (AttributeError, TypeError): - # If we can't access query_params, continue without them - pass - - openapi_schema["paths"][base_path] = { - "get": { - "summary": f"WebSocket: {route.name or base_path}", - "description": "WebSocket connection endpoint", - "operationId": f"websocket_{route.name or base_path.replace('/', '_')}", - "parameters": parameters, - "responses": {"101": {"description": "WebSocket Protocol Switched"}}, - "tags": ["WebSocket"], - } - } + # Add a synthetic GET stub for each so they render in Swagger UI, + # without clobbering existing HTTP operations on the same path. + openapi_schema = _inject_websocket_stubs_into_openapi_schema( + openapi_schema, websocket_routes + ) # Add LLM API request schema bodies for documentation from litellm.proxy.common_utils.custom_openapi_spec import CustomOpenAPISpec @@ -5937,10 +5951,20 @@ class ProxyConfig: verbose_proxy_logger.debug( "guardrails from the DB %s", str(guardrails_in_db) ) + db_guardrail_ids: set = set() for guardrail in guardrails_in_db: + guardrail_id = guardrail.get("guardrail_id") + if guardrail_id: + db_guardrail_ids.add(guardrail_id) IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db( guardrail=cast(Guardrail, guardrail), ) + + # Drop in-memory DB-backed entries whose row was deleted on another + # pod. Config-loaded entries are never touched. + IN_MEMORY_GUARDRAIL_HANDLER.reconcile_db_guardrails( + db_guardrail_ids=db_guardrail_ids + ) except Exception as e: verbose_proxy_logger.exception( "litellm.proxy.proxy_server.py::ProxyConfig:_init_guardrails_in_db - {}".format( diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 0577110a26f..2998555db1a 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -4977,10 +4977,10 @@ class ProxyUpdateSpend: timeout=timedelta(seconds=60) ) as transaction: async with transaction.batch_() as batcher: - for ( - end_user_id, - response_cost, - ) in end_user_list_transactions.items(): + # Sort by end_user_id for consistent lock ordering across pods to prevent deadlocks. + for end_user_id, response_cost in sorted( + end_user_list_transactions.items() + ): if litellm.max_end_user_budget is not None: pass batcher.litellm_endusertable.upsert( diff --git a/litellm/router.py b/litellm/router.py index 7512ee387dc..37295f1a7d2 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -7076,11 +7076,11 @@ class Router: _shared_model_info = { k: v for k, v in _model_info.items() if k not in _custom_pricing_fields } - litellm.register_model( - model_cost={ - _model_name: _shared_model_info, - } - ) + _backend_alias_cost = {_model_name: _shared_model_info} + if "responses/" in _model_name: + _stripped_model_name = _model_name.replace("responses/", "") + _backend_alias_cost[_stripped_model_name] = _shared_model_info + litellm.register_model(model_cost=_backend_alias_cost) ## Check if LLM Deployment is allowed for this deployment if ( @@ -7752,6 +7752,12 @@ class Router: # initialize client self._add_deployment(deployment=deployment) + _model_info_dict: dict = deployment.model_info.model_dump(exclude_none=True) + for field in CustomPricingLiteLLMParams.model_fields.keys(): + field_value = deployment.litellm_params.get(field) + if field_value is not None: + _model_info_dict[field] = field_value + # Register custom pricing in litellm.model_cost. # Mirrors _create_deployment() logic to ensure dynamically-added deployments # (e.g., loaded from DB) also have their custom pricing registered. @@ -7759,13 +7765,31 @@ class Router: # zero-cost models, causing budget checks to block free models. _model_id = deployment.model_info.id if _model_id is not None: - _model_info_dict: dict = deployment.model_info.model_dump(exclude_none=True) - for field in CustomPricingLiteLLMParams.model_fields.keys(): - field_value = deployment.litellm_params.get(field) - if field_value is not None: - _model_info_dict[field] = field_value litellm.register_model(model_cost={_model_id: _model_info_dict}) + ## REGISTER MODEL INFO IN LITELLM MODEL COST MAP + ## OLD MODEL REGISTRATION ## Kept to prevent breaking changes + _model_name = deployment.litellm_params.model + if deployment.litellm_params.custom_llm_provider is not None: + _model_name = ( + deployment.litellm_params.custom_llm_provider + "/" + _model_name + ) + + # For the shared backend key, strip custom pricing fields so that + # one deployment's pricing overrides don't pollute another + # deployment sharing the same backend model name. + # Each deployment's full pricing is already stored under its + # unique model_id above (when present). + _custom_pricing_fields = CustomPricingLiteLLMParams.model_fields.keys() + _shared_model_info = { + k: v for k, v in _model_info_dict.items() if k not in _custom_pricing_fields + } + _backend_alias_cost = {_model_name: _shared_model_info} + if "responses/" in _model_name: + _stripped_model_name = _model_name.replace("responses/", "") + _backend_alias_cost[_stripped_model_name] = _shared_model_info + litellm.register_model(model_cost=_backend_alias_cost) + # add to model names self._add_model_to_list_and_index_map( model=_deployment, model_id=deployment.model_info.id diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py index 64827db13f6..5db2a45054a 100644 --- a/litellm/types/llms/bedrock.py +++ b/litellm/types/llms/bedrock.py @@ -1042,3 +1042,10 @@ class BedrockInvokeAnthropicMessagesRequest(TypedDict, total=False): thinking: dict metadata: dict output_config: dict + + # `context_management` is allowed for Bedrock InvokeModel only when it + # carries `compact_20260112` edits paired with the `compact-2026-01-12` + # anthropic-beta header. The Invoke transformation filters edits to the + # supported subset and strips the field entirely when nothing remains, so + # other edit types (e.g. `clear_thinking_20251015`) never reach Bedrock. + context_management: dict diff --git a/litellm/types/proxy/management_endpoints/key_management_endpoints.py b/litellm/types/proxy/management_endpoints/key_management_endpoints.py index b1d25455d18..d214cdb4f5d 100644 --- a/litellm/types/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/types/proxy/management_endpoints/key_management_endpoints.py @@ -1,6 +1,7 @@ -from typing import Any, Dict, List, Optional +from datetime import datetime +from typing import Any, Dict, List, Literal, Optional -from pydantic import BaseModel +from pydantic import BaseModel, ConfigDict, model_validator class BulkUpdateKeyRequestItem(BaseModel): @@ -40,3 +41,78 @@ class BulkUpdateKeyResponse(BaseModel): total_requested: int successful_updates: List[SuccessfulKeyUpdate] failed_updates: List[FailedKeyUpdate] + + +class KeyUpdateFields(BaseModel): + """Allowlist of bulk-broadcastable fields for /team/key/bulk_update; `extra="forbid"` blocks RBAC/ownership/scope mutations even by team admins.""" + + model_config = ConfigDict(extra="forbid", protected_namespaces=()) + + # Budgets + max_budget: Optional[float] = None + budget_id: Optional[str] = None + budget_duration: Optional[str] = None + budget_limits: Optional[List[Any]] = None + model_max_budget: Optional[Dict[str, Any]] = None + + # Rate limits + tpm_limit: Optional[int] = None + rpm_limit: Optional[int] = None + model_tpm_limit: Optional[Dict[str, Any]] = None + model_rpm_limit: Optional[Dict[str, Any]] = None + max_parallel_requests: Optional[int] = None + rpm_limit_type: Optional[ + Literal["guaranteed_throughput", "best_effort_throughput", "dynamic"] + ] = None + tpm_limit_type: Optional[ + Literal["guaranteed_throughput", "best_effort_throughput", "dynamic"] + ] = None + + # Temporary budget grants (auto-expire). `spend` deliberately omitted — bulk-zeroing it bypasses budget enforcement; admin-only via /key/update. + temp_budget_increase: Optional[float] = None + temp_budget_expiry: Optional[datetime] = None + + # Expiry + duration: Optional[str] = None + + # Operational metadata + tags: Optional[List[str]] = None + metadata: Optional[Dict[str, Any]] = None + + @model_validator(mode="after") + def validate_temp_budget(self) -> "KeyUpdateFields": + if self.temp_budget_increase is not None or self.temp_budget_expiry is not None: + if self.temp_budget_increase is None or self.temp_budget_expiry is None: + raise ValueError( + "temp_budget_increase and temp_budget_expiry must be set together" + ) + return self + + @model_validator(mode="after") + def require_at_least_one_field(self) -> "KeyUpdateFields": + # Reject empty payload — would iterate every key with no-op writes. + if not self.model_fields_set: + raise ValueError("update_fields must specify at least one field to update.") + return self + + +class BulkUpdateTeamKeysRequest(BaseModel): + """Apply one update payload to many keys inside a team; provide either `key_ids` or `all_keys_in_team=True`.""" + + team_id: str + key_ids: Optional[List[str]] = None + all_keys_in_team: bool = False + update_fields: KeyUpdateFields + + @model_validator(mode="after") + def validate_selection(self) -> "BulkUpdateTeamKeysRequest": + has_key_ids = self.key_ids is not None and len(self.key_ids) > 0 + if has_key_ids and self.all_keys_in_team: + raise ValueError( + "Provide either `key_ids` or `all_keys_in_team=True`, not both." + ) + if not has_key_ids and not self.all_keys_in_team: + raise ValueError( + "Must provide either `key_ids` (non-empty) or `all_keys_in_team=True`." + ) + return self diff --git a/litellm/utils.py b/litellm/utils.py index 019fbc2add8..da80e4ae164 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1,3 +1,5 @@ +"""Utility helpers for LiteLLM core request handling and provider support.""" + # from __future__ import annotations must be the first non-comment statement from __future__ import annotations @@ -4405,6 +4407,10 @@ def get_optional_params( # noqa: PLR0915 else False ), ) + if bedrock_route == "claude_platform": + optional_params = BedrockModelInfo.map_claude_platform_auth_params( + passed_params=passed_params, optional_params=optional_params + ) elif custom_llm_provider == "cloudflare": optional_params = litellm.CloudflareChatConfig().map_openai_params( model=model, @@ -9490,6 +9496,49 @@ def get_non_default_completion_params(kwargs: dict) -> dict: return non_default_params +def peek_reasoning_summary_aliases(optional_params: dict) -> Optional[Any]: + """Read AI-SDK-style reasoning summary from optional_params or nested extra_body. + + Uses key membership (not ``or`` chains) so falsy values like ``""`` are not skipped. + """ + if "reasoningSummary" in optional_params: + return optional_params["reasoningSummary"] + if "reasoning_summary" in optional_params: + return optional_params["reasoning_summary"] + extra_body = optional_params.get("extra_body") + if isinstance(extra_body, dict): + if "reasoningSummary" in extra_body: + return extra_body["reasoningSummary"] + if "reasoning_summary" in extra_body: + return extra_body["reasoning_summary"] + return None + + +def strip_reasoning_summary_aliases_from_optional_params( + optional_params: dict, +) -> Tuple[dict, Optional[Any]]: + """Copy optional_params; remove reasoningSummary aliases from top-level and extra_body.""" + op = dict(optional_params) + rs_val = op.pop("reasoningSummary", None) + snake_rs_val = op.pop("reasoning_summary", None) + if rs_val is None: + rs_val = snake_rs_val + eb = op.get("extra_body") + if isinstance(eb, dict): + eb = dict(eb) + eb_rs_val = eb.pop("reasoningSummary", None) + eb_snake_rs_val = eb.pop("reasoning_summary", None) + if rs_val is None: + rs_val = eb_rs_val + if rs_val is None: + rs_val = eb_snake_rs_val + if eb: + op["extra_body"] = eb + else: + op.pop("extra_body", None) + return op, rs_val + + def get_non_default_transcription_params(kwargs: dict) -> dict: from litellm.constants import OPENAI_TRANSCRIPTION_PARAMS diff --git a/pyproject.toml b/pyproject.toml index 5cd83148d37..65557d8a9ee 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -275,6 +275,33 @@ filterwarnings = [ "ignore::DeprecationWarning:pytest_asyncio.plugin", ] +[tool.mutmut] +# Mutation-testing scope. Driven by the manually-triggered workflow at +# .github/workflows/mutation-test.yml. mutmut is not part of the project's +# default install; it is pulled in via `uv run --with mutmut==` in CI. +# `also_copy = ["litellm/"]` is required because mutmut runs in a `mutants/` +# sandbox and the test conftest imports from across the litellm package. +paths_to_mutate = [ + "litellm/proxy/management_endpoints/", +] +tests_dir = [ + "tests/test_litellm/proxy/management_endpoints/", +] +also_copy = [ + "litellm/", +] +# Disable rerun/parallel plugins for mutation runs: +# - pytest-retry triggers an `INTERNALERROR: no option named 'filtered_exceptions'` +# when invoked via mutmut's in-process `pytest.main()` call. +# - rerunning a "failed" test on a mutant would mask which mutants are killed +# vs. survive, so reruns are wrong for mutation testing regardless. +# - xdist is unnecessary inside mutmut (mutmut handles its own parallelism). +pytest_add_cli_args = [ + "-p", "no:retry", + "-p", "no:rerunfailures", + "-p", "no:xdist", +] + [tool.coverage.run] source = ["litellm"] relative_files = true diff --git a/scripts/mutation_report.py b/scripts/mutation_report.py new file mode 100644 index 00000000000..a606e3f71cf --- /dev/null +++ b/scripts/mutation_report.py @@ -0,0 +1,423 @@ +#!/usr/bin/env python3 +"""Generate an agent-actionable mutation testing report. + +Reads the mutmut sandbox state at `mutants/` and produces a single +`mutation-report.md` grouped by function. For each function with surviving +mutants, the report embeds the original function source (via AST), the +unified diff for each surviving mutation (via `mutmut show`), and the +existing test file(s) — followed by an ACH-style instruction asking the +reader to write tests that kill the survivors. + +Run after `mutmut run` and `mutmut export-cicd-stats`. Expects mutmut to be +invokable as `uv run --no-sync --with mutmut== mutmut `. +""" +from __future__ import annotations + +import ast +import json +import re +import subprocess +import sys +import tomllib +from collections import defaultdict +from difflib import SequenceMatcher +from pathlib import Path +from textwrap import dedent + +ROOT = Path(__file__).resolve().parent.parent +MUTMUT_INVOCATION = ["uv", "run", "--no-sync", "--with", "mutmut==3.5.0", "mutmut"] + + +def load_mutmut_config() -> dict: + with open(ROOT / "pyproject.toml", "rb") as f: + return tomllib.load(f)["tool"]["mutmut"] + + +def get_survivors() -> list[str]: + proc = subprocess.run( + [*MUTMUT_INVOCATION, "results"], capture_output=True, text=True, check=False + ) + survivors = [] + for line in proc.stdout.splitlines(): + m = re.match(r"\s*(\S+):\s*survived\s*$", line) + if m: + survivors.append(m.group(1)) + return survivors + + +def get_mutmut_show(mutant_name: str) -> str: + proc = subprocess.run( + [*MUTMUT_INVOCATION, "show", mutant_name], + capture_output=True, + text=True, + check=False, + ) + return proc.stdout.strip() or "(mutmut show produced no output)" + + +def parse_mutant_name(name: str) -> tuple[str, str, str]: + """Parse `.x___mutmut_` -> (module, function, N). + + mutmut prefixes mutated functions with `x_` (single underscore). For a + function named `foo`, mutants are `x_foo__mutmut_N`. For a function named + `_foo` (leading underscore), the mutant becomes `x__foo__mutmut_N` — so + the regex matches a single underscore after `x` and captures everything + (including any leading underscores) up to `__mutmut_`. + """ + m = re.match(r"^(.+)\.x_(.+)__mutmut_(\d+)$", name) + if not m: + return name, name, "?" + return m.group(1), m.group(2), m.group(3) + + +def function_anchor(module_path: str, function_name: str) -> str: + return re.sub(r"[^a-z0-9_-]+", "-", f"{module_path}-{function_name}".lower()).strip( + "-" + ) + + +def module_to_file(module_path: str) -> Path | None: + candidate = ROOT / Path(*module_path.split(".")).with_suffix(".py") + return candidate if candidate.exists() else None + + +def find_function_in_file( + file_path: Path, function_name: str +) -> tuple[int, int, str, list[int]] | None: + """Find a top-level or nested function by name; returns the first match. + + Returns ``(start_line, end_line, source, all_match_lines)`` or ``None``. + ``all_match_lines`` is the start line of every function (any nesting + level) in the file with this name. When ``len(all_match_lines) > 1`` the + file defines the same name in multiple places (e.g., a module-level + helper and a class method) — mutmut's mutant identifier does not carry + class context, so we can't determine which definition was mutated. + Callers surface a disambiguation note in that case. + """ + src = file_path.read_text() + tree = ast.parse(src) + matches = [ + node + for node in ast.walk(tree) + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) + and node.name == function_name + ] + if not matches: + return None + first = matches[0] + lines = src.splitlines() + return ( + first.lineno, + first.end_lineno, + "\n".join(lines[first.lineno - 1 : first.end_lineno]), + [m.lineno for m in matches], + ) + + +def collect_test_files(tests_dir: list[str]) -> list[Path]: + found: list[Path] = [] + for entry in tests_dir: + p = ROOT / entry + if p.is_file(): + found.append(p) + elif p.is_dir(): + found.extend(sorted(p.rglob("test_*.py"))) + return found + + +def _indent_of(line: str) -> str: + return line[: len(line) - len(line.lstrip())] + + +def render_meta_style_mutant( + module_path: str, function_name: str, mutant_num: str +) -> str | None: + """Render the mutated function with `# MUTANT START`/`# MUTANT END` delimiters. + + Reads `mutants/.py` (the trampoline file mutmut emits), finds + `x___mutmut_orig` and `x___mutmut_`, and renders the + mutated version with the lines that differ from `__mutmut_orig` wrapped + in `# MUTANT START`/`# MUTANT END` comments — the format from Meta's + ACH paper (arXiv 2501.12862, Table 1). + + The function header is rewritten to use the original function name so + the agent sees the source as it would appear in the file (rather than + mutmut's internal `x_*__mutmut_` name). + + Returns None if the trampoline file or either function cannot be found + (the caller falls back to the unified diff). + """ + trampoline = ROOT / "mutants" / Path(*module_path.split(".")).with_suffix(".py") + if not trampoline.exists(): + return None + + src = trampoline.read_text() + try: + tree = ast.parse(src) + except SyntaxError: + return None + file_lines = src.splitlines() + + orig_def = f"x_{function_name}__mutmut_orig" + mutant_def = f"x_{function_name}__mutmut_{mutant_num}" + + orig_node = mutated_node = None + for node in ast.walk(tree): + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): + if node.name == orig_def: + orig_node = node + elif node.name == mutant_def: + mutated_node = node + + if orig_node is None or mutated_node is None: + return None + + orig_lines = file_lines[orig_node.lineno - 1 : orig_node.end_lineno] + mutated_lines = file_lines[mutated_node.lineno - 1 : mutated_node.end_lineno] + if not orig_lines or not mutated_lines: + return None + + # Rewrite the def line to use the original (non-trampolined) function name + # so the agent sees the function as it appears in the source file. + orig_lines[0] = orig_lines[0].replace(orig_def, function_name, 1) + mutated_lines[0] = mutated_lines[0].replace(mutant_def, function_name, 1) + + matcher = SequenceMatcher(a=orig_lines, b=mutated_lines) + out: list[str] = [] + in_diff = False + + for op, i1, i2, j1, j2 in matcher.get_opcodes(): + if op == "equal": + if in_diff: + # Close the block at the indent of the line just inside it. + indent = _indent_of(out[-1]) if out else "" + out.append(f"{indent}# MUTANT END") + in_diff = False + out.extend(mutated_lines[j1:j2]) + else: + if not in_diff: + # Open the block at the indent of the first differing line. + if j1 < len(mutated_lines): + indent = _indent_of(mutated_lines[j1]) + elif i1 < len(orig_lines): + indent = _indent_of(orig_lines[i1]) + else: + indent = "" + out.append(f"{indent}# MUTANT START") + in_diff = True + if op == "delete": + # Mutation removed lines — surface what was deleted as a + # comment so the agent can see the intent of the change. + for deleted in orig_lines[i1:i2]: + indent = _indent_of(deleted) + out.append(f"{indent}# (deleted by mutation): {deleted.lstrip()}") + else: + # replace / insert: take from mutated_lines + out.extend(mutated_lines[j1:j2]) + + if in_diff: + indent = _indent_of(out[-1]) if out else "" + out.append(f"{indent}# MUTANT END") + + return "\n".join(out) + + +def render(config: dict, survivors: list[str], stats: dict | None) -> str: + by_function: dict[tuple[str, str], list[tuple[str, str]]] = defaultdict(list) + for survivor in survivors: + module_path, function_name, mutant_num = parse_mutant_name(survivor) + by_function[(module_path, function_name)].append((survivor, mutant_num)) + + out: list[str] = [] + out.append("# Mutation Test Report") + out.append("") + + out.append("## Summary") + out.append("") + if stats: + total = stats.get("total", 0) or sum( + stats.get(k, 0) + for k in ( + "killed", + "survived", + "no_tests", + "skipped", + "suspicious", + "timeout", + "segfault", + ) + ) + killed = stats.get("killed", 0) + survived = stats.get("survived", 0) + score = (killed / total * 100) if total else 0.0 + out.append(f"- Total mutants: **{total}**") + out.append(f"- Killed: **{killed}**") + out.append(f"- Survived: **{survived}**") + out.append(f"- Mutation score: **{score:.1f}%**") + for k in ("no_tests", "skipped", "suspicious", "timeout", "segfault"): + v = stats.get(k, 0) + if v: + out.append(f"- {k.replace('_', ' ').title()}: {v}") + else: + out.append(f"- Survivors found: **{len(survivors)}**") + out.append("- (mutmut-cicd-stats.json not available — full counts unavailable)") + out.append("") + + if not survivors: + out.append("**No surviving mutants — the test suite caught every mutation.**") + out.append("") + return "\n".join(out) + + out.append("## Surviving mutants by function") + out.append("") + for (module_path, function_name), items in by_function.items(): + anchor = function_anchor(module_path, function_name) + out.append( + f"- [`{function_name}`](#{anchor}) — {len(items)} mutant" + f"{'s' if len(items) != 1 else ''} ({module_path})" + ) + out.append("") + + for (module_path, function_name), items in by_function.items(): + anchor = function_anchor(module_path, function_name) + out.append(f'') + out.append(f"## `{module_path}.{function_name}`") + out.append("") + out.append(f"**Module:** `{module_path}`") + + file_path = module_to_file(module_path) + if file_path is None: + out.append("") + out.append(f"_(could not locate source file for module `{module_path}`)_") + out.append("") + else: + rel = file_path.relative_to(ROOT) + out.append(f"**File:** `{rel}`") + out.append("") + found = find_function_in_file(file_path, function_name) + if found: + start, end, fn_src, all_lines = found + out.append(f"### Original function (lines {start}-{end})") + out.append("") + if len(all_lines) > 1: + line_list = ", ".join(str(line) for line in all_lines) + out.append( + f"> **Note:** {len(all_lines)} functions named " + f"`{function_name}` are defined in this file at lines " + f"{line_list}. Showing the first match. mutmut's " + f"mutant identifier does not carry class context, so " + f"the body below may not correspond to the function " + f"that was actually mutated — verify manually before " + f"writing the killing test." + ) + out.append("") + out.append("```python") + out.append(fn_src) + out.append("```") + out.append("") + else: + out.append(f"_(could not locate `{function_name}` in {rel} via AST)_") + out.append("") + + out.append(f"### Surviving mutations ({len(items)})") + out.append("") + for i, (mutant_name, mutant_num) in enumerate(items, 1): + out.append(f"#### Mutation {i} of {len(items)} — `{mutant_name}`") + out.append("") + meta_style = render_meta_style_mutant( + module_path, function_name, mutant_num + ) + if meta_style is not None: + out.append( + "Mutated function (the bug is delimited by " + "`# MUTANT START` / `# MUTANT END`):" + ) + out.append("") + out.append("```python") + out.append(meta_style) + out.append("```") + out.append("") + out.append("
Unified diff (`mutmut show`)") + out.append("") + out.append("```diff") + out.append(get_mutmut_show(mutant_name)) + out.append("```") + out.append("") + out.append("
") + out.append("") + else: + # Fallback: trampoline file or function lookup failed. + out.append("```diff") + out.append(get_mutmut_show(mutant_name)) + out.append("```") + out.append("") + + test_files = collect_test_files(config.get("tests_dir", [])) + if test_files: + out.append("## Existing tests") + out.append("") + out.append( + "These are the test files that mutmut considered when classifying the " + "mutants above. New tests should be added here, matching existing " + "conventions, fixtures, and naming." + ) + out.append("") + for tf in test_files: + rel = tf.relative_to(ROOT) + out.append(f"### `{rel}`") + out.append("") + out.append("```python") + out.append(tf.read_text()) + out.append("```") + out.append("") + + out.append("## Task") + out.append("") + out.append( + dedent( + """\ + For each surviving mutant listed above, write a new test in the + existing test file (matching its conventions, fixtures, and naming + style) that: + + - **Fails** when the mutated version of the function is in place. + - **Passes** when the original (correct) version is in place. + + Aim for one test per surviving mutant. If multiple mutants in the + same function can be killed by a single test, that is fine — note + which mutant numbers in the test name or docstring. + + Do not modify the source file. Only add tests. + """ + ).strip() + ) + out.append("") + + return "\n".join(out) + + +def main() -> int: + config = load_mutmut_config() + + stats_file = ROOT / "mutants" / "mutmut-cicd-stats.json" + stats: dict | None = None + if stats_file.exists(): + try: + stats = json.loads(stats_file.read_text()) + except json.JSONDecodeError as exc: + print(f"warning: could not parse {stats_file}: {exc}", file=sys.stderr) + + survivors = get_survivors() + report = render(config, survivors, stats) + + out_path = ROOT / "mutation-report.md" + out_path.write_text(report) + print( + f"Wrote {out_path} ({len(survivors)} survivor" + f"{'s' if len(survivors) != 1 else ''}, {len(report)} chars)" + ) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tests/litellm_utils_tests/test_proxy_budget_reset.py b/tests/litellm_utils_tests/test_proxy_budget_reset.py index a64c6c7aa36..6240bedd3e6 100644 --- a/tests/litellm_utils_tests/test_proxy_budget_reset.py +++ b/tests/litellm_utils_tests/test_proxy_budget_reset.py @@ -233,6 +233,12 @@ async def test_reset_budget_endusers_partial_failure(): prisma_client.db.litellm_verificationtoken.update_many = AsyncMock( return_value={"count": 0} ) + # Mock db.litellm_organizationtable.update_many (used by reset_budget_for_orgs_linked_to_budgets) + prisma_client.db.litellm_organizationtable.update_many = AsyncMock( + return_value={"count": 0} + ) + # Mock db.litellm_tagtable.update_many (used by reset_budget_for_tags_linked_to_budgets) + prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0}) proxy_logging_obj = MagicMock() proxy_logging_obj.service_logging_obj = MagicMock() @@ -400,6 +406,12 @@ async def test_reset_budget_continues_other_categories_on_failure(): prisma_client.db.litellm_verificationtoken.update_many = AsyncMock( return_value={"count": 0} ) + # Mock db.litellm_organizationtable.update_many (used by reset_budget_for_orgs_linked_to_budgets) + prisma_client.db.litellm_organizationtable.update_many = AsyncMock( + return_value={"count": 0} + ) + # Mock db.litellm_tagtable.update_many (used by reset_budget_for_tags_linked_to_budgets) + prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0}) proxy_logging_obj = MagicMock() proxy_logging_obj.service_logging_obj = MagicMock() @@ -884,6 +896,12 @@ async def test_service_logger_endusers_success(): prisma_client.db.litellm_verificationtoken.update_many = AsyncMock( return_value={"count": 0} ) + # Mock db.litellm_organizationtable.update_many (used by reset_budget_for_orgs_linked_to_budgets) + prisma_client.db.litellm_organizationtable.update_many = AsyncMock( + return_value={"count": 0} + ) + # Mock db.litellm_tagtable.update_many (used by reset_budget_for_tags_linked_to_budgets) + prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0}) proxy_logging_obj = MagicMock() proxy_logging_obj.service_logging_obj = MagicMock() @@ -966,6 +984,12 @@ async def test_service_logger_endusers_failure(): prisma_client.db.litellm_verificationtoken.update_many = AsyncMock( return_value={"count": 0} ) + # Mock db.litellm_organizationtable.update_many (used by reset_budget_for_orgs_linked_to_budgets) + prisma_client.db.litellm_organizationtable.update_many = AsyncMock( + return_value={"count": 0} + ) + # Mock db.litellm_tagtable.update_many (used by reset_budget_for_tags_linked_to_budgets) + prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0}) proxy_logging_obj = MagicMock() proxy_logging_obj.service_logging_obj = MagicMock() @@ -1060,6 +1084,10 @@ async def test_reset_budget_for_litellm_team_members_called(): prisma_client.db.litellm_verificationtoken.update_many = AsyncMock( return_value={"count": 0} ) + prisma_client.db.litellm_organizationtable.update_many = AsyncMock( + return_value={"count": 0} + ) + prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0}) proxy_logging_obj = MagicMock() proxy_logging_obj.service_logging_obj = MagicMock() diff --git a/tests/llm_translation/test_openrouter.py b/tests/llm_translation/test_openrouter.py index 631b0770e3d..8fbb8803d11 100644 --- a/tests/llm_translation/test_openrouter.py +++ b/tests/llm_translation/test_openrouter.py @@ -11,7 +11,7 @@ import litellm def test_completion_openrouter_reasoning_content(): litellm._turn_on_debug() resp = litellm.completion( - model="openrouter/anthropic/claude-3.7-sonnet", + model="openrouter/anthropic/claude-sonnet-4", messages=[{"role": "user", "content": "Hello world"}], reasoning={"effort": "high"}, ) diff --git a/tests/local_testing/test_function_calling.py b/tests/local_testing/test_function_calling.py index 6fd253ee294..3c7e004b62e 100644 --- a/tests/local_testing/test_function_calling.py +++ b/tests/local_testing/test_function_calling.py @@ -268,51 +268,63 @@ def test_aaparallel_function_call_with_anthropic_thinking(model): from litellm.types.utils import ChatCompletionMessageToolCall, Function, Message +_PARALLEL_TOOL_HISTORY_MESSAGES = [ + { + "role": "user", + "content": "What's the weather like in San Francisco, Tokyo, and Paris? - give me 3 responses", + }, + Message( + content="Here are the current weather conditions for San Francisco, Tokyo, and Paris:", + role="assistant", + tool_calls=[ + ChatCompletionMessageToolCall( + index=1, + function=Function( + arguments='{"location": "San Francisco, CA", "unit": "fahrenheit"}', + name="get_current_weather", + ), + id="tooluse_Jj98qn6xQlOP_PiQr-w9iA", + type="function", + ) + ], + function_call=None, + ), + { + "tool_call_id": "tooluse_Jj98qn6xQlOP_PiQr-w9iA", + "role": "tool", + "name": "get_current_weather", + "content": '{"location": "San Francisco", "temperature": "72", "unit": "fahrenheit"}', + }, +] + + @pytest.mark.parametrize( - "model, provider", + "model, messages, expect_unsupported_params_error", [ + # Bedrock Converse still requires modify_params to inject the dummy tool. ( "anthropic.claude-3-sonnet-20240229-v1:0", - "bedrock", + _PARALLEL_TOOL_HISTORY_MESSAGES, + True, ), - ("claude-haiku-4-5-20251001", "anthropic"), - ], -) -@pytest.mark.parametrize( - "messages, expected_error_msg", - [ + # Anthropic Messages API: dummy tool is injected without modify_params. ( + "claude-haiku-4-5-20251001", + _PARALLEL_TOOL_HISTORY_MESSAGES, + False, + ), + ( + "anthropic.claude-3-sonnet-20240229-v1:0", [ { "role": "user", "content": "What's the weather like in San Francisco, Tokyo, and Paris? - give me 3 responses", - }, - Message( - content="Here are the current weather conditions for San Francisco, Tokyo, and Paris:", - role="assistant", - tool_calls=[ - ChatCompletionMessageToolCall( - index=1, - function=Function( - arguments='{"location": "San Francisco, CA", "unit": "fahrenheit"}', - name="get_current_weather", - ), - id="tooluse_Jj98qn6xQlOP_PiQr-w9iA", - type="function", - ) - ], - function_call=None, - ), - { - "tool_call_id": "tooluse_Jj98qn6xQlOP_PiQr-w9iA", - "role": "tool", - "name": "get_current_weather", - "content": '{"location": "San Francisco", "temperature": "72", "unit": "fahrenheit"}', - }, + } ], - True, + False, ), ( + "claude-haiku-4-5-20251001", [ { "role": "user", @@ -324,25 +336,26 @@ from litellm.types.utils import ChatCompletionMessageToolCall, Function, Message ], ) def test_parallel_function_call_anthropic_error_msg( - model, provider, messages, expected_error_msg + model, messages, expect_unsupported_params_error ): """ - Anthropic doesn't support tool calling without `tools=` param specified. + Tool history without an explicit ``tools`` param: - Ensure this error is thrown when `tools=` param is not specified. But tool call requests are made. + - Bedrock **Converse** still raises ``UnsupportedParamsError`` unless + ``litellm.modify_params`` is enabled (dummy tool is only added there). + - **Anthropic** (and Bedrock Invoke via ``AnthropicConfig.transform_request``) + always get a dummy tool so CLIs work with ``modify_params`` left off. Reference Issue: https://github.com/BerriAI/litellm/issues/5747, https://github.com/BerriAI/litellm/issues/5388 """ - # Ensure modify_params is False so UnsupportedParamsError is raised + # Ensure modify_params is False so Bedrock Converse path still raises. # (other tests in this file set it to True and don't reset it) original_modify_params = litellm.modify_params litellm.modify_params = False try: litellm.set_verbose = True - messages = messages - - if expected_error_msg: + if expect_unsupported_params_error: with pytest.raises(litellm.UnsupportedParamsError) as e: second_response = litellm.completion( model=model, diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index 70232f25c37..68aff36038b 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -587,12 +587,21 @@ def test_foward_litellm_user_info_to_backend_llm_call(): user_api_key_dict=user_api_key_dict, ) + # All header values must be str/bytes so httpx won't reject them when the + # downstream client builds the request (regression: #27458). + for k, v in data.items(): + assert isinstance(v, (str, bytes)), ( + f"header {k!r} has non-str value {v!r} ({type(v).__name__}); " + "httpx will raise 'Header value must be str or bytes' when the LLM " + "request is built." + ) + expected_data = { "x-litellm-user_api_key_user_id": "test_user_id", "x-litellm-user_api_key_org_id": "test_org_id", "x-litellm-user_api_key_hash": "test_api_key", - "x-litellm-user_api_key_spend": 0.0, - "x-litellm-user_api_key_auth_metadata": {}, + "x-litellm-user_api_key_spend": "0.0", + "x-litellm-user_api_key_auth_metadata": "{}", } assert json.dumps(data, sort_keys=True) == json.dumps(expected_data, sort_keys=True) diff --git a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py index e40543e01a0..697a9ebc720 100644 --- a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py +++ b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py @@ -2098,6 +2098,56 @@ def test_map_optional_params_preserves_reasoning_summary(): assert responses_api_request["reasoning"]["summary"] == "detailed" +def test_map_optional_params_tool_choice_chat_nested_to_responses_api(): + """Chat tool_choice must become Responses ToolChoiceFunction (top-level name).""" + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + LiteLLMResponsesTransformationHandler, + ) + from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams + + handler = LiteLLMResponsesTransformationHandler() + responses_api_request = ResponsesAPIOptionalRequestParams() + handler._map_optional_params_to_responses_api_request( + { + "stream": False, + "tool_choice": { + "type": "function", + "function": {"name": "Echo"}, + }, + }, + responses_api_request, + ) + assert responses_api_request["tool_choice"] == { + "type": "function", + "name": "Echo", + } + + +@pytest.mark.parametrize( + ("tool_choice", "expected"), + [ + ("auto", "auto"), + ("none", "none"), + ( + {"type": "function", "name": "Echo"}, + {"type": "function", "name": "Echo"}, + ), + ( + {"type": "function", "name": "foo", "function": {"name": "bar"}}, + {"type": "function", "name": "foo"}, + ), + ({"type": "required"}, {"type": "required"}), + ], +) +def test_normalize_tool_choice_for_responses_api(tool_choice, expected): + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + LiteLLMResponsesTransformationHandler, + ) + + handler = LiteLLMResponsesTransformationHandler() + assert handler._normalize_tool_choice_for_responses_api(tool_choice) == expected + + def test_convert_chat_completion_file_type_to_input_file(): """ Test that Chat Completion content with type 'file' is correctly mapped diff --git a/tests/test_litellm/integrations/test_opentelemetry.py b/tests/test_litellm/integrations/test_opentelemetry.py index bea718dfea0..898f42b45bf 100644 --- a/tests/test_litellm/integrations/test_opentelemetry.py +++ b/tests/test_litellm/integrations/test_opentelemetry.py @@ -3159,3 +3159,630 @@ class TestResponseIdFallback(unittest.TestCase): otel.set_attributes(mock_span, kwargs, response_obj) mock_span.set_attribute.assert_any_call("litellm.call_id", call_id) + + + +class TestOpenTelemetryResponsesAPI(unittest.TestCase): + """ + Tests for Responses API (/v1/responses) OTel span attributes. + + The Responses API uses ``output`` (list of output items) instead of + ``choices``, ``instructions`` instead of ``system_instructions``, and + ``status`` instead of per-choice ``finish_reason``. + + See: https://github.com/BerriAI/litellm/issues/25840 + """ + + def _base_kwargs(self, **overrides): + """Return minimal kwargs for set_attributes with Responses API defaults.""" + kwargs = { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "What is 2+2?"}], + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "openai"}, + "standard_logging_object": { + "id": "resp_abc123", + "call_type": "responses", + "metadata": {}, + }, + } + kwargs.update(overrides) + return kwargs + + def _responses_api_response_obj(self, text="The answer is 4.", status="completed"): + """Return a dict mimicking ResponsesAPIResponse with a message output.""" + return { + "id": "resp_abc123", + "model": "gpt-4o", + "status": status, + "output": [ + { + "type": "message", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": text, + } + ], + } + ], + "usage": { + "prompt_tokens": 10, + "completion_tokens": 20, + "total_tokens": 30, + }, + } + + def _get_attr(self, mock_span, attr_name): + """Extract the value set for a specific attribute name, or None.""" + calls = [ + call + for call in mock_span.set_attribute.call_args_list + if call[0][0] == attr_name + ] + if not calls: + return None + return calls[0][0][1] + + # ------------------------------------------------------------------ + # gen_ai.output.messages + # ------------------------------------------------------------------ + + def test_output_messages_populated_for_responses_api(self): + """gen_ai.output.messages must be set when response has output items.""" + otel = OpenTelemetry() + mock_span = MagicMock() + + kwargs = self._base_kwargs() + response_obj = self._responses_api_response_obj(text="The answer is 4.") + + otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj) + + raw = self._get_attr(mock_span, "gen_ai.output.messages") + self.assertIsNotNone(raw, "gen_ai.output.messages should be set") + + parsed = json.loads(raw) + self.assertIsInstance(parsed, list) + self.assertEqual(len(parsed), 1) + self.assertEqual(parsed[0]["role"], "assistant") + self.assertIn("parts", parsed[0]) + self.assertEqual(parsed[0]["parts"][0]["type"], "text") + self.assertEqual(parsed[0]["parts"][0]["content"], "The answer is 4.") + + def test_output_messages_with_multiple_content_items(self): + """Multiple output_text items in a single message should all appear as parts.""" + otel = OpenTelemetry() + mock_span = MagicMock() + + response_obj = { + "id": "resp_multi", + "model": "gpt-4o", + "status": "completed", + "output": [ + { + "type": "message", + "role": "assistant", + "content": [ + {"type": "output_text", "text": "First paragraph."}, + {"type": "output_text", "text": "Second paragraph."}, + ], + } + ], + } + + otel.set_attributes( + span=mock_span, kwargs=self._base_kwargs(), response_obj=response_obj + ) + + raw = self._get_attr(mock_span, "gen_ai.output.messages") + parsed = json.loads(raw) + self.assertEqual(len(parsed[0]["parts"]), 2) + self.assertEqual(parsed[0]["parts"][0]["content"], "First paragraph.") + self.assertEqual(parsed[0]["parts"][1]["content"], "Second paragraph.") + + def test_output_messages_with_function_call(self): + """function_call output items should appear as tool_call parts.""" + otel = OpenTelemetry() + mock_span = MagicMock() + + response_obj = { + "id": "resp_fc", + "model": "gpt-4o", + "status": "completed", + "output": [ + { + "type": "function_call", + "name": "get_weather", + "call_id": "call_abc", + "arguments": '{"location": "SF"}', + } + ], + } + + otel.set_attributes( + span=mock_span, kwargs=self._base_kwargs(), response_obj=response_obj + ) + + raw = self._get_attr(mock_span, "gen_ai.output.messages") + parsed = json.loads(raw) + self.assertEqual(len(parsed), 1) + self.assertEqual(parsed[0]["role"], "assistant") + self.assertEqual(parsed[0]["parts"][0]["type"], "tool_call") + self.assertEqual(parsed[0]["parts"][0]["name"], "get_weather") + self.assertEqual(parsed[0]["parts"][0]["arguments"], '{"location": "SF"}') + self.assertEqual(parsed[0]["parts"][0]["id"], "call_abc") + + def test_output_messages_mixed_message_and_function_call(self): + """Mixed output with both message and function_call items.""" + otel = OpenTelemetry() + mock_span = MagicMock() + + response_obj = { + "id": "resp_mixed", + "model": "gpt-4o", + "status": "completed", + "output": [ + { + "type": "message", + "role": "assistant", + "content": [ + {"type": "output_text", "text": "Let me check the weather."}, + ], + }, + { + "type": "function_call", + "name": "get_weather", + "call_id": "call_xyz", + "arguments": "{}", + }, + ], + } + + otel.set_attributes( + span=mock_span, kwargs=self._base_kwargs(), response_obj=response_obj + ) + + raw = self._get_attr(mock_span, "gen_ai.output.messages") + parsed = json.loads(raw) + self.assertEqual(len(parsed), 2) + self.assertEqual(parsed[0]["role"], "assistant") + self.assertEqual(parsed[0]["parts"][0]["content"], "Let me check the weather.") + self.assertEqual(parsed[1]["parts"][0]["type"], "tool_call") + + def test_output_messages_empty_text_skipped(self): + """Output items with empty text should not produce parts.""" + otel = OpenTelemetry() + mock_span = MagicMock() + + response_obj = { + "id": "resp_empty", + "model": "gpt-4o", + "status": "completed", + "output": [ + { + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": ""}], + } + ], + } + + otel.set_attributes( + span=mock_span, kwargs=self._base_kwargs(), response_obj=response_obj + ) + + # No output messages should be set since the text is empty + raw = self._get_attr(mock_span, "gen_ai.output.messages") + self.assertIsNone(raw, "Empty output text should not produce gen_ai.output.messages") + + def test_choices_still_work(self): + """Existing choices-based responses must still work (no regression).""" + otel = OpenTelemetry() + mock_span = MagicMock() + + kwargs = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "openai"}, + "standard_logging_object": { + "id": "test-id", + "call_type": "completion", + "metadata": {}, + }, + } + + response_obj = { + "id": "chatcmpl-123", + "model": "gpt-4", + "choices": [ + { + "finish_reason": "stop", + "message": {"role": "assistant", "content": "Hi there!"}, + } + ], + "usage": {"prompt_tokens": 5, "completion_tokens": 10, "total_tokens": 15}, + } + + otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj) + + raw = self._get_attr(mock_span, "gen_ai.output.messages") + parsed = json.loads(raw) + self.assertEqual(parsed[0]["parts"][0]["content"], "Hi there!") + self.assertEqual(parsed[0]["finish_reason"], "stop") + + # ------------------------------------------------------------------ + # gen_ai.response.finish_reasons + # ------------------------------------------------------------------ + + def test_finish_reasons_from_status(self): + """gen_ai.response.finish_reasons should use ResponsesAPIResponse.status.""" + otel = OpenTelemetry() + mock_span = MagicMock() + + otel.set_attributes( + span=mock_span, + kwargs=self._base_kwargs(), + response_obj=self._responses_api_response_obj(status="completed"), + ) + + raw = self._get_attr(mock_span, "gen_ai.response.finish_reasons") + self.assertIsNotNone(raw) + parsed = json.loads(raw) + self.assertEqual(parsed, ["completed"]) + + def test_finish_reasons_incomplete_status(self): + """Non-completed status values should still be captured.""" + otel = OpenTelemetry() + mock_span = MagicMock() + + otel.set_attributes( + span=mock_span, + kwargs=self._base_kwargs(), + response_obj=self._responses_api_response_obj(status="incomplete"), + ) + + raw = self._get_attr(mock_span, "gen_ai.response.finish_reasons") + parsed = json.loads(raw) + self.assertEqual(parsed, ["incomplete"]) + + # ------------------------------------------------------------------ + # gen_ai.system_instructions + # ------------------------------------------------------------------ + + def test_system_instructions_from_instructions_kwarg(self): + """Responses API passes system prompt as kwargs['instructions'].""" + otel = OpenTelemetry() + mock_span = MagicMock() + + kwargs = self._base_kwargs(instructions="You are a math tutor.") + response_obj = self._responses_api_response_obj() + + otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj) + + value = self._get_attr(mock_span, "gen_ai.system_instructions") + self.assertEqual(value, "You are a math tutor.") + + def test_system_instructions_from_system_kwarg(self): + """Anthropic Messages API passes system prompt as kwargs['system'].""" + otel = OpenTelemetry() + mock_span = MagicMock() + + kwargs = self._base_kwargs(system="You are a helpful assistant.") + response_obj = self._responses_api_response_obj() + + otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj) + + value = self._get_attr(mock_span, "gen_ai.system_instructions") + self.assertEqual(value, "You are a helpful assistant.") + + def test_system_instructions_from_system_instructions_kwarg(self): + """Vertex AI Gemini path uses kwargs['system_instructions'] (existing behavior).""" + otel = OpenTelemetry() + mock_span = MagicMock() + + kwargs = self._base_kwargs( + system_instructions=[{"role": "system", "content": "Be concise."}] + ) + response_obj = self._responses_api_response_obj() + + otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj) + + raw = self._get_attr(mock_span, "gen_ai.system_instructions") + self.assertIsNotNone(raw) + parsed = json.loads(raw) + self.assertEqual(parsed[0]["role"], "system") + self.assertIn("parts", parsed[0]) + + def test_system_instructions_precedence(self): + """system_instructions takes precedence over instructions and system.""" + otel = OpenTelemetry() + mock_span = MagicMock() + + kwargs = self._base_kwargs( + system_instructions="From Gemini", + instructions="From Responses API", + system="From Anthropic", + ) + response_obj = self._responses_api_response_obj() + + otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj) + + # system_instructions (string) should win — it's checked first + value = self._get_attr(mock_span, "gen_ai.system_instructions") + self.assertEqual(value, "From Gemini") + + def test_no_system_instructions_when_absent(self): + """No gen_ai.system_instructions attr when none of the kwargs are set.""" + otel = OpenTelemetry() + mock_span = MagicMock() + + kwargs = self._base_kwargs() + response_obj = self._responses_api_response_obj() + + otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj) + + value = self._get_attr(mock_span, "gen_ai.system_instructions") + self.assertIsNone(value) + + +class TestTransformResponsesAPIOutput(unittest.TestCase): + """ + Unit tests for _transform_responses_api_output_to_otel. + """ + + def test_message_with_output_text(self): + otel = OpenTelemetry() + output = [ + { + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": "Hello!"}], + } + ] + result = otel._transform_responses_api_output_to_otel(output) + self.assertEqual(len(result), 1) + self.assertEqual(result[0]["role"], "assistant") + self.assertEqual(result[0]["parts"], [{"type": "text", "content": "Hello!"}]) + + def test_function_call_item(self): + otel = OpenTelemetry() + output = [ + { + "type": "function_call", + "name": "search", + "call_id": "call_1", + "arguments": '{"q": "test"}', + } + ] + result = otel._transform_responses_api_output_to_otel(output) + self.assertEqual(len(result), 1) + self.assertEqual(result[0]["role"], "assistant") + self.assertEqual(result[0]["parts"][0]["type"], "tool_call") + self.assertEqual(result[0]["parts"][0]["name"], "search") + self.assertEqual(result[0]["parts"][0]["id"], "call_1") + + def test_function_call_without_call_id(self): + otel = OpenTelemetry() + output = [ + { + "type": "function_call", + "name": "search", + "arguments": "{}", + } + ] + result = otel._transform_responses_api_output_to_otel(output) + self.assertNotIn("id", result[0]["parts"][0]) + + def test_unknown_type_ignored(self): + otel = OpenTelemetry() + output = [{"type": "reasoning", "content": "thinking..."}] + result = otel._transform_responses_api_output_to_otel(output) + self.assertEqual(result, []) + + def test_non_dict_items_ignored(self): + otel = OpenTelemetry() + output = ["not a dict", 42, None] + result = otel._transform_responses_api_output_to_otel(output) + self.assertEqual(result, []) + + def test_empty_output(self): + otel = OpenTelemetry() + result = otel._transform_responses_api_output_to_otel([]) + self.assertEqual(result, []) + + def test_message_with_empty_text_skipped(self): + otel = OpenTelemetry() + output = [ + { + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": ""}], + } + ] + result = otel._transform_responses_api_output_to_otel(output) + self.assertEqual(result, []) + + def test_message_default_role(self): + """Messages without explicit role should default to assistant.""" + otel = OpenTelemetry() + output = [ + { + "type": "message", + "content": [{"type": "output_text", "text": "Hi"}], + } + ] + result = otel._transform_responses_api_output_to_otel(output) + self.assertEqual(result[0]["role"], "assistant") + + + def test_pydantic_like_objects_accepted(self): + """Items with .get() but not isinstance(dict) should be accepted.""" + + class FakeOutputItem: + """Mimics BaseLiteLLMOpenAIResponseObject duck-typing.""" + + def __init__(self, data): + self._data = data + + def get(self, key, default=None): + return self._data.get(key, default) + + class FakeContent: + def __init__(self, data): + self._data = data + + def get(self, key, default=None): + return self._data.get(key, default) + + otel = OpenTelemetry() + output = [ + FakeOutputItem( + { + "type": "message", + "role": "assistant", + "content": [ + FakeContent({"type": "output_text", "text": "Pydantic works!"}), + ], + } + ) + ] + result = otel._transform_responses_api_output_to_otel(output) + self.assertEqual(len(result), 1) + self.assertEqual(result[0]["parts"][0]["content"], "Pydantic works!") + + +class TestSystemInstructionsPrecedence(unittest.TestCase): + """Tests for the is-not-None precedence in system_instructions coalescing.""" + + def _get_attr(self, mock_span, attr_name): + calls = [ + call + for call in mock_span.set_attribute.call_args_list + if call[0][0] == attr_name + ] + if not calls: + return None + return calls[0][0][1] + + def _base_kwargs(self, **overrides): + kwargs = { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "Hi"}], + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "openai"}, + "standard_logging_object": { + "id": "test-id", + "call_type": "responses", + "metadata": {}, + }, + } + kwargs.update(overrides) + return kwargs + + def test_empty_list_system_instructions_does_not_fallthrough(self): + """An empty list for system_instructions should NOT fall through to instructions.""" + otel = OpenTelemetry() + mock_span = MagicMock() + + kwargs = self._base_kwargs( + system_instructions=[], + instructions="Should not be used", + ) + response_obj = {"id": "r1", "model": "gpt-4o"} + + otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj) + + # system_instructions is [] (falsy but not None), so it wins. + # Since it's an empty list, no attribute should be set (nothing to transform). + value = self._get_attr(mock_span, "gen_ai.system_instructions") + # The empty list is truthy for `is not None` but produces empty + # transformed output — the attribute should NOT contain "Should not be used". + if value is not None: + self.assertNotIn("Should not be used", str(value)) + + +class TestResponsesAPIToolCallSpanAttributes(unittest.TestCase): + """Tests for per-tool-call span attributes on Responses API function_call items.""" + + def _base_kwargs(self): + return { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "What is the weather?"}], + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "openai"}, + "standard_logging_object": { + "id": "resp_tc", + "call_type": "responses", + "metadata": {}, + }, + } + + def test_per_tool_call_attributes_emitted(self): + """function_call output items should produce per-tool-call span attributes.""" + otel = OpenTelemetry() + mock_span = MagicMock() + + response_obj = { + "id": "resp_tc", + "model": "gpt-4o", + "status": "completed", + "output": [ + { + "type": "function_call", + "name": "get_weather", + "call_id": "call_abc", + "arguments": '{"location": "SF"}', + } + ], + } + + otel.set_attributes(span=mock_span, kwargs=self._base_kwargs(), response_obj=response_obj) + + # Verify per-tool-call attributes were set (same format as choices branch) + attr_names = [call[0][0] for call in mock_span.set_attribute.call_args_list] + tool_call_attrs = [a for a in attr_names if "function_call" in a] + self.assertTrue(len(tool_call_attrs) > 0, "Per-tool-call span attributes should be emitted") + + # Verify the name attribute specifically + mock_span.set_attribute.assert_any_call( + "gen_ai.completion.0.function_call.name", "get_weather" + ) + mock_span.set_attribute.assert_any_call( + "gen_ai.completion.0.function_call.arguments", '{"location": "SF"}' + ) + + def test_multiple_tool_calls_indexed(self): + """Multiple function_call items should be indexed correctly.""" + otel = OpenTelemetry() + mock_span = MagicMock() + + response_obj = { + "id": "resp_tc2", + "model": "gpt-4o", + "status": "completed", + "output": [ + { + "type": "function_call", + "name": "get_weather", + "call_id": "call_1", + "arguments": "{}", + }, + { + "type": "function_call", + "name": "get_time", + "call_id": "call_2", + "arguments": "{}", + }, + ], + } + + otel.set_attributes(span=mock_span, kwargs=self._base_kwargs(), response_obj=response_obj) + + mock_span.set_attribute.assert_any_call( + "gen_ai.completion.0.function_call.name", "get_weather" + ) + mock_span.set_attribute.assert_any_call( + "gen_ai.completion.1.function_call.name", "get_time" + ) diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py index e38698c9100..a19752dc648 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py @@ -2135,6 +2135,53 @@ def test_validate_effort_for_model_centralises_per_model_gating( assert err is None +def test_transform_request_injects_dummy_tool_without_tools_param(): + """ + Anthropic rejects messages that contain tool turns when ``tools`` is omitted. + LiteLLM must inject a dummy tool without ``litellm.modify_params``. + """ + config = AnthropicConfig() + prev_modify_params = litellm.modify_params + litellm.modify_params = False + try: + messages = [ + {"role": "user", "content": "Hello"}, + { + "role": "assistant", + "content": "Calling tool", + "tool_calls": [ + { + "id": "toolu_test_dummy", + "type": "function", + "function": {"name": "get_x", "arguments": "{}"}, + } + ], + }, + { + "role": "tool", + "tool_call_id": "toolu_test_dummy", + "content": "{}", + }, + ] + result = config.transform_request( + model="claude-3-5-haiku-20241022", + messages=messages, + optional_params={"max_tokens": 256}, + litellm_params={}, + headers={}, + ) + finally: + litellm.modify_params = prev_modify_params + + assert "tools" in result + names = [ + t.get("name") + for t in result["tools"] + if isinstance(t, dict) and t.get("name") is not None + ] + assert "dummy_tool" in names + + def test_transform_request_uses_dynamic_max_tokens(): """ Test that transform_request uses dynamic max_tokens based on model diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py index 33628e1d19d..8deb3c86ad0 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py @@ -499,3 +499,98 @@ class TestThinkingSummaryPreservation: assert result == { "reasoning_effort": {"effort": "medium", "summary": "concise"} } + + +def test_anthropic_messages_handler_passes_litellm_params_timeout_to_base_handler(): + """The wrapper in this module must read `litellm_params.timeout` (set by + the router from the caller's `timeout` / `stream_timeout`) and forward it + to `BaseLLMHTTPHandler.anthropic_messages_handler` so the proxy honors + per-request timeouts on /v1/messages — same as /chat/completions does.""" + from litellm.llms.anthropic.experimental_pass_through.messages.handler import ( + anthropic_messages_handler, + base_llm_http_handler, + ) + + with patch.object( + base_llm_http_handler, + "anthropic_messages_handler", + return_value=MagicMock(), + ) as mock_base: + try: + anthropic_messages_handler( + max_tokens=100, + messages=[{"role": "user", "content": "hi"}], + model="anthropic/claude-sonnet-4-20250514", + api_key="test-key", + timeout=0.5, + ) + except (ValueError, TypeError, AttributeError): + # downstream initialization may fail under mocks; we only care + # about what was passed to the base handler + pass + + mock_base.assert_called_once() + assert mock_base.call_args.kwargs["timeout"] == 0.5 + + +def test_anthropic_messages_handler_coerces_string_timeout_to_float(): + """`GenericLiteLLMParams.timeout` allows `str`, but the httpx-level + handlers expect `float | httpx.Timeout`. The wrapper must coerce string + inputs to float before forwarding (mirrors the pattern in + `litellm/main.py::_sleep_for_timeout`).""" + from litellm.llms.anthropic.experimental_pass_through.messages.handler import ( + anthropic_messages_handler, + base_llm_http_handler, + ) + + with patch.object( + base_llm_http_handler, + "anthropic_messages_handler", + return_value=MagicMock(), + ) as mock_base: + try: + anthropic_messages_handler( + max_tokens=100, + messages=[{"role": "user", "content": "hi"}], + model="anthropic/claude-sonnet-4-20250514", + api_key="test-key", + timeout="0.5", + ) + except (ValueError, TypeError, AttributeError): + pass + + mock_base.assert_called_once() + assert mock_base.call_args.kwargs["timeout"] == 0.5 + assert isinstance(mock_base.call_args.kwargs["timeout"], float) + + +def test_anthropic_messages_handler_prefers_stream_timeout_for_streaming_calls(): + """When `stream=True` and `litellm_params.stream_timeout` is set, the + wrapper must forward `stream_timeout` (not `timeout`) to the base + handler — mirrors `Router._get_timeout` so direct (non-router) callers + get the same stream/non-stream timeout selection as proxy callers.""" + from litellm.llms.anthropic.experimental_pass_through.messages.handler import ( + anthropic_messages_handler, + base_llm_http_handler, + ) + + with patch.object( + base_llm_http_handler, + "anthropic_messages_handler", + return_value=MagicMock(), + ) as mock_base: + try: + anthropic_messages_handler( + max_tokens=100, + messages=[{"role": "user", "content": "hi"}], + model="anthropic/claude-sonnet-4-20250514", + api_key="test-key", + stream=True, + timeout=10.0, + stream_timeout=0.5, + ) + except (ValueError, TypeError, AttributeError): + pass + + mock_base.assert_called_once() + assert mock_base.call_args.kwargs["timeout"] == 0.5 diff --git a/tests/test_litellm/llms/bedrock/batches/test_handler.py b/tests/test_litellm/llms/bedrock/batches/test_handler.py new file mode 100644 index 00000000000..18780ccce0f --- /dev/null +++ b/tests/test_litellm/llms/bedrock/batches/test_handler.py @@ -0,0 +1,338 @@ +"""Unit tests for ``BedrockBatchesHandler._handle_model_invocation_job_status``. + +These cover the upstream support for retrieving Bedrock bulk batch jobs +(``arn:aws:bedrock:::model-invocation-job/``) — the ARN +type returned by ``CreateModelInvocationJob``. We mock the boto3 client so +the tests don't hit AWS. +""" + +from __future__ import annotations + +import os +import sys +from datetime import datetime, timezone +from unittest.mock import MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.bedrock.batches.handler import ( # noqa: E402 + BedrockBatchesHandler, + _extract_job_id_from_arn, + _extract_region_from_bedrock_arn, + _predict_output_file_uri, + _to_epoch, +) + +JOB_ID = "abc1234567" +JOB_ARN = f"arn:aws:bedrock:us-west-2:123456789012:model-invocation-job/{JOB_ID}" +INPUT_URI = "s3://my-bucket/inputs/qwen3-235b-a22b-2507-batch.jsonl" +OUTPUT_PREFIX = "s3://my-bucket/litellm-batch-outputs/litellm-bedrock-files-qwen-uuid/" +SUBMIT_TIME = datetime(2026, 4, 28, 12, 0, 0, tzinfo=timezone.utc) +END_TIME = datetime(2026, 4, 28, 12, 30, 0, tzinfo=timezone.utc) + + +def _fake_boto3_response(status: str = "Completed", end_time=END_TIME): + return { + "jobArn": JOB_ARN, + "jobName": "litellm-bedrock-files-qwen-uuid", + "modelId": "bedrock/qwen.qwen3-235b-a22b-2507-v1:0", + "status": status, + "submitTime": SUBMIT_TIME, + "lastModifiedTime": end_time, + "endTime": end_time, + "inputDataConfig": {"s3InputDataConfig": {"s3Uri": INPUT_URI}}, + "outputDataConfig": {"s3OutputDataConfig": {"s3Uri": OUTPUT_PREFIX}}, + } + + +@pytest.fixture +def patched_boto3(): + """Yield a stub bedrock client whose `get_model_invocation_job` is a MagicMock.""" + fake_client = MagicMock() + fake_client.get_model_invocation_job.return_value = _fake_boto3_response() + with ( + patch("boto3.client", return_value=fake_client) as boto_client_factory, + patch( + "litellm.llms.bedrock.batches.transformation.BedrockBatchesConfig.get_credentials", + return_value=MagicMock(access_key="AKIA", secret_key="SECRET", token=None), + ), + ): + yield fake_client, boto_client_factory + + +def test_extract_region_from_arn(): + assert _extract_region_from_bedrock_arn(JOB_ARN) == "us-west-2" + assert _extract_region_from_bedrock_arn("arn:aws:bedrock::123:foo/bar") is None + assert _extract_region_from_bedrock_arn("not-an-arn") is None + + +def test_extract_region_swallows_unexpected_split_errors(): + """Defensive `except Exception` branch — anything that isn't a plain str + should fall through to ``None`` rather than blow up.""" + + class WeirdArn: + def split(self, _sep): + raise RuntimeError("boom") + + assert _extract_region_from_bedrock_arn(WeirdArn()) is None # type: ignore[arg-type] + + +def test_predict_output_file_uri_returns_none_for_directory_input_uri(): + """Input URI ending in `/` has an empty basename — we must bail rather + than emit ``//.out``.""" + assert ( + _predict_output_file_uri(OUTPUT_PREFIX, "s3://bucket/inputs/", JOB_ID) is None + ) + + +_DT = datetime(2026, 4, 28, 12, 0, 0, tzinfo=timezone.utc) + + +@pytest.mark.parametrize( + "value,expected", + [ + (None, None), + (1730000000, 1730000000), + (1730000000.5, 1730000000), + (_DT, int(_DT.timestamp())), + ("2026-04-28T12:00:00Z", None), # strings aren't supported -> None + ], +) +def test_to_epoch_handles_supported_types(value, expected): + assert _to_epoch(value) == expected + + +def test_extract_job_id_from_arn(): + assert _extract_job_id_from_arn(JOB_ARN) == JOB_ID + assert ( + _extract_job_id_from_arn("arn:aws:bedrock:us-west-2:1:async-invoke/x") is None + ) + + +def test_predict_output_file_uri_happy_path(): + expected = f"{OUTPUT_PREFIX}{JOB_ID}/qwen3-235b-a22b-2507-batch.jsonl.out" + assert _predict_output_file_uri(OUTPUT_PREFIX, INPUT_URI, JOB_ID) == expected + + +def test_predict_output_file_uri_adds_trailing_slash(): + prefix_no_slash = OUTPUT_PREFIX.rstrip("/") + expected = f"{OUTPUT_PREFIX}{JOB_ID}/qwen3-235b-a22b-2507-batch.jsonl.out" + assert _predict_output_file_uri(prefix_no_slash, INPUT_URI, JOB_ID) == expected + + +@pytest.mark.parametrize( + "missing_arg", + [ + ("", INPUT_URI, JOB_ID), + (OUTPUT_PREFIX, "", JOB_ID), + (OUTPUT_PREFIX, INPUT_URI, None), + ], +) +def test_predict_output_file_uri_returns_none_when_missing_input(missing_arg): + assert _predict_output_file_uri(*missing_arg) is None + + +def test_handle_model_invocation_job_status_completed(patched_boto3): + fake_client, boto_client_factory = patched_boto3 + + batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN) + + fake_client.get_model_invocation_job.assert_called_once_with(jobIdentifier=JOB_ARN) + + # Region should be sniffed from the ARN. + _, kwargs = boto_client_factory.call_args + assert kwargs["region_name"] == "us-west-2" + + assert batch.id == JOB_ARN + assert batch.status == "completed" + assert batch.input_file_id == INPUT_URI + expected_out = f"{OUTPUT_PREFIX}{JOB_ID}/qwen3-235b-a22b-2507-batch.jsonl.out" + assert batch.output_file_id == expected_out + assert batch.completed_at == int(END_TIME.timestamp()) + assert batch.failed_at is None + assert batch.cancelled_at is None + # Per-record counts aren't reported by GetModelInvocationJob, so we leave + # them zeroed; consumers should parse manifest.json.out for accurate counts. + assert batch.request_counts.total == 0 + assert batch.metadata["job_arn"] == JOB_ARN + assert batch.metadata["output_file_uri"] == expected_out + assert batch.metadata["output_s3_uri"] == OUTPUT_PREFIX + + +@pytest.mark.parametrize( + "bedrock_status,openai_status", + [ + ("Submitted", "validating"), + ("Validating", "validating"), + ("Scheduled", "validating"), + ("InProgress", "in_progress"), + ("Stopping", "cancelling"), + ("Stopped", "cancelled"), + ("Completed", "completed"), + ("PartiallyCompleted", "completed"), + ("Failed", "failed"), + ("Expired", "expired"), + # Unknown/unmapped Bedrock status falls back to "in_progress" so we + # don't 500 on a future AWS-side enum addition. + ("MyBrandNewStatus", "in_progress"), + ], +) +def test_status_mapping(patched_boto3, bedrock_status, openai_status): + fake_client, _ = patched_boto3 + fake_client.get_model_invocation_job.return_value = _fake_boto3_response( + status=bedrock_status + ) + + batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN) + + assert batch.status == openai_status + # output_file_id is only populated for terminal-completed jobs, so callers + # don't accidentally try to download a non-existent file mid-run. + if openai_status == "completed": + assert batch.output_file_id is not None + else: + assert batch.output_file_id is None + + +def test_explicit_region_overrides_arn(patched_boto3): + _, boto_client_factory = patched_boto3 + BedrockBatchesHandler._handle_model_invocation_job_status( + batch_id=JOB_ARN, aws_region_name="eu-central-1" + ) + _, kwargs = boto_client_factory.call_args + assert kwargs["region_name"] == "eu-central-1" + + +def test_failure_message_propagates(patched_boto3): + fake_client, _ = patched_boto3 + failed_response = _fake_boto3_response(status="Failed") + failed_response["message"] = "Input file failed validation" + fake_client.get_model_invocation_job.return_value = failed_response + + batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN) + + assert batch.status == "failed" + assert batch.failed_at == int(END_TIME.timestamp()) + assert batch.metadata["failure_message"] == "Input file failed validation" + + +def test_completed_with_unpredictable_output_uri_stays_none(patched_boto3): + """ + Regression guard for the original NoSuchKey bug: if Bedrock's response is + missing pieces we need to compute the per-job output file path (here, the + input s3Uri), `output_file_id` must stay `None` rather than fall back to + the bare prefix. Falling back to the prefix is what produced the original + NoSuchKey error this PR fixes. + """ + fake_client, _ = patched_boto3 + incomplete_response = _fake_boto3_response(status="Completed") + incomplete_response["inputDataConfig"] = {"s3InputDataConfig": {"s3Uri": ""}} + fake_client.get_model_invocation_job.return_value = incomplete_response + + batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN) + + assert batch.status == "completed" + # output_file_id MUST be None (not the bare prefix) — that's the whole + # point of this regression test. Callers branch on this field. + assert batch.output_file_id is None + # The metadata field uses "" because OpenAI Batch metadata is dict[str, str]; + # callers should branch on `output_file_id` (above) instead. + assert batch.metadata["output_file_uri"] == "" + # The bare prefix is still preserved in metadata so callers can list it. + assert batch.metadata["output_s3_uri"] == OUTPUT_PREFIX + + +def test_cancelled_status_sets_cancelled_at(patched_boto3): + fake_client, _ = patched_boto3 + fake_client.get_model_invocation_job.return_value = _fake_boto3_response( + status="Stopped" + ) + + batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN) + + assert batch.status == "cancelled" + assert batch.cancelled_at == int(END_TIME.timestamp()) + assert batch.completed_at is None + assert batch.failed_at is None + assert batch.expired_at is None + + +def test_expired_status_sets_expired_at(patched_boto3): + fake_client, _ = patched_boto3 + fake_client.get_model_invocation_job.return_value = _fake_boto3_response( + status="Expired" + ) + + batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN) + + assert batch.status == "expired" + assert batch.expired_at == int(END_TIME.timestamp()) + assert batch.completed_at is None + assert batch.failed_at is None + assert batch.cancelled_at is None + + +def test_logging_obj_pre_and_post_call_invoked(patched_boto3): + """`pre_call` / `post_call` get called with sensible payloads when a + `logging_obj` is supplied.""" + _, _ = patched_boto3 + logging_obj = MagicMock() + + BedrockBatchesHandler._handle_model_invocation_job_status( + batch_id=JOB_ARN, logging_obj=logging_obj + ) + + logging_obj.pre_call.assert_called_once() + logging_obj.post_call.assert_called_once() + + pre_kwargs = logging_obj.pre_call.call_args.kwargs + assert pre_kwargs["input"] == JOB_ARN + assert pre_kwargs["additional_args"]["complete_input_dict"] == { + "jobIdentifier": JOB_ARN + } + # Logged URL must use the bare job id, not the full ARN, so it doesn't + # double the `model-invocation-job/` segment or embed colons in the path. + assert pre_kwargs["additional_args"]["api_base"] == ( + f"https://bedrock.us-west-2.amazonaws.com/model-invocation-job/{JOB_ID}" + ) + + post_kwargs = logging_obj.post_call.call_args.kwargs + assert post_kwargs["input"] == JOB_ARN + assert post_kwargs["original_response"]["jobArn"] == JOB_ARN + + +def test_missing_boto3_raises_helpful_import_error(): + """If boto3 isn't installed we should raise a clear, actionable + ImportError rather than letting a NameError escape.""" + real_import = ( + __builtins__["__import__"] + if isinstance(__builtins__, dict) + else __builtins__.__import__ + ) + + def fake_import(name, *args, **kwargs): + if name == "boto3": + raise ImportError("No module named 'boto3'") + return real_import(name, *args, **kwargs) + + with patch("builtins.__import__", side_effect=fake_import): + with pytest.raises(ImportError, match="pip install boto3"): + BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN) + + +def test_logging_url_uses_bare_id_when_only_id_passed(patched_boto3): + """If the caller passes just the trailing job id (also valid for + `GetModelInvocationJob`), the logged URL should use it as-is.""" + _, _ = patched_boto3 + logging_obj = MagicMock() + + BedrockBatchesHandler._handle_model_invocation_job_status( + batch_id=JOB_ID, aws_region_name="us-west-2", logging_obj=logging_obj + ) + + pre_kwargs = logging_obj.pre_call.call_args.kwargs + assert pre_kwargs["additional_args"]["api_base"] == ( + f"https://bedrock.us-west-2.amazonaws.com/model-invocation-job/{JOB_ID}" + ) diff --git a/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py b/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py index c98f7840343..9ecdad1fcff 100644 --- a/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py +++ b/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py @@ -867,10 +867,12 @@ def test_bedrock_messages_explicit_output_config_wins_over_reasoning_effort(): def test_bedrock_messages_strips_context_management(): """ Ensure context_management is stripped from the request before sending to - Bedrock Invoke, which doesn't support this Anthropic-specific parameter. + Bedrock Invoke when it carries only LiteLLM-internal edits (e.g. + clear_thinking_20251015, which is consumed via thinking injection). - Claude Code sends context_management on every request; leaving it in the body - causes a 400 "context_management: Extra inputs are not permitted" from Bedrock. + Claude Code sends context_management on every request; leaving such edits + in the body causes a 400 "context_management: Extra inputs are not + permitted" from Bedrock. """ from litellm.types.router import GenericLiteLLMParams @@ -897,6 +899,77 @@ def test_bedrock_messages_strips_context_management(): assert result.get("max_tokens") == 4096 +def test_bedrock_messages_preserves_compact_context_management_and_adds_beta(): + """ + Bedrock InvokeModel supports compaction when paired with the + ``compact-2026-01-12`` anthropic-beta header, even though the Converse API + does not. The transformation should: + 1. Keep ``context_management`` with compact_20260112 edits in the body + (Bedrock rejects unknown top-level fields, but accepts this one with + the right beta). + 2. Auto-inject ``compact-2026-01-12`` into ``anthropic_beta``. + + Ref: https://github.com/BerriAI/litellm/issues/27532 + """ + from litellm.types.router import GenericLiteLLMParams + + cfg = AmazonAnthropicClaudeMessagesConfig() + messages = [{"role": "user", "content": [{"type": "text", "text": "Hi"}]}] + optional_params = { + "max_tokens": 4096, + "context_management": { + "edits": [{"type": "compact_20260112"}] + }, + } + + result = cfg.transform_anthropic_messages_request( + model="anthropic.claude-sonnet-4-6-20250929-v1:0", + messages=messages, + anthropic_messages_optional_request_params=optional_params, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert result.get("context_management") == { + "edits": [{"type": "compact_20260112"}] + } + assert "compact-2026-01-12" in result.get("anthropic_beta", []) + assert result["max_tokens"] == 4096 + + +def test_bedrock_messages_filters_unsupported_context_management_edits(): + """ + Mixed edit lists must drop the LiteLLM-internal ``clear_thinking_20251015`` + entries while keeping ``compact_20260112`` and adding the compact beta. + """ + from litellm.types.router import GenericLiteLLMParams + + cfg = AmazonAnthropicClaudeMessagesConfig() + messages = [{"role": "user", "content": [{"type": "text", "text": "Hi"}]}] + optional_params = { + "max_tokens": 4096, + "context_management": { + "edits": [ + {"type": "clear_thinking_20251015", "keep": "all"}, + {"type": "compact_20260112"}, + ] + }, + } + + result = cfg.transform_anthropic_messages_request( + model="anthropic.claude-sonnet-4-6-20250929-v1:0", + messages=messages, + anthropic_messages_optional_request_params=optional_params, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert result.get("context_management") == { + "edits": [{"type": "compact_20260112"}] + } + assert "compact-2026-01-12" in result.get("anthropic_beta", []) + + def test_bedrock_messages_allowlist_filters_anthropic_only_fields(): """ Bedrock Invoke rejects any top-level body field it doesn't recognize with diff --git a/tests/test_litellm/llms/bedrock/test_claude_platform_provider.py b/tests/test_litellm/llms/bedrock/test_claude_platform_provider.py new file mode 100644 index 00000000000..74cfbb265cd --- /dev/null +++ b/tests/test_litellm/llms/bedrock/test_claude_platform_provider.py @@ -0,0 +1,312 @@ +import json +from unittest.mock import patch + +import httpx +import pytest + + +def _anthropic_response(url: str) -> httpx.Response: + return httpx.Response( + status_code=200, + json={ + "id": "msg_test", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-6", + "content": [{"type": "text", "text": "ok"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 1, "output_tokens": 1}, + }, + request=httpx.Request("POST", url), + ) + + +def _capture_request(url: str, headers: dict, data: bytes | str | None) -> dict: + raw_body = data.decode("utf-8") if isinstance(data, bytes) else data or "{}" + return { + "path": httpx.URL(url).path, + "headers": headers, + "body": json.loads(raw_body), + } + + +def test_claude_platform_builds_default_messages_url_from_region(): + from litellm.llms.bedrock.claude_platform.transformation import ( + BedrockClaudePlatformConfig, + ) + + config = BedrockClaudePlatformConfig() + + assert ( + config.get_complete_url( + api_base=None, + api_key=None, + model="claude-sonnet-4-6", + optional_params={"aws_region_name": "us-west-2"}, + litellm_params={}, + ) + == "https://aws-external-anthropic.us-west-2.api.aws/v1/messages" + ) + + +def test_claude_platform_ignores_standard_anthropic_base_url(monkeypatch): + from litellm.llms.bedrock.claude_platform.transformation import ( + BedrockClaudePlatformConfig, + ) + + monkeypatch.setenv("ANTHROPIC_BASE_URL", "https://api.anthropic.example") + monkeypatch.setenv("ANTHROPIC_API_BASE", "https://api.anthropic-api.example") + + config = BedrockClaudePlatformConfig() + + assert ( + config.get_complete_url( + api_base=None, + api_key=None, + model="claude-sonnet-4-6", + optional_params={"aws_region_name": "us-west-2"}, + litellm_params={}, + ) + == "https://aws-external-anthropic.us-west-2.api.aws/v1/messages" + ) + + +def test_claude_platform_uses_bedrock_subroute(): + import litellm + from litellm.llms.bedrock.common_utils import BedrockModelInfo + + model, provider, _, _ = litellm.get_llm_provider( + model="bedrock/claude_platform/claude-sonnet-4-6" + ) + + assert provider == "bedrock" + assert model == "claude_platform/claude-sonnet-4-6" + assert BedrockModelInfo.get_bedrock_route(model) == "claude_platform" + assert BedrockModelInfo.get_claude_platform_model(model) == "claude-sonnet-4-6" + + +def test_claude_platform_requires_workspace_header(): + from litellm import AuthenticationError + from litellm.llms.bedrock.claude_platform.transformation import ( + BedrockClaudePlatformConfig, + ) + + config = BedrockClaudePlatformConfig() + + with pytest.raises(AuthenticationError) as exc_info: + config.validate_environment( + api_key="fake-platform-key", + headers={}, + model="claude-sonnet-4-6", + messages=[{"role": "user", "content": "hello"}], + optional_params={}, + litellm_params={}, + ) + + assert "workspace" in str(exc_info.value).lower() + + +def test_claude_platform_api_key_auth_sets_workspace_and_key_headers(): + from litellm.llms.bedrock.claude_platform.transformation import ( + BedrockClaudePlatformConfig, + ) + + config = BedrockClaudePlatformConfig() + headers = config.validate_environment( + api_key="fake-platform-key", + headers={"anthropic-beta": "skills-2025-10-02"}, + model="claude-sonnet-4-6", + messages=[{"role": "user", "content": "hello"}], + optional_params={"workspace_id": "wrkspc_test"}, + litellm_params={}, + ) + + assert headers["x-api-key"] == "fake-platform-key" + assert headers["anthropic-workspace-id"] == "wrkspc_test" + assert headers["anthropic-beta"] == "skills-2025-10-02" + + +def test_claude_platform_does_not_use_standard_anthropic_api_key(monkeypatch): + from litellm.llms.bedrock.claude_platform.transformation import ( + BedrockClaudePlatformConfig, + ) + + monkeypatch.setenv("ANTHROPIC_API_KEY", "standard-anthropic-key") + + config = BedrockClaudePlatformConfig() + headers = config.validate_environment( + api_key=None, + headers={}, + model="claude-sonnet-4-6", + messages=[{"role": "user", "content": "hello"}], + optional_params={"workspace_id": "wrkspc_test"}, + litellm_params={}, + ) + + assert "x-api-key" not in headers + + +def test_claude_platform_sigv4_signs_transformed_request_body(): + from litellm.llms.bedrock.claude_platform.transformation import ( + BedrockClaudePlatformConfig, + ) + + config = BedrockClaudePlatformConfig() + request_body = { + "model": "claude-sonnet-4-6", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 10, + } + + with patch.object( + config, + "_sign_request", + return_value=({"Authorization": "signed"}, json.dumps(request_body).encode()), + ) as mock_sign_request: + headers, signed_body = config.sign_request( + headers={"anthropic-workspace-id": "wrkspc_test"}, + optional_params={"aws_region_name": "us-west-2"}, + request_data=request_body, + api_base="https://aws-external-anthropic.us-west-2.api.aws/v1/messages", + api_key=None, + model="claude-sonnet-4-6", + ) + + assert signed_body == json.dumps(request_body).encode() + assert headers["Authorization"] == "signed" + mock_sign_request.assert_called_once() + assert ( + mock_sign_request.call_args.kwargs["service_name"] == "aws-external-anthropic" + ) + assert mock_sign_request.call_args.kwargs["request_data"] == request_body + + +def test_claude_platform_standard_anthropic_api_key_does_not_skip_sigv4(monkeypatch): + from litellm.llms.bedrock.claude_platform.transformation import ( + BedrockClaudePlatformConfig, + ) + + monkeypatch.setenv("ANTHROPIC_API_KEY", "standard-anthropic-key") + config = BedrockClaudePlatformConfig() + request_body = { + "model": "claude-sonnet-4-6", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 10, + } + + with patch.object( + config, + "_sign_request", + return_value=({"Authorization": "signed"}, json.dumps(request_body).encode()), + ) as mock_sign_request: + headers, signed_body = config.sign_request( + headers={"anthropic-workspace-id": "wrkspc_test"}, + optional_params={"aws_region_name": "us-west-2"}, + request_data=request_body, + api_base="https://aws-external-anthropic.us-west-2.api.aws/v1/messages", + api_key=None, + model="claude-sonnet-4-6", + ) + + assert signed_body == json.dumps(request_body).encode() + assert headers["Authorization"] == "signed" + mock_sign_request.assert_called_once() + + +def test_bedrock_claude_platform_messages_config_round_trips_native_body(): + import litellm + from litellm.types.utils import LlmProviders + + config = litellm.ProviderConfigManager.get_provider_anthropic_messages_config( + model="claude_platform/claude-sonnet-4-6", + provider=LlmProviders.BEDROCK, + ) + + assert config is not None + headers, _ = config.validate_anthropic_messages_environment( + api_key="fake-platform-key", + headers={}, + model="claude_platform/claude-sonnet-4-6", + messages=[{"role": "user", "content": "hello"}], + optional_params={"max_tokens": 10}, + litellm_params={"workspace_id": "wrkspc_test"}, + ) + request_body = config.transform_anthropic_messages_request( + model="claude_platform/claude-sonnet-4-6", + messages=[{"role": "user", "content": "hello"}], + anthropic_messages_optional_request_params={"max_tokens": 10}, + litellm_params={}, + headers=headers, + ) + + assert headers["anthropic-workspace-id"] == "wrkspc_test" + assert headers["x-api-key"] == "fake-platform-key" + assert request_body == { + "model": "claude-sonnet-4-6", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 10, + } + + +def test_chat_completion_routes_bedrock_claude_platform_to_messages_api(): + import litellm + + requests = [] + + def mock_post(self, url, data=None, headers=None, **kwargs): + requests.append(_capture_request(url=url, headers=headers or {}, data=data)) + return _anthropic_response(url) + + with patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post", mock_post): + response = litellm.completion( + model="bedrock/claude_platform/claude-sonnet-4-6", + messages=[{"role": "user", "content": "hello"}], + max_tokens=10, + api_base="https://aws-external-anthropic.us-west-2.api.aws", + api_key="fake-platform-key", + workspace_id="wrkspc_test", + ) + + assert response.choices[0].message.content == "ok" + assert len(requests) == 1 + assert requests[0]["path"] == "/v1/messages" + assert requests[0]["headers"]["x-api-key"] == "fake-platform-key" + assert requests[0]["headers"]["anthropic-workspace-id"] == "wrkspc_test" + assert requests[0]["body"]["model"] == "claude-sonnet-4-6" + + +@pytest.mark.asyncio +async def test_anthropic_messages_routes_bedrock_claude_platform_to_messages_api(): + import litellm + + requests = [] + + async def mock_post(self, url, data=None, headers=None, **kwargs): + requests.append(_capture_request(url=url, headers=headers or {}, data=data)) + return _anthropic_response(url) + + try: + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new=mock_post, + ): + response = await litellm.anthropic_messages( + model="bedrock/claude_platform/claude-sonnet-4-6", + messages=[{"role": "user", "content": "hello"}], + max_tokens=10, + api_base="https://aws-external-anthropic.us-west-2.api.aws", + api_key="fake-platform-key", + workspace_id="wrkspc_test", + ) + finally: + await litellm.close_litellm_async_clients() + + assert response["content"][0]["text"] == "ok" + assert len(requests) == 1 + assert requests[0]["path"] == "/v1/messages" + assert requests[0]["headers"]["x-api-key"] == "fake-platform-key" + assert requests[0]["headers"]["anthropic-workspace-id"] == "wrkspc_test" + assert requests[0]["body"]["messages"] == [{"role": "user", "content": "hello"}] + assert requests[0]["body"]["max_tokens"] == 10 + assert requests[0]["body"]["model"] == "claude-sonnet-4-6" diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index b846cd600f0..20050a43355 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -422,3 +422,140 @@ def test_sync_delete_responses_omits_body_for_azure(): assert captured["url"].endswith( "/openai/responses/resp_xyz?api-version=2025-03-01-preview" ) + + +# Per-request `timeout` is plumbed through the /v1/messages chain +# (async handler -> retry helper -> httpx post), matching /chat/completions. + + +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.asyncio +async def test_async_anthropic_messages_handler_forwards_timeout_to_retry_helper( + stream, +): + """`async_anthropic_messages_handler` must forward `timeout` to + `_async_post_anthropic_messages_with_http_error_retry` for both + streaming and non-streaming calls.""" + handler = BaseLLMHTTPHandler() + + mock_config = Mock() + mock_config.validate_anthropic_messages_environment = Mock( + return_value=({"x-api-key": "test-key"}, "https://api.anthropic.com") + ) + mock_config.transform_anthropic_messages_request = Mock( + return_value={"model": "claude-sonnet-4-20250514", "messages": []} + ) + mock_config.get_complete_url = Mock( + return_value="https://api.anthropic.com/v1/messages" + ) + mock_config.sign_request = Mock(return_value=({}, b"{}")) + + mock_logging_obj = Mock() + mock_logging_obj.update_environment_variables = Mock() + mock_logging_obj.update_from_kwargs = Mock() + mock_logging_obj.pre_call = Mock() + mock_logging_obj.model_call_details = {} + mock_logging_obj.stream = False + + with patch.object( + handler, + "_async_post_anthropic_messages_with_http_error_retry", + new_callable=AsyncMock, + ) as mock_retry, patch( + "litellm.litellm_core_utils.get_provider_specific_headers.ProviderSpecificHeaderUtils.get_provider_specific_headers", + return_value=None, + ): + try: + await handler.async_anthropic_messages_handler( + model="claude-sonnet-4-20250514", + messages=[{"role": "user", "content": "hi"}], + anthropic_messages_provider_config=mock_config, + anthropic_messages_optional_request_params={}, + custom_llm_provider="anthropic", + litellm_params=GenericLiteLLMParams(), + logging_obj=mock_logging_obj, + stream=stream, + timeout=0.5, + kwargs={}, + ) + except Exception: + # Downstream response transformation may fail under mocks; we + # only care about what reached the retry helper. + pass + + mock_retry.assert_awaited_once() + assert mock_retry.await_args.kwargs["timeout"] == 0.5 + assert mock_retry.await_args.kwargs["stream"] is stream + + +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.asyncio +async def test_async_post_anthropic_messages_forwards_timeout_to_httpx_post(stream): + """`_async_post_anthropic_messages_with_http_error_retry` must forward + `timeout` to `async_httpx_client.post` for both streaming and + non-streaming requests (router resolves `stream_timeout` vs `timeout` + upstream into a single value).""" + handler = BaseLLMHTTPHandler() + + mock_response = Mock(status_code=200) + mock_response.raise_for_status = Mock(return_value=None) + + mock_client = AsyncMock() + mock_client.post = AsyncMock(return_value=mock_response) + + mock_provider_config = Mock() + mock_provider_config.max_retry_on_anthropic_messages_http_error = 1 + + await handler._async_post_anthropic_messages_with_http_error_retry( + async_httpx_client=mock_client, + request_url="https://api.anthropic.com/v1/messages", + headers={}, + signed_json_body=None, + request_body={"model": "claude-sonnet-4-20250514", "messages": []}, + stream=stream, + logging_obj=Mock(), + provider_config=mock_provider_config, + litellm_params=GenericLiteLLMParams(), + api_key=None, + model="claude-sonnet-4-20250514", + timeout=0.5, + ) + + mock_client.post.assert_awaited_once() + assert mock_client.post.await_args.kwargs["timeout"] == 0.5 + assert mock_client.post.await_args.kwargs["stream"] is stream + + +@pytest.mark.asyncio +async def test_async_post_anthropic_messages_defaults_timeout_to_none_when_unset(): + """When no `timeout` is supplied, the retry helper must call + `async_httpx_client.post(timeout=None)` so AsyncHTTPHandler can fall + back to its own default — preserves pre-patch behavior for callers + that do not opt into a per-request timeout.""" + handler = BaseLLMHTTPHandler() + + mock_response = Mock(status_code=200) + mock_response.raise_for_status = Mock(return_value=None) + + mock_client = AsyncMock() + mock_client.post = AsyncMock(return_value=mock_response) + + mock_provider_config = Mock() + mock_provider_config.max_retry_on_anthropic_messages_http_error = 1 + + await handler._async_post_anthropic_messages_with_http_error_retry( + async_httpx_client=mock_client, + request_url="https://api.anthropic.com/v1/messages", + headers={}, + signed_json_body=None, + request_body={"model": "claude-sonnet-4-20250514", "messages": []}, + stream=False, + logging_obj=Mock(), + provider_config=mock_provider_config, + litellm_params=GenericLiteLLMParams(), + api_key=None, + model="claude-sonnet-4-20250514", + ) + + mock_client.post.assert_awaited_once() + assert mock_client.post.await_args.kwargs["timeout"] is None diff --git a/tests/test_litellm/llms/github_copilot/test_github_copilot_transformation.py b/tests/test_litellm/llms/github_copilot/test_github_copilot_transformation.py index 678aa6b56c1..45ce5d58405 100644 --- a/tests/test_litellm/llms/github_copilot/test_github_copilot_transformation.py +++ b/tests/test_litellm/llms/github_copilot/test_github_copilot_transformation.py @@ -94,27 +94,33 @@ def test_github_copilot_config_get_openai_compatible_provider_info(): @patch("litellm.llms.github_copilot.authenticator.Authenticator.get_api_key") +@patch("litellm.main.openai_chat_completions.completion") @patch("litellm.llms.openai.openai.OpenAIChatCompletion.completion") -def test_completion_github_copilot_mock_response(mock_completion, mock_get_api_key): +def test_completion_github_copilot_mock_response( + mock_class_completion, mock_instance_completion, mock_get_api_key, monkeypatch +): """Test the completion function with GitHub Copilot provider.""" - # Mock the API key return value + # Force chat path through the patched openai_chat_completions instance even if + # a previous test left EXPERIMENTAL_OPENAI_BASE_LLM_HTTP_HANDLER set in the env. + monkeypatch.delenv("EXPERIMENTAL_OPENAI_BASE_LLM_HTTP_HANDLER", raising=False) + mock_api_key = "gh.test-key-123456789" mock_get_api_key.return_value = mock_api_key - # Mock completion response mock_response = MagicMock() mock_response.choices = [MagicMock()] mock_response.choices[0].message.content = "Hello, I'm GitHub Copilot!" - mock_completion.return_value = mock_response + # Patch both the class method and the live module-level instance to survive + # conftest module reloads that can swap which class object is in use. + mock_class_completion.return_value = mock_response + mock_instance_completion.return_value = mock_response - # Test non-streaming completion messages = [ {"role": "system", "content": "You're GitHub Copilot, an AI assistant."}, {"role": "user", "content": "Hello, who are you?"}, ] - # Create a properly formatted headers dictionary headers = { "editor-version": "Neovim/0.9.0", "Copilot-Integration-Id": "vscode-chat", @@ -128,19 +134,16 @@ def test_completion_github_copilot_mock_response(mock_completion, mock_get_api_k assert response is not None - # Verify the get_api_key call was made (can be called multiple times) assert mock_get_api_key.call_count >= 1 - # Verify the completion call was made with the expected params - mock_completion.assert_called_once() - args, kwargs = mock_completion.call_args + # Exactly one of the two patched targets should have been used. + invoked = [m for m in (mock_class_completion, mock_instance_completion) if m.called] + assert len(invoked) == 1 + invoked[0].assert_called_once() + _, kwargs = invoked[0].call_args - # Check that the proper authorization header is set assert "headers" in kwargs - # Check that the model name is correctly formatted - assert ( - kwargs.get("model") == "gpt-4" - ) # Model name should be without provider prefix + assert kwargs.get("model") == "gpt-4" assert kwargs.get("messages") == messages diff --git a/tests/test_litellm/llms/openai/test_gpt5_transformation.py b/tests/test_litellm/llms/openai/test_gpt5_transformation.py index ebf7681f2f3..d279b119efe 100644 --- a/tests/test_litellm/llms/openai/test_gpt5_transformation.py +++ b/tests/test_litellm/llms/openai/test_gpt5_transformation.py @@ -1,10 +1,15 @@ import pytest import litellm +import litellm.main as litellm_main from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map from litellm.llms.openai.chat.gpt_5_transformation import OpenAIGPT5Config from litellm.llms.openai.openai import OpenAIConfig -from litellm.utils import _is_explicitly_disabled_factory +from litellm.utils import ( + _is_explicitly_disabled_factory, + peek_reasoning_summary_aliases, + strip_reasoning_summary_aliases_from_optional_params, +) @pytest.fixture() @@ -1007,6 +1012,76 @@ def test_gpt5_search_drops_unsupported_params(config: OpenAIConfig): assert "tools" not in params +def test_gpt5_chat_strips_reasoning_summary_aliases_after_bridge_check( + monkeypatch: pytest.MonkeyPatch, +): + """Non-bridged GPT-5 chat calls strip Responses-only reasoning summary aliases.""" + captured_kwargs = {} + + def fake_openai_completion(**kwargs): + captured_kwargs.update(kwargs) + return {} + + monkeypatch.setattr( + litellm_main.openai_chat_completions, + "completion", + fake_openai_completion, + ) + + litellm.completion( + model="gpt-5", + messages=[{"role": "user", "content": "ok"}], + reasoningSummary="auto", + extra_body={"reasoning_summary": "ignored", "metadata": "ok"}, + api_key="fake-key", + ) + + optional_params = captured_kwargs["optional_params"] + assert "reasoningSummary" not in optional_params + assert "reasoning_summary" not in optional_params + assert optional_params["extra_body"] == {"metadata": "ok"} + + +def test_reasoning_summary_alias_helpers_preserve_falsy_and_strip_all_aliases(): + optional_params = {"reasoningSummary": False, "reasoning_summary": "ignored"} + + assert peek_reasoning_summary_aliases(optional_params) is False + stripped, rs_val = strip_reasoning_summary_aliases_from_optional_params( + optional_params + ) + + assert rs_val is False + assert stripped == {} + + optional_params = { + "extra_body": {"reasoningSummary": False, "reasoning_summary": "ignored"} + } + + assert peek_reasoning_summary_aliases(optional_params) is False + stripped, rs_val = strip_reasoning_summary_aliases_from_optional_params( + optional_params + ) + + assert rs_val is False + assert stripped == {} + + optional_params = { + "extra_body": { + "reasoningSummary": "auto", + "reasoning_summary": "ignored", + "metadata": "ok", + } + } + + assert peek_reasoning_summary_aliases(optional_params) == "auto" + stripped, rs_val = strip_reasoning_summary_aliases_from_optional_params( + optional_params + ) + + assert rs_val == "auto" + assert stripped == {"extra_body": {"metadata": "ok"}} + + # GPT-5 unsupported params audit (validated via direct API calls) def test_gpt5_rejects_params_unsupported_by_openai(config: OpenAIConfig): """Params that OpenAI rejects for all GPT-5 reasoning models.""" diff --git a/tests/test_litellm/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py b/tests/test_litellm/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py index 8cc46dc98d0..c8751fb2d95 100644 --- a/tests/test_litellm/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py +++ b/tests/test_litellm/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py @@ -54,3 +54,61 @@ def test_ovhcloud_audio_transcription_config_installed(): assert config is not None assert isinstance(config, BaseAudioTranscriptionConfig) + + + +class TestOVHCloudDurationFieldMigration: + """Tests for OVHCloud duration -> seconds field migration.""" + + def test_seconds_field_mapped_to_duration(self): + """New `seconds` field should be normalized to `duration`.""" + from litellm.llms.ovhcloud.audio_transcription.transformation import ( + OVHCloudAudioTranscriptionConfig, + ) + from unittest.mock import MagicMock + + config = OVHCloudAudioTranscriptionConfig() + mock_response = MagicMock() + mock_response.json.return_value = { + "text": "Hello world", + "seconds": 3.14, + } + + result = config.transform_audio_transcription_response(mock_response) + + assert result.text == "Hello world" + assert result._hidden_params["duration"] == 3.14 + + def test_legacy_duration_field_still_works(self): + """Legacy `duration` field should still be accepted.""" + from litellm.llms.ovhcloud.audio_transcription.transformation import ( + OVHCloudAudioTranscriptionConfig, + ) + from unittest.mock import MagicMock + + config = OVHCloudAudioTranscriptionConfig() + mock_response = MagicMock() + mock_response.json.return_value = { + "text": "Hello world", + "duration": 2.71, + } + + result = config.transform_audio_transcription_response(mock_response) + + assert result.text == "Hello world" + assert result._hidden_params["duration"] == 2.71 + + + + def test_seconds_zero_mapped_to_duration(self): + """seconds=0.0 must not be treated as falsy and lost.""" + from litellm.llms.ovhcloud.audio_transcription.transformation import ( + OVHCloudAudioTranscriptionConfig, + ) + from unittest.mock import MagicMock + + config = OVHCloudAudioTranscriptionConfig() + mock_response = MagicMock() + mock_response.json.return_value = {"text": "silence", "seconds": 0.0} + result = config.transform_audio_transcription_response(mock_response) + assert result._hidden_params["duration"] == 0.0 \ No newline at end of file diff --git a/tests/test_litellm/llms/ovhcloud/test_ovhcloud_chat_transformation.py b/tests/test_litellm/llms/ovhcloud/test_ovhcloud_chat_transformation.py index a1b3b31f786..40d57c76d02 100644 --- a/tests/test_litellm/llms/ovhcloud/test_ovhcloud_chat_transformation.py +++ b/tests/test_litellm/llms/ovhcloud/test_ovhcloud_chat_transformation.py @@ -292,3 +292,78 @@ def test_ovhcloud_with_custom_base_url(): if __name__ == "__main__": pytest.main([__file__, "-v"]) + + +class TestOVHCloudReasoningFieldMigration: + """Tests for OVHCloud reasoning_content -> reasoning field migration.""" + + def test_streaming_new_reasoning_field(self): + """New `reasoning` field should be mapped to `reasoning_content`.""" + handler = OVHCloudChatCompletionStreamingHandler( + streaming_response=iter([]), + sync_stream=True, + ) + chunk = { + "id": "test-id", + "created": 1234567890, + "model": "test-model", + "choices": [ + { + "delta": { + "role": "assistant", + "reasoning": "Let me think...", + }, + "index": 0, + } + ], + } + result = handler.chunk_parser(chunk) + assert result.choices[0]["delta"]["reasoning_content"] == "Let me think..." + + def test_streaming_legacy_reasoning_content_unchanged(self): + """Legacy `reasoning_content` field should pass through untouched.""" + handler = OVHCloudChatCompletionStreamingHandler( + streaming_response=iter([]), + sync_stream=True, + ) + chunk = { + "id": "test-id", + "created": 1234567890, + "model": "test-model", + "choices": [ + { + "delta": { + "role": "assistant", + "reasoning_content": "Already correct field.", + }, + "index": 0, + } + ], + } + result = handler.chunk_parser(chunk) + assert result.choices[0]["delta"]["reasoning_content"] == "Already correct field." + + def test_streaming_both_fields_legacy_wins(self): + """When both fields present, existing `reasoning_content` is not overwritten.""" + handler = OVHCloudChatCompletionStreamingHandler( + streaming_response=iter([]), + sync_stream=True, + ) + chunk = { + "id": "test-id", + "created": 1234567890, + "model": "test-model", + "choices": [ + { + "delta": { + "reasoning": "new field", + "reasoning_content": "legacy field", + }, + "index": 0, + } + ], + } + result = handler.chunk_parser(chunk) + assert result.choices[0]["delta"]["reasoning_content"] == "legacy field" + + diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py index 89992c510f5..9f004318488 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py @@ -1135,3 +1135,69 @@ def test_validate_loopback_redirect_uri_rejects_malformed_cleanly(): with pytest.raises(HTTPException) as exc: validate_loopback_redirect_uri("http://[not-an-ip]/cb") assert exc.value.status_code == 400 + + +def _mock_request_with_base_url(base_url: str): + req = MagicMock() + req.base_url = base_url + req.headers = {} + return req + + +def test_validate_trusted_redirect_uri_accepts_same_origin(): + """UI OAuth flow: redirect_uri on the proxy's own origin is allowed.""" + from litellm.proxy._experimental.mcp_server.oauth_utils import ( + validate_trusted_redirect_uri, + ) + + req = _mock_request_with_base_url("https://proxy.example.com/") + # Should not raise. + validate_trusted_redirect_uri( + req, "https://proxy.example.com/ui/mcp/oauth/callback" + ) + + +def test_validate_trusted_redirect_uri_accepts_loopback(): + """Native MCP client flow: loopback is still allowed.""" + from litellm.proxy._experimental.mcp_server.oauth_utils import ( + validate_trusted_redirect_uri, + ) + + req = _mock_request_with_base_url("https://proxy.example.com/") + validate_trusted_redirect_uri(req, "http://127.0.0.1:3000/cb") + validate_trusted_redirect_uri(req, "http://localhost:3000/cb") + + +def test_validate_trusted_redirect_uri_rejects_external_origin(): + """An attacker-controlled origin must still be rejected.""" + from litellm.proxy._experimental.mcp_server.oauth_utils import ( + validate_trusted_redirect_uri, + ) + + req = _mock_request_with_base_url("https://proxy.example.com/") + with pytest.raises(HTTPException) as exc: + validate_trusted_redirect_uri(req, "https://attacker.example.com/cb") + assert exc.value.status_code == 400 + + +def test_validate_trusted_redirect_uri_rejects_scheme_mismatch(): + """https→http (or vice versa) on the same host is not same-origin.""" + from litellm.proxy._experimental.mcp_server.oauth_utils import ( + validate_trusted_redirect_uri, + ) + + req = _mock_request_with_base_url("https://proxy.example.com/") + with pytest.raises(HTTPException) as exc: + validate_trusted_redirect_uri(req, "http://proxy.example.com/ui/callback") + assert exc.value.status_code == 400 + + +def test_validate_trusted_redirect_uri_rejects_fragment(): + from litellm.proxy._experimental.mcp_server.oauth_utils import ( + validate_trusted_redirect_uri, + ) + + req = _mock_request_with_base_url("https://proxy.example.com/") + with pytest.raises(HTTPException) as exc: + validate_trusted_redirect_uri(req, "https://proxy.example.com/ui/cb#code=1") + assert exc.value.status_code == 400 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 85d5d6ba466..581324d47d2 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -23,6 +23,20 @@ def mock_mcp_client_ip(): yield +def _mock_callback_request(base_url: str = "http://localhost:3000/"): + """Return a MagicMock Request for callback/authorize same-origin tests. + + The callback handler only uses ``request`` to compute the proxy's own + base URL via ``get_request_base_url`` (which reads ``request.base_url`` + and trusted ``X-Forwarded-*`` headers). A simple MagicMock with the + right attributes is sufficient. + """ + req = MagicMock() + req.base_url = base_url + req.headers = {} + return req + + @pytest.fixture def trust_xff(): """Force ``IPAddressUtils.is_request_from_trusted_proxy`` to True. @@ -1844,6 +1858,7 @@ async def test_oauth_callback_redirects_with_state(): # Call callback endpoint with code and state response = await callback( + request=_mock_callback_request(), code="test_authorization_code_12345", state="encrypted_state_value", ) @@ -1887,6 +1902,7 @@ async def test_oauth_callback_preserves_client_redirect_uri_query(): } response = await callback( + request=_mock_callback_request(), code="test_authorization_code_12345", state="encrypted_state_value", ) @@ -1917,6 +1933,7 @@ async def test_oauth_callback_handles_invalid_state(): # Call callback endpoint with invalid state response = await callback( + request=_mock_callback_request(), code="test_code", state="invalid_encrypted_state", ) @@ -1926,6 +1943,40 @@ async def test_oauth_callback_handles_invalid_state(): assert "Authentication incomplete" in response.body.decode() +@pytest.mark.asyncio +async def test_oauth_callback_accepts_same_origin_ui_redirect(): + """UI OAuth flow: the callback should redirect to the proxy's own UI + origin when the encrypted state carries a same-origin client_redirect_uri.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + callback, + ) + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash" + ) as mock_decode: + mock_decode.return_value = { + "base_url": "https://proxy.example.com/ui/mcp/oauth/callback", + "original_state": "state-123", + "code_challenge": None, + "code_challenge_method": None, + "client_redirect_uri": "https://proxy.example.com/ui/mcp/oauth/callback", + } + + response = await callback( + request=_mock_callback_request(base_url="https://proxy.example.com/"), + code="auth-code-123", + state="encrypted_state", + ) + + assert response.status_code == 302 + assert ( + "https://proxy.example.com/ui/mcp/oauth/callback" + in response.headers["location"] + ) + assert "code=auth-code-123" in response.headers["location"] + assert "state=state-123" in response.headers["location"] + + @pytest.mark.asyncio async def test_oauth_authorize_includes_scopes_from_server_config(): """Test that authorize endpoint includes scopes from server configuration.""" @@ -2307,7 +2358,11 @@ async def test_callback_revalidates_loopback_on_decoded_base_url(): "client_redirect_uri": "https://attacker.example.com/cb", } with pytest.raises(HTTPException) as exc_info: - await callback(code="stolen_code", state="encrypted_stale_state") + await callback( + request=_mock_callback_request(), + code="stolen_code", + state="encrypted_stale_state", + ) assert exc_info.value.status_code == 400 @@ -2329,7 +2384,11 @@ async def test_callback_revalidates_loopback_on_decoded_client_redirect_uri(): "client_redirect_uri": "https://attacker.example.com/cb", } with pytest.raises(HTTPException) as exc_info: - await callback(code="stolen_code", state="encrypted_stale_state") + await callback( + request=_mock_callback_request(), + code="stolen_code", + state="encrypted_stale_state", + ) assert exc_info.value.status_code == 400 @@ -2349,7 +2408,11 @@ async def test_callback_rejects_state_missing_redirect_uri(): "code_challenge_method": None, } with pytest.raises(HTTPException) as exc_info: - await callback(code="code", state="encrypted_malformed_state") + await callback( + request=_mock_callback_request(), + code="code", + state="encrypted_malformed_state", + ) assert exc_info.value.status_code == 400 diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index b7dba9c1d16..90c5d4f4fc5 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -1,5 +1,5 @@ from typing import Optional -from unittest.mock import AsyncMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -494,6 +494,80 @@ async def test_sync_user_role_and_teams_no_cache_write_when_nothing_changes(): mock_cache.async_set_cache.assert_not_called() +def test_get_all_jwt_team_ids_unions_singular_and_plural(): + """get_all_jwt_team_ids must include the singular team_id_jwt_field claim + in addition to the plural team_ids_jwt_field, deduplicated.""" + jwt_handler = JWTHandler() + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=MagicMock(), + litellm_jwtauth=LiteLLM_JWTAuth( + team_id_jwt_field="team_id", + team_ids_jwt_field="teams", + ), + ) + + # singular only — Okta/Auth0 default shape + assert jwt_handler.get_all_jwt_team_ids({"team_id": "team-low"}) == ["team-low"] + + # plural only — pre-fix shape + assert jwt_handler.get_all_jwt_team_ids({"teams": ["a", "b"]}) == ["a", "b"] + + # both populated, no overlap + assert jwt_handler.get_all_jwt_team_ids( + {"team_id": "primary", "teams": ["a", "b"]} + ) == ["a", "b", "primary"] + + # both populated with overlap — singular dedup'd + assert jwt_handler.get_all_jwt_team_ids({"team_id": "a", "teams": ["a", "b"]}) == [ + "a", + "b", + ] + + # singular field as multi-element list (some IdPs) — merge all, preserve plural-first order + assert jwt_handler.get_all_jwt_team_ids( + {"team_id": ["primary", "secondary"], "teams": ["a"]} + ) == ["a", "primary", "secondary"] + + # neither populated + assert jwt_handler.get_all_jwt_team_ids({}) == [] + + +def test_get_all_jwt_team_ids_does_not_use_team_id_default(): + """team_id_default is a JWT-bearer-flow auth-builder fallback, not a token + claim. It must NOT leak into get_all_jwt_team_ids — otherwise SSO logins + would silently start adding users to the default team for any tenant that + has team_id_default configured.""" + jwt_handler = JWTHandler() + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=MagicMock(), + litellm_jwtauth=LiteLLM_JWTAuth( + team_id_jwt_field="team_id", + team_ids_jwt_field="teams", + team_id_default="default-team", + ), + ) + + # team_id claim missing — must not fall back to default-team + assert jwt_handler.get_all_jwt_team_ids({"teams": []}) == [] + assert jwt_handler.get_all_jwt_team_ids({}) == [] + + # only the plural is populated — default still must not be added + assert jwt_handler.get_all_jwt_team_ids({"teams": ["a"]}) == ["a"] + + # team_id_jwt_field unset entirely + only default configured: still no default + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=MagicMock(), + litellm_jwtauth=LiteLLM_JWTAuth( + team_ids_jwt_field="teams", + team_id_default="default-team", + ), + ) + assert jwt_handler.get_all_jwt_team_ids({"teams": []}) == [] + + @pytest.mark.asyncio async def test_map_jwt_role_to_litellm_role(): """Test JWT role mapping to LiteLLM roles with various patterns""" diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index 9c1cd116e69..fa56383e9bb 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -53,14 +53,20 @@ def test_non_admin_config_update_route_rejected(): assert "Your role=internal_user" in str(exc_info.value) +@pytest.mark.parametrize( + "role", + [ + LitellmUserRoles.INTERNAL_USER.value, + LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value, + ], +) @pytest.mark.parametrize( "route", ["/compliance/eu-ai-act", "/compliance/gdpr"], ) -def test_compliance_routes_open_to_internal_user(route): +def test_compliance_routes_open_to_non_admin_roles(role, route): """Compliance routes are stateless validators on caller-supplied log data - - non-admin internal_user roles can call them.""" - role = LitellmUserRoles.INTERNAL_USER.value + — both non-admin internal_user roles can call them.""" user_obj = LiteLLM_UserTable( user_id="test_user", user_email="test@example.com", @@ -80,56 +86,6 @@ def test_compliance_routes_open_to_internal_user(route): ) -def test_health_test_connection_route_delegates_internal_user_auth_to_endpoint(): - """Team model test-connection requests are authorized by the endpoint.""" - role = LitellmUserRoles.INTERNAL_USER.value - user_obj = LiteLLM_UserTable( - user_id="test_user", - user_email="test@example.com", - user_role=role, - ) - valid_token = UserAPIKeyAuth(user_id="test_user", user_role=role) - request = MagicMock(spec=Request) - request.query_params = {} - - RouteChecks.non_proxy_admin_allowed_routes_check( - user_obj=user_obj, - _user_role=role, - route="/health/test_connection", - request=request, - valid_token=valid_token, - request_data={}, - ) - - -@pytest.mark.parametrize( - "route", - ["/compliance/eu-ai-act", "/compliance/gdpr"], -) -def test_compliance_routes_blocked_for_internal_user_view_only(route): - """Deprecated internal_user_viewer role must not gain compliance route access.""" - role = LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value - user_obj = LiteLLM_UserTable( - user_id="test_user", - user_email="test@example.com", - user_role=role, - ) - valid_token = UserAPIKeyAuth(user_id="test_user", user_role=role) - request = MagicMock(spec=Request) - request.query_params = {} - - with pytest.raises(Exception) as exc_info: - RouteChecks.non_proxy_admin_allowed_routes_check( - user_obj=user_obj, - _user_role=role, - route=route, - request=request, - valid_token=valid_token, - request_data={}, - ) - assert "Only proxy admin can be used" in str(exc_info.value) - - def test_proxy_admin_viewer_config_update_route_rejected(): """Test that proxy admin viewer users are rejected when trying to call /config/update""" @@ -1933,6 +1889,41 @@ def test_non_admin_non_team_admin_cannot_access_config_update_but_can_attempt_re assert "Only proxy admin can be used to generate" in str(exc_info.value) +@pytest.mark.parametrize( + "user_role", + [ + LitellmUserRoles.INTERNAL_USER.value, + LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value, + ], +) +@pytest.mark.parametrize("route", ["/tag/list", "/tag/daily/activity"]) +def test_internal_users_can_access_scoped_tag_usage_routes(user_role, route): + """ + Internal users can read tag usage endpoints because the endpoint handlers + scope results to the caller's own keys. + """ + user_obj = LiteLLM_UserTable( + user_id="test_user", + user_email="test@example.com", + user_role=user_role, + ) + valid_token = UserAPIKeyAuth( + user_id="test_user", + user_role=user_role, + ) + request = MagicMock(spec=Request) + request.query_params = {} + + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=user_obj, + _user_role=user_role, + route=route, + request=request, + valid_token=valid_token, + request_data={}, + ) + + @pytest.mark.parametrize( "user_role", [ diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 50c5f43b218..442625c75a7 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -31,8 +31,10 @@ from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy.auth.auth_checks import get_key_object, _cache_key_object from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.auth.user_api_key_auth import ( - _route_requires_auth_despite_public, + _matches_routing_override, _reserve_budget_after_common_checks, + _route_requires_auth_despite_public, + _routing_selector_matches_claim, _run_centralized_common_checks, _run_post_custom_auth_checks, get_api_key, @@ -594,6 +596,151 @@ def _assert_get_api_key_with_custom_litellm_key_header( ) == (api_key, passed_in_key) +@pytest.mark.parametrize( + "selector_value, claim_value, expected, split_space_delimited", + [ + (None, "any-value", True, False), + ("issuer.example.com", "issuer.example.com", True, False), + ("issuer.example.com", "other-issuer.example.com", False, False), + # iss (and other non-scope claims) must not match via space-split injection + ( + "trusted.example.com", + "trusted.example.com attacker.example.com", + False, + False, + ), + # Wildcard iss must not match space-containing claim strings (fnmatch * spans spaces) + ( + "trusted.*", + "trusted.example.com attacker.example.com", + False, + False, + ), + ("trusted.*", "trusted.example.com", True, False), + ( + ["issuer-a.example.com", "issuer-b.example.com"], + "issuer-b.example.com", + True, + False, + ), + ("*MID_LITELLM", "STREAM_MID_LITELLM", True, False), + ("*MID_LITELLM", "REDIS_LITELLM", False, False), + ("machine-??", "machine-01", True, False), + ("machine-??", "machine-001", False, False), + # Wildcard matching is case-sensitive (fnmatch.fnmatchcase) + ("*litellm", "BATCH_LITELLM", False, False), + ("*LITELLM", "BATCH_LITELLM", True, False), + ("App:LiteLLM", "App:LiteLLM openid", True, True), + ("App:*", "App:LiteLLM openid", True, True), + (["openid", "App:LiteLLM"], "openid profile", True, True), + (["service-*", "batch-*"], "batch-123", True, False), + (["service-*", "batch-*"], "other-123", False, False), + ("App:LiteLLM", ["openid", "App:LiteLLM"], True, False), + ("App:LiteLLM", None, False, False), + ], +) +def test_routing_selector_matches_claim_parametrized( + selector_value, claim_value, expected, split_space_delimited +): + assert ( + _routing_selector_matches_claim( + selector_value=selector_value, + claim_value=claim_value, + split_space_delimited=split_space_delimited, + ) + is expected + ) + + +@pytest.mark.parametrize( + "override, token_claims, expected", + [ + # Only iss selector is required and should match. + ( + JWTRoutingOverride(iss="oauth-issuer.example.com", path="oauth2"), + {"iss": "oauth-issuer.example.com"}, + True, + ), + # Scope selector narrows the match. + ( + JWTRoutingOverride( + iss="oauth-issuer.example.com", + scope="App:LiteLLM", + path="oauth2", + ), + {"iss": "oauth-issuer.example.com", "scope": "App:LiteLLM openid"}, + True, + ), + # client_id wildcard selector narrows the match. + ( + JWTRoutingOverride( + iss="oauth-issuer.example.com", + client_id="*MID_LITELLM", + path="oauth2", + ), + {"iss": "oauth-issuer.example.com", "client_id": "BATCH_MID_LITELLM"}, + True, + ), + ( + JWTRoutingOverride( + iss="oauth-issuer.example.com", + client_id="*MID_LITELLM", + path="oauth2", + ), + {"iss": "oauth-issuer.example.com", "client_id": "BATCH_PORTAL"}, + False, + ), + # aud selector still works with list claims. + ( + JWTRoutingOverride( + iss="oauth-issuer.example.com", + aud=["api://litellm", "api://fallback"], + path="oauth2", + ), + { + "iss": "oauth-issuer.example.com", + "aud": ["api://other", "api://litellm"], + }, + True, + ), + # All provided selectors are AND-ed. + ( + JWTRoutingOverride( + iss="oauth-issuer.example.com", + scope="App:LiteLLM", + client_id="*MID_LITELLM", + path="oauth2", + ), + { + "iss": "oauth-issuer.example.com", + "scope": "App:LiteLLM openid", + "client_id": "BATCH_MID_LITELLM", + }, + True, + ), + ( + JWTRoutingOverride( + iss="oauth-issuer.example.com", + scope="App:LiteLLM", + client_id="*MID_LITELLM", + path="oauth2", + ), + { + "iss": "oauth-issuer.example.com", + "scope": "App:Other openid", + "client_id": "BATCH_MID_LITELLM", + }, + False, + ), + ], +) +def test_matches_routing_override_parametrized(override, token_claims, expected): + assert ( + _matches_routing_override(token_claims=token_claims, override=override) + is expected + ) + + def test_get_api_key_with_custom_litellm_key_header_bearer_prefix(): token = "sk-" + "1" * 8 header = f"Bearer {token}" @@ -1601,6 +1748,206 @@ class TestJWTOAuth2Coexistence: mock_jwt_auth.assert_not_called() assert result.user_id == "machine-client-aud-list" + @pytest.mark.asyncio + async def test_routing_override_matches_scope_claim(self): + """ + Match routing override when scope selector is configured and scope claim matches. + """ + jwt_token = ( + "eyJhbGciOiJSUzI1NiJ9." + "eyJpc3MiOiJvYXV0aC1pc3N1ZXIuZXhhbXBsZS5jb20iLCJzY29wZSI6IkFwcDpMaXRlTExNIiwiY2xpZW50X2lkIjoiTUFDSElORV9NSURfTElURUxMTSJ9." + "c2ln" + ) + general_settings = { + "enable_oauth2_auth": False, + "enable_jwt_auth": True, + } + mock_oauth2_response = UserAPIKeyAuth( + api_key=jwt_token, + user_id="machine-client-scope-match", + ) + + mock_request = MagicMock() + mock_request.url.path = "/v1/chat/completions" + mock_request.headers = {"authorization": f"Bearer {jwt_token}"} + mock_request.query_params = {} + + with ( + patch("litellm.proxy.proxy_server.general_settings", general_settings), + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.proxy_server.master_key", "sk-master"), + patch("litellm.proxy.proxy_server.prisma_client", None), + patch( + "litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token", + new_callable=AsyncMock, + return_value=mock_oauth2_response, + ) as mock_oauth2, + patch( + "litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder", + new_callable=AsyncMock, + ) as mock_jwt_auth, + ): + litellm.proxy.proxy_server.jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=DualCache(), + litellm_jwtauth=LiteLLM_JWTAuth( + routing_overrides=[ + JWTRoutingOverride( + iss="oauth-issuer.example.com", + scope="App:LiteLLM", + path="oauth2", + ) + ] + ), + ) + + result = await user_api_key_auth( + request=mock_request, + api_key=f"Bearer {jwt_token}", + ) + + mock_oauth2.assert_called_once_with(token=jwt_token) + mock_jwt_auth.assert_not_called() + assert result.user_id == "machine-client-scope-match" + + @pytest.mark.asyncio + async def test_routing_override_scope_mismatch_falls_back_to_jwt(self): + """ + If scope selector does not match, continue default JWT flow. + """ + jwt_token = ( + "eyJhbGciOiJSUzI1NiJ9." + "eyJpc3MiOiJvYXV0aC1pc3N1ZXIuZXhhbXBsZS5jb20iLCJzY29wZSI6IkFwcDpPdGhlciIsImNsaWVudF9pZCI6IlBPUlRBTF9NSURfTElURUxMTSJ9." + "c2ln" + ) + general_settings = { + "enable_oauth2_auth": False, + "enable_jwt_auth": True, + } + mock_jwt_result = { + "is_proxy_admin": True, + "team_object": None, + "user_object": None, + "end_user_object": None, + "org_object": None, + "token": jwt_token, + "team_id": "jwt-team", + "user_id": "jwt-user-scope-mismatch", + "end_user_id": None, + "org_id": None, + "team_membership": None, + "jwt_claims": {"sub": "user1"}, + } + + mock_request = MagicMock() + mock_request.url.path = "/v1/chat/completions" + mock_request.headers = {"authorization": f"Bearer {jwt_token}"} + mock_request.query_params = {} + + with ( + patch("litellm.proxy.proxy_server.general_settings", general_settings), + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.proxy_server.master_key", "sk-master"), + patch("litellm.proxy.proxy_server.prisma_client", None), + patch( + "litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token", + new_callable=AsyncMock, + ) as mock_oauth2, + patch( + "litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder", + new_callable=AsyncMock, + return_value=mock_jwt_result, + ) as mock_jwt_auth, + ): + litellm.proxy.proxy_server.jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=DualCache(), + litellm_jwtauth=LiteLLM_JWTAuth( + routing_overrides=[ + JWTRoutingOverride( + iss="oauth-issuer.example.com", + scope="App:LiteLLM", + path="oauth2", + ) + ] + ), + ) + + result = await user_api_key_auth( + request=mock_request, + api_key=f"Bearer {jwt_token}", + ) + + mock_oauth2.assert_not_called() + mock_jwt_auth.assert_called_once() + assert result.user_id == "jwt-user-scope-mismatch" + + @pytest.mark.asyncio + async def test_routing_override_matches_scope_and_client_wildcard_when_scope_claim_is_space_delimited( + self, + ): + """ + Integration check: combined scope + wildcard selectors match on OAuth2 path + when scope claim is a space-delimited string. + """ + jwt_token = ( + "eyJhbGciOiJSUzI1NiJ9." + "eyJpc3MiOiJvYXV0aC1pc3N1ZXIuZXhhbXBsZS5jb20iLCJzY29wZSI6IkFwcDpMaXRlTExNIG9wZW5pZCIsImNsaWVudF9pZCI6IkJBVENIX01JRF9MSVRFTExNIn0." + "c2ln" + ) + general_settings = { + "enable_oauth2_auth": False, + "enable_jwt_auth": True, + } + mock_oauth2_response = UserAPIKeyAuth( + api_key=jwt_token, + user_id="machine-client-space-delimited-scope-match", + ) + + mock_request = MagicMock() + mock_request.url.path = "/v1/chat/completions" + mock_request.headers = {"authorization": f"Bearer {jwt_token}"} + mock_request.query_params = {} + + with ( + patch("litellm.proxy.proxy_server.general_settings", general_settings), + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.proxy_server.master_key", "sk-master"), + patch("litellm.proxy.proxy_server.prisma_client", None), + patch( + "litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token", + new_callable=AsyncMock, + return_value=mock_oauth2_response, + ) as mock_oauth2, + patch( + "litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder", + new_callable=AsyncMock, + ) as mock_jwt_auth, + ): + litellm.proxy.proxy_server.jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=DualCache(), + litellm_jwtauth=LiteLLM_JWTAuth( + routing_overrides=[ + JWTRoutingOverride( + iss="oauth-issuer.example.com", + scope="App:LiteLLM", + client_id="*MID_LITELLM", + path="oauth2", + ) + ] + ), + ) + + result = await user_api_key_auth( + request=mock_request, + api_key=f"Bearer {jwt_token}", + ) + + mock_oauth2.assert_called_once_with(token=jwt_token) + mock_jwt_auth.assert_not_called() + assert result.user_id == "machine-client-space-delimited-scope-match" + @pytest.mark.asyncio async def test_routing_override_routes_jwt_to_oauth2_when_oauth2_globally_disabled( self, diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py index 5c86f9057a1..8b0c76f836c 100644 --- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py +++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py @@ -39,6 +39,46 @@ class MockLiteLLMVerificationToken: return {"count": 1} +class MockLiteLLMOrganizationTable: + def __init__(self): + self.update_many_calls: List[Dict[str, Any]] = [] + self.find_many_calls: List[Dict[str, Any]] = [] + self._find_many_results: List[Any] = [] + + def set_find_many_results(self, results: List[Any]): + self._find_many_results = results + + async def find_many(self, where: Dict[str, Any]) -> List[Any]: + self.find_many_calls.append({"where": where}) + return self._find_many_results + + async def update_many( + self, where: Dict[str, Any], data: Dict[str, Any] + ) -> Dict[str, Any]: + self.update_many_calls.append({"where": where, "data": data}) + return {"count": 1} + + +class MockLiteLLMTagTable: + def __init__(self): + self.update_many_calls: List[Dict[str, Any]] = [] + self.find_many_calls: List[Dict[str, Any]] = [] + self._find_many_results: List[Any] = [] + + def set_find_many_results(self, results: List[Any]): + self._find_many_results = results + + async def find_many(self, where: Dict[str, Any]) -> List[Any]: + self.find_many_calls.append({"where": where}) + return self._find_many_results + + async def update_many( + self, where: Dict[str, Any], data: Dict[str, Any] + ) -> Dict[str, Any]: + self.update_many_calls.append({"where": where, "data": data}) + return {"count": 1} + + class MockLiteLLMEndUserTable: def __init__(self): self.find_many_calls: List[Dict[str, Any]] = [] @@ -57,6 +97,8 @@ class MockDB: self.litellm_teammembership = MockLiteLLMTeamMembership() self.litellm_verificationtoken = MockLiteLLMVerificationToken() self.litellm_endusertable = MockLiteLLMEndUserTable() + self.litellm_organizationtable = MockLiteLLMOrganizationTable() + self.litellm_tagtable = MockLiteLLMTagTable() class MockPrismaClient: @@ -459,6 +501,100 @@ def test_reset_budget_for_keys_linked_to_budgets_empty( assert len(calls) == 0 +def test_reset_budget_for_orgs_linked_to_budgets(reset_budget_job, mock_prisma_client): + """ + Test that when a budget tier is reset, orgs linked to that budget + (via budget_id) also get their spend reset. + """ + now = datetime.now(timezone.utc) + + test_budget = type( + "LiteLLM_BudgetTableFull", + (), + { + "max_budget": 100.0, + "budget_duration": "30d", + "budget_reset_at": now - timedelta(hours=1), + "budget_id": "30d-org-budget", + "created_at": now - timedelta(days=30), + }, + ) + + asyncio.run( + reset_budget_job.reset_budget_for_orgs_linked_to_budgets( + budgets_to_reset=[test_budget] + ) + ) + + calls = mock_prisma_client.db.litellm_organizationtable.update_many_calls + assert len(calls) == 1 + call = calls[0] + assert call["where"]["budget_id"] == {"in": ["30d-org-budget"]} + assert call["where"]["spend"] == {"gt": 0} + assert call["data"]["spend"] == 0 + + +def test_reset_budget_for_orgs_linked_to_budgets_empty( + reset_budget_job, mock_prisma_client +): + """ + Test that when there are no budgets to reset, no update is performed + on the organization table. + """ + asyncio.run( + reset_budget_job.reset_budget_for_orgs_linked_to_budgets(budgets_to_reset=[]) + ) + calls = mock_prisma_client.db.litellm_organizationtable.update_many_calls + assert len(calls) == 0 + + +def test_reset_budget_for_tags_linked_to_budgets(reset_budget_job, mock_prisma_client): + """ + Test that when a budget tier is reset, tags linked to that budget + (via budget_id) also get their spend reset. + """ + now = datetime.now(timezone.utc) + + test_budget = type( + "LiteLLM_BudgetTableFull", + (), + { + "max_budget": 50.0, + "budget_duration": "30d", + "budget_reset_at": now - timedelta(hours=1), + "budget_id": "30d-tag-budget", + "created_at": now - timedelta(days=30), + }, + ) + + asyncio.run( + reset_budget_job.reset_budget_for_tags_linked_to_budgets( + budgets_to_reset=[test_budget] + ) + ) + + calls = mock_prisma_client.db.litellm_tagtable.update_many_calls + assert len(calls) == 1 + call = calls[0] + assert call["where"]["budget_id"] == {"in": ["30d-tag-budget"]} + assert call["where"]["spend"] == {"gt": 0} + assert call["data"]["spend"] == 0 + + +def test_reset_budget_for_tags_linked_to_budgets_empty( + reset_budget_job, mock_prisma_client +): + """ + Test that when there are no budgets to reset, no update is performed + on the tag table. + """ + asyncio.run( + reset_budget_job.reset_budget_for_tags_linked_to_budgets(budgets_to_reset=[]) + ) + calls = mock_prisma_client.db.litellm_tagtable.update_many_calls + assert len(calls) == 0 + + @pytest.mark.parametrize( "budget_duration, expected_day, expected_month", [ @@ -618,6 +754,75 @@ def test_budget_table_reset_also_resets_linked_keys( assert calls[0]["data"]["spend"] == 0 +def test_budget_table_reset_also_resets_linked_orgs( + reset_budget_job, mock_prisma_client +): + """ + Integration-style test: when reset_budget_for_litellm_budget_table runs, + it should also reset spend for orgs linked to the expiring budget tiers + (in addition to end-users, team members, and keys). + """ + now = datetime.now(timezone.utc) + + test_budget = type( + "LiteLLM_BudgetTableFull", + (), + { + "max_budget": 100.0, + "budget_duration": "30d", + "budget_reset_at": now - timedelta(hours=1), + "budget_id": "30d-org-budget", + "created_at": now - timedelta(days=30), + }, + ) + + mock_prisma_client.data["budget"] = [test_budget] + + asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) + + calls = mock_prisma_client.db.litellm_organizationtable.update_many_calls + assert len(calls) == 1, ( + "Expected reset_budget_for_litellm_budget_table to also reset orgs " + f"linked to expiring budgets, but got {len(calls)} update_many calls" + ) + assert calls[0]["where"]["budget_id"] == {"in": ["30d-org-budget"]} + assert calls[0]["data"]["spend"] == 0 + + +def test_budget_table_reset_also_resets_linked_tags( + reset_budget_job, mock_prisma_client +): + """ + Integration-style test: when reset_budget_for_litellm_budget_table runs, + it should also reset spend for tags linked to the expiring budget tiers. + """ + now = datetime.now(timezone.utc) + + test_budget = type( + "LiteLLM_BudgetTableFull", + (), + { + "max_budget": 50.0, + "budget_duration": "30d", + "budget_reset_at": now - timedelta(hours=1), + "budget_id": "30d-tag-budget", + "created_at": now - timedelta(days=30), + }, + ) + + mock_prisma_client.data["budget"] = [test_budget] + + asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) + + calls = mock_prisma_client.db.litellm_tagtable.update_many_calls + assert len(calls) == 1, ( + "Expected reset_budget_for_litellm_budget_table to also reset tags " + f"linked to expiring budgets, but got {len(calls)} update_many calls" + ) + assert calls[0]["where"]["budget_id"] == {"in": ["30d-tag-budget"]} + assert calls[0]["data"]["spend"] == 0 + + def test_reset_budget_resets_endusers_with_null_budget_id( reset_budget_job, mock_prisma_client ): @@ -1057,16 +1262,26 @@ def test_reset_budget_windows_query_error_does_not_break_team_path(monkeypatch): def _make_counter_invalidation_job(monkeypatch): - """Stub spend_counter_cache so we can observe invalidation calls.""" + """Stub spend_counter_cache (and user_api_key_cache) so we can observe + invalidation calls. + + Both caches are looked up via ``from litellm.proxy.proxy_server import + `` inside the reset job, so we publish them on a fake module. + """ spend_counter_cache = MagicMock() spend_counter_cache.in_memory_cache.set_cache = MagicMock() spend_counter_cache.redis_cache = MagicMock() spend_counter_cache.redis_cache.async_set_cache = AsyncMock() + user_api_key_cache = MagicMock() + user_api_key_cache.async_delete_cache = AsyncMock() + fake_module = types.ModuleType("litellm.proxy.proxy_server") fake_module.spend_counter_cache = spend_counter_cache + fake_module.user_api_key_cache = user_api_key_cache monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_module) + spend_counter_cache.user_api_key_cache = user_api_key_cache return spend_counter_cache @@ -1205,3 +1420,136 @@ def test_reset_budget_for_keys_linked_to_budgets_invalidates_redis_counter(monke counter_cache.in_memory_cache.set_cache.assert_any_call( key="spend:key:sk-linked", value=0.0, ttl=60 ) + + +def test_reset_budget_for_orgs_linked_to_budgets_invalidates_redis_counter(monkeypatch): + """Resetting orgs via budget tier must clear each linked org's counter.""" + counter_cache = _make_counter_invalidation_job(monkeypatch) + + expired_budget = type("B", (), {"budget_id": "budget-1"}) + linked_org = type("Org", (), {"organization_id": "org-acme"}) + + prisma_client = MagicMock() + prisma_client.db.litellm_organizationtable.find_many = AsyncMock( + return_value=[linked_org] + ) + prisma_client.db.litellm_organizationtable.update_many = AsyncMock( + return_value={"count": 1} + ) + + job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) + asyncio.run(job.reset_budget_for_orgs_linked_to_budgets([expired_budget])) + + counter_cache.in_memory_cache.set_cache.assert_any_call( + key="spend:org:org-acme", value=0.0, ttl=60 + ) + counter_cache.redis_cache.async_set_cache.assert_any_await( + key="spend:org:org-acme", value=0.0, ttl=60 + ) + + +def test_reset_budget_for_tags_linked_to_budgets_invalidates_redis_counter(monkeypatch): + """Resetting tags via budget tier must clear each linked tag's counter.""" + counter_cache = _make_counter_invalidation_job(monkeypatch) + + expired_budget = type("B", (), {"budget_id": "budget-1"}) + linked_tag = type("Tag", (), {"tag_name": "tenant-42"}) + + prisma_client = MagicMock() + prisma_client.db.litellm_tagtable.find_many = AsyncMock(return_value=[linked_tag]) + prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 1}) + + job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) + asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget])) + + counter_cache.in_memory_cache.set_cache.assert_any_call( + key="spend:tag:tenant-42", value=0.0, ttl=60 + ) + counter_cache.redis_cache.async_set_cache.assert_any_await( + key="spend:tag:tenant-42", value=0.0, ttl=60 + ) + + +def test_reset_budget_for_tags_linked_to_budgets_invalidates_management_cache( + monkeypatch, +): + """Regression guard for the bug where tag spend stayed frozen across cycles. + + ``SpendCounterReseed.from_db`` returns ``None`` for ``spend:tag:*`` keys, + so once the spend counter expires the tag budget check falls back to the + cached ``LiteLLM_TagTable.spend``. If we don't drop the management cache + entry on reset, that cached object lingers (TTL 60s) with the pre-reset + spend, and ``_tag_max_budget_check`` keeps returning HTTP 400 even though + the DB row has been zeroed. + """ + counter_cache = _make_counter_invalidation_job(monkeypatch) + + expired_budget = type("B", (), {"budget_id": "budget-1"}) + linked_tag = type("Tag", (), {"tag_name": "tenant-42"}) + + prisma_client = MagicMock() + prisma_client.db.litellm_tagtable.find_many = AsyncMock(return_value=[linked_tag]) + prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 1}) + + job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) + asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget])) + + counter_cache.user_api_key_cache.async_delete_cache.assert_any_await( + key="tag:tenant-42" + ) + + +def test_reset_budget_for_tags_linked_to_budgets_invalidates_each_tag_management_cache( + monkeypatch, +): + """When multiple tags share the expired budget tier, every one of them + has its ``user_api_key_cache`` entry dropped — not just the first.""" + counter_cache = _make_counter_invalidation_job(monkeypatch) + + expired_budget = type("B", (), {"budget_id": "budget-1"}) + linked_tags = [ + type("Tag", (), {"tag_name": "tenant-a"}), + type("Tag", (), {"tag_name": "tenant-b"}), + type("Tag", (), {"tag_name": "tenant-c"}), + ] + + prisma_client = MagicMock() + prisma_client.db.litellm_tagtable.find_many = AsyncMock(return_value=linked_tags) + prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 3}) + + job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) + asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget])) + + deleted_keys = { + call.kwargs.get("key") + for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list + } + assert deleted_keys == {"tag:tenant-a", "tag:tenant-b", "tag:tenant-c"} + + +def test_reset_budget_for_keys_linked_to_budgets_does_not_touch_management_cache( + monkeypatch, +): + """Cache invalidation is opt-in: keys / orgs / team-members rely on + ``SpendCounterReseed.from_db`` (which DOES handle their counter keys), + so the cache_key_fn hook is intentionally not wired for them. This test + locks in that no-op so a future refactor doesn't accidentally start + clobbering the key cache (which would cost an extra DB round-trip per + reset cycle without fixing anything).""" + counter_cache = _make_counter_invalidation_job(monkeypatch) + + expired_budget = type("B", (), {"budget_id": "budget-1"}) + linked_key = type("Key", (), {"token": "sk-linked"}) + + prisma_client = MagicMock() + prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[linked_key] + ) + prisma_client.db.litellm_verificationtoken.update_many = AsyncMock( + return_value={"count": 1} + ) + + job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) + asyncio.run(job.reset_budget_for_keys_linked_to_budgets([expired_budget])) + + counter_cache.user_api_key_cache.async_delete_cache.assert_not_awaited() diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index c6017752814..79e6494eab0 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -1513,3 +1513,146 @@ async def test_commit_spend_updates_uses_pipeline(): mock_redis_update_buffer.get_all_daily_end_user_spend_update_transactions_from_redis_buffer.assert_not_called() mock_redis_update_buffer.get_all_daily_agent_spend_update_transactions_from_redis_buffer.assert_not_called() mock_redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer.assert_not_called() + + +@pytest.mark.parametrize( + "bucket_name,input_dict,table_attr,method_name,where_key,expected_order", + [ + pytest.param( + "user_list_transactions", + {"user_c": 0.1, "user_a": 0.2, "user_b": 0.3}, + "litellm_usertable", + "update_many", + "user_id", + ["user_a", "user_b", "user_c"], + id="user", + ), + pytest.param( + "key_list_transactions", + {"tok_c": 0.1, "tok_a": 0.2, "tok_b": 0.3}, + "litellm_verificationtoken", + "update_many", + "token", + ["tok_a", "tok_b", "tok_c"], + id="key", + ), + pytest.param( + "team_list_transactions", + {"team_c": 0.1, "team_a": 0.2, "team_b": 0.3}, + "litellm_teamtable", + "update_many", + "team_id", + ["team_a", "team_b", "team_c"], + id="team", + ), + pytest.param( + "team_member_list_transactions", + { + "team_id::team_c::user_id::user_x": 0.1, + "team_id::team_a::user_id::user_x": 0.2, + "team_id::team_b::user_id::user_x": 0.3, + }, + "litellm_teammembership", + "update_many", + "team_id", + ["team_a", "team_b", "team_c"], + id="team_member", + ), + pytest.param( + "org_list_transactions", + {"org_c": 0.1, "org_a": 0.2, "org_b": 0.3}, + "litellm_organizationtable", + "update_many", + "organization_id", + ["org_a", "org_b", "org_c"], + id="org", + ), + pytest.param( + "end_user_list_transactions", + {"eu_c": 0.1, "eu_a": 0.2, "eu_b": 0.3}, + "litellm_endusertable", + "upsert", + "user_id", + ["eu_a", "eu_b", "eu_c"], + id="end_user", + ), + pytest.param( + "tag_list_transactions", + {"prod": 0.1, "customer-x": 0.2, "test": 0.3}, + "litellm_tagtable", + "update_many", + "tag_name", + ["customer-x", "prod", "test"], + id="tag", + ), + pytest.param( + "agent_list_transactions", + {"agent_c": 0.1, "agent_a": 0.2, "agent_b": 0.3}, + "litellm_agentstable", + "update_many", + "agent_id", + ["agent_a", "agent_b", "agent_c"], + id="agent", + ), + ], +) +@pytest.mark.asyncio +async def test_commit_spend_updates_iterates_in_sorted_order( + bucket_name, input_dict, table_attr, method_name, where_key, expected_order +): + """ + Every spend-bucket code path in _commit_spend_updates_to_db must iterate + in sorted order so concurrent pods acquire row locks in the same order + and avoid PostgreSQL deadlocks. Covers the 5 direct loops (user/key/team/ + team_member/org), the end_user path in ProxyUpdateSpend.update_end_user_spend, + and the shared _update_entity_spend_in_db helper (tag, agent). + """ + db_writer = DBSpendUpdateWriter() + + captured_where_values = [] + + def capture(*, where, data): + captured_where_values.append(where[where_key]) + + mock_batcher = MagicMock() + table_mock = MagicMock() + setattr(table_mock, method_name, MagicMock(side_effect=capture)) + setattr(mock_batcher, table_attr, table_mock) + + mock_transaction = AsyncMock() + mock_transaction.__aenter__ = AsyncMock(return_value=mock_transaction) + mock_transaction.__aexit__ = AsyncMock(return_value=False) + mock_transaction.batch_ = MagicMock( + return_value=AsyncMock( + __aenter__=AsyncMock(return_value=mock_batcher), + __aexit__=AsyncMock(return_value=False), + ) + ) + + mock_prisma_client = MagicMock() + mock_prisma_client.db = MagicMock() + mock_prisma_client.db.tx = MagicMock(return_value=mock_transaction) + + mock_proxy_logging = MagicMock() + mock_proxy_logging.call_details = {} + + buckets = { + "user_list_transactions": {}, + "end_user_list_transactions": {}, + "key_list_transactions": {}, + "team_list_transactions": {}, + "team_member_list_transactions": {}, + "org_list_transactions": {}, + "tag_list_transactions": {}, + "agent_list_transactions": {}, + } + buckets[bucket_name] = input_dict + + await db_writer._commit_spend_updates_to_db( + prisma_client=mock_prisma_client, + n_retry_times=3, + proxy_logging_obj=mock_proxy_logging, + db_spend_update_transactions=buckets, + ) + + assert captured_where_values == expected_order diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index 033deb3ff42..0d7becd3e2e 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -106,9 +106,11 @@ def mock_in_memory_handler(mocker): mock_handler = mocker.Mock(spec=InMemoryGuardrailHandler) mock_handler.list_in_memory_guardrails.return_value = [MOCK_CONFIG_GUARDRAIL] mock_handler.get_guardrail_by_id.return_value = MOCK_CONFIG_GUARDRAIL + mock_handler.get_source.return_value = "config" mock_handler.initialize_guardrail = mocker.Mock() mock_handler.update_in_memory_guardrail = mocker.Mock() mock_handler.delete_in_memory_guardrail = mocker.Mock() + mock_handler.reconcile_db_guardrails = mocker.Mock(return_value=[]) return mock_handler @@ -162,6 +164,67 @@ async def test_list_guardrails_v2_with_db_and_config( assert isinstance(config_guardrail.litellm_params, BaseLitellmParams) +@pytest.mark.asyncio +async def test_list_guardrails_v2_skips_stale_db_backed_in_memory_entries(mocker): + """ + A guardrail that's still in this pod's memory tagged source='db' but is no + longer in the DB result (deleted on another pod, awaiting reconcile) must + NOT surface in the list response — pre-fix it leaked as 'config'. + """ + stale_guardrail = { + "guardrail_id": "stale-db-id", + "guardrail_name": "Stale DB Guardrail", + "litellm_params": {"guardrail": "bedrock", "mode": "pre_call"}, + "guardrail_info": {}, + } + mock_prisma_client = mocker.Mock() + mock_prisma_client.db = mocker.Mock() + mock_prisma_client.db.litellm_guardrailstable = mocker.Mock() + mock_prisma_client.db.litellm_guardrailstable.find_many = AsyncMock(return_value=[]) + + mock_in_memory_handler = mocker.Mock() + mock_in_memory_handler.list_in_memory_guardrails.return_value = [stale_guardrail] + mock_in_memory_handler.get_source.return_value = "db" + + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mocker.patch( + "litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER", + mock_in_memory_handler, + ) + + admin_auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + response = await list_guardrails_v2(user_api_key_dict=admin_auth) + + assert response.guardrails == [] + mock_in_memory_handler.get_source.assert_called_with("stale-db-id") + + +@pytest.mark.asyncio +async def test_get_guardrail_info_404s_stale_db_backed_entry( + mocker, mock_prisma_client, mock_in_memory_handler +): + """ + Stale DB-backed entry (in-memory but not in DB) must 404 instead of being + returned as if it were a config-loaded guardrail. + """ + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mocker.patch( + "litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER", + mock_in_memory_handler, + ) + mock_prisma_client.db.litellm_guardrailstable.find_unique = AsyncMock( + return_value=None + ) + # In-memory still has it, but it's tagged as 'db' (stale, awaiting reconcile) + mock_in_memory_handler.get_source.return_value = "db" + + with pytest.raises(HTTPException) as exc_info: + await get_guardrail_info("stale-db-id") + + assert exc_info.value.status_code == 404 + assert "not found" in str(exc_info.value.detail) + + @pytest.mark.asyncio async def test_list_guardrails_v2_masks_sensitive_data_in_db_guardrails(mocker): """Test that sensitive litellm_params are masked for DB guardrails in list response""" @@ -1160,6 +1223,7 @@ async def test_get_guardrail_info_endpoint_config_guardrail(mocker): # Mock IN_MEMORY_GUARDRAIL_HANDLER at its source to return config guardrail mock_in_memory_handler = mocker.Mock() mock_in_memory_handler.get_guardrail_by_id.return_value = MOCK_CONFIG_GUARDRAIL + mock_in_memory_handler.get_source.return_value = "config" mocker.patch( "litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER", mock_in_memory_handler, diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py index 1d70126681d..9f7173383b0 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py @@ -60,3 +60,123 @@ def test_update_in_memory_guardrail(): handler.guardrail_id_to_custom_guardrail["123"].event_hook is GuardrailEventHooks.pre_call ) + + +def _make_guardrail(guardrail_id: str, name: str = "g") -> Guardrail: + return Guardrail( + guardrail_id=guardrail_id, + guardrail_name=name, + litellm_params=LitellmParams(guardrail=name, mode="pre_call", default_on=False), + ) + + +def test_reconcile_db_guardrails_drops_stale_db_entries_only(): + """ + The reconcile pass must drop in-memory entries marked source='db' that are + missing from the DB result, and never touch source='config' entries. + Models the multi-pod case where another pod deleted a DB-backed guardrail. + """ + handler = InMemoryGuardrailHandler() + + # Two DB-backed entries on this pod (synced from earlier polling cycles) + handler.IN_MEMORY_GUARDRAILS["db-keep"] = _make_guardrail("db-keep") + handler.IN_MEMORY_GUARDRAILS["db-stale"] = _make_guardrail("db-stale") + handler._sources["db-keep"] = "db" + handler._sources["db-stale"] = "db" + + # One config-loaded entry that must survive reconciliation + handler.IN_MEMORY_GUARDRAILS["cfg"] = _make_guardrail("cfg") + handler._sources["cfg"] = "config" + + # The DB now only contains db-keep — db-stale was deleted on another pod. + removed = handler.reconcile_db_guardrails(db_guardrail_ids={"db-keep"}) + + assert removed == ["db-stale"] + assert "db-stale" not in handler.IN_MEMORY_GUARDRAILS + assert "db-stale" not in handler._sources + assert "db-keep" in handler.IN_MEMORY_GUARDRAILS + assert "cfg" in handler.IN_MEMORY_GUARDRAILS + assert handler._sources["cfg"] == "config" + + +def test_reconcile_does_not_drop_config_entries_missing_from_db(): + """A config-only guardrail (no DB row) must never be reconciled away.""" + handler = InMemoryGuardrailHandler() + handler.IN_MEMORY_GUARDRAILS["cfg-only"] = _make_guardrail("cfg-only") + handler._sources["cfg-only"] = "config" + + removed = handler.reconcile_db_guardrails(db_guardrail_ids=set()) + + assert removed == [] + assert "cfg-only" in handler.IN_MEMORY_GUARDRAILS + + +def test_get_source_returns_marker_set_at_insert(): + handler = InMemoryGuardrailHandler() + handler.IN_MEMORY_GUARDRAILS["a"] = _make_guardrail("a") + handler._sources["a"] = "db" + handler.IN_MEMORY_GUARDRAILS["b"] = _make_guardrail("b") + handler._sources["b"] = "config" + + assert handler.get_source("a") == "db" + assert handler.get_source("b") == "config" + assert handler.get_source("missing") is None + + +def test_delete_in_memory_guardrail_clears_source_marker(): + handler = InMemoryGuardrailHandler() + handler.IN_MEMORY_GUARDRAILS["a"] = _make_guardrail("a") + handler._sources["a"] = "db" + + handler.delete_in_memory_guardrail("a") + + assert "a" not in handler.IN_MEMORY_GUARDRAILS + assert "a" not in handler._sources + assert handler.get_source("a") is None + + +def test_initialize_guardrail_early_return_updates_source_marker(): + """ + When initialize_guardrail is called for a guardrail that already exists + in memory, the early-return path must still honor the caller's source. + Otherwise a racing polling tick that placed a DB entry in memory first + would leave a later config-init call wrongly marked as 'db' (or vice + versa), and the entry would be reconciled with the wrong classification. + """ + handler = InMemoryGuardrailHandler() + # Simulate a polling tick already placing the entry as DB-backed. + handler.IN_MEMORY_GUARDRAILS["collide"] = _make_guardrail("collide", name="bedrock") + handler._sources["collide"] = "db" + + # Config init re-visits the same id (e.g., hot-reload, or UUID collision). + g = Guardrail( + guardrail_id="collide", + guardrail_name="bedrock", + litellm_params=LitellmParams( + guardrail="bedrock", mode="pre_call", default_on=False + ), + ) + handler.initialize_guardrail(guardrail=g, source="config") + + assert handler.get_source("collide") == "config" + + # And the symmetric direction: db sync should override an entry left + # marked as 'config' from a stale init path. + handler.initialize_guardrail(guardrail=g, source="db") + assert handler.get_source("collide") == "db" + + +def test_sync_guardrail_from_db_marks_source_db_when_unchanged(): + """ + sync_guardrail_from_db must enforce source='db' even when params are + unchanged, so a config entry whose UUID happens to collide with a later + DB row gets re-tagged correctly. + """ + handler = InMemoryGuardrailHandler() + g = _make_guardrail("collide") + handler.IN_MEMORY_GUARDRAILS["collide"] = g + handler._sources["collide"] = "config" + + handler.sync_guardrail_from_db(g) + + assert handler.get_source("collide") == "db" diff --git a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py index 2edcb00c967..bcd7fcb37b3 100644 --- a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py @@ -466,6 +466,236 @@ async def test_test_model_connection_loads_config_from_router(): assert "result" in result +@pytest.mark.asyncio +async def test_test_model_connection_uses_model_info_id_to_disambiguate_duplicate_model_names(): + """ + When two deployments share the same `model_name` (e.g. wildcard + `openai/*`) but have different `api_base` values, clicking "Test + Connection" on a specific row in the UI must probe THAT row's + `api_base` — not whichever happens to be `deployments[0]`. + + The UI passes `model_info.id` to identify the deployment the user + actually clicked on. The backend must use that id to look up the + specific deployment rather than always grabbing the first match. + + Regression test for: silent fallback to deployments[0] when + multiple deployments share a wildcard model_name. + """ + from litellm.types.router import Deployment, LiteLLM_Params + + mock_request = MagicMock() + mock_user_api_key_dict = MagicMock() + mock_user_api_key_dict.user_id = "test-user" + mock_user_api_key_dict.token = "test-token" + + mock_prisma_client = MagicMock() + + deployment_a = { + "model_name": "openai/*", + "litellm_params": { + "model": "openai/*", + "api_base": "https://deployment-A-base.invalid/v1", + "api_key": "fake-key-A", + }, + "model_info": {"id": "deployment-A-id"}, + } + deployment_b = { + "model_name": "openai/*", + "litellm_params": { + "model": "openai/*", + "api_base": "https://deployment-B-base.invalid/v1", + "api_key": "fake-key-B", + }, + "model_info": {"id": "deployment-B-id"}, + } + + mock_router = MagicMock() + mock_router.get_model_list.return_value = [deployment_a, deployment_b] + + # Backend uses get_deployment(model_id=...) for O(1) lookup by id. + def _get_deployment_by_id(model_id): + if model_id == "deployment-A-id": + return Deployment( + model_name="openai/*", + litellm_params=LiteLLM_Params(**deployment_a["litellm_params"]), + model_info=deployment_a["model_info"], + ) + if model_id == "deployment-B-id": + return Deployment( + model_name="openai/*", + litellm_params=LiteLLM_Params(**deployment_b["litellm_params"]), + model_info=deployment_b["model_info"], + ) + return None + + mock_router.get_deployment.side_effect = _get_deployment_by_id + + mock_can_user_make_model_call = AsyncMock() + + mock_health_check_result = {"status": "healthy", "response_time_ms": 50} + mock_ahealth_check = AsyncMock(return_value=mock_health_check_result) + mock_run_with_timeout = AsyncMock(return_value=mock_health_check_result) + + def mock_update_params(model_info, litellm_params): + params = litellm_params.copy() + params["messages"] = [{"role": "user", "content": "test"}] + return params + + def mock_reject_os_environ(params): + return None + + with ( + patch( + "litellm.proxy.proxy_server.prisma_client", + mock_prisma_client, + ), + patch( + "litellm.proxy.proxy_server.llm_router", + mock_router, + ), + patch( + "litellm.proxy.proxy_server.premium_user", + False, + ), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + mock_can_user_make_model_call, + ), + patch( + "litellm.proxy.health_endpoints._health_endpoints.litellm.ahealth_check", + mock_ahealth_check, + ), + patch( + "litellm.proxy.health_endpoints._health_endpoints.run_with_timeout", + mock_run_with_timeout, + ), + patch( + "litellm.proxy.health_endpoints._health_endpoints._update_litellm_params_for_health_check", + mock_update_params, + ), + patch( + "litellm.proxy.health_endpoints._health_endpoints._reject_os_environ_references", + mock_reject_os_environ, + ), + ): + # Click "Test Connection" on deployment B (NOT the first one). + # The UI sends only `model` + `model_info.id` — it does NOT + # send `api_base`/`api_key`, so the backend must resolve them + # from the right deployment. + await health_test_model_connection( + request=mock_request, + mode="chat", + litellm_params={"model": "openai/*"}, + model_info={"id": "deployment-B-id"}, + user_api_key_dict=mock_user_api_key_dict, + ) + + # The outbound health check must hit deployment B's api_base. + ahealth_check_call_args = mock_ahealth_check.call_args + assert ahealth_check_call_args is not None + model_params = ahealth_check_call_args.kwargs.get("model_params", {}) + + assert model_params.get("api_base") == ( + "https://deployment-B-base.invalid/v1" + ), ( + "Expected /health/test_connection to probe deployment B's " + "api_base when model_info.id='deployment-B-id' was provided. " + f"Got: {model_params.get('api_base')!r}. This means the " + "backend silently fell back to deployments[0] (A) instead " + "of disambiguating by model_info.id." + ) + assert model_params.get("api_key") == "fake-key-B" + + +@pytest.mark.asyncio +async def test_test_model_connection_falls_back_to_deployments_zero_without_id(): + """ + Backwards-compat: when the request body does NOT include + `model_info.id`, the legacy behavior of using `deployments[0]` + is preserved (single-deployment case, or callers that haven't + been updated to pass an id). + """ + mock_request = MagicMock() + mock_user_api_key_dict = MagicMock() + mock_user_api_key_dict.user_id = "test-user" + mock_user_api_key_dict.token = "test-token" + + mock_prisma_client = MagicMock() + + deployment_a = { + "model_name": "openai/*", + "litellm_params": { + "model": "openai/*", + "api_base": "https://deployment-A-base.invalid/v1", + "api_key": "fake-key-A", + }, + "model_info": {"id": "deployment-A-id"}, + } + deployment_b = { + "model_name": "openai/*", + "litellm_params": { + "model": "openai/*", + "api_base": "https://deployment-B-base.invalid/v1", + "api_key": "fake-key-B", + }, + "model_info": {"id": "deployment-B-id"}, + } + + mock_router = MagicMock() + mock_router.get_model_list.return_value = [deployment_a, deployment_b] + + mock_can_user_make_model_call = AsyncMock() + mock_health_check_result = {"status": "healthy"} + mock_ahealth_check = AsyncMock(return_value=mock_health_check_result) + mock_run_with_timeout = AsyncMock(return_value=mock_health_check_result) + + def mock_update_params(model_info, litellm_params): + params = litellm_params.copy() + params["messages"] = [{"role": "user", "content": "test"}] + return params + + def mock_reject_os_environ(params): + return None + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), + patch("litellm.proxy.proxy_server.llm_router", mock_router), + patch("litellm.proxy.proxy_server.premium_user", False), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + mock_can_user_make_model_call, + ), + patch( + "litellm.proxy.health_endpoints._health_endpoints.litellm.ahealth_check", + mock_ahealth_check, + ), + patch( + "litellm.proxy.health_endpoints._health_endpoints.run_with_timeout", + mock_run_with_timeout, + ), + patch( + "litellm.proxy.health_endpoints._health_endpoints._update_litellm_params_for_health_check", + mock_update_params, + ), + patch( + "litellm.proxy.health_endpoints._health_endpoints._reject_os_environ_references", + mock_reject_os_environ, + ), + ): + await health_test_model_connection( + request=mock_request, + mode="chat", + litellm_params={"model": "openai/*"}, + model_info={}, # no id provided + user_api_key_dict=mock_user_api_key_dict, + ) + + # Without id, deployments[0] (A) should be used (legacy behavior). + model_params = mock_ahealth_check.call_args.kwargs.get("model_params", {}) + assert model_params.get("api_base") == "https://deployment-A-base.invalid/v1" + assert model_params.get("api_key") == "fake-key-A" + + @pytest.mark.asyncio async def test_health_services_endpoint_datadog_llm_observability(): """ diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index b292e8d0cae..66716400c4f 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -5689,7 +5689,7 @@ async def test_process_single_key_update(): "litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook" ): # Create update request - key_update_item = BulkUpdateKeyRequestItem( + update_key_request = UpdateKeyRequest( key="test-key-123", max_budget=100.0, tags=["production"], @@ -5703,7 +5703,7 @@ async def test_process_single_key_update(): # Call the function result = await _process_single_key_update( - key_update_item=key_update_item, + update_key_request=update_key_request, user_api_key_dict=user_api_key_dict, litellm_changed_by=None, prisma_client=mock_prisma_client, @@ -9855,9 +9855,6 @@ async def test_process_single_key_update_cache_invalidation_with_token_hash(): from litellm.proxy.management_endpoints.key_management_endpoints import ( _process_single_key_update, ) - from litellm.types.proxy.management_endpoints.key_management_endpoints import ( - BulkUpdateKeyRequestItem, - ) token_hash = "abc123def456" @@ -9900,7 +9897,7 @@ async def test_process_single_key_update_cache_invalidation_with_token_hash(): new_callable=AsyncMock, ), ): - key_update_item = BulkUpdateKeyRequestItem( + update_key_request = UpdateKeyRequest( key=token_hash, max_budget=100.0, ) @@ -9912,7 +9909,7 @@ async def test_process_single_key_update_cache_invalidation_with_token_hash(): ) await _process_single_key_update( - key_update_item=key_update_item, + update_key_request=update_key_request, user_api_key_dict=user_api_key_dict, litellm_changed_by=None, prisma_client=mock_prisma_client, @@ -10019,3 +10016,583 @@ async def test_execute_virtual_key_regeneration_cache_invalidation_with_token_ha call_kwargs = mock_delete_cache.call_args.kwargs # The token hash should be passed as-is, NOT double-hashed assert call_kwargs["hashed_token"] == token_hash + + +# --------------------------------------------------------------------------- +# /team/key/bulk_update tests +# --------------------------------------------------------------------------- + + +_BULK_PKG = "litellm.proxy.management_endpoints.key_management_endpoints" + + +def _make_team_key(token: str, team_id: str = "team-abc") -> LiteLLM_VerificationToken: + return LiteLLM_VerificationToken( + token=token, + user_id="user-123", + models=[], + team_id=team_id, + max_budget=None, + ) + + +def _admin() -> UserAPIKeyAuth: + return UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin" + ) + + +def _internal_user() -> UserAPIKeyAuth: + return UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, api_key="sk-iu", user_id="iu" + ) + + +def _updated(payload): + m = MagicMock() + m.model_dump.return_value = payload + return m + + +def _setup_team_keys_mocks( + monkeypatch, + *, + find_many=None, + find_unique=None, + update_data=None, + hash_identity=True, +): + """Set up mocks for bulk_update_team_keys; returns mock_prisma.""" + mock_prisma = AsyncMock() + mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[] if find_many is None else find_many + ) + if find_unique is not None: + mock_prisma.db.litellm_verificationtoken.find_unique = find_unique + if update_data is not None: + mock_prisma.update_data = update_data + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.user_custom_key_update", None) + monkeypatch.setattr( + f"{_BULK_PKG}.prepare_key_update_data", + AsyncMock(return_value={"max_budget": 50.0}), + ) + monkeypatch.setattr(f"{_BULK_PKG}._delete_cache_key_object", AsyncMock()) + monkeypatch.setattr( + f"{_BULK_PKG}.KeyManagementEventHooks.async_key_updated_hook", AsyncMock() + ) + monkeypatch.setattr(f"{_BULK_PKG}.get_team_object", AsyncMock(return_value=None)) + monkeypatch.setattr(f"{_BULK_PKG}._check_team_key_limits", AsyncMock()) + monkeypatch.setattr( + f"{_BULK_PKG}.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint", + AsyncMock(), + ) + if hash_identity: + # Tests use already-hashed tokens; the raw-sk regression opts out. + monkeypatch.setattr(f"{_BULK_PKG}._hash_token_if_needed", lambda token: token) + return mock_prisma + + +async def _call_as_admin(data): + from litellm.proxy.management_endpoints.key_management_endpoints import ( + bulk_update_team_keys, + ) + + return await bulk_update_team_keys( + data=data, user_api_key_dict=_admin(), litellm_changed_by=None + ) + + +# ---- happy paths ---------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_bulk_update_team_keys_success_with_key_ids(monkeypatch): + from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BulkUpdateTeamKeysRequest, + KeyUpdateFields, + ) + + keys = [_make_team_key("tok-a"), _make_team_key("tok-b")] + find_unique = AsyncMock(side_effect=keys) + mock = _setup_team_keys_mocks( + monkeypatch, + find_many=keys, + find_unique=find_unique, + update_data=AsyncMock( + side_effect=[{"data": _updated({"max_budget": 50.0})}] * 2 + ), + ) + + response = await _call_as_admin( + BulkUpdateTeamKeysRequest( + team_id="team-abc", + key_ids=["tok-a", "tok-b"], + update_fields=KeyUpdateFields(max_budget=50.0), + ) + ) + + assert len(response.successful_updates) == 2 + assert len(response.failed_updates) == 0 + where = mock.db.litellm_verificationtoken.find_many.await_args.kwargs["where"] + assert where["team_id"] == "team-abc" + assert where["token"] == {"in": ["tok-a", "tok-b"]} + find_unique.assert_not_called() + + +@pytest.mark.asyncio +async def test_bulk_update_team_keys_success_all_keys_in_team(monkeypatch): + from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BulkUpdateTeamKeysRequest, + KeyUpdateFields, + ) + + keys = [_make_team_key(f"tok-{i}") for i in range(3)] + find_unique = AsyncMock(side_effect=keys) + mock = _setup_team_keys_mocks( + monkeypatch, + find_many=keys, + find_unique=find_unique, + update_data=AsyncMock( + side_effect=[{"data": _updated({"max_budget": 50.0})}] * 3 + ), + ) + + response = await _call_as_admin( + BulkUpdateTeamKeysRequest( + team_id="team-abc", + all_keys_in_team=True, + update_fields=KeyUpdateFields(max_budget=50.0), + ) + ) + + assert len(response.successful_updates) == 3 + where = mock.db.litellm_verificationtoken.find_many.await_args.kwargs["where"] + # `blocked` is Boolean? with no default → /key/generate writes NULL. Prisma's + # NOT excludes NULLs, so the filter has to OR `false` with `null` explicitly. + blocked_or, expires_or = where["AND"][0]["OR"], where["AND"][1]["OR"] + assert {"blocked": False} in blocked_or and {"blocked": None} in blocked_or + assert {"expires": None} in expires_or + assert any( + "gt" in c.get("expires", {}) + for c in expires_or + if isinstance(c.get("expires"), dict) + ) + find_unique.assert_not_called() + + +@pytest.mark.asyncio +async def test_bulk_update_team_keys_key_not_in_team(monkeypatch): + from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BulkUpdateTeamKeysRequest, + KeyUpdateFields, + ) + + in_team = _make_team_key("tok-a") + _setup_team_keys_mocks( + monkeypatch, + find_many=[in_team], + find_unique=AsyncMock(return_value=in_team), + update_data=AsyncMock(return_value={"data": _updated({"max_budget": 50.0})}), + ) + + response = await _call_as_admin( + BulkUpdateTeamKeysRequest( + team_id="team-abc", + key_ids=["tok-a", "tok-foreign"], + update_fields=KeyUpdateFields(max_budget=50.0), + ) + ) + assert [u.key for u in response.successful_updates] == ["tok-a"] + assert [u.key for u in response.failed_updates] == ["tok-foreign"] + assert "not found in team" in response.failed_updates[0].failed_reason + + +# ---- error paths ---------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_bulk_update_team_keys_batch_size_cap(monkeypatch): + from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BulkUpdateTeamKeysRequest, + KeyUpdateFields, + ) + + _setup_team_keys_mocks( + monkeypatch, + find_many=[_make_team_key(f"tok-{i}") for i in range(501)], + ) + + with pytest.raises(HTTPException) as exc: + await _call_as_admin( + BulkUpdateTeamKeysRequest( + team_id="team-abc", + all_keys_in_team=True, + update_fields=KeyUpdateFields(max_budget=50.0), + ) + ) + assert exc.value.status_code == 400 + assert "more than 500" in exc.value.detail["error"] + + +@pytest.mark.asyncio +async def test_bulk_update_team_keys_empty_team_returns_404(monkeypatch): + from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BulkUpdateTeamKeysRequest, + KeyUpdateFields, + ) + + _setup_team_keys_mocks(monkeypatch, find_many=[]) + with pytest.raises(HTTPException) as exc: + await _call_as_admin( + BulkUpdateTeamKeysRequest( + team_id="team-empty", + all_keys_in_team=True, + update_fields=KeyUpdateFields(max_budget=50.0), + ) + ) + assert exc.value.status_code == 404 + + +# ---- auth ----------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_bulk_update_team_keys_team_member_with_permission(monkeypatch): + from litellm.proxy.management_endpoints.key_management_endpoints import ( + bulk_update_team_keys, + ) + from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BulkUpdateTeamKeysRequest, + KeyUpdateFields, + ) + + key_a = _make_team_key("tok-a") + _setup_team_keys_mocks( + monkeypatch, + find_many=[key_a], + find_unique=AsyncMock(return_value=key_a), + update_data=AsyncMock(return_value={"data": _updated({"max_budget": 50.0})}), + ) + auth_check = AsyncMock() + monkeypatch.setattr( + f"{_BULK_PKG}.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint", + auth_check, + ) + + response = await bulk_update_team_keys( + data=BulkUpdateTeamKeysRequest( + team_id="team-abc", + all_keys_in_team=True, + update_fields=KeyUpdateFields(max_budget=50.0), + ), + user_api_key_dict=_internal_user(), + litellm_changed_by=None, + ) + assert len(response.successful_updates) == 1 + # Upfront check + per-key check inside _process_single_key_update + assert auth_check.await_count == 2 + + +@pytest.mark.asyncio +async def test_bulk_update_team_keys_team_member_no_permission(monkeypatch): + from litellm.proxy.management_endpoints.key_management_endpoints import ( + bulk_update_team_keys, + ) + from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BulkUpdateTeamKeysRequest, + KeyUpdateFields, + ) + + mock = _setup_team_keys_mocks(monkeypatch, find_many=[_make_team_key("tok-a")]) + monkeypatch.setattr( + f"{_BULK_PKG}.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint", + AsyncMock( + side_effect=ProxyException( + message="not in team", + type="team_member_permission_error", + param="/key/update", + code=401, + ) + ), + ) + + with pytest.raises(ProxyException): + await bulk_update_team_keys( + data=BulkUpdateTeamKeysRequest( + team_id="team-abc", + all_keys_in_team=True, + update_fields=KeyUpdateFields(max_budget=1.0), + ), + user_api_key_dict=_internal_user(), + litellm_changed_by=None, + ) + mock.update_data.assert_not_called() + + +# ---- pydantic-layer validation ------------------------------------------- + + +def test_bulk_update_team_keys_request_validation(): + """Allowlist (extra='forbid'), empty-payload rejection, and selection XOR.""" + from pydantic import ValidationError + + from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BulkUpdateTeamKeysRequest, + KeyUpdateFields, + ) + + forbidden = [ + "key", + "key_alias", + "team_id", + "allowed_routes", + "allowed_passthrough_routes", + "permissions", + "object_permission", + "access_group_ids", + "user_id", + "organization_id", + "blocked", + "key_type", + "models", + "config", + "router_settings", + "spend", + ] + for f in forbidden: + with pytest.raises(ValidationError, match=f): + KeyUpdateFields(**{f: True}) + + with pytest.raises(ValidationError, match="at least one"): + KeyUpdateFields() + + assert KeyUpdateFields(max_budget=50.0, tags=["x"]).max_budget == 50.0 + + valid = KeyUpdateFields(max_budget=10) + with pytest.raises(ValidationError): + BulkUpdateTeamKeysRequest( + team_id="t", key_ids=["k"], all_keys_in_team=True, update_fields=valid + ) + with pytest.raises(ValidationError): + BulkUpdateTeamKeysRequest(team_id="t", update_fields=valid) + + +# ---- security regressions ------------------------------------------------ + + +@pytest.mark.asyncio +async def test_bulk_update_team_keys_hashes_raw_sk_key_ids(monkeypatch): + """Regression: raw sk-... key_ids must be hashed before the find_many lookup.""" + from litellm.proxy._types import hash_token + from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BulkUpdateTeamKeysRequest, + KeyUpdateFields, + ) + + raw_sk = "sk-rawkey1234567890" + hashed = hash_token(raw_sk) + row = LiteLLM_VerificationToken( + token=hashed, user_id="u", models=[], team_id="team-abc", max_budget=None + ) + mock = _setup_team_keys_mocks( + monkeypatch, + find_many=[row], + find_unique=AsyncMock(return_value=row), + update_data=AsyncMock(return_value={"data": _updated({"max_budget": 50.0})}), + hash_identity=False, + ) + + response = await _call_as_admin( + BulkUpdateTeamKeysRequest( + team_id="team-abc", + key_ids=[raw_sk], + update_fields=KeyUpdateFields(max_budget=50.0), + ) + ) + where = mock.db.litellm_verificationtoken.find_many.await_args.kwargs["where"] + assert where["token"] == {"in": [hashed]} + # Response reports the user-supplied form, not the hash. + assert response.successful_updates[0].key == raw_sk + + +@pytest.mark.asyncio +async def test_bulk_update_team_keys_auth_check_runs_when_no_keys_match(monkeypatch): + """Regression: non-admin with bogus key_ids must still hit the membership gate.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + bulk_update_team_keys, + ) + from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BulkUpdateTeamKeysRequest, + KeyUpdateFields, + ) + + mock = _setup_team_keys_mocks(monkeypatch, find_many=[]) + auth_check = AsyncMock( + side_effect=ProxyException( + message="not in team", + type="team_member_permission_error", + param="/key/update", + code=401, + ) + ) + monkeypatch.setattr( + f"{_BULK_PKG}.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint", + auth_check, + ) + + with pytest.raises(ProxyException): + await bulk_update_team_keys( + data=BulkUpdateTeamKeysRequest( + team_id="victim-team", + key_ids=["bogus-1", "bogus-2"], + update_fields=KeyUpdateFields(max_budget=1.0), + ), + user_api_key_dict=_internal_user(), + litellm_changed_by=None, + ) + # Anchored on data.team_id, not existing_keys[0]. + assert auth_check.await_args.kwargs["existing_key_row"].team_id == "victim-team" + mock.update_data.assert_not_called() + + +@pytest.mark.asyncio +async def test_bulk_update_team_keys_does_not_log_raw_sk_token_on_failure( + monkeypatch, caplog +): + """Regression: per-key failure must not log the raw sk-... (ERROR-level logs persist).""" + import logging + + from litellm.proxy._types import hash_token + from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BulkUpdateTeamKeysRequest, + KeyUpdateFields, + ) + + raw_sk = "sk-supersecret1234567890" + row = LiteLLM_VerificationToken( + token=hash_token(raw_sk), + user_id="u", + models=[], + team_id="team-abc", + max_budget=None, + ) + _setup_team_keys_mocks( + monkeypatch, + find_many=[row], + update_data=AsyncMock(side_effect=RuntimeError("boom")), + hash_identity=False, + ) + + with caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"): + response = await _call_as_admin( + BulkUpdateTeamKeysRequest( + team_id="team-abc", + key_ids=[raw_sk], + update_fields=KeyUpdateFields(max_budget=50.0), + ) + ) + assert len(response.failed_updates) == 1 + log_text = "\n".join(r.getMessage() for r in caplog.records) + assert raw_sk not in log_text + + +@pytest.mark.asyncio +async def test_bulk_update_team_keys_propagates_team_id_to_per_key_request(monkeypatch): + """Regression: per-key UpdateKeyRequest carries data.team_id (gates _check_team_key_limits).""" + from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BulkUpdateTeamKeysRequest, + KeyUpdateFields, + ) + + _setup_team_keys_mocks(monkeypatch, find_many=[_make_team_key("tok-a")]) + captured = [] + + async def fake_process(*, update_key_request, **kw): + captured.append(update_key_request) + return {"max_budget": update_key_request.max_budget} + + monkeypatch.setattr(f"{_BULK_PKG}._process_single_key_update", fake_process) + + await _call_as_admin( + BulkUpdateTeamKeysRequest( + team_id="team-abc", + key_ids=["tok-a"], + update_fields=KeyUpdateFields( + tpm_limit=10_000, tpm_limit_type="guaranteed_throughput" + ), + ) + ) + assert captured[0].team_id == "team-abc" + assert captured[0].tpm_limit_type == "guaranteed_throughput" + + +@pytest.mark.asyncio +async def test_bulk_update_team_keys_dedupes_key_ids(monkeypatch): + """Duplicate key_ids collapse to a single update (no redundant DB writes, no inflated counts).""" + from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BulkUpdateTeamKeysRequest, + KeyUpdateFields, + ) + + key_a = _make_team_key("tok-a") + update_data = AsyncMock(return_value={"data": _updated({"max_budget": 50.0})}) + _setup_team_keys_mocks( + monkeypatch, + find_many=[key_a], + find_unique=AsyncMock(return_value=key_a), + update_data=update_data, + ) + + response = await _call_as_admin( + BulkUpdateTeamKeysRequest( + team_id="team-abc", + key_ids=["tok-a", "tok-a", "tok-a"], + update_fields=KeyUpdateFields(max_budget=50.0), + ) + ) + + assert response.total_requested == 1 + assert len(response.successful_updates) == 1 + assert len(response.failed_updates) == 0 + update_data.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_bulk_update_team_keys_blocks_metadata_allowed_passthrough_routes( + monkeypatch, +): + """Non-admin can't grant passthrough access by smuggling allowed_passthrough_routes through metadata.""" + from fastapi import HTTPException + + from litellm.proxy.management_endpoints.key_management_endpoints import ( + bulk_update_team_keys, + ) + from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BulkUpdateTeamKeysRequest, + KeyUpdateFields, + ) + + mock = _setup_team_keys_mocks(monkeypatch, find_many=[_make_team_key("tok-a")]) + + request = BulkUpdateTeamKeysRequest( + team_id="team-abc", + all_keys_in_team=True, + update_fields=KeyUpdateFields( + metadata={"allowed_passthrough_routes": ["/admin/*"]} + ), + ) + + with pytest.raises(HTTPException) as exc: + await bulk_update_team_keys( + data=request, + user_api_key_dict=_internal_user(), + litellm_changed_by=None, + ) + + assert exc.value.status_code == 403 + assert "allowed_passthrough_routes" in str(exc.value.detail) + mock.update_data.assert_not_called() diff --git a/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py index 39ec6f075d7..76ba0e3dc67 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py @@ -380,6 +380,408 @@ async def test_list_tags_no_dynamic_tags(): app.dependency_overrides.clear() +async def test_internal_user_list_tags_only_returns_tags_used_by_their_keys(): + """ + Internal users can view tag usage, but the tag list must be scoped to tags + produced by API keys owned by the caller. + """ + from datetime import datetime + from unittest.mock import AsyncMock, Mock + + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + mock_user_auth = UserAPIKeyAuth( + api_key="current-owned-key", + user_id="internal-user-123", + user_role=LitellmUserRoles.INTERNAL_USER, + ) + app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth + + try: + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_db = Mock() + mock_prisma.db = mock_db + + owned_key_record = Mock() + owned_key_record.token = "owned-key" + mock_db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[owned_key_record] + ) + + mock_db.litellm_dailytagspend.group_by = AsyncMock( + return_value=[ + { + "tag": "stored-owned-tag", + "_min": {"created_at": "2025-02-01T00:00:00Z"}, + "_max": {"updated_at": "2025-03-01T00:00:00Z"}, + }, + { + "tag": "dynamic-owned-tag", + "_min": {"created_at": "2025-02-02T00:00:00Z"}, + "_max": {"updated_at": "2025-03-02T00:00:00Z"}, + }, + ] + ) + + stored_tag = Mock() + stored_tag.tag_name = "stored-owned-tag" + stored_tag.description = "A stored tag used by the caller" + stored_tag.models = ["model-1"] + stored_tag.model_info = {} + stored_tag.spend = 0.0 + stored_tag.budget_id = None + stored_tag.created_at = datetime(2025, 1, 1) + stored_tag.updated_at = datetime(2025, 1, 1) + stored_tag.created_by = "admin-user" + stored_tag.litellm_budget_table = None + mock_db.litellm_tagtable.find_many = AsyncMock(return_value=[stored_tag]) + + response = client.get( + "/tag/list", + headers={"Authorization": "Bearer test-key"}, + ) + + assert response.status_code == 200 + assert [tag["name"] for tag in response.json()] == [ + "stored-owned-tag", + "dynamic-owned-tag", + ] + mock_db.litellm_verificationtoken.find_many.assert_awaited_once_with( + where={"user_id": "internal-user-123"}, + select={"token": True}, + ) + mock_db.litellm_dailytagspend.group_by.assert_awaited_once_with( + by=["tag"], + where={ + "tag": {"not": None}, + "api_key": {"in": ["current-owned-key", "owned-key"]}, + }, + min={"created_at": True}, + max={"updated_at": True}, + ) + mock_db.litellm_tagtable.find_many.assert_awaited_once_with( + where={"tag_name": {"in": ["stored-owned-tag", "dynamic-owned-tag"]}}, + include={"litellm_budget_table": True}, + ) + + finally: + app.dependency_overrides.clear() + + +@pytest.mark.asyncio +async def test_list_tags_with_date_range_filters_dynamic_tags(): + """ + /tag/list?start_date=...&end_date=... should push the date window into + the dailytagspend group_by WHERE clause so large tables don't get scanned. + """ + from unittest.mock import AsyncMock, Mock + + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + mock_user_auth = UserAPIKeyAuth( + user_id="test-user-123", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth + + try: + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_db = Mock() + mock_prisma.db = mock_db + mock_db.litellm_tagtable.find_many = AsyncMock(return_value=[]) + group_by_mock = AsyncMock(return_value=[]) + mock_db.litellm_dailytagspend.group_by = group_by_mock + + headers = {"Authorization": "Bearer sk-1234"} + response = client.get( + "/tag/list?start_date=2026-04-01&end_date=2026-04-29", + headers=headers, + ) + + assert response.status_code == 200 + group_by_mock.assert_awaited_once() + where = group_by_mock.await_args.kwargs["where"] + assert where["tag"] == {"not": None} + assert where["date"] == {"gte": "2026-04-01", "lte": "2026-04-29"} + + finally: + app.dependency_overrides.clear() + + +@pytest.mark.asyncio +async def test_internal_user_tag_daily_activity_is_scoped_to_their_keys(): + """ + Internal users must not receive proxy-wide tag spend rows when viewing tag + usage daily activity. + """ + from unittest.mock import AsyncMock, Mock + + from litellm.proxy.management_endpoints.tag_management_endpoints import ( + get_tag_daily_activity, + ) + + mock_user_auth = UserAPIKeyAuth( + user_id="internal-user-123", + user_role=LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, + ) + + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch( + "litellm.proxy.management_endpoints.tag_management_endpoints.get_daily_activity", + new_callable=AsyncMock, + ) as mock_get_daily_activity, + ): + mock_db = Mock() + mock_prisma.db = mock_db + + owned_key_record = Mock() + owned_key_record.token = "owned-key" + mock_db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[owned_key_record] + ) + mock_get_daily_activity.return_value = "daily-activity-response" + + result = await get_tag_daily_activity( + start_date="2025-01-01", + end_date="2025-01-31", + user_api_key_dict=mock_user_auth, + ) + + assert result == "daily-activity-response" + mock_get_daily_activity.assert_awaited_once() + assert mock_get_daily_activity.await_args.kwargs["api_key"] == ["owned-key"] + + +@pytest.mark.asyncio +async def test_internal_user_tag_daily_activity_rejects_unowned_api_key_filter(): + """ + If an internal user filters tag usage by an API key they do not own, the + endpoint should return an empty scoped filter instead of exposing that key's + tag spend. + """ + from unittest.mock import AsyncMock, Mock + + from litellm.proxy.management_endpoints.tag_management_endpoints import ( + get_tag_daily_activity, + ) + + mock_user_auth = UserAPIKeyAuth( + user_id="internal-user-123", + user_role=LitellmUserRoles.INTERNAL_USER, + ) + + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch( + "litellm.proxy.management_endpoints.tag_management_endpoints.get_daily_activity", + new_callable=AsyncMock, + ) as mock_get_daily_activity, + ): + mock_db = Mock() + mock_prisma.db = mock_db + + owned_key_record = Mock() + owned_key_record.token = "owned-key" + mock_db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[owned_key_record] + ) + result = await get_tag_daily_activity( + start_date="2025-01-01", + end_date="2025-01-31", + api_key="unowned-key", + user_api_key_dict=mock_user_auth, + ) + + assert result.results == [] + assert result.metadata.total_spend == 0 + assert result.metadata.total_api_requests == 0 + mock_get_daily_activity.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_internal_user_tag_daily_activity_scopes_to_current_key_without_user_id(): + """ + If an internal-user token has no user_id, it should still scope tag usage to + the current request key instead of falling back to proxy-wide tag spend. + """ + from unittest.mock import AsyncMock, Mock + + from litellm.proxy.management_endpoints.tag_management_endpoints import ( + get_tag_daily_activity, + ) + + mock_user_auth = UserAPIKeyAuth( + api_key="current-owned-key", + user_id=None, + user_role=LitellmUserRoles.INTERNAL_USER, + ) + + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch( + "litellm.proxy.management_endpoints.tag_management_endpoints.get_daily_activity", + new_callable=AsyncMock, + ) as mock_get_daily_activity, + ): + mock_db = Mock() + mock_prisma.db = mock_db + mock_db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + mock_get_daily_activity.return_value = "daily-activity-response" + + result = await get_tag_daily_activity( + start_date="2025-01-01", + end_date="2025-01-31", + user_api_key_dict=mock_user_auth, + ) + + assert result == "daily-activity-response" + mock_db.litellm_verificationtoken.find_many.assert_not_awaited() + mock_get_daily_activity.assert_awaited_once() + assert mock_get_daily_activity.await_args.kwargs["api_key"] == [ + "current-owned-key" + ] + + +@pytest.mark.asyncio +async def test_internal_user_tag_daily_activity_without_any_scoped_keys_returns_empty(): + """ + If an internal-user token has neither user_id nor api_key, the endpoint must + return an empty response instead of dropping the API key filter. + """ + from unittest.mock import AsyncMock, Mock + + from litellm.proxy.management_endpoints.tag_management_endpoints import ( + get_tag_daily_activity, + ) + + mock_user_auth = UserAPIKeyAuth( + user_id=None, + user_role=LitellmUserRoles.INTERNAL_USER, + ) + + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch( + "litellm.proxy.management_endpoints.tag_management_endpoints.get_daily_activity", + new_callable=AsyncMock, + ) as mock_get_daily_activity, + ): + mock_db = Mock() + mock_prisma.db = mock_db + mock_db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + + result = await get_tag_daily_activity( + start_date="2025-01-01", + end_date="2025-01-31", + user_api_key_dict=mock_user_auth, + ) + + assert result.results == [] + assert result.metadata.total_spend == 0 + assert result.metadata.total_api_requests == 0 + mock_db.litellm_verificationtoken.find_many.assert_not_awaited() + mock_get_daily_activity.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_get_tag_daily_activity_requires_database_connection(): + """ + Tag daily activity should fail with the same explicit DB error used by other + tag endpoints instead of raising an AttributeError during scope resolution. + """ + from litellm.proxy.management_endpoints.tag_management_endpoints import ( + get_tag_daily_activity, + ) + + mock_user_auth = UserAPIKeyAuth( + user_id="internal-user-123", + user_role=LitellmUserRoles.INTERNAL_USER, + ) + + with patch("litellm.proxy.proxy_server.prisma_client", None): + with pytest.raises(HTTPException) as exc_info: + await get_tag_daily_activity( + start_date="2025-01-01", + end_date="2025-01-31", + user_api_key_dict=mock_user_auth, + ) + + assert exc_info.value.status_code == 500 + assert exc_info.value.detail == "Database not connected" + + +@pytest.mark.asyncio +async def test_list_tags_without_date_range_omits_date_filter(): + """When no date range is passed, the WHERE clause must not carry a date key.""" + from unittest.mock import AsyncMock, Mock + + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + mock_user_auth = UserAPIKeyAuth( + user_id="test-user-123", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth + + try: + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_db = Mock() + mock_prisma.db = mock_db + mock_db.litellm_tagtable.find_many = AsyncMock(return_value=[]) + group_by_mock = AsyncMock(return_value=[]) + mock_db.litellm_dailytagspend.group_by = group_by_mock + + headers = {"Authorization": "Bearer sk-1234"} + response = client.get("/tag/list", headers=headers) + + assert response.status_code == 200 + group_by_mock.assert_awaited_once() + where = group_by_mock.await_args.kwargs["where"] + assert "date" not in where + + finally: + app.dependency_overrides.clear() + + +@pytest.mark.parametrize( + "query, expected_detail_fragment", + [ + ("?start_date=2026-04-01", "must be provided together"), + ("?end_date=2026-04-29", "must be provided together"), + ("?start_date=2026-04-29&end_date=2026-04-01", "on or before end_date"), + ("?start_date=not-a-date&end_date=2026-04-29", "YYYY-MM-DD"), + ], +) +@pytest.mark.asyncio +async def test_list_tags_rejects_invalid_date_range(query, expected_detail_fragment): + from unittest.mock import AsyncMock, Mock + + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + mock_user_auth = UserAPIKeyAuth( + user_id="test-user-123", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth + + try: + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_db = Mock() + mock_prisma.db = mock_db + mock_db.litellm_tagtable.find_many = AsyncMock(return_value=[]) + mock_db.litellm_dailytagspend.group_by = AsyncMock(return_value=[]) + + headers = {"Authorization": "Bearer sk-1234"} + response = client.get(f"/tag/list{query}", headers=headers) + + assert response.status_code == 400 + assert expected_detail_fragment in response.json()["detail"] + + finally: + app.dependency_overrides.clear() + + @pytest.mark.asyncio async def test_get_deployments_by_model_id(): """ diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 92611431a15..b803dfb709a 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -21,6 +21,7 @@ from litellm.proxy.litellm_pre_call_utils import ( _get_enforced_params, _get_metadata_variable_name, _resolve_credential_from_model_config, + _resolve_provider_from_deployment, _update_model_if_key_alias_exists, add_guardrails_from_policy_engine, add_litellm_data_to_request, @@ -4043,3 +4044,174 @@ def test_get_guardrail_from_metadata_reads_litellm_metadata_when_no_metadata(): assert result == [ "my-guardrail" ], f"Expected guardrails from litellm_metadata fallback, got: {result}" + + +# ============================================================================ +# Tests for #27516: provider hint resolution from deployment when the +# user-facing model name has no provider prefix. +# ============================================================================ + + +def test_resolve_provider_from_deployment_uses_litellm_params_model(): + """When custom_llm_provider is unset, fall back to the prefix of model.""" + router = MagicMock() + deployment = MagicMock() + deployment.litellm_params.model = "bedrock/us.anthropic.claude-sonnet-4-6" + deployment.litellm_params.custom_llm_provider = None + router.get_deployment_by_model_group_name.return_value = deployment + + assert ( + _resolve_provider_from_deployment(router, "claude-sonnet-4.6") == "bedrock" + ) + + +def test_resolve_provider_from_deployment_prefers_custom_llm_provider(): + """Explicit custom_llm_provider on the deployment wins over model prefix.""" + router = MagicMock() + deployment = MagicMock() + deployment.litellm_params.model = "us.anthropic.claude-sonnet-4-6" + deployment.litellm_params.custom_llm_provider = "bedrock" + router.get_deployment_by_model_group_name.return_value = deployment + + assert ( + _resolve_provider_from_deployment(router, "claude-sonnet-4.6") == "bedrock" + ) + + +def test_resolve_provider_from_deployment_no_match(): + """No deployment for the model group -> None.""" + router = MagicMock() + router.get_deployment_by_model_group_name.return_value = None + assert _resolve_provider_from_deployment(router, "unknown-model") is None + + +def test_resolve_provider_from_deployment_router_raises(): + """Router exceptions must not propagate — fall back to None.""" + router = MagicMock() + router.get_deployment_by_model_group_name.side_effect = RuntimeError("boom") + assert _resolve_provider_from_deployment(router, "claude-sonnet-4.6") is None + + +def test_resolve_provider_from_deployment_falls_back_to_pre_alias(): + """If post-alias lookup fails, the pre-alias name is also tried.""" + router = MagicMock() + deployment = MagicMock() + deployment.litellm_params.model = "bedrock/anthropic.claude-sonnet-4-6" + deployment.litellm_params.custom_llm_provider = None + + def lookup(model_group_name): + if model_group_name == "pre-alias-name": + return deployment + return None + + router.get_deployment_by_model_group_name.side_effect = lookup + + result = _resolve_provider_from_deployment( + router, "post-alias-name", pre_alias_model_name="pre-alias-name" + ) + assert result == "bedrock" + + +def test_apply_overrides_multi_provider_default_picks_correct_provider( + setup_test_credentials, +): + """ + Regression for #27516: when defaultconfig has multiple providers and the + request model has no '/' prefix, the deployment's custom_llm_provider must + drive provider matching instead of falling through to dict insertion order. + """ + litellm.credential_list.append( + CredentialItem( + credential_name="bedrock-team-1", + credential_info={}, + credential_values={"api_key": "ABSK-bedrock-key-for-team-1"}, + ) + ) + litellm.credential_list.append( + CredentialItem( + credential_name="gemini-team-1", + credential_info={}, + credential_values={"api_key": "gemini-key-for-team-1"}, + ) + ) + + data = {"model": "claude-sonnet-4.6"} + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + team_metadata={ + "model_config": { + "defaultconfig": { + # gemini comes first in insertion order — the bug picked it. + "gemini": {"litellm_credentials": "gemini-team-1"}, + "bedrock": {"litellm_credentials": "bedrock-team-1"}, + } + } + }, + ) + + router = MagicMock() + deployment = MagicMock() + deployment.litellm_params.model = "us.anthropic.claude-sonnet-4-6" + deployment.litellm_params.custom_llm_provider = "bedrock" + router.get_deployment_by_model_group_name.return_value = deployment + + _apply_credential_overrides_from_model_config( + data=data, + user_api_key_dict=user_api_key_dict, + llm_router=router, + ) + assert data["api_key"] == "ABSK-bedrock-key-for-team-1" + + +def test_apply_overrides_no_router_keeps_legacy_behaviour(setup_test_credentials): + """ + Without a router, the function still works for the single-provider case + (the historical behaviour). Multi-provider configs with no '/' prefix + keep the legacy first-entry behaviour because there is no way to + disambiguate — this preserves backwards compatibility. + """ + data = {"model": "gpt-4"} + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + team_metadata={ + "model_config": { + "defaultconfig": { + "azure": {"litellm_credentials": "hotel-azure-eastus"} + } + } + }, + ) + _apply_credential_overrides_from_model_config( + data=data, user_api_key_dict=user_api_key_dict, llm_router=None + ) + assert data["api_base"] == "https://hotel-eastus.openai.azure.com/" + assert data["api_key"] == "key-hotel-eastus" + + +def test_apply_overrides_provider_prefix_in_model_skips_router_lookup( + setup_test_credentials, +): + """ + When the request model already has a 'provider/...' prefix, the router + lookup must be skipped — the explicit prefix is authoritative. + """ + data = {"model": "azure/gpt-4"} + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + team_metadata={ + "model_config": { + "defaultconfig": { + "azure": {"litellm_credentials": "hotel-azure-eastus"}, + "bedrock": {"litellm_credentials": "hotel-rec-azure"}, + } + } + }, + ) + + router = MagicMock() + _apply_credential_overrides_from_model_config( + data=data, user_api_key_dict=user_api_key_dict, llm_router=router + ) + assert data["api_base"] == "https://hotel-eastus.openai.azure.com/" + assert data["api_key"] == "key-hotel-eastus" + router.get_deployment_by_model_group_name.assert_not_called() diff --git a/tests/test_litellm/proxy/test_openapi_schema_validation.py b/tests/test_litellm/proxy/test_openapi_schema_validation.py index 68d537b2593..b44edc8a3bc 100644 --- a/tests/test_litellm/proxy/test_openapi_schema_validation.py +++ b/tests/test_litellm/proxy/test_openapi_schema_validation.py @@ -140,3 +140,110 @@ class TestCredentialEndpointsOpenAPISchema: assert ( "credential_name" in sig.parameters ), "get_credential_by_name must have a credential_name parameter" + + +class TestWebSocketStubInjection: + """ + Regression test for the v1.82.3 bug where adding a WebSocket route on a path + that already had an HTTP route silently dropped the HTTP operation from the + OpenAPI schema. + + Related case: 2026-05-05-madhu-swagger-responses-missing + """ + + def _make_fake_ws_route(self, path: str, name: str = "fake_ws"): + """Minimal stand-in for fastapi.routing.APIWebSocketRoute for the helper's purposes.""" + from types import SimpleNamespace + + return SimpleNamespace(path=path, name=name, dependant=None) + + def test_websocket_stub_does_not_clobber_existing_post(self): + """ + When a WebSocket route shares its path with an existing POST operation, + the POST must survive — the WebSocket stub is added alongside, not on top. + """ + from litellm.proxy.proxy_server import ( + _inject_websocket_stubs_into_openapi_schema, + ) + + schema = { + "paths": { + "/v1/responses": { + "post": {"summary": "responses_api", "operationId": "responses_api"} + } + } + } + ws_routes = [self._make_fake_ws_route("/v1/responses", name="responses_ws")] + + result = _inject_websocket_stubs_into_openapi_schema(schema, ws_routes) + + assert ( + "post" in result["paths"]["/v1/responses"] + ), "POST operation must be preserved when a WebSocket route shares the path" + assert ( + result["paths"]["/v1/responses"]["post"]["operationId"] == "responses_api" + ) + assert ( + "get" in result["paths"]["/v1/responses"] + ), "WebSocket stub should also be added under 'get'" + assert result["paths"]["/v1/responses"]["get"]["tags"] == ["WebSocket"] + + def test_websocket_stub_added_when_path_is_new(self): + """ + When a WebSocket route's path is not already in the schema, the stub + creates a fresh entry — preserving the original behavior for WebSocket-only + paths. + """ + from litellm.proxy.proxy_server import ( + _inject_websocket_stubs_into_openapi_schema, + ) + + schema = {"paths": {}} + ws_routes = [self._make_fake_ws_route("/ws_only", name="ws_only")] + + result = _inject_websocket_stubs_into_openapi_schema(schema, ws_routes) + + assert "/ws_only" in result["paths"] + assert "get" in result["paths"]["/ws_only"] + assert result["paths"]["/ws_only"]["get"]["tags"] == ["WebSocket"] + + def test_websocket_stub_skipped_when_existing_get(self): + """ + If a real GET is already documented on the path, the WebSocket stub is + skipped — a real operation always wins over the synthetic stub. This + closes the same trap for future GET-vs-WebSocket collisions. + """ + from litellm.proxy.proxy_server import ( + _inject_websocket_stubs_into_openapi_schema, + ) + + schema = { + "paths": { + "/health": { + "get": {"summary": "health_check", "operationId": "real_get"} + } + } + } + ws_routes = [self._make_fake_ws_route("/health", name="health_ws")] + + result = _inject_websocket_stubs_into_openapi_schema(schema, ws_routes) + + assert ( + result["paths"]["/health"]["get"]["operationId"] == "real_get" + ), "Real GET must take precedence over WebSocket stub" + + def test_responses_post_routes_registered_on_router(self): + """ + Sanity check: the three POST routes for the responses API are still wired + on the responses router. Guards against accidental removal at the source. + """ + from litellm.proxy.response_api_endpoints.endpoints import router + + post_paths = { + route.path + for route in router.routes + if hasattr(route, "methods") + and "POST" in (route.methods or set()) + and route.path in {"/v1/responses", "/responses", "/openai/v1/responses"} + } + assert post_paths == {"/v1/responses", "/responses", "/openai/v1/responses"} diff --git a/tests/test_litellm/proxy/test_proxy_cli.py b/tests/test_litellm/proxy/test_proxy_cli.py index 59b43330a25..327200a6a95 100644 --- a/tests/test_litellm/proxy/test_proxy_cli.py +++ b/tests/test_litellm/proxy/test_proxy_cli.py @@ -1,6 +1,6 @@ import os import sys -from unittest.mock import MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -387,6 +387,102 @@ class TestProxyInitializationHelpers: ), f"exit_code={result.exit_code}, output={result.output}" mock_uvicorn_run.assert_called_once() + @pytest.mark.parametrize( + "timeout_config,expected_timeout", + [ + ({"database_connection_timeout": 30}, 30), + ({"database_connection_pool_timeout": 45}, 45), + ( + { + "database_connection_timeout": 30, + "database_connection_pool_timeout": 45, + }, + 30, + ), + ], + ) + @patch("subprocess.run") + @patch("atexit.register") + @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database") + @patch( + "litellm.proxy.db.prisma_client.should_update_prisma_schema", return_value=False + ) + def test_db_timeout_settings_are_forwarded_to_pool_timeout( + self, + mock_should_update, + mock_setup_db, + mock_atexit_register, + mock_subprocess_run, + timeout_config, + expected_timeout, + ): + from click.testing import CliRunner + + from litellm.proxy.proxy_cli import run_server + + runner = CliRunner() + mock_subprocess_run.return_value = MagicMock(returncode=0) + + mock_proxy_module = MagicMock( + app=MagicMock(), + ProxyConfig=MagicMock(), + KeyManagementSettings=MagicMock(), + save_worker_config=MagicMock(), + ) + mock_proxy_module.ProxyConfig.return_value.get_config = AsyncMock( + return_value={ + "general_settings": { + "database_url": "postgresql://test:test@localhost:5432/test", + "database_connection_pool_limit": 5, + **timeout_config, + } + } + ) + + clean_env = { + k: v + for k, v in os.environ.items() + if k not in ("DATABASE_URL", "DIRECT_URL") + } + + with ( + patch.dict(os.environ, clean_env, clear=True), + patch.dict( + "sys.modules", + { + "proxy_server": mock_proxy_module, + "litellm.proxy.proxy_server": mock_proxy_module, + }, + ), + patch( + "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + ) as mock_get_args, + patch( + "litellm.proxy.proxy_cli.append_query_params", + side_effect=lambda url, params: ( + f"{url}?connection_limit={params['connection_limit']}&pool_timeout={params['pool_timeout']}" + ), + ) as mock_append_query_params, + ): + mock_get_args.return_value = { + "app": "litellm.proxy.proxy_server:app", + "host": "localhost", + "port": 8000, + } + + result = runner.invoke( + run_server, + ["--local", "--config", "test-config.yaml", "--skip_server_startup"], + ) + + assert ( + result.exit_code == 0 + ), f"exit_code={result.exit_code}, output={result.output}" + mock_append_query_params.assert_called() + appended_params = mock_append_query_params.call_args.args[1] + assert appended_params["connection_limit"] == 5 + assert appended_params["pool_timeout"] == expected_timeout + @patch("uvicorn.run") @patch("atexit.register") @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database") diff --git a/tests/test_litellm/proxy/test_shared_health_check.py b/tests/test_litellm/proxy/test_shared_health_check.py index 04099f2634b..1530d336085 100644 --- a/tests/test_litellm/proxy/test_shared_health_check.py +++ b/tests/test_litellm/proxy/test_shared_health_check.py @@ -322,13 +322,13 @@ class TestSharedHealthCheckManager: async def test_perform_shared_health_check_lock_failed_then_cache( self, shared_health_manager, mock_redis_cache ): - """Test performing shared health check when lock fails but cache becomes available""" + """Test performing shared health check when lock fails but cache becomes available during polling""" # First call: no cache, lock fails - # Second call: cache available + # Polling finds cache on first iteration mock_redis_cache.async_get_cache.side_effect = [ - None, # No cache initially + None, # No cache initially (get_cached_health_check_results) json.dumps( - { # Cache available after waiting + { # Cache available on first poll iteration "healthy_endpoints": [{"model": "cached-model"}], "unhealthy_endpoints": [], "healthy_count": 1, @@ -350,18 +350,68 @@ class TestSharedHealthCheckManager: ) ) - # Should wait and then get cached results - mock_sleep.assert_called_once_with(2) + # Should poll once (5s interval) and find cached results + mock_sleep.assert_called_once_with(5) assert healthy == [{"model": "cached-model"}] assert unhealthy == [] @pytest.mark.asyncio - async def test_perform_shared_health_check_fallback( + async def test_perform_shared_health_check_fallback(self, mock_redis_cache): + """Test performing shared health check with fallback to local health check""" + # Use short lock_ttl so the polling loop only runs 2 iterations + manager = SharedHealthCheckManager( + redis_cache=mock_redis_cache, + health_check_ttl=300, + lock_ttl=10, + ) + + # No cache ever, lock always held by another pod + mock_redis_cache.async_get_cache.side_effect = [ + None, # Initial cache check + None, # Iteration 1: cache check + "other_pod", # Iteration 1: lock check (still held) + None, # Iteration 2: cache check + "other_pod", # Iteration 2: lock check (still held) + ] + mock_redis_cache.async_set_cache.return_value = False # Lock acquisition fails + + model_list = [ + {"model_name": "test-model", "litellm_params": {"model": "test-model"}} + ] + expected_healthy = [{"model": "test-model", "status": "healthy"}] + expected_unhealthy = [] + + with ( + patch("asyncio.sleep") as mock_sleep, + patch( + "litellm.proxy.health_check_utils.shared_health_check_manager.perform_health_check" + ) as mock_perform, + ): + mock_perform.return_value = (expected_healthy, expected_unhealthy, {}) + + healthy, unhealthy, _ = await manager.perform_shared_health_check( + model_list, details=True + ) + + # Should poll twice (5s * 2 = 10s >= lock_ttl) then fall back + assert mock_sleep.call_count == 2 + mock_sleep.assert_called_with(5) + mock_perform.assert_called_once_with( + model_list=model_list, details=True, max_concurrency=None + ) + assert healthy == expected_healthy + assert unhealthy == expected_unhealthy + + @pytest.mark.asyncio + async def test_perform_shared_health_check_early_exit_orphaned_lock( self, shared_health_manager, mock_redis_cache ): - """Test performing shared health check with fallback to local health check""" - # No cache, lock fails, no cache after waiting - mock_redis_cache.async_get_cache.return_value = None + """Test that polling exits early when the lock disappears without a cache write (crash recovery)""" + mock_redis_cache.async_get_cache.side_effect = [ + None, # Initial cache check + None, # Iteration 1: cache check (still no cache) + None, # Iteration 1: lock check -> lock gone (holder crashed) + ] mock_redis_cache.async_set_cache.return_value = False # Lock acquisition fails model_list = [ @@ -384,8 +434,77 @@ class TestSharedHealthCheckManager: ) ) - # Should fall back to local health check - mock_sleep.assert_called_once_with(2) + # Should detect orphaned lock after 1 iteration and fall back immediately + mock_sleep.assert_called_once_with(5) + mock_perform.assert_called_once_with( + model_list=model_list, details=True, max_concurrency=None + ) + assert healthy == expected_healthy + assert unhealthy == expected_unhealthy + + @pytest.mark.asyncio + async def test_perform_shared_health_check_redis_error_during_polling( + self, shared_health_manager, mock_redis_cache + ): + """Test that a transient Redis error during lock polling doesn't crash the loop""" + cached_data = json.dumps( + { + "healthy_endpoints": [{"model": "cached-model"}], + "unhealthy_endpoints": [], + "healthy_count": 1, + "unhealthy_count": 0, + "timestamp": time.time() - 100, + } + ) + mock_redis_cache.async_get_cache.side_effect = [ + None, # Initial cache check + None, # Iteration 1: cache check + Exception("Redis connection lost"), # Iteration 1: lock check errors + cached_data, # Iteration 2: cache check -> found! + ] + mock_redis_cache.async_set_cache.return_value = False # Lock acquisition fails + + model_list = [ + {"model_name": "test-model", "litellm_params": {"model": "test-model"}} + ] + + with patch("asyncio.sleep") as mock_sleep: + healthy, unhealthy, _ = ( + await shared_health_manager.perform_shared_health_check( + model_list, details=True + ) + ) + + # Should survive the Redis error and find cache on iteration 2 + assert mock_sleep.call_count == 2 + assert healthy == [{"model": "cached-model"}] + assert unhealthy == [] + + @pytest.mark.asyncio + async def test_perform_shared_health_check_no_redis_skips_polling(self): + """Test that polling is skipped entirely when redis_cache is None""" + manager = SharedHealthCheckManager(redis_cache=None) + + model_list = [ + {"model_name": "test-model", "litellm_params": {"model": "test-model"}} + ] + expected_healthy = [{"model": "test-model", "status": "healthy"}] + expected_unhealthy = [] + + with ( + patch("asyncio.sleep") as mock_sleep, + patch( + "litellm.proxy.health_check_utils.shared_health_check_manager.perform_health_check" + ) as mock_perform, + ): + mock_perform.return_value = (expected_healthy, expected_unhealthy, {}) + + healthy, unhealthy, _ = await manager.perform_shared_health_check( + model_list, details=True + ) + + # Should NOT sleep at all — falls back to local health check immediately + mock_sleep.assert_not_called() mock_perform.assert_called_once_with( model_list=model_list, details=True, max_concurrency=None ) diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py index 4358d0dc193..76336a91fc3 100644 --- a/tests/test_litellm/test_main.py +++ b/tests/test_litellm/test_main.py @@ -757,6 +757,60 @@ def test_responses_api_bridge_check_gpt_5_4_tools_without_reasoning_stays_chat() assert model_info.get("mode") != "responses" +def test_responses_api_bridge_check_gpt_5_4_reasoning_summary_without_tools_routes_to_responses(): + """gpt-5.4+ with reasoning_effort + reasoningSummary but no tools should bridge (AI SDK).""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.4", + custom_llm_provider="openai", + tools=None, + reasoning_effort="medium", + reasoning_summary="auto", + ) + + assert model == "gpt-5.4" + assert model_info.get("mode") == "responses" + + +def test_responses_api_bridge_check_gpt_5_reasoning_summary_routes_to_responses(): + """Bare ``gpt-5`` with reasoning_effort + reasoningSummary should bridge (not 5.4+).""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5", + custom_llm_provider="openai", + tools=None, + reasoning_effort="medium", + reasoning_summary="auto", + ) + + assert model == "gpt-5" + assert model_info.get("mode") == "responses" + + +def test_responses_api_bridge_check_gpt_5_tools_without_summary_stays_chat(): + """gpt-5 with tools + reasoning_effort but no summary should stay on chat.""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5", + custom_llm_provider="openai", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort="medium", + reasoning_summary=None, + ) + + assert model == "gpt-5" + assert model_info.get("mode") != "responses" + + @patch("litellm.completion_extras.responses_api_bridge.completion") def test_gpt_5_4_responses_bridge_preserves_reasoning_summary_dict( mock_responses_completion, @@ -794,6 +848,93 @@ def test_gpt_5_4_responses_bridge_preserves_reasoning_summary_dict( } +@patch("litellm.completion_extras.responses_api_bridge.completion") +def test_gpt_5_4_responses_bridge_merges_reasoning_summary_kwarg_without_tools( + mock_responses_completion, +): + """reasoningSummary without tools should route and merge into reasoning_effort dict.""" + mock_responses_completion.return_value = MagicMock() + + import litellm + + litellm.completion( + model="gpt-5.4", + messages=[{"role": "user", "content": "ok"}], + reasoning_effort="medium", + reasoningSummary="auto", + api_key="fake-key", + ) + + assert mock_responses_completion.called is True + optional_params = mock_responses_completion.call_args.kwargs["optional_params"] + assert optional_params["reasoning_effort"] == { + "effort": "medium", + "summary": "auto", + } + assert "reasoningSummary" not in optional_params + assert "reasoning_summary" not in optional_params + + +@patch("litellm.completion_extras.responses_api_bridge.completion") +def test_responses_bridge_preserves_reasoning_summary_without_effort( + mock_responses_completion, +): + """Reasoning summary should survive responses routing even without effort.""" + mock_responses_completion.return_value = MagicMock() + + import litellm + + with patch.object(litellm, "route_all_chat_openai_to_responses", True): + litellm.completion( + model="gpt-4o", + messages=[{"role": "user", "content": "ok"}], + reasoningSummary="auto", + api_key="fake-key", + ) + + assert mock_responses_completion.called is True + optional_params = mock_responses_completion.call_args.kwargs["optional_params"] + assert optional_params["reasoning_effort"] == {"summary": "auto"} + assert "reasoningSummary" not in optional_params + assert "reasoning_summary" not in optional_params + + +@patch("litellm.completion_extras.responses_api_bridge.completion") +def test_gpt_5_responses_bridge_tools_and_reasoning_summary( + mock_responses_completion, +): + """Bare gpt-5 with tools + reasoningSummary should bridge (OpenCode-style).""" + mock_responses_completion.return_value = MagicMock() + + import litellm + + litellm.completion( + model="gpt-5", + messages=[{"role": "user", "content": "ok"}], + tools=[ + { + "type": "function", + "function": { + "name": "apply_patch", + "parameters": {"type": "object", "properties": {}}, + }, + } + ], + tool_choice="auto", + reasoning_effort="medium", + reasoningSummary="auto", + stream=True, + api_key="fake-key", + ) + + assert mock_responses_completion.called is True + optional_params = mock_responses_completion.call_args.kwargs["optional_params"] + assert optional_params.get("reasoning_effort") == { + "effort": "medium", + "summary": "auto", + } + + def test_responses_api_bridge_check_handles_exception(): """Test that responses_api_bridge_check handles exceptions and still processes responses/ models.""" from litellm.main import responses_api_bridge_check diff --git a/tests/test_litellm/test_main_module_header.py b/tests/test_litellm/test_main_module_header.py new file mode 100644 index 00000000000..a16e14e8c32 --- /dev/null +++ b/tests/test_litellm/test_main_module_header.py @@ -0,0 +1,13 @@ +from pathlib import Path + + +def test_main_py_starts_with_brief_file_description(): + repo_root = Path(__file__).resolve().parents[2] + main_py = repo_root / "litellm" / "main.py" + + first_two_lines = main_py.read_text(encoding="utf-8").splitlines()[:2] + + assert any( + "LiteLLM main module" in line and "entrypoints" in line + for line in first_two_lines + ) diff --git a/tests/test_litellm/test_retrieve_batch_bedrock_dispatch.py b/tests/test_litellm/test_retrieve_batch_bedrock_dispatch.py new file mode 100644 index 00000000000..9df18a9f0f0 --- /dev/null +++ b/tests/test_litellm/test_retrieve_batch_bedrock_dispatch.py @@ -0,0 +1,162 @@ +"""Cover the Bedrock-ARN dispatch in ``litellm.batches.main.retrieve_batch``. + +The dispatch picks one of two Bedrock handlers depending on the ARN +family in ``batch_id``: + +* ``:async-invoke/`` -> ``_handle_async_invoke_status`` (data plane) +* ``:model-invocation-job/`` -> ``_handle_model_invocation_job_status`` + (control plane, added in this PR) + +Anything else falls through to the generic ``provider_config`` retrieve +flow. We mock the two handlers so the tests don't hit AWS — the focus +here is purely the dispatch logic that lives in ``main.py``. +""" + +from __future__ import annotations + +import os +import sys +from unittest.mock import MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../..")) + +import litellm # noqa: E402 + +ASYNC_INVOKE_ARN = "arn:aws:bedrock:us-west-2:123456789012:async-invoke/abc123def456" +MIJ_ARN = "arn:aws:bedrock:us-west-2:123456789012:model-invocation-job/abc1234567" + + +@pytest.fixture +def mock_handlers(): + """Patch both Bedrock retrieve handlers and yield the mocks. + + We patch at the import site (litellm.batches.main) rather than the + definition site so the ``BedrockBatchesHandler`` reference inside + ``retrieve_batch`` resolves to our mocks. + """ + fake_batch = MagicMock(name="LiteLLMBatch") + with ( + patch( + "litellm.batches.main.BedrockBatchesHandler._handle_async_invoke_status", + return_value=fake_batch, + ) as async_invoke, + patch( + "litellm.batches.main.BedrockBatchesHandler._handle_model_invocation_job_status", + return_value=fake_batch, + ) as mij, + ): + yield async_invoke, mij, fake_batch + + +def test_async_invoke_arn_routes_to_async_invoke_handler(mock_handlers): + """``:async-invoke/`` ARNs go to the data-plane handler.""" + async_invoke, mij, fake_batch = mock_handlers + + result = litellm.retrieve_batch( + batch_id=ASYNC_INVOKE_ARN, + custom_llm_provider="bedrock", + aws_region_name="us-west-2", + ) + + assert result is fake_batch + async_invoke.assert_called_once() + mij.assert_not_called() + call_kwargs = async_invoke.call_args.kwargs + assert call_kwargs["batch_id"] == ASYNC_INVOKE_ARN + assert call_kwargs["aws_region_name"] == "us-west-2" + # Region must be stripped from the forwarded kwargs to avoid TypeError + # (it's already an explicit positional/keyword arg). + assert "aws_region_name" not in { + k + for k in call_kwargs + if k not in {"batch_id", "aws_region_name", "logging_obj"} + } + + +def test_async_invoke_arn_falls_back_to_default_region_when_unset(mock_handlers): + """If no ``aws_region_name`` is passed, the data-plane handler defaults + to ``us-east-1`` (preserving prior behavior on this branch).""" + async_invoke, _mij, _ = mock_handlers + + litellm.retrieve_batch( + batch_id=ASYNC_INVOKE_ARN, + custom_llm_provider="bedrock", + ) + + async_invoke.assert_called_once() + assert async_invoke.call_args.kwargs["aws_region_name"] == "us-east-1" + + +def test_model_invocation_job_arn_routes_to_mij_handler(mock_handlers): + """``:model-invocation-job/`` ARNs go to the new control-plane handler.""" + _async_invoke, mij, fake_batch = mock_handlers + + result = litellm.retrieve_batch( + batch_id=MIJ_ARN, + custom_llm_provider="bedrock", + aws_region_name="us-west-2", + ) + + assert result is fake_batch + mij.assert_called_once() + _async_invoke.assert_not_called() + call_kwargs = mij.call_args.kwargs + assert call_kwargs["batch_id"] == MIJ_ARN + assert call_kwargs["aws_region_name"] == "us-west-2" + + +def test_model_invocation_job_arn_with_no_region_passes_none(mock_handlers): + """The MIJ handler is responsible for sniffing region from the ARN + when none is explicitly provided. Dispatch must forward ``None`` + rather than substituting a default — otherwise per-region jobs in + other AWS regions would silently route to ``us-east-1``.""" + _async_invoke, mij, _ = mock_handlers + + litellm.retrieve_batch( + batch_id=MIJ_ARN, + custom_llm_provider="bedrock", + ) + + mij.assert_called_once() + assert mij.call_args.kwargs["aws_region_name"] is None + + +def test_unrelated_bedrock_arn_falls_through_to_provider_config(mock_handlers): + """Bedrock ARNs that aren't async-invoke or model-invocation-job + must NOT hit either special handler — they should fall through to + the existing generic provider_config path. We don't fully exercise + that path here (it requires a real provider config); we just assert + neither special handler is invoked.""" + async_invoke, mij, _ = mock_handlers + + # Use a plausible-but-unsupported Bedrock ARN family. + unrelated_arn = "arn:aws:bedrock:us-west-2:123456789012:provisioned-model/xyz" + + with pytest.raises(Exception): + # Will raise because no provider_config exists for this path — + # that's fine, we just need to assert neither bedrock handler ran + # before the failure. + litellm.retrieve_batch( + batch_id=unrelated_arn, + custom_llm_provider="bedrock", + ) + + async_invoke.assert_not_called() + mij.assert_not_called() + + +def test_non_bedrock_id_skips_bedrock_dispatch_entirely(mock_handlers): + """Plain (non-ARN) batch ids must not even enter the Bedrock dispatch + block — they belong to other providers' retrieve flows.""" + async_invoke, mij, _ = mock_handlers + + with pytest.raises(Exception): + litellm.retrieve_batch( + batch_id="batch_abc123", + custom_llm_provider="openai", + ) + + async_invoke.assert_not_called() + mij.assert_not_called() diff --git a/tests/test_litellm/test_router_model_cost_isolation.py b/tests/test_litellm/test_router_model_cost_isolation.py index 7a9d5acaa27..c3f93078557 100644 --- a/tests/test_litellm/test_router_model_cost_isolation.py +++ b/tests/test_litellm/test_router_model_cost_isolation.py @@ -18,6 +18,7 @@ sys.path.insert( import litellm from litellm import Router +from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo def test_should_not_pollute_shared_key_with_zero_cost_pricing(): @@ -266,3 +267,59 @@ def test_should_preserve_builtin_pricing_regardless_of_deployment_order(): f"Order should not matter. Expected {builtin_output_cost}, " f"got {info_std_2['output_cost_per_token']}" ) + + +def test_responses_prefix_stripped_alias_registered_for_model_list(): + """ + Register ``litellm.model_cost`` under the backend key with ``responses/`` and + under the stripped key (``responses_api_bridge_check`` removes that segment). + """ + uid = "responses-strip-alias-test-a1b2c3d4" + Router( + model_list=[ + { + "model_name": "azure-responses-strip-test", + "litellm_params": { + "model": "responses/gpt-strip-test-a1b2c3d4", + "custom_llm_provider": "azure", + "api_key": "fake-key-strip", + }, + "model_info": { + "id": uid, + "supports_native_streaming": True, + }, + } + ], + ) + assert "azure/responses/gpt-strip-test-a1b2c3d4" in litellm.model_cost + assert "azure/gpt-strip-test-a1b2c3d4" in litellm.model_cost + assert ( + litellm.model_cost["azure/gpt-strip-test-a1b2c3d4"].get( + "supports_native_streaming" + ) + is True + ) + + +def test_responses_prefix_stripped_alias_registered_for_add_deployment(): + """Dynamic ``add_deployment`` must mirror ``_create_deployment`` registration.""" + uid = "add-dep-responses-strip-e5f6a7b8" + router = Router(model_list=[]) + deployment = Deployment( + model_name="dyn-responses-strip", + litellm_params=LiteLLM_Params( + model="responses/gpt-add-strip-e5f6a7b8", + custom_llm_provider="azure", + api_key="fake-key-add", + ), + model_info=ModelInfo(id=uid, supports_native_streaming=True), + ) + router.add_deployment(deployment=deployment) + assert "azure/responses/gpt-add-strip-e5f6a7b8" in litellm.model_cost + assert "azure/gpt-add-strip-e5f6a7b8" in litellm.model_cost + assert ( + litellm.model_cost["azure/gpt-add-strip-e5f6a7b8"].get( + "supports_native_streaming" + ) + is True + ) diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 65305e1a81e..d07af922ea6 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -2817,6 +2817,128 @@ def test_generate_gcp_iam_access_token_import_error(): assert "pip install google-cloud-iam" in str(exc_info.value) +def test_generate_azure_ad_redis_token(): + """Test _generate_azure_ad_redis_token with mocked Azure credential.""" + from unittest.mock import Mock, patch + + expected_token = "azure-access-token-12345" + + mock_token = Mock() + mock_token.token = expected_token + + mock_credential = Mock() + mock_credential.get_token.return_value = mock_token + + mock_azure_identity = Mock() + mock_azure_identity.DefaultAzureCredential = Mock(return_value=mock_credential) + mock_azure_identity.ClientSecretCredential = Mock() + mock_azure_identity.ManagedIdentityCredential = Mock() + + with patch.dict( + "sys.modules", {"azure.identity": mock_azure_identity, "azure": Mock()} + ): + from litellm._redis import _generate_azure_ad_redis_token + + result = _generate_azure_ad_redis_token() + + assert result == expected_token + mock_credential.get_token.assert_called_once_with( + "https://redis.azure.com/.default" + ) + + +def test_generate_azure_ad_redis_token_service_principal(): + """Test _generate_azure_ad_redis_token with service principal credentials.""" + from unittest.mock import Mock, patch + + expected_token = "sp-access-token-67890" + + mock_token = Mock() + mock_token.token = expected_token + + mock_credential = Mock() + mock_credential.get_token.return_value = mock_token + + mock_client_secret_credential = Mock(return_value=mock_credential) + + mock_azure_identity = Mock() + mock_azure_identity.DefaultAzureCredential = Mock() + mock_azure_identity.ClientSecretCredential = mock_client_secret_credential + mock_azure_identity.ManagedIdentityCredential = Mock() + + with patch.dict( + "sys.modules", {"azure.identity": mock_azure_identity, "azure": Mock()} + ): + from litellm._redis import _generate_azure_ad_redis_token + + result = _generate_azure_ad_redis_token( + azure_client_id="test-client-id", + azure_tenant_id="test-tenant-id", + azure_client_secret="test-secret", + ) + + assert result == expected_token + mock_client_secret_credential.assert_called_once_with( + client_id="test-client-id", + tenant_id="test-tenant-id", + client_secret="test-secret", + ) + + +def test_generate_azure_ad_redis_token_import_error(): + """Test that _generate_azure_ad_redis_token raises ImportError when azure-identity is missing.""" + from unittest.mock import patch + from litellm._redis import _generate_azure_ad_redis_token + + with patch.dict("sys.modules", {"azure.identity": None}): + with pytest.raises(ImportError) as exc_info: + _generate_azure_ad_redis_token() + + assert "azure-identity is required" in str(exc_info.value) + + +def test_redis_client_logic_azure_ad_auth(): + """Test that _get_redis_client_logic sets up Azure AD auth when REDIS_AZURE_AD_TOKEN=true. + + Mocks ``azure.identity`` via ``sys.modules`` so the test does not require + the real ``azure-identity`` package to be installed in the CI environment. + """ + from unittest.mock import Mock, patch + + mock_credential = Mock() + mock_azure_identity = Mock() + mock_azure_identity.DefaultAzureCredential = Mock(return_value=mock_credential) + mock_azure_identity.ClientSecretCredential = Mock(return_value=mock_credential) + mock_azure_identity.ManagedIdentityCredential = Mock(return_value=mock_credential) + + with patch.dict( + "sys.modules", {"azure.identity": mock_azure_identity, "azure": Mock()} + ): + from litellm._redis import _get_redis_client_logic + + redis_kwargs = _get_redis_client_logic( + host="myredis.redis.cache.windows.net", + port="6380", + azure_redis_ad_token="true", + ssl=True, + ) + + assert "redis_connect_func" in redis_kwargs + # Marker for async paths to detect Azure AD auth + assert hasattr(redis_kwargs["redis_connect_func"], "_azure_redis_ad_token") + assert redis_kwargs["redis_connect_func"]._azure_redis_ad_token is True + # Live credential object (not raw secret) is exposed for async paths + assert hasattr(redis_kwargs["redis_connect_func"], "_azure_credential") + # Raw credentials must NOT be exposed on the function + assert not hasattr(redis_kwargs["redis_connect_func"], "_azure_client_secret") + assert not hasattr(redis_kwargs["redis_connect_func"], "_azure_client_id") + assert not hasattr(redis_kwargs["redis_connect_func"], "_azure_tenant_id") + + # Azure-specific kwargs should be removed from the dict passed to Redis + assert "azure_redis_ad_token" not in redis_kwargs + assert "azure_client_id" not in redis_kwargs + + if __name__ == "__main__": # Allow running this test file directly for debugging pytest.main([__file__, "-v"]) diff --git a/tests/test_litellm/test_utils_module_docstring.py b/tests/test_litellm/test_utils_module_docstring.py new file mode 100644 index 00000000000..ac99fb63fd4 --- /dev/null +++ b/tests/test_litellm/test_utils_module_docstring.py @@ -0,0 +1,11 @@ +import ast +from pathlib import Path + + +def test_utils_module_has_docstring(): + utils_path = Path(__file__).parents[2] / "litellm" / "utils.py" + module = ast.parse(utils_path.read_text()) + + assert ast.get_docstring(module) == ( + "Utility helpers for LiteLLM core request handling and provider support." + ) diff --git a/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.test.tsx b/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.test.tsx index 6da2f82a2f1..a57b7e2095a 100644 --- a/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.test.tsx +++ b/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.test.tsx @@ -535,6 +535,21 @@ describe("ModelSelect", () => { }); }); + it("should not render an empty optgroup when includeSpecialOptions is omitted", async () => { + renderWithProviders(); + + await waitFor(() => { + expect(screen.getByTestId("model-select")).toBeInTheDocument(); + }); + + const optgroups = document.querySelectorAll("optgroup"); + // Wildcard Options + Models — no blank leading group + expect(optgroups.length).toBe(2); + optgroups.forEach((g) => { + expect(g.getAttribute("label")).toBeTruthy(); + }); + }); + it("should render maxTagPlaceholder when many items are selected", async () => { // Create many models to trigger maxTagCount responsive behavior const manyModels: ProxyModel[] = Array.from({ length: 20 }, (_, i) => ({ diff --git a/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.tsx b/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.tsx index 74b2619f7f3..800fa86b165 100644 --- a/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.tsx +++ b/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.tsx @@ -141,36 +141,38 @@ export const ModelSelect = (props: ModelSelectProps) => { onChange={handleChange} style={style} options={[ - includeSpecialOptions - ? { - label: Special Options, - title: "Special Options", - options: [ - ...(shouldShowAllProxyModels - ? [ - { - label: All Proxy Models, - value: MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value, - disabled: - value.length > 0 && - value.some( - (v) => isSpecialOption(v) && v !== MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value, - ), - key: MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value, - }, - ] - : []), - { - label: No Default Models, - value: MODEL_SELECT_NO_DEFAULT_MODELS_SPECIAL_VALUE.value, - disabled: - value.length > 0 && - value.some((v) => isSpecialOption(v) && v !== MODEL_SELECT_NO_DEFAULT_MODELS_SPECIAL_VALUE.value), - key: MODEL_SELECT_NO_DEFAULT_MODELS_SPECIAL_VALUE.value, - }, - ], - } - : [], + ...(includeSpecialOptions + ? [ + { + label: Special Options, + title: "Special Options", + options: [ + ...(shouldShowAllProxyModels + ? [ + { + label: All Proxy Models, + value: MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value, + disabled: + value.length > 0 && + value.some( + (v) => isSpecialOption(v) && v !== MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value, + ), + key: MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value, + }, + ] + : []), + { + label: No Default Models, + value: MODEL_SELECT_NO_DEFAULT_MODELS_SPECIAL_VALUE.value, + disabled: + value.length > 0 && + value.some((v) => isSpecialOption(v) && v !== MODEL_SELECT_NO_DEFAULT_MODELS_SPECIAL_VALUE.value), + key: MODEL_SELECT_NO_DEFAULT_MODELS_SPECIAL_VALUE.value, + }, + ], + }, + ] + : []), ...(wildcard.length > 0 ? [ { diff --git a/ui/litellm-dashboard/src/components/UsagePage/components/UsagePageView.test.tsx b/ui/litellm-dashboard/src/components/UsagePage/components/UsagePageView.test.tsx index bbcddd572cd..b95f63368c7 100644 --- a/ui/litellm-dashboard/src/components/UsagePage/components/UsagePageView.test.tsx +++ b/ui/litellm-dashboard/src/components/UsagePage/components/UsagePageView.test.tsx @@ -57,7 +57,10 @@ vi.mock("./EndpointUsage/EndpointUsage", () => ({ vi.mock("./UsageViewSelect/UsageViewSelect", async () => { const React = await import("react"); - const UsageViewSelect = ({ value, onChange }: any) => { + const UsageViewSelect = ({ value, onChange, canViewTagUsage = false }: any) => { + const tagOption = canViewTagUsage + ? React.createElement("option", { value: "tag" }, "Tag Usage") + : null; return React.createElement( "select", { @@ -70,7 +73,7 @@ vi.mock("./UsageViewSelect/UsageViewSelect", async () => { React.createElement("option", { value: "team" }, "Team Usage"), React.createElement("option", { value: "organization" }, "Organization Usage"), React.createElement("option", { value: "customer" }, "Customer Usage"), - React.createElement("option", { value: "tag" }, "Tag Usage"), + tagOption, React.createElement("option", { value: "agent" }, "Agent Usage"), React.createElement("option", { value: "user-agent-activity" }, "User Agent Activity"), ); @@ -639,6 +642,29 @@ describe("UsagePage", () => { }); }); + it("should show tag usage selector option for internal users", async () => { + mockUseAuthorized.mockReturnValue({ + isLoading: false, + isAuthorized: true, + token: "mock-token", + accessToken: "test-token", + userId: "user-123", + userEmail: "test@example.com", + userRole: "internal_user", + premiumUser: true, + disabledPersonalKeyCreation: false, + showSSOBanner: false, + }); + + renderWithProviders(); + + await waitFor(() => { + expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalled(); + }); + + expect(screen.getByRole("option", { name: "Tag Usage" })).toBeInTheDocument(); + }); + it("should show organization usage banner and view for admins", async () => { renderWithProviders(); diff --git a/ui/litellm-dashboard/src/components/UsagePage/components/UsagePageView.tsx b/ui/litellm-dashboard/src/components/UsagePage/components/UsagePageView.tsx index 809f1d4e17b..50dd5e78f37 100644 --- a/ui/litellm-dashboard/src/components/UsagePage/components/UsagePageView.tsx +++ b/ui/litellm-dashboard/src/components/UsagePage/components/UsagePageView.tsx @@ -31,7 +31,7 @@ import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import { useCurrentUser } from "@/app/(dashboard)/hooks/users/useCurrentUser"; import { useInfiniteUsers } from "@/app/(dashboard)/hooks/users/useUsers"; import { formatNumberWithCommas } from "@/utils/dataUtils"; -import { all_admin_roles } from "../../../utils/roles"; +import { all_admin_roles, internalUserRoles } from "../../../utils/roles"; import { ActivityMetrics, processActivityData } from "../../activity_metrics"; import CloudZeroExportModal from "../../cloudzero_export_modal"; import EntityUsageExportModal from "../../EntityUsageExport"; @@ -84,6 +84,7 @@ const UsagePage: React.FC = ({ teams, organizations }) => { console.log(`currentUser: ${JSON.stringify(currentUser)}`); console.log(`currentUser max budget: ${currentUser?.max_budget}`); const isAdmin = all_admin_roles.includes(userRole || ""); + const canViewTagUsage = isAdmin || internalUserRoles.includes(userRole || ""); // Debounced search for user selector const [userSearchInput, setUserSearchInput] = useState(""); @@ -145,23 +146,6 @@ const UsagePage: React.FC = ({ teams, organizations }) => { const [topKeysLimit, setTopKeysLimit] = useState(5); const [topModelsLimit, setTopModelsLimit] = useState(5); const [showTokenBreakdown, setShowTokenBreakdown] = useState(false); - const getAllTags = async () => { - if (!accessToken) { - return; - } - const tags = await tagListCall(accessToken); - setAllTags( - Object.values(tags).map((tag: Tag) => ({ - label: tag.name, - value: tag.name, - })), - ); - }; - - useEffect(() => { - getAllTags(); - }, [accessToken]); - // Sync selectedUserId when auth state settles (isAdmin/userID may be null on initial render) useEffect(() => { if (!isAdmin && userID) { @@ -175,6 +159,30 @@ const UsagePage: React.FC = ({ teams, organizations }) => { const startTime = useMemo(() => (dateValue.from ? new Date(dateValue.from) : null), [dateValue.from]); const endTime = useMemo(() => (dateValue.to ? new Date(dateValue.to) : null), [dateValue.to]); + useEffect(() => { + if (!accessToken) return; + let cancelled = false; + (async () => { + try { + const tags = await tagListCall(accessToken, startTime, endTime); + if (cancelled) return; + setAllTags( + Object.values(tags).map((tag: Tag) => ({ + label: tag.name, + value: tag.name, + })), + ); + } catch (e) { + if (!cancelled) { + console.error("Failed to fetch tag list", e); + } + } + })(); + return () => { + cancelled = true; + }; + }, [accessToken, startTime, endTime]); + // Try aggregated endpoint first, fall back to paginated on failure const aggregatedFetchIdRef = useRef(0); useEffect(() => { @@ -437,7 +445,12 @@ const UsagePage: React.FC = ({ teams, organizations }) => {
- setUsageView(value)} isAdmin={isAdmin} /> + setUsageView(value)} + isAdmin={isAdmin} + canViewTagUsage={canViewTagUsage} + />
{paginatedResult.isFetchingMore && ( diff --git a/ui/litellm-dashboard/src/components/UsagePage/components/UsageViewSelect/UsageViewSelect.test.tsx b/ui/litellm-dashboard/src/components/UsagePage/components/UsageViewSelect/UsageViewSelect.test.tsx index 7bb80b424b5..b33129fcdb4 100644 --- a/ui/litellm-dashboard/src/components/UsagePage/components/UsageViewSelect/UsageViewSelect.test.tsx +++ b/ui/litellm-dashboard/src/components/UsagePage/components/UsageViewSelect/UsageViewSelect.test.tsx @@ -110,4 +110,16 @@ describe("UsageViewSelect", () => { expect(mockOnChange).toHaveBeenCalledWith("team"); }); + + it("should show Tag Usage for non-admin users with tag usage permission", () => { + render(); + + expect(screen.getByRole("option", { name: "Tag Usage" })).toBeInTheDocument(); + }); + + it("should hide Tag Usage for non-admin users without tag usage permission", () => { + render(); + + expect(screen.queryByRole("option", { name: "Tag Usage" })).not.toBeInTheDocument(); + }); }); diff --git a/ui/litellm-dashboard/src/components/UsagePage/components/UsageViewSelect/UsageViewSelect.tsx b/ui/litellm-dashboard/src/components/UsagePage/components/UsageViewSelect/UsageViewSelect.tsx index 11fdb8a7cf5..184fc000274 100644 --- a/ui/litellm-dashboard/src/components/UsagePage/components/UsageViewSelect/UsageViewSelect.tsx +++ b/ui/litellm-dashboard/src/components/UsagePage/components/UsageViewSelect/UsageViewSelect.tsx @@ -16,6 +16,7 @@ export interface UsageViewSelectProps { value: UsageOption; onChange: (value: UsageOption) => void; isAdmin: boolean; + canViewTagUsage?: boolean; title?: string; description?: string; "data-id"?: string; @@ -106,12 +107,16 @@ export const UsageViewSelect: React.FC = ({ value, onChange, isAdmin, + canViewTagUsage = false, title = "Usage View", description = "Select the usage data you want to view", "data-id": dataId, }) => { const getFilteredOptions = () => { return OPTIONS.filter((option) => { + if (option.value === "tag" && canViewTagUsage) { + return true; + } if (option.adminOnly && !isAdmin) { return false; } diff --git a/ui/litellm-dashboard/src/components/model_info_view.test.tsx b/ui/litellm-dashboard/src/components/model_info_view.test.tsx index 29eb8e0019b..54144d02eda 100644 --- a/ui/litellm-dashboard/src/components/model_info_view.test.tsx +++ b/ui/litellm-dashboard/src/components/model_info_view.test.tsx @@ -250,6 +250,33 @@ describe("ModelInfoView", () => { }); }); + it("should pass model_info.id to disambiguate duplicate model_name deployments", async () => { + // Regression test: when two deployments share `model_name` (e.g. + // wildcard `openai/*` with different `api_base` values), the UI + // must forward the clicked row's `model_info.id` to the backend. + // Otherwise /health/test_connection silently probes deployments[0] + // instead of the deployment the user actually selected. + const user = userEvent.setup(); + render(, { wrapper }); + + await waitFor(() => { + expect(screen.getByText("Model Settings")).toBeInTheDocument(); + }); + + const testButton = screen.getByRole("button", { name: /test connection/i }); + await user.click(testButton); + + await waitFor(() => { + expect(mockTestConnectionRequest).toHaveBeenCalled(); + }); + + const callArgs = mockTestConnectionRequest.mock.calls[0]; + // Signature: (accessToken, litellm_params, model_info, mode) + const modelInfoArg = callArgs[2] as Record; + expect(modelInfoArg).toBeDefined(); + expect(modelInfoArg.id).toBe("123"); + }); + it("should display error notification when connection test fails", async () => { const user = userEvent.setup(); mockTestConnectionRequest.mockRejectedValue(new Error("Connection failed")); diff --git a/ui/litellm-dashboard/src/components/model_info_view.tsx b/ui/litellm-dashboard/src/components/model_info_view.tsx index ea8af2fcd62..95a43862de5 100644 --- a/ui/litellm-dashboard/src/components/model_info_view.tsx +++ b/ui/litellm-dashboard/src/components/model_info_view.tsx @@ -379,6 +379,12 @@ export default function ModelInfoView({ model: localModelData.litellm_model_name, }, { + // `id` is required to disambiguate when multiple deployments + // share the same model_name (e.g. wildcard `openai/*` with two + // different `api_base` values for failover). Without it the + // backend silently falls back to deployments[0] and probes + // the wrong endpoint. + id: localModelData.model_info?.id, mode: localModelData.model_info?.mode, }, localModelData.model_info?.mode, diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 8bfe7c98691..7f28dbe6da9 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -7288,10 +7288,29 @@ export const tagInfoCall = async (accessToken: string, tagNames: string[]): Prom } }; -export const tagListCall = async (accessToken: string): Promise => { +const formatYmd = (value: Date): string => { + const year = value.getFullYear(); + const month = String(value.getMonth() + 1).padStart(2, "0"); + const day = String(value.getDate()).padStart(2, "0"); + return `${year}-${month}-${day}`; +}; + +export const tagListCall = async ( + accessToken: string, + startTime?: Date | null, + endTime?: Date | null, +): Promise => { try { let url = proxyBaseUrl ? `${proxyBaseUrl}/tag/list` : `/tag/list`; + if (startTime && endTime) { + const params = new URLSearchParams({ + start_date: formatYmd(startTime), + end_date: formatYmd(endTime), + }); + url = `${url}?${params.toString()}`; + } + const response = await fetch(url, { method: "GET", headers: { diff --git a/ui/litellm-dashboard/src/components/templates/key_edit_view.test.tsx b/ui/litellm-dashboard/src/components/templates/key_edit_view.test.tsx index 2e4d0d97e4c..1886075a9d9 100644 --- a/ui/litellm-dashboard/src/components/templates/key_edit_view.test.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_edit_view.test.tsx @@ -158,8 +158,8 @@ describe("KeyEditView", () => { const { getByText } = renderWithProviders( { }} - onSubmit={async () => { }} + onCancel={() => {}} + onSubmit={async () => {}} accessToken={""} userID={""} userRole={""} @@ -176,8 +176,8 @@ describe("KeyEditView", () => { const { getByText } = renderWithProviders( { }} - onSubmit={async () => { }} + onCancel={() => {}} + onSubmit={async () => {}} accessToken={""} userID={""} userRole={""} @@ -194,8 +194,8 @@ describe("KeyEditView", () => { const { getByLabelText } = renderWithProviders( { }} - onSubmit={async () => { }} + onCancel={() => {}} + onSubmit={async () => {}} accessToken={""} userID={""} userRole={""} @@ -219,7 +219,7 @@ describe("KeyEditView", () => { { }} + onSubmit={async () => {}} accessToken={""} userID={""} userRole={""} @@ -241,8 +241,8 @@ describe("KeyEditView", () => { renderWithProviders( { }} - onSubmit={async () => { }} + onCancel={() => {}} + onSubmit={async () => {}} accessToken={""} userID={""} userRole={""} @@ -259,8 +259,8 @@ describe("KeyEditView", () => { renderWithProviders( { }} - onSubmit={async () => { }} + onCancel={() => {}} + onSubmit={async () => {}} accessToken={""} userID={""} userRole={""} @@ -277,8 +277,8 @@ describe("KeyEditView", () => { renderWithProviders( { }} - onSubmit={async () => { }} + onCancel={() => {}} + onSubmit={async () => {}} accessToken={""} userID={""} userRole={""} @@ -295,8 +295,8 @@ describe("KeyEditView", () => { renderWithProviders( { }} - onSubmit={async () => { }} + onCancel={() => {}} + onSubmit={async () => {}} accessToken={""} userID={""} userRole={""} @@ -314,7 +314,7 @@ describe("KeyEditView", () => { renderWithProviders( { }} + onCancel={() => {}} onSubmit={onSubmitMock} accessToken={"test-token"} userID={"test-user"} @@ -344,8 +344,8 @@ describe("KeyEditView", () => { renderWithProviders( { }} - onSubmit={async () => { }} + onCancel={() => {}} + onSubmit={async () => {}} accessToken={""} userID={""} userRole={""} @@ -367,8 +367,8 @@ describe("KeyEditView", () => { renderWithProviders( { }} - onSubmit={async () => { }} + onCancel={() => {}} + onSubmit={async () => {}} accessToken={""} userID={""} userRole={""} @@ -385,8 +385,8 @@ describe("KeyEditView", () => { renderWithProviders( { }} - onSubmit={async () => { }} + onCancel={() => {}} + onSubmit={async () => {}} accessToken={"test-token"} userID={""} userRole={""} @@ -404,7 +404,7 @@ describe("KeyEditView", () => { renderWithProviders( { }} + onCancel={() => {}} onSubmit={onSubmitMock} accessToken={"test-token"} userID={"test-user"} @@ -434,10 +434,14 @@ describe("KeyEditView", () => { it("should handle empty allowed routes string on submit", async () => { const onSubmitMock = vi.fn().mockResolvedValue(undefined); + const keyDataWithRoutes = { + ...MOCK_KEY_DATA, + allowed_routes: ["llm_api_routes"], + }; renderWithProviders( { }} + keyData={keyDataWithRoutes} + onCancel={() => {}} onSubmit={onSubmitMock} accessToken={"test-token"} userID={"test-user"} @@ -463,6 +467,101 @@ describe("KeyEditView", () => { }); }); + it("should omit allowed_routes from submit when value is unchanged", async () => { + const onSubmitMock = vi.fn().mockResolvedValue(undefined); + const aiApisKeyData = { + ...MOCK_KEY_DATA, + allowed_routes: ["llm_api_routes"], + }; + renderWithProviders( + {}} + onSubmit={onSubmitMock} + accessToken={"test-token"} + userID={"test-user"} + userRole={"admin"} + premiumUser={false} + />, + ); + + await waitFor(() => { + expect(screen.getByRole("button", { name: /save changes/i })).toBeInTheDocument(); + }); + + const submitButton = screen.getByRole("button", { name: /save changes/i }); + await userEvent.click(submitButton); + + await waitFor(() => { + expect(onSubmitMock).toHaveBeenCalled(); + const callArgs = onSubmitMock.mock.calls[0][0]; + expect("allowed_routes" in callArgs).toBe(false); + }); + }); + + it("should omit allowed_routes from submit when keyData.allowed_routes is null and form is untouched", async () => { + const onSubmitMock = vi.fn().mockResolvedValue(undefined); + const keyDataNullRoutes = { + ...MOCK_KEY_DATA, + allowed_routes: null as unknown as string[], + }; + renderWithProviders( + {}} + onSubmit={onSubmitMock} + accessToken={"test-token"} + userID={"test-user"} + userRole={"admin"} + premiumUser={false} + />, + ); + + await waitFor(() => { + expect(screen.getByRole("button", { name: /save changes/i })).toBeInTheDocument(); + }); + + const submitButton = screen.getByRole("button", { name: /save changes/i }); + await userEvent.click(submitButton); + + await waitFor(() => { + expect(onSubmitMock).toHaveBeenCalled(); + const callArgs = onSubmitMock.mock.calls[0][0]; + expect("allowed_routes" in callArgs).toBe(false); + }); + }); + + it("should omit allowed_routes from submit when server returned routes in a different order", async () => { + const onSubmitMock = vi.fn().mockResolvedValue(undefined); + const keyDataReordered = { + ...MOCK_KEY_DATA, + allowed_routes: ["beta_routes", "alpha_routes"], + }; + renderWithProviders( + {}} + onSubmit={onSubmitMock} + accessToken={"test-token"} + userID={"test-user"} + userRole={"admin"} + premiumUser={false} + />, + ); + + await waitFor(() => { + expect(screen.getByRole("button", { name: /save changes/i })).toBeInTheDocument(); + }); + + const submitButton = screen.getByRole("button", { name: /save changes/i }); + await userEvent.click(submitButton); + + await waitFor(() => { + expect(onSubmitMock).toHaveBeenCalled(); + const callArgs = onSubmitMock.mock.calls[0][0]; + expect("allowed_routes" in callArgs).toBe(false); + }); + }); it("should pass access_group_ids to onSubmit when saving key with access groups", async () => { const onSubmitMock = vi.fn().mockResolvedValue(undefined); @@ -554,7 +653,7 @@ describe("KeyEditView", () => { renderWithProviders( { }} + onCancel={() => {}} onSubmit={onSubmitMock} accessToken={"test-token"} userID={"test-user"} @@ -576,10 +675,13 @@ describe("KeyEditView", () => { }); // Wait for the cancel button to actually be disabled (state update may take a moment) - await waitFor(() => { - const cancelButton = screen.getByRole("button", { name: /cancel/i }); - expect(cancelButton).toBeDisabled(); - }, { timeout: 3000 }); + await waitFor( + () => { + const cancelButton = screen.getByRole("button", { name: /cancel/i }); + expect(cancelButton).toBeDisabled(); + }, + { timeout: 3000 }, + ); // Clean up: resolve the promise to allow the form to complete if (resolveSubmit) { diff --git a/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx b/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx index 9b38d930a50..d5e410029a7 100644 --- a/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx @@ -78,7 +78,6 @@ const getKeyTypeFromRoutes = (allowedRoutes: string[] | null | undefined): strin return "default"; }; - export function KeyEditView({ keyData, onCancel, @@ -106,7 +105,7 @@ export function KeyEditView({ const [neverExpire, setNeverExpire] = useState(!keyData.expires); const [isKeySaving, setIsKeySaving] = useState(false); const [budgetLimits, setBudgetLimits] = useState( - Array.isArray(keyData.budget_limits) ? keyData.budget_limits : [] + Array.isArray(keyData.budget_limits) ? keyData.budget_limits : [], ); const { data: organizations, isLoading: isOrganizationsLoading } = useOrganizations(); const { data: projects } = useProjects(); @@ -116,9 +115,7 @@ export function KeyEditView({ const projectDisplay = (() => { if (!keyData.project_id) return null; const project = projects?.find((p) => p.project_id === keyData.project_id); - return project?.project_alias - ? `${project.project_alias} (${keyData.project_id})` - : keyData.project_id; + return project?.project_alias ? `${project.project_alias} (${keyData.project_id})` : keyData.project_id; })(); useEffect(() => { @@ -198,9 +195,10 @@ export function KeyEditView({ access_group_ids: keyData.access_group_ids || [], auto_rotate: keyData.auto_rotate || false, ...(keyData.rotation_interval && { rotation_interval: keyData.rotation_interval }), - allowed_routes: Array.isArray(keyData.allowed_routes) && keyData.allowed_routes.length > 0 - ? keyData.allowed_routes.join(", ") - : "", + allowed_routes: + Array.isArray(keyData.allowed_routes) && keyData.allowed_routes.length > 0 + ? keyData.allowed_routes.join(", ") + : "", }; useEffect(() => { @@ -226,9 +224,10 @@ export function KeyEditView({ access_group_ids: keyData.access_group_ids || [], auto_rotate: keyData.auto_rotate || false, ...(keyData.rotation_interval && { rotation_interval: keyData.rotation_interval }), - allowed_routes: Array.isArray(keyData.allowed_routes) && keyData.allowed_routes.length > 0 - ? keyData.allowed_routes.join(", ") - : "", + allowed_routes: + Array.isArray(keyData.allowed_routes) && keyData.allowed_routes.length > 0 + ? keyData.allowed_routes.join(", ") + : "", }); }, [keyData, form]); @@ -275,12 +274,25 @@ export function KeyEditView({ } // If it's already an array (shouldn't happen, but handle it), keep as is + // Backend rejects non-empty allowed_routes from non-admins, so re-sending + // an unchanged value 403s a team admin. Set compare tolerates reorder. + const originalRoutesSet = new Set(Array.isArray(keyData.allowed_routes) ? keyData.allowed_routes : []); + const submittedRoutesSet = new Set(Array.isArray(values.allowed_routes) ? values.allowed_routes : []); + const allowedRoutesUnchanged = + originalRoutesSet.size === submittedRoutesSet.size && + [...submittedRoutesSet].every((r) => originalRoutesSet.has(r)); + if (allowedRoutesUnchanged) { + delete values.allowed_routes; + } + if (neverExpire) { values.duration = null; } // Include multi-window budget limits (filter out incomplete entries) - const validWindows = budgetLimits.filter((w) => w.budget_duration && w.max_budget !== null && w.max_budget !== undefined); + const validWindows = budgetLimits.filter( + (w) => w.budget_duration && w.max_budget !== null && w.max_budget !== undefined, + ); values.budget_limits = validWindows.length > 0 ? validWindows : undefined; await onSubmit(values); @@ -305,9 +317,13 @@ export function KeyEditView({ {({ getFieldValue, setFieldValue }) => { const allowedRoutesValue = getFieldValue("allowed_routes") || ""; // Convert string to array for checking - const allowedRoutes = typeof allowedRoutesValue === "string" && allowedRoutesValue.trim() !== "" - ? allowedRoutesValue.split(",").map((r: string) => r.trim()).filter((r: string) => r.length > 0) - : []; + const allowedRoutes = + typeof allowedRoutesValue === "string" && allowedRoutesValue.trim() !== "" + ? allowedRoutesValue + .split(",") + .map((r: string) => r.trim()) + .filter((r: string) => r.length > 0) + : []; const isDisabled = allowedRoutes.includes("management_routes") || allowedRoutes.includes("info_routes"); const models = getFieldValue("models") || []; @@ -348,9 +364,13 @@ export function KeyEditView({ {({ getFieldValue, setFieldValue }) => { const allowedRoutesValue = getFieldValue("allowed_routes") || ""; // Convert string to array for getKeyTypeFromRoutes - const allowedRoutes = typeof allowedRoutesValue === "string" && allowedRoutesValue.trim() !== "" - ? allowedRoutesValue.split(",").map((r: string) => r.trim()).filter((r: string) => r.length > 0) - : []; + const allowedRoutes = + typeof allowedRoutesValue === "string" && allowedRoutesValue.trim() !== "" + ? allowedRoutesValue + .split(",") + .map((r: string) => r.trim()) + .filter((r: string) => r.length > 0) + : []; const keyTypeValue = getKeyTypeFromRoutes(allowedRoutes); return ( @@ -415,9 +435,7 @@ export function KeyEditView({ } name="allowed_routes" > - + @@ -442,10 +460,7 @@ export function KeyEditView({ } > - + @@ -579,7 +594,7 @@ export function KeyEditView({ !premiumUser ? "Premium feature - Upgrade to set allowed pass through routes by key" : Array.isArray(keyData.metadata?.allowed_passthrough_routes) && - keyData.metadata.allowed_passthrough_routes.length > 0 + keyData.metadata.allowed_passthrough_routes.length > 0 ? `Current: ${keyData.metadata.allowed_passthrough_routes.join(", ")}` : "Select or enter allowed pass through routes" } @@ -690,14 +705,13 @@ export function KeyEditView({ return team.team_alias?.toLowerCase().includes(input.toLowerCase()) ?? false; }} > - {(selectedOrganizationId - ? teams?.filter((t) => t.organization_id === selectedOrganizationId) - : teams - )?.map((team) => ( - - {`${team.team_alias} (${team.team_id})`} - - ))} + {(selectedOrganizationId ? teams?.filter((t) => t.organization_id === selectedOrganizationId) : teams)?.map( + (team) => ( + + {`${team.team_alias} (${team.team_id})`} + + ), + )} {enableProjectsUI && hasProject && ( diff --git a/ui/litellm-dashboard/src/components/templates/key_info_view.tsx b/ui/litellm-dashboard/src/components/templates/key_info_view.tsx index 492e43cbc81..65bd9d9eb95 100644 --- a/ui/litellm-dashboard/src/components/templates/key_info_view.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_info_view.tsx @@ -396,7 +396,7 @@ export default function KeyInfoView({ }; return ( -
+
- +
Key Settings {!isEditing && canModifyKey && ( diff --git a/ui/litellm-dashboard/src/components/view_logs/constants.ts b/ui/litellm-dashboard/src/components/view_logs/constants.ts index 57155feae23..5b0b1d0fee3 100644 --- a/ui/litellm-dashboard/src/components/view_logs/constants.ts +++ b/ui/litellm-dashboard/src/components/view_logs/constants.ts @@ -19,6 +19,7 @@ export const MCP_CALL_TYPES = ["call_mcp_tool", "list_mcp_tools"]; export const AGENT_CALL_TYPES = ["asend_message"]; export const QUICK_SELECT_OPTIONS: { label: string; value: number; unit: string }[] = [ + { label: "Last Minute", value: 1, unit: "minutes" }, { label: "Last 15 Minutes", value: 15, unit: "minutes" }, { label: "Last Hour", value: 1, unit: "hours" }, { label: "Last 4 Hours", value: 4, unit: "hours" }, diff --git a/ui/litellm-dashboard/src/components/view_logs/index.test.tsx b/ui/litellm-dashboard/src/components/view_logs/index.test.tsx index 937844e1c1d..7a9a541d3e0 100644 --- a/ui/litellm-dashboard/src/components/view_logs/index.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/index.test.tsx @@ -1,10 +1,12 @@ import { render, screen, waitFor } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; +import moment from "moment"; import { beforeEach, describe, expect, it, vi } from "vitest"; import SpendLogsTable, { RequestViewer } from "./index"; import type { LogEntry } from "./columns"; import type { Row } from "@tanstack/react-table"; import { renderWithProviders } from "../../../tests/test-utils"; +import { uiSpendLogsCall } from "../networking"; const mockHandleFilterResetFromHook = vi.fn(); vi.mock("./log_filter_logic", async (importOriginal) => { @@ -238,4 +240,52 @@ describe("SpendLogsTable", () => { expect(inputsAfterReset.length).toBe(0); }); }); + + describe("Quick Select time range", () => { + const waitForWindowSeconds = async (minMinutes: number) => { + let diff = -1; + await waitFor(() => { + const lastCall = vi.mocked(uiSpendLogsCall).mock.calls.at(-1)?.[0]; + if (!lastCall) throw new Error("uiSpendLogsCall was not called"); + diff = moment + .utc(lastCall.end_date, "YYYY-MM-DD HH:mm:ss") + .diff(moment.utc(lastCall.start_date, "YYYY-MM-DD HH:mm:ss"), "seconds"); + // start_date is rounded down to the minute boundary; end_date is current time + expect(diff).toBeGreaterThanOrEqual(minMinutes * 60); + expect(diff).toBeLessThan((minMinutes + 1) * 60); + }); + return diff; + }; + + it("should pass a ~1-minute window to uiSpendLogsCall when 'Last Minute' is selected", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await user.click(screen.getByRole("button", { name: /Last 24 Hours/i })); + await user.click(await screen.findByRole("button", { name: "Last Minute" })); + + await waitForWindowSeconds(1); + }); + + it("should pass a ~15-minute window to uiSpendLogsCall when 'Last 15 Minutes' is selected", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await user.click(screen.getByRole("button", { name: /Last 24 Hours/i })); + await user.click(await screen.findByRole("button", { name: "Last 15 Minutes" })); + + await waitForWindowSeconds(15); + }); + + it("should update the time-range button label to 'Last Minute' after selecting it", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await user.click(screen.getByRole("button", { name: /Last 24 Hours/i })); + await user.click(await screen.findByRole("button", { name: "Last Minute" })); + + expect(screen.getByRole("button", { name: "Last Minute" })).toBeInTheDocument(); + expect(screen.queryByRole("button", { name: /Last 24 Hours/i })).not.toBeInTheDocument(); + }); + }); }); diff --git a/ui/litellm-dashboard/src/utils/roles.ts b/ui/litellm-dashboard/src/utils/roles.ts index 11e31fd369c..bb23985d7c2 100644 --- a/ui/litellm-dashboard/src/utils/roles.ts +++ b/ui/litellm-dashboard/src/utils/roles.ts @@ -5,7 +5,7 @@ export const old_admin_roles = ["Admin", "Admin Viewer"]; export const v2_admin_role_names = ["proxy_admin", "proxy_admin_viewer", "org_admin"]; export const all_admin_roles = [...old_admin_roles, ...v2_admin_role_names]; -export const internalUserRoles = ["Internal User", "Internal Viewer"]; +export const internalUserRoles = ["Internal User", "Internal Viewer", "internal_user", "internal_user_viewer"]; export const rolesAllowedToSeeUsage = ["Admin", "Admin Viewer", "Internal User", "Internal Viewer"]; export const rolesWithWriteAccess = ["Internal User", "Admin", "proxy_admin"]; // Admin-tier read parity: Admin Viewer sees Models + Endpoints, Agents, and