refactor(python-bridge): drop OCR callback fast path and body re-aliasing

The bridge kept a Rust-side shadow of Logging's callback registries and
elided pre_call/post_call/success_handler when it thought nothing was
listening. That forked the logging contract and hid a bug: with callbacks
present, the during-call hook re-inserted the caller's original document
and unmapped optional params over the provider-transformed body.

Call the real Logging handlers unconditionally, like the Python wrapper
does, and hand callbacks a fresh dict built from core's composed body
instead of re-aliasing caller objects into it.
This commit is contained in:
Yujong Lee 2026-09-15 20:47:52 -07:00
parent 7167bd0ff0
commit 678e59c6b4
9 changed files with 24 additions and 381 deletions

View file

@ -19,34 +19,6 @@ 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")
@ -68,9 +40,6 @@ 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),
@ -86,19 +55,6 @@ 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")?
@ -129,9 +85,6 @@ 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")?
@ -155,9 +108,6 @@ 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")?
@ -226,13 +176,6 @@ 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

@ -312,9 +312,6 @@ 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,
@ -323,13 +320,6 @@ 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,
@ -340,9 +330,7 @@ impl PythonCallState {
HostPhase::Finalize => self.finalize(py)?,
HostPhase::Success => self.dispatch_success(py)?,
HostPhase::DeploymentFailure => {
if let Some(error) = &self.error
&& DeploymentHooks::needed(py)?
{
if let Some(error) = &self.error {
return Ok(HostStep::Suspend(DeploymentHooks::after_failure(
py,
&self.kwargs,
@ -455,15 +443,6 @@ impl PythonCallState {
end: self.end.as_ref().map(|value| value.clone_ref(py)),
};
if !self.asynchronous {
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
@ -473,9 +452,7 @@ impl PythonCallState {
.get_item("fallbacks")?
.is_none_or(|value| value.is_none())
{
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) {
if logger.defers_async_logging(py) {
logger.defer_success(
py,
Py::new(

View file

@ -88,13 +88,7 @@ impl PythonLogger {
kwargs.set_item("input", "OCR document processing")?;
kwargs.set_item("api_key", api_key)?;
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")?;
}
self.object(py).call_method("pre_call", (), Some(&kwargs))?;
Ok(())
}
@ -108,21 +102,11 @@ impl PythonLogger {
let additional = PyDict::new(py);
additional.set_item("complete_input_dict", body)?;
additional.set_item("headers", headers)?;
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),
)?;
}
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))?;
Ok(())
}
}

View file

@ -87,28 +87,9 @@ impl PythonOcrHost {
&projected.fields.secret_fields,
&request.url,
)?;
if !self.state.logger()?.callbacks_needed(py, "payload")? {
self.state
.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);
}
}
let body = to_py(py, &request.body)?
.into_bound(py)
.cast_into::<PyDict>()?;
if let Some(retained) = &self.projected()?.callback_inputs {
for name in &request.retained_fields {
if let Some(value) = retained.bind(py).get_item(name)? {
body.set_item(name, value)?;
}
}
}
let headers = PyDict::new(py);
for (name, value) in &request.headers {
headers.set_item(name, value)?;
@ -134,16 +115,13 @@ impl PythonOcrHost {
py: Python<'_>,
request: OcrPostCallRequest,
) -> PyResult<OcrPostCallRequest> {
let logger = self.state.logger()?;
if logger.callbacks_needed(py, "payload")? {
let projected = self.projected()?;
logger.post_ocr(
py,
&request.original_response,
projected.body.as_ref(),
projected.headers.as_ref(),
)?;
}
let projected = self.projected()?;
self.state.logger()?.post_ocr(
py,
&request.original_response,
projected.body.as_ref(),
projected.headers.as_ref(),
)?;
Ok(request)
}
}

View file

@ -573,7 +573,6 @@ 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

View file

@ -1,7 +1,6 @@
from __future__ import annotations
import datetime
import os
import uuid
from collections.abc import Awaitable, Mapping
from dataclasses import dataclass
@ -87,13 +86,10 @@ 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)
@ -122,84 +118,3 @@ 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

@ -119,69 +119,6 @@ def test_native_ocr_pre_call_header_rebinding_does_not_replace_execution_root(oc
assert "x-rebound" not in ocr_server.requests[0].headers
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"])
async def test_native_ocr_pre_call_nested_document_edit_updates_caller_callback_and_provider_references(
ocr_server: RecordingServer, asynchronous: bool
) -> None:
original: Final = dict(OCR_DOCUMENT)
replacement_url: Final = "data:application/pdf;base64,ZGVm"
retained: Final = []
aliases: Final = []
class Retain(CustomLogger):
def log_pre_api_call(self, model, messages, kwargs):
aliases.append(request_body(kwargs)["document"] is original)
retained.append(request_body(kwargs)["document"])
class Edit(CustomLogger):
def log_pre_api_call(self, model, messages, kwargs):
original["document_url"] = replacement_url
arguments: Final = {
"model": "mistral/mistral-ocr-latest",
"document": original,
"api_key": "test-key",
"api_base": ocr_server.base_url,
"callbacks": [Retain(), Edit()],
}
response: Final = (
await call_native_aocr(ocr_server, **arguments)
if asynchronous
else call_native_ocr(ocr_server, **arguments)
)
assert aliases == [True]
assert retained[0]["document_url"] == replacement_url
assert original["document_url"] == replacement_url
assert ocr_server.requests[0].body["document"]["document_url"] == replacement_url
assert response.pages[0].markdown == "native OCR response"
def test_native_ocr_pre_call_document_replacement_does_not_mutate_original_document(
ocr_server: RecordingServer,
) -> None:
original: Final = dict(OCR_DOCUMENT)
replacement: Final = {"type": "document_url", "document_url": "data:application/pdf;base64,ZGVm"}
retained: Final = []
class RetainAndReplace(CustomLogger):
def log_pre_api_call(self, model, messages, kwargs):
body = request_body(kwargs)
retained.append(body["document"])
body["document"] = replacement
call_native_ocr(
ocr_server,
document=original,
callbacks=[RetainAndReplace()],
)
assert retained[0] is original
assert original["document_url"] == OCR_DOCUMENT["document_url"]
assert ocr_server.requests[0].body["document"] == replacement
def test_native_ocr_pre_call_body_rebinding_is_visible_to_callbacks_but_not_provider(
ocr_server: RecordingServer,
) -> None:

View file

@ -79,7 +79,7 @@ async def test_native_bindings_run_callbacks_and_send_their_mutations(
class Observe(Logging):
def pre_call(self, input, api_key, additional_args):
body: Final = additional_args["complete_input_dict"]
assert body["pages"] is pages
assert body["pages"] == pages and body["pages"] is not pages
observed.append(self.model_call_details["litellm_params"]["metadata"]["marker"])
body["pages"].append(2)
additional_args["headers"]["x-callback"] = "native"
@ -98,6 +98,7 @@ async def test_native_bindings_run_callbacks_and_send_their_mutations(
assert response.pages[0].markdown == "native OCR response"
assert observed == [marker] and observed[0] is marker
assert pages == [0]
assert len(ocr_server.requests) == 1
assert ocr_server.requests[0].body["pages"] == [0, 2]
assert ocr_server.requests[0].headers["x-callback"] == "native"
@ -263,7 +264,6 @@ async def test_deployment_hook_replaces_complete_routing_request(ocr_server: Rec
assert response.pages[0].markdown == "native OCR response"
assert observed == [(replacement, "replacement-key")]
assert observed[0][0] is replacement
assert replacement == original
assert replacement is not original
assert original == {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}
@ -621,7 +621,7 @@ def test_sync_pre_call_can_make_nested_native_request(ocr_server: RecordingServe
@pytest.mark.asyncio
async def test_retained_argument_aliases_and_body_roots_survive_envelope_replacement(
async def test_callback_body_roots_survive_envelope_replacement(
ocr_server: RecordingServer,
) -> None:
pages: Final = [0]
@ -633,8 +633,8 @@ async def test_retained_argument_aliases_and_body_roots_survive_envelope_replace
def pre_call(self, input, api_key, additional_args):
body: Final = additional_args["complete_input_dict"]
headers: Final = additional_args["headers"]
observed.append((body["document"] is document, body["pages"] is pages))
pages.append(2)
observed.append((body["document"] == document, body["pages"] == pages))
body["pages"].append(2)
headers["x-retained"] = "yes"
additional_args["complete_input_dict"] = {"discarded": True}
additional_args["headers"] = {}
@ -672,6 +672,7 @@ async def test_retained_argument_aliases_and_body_roots_survive_envelope_replace
assert observed[0] == (False, False, True)
assert observed[1] == (True, True)
assert observed[3] == (True, True)
assert pages == [0]
assert ocr_server.requests[0].body["pages"] == [0, 2]
assert ocr_server.requests[0].headers["x-retained"] == "yes"
@ -933,72 +934,6 @@ async def test_response_limit_is_enforced_at_the_public_boundary(
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"]
@ -1076,30 +1011,3 @@ async def test_explicit_logging_consumers_keep_request_and_response_payloads(
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

View file

@ -12,6 +12,7 @@ from tests.test_litellm_rust.support.requests import (
OCR_RESPONSE,
call_native_aocr,
call_native_ocr,
request_body,
)
pytestmark = pytest.mark.requires_rust_extension
@ -535,7 +536,7 @@ async def test_native_ocr_inherits_named_credentials_without_overwriting_argumen
class Observer(RecordingLogger):
def log_pre_api_call(self, model, messages, kwargs):
super().log_pre_api_call(model, messages, kwargs)
pages.append(2)
request_body(kwargs)["pages"].append(2)
arguments: Final = {
"model": "mistral/mistral-ocr-latest",
@ -548,3 +549,4 @@ async def test_native_ocr_inherits_named_credentials_without_overwriting_argumen
assert response.pages[0].markdown == "native OCR response"
assert ocr_server.requests[0].headers["authorization"] == f"Bearer {expected_key}"
assert ocr_server.requests[0].body["pages"] == [0, 2]
assert pages == [0]