fix(rust): retain callback context across provider dispatch

This commit is contained in:
Yujong Lee 2026-09-08 11:55:42 -07:00
parent 77bf982da1
commit f35ce6ddb8
15 changed files with 900 additions and 470 deletions

View file

@ -25,12 +25,11 @@ use crate::marshal::optional_timeout;
struct ChatCompletionsState {
arguments: Option<Py<PyDict>>,
model: Option<String>,
messages: Option<Value>,
optional_params: Option<Map<String, Value>>,
body: Option<Py<PyDict>>,
headers: Option<Py<PyAny>>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
extra_headers: Option<Map<String, Value>>,
timeout: Option<std::time::Duration>,
terminal: Option<TerminalRecord>,
}
@ -38,13 +37,20 @@ struct ChatCompletionsState {
#[pymethods]
impl ChatCompletionsState {
fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.arguments)
visit.call(&self.arguments)?;
visit.call(&self.body)?;
visit.call(&self.headers)
}
fn __clear__(slf: &Bound<'_, Self>) {
let roots = {
let mut state = slf.borrow_mut();
(state.arguments.take(), state.terminal.take())
(
state.arguments.take(),
state.body.take(),
state.headers.take(),
state.terminal.take(),
)
};
drop(roots);
}
@ -168,7 +174,9 @@ fn prepare(
let admission = admission(bag)?;
let api_key = scalar(bag, "api_key")?;
let api_base = scalar(bag, "api_base")?;
let extra_headers = optional_map(bag, "extra_headers")?;
let headers = bag
.get_item("extra_headers")?
.filter(|value| !value.is_none());
let timeout = optional_timeout(
bag.get_item("timeout_seconds")?
.filter(|value| !value.is_none())
@ -178,13 +186,13 @@ fn prepare(
let complete_input = PyDict::new(py);
complete_input.set_item("model", &admission.model)?;
complete_input.set_item("messages", bag.get_item("messages")?)?;
for (name, value) in &admission.optional_params {
complete_input.set_item(name, Pythonized(value))?;
if let Some(optional_params) = bag.get_item("optional_params")? {
complete_input.call_method1("update", (optional_params,))?;
}
let additional = PyDict::new(py);
additional.set_item(COMPLETE_INPUT_DICT, complete_input)?;
additional.set_item(COMPLETE_INPUT_DICT, &complete_input)?;
additional.set_item(API_BASE, bag.get_item("api_base")?)?;
additional.set_item(HEADERS, bag.get_item("extra_headers")?)?;
additional.set_item(HEADERS, &headers)?;
let kwargs = PyDict::new(py);
kwargs.set_item(INPUT, bag.get_item("messages")?)?;
kwargs.set_item(API_KEY, bag.get_item("logging_api_key")?)?;
@ -197,12 +205,11 @@ fn prepare(
ChatCompletionsState {
arguments: Some(arguments),
model: Some(admission.model),
messages: Some(admission.messages),
optional_params: Some(admission.optional_params),
body: Some(complete_input.unbind()),
headers: headers.map(Bound::unbind),
api_key,
api_base,
custom_llm_provider: admission.custom_llm_provider,
extra_headers,
timeout,
terminal: None,
},
@ -222,24 +229,48 @@ struct OwnedRequest {
}
fn take_request(py: Python<'_>, state: &Py<ChatCompletionsState>) -> PyResult<OwnedRequest> {
let mut state = state.borrow_mut(py);
let call_id = state
.arguments
let (arguments, body, headers) = {
let state = state.borrow(py);
let arguments = state
.arguments
.as_ref()
.ok_or_else(|| PyRuntimeError::new_err("chat completions state was cleared"))?
.clone_ref(py);
let body = state
.body
.as_ref()
.ok_or_else(|| PyRuntimeError::new_err("chat completions body was cleared"))?
.clone_ref(py);
let headers = state.headers.as_ref().map(|value| value.clone_ref(py));
(arguments, body, headers)
};
let call_id = scalar(arguments.bind(py), "litellm_call_id")?.unwrap_or_default();
let messages = value(body.bind(py), "messages")?;
let optional_params = body
.bind(py)
.iter()
.filter_map(|(key, value)| match key.extract::<String>() {
Ok(key) if key == "model" || key == "messages" => None,
Ok(key) => Some(from_py(&value).map(|value| (key, value))),
Err(error) => Some(Err(error)),
})
.collect::<PyResult<Map<String, Value>>>()?;
let extra_headers = headers
.as_ref()
.ok_or_else(|| PyRuntimeError::new_err("chat completions state was cleared"))
.and_then(|arguments| scalar(arguments.bind(py), "litellm_call_id"))?
.unwrap_or_default();
.map(|value| from_py(value.bind(py)))
.transpose()?;
let mut state = state.borrow_mut(py);
Ok(OwnedRequest {
model: state
.model
.take()
.ok_or_else(|| PyRuntimeError::new_err("chat completions request was already sent"))?,
messages: state.messages.take().unwrap(),
optional_params: state.optional_params.take().unwrap(),
messages,
optional_params,
api_key: state.api_key.take(),
api_base: state.api_base.take(),
custom_llm_provider: state.custom_llm_provider.take(),
extra_headers: state.extra_headers.take(),
extra_headers,
timeout: state.timeout.take(),
call_id,
})
@ -399,3 +430,91 @@ pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
pub(super) fn register_trace(module: &Bound<'_, PyModule>) -> PyResult<()> {
register(module)
}
#[cfg(test)]
mod tests {
use super::*;
#[pyfunction]
fn snapshot(py: Python<'_>, state: Py<ChatCompletionsState>) -> PyResult<Py<PyAny>> {
let request = take_request(py, &state)?;
to_py(
py,
&(
request.messages,
request.optional_params,
request.extra_headers,
),
)
}
#[test]
#[ignore = "requires the Python SDK and its dependencies on PYTHONPATH"]
fn callback_roots_survive_rebinding_and_cycles_are_collected() {
Python::initialize();
Python::attach(|py| {
let module = PyModule::new(py, "chat_test").unwrap();
module
.add_function(wrap_pyfunction!(prepare, &module).unwrap())
.unwrap();
module
.add_function(wrap_pyfunction!(snapshot, &module).unwrap())
.unwrap();
let globals = PyDict::new(py);
globals.set_item("native", module).unwrap();
py.run(
c"
import gc
import weakref
class Opaque:
pass
class Logger:
def pre_call(self, **kwargs):
view = kwargs['additional_args']
self.body = view['complete_input_dict']
self.headers = view['headers']
assert kwargs['input'] is messages
assert self.body['messages'] is messages
assert self.body['stop'] is stops
assert self.headers is headers
messages[0]['content'] = 'edited'
stops.append('second')
self.headers['x-hook'] = 'edited'
view['complete_input_dict'] = {'replacement': True}
view['headers'] = {'replacement': 'true'}
messages = [{'role': 'user', 'content': 'original'}]
stops = ['first']
headers = {}
opaque = Opaque()
logger = Logger()
arguments = dict(model='claude-opus-5', messages=messages,
optional_params={'max_tokens': 16, 'stop': stops},
extra_headers=headers, api_key='test',
custom_llm_provider='anthropic', opaque=opaque,
litellm_logging_obj=logger)
state = native.prepare(arguments, logger)
wire_messages, wire_params, wire_headers = native.snapshot(state)
assert wire_messages[0]['content'] == 'edited'
assert wire_params['stop'] == ['first', 'second']
assert wire_headers == {'x-hook': 'edited'}
arguments['cycle'] = state
logger.body['cycle'] = state
headers['cycle'] = state
alive = weakref.ref(opaque)
del arguments, logger, opaque, headers
gc.collect()
assert alive() is not None
del state
gc.collect()
assert alive() is None
",
Some(&globals),
Some(&globals),
)
.unwrap();
});
}
}

View file

@ -407,26 +407,19 @@ class AnthropicChatCompletion(BaseLLM):
stream=stream,
)
if serves_via_rust:
rust_logging_args: Final = {
"complete_input_dict": {
"model": model,
"messages": messages,
**rust_optional_params,
},
"api_base": api_base,
"headers": headers,
}
logging_obj.pre_call(input=messages, api_key=api_key, additional_args=rust_logging_args)
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
logging_obj=logging_obj,
messages=messages,
api_key=api_key,
additional_args=rust_logging_args,
)
if acompletion is True:
async def python_fallback() -> "ModelResponse | CustomStreamWrapper":
fallback_headers, fallback_data = build_request()
logging_obj.pre_call(
input=messages,
api_key=api_key,
additional_args={
"complete_input_dict": fallback_data,
"api_base": api_base,
"headers": fallback_headers,
},
)
return await self.acompletion_function(
model=model,
messages=messages,
@ -460,8 +453,9 @@ class AnthropicChatCompletion(BaseLLM):
custom_llm_provider=custom_llm_provider,
extra_headers=headers,
timeout=timeout,
arguments={**litellm_params, "litellm_logging_obj": logging_obj},
on_response=log_rust_post_call,
logging_obj=logging_obj,
litellm_params=litellm_params,
lifecycle_owner=rust_chat_completions_bridge.LifecycleOwner.WRAPPER,
python_fallback=python_fallback,
)
rust_response: Final = rust_chat_completions_bridge.chat_completions(
@ -474,8 +468,9 @@ class AnthropicChatCompletion(BaseLLM):
custom_llm_provider=custom_llm_provider,
extra_headers=headers,
timeout=timeout,
arguments={**litellm_params, "litellm_logging_obj": logging_obj},
on_response=log_rust_post_call,
logging_obj=logging_obj,
litellm_params=litellm_params,
lifecycle_owner=rust_chat_completions_bridge.LifecycleOwner.WRAPPER,
)
if rust_response is not None:
return rust_response
@ -483,16 +478,15 @@ class AnthropicChatCompletion(BaseLLM):
headers, data = build_request()
## LOGGING
if not serves_via_rust:
logging_obj.pre_call(
input=messages,
api_key=api_key,
additional_args={
"complete_input_dict": data,
"api_base": api_base,
"headers": headers,
},
)
logging_obj.pre_call(
input=messages,
api_key=api_key,
additional_args={
"complete_input_dict": data,
"api_base": api_base,
"headers": headers,
},
)
print_verbose(f"_is_function_call: {_is_function_call}")
if acompletion is True:
if (

View file

@ -417,21 +417,6 @@ class BedrockConverseLLM(BaseAWSLLM):
stream=stream,
)
if serves_via_rust:
rust_logging_args: Final = {
"complete_input_dict": {
"messages": messages,
**optional_params,
},
"api_base": proxy_endpoint_url,
"headers": headers,
}
logging_obj.pre_call(input=messages, api_key="", additional_args=rust_logging_args)
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
logging_obj=logging_obj,
messages=messages,
api_key="",
additional_args=rust_logging_args,
)
if acompletion:
return rust_chat_completions_bridge.achat_completions_or_fallback(
model=model,
@ -443,9 +428,10 @@ class BedrockConverseLLM(BaseAWSLLM):
custom_llm_provider="bedrock",
extra_headers=headers,
timeout=timeout,
arguments={**litellm_params, "litellm_logging_obj": logging_obj},
logging_obj=logging_obj,
litellm_params=litellm_params,
lifecycle_owner=rust_chat_completions_bridge.LifecycleOwner.WRAPPER,
logging_api_key="",
on_response=log_rust_post_call,
python_fallback=lambda: self.async_completion(
model=model,
messages=messages,
@ -462,7 +448,7 @@ class BedrockConverseLLM(BaseAWSLLM):
client=client,
credentials=credentials,
api_key=api_key,
skip_pre_call_logging=True,
skip_pre_call_logging=False,
),
)
rust_response: Final = rust_chat_completions_bridge.chat_completions(
@ -475,9 +461,10 @@ class BedrockConverseLLM(BaseAWSLLM):
custom_llm_provider="bedrock",
extra_headers=headers,
timeout=timeout,
arguments={**litellm_params, "litellm_logging_obj": logging_obj},
logging_obj=logging_obj,
litellm_params=litellm_params,
lifecycle_owner=rust_chat_completions_bridge.LifecycleOwner.WRAPPER,
logging_api_key="",
on_response=log_rust_post_call,
)
if rust_response is not None:
return rust_response
@ -548,16 +535,15 @@ class BedrockConverseLLM(BaseAWSLLM):
)
## LOGGING
if not serves_via_rust:
logging_obj.pre_call(
input=messages,
api_key="",
additional_args={
"complete_input_dict": data,
"api_base": proxy_endpoint_url,
"headers": prepped.headers,
},
)
logging_obj.pre_call(
input=messages,
api_key="",
additional_args={
"complete_input_dict": data,
"api_base": proxy_endpoint_url,
"headers": prepped.headers,
},
)
if client is None or isinstance(client, AsyncHTTPHandler):
_params: Final = {}
if timeout is not None:

View file

@ -2257,12 +2257,9 @@ class BaseLLMHTTPHandler:
stream=stream or False,
custom_llm_provider=custom_llm_provider,
),
arguments={
**kwargs,
"messages": messages,
"litellm_logging_obj": logging_obj,
"litellm_params": litellm_params,
},
logging_obj=logging_obj,
request_arguments=kwargs,
messages=messages,
)
if rust_messages_response is not None:
if stream:
@ -2429,7 +2426,10 @@ class BaseLLMHTTPHandler:
headers: dict,
request_body: dict,
timeout: float | httpx.Timeout | None,
arguments: dict[str, object] | None = None, # mutable-ok: native bridge retains and updates the argument bag
arguments: dict[str, object] | None = None, # mutable-ok: retained legacy bridge input
logging_obj: object | None = None,
request_arguments: Mapping[str, object] | None = None,
messages: object = None,
) -> AnthropicMessagesResponse | None:
if custom_llm_provider not in ("azure_ai", "anthropic"):
return None
@ -2443,32 +2443,21 @@ class BaseLLMHTTPHandler:
from litellm.rust_bridge import messages as rust_messages_bridge
upstream_body: Final = {key: value for key, value in request_body.items() if key != "stream"}
try:
rust_response: Final = await rust_messages_bridge.amessages(
arguments=arguments or {},
model=model,
body=upstream_body,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=headers,
timeout=timeout,
)
except Exception as rust_error: # noqa: BLE001 # rollout-safety fallback: any Rust bridge failure must fall back to the Python path
from litellm.rust_bridge.bindings import native_exception_types
from litellm.rust_bridge.runtime import BridgeErrorContext, raise_upstream
exception_types: Final = native_exception_types()
if exception_types is not None and isinstance(rust_error, exception_types[1]):
raise_upstream(
rust_error,
BridgeErrorContext(route="messages", provider=custom_llm_provider, model=model),
)
verbose_logger.debug(
"Rust Anthropic messages bridge raised %s; falling back to Python path",
type(rust_error).__name__,
)
return None
rust_response: Final = await rust_messages_bridge.amessages(
arguments=arguments,
request_arguments=request_arguments,
logging_obj=logging_obj,
litellm_params=litellm_params,
messages=messages,
lifecycle_owner=rust_messages_bridge.LifecycleOwner.WRAPPER,
model=model,
body=upstream_body,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=headers,
timeout=timeout,
)
if rust_response is None:
return None

View file

@ -1,11 +1,21 @@
from __future__ import annotations
import traceback
from collections.abc import Awaitable, Callable, Mapping
from contextvars import copy_context
from datetime import datetime
from enum import IntEnum
from typing import TYPE_CHECKING, Final, Literal, Protocol
from enum import Enum, IntEnum
from typing import (
TYPE_CHECKING,
Final,
Literal,
Protocol,
cast, # noqa: TID251 # retained callback boundary accepts legacy logger interfaces
overload,
)
from uuid import uuid4
from pydantic import TypeAdapter
from pydantic import InstanceOf, TypeAdapter
if TYPE_CHECKING:
from litellm.types.utils import CallTypes
@ -20,11 +30,66 @@ TerminalAction = Literal[
"sync_failure",
"async_failure",
]
_OPTIONAL_ARGUMENTS_ADAPTER: Final[TypeAdapter[dict[str, object] | None]] = TypeAdapter(
dict[str, object] | None
_OPTIONAL_ARGUMENTS_ADAPTER: Final[TypeAdapter[InstanceOf[dict[str, object]] | None]] = TypeAdapter(
InstanceOf[dict[str, object]] | None
) # mutable-ok: native bridge retains and updates Python argument objects
LIFECYCLE_OWNER_KEY: Final = "_rust_lifecycle_owner"
LIFECYCLE_STARTED_KEY: Final = "_rust_lifecycle_started"
class LifecycleOwner(str, Enum):
BRIDGE = "bridge"
WRAPPER = "wrapper"
def build_call_arguments(
request_arguments: Mapping[str, object] | None,
route_arguments: Mapping[str, object],
*,
logging_obj: object | None = None,
litellm_params: Mapping[str, object] | None = None,
lifecycle_owner: LifecycleOwner = LifecycleOwner.BRIDGE,
) -> dict[str, object]:
return {
**(litellm_params if litellm_params is not None else {}),
**(request_arguments if request_arguments is not None else {}),
**route_arguments,
**({LOGGING_OBJECT_KEY: logging_obj} if logging_obj is not None else {}),
LIFECYCLE_OWNER_KEY: lifecycle_owner.value,
}
@overload
def map_native_error(error: None, arguments: Mapping[str, object], route: str) -> None: ...
@overload
def map_native_error(error: BaseException, arguments: Mapping[str, object], route: str) -> BaseException: ...
def map_native_error(error: BaseException | None, arguments: Mapping[str, object], route: str) -> BaseException | None:
from litellm.rust_bridge.bindings import native_exception_types
from litellm.rust_bridge.runtime import BridgeErrorContext, upstream_error
exceptions: Final = native_exception_types()
if error is None or exceptions is None or not isinstance(error, exceptions[1]):
return error
return upstream_error(
error,
BridgeErrorContext(
route=route,
provider=str(arguments.get("custom_llm_provider") or ""),
model=str(arguments.get("model") or ""),
),
)
def owns_lifecycle(arguments: Mapping[str, object]) -> bool:
return arguments.get(LIFECYCLE_OWNER_KEY, LifecycleOwner.BRIDGE.value) == LifecycleOwner.BRIDGE.value
class NativeOutcome(IntEnum):
SUCCESS = 0
FAILURE = 1
@ -161,11 +226,139 @@ async def drive_async(host: LifecycleHost) -> object:
def initialize_logging(
arguments: dict[str, object], asynchronous: bool, route: str
arguments: dict[str, object], asynchronous: bool, route: str = "ocr"
) -> object: # mutable-ok: native bridge retains and updates Python argument objects
from litellm.rust_bridge.ocr import initialize_logging as initialize_ocr_logging
import litellm
from litellm import utils
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils import litellm_logging
from litellm.litellm_core_utils.coroutine_checker import coroutine_checker
from litellm.litellm_core_utils.litellm_logging import Logging, set_callbacks
return initialize_ocr_logging(arguments, asynchronous, route)
supplied: Final = arguments.get(LOGGING_OBJECT_KEY)
if supplied is not None:
return supplied
callbacks: Final = tuple( # cast-ok: callback registry accepts heterogeneous legacy callback objects
dict.fromkeys(
utils.get_dynamic_callbacks(
cast( # cast-ok: callback registry accepts heterogeneous legacy callback objects
list, arguments.get("callbacks")
) # cast-ok: callback registry accepts heterogeneous legacy callback objects
) # cast-ok: callback registry accepts heterogeneous legacy callback objects
) # cast-ok: callback registry accepts heterogeneous legacy callback objects # mutable-ok: deduplication uses dict keys
)
success: Final = tuple( # cast-ok: per-call callback list is a legacy untyped boundary
dict.fromkeys(
(
*callbacks,
*cast( # cast-ok: per-call callback list is a legacy untyped boundary
list, arguments.get("success_callback") or ()
), # cast-ok: per-call callback list is a legacy untyped boundary
) # cast-ok: per-call callback list is a legacy untyped boundary
) # cast-ok: per-call callback list is a legacy untyped boundary # mutable-ok: deduplication uses dict keys
)
failure: Final = tuple( # cast-ok: per-call callback list is a legacy untyped boundary
dict.fromkeys(
(
*callbacks,
*cast( # cast-ok: per-call callback list is a legacy untyped boundary
list, arguments.get("failure_callback") or ()
), # cast-ok: per-call callback list is a legacy untyped boundary
) # cast-ok: per-call callback list is a legacy untyped boundary
) # cast-ok: per-call callback list is a legacy untyped boundary # mutable-ok: deduplication uses dict keys
)
configured: Final = tuple(
dict.fromkeys(
(
*litellm.input_callback,
*litellm.success_callback,
*litellm.failure_callback,
*litellm._async_success_callback, # pyright: ignore[reportPrivateUsage] # callback registry has no public accessor
*litellm._async_failure_callback, # pyright: ignore[reportPrivateUsage] # callback registry has no public accessor
*success,
*failure,
)
)
)
uninitialized: Final = [ # mutable-ok: set_callbacks requires a mutable callback list
cb
for cb in configured
if isinstance(cb, str)
and (
cb not in litellm._known_custom_logger_compatible_callbacks # pyright: ignore[reportPrivateUsage] # callback compatibility registry has no public accessor
or cb in litellm.input_callback + litellm.success_callback + litellm.failure_callback
)
and cb not in (utils.callback_list or ())
]
if uninitialized:
set_callbacks(uninitialized, function_id=arguments.get("id"))
utils.callback_list = list( # mutable-ok: global callback registry is mutable
dict.fromkeys((*(utils.callback_list or ()), *uninitialized))
) # mutable-ok: global callback registry is mutable
if litellm_logging.customLogger is None: # pyright: ignore[reportUnnecessaryComparison] # runtime plugin registry can be reset to None
set_callbacks(
[cb for cb in configured if callable(cb)], # mutable-ok: set_callbacks requires a mutable callback list
function_id=arguments.get("id"), # mutable-ok: set_callbacks requires a mutable callback list
) # mutable-ok: set_callbacks requires a mutable callback list
for event, registered, add_async in (
("input", litellm.input_callback, litellm.logging_callback_manager.add_litellm_input_callback),
("success", litellm.success_callback, litellm.logging_callback_manager.add_litellm_async_success_callback),
("failure", litellm.failure_callback, litellm.logging_callback_manager.add_litellm_async_failure_callback),
):
for cb in tuple(registered):
if coroutine_checker.is_async_callable(cb) or (event == "success" and cb in ("dynamodb", "openmeter")):
if cb not in getattr(litellm, f"_async_{event}_callback"):
add_async(cb)
registered.remove(cb)
elif event != "input" and isinstance(cb, str) and cb in litellm._known_custom_logger_compatible_callbacks: # pyright: ignore[reportPrivateUsage] # callback compatibility registry has no public accessor
utils._add_custom_logger_callback_to_specific_event(cb, event) # pyright: ignore[reportPrivateUsage] # callback manager only exposes this internal registration path
for event, registered, add_sync in (
("success", litellm._async_success_callback, litellm.logging_callback_manager.add_litellm_success_callback), # pyright: ignore[reportPrivateUsage] # callback registry has no public accessor
("failure", litellm._async_failure_callback, litellm.logging_callback_manager.add_litellm_failure_callback), # pyright: ignore[reportPrivateUsage] # callback registry has no public accessor
):
for cb in tuple(registered):
if callable(cb) and not isinstance(cb, CustomLogger) and not coroutine_checker.is_async_callable(cb):
if cb not in getattr(litellm, f"{event}_callback"):
add_sync(cb)
registered.remove(cb)
call_id: Final = str(arguments.get("litellm_call_id") or uuid4())
logger: Final = Logging(
model=str(arguments["model"]),
messages="default-message-value",
stream=False,
call_type=f"a{route}" if asynchronous else route,
start_time=datetime.now(), # noqa: DTZ005 # Logging preserves the legacy naive timestamp contract
litellm_call_id=call_id,
function_id=str(arguments.get("id") or ""),
litellm_trace_id=cast( # cast-ok: public call argument is validated by Logging
str | None, arguments.get("litellm_trace_id")
), # cast-ok: public call argument is validated by Logging
dynamic_input_callbacks=[ # mutable-ok: Logging callback configuration is mutable
cb for cb in callbacks if cb not in litellm.input_callback and not coroutine_checker.is_async_callable(cb)
],
dynamic_success_callbacks=[ # mutable-ok: Logging callback configuration is mutable
cb for cb in success if not coroutine_checker.is_async_callable(cb) and cb not in ("dynamodb", "s3")
],
dynamic_async_success_callbacks=[ # mutable-ok: Logging callback configuration is mutable
cb
for cb in success
if coroutine_checker.is_async_callable(cb) or isinstance(cb, CustomLogger) or cb in ("dynamodb", "s3")
],
dynamic_failure_callbacks=[ # mutable-ok: Logging callback configuration is mutable
cb for cb in failure if not coroutine_checker.is_async_callable(cb)
], # mutable-ok: Logging callback configuration is mutable
dynamic_async_failure_callbacks=[ # mutable-ok: Logging callback configuration is mutable
cb for cb in failure if coroutine_checker.is_async_callable(cb) or isinstance(cb, CustomLogger)
],
kwargs=arguments,
supports_correlation_logging=asynchronous,
)
logger.dynamic_input_callbacks = [ # mutable-ok: remove callbacks promoted to the global registry
cb for cb in dict.fromkeys(logger.dynamic_input_callbacks or ()) if cb not in litellm.input_callback
]
arguments["litellm_call_id"] = call_id
arguments[LOGGING_OBJECT_KEY] = logger
return logger
def invoke_terminal(
@ -174,9 +367,63 @@ def invoke_terminal(
logger: object,
record: Mapping[str, object] | None,
value: object,
start_time: datetime,
end_time: datetime,
fallback_start_time: datetime,
fallback_end_time: datetime,
) -> object:
from litellm.rust_bridge.ocr import invoke_terminal as invoke_ocr_terminal
from litellm import utils
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
return invoke_ocr_terminal(action, roots, logger, record, value, start_time, end_time)
logging: Final = cast( # cast-ok: callers may supply a Logging-compatible test or plugin implementation
Logging, logger
)
timing_value: Final = record.get("timing") if record is not None else None
timing: Final = timing_value if isinstance(timing_value, Mapping) else None
start_value: Final = timing.get("start_time") if timing is not None else None
end_value: Final = timing.get("end_time") if timing is not None else None
start_time: Final = (
datetime.fromtimestamp(start_value, tz=fallback_start_time.tzinfo)
if isinstance(start_value, (int, float))
else fallback_start_time
)
end_time: Final = (
datetime.fromtimestamp(end_value, tz=fallback_end_time.tzinfo)
if isinstance(end_value, (int, float))
else fallback_end_time
)
if action == "sync_success":
def run() -> None:
_retained: Final = roots
logging.success_handler(value, start_time, end_time)
return utils.executor.submit(copy_context().run, run)
if action == "async_success":
async def run_async() -> None:
_retained: Final = roots
await logging.async_success_handler(value, start_time, end_time)
def enqueue() -> None:
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(async_coroutine=run_async())
if getattr(logging, "_defer_async_logging", False) is True:
logging._enqueue_deferred_logging = enqueue # pyright: ignore[reportPrivateUsage] # preserves Logging's deferred callback contract
else:
enqueue()
return None
if action == "sync_success_if_needed":
if logging._should_run_sync_callbacks_for_async_calls(): # pyright: ignore[reportPrivateUsage] # preserves Logging's async callback policy
def run() -> None:
_retained: Final = roots
logging.success_handler(value, start_time, end_time)
return utils.executor.submit(copy_context().run, run)
return None
exception: Final = cast(Exception, value) # cast-ok: Rust routes terminal failure values as Python exceptions
trace: Final = "".join(traceback.format_exception(type(exception), exception, exception.__traceback__))
if action == "sync_failure":
logging.failure_handler(exception, trace, start_time, end_time)
return None
return logging.async_failure_handler(exception, trace, start_time, end_time)

View file

@ -34,12 +34,15 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
)
from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned
from litellm.rust_bridge._lifecycle import (
LIFECYCLE_STARTED_KEY,
LOGGING_OBJECT_KEY,
LifecycleOwner,
NativeLifecycle,
NativeLifecycleBindings,
NativeOutcome,
TerminalAction,
advance_host,
build_call_arguments,
deployment_failure,
deployment_pre,
deployment_success,
@ -47,6 +50,8 @@ from litellm.rust_bridge._lifecycle import (
drive_sync,
host_result,
invoke_terminal,
map_native_error,
owns_lifecycle,
restore_correlation_context,
)
from litellm.rust_bridge.configuration import rust_enabled
@ -319,6 +324,7 @@ def _reraise_or_decline(
*,
model: str,
custom_llm_provider: str | None,
lifecycle_started: bool = False,
) -> None:
"""Re-raise a failure the provider already saw, or return so the caller declines.
@ -329,11 +335,7 @@ def _reraise_or_decline(
"""
exceptions: Final = _rust_bridge_exceptions()
if exceptions is None:
verbose_logger.debug(
"Rust chat completions bridge raised %s; falling back to Python path",
type(rust_error).__name__,
)
return
raise rust_error
declined, upstream_failed = exceptions
if isinstance(rust_error, upstream_failed):
args: Final = rust_error.args
@ -345,7 +347,7 @@ def _reraise_or_decline(
llm_provider=custom_llm_provider or "",
model=model,
)
if not isinstance(rust_error, declined):
if lifecycle_started or not isinstance(rust_error, declined):
raise rust_error
verbose_logger.debug(
"Rust chat completions declined before calling the provider (%s); using the Python path",
@ -379,32 +381,37 @@ def chat_completions(
extra_headers: Mapping[str, object] | None,
timeout: float | httpx.Timeout | None,
arguments: dict[str, object] | None = None, # mutable-ok: native bridge retains and updates Python argument objects
logging_obj: object | None = None,
litellm_params: Mapping[str, object] | None = None,
lifecycle_owner: LifecycleOwner = LifecycleOwner.BRIDGE,
logging_api_key: str | None = None,
on_response: ResponseObserver | None = None,
) -> ModelResponse | None:
rust_chat_completions: Final = load_rust_chat_completions()
if rust_chat_completions is None:
return None
call_arguments: Final = _arguments(
arguments,
model,
messages,
optional_params,
model_response,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout,
logging_api_key,
logging_obj,
litellm_params,
lifecycle_owner,
)
try:
if _STATE.chat_completions is not None and _uses_argument_bag(rust_chat_completions):
argument_bag_call: Final = cast( # cast-ok: signature inspection selected the argument-bag callable
RustChatCompletions, rust_chat_completions
)
rust_result: Final = argument_bag_call(
_arguments(
arguments,
model,
messages,
optional_params,
model_response,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout,
logging_api_key,
)
)
rust_result: Final = argument_bag_call(call_arguments)
return rust_result
if _STATE.chat_completions is not None:
legacy: Final = cast( # cast-ok: signature inspection selected the legacy injected callable
@ -424,23 +431,14 @@ def chat_completions(
on_response(rust_response)
return build_model_response(rust_response, model_response)
native_call: Final = cast(RustChatCompletions, rust_chat_completions) # cast-ok: native ABI uses argument bag
return native_call(
_arguments(
arguments,
model,
messages,
optional_params,
model_response,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout,
logging_api_key,
)
)
return native_call(call_arguments)
except Exception as rust_error: # noqa: BLE001 # rollout safety: the helper re-raises anything the provider already saw
_reraise_or_decline(rust_error, model=model, custom_llm_provider=custom_llm_provider)
_reraise_or_decline(
rust_error,
model=model,
custom_llm_provider=custom_llm_provider,
lifecycle_started=call_arguments.get(LIFECYCLE_STARTED_KEY) is True,
)
return None
raise AssertionError("unreachable")
@ -457,32 +455,37 @@ async def achat_completions(
extra_headers: Mapping[str, object] | None,
timeout: float | httpx.Timeout | None,
arguments: dict[str, object] | None = None, # mutable-ok: native bridge retains and updates Python argument objects
logging_obj: object | None = None,
litellm_params: Mapping[str, object] | None = None,
lifecycle_owner: LifecycleOwner = LifecycleOwner.BRIDGE,
logging_api_key: str | None = None,
on_response: ResponseObserver | None = None,
) -> ModelResponse | None:
rust_achat_completions: Final = load_rust_achat_completions()
if rust_achat_completions is None:
return None
call_arguments: Final = _arguments(
arguments,
model,
messages,
optional_params,
model_response,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout,
logging_api_key,
logging_obj,
litellm_params,
lifecycle_owner,
)
try:
if _STATE.achat_completions is not None and _uses_argument_bag(rust_achat_completions):
argument_bag_call: Final = cast( # cast-ok: signature inspection selected the argument-bag callable
RustAchatCompletions, rust_achat_completions
)
rust_result: Final = await argument_bag_call(
_arguments(
arguments,
model,
messages,
optional_params,
model_response,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout,
logging_api_key,
)
)
rust_result: Final = await argument_bag_call(call_arguments)
return rust_result
if _STATE.achat_completions is not None:
legacy: Final = cast( # cast-ok: signature inspection selected the legacy injected callable
@ -504,23 +507,14 @@ async def achat_completions(
native_call: Final = cast( # cast-ok: native ABI uses argument bag
RustAchatCompletions, rust_achat_completions
)
return await native_call(
_arguments(
arguments,
model,
messages,
optional_params,
model_response,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout,
logging_api_key,
)
)
return await native_call(call_arguments)
except Exception as rust_error: # noqa: BLE001 # rollout safety: the helper re-raises anything the provider already saw
_reraise_or_decline(rust_error, model=model, custom_llm_provider=custom_llm_provider)
_reraise_or_decline(
rust_error,
model=model,
custom_llm_provider=custom_llm_provider,
lifecycle_started=call_arguments.get(LIFECYCLE_STARTED_KEY) is True,
)
return None
raise AssertionError("unreachable")
@ -538,6 +532,9 @@ async def achat_completions_or_fallback(
timeout: float | httpx.Timeout | None,
python_fallback: Callable[[], Awaitable[object]],
arguments: dict[str, object] | None = None, # mutable-ok: native bridge retains and updates Python argument objects
logging_obj: object | None = None,
litellm_params: Mapping[str, object] | None = None,
lifecycle_owner: LifecycleOwner = LifecycleOwner.BRIDGE,
logging_api_key: str | None = None,
on_response: ResponseObserver | None = None,
) -> object:
@ -560,6 +557,9 @@ async def achat_completions_or_fallback(
extra_headers=extra_headers,
timeout=timeout,
arguments=arguments,
logging_obj=logging_obj,
litellm_params=litellm_params,
lifecycle_owner=lifecycle_owner,
logging_api_key=logging_api_key,
on_response=on_response,
)
@ -588,20 +588,28 @@ def _arguments(
extra_headers: Mapping[str, object] | None,
timeout: float | httpx.Timeout | None,
logging_api_key: str | None,
logging_obj: object | None,
litellm_params: Mapping[str, object] | None,
lifecycle_owner: LifecycleOwner,
) -> dict[str, object]: # mutable-ok: native bridge retains and updates Python argument objects
return {
**(arguments or {}),
"model": model,
"messages": messages,
"optional_params": optional_params,
"model_response": model_response,
"api_key": api_key,
"api_base": api_base,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"timeout_seconds": timeout_to_seconds(timeout),
"logging_api_key": logging_api_key if logging_api_key is not None else api_key or "",
}
return build_call_arguments(
arguments,
{
"model": model,
"messages": messages,
"optional_params": optional_params,
"model_response": model_response,
"api_key": api_key,
"api_base": api_base,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"timeout_seconds": timeout_to_seconds(timeout),
"logging_api_key": logging_api_key if logging_api_key is not None else api_key or "",
},
logging_obj=logging_obj,
litellm_params=litellm_params,
lifecycle_owner=lifecycle_owner,
)
class _ChatCompletionsBindings(NativeLifecycleBindings, Protocol):
@ -634,6 +642,8 @@ class _ChatCompletionsHost:
arguments # mutable-ok: native bridge retains and updates Python argument objects
)
self.asynchronous: bool = asynchronous
self.lifecycle_owned: Final = owns_lifecycle(arguments)
self.arguments[LIFECYCLE_STARTED_KEY] = True
self.logger: object | None = arguments.get(LOGGING_OBJECT_KEY)
self.state: object | None = None
self.response: object = None
@ -649,6 +659,8 @@ class _ChatCompletionsHost:
self.arguments[LOGGING_OBJECT_KEY] = self.logger
async def deployment_pre(self) -> None:
if not self.lifecycle_owned:
return
self.current = await deployment_pre(self.current, "acompletion")
self.current[LOGGING_OBJECT_KEY] = self.logger
@ -672,14 +684,20 @@ class _ChatCompletionsHost:
self.end = datetime.now() # noqa: DTZ005 # Logging preserves the legacy naive timestamp contract
async def deployment_success(self) -> None:
if not self.lifecycle_owned:
return
from litellm.types.utils import CallTypes
self.response = await deployment_success(self.current, self.response, CallTypes.acompletion)
async def deployment_failure(self) -> None:
if not self.lifecycle_owned:
return
await deployment_failure(self.current, self.error, "acompletion")
def terminal(self, action: TerminalAction, value: object) -> object:
if not self.lifecycle_owned:
return None
if self.logger is None or self.end is None:
raise RuntimeError("chat completions terminal state was not initialized")
record: Final = self.bindings.terminal_record(self.state) if self.state is not None else None
@ -705,14 +723,16 @@ class _ChatCompletionsHost:
def sync_failure(self) -> object:
return self.terminal("sync_failure", self.error)
def async_failure(self) -> object:
return self.terminal("async_failure", self.error)
async def async_failure(self) -> object:
result: Final = self.terminal("async_failure", self.error)
return await result if isinstance(result, Awaitable) else result
def restore(self) -> None:
restore_correlation_context(self.logger)
if self.lifecycle_owned:
restore_correlation_context(self.logger)
def advance(self, outcome: NativeOutcome, error: BaseException | None = None) -> None:
advance_host(self, outcome, error)
advance_host(self, outcome, map_native_error(error, self.arguments, "chat completions"))
def result(self) -> object:
return host_result(self)

View file

@ -1,6 +1,6 @@
from __future__ import annotations
from collections.abc import Awaitable, Callable
from collections.abc import Awaitable, Callable, Mapping
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Final, Protocol, cast
@ -8,12 +8,15 @@ from typing import Final, Protocol, cast
import httpx
from litellm.rust_bridge._lifecycle import (
LIFECYCLE_STARTED_KEY,
LOGGING_OBJECT_KEY,
LifecycleOwner,
NativeLifecycle,
NativeLifecycleBindings,
NativeOutcome,
TerminalAction,
advance_host,
build_call_arguments,
deployment_failure,
deployment_pre,
deployment_success,
@ -21,16 +24,19 @@ from litellm.rust_bridge._lifecycle import (
drive_sync,
host_result,
invoke_terminal,
map_native_error,
owns_lifecycle,
restore_correlation_context,
)
from litellm.rust_bridge._lifecycle import (
initialize_logging as initialize_lifecycle_logging,
)
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.bindings import NativeBinding, native_exception_types
from litellm.rust_bridge.timeouts import timeout_to_seconds
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
from litellm.types.router import GenericLiteLLMParams
class RustMessages(Protocol):
@ -168,7 +174,7 @@ def retain_stream_response(
def _arguments(
arguments: dict[str, object], # mutable-ok: native bridge retains and updates Python argument objects
arguments: Mapping[str, object] | None,
model: str,
body: dict[str, object], # mutable-ok: native bridge retains and updates Python argument objects
api_key: str | None,
@ -176,17 +182,27 @@ def _arguments(
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None, # mutable-ok: native bridge retains and updates Python argument objects
timeout: float | httpx.Timeout | None,
logging_obj: object | None,
litellm_params: GenericLiteLLMParams | None,
messages: object,
lifecycle_owner: LifecycleOwner,
) -> dict[str, object]: # mutable-ok: native bridge retains and updates Python argument objects
return { # mutable-ok: the native bridge requires a concrete argument bag
**arguments,
"model": model,
"body": body,
"api_key": api_key,
"api_base": api_base,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"timeout_seconds": timeout_to_seconds(timeout),
}
return build_call_arguments(
arguments,
{
"model": model,
"body": body,
"api_key": api_key,
"api_base": api_base,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"timeout_seconds": timeout_to_seconds(timeout),
**({"messages": messages} if messages is not None else {}),
**({"litellm_params": litellm_params} if litellm_params is not None else {}),
},
logging_obj=logging_obj,
lifecycle_owner=lifecycle_owner,
)
def messages(
@ -199,15 +215,40 @@ def messages(
extra_headers: dict[str, object] | None, # mutable-ok: native bridge retains and updates Python argument objects
timeout: float | httpx.Timeout | None,
arguments: dict[str, object] | None = None, # mutable-ok: native bridge retains and updates Python argument objects
request_arguments: Mapping[str, object] | None = None,
logging_obj: object | None = None,
litellm_params: GenericLiteLLMParams | None = None,
messages: object = None,
lifecycle_owner: LifecycleOwner = LifecycleOwner.BRIDGE,
) -> AnthropicMessagesResponse | None:
implementation: Final = load_rust_messages()
if implementation is None:
return None
return implementation(
arguments=_arguments(
arguments or {}, model, body, api_key, api_base, custom_llm_provider, extra_headers, timeout
)
call_arguments: Final = _arguments(
request_arguments if request_arguments is not None else arguments,
model,
body,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout,
logging_obj,
litellm_params,
messages,
lifecycle_owner,
)
try:
return implementation(arguments=call_arguments)
except Exception as error: # noqa: BLE001 # only explicit declines before lifecycle setup may fall back
exceptions: Final = native_exception_types()
if (
exceptions is not None
and isinstance(error, exceptions[0])
and not call_arguments.get(LIFECYCLE_STARTED_KEY)
):
return None
raise map_native_error(error, call_arguments, "messages")
async def amessages(
@ -220,15 +261,40 @@ async def amessages(
extra_headers: dict[str, object] | None, # mutable-ok: native bridge retains and updates Python argument objects
timeout: float | httpx.Timeout | None,
arguments: dict[str, object] | None = None, # mutable-ok: native bridge retains and updates Python argument objects
request_arguments: Mapping[str, object] | None = None,
logging_obj: object | None = None,
litellm_params: GenericLiteLLMParams | None = None,
messages: object = None,
lifecycle_owner: LifecycleOwner = LifecycleOwner.BRIDGE,
) -> AnthropicMessagesResponse | None:
implementation: Final = load_rust_amessages()
if implementation is None:
return None
return await implementation(
arguments=_arguments(
arguments or {}, model, body, api_key, api_base, custom_llm_provider, extra_headers, timeout
)
call_arguments: Final = _arguments(
request_arguments if request_arguments is not None else arguments,
model,
body,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout,
logging_obj,
litellm_params,
messages,
lifecycle_owner,
)
try:
return await implementation(arguments=call_arguments)
except Exception as error: # noqa: BLE001 # only explicit declines before lifecycle setup may fall back
exceptions: Final = native_exception_types()
if (
exceptions is not None
and isinstance(error, exceptions[0])
and not call_arguments.get(LIFECYCLE_STARTED_KEY)
):
return None
raise map_native_error(error, call_arguments, "messages")
class _MessagesLifecycle(NativeLifecycle, Protocol):
@ -242,7 +308,6 @@ class _MessagesBindings(NativeLifecycleBindings, Protocol):
] # mutable-ok: native bridge retains and updates Python argument objects
send: Callable[[object], Awaitable[AnthropicMessagesResponse]]
send_sync: Callable[[object], AnthropicMessagesResponse]
committed_failure: Callable[[], None]
class _MessagesHost:
@ -261,7 +326,8 @@ class _MessagesHost:
)
self.asynchronous: bool = asynchronous
self.logger: object | None = arguments.get(LOGGING_OBJECT_KEY)
self.lifecycle_owned: bool = self.logger is None
self.lifecycle_owned: Final = owns_lifecycle(arguments)
self.arguments[LIFECYCLE_STARTED_KEY] = True
self.state: object | None = None
self.response: object = None
self.error: BaseException | None = None
@ -344,21 +410,20 @@ class _MessagesHost:
def sync_failure(self) -> object:
return self.terminal("sync_failure", self.error)
def async_failure(self) -> object:
return self.terminal("async_failure", self.error)
async def async_failure(self) -> object:
result: Final = self.terminal("async_failure", self.error)
return await result if isinstance(result, Awaitable) else result
def restore(self) -> None:
if not self.streaming and self.lifecycle_owned:
restore_correlation_context(self.logger)
def advance(self, outcome: NativeOutcome, error: BaseException | None = None) -> None:
advance_host(self, outcome, error)
advance_host(self, outcome, map_native_error(error, self.arguments, "messages"))
def result(self) -> object:
if self.machine.complete():
return self.response
if self.machine.failed_after_provider_response():
self.bindings.committed_failure()
return host_result(self)

View file

@ -2,12 +2,9 @@
from __future__ import annotations
import traceback
from collections.abc import Awaitable, Callable, Mapping
from contextvars import copy_context
from datetime import datetime
from typing import Final, Protocol, cast # noqa: TID251 # native extension exposes dynamically typed callables
from uuid import uuid4
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.rust_bridge._lifecycle import (
@ -23,6 +20,8 @@ from litellm.rust_bridge._lifecycle import (
drive_async,
drive_sync,
host_result,
initialize_logging,
invoke_terminal,
restore_correlation_context,
)
from litellm.rust_bridge.bindings import NativeBinding
@ -78,210 +77,6 @@ async def aocr(
return await implementation(arguments)
def initialize_logging(
arguments: dict[str, object], asynchronous: bool, route: str = "ocr"
) -> object: # mutable-ok: native bridge retains and updates Python argument objects
import litellm
from litellm import utils
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils import litellm_logging
from litellm.litellm_core_utils.coroutine_checker import coroutine_checker
from litellm.litellm_core_utils.litellm_logging import Logging, set_callbacks
supplied: Final = arguments.get(LOGGING_OBJECT_KEY)
if supplied is not None:
return supplied
callbacks: Final = tuple( # cast-ok: callback registry accepts heterogeneous legacy callback objects
dict.fromkeys(
utils.get_dynamic_callbacks(
cast( # cast-ok: callback registry accepts heterogeneous legacy callback objects
list, arguments.get("callbacks")
) # cast-ok: callback registry accepts heterogeneous legacy callback objects
) # cast-ok: callback registry accepts heterogeneous legacy callback objects
) # cast-ok: callback registry accepts heterogeneous legacy callback objects # mutable-ok: deduplication uses dict keys
)
success: Final = tuple( # cast-ok: per-call callback list is a legacy untyped boundary
dict.fromkeys(
(
*callbacks,
*cast( # cast-ok: per-call callback list is a legacy untyped boundary
list, arguments.get("success_callback") or ()
), # cast-ok: per-call callback list is a legacy untyped boundary
) # cast-ok: per-call callback list is a legacy untyped boundary
) # cast-ok: per-call callback list is a legacy untyped boundary # mutable-ok: deduplication uses dict keys
)
failure: Final = tuple( # cast-ok: per-call callback list is a legacy untyped boundary
dict.fromkeys(
(
*callbacks,
*cast( # cast-ok: per-call callback list is a legacy untyped boundary
list, arguments.get("failure_callback") or ()
), # cast-ok: per-call callback list is a legacy untyped boundary
) # cast-ok: per-call callback list is a legacy untyped boundary
) # cast-ok: per-call callback list is a legacy untyped boundary # mutable-ok: deduplication uses dict keys
)
configured: Final = tuple(
dict.fromkeys(
(
*litellm.input_callback,
*litellm.success_callback,
*litellm.failure_callback,
*litellm._async_success_callback, # pyright: ignore[reportPrivateUsage] # callback registry has no public accessor
*litellm._async_failure_callback, # pyright: ignore[reportPrivateUsage] # callback registry has no public accessor
*success,
*failure,
)
)
)
uninitialized: Final = [ # mutable-ok: set_callbacks requires a mutable callback list
cb
for cb in configured
if isinstance(cb, str)
and (
cb not in litellm._known_custom_logger_compatible_callbacks # pyright: ignore[reportPrivateUsage] # callback compatibility registry has no public accessor
or cb in litellm.input_callback + litellm.success_callback + litellm.failure_callback
)
and cb not in (utils.callback_list or ())
]
if uninitialized:
set_callbacks(uninitialized, function_id=arguments.get("id"))
utils.callback_list = list( # mutable-ok: global callback registry is mutable
dict.fromkeys((*(utils.callback_list or ()), *uninitialized))
) # mutable-ok: global callback registry is mutable
if litellm_logging.customLogger is None: # pyright: ignore[reportUnnecessaryComparison] # runtime plugin registry can be reset to None
set_callbacks(
[cb for cb in configured if callable(cb)], # mutable-ok: set_callbacks requires a mutable callback list
function_id=arguments.get("id"), # mutable-ok: set_callbacks requires a mutable callback list
) # mutable-ok: set_callbacks requires a mutable callback list
for event, registered, add_async in (
("input", litellm.input_callback, litellm.logging_callback_manager.add_litellm_input_callback),
("success", litellm.success_callback, litellm.logging_callback_manager.add_litellm_async_success_callback),
("failure", litellm.failure_callback, litellm.logging_callback_manager.add_litellm_async_failure_callback),
):
for cb in tuple(registered):
if coroutine_checker.is_async_callable(cb) or (event == "success" and cb in ("dynamodb", "openmeter")):
if cb not in getattr(litellm, f"_async_{event}_callback"):
add_async(cb)
registered.remove(cb)
elif event != "input" and isinstance(cb, str) and cb in litellm._known_custom_logger_compatible_callbacks: # pyright: ignore[reportPrivateUsage] # callback compatibility registry has no public accessor
utils._add_custom_logger_callback_to_specific_event(cb, event) # pyright: ignore[reportPrivateUsage] # callback manager only exposes this internal registration path
for event, registered, add_sync in (
("success", litellm._async_success_callback, litellm.logging_callback_manager.add_litellm_success_callback), # pyright: ignore[reportPrivateUsage] # callback registry has no public accessor
("failure", litellm._async_failure_callback, litellm.logging_callback_manager.add_litellm_failure_callback), # pyright: ignore[reportPrivateUsage] # callback registry has no public accessor
):
for cb in tuple(registered):
if callable(cb) and not isinstance(cb, CustomLogger) and not coroutine_checker.is_async_callable(cb):
if cb not in getattr(litellm, f"{event}_callback"):
add_sync(cb)
registered.remove(cb)
call_id: Final = str(arguments.get("litellm_call_id") or uuid4())
logger: Final = Logging(
model=str(arguments["model"]),
messages="default-message-value",
stream=False,
call_type=f"a{route}" if asynchronous else route,
start_time=datetime.now(), # noqa: DTZ005 # Logging preserves the legacy naive timestamp contract
litellm_call_id=call_id,
function_id=str(arguments.get("id") or ""),
litellm_trace_id=cast( # cast-ok: public call argument is validated by Logging
str | None, arguments.get("litellm_trace_id")
), # cast-ok: public call argument is validated by Logging
dynamic_input_callbacks=[ # mutable-ok: Logging callback configuration is mutable
cb for cb in callbacks if cb not in litellm.input_callback and not coroutine_checker.is_async_callable(cb)
],
dynamic_success_callbacks=[ # mutable-ok: Logging callback configuration is mutable
cb for cb in success if not coroutine_checker.is_async_callable(cb) and cb not in ("dynamodb", "s3")
],
dynamic_async_success_callbacks=[ # mutable-ok: Logging callback configuration is mutable
cb
for cb in success
if coroutine_checker.is_async_callable(cb) or isinstance(cb, CustomLogger) or cb in ("dynamodb", "s3")
],
dynamic_failure_callbacks=[ # mutable-ok: Logging callback configuration is mutable
cb for cb in failure if not coroutine_checker.is_async_callable(cb)
], # mutable-ok: Logging callback configuration is mutable
dynamic_async_failure_callbacks=[ # mutable-ok: Logging callback configuration is mutable
cb for cb in failure if coroutine_checker.is_async_callable(cb) or isinstance(cb, CustomLogger)
],
kwargs=arguments,
supports_correlation_logging=asynchronous,
)
logger.dynamic_input_callbacks = [ # mutable-ok: remove callbacks promoted to the global registry
cb for cb in dict.fromkeys(logger.dynamic_input_callbacks or ()) if cb not in litellm.input_callback
]
arguments["litellm_call_id"] = call_id
arguments[LOGGING_OBJECT_KEY] = logger
return logger
def invoke_terminal(
action: TerminalAction,
roots: object,
logger: object,
record: Mapping[str, object] | None,
value: object,
fallback_start_time: datetime,
fallback_end_time: datetime,
) -> object:
from litellm import utils
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
logging: Final = cast( # cast-ok: callers may supply a Logging-compatible test or plugin implementation
Logging, logger
)
timing_value: Final = record.get("timing") if record is not None else None
timing: Final = timing_value if isinstance(timing_value, Mapping) else None
start_value: Final = timing.get("start_time") if timing is not None else None
end_value: Final = timing.get("end_time") if timing is not None else None
start_time: Final = (
datetime.fromtimestamp(start_value, tz=fallback_start_time.tzinfo)
if isinstance(start_value, (int, float))
else fallback_start_time
)
end_time: Final = (
datetime.fromtimestamp(end_value, tz=fallback_end_time.tzinfo)
if isinstance(end_value, (int, float))
else fallback_end_time
)
if action == "sync_success":
def run() -> None:
_retained: Final = roots
logging.success_handler(value, start_time, end_time)
return utils.executor.submit(copy_context().run, run)
if action == "async_success":
async def run_async() -> None:
_retained: Final = roots
await logging.async_success_handler(value, start_time, end_time)
def enqueue() -> None:
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(async_coroutine=run_async())
if getattr(logging, "_defer_async_logging", False) is True:
logging._enqueue_deferred_logging = enqueue # pyright: ignore[reportPrivateUsage] # preserves Logging's deferred callback contract
else:
enqueue()
return None
if action == "sync_success_if_needed":
if logging._should_run_sync_callbacks_for_async_calls(): # pyright: ignore[reportPrivateUsage] # preserves Logging's async callback policy
def run() -> None:
_retained: Final = roots
logging.success_handler(value, start_time, end_time)
return utils.executor.submit(copy_context().run, run)
return None
exception: Final = cast(Exception, value) # cast-ok: Rust routes terminal failure values as Python exceptions
trace: Final = "".join(traceback.format_exception(type(exception), exception, exception.__traceback__))
if action == "sync_failure":
logging.failure_handler(exception, trace, start_time, end_time)
return None
return logging.async_failure_handler(exception, trace, start_time, end_time)
class _OcrLifecycle(NativeLifecycle, Protocol):
def identity(self) -> tuple[str, str | None]: ...
@ -395,8 +190,9 @@ class _OcrHost:
def sync_failure(self) -> object:
return self.terminal("sync_failure", self.error)
def async_failure(self) -> object:
return self.terminal("async_failure", self.error)
async def async_failure(self) -> object:
result: Final = self.terminal("async_failure", self.error)
return await result if isinstance(result, Awaitable) else result
def restore(self) -> None:
restore_correlation_context(self.logger)

View file

@ -137,20 +137,24 @@ def _required_reason(result: RustDeclined | RustUnavailable) -> str:
def raise_upstream(error: BaseException, context: BridgeErrorContext) -> NoReturn:
raise upstream_error(error, context) from error
def upstream_error(error: BaseException, context: BridgeErrorContext) -> Exception:
args: Final[tuple[object, ...]] = error.args
status_value: Final = args[0] if args else 0
message_value: Final = args[1] if len(args) > 1 else str(error)
status: Final = status_value if isinstance(status_value, int) else 0
message: Final = message_value if isinstance(message_value, str) else str(message_value)
if status == 500:
raise InternalServerError(
if status == 500 and context.route != "chat completions":
return InternalServerError(
message=f"litellm rust {context.route}: {message}",
llm_provider=context.provider,
model=context.model,
) from error
raise APIError(
)
return APIError(
status_code=status or 500,
message=f"litellm rust {context.route}: {message}",
llm_provider=context.provider,
model=context.model,
) from error
)

View file

@ -154,6 +154,7 @@ def test_messages_wrapper_forwards_args_and_converts_timeout():
assert response == FAKE_MESSAGES_RESPONSE
assert bridge.calls[0] == {
"_rust_lifecycle_owner": "bridge",
"model": "claude-sonnet-4-5",
"body": REQUEST_BODY,
"api_key": "sk-azure",
@ -222,14 +223,14 @@ async def test_gate_invokes_rust_and_marks_response_header():
@pytest.mark.asyncio
async def test_gate_falls_back_to_python_when_bridge_raises():
async def test_gate_propagates_unknown_errors_without_replaying():
bridge = RaisingAsyncMessages()
litellm.rust(True)
rust_messages.set_rust_messages(amessages=bridge)
response = await _gate()
with pytest.raises(RuntimeError, match="upstream request failed"):
await _gate()
assert response is None
assert bridge.calls == 1
@ -253,6 +254,31 @@ async def test_gate_does_not_fall_back_after_provider_commit(monkeypatch):
assert calls == 1
@pytest.mark.asyncio
@pytest.mark.parametrize("started", [False, True])
async def test_gate_only_accepts_declines_before_callback_setup(monkeypatch, started):
from litellm.rust_bridge._lifecycle import LIFECYCLE_STARTED_KEY
error = _DeclinedMessagesError("declined")
async def decline(arguments):
if started:
arguments[LIFECYCLE_STARTED_KEY] = True
raise error
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _NativeExceptions())
rust_messages.set_rust_messages(amessages=decline)
litellm.rust(True)
if not started:
assert await _gate() is None
return
with pytest.raises(_DeclinedMessagesError) as raised:
await _gate()
assert raised.value is error
@pytest.mark.asyncio
async def test_gate_skips_rust_when_flag_absent():
bridge = ExplodingAsyncMessages()

View file

@ -10,8 +10,8 @@ from __future__ import annotations
import pytest
import litellm
from litellm.rust_bridge import configuration
from litellm.rust_bridge import chat_completions as bridge
from litellm.rust_bridge import configuration
from litellm.types.utils import ModelResponse
RUST_RESPONSE = {
@ -241,6 +241,32 @@ def _call_kwargs(model_response: ModelResponse) -> dict:
class TestSyncCall:
def test_retains_explicit_logging_and_request_context(self):
native = _RecordingCall()
bridge.set_rust_chat_completions(chat_completions=native)
logger = object()
metadata = {"trace": []}
response = ModelResponse()
result = bridge.chat_completions(
**_call_kwargs(response), logging_obj=logger, litellm_params={"metadata": metadata}
)
assert native.calls[0]["litellm_logging_obj"] is logger
assert native.calls[0]["metadata"] is metadata
assert native.calls[0]["messages"] is MESSAGES
assert result is response
def test_unknown_failure_without_native_exception_types_is_not_replayed(self, monkeypatch):
_hide_native_bridge(monkeypatch)
error = ValueError("callback failed")
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=error))
with pytest.raises(ValueError, match="callback failed") as raised:
bridge.chat_completions(**_call_kwargs(ModelResponse()))
assert raised.value is error
def test_builds_a_model_response_and_stamps_the_rust_header(self):
native = _RecordingCall()
bridge.set_rust_chat_completions(chat_completions=native)
@ -357,6 +383,22 @@ class TestFailureClassification:
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming")))
assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None
def test_a_decline_after_callback_setup_is_not_replayed(self):
from litellm.rust_bridge._lifecycle import LIFECYCLE_STARTED_KEY
error = _FakeDeclined("raised by callback")
def started(arguments):
arguments[LIFECYCLE_STARTED_KEY] = True
raise error
bridge.set_rust_chat_completions(chat_completions=started)
with pytest.raises(_FakeDeclined) as raised:
bridge.chat_completions(**_call_kwargs(ModelResponse()))
assert raised.value is error
def test_an_upstream_failure_is_surfaced_with_its_status(self):
from litellm.exceptions import APIError

View file

@ -6,7 +6,7 @@ from typing import Final
import pytest
from litellm.rust_bridge._lifecycle import NativeOutcome, drive_async, drive_sync
from litellm.rust_bridge._lifecycle import NativeOutcome, deployment_pre, drive_async, drive_sync
class _Machine:
@ -89,3 +89,21 @@ async def test_drive_async_awaits_the_selected_operation() -> None:
assert await drive_async(host) == "complete"
assert completed.is_set()
assert host.machine.outcome is NativeOutcome.SUCCESS
@pytest.mark.asyncio
async def test_deployment_pre_retains_the_hooks_replacement(monkeypatch: pytest.MonkeyPatch) -> None:
import litellm
from litellm.integrations.custom_logger import CustomLogger
replacement: Final[dict[str, object]] = {"metadata": {"observed": False}}
class ReplacingLogger(CustomLogger):
async def async_pre_call_deployment_hook(self, kwargs, call_type):
return replacement
monkeypatch.setattr(litellm, "callbacks", [ReplacingLogger()])
result: Final = await deployment_pre({}, "acompletion")
assert result is replacement

View file

@ -63,13 +63,7 @@ async def test_generic_api_logger_exports_success_over_http(route: Route, provid
"route",
(
OCR_ASYNC,
pytest.param(
MESSAGES_ROUTE,
marks=pytest.mark.xfail(
reason="Rust Messages retries HTTP 500 despite num_retries=0",
strict=True,
),
),
MESSAGES_ROUTE,
),
ids=route_id,
)

View file

@ -192,10 +192,6 @@ async def test_messages_logging_drain_waits_for_suspended_callback(messages_serv
@pytest.mark.asyncio
@pytest.mark.xfail(
reason="Rust Messages sends two provider requests before invoking failure callbacks",
strict=True,
)
async def test_messages_failure_callbacks_receive_original_provider_error(messages_server: RecordingServer) -> None:
messages_server.default_response = ResponseSpec(body={"error": {"message": "provider unavailable"}}, status=500)
messages_server.expected_requests = None
@ -391,3 +387,53 @@ async def test_messages_cancelled_call_runs_no_terminal_callbacks(messages_serve
assert "async_log_success_event" not in recorder.names
assert "log_failure_event" not in recorder.names
assert "async_log_failure_event" not in recorder.names
@pytest.mark.asyncio
@pytest.mark.parametrize("status", [200, 500])
async def test_whole_call_with_supplied_logger_still_owns_terminal_callbacks(
messages_server: RecordingServer, status: int
) -> None:
from litellm.rust_bridge import messages as bridge
events: Final = []
class SuppliedLogger:
def pre_call(self, **kwargs):
events.append(("pre", kwargs))
async def async_success_handler(self, response, start, end):
events.append(("success", response))
def _should_run_sync_callbacks_for_async_calls(self):
return False
def failure_handler(self, error, trace, start, end):
events.append(("sync_failure", error))
async def async_failure_handler(self, error, trace, start, end):
events.append(("async_failure", error))
messages_server.default_response = ResponseSpec(body=MESSAGES_RESPONSE, status=status)
arguments: Final = {
"model": "claude-opus-5",
"body": {"model": "claude-opus-5", "messages": MESSAGES, "max_tokens": 64},
"api_key": "test",
"api_base": messages_server.base_url,
"custom_llm_provider": "anthropic",
"extra_headers": {},
"timeout": 5.0,
"logging_obj": SuppliedLogger(),
}
if status == 200:
response: Final = await bridge.amessages(**arguments)
await drain_logging()
assert [name for name, value in events] == ["pre", "success"]
assert events[1][1] is response
return
with pytest.raises(litellm.InternalServerError) as raised:
await bridge.amessages(**arguments)
assert [name for name, value in events] == ["pre", "sync_failure", "async_failure"]
assert events[1][1] is raised.value
assert events[2][1] is raised.value

View file

@ -64,3 +64,87 @@ def test_public_ocr_uses_python_transport_when_disabled(ocr_server: RecordingSer
assert response.pages[0].markdown == "native OCR response"
assert ocr_server.requests[0].headers["accept-encoding"] != "identity"
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("provider", ["anthropic", "bedrock"])
@pytest.mark.parametrize("rebind_logging_view", [False, True])
@pytest.mark.parametrize("status", [200, 429])
async def test_chat_retains_callback_edits_through_public_dispatch(
recording_server: RecordingServer, asynchronous: bool, provider: str, rebind_logging_view: bool, status: int
) -> None:
import threading
from tests.test_litellm_rust.callback_recorder import RecordingLogger
from tests.test_litellm_rust.contracts import MESSAGES_RESPONSE, request_body, request_headers
recording_server.default_response = ResponseSpec(
status=status,
body=(
MESSAGES_RESPONSE
if provider == "anthropic"
else {
"output": {"message": {"role": "assistant", "content": [{"text": "native chat"}]}},
"stopReason": "end_turn",
"usage": {"inputTokens": 5, "outputTokens": 4, "totalTokens": 9},
"metrics": {"latencyMs": 1},
}
),
)
caller_thread: Final = threading.current_thread()
observations: Final = []
class EditingLogger(RecordingLogger):
def log_pre_api_call(self, model, messages, kwargs):
super().log_pre_api_call(model, messages, kwargs)
body = request_body(kwargs)
headers = request_headers(kwargs)
body["messages"][0]["content"] = "edited by callback"
body["max_tokens" if provider == "anthropic" else "maxTokens"] = 32
headers["x-retained-callback"] = "original"
observations.append((threading.current_thread(), body, headers))
if rebind_logging_view:
kwargs["additional_args"]["complete_input_dict"] = {"replacement": True}
kwargs["additional_args"]["headers"] = {"x-retained-callback": "replacement"}
recorder: Final = EditingLogger()
kwargs: Final = {
"model": "anthropic/claude-opus-5" if provider == "anthropic" else "bedrock/anthropic.claude-opus-5",
"messages": [{"role": "user", "content": "original"}],
"max_tokens": 64,
"api_key": "test-key",
"api_base": recording_server.base_url,
"callbacks": [recorder],
"num_retries": 0,
**(
{"aws_access_key_id": "test", "aws_secret_access_key": "test", "aws_region_name": "us-east-1"}
if provider == "bedrock"
else {}
),
}
if status != 200:
with pytest.raises((litellm.APIError, litellm.RateLimitError)) as raised:
await litellm.acompletion(**kwargs) if asynchronous else litellm.completion(**kwargs)
failure_event: Final = "async_log_failure_event" if asynchronous else "log_failure_event"
failures: Final = await recorder.wait_for_async(failure_event)
assert raised.value.status_code == status
assert failures[0].kwargs["exception"] is raised.value
assert recorder.names.count(failure_event) == 1
assert recorder.names.count("log_pre_api_call") == 1
assert len(recording_server.requests) == 1
return
response: Final = await litellm.acompletion(**kwargs) if asynchronous else litellm.completion(**kwargs)
event_name: Final = "async_log_success_event" if asynchronous else "log_success_event"
events: Final = await recorder.wait_for_async(event_name)
assert response._hidden_params["additional_headers"]["x-litellm-rust"] == "true"
assert recorder.names.count("log_pre_api_call") == 1
assert recorder.names.count(event_name) == 1
assert observations[0][0] is caller_thread
assert events[0].response is response
assert recording_server.requests[0].body["messages"][0]["content"][0]["text"] == "edited by callback"
assert recording_server.requests[0].headers["x-retained-callback"] == "original"
body: Final = recording_server.requests[0].body
assert (body["max_tokens"] if provider == "anthropic" else body["inferenceConfig"]["maxTokens"]) == 32