diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs b/litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs index f78d4f87c62..1a1ecbce116 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs @@ -1,4 +1,3 @@ -use pyo3::exceptions::PyBaseException; use pyo3::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; use pyo3::types::PyDict; @@ -56,42 +55,3 @@ pub(super) fn is_internal_call(py: Python<'_>) -> PyResult { .call_method0("get")? .extract() } - -pub(super) struct DeploymentHooks; - -impl DeploymentHooks { - pub(super) fn before_call( - py: Python<'_>, - kwargs: &Py, - call_type: &str, - ) -> PyResult> { - py.import("litellm.utils")? - .getattr("async_pre_call_deployment_hook")? - .call1((kwargs, call_type)) - .map(Bound::unbind) - } - - pub(super) fn after_success( - py: Python<'_>, - kwargs: &Py, - response: &Option>, - call_type: &str, - ) -> PyResult> { - py.import("litellm.utils")? - .getattr("async_post_call_success_deployment_hook")? - .call1((kwargs, response, call_type)) - .map(Bound::unbind) - } - - pub(super) fn after_failure( - py: Python<'_>, - kwargs: &Py, - error: &Py, - call_type: &str, - ) -> PyResult> { - py.import("litellm.utils")? - .getattr("async_post_call_failure_deployment_hook")? - .call1((kwargs, error, call_type)) - .map(Bound::unbind) - } -} diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/dispatch.rs b/litellm-rust/crates/python-bridge/src/lifecycle/dispatch.rs index 4215b8e6c2d..706a31ba7d7 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle/dispatch.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/dispatch.rs @@ -340,7 +340,6 @@ impl Runner { let logger = self.job.logger.object(py); let target = self.job.targets.object(py, invocation.target); let kind = self.job.targets.kind(invocation.target); - let awaits = matches!(invocation.delivery, Delivery::Await | Delivery::Background); let value = match (invocation.method, kind) { (CallbackMethod::LoggingHook, CallbackKind::CustomLogger) => { let replaced = @@ -402,10 +401,19 @@ impl Runner { ))?, _ => return Ok(None), }; - if awaits && !value.is_none() { - return Ok(Some(value.unbind())); + match invocation.delivery { + Delivery::Inline | Delivery::Worker => Ok(None), + Delivery::Await | Delivery::Background if value.is_none() => Ok(None), + Delivery::Await | Delivery::Background if value.hasattr("__await__")? => { + Ok(Some(value.unbind())) + } + Delivery::Await | Delivery::Background => Err(PyRuntimeError::new_err(format!( + "{:?} leaf for {:?} returned a non-awaitable {}", + invocation.method, + self.job.family, + value.get_type().name()? + ))), } - Ok(None) } fn accept( @@ -648,6 +656,188 @@ pub(super) fn family_targets( Ok((targets, ordered)) } +pub(super) enum DeploymentEvent { + PreCall, + PostCall, + Failure { + request: Py, + exception: Py, + fallback_depth: Py, + }, +} + +pub(super) struct DeploymentBody { + logger: PythonLogger, + targets: Targets, + cursor: DispatchCursor, + event: DeploymentEvent, + call_type: &'static str, + current: Py, + kwargs: Py, + pending: Option, +} + +impl DeploymentBody { + pub(super) fn start( + py: Python<'_>, + logger: &PythonLogger, + family: CallbackFamily, + call_type: &'static str, + kwargs: &Py, + current: Py, + error: Option<&Py>, + ) -> PyResult { + let (targets, ids) = family_targets(py, logger, family)?; + let event = match family { + CallbackFamily::DeploymentPreCall => DeploymentEvent::PreCall, + CallbackFamily::DeploymentPostCall => DeploymentEvent::PostCall, + CallbackFamily::DeploymentFailure => { + let exception = error.ok_or_else(super::missing_state)?; + let view = leaves(py)? + .getattr("failure_deployment_hook_view")? + .call1((kwargs, exception))?; + let (request, snapshot, fallback_depth): (Py, Py, Py) = + view.extract()?; + DeploymentEvent::Failure { + request, + exception: snapshot + .into_bound(py) + .cast_into::()? + .unbind(), + fallback_depth, + } + } + _ => return Err(super::missing_state()), + }; + Ok(Self { + logger: logger.clone_ref(py), + targets, + cursor: DispatchCursor::start(family, ids, false, false), + event, + call_type, + current, + kwargs: kwargs.clone_ref(py), + pending: None, + }) + } + + fn invoke(&self, py: Python<'_>, target: CallbackId) -> PyResult> { + let leaves = leaves(py)?; + let object = self.targets.object(py, target); + let awaitable = match &self.event { + DeploymentEvent::PreCall => leaves.getattr("pre_call_deployment_hook")?.call1(( + object, + &self.current, + self.call_type, + ))?, + DeploymentEvent::PostCall => leaves + .getattr("post_call_success_deployment_hook")? + .call1((object, &self.kwargs, &self.current, self.call_type))?, + DeploymentEvent::Failure { + request, + exception, + fallback_depth, + } => leaves + .getattr("post_call_failure_deployment_hook")? + .call1((object, request, exception, self.call_type, fallback_depth))?, + }; + Ok(awaitable.unbind()) + } + + fn accept( + &mut self, + py: Python<'_>, + target: CallbackId, + result: PyResult>, + ) -> PyResult<()> { + match result { + Ok(value) => { + if !value.is_none(py) && !matches!(self.event, DeploymentEvent::Failure { .. }) { + self.current = value; + } + self.cursor.accept(InvocationOutcome::Completed); + Ok(()) + } + Err(error) + if matches!(self.event, DeploymentEvent::Failure { .. }) + && error.is_instance_of::(py) => + { + let object = self.targets.object(py, target); + leaves(py)? + .getattr("report_deployment_failure_hook_error")? + .call1((object, error.value(py)))?; + self.cursor.accept(InvocationOutcome::Failed); + Ok(()) + } + Err(error) => Err(error), + } + } +} + +impl super::handle::ExecutionBody for DeploymentBody { + fn resume( + &mut self, + result: Option>>, + ) -> PyResult { + Python::attach(|py| { + if let Some(result) = result { + let target = self.pending.take().ok_or_else(super::missing_state)?; + self.accept(py, target, result)?; + } + let mut facts = DeploymentEligibility(&self.targets); + loop { + match self.cursor.next(&mut facts) { + DispatchStep::Invoke(invocation) => { + let awaitable = self.invoke(py, invocation.target)?; + self.pending = Some(invocation.target); + return Ok(super::handle::ExecutionStep::Await(awaitable)); + } + DispatchStep::Complete { .. } => { + return Ok(super::handle::ExecutionStep::Return( + self.current.clone_ref(py), + )); + } + DispatchStep::PrepareLogging | DispatchStep::MarkLogged(_) => {} + } + } + }) + } + + fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + self.logger.traverse(visit)?; + self.targets.traverse(visit)?; + visit.call(&self.current)?; + visit.call(&self.kwargs)?; + if let DeploymentEvent::Failure { + request, + exception, + fallback_depth, + } = &self.event + { + visit.call(request)?; + visit.call(exception)?; + visit.call(fallback_depth)?; + } + Ok(()) + } +} + +struct DeploymentEligibility<'a>(&'a Targets); + +impl DispatchFacts for DeploymentEligibility<'_> { + fn eligible(&mut self, target: CallbackId, _: CallbackMethod) -> bool { + self.0.kind(target) == CallbackKind::CustomLogger + } +} + +pub(super) fn deployment_coroutine(py: Python<'_>, body: DeploymentBody) -> PyResult> { + let execution = Py::new(py, super::handle::Execution::new(body))?; + py.import("litellm.rust_bridge.lifecycle")? + .getattr("drive")? + .call1((execution,)) + .map(Bound::unbind) +} + #[cfg(test)] mod tests { use super::*; diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs b/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs index ca247138083..fa07262fe7b 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs @@ -26,7 +26,6 @@ mod setup; use crate::execution::{poll_async_value, run_async_value, run_sync_value}; pub(crate) use arguments::{BoundArguments, Signature}; -use bindings::DeploymentHooks; pub(crate) use bindings::PythonLogger; use handle::{Execution, ExecutionBody, ExecutionStep}; @@ -324,30 +323,34 @@ impl PythonCallState { match phase { HostPhase::Setup => self.setup(py)?, HostPhase::DeploymentPreCall => { - return Ok(HostStep::Suspend(DeploymentHooks::before_call( + let copied = self.kwargs.bind(py).copy()?.into_any().unbind(); + return Ok(HostStep::Suspend(self.deployment( py, - &self.kwargs, - self.call_type, + CallbackFamily::DeploymentPreCall, + copied, )?)); } HostPhase::Prepare => self.prepare(py)?, HostPhase::DeploymentPostCall => { - return Ok(HostStep::Suspend(DeploymentHooks::after_success( + let response = self + .response + .as_ref() + .map(|value| value.clone_ref(py)) + .unwrap_or_else(|| py.None()); + return Ok(HostStep::Suspend(self.deployment( py, - &self.kwargs, - &self.response, - self.call_type, + CallbackFamily::DeploymentPostCall, + response, )?)); } HostPhase::Finalize => self.finalize(py)?, HostPhase::Success => self.dispatch_success(py)?, HostPhase::DeploymentFailure => { - if let Some(error) = &self.error { - return Ok(HostStep::Suspend(DeploymentHooks::after_failure( + if self.error.is_some() { + return Ok(HostStep::Suspend(self.deployment( py, - &self.kwargs, - error, - self.call_type, + CallbackFamily::DeploymentFailure, + py.None(), )?)); } } @@ -366,6 +369,24 @@ impl PythonCallState { Ok(HostStep::Ready(py.None())) } + fn deployment( + &self, + py: Python<'_>, + family: CallbackFamily, + current: Py, + ) -> PyResult> { + let body = dispatch::DeploymentBody::start( + py, + self.logger()?, + family, + self.call_type, + &self.kwargs, + current, + self.error.as_ref(), + )?; + dispatch::deployment_coroutine(py, body) + } + fn accept(&mut self, py: Python<'_>, phase: HostPhase, value: Py) -> PyResult<()> { match phase { HostPhase::DeploymentPreCall => { diff --git a/litellm/rust_bridge/leaves.py b/litellm/rust_bridge/leaves.py index 6b9020aba78..7707a8938f1 100644 --- a/litellm/rust_bridge/leaves.py +++ b/litellm/rust_bridge/leaves.py @@ -31,15 +31,15 @@ from litellm.integrations.custom_logger import CustomLogger if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging + from litellm.types.utils import StandardCallbackDynamicParams TerminalFamily: TypeAlias = Literal["sync_success", "async_success", "sync_failure", "async_failure"] -Family: TypeAlias = Literal["request", TerminalFamily] +Family: TypeAlias = Literal["request", "deployment", TerminalFamily] Details: TypeAlias = dict[ str, object ] # mutable-ok: model_call_details is the shared mutable envelope callbacks write to Timestamp: TypeAlias = datetime.datetime LegacyCall: TypeAlias = Callable[..., object] -LegacyAsyncCall: TypeAlias = Callable[..., Awaitable[None]] class LoggerView(Protocol): @@ -51,7 +51,7 @@ class LoggerView(Protocol): completion_start_time: Timestamp | None model_call_details: Details log_raw_request_response: bool - standard_callback_dynamic_params: object + standard_callback_dynamic_params: StandardCallbackDynamicParams standard_built_in_tools_params: object def record_api_call_start_time(self) -> None: ... @@ -176,12 +176,10 @@ def _integration(callback: CustomLogger) -> IntegrationView: return cast(IntegrationView, callback) # cast-ok: legacy CustomLogger methods are untyped -def _legacy_module() -> Mapping[str, object]: +def _singleton(name: str) -> object: from litellm.litellm_core_utils import litellm_logging - return cast( # cast-ok: module globals hold the legacy integration singletons - Mapping[str, object], vars(litellm_logging) - ) + return getattr(litellm_logging, name, None) # pyright: ignore[reportAny] # legacy module globals are rebound at runtime def _print_verbose() -> LegacyCall: @@ -190,14 +188,6 @@ def _print_verbose() -> LegacyCall: return cast(LegacyCall, litellm_logging.print_verbose) # cast-ok: legacy debug printer is untyped -def _method(target: object, name: str) -> LegacyCall: - return _call_of(_attribute(target, name)) - - -def _async_method(target: object, name: str) -> LegacyAsyncCall: - return cast(LegacyAsyncCall, _attribute(target, name)) # cast-ok: legacy integration singletons are untyped - - def _redact_string(value: str) -> str: from litellm.litellm_core_utils import litellm_logging @@ -303,25 +293,27 @@ def log_post_api_call(logger: Logging, callback: CustomLogger) -> None: def dispatch_named_request(logger: Logging, name: str, event: Literal["pre_api_call", "post_api_call"]) -> None: + from litellm.integrations.supabase import Supabase + view: Final = _logger(logger) - module: Final = _legacy_module() - if name == "supabase" and event == "pre_api_call" and (client := module.get("supabaseClient")) is not None: - details: Final = view.model_call_details - _method(client, "input_log_event")( - model=view.model, - messages=view.messages, - end_user=details.get("user", "default"), - litellm_call_id=details["litellm_call_id"], - print_verbose=_print_verbose(), - ) - if name == "sentry" and (add_breadcrumb := module.get("add_breadcrumb")) is not None: - cast(LegacyCall, add_breadcrumb)( # cast-ok: legacy sentry hook - category="litellm.llm_call", message=f"Model Call Details {event}: {view.model_call_details}", level="info" - ) + details: Final = view.model_call_details + match name: + case "supabase" if event == "pre_api_call" and isinstance(client := _singleton("supabaseClient"), Supabase): + client.input_log_event( + model=view.model, + messages=view.messages, + end_user=details.get("user", "default"), + litellm_call_id=details["litellm_call_id"], + print_verbose=_print_verbose(), + ) + case "sentry" if callable(add_breadcrumb := _singleton("add_breadcrumb")): + add_breadcrumb(category="litellm.llm_call", message=f"Model Call Details {event}: {details}", level="info") + case _: + return def dispatch_callable_request(logger: Logging, callback: LegacyCall) -> None: - custom: Final = _legacy_module().get("customLogger") + custom: Final = _singleton("customLogger") if not isinstance(custom, CustomLogger): return view: Final = _logger(logger) @@ -334,6 +326,66 @@ def dispatch_callable_request(logger: Logging, callback: LegacyCall) -> None: ) +def _typed_call_type(call_type: str) -> object: + from litellm.types.utils import CallTypes + + try: + return CallTypes(call_type) + except ValueError: + return None + + +def _awaitable(value: object) -> Awaitable[object]: + return cast(Awaitable[object], value) # cast-ok: legacy async hooks return coroutines + + +def pre_call_deployment_hook(callback: CustomLogger, kwargs: Details, call_type: str) -> Awaitable[object]: + hook: Final = _call_of(_attribute(callback, "async_pre_call_deployment_hook")) + return _awaitable(hook(kwargs, _typed_call_type(call_type))) + + +def post_call_success_deployment_hook( + callback: CustomLogger, kwargs: Details, response: object, call_type: str +) -> Awaitable[object]: + hook: Final = _call_of(_attribute(callback, "async_post_call_success_deployment_hook")) + return _awaitable(hook(kwargs, response, _typed_call_type(call_type))) + + +def failure_deployment_hook_view( + kwargs: Details, exception: BaseException +) -> tuple[Mapping[str, object], BaseException, int | None]: + from litellm import utils + + raw_depth: Final = kwargs.get("fallback_depth") + depth: Final = raw_depth if isinstance(raw_depth, int) else None + safe_request: Final = MappingProxyType({key: value for key, value in kwargs.items() if key != "attempted_targets"}) + snapshot: Final = _call_of(utils._snapshot_exception_for_hook)(exception) # pyright: ignore[reportPrivateUsage] # legacy snapshot helper + return safe_request, cast(BaseException, snapshot), depth # cast-ok: snapshot is a same-class copy of the exception + + +def post_call_failure_deployment_hook( + callback: CustomLogger, + request: Mapping[str, object], + exception: BaseException, + call_type: str, + fallback_depth: int | None, +) -> Awaitable[object]: + from litellm import utils + + hook: Final = _call_of(_attribute(callback, "async_post_call_failure_deployment_hook")) + accepts_depth: Final = _call_of(utils._accepts_fallback_depth_kwarg_for_class)(type(callback)) # pyright: ignore[reportPrivateUsage] # legacy signature probe + typed: Final = _typed_call_type(call_type) + if accepts_depth: + return _awaitable(hook(request, exception, typed, fallback_depth=fallback_depth)) + return _awaitable(hook(request, exception, typed)) + + +def report_deployment_failure_hook_error(callback: object, error: BaseException) -> None: + from litellm._logging import verbose_logger + + verbose_logger.debug("async_post_call_failure_deployment_hook error in %s: %s", type(callback).__name__, error) + + def report_target_failure(logger: Logging, callback: object, family: Family, error: BaseException) -> None: from litellm._logging import verbose_logger @@ -343,10 +395,10 @@ def report_target_failure(logger: Logging, callback: object, family: Family, err callback, "".join(traceback.format_exception(error)), ) - capture: Final = _legacy_module().get("capture_exception") - if capture is not None and family in ("request", "sync_success", "sync_failure"): - cast(LegacyCall, capture)(error) # cast-ok: legacy sentry hook - if family not in ("request", "sync_failure"): + capture: Final = _singleton("capture_exception") + if callable(capture) and family in ("request", "sync_success", "sync_failure"): + capture(error) + if family not in ("request", "deployment", "sync_failure"): _logger(logger)._handle_callback_failure(callback=callback) # pyright: ignore[reportPrivateUsage] # legacy prometheus counter @@ -524,7 +576,7 @@ def async_log_failure_event( def _custom_logger_singleton() -> IntegrationView: from litellm.litellm_core_utils import litellm_logging - existing: Final = _legacy_module().get("customLogger") + existing: Final = _singleton("customLogger") if isinstance(existing, CustomLogger): return _integration(existing) created: Final = CustomLogger() @@ -564,63 +616,71 @@ def dispatch_callable( ) -_SUCCESS_SINGLETONS: Final[Mapping[str, str]] = MappingProxyType( - { - "promptlayer": "promptLayerLogger", - "supabase": "supabaseClient", - "wandb": "weightsBiasesLogger", - "logfire": "logfireLogger", - "lunary": "lunaryLogger", - "helicone": "heliconeLogger", - "greenscale": "greenscaleLogger", - "athina": "athinaLogger", - "traceloop": "traceloopLogger", - "s3": "s3Logger", - "openmeter": "openMeterLogger", - } -) - - def dispatch_named_success( logger: Logging, name: str, response: object, start_time: Timestamp, end_time: Timestamp ) -> Awaitable[None] | None: + from litellm.integrations.athina import AthinaLogger + from litellm.integrations.greenscale import GreenscaleLogger + from litellm.integrations.helicone import HeliconeLogger + from litellm.integrations.logfire_logger import LogfireLevel, LogfireLogger + from litellm.integrations.lunary import LunaryLogger + from litellm.integrations.openmeter import OpenMeterLogger + from litellm.integrations.prompt_layer import PromptLayerLogger + from litellm.integrations.s3 import S3Logger + from litellm.integrations.supabase import Supabase + from litellm.integrations.traceloop import TraceloopLogger + from litellm.integrations.weights_biases import WeightsBiasesLogger + view: Final = _logger(logger) details: Final = view.model_call_details print_verbose: Final = _print_verbose() - integration: Final = _legacy_module().get(_SUCCESS_SINGLETONS.get(name, "")) without_response: Final = { # mutable-ok: legacy integrations receive a private mutable copy key: value for key, value in details.items() if key != "original_response" } match name: - case "promptlayer" | "wandb" | "athina" if integration is not None: - _method(integration, "log_event")( + case "promptlayer" if isinstance(promptlayer := _singleton("promptLayerLogger"), PromptLayerLogger): + promptlayer.log_event( kwargs=details, response_obj=response, start_time=start_time, end_time=end_time, print_verbose=print_verbose, ) - case "logfire" if integration is not None: - from litellm.integrations.logfire_logger import LogfireLevel - - _method(integration, "log_event")( + case "wandb" if isinstance(wandb := _singleton("weightsBiasesLogger"), WeightsBiasesLogger): + wandb.log_event( + kwargs=details, + response_obj=response, + start_time=start_time, + end_time=end_time, + print_verbose=print_verbose, + ) + case "athina" if isinstance(athina := _singleton("athinaLogger"), AthinaLogger): + athina.log_event( + kwargs=details, + response_obj=response, + start_time=start_time, + end_time=end_time, + print_verbose=print_verbose, + ) + case "logfire" if isinstance(logfire := _singleton("logfireLogger"), LogfireLogger): + logfire.log_event( kwargs=without_response, response_obj=response, start_time=start_time, end_time=end_time, print_verbose=print_verbose, - level=LogfireLevel.INFO.value, + level=LogfireLevel.INFO, ) - case "greenscale" if integration is not None: - _method(integration, "log_event")( + case "greenscale" if isinstance(greenscale := _singleton("greenscaleLogger"), GreenscaleLogger): + greenscale.log_event( kwargs=without_response, response_obj=response, start_time=start_time, end_time=end_time, print_verbose=print_verbose, ) - case "supabase" if integration is not None: - _method(integration, "log_event")( + case "supabase" if isinstance(supabase := _singleton("supabaseClient"), Supabase): + supabase.log_event( model=view.model, messages=view.messages, end_user=details.get("user", "default"), @@ -630,8 +690,8 @@ def dispatch_named_success( litellm_call_id=details["litellm_call_id"], print_verbose=print_verbose, ) - case "lunary" if integration is not None: - _method(integration, "log_event")( + case "lunary" if isinstance(lunary := _singleton("lunaryLogger"), LunaryLogger): + lunary.log_event( kwargs=details, type="llm", event="end", @@ -644,8 +704,8 @@ def dispatch_named_success( run_id=view.litellm_call_id, print_verbose=print_verbose, ) - case "helicone" if integration is not None: - _method(integration, "log_success")( + case "helicone" if isinstance(helicone := _singleton("heliconeLogger"), HeliconeLogger): + helicone.log_success( model=view.model, messages=view.messages, response_obj=response, @@ -658,8 +718,8 @@ def dispatch_named_success( _langfuse( logger, response=response, start_time=start_time, end_time=end_time, level=None, status_message=None ) - case "traceloop" if integration is not None: - _method(integration, "log_event")( + case "traceloop" if isinstance(traceloop := _singleton("traceloopLogger"), TraceloopLogger): + traceloop.log_event( kwargs=details, response_obj=response, start_time=start_time, @@ -667,16 +727,16 @@ def dispatch_named_success( user_id=details.get("user", None), print_verbose=print_verbose, ) - case "s3" if integration is not None: - _method(integration, "log_event")( + case "s3" if isinstance(s3 := _singleton("s3Logger"), S3Logger): + s3.log_event( kwargs=details, response_obj=response, start_time=start_time, end_time=end_time, print_verbose=print_verbose, ) - case "openmeter" if integration is not None: - return _async_method(integration, "async_log_success_event")( + case "openmeter" if isinstance(openmeter := _singleton("openMeterLogger"), OpenMeterLogger): + return openmeter.async_log_success_event( kwargs=details, response_obj=response, start_time=start_time, end_time=end_time ) case "dynamodb": @@ -692,10 +752,10 @@ def _dynamodb( from litellm.integrations.dynamodb import DyanmoDBLogger from litellm.litellm_core_utils import litellm_logging - existing: Final = _legacy_module().get("dynamoLogger") + existing: Final = _singleton("dynamoLogger") dynamo: Final = existing if isinstance(existing, DyanmoDBLogger) else DyanmoDBLogger() litellm_logging.dynamoLogger = dynamo # pyright: ignore[reportAttributeAccessIssue] # legacy module global - return _async_method(dynamo, "_async_log_event")( + return dynamo._async_log_event( # pyright: ignore[reportPrivateUsage] # legacy async entry point kwargs=details, response_obj=response, start_time=start_time, end_time=end_time, print_verbose=print_verbose ) @@ -709,39 +769,36 @@ def _langfuse( level: str | None, status_message: str | None, ) -> None: - from litellm.integrations.langfuse import langfuse_handler + from litellm.integrations.langfuse.langfuse import LangFuseLogger + from litellm.integrations.langfuse.langfuse_handler import LangFuseHandler + from litellm.litellm_core_utils import litellm_logging + from litellm.types.utils import ModelResponse view: Final = _logger(logger) - module: Final = _legacy_module() kwargs: Final = { # mutable-ok: langfuse receives a private mutable copy key: value for key, value in view.model_call_details.items() if key != "original_response" } - select: Final = cast( # cast-ok: legacy factory - LegacyCall, langfuse_handler.LangFuseHandler.get_langfuse_logger_for_request - ) - handler: Final = select( - globalLangfuseLogger=module.get("langFuseLogger"), + global_logger: Final = _singleton("langFuseLogger") + handler: Final = LangFuseHandler.get_langfuse_logger_for_request( + globalLangfuseLogger=global_logger if isinstance(global_logger, LangFuseLogger) else None, standard_callback_dynamic_params=view.standard_callback_dynamic_params, - in_memory_dynamic_logger_cache=module["in_memory_dynamic_logger_cache"], + in_memory_dynamic_logger_cache=litellm_logging.in_memory_dynamic_logger_cache, ) - if handler is None: - return - extra: Final[Mapping[str, object]] = ( - MappingProxyType({"level": level, "status_message": status_message}) if level is not None else _EMPTY - ) - result: Final = _method(handler, "log_event_on_langfuse")( + user: Final = kwargs.get("user") + result: Final = handler.log_event_on_langfuse( kwargs=kwargs, - response_obj=response, + response_obj=cast( + ModelResponse, response + ), # cast-ok: OCR responses sit outside the legacy union langfuse names start_time=start_time, end_time=end_time, - user_id=kwargs.get("user", None), - **extra, + user_id=user if isinstance(user, str) else None, + level="DEFAULT" if level is None else level, + status_message=status_message, ) - trace_id: Final = ( - cast(Details, result).get("trace_id") if isinstance(result, dict) else None # cast-ok: legacy response dict - ) # cast-ok: legacy response dict - if trace_id is not None: - _method(module["in_memory_trace_id_cache"], "set_cache")( + trace_id: Final = result.get("trace_id") + if isinstance(trace_id, str): + litellm_logging.in_memory_trace_id_cache.set_cache( litellm_call_id=view.litellm_call_id, service_name="langfuse", trace_id=trace_id ) @@ -754,16 +811,17 @@ def dispatch_named_failure( start_time: Timestamp, end_time: Timestamp, ) -> None: + from litellm.integrations.logfire_logger import LogfireLevel, LogfireLogger + from litellm.integrations.lunary import LunaryLogger + from litellm.integrations.supabase import Supabase + from litellm.integrations.traceloop import TraceloopLogger + view: Final = _logger(logger) - module: Final = _legacy_module() details: Final = view.model_call_details print_verbose: Final = _print_verbose() - without_response: Final = MappingProxyType( - {key: value for key, value in details.items() if key != "original_response"} - ) match name: - case "lunary" if (lunary := module.get("lunaryLogger")) is not None: - _method(lunary, "log_event")( + case "lunary" if isinstance(lunary := _singleton("lunaryLogger"), LunaryLogger): + lunary.log_event( kwargs=details, type="llm", event="error", @@ -776,10 +834,10 @@ def dispatch_named_failure( end_time=end_time, print_verbose=print_verbose, ) - case "sentry" if (capture := module.get("capture_exception")) is not None: - cast(LegacyCall, capture)(exception) # cast-ok: legacy sentry hook - case "supabase" if (supabase := module.get("supabaseClient")) is not None: - _method(supabase, "log_event")( + case "sentry" if callable(capture := _singleton("capture_exception")): + capture(exception) + case "supabase" if isinstance(supabase := _singleton("supabaseClient"), Supabase): + supabase.log_event( model=view.model, messages=view.messages, end_user=details.get("user", "default"), @@ -798,8 +856,8 @@ def dispatch_named_failure( level="ERROR", status_message=str(exception), ) - case "traceloop" if (traceloop := module.get("traceloopLogger")) is not None: - _method(traceloop, "log_event")( + case "traceloop" if isinstance(traceloop := _singleton("traceloopLogger"), TraceloopLogger): + traceloop.log_event( start_time=start_time, end_time=end_time, response_obj=None, @@ -809,18 +867,16 @@ def dispatch_named_failure( level="ERROR", kwargs=details, ) - case "logfire" if (logfire := module.get("logfireLogger")) is not None: - from litellm.integrations.logfire_logger import LogfireLevel - - _method(logfire, "log_event")( + case "logfire" if isinstance(logfire := _singleton("logfireLogger"), LogfireLogger): + logfire.log_event( kwargs={ # mutable-ok: logfire receives a private mutable copy - **without_response, - "exception": exception, - }, + key: value for key, value in details.items() if key != "original_response" + } + | {"exception": exception}, response_obj=None, start_time=start_time, end_time=end_time, - level=LogfireLevel.ERROR.value, + level=LogfireLevel.ERROR, print_verbose=print_verbose, ) case _: diff --git a/tests/test_litellm_rust/ocr/test_callbacks.py b/tests/test_litellm_rust/ocr/test_callbacks.py index 9487550c80e..904504150a7 100644 --- a/tests/test_litellm_rust/ocr/test_callbacks.py +++ b/tests/test_litellm_rust/ocr/test_callbacks.py @@ -506,20 +506,52 @@ def legacy_orchestration_disabled(monkeypatch: pytest.MonkeyPatch) -> list[str]: for name in FORBIDDEN_ORCHESTRATION: monkeypatch.setattr(Logging, name, forbid(name)) - def forbidden_setup(*args, **kwargs): - reached.append("function_setup") - raise AssertionError("legacy orchestration reached: function_setup") + def forbid_utils(name: str): + async def hook(*args, **kwargs): + reached.append(name) + raise AssertionError(f"legacy orchestration reached: {name}") - monkeypatch.setattr(utils, "function_setup", forbidden_setup) + def setup(*args, **kwargs): + reached.append(name) + raise AssertionError(f"legacy orchestration reached: {name}") + + return setup if name == "function_setup" else hook + + for name in ( + "function_setup", + "async_pre_call_deployment_hook", + "async_post_call_success_deployment_hook", + "async_post_call_failure_deployment_hook", + ): + monkeypatch.setattr(utils, name, forbid_utils(name)) return reached +class DeploymentRecorder(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.hooks: list[str] = [] + + async def async_pre_call_deployment_hook(self, kwargs, call_type): + self.hooks.append(f"pre:{call_type.value}") + return {**kwargs, "metadata": {**(kwargs.get("metadata") or {}), "deployment": "seen"}} + + async def async_post_call_success_deployment_hook(self, request_data, response, call_type): + self.hooks.append(f"post:{request_data['metadata']['deployment']}") + return response + + async def async_post_call_failure_deployment_hook(self, request_data, exception, call_type, fallback_depth=None): + self.hooks.append(f"failure:{type(exception).__name__}:{fallback_depth}") + + @pytest.mark.asyncio @pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"]) async def test_native_ocr_success_runs_integrations_without_legacy_orchestration( ocr_server: RecordingServer, legacy_orchestration_disabled: list[str], asynchronous: bool ) -> None: recorder: Final = RecordingLogger() + deployment: Final = DeploymentRecorder() + litellm.callbacks.append(deployment) arguments: Final = {"callbacks": [recorder]} response: Final = ( await call_native_aocr(ocr_server, **arguments) if asynchronous else call_native_ocr(ocr_server, **arguments) @@ -528,6 +560,7 @@ async def test_native_ocr_success_runs_integrations_without_legacy_orchestration success_event: Final = "async_log_success_event" if asynchronous else "log_success_event" events: Final = await recorder.wait_for_async(success_event) assert legacy_orchestration_disabled == [] + assert deployment.hooks == (["pre:aocr", "post:seen"] if asynchronous else []) assert recorder.names.count("log_pre_api_call") == 1 assert events[0].kwargs["standard_logging_object"]["status"] == "success" assert events[0].kwargs["response_cost"] is not None @@ -541,10 +574,13 @@ async def test_native_ocr_failure_runs_integrations_without_legacy_orchestration ) -> None: ocr_server.enqueue(ResponseSpec(body={"message": "provider unavailable"}, status=500)) recorder: Final = RecordingLogger() + deployment: Final = DeploymentRecorder() + litellm.callbacks.append(deployment) arguments: Final = {"callbacks": [recorder]} with pytest.raises(litellm.InternalServerError) as caught: await call_native_aocr(ocr_server, **arguments) if asynchronous else call_native_ocr(ocr_server, **arguments) assert legacy_orchestration_disabled == [] + assert deployment.hooks == (["pre:aocr", "failure:InternalServerError:None"] if asynchronous else []) failures: Final = tuple(event for event in recorder.events if event.name.endswith("log_failure_event")) assert [event.name for event in failures] == ( ["log_failure_event", "async_log_failure_event"] if asynchronous else ["log_failure_event"]