mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge branch 'litellm_october_alexsander_stanging' into litellm_router_index_change
This commit is contained in:
commit
1846363b94
7 changed files with 144 additions and 26 deletions
|
|
@ -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:
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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)}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue