From c2679757b3ad93bd9f8def74021e78be42a025c7 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 15:36:42 -0700 Subject: [PATCH] refactor(rust): put stream billing on the legacy surface --- .../callbacks-legacy/python_contract.json | 3 ++ .../crates/callbacks-legacy/src/adapter.rs | 37 +++++++++++++++---- .../crates/callbacks-legacy/src/lib.rs | 2 +- .../crates/callbacks-legacy/tests/support.rs | 1 + .../crates/host-python/src/adapter.rs | 6 +-- litellm-rust/crates/host-python/src/driver.rs | 6 +-- .../python-bridge/src/routes/messages/mod.rs | 6 ++- .../python-bridge/src/routes/ocr/mod.rs | 1 + litellm/rust_bridge/legacy_callbacks.py | 14 +++++-- 9 files changed, 53 insertions(+), 23 deletions(-) diff --git a/litellm-rust/crates/callbacks-legacy/python_contract.json b/litellm-rust/crates/callbacks-legacy/python_contract.json index a09bdc711a3..8a7f3b98f47 100644 --- a/litellm-rust/crates/callbacks-legacy/python_contract.json +++ b/litellm-rust/crates/callbacks-legacy/python_contract.json @@ -100,6 +100,8 @@ ], "stream_success": [ "logger", + "url_route", + "endpoint_type", "request_body", "chunks", "start", @@ -108,6 +110,7 @@ ], "stream_failure": [ "logger", + "endpoint_type", "request_body", "chunks", "error" diff --git a/litellm-rust/crates/callbacks-legacy/src/adapter.rs b/litellm-rust/crates/callbacks-legacy/src/adapter.rs index a67da2188fb..4e9167100e0 100644 --- a/litellm-rust/crates/callbacks-legacy/src/adapter.rs +++ b/litellm-rust/crates/callbacks-legacy/src/adapter.rs @@ -30,6 +30,16 @@ pub struct LegacySurface { pub call_type: &'static str, /// What `Logging.pre_call` is told the input was. pub input_description: &'static str, + /// How a streamed response is billed; `None` for a route that never streams. + pub stream: Option, +} + +/// The pass-through billing a streamed response goes through once its chunks are in. +#[derive(Clone, Copy, Debug)] +pub struct PassThroughStream { + pub url_route: &'static str, + /// A value of Python's `EndpointType`. + pub endpoint_type: &'static str, } /// What the Messages stream iterator keeps for its end-of-stream billing. @@ -177,10 +187,13 @@ impl LegacyLogging { fn stream_success(&self, py: Python<'_>, stream: &DeliveredStream) -> PyResult<()> { let logger = self.logger()?; + let billing = self.surface.stream.ok_or_else(missing_state)?; let billed = Streaming::Success.call( py, ( logger.object(py), + billing.url_route, + billing.endpoint_type, &self.body, &stream.chunks, &self.start, @@ -201,14 +214,25 @@ impl LegacyLogging { /// partial usage. The sync path has no loop to schedule that on, so it falls back to /// the plain failure handler. fn stream_failure(&mut self, py: Python<'_>) -> PyResult { - let (Some(logger), Some(error), Some(stream)) = (&self.logger, &self.error, &self.stream) + let (Some(logger), Some(error), Some(stream), Some(billing)) = + (&self.logger, &self.error, &self.stream, self.surface.stream) else { return Ok(LifecycleStep::Done); }; if !self.asynchronous { return self.dispatch_failure(py); } - match Streaming::Failure.call(py, (logger.object(py), &self.body, &stream.chunks, error)) { + let scheduled = Streaming::Failure.call( + py, + ( + logger.object(py), + billing.endpoint_type, + &self.body, + &stream.chunks, + error, + ), + ); + match scheduled { Ok(awaitable) => { self.pending = Some(Pending::AsyncFailure); Ok(LifecycleStep::Await(awaitable.unbind())) @@ -343,11 +367,7 @@ impl PythonLifecycle for LegacyLogging { self.finalize(py) } - fn emit( - &mut self, - py: Python<'_>, - event: LifecycleEvent<'_>, - ) -> PyResult { + fn emit(&mut self, py: Python<'_>, event: LifecycleEvent<'_>) -> PyResult { match event { LifecycleEvent::Started { .. } => Ok(LifecycleStep::Done), LifecycleEvent::Machine(MachineEvent::ResponseReceived { raw }) => { @@ -403,6 +423,9 @@ impl PythonLifecycle for LegacyLogging { } fn opened(&mut self, py: Python<'_>) -> PyResult<()> { + if self.surface.stream.is_none() { + return Err(missing_state()); + } Streaming::Opened.call(py, (self.logger()?.object(py),))?; self.stream = Some(DeliveredStream { chunks: PyList::empty(py).unbind(), diff --git a/litellm-rust/crates/callbacks-legacy/src/lib.rs b/litellm-rust/crates/callbacks-legacy/src/lib.rs index 42ffd545e2b..eaa1a8b714e 100644 --- a/litellm-rust/crates/callbacks-legacy/src/lib.rs +++ b/litellm-rust/crates/callbacks-legacy/src/lib.rs @@ -21,7 +21,7 @@ mod preparation; mod test_support; pub(crate) use adapter::LegacyLogging; -pub use adapter::LegacySurface; +pub use adapter::{LegacySurface, PassThroughStream}; pub use call::{PublicCall, run_legacy_call}; pub(crate) use callbacks::{LegacyCallbacks, is_internal_call}; pub(crate) use logger::{DeploymentHooks, PythonLogger, finalize, setup}; diff --git a/litellm-rust/crates/callbacks-legacy/tests/support.rs b/litellm-rust/crates/callbacks-legacy/tests/support.rs index 42ca184eb16..d3cc32e301f 100644 --- a/litellm-rust/crates/callbacks-legacy/tests/support.rs +++ b/litellm-rust/crates/callbacks-legacy/tests/support.rs @@ -197,6 +197,7 @@ pub(crate) fn legacy_call( LegacySurface { call_type: "test", input_description: "test input", + stream: None, }, call, asynchronous, diff --git a/litellm-rust/crates/host-python/src/adapter.rs b/litellm-rust/crates/host-python/src/adapter.rs index f946dbc763a..e70e4a8c58f 100644 --- a/litellm-rust/crates/host-python/src/adapter.rs +++ b/litellm-rust/crates/host-python/src/adapter.rs @@ -69,11 +69,7 @@ pub trait PythonLifecycle: Send + Sync { timing: Timing, ) -> PyResult; - fn emit( - &mut self, - py: Python<'_>, - event: LifecycleEvent<'_>, - ) -> PyResult; + fn emit(&mut self, py: Python<'_>, event: LifecycleEvent<'_>) -> PyResult; /// The call streams and its stream was handed to the caller. The caller is not /// inside an await here, so this step and `delivered` cannot suspend. diff --git a/litellm-rust/crates/host-python/src/driver.rs b/litellm-rust/crates/host-python/src/driver.rs index 25f78d7013e..d6b31b27314 100644 --- a/litellm-rust/crates/host-python/src/driver.rs +++ b/litellm-rust/crates/host-python/src/driver.rs @@ -793,11 +793,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } } - fn emit( - &mut self, - py: Python<'_>, - event: LifecycleEvent<'_>, - ) -> PyResult { + fn emit(&mut self, py: Python<'_>, event: LifecycleEvent<'_>) -> PyResult { self.log.push(match event { LifecycleEvent::Started { .. } => "started".into(), LifecycleEvent::Machine(MachineEvent::ResponseReceived { raw }) => { diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs index b4259aca5ba..8c42315ac59 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -1,7 +1,7 @@ mod host; use host::MessagesRouteHost; -use litellm_callbacks_legacy::{LegacySurface, PublicCall, run_legacy_call}; +use litellm_callbacks_legacy::{LegacySurface, PassThroughStream, PublicCall, run_legacy_call}; use litellm_core::messages::route::{messages_machine, supports}; use pyo3::{ prelude::*, @@ -13,6 +13,10 @@ use crate::errors::RustBridgeDeclined; const SURFACE: LegacySurface = LegacySurface { call_type: "anthropic_messages", input_description: "Messages", + stream: Some(PassThroughStream { + url_route: "/v1/messages", + endpoint_type: "anthropic", + }), }; fn run_messages( diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index b5bb941708d..8afa1e2a906 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -15,6 +15,7 @@ use pyo3::{ const SURFACE: LegacySurface = LegacySurface { call_type: "ocr", input_description: "OCR document processing", + stream: None, }; const ASYNC_SURFACE: LegacySurface = LegacySurface { diff --git a/litellm/rust_bridge/legacy_callbacks.py b/litellm/rust_bridge/legacy_callbacks.py index fd0fac1799a..bac40442ce5 100644 --- a/litellm/rust_bridge/legacy_callbacks.py +++ b/litellm/rust_bridge/legacy_callbacks.py @@ -331,6 +331,8 @@ def stream_opened(logger: Logging) -> None: def stream_success( logger: Logging, + url_route: str, + endpoint_type: str, request_body: dict[str, object], chunks: list[bytes], start: datetime.datetime, @@ -353,9 +355,9 @@ def stream_success( coroutine: Final = build( litellm_logging_obj=logger, passthrough_success_handler_obj=GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ, - url_route="/v1/messages", + url_route=url_route, request_body=request_body, - endpoint_type=EndpointType.ANTHROPIC, + endpoint_type=EndpointType(endpoint_type), start_time=start, raw_bytes=chunks, end_time=end, @@ -374,14 +376,18 @@ def stream_success( def stream_failure( - logger: Logging, request_body: dict[str, object], chunks: list[bytes], error: Exception + logger: Logging, + endpoint_type: str, + request_body: dict[str, object], + chunks: list[bytes], + error: Exception, ) -> Coroutine[object, object, None]: from litellm.proxy.pass_through_endpoints.streaming_handler import PassThroughStreamingHandler from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType return PassThroughStreamingHandler.schedule_stream_failure_logging( litellm_logging_obj=logger, - endpoint_type=EndpointType.ANTHROPIC, + endpoint_type=EndpointType(endpoint_type), request_body=request_body, raw_bytes=chunks, exception=error,