diff --git a/.circleci/config.yml b/.circleci/config.yml index 4a442426976..0e53cfc0edb 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -1521,6 +1521,7 @@ jobs: - run: python ./tests/code_coverage_tests/prevent_key_leaks_in_exceptions.py - run: python ./tests/code_coverage_tests/check_unsafe_enterprise_import.py - run: python ./tests/code_coverage_tests/ban_copy_deepcopy_kwargs.py + - run: python ./tests/code_coverage_tests/check_fastuuid_usage.py - run: helm lint ./deploy/charts/litellm-helm db_migration_disable_update_check: diff --git a/litellm/_uuid.py b/litellm/_uuid.py new file mode 100644 index 00000000000..05b1adbf75b --- /dev/null +++ b/litellm/_uuid.py @@ -0,0 +1,23 @@ +""" +Internal unified UUID helper. + +Tries to use fastuuid (performance) and falls back to stdlib uuid if unavailable. +""" + +FASTUUID_AVAILABLE = False + +try: + import fastuuid as _uuid # type: ignore + + FASTUUID_AVAILABLE = True +except Exception: # pragma: no cover - fallback path + import uuid as _uuid # type: ignore + + +# Expose a module-like alias so callers can use: uuid.uuid4() +uuid = _uuid + + +def uuid4(): + """Return a UUID4 using the selected backend.""" + return uuid.uuid4() diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 65f665041f9..64986970d00 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -26,7 +26,6 @@ from typing import ( cast, ) -import fastuuid as uuid from httpx import Response from pydantic import BaseModel @@ -38,6 +37,7 @@ from litellm import ( turn_off_message_logging, ) from litellm._logging import _is_debugging_on, verbose_logger +from litellm._uuid import uuid from litellm.batches.batch_utils import _handle_completed_batch from litellm.caching.caching import DualCache, InMemoryCache from litellm.caching.caching_handler import LLMCachingHandler diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 322691e28b4..bc0fbdf5c11 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -1619,11 +1619,12 @@ class CustomStreamWrapper: completion_start_time=datetime.datetime.now() ) ## LOGGING - executor.submit( - self.run_success_logging_and_cache_storage, - response, - cache_hit, - ) # log response + if not litellm.disable_streaming_logging: + executor.submit( + self.run_success_logging_and_cache_storage, + response, + cache_hit, + ) # log response choice = response.choices[0] if isinstance(choice, StreamingChoices): self.response_uptil_now += choice.delta.get("content", "") or "" diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index e54aeaf995f..6d41655674c 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -804,7 +804,15 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): if content.get("citations") is not None: if citations is None: citations = [] - citations.append(content["citations"]) + citations.append( + [ + { + **citation, + "supported_text": content.get("text", ""), + } + for citation in content["citations"] + ] + ) if thinking_blocks is not None: reasoning_content = "" for block in thinking_blocks: diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index c7f7acf331d..63c366c9480 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -774,7 +774,7 @@ class CommonBatchFilesUtils: Returns: Unique job name (≤ 63 characters for Bedrock compatibility) """ - import fastuuid as uuid + from litellm._uuid import uuid unique_id = str(uuid.uuid4())[:8] # Format: {prefix}-batch-{model}-{uuid} # Example: litellm-batch-claude-266c398e diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 706b09bb098..3a67d5127b7 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -935,7 +935,7 @@ class HTTPHandler: if litellm.force_ipv4: return HTTPTransport(local_address="0.0.0.0") else: - return None + return getattr(litellm, 'sync_transport', None) def get_async_httpx_client( diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 5739e652043..95c84b914b6 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -14,7 +14,6 @@ from typing import ( Union, ) -import fastuuid as uuid import httpx import orjson from fastapi import HTTPException, Request, status @@ -22,6 +21,7 @@ from fastapi.responses import Response, StreamingResponse import litellm from litellm._logging import verbose_proxy_logger +from litellm._uuid import uuid from litellm.constants import ( DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE, STREAM_SSE_DATA_PREFIX, diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 5e171af5252..a834a7a13c3 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -459,6 +459,14 @@ async def anthropic_proxy_route( region_name=None, ) + custom_headers = {} + if ( + "authorization" not in request.headers + and "x-api-key" not in request.headers + and anthropic_api_key is not None + ): + custom_headers["x-api-key"] = "{}".format(anthropic_api_key) + ## check for streaming is_streaming_request = await is_streaming_request_fn(request) @@ -466,7 +474,7 @@ async def anthropic_proxy_route( endpoint_func = create_pass_through_route( endpoint=endpoint, target=str(updated_url), - custom_headers={"x-api-key": "{}".format(anthropic_api_key)}, + custom_headers=custom_headers, _forward_headers=True, ) # dynamically construct pass-through endpoint based on incoming path received_value = await endpoint_func( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 52f42c9f310..472aeb140cf 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -35,6 +35,7 @@ from litellm.constants import ( LITELLM_SETTINGS_SAFE_DB_OVERRIDES, ) from litellm.litellm_core_utils.safe_json_dumps import safe_dumps +from litellm.utils import load_credentials_from_list from litellm.types.utils import ( ModelResponse, ModelResponseStream, @@ -5894,6 +5895,7 @@ async def token_counter(request: TokenCountRequest, call_endpoint: bool = False) pass if deployment is not None: litellm_model_name = deployment.get("litellm_params", {}).get("model") + load_credentials_from_list(deployment.get("litellm_params", {})) # remove the custom_llm_provider_prefix in the litellm_model_name if "/" in litellm_model_name: litellm_model_name = litellm_model_name.split("/", 1)[1] diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index 83acf0e8226..e2ca3078384 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -493,21 +493,32 @@ async def update_sso_settings(sso_config: SSOConfig): config["general_settings"] = {} # Update environment variables in config and in memory - sso_data = sso_config.model_dump(exclude_none=True) + sso_data = sso_config.model_dump() for field_name, value in sso_data.items(): - - if field_name == "user_email" and value is not None: - # Store user_email in general_settings instead of environment variables - config["general_settings"]["proxy_admin_email"] = value - elif field_name == "ui_access_mode" and value is not None: - - config["general_settings"]["ui_access_mode"] = value - elif field_name in env_var_mapping and value is not None: + if field_name == "user_email": + if value: + # Store user_email in general_settings instead of environment variables + config["general_settings"]["proxy_admin_email"] = value + else: + # Clear user_email if null/empty + config["general_settings"].pop("proxy_admin_email", None) + elif field_name == "ui_access_mode": + if value: + config["general_settings"]["ui_access_mode"] = value + else: + # Clear ui_access_mode if null/empty + config["general_settings"].pop("ui_access_mode", None) + elif field_name in env_var_mapping and value: env_var_name = env_var_mapping[field_name] # Update in config config["environment_variables"][env_var_name] = value # Update in runtime environment os.environ[env_var_name] = value + elif field_name in env_var_mapping: + # Clear environment variable if value is null/empty + env_var_name = env_var_mapping[field_name] + config["environment_variables"].pop(env_var_name, None) + os.environ.pop(env_var_name, None) stored_config = config if len(config["environment_variables"]) > 0: diff --git a/litellm/types/utils.py b/litellm/types/utils.py index afa545d7fb9..d37978f3e21 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -13,7 +13,6 @@ from typing import ( Union, ) -import fastuuid as uuid from aiohttp import FormData from openai._models import BaseModel as OpenAIObject from openai.types.audio.transcription_create_params import FileTypes # type: ignore @@ -33,6 +32,7 @@ from pydantic import BaseModel, ConfigDict, Field, PrivateAttr, model_validator from typing_extensions import Callable, Dict, Required, TypedDict, override import litellm +from litellm._uuid import uuid from litellm.types.llms.base import ( BaseLiteLLMOpenAIResponseObject, LiteLLMPydanticObjectBase, diff --git a/litellm/utils.py b/litellm/utils.py index 3c3ab0832d7..0721b023d2b 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -40,7 +40,6 @@ from os.path import abspath, dirname, join import aiohttp import dotenv -import fastuuid as uuid import httpx import openai import tiktoken @@ -59,6 +58,7 @@ import litellm.litellm_core_utils.audio_utils.utils import litellm.litellm_core_utils.json_validation_rule import litellm.llms import litellm.llms.gemini +from litellm._uuid import uuid from litellm.caching._internal_lru_cache import lru_cache_wrapper from litellm.caching.caching import DualCache from litellm.caching.caching_handler import CachingHandlerResponse, LLMCachingHandler diff --git a/poetry.lock b/poetry.lock index 4d0e3bea696..a8ef9ee79df 100644 --- a/poetry.lock +++ b/poetry.lock @@ -555,7 +555,7 @@ description = "Foreign Function Interface for Python calling C code." optional = false python-versions = ">=3.8" groups = ["main", "dev", "proxy-dev"] -markers = "python_version < \"3.14\" and platform_python_implementation != \"PyPy\"" +markers = "python_version < \"3.10\" and platform_python_implementation != \"PyPy\"" files = [ {file = "cffi-1.17.1-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:df8b1c11f177bc2313ec4b2d46baec87a5f3e71fc8b45dab2ee7cae86d9aba14"}, {file = "cffi-1.17.1-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:8f2cdc858323644ab277e9bb925ad72ae0e67f69e804f4898c070998d50b1a67"}, @@ -636,7 +636,7 @@ description = "Foreign Function Interface for Python calling C code." optional = false python-versions = ">=3.9" groups = ["main", "dev", "proxy-dev"] -markers = "platform_python_implementation != \"PyPy\" and python_version >= \"3.14\"" +markers = "python_version >= \"3.10\" and platform_python_implementation != \"PyPy\"" files = [ {file = "cffi-2.0.0-cp310-cp310-macosx_10_13_x86_64.whl", hash = "sha256:0cf2d91ecc3fcc0625c2c530fe004f82c110405f101548512cce44322fa8ac44"}, {file = "cffi-2.0.0-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:f73b96c41e3b2adedc34a7356e64c8eb96e03a3782b535e043a986276ce12a49"}, @@ -1339,9 +1339,10 @@ zstandard = ["zstandard"] name = "fastuuid" version = "0.12.0" description = "Python bindings to Rust's UUID library." -optional = false +optional = true python-versions = ">=3.8" groups = ["main"] +markers = "extra == \"proxy\"" files = [ {file = "fastuuid-0.12.0-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:22a900ef0956aacf862b460e20541fdae2d7c340594fe1bd6fdcb10d5f0791a9"}, {file = "fastuuid-0.12.0-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:0302f5acf54dc75de30103025c5a95db06d6c2be36829043a0aa16fc170076bc"}, @@ -4361,7 +4362,7 @@ description = "C parser in Python" optional = false python-versions = ">=3.8" groups = ["main", "dev", "proxy-dev"] -markers = "platform_python_implementation != \"PyPy\" and (python_version < \"3.14\" or implementation_name != \"PyPy\")" +markers = "platform_python_implementation != \"PyPy\" and (implementation_name != \"PyPy\" or python_version < \"3.10\")" files = [ {file = "pycparser-2.23-py3-none-any.whl", hash = "sha256:e5c6e8d3fbad53479cab09ac03729e0a9faf2bee3db8208a550daf5af81a5934"}, {file = "pycparser-2.23.tar.gz", hash = "sha256:78816d4f24add8f10a06d6f05b4d424ad9e96cfebf68a4ddc99c65c0720d00c2"}, @@ -6746,11 +6747,11 @@ type = ["pytest-mypy"] caching = ["diskcache"] extra-proxy = ["azure-identity", "azure-keyvault-secrets", "google-cloud-iam", "google-cloud-kms", "prisma", "redisvl", "resend"] mlflow = ["mlflow"] -proxy = ["PyJWT", "apscheduler", "azure-identity", "azure-storage-blob", "backoff", "boto3", "cryptography", "fastapi", "fastapi-sso", "gunicorn", "litellm-enterprise", "litellm-proxy-extras", "mcp", "orjson", "polars", "pynacl", "python-multipart", "pyyaml", "rich", "rq", "uvicorn", "uvloop", "websockets"] +proxy = ["PyJWT", "apscheduler", "azure-identity", "azure-storage-blob", "backoff", "boto3", "cryptography", "fastapi", "fastapi-sso", "fastuuid", "gunicorn", "litellm-enterprise", "litellm-proxy-extras", "mcp", "orjson", "polars", "pynacl", "python-multipart", "pyyaml", "rich", "rq", "uvicorn", "uvloop", "websockets"] semantic-router = ["semantic-router"] utils = ["numpydoc"] [metadata] lock-version = "2.1" python-versions = ">=3.8.1,<4.0, !=3.9.7" -content-hash = "dcec654bc4b233d2f0d160341d8adf1ba2017ccb74c104bff9ba4cf027ba0186" +content-hash = "75004c6a23b70be86622fa417fd0d62fa3843e6e61c8dff8507ae5c967b7205d" diff --git a/pyproject.toml b/pyproject.toml index dc22000f457..b11b1a1c2e2 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -20,7 +20,6 @@ Documentation = "https://docs.litellm.ai" [tool.poetry.dependencies] python = ">=3.8.1,<4.0, !=3.9.7" -fastuuid = ">=0.12.0" httpx = ">=0.23.0" openai = ">=1.99.5" python-dotenv = ">=0.2.0" @@ -34,6 +33,7 @@ pydantic = "^2.5.0" jsonschema = "^4.22.0" pondpond = "^1.4.1" numpydoc = {version = "*", optional = true} # used in utils.py +fastuuid = {version = ">=0.12.0", optional = true} uvicorn = {version = "^0.29.0", optional = true} uvloop = {version = "^0.21.0", optional = true, markers="sys_platform != 'win32'"} @@ -93,6 +93,7 @@ proxy = [ "litellm-enterprise", "rich", "polars", + "fastuuid", ] extra_proxy = [ @@ -115,6 +116,7 @@ semantic-router = ["semantic-router"] mlflow = ["mlflow"] + [tool.isort] profile = "black" diff --git a/tests/code_coverage_tests/check_fastuuid_usage.py b/tests/code_coverage_tests/check_fastuuid_usage.py new file mode 100644 index 00000000000..e0433454371 --- /dev/null +++ b/tests/code_coverage_tests/check_fastuuid_usage.py @@ -0,0 +1,87 @@ +import ast +import os +from typing import List, Dict, Any + + +ALLOWED_FILE = os.path.normpath("litellm/_uuid.py") + + +def _to_module_path(relative_path: str) -> str: + module = os.path.splitext(relative_path)[0].replace(os.sep, ".") + if module.endswith(".__init__"): + return module[: -len(".__init__")] + return module + + +def _find_fastuuid_imports_in_file( + file_path: str, base_dir: str +) -> List[Dict[str, Any]]: + results: List[Dict[str, Any]] = [] + try: + with open(file_path, "r", encoding="utf-8") as f: + source = f.read() + tree = ast.parse(source, filename=file_path) + except Exception: + return results + + relative = os.path.normpath(os.path.relpath(file_path, base_dir)) + if relative == ALLOWED_FILE: + return results + + module = _to_module_path(relative) + for node in ast.walk(tree): + if isinstance(node, ast.Import): + for alias in node.names: + if alias.name == "fastuuid": + results.append( + { + "file": relative, + "line": getattr(node, "lineno", 0), + "import": f"import {alias.name}", + "module": module, + } + ) + elif isinstance(node, ast.ImportFrom) and node.module == "fastuuid": + names = ", ".join([a.name for a in node.names]) + results.append( + { + "file": relative, + "line": getattr(node, "lineno", 0), + "import": f"from fastuuid import {names}", + "module": module, + } + ) + + return results + + +def scan_directory_for_fastuuid(base_dir: str) -> List[Dict[str, Any]]: + violations: List[Dict[str, Any]] = [] + scan_root = os.path.join(base_dir, "litellm") + for root, _, files in os.walk(scan_root): + for filename in files: + if filename.endswith(".py"): + file_path = os.path.join(root, filename) + violations.extend(_find_fastuuid_imports_in_file(file_path, base_dir)) + return violations + + +def main() -> None: + base_dir = "." # tests run from repo root in CI + violations = scan_directory_for_fastuuid(base_dir) + if violations: + print( + "\n🚨 fastuuid must only be imported inside litellm/_uuid.py. Found violations:" + ) + for v in violations: + print(f"* {v['module']} ({v['file']}:{v['line']}) -> {v['import']}") + print("\n") + raise Exception( + "Found fastuuid imports outside litellm/_uuid.py. Use litellm._uuid.uuid or litellm._uuid.uuid4 instead." + ) + else: + print("✅ No invalid fastuuid imports found.") + + +if __name__ == "__main__": + main() diff --git a/tests/llm_translation/test_anthropic_completion.py b/tests/llm_translation/test_anthropic_completion.py index 45702a261e2..307c429fc62 100644 --- a/tests/llm_translation/test_anthropic_completion.py +++ b/tests/llm_translation/test_anthropic_completion.py @@ -256,7 +256,6 @@ def test_anthropic_tool_streaming(): for chunk in anthropic_chunk_list: parsed_chunk = response_iter.chunk_parser(chunk) if tool_use := parsed_chunk.get("tool_use"): - # We only increment when a new block starts if tool_use.get("id") is not None: correct_tool_index += 1 @@ -920,6 +919,14 @@ def test_anthropic_citations_api(): citations = resp.choices[0].message.provider_specific_fields["citations"] assert citations is not None + if citations: + citation = citations[0][0] + assert "supported_text" in citation + assert "cited_text" in citation + assert "document_index" in citation + assert "document_title" in citation + assert "start_char_index" in citation + assert "end_char_index" in citation def test_anthropic_citations_api_streaming(): @@ -955,11 +962,9 @@ def test_anthropic_citations_api_streaming(): has_citations = False for chunk in resp: print(f"returned chunk: {chunk}") - if ( - chunk.choices[0].delta.provider_specific_fields - and "citation" in chunk.choices[0].delta.provider_specific_fields - ): - has_citations = True + if provider_specific_fields := chunk.choices[0].delta.provider_specific_fields: + if "citation" in provider_specific_fields: + has_citations = True assert has_citations 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 ae44198af11..deb442e8075 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 @@ -183,7 +183,30 @@ def test_extract_response_content_with_citations(): } _, citations, _, _, _ = config.extract_response_content(completion_response) - assert citations is not None + assert citations == [ + [ + { + "type": "char_location", + "cited_text": "The grass is green. ", + "document_index": 0, + "document_title": "My Document", + "start_char_index": 0, + "end_char_index": 20, + "supported_text": "the grass is green", + }, + ], + [ + { + "type": "char_location", + "cited_text": "The sky is blue.", + "document_index": 0, + "document_title": "My Document", + "start_char_index": 20, + "end_char_index": 36, + "supported_text": "the sky is blue", + }, + ], + ] def test_map_tool_helper(): diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_anthropic_auth_headers.py b/tests/test_litellm/proxy/pass_through_endpoints/test_anthropic_auth_headers.py new file mode 100644 index 00000000000..9872ed2c9d1 --- /dev/null +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_anthropic_auth_headers.py @@ -0,0 +1,218 @@ +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +sys.path.insert( + 0, os.path.abspath("../../..") +) # Adds the parent directory to the system path + +from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + anthropic_proxy_route, +) + + +class TestAnthropicAuthHeaders: + """Test authentication header handling in anthropic_proxy_route.""" + + @pytest.fixture + def mock_request(self): + """Create a mock request object.""" + request = MagicMock() + request.method = "POST" + request.headers = {} + return request + + @pytest.fixture + def mock_response(self): + """Create a mock FastAPI response object.""" + return MagicMock() + + @pytest.fixture + def mock_user_api_key_dict(self): + """Create a mock user API key dict.""" + return {"user_id": "test_user"} + + @pytest.mark.asyncio + @patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route") + @patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_streaming_request_fn") + @patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router") + async def test_client_authorization_header_priority( + self, + mock_router, + mock_streaming, + mock_create_route, + mock_request, + mock_response, + mock_user_api_key_dict, + ): + """Test that client Authorization header takes priority over server key.""" + # Setup + mock_request.headers = {"authorization": "Bearer client-key-123"} + mock_router.get_credentials.return_value = "server-key-456" + mock_streaming.return_value = False + mock_endpoint_func = AsyncMock(return_value="test_response") + mock_create_route.return_value = mock_endpoint_func + + # Act + await anthropic_proxy_route( + endpoint="v1/messages", + request=mock_request, + fastapi_response=mock_response, + user_api_key_dict=mock_user_api_key_dict, + ) + + # Assert + mock_create_route.assert_called_once() + call_kwargs = mock_create_route.call_args[1] + + assert call_kwargs["custom_headers"] == {} + assert call_kwargs["_forward_headers"] is True + + @pytest.mark.asyncio + @patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route") + @patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_streaming_request_fn") + @patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router") + async def test_client_x_api_key_header_priority( + self, + mock_router, + mock_streaming, + mock_create_route, + mock_request, + mock_response, + mock_user_api_key_dict, + ): + """Test that client x-api-key header takes priority over server key.""" + # Setup + mock_request.headers = {"x-api-key": "client-x-api-key-123"} + mock_router.get_credentials.return_value = "server-key-456" + mock_streaming.return_value = False + mock_endpoint_func = AsyncMock(return_value="test_response") + mock_create_route.return_value = mock_endpoint_func + + # Act + await anthropic_proxy_route( + endpoint="v1/messages", + request=mock_request, + fastapi_response=mock_response, + user_api_key_dict=mock_user_api_key_dict, + ) + + # Assert + mock_create_route.assert_called_once() + call_kwargs = mock_create_route.call_args[1] + + assert call_kwargs["custom_headers"] == {} + assert call_kwargs["_forward_headers"] is True + + @pytest.mark.asyncio + @patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route") + @patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_streaming_request_fn") + @patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router") + async def test_server_api_key_fallback( + self, + mock_router, + mock_streaming, + mock_create_route, + mock_request, + mock_response, + mock_user_api_key_dict, + ): + """Test that server API key is used when no client authentication is provided.""" + # Setup + mock_request.headers = {} # No authentication headers + mock_router.get_credentials.return_value = "server-key-456" + mock_streaming.return_value = False + mock_endpoint_func = AsyncMock(return_value="test_response") + mock_create_route.return_value = mock_endpoint_func + + # Act + await anthropic_proxy_route( + endpoint="v1/messages", + request=mock_request, + fastapi_response=mock_response, + user_api_key_dict=mock_user_api_key_dict, + ) + + # Assert + mock_create_route.assert_called_once() + call_kwargs = mock_create_route.call_args[1] + + assert call_kwargs["custom_headers"] == {"x-api-key": "server-key-456"} + assert call_kwargs["_forward_headers"] is True + + @pytest.mark.asyncio + @patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route") + @patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_streaming_request_fn") + @patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router") + async def test_no_authentication_available( + self, + mock_router, + mock_streaming, + mock_create_route, + mock_request, + mock_response, + mock_user_api_key_dict, + ): + """Test that no x-api-key header is added when no authentication is available.""" + # Setup + mock_request.headers = {} # No authentication headers + mock_router.get_credentials.return_value = None # No server key + mock_streaming.return_value = False + mock_endpoint_func = AsyncMock(return_value="test_response") + mock_create_route.return_value = mock_endpoint_func + + # Act + await anthropic_proxy_route( + endpoint="v1/messages", + request=mock_request, + fastapi_response=mock_response, + user_api_key_dict=mock_user_api_key_dict, + ) + + # Assert + mock_create_route.assert_called_once() + call_kwargs = mock_create_route.call_args[1] + + assert call_kwargs["custom_headers"] == {} + assert call_kwargs["_forward_headers"] is True + + @pytest.mark.asyncio + @patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route") + @patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_streaming_request_fn") + @patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router") + async def test_both_client_headers_present( + self, + mock_router, + mock_streaming, + mock_create_route, + mock_request, + mock_response, + mock_user_api_key_dict, + ): + """Test that no server key is added when client has both auth headers.""" + # Setup + mock_request.headers = { + "authorization": "Bearer client-auth-key", + "x-api-key": "client-x-api-key" + } + mock_router.get_credentials.return_value = "server-key-456" + mock_streaming.return_value = False + mock_endpoint_func = AsyncMock(return_value="test_response") + mock_create_route.return_value = mock_endpoint_func + + # Act + await anthropic_proxy_route( + endpoint="v1/messages", + request=mock_request, + fastapi_response=mock_response, + user_api_key_dict=mock_user_api_key_dict, + ) + + # Assert + mock_create_route.assert_called_once() + call_kwargs = mock_create_route.call_args[1] + + assert call_kwargs["custom_headers"] == {} + assert call_kwargs["_forward_headers"] is True \ No newline at end of file diff --git a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index 395f12ca111..ef733eaa887 100644 --- a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -322,3 +322,207 @@ class TestProxySettingEndpoints: # Verify save_config was called exactly once assert mock_proxy_config["save_call_count"]() == 1 + + def test_update_sso_settings_with_null_values_clears_env_vars( + self, mock_proxy_config, mock_auth, monkeypatch + ): + """Test that updating SSO settings with null values clears environment variables""" + monkeypatch.setenv("LITELLM_SALT_KEY", "test_salt_key") + + # First, verify we have existing environment variables + initial_config = mock_proxy_config["config"] + assert "GOOGLE_CLIENT_ID" in initial_config["environment_variables"] + assert "MICROSOFT_CLIENT_ID" in initial_config["environment_variables"] + + # Set some initial environment variables for runtime testing + monkeypatch.setenv("GOOGLE_CLIENT_ID", "test_existing_google_id") + monkeypatch.setenv("MICROSOFT_CLIENT_ID", "test_existing_microsoft_id") + + # Send SSO settings with null values to clear them + clear_sso_settings = { + "google_client_id": None, + "google_client_secret": None, + "microsoft_client_id": None, + "microsoft_client_secret": None, + "microsoft_tenant": None, + "generic_client_id": None, + "generic_client_secret": None, + "generic_authorization_endpoint": None, + "generic_token_endpoint": None, + "generic_userinfo_endpoint": None, + "proxy_base_url": None, + "user_email": None, + "sso_provider": None, + } + + response = client.patch("/update/sso_settings", json=clear_sso_settings) + + assert response.status_code == 200 + data = response.json() + assert data["status"] == "success" + + # Verify that environment variables were cleared from config + updated_config = mock_proxy_config["config"] + + # These should be removed from environment_variables + assert "GOOGLE_CLIENT_ID" not in updated_config["environment_variables"] + assert "GOOGLE_CLIENT_SECRET" not in updated_config["environment_variables"] + assert "MICROSOFT_CLIENT_ID" not in updated_config["environment_variables"] + assert "MICROSOFT_CLIENT_SECRET" not in updated_config["environment_variables"] + assert "MICROSOFT_TENANT" not in updated_config["environment_variables"] + assert "PROXY_BASE_URL" not in updated_config["environment_variables"] + + # Verify that runtime environment variables were cleared + assert "GOOGLE_CLIENT_ID" not in os.environ + assert "MICROSOFT_CLIENT_ID" not in os.environ + + # Verify user_email was cleared from general_settings + assert updated_config["general_settings"].get("proxy_admin_email") is None + + # Verify save_config was called + assert mock_proxy_config["save_call_count"]() == 1 + + def test_update_sso_settings_with_empty_strings_clears_env_vars( + self, mock_proxy_config, mock_auth, monkeypatch + ): + """Test that updating SSO settings with empty strings also clears environment variables""" + monkeypatch.setenv("LITELLM_SALT_KEY", "test_salt_key") + + # Set some initial environment variables for runtime testing + monkeypatch.setenv("GOOGLE_CLIENT_ID", "test_existing_google_id") + monkeypatch.setenv("MICROSOFT_CLIENT_SECRET", "test_existing_microsoft_secret") + + # Send SSO settings with empty strings to clear them + clear_sso_settings = { + "google_client_id": "", + "google_client_secret": "", + "microsoft_client_secret": "", + "proxy_base_url": "", + "user_email": "", + } + + response = client.patch("/update/sso_settings", json=clear_sso_settings) + + assert response.status_code == 200 + data = response.json() + assert data["status"] == "success" + + # Verify that environment variables with empty strings were cleared from config + updated_config = mock_proxy_config["config"] + assert "GOOGLE_CLIENT_ID" not in updated_config["environment_variables"] + assert "GOOGLE_CLIENT_SECRET" not in updated_config["environment_variables"] + assert "MICROSOFT_CLIENT_SECRET" not in updated_config["environment_variables"] + assert "PROXY_BASE_URL" not in updated_config["environment_variables"] + + # Verify that runtime environment variables were cleared + assert "GOOGLE_CLIENT_ID" not in os.environ + assert "MICROSOFT_CLIENT_SECRET" not in os.environ + + # Verify user_email was cleared from general_settings + assert updated_config["general_settings"].get("proxy_admin_email") is None + + # Verify save_config was called + assert mock_proxy_config["save_call_count"]() == 1 + + def test_update_sso_settings_mixed_null_and_valid_values( + self, mock_proxy_config, mock_auth, monkeypatch + ): + """Test updating SSO settings with mix of null and valid values""" + monkeypatch.setenv("LITELLM_SALT_KEY", "test_salt_key") + + # Set some initial environment variables + monkeypatch.setenv("GOOGLE_CLIENT_ID", "old_google_id") + monkeypatch.setenv("MICROSOFT_CLIENT_ID", "old_microsoft_id") + monkeypatch.setenv("PROXY_BASE_URL", "old_proxy_url") + + # Send mixed SSO settings - some null, some valid + mixed_sso_settings = { + "google_client_id": "new_google_client_id", # Valid value + "google_client_secret": None, # Null to clear + "microsoft_client_id": None, # Null to clear + "microsoft_client_secret": "new_microsoft_secret", # Valid value + "proxy_base_url": "https://newproxy.com", # Valid value + "user_email": None, # Null to clear + } + + response = client.patch("/update/sso_settings", json=mixed_sso_settings) + + assert response.status_code == 200 + data = response.json() + assert data["status"] == "success" + + # Verify the config was updated correctly + updated_config = mock_proxy_config["config"] + + # Valid values should be set + assert ( + updated_config["environment_variables"]["GOOGLE_CLIENT_ID"] + != "new_google_client_id" + ) # Encrypted + assert ( + updated_config["environment_variables"]["MICROSOFT_CLIENT_SECRET"] + != "new_microsoft_secret" + ) # Encrypted + assert ( + updated_config["environment_variables"]["PROXY_BASE_URL"] + != "https://newproxy.com" + ) # Encrypted + + # Null values should be cleared + assert "GOOGLE_CLIENT_SECRET" not in updated_config["environment_variables"] + assert "MICROSOFT_CLIENT_ID" not in updated_config["environment_variables"] + + # Verify runtime environment variables + assert os.environ.get("GOOGLE_CLIENT_ID") == "new_google_client_id" + assert os.environ.get("MICROSOFT_CLIENT_SECRET") == "new_microsoft_secret" + assert "GOOGLE_CLIENT_SECRET" not in os.environ + assert "MICROSOFT_CLIENT_ID" not in os.environ + + # Verify user_email was cleared from general_settings + assert updated_config["general_settings"].get("proxy_admin_email") is None + + # Verify save_config was called + assert mock_proxy_config["save_call_count"]() == 1 + + def test_update_sso_settings_ui_access_mode_handling( + self, mock_proxy_config, mock_auth, monkeypatch + ): + """Test that ui_access_mode is handled correctly in general_settings""" + monkeypatch.setenv("LITELLM_SALT_KEY", "test_salt_key") + + # Test setting ui_access_mode + sso_settings_with_ui_mode = { + "ui_access_mode": "admin_only", + "user_email": "admin@test.com", + } + + response = client.patch("/update/sso_settings", json=sso_settings_with_ui_mode) + + assert response.status_code == 200 + data = response.json() + assert data["status"] == "success" + + # Verify ui_access_mode was set in general_settings (not environment_variables) + updated_config = mock_proxy_config["config"] + assert updated_config["general_settings"]["ui_access_mode"] == "admin_only" + assert ( + updated_config["general_settings"]["proxy_admin_email"] == "admin@test.com" + ) + + # Verify ui_access_mode is NOT in environment_variables + assert "ui_access_mode" not in updated_config["environment_variables"] + + # Test clearing ui_access_mode + clear_ui_mode = {"ui_access_mode": None, "user_email": None} + + response = client.patch("/update/sso_settings", json=clear_ui_mode) + + assert response.status_code == 200 + + # Verify ui_access_mode and user_email were cleared + updated_config = mock_proxy_config["config"] + assert updated_config["general_settings"].get("ui_access_mode") is None + assert updated_config["general_settings"].get("proxy_admin_email") is None + + # Verify save_config was called twice (once for each update) + assert mock_proxy_config["save_call_count"]() == 2 diff --git a/tests/test_litellm/test_uuid_fallback.py b/tests/test_litellm/test_uuid_fallback.py new file mode 100644 index 00000000000..e18bdabda8a --- /dev/null +++ b/tests/test_litellm/test_uuid_fallback.py @@ -0,0 +1,11 @@ +import importlib + + +def test_fastuuid_flag_exposed(): + mod = importlib.import_module("litellm._uuid") + assert hasattr(mod, "FASTUUID_AVAILABLE") + assert hasattr(mod, "uuid4") + # Ensure uuid4 returns something that looks like a UUID string + val = str(mod.uuid4()) + assert isinstance(val, str) + assert len(val) >= 8 diff --git a/ui/litellm-dashboard/src/components/SSOModals.tsx b/ui/litellm-dashboard/src/components/SSOModals.tsx index 4920cec03b3..56ec9393fa4 100644 --- a/ui/litellm-dashboard/src/components/SSOModals.tsx +++ b/ui/litellm-dashboard/src/components/SSOModals.tsx @@ -186,21 +186,21 @@ const SSOModals: React.FC = ({ } try { - // Clear all SSO settings by sending empty values + // Clear all SSO settings const clearSettings = { - google_client_id: '', - google_client_secret: '', - microsoft_client_id: '', - microsoft_client_secret: '', - microsoft_tenant: '', - generic_client_id: '', - generic_client_secret: '', - generic_authorization_endpoint: '', - generic_token_endpoint: '', - generic_userinfo_endpoint: '', - proxy_base_url: '', - user_email: '', - sso_provider: '', + google_client_id: null, + google_client_secret: null, + microsoft_client_id: null, + microsoft_client_secret: null, + microsoft_tenant: null, + generic_client_id: null, + generic_client_secret: null, + generic_authorization_endpoint: null, + generic_token_endpoint: null, + generic_userinfo_endpoint: null, + proxy_base_url: null, + user_email: null, + sso_provider: null, }; await updateSSOSettings(accessToken, clearSettings);