mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
refactor(rust): put stream billing on the legacy surface
This commit is contained in:
parent
3a5b7c12ef
commit
c2679757b3
9 changed files with 53 additions and 23 deletions
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
|
|
@ -197,6 +197,7 @@ pub(crate) fn legacy_call(
|
|||
LegacySurface {
|
||||
call_type: "test",
|
||||
input_description: "test input",
|
||||
stream: None,
|
||||
},
|
||||
call,
|
||||
asynchronous,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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 }) => {
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue