mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
Optimize embeddings latency: ORJSONResponse bypass, orjson required, backward compat tests
- common_request_processing: bypass FastAPI jsonable_encoder via _response_to_json_serializable + ORJSONResponse for non-streaming responses - pyproject.toml: make orjson a required dependency (was optional) - tests: fix EmbeddingResponse.usage type (Usage instance); add 4 backward compatibility tests comparing new path to legacy jsonable_encoder output - CHANGES_SUMMARY.md: document all changes
This commit is contained in:
parent
6c5eedc836
commit
481a57ea6c
5 changed files with 152 additions and 43 deletions
21
CHANGES_SUMMARY.md
Normal file
21
CHANGES_SUMMARY.md
Normal file
|
|
@ -0,0 +1,21 @@
|
|||
# Changes Summary
|
||||
|
||||
## 1. Large payload truncation
|
||||
|
||||
> Addressed the block from the right side
|
||||
|
||||
Truncate large payloads before `json.dumps` in logging to cut ~3–4s serialization cost.
|
||||
|
||||
- **File**: `litellm/litellm_core_utils/litellm_logging.py`
|
||||
- **Details**: See [truncation_change.md](./truncation_change.md)
|
||||
- **Tests**: `tests/test_litellm/litellm_core_utils/test_litellm_logging.py::TestTruncateLargePayloadForLogging` (15 tests)
|
||||
|
||||
## 2. ORJSONResponse bypass for serialize_response
|
||||
|
||||
> Addressed the request block (~5.5s serialize_response / jsonable_encoder overhead)
|
||||
|
||||
Bypass FastAPI `jsonable_encoder` for non-streaming responses by converting Pydantic/dict responses to plain dicts and returning `ORJSONResponse`. This eliminates the major serialize_response cost for large embeddings and completions.
|
||||
|
||||
- **File**: `litellm/proxy/common_request_processing.py`
|
||||
- **Change**: `_response_to_json_serializable()` converts responses to JSON-serializable dicts; when applicable, we return `ORJSONResponse` directly instead of letting FastAPI serialize the response.
|
||||
- **Tests**: `tests/test_litellm/proxy/test_common_request_processing.py::TestResponseToJsonSerializable` (7 tests, including 4 backward compatibility tests)
|
||||
|
|
@ -17,7 +17,7 @@ from typing import (
|
|||
import httpx
|
||||
import orjson
|
||||
from fastapi import HTTPException, Request, status
|
||||
from fastapi.responses import JSONResponse, Response, StreamingResponse
|
||||
from fastapi.responses import JSONResponse, ORJSONResponse, Response, StreamingResponse
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -53,6 +53,27 @@ from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
|
|||
from litellm.types.utils import ModelResponse, ModelResponseStream, Usage
|
||||
|
||||
|
||||
def _response_to_json_serializable(response: Any) -> Optional[dict]:
|
||||
"""
|
||||
Convert LLM response to a JSON-serializable dict for ORJSONResponse.
|
||||
Bypasses FastAPI jsonable_encoder to eliminate serialize_response overhead.
|
||||
Returns None if conversion is not applicable (caller should return response as-is).
|
||||
"""
|
||||
if isinstance(response, dict):
|
||||
return response
|
||||
if hasattr(response, "model_dump"):
|
||||
try:
|
||||
return response.model_dump(mode="json")
|
||||
except Exception:
|
||||
pass
|
||||
if hasattr(response, "dict"):
|
||||
try:
|
||||
return response.dict()
|
||||
except Exception:
|
||||
pass
|
||||
return None
|
||||
|
||||
|
||||
async def _parse_event_data_for_error(event_line: Union[str, bytes]) -> Optional[int]:
|
||||
"""Parses an event line and returns an error code if present, else None."""
|
||||
event_line = (
|
||||
|
|
@ -809,6 +830,16 @@ class ProxyBaseLLMRequestProcessing:
|
|||
)
|
||||
await check_response_size_is_safe(response=response)
|
||||
|
||||
# Bypass FastAPI jsonable_encoder: return ORJSONResponse with pre-serialized content
|
||||
# to eliminate serialize_response overhead (~5.5s for large embeddings/completions).
|
||||
content = _response_to_json_serializable(response)
|
||||
if content is not None:
|
||||
return ORJSONResponse(
|
||||
content=content,
|
||||
status_code=status.HTTP_200_OK,
|
||||
headers=dict(fastapi_response.headers),
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
async def base_passthrough_process_llm_request(
|
||||
|
|
|
|||
47
poetry.lock
generated
47
poetry.lock
generated
|
|
@ -1,37 +1,5 @@
|
|||
# This file is automatically @generated by Poetry 2.2.1 and should not be changed by hand.
|
||||
|
||||
[[package]]
|
||||
name = "a2a-sdk"
|
||||
version = "0.3.22"
|
||||
description = "A2A Python SDK"
|
||||
optional = true
|
||||
python-versions = ">=3.10"
|
||||
groups = ["main"]
|
||||
markers = "python_version >= \"3.10\" and extra == \"extra-proxy\""
|
||||
files = [
|
||||
{file = "a2a_sdk-0.3.22-py3-none-any.whl", hash = "sha256:b98701135bb90b0ff85d35f31533b6b7a299bf810658c1c65f3814a6c15ea385"},
|
||||
{file = "a2a_sdk-0.3.22.tar.gz", hash = "sha256:77a5694bfc4f26679c11b70c7f1062522206d430b34bc1215cfbb1eba67b7e7d"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
google-api-core = ">=1.26.0"
|
||||
httpx = ">=0.28.1"
|
||||
httpx-sse = ">=0.4.0"
|
||||
protobuf = ">=5.29.5"
|
||||
pydantic = ">=2.11.3"
|
||||
|
||||
[package.extras]
|
||||
all = ["cryptography (>=43.0.0)", "fastapi (>=0.115.2)", "grpcio (>=1.60)", "grpcio-reflection (>=1.7.0)", "grpcio-tools (>=1.60)", "opentelemetry-api (>=1.33.0)", "opentelemetry-sdk (>=1.33.0)", "pyjwt (>=2.0.0)", "sqlalchemy[aiomysql,asyncio] (>=2.0.0)", "sqlalchemy[aiosqlite,asyncio] (>=2.0.0)", "sqlalchemy[asyncio,postgresql-asyncpg] (>=2.0.0)", "sse-starlette", "starlette"]
|
||||
encryption = ["cryptography (>=43.0.0)"]
|
||||
grpc = ["grpcio (>=1.60)", "grpcio-reflection (>=1.7.0)", "grpcio-tools (>=1.60)"]
|
||||
http-server = ["fastapi (>=0.115.2)", "sse-starlette", "starlette"]
|
||||
mysql = ["sqlalchemy[aiomysql,asyncio] (>=2.0.0)"]
|
||||
postgresql = ["sqlalchemy[asyncio,postgresql-asyncpg] (>=2.0.0)"]
|
||||
signing = ["pyjwt (>=2.0.0)"]
|
||||
sql = ["sqlalchemy[aiomysql,asyncio] (>=2.0.0)", "sqlalchemy[aiosqlite,asyncio] (>=2.0.0)", "sqlalchemy[asyncio,postgresql-asyncpg] (>=2.0.0)"]
|
||||
sqlite = ["sqlalchemy[aiosqlite,asyncio] (>=2.0.0)"]
|
||||
telemetry = ["opentelemetry-api (>=1.33.0)", "opentelemetry-sdk (>=1.33.0)"]
|
||||
|
||||
[[package]]
|
||||
name = "aiofiles"
|
||||
version = "24.1.0"
|
||||
|
|
@ -3113,15 +3081,15 @@ files = [
|
|||
|
||||
[[package]]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.34"
|
||||
version = "0.4.23"
|
||||
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
|
||||
optional = true
|
||||
python-versions = "!=2.7.*,!=3.0.*,!=3.1.*,!=3.2.*,!=3.3.*,!=3.4.*,!=3.5.*,!=3.6.*,!=3.7.*,>=3.8"
|
||||
groups = ["main"]
|
||||
markers = "extra == \"proxy\""
|
||||
files = [
|
||||
{file = "litellm_proxy_extras-0.4.34-py3-none-any.whl", hash = "sha256:d455eb54f82e7c92f4f68a921240822df23158aad05fcdda7245887db7c30b90"},
|
||||
{file = "litellm_proxy_extras-0.4.34.tar.gz", hash = "sha256:39fa6c2295acc449320b5a710d150295fd0bf5f8c0d1742b5e9ae361d7bd3ed2"},
|
||||
{file = "litellm_proxy_extras-0.4.23-py3-none-any.whl", hash = "sha256:dfda21203dde9fd97cf364396a9b5be0cfdf00fa9846439ee33ce11b7a52f9ce"},
|
||||
{file = "litellm_proxy_extras-0.4.23.tar.gz", hash = "sha256:8e3f95576dc2a296e7f73d8c87e73628bd899b4644c45863960fe3c3762d8f64"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -4210,10 +4178,9 @@ opentelemetry-api = "1.25.0"
|
|||
name = "orjson"
|
||||
version = "3.11.4"
|
||||
description = "Fast, correct Python JSON library supporting dataclasses, datetimes, and numpy"
|
||||
optional = true
|
||||
optional = false
|
||||
python-versions = ">=3.9"
|
||||
groups = ["main"]
|
||||
markers = "extra == \"proxy\""
|
||||
files = [
|
||||
{file = "orjson-3.11.4-cp310-cp310-macosx_10_15_x86_64.macosx_11_0_arm64.macosx_10_15_universal2.whl", hash = "sha256:e3aa2118a3ece0d25489cbe48498de8a5d580e42e8d9979f65bf47900a15aba1"},
|
||||
{file = "orjson-3.11.4-cp310-cp310-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:a69ab657a4e6733133a3dca82768f2f8b884043714e8d2b9ba9f52b6efef5c44"},
|
||||
|
|
@ -5310,7 +5277,7 @@ description = "Pyroscope Python integration"
|
|||
optional = false
|
||||
python-versions = "*"
|
||||
groups = ["main"]
|
||||
markers = "extra == \"proxy\" and sys_platform != \"win32\""
|
||||
markers = "extra == \"proxy\""
|
||||
files = [
|
||||
{file = "pyroscope_io-0.8.16-py2.py3-none-macosx_11_0_arm64.whl", hash = "sha256:e07edcfd59f5bdce42948b92c9b118c824edbd551730305f095a6b9af401a9e8"},
|
||||
{file = "pyroscope_io-0.8.16-py2.py3-none-macosx_11_0_x86_64.whl", hash = "sha256:dc98355e27c0b7b61f27066500fe1045b70e9459bb8b9a3082bc4755cb6392b6"},
|
||||
|
|
@ -8024,11 +7991,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", "pyroscope-io", "python-multipart", "pyyaml", "rich", "rq", "soundfile", "uvicorn", "uvloop", "websockets"]
|
||||
proxy = ["PyJWT", "apscheduler", "azure-identity", "azure-storage-blob", "backoff", "boto3", "cryptography", "fastapi", "fastapi-sso", "gunicorn", "litellm-enterprise", "litellm-proxy-extras", "mcp", "polars", "pynacl", "pyroscope-io", "python-multipart", "pyyaml", "rich", "rq", "soundfile", "uvicorn", "uvloop", "websockets"]
|
||||
semantic-router = ["semantic-router"]
|
||||
utils = ["numpydoc"]
|
||||
|
||||
[metadata]
|
||||
lock-version = "2.1"
|
||||
python-versions = ">=3.9,<4.0"
|
||||
content-hash = "75f2d51e0e399eaa79f0211bee13f4c887c0e5a76994b21f5ad7237c273a72a0"
|
||||
content-hash = "b916879f4d105891f7088c51aec995a724a9bc81245fff59a0550f07ee1665a4"
|
||||
|
|
|
|||
|
|
@ -32,6 +32,7 @@ jinja2 = "^3.1.2"
|
|||
aiohttp = ">=3.10"
|
||||
pydantic = "^2.5.0"
|
||||
jsonschema = ">=4.23.0,<5.0.0"
|
||||
orjson = ">=3.9.7"
|
||||
numpydoc = {version = "*", optional = true} # used in utils.py
|
||||
|
||||
uvicorn = {version = "^0.31.1", optional = true}
|
||||
|
|
@ -41,7 +42,6 @@ fastapi = {version = ">=0.120.1", optional = true}
|
|||
backoff = {version = "*", optional = true}
|
||||
pyyaml = {version = "^6.0.1", optional = true}
|
||||
rq = {version = "*", optional = true}
|
||||
orjson = {version = "^3.9.7", optional = true}
|
||||
apscheduler = {version = "^3.10.4", optional = true}
|
||||
fastapi-sso = { version = "^0.16.0", optional = true }
|
||||
PyJWT = { version = "^2.10.1", optional = true, python = ">=3.9" }
|
||||
|
|
@ -86,7 +86,6 @@ proxy = [
|
|||
"backoff",
|
||||
"pyyaml",
|
||||
"rq",
|
||||
"orjson",
|
||||
"apscheduler",
|
||||
"fastapi-sso",
|
||||
"PyJWT",
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
import copy
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import orjson
|
||||
import pytest
|
||||
from fastapi import Request, status
|
||||
from fastapi.responses import JSONResponse, StreamingResponse
|
||||
|
|
@ -14,6 +16,7 @@ from litellm.proxy.common_request_processing import (
|
|||
_extract_error_from_sse_chunk,
|
||||
_get_cost_breakdown_from_logging_obj,
|
||||
_parse_event_data_for_error,
|
||||
_response_to_json_serializable,
|
||||
create_response,
|
||||
)
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
|
@ -1024,3 +1027,91 @@ class TestExtractErrorFromSSEChunk:
|
|||
# Other fields should be obtained from the original error object (if exists)
|
||||
|
||||
|
||||
def _parsed_via_legacy_path(obj) -> dict:
|
||||
"""Simulate legacy path: jsonable_encoder + json.dumps -> parse. Returns parsed dict."""
|
||||
from fastapi.encoders import jsonable_encoder
|
||||
|
||||
encoded = jsonable_encoder(obj)
|
||||
return json.loads(json.dumps(encoded))
|
||||
|
||||
|
||||
def _parsed_via_new_path(obj) -> dict:
|
||||
"""New path: _response_to_json_serializable + orjson.dumps -> parse. Returns parsed dict."""
|
||||
content = _response_to_json_serializable(obj)
|
||||
if content is None:
|
||||
raise ValueError("_response_to_json_serializable returned None")
|
||||
return json.loads(orjson.dumps(content).decode())
|
||||
|
||||
|
||||
class TestResponseToJsonSerializable:
|
||||
"""Tests for _response_to_json_serializable (ORJSONResponse bypass for serialize_response)."""
|
||||
|
||||
def test_dict_passthrough(self):
|
||||
"""Plain dict should be returned as-is."""
|
||||
d = {"data": [{"embedding": [0.1] * 100}], "model": "text-embedding-3"}
|
||||
assert _response_to_json_serializable(d) is d
|
||||
|
||||
def test_pydantic_model_dump(self):
|
||||
"""Pydantic model should use model_dump(mode='json')."""
|
||||
from litellm.types.utils import EmbeddingResponse, Usage
|
||||
|
||||
resp = EmbeddingResponse(
|
||||
data=[{"embedding": [0.1] * 100, "index": 0}],
|
||||
model="text-embedding-3",
|
||||
usage=Usage(prompt_tokens=10, total_tokens=10),
|
||||
)
|
||||
result = _response_to_json_serializable(resp)
|
||||
assert result is not None
|
||||
assert result["model"] == "text-embedding-3"
|
||||
assert len(result["data"]) == 1
|
||||
assert result["data"][0]["embedding"] == [0.1] * 100
|
||||
|
||||
def test_none_for_unsupported_type(self):
|
||||
"""Non-dict, non-Pydantic should return None."""
|
||||
assert _response_to_json_serializable("string") is None
|
||||
assert _response_to_json_serializable(42) is None
|
||||
assert _response_to_json_serializable(None) is None
|
||||
|
||||
# --- Backward compatibility: new ORJSON path should produce equivalent JSON to legacy jsonable_encoder ---
|
||||
|
||||
def test_backward_compat_embedding_response_same_as_legacy(self):
|
||||
"""EmbeddingResponse via new path should produce JSON structurally equivalent to jsonable_encoder."""
|
||||
from litellm.types.utils import EmbeddingResponse, Usage
|
||||
|
||||
resp = EmbeddingResponse(
|
||||
data=[{"embedding": [0.1, 0.2, 0.3] * 50, "index": 0}],
|
||||
model="text-embedding-3",
|
||||
usage=Usage(prompt_tokens=10, total_tokens=10),
|
||||
)
|
||||
legacy = _parsed_via_legacy_path(resp)
|
||||
new_path = _parsed_via_new_path(resp)
|
||||
assert legacy == new_path
|
||||
|
||||
def test_backward_compat_plain_dict_same_as_legacy(self):
|
||||
"""Plain dict via new path should produce identical parsed JSON to legacy."""
|
||||
d = {"data": [{"embedding": [0.1] * 100}], "model": "text-embedding-3"}
|
||||
legacy = _parsed_via_legacy_path(d)
|
||||
new_path = _parsed_via_new_path(d)
|
||||
assert legacy == new_path
|
||||
|
||||
def test_backward_compat_unicode_preserved(self):
|
||||
"""Unicode in response should be preserved identically by both paths."""
|
||||
d = {"model": "text-embedding-3", "text": "café résumé 日本語 🔥"}
|
||||
legacy = _parsed_via_legacy_path(d)
|
||||
new_path = _parsed_via_new_path(d)
|
||||
assert legacy == new_path
|
||||
|
||||
def test_backward_compat_nested_structure_same_as_legacy(self):
|
||||
"""Nested dicts/lists should produce equivalent JSON."""
|
||||
d = {
|
||||
"object": "list",
|
||||
"data": [
|
||||
{"embedding": [0.1, 0.2], "index": 0},
|
||||
{"embedding": [0.3, 0.4], "index": 1},
|
||||
],
|
||||
"usage": {"prompt_tokens": 5, "total_tokens": 5},
|
||||
"model": "embedding-model",
|
||||
}
|
||||
legacy = _parsed_via_legacy_path(d)
|
||||
new_path = _parsed_via_new_path(d)
|
||||
assert legacy == new_path
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue