mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
perf(ocr): skip unused callback work and benchmark callback overhead
This commit is contained in:
parent
5cd67f38d6
commit
e3e5b91016
8 changed files with 702 additions and 35 deletions
|
|
@ -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>,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(())
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
289
scripts/benchmark_ocr_callbacks.py
Normal file
289
scripts/benchmark_ocr_callbacks.py
Normal 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()))
|
||||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue