perf(ocr): skip unused callback work and benchmark callback overhead

This commit is contained in:
Yujong Lee 2026-09-12 10:20:39 -07:00
parent 5cd67f38d6
commit e3e5b91016
8 changed files with 702 additions and 35 deletions

View file

@ -19,6 +19,34 @@ impl PythonLogger {
visit.call(&self.0)
}
pub(crate) fn callbacks_needed(&self, py: Python<'_>, phase: &str) -> PyResult<bool> {
if !self
.object(py)
.getattr("_native_callback_fast_path")
.is_ok_and(|value| value.is_truthy().unwrap_or(false))
{
return Ok(true);
}
py.import("litellm.rust_bridge.lifecycle")?
.getattr("callbacks_needed")?
.call1((self.object(py), phase))?
.extract()
}
pub(super) fn success_bookkeeping(
&self,
py: Python<'_>,
response: &Option<Py<PyAny>>,
start: &Py<PyAny>,
end: &Option<Py<PyAny>>,
asynchronous: bool,
) -> PyResult<()> {
py.import("litellm.rust_bridge.lifecycle")?
.getattr("success_bookkeeping")?
.call1((self.object(py), response, start, end, asynchronous))?;
Ok(())
}
pub(super) fn defers_async_logging(&self, py: Python<'_>) -> bool {
self.object(py)
.getattr("_defer_async_logging")
@ -40,6 +68,9 @@ impl PythonLogger {
start: &Py<PyAny>,
end: &Option<Py<PyAny>>,
) -> PyResult<()> {
if !self.callbacks_needed(py, "sync_success_async")? {
return Ok(());
}
self.object(py).call_method1(
"handle_sync_success_callbacks_for_async_calls",
(response, start, end),
@ -55,6 +86,19 @@ impl PythonLogger {
end: &Option<Py<PyAny>>,
asynchronous: bool,
) -> PyResult<Option<Py<PyAny>>> {
if !self.callbacks_needed(
py,
if asynchronous {
"async_failure"
} else {
"sync_failure"
},
)? {
py.import("litellm.rust_bridge.lifecycle")?
.getattr("failure_bookkeeping")?
.call1((self.object(py), error, start, end, asynchronous))?;
return Ok(None);
}
let trace = py
.import("traceback")?
.getattr("format_exception")?
@ -85,6 +129,9 @@ impl PythonLogger {
start: &Py<PyAny>,
end: &Option<Py<PyAny>>,
) -> PyResult<()> {
if !self.callbacks_needed(py, "sync_success")? {
return self.success_bookkeeping(py, response, start, end, false);
}
let context = py.import("contextvars")?.call_method0("copy_context")?;
py.import("litellm.litellm_core_utils.litellm_logging")?
.getattr("executor")?
@ -108,6 +155,9 @@ impl PythonLogger {
start: &Py<PyAny>,
end: &Option<Py<PyAny>>,
) -> PyResult<()> {
if !self.callbacks_needed(py, "async_success")? {
return self.success_bookkeeping(py, response, start, end, true);
}
let context = py.import("contextvars")?.call_method0("copy_context")?;
let worker = py
.import("litellm.litellm_core_utils.logging_worker")?
@ -176,6 +226,13 @@ pub(super) fn is_internal_call(py: Python<'_>) -> PyResult<bool> {
pub(super) struct DeploymentHooks;
impl DeploymentHooks {
pub(super) fn needed(py: Python<'_>) -> PyResult<bool> {
py.import("litellm.rust_bridge.lifecycle")?
.getattr("deployment_callbacks_needed")?
.call0()?
.extract()
}
pub(super) fn before_call(
py: Python<'_>,
kwargs: &Py<PyDict>,

View file

@ -324,6 +324,9 @@ impl PythonCallState {
match phase {
HostPhase::Setup => self.setup(py)?,
HostPhase::DeploymentPreCall => {
if !DeploymentHooks::needed(py)? {
return Ok(HostStep::Ready(self.kwargs.clone_ref(py).into_any()));
}
return Ok(HostStep::Suspend(DeploymentHooks::before_call(
py,
&self.kwargs,
@ -332,6 +335,13 @@ impl PythonCallState {
}
HostPhase::Prepare => self.prepare(py)?,
HostPhase::DeploymentPostCall => {
if !DeploymentHooks::needed(py)? {
return self
.response
.as_ref()
.map(|value| HostStep::Ready(value.clone_ref(py)))
.ok_or_else(missing_state);
}
return Ok(HostStep::Suspend(DeploymentHooks::after_success(
py,
&self.kwargs,
@ -342,7 +352,9 @@ impl PythonCallState {
HostPhase::Finalize => self.finalize(py)?,
HostPhase::Success => self.dispatch_success(py)?,
HostPhase::DeploymentFailure => {
if let Some(error) = &self.error {
if let Some(error) = &self.error
&& DeploymentHooks::needed(py)?
{
return Ok(HostStep::Suspend(DeploymentHooks::after_failure(
py,
&self.kwargs,
@ -448,14 +460,23 @@ impl PythonCallState {
fn try_dispatch_success(&self, py: Python<'_>) -> PyResult<()> {
let logger = self.logger()?;
let pending = PendingSuccess {
let pending = || PendingSuccess {
logger: logger.clone_ref(py),
response: self.response.as_ref().map(|value| value.clone_ref(py)),
start: self.start.clone_ref(py),
end: self.end.as_ref().map(|value| value.clone_ref(py)),
};
if !self.asynchronous {
pending.sync(py)
if !logger.callbacks_needed(py, "sync_success")? {
return logger.success_bookkeeping(
py,
&self.response,
&self.start,
&self.end,
false,
);
}
pending().sync(py)
} else {
if !self.internal
&& self
@ -464,18 +485,20 @@ impl PythonCallState {
.get_item("fallbacks")?
.is_none_or(|value| value.is_none())
{
if logger.defers_async_logging(py) {
if !logger.callbacks_needed(py, "async_success")? {
logger.success_bookkeeping(py, &self.response, &self.start, &self.end, true)?;
} else if logger.defers_async_logging(py) {
logger.defer_success(
py,
Py::new(
py,
PendingLogging {
pending: Some(pending),
pending: Some(pending()),
},
)?,
)?;
} else {
pending.asynchronous(py)?;
pending().asynchronous(py)?;
}
}
logger.sync_success_for_async_call(py, &self.response, &self.start, &self.end)

View file

@ -51,6 +51,11 @@ impl PythonLogger {
kwargs.bind(py).get_item("litellm_call_id")?,
)?;
params.set_item("api_base", url)?;
for name in ["logger_fn", "litellm_request_debug"] {
if let Some(value) = kwargs.bind(py).get_item(name)? {
params.set_item(name, value)?;
}
}
update.set_item("litellm_params", params)?;
update.set_item("custom_llm_provider", &pre_call.custom_llm_provider)?;
self.object(py)
@ -73,8 +78,14 @@ impl PythonLogger {
let kwargs = PyDict::new(py);
kwargs.set_item("input", "OCR document processing")?;
kwargs.set_item("api_key", api_key)?;
kwargs.set_item("additional_args", additional)?;
self.object(py).call_method("pre_call", (), Some(&kwargs))?;
kwargs.set_item("additional_args", &additional)?;
if self.callbacks_needed(py, "input")? {
self.object(py).call_method("pre_call", (), Some(&kwargs))?;
} else {
self.object(py)
.call_method("_pre_call", (), Some(&kwargs))?;
self.object(py).call_method0("record_api_call_start_time")?;
}
Ok(())
}
@ -85,14 +96,24 @@ impl PythonLogger {
body: &Option<Py<PyDict>>,
headers: &Option<Py<PyDict>>,
) -> PyResult<()> {
let kwargs = PyDict::new(py);
kwargs.set_item("original_response", to_py(py, original_response)?)?;
let additional = PyDict::new(py);
additional.set_item("complete_input_dict", body)?;
additional.set_item("headers", headers)?;
kwargs.set_item("additional_args", additional)?;
self.object(py)
.call_method("post_call", (), Some(&kwargs))?;
if self.callbacks_needed(py, "input")? {
let kwargs = PyDict::new(py);
kwargs.set_item("original_response", to_py(py, original_response)?)?;
kwargs.set_item("additional_args", &additional)?;
self.object(py)
.call_method("post_call", (), Some(&kwargs))?;
} else {
let response = py
.import("json")?
.call_method1("dumps", (to_py(py, original_response)?,))?;
self.object(py).call_method1(
"record_post_call",
(response, py.None(), py.None(), additional),
)?;
}
Ok(())
}
}

View file

@ -86,6 +86,14 @@ impl PythonOcrHost {
mut request: OcrDuringCallRequest,
) -> PyResult<OcrDuringCallRequest> {
let pre_call = self.pre_call.as_ref().ok_or_else(missing_state)?;
let logger = self.state.logger()?;
logger.update_ocr(py, &self.state.kwargs, pre_call, &request.url)?;
if !logger.callbacks_needed(py, "payload")? {
logger
.object(py)
.call_method0("record_api_call_start_time")?;
return Ok(request);
}
if let Some(body) = request.body.as_object_mut() {
for name in &request.retained_fields {
body.remove(name);
@ -107,8 +115,6 @@ impl PythonOcrHost {
}
self.body = Some(body.clone().unbind());
self.headers = Some(headers.clone().unbind());
let logger = self.state.logger()?;
logger.update_ocr(py, &self.state.kwargs, pre_call, &request.url)?;
logger.pre_ocr(py, &self.api_key, &body, &headers, &request.url)?;
let headers = headers
.iter()
@ -124,9 +130,10 @@ impl PythonOcrHost {
py: Python<'_>,
request: OcrPostCallRequest,
) -> PyResult<OcrPostCallRequest> {
self.state
.logger()?
.post_ocr(py, &request.original_response, &self.body, &self.headers)?;
let logger = self.state.logger()?;
if logger.callbacks_needed(py, "payload")? {
logger.post_ocr(py, &request.original_response, &self.body, &self.headers)?;
}
Ok(request)
}
}

View file

@ -23,7 +23,6 @@ import litellm
from litellm import (
_custom_logger_compatible_callbacks_literal,
json_logs,
log_raw_request_response,
turn_off_message_logging,
)
from litellm._logging import (
@ -563,6 +562,7 @@ class Logging(LiteLLMLoggingBaseClass):
self.streaming_chunks: list[Any] = [] # for generating complete stream response
self.sync_streaming_chunks: list[Any] = [] # for generating complete stream response
self.log_raw_request_response = log_raw_request_response
self._native_callback_fast_path: bool = False
# Initialize dynamic callbacks
self.dynamic_input_callbacks: list[str | Callable | CustomLogger] | None = dynamic_input_callbacks
@ -1236,6 +1236,11 @@ class Logging(LiteLLMLoggingBaseClass):
additional_args.get("api_base", "")
)
def record_api_call_start_time(self) -> None:
self.model_call_details["api_call_start_time"] = datetime.datetime.now()
if self.model_call_details.get("first_api_call_start_time") is None:
self.model_call_details["first_api_call_start_time"] = self.model_call_details["api_call_start_time"]
def pre_call(self, input, api_key, model=None, additional_args={}):
# Log the exact input to the LLM API
try:
@ -1253,7 +1258,7 @@ class Logging(LiteLLMLoggingBaseClass):
additional_args=additional_args,
)
# log raw request to provider (like LangFuse) -- if opted in.
if self.log_raw_request_response is True or log_raw_request_response is True:
if self.log_raw_request_response is True or litellm.log_raw_request_response is True:
_litellm_params: Final = self.model_call_details.get("litellm_params", {})
_metadata: Final = _litellm_params.get("metadata", {}) or {}
try:
@ -1300,15 +1305,7 @@ class Logging(LiteLLMLoggingBaseClass):
"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging %s", e
)
self.model_call_details["api_call_start_time"] = datetime.datetime.now()
# Set-once first provider-handoff instant. api_call_start_time
# is overwritten on every retry, so it can't measure one-time
# preprocessing; pinning the first attempt excludes retry loops
# + backoff. Logging object only — must NOT go into
# litellm_params["metadata"] (caller request metadata, typed
# Dict[str, str], echoed downstream; a datetime breaks it).
if self.model_call_details.get("first_api_call_start_time") is None:
self.model_call_details["first_api_call_start_time"] = self.model_call_details["api_call_start_time"]
self.record_api_call_start_time()
# Input Integration Logging -> If you want to log the fact that an attempt to call the model was made
callbacks: Final = litellm.input_callback + (self.dynamic_input_callbacks or [])
for callback in callbacks:
@ -1442,16 +1439,21 @@ class Logging(LiteLLMLoggingBaseClass):
"""
return _get_masked_values(headers, ignore_sensitive_values=ignore_sensitive_headers)
def record_post_call(
self, original_response: object, input: object, api_key: object, additional_args: dict[str, object]
) -> None:
self.model_call_details["input"] = input
self.model_call_details["api_key"] = api_key
self.model_call_details["original_response"] = original_response
self.model_call_details["additional_args"] = additional_args
self.model_call_details["log_event_type"] = "post_api_call"
def post_call(self, original_response, input=None, api_key=None, additional_args={}):
# Log the exact result from the LLM API, for streaming - log the type of response received
if isinstance(original_response, dict):
original_response = json.dumps(original_response, default=str)
try:
self.model_call_details["input"] = input
self.model_call_details["api_key"] = api_key
self.model_call_details["original_response"] = original_response
self.model_call_details["additional_args"] = additional_args
self.model_call_details["log_event_type"] = "post_api_call"
self.record_post_call(original_response, input, api_key, additional_args)
attr: Literal["warning", "debug"]
if self.litellm_request_debug:
@ -2116,6 +2118,7 @@ class Logging(LiteLLMLoggingBaseClass):
logging_result,
start_time,
end_time,
build_logging_payload: bool = True,
):
"""Resolve hidden params, compute response cost, and emit the standard logging payload."""
hidden_params: Final = getattr(logging_result, "_hidden_params", {})
@ -2140,6 +2143,9 @@ class Logging(LiteLLMLoggingBaseClass):
else:
self.model_call_details["response_cost"] = self._response_cost_calculator(result=logging_result)
if not build_logging_payload:
return
self.model_call_details["standard_logging_object"] = self._build_standard_logging_payload(
logging_result, start_time, end_time
)
@ -2201,6 +2207,7 @@ class Logging(LiteLLMLoggingBaseClass):
end_time=None,
cache_hit=None,
standard_logging_object: StandardLoggingPayload | None = None,
build_logging_payload: bool = True,
):
try:
if start_time is None:
@ -2238,6 +2245,7 @@ class Logging(LiteLLMLoggingBaseClass):
logging_result=logging_result,
start_time=start_time,
end_time=end_time,
build_logging_payload=build_logging_payload,
)
elif standard_logging_object is not None:
self.model_call_details["standard_logging_object"] = standard_logging_object
@ -3261,7 +3269,9 @@ class Logging(LiteLLMLoggingBaseClass):
except Exception as e:
verbose_logger.debug("Error in _handle_callback_failure: %s", e)
def _failure_handler_helper_fn(self, exception, traceback_exception, start_time=None, end_time=None):
def _failure_handler_helper_fn(
self, exception, traceback_exception, start_time=None, end_time=None, build_logging_payload: bool = True
):
if start_time is None:
start_time = self.start_time
if end_time is None:
@ -3296,6 +3306,9 @@ class Logging(LiteLLMLoggingBaseClass):
metadata: Final = self.model_call_details["litellm_params"].get("metadata", {}) or {}
metadata.update(exception.headers)
if not build_logging_payload:
return start_time, end_time
## STANDARDIZED LOGGING PAYLOAD
self.model_call_details["standard_logging_object"] = get_standard_logging_object_payload(

View file

@ -1,6 +1,7 @@
from __future__ import annotations
import datetime
import os
import uuid
from collections.abc import Awaitable, Mapping
from dataclasses import dataclass
@ -86,10 +87,13 @@ def setup(
}
supplied: Final = arguments.get("litellm_logging_obj")
if isinstance(supplied, Logging):
supplied._native_callback_fast_path = False # pyright: ignore[reportPrivateUsage] # supplied loggers retain all dispatch contracts
return CallSetup(supplied, arguments)
logger, prepared = utils.function_setup(
call_type, utils.Rules(), start_time, *args, is_async_call=asynchronous, **arguments
)
if type(logger) is Logging and call_type in ("ocr", "aocr"):
logger._native_callback_fast_path = True # pyright: ignore[reportPrivateUsage] # only bridge-created OCR loggers opt into callback elision
return CallSetup(logger, prepared)
@ -128,3 +132,84 @@ def finalize(
MetadataUpdater, response_metadata.update_response_metadata
)
update(response, logger, model if isinstance(model, str) else None, kwargs, start_time, end_time)
def deployment_callbacks_needed() -> bool:
import litellm
from litellm.integrations.custom_logger import CustomLogger
return any(isinstance(callback, CustomLogger) for callback in litellm.callbacks)
def callbacks_needed(logger: Logging, phase: str) -> bool:
import litellm
from litellm._logging import (
_is_debugging_on, # pyright: ignore[reportPrivateUsage] # use the same debug gate as Logging
)
if (
_is_debugging_on()
or getattr(logger, "litellm_request_debug", False)
or os.getenv("LITELLM_PRINT_STANDARD_LOGGING_PAYLOAD")
):
return True
input_needed: Final = bool(
litellm.input_callback
or litellm._async_input_callback # pyright: ignore[reportPrivateUsage] # live async registries have no public accessor
or logger.dynamic_input_callbacks
or callable(getattr(logger, "logger_fn", None))
or logger.log_raw_request_response
or litellm.log_raw_request_response
)
match phase:
case "input":
return input_needed
case "sync_success":
return bool(litellm.success_callback or logger.dynamic_success_callbacks)
case "sync_success_async":
return bool(
(litellm.success_callback or logger.dynamic_success_callbacks)
and logger._should_run_sync_callbacks_for_async_calls() # pyright: ignore[reportPrivateUsage] # preserve async call filtering of sync callbacks
)
case "async_success":
return bool(litellm._async_success_callback or logger.dynamic_async_success_callbacks) # pyright: ignore[reportPrivateUsage] # live async registries have no public accessor
case "sync_failure":
return bool(litellm.failure_callback or logger.dynamic_failure_callbacks)
case "async_failure":
return bool(litellm._async_failure_callback or logger.dynamic_async_failure_callbacks) # pyright: ignore[reportPrivateUsage] # live async registries have no public accessor
case "payload":
return bool(
input_needed
or litellm.success_callback
or litellm.failure_callback
or litellm._async_success_callback # pyright: ignore[reportPrivateUsage] # live async registries have no public accessor
or litellm._async_failure_callback # pyright: ignore[reportPrivateUsage] # live async registries have no public accessor
or logger.dynamic_success_callbacks
or logger.dynamic_async_success_callbacks
or logger.dynamic_failure_callbacks
or logger.dynamic_async_failure_callbacks
)
case _:
return True
def success_bookkeeping(
logger: Logging, response: object, start: datetime.datetime, end: datetime.datetime, asynchronous: bool
) -> None:
phase: Final = "async_success" if asynchronous else "sync_success"
if logger.should_run_logging(phase):
logger._success_handler_helper_fn( # pyright: ignore[reportPrivateUsage] # retain success bookkeeping without constructing a callback payload
result=response, start_time=start, end_time=end, build_logging_payload=False
)
logger.has_run_logging(phase)
def failure_bookkeeping(
logger: Logging, error: BaseException, start: datetime.datetime, end: datetime.datetime, asynchronous: bool
) -> None:
phase: Final = "async_failure" if asynchronous else "sync_failure"
if logger.should_run_logging(phase):
logger._failure_handler_helper_fn( # pyright: ignore[reportPrivateUsage] # retain failure accounting without formatting an unused traceback or payload
error, "", start, end, build_logging_payload=False
)
logger.has_run_logging(phase)

View file

@ -0,0 +1,289 @@
#!/usr/bin/env python3
"""Measure serial sync/async OCR latency through a loopback HTTP provider
Run each callback mode in a fresh process against an installed release wheel:
python -I scripts/benchmark_ocr_callbacks.py --callbacks none --label before \
--expected-transport rust --iterations 200 --warmup 20 --output before-none.json
Repeat with --callbacks noop and with the candidate wheel in a separate venv
"""
from __future__ import annotations
import argparse
import asyncio
import base64
import hashlib
import importlib.metadata
import json
import statistics
import sys
import threading
import time
from collections.abc import Sequence
from dataclasses import asdict, dataclass
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from typing import Final, cast
SIZES: Final = (
1024,
4 * 1024,
16 * 1024,
64 * 1024,
256 * 1024,
1024 * 1024,
)
MODEL: Final = "mistral/mistral-ocr-latest"
EXPECTED_MARKDOWN: Final = "mock remote OCR response"
RESPONSE: Final = json.dumps(
{
"pages": [{"index": 0, "markdown": EXPECTED_MARKDOWN, "images": [], "dimensions": None}],
"model": "mistral-ocr-latest",
"usage_info": {"pages_processed": 1},
},
separators=(",", ":"),
).encode()
class Server(ThreadingHTTPServer):
daemon_threads = True
def __init__(self) -> None:
super().__init__(("127.0.0.1", 0), Handler)
self.user_agents: set[str] = set()
class Handler(BaseHTTPRequestHandler):
protocol_version = "HTTP/1.1"
def do_POST(self) -> None:
server: Final = cast(Server, self.server)
server.user_agents.add(self.headers.get("User-Agent", ""))
length: Final = int(self.headers["Content-Length"])
body: Final = self.rfile.read(length)
request: Final = json.loads(body)
if self.path != "/v1/ocr" or request.get("model") != "mistral-ocr-latest":
self.send_error(400)
return
document: Final = request.get("document", {})
if not isinstance(document, dict) or not str(document.get("document_url", "")).startswith(
"data:application/pdf;base64,"
):
self.send_error(400)
return
self.send_response(200)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(RESPONSE)))
self.end_headers()
self.wfile.write(RESPONSE)
def log_message(self, format: str, *args: object) -> None:
return
@dataclass(frozen=True, slots=True)
class Result:
label: str
mode: str
size: int
iterations: int
median_ms: float
mean_ms: float
p95_ms: float
requests_per_second: float
def document(size: int) -> dict[str, str]:
payload: Final = b"%PDF-1.4\n" + b"x" * max(0, size - 9)
encoded: Final = base64.b64encode(payload[:size]).decode("ascii")
return {"type": "document_url", "document_url": f"data:application/pdf;base64,{encoded}"}
def percentile(values: Sequence[float], quantile: float) -> float:
ordered: Final = sorted(values)
index: Final = min(len(ordered) - 1, round((len(ordered) - 1) * quantile))
return ordered[index]
def verify(response: object) -> None:
pages: Final = getattr(response, "pages", ())
if len(pages) != 1 or getattr(pages[0], "markdown", None) != EXPECTED_MARKDOWN:
raise RuntimeError(f"unexpected OCR response: {response!r}")
def summarize(label: str, mode: str, size: int, samples: Sequence[float]) -> Result:
median: Final = statistics.median(samples)
return Result(
label=label,
mode=mode,
size=size,
iterations=len(samples),
median_ms=median * 1000,
mean_ms=statistics.fmean(samples) * 1000,
p95_ms=percentile(samples, 0.95) * 1000,
requests_per_second=1 / median,
)
def sync_samples(litellm: object, url: str, request_document: dict[str, str], count: int) -> tuple[float, ...]:
samples: list[float] = []
for _ in range(count):
started: Final = time.perf_counter()
response: Final = litellm.ocr(
model=MODEL, document=request_document, api_base=url, api_key="mock-key", timeout=30
)
samples.append(time.perf_counter() - started)
verify(response)
return tuple(samples)
async def async_samples(litellm: object, url: str, request_document: dict[str, str], count: int) -> tuple[float, ...]:
samples: list[float] = []
for _ in range(count):
started: Final = time.perf_counter()
response: Final = await litellm.aocr(
model=MODEL, document=request_document, api_base=url, api_key="mock-key", timeout=30
)
samples.append(time.perf_counter() - started)
verify(response)
return tuple(samples)
async def main() -> int:
parser: Final = argparse.ArgumentParser(description="E2E OCR benchmark against a local remote-style HTTP server")
parser.add_argument("--callbacks", choices=("none", "noop"), required=True)
parser.add_argument("--label", required=True)
parser.add_argument("--expected-transport", choices=("python", "rust"), required=True)
parser.add_argument("--iterations", type=int, default=30)
parser.add_argument("--warmup", type=int, default=5)
parser.add_argument("--sizes", type=int, nargs="+", default=SIZES)
parser.add_argument("--output", type=Path, required=True)
args: Final = parser.parse_args()
import litellm
from litellm.integrations.custom_logger import CustomLogger
class NoopCallback(CustomLogger):
def __init__(self) -> None:
super().__init__()
self.pre_calls = 0
self.sync_successes = 0
self.async_successes = 0
def log_pre_api_call(self, model, messages, kwargs):
self.pre_calls += 1
def log_success_event(self, kwargs, response_obj, start_time, end_time):
self.sync_successes += 1
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
self.async_successes += 1
registry_names: Final = (
"callbacks",
"input_callback",
"success_callback",
"failure_callback",
"_async_input_callback",
"_async_success_callback",
"_async_failure_callback",
)
if any(getattr(litellm, name) for name in registry_names):
raise RuntimeError("benchmark requires initially empty callback registrations")
callback: Final = NoopCallback()
if args.callbacks == "noop":
litellm.callbacks.append(callback)
rust_toggle: Final = getattr(litellm, "rust", None)
if callable(rust_toggle):
rust_toggle(False)
package: Final = Path(litellm.__file__).resolve()
version: Final = importlib.metadata.version("litellm")
native_path: str | None = None
native_sha256: str | None = None
try:
from litellm.rust_bridge import _native
native: Final = Path(_native.__file__).resolve()
native_path = str(native)
native_sha256 = hashlib.file_digest(native.open("rb"), "sha256").hexdigest()
except ImportError:
pass
server: Final = Server()
thread: Final = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
url: Final = f"http://127.0.0.1:{server.server_port}"
results: list[Result] = []
try:
for size in args.sizes:
request_document: Final = document(size)
sync_samples(litellm, url, request_document, args.warmup)
sync_result: Final = summarize(
args.label, "sync", size, sync_samples(litellm, url, request_document, args.iterations)
)
results.append(sync_result)
await async_samples(litellm, url, request_document, args.warmup)
async_result: Final = summarize(
args.label,
"async",
size,
await async_samples(litellm, url, request_document, args.iterations),
)
results.append(async_result)
sys.stdout.write(json.dumps(asdict(sync_result)) + "\n")
sys.stdout.write(json.dumps(asdict(async_result)) + "\n")
sys.stdout.flush()
finally:
server.shutdown()
server.server_close()
thread.join()
from litellm.litellm_core_utils.litellm_logging import executor
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
await GLOBAL_LOGGING_WORKER.flush()
await asyncio.to_thread(executor.shutdown, wait=True)
per_mode: Final = len(args.sizes) * (args.iterations + args.warmup)
if args.callbacks == "noop":
if (callback.pre_calls, callback.sync_successes, callback.async_successes) != (
2 * per_mode,
per_mode,
per_mode,
):
raise RuntimeError(f"callback delivery mismatch: {vars(callback)}")
elif any(getattr(litellm, name) for name in registry_names):
raise RuntimeError("callback registrations appeared in the no-callback case")
await GLOBAL_LOGGING_WORKER.stop()
user_agents: Final = tuple(sorted(server.user_agents))
python_transport: Final = any(
value.startswith("python-httpx") or value.startswith("litellm/") for value in user_agents
)
if (args.expected_transport == "python") != python_transport:
raise RuntimeError(f"unexpected transport for {args.label}: user_agents={user_agents}")
artifact: Final = {
"label": args.label,
"callbacks": args.callbacks,
"python": sys.executable,
"callback_counts": {
"pre": callback.pre_calls,
"sync_success": callback.sync_successes,
"async_success": callback.async_successes,
},
"version": version,
"package": str(package),
"native": native_path,
"native_sha256": native_sha256,
"user_agents": user_agents,
"results": tuple(asdict(result) for result in results),
}
args.output.write_text(json.dumps(artifact, indent=2) + "\n")
sys.stdout.write(json.dumps({key: artifact[key] for key in ("label", "version", "package", "user_agents")}) + "\n")
sys.stdout.write(f"results={args.output}\n")
sys.stdout.flush()
return 0
if __name__ == "__main__":
raise SystemExit(asyncio.run(main()))

View file

@ -822,3 +822,175 @@ async def test_response_limit_is_enforced_at_the_public_boundary(
body: Final = ocr_server.requests[0].body
assert isinstance(body, dict)
assert "max_response_bytes" not in body
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("failure", [False, True])
async def test_empty_callbacks_keep_bookkeeping_without_optional_dispatch(
ocr_server: RecordingServer,
monkeypatch: pytest.MonkeyPatch,
asynchronous: bool,
failure: bool,
created_loggers: list[Logging],
) -> None:
from litellm import utils
from litellm.litellm_core_utils import litellm_logging, logging_worker
class DispatchProbe:
deployments = 0
submissions = 0
enqueues = 0
def deployment(self, *args: object, **kwargs: object) -> None:
self.deployments += 1
def submit(self, *args: object, **kwargs: object) -> None:
self.submissions += 1
def ensure_initialized_and_enqueue(self, coroutine: Coroutine[object, object, object]) -> None:
self.enqueues += 1
coroutine.close()
probe: Final = DispatchProbe()
for name in (
"async_pre_call_deployment_hook",
"async_post_call_success_deployment_hook",
"async_post_call_failure_deployment_hook",
):
monkeypatch.setattr(utils, name, probe.deployment)
monkeypatch.setattr(litellm_logging, "executor", probe)
monkeypatch.setattr(logging_worker, "GLOBAL_LOGGING_WORKER", probe)
if failure:
ocr_server.enqueue(ResponseSpec(body={"message": "provider failed"}, status=500))
trace_id_var.set("callback-free-parent")
arguments: Final = {"litellm_trace_id": "callback-free-call", "litellm_call_id": "callback-free-id"}
if failure:
with pytest.raises(litellm.InternalServerError):
await call_aocr(ocr_server, **arguments) if asynchronous else call_ocr(ocr_server, **arguments)
else:
response: Final = (
await call_aocr(ocr_server, **arguments) if asynchronous else call_ocr(ocr_server, **arguments)
)
assert response.pages[0].markdown == "native OCR response"
assert response._hidden_params["litellm_call_id"] == "callback-free-id"
assert response._hidden_params["response_cost"] is not None
assert response._hidden_params["_response_ms"] > 0
assert trace_id_var.get() == "callback-free-parent"
assert probe.deployments == probe.submissions == probe.enqueues == 0
assert len(created_loggers) == 1
logger: Final = created_loggers[0]
assert not hasattr(logger, "_native_pending_logging")
assert logger.model_call_details["first_api_call_start_time"] <= logger.model_call_details["end_time"]
assert "standard_logging_object" not in logger.model_call_details
assert (
"original_response" not in logger.model_call_details or logger.model_call_details["original_response"] is None
)
assert "complete_input_dict" not in logger.model_call_details.get("additional_args", {})
assert logger.model_call_details["response_cost"] == (0 if failure else response._hidden_params["response_cost"])
@pytest.mark.asyncio
@pytest.mark.parametrize(
"registration", ["success_callback", "_async_success_callback", "failure_callback", "_async_failure_callback"]
)
async def test_terminal_registration_added_during_http_is_observed(
ocr_server: RecordingServer, registration: str
) -> None:
failure: Final = "failure" in registration
observer: Final = RecordingLogger()
ocr_server.enqueue(
ResponseSpec(
body={"message": "provider failed"} if failure else OCR_RESPONSE, status=500 if failure else 200, delay=0.1
)
)
task: Final = asyncio.create_task(
asyncio.to_thread(call_ocr, ocr_server) if registration == "success_callback" else call_aocr(ocr_server)
)
await ocr_server.wait_for_requests(1)
getattr(litellm, registration).append(observer)
if failure:
with pytest.raises(litellm.InternalServerError):
await task
else:
await task
event: Final = ("async_" if registration.startswith("_async") else "") + (
"log_failure_event" if failure else "log_success_event"
)
await observer.wait_for_async(event)
assert event in observer.names
@pytest.fixture
def created_loggers(monkeypatch: pytest.MonkeyPatch) -> list[Logging]:
from litellm import utils
original_setup: Final = utils.function_setup
loggers: Final[list[Logging]] = []
def setup(
call_type: str,
rules: utils.Rules,
start: datetime.datetime,
*args: object,
is_async_call: bool = True,
**kwargs: object,
) -> tuple[Logging, dict[str, object]]:
logger, prepared = original_setup(call_type, rules, start, *args, is_async_call=is_async_call, **kwargs)
assert isinstance(logger, Logging)
setattr(logger, "_defer_async_logging", True)
loggers.append(logger)
return logger, prepared
monkeypatch.setattr(utils, "function_setup", setup)
return loggers
@pytest.mark.asyncio
@pytest.mark.parametrize("consumer", ["logger_fn", "raw_global", "request_debug"])
async def test_explicit_logging_consumers_keep_request_and_response_payloads(
ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch, created_loggers: list[Logging], consumer: str
) -> None:
snapshots: Final[list[dict[str, object]]] = []
if consumer == "raw_global":
monkeypatch.setattr(litellm, "log_raw_request_response", True)
arguments: Final = {
"logger_fn": {"logger_fn": lambda details: snapshots.append(dict(details))},
"raw_global": {},
"request_debug": {"litellm_request_debug": True},
}[consumer]
response: Final = await call_aocr(ocr_server, **arguments)
details: Final = created_loggers[0].model_call_details
assert details["additional_args"]["complete_input_dict"]["model"] == "mistral-ocr-latest"
assert json.loads(details["original_response"])["pages"][0]["markdown"] == response.pages[0].markdown
if consumer.startswith("raw_"):
assert details["raw_request_typed_dict"]["raw_request_body"]["model"] == "mistral-ocr-latest"
if consumer == "logger_fn":
assert [item["log_event_type"] for item in snapshots] == ["pre_api_call", "post_api_call"]
@pytest.mark.asyncio
async def test_registration_removed_before_deferred_release_skips_queue(
ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch, created_loggers: list[Logging]
) -> None:
from litellm.litellm_core_utils import logging_worker
class QueueProbe:
enqueues = 0
def ensure_initialized_and_enqueue(self, coroutine: Coroutine[object, object, object]) -> None:
self.enqueues += 1
coroutine.close()
observer: Final = RecordingLogger()
litellm._async_success_callback.append(observer)
await call_aocr(ocr_server)
logger: Final = created_loggers[0]
assert hasattr(logger, "_native_pending_logging")
litellm._async_success_callback.clear()
probe: Final = QueueProbe()
monkeypatch.setattr(logging_worker, "GLOBAL_LOGGING_WORKER", probe)
ProxyBaseLLMRequestProcessing._flush_deferred_async_logging(logger, False)
assert probe.enqueues == 0
assert not observer.names
assert logger.model_call_details["response_cost"] is not None