mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix: critical bugs in GIL-removal branch
- orjson: wrap imports with try/except fallback in core library files (llm_http_handler, openai_like handler/transformation) so litellm works without proxy extras installed - sidecar: support all HTTP methods (GET/PUT/PATCH/DELETE), not just POST - sidecar: forward all provider headers, not just 3 hardcoded ones - sidecar: fix stream error swallowing (log + empty frame vs silent drop) - sidecar: fix DashMap TOCTOU race with entry() API - sidecar transport: fix aclose() to use release() instead of sync close() - sidecar client: capture subprocess stdout/stderr instead of /dev/null - sidecar client: use lazy %-formatting in logger calls - mock server: pre-serialize response for lower overhead in load tests Co-authored-by: Krish Dholakia <krrishdholakia@gmail.com>
This commit is contained in:
parent
dde4042e2d
commit
c098e56082
18 changed files with 220 additions and 260 deletions
|
|
@ -33,18 +33,18 @@ impl Sidecar {
|
|||
}
|
||||
|
||||
fn get_or_create_client(&self, host: &str) -> Client {
|
||||
if let Some(client) = self.pools.get(host) {
|
||||
return client.clone();
|
||||
}
|
||||
let client = Client::builder()
|
||||
.pool_max_idle_per_host(200)
|
||||
.pool_idle_timeout(std::time::Duration::from_secs(90))
|
||||
.tcp_keepalive(std::time::Duration::from_secs(60))
|
||||
.tcp_nodelay(true)
|
||||
.build()
|
||||
.expect("Failed to build reqwest client");
|
||||
self.pools.insert(host.to_string(), client.clone());
|
||||
client
|
||||
self.pools
|
||||
.entry(host.to_string())
|
||||
.or_insert_with(|| {
|
||||
Client::builder()
|
||||
.pool_max_idle_per_host(200)
|
||||
.pool_idle_timeout(std::time::Duration::from_secs(90))
|
||||
.tcp_keepalive(std::time::Duration::from_secs(60))
|
||||
.tcp_nodelay(true)
|
||||
.build()
|
||||
.expect("Failed to build reqwest client")
|
||||
})
|
||||
.clone()
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -117,6 +117,13 @@ async fn handle_request(
|
|||
.map(|v| v == "true")
|
||||
.unwrap_or(false);
|
||||
|
||||
let method_str = req
|
||||
.headers()
|
||||
.get("x-litellm-method")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.unwrap_or("POST")
|
||||
.to_uppercase();
|
||||
|
||||
let content_type = req
|
||||
.headers()
|
||||
.get("content-type")
|
||||
|
|
@ -157,11 +164,21 @@ async fn handle_request(
|
|||
let client = sidecar.get_or_create_client(&host);
|
||||
let full_url = format!("{}{}", provider_url.trim_end_matches('/'), request_path);
|
||||
|
||||
let mut req_builder = client
|
||||
.post(&full_url)
|
||||
let mut req_builder = match method_str.as_str() {
|
||||
"GET" => client.get(&full_url),
|
||||
"PUT" => client.put(&full_url),
|
||||
"PATCH" => client.patch(&full_url),
|
||||
"DELETE" => client.delete(&full_url),
|
||||
_ => client.post(&full_url),
|
||||
};
|
||||
|
||||
req_builder = req_builder
|
||||
.header("content-type", &content_type)
|
||||
.timeout(std::time::Duration::from_secs(timeout_secs))
|
||||
.body(body_bytes.to_vec());
|
||||
.timeout(std::time::Duration::from_secs(timeout_secs));
|
||||
|
||||
if method_str != "GET" {
|
||||
req_builder = req_builder.body(body_bytes.to_vec());
|
||||
}
|
||||
|
||||
if !api_key.is_empty() {
|
||||
req_builder = req_builder.header("authorization", format!("Bearer {}", api_key));
|
||||
|
|
@ -197,7 +214,13 @@ async fn handle_request(
|
|||
|
||||
if is_stream {
|
||||
let stream = provider_response.bytes_stream().map(|result| {
|
||||
Ok::<_, Infallible>(Frame::data(result.unwrap_or_default()))
|
||||
match result {
|
||||
Ok(bytes) => Ok::<_, Infallible>(Frame::data(bytes)),
|
||||
Err(e) => {
|
||||
tracing::warn!("Stream chunk error: {}", e);
|
||||
Ok(Frame::data(Bytes::new()))
|
||||
}
|
||||
}
|
||||
});
|
||||
let elapsed = start.elapsed().as_micros() as u64;
|
||||
sidecar.total_latency_us.fetch_add(elapsed, Ordering::Relaxed);
|
||||
|
|
|
|||
|
|
@ -887,7 +887,7 @@ class PrometheusLogger(CustomLogger):
|
|||
from litellm.types.utils import StandardLoggingPayload
|
||||
|
||||
verbose_logger.debug(
|
||||
"prometheus Logging - Enters success logging function"
|
||||
f"prometheus Logging - Enters success logging function for kwargs {kwargs}"
|
||||
)
|
||||
|
||||
# unpack kwargs
|
||||
|
|
@ -943,10 +943,9 @@ class PrometheusLogger(CustomLogger):
|
|||
else:
|
||||
_tags = []
|
||||
|
||||
if litellm.set_verbose:
|
||||
print_verbose(
|
||||
f"inside track_prometheus_metrics, model {model}, response_cost {response_cost}, tokens_used {tokens_used}"
|
||||
)
|
||||
print_verbose(
|
||||
f"inside track_prometheus_metrics, model {model}, response_cost {response_cost}, tokens_used {tokens_used}, end_user_id {end_user_id}, user_api_key {user_api_key}"
|
||||
)
|
||||
|
||||
enum_values = UserAPIKeyLabelValues(
|
||||
end_user=end_user_id,
|
||||
|
|
@ -3057,13 +3056,15 @@ def prometheus_label_factory(
|
|||
|
||||
Ensures end_user param is not sent to prometheus if it is not supported.
|
||||
"""
|
||||
enum_dict = enum_values.get_label_dict()
|
||||
# Extract dictionary from Pydantic object
|
||||
enum_dict = enum_values.model_dump()
|
||||
|
||||
supported_set = frozenset(supported_enum_labels) if not isinstance(supported_enum_labels, (set, frozenset)) else supported_enum_labels
|
||||
# Filter supported labels and sanitize values to prevent breaking
|
||||
# the Prometheus text format (e.g. U+2028 Line Separator in label values)
|
||||
filtered_labels = {
|
||||
label: _sanitize_prometheus_label_value(value)
|
||||
for label, value in enum_dict.items()
|
||||
if label in supported_set
|
||||
if label in supported_enum_labels
|
||||
}
|
||||
|
||||
if UserAPIKeyLabelNames.END_USER.value in filtered_labels:
|
||||
|
|
@ -3078,14 +3079,14 @@ def prometheus_label_factory(
|
|||
for key, value in enum_values.custom_metadata_labels.items():
|
||||
# check sanitized key
|
||||
sanitized_key = _sanitize_prometheus_label_name(key)
|
||||
if sanitized_key in supported_set:
|
||||
if sanitized_key in supported_enum_labels:
|
||||
filtered_labels[sanitized_key] = _sanitize_prometheus_label_value(value)
|
||||
|
||||
# Add custom tags if configured
|
||||
if enum_values.tags is not None:
|
||||
custom_tag_labels = get_custom_labels_from_tags(enum_values.tags)
|
||||
for key, value in custom_tag_labels.items():
|
||||
if key in supported_set:
|
||||
if key in supported_enum_labels:
|
||||
filtered_labels[key] = _sanitize_prometheus_label_value(value)
|
||||
|
||||
for label in supported_enum_labels:
|
||||
|
|
|
|||
|
|
@ -1476,7 +1476,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
**response_cost_calculator_kwargs
|
||||
)
|
||||
|
||||
verbose_logger.debug("response_cost: %s", response_cost)
|
||||
verbose_logger.debug(f"response_cost: {response_cost}")
|
||||
return response_cost
|
||||
except Exception as e: # error calculating cost
|
||||
debug_info = StandardLoggingModelCostFailureDebugInformation(
|
||||
|
|
|
|||
|
|
@ -387,12 +387,11 @@ class StandardBuiltInToolCostTracking:
|
|||
message: Optional[Message] = getattr(choice, "message", None)
|
||||
if message is None:
|
||||
continue
|
||||
annotations = getattr(message, "annotations", None)
|
||||
if annotations:
|
||||
for annotation in annotations:
|
||||
_type = annotation.get("type") if isinstance(annotation, dict) else getattr(annotation, "type", None)
|
||||
if _type == annotation_type:
|
||||
return True
|
||||
if annotations := getattr(message, "annotations", None):
|
||||
if len(annotations) > 0:
|
||||
for annotation in annotations:
|
||||
if annotation.get("type", None) == annotation_type:
|
||||
return True
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -5,13 +5,6 @@ from pydantic import BaseModel
|
|||
|
||||
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
|
||||
|
||||
try:
|
||||
import orjson
|
||||
|
||||
_has_orjson = True
|
||||
except ImportError:
|
||||
_has_orjson = False
|
||||
|
||||
|
||||
def safe_dumps(data: Any, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH) -> str:
|
||||
"""
|
||||
|
|
@ -20,10 +13,13 @@ def safe_dumps(data: Any, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH) -> str:
|
|||
"""
|
||||
|
||||
def _serialize(obj: Any, seen: set, depth: int) -> Any:
|
||||
# Check for maximum depth.
|
||||
if depth > max_depth:
|
||||
return "MaxDepthExceeded"
|
||||
# Base-case: if it is a primitive, simply return it.
|
||||
if isinstance(obj, (str, int, float, bool, type(None))):
|
||||
return obj
|
||||
# Check for circular reference.
|
||||
if id(obj) in seen:
|
||||
return "CircularReference Detected"
|
||||
seen.add(id(obj))
|
||||
|
|
@ -31,7 +27,7 @@ def safe_dumps(data: Any, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH) -> str:
|
|||
if isinstance(obj, dict):
|
||||
result = {}
|
||||
for k, v in obj.items():
|
||||
if isinstance(k, str):
|
||||
if isinstance(k, (str)):
|
||||
result[k] = _serialize(v, seen, depth + 1)
|
||||
seen.remove(id(obj))
|
||||
return result
|
||||
|
|
@ -53,12 +49,11 @@ def safe_dumps(data: Any, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH) -> str:
|
|||
seen.remove(id(obj))
|
||||
return result
|
||||
else:
|
||||
# Fall back to string conversion for non-serializable objects.
|
||||
try:
|
||||
return str(obj)
|
||||
except Exception:
|
||||
return "Unserializable Object"
|
||||
|
||||
safe_data = _serialize(data, set(), 0)
|
||||
if _has_orjson:
|
||||
return orjson.dumps(safe_data, default=str).decode()
|
||||
return json.dumps(safe_data, default=str)
|
||||
|
|
|
|||
|
|
@ -1,5 +1,4 @@
|
|||
import asyncio
|
||||
import functools
|
||||
import os
|
||||
import ssl
|
||||
import sys
|
||||
|
|
@ -52,21 +51,6 @@ try:
|
|||
except Exception:
|
||||
version = "0.0.0"
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=64)
|
||||
def _parse_url(url: str) -> httpx.URL:
|
||||
"""Pre-parse a URL string into an httpx.URL to avoid regex-heavy
|
||||
parsing inside httpx._merge_url on every request (~7μs → ~0.4μs).
|
||||
|
||||
Safe to use with ``build_request(params=...)``: httpx replaces the
|
||||
query string entirely when ``params`` is non-None, so any query
|
||||
params baked into the cached URL are harmless in that case. When
|
||||
``params`` is None the cached URL preserves the original query
|
||||
string, which is the correct behaviour.
|
||||
"""
|
||||
return httpx.URL(url)
|
||||
|
||||
|
||||
def get_default_headers() -> dict:
|
||||
"""
|
||||
Get default headers for HTTP requests.
|
||||
|
|
@ -440,7 +424,7 @@ class AsyncHTTPHandler:
|
|||
params.update(HTTPHandler.extract_query_params(url))
|
||||
|
||||
response = await self.client.get(
|
||||
_parse_url(url), params=params, headers=headers, follow_redirects=_follow_redirects # type: ignore
|
||||
url, params=params, headers=headers, follow_redirects=_follow_redirects # type: ignore
|
||||
)
|
||||
return response
|
||||
|
||||
|
|
@ -468,10 +452,9 @@ class AsyncHTTPHandler:
|
|||
data, content
|
||||
)
|
||||
|
||||
parsed_url = _parse_url(url)
|
||||
req = self.client.build_request(
|
||||
"POST",
|
||||
parsed_url,
|
||||
url,
|
||||
data=request_data,
|
||||
json=json,
|
||||
params=params,
|
||||
|
|
@ -550,7 +533,7 @@ class AsyncHTTPHandler:
|
|||
)
|
||||
|
||||
req = self.client.build_request(
|
||||
"PUT", _parse_url(url), data=request_data, json=json, params=params, headers=headers, timeout=timeout, content=request_content # type: ignore
|
||||
"PUT", url, data=request_data, json=json, params=params, headers=headers, timeout=timeout, content=request_content # type: ignore
|
||||
)
|
||||
response = await self.client.send(req)
|
||||
response.raise_for_status()
|
||||
|
|
@ -616,7 +599,7 @@ class AsyncHTTPHandler:
|
|||
)
|
||||
|
||||
req = self.client.build_request(
|
||||
"PATCH", _parse_url(url), data=request_data, json=json, params=params, headers=headers, timeout=timeout, content=request_content # type: ignore
|
||||
"PATCH", url, data=request_data, json=json, params=params, headers=headers, timeout=timeout, content=request_content # type: ignore
|
||||
)
|
||||
response = await self.client.send(req)
|
||||
response.raise_for_status()
|
||||
|
|
@ -682,7 +665,7 @@ class AsyncHTTPHandler:
|
|||
)
|
||||
|
||||
req = self.client.build_request(
|
||||
"DELETE", _parse_url(url), data=request_data, json=json, params=params, headers=headers, timeout=timeout, content=request_content # type: ignore
|
||||
"DELETE", url, data=request_data, json=json, params=params, headers=headers, timeout=timeout, content=request_content # type: ignore
|
||||
)
|
||||
response = await self.client.send(req, stream=stream)
|
||||
response.raise_for_status()
|
||||
|
|
@ -734,7 +717,7 @@ class AsyncHTTPHandler:
|
|||
request_data, request_content = _prepare_request_data_and_content(data, content)
|
||||
|
||||
req = client.build_request(
|
||||
"POST", _parse_url(url), data=request_data, json=json, params=params, headers=headers, content=request_content # type: ignore
|
||||
"POST", url, data=request_data, json=json, params=params, headers=headers, content=request_content # type: ignore
|
||||
)
|
||||
response = await client.send(req, stream=stream)
|
||||
response.raise_for_status()
|
||||
|
|
@ -754,17 +737,18 @@ class AsyncHTTPHandler:
|
|||
) -> Optional[Union[LiteLLMAiohttpTransport, AsyncHTTPTransport]]:
|
||||
"""
|
||||
- Creates a transport for httpx.AsyncClient
|
||||
- if USE_SIDECAR is set, use the Rust sidecar transport
|
||||
- if litellm.force_ipv4 is True, it will return AsyncHTTPTransport with local_address="0.0.0.0"
|
||||
- [Default] It will return AiohttpTransport
|
||||
- Users can opt out of using AiohttpTransport by setting litellm.use_aiohttp_transport to False
|
||||
"""
|
||||
#########################################################
|
||||
# SIDECAR TRANSPORT — Rust-based HTTP forwarding
|
||||
#########################################################
|
||||
if AsyncHTTPHandler._should_use_sidecar_transport():
|
||||
return AsyncHTTPHandler._create_sidecar_transport()
|
||||
|
||||
|
||||
Notes on this handler:
|
||||
- Why AiohttpTransport?
|
||||
- By default, we use AiohttpTransport since it offers much higher throughput and lower latency than httpx.
|
||||
|
||||
- Why force ipv4?
|
||||
- Some users have seen httpx ConnectionError when using ipv6 - forcing ipv4 resolves the issue for them
|
||||
"""
|
||||
#########################################################
|
||||
# AIOHTTP TRANSPORT is off by default
|
||||
#########################################################
|
||||
|
|
@ -780,23 +764,6 @@ class AsyncHTTPHandler:
|
|||
#########################################################
|
||||
return AsyncHTTPHandler._create_httpx_transport()
|
||||
|
||||
@staticmethod
|
||||
def _should_use_sidecar_transport() -> bool:
|
||||
"""Check if the Rust sidecar transport is enabled via env var."""
|
||||
return os.getenv("USE_SIDECAR", "").lower() == "true"
|
||||
|
||||
@staticmethod
|
||||
def _create_sidecar_transport():
|
||||
"""Create a sidecar transport that forwards requests through the Rust binary."""
|
||||
from litellm.llms.custom_httpx.sidecar_transport import (
|
||||
LiteLLMSidecarTransport,
|
||||
)
|
||||
|
||||
port = int(os.getenv("SIDECAR_PORT", "8787"))
|
||||
sidecar_url = f"http://127.0.0.1:{port}"
|
||||
verbose_logger.info("Using Sidecar transport → %s", sidecar_url)
|
||||
return LiteLLMSidecarTransport(sidecar_url=sidecar_url)
|
||||
|
||||
@staticmethod
|
||||
def _should_use_aiohttp_transport() -> bool:
|
||||
"""
|
||||
|
|
@ -1001,7 +968,7 @@ class HTTPHandler:
|
|||
params.update(self.extract_query_params(url))
|
||||
|
||||
response = self.client.get(
|
||||
_parse_url(url),
|
||||
url,
|
||||
params=params,
|
||||
headers=headers,
|
||||
)
|
||||
|
|
@ -1040,11 +1007,10 @@ class HTTPHandler:
|
|||
data, content
|
||||
)
|
||||
|
||||
parsed_url = _parse_url(url)
|
||||
if timeout is not None:
|
||||
req = self.client.build_request(
|
||||
"POST",
|
||||
parsed_url,
|
||||
url,
|
||||
data=request_data, # type: ignore
|
||||
json=json,
|
||||
params=params,
|
||||
|
|
@ -1055,7 +1021,7 @@ class HTTPHandler:
|
|||
)
|
||||
else:
|
||||
req = self.client.build_request(
|
||||
"POST", parsed_url, data=request_data, json=json, params=params, headers=headers, files=files, content=request_content # type: ignore
|
||||
"POST", url, data=request_data, json=json, params=params, headers=headers, files=files, content=request_content # type: ignore
|
||||
)
|
||||
response = self.client.send(req, stream=stream)
|
||||
response.raise_for_status()
|
||||
|
|
@ -1097,14 +1063,13 @@ class HTTPHandler:
|
|||
data, content
|
||||
)
|
||||
|
||||
parsed_url = _parse_url(url)
|
||||
if timeout is not None:
|
||||
req = self.client.build_request(
|
||||
"PATCH", parsed_url, data=request_data, json=json, params=params, headers=headers, timeout=timeout, content=request_content # type: ignore
|
||||
"PATCH", url, data=request_data, json=json, params=params, headers=headers, timeout=timeout, content=request_content # type: ignore
|
||||
)
|
||||
else:
|
||||
req = self.client.build_request(
|
||||
"PATCH", parsed_url, data=request_data, json=json, params=params, headers=headers, content=request_content # type: ignore
|
||||
"PATCH", url, data=request_data, json=json, params=params, headers=headers, content=request_content # type: ignore
|
||||
)
|
||||
response = self.client.send(req, stream=stream)
|
||||
response.raise_for_status()
|
||||
|
|
@ -1147,14 +1112,13 @@ class HTTPHandler:
|
|||
data, content
|
||||
)
|
||||
|
||||
parsed_url = _parse_url(url)
|
||||
if timeout is not None:
|
||||
req = self.client.build_request(
|
||||
"PUT", parsed_url, data=request_data, json=json, params=params, headers=headers, timeout=timeout, content=request_content # type: ignore
|
||||
"PUT", url, data=request_data, json=json, params=params, headers=headers, timeout=timeout, content=request_content # type: ignore
|
||||
)
|
||||
else:
|
||||
req = self.client.build_request(
|
||||
"PUT", parsed_url, data=request_data, json=json, params=params, headers=headers, content=request_content # type: ignore
|
||||
"PUT", url, data=request_data, json=json, params=params, headers=headers, content=request_content # type: ignore
|
||||
)
|
||||
response = self.client.send(req, stream=stream)
|
||||
return response
|
||||
|
|
@ -1184,14 +1148,13 @@ class HTTPHandler:
|
|||
data, content
|
||||
)
|
||||
|
||||
parsed_url = _parse_url(url)
|
||||
if timeout is not None:
|
||||
req = self.client.build_request(
|
||||
"DELETE", parsed_url, data=request_data, json=json, params=params, headers=headers, timeout=timeout, content=request_content # type: ignore
|
||||
"DELETE", url, data=request_data, json=json, params=params, headers=headers, timeout=timeout, content=request_content # type: ignore
|
||||
)
|
||||
else:
|
||||
req = self.client.build_request(
|
||||
"DELETE", parsed_url, data=request_data, json=json, params=params, headers=headers, content=request_content # type: ignore
|
||||
"DELETE", url, data=request_data, json=json, params=params, headers=headers, content=request_content # type: ignore
|
||||
)
|
||||
response = self.client.send(req, stream=stream)
|
||||
response.raise_for_status()
|
||||
|
|
|
|||
|
|
@ -1,6 +1,11 @@
|
|||
import json
|
||||
import orjson
|
||||
import ssl
|
||||
|
||||
try:
|
||||
import orjson
|
||||
_has_orjson = True
|
||||
except ImportError:
|
||||
_has_orjson = False
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
|
|
@ -181,7 +186,7 @@ class BaseLLMHTTPHandler:
|
|||
data=(
|
||||
signed_json_body
|
||||
if signed_json_body is not None
|
||||
else orjson.dumps(data)
|
||||
else (orjson.dumps(data) if _has_orjson else json.dumps(data).encode())
|
||||
),
|
||||
timeout=timeout,
|
||||
stream=stream,
|
||||
|
|
@ -241,7 +246,7 @@ class BaseLLMHTTPHandler:
|
|||
data=(
|
||||
signed_json_body
|
||||
if signed_json_body is not None
|
||||
else orjson.dumps(data)
|
||||
else (orjson.dumps(data) if _has_orjson else json.dumps(data).encode())
|
||||
),
|
||||
timeout=timeout,
|
||||
stream=stream,
|
||||
|
|
|
|||
|
|
@ -27,7 +27,7 @@ class SidecarResponseStream(httpx.AsyncByteStream):
|
|||
yield chunk
|
||||
|
||||
async def aclose(self) -> None:
|
||||
self._response.close()
|
||||
self._response.release()
|
||||
|
||||
|
||||
class LiteLLMSidecarTransport(httpx.AsyncBaseTransport):
|
||||
|
|
@ -94,11 +94,17 @@ class LiteLLMSidecarTransport(httpx.AsyncBaseTransport):
|
|||
"Content-Type": request.headers.get("content-type", "application/json"),
|
||||
}
|
||||
|
||||
# Copy provider-specific headers the sidecar should forward
|
||||
for key in ("x-api-key", "anthropic-version", "x-goog-api-key"):
|
||||
val = request.headers.get(key)
|
||||
if val:
|
||||
headers[f"X-LiteLLM-Fwd-{key}"] = val
|
||||
_skip_headers = frozenset({
|
||||
"host", "content-length", "transfer-encoding", "connection",
|
||||
"authorization", "content-type", "accept", "user-agent",
|
||||
"accept-encoding",
|
||||
})
|
||||
for key, val in request.headers.items():
|
||||
lower = key.lower()
|
||||
if lower not in _skip_headers and not lower.startswith("x-litellm-"):
|
||||
headers[f"X-LiteLLM-Fwd-{lower}"] = val
|
||||
|
||||
headers["X-LiteLLM-Method"] = str(request.method)
|
||||
|
||||
session = self._get_session()
|
||||
resp = await session.post(
|
||||
|
|
|
|||
|
|
@ -4,9 +4,15 @@ OpenAI-like chat completion handler
|
|||
For handling OpenAI-like chat completions, like IBM WatsonX, etc.
|
||||
"""
|
||||
|
||||
import orjson
|
||||
import json
|
||||
from typing import Any, Callable, Optional, Union
|
||||
|
||||
try:
|
||||
import orjson
|
||||
_has_orjson = True
|
||||
except ImportError:
|
||||
_has_orjson = False
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
|
|
@ -139,7 +145,7 @@ class OpenAILikeChatHandler(OpenAILikeBase):
|
|||
client=client,
|
||||
api_base=api_base,
|
||||
headers=headers,
|
||||
data=orjson.dumps(data),
|
||||
data=(orjson.dumps(data) if _has_orjson else json.dumps(data).encode()),
|
||||
model=model,
|
||||
messages=messages,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -185,7 +191,7 @@ class OpenAILikeChatHandler(OpenAILikeBase):
|
|||
|
||||
try:
|
||||
response = await client.post(
|
||||
api_base, headers=headers, data=orjson.dumps(data), timeout=timeout
|
||||
api_base, headers=headers, data=(orjson.dumps(data) if _has_orjson else json.dumps(data).encode()), timeout=timeout
|
||||
)
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as e:
|
||||
|
|
@ -350,7 +356,7 @@ class OpenAILikeChatHandler(OpenAILikeBase):
|
|||
),
|
||||
api_base=api_base,
|
||||
headers=headers,
|
||||
data=orjson.dumps(data),
|
||||
data=(orjson.dumps(data) if _has_orjson else json.dumps(data).encode()),
|
||||
model=model,
|
||||
messages=messages,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -370,7 +376,7 @@ class OpenAILikeChatHandler(OpenAILikeBase):
|
|||
client = HTTPHandler(timeout=timeout) # type: ignore
|
||||
try:
|
||||
response = client.post(
|
||||
url=api_base, headers=headers, data=orjson.dumps(data)
|
||||
url=api_base, headers=headers, data=(orjson.dumps(data) if _has_orjson else json.dumps(data).encode())
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
|
|
|
|||
|
|
@ -5,7 +5,12 @@ OpenAI-like chat completion transformation
|
|||
from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union
|
||||
|
||||
import httpx
|
||||
import orjson
|
||||
|
||||
try:
|
||||
import orjson
|
||||
_has_orjson = True
|
||||
except ImportError:
|
||||
_has_orjson = False
|
||||
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionAssistantMessage
|
||||
|
|
@ -95,7 +100,7 @@ class OpenAILikeChatConfig(OpenAIGPTConfig):
|
|||
custom_llm_provider: Optional[str],
|
||||
base_model: Optional[str],
|
||||
) -> ModelResponse:
|
||||
response_json = orjson.loads(response.content)
|
||||
response_json = orjson.loads(response.content) if _has_orjson else response.json()
|
||||
logging_obj.post_call(
|
||||
input=messages,
|
||||
api_key="",
|
||||
|
|
|
|||
|
|
@ -102,15 +102,10 @@ def get_chain_id_from_headers(headers: Optional[Dict[str, str]]) -> Optional[str
|
|||
"""
|
||||
if not headers:
|
||||
return None
|
||||
session_id = None
|
||||
for k, v in headers.items():
|
||||
if isinstance(k, str):
|
||||
k_lower = k.lower()
|
||||
if k_lower == "x-litellm-trace-id":
|
||||
return v
|
||||
elif session_id is None and k_lower == "x-litellm-session-id":
|
||||
session_id = v
|
||||
return session_id
|
||||
normalized = {k.lower(): v for k, v in headers.items() if isinstance(k, str)}
|
||||
return normalized.get("x-litellm-trace-id") or normalized.get(
|
||||
"x-litellm-session-id"
|
||||
)
|
||||
|
||||
|
||||
def safe_add_api_version_from_query_params(data: dict, request: Request):
|
||||
|
|
|
|||
|
|
@ -948,33 +948,9 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915
|
|||
## Initialize shared aiohttp session for connection reuse
|
||||
shared_aiohttp_session = await _initialize_shared_aiohttp_session()
|
||||
|
||||
## Initialize Rust sidecar client (optional, for high-perf forwarding)
|
||||
_use_sidecar = os.environ.get("USE_SIDECAR", "").lower() == "true" or general_settings.get("use_sidecar", False)
|
||||
if _use_sidecar:
|
||||
from litellm.proxy.sidecar_client import init_sidecar_client
|
||||
|
||||
_sidecar_port = int(os.environ.get("SIDECAR_PORT", general_settings.get("sidecar_port", 8787)))
|
||||
_sidecar_binary = os.environ.get("SIDECAR_BINARY", general_settings.get("sidecar_binary", ""))
|
||||
await init_sidecar_client(
|
||||
port=_sidecar_port,
|
||||
binary=_sidecar_binary or None,
|
||||
auto_start=bool(_sidecar_binary),
|
||||
)
|
||||
|
||||
# End of startup event
|
||||
yield
|
||||
|
||||
# Shutdown event - close sidecar client
|
||||
try:
|
||||
from litellm.proxy.sidecar_client import get_sidecar_client
|
||||
|
||||
_sc = get_sidecar_client()
|
||||
if _sc is not None:
|
||||
await _sc.close()
|
||||
verbose_proxy_logger.info("Sidecar client closed")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Shutdown event - close shared aiohttp session
|
||||
if shared_aiohttp_session is not None:
|
||||
try:
|
||||
|
|
@ -6639,25 +6615,21 @@ async def model_info(
|
|||
"/v1/chat/completions",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
tags=["chat/completions"],
|
||||
response_class=ORJSONResponse,
|
||||
)
|
||||
@router.post(
|
||||
"/chat/completions",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
tags=["chat/completions"],
|
||||
response_class=ORJSONResponse,
|
||||
)
|
||||
@router.post(
|
||||
"/engines/{model:path}/chat/completions",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
tags=["chat/completions"],
|
||||
response_class=ORJSONResponse,
|
||||
)
|
||||
@router.post(
|
||||
"/openai/deployments/{model:path}/chat/completions",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
tags=["chat/completions"],
|
||||
response_class=ORJSONResponse,
|
||||
responses={200: {"description": "Successful response"}, **ERROR_RESPONSES},
|
||||
) # azure compatible endpoint
|
||||
async def chat_completion( # noqa: PLR0915
|
||||
|
|
|
|||
|
|
@ -59,11 +59,11 @@ class SidecarClient:
|
|||
self._healthy = await self._check_health()
|
||||
if self._healthy:
|
||||
verbose_proxy_logger.info(
|
||||
f"Sidecar client connected to {self.sidecar_url}"
|
||||
"Sidecar client connected to %s", self.sidecar_url
|
||||
)
|
||||
else:
|
||||
verbose_proxy_logger.warning(
|
||||
f"Sidecar not available at {self.sidecar_url}, will use fallback"
|
||||
"Sidecar not available at %s, will use fallback", self.sidecar_url
|
||||
)
|
||||
|
||||
async def _start_sidecar(self):
|
||||
|
|
@ -74,8 +74,8 @@ class SidecarClient:
|
|||
self._process = subprocess.Popen(
|
||||
[self.sidecar_binary],
|
||||
env=env,
|
||||
stdout=subprocess.DEVNULL,
|
||||
stderr=subprocess.DEVNULL,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
)
|
||||
for _ in range(50):
|
||||
await asyncio.sleep(0.1)
|
||||
|
|
|
|||
|
|
@ -632,37 +632,61 @@ def _sanitize_request_body_for_spend_logs_payload(
|
|||
Recursively sanitize request body to prevent logging large base64 strings or other large values.
|
||||
Truncates strings longer than MAX_STRING_LENGTH_PROMPT_IN_DB characters and handles nested dictionaries.
|
||||
"""
|
||||
from litellm.constants import (
|
||||
LITELLM_TRUNCATED_PAYLOAD_FIELD,
|
||||
LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE,
|
||||
)
|
||||
|
||||
if visited is None:
|
||||
visited = set()
|
||||
if max_string_length_prompt_in_db is None:
|
||||
max_string_length_prompt_in_db = _get_max_string_length_prompt_in_db()
|
||||
|
||||
# Get the object's memory address to track visited objects
|
||||
obj_id = id(request_body)
|
||||
if obj_id in visited:
|
||||
return {}
|
||||
visited.add(obj_id)
|
||||
|
||||
_max_len = max_string_length_prompt_in_db
|
||||
_start_chars = int(_max_len * 0.35)
|
||||
_end_chars = min(int(_max_len * 0.65), _max_len - _start_chars)
|
||||
|
||||
def _sanitize_value(value: Any) -> Any:
|
||||
if isinstance(value, str):
|
||||
if len(value) > _max_len:
|
||||
skipped_chars = len(value) - _start_chars - _end_chars
|
||||
return (
|
||||
f"{value[:_start_chars]}"
|
||||
f"... ({LITELLM_TRUNCATED_PAYLOAD_FIELD} skipped {skipped_chars} chars. "
|
||||
f"{LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE}) ..."
|
||||
f"{value[-_end_chars:]}"
|
||||
)
|
||||
return value
|
||||
elif isinstance(value, dict):
|
||||
if isinstance(value, dict):
|
||||
return _sanitize_request_body_for_spend_logs_payload(
|
||||
value, visited, _max_len
|
||||
value, visited, max_string_length_prompt_in_db
|
||||
)
|
||||
elif isinstance(value, list):
|
||||
return [_sanitize_value(item) for item in value]
|
||||
elif isinstance(value, str):
|
||||
if len(value) > max_string_length_prompt_in_db:
|
||||
# Keep 35% from beginning and 65% from end (end is usually more important)
|
||||
# This split ensures we keep more context from the end of conversations
|
||||
start_ratio = 0.35
|
||||
end_ratio = 0.65
|
||||
|
||||
# Calculate character distribution
|
||||
start_chars = int(max_string_length_prompt_in_db * start_ratio)
|
||||
end_chars = int(max_string_length_prompt_in_db * end_ratio)
|
||||
|
||||
# Ensure we don't exceed the total limit
|
||||
total_keep = start_chars + end_chars
|
||||
if total_keep > max_string_length_prompt_in_db:
|
||||
end_chars = max_string_length_prompt_in_db - start_chars
|
||||
|
||||
# If the string length is less than what we want to keep, just truncate normally
|
||||
if len(value) <= max_string_length_prompt_in_db:
|
||||
return value
|
||||
|
||||
# Calculate how many characters are being skipped
|
||||
skipped_chars = len(value) - total_keep
|
||||
|
||||
# Build the truncated string: beginning + truncation marker + end
|
||||
truncated_value = (
|
||||
f"{value[:start_chars]}"
|
||||
f"... ({LITELLM_TRUNCATED_PAYLOAD_FIELD} skipped {skipped_chars} chars. "
|
||||
f"{LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE}) ..."
|
||||
f"{value[-end_chars:]}"
|
||||
)
|
||||
return truncated_value
|
||||
return value
|
||||
return value
|
||||
|
||||
return {k: _sanitize_value(v) for k, v in request_body.items()}
|
||||
|
|
|
|||
|
|
@ -1314,30 +1314,24 @@ class Router:
|
|||
|
||||
def print_deployment(self, deployment: dict):
|
||||
"""
|
||||
Returns a lightweight dict with model_name + litellm_params (api key masked).
|
||||
returns a copy of the deployment with the api key masked
|
||||
|
||||
Only includes model_name and litellm_params to avoid deep-copying
|
||||
the full deployment dict on every log call.
|
||||
Only returns 2 characters of the api key and masks the rest with * (10 *).
|
||||
"""
|
||||
try:
|
||||
litellm_params: dict = deployment.get("litellm_params", {})
|
||||
_deployment_copy = copy.deepcopy(deployment)
|
||||
litellm_params: dict = _deployment_copy["litellm_params"]
|
||||
|
||||
if litellm.redact_user_api_key_info:
|
||||
masker = SensitiveDataMasker(visible_prefix=2, visible_suffix=0)
|
||||
masked_params = masker.mask_dict(dict(litellm_params))
|
||||
else:
|
||||
masked_params = dict(litellm_params)
|
||||
if "api_key" in masked_params:
|
||||
api_key = masked_params["api_key"]
|
||||
masked_params["api_key"] = (
|
||||
api_key[:2] + "*" * 10 if api_key else api_key
|
||||
)
|
||||
return {
|
||||
"model_name": deployment.get("model_name"),
|
||||
"litellm_params": masked_params,
|
||||
}
|
||||
_deployment_copy["litellm_params"] = masker.mask_dict(litellm_params)
|
||||
elif "api_key" in litellm_params:
|
||||
litellm_params["api_key"] = litellm_params["api_key"][:2] + "*" * 10
|
||||
|
||||
return _deployment_copy
|
||||
except Exception as e:
|
||||
verbose_router_logger.debug(
|
||||
"Error occurred while printing deployment - %s", str(e)
|
||||
f"Error occurred while printing deployment - {str(e)}"
|
||||
)
|
||||
raise e
|
||||
|
||||
|
|
@ -8931,13 +8925,9 @@ class Router:
|
|||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
raise exception
|
||||
if verbose_router_logger.isEnabledFor(logging.INFO):
|
||||
verbose_router_logger.info(
|
||||
"get_available_deployment for model: %s, Selected deployment: %s for model: %s",
|
||||
model,
|
||||
self.print_deployment(deployment),
|
||||
model,
|
||||
)
|
||||
verbose_router_logger.info(
|
||||
f"get_available_deployment for model: {model}, Selected deployment: {self.print_deployment(deployment)} for model: {model}"
|
||||
)
|
||||
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
|
|
@ -9087,12 +9077,9 @@ class Router:
|
|||
)
|
||||
raise exception
|
||||
|
||||
if verbose_router_logger.isEnabledFor(logging.INFO):
|
||||
verbose_router_logger.info(
|
||||
"async_get_available_deployment_for_pass_through model: %s, selected deployment: %s",
|
||||
model,
|
||||
self.print_deployment(deployment),
|
||||
)
|
||||
verbose_router_logger.info(
|
||||
f"async_get_available_deployment_for_pass_through model: {model}, selected deployment: {self.print_deployment(deployment)}"
|
||||
)
|
||||
|
||||
end_time = time.perf_counter()
|
||||
_duration = end_time - start_time
|
||||
|
|
@ -9281,13 +9268,9 @@ class Router:
|
|||
enable_pre_call_checks=self.enable_pre_call_checks,
|
||||
cooldown_list=_cooldown_list,
|
||||
)
|
||||
if verbose_router_logger.isEnabledFor(logging.INFO):
|
||||
verbose_router_logger.info(
|
||||
"get_available_deployment for model: %s, Selected deployment: %s for model: %s",
|
||||
model,
|
||||
self.print_deployment(deployment),
|
||||
model,
|
||||
)
|
||||
verbose_router_logger.info(
|
||||
f"get_available_deployment for model: {model}, Selected deployment: {self.print_deployment(deployment)} for model: {model}"
|
||||
)
|
||||
return deployment
|
||||
|
||||
def get_available_deployment_for_pass_through(
|
||||
|
|
@ -9448,12 +9431,9 @@ class Router:
|
|||
cooldown_list=_cooldown_list,
|
||||
)
|
||||
|
||||
if verbose_router_logger.isEnabledFor(logging.INFO):
|
||||
verbose_router_logger.info(
|
||||
"get_available_deployment_for_pass_through model: %s, selected deployment: %s",
|
||||
model,
|
||||
self.print_deployment(deployment),
|
||||
)
|
||||
verbose_router_logger.info(
|
||||
f"get_available_deployment_for_pass_through model: {model}, selected deployment: {self.print_deployment(deployment)}"
|
||||
)
|
||||
return deployment
|
||||
|
||||
def _filter_cooldown_deployments(
|
||||
|
|
|
|||
|
|
@ -5,7 +5,6 @@ If weights are provided, it will return a deployment based on the weights.
|
|||
|
||||
"""
|
||||
|
||||
import logging
|
||||
import random
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Union
|
||||
|
||||
|
|
@ -53,13 +52,9 @@ def simple_shuffle(
|
|||
selected_index = random.choices(range(len(weights)), weights=weights)[0]
|
||||
verbose_router_logger.debug(f"\n selected index, {selected_index}")
|
||||
deployment = healthy_deployments[selected_index]
|
||||
if verbose_router_logger.isEnabledFor(logging.INFO):
|
||||
verbose_router_logger.info(
|
||||
"get_available_deployment for model: %s, Selected deployment: %s for model: %s",
|
||||
model,
|
||||
llm_router_instance.print_deployment(deployment) or deployment[0],
|
||||
model,
|
||||
)
|
||||
verbose_router_logger.info(
|
||||
f"get_available_deployment for model: {model}, Selected deployment: {llm_router_instance.print_deployment(deployment) or deployment[0]} for model: {model}"
|
||||
)
|
||||
return deployment or deployment[0]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ from dataclasses import dataclass
|
|||
from enum import Enum
|
||||
from typing import Any, Dict, List, Literal, Optional, Tuple
|
||||
|
||||
from pydantic import BaseModel, Field, PrivateAttr, field_validator
|
||||
from pydantic import BaseModel, Field, field_validator
|
||||
from typing_extensions import Annotated
|
||||
|
||||
import litellm
|
||||
|
|
@ -722,8 +722,6 @@ class UserAPIKeyLabelValues(BaseModel):
|
|||
Optional[str], Field(..., alias=UserAPIKeyLabelNames.STREAM.value)
|
||||
] = None
|
||||
|
||||
_cached_dump: Optional[Dict[str, Any]] = PrivateAttr(default=None)
|
||||
|
||||
@field_validator("stream", mode="before")
|
||||
@classmethod
|
||||
def coerce_stream_to_str(cls, v: Any) -> Optional[str]:
|
||||
|
|
@ -731,12 +729,6 @@ class UserAPIKeyLabelValues(BaseModel):
|
|||
return None
|
||||
return str(v)
|
||||
|
||||
def get_label_dict(self) -> Dict[str, Any]:
|
||||
"""Return cached model_dump() dict to avoid re-serializing on every prometheus_label_factory call."""
|
||||
if self._cached_dump is None:
|
||||
self._cached_dump = self.model_dump()
|
||||
return self._cached_dump
|
||||
|
||||
|
||||
class PrometheusMetricsConfig(BaseModel):
|
||||
"""Configuration for filtering Prometheus metrics"""
|
||||
|
|
|
|||
|
|
@ -1,42 +1,41 @@
|
|||
"""
|
||||
Ultra-fast mock OpenAI server for load testing.
|
||||
Returns minimal valid responses with near-zero processing time.
|
||||
Pre-serializes the response body so the mock adds near-zero overhead.
|
||||
"""
|
||||
import time
|
||||
import orjson
|
||||
import uvicorn
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.responses import JSONResponse, StreamingResponse
|
||||
from fastapi import FastAPI, Request, Response
|
||||
|
||||
app = FastAPI()
|
||||
|
||||
MOCK_RESPONSE = {
|
||||
"id": "chatcmpl-mock-loadtest",
|
||||
"object": "chat.completion",
|
||||
"created": 0,
|
||||
"model": "fake-model",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "Mock response for load testing."},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 8, "total_tokens": 18},
|
||||
}
|
||||
_BODY = orjson.dumps(
|
||||
{
|
||||
"id": "chatcmpl-mock-loadtest",
|
||||
"object": "chat.completion",
|
||||
"created": 1700000000,
|
||||
"model": "fake-model",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "Mock."},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12},
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@app.post("/v1/chat/completions")
|
||||
@app.post("/chat/completions")
|
||||
async def chat_completions(request: Request):
|
||||
body = await request.body()
|
||||
response = MOCK_RESPONSE.copy()
|
||||
response["created"] = int(time.time())
|
||||
return JSONResponse(response)
|
||||
await request.body()
|
||||
return Response(content=_BODY, media_type="application/json")
|
||||
|
||||
|
||||
@app.get("/health")
|
||||
async def health():
|
||||
return {"status": "ok"}
|
||||
return Response(content=b'{"status":"ok"}', media_type="application/json")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue