Merge branch 'litellm_october_alexsander_stanging' into litellm_router_index_change

This commit is contained in:
Alexsander Hamir 2025-10-16 09:07:15 -07:00 • committed by GitHub
commit 1846363b94
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 144 additions and 26 deletions

View file

@ -117,10 +117,52 @@ litellm_settings:
```bash
export SSL_CERTIFICATE="/path/to/certificate.pem"
```
</TabItem>
</Tabs>
## 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.
<Tabs>
<TabItem value="sdk" label="SDK">
```python
import litellm
litellm.ssl_ecdh_curve = "X25519" # Disables PQC for better performance
```
</TabItem>
<TabItem value="proxy" label="PROXY">
```yaml
litellm_settings:
ssl_ecdh_curve: "X25519"
```
</TabItem>
<TabItem value="env_var" label="Environment Variables">
```bash
export SSL_ECDH_CURVE="X25519"
```
</TabItem>
</Tabs>
**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:

View file

@ -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

View file

@ -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

View file

@ -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(

View file

@ -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}")

View file

@ -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)}"
)

View file

@ -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