refactor(rust): put stream billing on the legacy surface

This commit is contained in:
Yujong Lee 2026-09-18 15:36:42 -07:00
parent 3a5b7c12ef
commit c2679757b3
9 changed files with 53 additions and 23 deletions

View file

@ -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"

View file

@ -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<PassThroughStream>,
}
/// 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<LifecycleStep> {
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<LifecycleStep> {
fn emit(&mut self, py: Python<'_>, event: LifecycleEvent<'_>) -> PyResult<LifecycleStep> {
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(),

View file

@ -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};

View file

@ -197,6 +197,7 @@ pub(crate) fn legacy_call(
LegacySurface {
call_type: "test",
input_description: "test input",
stream: None,
},
call,
asynchronous,

View file

@ -69,11 +69,7 @@ pub trait PythonLifecycle: Send + Sync {
timing: Timing,
) -> PyResult<LifecycleStep>;
fn emit(
&mut self,
py: Python<'_>,
event: LifecycleEvent<'_>,
) -> PyResult<LifecycleStep>;
fn emit(&mut self, py: Python<'_>, event: LifecycleEvent<'_>) -> PyResult<LifecycleStep>;
/// 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.

View file

@ -793,11 +793,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
}
}
fn emit(
&mut self,
py: Python<'_>,
event: LifecycleEvent<'_>,
) -> PyResult<LifecycleStep> {
fn emit(&mut self, py: Python<'_>, event: LifecycleEvent<'_>) -> PyResult<LifecycleStep> {
self.log.push(match event {
LifecycleEvent::Started { .. } => "started".into(),
LifecycleEvent::Machine(MachineEvent::ResponseReceived { raw }) => {

View file

@ -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(

View file

@ -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 {

View file

@ -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,