diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index f61267c86fa..a60b49d9f2c 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -3824,9 +3824,11 @@ dependencies = [ "base64 0.22.1", "bytes", "futures-util", + "litellm-accounting", "litellm-auth", "litellm-core", "litellm-gateway-auth", + "litellm-host", "litellm-host-http", "litellm-http", "litellm-llms", @@ -3834,10 +3836,13 @@ dependencies = [ "litellm-secrets", "litellm-types", "rstest", + "rusty-money", "serde_json", "thiserror 2.0.19", "tokio", + "tokio-util", "tower", + "tracing", "wiremock", ] @@ -3914,6 +3919,7 @@ dependencies = [ "litellm-coroutine", "rstest", "serde_json", + "thiserror 2.0.19", "tokio", ] @@ -7235,6 +7241,7 @@ dependencies = [ "bytes", "futures-core", "futures-sink", + "futures-util", "pin-project-lite", "tokio", ] diff --git a/litellm-rust/crates/accounting/README.md b/litellm-rust/crates/accounting/README.md index 3b1f4e9d500..8b1618b093f 100644 --- a/litellm-rust/crates/accounting/README.md +++ b/litellm-rust/crates/accounting/README.md @@ -1,6 +1,6 @@ # Accounting API -This first version implements a per-call settlement state machine with injected storage operations. It does not implement prices, admission policy, Redis operations, database writes, or Python callback dispatch. Production calls are not connected to it yet +This version implements a per-call settlement state machine and an optional injected consumer for native Rust Responses calls. It does not supply default prices, admission policy, Redis operations, database writes, or Python callback dispatch ## Ownership @@ -32,7 +32,7 @@ Coordination stores for budgets and rate limits must be injected independently o `Terminal` records success, failure, or cancellation, reported usage, chargeable provider work, and charges. `ProviderWork::NotStarted` requires an established fact that this call incurred no provider generation work. It makes the provider charge zero while retaining usage and service charges. A response-cache hit can use this classification even when the cached response reports usage. Provider prompt-cache usage still represents started provider work -The inspected main branch has no shared provider/model/source/cache-hit execution-facts contract yet. Production wiring requires a core fact establishing whether this call started provider generation, retained reported usage on cache hits, and terminal facts for partial failure and cancellation. Reuse the cache work’s shared contract when it lands, without importing that branch or duplicating its provider/model/source/cache-key types here. Missing facts must not be classified as avoided provider work +The inspected main branch has no shared provider/model/source/cache-hit execution-facts contract yet. Cache-aware production wiring requires a core fact establishing whether this call started provider generation, retained reported usage on cache hits, and terminal facts for partial failure and cancellation. Reuse the cache work’s shared contract when it lands, without importing that branch or duplicating its provider/model/source/cache-key types here. `ProviderWork::Unknown` preserves supplied charges without asserting that provider generation started or was avoided. Missing facts must use this classification rather than avoided provider work Failure and cancellation never automatically zero charges. Pricing and admission policy must supply incurred costs, including partial-stream costs or an input-cost estimate where the established policy requires one. This API preserves those decisions rather than guessing charges from a terminal outcome @@ -42,7 +42,7 @@ Failure and cancellation never automatically zero charges. Pricing and admission Budget reconciliation and release run only when a monetary budget receipt exists. Spend recording runs even without a monetary reservation. This API has no response-cache dependency or callback registry -`Backend::apply` receives the terminal, admission receipts, and read-only progress of all effects. Its future is awaited inline and need not be `Send`, allowing a Python adapter to preserve the caller's task and context +`Backend::apply` receives the terminal, admission receipts, and read-only progress of all effects. Its future is awaited inline and need not be `Send`, allowing a Python adapter to preserve the caller's task and context. Native adapters can use `next_effect` and `complete_effect` to await their own `Send` backend futures. `next_effect` marks an effect in flight before returning its read-only settlement inputs; `complete_effect` accepts a result only for an in-flight effect. The same state machine owns order, progress, and duplicate protection for both drivers | Backend result | Meaning | Session behavior | | --- | --- | --- | @@ -63,7 +63,7 @@ Release must use the supplied progress to release remaining resources without un The state becomes `InFlight` before invoking a backend. Dropping the settlement future leaves that effect in flight because its external result is unknown. A later `settle` attempts only the remaining pending effects, including release, and does not replay the interrupted operation -The host must retain the session and resume remaining cleanup through its cancellation-safe finalization path. Dropping the entire session performs no asynchronous cleanup. Cancellation during release itself leaves release uncertain and requires backend-specific recovery +The host must retain the session and resume remaining cleanup through its cancellation handling. Dropping the entire session performs no asynchronous cleanup. Cancellation during release itself leaves release uncertain and requires backend-specific recovery This is local protection within one retained session. There is no persistent settlement key, remote deduplication, crash recovery, queue acknowledgement reconciliation, or automatic resolution of unknown prices. A new session can repeat the same remote operations. Persistent idempotency and recovery require a later backend implementation @@ -71,6 +71,32 @@ This is local protection within one retained session. There is no persistent set Public API tests inject backend outcomes at every effect, drop futures at every suspension boundary, and verify completed work is not replayed. They cover optional monetary admission receipts, unknown and zero values, partial usage on failure and cancellation, avoided provider work with service charges, overflow, queue acceptance, and release progress -The next Python adapter must consume shared execution facts and separate existing built-in settlement from CustomLogger delivery. Native gateway accounting later supplies pricing, monetary budget policy, storage operations, and persistent idempotency behind the same contract. Both paths must avoid duplicate cost calculation and spend updates +The next Python adapter must consume shared execution facts and separate existing built-in settlement from CustomLogger delivery. Both paths must avoid duplicate cost calculation and spend updates -Python callback compatibility, real pricing, coordination-store behavior, and native gateway accounting require production adapters and their own integration tests. The API tests do not claim those paths are migrated or validated +## Native Responses integration + +`gateway-inference::accounting::ResponsesAccounting` accepts an injected `AccountingService`. `Gateway::with_responses_accounting` connects it to `/responses` and `/v1/responses`, including streaming. All other endpoints and unconfigured calls retain their existing execution paths. The server binary has no default accounting configuration or backend + +`AccountingService::admit` receives the authenticated caller, public and deployment models, and original request. It returns a plugin-owned call identity, optional monetary reservation receipt, and per-call backend. It runs after authorization and before inference execution. The plugin must make failed admission leave no reservation behind. Once admission starts, the tracked task retains its future even if the caller cancels, then settles any acquired reservation as cancelled + +The backend receives existing prepared-request context, raw provider responses, typed Responses values, and untouched streaming bytes. Resolved provider credentials and declared secret fields are removed from prepared-request context. Raw responses and request content remain trusted plugin inputs, not diagnostic logs. Prepared context establishes preparation, not provider generation or a cache hit. Stream bytes can split SSE records, so a plugin must retain framing state when extracting usage. Reuse shared core execution facts when the cache prerequisite lands rather than parsing cache behavior here + +`AccountingBackend::inputs` supplies reported usage and work classification independently of `price`. `price` computes charges without financial writes. Pricing or collection failure produces unknown charges and a diagnostic while retaining the plugin's available usage and work inputs. Settlement still attempts budget reconciliation, spend recording, and reservation release. Unknown prices never become zero or successful accounting + +`AccountingBackend::apply` implements the financial effects and their persistence contract. It must honor progress when releasing a reservation after a failed or indeterminate adjustment. Coordination dependencies belong to the injected service, independently of response-cache storage. This integration has no response-cache or rate-limit backend selection + +Native `CallInterceptors` extends `ProviderInterceptors` with response transformation, stream chunks, terminal delivery, and synchronous cancellation notification. `host-http::serve_with_hooks` accepts one adapter for these stages. Existing `serve` and `serve_unary` entrypoints remain compatible. Interceptor and observer contracts live under `host/src/hooks/`; root interception and observation import paths remain re-exports. `Interceptors` remains a compatibility alias for `ProviderInterceptors`. `CallHooks` continues to serve Python runtime mechanics unchanged. Observers receive best-effort notifications; interceptors participate in execution and may inspect, change, or reject it. Names describe scope rather than whether the runtime waits + +Accounting facts use a bounded per-call delivery channel with backpressure. The task processes accounting inputs and effects, never produces inference chunks or reads the provider stream ahead of HTTP body demand. The response body retains the hook adapter until exhaustion, failure, or drop + +Terminal outcomes describe host completion, execution or encoding failure, and cancellation. Provider-declared failures inside successful JSON or SSE payloads remain input facts for the plugin; this adapter does not add a provider protocol parser. The shared execution-facts prerequisite must include those provider terminal facts + +Normal completion awaits accounting before returning a response or ending the stream. Required accounting errors turn an otherwise successful unary response into HTTP 500, or append one error SSE frame after streaming headers have been sent. Admission failures also return HTTP 500. An existing inference or encoding error retains its original HTTP error, with accounting failure exposed in the report and diagnostic event + +An optional awaited `AccountingCallback` receives the call identity, selected terminal, settlement status, effect progress, and assessment diagnostic after all financial effects have been attempted. Ordinary callback errors produce a diagnostic and cannot skip settlement or release. These callbacks are distinct from detached, best-effort `ObservationSender` telemetry. Neither callback delivery nor queue acceptance proves a committed financial write + +Client cancellation closes fact delivery when the call or body releases its adapters. The tracked task drains delivered facts, selects cancellation if no terminal was delivered, and attempts settlement. Cancellation after a terminal was delivered preserves that terminal and lets the ongoing financial operation finish without replay. Hosts must stop admission and drain or drop response bodies before awaiting `ResponsesAccounting::shutdown`; shutdown prevents new admission and waits for the accounting tasks. Plugin operations and callbacks must have bounded completion. Process termination, task panics, hung plugins, and persistent recovery remain outside these guarantees + +Native HTTP tests inject pricing and financial operations and verify payload preservation, accounting without custom callbacks, release after financial or callback errors, partial-stream transport failure, cancellation during admission and settlement, separate call sessions, and shutdown. Declared avoided work retains reported usage and service charges, but this is not a test of an actual native response-cache hit + +Python compatibility, actual response-cache hits with authoritative shared facts, production pricing and storage, limiter composition, and persistent idempotency remain follow-up work. The Python implementation is unchanged, and its existing callback compatibility tests remain required diff --git a/litellm-rust/crates/accounting/src/charges.rs b/litellm-rust/crates/accounting/src/charges.rs index 143a5ed37ba..e9340d08349 100644 --- a/litellm-rust/crates/accounting/src/charges.rs +++ b/litellm-rust/crates/accounting/src/charges.rs @@ -29,6 +29,7 @@ impl Cost { #[derive(Clone, Copy, Debug, Eq, PartialEq)] pub enum ProviderWork { + Unknown, NotStarted, Started, } @@ -71,7 +72,7 @@ impl Charges { pub(crate) fn for_work(self, work: ProviderWork) -> Self { match work { - ProviderWork::Started => self, + ProviderWork::Unknown | ProviderWork::Started => self, ProviderWork::NotStarted => Self { provider: Cost::Known(Money::from_major(0, iso::USD)), services: self.services, diff --git a/litellm-rust/crates/accounting/src/error.rs b/litellm-rust/crates/accounting/src/error.rs index a2259a1d259..83ef8ae35f6 100644 --- a/litellm-rust/crates/accounting/src/error.rs +++ b/litellm-rust/crates/accounting/src/error.rs @@ -12,6 +12,8 @@ pub enum Error { AlreadyTerminal, #[error("the call has no terminal outcome")] NotTerminal, + #[error("{effect:?} has no in-flight application")] + NotInFlight { effect: Effect }, #[error("{effect:?} can only be retried after a confirmed not-applied result")] UnsafeRetry { effect: Effect }, } diff --git a/litellm-rust/crates/accounting/src/lib.rs b/litellm-rust/crates/accounting/src/lib.rs index 425911637f7..bb2947c20c2 100644 --- a/litellm-rust/crates/accounting/src/lib.rs +++ b/litellm-rust/crates/accounting/src/lib.rs @@ -5,6 +5,6 @@ mod session; pub use charges::{Charges, Cost, ProviderWork, ReportedUsage, Usd}; pub use error::Error; pub use session::{ - ApplyResult, Backend, BudgetAdmission, Effect, EffectState, Outcome, Progress, Session, - Settlement, SettlementStatus, Terminal, + ApplyResult, Backend, BudgetAdmission, Effect, EffectState, Outcome, PendingEffect, Progress, + Session, Settlement, SettlementStatus, Terminal, }; diff --git a/litellm-rust/crates/accounting/src/session.rs b/litellm-rust/crates/accounting/src/session.rs index 5e586e4130c..ca38dccf20c 100644 --- a/litellm-rust/crates/accounting/src/session.rs +++ b/litellm-rust/crates/accounting/src/session.rs @@ -119,6 +119,7 @@ impl From> for EffectState { } } +#[derive(Clone)] pub struct Progress { effects: [EffectState; 3], } @@ -244,24 +245,43 @@ impl Session { &mut self, backend: &mut B, ) -> Result { - let terminal = self.terminal.as_ref().ok_or(Error::NotTerminal)?; - for effect in Effect::ORDER { - if !matches!(self.progress.effects[effect.index()], EffectState::Pending) { - continue; - } - self.progress.effects[effect.index()] = EffectState::InFlight; - let result = backend - .apply( - effect, - Settlement { - admission: &self.admission, - terminal, - progress: &self.progress, - }, - ) - .await; - self.progress.effects[effect.index()] = result.into(); + while let Some(pending) = self.next_effect()? { + let PendingEffect { effect, settlement } = pending; + let result = backend.apply(effect, settlement).await; + self.complete_effect(effect, result)?; } Ok(self.status()) } + + pub fn next_effect(&mut self) -> Result>, Error> { + let terminal = self.terminal.as_ref().ok_or(Error::NotTerminal)?; + let Some(effect) = Effect::ORDER + .into_iter() + .find(|effect| matches!(self.progress.effects[effect.index()], EffectState::Pending)) + else { + return Ok(None); + }; + self.progress.effects[effect.index()] = EffectState::InFlight; + Ok(Some(PendingEffect { + effect, + settlement: Settlement { + admission: &self.admission, + terminal, + progress: &self.progress, + }, + })) + } + + pub fn complete_effect(&mut self, effect: Effect, result: ApplyResult) -> Result<(), Error> { + if !matches!(self.progress.effect(effect), EffectState::InFlight) { + return Err(Error::NotInFlight { effect }); + } + self.progress.effects[effect.index()] = result.into(); + Ok(()) + } +} + +pub struct PendingEffect<'a, R, U, E> { + pub effect: Effect, + pub settlement: Settlement<'a, R, U, E>, } diff --git a/litellm-rust/crates/accounting/tests/charges.rs b/litellm-rust/crates/accounting/tests/charges.rs index d7a895db6a1..94d4b5932ce 100644 --- a/litellm-rust/crates/accounting/tests/charges.rs +++ b/litellm-rust/crates/accounting/tests/charges.rs @@ -133,3 +133,21 @@ fn terminal_outcomes_do_not_erase_incurred_charges(services: Cost, #[case] outco assert_eq!(terminal.charges(), charges); assert_eq!(terminal.usage(), &ReportedUsage::Known((10, 1))); } + +#[rstest] +#[case::known(costs_for_unknown_work())] +#[case::unknown(Charges::new(Cost::Unknown, Cost::Known(usd("1"))).unwrap())] +fn unknown_provider_work_does_not_imply_free_generation(#[case] charges: Charges) { + let terminal = Terminal::<()>::new( + Outcome::Failed, + ProviderWork::Unknown, + ReportedUsage::Unknown, + charges, + ) + .unwrap(); + assert_eq!(terminal.charges(), charges); +} + +fn costs_for_unknown_work() -> Charges { + Charges::new(Cost::Known(usd("3")), Cost::Known(usd("1"))).unwrap() +} diff --git a/litellm-rust/crates/accounting/tests/session.rs b/litellm-rust/crates/accounting/tests/session.rs index 173b707e9a1..beebdbe7b68 100644 --- a/litellm-rust/crates/accounting/tests/session.rs +++ b/litellm-rust/crates/accounting/tests/session.rs @@ -489,3 +489,39 @@ async fn unknown_charges_cannot_become_successful_settlement_through_queue_accep &EffectState::Committed ); } + +#[rstest] +fn externally_driven_effects_preserve_progress_and_reject_invalid_completions( + mut session: Session, + terminal: Terminal, +) { + assert!(matches!(session.next_effect(), Err(Error::NotTerminal))); + session.finish(terminal).unwrap(); + assert_eq!( + session.complete_effect(Effect::RecordSpend, ApplyResult::Committed), + Err(Error::NotInFlight { + effect: Effect::RecordSpend + }) + ); + for effect in EFFECTS { + let pending = session.next_effect().unwrap().unwrap(); + assert_eq!(pending.effect, effect); + assert_eq!( + pending.settlement.progress.effect(effect), + &EffectState::InFlight + ); + assert_eq!( + pending.settlement.admission.budget().unwrap(), + "budget-receipt" + ); + session + .complete_effect(effect, ApplyResult::Committed) + .unwrap(); + assert_eq!( + session.complete_effect(effect, ApplyResult::Committed), + Err(Error::NotInFlight { effect }) + ); + } + assert!(session.next_effect().unwrap().is_none()); + assert_eq!(session.status(), SettlementStatus::Committed); +} diff --git a/litellm-rust/crates/gateway-inference/Cargo.toml b/litellm-rust/crates/gateway-inference/Cargo.toml index fd7c99204f7..58b1c7e8080 100644 --- a/litellm-rust/crates/gateway-inference/Cargo.toml +++ b/litellm-rust/crates/gateway-inference/Cargo.toml @@ -9,6 +9,8 @@ repository.workspace = true axum = { workspace = true, features = ["json", "multipart", "original-uri"] } base64.workspace = true bytes.workspace = true +litellm-accounting.workspace = true +litellm-host.workspace = true litellm-auth.workspace = true litellm-gateway-auth.workspace = true litellm-core.workspace = true @@ -20,8 +22,12 @@ litellm-secrets.workspace = true litellm-types.workspace = true serde_json.workspace = true thiserror.workspace = true +tokio = { workspace = true, features = ["sync"] } +tokio-util = { version = "0.7", features = ["rt"] } +tracing.workspace = true [dev-dependencies] +rusty-money.workspace = true futures-util.workspace = true tokio = { workspace = true, features = ["io-util"] } rstest.workspace = true diff --git a/litellm-rust/crates/gateway-inference/src/accounting.rs b/litellm-rust/crates/gateway-inference/src/accounting.rs new file mode 100644 index 00000000000..e0c4820188b --- /dev/null +++ b/litellm-rust/crates/gateway-inference/src/accounting.rs @@ -0,0 +1,347 @@ +use crate::{AccountingError, PluginError}; +use bytes::Bytes; +use litellm_accounting::{ + ApplyResult, BudgetAdmission, Charges, Cost, Effect, Outcome, Progress, ProviderWork, + ReportedUsage, Session, Settlement, SettlementStatus, Terminal, +}; +use litellm_core::{RouteError, responses::route::Responses}; +use litellm_gateway_auth::AuthenticatedCaller; +use litellm_host::{ + HookError, + hooks::{CallInterceptors, CallOutcome}, + interceptors::{ProviderInterceptors, RawResponse, RequestContext, WireRequest}, +}; +use litellm_types::responses::main::ResponsesApiResponse; +use serde_json::{Map, Value}; +use std::{ + future::Future, + pin::Pin, + sync::{ + Arc, + atomic::{AtomicBool, Ordering}, + }, +}; +use tokio::sync::{mpsc, oneshot}; +use tokio_util::task::TaskTracker; +use tracing::{Instrument, instrument::WithSubscriber}; + +pub type AccountingFuture<'a, T> = Pin + Send + 'a>>; +pub type AccountingSettlement<'a> = Settlement<'a, String, Value, PluginError>; + +pub struct AdmissionRequest { + pub caller: AuthenticatedCaller, + pub public_model: String, + pub deployment_model: String, + pub request: Map, +} + +pub struct AdmittedCall { + pub call_id: String, + pub budget_receipt: Option, + pub backend: Box, +} + +pub trait AccountingService: Send + Sync { + fn admit( + &self, + request: AdmissionRequest, + ) -> AccountingFuture<'_, Result>; +} + +pub enum Fact { + Prepared(RequestContext), + ProviderResponse(RawResponse), + Response(ResponsesApiResponse), + Chunk(Bytes), +} + +pub struct AccountingInputs { + pub work: ProviderWork, + pub usage: ReportedUsage, +} + +pub trait AccountingBackend: Send { + fn observe(&mut self, fact: Fact) -> Result<(), PluginError>; + fn inputs(&self) -> AccountingInputs; + fn price(&mut self, outcome: Outcome) -> AccountingFuture<'_, Result>; + fn apply<'a>( + &'a mut self, + effect: Effect, + settlement: AccountingSettlement<'a>, + ) -> AccountingFuture<'a, ApplyResult>; +} + +#[derive(Clone)] +pub struct Report { + pub call_id: String, + pub terminal: Terminal, + pub status: SettlementStatus, + pub progress: Progress, + pub assessment_error: Option, +} + +pub trait AccountingCallback: Send + Sync { + fn completed(&self, report: Report) -> AccountingFuture<'_, Result<(), PluginError>>; +} + +#[derive(Clone)] +pub struct ResponsesAccounting { + service: Arc, + callback: Option>, + tasks: TaskTracker, + closed: Arc, +} + +impl ResponsesAccounting { + pub fn new(service: Arc) -> Self { + Self { + service, + callback: None, + tasks: TaskTracker::new(), + closed: Arc::new(AtomicBool::new(false)), + } + } + + pub fn with_callback(mut self, callback: Arc) -> Self { + self.callback = Some(callback); + self + } + + pub async fn shutdown(&self) { + self.closed.store(true, Ordering::SeqCst); + self.tasks.close(); + self.tasks.wait().await; + } + + pub(crate) async fn begin(&self, request: AdmissionRequest) -> Result { + let token = self.tasks.token(); + if self.closed.load(Ordering::SeqCst) { + return Err(AccountingError::Closed); + } + let (sender, receiver) = mpsc::channel(1); + let (ready, admitted) = oneshot::channel(); + let service = self.service.clone(); + let callback = self.callback.clone(); + tokio::spawn(async move { + let _token = token; + let admission = match service.admit(request).await { + Ok(admission) => admission, + Err(error) => { + let _ = ready.send(Err(AccountingError::Admission(error))); + return; + } + }; + let mut session = Session::new(BudgetAdmission::new(admission.budget_receipt)); + let mut backend = admission.backend; + let delivered = ready.send(Ok(())).is_ok(); + let (outcome, reply, assessment_error) = if delivered { + collect(receiver, backend.as_mut()).await + } else { + (Outcome::Cancelled, None, None) + }; + let result = settle(&mut session, backend.as_mut(), admission.call_id, outcome, assessment_error).await; + match result { + Ok(report) => { + if let Some(callback) = callback + && callback.completed(report.clone()).await.is_err() { + tracing::warn!("accounting terminal callback failed"); + } + let delivery = report_result(&report); + if delivery.is_err() { + tracing::error!(status = ?report.status, "call accounting needs attention"); + } + if let Some(reply) = reply { + let _ = reply.send(delivery); + } + } + Err(error) => { + tracing::error!("accounting session contract failed"); + if let Some(reply) = reply { let _ = reply.send(Err(error)); } + } + } + }.with_current_subscriber().in_current_span()); + admitted.await.map_err(|_| AccountingError::Closed)??; + Ok(Call { + sender: Some(sender), + }) + } +} + +type TerminalReply = oneshot::Sender>; +enum Delivery { + Fact(Fact), + Terminal { + outcome: Outcome, + reply: TerminalReply, + }, +} + +async fn collect( + mut receiver: mpsc::Receiver, + backend: &mut dyn AccountingBackend, +) -> (Outcome, Option, Option) { + let mut error = None; + while let Some(delivery) = receiver.recv().await { + match delivery { + Delivery::Fact(fact) => { + if let Err(failure) = backend.observe(fact) { + error.get_or_insert(failure); + } + } + Delivery::Terminal { outcome, reply } => return (outcome, Some(reply), error), + } + } + (Outcome::Cancelled, None, error) +} + +async fn settle( + session: &mut Session, + backend: &mut dyn AccountingBackend, + call_id: String, + outcome: Outcome, + observation_error: Option, +) -> Result { + let inputs = backend.inputs(); + let pricing = match observation_error { + Some(error) => Err(error), + None => backend.price(outcome).await, + }; + let checked = pricing.and_then(|charges| { + Terminal::new(outcome, inputs.work, inputs.usage.clone(), charges) + .map_err(|error| Arc::new(error) as PluginError) + }); + let (terminal, assessment_error) = match checked { + Ok(terminal) => (terminal, None), + Err(error) => ( + Terminal::new( + outcome, + inputs.work, + inputs.usage, + Charges::new(Cost::Unknown, Cost::Unknown).map_err(AccountingError::Contract)?, + ) + .map_err(AccountingError::Contract)?, + Some(error), + ), + }; + session + .finish(terminal.clone()) + .map_err(AccountingError::Contract)?; + while let Some(pending) = session.next_effect().map_err(AccountingError::Contract)? { + let litellm_accounting::PendingEffect { effect, settlement } = pending; + let result = backend.apply(effect, settlement).await; + session + .complete_effect(effect, result) + .map_err(AccountingError::Contract)?; + } + Ok(Report { + call_id, + terminal, + status: session.status(), + progress: session.progress().clone(), + assessment_error, + }) +} + +fn report_result(report: &Report) -> Result<(), AccountingError> { + if let Some(error) = &report.assessment_error { + return Err(AccountingError::Assessment(error.clone())); + } + match report.status { + SettlementStatus::Committed | SettlementStatus::Accepted => Ok(()), + status => Err(AccountingError::Incomplete { status }), + } +} + +#[derive(Clone)] +pub(crate) struct Call { + sender: Option>, +} +impl Call { + async fn fact(&self, fact: Fact) -> Result<(), AccountingError> { + self.sender + .as_ref() + .ok_or(AccountingError::Closed)? + .send(Delivery::Fact(fact)) + .await + .map_err(|_| AccountingError::Closed) + } +} + +impl ProviderInterceptors for Call { + async fn before_provider_request( + &self, + wire: WireRequest, + context: RequestContext, + ) -> Result { + let optional_params = match context.optional_params { + Value::Object(params) => Value::Object( + params + .into_iter() + .filter(|(key, _)| !context.secret_fields.contains(key)) + .collect(), + ), + _ => Value::Null, + }; + self.fact(Fact::Prepared(RequestContext { + api_key: None, + optional_params, + ..context + })) + .await + .map_err(hook_route_error)?; + Ok(wire) + } + async fn after_provider_response(&self, raw: RawResponse) -> Result<(), RouteError> { + self.fact(Fact::ProviderResponse(raw)) + .await + .map_err(hook_route_error) + } +} +fn hook_route_error(error: AccountingError) -> RouteError { + RouteError::InvalidResponse(litellm_llms::ErrorDetail::failed( + "accounting fact delivery", + error, + )) +} + +impl CallInterceptors for Call { + async fn transform_response( + &mut self, + response: ResponsesApiResponse, + ) -> Result { + self.fact(Fact::Response(response.clone())) + .await + .map_err(HookError::new)?; + Ok(response) + } + fn on_stream_chunk( + &mut self, + chunk: &Bytes, + ) -> impl Future> + Send { + let chunk = chunk.clone(); + async move { self.fact(Fact::Chunk(chunk)).await.map_err(HookError::new) } + } + async fn on_terminal(&mut self, outcome: CallOutcome) -> Result<(), HookError> { + let Some(sender) = self.sender.take() else { + return Ok(()); + }; + let outcome = match outcome { + CallOutcome::Succeeded => Outcome::Succeeded, + CallOutcome::Failed => Outcome::Failed, + CallOutcome::Cancelled => Outcome::Cancelled, + }; + let (reply, settled) = oneshot::channel(); + sender + .send(Delivery::Terminal { outcome, reply }) + .await + .map_err(|_| HookError::new(AccountingError::Closed))?; + drop(sender); + settled + .await + .map_err(|_| HookError::new(AccountingError::Closed))? + .map_err(HookError::new) + } + fn on_cancel(&mut self) { + self.sender.take(); + } +} diff --git a/litellm-rust/crates/gateway-inference/src/error.rs b/litellm-rust/crates/gateway-inference/src/error.rs index 6b99069dc9d..c151194785d 100644 --- a/litellm-rust/crates/gateway-inference/src/error.rs +++ b/litellm-rust/crates/gateway-inference/src/error.rs @@ -28,6 +28,28 @@ pub enum Error { BodyTooLarge, #[error("{0}")] Internal(String), + #[error(transparent)] + Hook(#[from] litellm_host::HookError), + #[error(transparent)] + Accounting(#[from] AccountingError), +} + +pub type PluginError = std::sync::Arc; + +#[derive(Debug, thiserror::Error)] +pub enum AccountingError { + #[error("accounting admission failed: {0}")] + Admission(#[source] PluginError), + #[error("accounting assessment failed: {0}")] + Assessment(#[source] PluginError), + #[error("accounting service is closed")] + Closed, + #[error("accounting settlement ended with {status:?}")] + Incomplete { + status: litellm_accounting::SettlementStatus, + }, + #[error(transparent)] + Contract(#[from] litellm_accounting::Error), } impl IntoResponse for Error { @@ -40,6 +62,7 @@ impl From> for Error { fn from(error: litellm_host_http::Error) -> Self { match error { litellm_host_http::Error::Call(error) => Self::Route(error), + litellm_host_http::Error::Hook(error) => Self::Hook(error), litellm_host_http::Error::Protocol => Self::Internal(error.to_string()), } } @@ -71,7 +94,9 @@ impl Error { StatusCode::UNAUTHORIZED } Self::Route(error) if error.is_request() => StatusCode::BAD_REQUEST, - Self::Route(_) | Self::Internal(_) => StatusCode::INTERNAL_SERVER_ERROR, + Self::Route(_) | Self::Internal(_) | Self::Hook(_) | Self::Accounting(_) => { + StatusCode::INTERNAL_SERVER_ERROR + } } } diff --git a/litellm-rust/crates/gateway-inference/src/lib.rs b/litellm-rust/crates/gateway-inference/src/lib.rs index e9ffd14c257..1e63e66cb7d 100644 --- a/litellm-rust/crates/gateway-inference/src/lib.rs +++ b/litellm-rust/crates/gateway-inference/src/lib.rs @@ -3,6 +3,7 @@ //! Authentication, rate limiting and logging are the mounting server's layers; this crate //! maps a public model name to its deployment and runs the core route. +pub mod accounting; mod audio_transcription; mod chat_completions; mod error; @@ -22,11 +23,12 @@ use litellm_http::{ClientVariant, HttpClientConfig, media::UrlPolicy}; use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings}; use litellm_secrets::source::SecretSource; -pub use error::Error; +pub use error::{AccountingError, Error, PluginError}; pub use litellm_router::{Deployment, Router as ModelRouter}; pub use request::{JsonObject, RequestId}; pub struct Gateway { + responses_accounting: Option, pub audio_transcription: AudioTranscriptionRoute, pub chat_completions: ChatCompletionsRoute, pub messages: MessagesRoute, @@ -48,6 +50,7 @@ impl Gateway { let provider = resources.pool.client(&http, ClientVariant::Provider)?; let auth = resources.auth.clone(); Ok(Self { + responses_accounting: None, audio_transcription: AudioTranscriptionRoute::new( provider.clone(), auth.clone(), @@ -74,6 +77,13 @@ impl Gateway { http, }) } + pub fn with_responses_accounting( + mut self, + accounting: accounting::ResponsesAccounting, + ) -> Self { + self.responses_accounting = Some(accounting); + self + } } pub fn router(gateway: Arc) -> Router { diff --git a/litellm-rust/crates/gateway-inference/src/responses.rs b/litellm-rust/crates/gateway-inference/src/responses.rs index 3aca2a0c7f7..2d6c1c7c458 100644 --- a/litellm-rust/crates/gateway-inference/src/responses.rs +++ b/litellm-rust/crates/gateway-inference/src/responses.rs @@ -15,6 +15,23 @@ pub(crate) async fn create( ) -> Result { let deployment = request::resolve_deployment(&gateway, &body)?; request::authorize_model(&identity, deployment, &body).await?; + let accounting = match &gateway.responses_accounting { + Some(service) => Some( + service + .begin(crate::accounting::AdmissionRequest { + caller: identity.caller().clone(), + public_model: body + .get("model") + .and_then(serde_json::Value::as_str) + .unwrap_or_default() + .to_owned(), + deployment_model: deployment.model.clone(), + request: body.clone(), + }) + .await?, + ), + None => None, + }; let call = ResponsesCall { model: deployment.model.clone(), input: body.get("input").cloned().unwrap_or_default(), @@ -36,5 +53,10 @@ pub(crate) async fn create( json!({"type": "error", "code": error.status().as_u16().to_string(), "message": error.to_string(), "param": null}) )) }); - Ok(litellm_host_http::serve(machine, (), (), stream, None).await?) + match accounting { + Some(hooks) => { + Ok(litellm_host_http::serve_with_hooks(machine, (), hooks, stream, None).await?) + } + None => Ok(litellm_host_http::serve(machine, (), (), stream, None).await?), + } } diff --git a/litellm-rust/crates/gateway-inference/tests/accounting.rs b/litellm-rust/crates/gateway-inference/tests/accounting.rs new file mode 100644 index 00000000000..08463e7712b --- /dev/null +++ b/litellm-rust/crates/gateway-inference/tests/accounting.rs @@ -0,0 +1,687 @@ +mod accounting_support; +mod support; + +use accounting_support::{Gate, Ledger, Policy, cost, failure, ledger, runtime}; +use axum::body::to_bytes; +use futures_util::StreamExt; +use litellm_accounting::{ + ApplyResult, Cost, Effect, EffectState, Outcome, ProviderWork, ReportedUsage, SettlementStatus, +}; +use rstest::rstest; +use serde_json::{Value, json}; +use std::{ + sync::{Arc, Mutex}, + time::Duration, +}; +use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method}; + +fn accounted_app( + model: &str, + base: &str, + accounting: Option, +) -> axum::Router { + support::app_with_accounting( + model, + base, + accounting, + litellm_gateway_auth::Permissions::All, + ) +} + +const EFFECTS: [Effect; 3] = [ + Effect::ReconcileBudget, + Effect::RecordSpend, + Effect::ReleaseBudgetReservation, +]; + +async fn upstream(stream: bool) -> (MockServer, Value, String) { + let server = MockServer::start().await; + let completed = json!({"id": "response", "model": "test-model", "output": [], "usage": {"input_tokens": 10, "output_tokens": 2}, "custom": true}); + let events = format!( + "event: response.completed\ndata: {}\n\n", + json!({"type": "response.completed", "response": completed}) + ); + let template = if stream { + ResponseTemplate::new(200) + .insert_header("content-type", "text/event-stream") + .set_body_string(&events) + } else { + ResponseTemplate::new(200).set_body_json(&completed) + }; + Mock::given(method("POST")) + .respond_with(template) + .mount(&server) + .await; + (server, completed, events) +} + +fn request(stream: bool) -> Value { + json!({"model": "public/model", "input": "hello", "stream": stream}) +} + +#[rstest] +#[case::ordinary(false)] +#[case::streaming(true)] +#[tokio::test] +async fn accounting_without_callbacks_preserves_payload_and_records_once( + ledger: Arc>, + #[case] stream: bool, + #[values("/responses", "/v1/responses")] route: &str, +) { + let (server, completed, events) = upstream(stream).await; + let accounting = runtime(&ledger, Policy::default(), None); + let app = accounted_app("openai/test-model", &server.uri(), Some(accounting.clone())); + let response = support::post(app, route, request(stream)).await; + assert_eq!(response.status(), 200); + if stream { + assert_eq!( + to_bytes(response.into_body(), 4096).await.unwrap().as_ref(), + events.as_bytes() + ); + } else { + assert_eq!(support::json(response).await, completed); + } + accounting.shutdown().await; + let state = ledger.lock().unwrap(); + assert_eq!(state.admitted.len(), 1); + assert_eq!(state.admitted[0].0, "public/model"); + assert_eq!(state.admitted[0].1, "openai/test-model"); + assert!(!state.admitted[0].2.is_empty()); + assert_eq!(state.effects, EFFECTS); + assert_eq!(state.released, ["reservation"]); + assert_eq!(state.spend.len(), 1); + assert_eq!( + state.spend[0].usage(), + &ReportedUsage::Known(completed["usage"].clone()) + ); + assert_eq!(state.spend[0].charges().total(), Ok(cost(4))); + assert_eq!(state.prepared_credentials, [false]); + assert!(state.reports.is_empty()); +} + +#[rstest] +#[case::healthy(false)] +#[case::callback_failure(true)] +#[tokio::test] +async fn terminal_callback_receives_accounting_after_reservation_release( + ledger: Arc>, + #[case] fail: bool, +) { + let (server, _, _) = upstream(false).await; + let accounting = runtime(&ledger, Policy::default(), Some(fail)); + let app = accounted_app("openai/test-model", &server.uri(), Some(accounting.clone())); + let response = support::post(app, "/responses", request(false)).await; + assert_eq!(response.status(), 200); + accounting.shutdown().await; + let state = ledger.lock().unwrap(); + assert_eq!(state.effects, EFFECTS); + assert_eq!(state.released, ["reservation"]); + assert_eq!(state.spend.len(), 1); + assert_eq!(state.reports.len(), 1); + assert_eq!(state.reports[0].status, SettlementStatus::Committed); + assert_eq!(state.callback_release_counts, [1]); + assert!(EFFECTS.iter().all(|effect| matches!( + state.reports[0].progress.effect(*effect), + EffectState::Committed + ))); +} + +#[rstest] +#[case::budget(Effect::ReconcileBudget)] +#[case::spend(Effect::RecordSpend)] +#[case::release(Effect::ReleaseBudgetReservation)] +#[tokio::test] +async fn financial_failure_is_visible_and_all_effects_are_attempted( + ledger: Arc>, + #[case] failed: Effect, + #[values(false, true)] uncertain: bool, +) { + let (server, _, _) = upstream(false).await; + let result = if uncertain { + ApplyResult::Indeterminate(failure()) + } else { + ApplyResult::NotApplied(failure()) + }; + let accounting = runtime( + &ledger, + Policy { + result: Some((failed, result)), + ..Default::default() + }, + Some(false), + ); + let response = support::post( + accounted_app("openai/test-model", &server.uri(), Some(accounting.clone())), + "/responses", + request(false), + ) + .await; + assert_eq!(response.status(), 500); + assert_eq!(support::json(response).await["error"]["code"], 500); + accounting.shutdown().await; + let state = ledger.lock().unwrap(); + assert_eq!(state.effects, EFFECTS); + assert_eq!(state.reports[0].status, SettlementStatus::NeedsAttention); + assert_eq!( + state.released.len(), + usize::from(failed != Effect::ReleaseBudgetReservation) + ); +} + +#[rstest] +#[case::fact_failure(Policy { bad_fact: true, ..Default::default() })] +#[case::price_failure(Policy { bad_price: true, ..Default::default() })] +#[case::unknown_price(Policy { unknown_price: true, ..Default::default() })] +#[tokio::test] +async fn unknown_accounting_is_not_zero_or_success_and_still_releases( + ledger: Arc>, + #[case] policy: Policy, +) { + let (server, _, _) = upstream(false).await; + let accounting = runtime(&ledger, policy, Some(false)); + let response = support::post( + accounted_app("openai/test-model", &server.uri(), Some(accounting.clone())), + "/responses", + request(false), + ) + .await; + assert_eq!(response.status(), 500); + accounting.shutdown().await; + let state = ledger.lock().unwrap(); + assert_eq!(state.effects, EFFECTS); + assert_eq!(state.released, ["reservation"]); + assert_eq!(state.reports[0].status, SettlementStatus::Unpriced); + assert_eq!(state.spend[0].charges().total(), Ok(Cost::Unknown)); + assert_eq!(state.spend[0].work(), ProviderWork::Unknown); +} + +#[rstest] +#[tokio::test] +async fn avoided_provider_work_preserves_usage_and_independent_service_cost( + ledger: Arc>, +) { + let (server, completed, _) = upstream(false).await; + let accounting = runtime( + &ledger, + Policy { + avoided: true, + ..Default::default() + }, + None, + ); + let response = support::post( + accounted_app("openai/test-model", &server.uri(), Some(accounting.clone())), + "/responses", + request(false), + ) + .await; + assert_eq!(response.status(), 200); + accounting.shutdown().await; + let state = ledger.lock().unwrap(); + assert_eq!( + state.spend[0].usage(), + &ReportedUsage::Known(completed["usage"].clone()) + ); + assert_eq!(state.spend[0].charges().provider(), cost(0)); + assert_eq!(state.spend[0].charges().services(), cost(1)); + assert_eq!(state.spend[0].charges().total(), Ok(cost(1))); +} + +#[rstest] +#[tokio::test] +async fn admission_rejection_never_reaches_provider(ledger: Arc>) { + let (server, _, _) = upstream(false).await; + let accounting = runtime( + &ledger, + Policy { + reject: true, + ..Default::default() + }, + None, + ); + let response = support::post( + accounted_app("openai/test-model", &server.uri(), Some(accounting.clone())), + "/responses", + request(false), + ) + .await; + assert_eq!(response.status(), 500); + accounting.shutdown().await; + assert!(server.received_requests().await.unwrap().is_empty()); + let state = ledger.lock().unwrap(); + assert!(state.admitted.is_empty()); + assert!(state.effects.is_empty()); +} + +#[rstest] +#[tokio::test] +async fn provider_failure_preserves_http_error_and_accounting_cleanup(ledger: Arc>) { + let server = MockServer::start().await; + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(429).set_body_string("slow down")) + .mount(&server) + .await; + let accounting = runtime( + &ledger, + Policy { + result: Some((Effect::RecordSpend, ApplyResult::NotApplied(failure()))), + ..Default::default() + }, + Some(false), + ); + let response = support::post( + accounted_app("openai/test-model", &server.uri(), Some(accounting.clone())), + "/responses", + request(false), + ) + .await; + assert_eq!(response.status(), 429); + assert!( + support::json(response).await["error"]["message"] + .as_str() + .unwrap() + .contains("slow down") + ); + accounting.shutdown().await; + let state = ledger.lock().unwrap(); + assert_eq!(state.effects, EFFECTS); + assert_eq!(state.released, ["reservation"]); + assert_eq!(state.reports[0].terminal.outcome(), Outcome::Failed); + assert_eq!(state.reports[0].status, SettlementStatus::NeedsAttention); +} + +#[rstest] +#[case::before_chunks(false)] +#[case::after_usage(true)] +#[tokio::test] +async fn dropping_stream_retains_delivered_usage_and_releases( + ledger: Arc>, + #[case] consume: bool, +) { + let (server, completed, events) = upstream(true).await; + let accounting = runtime(&ledger, Policy::default(), Some(false)); + let response = support::post( + accounted_app("openai/test-model", &server.uri(), Some(accounting.clone())), + "/responses", + request(true), + ) + .await; + let mut body = response.into_body().into_data_stream(); + if consume { + assert_eq!( + body.next().await.unwrap().unwrap().as_ref(), + events.as_bytes() + ); + } + drop(body); + tokio::time::timeout(Duration::from_secs(5), accounting.shutdown()) + .await + .unwrap(); + let state = ledger.lock().unwrap(); + assert_eq!(state.effects, EFFECTS); + assert_eq!(state.released, ["reservation"]); + assert_eq!(state.spend[0].outcome(), Outcome::Cancelled); + assert_eq!(state.spend[0].charges().total(), Ok(cost(4))); + let usage = if consume { + ReportedUsage::Known(completed["usage"].clone()) + } else { + ReportedUsage::Unknown + }; + assert_eq!(state.spend[0].usage(), &usage); + assert_eq!(state.reports.len(), 1); +} + +#[rstest] +#[case::admission(false)] +#[case::settlement(true)] +#[tokio::test] +async fn cancelling_a_wait_does_not_cancel_accounting_or_repeat_writes( + ledger: Arc>, + #[case] settlement: bool, +) { + let (server, _, _) = upstream(false).await; + let gate = Arc::new(Gate::new()); + let policy = if settlement { + Policy { + spend_gate: Some(gate.clone()), + ..Default::default() + } + } else { + Policy { + admission_gate: Some(gate.clone()), + ..Default::default() + } + }; + let accounting = runtime(&ledger, policy, Some(false)); + let app = accounted_app("openai/test-model", &server.uri(), Some(accounting.clone())); + let call = tokio::spawn(support::post(app, "/responses", request(false))); + tokio::time::timeout(Duration::from_secs(5), gate.entered.notified()) + .await + .unwrap(); + call.abort(); + assert!(call.await.unwrap_err().is_cancelled()); + gate.resume.add_permits(1); + tokio::time::timeout(Duration::from_secs(5), accounting.shutdown()) + .await + .unwrap(); + let state = ledger.lock().unwrap(); + assert_eq!(state.effects, EFFECTS); + assert_eq!(state.released, ["reservation"]); + assert_eq!(state.spend.len(), 1); + assert_eq!(state.reports.len(), 1); + assert_eq!( + state.spend[0].outcome(), + if settlement { + Outcome::Succeeded + } else { + Outcome::Cancelled + } + ); +} + +#[rstest] +#[tokio::test] +async fn shutdown_stops_new_admissions(ledger: Arc>) { + let (server, _, _) = upstream(false).await; + let accounting = runtime(&ledger, Policy::default(), None); + accounting.shutdown().await; + let response = support::post( + accounted_app("openai/test-model", &server.uri(), Some(accounting)), + "/responses", + request(false), + ) + .await; + assert_eq!(response.status(), 500); + assert!(ledger.lock().unwrap().admitted.is_empty()); + assert!(server.received_requests().await.unwrap().is_empty()); +} + +#[rstest] +#[tokio::test] +async fn authorization_precedes_accounting_admission(ledger: Arc>) { + let (server, _, _) = upstream(false).await; + let accounting = runtime(&ledger, Policy::default(), None); + let app = support::app_with_accounting( + "openai/test-model", + &server.uri(), + Some(accounting.clone()), + litellm_gateway_auth::Permissions::None, + ); + let response = support::post(app, "/responses", request(false)).await; + assert_eq!(response.status(), 403); + accounting.shutdown().await; + assert!(ledger.lock().unwrap().admitted.is_empty()); + assert!(server.received_requests().await.unwrap().is_empty()); +} + +#[rstest] +#[tokio::test] +async fn missing_accounting_service_preserves_existing_response() { + let (server, completed, _) = upstream(false).await; + let response = support::post( + support::app("openai/test-model", &server.uri()), + "/responses", + request(false), + ) + .await; + assert_eq!(response.status(), 200); + assert_eq!(support::json(response).await, completed); +} + +#[rstest] +#[case::queue_accepted(ApplyResult::Accepted, SettlementStatus::Accepted)] +#[case::committed(ApplyResult::Committed, SettlementStatus::Committed)] +#[tokio::test] +async fn terminal_callback_distinguishes_queue_acceptance_from_commit( + ledger: Arc>, + #[case] result: ApplyResult, + #[case] status: SettlementStatus, +) { + let (server, _, _) = upstream(false).await; + let accounting = runtime( + &ledger, + Policy { + result: Some((Effect::RecordSpend, result)), + ..Default::default() + }, + Some(false), + ); + let response = support::post( + accounted_app("openai/test-model", &server.uri(), Some(accounting.clone())), + "/responses", + request(false), + ) + .await; + assert_eq!(response.status(), 200); + accounting.shutdown().await; + let state = ledger.lock().unwrap(); + assert_eq!(state.effects, EFFECTS); + assert_eq!(state.released, ["reservation"]); + assert_eq!(state.reports[0].status, status); +} + +#[rstest] +#[tokio::test] +async fn stream_settlement_failure_is_one_sse_error_after_original_chunks( + ledger: Arc>, +) { + let (server, _, events) = upstream(true).await; + let accounting = runtime( + &ledger, + Policy { + result: Some((Effect::RecordSpend, ApplyResult::NotApplied(failure()))), + ..Default::default() + }, + Some(false), + ); + let response = support::post( + accounted_app("openai/test-model", &server.uri(), Some(accounting.clone())), + "/responses", + request(true), + ) + .await; + assert_eq!(response.status(), 200); + let body = to_bytes(response.into_body(), 8192).await.unwrap(); + let text = std::str::from_utf8(&body).unwrap(); + assert!(text.starts_with(&events)); + assert_eq!(text.matches("event: error").count(), 1); + assert!(text.contains("NeedsAttention")); + accounting.shutdown().await; + let state = ledger.lock().unwrap(); + assert_eq!(state.effects, EFFECTS); + assert_eq!(state.released, ["reservation"]); + assert_eq!(state.reports[0].terminal.outcome(), Outcome::Succeeded); +} + +#[rstest] +#[tokio::test] +async fn partial_stream_transport_failure_preserves_usage_and_charge_policy( + ledger: Arc>, +) { + let gate = Arc::new(Gate::new()); + let partial = json!({"type": "response.in_progress", "response": {"usage": {"input_tokens": 10, "output_tokens": 1}}}); + let event = format!("event: response.in_progress\ndata: {partial}\n\n"); + let frame = event.clone(); + let upstream_gate = gate.clone(); + let app = axum::Router::new().route( + "/responses", + axum::routing::post(move || { + let frame = frame.clone(); + let gate = upstream_gate.clone(); + async move { + let chunks = futures_util::stream::once(async move { + Ok::<_, std::io::Error>(bytes::Bytes::from(frame)) + }) + .chain(futures_util::stream::once(async move { + gate.resume.acquire().await.unwrap().forget(); + Err(std::io::Error::other("upstream disconnected")) + })); + ( + [("content-type", "text/event-stream")], + axum::body::Body::from_stream(chunks), + ) + } + }), + ); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let url = format!("http://{}", listener.local_addr().unwrap()); + let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let accounting = runtime(&ledger, Policy::default(), Some(false)); + let response = support::post( + accounted_app("openai/test-model", &url, Some(accounting.clone())), + "/responses", + request(true), + ) + .await; + assert_eq!(response.status(), 200); + let mut body = response.into_body().into_data_stream(); + assert_eq!(body.next().await.unwrap().unwrap(), event); + gate.resume.add_permits(1); + let error = tokio::time::timeout(Duration::from_secs(5), body.next()) + .await + .unwrap() + .unwrap() + .unwrap(); + assert!( + std::str::from_utf8(&error) + .unwrap() + .starts_with("event: error\n") + ); + assert!(body.next().await.is_none()); + drop(body); + accounting.shutdown().await; + server.abort(); + let state = ledger.lock().unwrap(); + assert_eq!(state.effects, EFFECTS); + assert_eq!(state.released, ["reservation"]); + assert_eq!(state.spend.len(), 1); + assert_eq!(state.spend[0].outcome(), Outcome::Failed); + assert_eq!( + state.spend[0].usage(), + &ReportedUsage::Known(partial["response"]["usage"].clone()) + ); + assert_eq!(state.spend[0].charges().total(), Ok(cost(4))); +} + +#[rstest] +#[tokio::test] +async fn pricing_failure_preserves_reported_usage(ledger: Arc>) { + let (server, completed, _) = upstream(false).await; + let accounting = runtime( + &ledger, + Policy { + bad_price: true, + ..Default::default() + }, + Some(false), + ); + let response = support::post( + accounted_app("openai/test-model", &server.uri(), Some(accounting.clone())), + "/responses", + request(false), + ) + .await; + assert_eq!(response.status(), 500); + accounting.shutdown().await; + let state = ledger.lock().unwrap(); + assert_eq!( + state.spend[0].usage(), + &ReportedUsage::Known(completed["usage"].clone()) + ); + assert_eq!(state.reports[0].terminal.usage(), state.spend[0].usage()); + assert!(state.reports[0].assessment_error.is_some()); + assert_eq!(state.spend[0].charges().total(), Ok(Cost::Unknown)); + assert_eq!(state.released, ["reservation"]); +} + +#[rstest] +#[tokio::test] +async fn separate_calls_have_separate_sessions_and_callback_identity(ledger: Arc>) { + let (server, _, _) = upstream(false).await; + let accounting = runtime(&ledger, Policy::default(), Some(false)); + let app = accounted_app("openai/test-model", &server.uri(), Some(accounting.clone())); + let (first, second) = tokio::join!( + support::post(app.clone(), "/responses", request(false)), + support::post(app, "/responses", request(false)) + ); + assert_eq!(first.status(), 200); + assert_eq!(second.status(), 200); + accounting.shutdown().await; + let state = ledger.lock().unwrap(); + assert_eq!(state.admitted.len(), 2); + assert_eq!(state.spend.len(), 2); + assert_eq!(state.released.len(), 2); + assert_eq!(state.reports.len(), 2); + assert_ne!(state.reports[0].call_id, state.reports[1].call_id); + for effect in EFFECTS { + assert_eq!( + state + .effects + .iter() + .filter(|actual| **actual == effect) + .count(), + 2 + ); + } +} + +#[rstest] +#[tokio::test] +async fn unreserved_native_calls_still_record_spend(ledger: Arc>) { + let (server, completed, _) = upstream(false).await; + let accounting = runtime( + &ledger, + Policy { + no_reservation: true, + ..Default::default() + }, + Some(false), + ); + let response = support::post( + accounted_app("openai/test-model", &server.uri(), Some(accounting.clone())), + "/responses", + request(false), + ) + .await; + assert_eq!(response.status(), 200); + accounting.shutdown().await; + let state = ledger.lock().unwrap(); + assert_eq!(state.effects, [Effect::RecordSpend]); + assert!(state.released.is_empty()); + assert_eq!( + state.spend[0].usage(), + &ReportedUsage::Known(completed["usage"].clone()) + ); + assert_eq!(state.spend[0].charges().total(), Ok(cost(4))); + assert_eq!(state.reports[0].status, SettlementStatus::Committed); + assert!(matches!( + state.reports[0] + .progress + .effect(Effect::ReleaseBudgetReservation), + EffectState::NotRequired + )); +} + +#[rstest] +#[tokio::test] +async fn core_validation_failure_releases_admission_without_provider_execution( + ledger: Arc>, +) { + let (server, _, _) = upstream(false).await; + let accounting = runtime(&ledger, Policy::default(), Some(false)); + let response = support::post( + accounted_app("openai/test-model", &server.uri(), Some(accounting.clone())), + "/responses", + json!({"model": "public/model"}), + ) + .await; + assert_eq!(response.status(), 400); + accounting.shutdown().await; + assert!(server.received_requests().await.unwrap().is_empty()); + let state = ledger.lock().unwrap(); + assert_eq!(state.admitted.len(), 1); + assert_eq!(state.effects, EFFECTS); + assert_eq!(state.released, ["reservation"]); + assert_eq!(state.reports[0].terminal.outcome(), Outcome::Failed); + assert_eq!(state.reports[0].terminal.usage(), &ReportedUsage::Unknown); +} diff --git a/litellm-rust/crates/gateway-inference/tests/accounting_support/mod.rs b/litellm-rust/crates/gateway-inference/tests/accounting_support/mod.rs new file mode 100644 index 00000000000..63d3dcb02bf --- /dev/null +++ b/litellm-rust/crates/gateway-inference/tests/accounting_support/mod.rs @@ -0,0 +1,239 @@ +use std::sync::{Arc, Mutex}; + +use litellm_accounting::{ + ApplyResult, Charges, Cost, Effect, Outcome, ProviderWork, ReportedUsage, Usd, +}; +use litellm_gateway_inference::{ + PluginError, + accounting::{ + AccountingBackend, AccountingCallback, AccountingFuture, AccountingInputs, + AccountingService, AccountingSettlement, AdmissionRequest, AdmittedCall, Fact, Report, + ResponsesAccounting, + }, +}; +use rstest::fixture; +use rusty_money::iso; +use serde_json::Value; +use tokio::sync::{Notify, Semaphore}; + +#[derive(Default)] +pub struct Ledger { + pub admitted: Vec<(String, String, String)>, + pub effects: Vec, + pub spend: Vec>, + pub released: Vec, + pub reports: Vec, + pub callback_release_counts: Vec, + pub prepared_credentials: Vec, +} + +#[fixture] +pub fn ledger() -> Arc> { + Arc::new(Mutex::new(Ledger::default())) +} + +pub struct Gate { + pub entered: Notify, + pub resume: Semaphore, +} +impl Gate { + pub fn new() -> Self { + Self { + entered: Notify::new(), + resume: Semaphore::new(0), + } + } + async fn wait(&self) { + self.entered.notify_one(); + self.resume.acquire().await.unwrap().forget(); + } +} + +#[derive(Default)] +pub struct Policy { + pub reject: bool, + pub no_reservation: bool, + pub bad_fact: bool, + pub bad_price: bool, + pub unknown_price: bool, + pub avoided: bool, + pub result: Option<(Effect, ApplyResult)>, + pub admission_gate: Option>, + pub spend_gate: Option>, +} + +pub fn failure() -> PluginError { + Arc::new(std::io::Error::other("injected failure")) +} +pub fn cost(amount: i64) -> Cost { + Cost::Known(Usd::from_major(amount, iso::USD)) +} + +pub struct Service { + pub ledger: Arc>, + pub policy: Arc, +} +impl AccountingService for Service { + fn admit( + &self, + request: AdmissionRequest, + ) -> AccountingFuture<'_, Result> { + Box::pin(async move { + if self.policy.reject { + return Err(failure()); + } + let call_id = { + let mut ledger = self.ledger.lock().unwrap(); + ledger.admitted.push(( + request.public_model, + request.deployment_model, + request.caller.principal().subject().to_owned(), + )); + format!("call-{}", ledger.admitted.len()) + }; + if let Some(gate) = &self.policy.admission_gate { + gate.wait().await; + } + Ok(AdmittedCall { + call_id, + budget_receipt: (!self.policy.no_reservation).then(|| "reservation".into()), + backend: Box::new(Store { + ledger: self.ledger.clone(), + policy: self.policy.clone(), + usage: ReportedUsage::Unknown, + }), + }) + }) + } +} + +struct Store { + ledger: Arc>, + policy: Arc, + usage: ReportedUsage, +} +impl AccountingBackend for Store { + fn observe(&mut self, fact: Fact) -> Result<(), PluginError> { + if self.policy.bad_fact { + return Err(failure()); + } + let usage = match fact { + Fact::Prepared(context) => { + self.ledger + .lock() + .unwrap() + .prepared_credentials + .push(context.api_key.is_some()); + None + } + Fact::ProviderResponse(raw) => serde_json::from_str::(&raw.body) + .ok() + .and_then(|value| value.get("usage").cloned()), + Fact::Response(response) => response.extra.get("usage").cloned(), + Fact::Chunk(chunk) => std::str::from_utf8(&chunk) + .ok() + .and_then(|text| text.lines().find_map(|line| line.strip_prefix("data: "))) + .and_then(|data| serde_json::from_str::(data).ok()) + .and_then(|event| { + event + .get("response") + .and_then(|response| response.get("usage")) + .cloned() + }), + }; + if let Some(usage) = usage { + self.usage = ReportedUsage::Known(usage); + } + Ok(()) + } + fn inputs(&self) -> AccountingInputs { + AccountingInputs { + work: if self.policy.avoided { + ProviderWork::NotStarted + } else { + ProviderWork::Unknown + }, + usage: self.usage.clone(), + } + } + fn price(&mut self, _: Outcome) -> AccountingFuture<'_, Result> { + Box::pin(async move { + if self.policy.bad_price { + return Err(failure()); + } + Ok(Charges::new( + if self.policy.unknown_price { + Cost::Unknown + } else { + cost(3) + }, + cost(1), + ) + .unwrap()) + }) + } + fn apply<'a>( + &'a mut self, + effect: Effect, + settlement: AccountingSettlement<'a>, + ) -> AccountingFuture<'a, ApplyResult> { + Box::pin(async move { + self.ledger.lock().unwrap().effects.push(effect); + if effect == Effect::RecordSpend + && let Some(gate) = &self.policy.spend_gate + { + gate.wait().await; + } + if let Some((target, result)) = &self.policy.result + && effect == *target + { + return result.clone(); + } + let mut ledger = self.ledger.lock().unwrap(); + match effect { + Effect::RecordSpend => ledger.spend.push(settlement.terminal.clone()), + Effect::ReleaseBudgetReservation => ledger + .released + .push(settlement.admission.budget().unwrap().clone()), + Effect::ReconcileBudget => {} + } + ApplyResult::Committed + }) + } +} + +pub struct Callback { + pub ledger: Arc>, + pub fail: bool, +} +impl AccountingCallback for Callback { + fn completed(&self, report: Report) -> AccountingFuture<'_, Result<(), PluginError>> { + Box::pin(async move { + { + let mut ledger = self.ledger.lock().unwrap(); + let released = ledger.released.len(); + ledger.callback_release_counts.push(released); + ledger.reports.push(report); + } + if self.fail { Err(failure()) } else { Ok(()) } + }) + } +} + +pub fn runtime( + ledger: &Arc>, + policy: Policy, + callback: Option, +) -> ResponsesAccounting { + let accounting = ResponsesAccounting::new(Arc::new(Service { + ledger: ledger.clone(), + policy: Arc::new(policy), + })); + match callback { + Some(fail) => accounting.with_callback(Arc::new(Callback { + ledger: ledger.clone(), + fail, + })), + None => accounting, + } +} diff --git a/litellm-rust/crates/gateway-inference/tests/support/mod.rs b/litellm-rust/crates/gateway-inference/tests/support/mod.rs index b59e335334e..4695f109f0a 100644 --- a/litellm-rust/crates/gateway-inference/tests/support/mod.rs +++ b/litellm-rust/crates/gateway-inference/tests/support/mod.rs @@ -33,32 +33,43 @@ pub fn app_with_permissions( model: &str, api_base: &str, permissions: litellm_gateway_auth::Permissions, +) -> Router { + app_with_accounting(model, api_base, None, permissions) +} + +pub fn app_with_accounting( + model: &str, + api_base: &str, + accounting: Option, + permissions: litellm_gateway_auth::Permissions, ) -> Router { let pool = Arc::new(HttpClientPool::new(Arc::new(PublicDnsResolver))); let http = Resolution::from(&HttpSettings::default()).config; let secrets = Arc::new(NoSecrets); let resources = CoreResources::new(pool); - router(Arc::new( - Gateway::new( - resources, - http, - secrets, - [( - "public/model".into(), - Deployment { - model: model.into(), - api_base: Some(api_base.into()), - api_key: Some("test-key".into()), - timeout: Some(Duration::from_secs(5)), - ..Default::default() - }, - )] - .into_iter() - .collect(), - ) - .unwrap(), - )) - .layer(axum::middleware::from_fn_with_state( + let gateway = Gateway::new( + resources, + http, + secrets, + [( + "public/model".into(), + Deployment { + model: model.into(), + api_base: Some(api_base.into()), + api_key: Some("test-key".into()), + timeout: Some(Duration::from_secs(5)), + ..Default::default() + }, + )] + .into_iter() + .collect(), + ) + .unwrap(); + let gateway = match accounting { + Some(accounting) => gateway.with_responses_accounting(accounting), + None => gateway, + }; + router(Arc::new(gateway)).layer(axum::middleware::from_fn_with_state( permissions, test_identity, )) diff --git a/litellm-rust/crates/host-http/src/driver.rs b/litellm-rust/crates/host-http/src/driver.rs index bd2e919d1db..6c58a483c13 100644 --- a/litellm-rust/crates/host-http/src/driver.rs +++ b/litellm-rust/crates/host-http/src/driver.rs @@ -7,6 +7,7 @@ use futures_util::{StreamExt, stream}; use litellm_host::{ call::{CallOutput, HostedCompletion, HostedMachine}, + hooks::{CallInterceptors, CallOutcome}, interceptors::Interceptors, lifecycle::{observe_call, observe_unary}, machine::MachineFault, @@ -59,9 +60,49 @@ where H: Interceptors + 'static, S: HostCallHandler

+ 'static, { - let encoder = Arc::new(encoder); let driver = Driver::new(machine, services, interceptors); - match observe_call(observers, start(driver, encoder.clone())).await? { + serve_driver(driver, encoder, observers, ()).await +} + +pub async fn serve_with_hooks( + machine: HostedMachine

, + services: S, + hooks: H, + encoder: A, + observers: Option, +) -> Result> +where + P: Protocol, + P::Error: From, + A: StreamEncoder, + H: CallInterceptors

+ Clone + 'static, + S: HostCallHandler

+ 'static, +{ + let driver = Driver::new(machine, services, hooks.clone()); + serve_driver(driver, encoder, observers, hooks).await +} + +async fn serve_driver( + driver: HostedDriver, + encoder: A, + observers: Option, + hooks: L, +) -> Result> +where + P: Protocol, + P::Error: From, + A: StreamEncoder, + H: Interceptors + 'static, + S: HostCallHandler

+ 'static, + L: CallInterceptors

+ 'static, +{ + let encoder = Arc::new(encoder); + let hooks = HookGuard:: { + hooks, + finished: false, + protocol: std::marker::PhantomData, + }; + match observe_call(observers, start(driver, encoder.clone(), hooks)).await? { CallOutput::Complete(response) => Ok(response), CallOutput::Stream { head, chunks } => { let body = chunks.map(move |chunk| { @@ -74,9 +115,36 @@ where } } -async fn start( +struct HookGuard> { + hooks: H, + finished: bool, + protocol: std::marker::PhantomData P>, +} + +impl> HookGuard { + async fn terminal(&mut self, outcome: CallOutcome) -> Result<(), litellm_host::HookError> { + let result = self.hooks.on_terminal(outcome).await; + self.finished = true; + result + } + + async fn failed(&mut self) { + let _ = self.terminal(CallOutcome::Failed).await; + } +} + +impl> Drop for HookGuard { + fn drop(&mut self) { + if !self.finished { + self.hooks.on_cancel(); + } + } +} + +async fn start( mut driver: HostedDriver, encoder: Arc, + mut hooks: HookGuard, ) -> Result>, Error> where P: Protocol, @@ -84,28 +152,84 @@ where A: StreamEncoder, H: Interceptors + 'static, S: HostCallHandler

+ 'static, + L: CallInterceptors

+ 'static, { - match driver.advance().await.map_err(Error::Call)? { - Boundary::Complete(HostedCompletion::Complete(value)) => encoder - .encode_response(value) - .map(CallOutput::Complete) - .map_err(Error::Call), + let boundary = match driver.advance().await { + Ok(boundary) => boundary, + Err(error) => { + hooks.failed().await; + return Err(Error::Call(error)); + } + }; + match boundary { + Boundary::Complete(HostedCompletion::Complete(value)) => { + let value = match hooks.hooks.transform_response(value).await { + Ok(value) => value, + Err(error) => { + hooks.failed().await; + return Err(Error::Hook(error)); + } + }; + let response = match encoder.encode_response(value) { + Ok(response) => response, + Err(error) => { + hooks.failed().await; + return Err(Error::Call(error)); + } + }; + hooks.terminal(CallOutcome::Succeeded).await?; + Ok(CallOutput::Complete(response)) + } Boundary::Open(head) => { - let head = encoder.encode_stream_head(head).map_err(Error::Call)?; - let chunks = - stream::try_unfold((driver, encoder), |(mut driver, encoder)| async move { - match driver.advance().await.map_err(Error::Call)? { - Boundary::Chunk(chunk) => { - let bytes = encoder.encode_chunk(chunk).map_err(Error::Call)?; - Ok(Some((bytes, (driver, encoder)))) + let head = match encoder.encode_stream_head(head) { + Ok(head) => head, + Err(error) => { + hooks.failed().await; + return Err(Error::Call(error)); + } + }; + let chunks = stream::try_unfold( + (driver, encoder, hooks), + |(mut driver, encoder, mut hooks)| async move { + let boundary = match driver.advance().await { + Ok(boundary) => boundary, + Err(error) => { + hooks.failed().await; + return Err(Error::Call(error)); + } + }; + match boundary { + Boundary::Chunk(chunk) => { + if let Err(error) = hooks.hooks.on_stream_chunk(&chunk).await { + hooks.failed().await; + return Err(Error::Hook(error)); + } + let bytes = match encoder.encode_chunk(chunk) { + Ok(bytes) => bytes, + Err(error) => { + hooks.failed().await; + return Err(Error::Call(error)); + } + }; + Ok(Some((bytes, (driver, encoder, hooks)))) + } + Boundary::Complete(HostedCompletion::StreamEnded) => { + hooks.terminal(CallOutcome::Succeeded).await?; + Ok(None) + } + _ => { + hooks.failed().await; + Err(Error::Protocol) } - Boundary::Complete(HostedCompletion::StreamEnded) => Ok(None), - _ => Err(Error::Protocol), } - }) - .boxed(); + }, + ) + .boxed(); Ok(CallOutput::Stream { head, chunks }) } - _ => Err(Error::Protocol), + _ => { + hooks.failed().await; + Err(Error::Protocol) + } } } diff --git a/litellm-rust/crates/host-http/src/error.rs b/litellm-rust/crates/host-http/src/error.rs index 285d08d2fe8..b944436c023 100644 --- a/litellm-rust/crates/host-http/src/error.rs +++ b/litellm-rust/crates/host-http/src/error.rs @@ -4,4 +4,6 @@ pub enum Error { Call(E), #[error("unexpected HTTP host operation")] Protocol, + #[error(transparent)] + Hook(#[from] litellm_host::HookError), } diff --git a/litellm-rust/crates/host-http/src/lib.rs b/litellm-rust/crates/host-http/src/lib.rs index a4714ef72b2..918eddfa8d7 100644 --- a/litellm-rust/crates/host-http/src/lib.rs +++ b/litellm-rust/crates/host-http/src/lib.rs @@ -3,7 +3,7 @@ mod encoding; mod error; mod sse; -pub use driver::{serve, serve_unary}; +pub use driver::{serve, serve_unary, serve_with_hooks}; pub use encoding::{ResponseEncoder, StreamEncoder, Unary}; pub use error::Error; pub use sse::Sse; diff --git a/litellm-rust/crates/host-http/tests/serve.rs b/litellm-rust/crates/host-http/tests/serve.rs index 2ba39ad4cd2..bd5b3c3b226 100644 --- a/litellm-rust/crates/host-http/tests/serve.rs +++ b/litellm-rust/crates/host-http/tests/serve.rs @@ -850,3 +850,357 @@ async fn unary_custom_operations_and_conversion_finish_before_terminal_observati )); } } + +#[derive(Clone, Debug, PartialEq)] +enum ActiveEvent { + Prepared(String), + Response(Bytes), + Chunk(Bytes), + Terminal(litellm_host::hooks::CallOutcome), + Cancelled, +} + +#[derive(Clone, Copy, Eq, PartialEq)] +enum ActiveRejection { + None, + Response, + Chunk, + Terminal, +} + +#[derive(Clone)] +struct Awaited { + events: Arc>>, + gate: Option>, + entered: Arc, + rejection: ActiveRejection, +} + +#[fixture] +fn awaited() -> Awaited { + Awaited { + events: Arc::new(Mutex::new(Vec::new())), + gate: None, + entered: Arc::new(tokio::sync::Notify::new()), + rejection: ActiveRejection::None, + } +} + +impl Interceptors for Awaited { + async fn before_provider_request( + &self, + mut wire: WireRequest, + _: RequestContext, + ) -> Result { + wire.url.push_str("/rewritten"); + self.events + .lock() + .unwrap() + .push(ActiveEvent::Prepared(wire.url.clone())); + Ok(wire) + } + async fn after_provider_response(&self, _: RawResponse) -> Result<(), TestError> { + Ok(()) + } +} + +impl litellm_host::hooks::CallInterceptors for Awaited { + async fn transform_response( + &mut self, + response: Bytes, + ) -> Result { + self.events + .lock() + .unwrap() + .push(ActiveEvent::Response(response.clone())); + if self.rejection == ActiveRejection::Response { + return Err(litellm_host::HookError::new(std::io::Error::other( + "response rejection", + ))); + } + Ok(response) + } + fn on_stream_chunk( + &mut self, + chunk: &Bytes, + ) -> impl std::future::Future> + Send { + self.events + .lock() + .unwrap() + .push(ActiveEvent::Chunk(chunk.clone())); + std::future::ready(if self.rejection == ActiveRejection::Chunk { + Err(litellm_host::HookError::new(std::io::Error::other( + "chunk rejection", + ))) + } else { + Ok(()) + }) + } + async fn on_terminal( + &mut self, + outcome: litellm_host::hooks::CallOutcome, + ) -> Result<(), litellm_host::HookError> { + self.events + .lock() + .unwrap() + .push(ActiveEvent::Terminal(outcome)); + self.entered.notify_one(); + if let Some(gate) = &self.gate { + gate.acquire().await.unwrap().forget(); + } + if self.rejection == ActiveRejection::Terminal { + return Err(litellm_host::HookError::new(std::io::Error::other( + "terminal rejection", + ))); + } + Ok(()) + } + fn on_cancel(&mut self) { + self.events.lock().unwrap().push(ActiveEvent::Cancelled); + } +} + +#[rstest] +#[case::success(false)] +#[case::terminal_rejection(true)] +#[tokio::test] +async fn required_terminal_hook_is_awaited_before_returning_a_response( + mut awaited: Awaited, + #[case] reject: bool, +) { + awaited.rejection = if reject { + ActiveRejection::Terminal + } else { + ActiveRejection::None + }; + let gate = Arc::new(tokio::sync::Semaphore::new(0)); + awaited.gate = Some(gate.clone()); + let inspected = awaited.clone(); + let machine = hosted_call::("request", None, |_, _, hooks, _| async move { + let wire = hooks + .before_provider_request( + WireRequest { + url: "request".into(), + headers: Vec::new(), + body: json!({}), + }, + RequestContext { + model: "m".into(), + custom_llm_provider: "p".into(), + optional_params: json!({}), + secret_fields: Vec::new(), + api_key: None, + }, + ) + .await?; + Ok(CallOutput::Complete(Bytes::from(wire.url))) + }); + let call = tokio::spawn(litellm_host_http::serve_with_hooks( + machine, + Adapter(Rejection::None), + awaited, + Adapter(Rejection::None), + None, + )); + tokio::time::timeout( + std::time::Duration::from_secs(5), + inspected.entered.notified(), + ) + .await + .unwrap(); + assert!(!call.is_finished()); + assert_eq!( + inspected.events.lock().unwrap().as_slice(), + [ + ActiveEvent::Prepared("request/rewritten".into()), + ActiveEvent::Response(Bytes::from_static(b"request/rewritten")), + ActiveEvent::Terminal(litellm_host::hooks::CallOutcome::Succeeded) + ] + ); + gate.add_permits(1); + let result = call.await.unwrap(); + if reject { + assert!(matches!(result, Err(Error::Hook(_)))); + } else { + assert_eq!( + to_bytes(result.unwrap().into_body(), 1024).await.unwrap(), + "request/rewritten" + ); + } + assert_eq!(inspected.events.lock().unwrap().len(), 3); +} + +#[rstest] +#[case::exhausted(None, false)] +#[case::partial_failure(None, true)] +#[case::cancel_before_chunks(Some(0), false)] +#[case::cancel_after_chunk(Some(1), false)] +#[tokio::test] +async fn awaited_stream_hooks_follow_demand_and_select_one_terminal( + awaited: Awaited, + #[case] drop_after: Option, + #[case] failed: bool, +) { + let events = awaited.events.clone(); + let machine = + hosted_call::("request", None, move |_, _, _, _| async move { + let last = if failed { + Err(TestError::Provider) + } else { + Ok(Bytes::from_static(b"last")) + }; + Ok(CallOutput::Stream { + head: "text/event-stream", + chunks: stream::iter([Ok(Bytes::from_static(b"first")), last]).boxed(), + }) + }); + let response = litellm_host_http::serve_with_hooks( + machine, + Adapter(Rejection::None), + awaited, + Adapter(Rejection::None), + None, + ) + .await + .unwrap(); + assert!(events.lock().unwrap().is_empty()); + let mut body = response.into_body().into_data_stream(); + if let Some(count) = drop_after { + for _ in 0..count { + assert_eq!(body.next().await.unwrap().unwrap(), "first"); + } + } else { + assert_eq!(body.next().await.unwrap().unwrap(), "first"); + let last = body.next().await.unwrap().unwrap(); + if failed { + assert_eq!(last, "event: error\ndata: Call(Provider)\n\n"); + } else { + assert_eq!(last, "last"); + } + assert!(body.next().await.is_none()); + } + drop(body); + let recorded = events.lock().unwrap(); + let expected = if drop_after.is_some() { + ActiveEvent::Cancelled + } else { + ActiveEvent::Terminal(if failed { + litellm_host::hooks::CallOutcome::Failed + } else { + litellm_host::hooks::CallOutcome::Succeeded + }) + }; + assert_eq!(recorded.last(), Some(&expected)); + assert_eq!( + recorded + .iter() + .filter(|event| matches!(event, ActiveEvent::Terminal(_) | ActiveEvent::Cancelled)) + .count(), + 1 + ); + assert_eq!( + recorded + .iter() + .filter(|event| matches!(event, ActiveEvent::Chunk(_))) + .count(), + drop_after.unwrap_or(if failed { 1 } else { 2 }) + ); +} + +#[rstest] +#[case::provider(Rejection::None)] +#[case::encoding(Rejection::Complete)] +#[tokio::test] +async fn original_failure_survives_a_terminal_hook_failure( + mut awaited: Awaited, + #[case] rejection: Rejection, +) { + awaited.rejection = ActiveRejection::Terminal; + let events = awaited.events.clone(); + let machine = + hosted_call::("request", None, move |_, _, _, _| async move { + if rejection == Rejection::None { + return Err(TestError::Provider); + } + Ok(CallOutput::Complete(Bytes::new())) + }); + let error = litellm_host_http::serve_with_hooks( + machine, + Adapter(rejection), + awaited, + Adapter(rejection), + None, + ) + .await + .unwrap_err(); + assert_eq!( + error, + Error::Call(if rejection == Rejection::None { + TestError::Provider + } else { + TestError::Adapter + }) + ); + let recorded = events.lock().unwrap(); + assert_eq!( + recorded.last(), + Some(&ActiveEvent::Terminal( + litellm_host::hooks::CallOutcome::Failed + )) + ); + assert!(!recorded.contains(&ActiveEvent::Cancelled)); +} + +#[rstest] +#[case::response(false)] +#[case::stream_chunk(true)] +#[tokio::test] +async fn hook_failure_prevents_delivery_and_still_dispatches_failed_terminal( + mut awaited: Awaited, + #[case] streaming: bool, +) { + awaited.rejection = if streaming { + ActiveRejection::Chunk + } else { + ActiveRejection::Response + }; + let events = awaited.events.clone(); + let machine = + hosted_call::("request", None, move |_, _, _, _| async move { + if streaming { + return Ok(CallOutput::Stream { + head: "text/event-stream", + chunks: stream::iter([Ok(Bytes::from_static(b"must not be delivered"))]) + .boxed(), + }); + } + Ok(CallOutput::Complete(Bytes::from_static( + b"must not be delivered", + ))) + }); + let result = litellm_host_http::serve_with_hooks( + machine, + Adapter(Rejection::None), + awaited, + Adapter(Rejection::None), + None, + ) + .await; + if streaming { + let body = to_bytes(result.unwrap().into_body(), 1024).await.unwrap(); + let text = std::str::from_utf8(&body).unwrap(); + assert!(text.starts_with("event: error\n")); + assert!(!text.contains("must not be delivered")); + assert_eq!(text.matches("event: error").count(), 1); + } else { + assert!(matches!(result, Err(Error::Hook(_)))); + } + let recorded = events.lock().unwrap(); + assert_eq!(recorded.len(), 2); + assert_eq!( + recorded.last(), + Some(&ActiveEvent::Terminal( + litellm_host::hooks::CallOutcome::Failed + )) + ); +} diff --git a/litellm-rust/crates/host/AGENTS.md b/litellm-rust/crates/host/AGENTS.md index e2034e76c6f..7d9d68840b9 100644 --- a/litellm-rust/crates/host/AGENTS.md +++ b/litellm-rust/crates/host/AGENTS.md @@ -14,9 +14,9 @@ Core route constructors prepare their dependencies and return a closure acceptin Each driver owns terminal dispatch. Hooks can change or fail execution; passive observers return no result. HTTP observes success after response conversion or stream exhaustion, failure on errors, and cancellation on body drop. Python preserves exception identity and maps native failures once. Explicit Python stream close reports success for delivered chunks; cancellation stops further callback dispatch -`interceptors.rs` owns `Interceptors` and its request/response payload types. `lifecycle.rs` owns `CallObserver`, `CallEvent`, `ExecutionEvent`, timing, failure origin, and the observation wrappers. Event payloads are generic so a runtime can retain its own response, exception and raw-response references without introducing a language dependency. `snapshot()` projects them into the owned observation contract without retaining runtime objects. Pass interceptors and observers separately at direct route and HTTP entrypoints. Routes publish execution events independently of interception +`hooks/interceptors.rs` owns `Interceptors` and its request/response payload types. `lifecycle.rs` owns `CallObserver`, `CallEvent`, `ExecutionEvent`, timing, failure origin, and the observation wrappers. Event payloads are generic so a runtime can retain its own response, exception and raw-response references without introducing a language dependency. `snapshot()` projects them into the owned observation contract without retaining runtime objects. Pass interceptors and observers separately at direct route and HTTP entrypoints. Routes publish execution events independently of interception -`hooks.rs` owns the call-stage interface and its runtime-associated types. It contains no Python types or legacy callback policy. A runtime supplies its context and continuation representation through `HookRuntime` +`hooks/runtime.rs` owns the runtime call-stage interface and its associated types. `hooks/call.rs` defines native call interceptors, extending provider interception with response, chunk, terminal, and synchronous cancellation delivery. `hooks/observation.rs` owns the optional nonblocking observation channel. `hooks/mod.rs` exports these contracts; the existing root interception and observation module paths remain re-exports. It contains no Python types or legacy callback policy. A runtime supplies its context and continuation representation through `HookRuntime` `protocol.rs` owns `Protocol` and suspension messages, including `InterceptRequest` and `StreamDelivery`. `call.rs` owns route outputs and their adaptation into a hosted machine. Rust service handling belongs in `host-native::services`; coroutine channel handles stay in `machine/context.rs` @@ -25,3 +25,9 @@ Rust handlers answer suspensions through `litellm-host-native::Driver`, which `l Keep API policy in gateway-inference and python-bridge, and legacy callback policy in callbacks-legacy-python. Python bindings and hooks expose retained references through `PythonOwned`, with idempotent close and GC traversal. Runtime machinery stays in driver, native, handle and runtime modules Interceptors run inline and can rewrite values or fail execution. Observers consume owned `CallEvent` snapshots from `observation_channel`; its bounded `ObservationSender` never waits for delivery and counts events dropped when the queue is full or closed. The host owns receiver processing and draining. Pass the same publisher to machine construction and the driver when one receiver should collect execution and lifecycle events. Legacy Python callbacks retain their existing awaited, fallible behavior through the Python adapter + +Native `CallInterceptors

` extends `ProviderInterceptors` with response transformation, stream chunks, terminal delivery, and cancellation notification. HTTP drives these interceptors inline through `serve_with_hooks`; passive `ObservationSender` delivery remains separate. Response conversion completes before success terminal delivery, and stream success follows exhaustion rather than headers. Preserve the original execution or encoding error if terminal hooks also fail. Interceptor implementations retain and report their own required operation failures + +`on_cancel` is synchronous because a dropped future cannot await cleanup. Runtime adapters retain asynchronous cleanup work independently of the cancelled caller. Do not execute financial policy or spawn inference producers in the host. Native accounting owns its session and effect progress in the gateway adapter, with financial operations supplied by its injected backend + +Observers receive best-effort notifications; interceptors participate in execution and may inspect, change, or reject it. `ProviderInterceptors` covers core-requested provider stages, and `CallInterceptors` covers the native driver's broader lifecycle. `Interceptors` remains an alias for the provider contract. Python's `CallHooks` represents the same active participation using runtime-associated continuation types; preserve its public contract rather than forcing Python through native `Send` futures diff --git a/litellm-rust/crates/host/Cargo.toml b/litellm-rust/crates/host/Cargo.toml index 51015e99906..e02fc239159 100644 --- a/litellm-rust/crates/host/Cargo.toml +++ b/litellm-rust/crates/host/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +thiserror.workspace = true futures-util.workspace = true litellm-auth.workspace = true litellm-coroutine.workspace = true diff --git a/litellm-rust/crates/host/src/error.rs b/litellm-rust/crates/host/src/error.rs new file mode 100644 index 00000000000..866145c4d1d --- /dev/null +++ b/litellm-rust/crates/host/src/error.rs @@ -0,0 +1,17 @@ +use std::sync::Arc; + +#[derive(Clone, Debug, thiserror::Error)] +#[error("call interceptor failed: {0}")] +pub struct HookError(#[source] Arc); + +impl HookError { + pub fn new(error: impl std::error::Error + Send + Sync + 'static) -> Self { + Self(Arc::new(error)) + } +} + +impl PartialEq for HookError { + fn eq(&self, other: &Self) -> bool { + Arc::ptr_eq(&self.0, &other.0) + } +} diff --git a/litellm-rust/crates/host/src/hooks/call.rs b/litellm-rust/crates/host/src/hooks/call.rs new file mode 100644 index 00000000000..ee66743070b --- /dev/null +++ b/litellm-rust/crates/host/src/hooks/call.rs @@ -0,0 +1,35 @@ +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum CallOutcome { + Succeeded, + Failed, + Cancelled, +} + +pub trait CallInterceptors: + super::interceptors::ProviderInterceptors +{ + fn transform_response( + &mut self, + response: P::Response, + ) -> impl std::future::Future> + Send { + std::future::ready(Ok(response)) + } + + fn on_stream_chunk( + &mut self, + _chunk: &P::Chunk, + ) -> impl std::future::Future> + Send { + std::future::ready(Ok(())) + } + + fn on_terminal( + &mut self, + _outcome: CallOutcome, + ) -> impl std::future::Future> + Send { + std::future::ready(Ok(())) + } + + fn on_cancel(&mut self) {} +} + +impl CallInterceptors

for () {} diff --git a/litellm-rust/crates/host/src/interceptors.rs b/litellm-rust/crates/host/src/hooks/interceptors.rs similarity index 89% rename from litellm-rust/crates/host/src/interceptors.rs rename to litellm-rust/crates/host/src/hooks/interceptors.rs index 0044e842e4a..a362c0386df 100644 --- a/litellm-rust/crates/host/src/interceptors.rs +++ b/litellm-rust/crates/host/src/hooks/interceptors.rs @@ -30,7 +30,7 @@ pub struct RawResponse { pub body: String, } -pub trait Interceptors: Send + Sync { +pub trait ProviderInterceptors: Send + Sync { fn before_provider_request( &self, wire: WireRequest, @@ -43,7 +43,7 @@ pub trait Interceptors: Send + Sync { ) -> impl Future> + Send; } -impl + ?Sized> Interceptors for &T { +impl + ?Sized> ProviderInterceptors for &T { fn before_provider_request( &self, wire: WireRequest, @@ -60,7 +60,7 @@ impl + ?Sized> Interceptors for &T { } } -impl Interceptors for () { +impl ProviderInterceptors for () { async fn before_provider_request( &self, wire: WireRequest, @@ -74,6 +74,8 @@ impl Interceptors for () { } } +pub use ProviderInterceptors as Interceptors; + #[cfg(test)] mod tests { use std::convert::Infallible; @@ -130,13 +132,13 @@ mod tests { async fn the_channel_yields_each_hook_as_its_op_and_returns_the_answer() { let mut machine = CallMachine::::new(None, |channel| { Box::pin(async move { - let sent = Interceptors::before_provider_request( + let sent = ProviderInterceptors::before_provider_request( &channel.interceptors, wire("prepared"), context(), ) .await?; - Interceptors::after_provider_response( + ProviderInterceptors::after_provider_response( &channel.interceptors, RawResponse { body: "raw".into() }, ) @@ -175,9 +177,13 @@ mod tests { #[rstest::rstest] #[tokio::test] async fn no_hooks_pass_the_wire_request_through() { - let sent = Interceptors::::before_provider_request(&(), wire("prepared"), context()) - .await - .unwrap(); + let sent = ProviderInterceptors::::before_provider_request( + &(), + wire("prepared"), + context(), + ) + .await + .unwrap(); assert_eq!(sent.url, "prepared"); } } diff --git a/litellm-rust/crates/host/src/hooks/mod.rs b/litellm-rust/crates/host/src/hooks/mod.rs new file mode 100644 index 00000000000..74ab8a82232 --- /dev/null +++ b/litellm-rust/crates/host/src/hooks/mod.rs @@ -0,0 +1,8 @@ +mod call; +pub mod interceptors; +pub mod observation; +mod runtime; + +pub use call::{CallInterceptors, CallOutcome}; +pub use interceptors::ProviderInterceptors; +pub use runtime::{CallHooks, HookRuntime, RuntimeCallEvent}; diff --git a/litellm-rust/crates/host/src/observation.rs b/litellm-rust/crates/host/src/hooks/observation.rs similarity index 100% rename from litellm-rust/crates/host/src/observation.rs rename to litellm-rust/crates/host/src/hooks/observation.rs diff --git a/litellm-rust/crates/host/src/hooks.rs b/litellm-rust/crates/host/src/hooks/runtime.rs similarity index 100% rename from litellm-rust/crates/host/src/hooks.rs rename to litellm-rust/crates/host/src/hooks/runtime.rs diff --git a/litellm-rust/crates/host/src/lib.rs b/litellm-rust/crates/host/src/lib.rs index 76fce582fb3..0da041259e0 100644 --- a/litellm-rust/crates/host/src/lib.rs +++ b/litellm-rust/crates/host/src/lib.rs @@ -7,9 +7,11 @@ //! may rewrite the wire request before it is sent. pub mod call; +mod error; pub mod hooks; -pub mod interceptors; +pub use error::HookError; +pub use hooks::interceptors; pub mod lifecycle; pub mod machine; -pub mod observation; +pub use hooks::observation; pub mod protocol;