diff --git a/docs/my-website/docs/guides/security_settings.md b/docs/my-website/docs/guides/security_settings.md index 7995f6c3c9c..d6397a7c197 100644 --- a/docs/my-website/docs/guides/security_settings.md +++ b/docs/my-website/docs/guides/security_settings.md @@ -117,10 +117,52 @@ litellm_settings: ```bash export SSL_CERTIFICATE="/path/to/certificate.pem" ``` + -## 5. Use HTTP_PROXY environment variable +## 5. Configure ECDH Curve for SSL/TLS Performance + +The `ssl_ecdh_curve` setting allows you to configure the Elliptic Curve Diffie-Hellman (ECDH) curve used for SSL/TLS key exchange. This is particularly useful for disabling Post-Quantum Cryptography (PQC) to improve performance in environments where PQC is not required. + +**Use Case:** Some OpenSSL 3.x systems enable PQC by default, which can slow down TLS handshakes. Setting the ECDH curve to `X25519` disables PQC and can significantly improve connection performance. + + + + +```python +import litellm +litellm.ssl_ecdh_curve = "X25519" # Disables PQC for better performance +``` + + + + +```yaml +litellm_settings: + ssl_ecdh_curve: "X25519" +``` + + + + +```bash +export SSL_ECDH_CURVE="X25519" +``` + + + + +**Common Valid Curves:** + +- `X25519` - Modern, fast curve (recommended for disabling PQC) +- `prime256v1` - NIST P-256 curve +- `secp384r1` - NIST P-384 curve +- `secp521r1` - NIST P-521 curve + +**Note:** If an invalid curve name is provided or if your Python/OpenSSL version doesn't support this feature, LiteLLM will log a warning and continue with default curves. + +## 6. Use HTTP_PROXY environment variable Both httpx and aiohttp libraries use `urllib.request.getproxies` from environment variables. Before client initialization, you may set proxy (and optional SSL_CERT_FILE) by setting the environment variables: diff --git a/litellm/__init__.py b/litellm/__init__.py index e461c88efd6..95a761cd464 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -263,6 +263,7 @@ use_client: bool = False ssl_verify: Union[str, bool] = True ssl_security_level: Optional[str] = None ssl_certificate: Optional[str] = None +ssl_ecdh_curve: Optional[str] = None # Set to 'X25519' to disable PQC and improve performance disable_streaming_logging: bool = False disable_token_counter: bool = False disable_add_transform_inline_image_block: bool = False diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index a3ad2c67272..accdddbc4dd 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -1,6 +1,7 @@ import asyncio import os import ssl +import sys import time from typing import TYPE_CHECKING, Any, Callable, Dict, List, Mapping, Optional, Union @@ -114,6 +115,28 @@ def get_ssl_configuration( # but falls back to widely compatible ones custom_ssl_context.set_ciphers(DEFAULT_SSL_CIPHERS) + # Configure ECDH curve for key exchange (e.g., to disable PQC and improve performance) + # Set SSL_ECDH_CURVE env var or litellm.ssl_ecdh_curve to 'X25519' to disable PQC + # Common valid curves: X25519, prime256v1, secp384r1, secp521r1 + ssl_ecdh_curve = os.getenv("SSL_ECDH_CURVE", litellm.ssl_ecdh_curve) + if ssl_ecdh_curve and isinstance(ssl_ecdh_curve, str): + try: + custom_ssl_context.set_ecdh_curve(ssl_ecdh_curve) + verbose_logger.debug(f"SSL ECDH curve set to: {ssl_ecdh_curve}") + except AttributeError: + verbose_logger.warning( + f"SSL ECDH curve configuration not supported. " + f"Python version: {sys.version.split()[0]}, OpenSSL version: {ssl.OPENSSL_VERSION}. " + f"Requested curve: {ssl_ecdh_curve}. Continuing with default curves." + ) + except ValueError as e: + # Invalid curve name + verbose_logger.warning( + f"Invalid SSL ECDH curve name: '{ssl_ecdh_curve}'. {e}. " + f"Common valid curves: X25519, prime256v1, secp384r1, secp521r1. " + f"Continuing with default curves (including PQC)." + ) + # Use our custom SSL context instead of the original ssl_verify value return custom_ssl_context diff --git a/litellm/router.py b/litellm/router.py index 82c9d28bb9c..a41b7873298 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -189,7 +189,7 @@ class RoutingArgs(enum.Enum): class Router: - model_names: List = [] + model_names: set = set() cache_responses: Optional[bool] = False default_cache_time_seconds: int = 1 * 60 * 60 # 1 hour tenacity = None @@ -1057,7 +1057,7 @@ class Router: self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs) request_priority = kwargs.get("priority") or self.default_priority - start_time = time.time() + start_time = time.perf_counter() _is_prompt_management_model = self._is_prompt_management_model(model) if _is_prompt_management_model: @@ -1070,7 +1070,7 @@ class Router: response = await self.schedule_acompletion(**kwargs) else: response = await self.async_function_with_fallbacks(**kwargs) - end_time = time.time() + end_time = time.perf_counter() _duration = end_time - start_time asyncio.create_task( self.service_logger_obj.async_service_success_hook( @@ -1245,7 +1245,7 @@ class Router: input_kwargs_for_streaming_fallback["model"] = model parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs) - start_time = time.time() + start_time = time.perf_counter() deployment = await self.async_get_available_deployment( model=model, messages=messages, @@ -1254,7 +1254,7 @@ class Router: ) _timeout_debug_deployment_dict = deployment - end_time = time.time() + end_time = time.perf_counter() _duration = end_time - start_time asyncio.create_task( self.service_logger_obj.async_service_success_hook( @@ -1834,8 +1834,8 @@ class Router: await self.scheduler.add_request(request=item) ## POLL QUEUE - end_time = time.time() + self.timeout - curr_time = time.time() + end_time = time.monotonic() + self.timeout + curr_time = time.monotonic() poll_interval = self.scheduler.polling_interval # poll every 3ms make_request = False @@ -1852,7 +1852,7 @@ class Router: break else: ## ELSE -> loop till default_timeout await asyncio.sleep(poll_interval) - curr_time = time.time() + curr_time = time.monotonic() if make_request: try: @@ -1896,8 +1896,8 @@ class Router: await self.scheduler.add_request(request=item) ## POLL QUEUE - end_time = time.time() + self.timeout - curr_time = time.time() + end_time = time.monotonic() + self.timeout + curr_time = time.monotonic() poll_interval = self.scheduler.polling_interval # poll every 3ms make_request = False @@ -1914,7 +1914,7 @@ class Router: break else: ## ELSE -> loop till default_timeout await asyncio.sleep(poll_interval) - curr_time = time.time() + curr_time = time.monotonic() if make_request: try: @@ -4915,22 +4915,25 @@ class Router: - hash - use hash as id """ - concat_str = model_group + # Optimized: Use list and join instead of string concatenation in loop + # This avoids creating many temporary string objects (O(n) vs O(n²) complexity) + parts = [model_group] for k, v in litellm_params.items(): if isinstance(k, str): - concat_str += k + parts.append(k) elif isinstance(k, dict): - concat_str += json.dumps(k) + parts.append(json.dumps(k)) else: - concat_str += str(k) + parts.append(str(k)) if isinstance(v, str): - concat_str += v + parts.append(v) elif isinstance(v, dict): - concat_str += json.dumps(v) + parts.append(json.dumps(v)) else: - concat_str += str(v) + parts.append(str(v)) + concat_str = "".join(parts) hash_object = hashlib.sha256(concat_str.encode()) return hash_object.hexdigest() @@ -5154,7 +5157,7 @@ class Router: verbose_router_logger.debug( f"\nInitialized Model List {self.get_model_names()}" ) - self.model_names = [m["model_name"] for m in model_list] + self.model_names = {m["model_name"] for m in model_list} # Build model_name index for O(1) lookups self._build_model_name_index(self.model_list) @@ -5360,7 +5363,7 @@ class Router: self._add_model_to_list_and_index_map( model=_deployment, model_id=deployment.model_info.id ) - self.model_names.append(deployment.model_name) + self.model_names.add(deployment.model_name) return deployment def _update_deployment_indices_after_removal( @@ -6281,7 +6284,9 @@ class Router: model_name=model_name, model=model, team_id=team_id ): if model_alias is not None: - alias_model = copy.deepcopy(model) + # Optimized: Use shallow copy since we only modify top-level model_name + # This is much faster than deepcopy for nested dict structures + alias_model = model.copy() alias_model["model_name"] = model_alias returned_models.append(alias_model) else: @@ -6295,7 +6300,8 @@ class Router: model_name=model_name, model=model, team_id=team_id ): if model_alias is not None: - alias_model = copy.deepcopy(model) + # Optimized: Use shallow copy since we only modify top-level model_name + alias_model = model.copy() alias_model["model_name"] = model_alias returned_models.append(alias_model) else: @@ -7094,7 +7100,7 @@ class Router: if isinstance(healthy_deployments, dict): return healthy_deployments - start_time = time.time() + start_time = time.perf_counter() if ( self.routing_strategy == "usage-based-routing-v2" and self.lowesttpm_logger_v2 is not None @@ -7161,7 +7167,7 @@ class Router: f"get_available_deployment for model: {model}, Selected deployment: {self.print_deployment(deployment)} for model: {model}" ) - end_time = time.time() + end_time = time.perf_counter() _duration = end_time - start_time asyncio.create_task( self.service_logger_obj.async_service_success_hook( diff --git a/tests/router_unit_tests/test_router_helper_utils.py b/tests/router_unit_tests/test_router_helper_utils.py index a31c4d8210f..c2339d9eec5 100644 --- a/tests/router_unit_tests/test_router_helper_utils.py +++ b/tests/router_unit_tests/test_router_helper_utils.py @@ -1411,7 +1411,8 @@ def test_generate_model_id_with_deployment_model_name(model_list): "Expected TypeError when model_group is None - this confirms our fix is needed" ) except TypeError as e: - assert "unsupported operand type(s) for +=" in str(e) + # After optimization, error message changed but still fails appropriately on None + assert "unsupported operand type(s) for +=" in str(e) or "expected str instance, NoneType found" in str(e) print(f"✓ Correctly failed with None model_group (as expected): {e}") except Exception as e: pytest.fail(f"Unexpected error with None model_group: {e}") diff --git a/tests/router_unit_tests/test_router_index_management.py b/tests/router_unit_tests/test_router_index_management.py index 28a48604a01..239ee18afe2 100644 --- a/tests/router_unit_tests/test_router_index_management.py +++ b/tests/router_unit_tests/test_router_index_management.py @@ -265,3 +265,10 @@ class TestRouterIndexManagement: f"ALLOWED_METHODS in this test method.\n" f"{'='*70}\n" ) + def test_model_names_is_set(self): + """Verify that model_names uses a set for O(1) lookups, not a list (O(n))""" + router = Router(model_list=[]) + + assert isinstance(router.model_names, set), ( + f"model_names should be a set for O(1) lookups, but got {type(router.model_names)}" + ) diff --git a/tests/test_litellm/llms/custom_httpx/test_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_http_handler.py index 0e31699fd83..09fc31d18b7 100644 --- a/tests/test_litellm/llms/custom_httpx/test_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_http_handler.py @@ -399,3 +399,41 @@ async def test_session_validation(): mock_valid_session = MockClientSession() transport3 = AsyncHTTPHandler._create_aiohttp_transport(shared_session=mock_valid_session) # type: ignore assert transport3.client is mock_valid_session # Should reuse session + + +@pytest.mark.parametrize( + "env_curve,litellm_curve,expected_curve,should_call", + [ + # env_curve: SSL_ECDH_CURVE env var | litellm_curve: litellm.ssl_ecdh_curve variable + # expected_curve: curve that should be set | should_call: whether set_ecdh_curve() should be called + + # Valid configurations + ("X25519", None, "X25519", True), # Env var only + ("prime256v1", None, "prime256v1", True), # Different valid curve + (None, "secp384r1", "secp384r1", True), # litellm variable only + ("X25519", "secp521r1", "X25519", True), # Env var takes precedence + # Empty/None configurations - should skip + ("", None, None, False), # Empty string - skip configuration + (None, None, None, False), # None value - skip configuration + ] +) +def test_ssl_ecdh_curve(env_curve, litellm_curve, expected_curve, should_call, monkeypatch): + """Test SSL ECDH curve configuration with valid curves and precedence""" + with patch.dict(os.environ, clear=True): + if env_curve: + monkeypatch.setenv("SSL_ECDH_CURVE", env_curve) + + original_value = litellm.ssl_ecdh_curve + try: + litellm.ssl_ecdh_curve = litellm_curve + + with patch.object(ssl.SSLContext, 'set_ecdh_curve') as mock_set_curve: + ssl_context = get_ssl_configuration() + + if should_call: + mock_set_curve.assert_called_once_with(expected_curve) + else: + mock_set_curve.assert_not_called() + assert isinstance(ssl_context, ssl.SSLContext) + finally: + litellm.ssl_ecdh_curve = original_value