diff --git a/enterprise/pyproject.toml b/enterprise/pyproject.toml index 8509600ad96..e5e54a3df2c 100644 --- a/enterprise/pyproject.toml +++ b/enterprise/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "litellm-enterprise" -version = "0.1.70" +version = "0.1.71" description = "Package for LiteLLM Enterprise features" readme = "README.md" requires-python = ">=3.9" @@ -26,7 +26,7 @@ required-version = ">=0.10.9" module-root = "" [tool.commitizen] -version = "0.1.70" +version = "0.1.71" version_files = [ "pyproject.toml:^version", "../pyproject.toml:litellm-enterprise==", diff --git a/litellm-proxy-extras/pyproject.toml b/litellm-proxy-extras/pyproject.toml index e9a4ff90b9e..2835715ef30 100644 --- a/litellm-proxy-extras/pyproject.toml +++ b/litellm-proxy-extras/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "litellm-proxy-extras" -version = "0.4.101" +version = "0.4.102" description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package." readme = "README.md" requires-python = ">=3.9" @@ -26,7 +26,7 @@ required-version = ">=0.10.9" module-root = "" [tool.commitizen] -version = "0.4.101" +version = "0.4.102" version_files = [ "pyproject.toml:^version", "../pyproject.toml:litellm-proxy-extras==", diff --git a/litellm-rust/AGENTS.md b/litellm-rust/AGENTS.md index e5ffcd1c57a..70fcc367905 100644 --- a/litellm-rust/AGENTS.md +++ b/litellm-rust/AGENTS.md @@ -8,3 +8,11 @@ - Split a mixed test file along that line instead of widening visibility to move it - A test for another crate's item belongs in that crate, not in a downstream one - Never set `autotests = false` or hand-list `[[test]]` targets; every file directly under `tests/` is discovered by cargo, and a shared helper goes in `tests//mod.rs` or `tests//support.rs` so it is not picked up as a test crate of its own + +## Error definitions + +- A crate's errors live in `src/error.rs`, defined with `thiserror`, and re-exported from `lib.rs` +- Default to one top-level `Error` enum per crate, with one variant per failure mode and a `#[error(...)]` message on each +- Wrap a lower-level error as a variant with `#[from]` or `#[source]` instead of flattening it to a string +- Exception: split into separate types when different functions fail in disjoint ways, especially when different callers see them. A shared enum would force every caller to match variants its function can never return +- Name a split type after what went wrong (a unit struct is fine for a single failure mode), not after the function that returns it diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 0a91f0759c2..c522bf205b4 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -3013,6 +3013,15 @@ dependencies = [ "url", ] +[[package]] +name = "litellm-coroutine" +version = "0.1.0" +dependencies = [ + "rstest", + "thiserror 2.0.19", + "tokio", +] + [[package]] name = "litellm-cost" version = "0.1.0" @@ -3040,6 +3049,7 @@ name = "litellm-host" version = "0.1.0" dependencies = [ "litellm-auth", + "litellm-coroutine", "rstest", "serde_json", "tokio", @@ -3049,6 +3059,7 @@ dependencies = [ name = "litellm-host-python" version = "0.1.0" dependencies = [ + "bytes", "futures-util", "litellm-host", "pyo3", diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 9d05c8d2b98..0c7236e807e 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -12,6 +12,7 @@ repository = "https://github.com/BerriAI/litellm" litellm-tracing = { path = "crates/tracing" } tracing = "0.1" litellm-core = { path = "crates/core" } +litellm-coroutine = { path = "crates/coroutine" } litellm-host = { path = "crates/host" } litellm-callbacks-legacy-python = { path = "crates/callbacks-legacy-python" } litellm-framing = { path = "crates/framer" } diff --git a/litellm-rust/crates/callbacks-legacy-python/src/call.rs b/litellm-rust/crates/callbacks-legacy-python/src/call.rs index b37790f60a8..9b921070839 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/call.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/call.rs @@ -3,8 +3,8 @@ //! lifetime. No other callback host has that obligation, which is why nothing outside //! this crate holds them. -use litellm_host::{machine::Machine, route::Route}; -use litellm_host_python::{RouteHost, lookup, run_call}; +use litellm_host::{machine::Machine, protocol::Protocol}; +use litellm_host_python::{ProtocolHost, lookup, run_call}; use pyo3::{ gc::{PyTraverseError, PyVisit}, prelude::*, @@ -63,25 +63,25 @@ impl PublicCall { } } -/// Runs one native call under the legacy `Logging` contract: the route host projects from +/// Runs one native call under the legacy `Logging` contract: the protocol host projects from /// the keyword view the contract prepares, and the contract observes the call. pub fn run_legacy_call( py: Python<'_>, surface: LegacySurface, call: PublicCall, machine: M, - route: H, + host: H, asynchronous: bool, ) -> PyResult> where - H: RouteHost + 'static, - M: Machine::Response> + 'static, + H: ProtocolHost + 'static, + M: Machine::Response> + 'static, { let arguments = call.kwargs.clone_ref(py); run_call( py, machine, - route, + host, Box::new(LegacyLogging::new(py, surface, call, asynchronous)), arguments, asynchronous, diff --git a/litellm-rust/crates/core/src/messages/route.rs b/litellm-rust/crates/core/src/messages/route.rs index 8cd3eaf3aa3..fc1a9b63252 100644 --- a/litellm-rust/crates/core/src/messages/route.rs +++ b/litellm-rust/crates/core/src/messages/route.rs @@ -1,4 +1,5 @@ use std::{ + convert::Infallible, sync::{Arc, Mutex}, time::Duration, }; @@ -9,8 +10,8 @@ use litellm_core_utils::get_llm_provider_logic::get_custom_llm_provider; use litellm_host::{ event::{MachineEvent, RawResponse, RequestContext, WireRequest}, host::{Demand, Host}, - machine::{HostChannel, MachineFault, RouteMachine}, - route::Route, + machine::{CallMachine, HostChannel, MachineFault}, + protocol::Protocol, }; use litellm_secrets::source::SecretSource; use litellm_types::{ @@ -28,15 +29,6 @@ use super::{ }; use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum MessagesOp { - ProjectRequest, -} - -pub enum MessagesOpResult { - Request(Box), -} - /// The caller's request as the host projects it. pub struct MessagesCall { pub model: String, @@ -64,11 +56,11 @@ pub enum MessagesOutput { pub struct Messages; -impl Route for Messages { +impl Protocol for Messages { type Response = MessagesOutput; type Error = Error; - type Op = MessagesOp; - type OpResult = MessagesOpResult; + type Projection = MessagesCall; + type Op = Infallible; type Chunk = Bytes; type StreamHead = (); } @@ -78,13 +70,12 @@ impl From for Error { Self::InvalidRequest(match fault { MachineFault::Abandoned => "messages host driver was abandoned".into(), MachineFault::Protocol(message) => format!("messages {message}"), - MachineFault::Mismatch => "invalid messages host operation result".into(), }) } } pub type MessagesHost = HostChannel; -pub type MessagesMachine = RouteMachine; +pub type MessagesMachine = CallMachine; /// Whether this route serves the request, decided before any callback runs so a host /// can still run its own path. @@ -114,30 +105,28 @@ impl LocalMessagesHost { } impl Host for LocalMessagesHost { - async fn route(&self, op: MessagesOp) -> Result { - match op { - MessagesOp::ProjectRequest => self - .call - .lock() - .unwrap_or_else(|error| error.into_inner()) - .take() - .map(|call| MessagesOpResult::Request(Box::new(call))) - .ok_or_else(|| { - Error::InvalidRequest("messages request was already projected".into()) - }), - } + async fn project(&self) -> Result { + self.call + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + .ok_or_else(|| Error::InvalidRequest("messages request was already projected".into())) + } + + async fn custom_op(&self, op: Infallible) -> Result<(), Error> { + match op {} } } pub fn messages_machine(secrets: Arc) -> MessagesMachine { - RouteMachine::new(move |host| Box::pin(execute(host, secrets.clone()))) + CallMachine::new(move |host| Box::pin(execute(host, secrets.clone()))) } async fn execute( host: MessagesHost, secrets: Arc, ) -> Result { - let MessagesOpResult::Request(call) = host.route(MessagesOp::ProjectRequest).await?; + let call = host.project().await?; let stream = call.streams(); let resolved = resolve_provider(&call.model, call.custom_llm_provider.as_deref())?; let secrets = secrets.resolve(resolved.config.secret_names()).await?; diff --git a/litellm-rust/crates/core/src/ocr/document.rs b/litellm-rust/crates/core/src/ocr/document.rs index ffa4f045e8e..b78c09298de 100644 --- a/litellm-rust/crates/core/src/ocr/document.rs +++ b/litellm-rust/crates/core/src/ocr/document.rs @@ -23,9 +23,6 @@ pub fn prepare_document(input: OcrDocumentInput) -> Result { file_name.as_deref(), mime_type.as_deref(), )?), - OcrDocumentInput::HostReader { .. } => Err(Error::InvalidRequest( - "OCR file reader was not read by the host".into(), - )), } } @@ -207,7 +204,7 @@ mod tests { } #[test] - fn byte_documents_are_encoded_and_host_readers_must_be_read_first() { + fn byte_documents_are_encoded() { assert_eq!( prepare_document(OcrDocumentInput::Bytes { bytes: b"abc".as_slice().into(), @@ -217,7 +214,6 @@ mod tests { .unwrap(), document("data:application/pdf;base64,YWJj") ); - assert!(prepare_document(OcrDocumentInput::HostReader { mime_type: None }).is_err()); } #[test] diff --git a/litellm-rust/crates/core/src/ocr/route.rs b/litellm-rust/crates/core/src/ocr/route.rs index 7f83291bdab..e3e57bd2d77 100644 --- a/litellm-rust/crates/core/src/ocr/route.rs +++ b/litellm-rust/crates/core/src/ocr/route.rs @@ -3,102 +3,73 @@ use std::sync::{Arc, Mutex}; use litellm_auth::ResolvedCredential; use litellm_host::{ event::{CallEvent, RequestContext, WireRequest}, - machine::{HostChannel, HostTokenProvider, MachineFault, RouteMachine, TokenRoute}, - route::Route, + host::Reply, + machine::{CallMachine, HostChannel, HostTokenProvider, TokenProtocol}, + protocol::Protocol, }; use litellm_llms::base_llm::ocr::{ error::Error, handler::OcrClient, transformation::LiteLLMOcrResponse, }; use super::handler::perform_ocr_request; -use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput, OcrFileContent, ResolvedOcrRequest}; +use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput, ResolvedOcrRequest}; -#[derive(Clone, Copy, Debug, PartialEq, Eq)] pub enum OcrOp { - ProjectRequest, - ReadDocument, - AcquireAzureAdToken, + AcquireAzureAdToken(Reply), } -pub enum OcrOpResult { - Request { - request: Box>, - caller_token: bool, - }, - Document(OcrFileContent), - AzureAdToken(ResolvedCredential), +/// The caller's request as the host projects it. +pub struct OcrProjection { + pub request: LiteLLMOcrRequest, + /// The caller passed its own Azure AD token provider, which the host keeps. + pub caller_token: bool, } pub struct Ocr; -impl Route for Ocr { +impl Protocol for Ocr { type Response = LiteLLMOcrResponse; type Error = Error; + type Projection = OcrProjection; type Op = OcrOp; - type OpResult = OcrOpResult; type Chunk = std::convert::Infallible; type StreamHead = std::convert::Infallible; } -impl TokenRoute for Ocr { - fn acquire_token_op() -> OcrOp { - OcrOp::AcquireAzureAdToken - } - - fn token_credential(result: OcrOpResult) -> Option { - match result { - OcrOpResult::AzureAdToken(credential) => Some(credential), - _ => None, - } +impl TokenProtocol for Ocr { + fn acquire_token_op(reply: Reply) -> OcrOp { + OcrOp::AcquireAzureAdToken(reply) } } pub type OcrHost = HostChannel; -pub type OcrMachine = RouteMachine; +pub type OcrMachine = CallMachine; -/// The OCR call as a machine: projection, document reading and token acquisition are -/// host operations; everything else runs in Rust. +/// The OCR call as a machine: projection and token acquisition are host operations; +/// everything else runs in Rust. pub fn ocr_machine(client: OcrClient) -> OcrMachine { - RouteMachine::new(move |host| Box::pin(execute(client, host))) + CallMachine::new(move |host| Box::pin(execute(client, host))) } async fn execute(client: OcrClient, host: OcrHost) -> Result { - let OcrOpResult::Request { + let OcrProjection { request, caller_token, - } = host.route(OcrOp::ProjectRequest).await? - else { - return Err(MachineFault::Mismatch.into()); - }; + } = host.project().await?; let request = LiteLLMOcrRequest { azure_ad_token_provider: caller_token .then(|| HostTokenProvider::handle(host.clone())) .or(request.azure_ad_token_provider), - ..*request + ..request }; let caller_document = matches!(request.document, OcrDocumentInput::Document(_)); - let request = prepare_request_document(request, &host).await?; + let request = prepare_request_document(request).await?; perform_ocr_request(&client, request, &host, caller_document).await } async fn prepare_request_document( request: LiteLLMOcrRequest, - host: &OcrHost, ) -> Result { - let request = match &request.document { - OcrDocumentInput::HostReader { mime_type } => { - let mime_type = mime_type.clone(); - let OcrOpResult::Document(content) = host.route(OcrOp::ReadDocument).await? else { - return Err(MachineFault::Mismatch.into()); - }; - request.with_document(OcrDocumentInput::Bytes { - bytes: content.bytes, - file_name: content.file_name, - mime_type, - }) - } - _ => request, - }; if let OcrDocumentInput::Document(_) = &request.document { return request.map_document(super::document::prepare_document); } @@ -107,7 +78,6 @@ async fn prepare_request_document( .map_err(|error| Error::DocumentTask(Arc::new(error)))? } -type Reader = Box Result + Send + Sync>; type BeforeSend = Box Result + Send + Sync>; type Observer = Box; @@ -116,7 +86,6 @@ type Observer = Box; /// projection, and the optional observer sees and may rewrite the wire request. pub struct LocalOcrHost { request: Mutex>>, - reader: Option, before_send: Option, observer: Option, } @@ -125,22 +94,11 @@ impl LocalOcrHost { pub fn new(request: LiteLLMOcrRequest) -> Self { Self { request: Mutex::new(Some(request)), - reader: None, before_send: None, observer: None, } } - pub fn with_reader( - self, - reader: impl Fn() -> Result + Send + Sync + 'static, - ) -> Self { - Self { - reader: Some(Box::new(reader)), - ..self - } - } - pub fn with_before_send( self, before_send: impl Fn(WireRequest, &RequestContext) -> Result @@ -163,25 +121,21 @@ impl LocalOcrHost { } impl litellm_host::host::Host for LocalOcrHost { - async fn route(&self, op: OcrOp) -> Result { + async fn project(&self) -> Result { + self.request + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + .map(|request| OcrProjection { + request, + caller_token: false, + }) + .ok_or_else(|| Error::InvalidRequest("OCR request was already projected".into())) + } + + async fn custom_op(&self, op: OcrOp) -> Result<(), Error> { match op { - OcrOp::ProjectRequest => self - .request - .lock() - .unwrap_or_else(|error| error.into_inner()) - .take() - .map(|request| OcrOpResult::Request { - request: Box::new(request), - caller_token: false, - }) - .ok_or_else(|| Error::InvalidRequest("OCR request was already projected".into())), - OcrOp::ReadDocument => self - .reader - .as_ref() - .ok_or_else(|| Error::InvalidRequest("OCR host has no document reader".into())) - .and_then(|reader| reader()) - .map(OcrOpResult::Document), - OcrOp::AcquireAzureAdToken => { + OcrOp::AcquireAzureAdToken(_) => { Err(Error::Auth(litellm_auth::Error::AzureTokenAcquisition( "OCR host has no Azure AD token provider".into(), ))) @@ -2757,7 +2711,7 @@ pub(crate) mod tests { use litellm_auth_gcp::VertexAuth; use litellm_host::{ event::{CallEvent, MachineEvent, WireRequest}, - host::{Host, HostOp, HostResult}, + host::{Host, HostOp}, machine::{HostFailure, Machine, MachineStep}, }; use litellm_http::{ @@ -2776,7 +2730,7 @@ pub(crate) mod tests { use rstest::rstest; use serde_json::{Value, json}; - use crate::ocr::route::{LocalOcrHost, OcrOp, OcrOpResult, ocr_machine}; + use crate::ocr::route::{LocalOcrHost, OcrOp, OcrProjection, ocr_machine}; use crate::ocr::{ test_support::{ MockResponse, mock_server, ocr_client, perform_ocr, perform_ocr_with, wire_request, @@ -3212,42 +3166,42 @@ pub(crate) mod tests { crate::ocr::route::OcrMachine, ) { let mut machine = ocr_machine(client); - let mut result = None; let mut ops = Vec::new(); let outcome = loop { - let op = match machine.resume(result.take()).await { + let op = match machine.resume().await { Ok(MachineStep::Host(op)) => op, Ok(MachineStep::Complete(response)) => break Ok(response), Err(error) => break Err(error), }; let answer = match op { - HostOp::Route(op) => { - ops.push(match op { - OcrOp::ProjectRequest => "ProjectRequest", - OcrOp::ReadDocument => "ReadDocument", - OcrOp::AcquireAzureAdToken => "AcquireAzureAdToken", - }); - host.route(op) + HostOp::Project(reply) => { + ops.push("Project"); + host.project() .await - .map(HostResult::Route) + .map(|projection| reply.send(projection)) .map_err(HostFailure::Error) } - HostOp::BeforeSend { wire, .. } => { - ops.push("BeforeSend"); - intercept(*wire).map(|wire| HostResult::BeforeSend(Box::new(wire))) + HostOp::Custom(op) => { + ops.push(match op { + OcrOp::AcquireAzureAdToken(_) => "AcquireAzureAdToken", + }); + host.custom_op(op).await.map_err(HostFailure::Error) } - HostOp::Emit(event) => { + HostOp::BeforeSend { wire, reply, .. } => { + ops.push("BeforeSend"); + intercept(*wire).map(|wire| reply.send(wire)) + } + HostOp::Emit(event, reply) => { let event = CallEvent::Machine(event); ops.push(event_name(&event)); host.emit(&event) .await - .map(|()| HostResult::Emitted) + .map(|()| reply.send(())) .map_err(HostFailure::Error) } }; - match answer { - Ok(answer) => result = Some(answer), - Err(failure) => break machine.interrupt(failure).await, + if let Err(failure) = answer { + break machine.interrupt(failure).await; } }; (outcome, ops, machine) @@ -3269,8 +3223,8 @@ pub(crate) mod tests { assert!( matches!(outcome, Err(OcrError::InvalidRequest(message)) if message == "before_send failed") ); - assert_eq!(ops, ["ProjectRequest", "BeforeSend"]); - assert!(machine.resume(None).await.is_err()); + assert_eq!(ops, ["Project", "BeforeSend"]); + assert!(machine.resume().await.is_err()); } #[tokio::test] @@ -3306,80 +3260,24 @@ pub(crate) mod tests { server.await.unwrap(); assert_eq!(outcome.unwrap().pages[0].markdown, "native"); assert_eq!(seen.lock().unwrap().len(), 1); - assert_eq!(ops, ["ProjectRequest", "BeforeSend", "response"]); + assert_eq!(ops, ["Project", "BeforeSend", "response"]); assert!(matches!( - machine.resume(None).await, + machine.resume().await, Err(OcrError::InvalidRequest(_)) )); } - async fn drive_native_file_call( - request: crate::ocr::types::LiteLLMOcrRequest, - content: Result, - ) -> (Result, usize) { - let reads = Arc::new(Mutex::new(0)); - let counted = reads.clone(); - let content = Mutex::new(Some(content)); - let host = LocalOcrHost::new(request).with_reader(move || { - *counted.lock().unwrap() += 1; - content.lock().unwrap().take().unwrap() - }); - let outcome = perform_ocr_with(host).await; - let reads = *reads.lock().unwrap(); - (outcome, reads) - } - #[tokio::test] - async fn host_reader_documents_are_read_once_at_the_core_selected_point_and_encoded() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ - "pages":[{"index":0,"markdown":"file"}] - }))]) - .await; - let request = wire_request("mistral/model", &base, json!({})).with_document( - crate::ocr::types::OcrDocumentInput::HostReader { - mime_type: Some("application/pdf".into()), - }, - ); - let (response, reads) = drive_native_file_call( - request, - Ok(crate::ocr::types::OcrFileContent { - bytes: b"abc".as_slice().into(), - file_name: Some("scan.png".into()), - }), - ) - .await; - server.await.unwrap(); - assert_eq!(response.unwrap().pages[0].markdown, "file"); - assert_eq!(reads, 1); - assert!(seen.lock().unwrap()[0].contains("data:application/pdf;base64,YWJj")); - } - - #[tokio::test] - async fn host_reader_failures_and_empty_files_fail_before_the_provider_is_called() { + async fn empty_byte_documents_fail_before_the_provider_is_called() { let (base, seen, _server) = mock_server(vec![]).await; - let request = wire_request("mistral/model", &base, json!({})); - let failure = OcrError::InvalidRequest("reader exploded".into()); - let (response, reads) = drive_native_file_call( - request - .with_document(crate::ocr::types::OcrDocumentInput::HostReader { mime_type: None }), - Err(failure.clone()), - ) - .await; - assert!( - matches!(response.unwrap_err(), OcrError::InvalidRequest(message) if message == "reader exploded") - ); - assert_eq!(reads, 1); - - let request = wire_request("mistral/model", &base, json!({})); - let (response, _) = drive_native_file_call( - request - .with_document(crate::ocr::types::OcrDocumentInput::HostReader { mime_type: None }), - Ok(crate::ocr::types::OcrFileContent { + let request = wire_request("mistral/model", &base, json!({})).with_document( + crate::ocr::types::OcrDocumentInput::Bytes { bytes: Default::default(), file_name: None, - }), - ) - .await; + mime_type: None, + }, + ); + let response = perform_ocr_with(LocalOcrHost::new(request)).await; assert!(matches!(response.unwrap_err(), OcrError::EmptyFile)); assert!(seen.lock().unwrap().is_empty()); } @@ -3400,24 +3298,21 @@ pub(crate) mod tests { mime_type: None, }, ); - let (response, reads) = - drive_native_file_call(request, Err(OcrError::InvalidRequest("unused".into()))).await; + let (response, ops, _) = drive_until(ocr_client(), &LocalOcrHost::new(request), Ok).await; server.await.unwrap(); std::fs::remove_dir_all(&dir).unwrap(); assert_eq!(response.unwrap().pages[0].markdown, "path"); - assert_eq!(reads, 0); + assert_eq!(ops, ["Project", "BeforeSend", "response"]); assert!(seen.lock().unwrap()[0].contains("data:image/png;base64,YWJj")); let (base, seen, _server) = mock_server(vec![]).await; - let request = wire_request("mistral/model", &base, json!({})); - let (response, _) = drive_native_file_call( - request.with_document(crate::ocr::types::OcrDocumentInput::Path { + let request = wire_request("mistral/model", &base, json!({})).with_document( + crate::ocr::types::OcrDocumentInput::Path { path: path.clone(), mime_type: None, - }), - Err(OcrError::InvalidRequest("unused".into())), - ) - .await; + }, + ); + let response = perform_ocr_with(LocalOcrHost::new(request)).await; assert!(matches!( response.unwrap_err(), OcrError::FileRead { path: failed, source } if failed == path && source.kind() == std::io::ErrorKind::NotFound @@ -3441,28 +3336,25 @@ pub(crate) mod tests { assert!( matches!(outcome, Err(OcrError::InvalidRequest(message)) if message == "cancelled") ); - assert_eq!(ops, ["ProjectRequest", "BeforeSend"]); - assert!(machine.resume(Some(HostResult::Emitted)).await.is_err()); + assert_eq!(ops, ["Project", "BeforeSend"]); + assert!(machine.resume().await.is_err()); } #[tokio::test] - async fn missing_host_result_preserves_pending_operation() { + async fn resuming_before_answering_preserves_pending_operation() { let request = wire_request("mistral/model", "http://127.0.0.1:1", json!({})); let mut machine = ocr_machine(ocr_client()); + let Ok(MachineStep::Host(HostOp::Project(reply))) = machine.resume().await else { + panic!("expected the projection op first"); + }; + assert!(machine.resume().await.is_err()); + reply.send(OcrProjection { + request, + caller_token: false, + }); assert!(matches!( - machine.resume(None).await.unwrap(), - MachineStep::Host(HostOp::Route(OcrOp::ProjectRequest)) - )); - assert!(machine.resume(None).await.is_err()); - assert!(matches!( - machine - .resume(Some(HostResult::Route(OcrOpResult::Request { - request: Box::new(request), - caller_token: false, - }))) - .await - .unwrap(), - MachineStep::Host(HostOp::BeforeSend { .. }) + machine.resume().await, + Ok(MachineStep::Host(HostOp::BeforeSend { .. })) )); } @@ -3623,20 +3515,18 @@ pub(crate) mod tests { }; let host = LocalOcrHost::new(request); let mut machine = ocr_machine(ocr_client()); - let mut result = None; tokio::time::timeout(std::time::Duration::from_secs(2), async { loop { tokio::select! { _ = entered.notified() => break, - step = machine.resume(result.take()) => { - result = Some(match step.unwrap() { - MachineStep::Host(HostOp::Route(op)) => HostResult::Route(host.route(op).await.unwrap()), - MachineStep::Host(HostOp::BeforeSend { wire, .. }) => { - HostResult::BeforeSend(wire) - } - MachineStep::Host(HostOp::Emit(_)) => HostResult::Emitted, + step = machine.resume() => { + match step.unwrap() { + MachineStep::Host(HostOp::Project(reply)) => reply.send(host.project().await.unwrap()), + MachineStep::Host(HostOp::Custom(op)) => host.custom_op(op).await.unwrap(), + MachineStep::Host(HostOp::BeforeSend { wire, reply, .. }) => reply.send(*wire), + MachineStep::Host(HostOp::Emit(_, reply)) => reply.send(()), MachineStep::Complete(_) => panic!("pending provider completed"), - }); + } } } } @@ -3661,24 +3551,23 @@ pub(crate) mod tests { } impl Host for CallerTokenHost { - async fn route(&self, op: OcrOp) -> Result { + async fn project(&self) -> Result { + self.trace.lock().unwrap().push("project".into()); + Ok(OcrProjection { + request: self.request.lock().unwrap().take().unwrap(), + caller_token: true, + }) + } + + async fn custom_op(&self, op: OcrOp) -> Result<(), OcrError> { match op { - OcrOp::ProjectRequest => { - self.trace.lock().unwrap().push("project".into()); - Ok(OcrOpResult::Request { - request: Box::new(self.request.lock().unwrap().take().unwrap()), - caller_token: true, - }) - } - OcrOp::AcquireAzureAdToken => { + OcrOp::AcquireAzureAdToken(reply) => { self.trace.lock().unwrap().push("token".into()); - Ok(OcrOpResult::AzureAdToken( - litellm_auth::ResolvedCredential::Static(litellm_auth::SecretValue::new( - "caller-token", - )), - )) + reply.send(litellm_auth::ResolvedCredential::Static( + litellm_auth::SecretValue::new("caller-token"), + )); + Ok(()) } - OcrOp::ReadDocument => Err(OcrError::InvalidRequest("no reader".into())), } } @@ -3761,18 +3650,18 @@ pub(crate) mod tests { }); let host = LocalOcrHost::new(wire_request("mistral/model", &base, json!({}))); let mut machine = ocr_machine(ocr_client()); - let mut result = None; tokio::time::timeout(std::time::Duration::from_secs(2), async { loop { tokio::select! { _ = received.notified() => break, - step = machine.resume(result.take()) => { - result = Some(match step.unwrap() { - MachineStep::Host(HostOp::Route(op)) => HostResult::Route(host.route(op).await.unwrap()), - MachineStep::Host(HostOp::BeforeSend { wire, .. }) => HostResult::BeforeSend(wire), - MachineStep::Host(HostOp::Emit(_)) => HostResult::Emitted, + step = machine.resume() => { + match step.unwrap() { + MachineStep::Host(HostOp::Project(reply)) => reply.send(host.project().await.unwrap()), + MachineStep::Host(HostOp::Custom(op)) => host.custom_op(op).await.unwrap(), + MachineStep::Host(HostOp::BeforeSend { wire, reply, .. }) => reply.send(*wire), + MachineStep::Host(HostOp::Emit(_, reply)) => reply.send(()), MachineStep::Complete(_) => panic!("the stalled provider completed"), - }); + } } } } diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs index 59c9cec8da9..20a21e43676 100644 --- a/litellm-rust/crates/core/src/ocr/types.rs +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -25,9 +25,6 @@ pub enum OcrDocumentInput { file_name: Option, mime_type: Option, }, - HostReader { - mime_type: Option, - }, } impl From for OcrDocumentInput { @@ -45,12 +42,6 @@ impl From for OcrDocumentInput { } } -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct OcrFileContent { - pub bytes: Bytes, - pub file_name: Option, -} - /// Caller-supplied connection overrides for a [`LiteLLMOcrRequest`], in the /// shape hosts receive them: JSON-ish headers, optional timeout, optional /// credentials, and per-field provenance in `input_sources`. diff --git a/litellm-rust/crates/coroutine/AGENTS.md b/litellm-rust/crates/coroutine/AGENTS.md new file mode 100644 index 00000000000..fcb4f4df47f --- /dev/null +++ b/litellm-rust/crates/coroutine/AGENTS.md @@ -0,0 +1,31 @@ +# Requirements + +Core must pause mid-call to ask the host for things it cannot do itself (Python callbacks, secret and token reads, `before_send` rewrites, stream demand), then continue where it stopped. Any change to this crate must keep every requirement below; the alternatives section says which one each rejected design breaks + +- R1 Core never calls the host: it names an op and waits for the answer, so it stays free of PyO3 and of any other host runtime +- R2 Async host work is awaited by the host's own driver in the caller's asyncio task (`litellm/rust_bridge/lifecycle.py`), so `contextvars` writes reach the caller; a Rust-side `into_future` would run it in a copied context +- R3 The body awaits real I/O (HTTP, `spawn_blocking`, timers) between yields, so `resume` is itself a future driven by the caller's runtime +- R4 Route code stays straight-line async (`host.route(OcrOp::ReadDocument).await?`) instead of hand-written states +- R5 Each op fixes its answer type at compile time: a host cannot answer `ReadDocument` with a token, and core never matches a result variant it did not ask for +- R6 A yield the body makes while being resumed is returned by that same poll, so the host driver's inline first poll needs no extra event-loop turn per op +- R7 No task is spawned: `cancel`, or dropping the coroutine, drops the body, and nothing waits forever on an answer that cannot come +- R8 Several yields can be pending at once, since route code hands clones of its `Co` to token providers and hooks +- R9 Stable Rust + +# Other implementations and why they do not fit + +- Nightly `std::ops::Coroutine`: breaks R9, and its body cannot await futures between yields (R3) +- `genawaiter`: resumes async bodies only with a noop waker, so the body cannot await real I/O (R3) +- `simple_coro`: typestate `Coro` makes answering before resuming a compile-time rule, but its body cannot await arbitrary futures (R3) and its reply type `R` is fixed per coroutine (R5) +- `corosensei` and other stackful coroutines: sync bodies on their own stack, no async I/O inside (R3) +- A hand-written phase enum with an `advance` match (the old `HostPhase`): every await point becomes a state (R4) +- An injected host trait with `async fn`s: core would call the host itself (R1, R2) +- Sans-IO, where core does no I/O and HTTP becomes one more host op: keeps every requirement and makes `resume` a pure step function, but HTTP, streaming, retries and timeouts would move out of core into every bridge; the one real alternative, not taken +- Temporal's Rust workflow SDK (`WorkflowFuture`, `WfContext`) is the closest precedent: an `async fn` polled in place, commands sent over a channel with a oneshot to unblock them. Roles are inverted there (the language SDK owns the program, core answers), and its workflow body may not do real I/O + +# Tradeoffs accepted + +- A tokio `mpsc` channel plus a `oneshot` per yield instead of compiler-generated states +- Protocol mistakes (resuming before answering, resuming after the end) are runtime `ResumeError`s, not compile errors +- Pending yields come out one per `resume`, in the order they were made, and each reply goes back to the yield that made it (R8) +- An answer sent after its yield stopped waiting (for example the body timed out on it) is discarded, since the body already moved on diff --git a/litellm-rust/crates/coroutine/Cargo.toml b/litellm-rust/crates/coroutine/Cargo.toml new file mode 100644 index 00000000000..3ff79ac5f2c --- /dev/null +++ b/litellm-rust/crates/coroutine/Cargo.toml @@ -0,0 +1,15 @@ +[package] +name = "litellm-coroutine" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true +description = "Async coroutines on stable Rust whose every yield carries its own typed reply" + +[dependencies] +thiserror.workspace = true +tokio = { workspace = true, features = ["sync"] } + +[dev-dependencies] +rstest.workspace = true +tokio = { workspace = true, features = ["rt", "macros", "time"] } diff --git a/litellm-rust/crates/coroutine/src/co.rs b/litellm-rust/crates/coroutine/src/co.rs new file mode 100644 index 00000000000..d84b041931c --- /dev/null +++ b/litellm-rust/crates/coroutine/src/co.rs @@ -0,0 +1,42 @@ +use std::sync::Weak; + +use tokio::sync::mpsc; + +use crate::{Abandoned, Reply, reply}; + +pub(crate) struct Request { + pub(crate) value: Y, + pub(crate) outstanding: Weak<()>, +} + +/// The body's handle for yielding, `genawaiter`'s `Co`. +pub struct Co { + yields: mpsc::UnboundedSender>, +} + +impl Clone for Co { + fn clone(&self) -> Self { + Self { + yields: self.yields.clone(), + } + } +} + +impl Co { + pub(crate) fn new(yields: mpsc::UnboundedSender>) -> Self { + Self { yields } + } + + /// Yields the value `ask` builds around a fresh [`Reply`] and waits for its answer. + pub async fn yield_(&self, ask: impl FnOnce(Reply) -> Y) -> Result { + let (reply, answer) = reply(); + let outstanding = reply.outstanding(); + self.yields + .send(Request { + value: ask(reply), + outstanding, + }) + .map_err(|_| Abandoned)?; + answer.await + } +} diff --git a/litellm-rust/crates/coroutine/src/coroutine.rs b/litellm-rust/crates/coroutine/src/coroutine.rs new file mode 100644 index 00000000000..fa34f816a37 --- /dev/null +++ b/litellm-rust/crates/coroutine/src/coroutine.rs @@ -0,0 +1,94 @@ +use std::{ + future::{Future, poll_fn}, + pin::Pin, + sync::Weak, + task::{Context, Poll}, +}; + +use tokio::sync::mpsc; + +use crate::{Co, ResumeError, co::Request}; + +/// What one `resume` produced, as in [`std::ops::CoroutineState`]. +#[derive(Debug, PartialEq, Eq)] +pub enum CoroutineState { + Yielded(Y), + Complete(C), +} + +type Body = Pin + Send>>; + +enum Step { + Yielded(Request), + Complete(C), +} + +fn queued( + yields: &mut mpsc::UnboundedReceiver>, + context: &mut Context<'_>, +) -> Option> { + match yields.poll_recv(context) { + Poll::Ready(request) => request, + Poll::Pending => None, + } +} + +pub struct Coroutine { + body: Option>, + yields: mpsc::UnboundedReceiver>, + outstanding: Weak<()>, +} + +impl Coroutine { + /// Builds the body from `producer`. Nothing runs until the first `resume`. + pub fn new(producer: impl FnOnce(Co) -> F) -> Self + where + F: Future + Send + 'static, + { + let (sender, yields) = mpsc::unbounded_channel(); + Self { + body: Some(Box::pin(producer(Co::new(sender)))), + yields, + outstanding: Weak::new(), + } + } + + pub async fn resume(&mut self) -> Result, ResumeError> { + let Some(body) = self.body.as_mut() else { + return Err(ResumeError::Finished); + }; + if self.outstanding.strong_count() > 0 { + return Err(ResumeError::Unanswered); + } + let yields = &mut self.yields; + let step = poll_fn(|context| { + if let Some(request) = queued(yields, context) { + return Poll::Ready(Step::Yielded(request)); + } + if let Poll::Ready(output) = body.as_mut().poll(context) { + return Poll::Ready(Step::Complete(output)); + } + queued(yields, context) + .map_or(Poll::Pending, |request| Poll::Ready(Step::Yielded(request))) + }) + .await; + match step { + Step::Yielded(Request { value, outstanding }) => { + self.outstanding = outstanding; + Ok(CoroutineState::Yielded(value)) + } + Step::Complete(output) => { + self.cancel(); + Ok(CoroutineState::Complete(output)) + } + } + } + + /// Drops the body and fails every yield still waiting, or yet to be made, with + /// [`Abandoned`](crate::Abandoned). + pub fn cancel(&mut self) { + self.body = None; + self.yields.close(); + while self.yields.try_recv().is_ok() {} + } +} diff --git a/litellm-rust/crates/coroutine/src/error.rs b/litellm-rust/crates/coroutine/src/error.rs new file mode 100644 index 00000000000..b8fded23fdb --- /dev/null +++ b/litellm-rust/crates/coroutine/src/error.rs @@ -0,0 +1,14 @@ +/// A `resume` the coroutine refused, leaving it as it was. +#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)] +pub enum ResumeError { + #[error("coroutine resumed after it finished")] + Finished, + #[error("coroutine resumed before the reply to its last yield was sent or dropped")] + Unanswered, +} + +/// No answer will come to a yield: its [`Reply`](crate::Reply) was dropped unsent, or the +/// coroutine it was sent to is gone. +#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)] +#[error("the yield was abandoned before it was answered")] +pub struct Abandoned; diff --git a/litellm-rust/crates/coroutine/src/lib.rs b/litellm-rust/crates/coroutine/src/lib.rs new file mode 100644 index 00000000000..636aaf1b2b6 --- /dev/null +++ b/litellm-rust/crates/coroutine/src/lib.rs @@ -0,0 +1,12 @@ +//! Async coroutines on stable Rust whose every yield carries its own typed [`Reply`]. +//! See `AGENTS.md` for the requirement, the alternatives and the contracts. + +mod co; +mod coroutine; +mod error; +mod reply; + +pub use co::Co; +pub use coroutine::{Coroutine, CoroutineState}; +pub use error::{Abandoned, ResumeError}; +pub use reply::{Answer, Reply, reply}; diff --git a/litellm-rust/crates/coroutine/src/reply.rs b/litellm-rust/crates/coroutine/src/reply.rs new file mode 100644 index 00000000000..b3cb7da2e9d --- /dev/null +++ b/litellm-rust/crates/coroutine/src/reply.rs @@ -0,0 +1,60 @@ +use std::{ + fmt, + future::Future, + pin::Pin, + sync::{Arc, Weak}, + task::{Context, Poll}, +}; + +use tokio::sync::oneshot; + +use crate::Abandoned; + +/// The one way to answer a yield. Sending or dropping it settles the yield. +pub struct Reply { + slot: oneshot::Sender, + outstanding: Arc<()>, +} + +impl Reply { + /// An answer the yield no longer awaits is discarded. + pub fn send(self, answer: A) { + let _ = self.slot.send(answer); + } + + /// Alive until this reply is sent or dropped. + pub(crate) fn outstanding(&self) -> Weak<()> { + Arc::downgrade(&self.outstanding) + } +} + +impl fmt::Debug for Reply { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("Reply") + } +} + +/// The waiting end of a [`Reply`]. +pub struct Answer { + slot: oneshot::Receiver, +} + +impl Future for Answer { + type Output = Result; + + fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll { + Pin::new(&mut self.slot) + .poll(context) + .map(|answer| answer.map_err(|_| Abandoned)) + } +} + +/// A reply outside any coroutine, for answering a host operation directly. +pub fn reply() -> (Reply, Answer) { + let (slot, answer) = oneshot::channel(); + let reply = Reply { + slot, + outstanding: Arc::new(()), + }; + (reply, Answer { slot: answer }) +} diff --git a/litellm-rust/crates/coroutine/tests/coroutine.rs b/litellm-rust/crates/coroutine/tests/coroutine.rs new file mode 100644 index 00000000000..91d195503df --- /dev/null +++ b/litellm-rust/crates/coroutine/tests/coroutine.rs @@ -0,0 +1,256 @@ +use std::{ + future::Future, + sync::{Arc, Mutex}, + time::Duration, +}; + +use litellm_coroutine::{Abandoned, Co, Coroutine, CoroutineState, Reply, ResumeError, reply}; +use rstest::rstest; +use tokio::time::timeout; + +#[derive(Debug)] +enum Ask { + Name(Reply<&'static str>), + Count(Reply), +} + +type Test = Coroutine; + +fn yielded(state: Result, ResumeError>) -> Ask { + match state { + Ok(CoroutineState::Yielded(ask)) => ask, + Ok(CoroutineState::Complete(_)) => panic!("expected a yield, the body returned"), + Err(error) => panic!("expected a yield, resume failed: {error}"), + } +} + +fn complete(state: Result, ResumeError>) -> C { + match state { + Ok(CoroutineState::Complete(output)) => output, + Ok(CoroutineState::Yielded(ask)) => panic!("expected completion, got {ask:?}"), + Err(error) => panic!("expected completion, resume failed: {error}"), + } +} + +fn name(ask: Ask) -> Reply<&'static str> { + match ask { + Ask::Name(reply) => reply, + other => panic!("expected a name ask, got {other:?}"), + } +} + +fn count(ask: Ask) -> Reply { + match ask { + Ask::Count(reply) => reply, + other => panic!("expected a count ask, got {other:?}"), + } +} + +/// A body parked at one name ask, with nothing else going on. +fn suspended_once() -> Test> { + Coroutine::new(|co| async move { co.yield_(Ask::Name).await }) +} + +#[tokio::test] +async fn each_typed_answer_resumes_the_yield_that_asked_for_it() { + let mut coroutine: Test = Coroutine::new(|co| async move { + let first = co.yield_(Ask::Name).await.unwrap(); + let second = co.yield_(Ask::Count).await.unwrap(); + format!("{first}+{second}") + }); + + name(yielded(coroutine.resume().await)).send("a"); + count(yielded(coroutine.resume().await)).send(2); + + assert_eq!(complete(coroutine.resume().await), "a+2"); +} + +/// A driver that polls `resume` once, inline, sees every yield the body makes during +/// that poll instead of being sent back to its event loop. +#[test] +fn a_yield_made_while_resuming_is_returned_by_that_same_poll() { + let mut coroutine: Test = Coroutine::new(|co| async move { + let first = co.yield_(Ask::Count).await.unwrap(); + let second = co.yield_(Ask::Count).await.unwrap(); + first + second + }); + let mut context = std::task::Context::from_waker(std::task::Waker::noop()); + let mut poll_once = + |coroutine: &mut Test| match std::pin::pin!(coroutine.resume()).poll(&mut context) { + std::task::Poll::Ready(state) => state, + std::task::Poll::Pending => panic!("resume needed a second poll"), + }; + + count(yielded(poll_once(&mut coroutine))).send(1); + count(yielded(poll_once(&mut coroutine))).send(2); + + assert_eq!(complete(poll_once(&mut coroutine)), 3); +} + +#[tokio::test] +async fn the_body_awaits_real_futures_between_yields() { + let mut coroutine: Test = Coroutine::new(|co| async move { + tokio::time::sleep(Duration::from_millis(5)).await; + co.yield_(Ask::Count).await.unwrap() + }); + + count(yielded(coroutine.resume().await)).send(7); + + assert_eq!(complete(coroutine.resume().await), 7); +} + +#[tokio::test] +async fn concurrent_yields_come_out_in_order_and_are_answered_separately() { + let mut coroutine: Test<(&str, u32)> = Coroutine::new(|co| async move { + let (first, second) = tokio::join!(co.yield_(Ask::Name), co.yield_(Ask::Count)); + (first.unwrap(), second.unwrap()) + }); + + name(yielded(coroutine.resume().await)).send("one"); + count(yielded(coroutine.resume().await)).send(2); + + assert_eq!(complete(coroutine.resume().await), ("one", 2)); +} + +#[tokio::test] +async fn resuming_before_the_reply_is_settled_is_refused_and_keeps_the_yield_waiting() { + let mut coroutine = suspended_once(); + let reply = name(yielded(coroutine.resume().await)); + + assert_eq!( + coroutine.resume().await.unwrap_err(), + ResumeError::Unanswered + ); + + reply.send("real"); + assert_eq!(complete(coroutine.resume().await), Ok("real")); +} + +#[tokio::test] +async fn a_dropped_reply_abandons_its_yield() { + let mut coroutine = suspended_once(); + drop(yielded(coroutine.resume().await)); + + assert_eq!(complete(coroutine.resume().await), Err(Abandoned)); +} + +#[tokio::test] +async fn an_answer_the_yield_no_longer_awaits_is_discarded() { + let mut coroutine: Test<&str> = Coroutine::new(|co| async move { + tokio::select! { + biased; + _ = co.yield_(Ask::Name) => unreachable!("the answer comes after the body moved on"), + () = std::future::ready(()) => {} + } + co.yield_(Ask::Name).await.unwrap() + }); + let stale = name(yielded(coroutine.resume().await)); + stale.send("stale"); + + name(yielded(coroutine.resume().await)).send("fresh"); + + assert_eq!(complete(coroutine.resume().await), "fresh"); +} + +#[rstest] +#[case::returned(false)] +#[case::cancelled(true)] +#[tokio::test] +async fn a_finished_coroutine_refuses_to_resume(#[case] cancel: bool) { + let mut coroutine = suspended_once(); + let reply = name(yielded(coroutine.resume().await)); + if cancel { + coroutine.cancel(); + } else { + reply.send("done"); + complete(coroutine.resume().await).unwrap(); + } + + assert_eq!(coroutine.resume().await.unwrap_err(), ResumeError::Finished); +} + +#[tokio::test] +async fn a_dropped_resume_leaves_the_coroutine_resumable() { + let mut coroutine: Test = Coroutine::new(|co| async move { + tokio::time::sleep(Duration::from_millis(20)).await; + co.yield_(Ask::Count).await.unwrap() + }); + assert!( + timeout(Duration::from_millis(1), coroutine.resume()) + .await + .is_err() + ); + + count(yielded(coroutine.resume().await)).send(3); + + assert_eq!(complete(coroutine.resume().await), 3); +} + +struct Dropped(Arc>); + +impl Drop for Dropped { + fn drop(&mut self) { + *self.0.lock().unwrap() = true; + } +} + +#[tokio::test] +async fn cancel_drops_the_body() { + let dropped = Arc::new(Mutex::new(false)); + let guard = Dropped(Arc::clone(&dropped)); + let mut coroutine: Test<()> = Coroutine::new(|co| async move { + let _guard = guard; + co.yield_(Ask::Count).await.unwrap(); + }); + let _reply = yielded(coroutine.resume().await); + + coroutine.cancel(); + + assert!(*dropped.lock().unwrap()); +} + +#[rstest] +#[case::cancelled(true)] +#[case::dropped(false)] +#[tokio::test] +async fn a_co_that_escaped_the_body_is_abandoned_once_the_coroutine_ends(#[case] cancel: bool) { + let escaped: Arc>>> = Arc::default(); + let slot = Arc::clone(&escaped); + let mut coroutine: Test<()> = Coroutine::new(move |co| { + *slot.lock().unwrap() = Some(co.clone()); + async move { + co.yield_(Ask::Count).await.unwrap(); + } + }); + let _reply = yielded(coroutine.resume().await); + let co = escaped.lock().unwrap().take().unwrap(); + let waiting = tokio::spawn(async move { co.yield_(Ask::Name).await }); + tokio::task::yield_now().await; + + if cancel { + coroutine.cancel(); + } else { + drop(coroutine); + } + + let outcome = timeout(Duration::from_secs(1), waiting) + .await + .expect("an escaped yield waits forever") + .unwrap(); + assert_eq!(outcome, Err(Abandoned)); +} + +#[rstest] +#[case::sent(true)] +#[case::dropped(false)] +#[tokio::test] +async fn a_detached_reply_settles_its_answer(#[case] send: bool) { + let (reply, answer) = reply::(); + if send { + reply.send(5); + } else { + drop(reply); + } + + assert_eq!(answer.await, if send { Ok(5) } else { Err(Abandoned) }); +} diff --git a/litellm-rust/crates/host-python/AGENTS.md b/litellm-rust/crates/host-python/AGENTS.md index 5aca13eeb18..7c1919f9f39 100644 --- a/litellm-rust/crates/host-python/AGENTS.md +++ b/litellm-rust/crates/host-python/AGENTS.md @@ -1,8 +1,8 @@ - Target invariants; implementation and runtime validation may lag these rules -- Keep this crate the CPython runtime adapter and nothing more: Serde marshalling, interpreter detachment, tokio/asyncio glue, the `Execution` handle, the call driver and the `PythonLifecycle`/`RouteHost` traits +- Keep this crate the CPython runtime adapter and nothing more: Serde marshalling, interpreter detachment, tokio/asyncio glue, the `Execution` handle, the call driver and the `PythonLifecycle`/`ProtocolHost` traits - No LiteLLM domain dependencies beyond `litellm-host`: no route types, no `Logging` policy, no public API registration, no cdylib build features - The driver emits `Succeeded` or `Failed` exactly once and never dispatches after a cancellation; which Python objects consume those events is the adapter's business - - `RouteHost::invoke` receives the keyword view the adapter's `begin` returned, not the caller's dict; a route host that projects from it inherits that adapter's rewrites (for the legacy adapter: setup, deployment hooks, credential inheritance) + - `ProtocolHost::project` receives the keyword view the adapter's `begin` returned, not the caller's dict; a protocol host that projects from it inherits that adapter's rewrites (for the legacy adapter: setup, deployment hooks, credential inheritance) - A native failure, including one a host op returns as `InvokeError::Native`, is classified exactly once through the route's `classify`; a Python exception raised inside the call, and a failure in `begin` or `after_success`, is raised as is - A failing `classify` is raised with the native error's text as its `__context__`, never swallowed - Use standard PyO3 ownership and conversion APIs diff --git a/litellm-rust/crates/host-python/Cargo.toml b/litellm-rust/crates/host-python/Cargo.toml index fb6379dc35a..c1b35c0f69d 100644 --- a/litellm-rust/crates/host-python/Cargo.toml +++ b/litellm-rust/crates/host-python/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +bytes.workspace = true futures-util.workspace = true litellm-host.workspace = true pyo3.workspace = true diff --git a/litellm-rust/crates/host-python/src/adapter.rs b/litellm-rust/crates/host-python/src/adapter.rs index 3a4cb49be4d..87481aa89b7 100644 --- a/litellm-rust/crates/host-python/src/adapter.rs +++ b/litellm-rust/crates/host-python/src/adapter.rs @@ -1,5 +1,5 @@ use litellm_host::event::{FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest}; -use litellm_host::route::Route; +use litellm_host::protocol::Protocol; use pyo3::exceptions::PyRuntimeError; use pyo3::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; @@ -85,7 +85,7 @@ pub trait PythonLifecycle: Send + Sync { fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>; } -/// Why a route operation the host answered did not produce a result: the route's own code +/// Why a custom operation the host answered did not produce a result: the route's own code /// rejected it, which the route classifies like any other native failure, or Python code /// raised, which reaches the caller as it was raised. #[derive(Debug)] @@ -100,45 +100,54 @@ impl From for InvokeError { } } -/// The Python side of one route: answers the route's own operations, builds the public +/// The Python side of one protocol: answers its custom operations, builds the public /// response and classifies native failures into public exceptions. -pub trait RouteHost: Send + Sync { - type Route: Route; +pub trait ProtocolHost: Send + Sync { + type Protocol: Protocol; /// The public exception a native failure maps to, kept as a value until the driver /// raises it. type Failure: Into; - /// `arguments` is the keyword view the lifecycle's `begin` produced, not the - /// caller's own dict. A route host that projects from it inherits whatever that - /// adapter rewrote. - fn invoke( + /// Projects the call's request. `arguments` is the keyword view the lifecycle's + /// `begin` produced, not the caller's own dict, so the projection inherits whatever + /// that adapter rewrote. + fn project( &mut self, py: Python<'_>, arguments: &Bound<'_, PyDict>, - op: ::Op, - ) -> Result<::OpResult, InvokeError<::Error>>; + ) -> Result< + ::Projection, + InvokeError<::Error>, + >; + + /// Answers `op` through its reply. + fn invoke( + &mut self, + py: Python<'_>, + op: ::Op, + ) -> Result<(), InvokeError<::Error>>; fn complete( &mut self, py: Python<'_>, - response: ::Response, + response: ::Response, ) -> PyResult>; /// One streamed chunk as the caller receives it. fn chunk( &mut self, py: Python<'_>, - chunk: ::Chunk, + chunk: ::Chunk, ) -> PyResult>; fn classify( &self, py: Python<'_>, - error: ::Error, + error: ::Error, ) -> PyResult; - fn host_error(error: &PyErr) -> ::Error; + fn host_error(error: &PyErr) -> ::Error; fn close(&mut self, py: Python<'_>); diff --git a/litellm-rust/crates/host-python/src/driver.rs b/litellm-rust/crates/host-python/src/driver.rs index 77a294d274b..aaa0752522b 100644 --- a/litellm-rust/crates/host-python/src/driver.rs +++ b/litellm-rust/crates/host-python/src/driver.rs @@ -2,10 +2,11 @@ use std::sync::Arc; use std::task::Poll; use futures_util::future::{AbortHandle, Abortable}; +use litellm_host::event::WireRequest; use litellm_host::event::{FailureOrigin, Timing, epoch_seconds}; -use litellm_host::host::{Demand, HostOp, HostResult, HostStep}; +use litellm_host::host::{Demand, HostOp, HostStep, Reply}; use litellm_host::machine::{HostFailure, Machine, MachineStep}; -use litellm_host::route::Route; +use litellm_host::protocol::Protocol; use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError}; use pyo3::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; @@ -13,21 +14,21 @@ use pyo3::types::PyDict; use tokio::sync::Mutex; use crate::adapter::{ - InvokeError, LifecycleEvent, LifecycleStep, PythonLifecycle, RouteHost, missing_state, + InvokeError, LifecycleEvent, LifecycleStep, ProtocolHost, PythonLifecycle, missing_state, }; use crate::execution::{poll_async_value, run_async_value, run_sync_value}; use crate::handle::{Execution, ExecutionBody, ExecutionStep}; -type RouteOf = ::Route; -type ErrorOf = as Route>::Error; -type ResponseOf = as Route>::Response; -type NativeStep = MachineStep, ResponseOf>; +type ProtocolOf = ::Protocol; +type ErrorOf = as Protocol>::Error; +type ResponseOf = as Protocol>::Response; +type NativeStep = MachineStep, ResponseOf>; type NativeResult = Result, ErrorOf>; -type NativeResume = Option>, HostFailure>>>; +type Interruption = Option>>; type MachineResult = Result< - MachineStep<::Route, ::Complete>, - <::Route as Route>::Error, + MachineStep<::Protocol, ::Complete>, + <::Protocol as Protocol>::Error, >; struct MachineState { @@ -44,12 +45,11 @@ enum Stage { Failed(Py), } -#[derive(Clone, Copy)] enum Expect { Started, Arguments, - Wire, - Emitted, + Wire(Reply), + Emitted(Reply<()>), Response, Terminal, } @@ -58,20 +58,30 @@ enum Pending { Native, Adapter(Expect), /// The stream handed to the caller waits for its next read or its close. - Consumer, + Consumer(Reply), } -enum Next { +/// A route answer as the driver resumes on it: a Python exception interrupts the call as +/// raised, a native rejection resumes the machine with it. +fn answered(answer: Result<(), InvokeError>) -> PyResult> { + match answer { + Ok(()) => Ok(Ok(())), + Err(InvokeError::Native(error)) => Ok(Err(error)), + Err(InvokeError::Python(error)) => Err(error), + } +} + +enum Next { Return(ExecutionStep), Continue(HostStep, Py>), } struct PythonDriver where - H: RouteHost, - M: Machine> + 'static, + H: ProtocolHost, + M: Machine> + 'static, { - route: H, + host: H, adapter: Box, machine: Option>>>, arguments: Option>, @@ -89,17 +99,17 @@ where pub fn run_call( py: Python<'_>, machine: M, - route: H, + host: H, adapter: Box, arguments: Py, asynchronous: bool, ) -> PyResult> where - H: RouteHost + 'static, - M: Machine> + 'static, + H: ProtocolHost + 'static, + M: Machine> + 'static, { let mut driver = PythonDriver { - route, + host, adapter, machine: Some(Arc::new(Mutex::new(MachineState { machine, @@ -141,8 +151,8 @@ fn is_cancellation(py: Python<'_>, error: &PyErr) -> bool { impl PythonDriver where - H: RouteHost, - M: Machine> + 'static, + H: ProtocolHost, + M: Machine> + 'static, { fn timing(&self) -> Timing { Timing { @@ -172,13 +182,13 @@ where self.run_steps(py, HostStep::Ready(result)) } (Some(Pending::Native), Some(Err(error))) => self.interrupt(py, error), - (Some(Pending::Consumer), Some(read)) => { - let demand = if read.is_ok() { + (Some(Pending::Consumer(reply)), Some(read)) => { + reply.send(if read.is_ok() { Demand::More } else { Demand::Detached - }; - self.resume_machine(py, Some(Ok(HostResult::Demand(demand)))) + }); + self.resume_machine(py, None) } (Some(Pending::Adapter(expect)), Some(result)) => { match self.adapter.resume(py, result) { @@ -196,22 +206,24 @@ where step: LifecycleStep, expect: Expect, ) -> PyResult { + if let LifecycleStep::Await(awaitable) = step { + self.pending = Some(Pending::Adapter(expect)); + return Ok(ExecutionStep::Await(awaitable)); + } match (expect, step) { - (_, LifecycleStep::Await(awaitable)) => { - self.pending = Some(Pending::Adapter(expect)); - Ok(ExecutionStep::Await(awaitable)) - } (Expect::Started, LifecycleStep::Done) => self.begin(py), (Expect::Arguments, LifecycleStep::Arguments(arguments)) => { self.arguments = Some(arguments); self.stage = Stage::Call; self.resume_machine(py, None) } - (Expect::Wire, LifecycleStep::Wire(wire)) => { - self.resume_machine(py, Some(Ok(HostResult::BeforeSend(wire)))) + (Expect::Wire(reply), LifecycleStep::Wire(wire)) => { + reply.send(*wire); + self.resume_machine(py, None) } - (Expect::Emitted, LifecycleStep::Done) => { - self.resume_machine(py, Some(Ok(HostResult::Emitted))) + (Expect::Emitted(reply), LifecycleStep::Done) => { + reply.send(()); + self.resume_machine(py, None) } (Expect::Response, LifecycleStep::Response(response)) => self.succeeded(py, response), (Expect::Terminal, LifecycleStep::Done) => match &self.stage { @@ -242,9 +254,9 @@ where fn resume_machine( &mut self, py: Python<'_>, - result: NativeResume, + interruption: Interruption, ) -> PyResult { - let step = self.resume_core(py, result)?; + let step = self.resume_core(py, interruption)?; self.run_steps(py, step) } @@ -277,53 +289,62 @@ where } Err(error) => return self.machine_failed(py, error).map(Next::Return), }; - let answer = match op { - HostOp::Route(op) => { + let answered = match op { + HostOp::Project(reply) => { let arguments = self.arguments.as_ref().ok_or_else(missing_state)?; - match self.route.invoke(py, arguments.bind(py), op) { - Ok(result) => Ok(HostResult::Route(result)), - Err(InvokeError::Native(error)) => { - return self - .resume_core(py, Some(Err(HostFailure::Error(error)))) - .map(Next::Continue); - } - Err(InvokeError::Python(error)) => Err(error), - } + let projected = self.host.project(py, arguments.bind(py)); + answered(projected.map(|projection| reply.send(projection))) } - HostOp::BeforeSend { wire, context } => { - match self.adapter.before_send(py, wire, &context) { - Ok(LifecycleStep::Wire(wire)) => Ok(HostResult::BeforeSend(wire)), + HostOp::Custom(op) => answered(self.host.invoke(py, op)), + HostOp::BeforeSend { + wire, + context, + reply, + } => match self.adapter.before_send(py, wire, &context) { + Ok(LifecycleStep::Wire(wire)) => { + reply.send(*wire); + Ok(Ok(())) + } + Ok(LifecycleStep::Await(awaitable)) => { + self.pending = Some(Pending::Adapter(Expect::Wire(reply))); + return Ok(Next::Return(ExecutionStep::Await(awaitable))); + } + Ok(_) => return Err(missing_state()), + Err(error) => Err(error), + }, + HostOp::Open(_, reply) => return self.opened(py, reply).map(Next::Return), + HostOp::Deliver(chunk, reply) => { + return self.delivered(py, chunk, reply).map(Next::Return); + } + HostOp::Emit(event, reply) => { + match self.adapter.emit(py, LifecycleEvent::Machine(&event)) { + Ok(LifecycleStep::Done) => { + reply.send(()); + Ok(Ok(())) + } Ok(LifecycleStep::Await(awaitable)) => { - self.pending = Some(Pending::Adapter(Expect::Wire)); + self.pending = Some(Pending::Adapter(Expect::Emitted(reply))); return Ok(Next::Return(ExecutionStep::Await(awaitable))); } Ok(_) => return Err(missing_state()), Err(error) => Err(error), } } - HostOp::Open(_) => return self.opened(py).map(Next::Return), - HostOp::Deliver(chunk) => return self.delivered(py, chunk).map(Next::Return), - HostOp::Emit(event) => match self.adapter.emit(py, LifecycleEvent::Machine(&event)) { - Ok(LifecycleStep::Done) => Ok(HostResult::Emitted), - Ok(LifecycleStep::Await(awaitable)) => { - self.pending = Some(Pending::Adapter(Expect::Emitted)); - return Ok(Next::Return(ExecutionStep::Await(awaitable))); - } - Ok(_) => return Err(missing_state()), - Err(error) => Err(error), - }, }; - match answer { - Ok(answer) => self.resume_core(py, Some(Ok(answer))).map(Next::Continue), + match answered { + Ok(Ok(())) => self.resume_core(py, None).map(Next::Continue), + Ok(Err(native)) => self + .resume_core(py, Some(HostFailure::Error(native))) + .map(Next::Continue), Err(error) => self.interrupt(py, error).map(Next::Return), } } - fn opened(&mut self, py: Python<'_>) -> PyResult { + fn opened(&mut self, py: Python<'_>, reply: Reply) -> PyResult { self.stage = Stage::Streaming; match self.adapter.opened(py) { Ok(()) => { - self.pending = Some(Pending::Consumer); + self.pending = Some(Pending::Consumer(reply)); Ok(ExecutionStep::Open) } Err(error) => self.interrupt(py, error), @@ -333,15 +354,16 @@ where fn delivered( &mut self, py: Python<'_>, - chunk: as Route>::Chunk, + chunk: as Protocol>::Chunk, + reply: Reply, ) -> PyResult { - let chunk = match self.route.chunk(py, chunk) { + let chunk = match self.host.chunk(py, chunk) { Ok(chunk) => chunk, Err(error) => return self.interrupt(py, error), }; match self.adapter.delivered(py, &chunk) { Ok(()) => { - self.pending = Some(Pending::Consumer); + self.pending = Some(Pending::Consumer(reply)); Ok(ExecutionStep::Yield(chunk)) } Err(error) => self.interrupt(py, error), @@ -357,25 +379,24 @@ where } else { HostFailure::Error(native) }; - self.resume_machine(py, Some(Err(failure))) + self.resume_machine(py, Some(failure)) } fn resume_core( &mut self, py: Python<'_>, - result: NativeResume, + interruption: Interruption, ) -> PyResult, Py>> { let state = Arc::clone(self.machine.as_ref().ok_or_else(missing_state)?); let future = async move { let mut state = state.lock().await; - let result = match result { - Some(Err(failure)) => state + let result = match interruption { + Some(failure) => state .machine .interrupt(failure) .await .map(MachineStep::Complete), - Some(Ok(result)) => state.machine.resume(Some(result)).await, - None => state.machine.resume(None).await, + None => state.machine.resume().await, }; state.result = Some(result); Ok(()) @@ -414,7 +435,7 @@ where fn completed(&mut self, py: Python<'_>, response: ResponseOf) -> PyResult { self.ended_at = Some(epoch_seconds()); - let public = match self.route.complete(py, response) { + let public = match self.host.complete(py, response) { Ok(public) => public, Err(error) => return self.failure(py, error, FailureOrigin::Call), }; @@ -441,7 +462,7 @@ where /// fails, that failure is raised with the native error's text as its `__context__`. fn classified(&self, py: Python<'_>, error: ErrorOf) -> PyErr { let native = error.to_string(); - let classifier_error = match self.route.classify(py, error) { + let classifier_error = match self.host.classify(py, error) { Ok(failure) => return failure.into(), Err(classifier_error) => classifier_error, }; @@ -486,7 +507,7 @@ where if self.machine.take().is_some() { Python::attach(|py| { self.adapter.close(py); - self.route.close(py); + self.host.close(py); }); } } @@ -494,15 +515,15 @@ where impl ExecutionBody for PythonDriver where - H: RouteHost, - M: Machine> + 'static, + H: ProtocolHost, + M: Machine> + 'static, { fn resume(&mut self, result: Option>>) -> PyResult { Python::attach(|py| self.drive(py, result)) } fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - self.route.traverse(visit)?; + self.host.traverse(visit)?; self.adapter.traverse(visit)?; visit.call(&self.arguments)?; visit.call(&self.interrupted)?; @@ -516,8 +537,8 @@ where impl Drop for PythonDriver where - H: RouteHost, - M: Machine> + 'static, + H: ProtocolHost, + M: Machine> + 'static, { fn drop(&mut self) { self.clear(); @@ -528,8 +549,8 @@ where mod tests { use std::sync::{Arc, Mutex}; - use litellm_host::event::{MachineEvent, RequestContext, WireRequest}; - use litellm_host::machine::{Interrupted, Step}; + use litellm_host::event::{MachineEvent, RawResponse, RequestContext}; + use litellm_host::machine::{CallMachine, MachineFault}; use pyo3::exceptions::{PyBaseException, PyValueError}; use pyo3::types::PyDict; @@ -573,22 +594,21 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } } - struct Synthetic; - - impl Route for Synthetic { - type Response = String; - type Error = Error; - type Op = &'static str; - type OpResult = String; - type Chunk = std::convert::Infallible; - type StreamHead = std::convert::Infallible; + impl From for Error { + fn from(fault: MachineFault) -> Self { + Self(format!("{fault:?}")) + } } - /// Yields the scripted ops in order, then completes or fails as scripted. - struct ScriptedMachine { - ops: Vec>, - outcome: Option>, - answers: Vec, + struct Synthetic; + + impl Protocol for Synthetic { + type Response = String; + type Error = Error; + type Projection = String; + type Op = (&'static str, Reply); + type Chunk = std::convert::Infallible; + type StreamHead = std::convert::Infallible; } fn wire() -> WireRequest { @@ -609,37 +629,6 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } } - impl Machine for ScriptedMachine { - type Route = Synthetic; - type Complete = String; - - fn resume(&mut self, result: Option>) -> Step<'_, Self> { - Box::pin(async move { - if let Some(result) = result { - self.answers.push(match result { - HostResult::Route(value) => value, - HostResult::BeforeSend(wire) => wire.url, - HostResult::Emitted => "emitted".into(), - HostResult::Demand(demand) => format!("{demand:?}"), - }); - } - if !self.ops.is_empty() { - return Ok(MachineStep::Host(self.ops.remove(0))); - } - self.outcome - .take() - .ok_or_else(|| Error("resumed after completion".into()))? - .map(MachineStep::Complete) - }) - } - - fn interrupt(&mut self, failure: HostFailure) -> Interrupted<'_, Self> { - self.ops.clear(); - self.outcome = None; - Box::pin(async move { Err(failure.into_error()) }) - } - } - #[derive(Default)] struct Log(Arc>>); @@ -677,22 +666,37 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } } - impl RouteHost for SyntheticHost { - type Route = Synthetic; + impl SyntheticHost { + fn answer(&self, value: impl FnOnce() -> String) -> Result> { + match self.op { + OpScript::Answer => Ok(value()), + OpScript::RaisePython => Err(PyValueError::new_err("op failed").into()), + OpScript::RejectNatively => Err(InvokeError::Native(Error("op rejected".into()))), + } + } + } + + impl ProtocolHost for SyntheticHost { + type Protocol = Synthetic; type Failure = Classified; + fn project( + &mut self, + _: Python<'_>, + arguments: &Bound<'_, PyDict>, + ) -> Result> { + self.log.push("project"); + self.answer(|| format!("project:{}", arguments.len())) + } + fn invoke( &mut self, _: Python<'_>, - arguments: &Bound<'_, PyDict>, - op: &'static str, - ) -> Result> { - self.log.push(format!("route:{op}")); - match self.op { - OpScript::Answer => Ok(format!("{op}:{}", arguments.len())), - OpScript::RaisePython => Err(PyValueError::new_err("op failed").into()), - OpScript::RejectNatively => Err(InvokeError::Native(Error("op rejected".into()))), - } + (op, reply): (&'static str, Reply), + ) -> Result<(), InvokeError> { + self.log.push(format!("op:{op}")); + self.answer(|| op.to_string()) + .map(|answer| reply.send(answer)) } fn chunk(&mut self, _: Python<'_>, chunk: std::convert::Infallible) -> PyResult> { @@ -719,7 +723,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } fn close(&mut self, _: Python<'_>) { - self.log.push("route.close"); + self.log.push("host.close"); } fn traverse(&self, _: &PyVisit<'_>) -> Result<(), PyTraverseError> { @@ -828,7 +832,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri fn run_scripted( py: Python<'_>, - machine: ScriptedMachine, + machine: CallMachine, op: OpScript, script: AdapterScript, asynchronous: bool, @@ -848,12 +852,12 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri fn run_hosted( py: Python<'_>, - machine: ScriptedMachine, - route: SyntheticHost, + machine: CallMachine, + host: SyntheticHost, script: AdapterScript, asynchronous: bool, ) -> (PyResult>, Vec) { - let log = Log(route.log.0.clone()); + let log = Log(host.log.0.clone()); let adapter = SyntheticAdapter { log: Log(log.0.clone()), script, @@ -863,7 +867,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri let result = run_call( py, machine, - route, + host, Box::new(adapter), arguments.unbind(), asynchronous, @@ -884,21 +888,21 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri (result, log.entries()) } - fn success_machine() -> ScriptedMachine { - ScriptedMachine { - ops: vec![ - HostOp::Route("project"), - HostOp::BeforeSend { - wire: Box::new(wire()), - context: Box::new(context()), - }, - HostOp::Emit(MachineEvent::ResponseReceived { - raw: litellm_host::event::RawResponse { body: "raw".into() }, - }), - ], - outcome: Some(Ok("done".into())), - answers: Vec::new(), - } + /// Answers to projection, to the route op and to `before_send` all reach the + /// response, so a driver that misroutes a reply changes what the call returns. + fn success_machine() -> CallMachine { + CallMachine::new(|host| { + Box::pin(async move { + let projected = host.project().await?; + let signed = host.custom_op(|reply| ("sign", reply)).await?; + let wire = host.before_send(wire(), context()).await?; + host.emit(MachineEvent::ResponseReceived { + raw: RawResponse { body: "raw".into() }, + }) + .await?; + Ok(format!("{projected}|{signed}|{}", wire.url)) + }) + }) } #[test] @@ -917,32 +921,37 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri AdapterScript::Plain, asynchronous, ); - assert_eq!(result.unwrap().extract::(py).unwrap(), "done"); + assert_eq!( + result.unwrap().extract::(py).unwrap(), + "project:1|sign|rewritten" + ); assert_eq!( log, [ "started", "begin", - "route:project", + "project", + "op:sign", "before_send", "response:raw", "complete", "after_success", - "succeeded:done", + "succeeded:project:1|sign|rewritten", "adapter.close", - "route.close", + "host.close", ] ); } }); } - fn failing_machine() -> ScriptedMachine { - ScriptedMachine { - ops: vec![HostOp::Route("project")], - outcome: Some(Err(Error("provider exploded".into()))), - answers: Vec::new(), - } + fn failing_machine() -> CallMachine { + CallMachine::new(|host| { + Box::pin(async move { + host.project().await?; + Err(Error("provider exploded".into())) + }) + }) } #[test] @@ -969,11 +978,11 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri [ "started", "begin", - "route:project", + "project", "classify:provider exploded", "failed:Call:classified: provider exploded", "adapter.close", - "route.close", + "host.close", ] ); } @@ -1003,11 +1012,11 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri [ "started", "begin", - "route:project", + "project", "classify:op rejected", "failed:Call:classified: op rejected", "adapter.close", - "route.close", + "host.close", ] ); }); @@ -1035,10 +1044,10 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri [ "started", "begin", - "route:project", + "project", "failed:Call:op failed", "adapter.close", - "route.close", + "host.close", ] ); }); @@ -1073,11 +1082,11 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri [ "started", "begin", - "route:project", + "project", "classify:provider exploded", "failed:Call:classifier failed", "adapter.close", - "route.close", + "host.close", ] ); }); @@ -1106,7 +1115,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri "begin", "failed:Host:begin failed", "adapter.close", - "route.close" + "host.close" ] ); }); @@ -1130,7 +1139,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri ); assert_eq!(result.unwrap().extract::(py).unwrap(), "replaced"); assert!(log.contains(&"succeeded:replaced".to_string())); - assert!(!log.contains(&"succeeded:done".to_string())); + assert!(!log.contains(&"succeeded:project:1|rewritten".to_string())); } }); } @@ -1159,7 +1168,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri "after_success", "failed:Host:after_success failed", "adapter.close", - "route.close" + "host.close" ] ); assert!(!log.iter().any(|entry| entry.starts_with("succeeded"))); @@ -1175,18 +1184,24 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri crate::initialize_python(); Python::attach(|py| { struct Cancelling(Log); - impl RouteHost for Cancelling { - type Route = Synthetic; + impl ProtocolHost for Cancelling { + type Protocol = Synthetic; type Failure = Classified; - fn invoke( + fn project( &mut self, _: Python<'_>, _: &Bound<'_, PyDict>, - _: &'static str, ) -> Result> { - self.0.push("route"); + self.0.push("project"); Err(pyo3::exceptions::asyncio::CancelledError::new_err(()).into()) } + fn invoke( + &mut self, + _: Python<'_>, + _: (&'static str, Reply), + ) -> Result<(), InvokeError> { + Err(missing_state().into()) + } fn chunk( &mut self, _: Python<'_>, @@ -1210,7 +1225,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } } let log = Log::default(); - let route = Cancelling(Log(log.0.clone())); + let host = Cancelling(Log(log.0.clone())); let adapter = SyntheticAdapter { log: Log(log.0.clone()), script: AdapterScript::Plain, @@ -1218,7 +1233,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri let error = run_call( py, success_machine(), - route, + host, Box::new(adapter), PyDict::new(py).unbind(), false, @@ -1227,7 +1242,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri assert!(!error.is_instance_of::(py)); assert_eq!( log.entries(), - ["started", "begin", "route", "adapter.close"] + ["started", "begin", "project", "adapter.close"] ); }); } diff --git a/litellm-rust/crates/host-python/src/file_reader.rs b/litellm-rust/crates/host-python/src/file_reader.rs new file mode 100644 index 00000000000..bbc7a233b28 --- /dev/null +++ b/litellm-rust/crates/host-python/src/file_reader.rs @@ -0,0 +1,241 @@ +//! A caller's file-like object: anything with a callable `read`, kept as a handle and read +//! once, on the host's thread, into bytes Rust owns. + +use bytes::Bytes; +use pyo3::{ + exceptions::PyTypeError, + gc::{PyTraverseError, PyVisit}, + prelude::*, + pybacked::PyBackedBytes, + types::{PyBytes, PyString}, +}; + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct FileContent { + pub bytes: Bytes, + pub file_name: Option, +} + +#[derive(Debug)] +pub struct PythonFileReader { + reader: Py, + name: Option, +} + +impl PythonFileReader { + /// `None` when `file` has no callable `read`. The object's `name` is read now, its + /// contents only on [`read`](Self::read). + pub fn from_file_like(file: &Bound<'_, PyAny>) -> PyResult> { + let reader = file + .getattr_opt("read")? + .filter(|value| value.is_callable()); + let Some(reader) = reader else { + return Ok(None); + }; + let name = file + .getattr_opt("name")? + .filter(|value| !value.is_none()) + .map(|value| value.extract::()) + .transpose()?; + Ok(Some(Self { + reader: reader.unbind(), + name, + })) + } + + pub fn read(&self, py: Python<'_>) -> PyResult { + let value = self.reader.bind(py).call0()?; + let bytes = if value.is_instance_of::() { + Bytes::from(value.extract::()?) + } else if value.is_instance_of::() { + py_bytes(&value)? + } else { + return Err(PyTypeError::new_err(format!( + "file read must return bytes or str, got {}", + value.get_type(), + ))); + }; + Ok(FileContent { + bytes, + file_name: self.name.clone(), + }) + } + + pub fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + visit.call(&self.reader) + } +} + +/// An exact `bytes` object is shared without copying and keeps the Python object alive; +/// a `bytes` subclass is copied. +pub fn py_bytes(value: &Bound<'_, PyAny>) -> PyResult { + if value.is_exact_instance_of::() { + return Ok(Bytes::from_owner(value.extract::()?)); + } + Ok(Bytes::copy_from_slice( + value.extract::()?.as_ref(), + )) +} + +#[cfg(test)] +mod tests { + use pyo3::{exceptions::PyTypeError, types::PyDict}; + + use super::*; + + fn eval<'py>(py: Python<'py>, source: &std::ffi::CStr) -> Bound<'py, PyDict> { + let locals = PyDict::new(py); + py.run(source, Some(&locals), Some(&locals)).unwrap(); + locals + } + + fn reader<'py>(locals: &Bound<'py, PyDict>, name: &str) -> PythonFileReader { + PythonFileReader::from_file_like(&locals.get_item(name).unwrap().unwrap()) + .unwrap() + .unwrap() + } + + #[test] + fn objects_without_a_callable_read_are_not_readers() { + Python::initialize(); + Python::attach(|py| { + let locals = eval( + py, + c" +class Attribute: + read = 'not callable' +plain = object() +attribute = Attribute() +", + ); + for name in ["plain", "attribute"] { + let file = locals.get_item(name).unwrap().unwrap(); + assert!(PythonFileReader::from_file_like(&file).unwrap().is_none()); + } + }); + } + + #[test] + fn the_name_is_taken_up_front_and_the_contents_only_on_read() { + Python::initialize(); + Python::attach(|py| { + let locals = eval( + py, + c" +class Reader: + name = 'scan.png' + def __init__(self): + self.reads = 0 + def read(self): + self.reads += 1 + return b'abc' +file = Reader() +", + ); + let reads = || { + locals + .get_item("file") + .unwrap() + .unwrap() + .getattr("reads") + .unwrap() + .extract::() + .unwrap() + }; + let file = reader(&locals, "file"); + assert_eq!(reads(), 0); + let content = file.read(py).unwrap(); + assert_eq!(reads(), 1); + assert_eq!( + content, + FileContent { + bytes: b"abc".as_slice().into(), + file_name: Some("scan.png".into()), + } + ); + }); + } + + #[test] + fn read_results_are_normalized_and_exceptions_keep_their_identity() { + Python::initialize(); + Python::attach(|py| { + let locals = eval( + py, + c" +failure = KeyError('reader failed') +class Raising: + def read(self): + raise failure +class Text: + def read(self): + return 'héllo' +class Wrong: + def read(self): + return 7 +raising = Raising() +text = Text() +wrong = Wrong() +", + ); + let error = reader(&locals, "raising").read(py).unwrap_err(); + assert!( + error + .value(py) + .is(locals.get_item("failure").unwrap().unwrap()) + ); + assert_eq!( + reader(&locals, "text").read(py).unwrap().bytes.as_ref(), + "héllo".as_bytes() + ); + let error = reader(&locals, "wrong").read(py).unwrap_err(); + assert!(error.is_instance_of::(py)); + assert!(error.to_string().contains("bytes or str")); + }); + } + + #[rstest::rstest] + #[case::read("read")] + #[case::name("name")] + fn attribute_failures_keep_their_identity(#[case] attribute: &str) { + Python::initialize(); + Python::attach(|py| { + let locals = eval( + py, + c" +failure = LookupError('file property failed') +class File: + def __getattribute__(self, name): + if name == attribute: + raise failure + return super().__getattribute__(name) + name = 'scan.pdf' + def read(self): + return b'abc' +file = File() +", + ); + locals.set_item("attribute", attribute).unwrap(); + let error = + PythonFileReader::from_file_like(&locals.get_item("file").unwrap().unwrap()) + .unwrap_err(); + assert!( + error + .value(py) + .is(locals.get_item("failure").unwrap().unwrap()) + ); + }); + } + + #[test] + fn exact_python_bytes_transfer_without_copying_and_outlive_the_input() { + Python::initialize(); + let (bytes, pointer) = Python::attach(|py| { + let value = PyBytes::new(py, b"document bytes"); + let pointer = value.as_bytes().as_ptr() as usize; + (py_bytes(value.as_any()).unwrap(), pointer) + }); + assert_eq!(bytes.as_ptr() as usize, pointer); + assert_eq!(bytes.as_ref(), b"document bytes"); + } +} diff --git a/litellm-rust/crates/host-python/src/lib.rs b/litellm-rust/crates/host-python/src/lib.rs index 55b27e34b46..7e17c4da51e 100644 --- a/litellm-rust/crates/host-python/src/lib.rs +++ b/litellm-rust/crates/host-python/src/lib.rs @@ -1,6 +1,6 @@ //! The CPython runtime adapter: value marshalling, interpreter detachment, the tokio and //! asyncio glue, and the driver that runs a native [`Machine`](litellm_host::machine::Machine) -//! against a Python route host and a Python lifecycle. Everything here is Python-specific by +//! against a Python protocol host and a Python lifecycle. Everything here is Python-specific by //! construction; another host language gets its own crate of the same shape. mod adapter; @@ -8,13 +8,14 @@ mod argument; mod callable; mod driver; mod execution; +mod file_reader; mod fork_gate; mod gil; mod handle; mod marshal; pub use adapter::{ - InvokeError, LifecycleEvent, LifecycleStep, PythonLifecycle, RouteHost, missing_state, + InvokeError, LifecycleEvent, LifecycleStep, ProtocolHost, PythonLifecycle, missing_state, }; pub use argument::lookup; pub use callable::wrap_failure; @@ -24,6 +25,7 @@ pub use execution::{ reserve_process_for_forking, run_async, run_async_value, run_sync, run_sync_value, runtime_started, }; +pub use file_reader::{FileContent, PythonFileReader, py_bytes}; pub use fork_gate::RuntimeAlreadyStarted; pub use gil::{PythonContext, attach_blocking, release_count, release_gil}; pub use handle::{Execution, ExecutionBody, ExecutionStep}; diff --git a/litellm-rust/crates/host/Cargo.toml b/litellm-rust/crates/host/Cargo.toml index 0c7c46192b5..bbbed68f345 100644 --- a/litellm-rust/crates/host/Cargo.toml +++ b/litellm-rust/crates/host/Cargo.toml @@ -7,6 +7,7 @@ repository.workspace = true [dependencies] litellm-auth.workspace = true +litellm-coroutine.workspace = true serde_json.workspace = true tokio = { workspace = true, features = ["sync"] } diff --git a/litellm-rust/crates/host/src/host.rs b/litellm-rust/crates/host/src/host.rs index aba35185a18..9714b9470a3 100644 --- a/litellm-rust/crates/host/src/host.rs +++ b/litellm-rust/crates/host/src/host.rs @@ -1,28 +1,27 @@ use std::future::Future; -use crate::event::{CallEvent, MachineEvent, RequestContext, WireRequest}; -use crate::route::Route; +pub use litellm_coroutine::{Abandoned, Answer, Reply, reply}; -/// One suspension point of a native call, performed by the host. -pub enum HostOp { - Route(R::Op), +use crate::event::{CallEvent, MachineEvent, RequestContext, WireRequest}; +use crate::protocol::Protocol; + +/// One suspension point of a native call, performed by the host and answered through the +/// [`Reply`] it carries. +pub enum HostOp { + /// The first op of every call: the caller's request as the host projects it. + Project(Reply), + Custom(R::Op), BeforeSend { wire: Box, context: Box, + reply: Reply, }, - Emit(MachineEvent), + Emit(MachineEvent, Reply<()>), /// The response streams: the host hands the caller a stream and answers once the /// caller asks for the first chunk or goes away. - Open(R::StreamHead), + Open(R::StreamHead, Reply), /// The next chunk of an open stream, answered once the caller asks for the one after. - Deliver(R::Chunk), -} - -pub enum HostResult { - Route(R::OpResult), - BeforeSend(Box), - Emitted, - Demand(Demand), + Deliver(R::Chunk, Reply), } /// Whether the caller of a streamed call still reads it. @@ -39,10 +38,13 @@ pub enum HostStep { Suspend(S), } -/// An in-process host: answers route operations and observes the call without leaving +/// An in-process host: answers custom operations and observes the call without leaving /// the Rust runtime. Language hosts implement their own driver instead. -pub trait Host: Send + Sync { - fn route(&self, op: R::Op) -> impl Future> + Send; +pub trait Host: Send + Sync { + fn project(&self) -> impl Future> + Send; + + /// Answers `op` through its reply, or fails the call. + fn custom_op(&self, op: R::Op) -> impl Future> + Send; fn before_send( &self, diff --git a/litellm-rust/crates/host/src/lib.rs b/litellm-rust/crates/host/src/lib.rs index 65479c2380f..c6b9e59b65a 100644 --- a/litellm-rust/crates/host/src/lib.rs +++ b/litellm-rust/crates/host/src/lib.rs @@ -1,12 +1,13 @@ //! The contract between a native call and the host runtime that drives it. //! //! A host is whatever sits on the far side of the language boundary: CPython today, -//! another runtime later. Core runs each route on a [`machine::RouteMachine`] and never learns +//! another runtime later. Core runs each route on a [`machine::CallMachine`] and never learns //! which host is on the other end. The machine yields [`host::HostOp`]s; a driver answers -//! them, observes [`event::CallEvent`]s and may rewrite the wire request before it is sent. +//! each through the typed [`host::Reply`] it carries, observes [`event::CallEvent`]s and +//! may rewrite the wire request before it is sent. pub mod event; pub mod host; pub mod machine; -pub mod route; +pub mod protocol; pub mod run; diff --git a/litellm-rust/crates/host/src/machine/auth.rs b/litellm-rust/crates/host/src/machine/auth.rs index ba7e242e766..76e3504ca28 100644 --- a/litellm-rust/crates/host/src/machine/auth.rs +++ b/litellm-rust/crates/host/src/machine/auth.rs @@ -1,22 +1,21 @@ use std::sync::Arc; use super::{HostChannel, MachineFault}; -use crate::route::Route; +use crate::{host::Reply, protocol::Protocol}; use litellm_auth::{Error, ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle}; -/// A route whose host can mint credentials on the call's behalf. -pub trait TokenRoute: Route { - fn acquire_token_op() -> Self::Op; - fn token_credential(result: Self::OpResult) -> Option; +/// A protocol whose host can mint credentials on the call's behalf. +pub trait TokenProtocol: Protocol { + fn acquire_token_op(reply: Reply) -> Self::Op; } /// A [`TokenProvider`] that asks the host for each credential through the call's own /// operation channel, so the host answers it on the caller's thread and context. -pub struct HostTokenProvider { +pub struct HostTokenProvider { channel: HostChannel, } -impl std::fmt::Debug for HostTokenProvider { +impl std::fmt::Debug for HostTokenProvider { fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { formatter.write_str("HostTokenProvider") } @@ -24,7 +23,7 @@ impl std::fmt::Debug for HostTokenProvider { impl HostTokenProvider where - R: TokenRoute, + R: TokenProtocol, R::Error: From + std::fmt::Display, { pub fn handle(channel: HostChannel) -> TokenProviderHandle { @@ -34,19 +33,15 @@ where impl TokenProvider for HostTokenProvider where - R: TokenRoute, + R: TokenProtocol, R::Error: From + std::fmt::Display, { fn acquire(&self) -> TokenFuture<'_> { Box::pin(async move { - let result = self - .channel - .route(R::acquire_token_op()) + self.channel + .custom_op(R::acquire_token_op) .await - .map_err(|error| Error::AzureTokenAcquisition(error.to_string()))?; - R::token_credential(result).ok_or_else(|| { - Error::AzureTokenAcquisition("invalid token provider host result".into()) - }) + .map_err(|error| Error::AzureTokenAcquisition(error.to_string())) }) } } diff --git a/litellm-rust/crates/host/src/machine/call_machine.rs b/litellm-rust/crates/host/src/machine/call_machine.rs new file mode 100644 index 00000000000..af0bc50fbe6 --- /dev/null +++ b/litellm-rust/crates/host/src/machine/call_machine.rs @@ -0,0 +1,137 @@ +//! The one machine every route runs on: the route's provider future as a +//! [`Coroutine`] that yields [`HostOp`]s, each answered through its own typed reply. No +//! task is spawned; dropping the machine drops the in-flight call. + +use std::{future::Future, pin::Pin}; + +use litellm_coroutine::{Co, Coroutine, CoroutineState, ResumeError}; + +use super::{HostFailure, Interrupted, Machine, MachineStep, Step}; +use crate::{ + event::{MachineEvent, RequestContext, WireRequest}, + host::{Demand, HostOp, Reply}, + protocol::Protocol, +}; + +/// The machine's own failures, distinct from anything the provider call reports. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum MachineFault { + /// The host dropped an op's reply unanswered, or went away while the call waited. + Abandoned, + /// The host resumed the call out of turn. + Protocol(ResumeError), +} + +pub type ExecuteFuture = + Pin::Response, ::Error>> + Send>>; + +/// The provider side of the machine: how the in-flight call reaches its host. +pub struct HostChannel { + co: Co>, +} + +impl Clone for HostChannel { + fn clone(&self) -> Self { + Self { + co: self.co.clone(), + } + } +} + +impl HostChannel +where + R::Error: From, +{ + async fn yield_( + &self, + ask: impl FnOnce(Reply) -> HostOp + Send, + ) -> Result { + self.co + .yield_(ask) + .await + .map_err(|_| MachineFault::Abandoned.into()) + } + + pub async fn project(&self) -> Result { + self.yield_(HostOp::Project).await + } + + /// Asks the host to perform the custom operation `ask` builds around its reply, as in + /// `host.custom_op(OcrOp::AcquireAzureAdToken)`. + pub async fn custom_op( + &self, + ask: impl FnOnce(Reply) -> R::Op + Send, + ) -> Result { + self.yield_(|reply| HostOp::Custom(ask(reply))).await + } + + pub async fn before_send( + &self, + wire: WireRequest, + context: RequestContext, + ) -> Result { + self.yield_(|reply| HostOp::BeforeSend { + wire: Box::new(wire), + context: Box::new(context), + reply, + }) + .await + } + + pub async fn emit(&self, event: MachineEvent) -> Result<(), R::Error> { + self.yield_(|reply| HostOp::Emit(event, reply)).await + } + + pub async fn open(&self, head: R::StreamHead) -> Result { + self.yield_(|reply| HostOp::Open(head, reply)).await + } + + pub async fn deliver(&self, chunk: R::Chunk) -> Result { + self.yield_(|reply| HostOp::Deliver(chunk, reply)).await + } +} + +type CallCoroutine = + Coroutine, Result<::Response, ::Error>>; + +pub struct CallMachine { + coroutine: CallCoroutine, +} + +impl CallMachine +where + R::Error: From, +{ + pub fn new(execute: impl FnOnce(HostChannel) -> ExecuteFuture + Send + 'static) -> Self { + Self { + coroutine: Coroutine::new(|co| execute(HostChannel { co })), + } + } +} + +impl Machine for CallMachine +where + R::Error: From, +{ + type Protocol = R; + type Complete = R::Response; + + fn resume(&mut self) -> Step<'_, Self> { + Box::pin(async move { + match self + .coroutine + .resume() + .await + .map_err(MachineFault::Protocol)? + { + CoroutineState::Yielded(op) => Ok(MachineStep::Host(op)), + CoroutineState::Complete(outcome) => outcome.map(MachineStep::Complete), + } + }) + } + + fn interrupt(&mut self, failure: HostFailure) -> Interrupted<'_, Self> { + self.coroutine.cancel(); + Box::pin(async move { Err(failure.into_error()) }) + } +} diff --git a/litellm-rust/crates/host/src/machine/mod.rs b/litellm-rust/crates/host/src/machine/mod.rs index 2c26db61582..0c7501633fa 100644 --- a/litellm-rust/crates/host/src/machine/mod.rs +++ b/litellm-rust/crates/host/src/machine/mod.rs @@ -1,16 +1,16 @@ mod auth; -mod route_machine; +mod call_machine; use std::future::Future; use std::pin::Pin; -pub use auth::{HostTokenProvider, TokenRoute}; -pub use route_machine::{ExecuteFuture, HostChannel, MachineFault, RouteMachine}; +pub use auth::{HostTokenProvider, TokenProtocol}; +pub use call_machine::{CallMachine, ExecuteFuture, HostChannel, MachineFault}; -use crate::host::{HostOp, HostResult}; -use crate::route::Route; +use crate::host::HostOp; +use crate::protocol::Protocol; -pub enum MachineStep { +pub enum MachineStep { Host(HostOp), Complete(C), } @@ -19,8 +19,8 @@ pub type Step<'a, M> = Pin< Box< dyn Future< Output = Result< - MachineStep<::Route, ::Complete>, - <::Route as Route>::Error, + MachineStep<::Protocol, ::Complete>, + <::Protocol as Protocol>::Error, >, > + Send + 'a, @@ -30,7 +30,10 @@ pub type Step<'a, M> = Pin< pub type Interrupted<'a, M> = Pin< Box< dyn Future< - Output = Result<::Complete, <::Route as Route>::Error>, + Output = Result< + ::Complete, + <::Protocol as Protocol>::Error, + >, > + Send + 'a, >, @@ -51,19 +54,18 @@ impl HostFailure { } /// A resumable call. Core implements it per route; a host drives it. Every suspension -/// point is an op the host performs and answers with a result. +/// point is an op the host performs and answers through the op's own reply before it +/// resumes the call again. pub trait Machine: Send { - type Route: Route; + type Protocol: Protocol; type Complete: Send + 'static; - /// `None` on the first call and whenever the previous step completed without - /// yielding an op; otherwise the result of the op last yielded. - fn resume(&mut self, result: Option>) -> Step<'_, Self>; + fn resume(&mut self) -> Step<'_, Self>; /// The host failed to perform the pending op, or the caller cancelled. The call /// yields no further ops. fn interrupt( &mut self, - failure: HostFailure<::Error>, + failure: HostFailure<::Error>, ) -> Interrupted<'_, Self>; } diff --git a/litellm-rust/crates/host/src/machine/route_machine.rs b/litellm-rust/crates/host/src/machine/route_machine.rs deleted file mode 100644 index 38a0b8bc16a..00000000000 --- a/litellm-rust/crates/host/src/machine/route_machine.rs +++ /dev/null @@ -1,199 +0,0 @@ -//! The one machine every route runs on: it owns the route's provider future, polls it in -//! place, and turns the host operations that future requests into [`Machine`] steps. No -//! task is spawned; dropping the machine drops the in-flight call. - -use std::{future::Future, pin::Pin}; - -use tokio::sync::{mpsc, oneshot}; - -use super::{HostFailure, Interrupted, Machine, MachineStep, Step}; -use crate::{ - event::{MachineEvent, RequestContext, WireRequest}, - host::{Demand, HostOp, HostResult}, - route::Route, -}; - -/// The machine's own failures, distinct from anything the provider call reports. -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum MachineFault { - /// The host driver went away while the call was waiting on it. - Abandoned, - /// The host answered out of turn: a result with nothing pending, or nothing when a - /// result was pending. - Protocol(&'static str), - /// The host answered a route operation with the wrong result variant. - Mismatch, -} - -pub type ExecuteFuture = - Pin::Response, ::Error>> + Send>>; - -struct PendingOp { - op: HostOp, - reply: oneshot::Sender>, -} - -/// The provider side of the machine: how the in-flight call reaches its host. -pub struct HostChannel { - ops: mpsc::UnboundedSender>, -} - -impl Clone for HostChannel { - fn clone(&self) -> Self { - Self { - ops: self.ops.clone(), - } - } -} - -impl HostChannel -where - R::Error: From, -{ - async fn invoke(&self, op: HostOp) -> Result, R::Error> { - let (reply, answer) = oneshot::channel(); - self.ops - .send(PendingOp { op, reply }) - .map_err(|_| MachineFault::Abandoned)?; - answer.await.map_err(|_| MachineFault::Abandoned.into()) - } - - pub async fn route(&self, op: R::Op) -> Result { - match self.invoke(HostOp::Route(op)).await? { - HostResult::Route(result) => Ok(result), - _ => Err(MachineFault::Mismatch.into()), - } - } - - pub async fn before_send( - &self, - wire: WireRequest, - context: RequestContext, - ) -> Result { - let op = HostOp::BeforeSend { - wire: Box::new(wire), - context: Box::new(context), - }; - match self.invoke(op).await? { - HostResult::BeforeSend(wire) => Ok(*wire), - _ => Err(MachineFault::Mismatch.into()), - } - } - - pub async fn emit(&self, event: MachineEvent) -> Result<(), R::Error> { - match self.invoke(HostOp::Emit(event)).await? { - HostResult::Emitted => Ok(()), - _ => Err(MachineFault::Mismatch.into()), - } - } - - pub async fn open(&self, head: R::StreamHead) -> Result { - self.demand(HostOp::Open(head)).await - } - - pub async fn deliver(&self, chunk: R::Chunk) -> Result { - self.demand(HostOp::Deliver(chunk)).await - } - - async fn demand(&self, op: HostOp) -> Result { - match self.invoke(op).await? { - HostResult::Demand(demand) => Ok(demand), - _ => Err(MachineFault::Mismatch.into()), - } - } -} - -enum Execution { - Unstarted(Box) -> ExecuteFuture + Send>), - Running(ExecuteFuture), - Done, -} - -pub struct RouteMachine { - execution: Execution, - ops: mpsc::UnboundedReceiver>, - channel: HostChannel, - reply: Option>>, -} - -impl RouteMachine -where - R::Error: From, -{ - pub fn new(execute: impl FnOnce(HostChannel) -> ExecuteFuture + Send + 'static) -> Self { - let (ops_tx, ops) = mpsc::unbounded_channel(); - Self { - execution: Execution::Unstarted(Box::new(execute)), - ops, - channel: HostChannel { ops: ops_tx }, - reply: None, - } - } - - async fn step( - &mut self, - result: Option>, - ) -> Result, R::Error> { - match (self.reply.take(), result) { - (Some(reply), Some(result)) => { - reply - .send(result) - .map_err(|_| MachineFault::Protocol("the call stopped waiting on the host"))?; - } - (None, None) if matches!(self.execution, Execution::Unstarted(_)) => {} - (Some(reply), None) => { - self.reply = Some(reply); - return Err(MachineFault::Protocol("host operation result is required").into()); - } - (None, Some(_)) => { - return Err(MachineFault::Protocol("unexpected host operation result").into()); - } - (None, None) => { - return Err( - MachineFault::Protocol("call cannot be resumed after completion").into(), - ); - } - } - if let Execution::Unstarted(_) = self.execution { - let Execution::Unstarted(start) = - std::mem::replace(&mut self.execution, Execution::Done) - else { - unreachable!() - }; - self.execution = Execution::Running(start(self.channel.clone())); - } - let Execution::Running(future) = &mut self.execution else { - return Err(MachineFault::Protocol("call cannot be resumed after completion").into()); - }; - tokio::select! { - biased; - pending = self.ops.recv() => { - let pending = pending.ok_or(MachineFault::Abandoned)?; - self.reply = Some(pending.reply); - Ok(MachineStep::Host(pending.op)) - } - outcome = future => { - self.execution = Execution::Done; - outcome.map(MachineStep::Complete) - } - } - } -} - -impl Machine for RouteMachine -where - R::Error: From, -{ - type Route = R; - type Complete = R::Response; - - fn resume(&mut self, result: Option>) -> Step<'_, Self> { - Box::pin(self.step(result)) - } - - fn interrupt(&mut self, failure: HostFailure) -> Interrupted<'_, Self> { - self.reply = None; - self.execution = Execution::Done; - Box::pin(async move { Err(failure.into_error()) }) - } -} diff --git a/litellm-rust/crates/host/src/protocol.rs b/litellm-rust/crates/host/src/protocol.rs new file mode 100644 index 00000000000..a7c0f3470b2 --- /dev/null +++ b/litellm-rust/crates/host/src/protocol.rs @@ -0,0 +1,17 @@ +/// One public call surface: what a completed call produces, how it fails, what the host +/// projects the caller's request into, and the protocol-specific operations only its host +/// can perform mid-call (token acquisition, for one). +pub trait Protocol: Send + Sync + 'static { + type Response: Send + 'static; + type Error: Clone + Send + Sync + 'static; + /// The caller's request as the host projects it, answered once before anything else. + type Projection: Send + 'static; + /// Each operation carries the [`Reply`](crate::host::Reply) its answer goes through. + /// A protocol with no operations of its own uses `Infallible`. + type Op: Send + 'static; + /// One piece of a streamed response, handed to the caller as it arrives. A protocol + /// that never streams uses `Infallible`. + type Chunk: Send + 'static; + /// What the call knows once a streamed response starts, before its first chunk. + type StreamHead: Send + 'static; +} diff --git a/litellm-rust/crates/host/src/route.rs b/litellm-rust/crates/host/src/route.rs deleted file mode 100644 index 8ab2b125760..00000000000 --- a/litellm-rust/crates/host/src/route.rs +++ /dev/null @@ -1,14 +0,0 @@ -/// One public call surface: what a completed call produces, how it fails, and the -/// route-specific operations only its host can perform (request projection, file reads, -/// token acquisition). -pub trait Route: Send + Sync + 'static { - type Response: Send + 'static; - type Error: Clone + Send + Sync + 'static; - type Op: Send + 'static; - type OpResult: Send + 'static; - /// One piece of a streamed response, handed to the caller as it arrives. A route - /// that never streams uses `Infallible`. - type Chunk: Send + 'static; - /// What the route knows once a streamed response starts, before its first chunk. - type StreamHead: Send + 'static; -} diff --git a/litellm-rust/crates/host/src/run.rs b/litellm-rust/crates/host/src/run.rs index 6a0c08fba68..baa3b58e058 100644 --- a/litellm-rust/crates/host/src/run.rs +++ b/litellm-rust/crates/host/src/run.rs @@ -1,40 +1,28 @@ use crate::event::{CallEvent, FailureOrigin, Timing, epoch_seconds}; -use crate::host::{Host, HostOp, HostResult}; +use crate::host::{Host, HostOp}; use crate::machine::{HostFailure, Machine, MachineStep}; -use crate::route::Route; +use crate::protocol::Protocol; /// Drives a machine to completion against an in-process host and emits exactly one /// terminal event. -pub async fn run(mut machine: M, host: &H) -> Result::Error> +pub async fn run( + mut machine: M, + host: &H, +) -> Result::Error> where M: Machine, - H: Host, + H: Host, { let start_time = epoch_seconds(); let _ = host.emit(&CallEvent::Started { start_time }).await; - let mut result = None; let outcome = loop { - let step = match machine.resume(result.take()).await { + let op = match machine.resume().await { Ok(MachineStep::Complete(complete)) => break Ok(complete), Ok(MachineStep::Host(op)) => op, Err(error) => break Err(error), }; - let answer = match step { - HostOp::Route(op) => host.route(op).await.map(HostResult::Route), - HostOp::BeforeSend { wire, context } => host - .before_send(*wire, &context) - .await - .map(|wire| HostResult::BeforeSend(Box::new(wire))), - HostOp::Emit(event) => host - .emit(&CallEvent::Machine(event)) - .await - .map(|()| HostResult::Emitted), - HostOp::Open(head) => host.open(head).await.map(HostResult::Demand), - HostOp::Deliver(chunk) => host.deliver(chunk).await.map(HostResult::Demand), - }; - match answer { - Ok(answer) => result = Some(answer), - Err(error) => break machine.interrupt(HostFailure::Error(error)).await, + if let Err(error) = perform(host, op).await { + break machine.interrupt(HostFailure::Error(error)).await; } }; let timing = Timing { @@ -52,44 +40,52 @@ where outcome } +async fn perform>(host: &H, op: HostOp) -> Result<(), R::Error> { + match op { + HostOp::Project(reply) => host + .project() + .await + .map(|projection| reply.send(projection)), + HostOp::Custom(op) => host.custom_op(op).await, + HostOp::BeforeSend { + wire, + context, + reply, + } => host + .before_send(*wire, &context) + .await + .map(|wire| reply.send(wire)), + HostOp::Emit(event, reply) => host + .emit(&CallEvent::Machine(event)) + .await + .map(|()| reply.send(())), + HostOp::Open(head, reply) => host.open(head).await.map(|demand| reply.send(demand)), + HostOp::Deliver(chunk, reply) => host.deliver(chunk).await.map(|demand| reply.send(demand)), + } +} + #[cfg(test)] mod tests { use std::sync::Mutex; use super::*; - use crate::machine::{Interrupted, Step}; + use crate::host::Reply; + use crate::machine::{CallMachine, MachineFault}; struct Unit; - impl Route for Unit { + impl Protocol for Unit { type Response = (); type Error = &'static str; - type Op = &'static str; - type OpResult = (); + type Projection = (); + type Op = (&'static str, Reply<()>); type Chunk = std::convert::Infallible; type StreamHead = std::convert::Infallible; } - struct Scripted { - ops: Vec<&'static str>, - outcome: Result<(), &'static str>, - } - - impl Machine for Scripted { - type Route = Unit; - type Complete = (); - - fn resume(&mut self, _: Option>) -> Step<'_, Self> { - Box::pin(async move { - if !self.ops.is_empty() { - return Ok(MachineStep::Host(HostOp::Route(self.ops.remove(0)))); - } - self.outcome.map(MachineStep::Complete) - }) - } - - fn interrupt(&mut self, failure: HostFailure<&'static str>) -> Interrupted<'_, Self> { - Box::pin(async move { Err(failure.into_error()) }) + impl From for &'static str { + fn from(_: MachineFault) -> Self { + "machine fault" } } @@ -100,12 +96,21 @@ mod tests { } impl Host for Recording { - async fn route(&self, op: &'static str) -> Result<(), &'static str> { - self.seen.lock().unwrap().push(format!("route:{op}")); - match self.fail { - Some(failing) if failing == op => Err("host failed"), - _ => Ok(()), + async fn project(&self) -> Result<(), &'static str> { + self.seen.lock().unwrap().push("project".into()); + Ok(()) + } + + async fn custom_op( + &self, + (op, reply): (&'static str, Reply<()>), + ) -> Result<(), &'static str> { + self.seen.lock().unwrap().push(format!("op:{op}")); + if self.fail == Some(op) { + return Err("host failed"); } + reply.send(()); + Ok(()) } async fn emit(&self, event: &CallEvent) -> Result<(), &'static str> { @@ -119,21 +124,29 @@ mod tests { } } - fn scripted(ops: &[&'static str], outcome: Result<(), &'static str>) -> Scripted { - Scripted { - ops: ops.to_vec(), - outcome, - } + fn scripted( + ops: &'static [&'static str], + outcome: Result<(), &'static str>, + ) -> CallMachine { + CallMachine::new(move |host| { + Box::pin(async move { + host.project().await?; + for op in ops { + host.custom_op(|reply| (*op, reply)).await?; + } + outcome + }) + }) } #[tokio::test] async fn forwards_every_op_then_emits_one_succeeded() { let host = Recording::default(); - let outcome = run(scripted(&["project", "send"], Ok(())), &host).await; + let outcome = run(scripted(&["sign", "send"], Ok(())), &host).await; assert_eq!(outcome, Ok(())); assert_eq!( *host.seen.lock().unwrap(), - ["started", "route:project", "route:send", "succeeded"] + ["started", "project", "op:sign", "op:send", "succeeded"] ); } @@ -142,24 +155,32 @@ mod tests { let host = Recording::default(); let outcome = run(scripted(&[], Err("boom")), &host).await; assert_eq!(outcome, Err("boom")); - assert_eq!(*host.seen.lock().unwrap(), ["started", "failed"]); + assert_eq!(*host.seen.lock().unwrap(), ["started", "project", "failed"]); let host = Recording { fail: Some("send"), ..Recording::default() }; - let outcome = run(scripted(&["project", "send", "never"], Ok(())), &host).await; + let outcome = run(scripted(&["sign", "send", "never"], Ok(())), &host).await; assert_eq!(outcome, Err("host failed")); assert_eq!( *host.seen.lock().unwrap(), - ["started", "route:project", "route:send", "failed"] + ["started", "project", "op:sign", "op:send", "failed"] ); } struct StartTimes(Mutex>); impl Host for StartTimes { - async fn route(&self, _: &'static str) -> Result<(), &'static str> { + async fn project(&self) -> Result<(), &'static str> { + Ok(()) + } + + async fn custom_op( + &self, + (_, reply): (&'static str, Reply<()>), + ) -> Result<(), &'static str> { + reply.send(()); Ok(()) } @@ -178,7 +199,7 @@ mod tests { #[tokio::test] async fn started_opens_the_call_at_the_terminal_start_time_and_cannot_fail_it() { let host = StartTimes(Mutex::default()); - assert_eq!(run(scripted(&["project"], Ok(())), &host).await, Ok(())); + assert_eq!(run(scripted(&["send"], Ok(())), &host).await, Ok(())); let times = host.0.lock().unwrap(); assert_eq!(times.len(), 2); assert_eq!(times[0], times[1]); diff --git a/litellm-rust/crates/llms/src/base_llm/ocr/error.rs b/litellm-rust/crates/llms/src/base_llm/ocr/error.rs index b3df8fc18c8..5d7be8b10df 100644 --- a/litellm-rust/crates/llms/src/base_llm/ocr/error.rs +++ b/litellm-rust/crates/llms/src/base_llm/ocr/error.rs @@ -118,7 +118,6 @@ impl From for Error { Self::InvalidRequest(match fault { MachineFault::Abandoned => "OCR host driver was abandoned".into(), MachineFault::Protocol(message) => format!("OCR {message}"), - MachineFault::Mismatch => "invalid OCR host operation result".into(), }) } } diff --git a/litellm-rust/crates/python-bridge/src/logger/machine.rs b/litellm-rust/crates/python-bridge/src/logger/machine.rs index 7234308e67e..54f8f3d4b3f 100644 --- a/litellm-rust/crates/python-bridge/src/logger/machine.rs +++ b/litellm-rust/crates/python-bridge/src/logger/machine.rs @@ -1,9 +1,8 @@ use std::sync::OnceLock; use litellm_host::{ - host::HostResult, machine::{HostFailure, Interrupted, Machine, Step}, - route::Route, + protocol::Protocol, }; use litellm_tracing::Logger; use pyo3::Python; @@ -23,17 +22,17 @@ impl LoggedMachine { } impl Machine for LoggedMachine { - type Route = M::Route; + type Protocol = M::Protocol; type Complete = M::Complete; - fn resume(&mut self, result: Option>) -> Step<'_, Self> { + fn resume(&mut self) -> Step<'_, Self> { let logger = self.logger.get_or_init(|| Python::attach(super::capture)); - Box::pin(logger.instrument(logger.scope(|| self.machine.resume(result)))) + Box::pin(logger.instrument(logger.scope(|| self.machine.resume()))) } fn interrupt( &mut self, - failure: HostFailure<::Error>, + failure: HostFailure<::Error>, ) -> Interrupted<'_, Self> { let logger = self.logger.get_or_init(|| Python::attach(super::capture)); Box::pin(logger.instrument(logger.scope(|| self.machine.interrupt(failure)))) diff --git a/litellm-rust/crates/python-bridge/src/logger/tests.rs b/litellm-rust/crates/python-bridge/src/logger/tests.rs index 9312d4c187c..b65e7d37023 100644 --- a/litellm-rust/crates/python-bridge/src/logger/tests.rs +++ b/litellm-rust/crates/python-bridge/src/logger/tests.rs @@ -1,29 +1,28 @@ use std::{process::Command, task::Poll}; use litellm_host::{ - host::HostResult, machine::{HostFailure, Interrupted, Machine, MachineStep, Step}, - route::Route, + protocol::Protocol, }; use pyo3::{prelude::*, types::PyDict}; struct DiagnosticMachine; -impl Route for DiagnosticMachine { +impl Protocol for DiagnosticMachine { type Response = (); type Error = String; + type Projection = (); type Op = (); - type OpResult = (); type Chunk = (); type StreamHead = (); } impl Machine for DiagnosticMachine { - type Route = Self; + type Protocol = Self; type Complete = (); - fn resume(&mut self, _: Option>) -> Step<'_, Self> { + fn resume(&mut self) -> Step<'_, Self> { litellm_tracing::warn!("machine started"); Box::pin(async { tokio::task::yield_now().await; @@ -45,7 +44,7 @@ fn machine_warning(py: Python<'_>) -> PyResult> { let mut machine = super::LoggedMachine::new(DiagnosticMachine); let mut future = Box::pin(async move { machine - .resume(None) + .resume() .await .map_err(pyo3::exceptions::PyValueError::new_err)?; machine diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs index 9d97094aeda..6de4e1320e1 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs @@ -1,10 +1,12 @@ +use std::convert::Infallible; + use bytes::Bytes; use litellm_core::messages::{ Error, - route::{Messages, MessagesCall, MessagesOp, MessagesOpResult, MessagesOutput}, + route::{Messages, MessagesCall, MessagesOutput}, types::MessagesShaping, }; -use litellm_host_python::{InvokeError, RouteHost, from_py, lookup, to_py}; +use litellm_host_python::{InvokeError, ProtocolHost, from_py, lookup, to_py}; use litellm_http::transport::Error as TransportError; use litellm_types::utils::ProviderSpecificHeaders; use pyo3::{ @@ -80,16 +82,16 @@ fn native_error(py: Python<'_>, error: Error) -> PyResult { /// The Python side of the Messages route: projects the prepared arguments and builds the /// public response, chunks and exceptions. -pub(super) struct MessagesRouteHost { +pub(super) struct MessagesPythonHost { request: Py, } -impl MessagesRouteHost { +impl MessagesPythonHost { pub(super) fn new(request: Py) -> Self { Self { request } } - fn project(&self, py: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult { + fn projection(&self, py: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult { let request = self.request.bind(py); let argument = |name: &str| -> PyResult>> { Ok(lookup(arguments, request, name)?.filter(|value| !value.is_none())) @@ -208,22 +210,21 @@ impl MessagesRouteHost { } } -impl RouteHost for MessagesRouteHost { - type Route = Messages; +impl ProtocolHost for MessagesPythonHost { + type Protocol = Messages; type Failure = PyErr; - fn invoke( + fn project( &mut self, py: Python<'_>, arguments: &Bound<'_, PyDict>, - op: MessagesOp, - ) -> Result> { - match op { - MessagesOp::ProjectRequest => self - .project(py, arguments) - .map(|call| MessagesOpResult::Request(Box::new(call))) - .map_err(|error| InvokeError::Python(self.map_failure(py, error))), - } + ) -> Result> { + self.projection(py, arguments) + .map_err(|error| InvokeError::Python(self.map_failure(py, error))) + } + + fn invoke(&mut self, _: Python<'_>, op: Infallible) -> Result<(), InvokeError> { + match op {} } fn complete(&mut self, py: Python<'_>, response: MessagesOutput) -> PyResult> { diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs index fd474e6b2d4..65040f31684 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -1,6 +1,6 @@ mod host; -use host::MessagesRouteHost; +use host::MessagesPythonHost; use litellm_callbacks_legacy_python::{ LegacySurface, PassThroughStream, PublicCall, run_legacy_call, }; @@ -45,7 +45,7 @@ fn run_messages( SURFACE, PublicCall::capture(&request, &args, &kwargs)?, crate::logger::LoggedMachine::new(messages_machine(secrets)), - MessagesRouteHost::new(request.unbind()), + MessagesPythonHost::new(request.unbind()), asynchronous, ) } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs index a928e62d5b7..e6821241c89 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs @@ -1,58 +1,38 @@ use std::path::PathBuf; -use bytes::Bytes; -use litellm_core::ocr::types::{OcrDocumentInput, OcrFileContent}; +use litellm_core::ocr::types::OcrDocumentInput; +use litellm_host_python::{PythonFileReader, py_bytes}; use pyo3::{ - exceptions::{PyTypeError, PyValueError}, - gc::{PyTraverseError, PyVisit}, + exceptions::PyValueError, prelude::*, - pybacked::PyBackedBytes, sync::PyOnceLock, types::{PyBytes, PyString, PyType}, }; -#[derive(Debug)] -pub(super) struct PythonFileReader { - reader: Py, - name: Option, +/// A `type='file'` document as projected: paths and bytes are typed inputs already; a +/// file-like object is a reader the projection consumes once every other field is read. +pub(super) enum FileDocumentInput { + Ready(OcrDocumentInput), + Deferred { + reader: PythonFileReader, + mime_type: Option, + }, } -impl PythonFileReader { - pub(super) fn read(&self, py: Python<'_>) -> PyResult { - let value = self.reader.bind(py).call0()?; - let bytes = if value.is_instance_of::() { - Bytes::from(value.extract::()?) - } else if value.is_instance_of::() { - extract_bytes(&value)? - } else { - return Err(PyTypeError::new_err(format!( - "OCR file read must return bytes or str, got {}", - value.get_type(), - ))); - }; - Ok(OcrFileContent { - bytes, - file_name: self.name.clone(), - }) +impl FileDocumentInput { + pub(super) fn resolve(self, py: Python<'_>) -> PyResult { + match self { + Self::Ready(input) => Ok(input), + Self::Deferred { reader, mime_type } => { + let content = reader.read(py)?; + Ok(OcrDocumentInput::Bytes { + bytes: content.bytes, + file_name: content.file_name, + mime_type, + }) + } + } } - - pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - visit.call(&self.reader) - } -} - -fn extract_bytes(value: &Bound<'_, PyAny>) -> PyResult { - if value.is_exact_instance_of::() { - return Ok(Bytes::from_owner(value.extract::()?)); - } - Ok(Bytes::copy_from_slice( - value.extract::()?.as_ref(), - )) -} - -pub(super) struct FileDocumentInput { - pub input: OcrDocumentInput, - pub reader: Option, } impl FromPyObject<'_, '_> for FileDocumentInput { @@ -87,51 +67,31 @@ impl FromPyObject<'_, '_> for FileDocumentInput { } static PATH_LIKE: PyOnceLock> = PyOnceLock::new(); if file.is_instance(PATH_LIKE.import(py, "os", "PathLike")?)? { - return Ok(Self { - input: OcrDocumentInput::Path { - path: file.extract::()?, - mime_type, - }, - reader: None, - }); + return Ok(Self::Ready(OcrDocumentInput::Path { + path: file.extract::()?, + mime_type, + })); } if file.is_instance_of::() { - return Ok(Self { - input: OcrDocumentInput::Bytes { - bytes: extract_bytes(&file)?, - file_name: None, - mime_type, - }, - reader: None, - }); + return Ok(Self::Ready(OcrDocumentInput::Bytes { + bytes: py_bytes(&file)?, + file_name: None, + mime_type, + })); } - let reader = file - .getattr_opt("read")? - .filter(|value| value.is_callable()); - let Some(reader) = reader else { - return Err(PyValueError::new_err(format!( + match PythonFileReader::from_file_like(&file)? { + Some(reader) => Ok(Self::Deferred { reader, mime_type }), + None => Err(PyValueError::new_err(format!( "Unsupported file input type: {}. Expected pathlib.Path, bytes, or a file-like object.", file.get_type(), - ))); - }; - let name = file - .getattr_opt("name")? - .filter(|value| !value.is_none()) - .map(|value| value.extract::()) - .transpose()?; - Ok(Self { - input: OcrDocumentInput::HostReader { mime_type }, - reader: Some(PythonFileReader { - reader: reader.unbind(), - name, - }), - }) + ))), + } } } #[cfg(test)] mod tests { - use pyo3::types::PyDict; + use pyo3::{exceptions::PyTypeError, types::PyDict}; use super::*; @@ -141,6 +101,13 @@ mod tests { locals } + fn ready(input: FileDocumentInput) -> OcrDocumentInput { + match input { + FileDocumentInput::Ready(input) => input, + FileDocumentInput::Deferred { .. } => panic!("expected a ready document"), + } + } + #[test] fn extraction_validates_required_file_and_optional_mime_type() { Python::initialize(); @@ -167,13 +134,19 @@ mod tests { .unwrap(); assert!(error.is_instance_of::(py)); assert!(error.to_string().contains("bare str")); + let error = py + .eval(c"{'file': object()}", None, None) + .unwrap() + .extract::() + .err() + .unwrap(); + assert!(error.is_instance_of::(py)); + assert!(error.to_string().contains("Unsupported file input type")); let document = py .eval(c"{'file': b'abc', 'mime_type': 'image/png'}", None, None) .unwrap(); - let input: FileDocumentInput = document.extract().unwrap(); - assert!(input.reader.is_none()); assert_eq!( - input.input, + ready(document.extract().unwrap()), OcrDocumentInput::Bytes { bytes: b"abc".as_slice().into(), file_name: None, @@ -199,7 +172,7 @@ class Reader: return b'abc' reader = Reader() document = {'file': reader, 'mime_type': 7} -reader_document = {'file': reader} +reader_document = {'file': reader, 'mime_type': 'application/pdf'} path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf'), 'mime_type': 'image/png'}", ); let document = locals.get_item("document").unwrap().unwrap(); @@ -208,10 +181,6 @@ path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf'), 'mime_typ let document = locals.get_item("reader_document").unwrap().unwrap(); let input: FileDocumentInput = document.extract().unwrap(); - assert_eq!( - input.input, - OcrDocumentInput::HostReader { mime_type: None } - ); let reads = || { locals .get_item("reader") @@ -223,21 +192,20 @@ path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf'), 'mime_typ .unwrap() }; assert_eq!(reads(), 0); - let content = input.reader.unwrap().read(py).unwrap(); + let resolved = input.resolve(py).unwrap(); assert_eq!(reads(), 1); assert_eq!( - content, - OcrFileContent { + resolved, + OcrDocumentInput::Bytes { bytes: b"abc".as_slice().into(), file_name: Some("scan.png".into()), + mime_type: Some("application/pdf".into()), } ); let document = locals.get_item("path_document").unwrap().unwrap(); - let input: FileDocumentInput = document.extract().unwrap(); - assert!(input.reader.is_none()); assert_eq!( - input.input, + ready(document.extract().unwrap()), OcrDocumentInput::Path { path: PathBuf::from("/nonexistent/ocr-projection-test.pdf"), mime_type: Some("image/png".into()), @@ -245,97 +213,4 @@ path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf'), 'mime_typ ); }); } - - #[test] - fn reader_results_are_normalized_and_exceptions_keep_their_identity() { - Python::initialize(); - Python::attach(|py| { - let locals = eval( - py, - c"failure = KeyError('reader failed') -class Raising: - def read(self): - raise failure -class Text: - def read(self): - return 'héllo' -class Wrong: - def read(self): - return 7 -raising = {'file': Raising()} -text = {'file': Text()} -wrong = {'file': Wrong()}", - ); - let reader = |name: &str| { - locals - .get_item(name) - .unwrap() - .unwrap() - .extract::() - .unwrap() - .reader - .unwrap() - }; - let error = reader("raising").read(py).unwrap_err(); - assert!( - error - .value(py) - .is(locals.get_item("failure").unwrap().unwrap()) - ); - assert_eq!( - reader("text").read(py).unwrap().bytes.as_ref(), - "héllo".as_bytes() - ); - let error = reader("wrong").read(py).unwrap_err(); - assert!(error.is_instance_of::(py)); - assert!(error.to_string().contains("bytes or str")); - }); - } - - #[rstest::rstest] - #[case::read("read")] - #[case::name("name")] - fn reader_attribute_failures_keep_their_identity(#[case] attribute: &str) { - Python::initialize(); - Python::attach(|py| { - let locals = eval( - py, - c"failure = LookupError('file property failed') -class File: - def __getattribute__(self, name): - if name == attribute: - raise failure - return super().__getattribute__(name) - name = 'scan.pdf' - def read(self): - return b'abc' -document = {'file': File()}", - ); - locals.set_item("attribute", attribute).unwrap(); - let error = locals - .get_item("document") - .unwrap() - .unwrap() - .extract::() - .err() - .unwrap(); - assert!( - error - .value(py) - .is(locals.get_item("failure").unwrap().unwrap()) - ); - }); - } - - #[test] - fn exact_python_bytes_transfer_without_copying_and_outlive_the_input() { - Python::initialize(); - let (bytes, pointer) = Python::attach(|py| { - let value = PyBytes::new(py, b"document bytes"); - let pointer = value.as_bytes().as_ptr() as usize; - (extract_bytes(value.as_any()).unwrap(), pointer) - }); - assert_eq!(bytes.as_ptr() as usize, pointer); - assert_eq!(bytes.as_ref(), b"document bytes"); - } } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs index 8bf99cd355f..5a3806e61e3 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs @@ -1,6 +1,6 @@ use litellm_auth::ResolvedCredential; -use litellm_core::ocr::route::{Ocr, OcrOp, OcrOpResult}; -use litellm_host_python::{InvokeError, RouteHost, missing_state, to_py}; +use litellm_core::ocr::route::{Ocr, OcrOp, OcrProjection}; +use litellm_host_python::{InvokeError, ProtocolHost, missing_state, to_py}; use litellm_llms::base_llm::ocr::{error::Error, transformation::LiteLLMOcrResponse}; use pyo3::{ exceptions::{PyBaseException, PyException}, @@ -20,14 +20,15 @@ enum OcrHostData { Released, } -/// The Python side of the OCR route: projects the prepared arguments, reads file-like -/// documents, acquires Azure AD tokens, and builds the public response and exception. -pub(super) struct OcrRouteHost { +/// The Python side of the OCR route: projects the prepared arguments (reading a file-like +/// document as it goes), acquires Azure AD tokens, and builds the public response and +/// exception. +pub(super) struct OcrPythonHost { request: Py, data: OcrHostData, } -impl OcrRouteHost { +impl OcrPythonHost { pub(super) fn new(request: Py) -> Self { Self { request, @@ -42,14 +43,6 @@ impl OcrRouteHost { } } - fn read_document(&self, py: Python<'_>) -> PyResult { - self.handles()? - .reader - .as_ref() - .ok_or_else(missing_state)? - .read(py) - } - fn acquire_azure_ad_token(&self, py: Python<'_>) -> PyResult { self.handles()? .azure_ad_token_provider @@ -58,30 +51,21 @@ impl OcrRouteHost { .acquire(py) } - fn answer( + fn projection( &mut self, py: Python<'_>, arguments: &Bound<'_, PyDict>, - op: OcrOp, - ) -> PyResult { - match op { - OcrOp::ProjectRequest => { - let OcrHostData::Unprojected = self.data else { - return Err(missing_state()); - }; - let (request, handles) = project_request(self.request.bind(py), arguments)?; - let caller_token = handles.azure_ad_token_provider.is_some(); - self.data = OcrHostData::Projected(Box::new(handles)); - Ok(OcrOpResult::Request { - request: Box::new(request), - caller_token, - }) - } - OcrOp::ReadDocument => self.read_document(py).map(OcrOpResult::Document), - OcrOp::AcquireAzureAdToken => self - .acquire_azure_ad_token(py) - .map(OcrOpResult::AzureAdToken), - } + ) -> PyResult { + let OcrHostData::Unprojected = self.data else { + return Err(missing_state()); + }; + let (request, handles) = project_request(self.request.bind(py), arguments)?; + let caller_token = handles.azure_ad_token_provider.is_some(); + self.data = OcrHostData::Projected(Box::new(handles)); + Ok(OcrProjection { + request, + caller_token, + }) } fn map_failure(&self, py: Python<'_>, error: PyErr) -> PyErr { @@ -104,20 +88,28 @@ impl OcrRouteHost { } } -impl RouteHost for OcrRouteHost { - type Route = Ocr; +impl ProtocolHost for OcrPythonHost { + type Protocol = Ocr; type Failure = PyErr; - fn invoke( + fn project( &mut self, py: Python<'_>, arguments: &Bound<'_, PyDict>, - op: OcrOp, - ) -> Result> { - self.answer(py, arguments, op) + ) -> Result> { + self.projection(py, arguments) .map_err(|error| InvokeError::Python(self.map_failure(py, error))) } + fn invoke(&mut self, py: Python<'_>, op: OcrOp) -> Result<(), InvokeError> { + match op { + OcrOp::AcquireAzureAdToken(reply) => self + .acquire_azure_ad_token(py) + .map(|token| reply.send(token)) + .map_err(|error| InvokeError::Python(self.map_failure(py, error))), + } + } + fn complete(&mut self, py: Python<'_>, response: LiteLLMOcrResponse) -> PyResult> { py.import("litellm.rust_bridge.ocr.route_host")? .getattr("response")? @@ -148,13 +140,10 @@ impl RouteHost for OcrRouteHost { fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { visit.call(&self.request)?; - if let OcrHostData::Projected(handles) = &self.data { - if let Some(reader) = &handles.reader { - reader.traverse(visit)?; - } - if let Some(provider) = &handles.azure_ad_token_provider { - provider.traverse(visit)?; - } + if let OcrHostData::Projected(handles) = &self.data + && let Some(provider) = &handles.azure_ad_token_provider + { + provider.traverse(visit)?; } Ok(()) } @@ -205,20 +194,13 @@ del provider .unwrap() .cast_into::() .unwrap(); - let mut host = OcrRouteHost::new(py.None()); - let projected = host.invoke(py, &kwargs, OcrOp::ProjectRequest).unwrap(); - assert!(matches!( - projected, - OcrOpResult::Request { - caller_token: true, - .. - } - )); + let mut host = OcrPythonHost::new(py.None()); + assert!(host.project(py, &kwargs).unwrap().caller_token); locals.del_item("kwargs").unwrap(); drop(kwargs); + let (reply, _) = litellm_host::host::reply(); assert_eq!( - host.invoke(py, &PyDict::new(py), OcrOp::AcquireAzureAdToken) - .is_ok(), + host.invoke(py, OcrOp::AcquireAzureAdToken(reply)).is_ok(), succeeds ); let alive = || { diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index b54316b258b..a4f2bf851d7 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -5,7 +5,7 @@ mod project; use std::sync::LazyLock; -use host::OcrRouteHost; +use host::OcrPythonHost; use litellm_auth_gcp::VertexAuth; use litellm_callbacks_legacy_python::{LegacySurface, PublicCall, run_legacy_call}; use litellm_core::ocr::{provider_config, route::ocr_machine}; @@ -69,7 +69,7 @@ fn run_ocr( if asynchronous { ASYNC_SURFACE } else { SURFACE }, PublicCall::capture(&request, &args, &kwargs)?, crate::logger::LoggedMachine::new(ocr_machine(client)), - OcrRouteHost::new(request.unbind()), + OcrPythonHost::new(request.unbind()), asynchronous, ) } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs index 697b935a1d4..be43a1b7711 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs @@ -8,19 +8,15 @@ use litellm_llms::base_llm::ocr::error::Error; use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict}; use serde_json::{Map, Value}; -use super::{ - document::{FileDocumentInput, PythonFileReader}, - errors::to_pyerr as ocr_error_to_pyerr, -}; +use super::{document::FileDocumentInput, errors::to_pyerr as ocr_error_to_pyerr}; use crate::{ credentials::{self, CallerTokenProvider}, marshal::{project_optional_fields, python_timeout_seconds, request_input_sources}, }; -/// What the host keeps after projection: the caller's callables that answer the document -/// read and token operations, and the provider name the failure mapping reports. +/// What the host keeps after projection: the caller's token callable that answers the +/// token operation, and the provider name the failure mapping reports. pub(super) struct OcrHostHandles { - pub reader: Option, pub azure_ad_token_provider: Option, pub provider: &'static str, } @@ -104,13 +100,11 @@ impl ProjectedDocument { Ok(Self::File(document.extract()?)) } - fn into_parts(self) -> PyResult<(OcrDocumentInput, Option)> { + /// Reads a file-like document now, so it runs after every other argument was read. + fn resolve(self, py: Python<'_>) -> PyResult { match self { - Self::File(FileDocumentInput { input, reader }) => Ok((input, reader)), - Self::Other(wire) => Ok(( - decode_document(wire).map_err(ocr_error_to_pyerr)?.into(), - None, - )), + Self::File(file) => file.resolve(py), + Self::Other(wire) => Ok(decode_document(wire).map_err(ocr_error_to_pyerr)?.into()), } } } @@ -136,24 +130,25 @@ pub(super) fn project_request( .chain(["api_key", "api_base", "extra_headers"]), )?; let azure_ad_token_provider = credentials::azure_ad_token_provider(kwargs)?; - let (document, reader) = document.into_parts()?; + let api_base = arguments.api_base()?; + let extra_headers = arguments.extra_headers()?; + let timeout_seconds = arguments.timeout_seconds()?; let wire = OcrWireRequest { model, - document, + document: document.resolve(request.py())?, api_key, - api_base: arguments.api_base()?, + api_base, custom_llm_provider, - extra_headers: arguments.extra_headers()?, + extra_headers, optional_params, input_sources, - timeout_seconds: arguments.timeout_seconds()?, + timeout_seconds, }; let request = decode_request_input(wire).map_err(ocr_error_to_pyerr)?; let provider = request.provider_name(); Ok(( request, OcrHostHandles { - reader, azure_ad_token_provider, provider, }, @@ -180,10 +175,8 @@ mod tests { OcrArguments { request, kwargs } } - fn project_document( - document: &Bound<'_, PyAny>, - ) -> PyResult<(OcrDocumentInput, Option)> { - ProjectedDocument::project(document)?.into_parts() + fn project_document(document: &Bound<'_, PyAny>) -> PyResult { + ProjectedDocument::project(document)?.resolve(document.py()) } fn url_document(url: &str) -> OcrDocumentInput { @@ -342,8 +335,11 @@ kwargs = {} }); } + /// A reader that rewrites the request while it runs shows which arguments projection + /// read before it and which after: every other argument is read first, and the read + /// happens exactly once. #[test] - fn document_readers_are_not_consumed_during_projection() { + fn document_readers_are_read_once_after_every_other_argument() { Python::initialize(); Python::attach(|py| { stub_timeout_conversion(py); @@ -351,17 +347,24 @@ kwargs = {} py, c" class Request: - api_base = 'original' + model = 'mistral/mistral-ocr-latest' + custom_llm_provider = None + api_key = None + api_base = 'https://original.example.com' + extra_headers = {'x-source': 'original'} timeout = 1 @property def document(self): return document class Reader: + reads = 0 def read(self): - Request.api_base = 'mutated' + Reader.reads += 1 + Request.api_base = 'https://mutated.example.com' + Request.extra_headers = {'x-source': 'mutated'} Request.timeout = 9 return b'abc' -document = {'type': 'file', 'file': Reader()} +document = {'type': 'file', 'file': Reader(), 'mime_type': 'application/pdf'} request = Request() kwargs = {} ", @@ -373,15 +376,38 @@ kwargs = {} .unwrap() .cast_into::() .unwrap(); - let arguments = arguments(&request, &kwargs); - let document = arguments.document().unwrap(); - let (input, reader) = project_document(&document).unwrap(); - assert_eq!(input, OcrDocumentInput::HostReader { mime_type: None }); - assert_eq!(arguments.api_base().unwrap().as_deref(), Some("original")); - assert_eq!(arguments.timeout_seconds().unwrap(), Some(1.0)); - reader.unwrap().read(py).unwrap(); - assert_eq!(arguments.api_base().unwrap().as_deref(), Some("mutated")); - assert_eq!(arguments.timeout_seconds().unwrap(), Some(9.0)); + let (projected, _) = project_request(&request, &kwargs).unwrap(); + assert_eq!( + py.eval(c"Reader.reads", Some(&locals), Some(&locals)) + .unwrap() + .extract::() + .unwrap(), + 1 + ); + assert_eq!( + projected.document, + OcrDocumentInput::Bytes { + bytes: b"abc".as_slice().into(), + file_name: None, + mime_type: Some("application/pdf".into()), + } + ); + assert_eq!( + projected + .credentials + .api_base + .as_ref() + .map(|base| base.value().as_str()), + Some("https://original.example.com") + ); + assert_eq!( + projected.transport.extra_headers, + [("x-source".to_string(), "original".to_string())] + ); + assert_eq!( + projected.transport.timeout, + Some(std::time::Duration::from_secs(1)) + ); }); } @@ -396,16 +422,14 @@ kwargs = {} None, ) .unwrap(); - let (input, reader) = project_document(&file).unwrap(); assert_eq!( - input, + project_document(&file).unwrap(), OcrDocumentInput::Bytes { bytes: b"%PDF-1.4".as_slice().into(), file_name: None, mime_type: Some("application/pdf".into()), } ); - assert!(reader.is_none()); let original = py .eval( @@ -414,8 +438,10 @@ kwargs = {} None, ) .unwrap(); - let (input, _) = project_document(&original).unwrap(); - assert_eq!(input, url_document("https://example.com/a.pdf")); + assert_eq!( + project_document(&original).unwrap(), + url_document("https://example.com/a.pdf") + ); }); } @@ -617,7 +643,7 @@ document = Document() ", ); let document = locals.get_item("document").unwrap().unwrap(); - let (input, _) = project_document(&document).unwrap(); + let input = project_document(&document).unwrap(); assert!(matches!(input, OcrDocumentInput::Bytes { .. })); let reads: Vec = document.getattr("reads").unwrap().extract().unwrap(); assert_eq!(reads, ["type", "mime_type", "file"]); diff --git a/litellm/constants.py b/litellm/constants.py index a86be55d654..67021ae2abc 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -167,6 +167,9 @@ MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE: Final = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_M MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL: Final = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL", "3600")) MCP_SSO_ASSERTION_CACHE_TTL_SECONDS: Final = int(os.getenv("MCP_SSO_ASSERTION_CACHE_TTL_SECONDS", "60")) +# mcp_tool_permissions entry that grants every current and future tool on a server +MCP_ALL_TOOLS_WILDCARD: Final = "*" + # Default npm cache directory for STDIO MCP servers. # npm/npx needs a writable cache dir; in containers the default (~/.npm) # may not exist or be read-only. /tmp is always writable. diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 6cc0d9444cd..a279b9f0903 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -790,6 +790,16 @@ def _get_hidden_str_for_cost_calc(hidden_params: object, key: str) -> str | None return value if isinstance(value, str) and value else None +_NON_TOKEN_RATE_FIELDS: Final = frozenset({"input_cost_per_second", "input_cost_per_query", "tiered_pricing"}) + + +def _cost_map_entry_prices_anything(entry: Mapping[str, object]) -> bool: + return any( + value is not None and (field in _NON_TOKEN_RATE_FIELDS or ("cost_per" in field and "token" in field)) + for field, value in entry.items() + ) + + def _select_model_name_for_cost_calc( model: str | None, completion_response: object | None, @@ -828,12 +838,7 @@ def _select_model_name_for_cost_calc( if custom_pricing is True: if router_model_id is not None and router_model_id in litellm.model_cost: entry: Final = litellm.model_cost[router_model_id] - if ( - entry.get("input_cost_per_token") is not None - or entry.get("input_cost_per_second") is not None - or entry.get("input_cost_per_query") is not None - or entry.get("tiered_pricing") is not None - ): + if _cost_map_entry_prices_anything(entry): return_model = router_model_id else: return_model = model @@ -1699,6 +1704,8 @@ def completion_cost( litellm_model_name=model, data_residency=data_residency, litellm_logging_obj=litellm_logging_obj, + custom_pricing_model=selected_model if custom_pricing else None, + base_pricing_model=(selected_model if base_model is not None and not custom_pricing else None), ) elif call_type == _MCP_CALL_TYPE: from litellm.proxy._experimental.mcp_server.cost_calculator import ( @@ -2870,14 +2877,20 @@ def _candidate_realtime_token_costs( def _cost_map_entry_declares_pricing(model_name: str, custom_llm_provider: str) -> bool: + """Whether the entry behind ``model_name`` sets any rate of its own, even a zero one. + + The name is resolved the way ``get_model_info`` resolves it before the raw entry is read, + because a deployment-scoped name arrives here already carrying its provider prefix. Two raw + lookups cannot strip that prefix, so a zero-rated override read as declaring nothing, and a + session that should bill nothing fell through to the public rates instead. + """ + resolved: Final = _get_model_info_or_none(model_name, custom_llm_provider) entries: Final = ( + litellm.model_cost.get(resolved.get("key")) if resolved is not None else None, litellm.model_cost.get(model_name), litellm.model_cost.get(f"{custom_llm_provider}/{model_name}"), ) - return any( - entry is not None and any("cost_per" in field and value is not None for field, value in entry.items()) - for entry in entries - ) + return any(entry is not None and _cost_map_entry_prices_anything(entry) for entry in entries) def _first_priced_realtime_token_costs( @@ -2917,6 +2930,8 @@ def handle_realtime_stream_cost_calculation( litellm_model_name: str, data_residency: str | None = None, litellm_logging_obj: LitellmLoggingObject | None = None, + custom_pricing_model: str | None = None, + base_pricing_model: str | None = None, ) -> float: """ Handles the cost calculation for realtime stream responses. @@ -2925,9 +2940,13 @@ def handle_realtime_stream_cost_calculation( Args: results: A list of OpenAIRealtimeStreamBaseObject objects + custom_pricing_model: deployment-scoped pricing key from the deployment's + custom rates, tried ahead of the session-reported model + base_pricing_model: the deployment's resolved base_model, tried ahead of the + session-reported model but after custom rates """ received_model = None - potential_model_names: Final = [] + potential_model_names: Final = [custom_pricing_model, base_pricing_model] for result in results: if result["type"] == "session.created": received_model = cast(OpenAIRealtimeStreamSessionEvents, result)["session"].get("model", None) @@ -2945,6 +2964,7 @@ def handle_realtime_stream_cost_calculation( results=results, custom_llm_provider=custom_llm_provider, litellm_model_name=litellm_model_name, + custom_pricing_model=custom_pricing_model, ) if any(r.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE for r in results) else 0.0 @@ -2968,6 +2988,7 @@ def handle_realtime_transcription_cost_calculation( results: OpenAIRealtimeStreamList, custom_llm_provider: str, litellm_model_name: str, + custom_pricing_model: str | None = None, ) -> float: """ Cost for realtime transcription sessions (e.g. gpt-realtime-whisper). @@ -2985,15 +3006,15 @@ def handle_realtime_transcription_cost_calculation( return 0.0 model_name: Final = _get_transcription_model_name_from_results(results) or litellm_model_name - try: - model_info = litellm.get_model_info(model=model_name, custom_llm_provider=custom_llm_provider) - except Exception: - model_info = None + model_info: Final = _get_model_info_or_none(model_name, custom_llm_provider) + override_info: Final = ( + _get_model_info_or_none(custom_pricing_model, custom_llm_provider) if custom_pricing_model is not None else None + ) total_cost = 0.0 for event in completed_events: usage = event.get("usage") or {} - total_cost += _transcription_usage_cost(usage, model_info) + total_cost += _transcription_usage_cost(usage, model_info, override_info) return total_cost @@ -3018,23 +3039,57 @@ def _get_transcription_model_name_from_results( return None -def _transcription_usage_cost(usage: dict, model_info: ModelInfo | None) -> float: - if model_info is None: +def _get_model_info_or_none(model: str, custom_llm_provider: str) -> ModelInfo | None: + try: + return litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider) + except Exception: + return None + + +def _declared_transcription_rate(info: ModelInfo | None, keys: tuple[str, ...]) -> float | None: + """First of ``keys`` this entry prices, read off the raw ``litellm.model_cost`` entry + because ``get_model_info`` synthesizes zero token rates for entries that omit them.""" + if info is None: + return None + declared: Final = litellm.model_cost.get(info.get("key")) + if declared is None: + return None + return next( + (float(value) for key in keys if declared.get(key) is not None and (value := info.get(key)) is not None), + None, + ) + + +def _transcription_rate(keys: tuple[str, ...], override: ModelInfo | None, base: ModelInfo | None) -> float: + rates: Final = (_declared_transcription_rate(info, keys) for info in (override, base)) + return next((rate for rate in rates if rate is not None), 0.0) + + +def _transcription_usage_cost( + usage: dict, + model_info: ModelInfo | None, + override_info: ModelInfo | None = None, +) -> float: + if model_info is None and override_info is None: return 0.0 + usage_type: Final = usage.get("type") if usage_type == "duration": seconds: Final = usage.get("seconds") or 0.0 - per_second: Final = model_info.get("input_cost_per_second") or 0.0 - return float(seconds) * float(per_second) + return float(seconds) * _transcription_rate(("input_cost_per_second",), override_info, model_info) if usage_type == "tokens": input_token_details: Final = usage.get("input_token_details") or {} audio_tokens: Final = input_token_details.get("audio_tokens") or 0 text_tokens: Final = input_token_details.get("text_tokens") or 0 output_tokens: Final = usage.get("output_tokens") or 0 - audio_cost: Final = float(audio_tokens) * float( - model_info.get("input_cost_per_audio_token") or model_info.get("input_cost_per_token") or 0.0 + audio_cost: Final = float(audio_tokens) * _transcription_rate( + ("input_cost_per_audio_token", "input_cost_per_token"), override_info, model_info + ) + text_cost: Final = float(text_tokens) * _transcription_rate( + ("input_cost_per_token",), override_info, model_info + ) + output_cost: Final = float(output_tokens) * _transcription_rate( + ("output_cost_per_token",), override_info, model_info ) - text_cost: Final = float(text_tokens) * float(model_info.get("input_cost_per_token") or 0.0) - output_cost: Final = float(output_tokens) * float(model_info.get("output_cost_per_token") or 0.0) return audio_cost + text_cost + output_cost return 0.0 diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py index 14aebcaabaf..6d050d5a856 100644 --- a/litellm/llms/vertex_ai/common_utils.py +++ b/litellm/llms/vertex_ai/common_utils.py @@ -25,6 +25,21 @@ from litellm.types.llms.vertex_ai import ( from litellm.types.utils import TokenCountResponse from litellm.utils import supports_response_schema, supports_system_messages +VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS: Final = frozenset( + { + "audio", + "max_retries", + "modalities", + "prediction", + "prompt_cache_key", + "prompt_cache_retention", + "safety_identifier", + "service_tier", + "store", + "web_search_options", + } +) + class VertexAILyriaModelInfo(TypedDict): vertex_ai_audio_api: ReadOnly[Literal["lyria_predict", "lyria_interactions"]] @@ -370,6 +385,27 @@ def get_vertex_base_model_name(model: str) -> str: return model +def vertex_model_garden_model_id_in_json_body(model: str) -> bool: + """ + Vertex catalog / publisher models are addressed as publisher/model (e.g. + xai/grok-4.1-fast-reasoning) on the shared OpenAPI URL, with the id in the JSON body. + + Deployed Model Garden endpoints are typically a single segment (often numeric) + and use .../endpoints/{ENDPOINT_ID}/chat/completions with an empty model field. + """ + return "/" in model + + +def is_vertex_self_deployed_openai_compatible_endpoint(model: str) -> bool: + local_model: Final = model.removeprefix("vertex_ai/") + route: Final = get_vertex_ai_model_route(local_model) + if route == VertexAIModelRoute.GEMMA: + return True + return route == VertexAIModelRoute.MODEL_GARDEN and not vertex_model_garden_model_id_in_json_body( + get_vertex_base_model_name(local_model) + ) + + def get_vertex_ai_fine_tuned_endpoint_id(model: str) -> str | None: """ Fine-tuned Gemini deployments are addressed by a numeric endpoint id, diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py index f2d2c0896d2..ca0bcb74906 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py @@ -18,7 +18,11 @@ from litellm.types.utils import ( Usage, ) -from ...common_utils import VertexAIError +from ...common_utils import ( + VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS, + VertexAIError, + is_vertex_self_deployed_openai_compatible_endpoint, +) if TYPE_CHECKING: from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer @@ -66,13 +70,15 @@ class VertexAILlama3Config(OpenAIGPTConfig): and v is not None } - def get_supported_openai_params(self, model: str): - supported_params: Final = super().get_supported_openai_params(model=model) - try: - supported_params.remove("max_retries") - except KeyError: - pass - return supported_params + def get_supported_openai_params(self, model: str) -> list[str]: + unsupported_params: Final = ( + VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS + if is_vertex_self_deployed_openai_compatible_endpoint(model) + else frozenset({"max_retries"}) + ) + return [ # mutable-ok: get_optional_params extends the returned list with allowed_openai_params + param for param in super().get_supported_openai_params(model=model) if param not in unsupported_params + ] def map_openai_params( self, diff --git a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py index a774dba6cf2..ea97f0a0a9a 100644 --- a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py +++ b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py @@ -21,6 +21,7 @@ from litellm.llms.custom_httpx.http_handler import ( _get_httpx_client, ) from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig +from litellm.llms.vertex_ai.common_utils import VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS from litellm.types.llms.openai import AllMessageValues from litellm.types.llms.vertex_ai_gemma import VertexGemmaContainerError from litellm.types.utils import ModelResponse @@ -49,6 +50,13 @@ class VertexGemmaConfig(OpenAIGPTConfig): def __init__(self) -> None: super().__init__() + def get_supported_openai_params(self, model: str) -> list[str]: + return [ # mutable-ok: get_optional_params extends the returned list with allowed_openai_params + param + for param in super().get_supported_openai_params(model=model) + if param not in VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS + ] + def should_fake_stream( self, model: str | None, diff --git a/litellm/llms/vertex_ai/vertex_model_garden/main.py b/litellm/llms/vertex_ai/vertex_model_garden/main.py index f5c9ac623a1..84907f01685 100644 --- a/litellm/llms/vertex_ai/vertex_model_garden/main.py +++ b/litellm/llms/vertex_ai/vertex_model_garden/main.py @@ -24,21 +24,14 @@ import httpx from litellm.llms.vertex_ai.common_utils import get_vertex_base_url from litellm.utils import ModelResponse -from ..common_utils import VertexAIError, get_vertex_base_model_name +from ..common_utils import ( + VertexAIError, + get_vertex_base_model_name, + vertex_model_garden_model_id_in_json_body, +) from ..vertex_llm_base import VertexBase -def _vertex_model_garden_model_id_in_json_body(model: str) -> bool: - """ - Vertex catalog / publisher models are addressed as publisher/model (e.g. - xai/grok-4.1-fast-reasoning) on the shared OpenAPI URL, with the id in the JSON body. - - Deployed Model Garden endpoints are typically a single segment (often numeric) - and use .../endpoints/{ENDPOINT_ID}/chat/completions with an empty model field. - """ - return "/" in model - - def create_vertex_url( vertex_location: str, vertex_project: str, @@ -48,7 +41,7 @@ def create_vertex_url( ) -> str: """Return the api base for vertex model garden (without /chat/completions).""" base_url: Final = get_vertex_base_url(vertex_location) - if _vertex_model_garden_model_id_in_json_body(model): + if vertex_model_garden_model_id_in_json_body(model): return f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/endpoints/openapi" return f"{base_url}/v1beta1/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}" @@ -124,7 +117,7 @@ class VertexAIModelGardenModels(VertexBase): ) # Publisher/catalog models: model id must be sent in the JSON body (OpenAPI route). # Single-segment endpoint ids: model is encoded in the URL path; body model stays empty. - if not _vertex_model_garden_model_id_in_json_body(model): + if not vertex_model_garden_model_id_in_json_body(model): model = "" return openai_like_chat_completions.completion( model=model, diff --git a/litellm/main.py b/litellm/main.py index 98bb5126a90..72c9afad36c 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -126,6 +126,7 @@ from litellm.types.completion import ( _CompletionDispatchContext, _CompletionDispatchResult, ) +from litellm.types.litellm_params import RetryStrategy from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ( CustomPricingLiteLLMParams, @@ -6026,9 +6027,7 @@ def completion_with_retries(*args, **kwargs): # reset retries in .completion() kwargs["max_retries"] = 0 kwargs["num_retries"] = 0 - retry_strategy: Final[Literal["exponential_backoff_retry", "constant_retry"]] = kwargs.pop( - "retry_strategy", "constant_retry" - ) + retry_strategy: Final[RetryStrategy] = kwargs.pop("retry_strategy", "constant_retry") original_function: Final = kwargs.pop("original_function", completion) if retry_strategy == "exponential_backoff_retry": retryer = tenacity.Retrying( @@ -6054,7 +6053,7 @@ async def acompletion_with_retries(*args, **kwargs): num_retries: Final = kwargs.pop("num_retries", 3) kwargs["max_retries"] = 0 kwargs["num_retries"] = 0 - retry_strategy: Final = kwargs.pop("retry_strategy", "constant_retry") + retry_strategy: Final[RetryStrategy] = kwargs.pop("retry_strategy", "constant_retry") original_function: Final = kwargs.pop("original_function", completion) if retry_strategy == "exponential_backoff_retry": retryer = tenacity.AsyncRetrying( @@ -6082,9 +6081,7 @@ def responses_with_retries(*args, **kwargs): # reset retries in .responses() kwargs["max_retries"] = 0 kwargs["num_retries"] = 0 - retry_strategy: Final[Literal["exponential_backoff_retry", "constant_retry"]] = kwargs.pop( - "retry_strategy", "constant_retry" - ) + retry_strategy: Final[RetryStrategy] = kwargs.pop("retry_strategy", "constant_retry") original_function: Final = kwargs.pop("original_function", responses) if retry_strategy == "exponential_backoff_retry": retryer = tenacity.Retrying( @@ -6111,7 +6108,7 @@ async def aresponses_with_retries(*args, **kwargs): num_retries: Final = kwargs.pop("num_retries", 3) kwargs["max_retries"] = 0 kwargs["num_retries"] = 0 - retry_strategy: Final = kwargs.pop("retry_strategy", "constant_retry") + retry_strategy: Final[RetryStrategy] = kwargs.pop("retry_strategy", "constant_retry") original_function: Final = kwargs.pop("original_function", aresponses) if retry_strategy == "exponential_backoff_retry": retryer = tenacity.AsyncRetrying( diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index b529a225f84..f875936b0bf 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -3342,6 +3342,7 @@ "supports_vision": true }, "azure/command-r-plus": { + "deprecation_date": "2025-06-30", "input_cost_per_token": 3e-06, "litellm_provider": "azure", "max_input_tokens": 128000, @@ -3349,6 +3350,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.5e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_function_calling": true }, "azure_ai/claude-haiku-4-5": { @@ -5194,6 +5196,7 @@ "output_cost_per_token": 2e-06 }, "azure/gpt-4": { + "deprecation_date": "2025-06-06", "input_cost_per_token": 3e-05, "litellm_provider": "azure", "max_input_tokens": 8192, @@ -5201,6 +5204,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_function_calling": true, "supports_tool_choice": true }, @@ -69074,6 +69078,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/google/gemma-4-31B-it": { + "deprecation_date": "2026-09-14", "input_cost_per_token": 3.9e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -69203,6 +69208,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/openai/gpt-oss-20b": { + "deprecation_date": "2026-09-14", "input_cost_per_token": 5e-08, "litellm_provider": "together_ai", "max_input_tokens": 131072, diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 793e558c943..e789a001878 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -13,6 +13,7 @@ from typing_extensions import assert_never import litellm from litellm._logging import verbose_logger +from litellm.constants import MCP_ALL_TOOLS_WILDCARD from litellm.proxy._experimental.mcp_server.oauth_utils import ( get_passthrough_resource_metadata_url, get_passthrough_www_authenticate, @@ -92,7 +93,9 @@ def level_allowed_tools( answer from this level and intersects with the rest. 1. A legacy ``mcp_tool_permissions`` entry for the server stays a closed - allowlist (``[]`` denies all): allowed = legacy ∪ toolset tools. + allowlist (``[]`` denies all): allowed = legacy ∪ toolset tools. An + entry containing ``MCP_ALL_TOOLS_WILDCARD`` grants every current and + future tool, so the level places no restriction at all. 2. An unconverted row (``mcp_permission_version`` falsy) keeps pre-overrides behavior: unrestricted unless a toolset names the server. 3. A converted row that does not grant the server places no restriction. @@ -109,6 +112,8 @@ def level_allowed_tools( toolset: Final[frozenset[str]] = frozenset(toolset_tools or ()) legacy: Final = global_mcp_server_manager.expand_tool_permissions(row.mcp_tool_permissions).get(server_id) if legacy is not None: + if MCP_ALL_TOOLS_WILDCARD in legacy: + return None return frozenset(legacy) | toolset if not row.mcp_permission_version: return frozenset(toolset) if toolset_tools is not None else None @@ -2196,7 +2201,11 @@ class MCPRequestHandler: via_toolsets: Sequence[str] | None, ) -> Sequence[str] | None: """Union of one level's direct tool grants and its toolset-granted tools on one server, - ``None`` when neither source restricts (allow-all from this level).""" + ``None`` when neither source restricts (allow-all from this level). A direct grant + containing ``MCP_ALL_TOOLS_WILDCARD`` makes the level unrestricted, so it returns + ``None`` whatever the toolsets name.""" + if direct is not None and MCP_ALL_TOOLS_WILDCARD in direct: + return None if direct is None and via_toolsets is None: return None return tuple({*(direct or ()), *(via_toolsets or ())}) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 44c1dee63e5..121ddc67753 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -27,7 +27,8 @@ from collections.abc import ( from contextlib import asynccontextmanager from dataclasses import dataclass, replace from functools import lru_cache -from itertools import chain +from itertools import chain, groupby +from operator import itemgetter from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Generic, Literal, TypeAlias, TypedDict, TypeVar, cast from urllib.parse import ParseResult, urlparse @@ -6857,9 +6858,11 @@ class MCPServerManager: """ Rewrite an ``mcp_tool_permissions`` dict keyed by id/name/alias so every key is a concrete server_id where possible. Tool lists from - keys that point at the same server are unioned, matching the - "duplicate names grant access to all matches" semantics of - ``expand_permission_list``. + keys that point at the same server are unioned and deduplicated + first-seen, matching the "duplicate names grant access to all + matches" semantics of ``expand_permission_list``; the + ``MCP_ALL_TOOLS_WILDCARD`` entry is preserved as an ordinary list + entry for the caller to interpret. Required so name-based keys don't silently drop their tool restrictions when the lookup uses the resolved server_id. Unresolved @@ -6868,11 +6871,15 @@ class MCPServerManager: """ if not tool_permissions: return {} - result: Final[dict[str, list[str]]] = {} - for key, tools in tool_permissions.items(): - for server_id in self.expand_permission_list((key,)): - result.setdefault(server_id, []).extend(tools or []) - return result + expanded: Final = tuple( + (server_id, tuple(tools or ())) + for key, tools in tool_permissions.items() + for server_id in self.expand_permission_list([key]) + ) + return { + server_id: list(dict.fromkeys(tool for _, tools in group for tool in tools)) + for server_id, group in groupby(sorted(expanded, key=itemgetter(0)), key=itemgetter(0)) + } def expand_tool_overrides( self, diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 67950e603c0..f19a8055ae6 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -89,6 +89,7 @@ from litellm.proxy.common_utils.http_parsing_utils import ( _safe_get_request_headers, _safe_get_request_query_params, ) +from litellm.proxy.common_utils.model_listing_utils import alias_map from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.common_utils.user_api_key_cache import ( END_USER_RESTRICTED_REGISTRY_OVERFLOW_SENTINEL, @@ -722,6 +723,7 @@ async def _run_project_checks( model=_model, project_object=project_object, llm_router=llm_router, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) if not skip_budget_checks: @@ -1018,6 +1020,7 @@ async def common_checks( team_object=team_object, llm_router=llm_router, team_model_aliases=(valid_token.team_model_aliases if valid_token else None), + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) except ProxyException as team_denial: if team_denial.type != ProxyErrorTypes.team_model_access_denied: @@ -1027,6 +1030,7 @@ async def common_checks( valid_token=valid_token, team_object=team_object, llm_router=llm_router, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ): raise @@ -1043,6 +1047,7 @@ async def common_checks( proxy_logging_obj=proxy_logging_obj, team_membership=loaded_team_membership, team_membership_loaded=team_membership_loaded, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) # Require trace id for agent keys when agent has require_trace_id_on_calls_by_agent @@ -1081,6 +1086,7 @@ async def common_checks( model=_model, llm_router=llm_router, user_object=user_object, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) # 1.1 - 2.2 - 3.0.2 - 3.0.3: Project checks (blocked, model access, budget) @@ -4349,6 +4355,7 @@ def _can_object_call_model( models: list[str], team_model_aliases: dict[str, str] | None = None, team_id: str | None = None, + key_model_aliases: Mapping[str, str] | None = None, object_type: Literal["user", "team", "key", "org", "project", "agent"] = "user", fallback_depth: int = 0, ) -> Literal[True]: @@ -4378,6 +4385,7 @@ def _can_object_call_model( models=models, team_model_aliases=team_model_aliases, team_id=team_id, + key_model_aliases=key_model_aliases, object_type=object_type, fallback_depth=fallback_depth + 1, ) @@ -4386,13 +4394,32 @@ def _can_object_call_model( from litellm.router_strategy.complexity_router.context_compaction import native_compaction_parent compaction_parent: Final = native_compaction_parent(model) - potential_models: Final = [model, compaction_parent] if compaction_parent is not None else [model] - if model in litellm.model_alias_map: - potential_models.append(litellm.model_alias_map[model]) - elif llm_router and model in llm_router.model_group_alias: - _model: Final = llm_router._get_model_from_alias(model) - if _model: - potential_models.append(_model) + global_or_router_alias_target: Final = ( + litellm.model_alias_map[model] + if model in litellm.model_alias_map + else ( + llm_router._get_model_from_alias(model) + if llm_router is not None and model in llm_router.model_group_alias + else None + ) + ) + after_team_alias: Final = team_model_aliases.get(model, model) if team_model_aliases else model + after_key_alias: Final = ( + key_model_aliases.get(after_team_alias, after_team_alias) if key_model_aliases else after_team_alias + ) + after_global_alias: Final = litellm.model_alias_map.get(after_key_alias, after_key_alias) + dispatched_model: Final = ( + key_model_aliases.get(after_global_alias, after_global_alias) if key_model_aliases else after_global_alias + ) + key_alias_applied: Final = after_key_alias != after_team_alias or dispatched_model != after_global_alias + potential_models: Final = ( + (dispatched_model,) + if key_alias_applied + else ( + *((model, compaction_parent) if compaction_parent is not None else (model,)), + *((global_or_router_alias_target,) if global_or_router_alias_target else ()), + ) + ) ## check model access for alias + underlying model - allow if either is in allowed models for m in potential_models: @@ -4418,6 +4445,35 @@ def _can_object_call_model( ) +def _resolve_team_alias( + model: str | list[str], + team_model_aliases: dict[str, str] | None, + team_id: str | None, + llm_router: Router | None, +) -> str | list[str]: + if not team_model_aliases: + return model + if isinstance(model, str): + return _live_team_alias_target(model, team_model_aliases, team_id, llm_router) + return [ # mutable-ok: _can_object_call_model takes list[str] + _live_team_alias_target(name, team_model_aliases, team_id, llm_router) for name in model + ] + + +def _live_team_alias_target( + model: str, team_model_aliases: dict[str, str], team_id: str | None, llm_router: Router | None +) -> str: + target: Final = team_model_aliases.get(model) + if target is None: + return model + deleted_team_deployment: Final = ( + llm_router is not None + and target.startswith(f"model_name_{team_id}_") + and target not in llm_router.model_name_to_deployment_indices + ) + return model if deleted_team_deployment else target + + async def _check_agent_access_group_model_access( model: str | list[str] | None, # mutable-ok: _can_object_call_model and the client message helper take list[str] valid_token: UserAPIKeyAuth | None, @@ -4438,12 +4494,14 @@ async def _check_agent_access_group_model_access( param="model", code=status.HTTP_403_FORBIDDEN, ) + dispatched: Final = _resolve_team_alias(model, valid_token.team_model_aliases, valid_token.team_id, llm_router) return _can_object_call_model( - model=model, + model=dispatched, llm_router=llm_router, models=sorted(ceiling.models), team_id=valid_token.team_id, object_type="agent", + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) @@ -4471,12 +4529,14 @@ async def _check_agent_caller_model_access( if caller_auth is None: return caller_team: Final = await load_team(valid_token) + caller_key_model_aliases: Final = key_model_aliases_for_auth_check(valid_token) if caller_team is not None: await can_team_access_model( model=model, team_object=caller_team, llm_router=llm_router, prisma_client=prisma_client, + key_model_aliases=caller_key_model_aliases, ) await _check_team_member_model_access( model=model, @@ -4486,12 +4546,18 @@ async def _check_agent_caller_model_access( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + key_model_aliases=caller_key_model_aliases, ) return caller_user: Final = await load_user(valid_token) if caller_user is None: return - await can_user_call_model(model=model, llm_router=llm_router, user_object=caller_user) + await can_user_call_model( + model=model, + llm_router=llm_router, + user_object=caller_user, + key_model_aliases=caller_key_model_aliases, + ) def _model_in_team_aliases(model: str, team_model_aliases: dict[str, str] | None = None) -> bool: @@ -4512,6 +4578,10 @@ def _model_in_team_aliases(model: str, team_model_aliases: dict[str, str] | None return False +def key_model_aliases_for_auth_check(valid_token: UserAPIKeyAuth | None) -> Mapping[str, str] | None: + return alias_map(valid_token.aliases) if valid_token is not None and valid_token.aliases else None + + def _resolve_key_models_for_auth_check(valid_token: UserAPIKeyAuth) -> list[str]: """ Expand key model sentinels before auth checks. @@ -4831,6 +4901,7 @@ async def can_key_call_model( models=key_models, team_model_aliases=valid_token.team_model_aliases, team_id=valid_token.team_id, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), object_type="key", ) except ProxyException: @@ -4848,6 +4919,7 @@ async def can_key_call_model( models=models_from_groups, team_model_aliases=valid_token.team_model_aliases, team_id=valid_token.team_id, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), object_type="key", ) raise @@ -4906,6 +4978,7 @@ async def can_key_call_resolved_model( team_object=team_object, llm_router=llm_router, team_model_aliases=valid_token.team_model_aliases, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) except ProxyException as team_denial: if team_denial.type != ProxyErrorTypes.team_model_access_denied: @@ -4915,6 +4988,7 @@ async def can_key_call_resolved_model( valid_token=valid_token, team_object=team_object, llm_router=llm_router, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ): raise @@ -4927,6 +5001,7 @@ async def can_key_call_resolved_model( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) if valid_token.project_id is not None: @@ -4941,6 +5016,7 @@ async def can_key_call_resolved_model( model=model, project_object=project_object, llm_router=llm_router, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) @@ -4968,6 +5044,7 @@ async def can_team_access_model( team_object: LiteLLM_TeamTable | None, llm_router: Router | None, team_model_aliases: dict[str, str] | None = None, + key_model_aliases: Mapping[str, str] | None = None, prisma_client: DatabaseClient | None = None, ) -> Literal[True]: """ @@ -4983,6 +5060,7 @@ async def can_team_access_model( models=team_object.models if team_object else [], team_model_aliases=team_model_aliases, team_id=team_object.team_id if team_object else None, + key_model_aliases=key_model_aliases, object_type="team", ) except ProxyException: @@ -5000,6 +5078,7 @@ async def can_team_access_model( models=list(dict.fromkeys([*(team_object.models if team_object else []), *models_from_groups])), team_model_aliases=team_model_aliases, team_id=team_object.team_id if team_object else None, + key_model_aliases=key_model_aliases, object_type="team", ) raise @@ -5058,6 +5137,7 @@ async def _key_access_group_grants_model( valid_token: UserAPIKeyAuth | None, team_object: LiteLLM_TeamTable | None, llm_router: Router | None, + key_model_aliases: Mapping[str, str] | None = None, ) -> bool: """ Returns True if the key's `access_group_ids` expand to models that grant @@ -5078,6 +5158,7 @@ async def _key_access_group_grants_model( models=authorized_models, team_model_aliases=valid_token.team_model_aliases if valid_token else None, team_id=valid_token.team_id if valid_token else None, + key_model_aliases=key_model_aliases, object_type="key", ) return True @@ -5089,6 +5170,7 @@ def can_project_access_model( model: str | list[str], project_object: LiteLLM_ProjectTable, llm_router: Router | None, + key_model_aliases: Mapping[str, str] | None = None, ) -> Literal[True]: """ Returns True if the project can access a specific model. @@ -5099,6 +5181,7 @@ def can_project_access_model( model=model, llm_router=llm_router, models=project_object.models if project_object else [], + key_model_aliases=key_model_aliases, object_type="project", ) @@ -5107,6 +5190,7 @@ async def can_user_call_model( model: str | list[str], llm_router: Router | None, user_object: LiteLLM_UserTable | None, + key_model_aliases: Mapping[str, str] | None = None, ) -> Literal[True]: if user_object is None: return True @@ -5128,6 +5212,7 @@ async def can_user_call_model( model=model, llm_router=llm_router, models=user_object.models, + key_model_aliases=key_model_aliases, object_type="user", ) @@ -5682,6 +5767,7 @@ async def _check_team_member_model_access( proxy_logging_obj: ProxyLogging, team_membership: LiteLLM_TeamMembership | None = None, team_membership_loaded: bool = False, + key_model_aliases: Mapping[str, str] | None = None, ) -> None: """ Check if a team member's per-member model scope allows access to the requested model. @@ -5717,6 +5803,7 @@ async def _check_team_member_model_access( models=member_allowed_models, object_type="team", team_id=team_object.team_id, + key_model_aliases=key_model_aliases, ) except ProxyException: internal_message: Final = ( diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index ae95e94dd2d..22c3a248b9d 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -66,6 +66,7 @@ from litellm.proxy.auth.auth_checks import ( get_user_object, is_valid_fallback_model, jwt_key_mapping_cache_key, + key_model_aliases_for_auth_check, resolve_and_validate_end_user_id, resolve_default_end_user_budget, ) @@ -469,6 +470,7 @@ async def _check_key_model_budget_with_fallback( models=valid_token.team_models, team_model_aliases=valid_token.team_model_aliases, team_id=valid_token.team_id, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), object_type="team", ) except ProxyException: diff --git a/litellm/proxy/common_utils/model_listing_utils.py b/litellm/proxy/common_utils/model_listing_utils.py index 3c6555662e6..8958fb20918 100644 --- a/litellm/proxy/common_utils/model_listing_utils.py +++ b/litellm/proxy/common_utils/model_listing_utils.py @@ -180,7 +180,7 @@ def caller_alias_maps( return CallerAliases((team_aliases, key_aliases), (team_aliases, key_aliases, litellm.model_alias_map, key_aliases)) -def _alias_map(aliases: object) -> Mapping[str, str]: +def alias_map(aliases: object) -> Mapping[str, str]: try: entries: Final = _ALIAS_ENTRIES.validate_python(aliases, strict=True) except ValidationError: @@ -204,7 +204,7 @@ def alias_target(model_id: str, aliases: CallerAliases, listed: Container[str] = already `listed` keeps its own row, so it is never rewritten.""" if model_id in listed: return None - return _rewrite(model_id, tuple(_alias_map(alias_map) for alias_map in aliases.rewrite)) + return _rewrite(model_id, tuple(alias_map(raw) for raw in aliases.rewrite)) def alias_listing_entries( @@ -213,8 +213,8 @@ def alias_listing_entries( ) -> tuple[tuple[str, str], ...]: """`entries` plus one `(alias, lookup_id)` row per key or team alias whose target is listed. An alias colliding with a listed id keeps the listed entry.""" - maps: Final = tuple(_alias_map(alias_map) for alias_map in aliases.rewrite) - own: Final = tuple(_alias_map(alias_map) for alias_map in aliases.own) + maps: Final = tuple(alias_map(raw) for raw in aliases.rewrite) + own: Final = tuple(alias_map(raw) for raw in aliases.own) lookup_by_response: Final = MappingProxyType(dict(entries)) lookup_ids: Final = frozenset(lookup_by_response.values()) targets: Final = MappingProxyType( diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 7d6db30e3e3..e0a4184291e 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -105,6 +105,8 @@ from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.proxy.utils import normalize_route_for_root_path from litellm.repositories.team_repository import TeamRepository from litellm.secret_managers.main import get_secret_str +from litellm.types import utils as types_utils +from litellm.types.litellm_params import ProxyRequestState, wire_names from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.passthrough_endpoints.pass_through_endpoints import ( LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, @@ -133,6 +135,9 @@ router: Final = APIRouter() pass_through_endpoint_logging: Final = PassThroughEndpointLogging() +_METADATA_KEYS: Final = frozenset(("litellm_metadata", "metadata")) +_KEPT_OUT_OF_LITELLM_PARAMS: Final = _METADATA_KEYS | frozenset(wire_names(ProxyRequestState)) + # Global registry to track registered pass-through routes and prevent memory leaks _registered_pass_through_routes: Final[dict[str, dict[str, str | bool | list[str] | Mapping[str, object]]]] = {} @@ -578,21 +583,21 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): """ Filter out litellm params from the request body """ - from litellm.types.utils import all_litellm_params - _parsed_body = _parsed_body or {} - litellm_params_in_body: Final = {} - for k in all_litellm_params: - if k in _parsed_body: - litellm_params_in_body[k] = _parsed_body.pop(k, None) + litellm_keys_in_body: Final = MappingProxyType( + {k: _parsed_body.pop(k) for k in types_utils.all_litellm_params if k in _parsed_body} + ) + litellm_params_in_body: Final = MappingProxyType( + {k: v for k, v in litellm_keys_in_body.items() if k not in _KEPT_OUT_OF_LITELLM_PARAMS} + ) _metadata = dict( LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict) ) - litellm_metadata: Final = litellm_params_in_body.pop("litellm_metadata", None) - metadata: Final = litellm_params_in_body.pop("metadata", None) + litellm_metadata: Final = litellm_keys_in_body.get("litellm_metadata") + metadata: Final = litellm_keys_in_body.get("metadata") if litellm_metadata: _metadata.update(litellm_metadata) if metadata: diff --git a/litellm/router.py b/litellm/router.py index 9bf8c410bcb..023b99cd64e 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -259,6 +259,7 @@ from litellm.router_utils.routing_groups import ( validate_routing_strategy, ) from litellm.scheduler import FlowItem, Scheduler +from litellm.types.litellm_params import RoutingStrategyName from litellm.types.llms.openai import ( AllMessageValues, ChatCompletionToolParam, @@ -796,15 +797,7 @@ class Router: allowed_fails_policy: AllowedFailsPolicy | None = None, # set custom allowed fails policy cooldown_time: float | None = None, # (seconds) time to cooldown a deployment after failure disable_cooldowns: bool | None = None, - routing_strategy: Literal[ - "simple-shuffle", - "least-busy", - "usage-based-routing", - "latency-based-routing", - "cost-based-routing", - "usage-based-routing-v2", - "lar1", - ] = "simple-shuffle", + routing_strategy: RoutingStrategyName = "simple-shuffle", optional_pre_call_checks: OptionalPreCallChecks | None = None, routing_strategy_args: dict = {}, # just for latency-based routing_groups: list[RoutingGroup | dict] | None = None, @@ -2940,7 +2933,7 @@ class Router: self._update_kwargs_before_fallbacks(model=model_group, kwargs=initial_kwargs) fallback_response = await self.async_function_with_fallbacks_common_utils( e=e, - disable_fallbacks=False, + disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs), fallbacks=fallbacks, context_window_fallbacks=context_window_fallbacks, content_policy_fallbacks=content_policy_fallbacks, @@ -3384,7 +3377,7 @@ class Router: ) fallback_response = await self.async_function_with_fallbacks_common_utils( e=fallback_trigger, - disable_fallbacks=False, + disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs), fallbacks=fallbacks, context_window_fallbacks=context_window_fallbacks, content_policy_fallbacks=content_policy_fallbacks, @@ -3475,8 +3468,9 @@ class Router: for item in model_response: yield item except MidStreamFallbackError as e: - if not e.is_pre_first_chunk and ( - e.generated_content or _stream_chunks_have_generated_content(model_response.chunks) + if fallbacks_disabled_for_request(initial_kwargs) or ( + not e.is_pre_first_chunk + and (e.generated_content or _stream_chunks_have_generated_content(model_response.chunks)) ): if e.original_exception is not None: raise e.original_exception from e @@ -5611,7 +5605,7 @@ class Router: ) fallback_response = await self.async_function_with_fallbacks_common_utils( # rebind-ok: set on success e=fallback_trigger, - disable_fallbacks=False, + disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs), fallbacks=fallbacks, context_window_fallbacks=context_window_fallbacks, content_policy_fallbacks=content_policy_fallbacks, diff --git a/litellm/types/integrations/custom_logger.py b/litellm/types/integrations/custom_logger.py index 5de58a20242..9a9f3ae34ce 100644 --- a/litellm/types/integrations/custom_logger.py +++ b/litellm/types/integrations/custom_logger.py @@ -3,8 +3,10 @@ from typing import Any, Final from pydantic import BaseModel, Field -CHAT_COMPLETION_AGENTIC_SURFACE: Final = "chat_completions" -RESPONSES_AGENTIC_SURFACE: Final = "responses" +from litellm.types.litellm_params import AgenticSurface + +CHAT_COMPLETION_AGENTIC_SURFACE: Final[AgenticSurface] = "chat_completions" +RESPONSES_AGENTIC_SURFACE: Final[AgenticSurface] = "responses" CODE_INTERPRETER_INTERCEPTION_PREFIX: Final = "_code_interpreter_interception" HEADROOM_INTERCEPTION_PREFIX: Final = "_headroom_interception" HEADROOM_CONVERTED_STREAM_KEY: Final = f"{HEADROOM_INTERCEPTION_PREFIX}_converted_stream" diff --git a/litellm/types/litellm_params.py b/litellm/types/litellm_params.py new file mode 100644 index 00000000000..83a42c235f9 --- /dev/null +++ b/litellm/types/litellm_params.py @@ -0,0 +1,364 @@ +"""LiteLLM-owned request kwargs declared as typed fields; types/utils.py splices these with the callback and pricing +models and KWARG_ARTIFACTS into all_litellm_params.""" + +from collections.abc import Callable, Iterator, Mapping, MutableMapping, Sequence +from dataclasses import dataclass, field, fields, is_dataclass +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, Literal, TypeAlias + +if TYPE_CHECKING: + import httpx + from aiohttp import ClientSession + from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI + + from litellm.litellm_core_utils.litellm_logging import Logging + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler + from litellm.router_strategy.complexity_router.context_compaction import CompactionState + from litellm.router_utils.fallback_event_handlers import AttemptedFallbackTargets + from litellm.types.caching import DynamicCacheControl + from litellm.types.llms.openai import ChatCompletionAssistantMessage, ChatCompletionUserMessage + from litellm.types.proxy.litellm_pre_call_utils import SecretFields + from litellm.types.router import ConfigurableClientsideParamsCustomAuth, DeploymentTypedDict, RetryPolicy + from litellm.types.router_weights import RouterWeights + from litellm.types.utils import ModelResponse, ModelResponseStream, ProviderSpecificHeader + + ProviderClient: TypeAlias = ( + OpenAI + | AsyncOpenAI + | AzureOpenAI + | AsyncAzureOpenAI + | HTTPHandler + | AsyncHTTPHandler + | httpx.Client + | httpx.AsyncClient + ) + MockResponse: TypeAlias = ( + str | Exception | Mapping[str, object] | Sequence[float] | ModelResponse | ModelResponseStream + ) + +RetryStrategy: TypeAlias = Literal["constant_retry", "exponential_backoff_retry"] +AgenticSurface: TypeAlias = Literal["chat_completions", "responses"] +RoutingStrategyName: TypeAlias = Literal[ + "simple-shuffle", + "least-busy", + "usage-based-routing", + "latency-based-routing", + "cost-based-routing", + "usage-based-routing-v2", + "lar1", +] + +TRUSTED_CALLBACK_VARS_FIELD: Final = "litellm_trusted_callback_vars" +ADDRESSED_RESPONSE_ID_FIELD: Final = "_litellm_addressed_response_id" + +WIRE_NAME: Final = "wire_name" + + +def wire(name: str) -> Mapping[str, str]: + return MappingProxyType({WIRE_NAME: name}) + + +@dataclass(frozen=True, slots=True, kw_only=True) +class ProviderConnection: + api_key: str | None = None + api_base: str | None = None + api_version: str | None = None + region_name: str | None = None + headers: Mapping[str, str] | None = None + provider_specific_header: "ProviderSpecificHeader | Sequence[ProviderSpecificHeader] | None" = None + client: "ProviderClient | None" = None + shared_session: "ClientSession | None" = None + ssl_verify: bool | str | None = None + request_timeout: float | None = None + force_timeout: float | None = None + stream_timeout: float | str | None = None + max_retries: int | None = None + tenant_id: str | None = None + client_id: str | None = None + client_secret: str | None = None + azure_username: str | None = None + azure_password: str | None = None + azure_scope: str | None = None + azure_ad_token_provider: Callable[[], str] | None = None + litellm_credential_name: str | None = None + configurable_clientside_auth_params: "Sequence[str | ConfigurableClientsideParamsCustomAuth] | None" = None + use_xai_oauth: bool | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class BedrockBatchConnection: + # Bedrock rejects these names in request bodies, so register them as LiteLLM-owned + aws_batch_role_arn: str | None = None + s3_bucket_name: str | None = None + s3_region_name: str | None = None + s3_endpoint_url: str | None = None + s3_output_bucket_name: str | None = None + s3_bucket_owner: str | None = None + s3_access_key_id: str | None = None + s3_secret_access_key: str | None = None + s3_encryption_key_id: str | None = None + bedrock_tags: Sequence[Mapping[str, str]] | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class ConnectionSettings: + provider: ProviderConnection + bedrock_batch: BedrockBatchConnection + + +@dataclass(frozen=True, slots=True, kw_only=True) +class DispatchOptions: + custom_llm_provider: str | None = None + azure: bool | None = None + use_litellm_proxy: bool | None = None + use_chat_completions_api: bool | None = None + use_in_pass_through: bool | None = None + allowed_openai_params: Sequence[str] | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class RoutingOptions: + fallbacks: Sequence[str | Mapping[str, object]] | None = None + context_window_fallback_dict: Mapping[str, str] | None = None + num_retries: int | None = None + retry_policy: "RetryPolicy | Mapping[str, object] | None" = None + retry_strategy: RetryStrategy | None = None + routing_strategy: RoutingStrategyName | None = None + cooldown_time: float | None = None + allowed_model_region: str | None = None + enable_tag_filtering: bool | None = None + fastest_response: bool | None = None + provider_affinity_header: str | None = None + search_tool_name: str | None = None + model_list: "Sequence[DeploymentTypedDict] | None" = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class DeploymentOptions: + model_info: Mapping[str, object] | None = None + rpm: int | None = None + tpm: int | None = None + itpm: int | None = None + otpm: int | None = None + default_api_key_rpm_limit: int | None = None + default_api_key_tpm_limit: int | None = None + max_parallel_requests: int | None = None + weight: int | None = None + order: int | None = None + tag_regex: Sequence[str] | None = None + max_file_size_mb: float | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class SpecializedRouterOptions: + auto_router_config_path: str | None = None + auto_router_config: str | None = None + auto_router_default_model: str | None = None + auto_router_embedding_model: str | None = None + auto_router_max_input_chars: int | None = None + auto_router_routing_compression: str | None = None + auto_router_model_compression: str | None = None + complexity_router_config: Mapping[str, object] | None = None + complexity_router_default_model: str | None = None + adaptive_router_config: Mapping[str, object] | None = None + adaptive_router_default_model: str | None = None + quality_router_config: Mapping[str, object] | None = None + quality_router_default_model: str | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class CachingOptions: + caching: bool | None = None + cache: "DynamicCacheControl | None" = None + ttl: float | None = None + enable_prompt_caching: bool | None = None + caching_groups: Sequence[Sequence[str]] | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class CostOptions: + cost_per_query: float | None = None + base_model: str | None = None + max_budget: float | None = None + budget_duration: str | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class ObservabilityOptions: + id: str | None = None + metadata: MutableMapping[str, object] | None = None # mutable-ok: the router and logging write keys into it + litellm_metadata: MutableMapping[str, object] | None = None # mutable-ok: the proxy writes keys into it + tags: Sequence[str] | None = None + litellm_trace_id: str | None = None + litellm_session_id: str | None = None + litellm_request_debug: bool | None = None + logger_fn: Callable[[Mapping[str, object]], None] | None = None + verbose: bool | None = None + no_log: bool | None = field(default=None, metadata=wire("no-log")) + + +@dataclass(frozen=True, slots=True, kw_only=True) +class AgenticLoopOptions: + max_agentic_loops: int | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class GuardrailOptions: + guardrails: Sequence[str] | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class PromptOptions: + prompt_id: str | None = None + prompt_variables: Mapping[str, object] | None = None + prompt_version: str | None = None + prompt_environment: str | None = None + prompt_label: str | None = None + litellm_system_prompt: str | None = None + custom_prompt_dict: Mapping[str, object] | None = None + roles: Mapping[str, object] | None = None + final_prompt_value: str | None = None + bos_token: str | None = None + eos_token: str | None = None + hf_model_name: str | None = None + supports_system_message: bool | None = None + ensure_alternating_roles: bool | None = None + user_continue_message: "ChatCompletionUserMessage | None" = None + assistant_continue_message: "ChatCompletionAssistantMessage | None" = None + disable_add_transform_inline_image_block: bool | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class ResponseOptions: + merge_reasoning_content_in_choices: bool | None = None + enable_json_schema_validation: bool | None = None + complete_response: bool | None = None + stream_chunk_size: int | None = None + keepalive_seconds: float | None = None + allow_client_keepalive_override: bool | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class MockOptions: + mock_response: "MockResponse | None" = None + mock_timeout: bool | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class LiteLLMOptions: + dispatch: DispatchOptions + routing: RoutingOptions + deployment: DeploymentOptions + specialized_routers: SpecializedRouterOptions + caching: CachingOptions + cost: CostOptions + observability: ObservabilityOptions + agentic_loop: AgenticLoopOptions + guardrails: GuardrailOptions + prompt: PromptOptions + response: ResponseOptions + mock: MockOptions + + +@dataclass(frozen=True, slots=True, kw_only=True) +class CallState: + litellm_call_id: str | None = None + completion_call_id: str | None = None + model_alias_map: Mapping[str, str] | None = None + data_residency: str | None = None + litellm_logging_obj: "Logging | None" = None + preset_cache_key: str | None = None + cache_key: str | None = None + stream_response: "Mapping[str, ModelResponse] | None" = None + context_compaction_state: "CompactionState | None" = field(default=None, metadata=wire("_context_compaction_state")) + + +@dataclass(frozen=True, slots=True, kw_only=True) +class AgenticLoopState: + depth: int | None = field(default=None, metadata=wire("_agentic_loop_depth")) + fingerprints: Sequence[str] | None = field(default=None, metadata=wire("_agentic_loop_fingerprints")) + api_surface: Literal["chat_completions", "responses"] | None = field( + default=None, metadata=wire("_agentic_loop_api_surface") + ) + code_interpreter_active: bool | None = field(default=None, metadata=wire("_code_interpreter_interception_active")) + code_interpreter_sandbox_key: str | None = field( + default=None, metadata=wire("_code_interpreter_interception_sandbox_key") + ) + code_interpreter_session_scoped: bool | None = field( + default=None, metadata=wire("_code_interpreter_interception_session_scoped") + ) + code_interpreter_converted_stream: bool | None = field( + default=None, metadata=wire("_code_interpreter_interception_converted_stream") + ) + websearch_emit_native_blocks: bool | None = field( + default=None, metadata=wire("_websearch_interception_emit_native_blocks") + ) + websearch_converted_stream: bool | None = field( + default=None, metadata=wire("_websearch_interception_converted_stream") + ) + headroom_converted_stream: bool | None = field( + default=None, metadata=wire("_headroom_interception_converted_stream") + ) + + +@dataclass(frozen=True, slots=True, kw_only=True) +class RouterState: + weights: "RouterWeights | None" = field(default=None, metadata=wire("_router_weights")) + fallback_depth: int | None = None + max_fallbacks: int | None = None + attempted_targets: "AttemptedFallbackTargets | None" = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class ProxyRequestState: + proxy_server_request: Mapping[str, object] | None = None + secret_fields: "SecretFields | None" = None + trusted_callback_vars: Mapping[str, str] | None = field(default=None, metadata=wire(TRUSTED_CALLBACK_VARS_FIELD)) + addressed_response_id: str | None = field(default=None, metadata=wire(ADDRESSED_RESPONSE_ID_FIELD)) + strip_stream_usage: bool | None = field(default=None, metadata=wire("_litellm_strip_stream_usage")) + client_side_timeout: bool | None = None + model_file_id_mapping: Mapping[str, Mapping[str, str]] | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class EntrypointState: + acompletion: bool | None = None + aembedding: bool | None = None + aimg_generation: bool | None = None + atext_completion: bool | None = None + text_completion: bool | None = None + allm_passthrough_route: bool | None = None + async_call: bool | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class InternalState: + call: CallState + agentic_loop: AgenticLoopState + router: RouterState + proxy: ProxyRequestState + entrypoint: EntrypointState + + +KWARG_ARTIFACTS: Final[tuple[str, ...]] = ("self", "use_client", "model_config", "rust") + +LITELLM_OWNED_ROOTS: Final = (ConnectionSettings, LiteLLMOptions, InternalState) + + +def wire_names(owner: type) -> tuple[str, ...]: + return tuple(owned.metadata.get(WIRE_NAME, owned.name) for owned in fields(owner)) + + +def owned_wire_names(root: type) -> tuple[str, ...]: + def names() -> Iterator[str]: + for leaf in fields(root): + if not is_dataclass(leaf.type): + raise TypeError(f"{root.__name__}.{leaf.name} is not a dataclass leaf") + yield from wire_names(leaf.type) # pyright: ignore[reportArgumentType] # Field.type admits str + + return tuple(names()) + + +OWNED_KWARG_NAMES: Final = tuple(name for root in LITELLM_OWNED_ROOTS for name in owned_wire_names(root)) +AGENTIC_LOOP_KWARG_NAMES: Final = (*wire_names(AgenticLoopState), *wire_names(AgenticLoopOptions)) +BEDROCK_BATCH_KWARG_NAMES: Final = wire_names(BedrockBatchConnection) diff --git a/litellm/types/router.py b/litellm/types/router.py index c0f724584fd..b72809f625f 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -24,6 +24,7 @@ if TYPE_CHECKING: from .completion import CompletionRequest from .embedding import EmbeddingRequest +from .litellm_params import RoutingStrategyName from .llms.bedrock import AwsSessionTag from .llms.openai import OpenAIFileObject from .search import SearchProvider @@ -104,12 +105,7 @@ class RouterConfig(BaseModel): context_window_fallbacks: list | None = [] model_group_alias: dict[str, list[str]] | None = {} retry_after: int | None = 0 - routing_strategy: Literal[ - "simple-shuffle", - "least-busy", - "usage-based-routing", - "latency-based-routing", - ] = "simple-shuffle" + routing_strategy: RoutingStrategyName = "simple-shuffle" routing_groups: list[RoutingGroup] | None = None model_config = ConfigDict(protected_namespaces=()) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index caf88e5d517..7aaf11faa5d 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -56,8 +56,15 @@ from litellm.types.llms.base import ( from litellm.types.mcp import MCPServerCostInfo from ..litellm_core_utils.core_helpers import map_finish_reason, process_response_headers +from . import litellm_params as _litellm_params from .agents import LiteLLMSendMessageResponse from .guardrails import GuardrailEventHooks +from .litellm_params import ( + AGENTIC_LOOP_KWARG_NAMES, + BEDROCK_BATCH_KWARG_NAMES, + KWARG_ARTIFACTS, + OWNED_KWARG_NAMES, +) from .llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse from .llms.base import HiddenParams from .llms.openai import ( @@ -3901,205 +3908,20 @@ def pricing_override_fields(*sources: Mapping[str, object]) -> tuple[str, ...]: ) -# Server-controlled fields that bound or drive an interceptor's agentic loop -# (depth, cycle fingerprints, ceiling, code-interpreter sandbox state). Listed -# in all_litellm_params so they are treated as LiteLLM-level and excluded from -# get_non_default_completion_params; otherwise the OpenAI param builder sweeps -# any unrecognized top-level key into extra_body and leaks them to the provider. -# This is what lets the loop carry state across rerun calls without a provider -# scrubber. -agentic_loop_internal_litellm_params: Final = [ - "_agentic_loop_depth", - "_agentic_loop_fingerprints", - "_agentic_loop_api_surface", - "max_agentic_loops", - "_code_interpreter_interception_active", - "_code_interpreter_interception_sandbox_key", - "_code_interpreter_interception_session_scoped", - "_code_interpreter_interception_converted_stream", - "_websearch_interception_emit_native_blocks", - "_websearch_interception_converted_stream", - "_headroom_interception_converted_stream", +agentic_loop_internal_litellm_params: Final = list(AGENTIC_LOOP_KWARG_NAMES) # mutable-ok: public type stays a list + +bedrock_batch_litellm_params: Final = BEDROCK_BATCH_KWARG_NAMES + +TRUSTED_CALLBACK_VARS_FIELD: Final = _litellm_params.TRUSTED_CALLBACK_VARS_FIELD +ADDRESSED_RESPONSE_ID_FIELD: Final = _litellm_params.ADDRESSED_RESPONSE_ID_FIELD + +all_litellm_params = [ # rebind-ok: two star imports in litellm/__init__.py re-bind it # mutable-ok: callers concat + *OWNED_KWARG_NAMES, + *KWARG_ARTIFACTS, + *StandardCallbackDynamicParams.__annotations__, + *CustomPricingLiteLLMParams.model_fields, ] -# Proxy-owned callback credentials, stamped from admin-configured team/key callback -# settings. Listed in all_litellm_params for the same reason as the agentic-loop -# fields above: an unrecognized top-level key is swept into extra_body and sent to -# the provider. -TRUSTED_CALLBACK_VARS_FIELD: Final = "litellm_trusted_callback_vars" - -ADDRESSED_RESPONSE_ID_FIELD: Final = "_litellm_addressed_response_id" - -# Bedrock managed-batch deployment config, read from litellm_params by the batch and -# files transformations. Listed for the same reason as the fields above: these sit on -# a deployment that also serves chat, so leaking them into extra_body makes Bedrock -# reject every non-batch request to that deployment. -bedrock_batch_litellm_params: Final = ( - "aws_batch_role_arn", - "s3_bucket_name", - "s3_region_name", - "s3_endpoint_url", - "s3_output_bucket_name", - "s3_bucket_owner", - "s3_access_key_id", - "s3_secret_access_key", - "s3_encryption_key_id", - "bedrock_tags", -) - -all_litellm_params = ( - agentic_loop_internal_litellm_params - + [TRUSTED_CALLBACK_VARS_FIELD, ADDRESSED_RESPONSE_ID_FIELD, *bedrock_batch_litellm_params] - + [ - "_context_compaction_state", - "metadata", - "litellm_metadata", - "keepalive_seconds", - "allow_client_keepalive_override", - "litellm_trace_id", - "litellm_request_debug", - "guardrails", - "tags", - "acompletion", - "aimg_generation", - "atext_completion", - "text_completion", - "caching", - "mock_response", - "mock_timeout", - "disable_add_transform_inline_image_block", - "api_key", - "api_version", - "prompt_id", - "prompt_variables", - "litellm_system_prompt", - "provider_specific_header", - "prompt_version", - "prompt_environment", - "api_base", - "force_timeout", - "logger_fn", - "verbose", - "custom_llm_provider", - "model_file_id_mapping", - "litellm_logging_obj", - "litellm_call_id", - "completion_call_id", - "model_alias_map", - "custom_prompt_dict", - "stream_response", - "cost_per_query", - "ssl_verify", - "data_residency", - "async_call", - "aembedding", - "allm_passthrough_route", - "_litellm_strip_stream_usage", - "use_client", - "id", - "fallbacks", - "routing_strategy", - "_router_weights", - "azure", - "headers", - "model_list", - "num_retries", - "context_window_fallback_dict", - "retry_policy", - "retry_strategy", - "roles", - "final_prompt_value", - "bos_token", - "eos_token", - "request_timeout", - "client_side_timeout", - "complete_response", - "self", - "client", - "rpm", - "tpm", - "default_api_key_rpm_limit", - "default_api_key_tpm_limit", - "itpm", - "otpm", - "max_parallel_requests", - "input_cost_per_token", - "output_cost_per_token", - "input_cost_per_second", - "output_cost_per_second", - "hf_model_name", - "model_info", - "proxy_server_request", - "secret_fields", - "preset_cache_key", - "caching_groups", - "ttl", - "cache", - "enable_prompt_caching", - "no-log", - "base_model", - "stream_timeout", - "stream_chunk_size", - "supports_system_message", - "region_name", - "allowed_model_region", - "model_config", - "fastest_response", - "cooldown_time", - "cache_key", - "max_retries", - "azure_ad_token_provider", - "tenant_id", - "client_id", - "azure_username", - "azure_password", - "azure_scope", - "client_secret", - "user_continue_message", - "configurable_clientside_auth_params", - "weight", - "ensure_alternating_roles", - "assistant_continue_message", - "user_continue_message", - "fallback_depth", - "max_fallbacks", - "attempted_targets", - "max_budget", - "budget_duration", - "use_in_pass_through", - "merge_reasoning_content_in_choices", - "litellm_credential_name", - "allowed_openai_params", - "litellm_session_id", - "provider_affinity_header", - "use_litellm_proxy", - "use_chat_completions_api", - "rust", - "prompt_label", - "shared_session", - "search_tool_name", - "order", - "enable_tag_filtering", - "enable_json_schema_validation", - "use_xai_oauth", - "auto_router_config_path", - "auto_router_config", - "auto_router_default_model", - "auto_router_embedding_model", - "auto_router_max_input_chars", - "auto_router_routing_compression", - "auto_router_model_compression", - "complexity_router_config", - "complexity_router_default_model", - "adaptive_router_config", - "adaptive_router_default_model", - "quality_router_config", - "quality_router_default_model", - ] - + list(StandardCallbackDynamicParams.__annotations__.keys()) - + list(CustomPricingLiteLLMParams.model_fields.keys()) -) - class KeyGenerationConfig(TypedDict, total=False): required_params: list[str] # specify params that must be present in the key generation request diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index b529a225f84..f875936b0bf 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -3342,6 +3342,7 @@ "supports_vision": true }, "azure/command-r-plus": { + "deprecation_date": "2025-06-30", "input_cost_per_token": 3e-06, "litellm_provider": "azure", "max_input_tokens": 128000, @@ -3349,6 +3350,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.5e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_function_calling": true }, "azure_ai/claude-haiku-4-5": { @@ -5194,6 +5196,7 @@ "output_cost_per_token": 2e-06 }, "azure/gpt-4": { + "deprecation_date": "2025-06-06", "input_cost_per_token": 3e-05, "litellm_provider": "azure", "max_input_tokens": 8192, @@ -5201,6 +5204,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_function_calling": true, "supports_tool_choice": true }, @@ -69074,6 +69078,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/google/gemma-4-31B-it": { + "deprecation_date": "2026-09-14", "input_cost_per_token": 3.9e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -69203,6 +69208,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/openai/gpt-oss-20b": { + "deprecation_date": "2026-09-14", "input_cost_per_token": 5e-08, "litellm_provider": "together_ai", "max_input_tokens": 131072, diff --git a/pyproject.toml b/pyproject.toml index ba72378989a..15eb8f0c4f4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -75,8 +75,8 @@ proxy = [ "mcp>=2.2.0,<3", "httpx2>=2.5.0,<3", "pydantic>=2.12.0,<3", - "litellm-proxy-extras==0.4.101", - "litellm-enterprise==0.1.70", + "litellm-proxy-extras==0.4.102", + "litellm-enterprise==0.1.71", "RestrictedPython>=8.5,<9.0", "rich>=13.9.4,<14.0", "InquirerPy>=0.3.4,<1.0", diff --git a/tests/integration/authorization/test_key_alias_model_access.py b/tests/integration/authorization/test_key_alias_model_access.py new file mode 100644 index 00000000000..50fc53bd4a9 --- /dev/null +++ b/tests/integration/authorization/test_key_alias_model_access.py @@ -0,0 +1,72 @@ +import uuid +from typing import Final + +import httpx + +from tests.integration._support.client import Gateway, eventually, object_value, string_value + + +def _listed_model_ids(response: httpx.Response) -> frozenset[str]: + entries: Final = response.json()["data"] + assert isinstance(entries, list), response.text + return frozenset(string_value(object_value(entry)["id"]) for entry in entries) + + +def _listed_and_callable(gateway: Gateway, key: str, model: str, alias: str) -> None: + """Every id /v1/models lists for this key must be callable by the same key.""" + response: Final = eventually( + lambda: gateway.request("GET", "/v1/models", key=key), + lambda value: value.status_code == 200 and _listed_model_ids(value) == frozenset({model, alias}), + return_last_on_timeout=True, + ) + assert response.status_code == 200, response.text + listed: Final = _listed_model_ids(response) + assert listed == frozenset({model, alias}), response.text + for model_id in sorted(listed): + called: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model_id, "messages": [{"role": "user", "content": "ping"}]}, + key=key, + ) + assert called.status_code == 200, f"listed id {model_id} is not callable: {called.status_code} {called.text}" + + +def test_key_alias_listed_by_v1_models_is_callable(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + alias: Final = f"integration-alias-{uuid.uuid4().hex}" + key: Final = scenario.key(models=[model], aliases={alias: model}) + _listed_and_callable(gateway, key, model, alias) + + +def test_team_key_alias_listed_by_v1_models_is_callable(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + team_id: Final = scenario.team(models=[model]) + alias: Final = f"integration-alias-{uuid.uuid4().hex}" + key: Final = scenario.key(team_id=team_id, aliases={alias: model}) + _listed_and_callable(gateway, key, model, alias) + + +def test_key_alias_to_model_outside_key_allowlist_is_hidden_and_denied(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + allowed: Final = scenario.model() + hidden: Final = scenario.model() + alias: Final = f"integration-alias-{uuid.uuid4().hex}" + key: Final = scenario.key(models=[allowed], aliases={alias: hidden}) + response: Final = eventually( + lambda: gateway.request("GET", "/v1/models", key=key), + lambda value: value.status_code == 200 and _listed_model_ids(value) == frozenset({allowed}), + return_last_on_timeout=True, + ) + assert response.status_code == 200, response.text + assert _listed_model_ids(response) == frozenset({allowed}), response.text + called: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": alias, "messages": [{"role": "user", "content": "ping"}]}, + key=key, + ) + assert called.status_code == 403, called.text + assert "key_model_access_denied" in called.text, called.text diff --git a/tests/integration/pricing/test_configured_prices.py b/tests/integration/pricing/test_configured_prices.py index 0e4efea3a15..93290a404fc 100644 --- a/tests/integration/pricing/test_configured_prices.py +++ b/tests/integration/pricing/test_configured_prices.py @@ -1,6 +1,9 @@ +import asyncio import json +import os import uuid from collections.abc import Iterator, Mapping +from hashlib import sha256 from pathlib import Path from typing import Final @@ -12,6 +15,96 @@ from litellm import get_model_info from tests.integration._support.client import Gateway, eventually, object_value, string_value from tests.integration._support.database import read_rows from tests.integration._support.process import owned_proxy +from tests.integration._support.upstream import delete_scenario, register_scenario +from tests.integration.cost_calculation.cost_tracking_case import RealtimeResponse +from tests.integration.pricing.test_realtime_cached_audio_pricing import one_realtime_turn + +REALTIME_MODEL: Final = "gpt-realtime-2" +REALTIME_INPUT_TEXT_TOKENS: Final = 10 +REALTIME_INPUT_AUDIO_TOKENS: Final = 20 +REALTIME_OUTPUT_TEXT_TOKENS: Final = 5 +REALTIME_OUTPUT_AUDIO_TOKENS: Final = 7 + + +def _realtime_response_done() -> RealtimeResponse: + return RealtimeResponse( + content_type="application/x-realtime", + events=( + { + "type": "response.done", + "event_id": "evt_$REQUEST_ID", + "response": { + "id": "resp_$REQUEST_ID", + "object": "realtime.response", + "status": "completed", + "output": [], + "usage": { + "total_tokens": REALTIME_INPUT_TEXT_TOKENS + + REALTIME_INPUT_AUDIO_TOKENS + + REALTIME_OUTPUT_TEXT_TOKENS + + REALTIME_OUTPUT_AUDIO_TOKENS, + "input_tokens": REALTIME_INPUT_TEXT_TOKENS + REALTIME_INPUT_AUDIO_TOKENS, + "output_tokens": REALTIME_OUTPUT_TEXT_TOKENS + REALTIME_OUTPUT_AUDIO_TOKENS, + "input_token_details": { + "text_tokens": REALTIME_INPUT_TEXT_TOKENS, + "audio_tokens": REALTIME_INPUT_AUDIO_TOKENS, + "cached_tokens": 0, + }, + "output_token_details": { + "text_tokens": REALTIME_OUTPUT_TEXT_TOKENS, + "audio_tokens": REALTIME_OUTPUT_AUDIO_TOKENS, + }, + }, + }, + }, + ), + ) + + +@pytest.mark.parametrize( + ("input_text_rate", "input_audio_rate", "output_text_rate", "output_audio_rate"), + ((0.001, 0.002, 0.003, 0.004), (0.0, 0.0, 0.0, 0.0)), + ids=("custom_rates", "zero_rated"), +) +def test_realtime_session_is_charged_at_the_deployment_configured_rates( + gateway: Gateway, + input_text_rate: float, + input_audio_rate: float, + output_text_rate: float, + output_audio_rate: float, +) -> None: + with gateway.scenario() as scenario: + scenario_id: Final = f"realtime-configured-price-{uuid.uuid4().hex[:12]}" + handle: Final = register_scenario(scenario_id, _realtime_response_done()) + scenario.cleanups.callback(delete_scenario, handle) + key: Final = scenario.key() + model: Final = scenario.model( + model=f"openai/{REALTIME_MODEL}", + api_key=scenario_id, + api_base=gateway.upstream_url.rstrip("/"), + input_cost_per_token=input_text_rate, + input_cost_per_audio_token=input_audio_rate, + output_cost_per_token=output_text_rate, + output_cost_per_audio_token=output_audio_rate, + ) + session: Final = asyncio.run(one_realtime_turn(os.environ["INTEGRATION_PROXY_URL"].rstrip("/"), key, model)) + assert session.get("type") == "session.created", session + rows: Final = eventually( + lambda: read_rows( + 'SELECT spend, call_type FROM "LiteLLM_SpendLogs" WHERE api_key = %s', + (sha256(key.encode()).hexdigest(),), + ), + lambda values: len(values) == 1, + seconds=70, + ) + assert rows[0]["call_type"] == "_arealtime", rows + assert float(str(rows[0]["spend"])) == pytest.approx( + REALTIME_INPUT_TEXT_TOKENS * input_text_rate + + REALTIME_INPUT_AUDIO_TOKENS * input_audio_rate + + REALTIME_OUTPUT_TEXT_TOKENS * output_text_rate + + REALTIME_OUTPUT_AUDIO_TOKENS * output_audio_rate, + abs=1e-9, + ), rows @pytest.mark.covers("quota_management.spend_tracking.custom_price.matches_input_rates") diff --git a/tests/integration/pricing/test_realtime_cached_audio_pricing.py b/tests/integration/pricing/test_realtime_cached_audio_pricing.py index 4a7598d0cbf..42e90cf2c3e 100644 --- a/tests/integration/pricing/test_realtime_cached_audio_pricing.py +++ b/tests/integration/pricing/test_realtime_cached_audio_pricing.py @@ -92,7 +92,7 @@ def cached_audio_response_done() -> RealtimeResponse: ) -async def _one_realtime_turn(proxy_url: str, key: str, model: str) -> dict[str, JsonValue]: +async def one_realtime_turn(proxy_url: str, key: str, model: str) -> dict[str, JsonValue]: async with websockets.connect( f"{proxy_url.replace('http://', 'ws://').replace('https://', 'wss://')}/v1/realtime?model={model}", additional_headers={"Authorization": f"Bearer {key}"}, @@ -115,7 +115,7 @@ def test_realtime_cached_audio_tokens_bill_at_audio_cache_read_rate_not_full_aud model: Final = scenario.model( model=f"openai/{MODEL}", api_key=scenario_id, api_base=gateway.upstream_url.rstrip("/") ) - session: Final = asyncio.run(_one_realtime_turn(os.environ["INTEGRATION_PROXY_URL"].rstrip("/"), key, model)) + session: Final = asyncio.run(one_realtime_turn(os.environ["INTEGRATION_PROXY_URL"].rstrip("/"), key, model)) assert session.get("type") == "session.created", session rows: Final = eventually( lambda: read_rows( diff --git a/tests/test_keys.py b/tests/test_keys.py index 7a5b2502cfd..c1785b88822 100644 --- a/tests/test_keys.py +++ b/tests/test_keys.py @@ -834,12 +834,12 @@ async def test_key_model_list(model_access, model_access_level, model_endpoint): assert len(model_list["data"]) > 0 if model_access == "gpt-3.5-turbo": if model_endpoint == "/v1/models": - assert ( - len(model_list["data"]) == 1 - ), "model_access={}, model_access_level={}".format( + assert {entry["id"] for entry in model_list["data"]} == { + model_access, + "mistral-7b", + }, "generate_key sets alias mistral-7b -> gpt-3.5-turbo, so /v1/models lists both; model_access={}, model_access_level={}".format( model_access, model_access_level ) - assert model_list["data"][0]["id"] == model_access elif model_endpoint == "/model/info": assert isinstance(model_list["data"], list) assert len(model_list["data"]) == 1 diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py index d03174bc2c6..2fe22ba2620 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py @@ -1046,11 +1046,26 @@ def test_translate_anthropic_to_openai_sets_prompt_cache_key_for_azure(model: st assert openai_request["prompt_cache_key"] == "session-abc" +@pytest.mark.parametrize( + "model", + [ + "vertex_ai/meta/llama-4-maverick-17b-128e-instruct-maas", + "vertex_ai/moonshotai/kimi-k2-thinking-maas", + "vertex_ai/xai/grok-4.1-fast-non-reasoning", + ], +) +def test_translate_anthropic_to_openai_sets_prompt_cache_key_for_vertex_maas_models(model: str): + openai_request = _translate_with_metadata(model, {"user_id": CLAUDE_CODE_USER_ID}, "vertex_ai") + assert openai_request["prompt_cache_key"] == "session-abc" + + @pytest.mark.parametrize( "model, custom_llm_provider", [ ("gemini/gemini-2.5-pro", "gemini"), ("vertex_ai/gemini-2.5-pro", "vertex_ai"), + ("vertex_ai/gemma/gemma-2-2b-it", "vertex_ai"), + ("vertex_ai/openai/mg-endpoint-lit8592", "vertex_ai"), ("anthropic/claude-sonnet-4-5", "anthropic"), ("bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", "bedrock"), ("no-such-model-lit5875", "no-such-provider-lit5875"), diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py b/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py index 0dcaa4c72c2..b80d4714253 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py @@ -8,10 +8,10 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest import litellm -from litellm.llms.vertex_ai.vertex_model_garden.main import ( - _vertex_model_garden_model_id_in_json_body, - create_vertex_url, +from litellm.llms.vertex_ai.common_utils import ( + vertex_model_garden_model_id_in_json_body, ) +from litellm.llms.vertex_ai.vertex_model_garden.main import create_vertex_url @pytest.mark.parametrize( @@ -43,11 +43,8 @@ def test_create_vertex_url_openapi_vs_deployed_endpoint( def test_model_id_in_json_body_heuristic() -> None: - assert ( - _vertex_model_garden_model_id_in_json_body("xai/grok-4.1-fast-reasoning") - is True - ) - assert _vertex_model_garden_model_id_in_json_body("5464397967697903616") is False + assert vertex_model_garden_model_id_in_json_body("xai/grok-4.1-fast-reasoning") is True + assert vertex_model_garden_model_id_in_json_body("5464397967697903616") is False @pytest.fixture diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 685d4f4fe78..2b37abeb1ee 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -346,6 +346,15 @@ class TestMCPRequestHandler: mock_manager.discovered_inventory = MagicMock(return_value=inventory or {}) return mock_manager + def _real_manager_with_toolsets(self, toolset_perms): + """A real MCPServerManager so the real expand_tool_permissions runs; + only the DB-backed toolset lookup is stubbed""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + + manager = MCPServerManager() + manager.resolve_toolset_tool_permissions = AsyncMock(return_value=toolset_perms) + return manager + async def test_get_allowed_mcp_servers_for_key_includes_toolset_servers(self): """A key granted only mcp_toolsets must reach the toolset's servers on every path (list, call, REST); regression for the list-ok/call-403 bug""" @@ -538,6 +547,141 @@ class TestMCPRequestHandler: assert result is None + @pytest.mark.parametrize( + "direct,via_toolsets,expected", + [ + (["*"], None, None), + (["*"], ["read_file"], None), + (None, None, None), + ([], None, ()), + (None, ["read_file"], ("read_file",)), + ], + ) + def test_union_tool_grants_wildcard_and_union_cases(self, direct, via_toolsets, expected): + """A direct ["*"] makes the level unrestricted even beside a toolset + list (regression: mapping ["*"] to None in expand_tool_permissions let + a same-level toolset list deny every other tool)""" + result = MCPRequestHandler._union_tool_grants(direct, via_toolsets) + + if expected is None: + assert result is None + else: + assert result is not None + assert set(result) == set(expected) + + def test_union_tool_grants_unions_two_concrete_lists(self): + result = MCPRequestHandler._union_tool_grants(["read_file"], ["write_file"]) + + assert result is not None + assert set(result) == {"read_file", "write_file"} + + async def test_key_wildcard_allows_a_tool_never_enumerated(self): + """End to end at the key level: object_permission sits on the auth + object already, no team named, so no patching is needed; the real + global manager expands ["*"] and the level reads unrestricted""" + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="perm-1", + mcp_tool_permissions={"server-a": ["*"]}, + ), + ) + + allowed = await MCPRequestHandler.get_allowed_tools_for_server( + server_id="server-a", user_api_key_auth=user_api_key_auth + ) + brand_new_tool_allowed = await MCPRequestHandler.is_tool_allowed_for_server( + tool_name="brand_new_tool", server_id="server-a", user_api_key_auth=user_api_key_auth + ) + + assert allowed is None + assert brand_new_tool_allowed is True + + async def test_key_wildcard_stays_capped_by_team_allowlist(self): + """A wildcard on the key must never widen a team's enumerated ceiling: + the intersection keeps only the team's named tools""" + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + team_id="team-1", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="perm-1", + mcp_tool_permissions={"server-a": ["*"]}, + ), + ) + team_object_permission = self._toolset_only_object_permission([]) + team_object_permission.mcp_tool_permissions = {"server-a": ["read_file"]} + manager = self._real_manager_with_toolsets({}) + + with ( + patch.object( # test-quality-ok: stub the DB team loader to drive the real team-server resolution path + MCPRequestHandler, "_get_team_object_permission", AsyncMock(return_value=team_object_permission) + ), + patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + manager, + ), + ): + allowed = await MCPRequestHandler.get_allowed_tools_for_server( + server_id="server-a", user_api_key_auth=user_api_key_auth + ) + brand_new_tool_allowed = await MCPRequestHandler.is_tool_allowed_for_server( + tool_name="brand_new_tool", server_id="server-a", user_api_key_auth=user_api_key_auth + ) + + assert allowed == ["read_file"] + assert brand_new_tool_allowed is False + + async def test_team_wildcard_stays_capped_by_key_allowlist(self): + """A wildcard on the team leaves the key's enumerated list as the + effective ceiling""" + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + team_id="team-1", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="perm-1", + mcp_tool_permissions={"server-a": ["read_file"]}, + ), + ) + team_object_permission = self._toolset_only_object_permission([]) + team_object_permission.mcp_tool_permissions = {"server-a": ["*"]} + manager = self._real_manager_with_toolsets({}) + + with ( + patch.object( # test-quality-ok: stub the DB team loader to drive the real team-server resolution path + MCPRequestHandler, "_get_team_object_permission", AsyncMock(return_value=team_object_permission) + ), + patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + manager, + ), + ): + allowed = await MCPRequestHandler.get_allowed_tools_for_server( + server_id="server-a", user_api_key_auth=user_api_key_auth + ) + + assert allowed == ["read_file"] + + async def test_key_empty_tool_list_stays_deny_all(self): + """[] on the key is deny-all, distinct from the wildcard: it must not + be widened into allow-all""" + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="perm-1", + mcp_tool_permissions={"server-a": []}, + ), + ) + + allowed = await MCPRequestHandler.get_allowed_tools_for_server( + server_id="server-a", user_api_key_auth=user_api_key_auth + ) + read_file_allowed = await MCPRequestHandler.is_tool_allowed_for_server( + tool_name="read_file", server_id="server-a", user_api_key_auth=user_api_key_auth + ) + + assert allowed == [] + assert read_file_allowed is False + # ------------------------------------------------------------------ # LIT-5749: toolsets attached to a TEAM, ORG, or internal USER must be # enforced exactly like inline tool allowlists, on both axes diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index f7f114dda30..d7aa2f5c3c9 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -8783,6 +8783,35 @@ class TestMCPServerManagerExpandToolPermissions: result = manager.expand_tool_permissions({"uuid-a": ["read_file"], "alias-a": ["write_file"]}) assert sorted(result["uuid-a"]) == ["read_file", "write_file"] + def test_wildcard_survives_expansion_as_list_entry(self): + """["*"] stays in the expanded list so the caller's wildcard check + (``_union_tool_grants``) can read it; this function only normalizes + keys and never maps grants to None.""" + manager = MCPServerManager() + manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="alpha") + + result = manager.expand_tool_permissions({"uuid-a": ["*"]}) + assert result == {"uuid-a": ["*"]} + + def test_wildcard_unions_with_concrete_names_across_keys_for_same_server(self): + """An alias key carrying ["*"] unioned with an id key naming one tool + keeps both entries; interpretation of the wildcard belongs to the + caller, not the expansion.""" + manager = MCPServerManager() + manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="alias-a", alias="alias-a") + + result = manager.expand_tool_permissions({"uuid-a": ["read_file"], "alias-a": ["*"]}) + assert sorted(result["uuid-a"]) == ["*", "read_file"] + + def test_empty_list_stays_deny_all(self): + """[] is deny-all, a distinct meaning from no entry (unrestricted); + the key must survive expansion rather than disappear.""" + manager = MCPServerManager() + manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="alpha") + + result = manager.expand_tool_permissions({"uuid-a": []}) + assert result == {"uuid-a": []} + class TestOAuthDiscoverySSRFGuard: """SSRF guard for the OAuth metadata discovery follow-up fetches. diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 30f5abdbb98..b811d4453ca 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -1783,6 +1783,336 @@ def test_can_object_call_model_access_via_alias_only(): assert result is True +def test_can_object_call_model_key_alias_to_allowed_target_is_allowed(): + """A key alias whose target is on the key allowlist resolves like a team alias.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + result = _can_object_call_model( + model="mistral-7b", + llm_router=None, + models=["gpt-4o-mini"], + key_model_aliases={"mistral-7b": "gpt-4o-mini"}, + object_type="key", + fallback_depth=0, + ) + + assert result is True + + +def test_can_object_call_model_key_alias_to_disallowed_target_is_denied(): + """A key alias whose target is outside the key allowlist stays denied.""" + from litellm.proxy._types import ProxyErrorTypes, ProxyException + from litellm.proxy.auth.auth_checks import _can_object_call_model + + with pytest.raises(ProxyException) as exc_info: + _can_object_call_model( + model="mistral-7b", + llm_router=None, + models=["gpt-4o-mini"], + key_model_aliases={"mistral-7b": "gpt-4"}, + object_type="key", + fallback_depth=0, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + assert exc_info.value.code == "403" + + +@pytest.mark.asyncio +async def test_can_team_access_model_honors_key_alias(): + """A key on a team can call a model through its own alias when the target is on the team allowlist.""" + from litellm.proxy.auth.auth_checks import can_team_access_model + + team_object = LiteLLM_TeamTable( + team_id="team-123", + models=["gpt-4o-mini"], + ) + + assert ( + await can_team_access_model( + model="mistral-7b", + team_object=team_object, + llm_router=None, + key_model_aliases={"mistral-7b": "gpt-4o-mini"}, + ) + is True + ) + + with pytest.raises(ProxyException) as exc_info: + await can_team_access_model( + model="mistral-7b", + team_object=team_object, + llm_router=None, + ) + + assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied + + +@pytest.mark.asyncio +async def test_can_key_call_model_honors_key_alias(): + """The real key entry point resolves a key alias to its target before the allowlist check.""" + from litellm.proxy.auth.auth_checks import can_key_call_model + + allowed_token = UserAPIKeyAuth( + api_key="sk-test", + models=["gpt-4o-mini"], + aliases={"mistral-7b": "gpt-4o-mini"}, + ) + + assert ( + await can_key_call_model( + model="mistral-7b", + llm_model_list=None, + valid_token=allowed_token, + llm_router=None, + ) + is True + ) + + denied_token = UserAPIKeyAuth( + api_key="sk-test", + models=["gpt-4o-mini"], + aliases={"mistral-7b": "gpt-4"}, + ) + + with pytest.raises(ProxyException) as exc_info: + await can_key_call_model( + model="mistral-7b", + llm_model_list=None, + valid_token=denied_token, + llm_router=None, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + + +def test_can_object_call_model_key_alias_applies_before_global_alias(monkeypatch): + """The key alias rewrite precedes the global one at dispatch, so the key target is authorized.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"}) + + assert ( + _can_object_call_model( + model="foo", + llm_router=None, + models=["baz"], + key_model_aliases={"foo": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + with pytest.raises(ProxyException) as exc_info: + _can_object_call_model( + model="foo", + llm_router=None, + models=["bar"], + key_model_aliases={"foo": "baz"}, + object_type="key", + fallback_depth=0, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + + +def test_can_object_call_model_key_alias_matches_global_rewritten_name(monkeypatch): + """A key alias on the globally rewritten name resolves the same way the request chain does.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"}) + + assert ( + _can_object_call_model( + model="foo", + llm_router=None, + models=["baz"], + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + +def test_can_object_call_model_chained_alias_requires_final_target(monkeypatch): + """When a key alias fires on the globally rewritten name, only the final target is dispatched.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"}) + + with pytest.raises(ProxyException) as exc_info: + _can_object_call_model( + model="foo", + llm_router=None, + models=["bar"], + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + + assert ( + _can_object_call_model( + model="foo", + llm_router=None, + models=["baz"], + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + +def test_can_object_call_model_key_alias_name_alone_is_not_enough(): + """A key that may call the alias name but not its target cannot call the alias.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + with pytest.raises(ProxyException) as exc_info: + _can_object_call_model( + model="bar", + llm_router=None, + models=["bar"], + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + + assert ( + _can_object_call_model( + model="bar", + llm_router=None, + models=["baz"], + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + +def test_can_object_call_model_team_alias_applies_before_key_alias(): + """A key alias on the raw name loses to the team alias that rewrites it first at dispatch.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + assert ( + _can_object_call_model( + model="foo", + llm_router=None, + models=["bar"], + team_model_aliases={"foo": "bar"}, + key_model_aliases={"foo": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + +def test_can_object_call_model_key_alias_on_team_alias_target(): + """A key alias on the team-rewritten name resolves like the dispatch chain does.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + assert ( + _can_object_call_model( + model="foo", + llm_router=None, + models=["baz"], + team_model_aliases={"foo": "bar"}, + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + with pytest.raises(ProxyException) as exc_info: + _can_object_call_model( + model="foo", + llm_router=None, + models=["bar"], + team_model_aliases={"foo": "bar"}, + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + + +@pytest.mark.asyncio +async def test_can_user_call_model_honors_key_alias(): + """A personal-scope key alias resolves to its target before the user allowlist check.""" + from litellm.proxy.auth.auth_checks import can_user_call_model + + user_object = LiteLLM_UserTable(user_id="test-user", models=["gpt-4o-mini"]) + + assert ( + await can_user_call_model( + model="mistral-7b", + llm_router=None, + user_object=user_object, + key_model_aliases={"mistral-7b": "gpt-4o-mini"}, + ) + is True + ) + + with pytest.raises(ProxyException) as exc_info: + await can_user_call_model( + model="mistral-7b", + llm_router=None, + user_object=user_object, + ) + + assert exc_info.value.type == ProxyErrorTypes.user_model_access_denied + + +@pytest.mark.asyncio +async def test_check_team_member_model_access_honors_key_alias(): + """A key alias resolves against the member allowlist, not just the raw alias name.""" + from litellm.proxy._types import LiteLLM_TeamMembership + from litellm.proxy.auth.auth_checks import _check_team_member_model_access + + membership = LiteLLM_TeamMembership( + user_id="alice", + team_id="team-a", + litellm_budget_table=LiteLLM_BudgetTable(allowed_models=["gpt-4o-mini"]), + ) + + await _check_team_member_model_access( + model="mistral-7b", + team_object=LiteLLM_TeamTable(team_id="team-a"), + valid_token=UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a"), + llm_router=None, + prisma_client=None, + user_api_key_cache=UserApiKeyCache(), + proxy_logging_obj=MagicMock(), + team_membership=membership, + team_membership_loaded=True, + key_model_aliases={"mistral-7b": "gpt-4o-mini"}, + ) + + with pytest.raises(ProxyException) as exc_info: + await _check_team_member_model_access( + model="mistral-7b", + team_object=LiteLLM_TeamTable(team_id="team-a"), + valid_token=UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a"), + llm_router=None, + prisma_client=None, + user_api_key_cache=UserApiKeyCache(), + proxy_logging_obj=MagicMock(), + team_membership=membership, + team_membership_loaded=True, + ) + + assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied + + def test_can_object_call_model_access_via_underlying_model_only(): """ Test that a key can access a model via underlying model even when using an alias. @@ -9139,6 +9469,50 @@ async def test_agent_access_groups_cap_models_even_when_key_allows_them(): assert asked == ["agent-1", "agent-1"] +@pytest.mark.asyncio +async def test_agent_access_group_ceiling_admits_the_key_alias_target(): + agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=["gpt-5"]) + agent_key.aliases = {"fast": "gpt-5"} + resolve, asked = _agent_model_ceiling_resolver(frozenset({"gpt-5"})) + + assert await _check_agent_access_group_model_access("fast", agent_key, None, resolve) is True + assert asked == ["agent-1"] + + +@pytest.mark.asyncio +async def test_agent_access_group_ceiling_checks_the_team_alias_target(): + agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", team_id="team-1", models=["gpt-5"]) + agent_key.team_model_aliases = {"foo": "gpt-5"} + resolve, asked = _agent_model_ceiling_resolver(frozenset({"gpt-5"})) + + assert await _check_agent_access_group_model_access("foo", agent_key, None, resolve) is True + assert asked == ["agent-1"] + + +@pytest.mark.asyncio +async def test_agent_access_group_ceiling_denies_a_team_alias_outside_the_ceiling(): + agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", team_id="team-1", models=["gpt-5"]) + agent_key.team_model_aliases = {"foo": "claude-sonnet-4-5"} + resolve, _ = _agent_model_ceiling_resolver(frozenset({"gpt-5"})) + + with pytest.raises(ModelAccessDeniedProxyException) as exc: + await _check_agent_access_group_model_access("foo", agent_key, None, resolve) + assert exc.value.type == ProxyErrorTypes.agent_model_access_denied + + +@pytest.mark.asyncio +async def test_agent_access_group_ceiling_keeps_the_name_for_a_deleted_team_deployment(): + from litellm.router import Router + + agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", team_id="team-1", models=["foo"]) + agent_key.team_model_aliases = {"foo": "model_name_team-1_deadbeef"} + router: Final = Router(model_list=[]) + resolve, asked = _agent_model_ceiling_resolver(frozenset({"foo"})) + + assert await _check_agent_access_group_model_access("foo", agent_key, router, resolve) is True + assert asked == ["agent-1"] + + @pytest.mark.asyncio async def test_agent_access_groups_naming_no_model_deny_every_model(): agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=[]) @@ -9437,6 +9811,18 @@ async def test_agent_key_acting_for_a_teamless_user_is_capped_at_that_users_mode assert asked == ["team:None", "user:alice", "team:None", "user:alice"] +@pytest.mark.asyncio +async def test_agent_key_alias_resolves_against_the_echoed_teams_models(): + agent_key: Final = _agent_key_acting_for(user_id="alice", team_id="team-a") + agent_key.aliases = {"foo": "bar"} + load_team, load_user, asked = _caller_loaders(LiteLLM_TeamTable(team_id="team-a", models=["bar"]), None) + cache: Final = await _cache_with_membership("alice", "team-a", allowed_models=None) + + await _check_caller_models(agent_key, "foo", load_team, load_user, cache) + + assert asked == ["team:team-a"] + + @pytest.mark.asyncio async def test_agent_key_without_an_echoed_caller_keeps_its_own_models(): agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=["gpt-5", "claude-sonnet"]) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 2d929a832a5..a40741c8fdb 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -5,10 +5,11 @@ import logging import os import sys import zlib -from collections.abc import Callable +from collections.abc import Callable, Mapping from contextlib import ExitStack, contextmanager +from dataclasses import dataclass from io import BytesIO -from types import SimpleNamespace +from types import MappingProxyType, SimpleNamespace from typing import Final from unittest.mock import AsyncMock, MagicMock, patch @@ -16,7 +17,7 @@ import httpx import pytest from fastapi import HTTPException, Request, Response, UploadFile from fastapi.responses import StreamingResponse -from pydantic import ValidationError +from pydantic import TypeAdapter, ValidationError from starlette.datastructures import FormData, Headers, QueryParams from starlette.datastructures import UploadFile as StarletteUploadFile @@ -45,6 +46,7 @@ from litellm.proxy.pass_through_endpoints.success_handler import ( PassThroughEndpointLogging, ) from litellm.proxy.route_llm_request import ProxyModelNotFoundError +from litellm.types import utils as types_utils from litellm.types.passthrough_endpoints.pass_through_endpoints import ( LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, @@ -7305,6 +7307,156 @@ def test_passthrough_logs_the_resolved_deployment_model_info_over_the_request_bo assert kwargs["litellm_params"]["metadata"]["model_info"] == {"id": "vertex-gemini-38-flash-dep"} +@dataclass(frozen=True, slots=True, kw_only=True) +class _PassThroughSplit: + litellm_params: Mapping[str, object] + forwarded_body: Mapping[str, object] + + +_LITELLM_PARAMS: Final = TypeAdapter(dict[str, object]) +_PROXY_SERVER_REQUEST: Final = TypeAdapter(dict[str, object]) + + +def _split_pass_through_body(body: str) -> _PassThroughSplit: + mock_request: Final = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent" + mock_request.headers = Headers() + mock_request.scope = MappingProxyType({}) + + init_kwargs_for_pass_through_endpoint: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint # pyright: ignore[reportUnknownVariableType, reportUnknownMemberType] # untyped legacy helper + kwargs: Final = init_kwargs_for_pass_through_endpoint( + request=mock_request, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + passthrough_logging_payload=MagicMock(), + logging_obj=MagicMock(), + _parsed_body=json.loads(body), + litellm_call_id="lit-owned-keys-call-id", + ) + validate_litellm_params: Final = _LITELLM_PARAMS.validate_python # pyright: ignore[reportUnknownArgumentType] # untyped legacy helper + litellm_params: Final = validate_litellm_params(kwargs["litellm_params"]) + return _PassThroughSplit( + litellm_params=MappingProxyType(litellm_params), + forwarded_body=MappingProxyType( + _LITELLM_PARAMS.validate_python( + _PROXY_SERVER_REQUEST.validate_python(litellm_params["proxy_server_request"])["body"] + ) + ), + ) + + +GEMINI_BODY: Final = '{"contents": [{"parts": [{"text": "hi"}]}], "generationConfig": {"temperature": 0}}' + + +def _metadata_of(split: _PassThroughSplit) -> Mapping[str, object]: + return MappingProxyType(_LITELLM_PARAMS.validate_python(split.litellm_params["metadata"])) + + +def test_passthrough_moves_every_litellm_owned_key_from_the_forwarded_body_into_litellm_params() -> None: + split: Final = _split_pass_through_body( + '{"ttl": 30, "contents": [{"parts": [{"text": "hi"}]}], "num_retries": 2,' + ' "generationConfig": {"temperature": 0}, "litellm_trace_id": "trace-a"}' + ) + + assert frozenset(split.litellm_params) == frozenset( + ("ttl", "num_retries", "litellm_trace_id", "metadata", "proxy_server_request") + ) + assert tuple(split.litellm_params[k] for k in ("ttl", "num_retries", "litellm_trace_id")) == (30, 2, "trace-a") + assert split.forwarded_body == json.loads(GEMINI_BODY) + + +PROXY_STAMPED_NAMES: Final = frozenset( + ( + "proxy_server_request", + "secret_fields", + "litellm_trusted_callback_vars", + "_litellm_addressed_response_id", + "_litellm_strip_stream_usage", + "client_side_timeout", + "model_file_id_mapping", + ) +) + + +@pytest.mark.parametrize( + "name", + sorted(frozenset(litellm.all_litellm_params) - frozenset(("metadata", "litellm_metadata")) - PROXY_STAMPED_NAMES), +) +def test_passthrough_keeps_each_registered_litellm_owned_name_out_of_the_forwarded_body(name: str) -> None: + split: Final = _split_pass_through_body(json.dumps({name: "owned", **json.loads(GEMINI_BODY)})) + + assert frozenset(split.litellm_params) == frozenset((name, "metadata", "proxy_server_request")) + assert split.litellm_params[name] == "owned" + assert split.forwarded_body == json.loads(GEMINI_BODY) + + +@pytest.mark.parametrize("name", sorted(PROXY_STAMPED_NAMES)) +def test_passthrough_drops_a_client_supplied_proxy_stamped_name(name: str) -> None: + split: Final = _split_pass_through_body(json.dumps({name: {"forged": "by-client"}, **json.loads(GEMINI_BODY)})) + + assert frozenset(split.litellm_params) == frozenset(("metadata", "proxy_server_request")) + assert split.litellm_params["proxy_server_request"] != {"forged": "by-client"}, split.litellm_params + assert split.forwarded_body == json.loads(GEMINI_BODY) + + +def test_passthrough_merges_both_metadata_carriers_from_the_body_into_one_metadata_key() -> None: + split: Final = _split_pass_through_body( + '{"metadata": {"client_tag": "a"}, "contents": [{"parts": [{"text": "hi"}]}], "ttl": 30,' + ' "litellm_metadata": {"lm": "b"}, "generationConfig": {"temperature": 0}}' + ) + + assert frozenset(split.litellm_params) == frozenset(("ttl", "metadata", "proxy_server_request")) + assert _metadata_of(split) == {**_metadata_of(_split_pass_through_body(GEMINI_BODY)), "client_tag": "a", "lm": "b"} + assert split.forwarded_body == json.loads(GEMINI_BODY) + + +def test_passthrough_lets_metadata_win_over_litellm_metadata_on_a_shared_key() -> None: + split: Final = _split_pass_through_body( + '{"litellm_metadata": {"shared": "from-litellm-metadata", "lm": "b"},' + ' "metadata": {"shared": "from-metadata", "client_tag": "a"}, "contents": []}' + ) + + assert _metadata_of(split) == { + **_metadata_of(_split_pass_through_body('{"contents": []}')), + "shared": "from-metadata", + "lm": "b", + "client_tag": "a", + } + + +def test_passthrough_orders_extracted_litellm_params_by_the_registry() -> None: + body: Final = json.dumps({"ttl": 30, "tags": ["team-a"], "num_retries": 2, "contents": []}) + split: Final = _split_pass_through_body(body) + body_keys: Final = frozenset(json.loads(body)) + + assert tuple(k for k in split.litellm_params if k in body_keys) == tuple( + k for k in types_utils.all_litellm_params if k in body_keys + ) + + +LATE_REGISTERED_BODY: Final = '{"registered_later": 1, "contents": [{"parts": [{"text": "hi"}]}]}' + + +def test_passthrough_sees_a_name_appended_to_the_public_list_after_import() -> None: + litellm.all_litellm_params.append("registered_later") + try: + split: Final = _split_pass_through_body(LATE_REGISTERED_BODY) + finally: + litellm.all_litellm_params.remove("registered_later") + + assert frozenset(split.litellm_params) == frozenset(("registered_later", "metadata", "proxy_server_request")) + assert split.forwarded_body == {"contents": [{"parts": [{"text": "hi"}]}]} + + +def test_passthrough_sees_the_public_list_rebound_after_import(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(types_utils, "all_litellm_params", (*litellm.all_litellm_params, "registered_later")) + + split: Final = _split_pass_through_body(LATE_REGISTERED_BODY) + + assert frozenset(split.litellm_params) == frozenset(("registered_later", "metadata", "proxy_server_request")) + assert split.forwarded_body == {"contents": [{"parts": [{"text": "hi"}]}]} + + @pytest.mark.asyncio async def test_chat_completion_pass_through_endpoint_answers_an_openai_typed_error_for_an_unknown_model( monkeypatch: pytest.MonkeyPatch, diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index f7d6cfaf079..99dea6366f9 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -309,6 +309,243 @@ def test_realtime_logging_object_does_not_validate_unknown_event_types(): assert len(dumped["results"]) == len(results) +def test_realtime_transcription_honors_deployment_pricing_override(monkeypatch: pytest.MonkeyPatch) -> None: + """A deployment's pricing override must reach transcription events too. + + Transcription is billed separately from response usage inside the same realtime + session, so a deployment registered at zero rates has to zero both. Resolving + transcription against the public ASR model instead billed a zero-rated + deployment for every .completed event. + """ + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + deployment_id = "deployment-hash-zero-rated-asr" + litellm.register_model( + model_cost={ + deployment_id: { + "litellm_provider": "openai", + "mode": "realtime", + "input_cost_per_second": 0.0, + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "input_cost_per_audio_token": 0.0, + } + } + ) + + results: OpenAIRealtimeStreamList = [ + { + "type": "session.created", + "session": { + "type": "transcription", + "audio": {"input": {"transcription": {"model": "gpt-realtime-whisper"}}}, + }, + }, + { + "type": "conversation.item.input_audio_transcription.completed", + "usage": {"type": "duration", "seconds": 120.0}, + }, + ] + + public_rate_cost = 120.0 * litellm.model_cost["gpt-realtime-whisper"]["input_cost_per_second"] + assert public_rate_cost > 0, "the public ASR rate must be non-zero for this test to mean anything" + + without_override = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="gpt-realtime-whisper", + ) + assert abs(without_override - public_rate_cost) < 1e-9 + + with_override = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="gpt-realtime-whisper", + custom_pricing_model=deployment_id, + ) + assert with_override == 0.0, "the zero-rated deployment must not be billed for transcription" + + +def test_realtime_transcription_partial_override_keeps_unset_rates(monkeypatch: pytest.MonkeyPatch) -> None: + """An override must not blank the rates it does not set. + + A deployment that prices tokens but omits input_cost_per_second would otherwise + bill duration-based transcription at nothing, because the cost helpers read + `.get(key) or 0.0`. Only the fields the operator actually set may win. + """ + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + deployment_id = "deployment-hash-tokens-only" + litellm.register_model( + model_cost={ + deployment_id: { + "litellm_provider": "openai", + "mode": "realtime", + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + } + } + ) + + results: OpenAIRealtimeStreamList = [ + { + "type": "session.created", + "session": { + "type": "transcription", + "audio": {"input": {"transcription": {"model": "gpt-realtime-whisper"}}}, + }, + }, + { + "type": "conversation.item.input_audio_transcription.completed", + "usage": {"type": "duration", "seconds": 120.0}, + }, + ] + + cost = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="gpt-realtime-whisper", + custom_pricing_model=deployment_id, + ) + + expected = 120.0 * litellm.model_cost["gpt-realtime-whisper"]["input_cost_per_second"] + assert expected > 0, "the public ASR per-second rate must be non-zero for this test to mean anything" + assert cost == pytest.approx(expected, rel=1e-9), ( + "duration must keep the ASR per-second rate the override left unset" + ) + + +@pytest.mark.parametrize( + "label,override,expected_audio_rate,expected_per_second", + [ + ("tokens only", {"input_cost_per_token": 0.0}, 0.0, 0.017 / 60), + ("audio zeroed", {"input_cost_per_audio_token": 0.0}, 0.0, 0.017 / 60), + ("per second only", {"input_cost_per_second": 0.001}, 6e-06, 0.001), + ("empty override", {}, 6e-06, 0.017 / 60), + ("no override", None, 6e-06, 0.017 / 60), + ], +) +def test_transcription_rate_precedence( + monkeypatch: pytest.MonkeyPatch, + label: str, + override: dict[str, float] | None, + expected_audio_rate: float, + expected_per_second: float, +) -> None: + """Rates resolve within one entry before moving to the next, and zero is a real value. + + An override that prices only tokens must apply its own token rate to audio rather + than reaching past itself for the public audio rate, a deliberate zero must win + instead of being treated as unset, and a rate the override never mentions must keep + the base entry's value. + """ + from litellm.cost_calculator import handle_realtime_transcription_cost_calculation + + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + base_model = "asr-precedence-base" + deployment_id = "asr-precedence-deployment" + litellm.register_model( + model_cost={ + base_model: { + "litellm_provider": "openai", + "mode": "audio_transcription", + "input_cost_per_audio_token": 6e-06, + "input_cost_per_token": 2.5e-06, + "input_cost_per_second": 0.017 / 60, + } + } + ) + if override is not None: + litellm.register_model( + model_cost={deployment_id: {"litellm_provider": "openai", "mode": "audio_transcription", **override}} + ) + + def cost_for(usage: dict[str, object]) -> float: + return handle_realtime_transcription_cost_calculation( + results=[ + {"type": "transcription_session.created", "session": {"model": base_model}}, + {"type": "conversation.item.input_audio_transcription.completed", "usage": usage}, + ], + custom_llm_provider="openai", + litellm_model_name=base_model, + custom_pricing_model=deployment_id if override is not None else None, + ) + + audio_cost = cost_for({"type": "tokens", "input_token_details": {"audio_tokens": 100}}) + assert audio_cost == pytest.approx(100 * expected_audio_rate, rel=1e-9), f"{label}: audio rate" + + per_second_cost = cost_for({"type": "duration", "seconds": 120.0}) + assert per_second_cost == pytest.approx(120.0 * expected_per_second, rel=1e-9), ( + f"{label}: an override must never blank a rate it does not set" + ) + + +def test_realtime_transcription_per_second_override_keeps_public_token_rates( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """A per-second override must not zero the token rates ``get_model_info`` synthesizes. + + ``get_model_info`` defaults input_cost_per_token and output_cost_per_token to 0 for entries + that omit them, so a deployment priced only per second looked like it had declared token + rates of 0. Token-shaped transcription then billed nothing instead of falling through to the + public ASR rates, while the per-second rate the operator did set stayed in force. + """ + from litellm.cost_calculator import handle_realtime_transcription_cost_calculation + + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + asr_model = "gpt-4o-transcribe" + per_second_rate = 0.001 + deployment_id = "deployment-hash-per-second-only" + litellm.register_model( + model_cost={ + deployment_id: { + "litellm_provider": "openai", + "mode": "audio_transcription", + "input_cost_per_second": per_second_rate, + } + } + ) + + public = litellm.model_cost[asr_model] + session_event = {"type": "transcription_session.created", "session": {"model": asr_model}} + + def cost_for(usage: dict[str, object]) -> float: + return handle_realtime_transcription_cost_calculation( + results=[session_event, {"type": "conversation.item.input_audio_transcription.completed", "usage": usage}], + custom_llm_provider="openai", + litellm_model_name=asr_model, + custom_pricing_model=deployment_id, + ) + + token_cost = cost_for( + { + "type": "tokens", + "input_token_details": {"audio_tokens": 400, "text_tokens": 12}, + "output_tokens": 30, + } + ) + expected_token_cost = ( + 400 * public["input_cost_per_audio_token"] + + 12 * public["input_cost_per_token"] + + 30 * public["output_cost_per_token"] + ) + assert expected_token_cost > 0, "the public ASR token rates must be non-zero for this test to mean anything" + assert token_cost == pytest.approx(expected_token_cost, rel=1e-9), ( + "an override that prices only seconds must leave the public token rates in place" + ) + + assert cost_for({"type": "duration", "seconds": 120.0}) == pytest.approx(120.0 * per_second_rate, rel=1e-9) + + def test_realtime_transcription_no_completed_events_is_zero(monkeypatch): """A realtime stream without transcription completed events adds no extra cost.""" monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") @@ -4635,6 +4872,391 @@ def test_gemini_live_native_audio_limits_and_capabilities_match_vendor_model_car assert info["supports_pdf_input"] is False +def test_realtime_honours_deployment_custom_pricing(monkeypatch: pytest.MonkeyPatch) -> None: + """Regression: a deployment's pricing override never reached realtime costing. + + `model_info` overrides are registered under the deployment's own model_id, and + only `_select_model_name_for_cost_calc` knows to look there. The realtime branch + discarded that result and priced by the model the session reported, so a config + that zeroes a realtime deployment was billed at the public rate anyway. Audio is + the bulk of a voice call, so the gap was most of the cost. + """ + from litellm.types.utils import CompletionTokensDetailsWrapper + + model = "gemini-3.1-flash-live-preview" + deployment_key = "deployment-id-for-a-zero-rated-realtime-group" + paid = litellm.model_cost[model] + monkeypatch.setitem( + litellm.model_cost, + deployment_key, + { + **paid, + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "input_cost_per_audio_token": 0.0, + "output_cost_per_audio_token": 0.0, + "cache_read_input_token_cost": 0.0, + }, + ) + + results: OpenAIRealtimeStreamList = [ + {"type": "session.created", "session": {"model": model}}, + { + "type": "response.done", + "response": {"usage": {"input_tokens": 10, "output_tokens": 200, "total_tokens": 210}}, + }, + ] + usage = Usage( + prompt_tokens=10, + completion_tokens=200, + total_tokens=210, + prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=10, cached_tokens=0), + completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=20, audio_tokens=180), + ) + + paid_cost = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=usage, + custom_llm_provider="gemini", + litellm_model_name=model, + ) + expected_paid = ( + 10 * paid["input_cost_per_token"] + + 20 * paid["output_cost_per_token"] + + 180 * paid["output_cost_per_audio_token"] + ) + assert paid_cost == pytest.approx(expected_paid, rel=1e-9) + assert paid_cost > 0 + + zero_rated_cost = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=usage, + custom_llm_provider="gemini", + litellm_model_name=model, + custom_pricing_model=deployment_key, + ) + assert zero_rated_cost == 0.0 + + +def test_realtime_honours_a_provider_prefixed_zero_rated_deployment(monkeypatch: pytest.MonkeyPatch) -> None: + """Regression: the override arrived provider-prefixed and was read as pricing nothing. + + `_select_model_name_for_cost_calc` hands back `/`, so the name reaching + the pricing guard carries a prefix the raw cost-map lookups cannot strip. The rates resolved + correctly through `get_model_info`, then the guard rejected them as undeclared and the session + billed the public rates. A zero-rated deployment must stay at zero however its name arrives. + """ + from litellm.types.utils import CompletionTokensDetailsWrapper + + model = "gemini-live-2.5-flash-native-audio" + deployment_key = "deployment-id-for-a-prefixed-zero-rated-realtime-group" + paid = litellm.model_cost[model] + monkeypatch.setitem( + litellm.model_cost, + deployment_key, + { + **paid, + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "input_cost_per_audio_token": 0.0, + "output_cost_per_audio_token": 0.0, + }, + ) + + results: OpenAIRealtimeStreamList = [ + {"type": "session.created", "session": {"model": model}}, + { + "type": "response.done", + "response": {"usage": {"input_tokens": 219, "output_tokens": 81, "total_tokens": 300}}, + }, + ] + usage = Usage( + prompt_tokens=219, + completion_tokens=81, + total_tokens=300, + prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=16, audio_tokens=203), + completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=23, audio_tokens=58), + ) + + paid_cost = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=usage, + custom_llm_provider="vertex_ai", + litellm_model_name=model, + ) + assert paid_cost > 0 + + zero_rated_cost = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=usage, + custom_llm_provider="vertex_ai", + litellm_model_name=model, + custom_pricing_model=f"vertex_ai/{deployment_key}", + ) + assert zero_rated_cost == 0.0 + + +def test_unpriced_deployment_entry_still_falls_through_to_the_session_model( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """The guard's own purpose must survive: an entry that prices nothing is not an override. + + Deployments are auto-registered under their model_id with no rates at all, and those must + keep billing at the session model's public rates rather than silently costing nothing. + """ + from litellm.types.utils import CompletionTokensDetailsWrapper + + model = "gemini-live-2.5-flash-native-audio" + deployment_key = "deployment-id-with-no-declared-rates" + monkeypatch.setitem( + litellm.model_cost, + deployment_key, + {key: value for key, value in litellm.model_cost[model].items() if "cost_per" not in key}, + ) + + results: OpenAIRealtimeStreamList = [ + {"type": "session.created", "session": {"model": model}}, + { + "type": "response.done", + "response": {"usage": {"input_tokens": 219, "output_tokens": 81, "total_tokens": 300}}, + }, + ] + usage = Usage( + prompt_tokens=219, + completion_tokens=81, + total_tokens=300, + prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=16, audio_tokens=203), + completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=23, audio_tokens=58), + ) + + with_unpriced_override = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=usage, + custom_llm_provider="vertex_ai", + litellm_model_name=model, + custom_pricing_model=f"vertex_ai/{deployment_key}", + ) + without_override = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=usage, + custom_llm_provider="vertex_ai", + litellm_model_name=model, + ) + assert with_unpriced_override == pytest.approx(without_override, rel=1e-9) + assert with_unpriced_override > 0 + + +def test_realtime_audio_only_override_bills_audio_at_the_deployment_rate( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Regression: an audio-only pricing override was never selected as the pricing key. + + The deployment-selection guard recognised only text, per-second, per-query and + tiered rates, so a deployment that priced just the audio meters was passed over + and the session kept billing the public rates for the exact tokens it priced. + """ + from litellm.types.utils import CompletionTokensDetailsWrapper + + model = "gemini-live-2.5-flash-native-audio" + deployment_key = "deployment-id-for-an-audio-only-realtime-group" + monkeypatch.setitem( + litellm.model_cost, + deployment_key, + { + "litellm_provider": "vertex_ai", + "mode": "realtime", + "input_cost_per_audio_token": 0.0, + "output_cost_per_audio_token": 0.0, + }, + ) + + logging_object = LiteLLMRealtimeStreamLoggingObject( + usage=Usage( + prompt_tokens=203, + completion_tokens=58, + total_tokens=261, + prompt_tokens_details=PromptTokensDetailsWrapper(audio_tokens=203), + completion_tokens_details=CompletionTokensDetailsWrapper(audio_tokens=58), + ), + results=[ + {"type": "session.created", "session": {"model": model}}, + { + "type": "response.done", + "response": {"usage": {"input_tokens": 203, "output_tokens": 58, "total_tokens": 261}}, + }, + ], + ) + + public_cost = completion_cost( + completion_response=logging_object, + model=model, + call_type=CallTypes.arealtime.value, + custom_llm_provider="vertex_ai", + ) + assert public_cost > 0 + + overridden_cost = completion_cost( + completion_response=logging_object, + model=model, + call_type=CallTypes.arealtime.value, + custom_llm_provider="vertex_ai", + custom_pricing=True, + router_model_id=deployment_key, + ) + assert overridden_cost == pytest.approx(0.0) + + +def test_realtime_session_falls_back_to_base_model_pricing(monkeypatch: pytest.MonkeyPatch) -> None: + """Regression: a priced base_model was discarded for realtime sessions. + + The resolved base model only reached the realtime cost path when custom pricing + was on, so a session reporting an alias unmapped in the cost map recorded zero + instead of the base model's published price. + """ + from litellm.types.utils import CompletionTokensDetailsWrapper + + base_model = "gemini-live-2.5-flash-native-audio" + logging_object = LiteLLMRealtimeStreamLoggingObject( + usage=Usage( + prompt_tokens=219, + completion_tokens=81, + total_tokens=300, + prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=16, audio_tokens=203), + completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=23, audio_tokens=58), + ), + results=[ + { + "type": "session.created", + "session": {"model": "my-voice-alias"}, + }, + { + "type": "response.done", + "response": {"usage": {"input_tokens": 219, "output_tokens": 81, "total_tokens": 300}}, + }, + ], + ) + + aliased_cost = completion_cost( + completion_response=logging_object, + model="my-voice-alias", + call_type=CallTypes.arealtime.value, + custom_llm_provider="vertex_ai", + base_model=base_model, + ) + base_cost = completion_cost( + completion_response=logging_object, + model=base_model, + call_type=CallTypes.arealtime.value, + custom_llm_provider="vertex_ai", + ) + assert aliased_cost == pytest.approx(base_cost, rel=1e-9) + assert aliased_cost > 0 + + +def test_base_model_does_not_override_transcription_rates(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + base_model = "gpt-realtime-2" + asr_model = "gpt-4o-transcribe" + logging_object = LiteLLMRealtimeStreamLoggingObject( + usage=Usage(), + results=[ + { + "type": "session.created", + "session": { + "model": "my-voice-alias", + "audio": {"input": {"transcription": {"model": asr_model}}}, + }, + }, + { + "type": "conversation.item.input_audio_transcription.completed", + "usage": { + "type": "tokens", + "input_token_details": {"audio_tokens": 400, "text_tokens": 12}, + "output_tokens": 30, + }, + }, + ], + ) + + with_base_model = completion_cost( + completion_response=logging_object, + model="my-voice-alias", + call_type=CallTypes.arealtime.value, + custom_llm_provider="openai", + base_model=base_model, + ) + asr_priced = completion_cost( + completion_response=logging_object, + model="my-voice-alias", + call_type=CallTypes.arealtime.value, + custom_llm_provider="openai", + ) + realtime_card = litellm.model_cost[base_model] + billed_at_realtime = ( + 400 * realtime_card["input_cost_per_audio_token"] + + 12 * realtime_card["input_cost_per_token"] + + 30 * realtime_card["output_cost_per_audio_token"] + ) + assert billed_at_realtime != pytest.approx(asr_priced, rel=1e-9) + assert with_base_model == pytest.approx(asr_priced, rel=1e-9) + assert with_base_model > 0 + + +def test_realtime_base_model_outranks_the_session_reported_model(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + from litellm.types.utils import CompletionTokensDetailsWrapper + + session_model = "gpt-realtime-mini" + base_model = "gpt-realtime-2" + + def logging_object_for(session: str) -> LiteLLMRealtimeStreamLoggingObject: + return LiteLLMRealtimeStreamLoggingObject( + usage=Usage( + prompt_tokens=120, + completion_tokens=60, + total_tokens=180, + prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=20, audio_tokens=100), + completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=10, audio_tokens=50), + ), + results=[ + { + "type": "session.created", + "session": {"model": session}, + }, + { + "type": "response.done", + "response": {"usage": {"input_tokens": 120, "output_tokens": 60, "total_tokens": 180}}, + }, + ], + ) + + with_base_model = completion_cost( + completion_response=logging_object_for(session_model), + model=session_model, + call_type=CallTypes.arealtime.value, + custom_llm_provider="openai", + base_model=base_model, + ) + base_priced = completion_cost( + completion_response=logging_object_for(base_model), + model=base_model, + call_type=CallTypes.arealtime.value, + custom_llm_provider="openai", + ) + session_priced = completion_cost( + completion_response=logging_object_for(session_model), + model=session_model, + call_type=CallTypes.arealtime.value, + custom_llm_provider="openai", + ) + assert base_priced != pytest.approx(session_priced, rel=1e-9) + assert with_base_model == pytest.approx(base_priced, rel=1e-9) + + def test_baseten_glm_5_3_fast_is_priced_from_registry(_local_model_cost_map: None) -> None: model: Final = "baseten/zai-org/GLM-5.3-Fast" prompt_tokens: Final = 1000 diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 82122da15dc..80131534183 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -25,6 +25,7 @@ from litellm import Router from litellm.caching.caching import DualCache from litellm.caching.redis_cache import _redis_circuit_breaker_guard from litellm.exceptions import MidStreamFallbackError +from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging @@ -49,6 +50,7 @@ from litellm.router import ( from litellm.router_strategy import simple_shuffle from litellm.router_utils.client_initalization_utils import MaxParallelRequestsLimit from litellm.router_utils.cooldown_handlers import _async_get_cooldown_deployments +from litellm.router_utils.fallback_event_handlers import DISABLE_FALLBACKS_METADATA_KEY from litellm.router_utils.router_callbacks.track_deployment_metrics import get_deployment_successes_for_current_minute from litellm.types.llms.openai import ChatCompletionRequest from litellm.types.router import Deployment, DeploymentTypedDict, LiteLLM_Params, ModelInfo, PreRoutingHookResponse, RetryPolicy @@ -14392,6 +14394,193 @@ async def test_anthropic_messages_fallback_also_catches_raised_midstream_error() assert mock_fallback.await_args.kwargs["e"] is raised_error +_MID_STREAM_OPT_OUT_SHAPES: Final = ( + pytest.param({"disable_fallbacks": True}, id="raw-kwarg"), + pytest.param({"metadata": {DISABLE_FALLBACKS_METADATA_KEY: True}}, id="metadata-stamp"), + pytest.param({"litellm_metadata": {DISABLE_FALLBACKS_METADATA_KEY: True}}, id="litellm_metadata-stamp"), +) + + +def _mid_stream_opt_out_router() -> Router: + return Router( + model_list=[ + {"model_name": "primary", "litellm_params": {"model": "openai/gpt-5.4", "api_key": "k1"}}, + {"model_name": "fallback", "litellm_params": {"model": "openai/gpt-5.4-mini", "api_key": "k2"}}, + ], + fallbacks=[{"primary": ["fallback"]}], + ) + + +def _mid_stream_opt_out_primary_error() -> litellm.InternalServerError: + return litellm.InternalServerError(message="primary failed at stream start", llm_provider="openai", model="primary") + + +def _mid_stream_opt_out_trigger(primary_error: Exception) -> MidStreamFallbackError: + return MidStreamFallbackError( + message=str(primary_error), + model="primary", + llm_provider="openai", + original_exception=primary_error, + is_pre_first_chunk=True, + ) + + +class _MidStreamOptOutChatStream(CustomStreamWrapper): + """A chat deployment stream, as the router sees one, that dies before its first chunk.""" + + def __init__(self, error: Exception, model: str = "primary") -> None: + super().__init__(completion_stream=object(), model=model, custom_llm_provider="openai", logging_obj=MagicMock()) + self._error: Final = error + + def __aiter__(self): + return self + + async def __anext__(self) -> object: + raise self._error + + def __iter__(self): + return self + + def __next__(self) -> object: + raise self._error + + +@pytest.mark.asyncio +@pytest.mark.parametrize("opt_out", _MID_STREAM_OPT_OUT_SHAPES) +async def test_acompletion_streaming_iterator_honors_disable_fallbacks(opt_out): + """A chat stream that fails before its first chunk on a request that opted out of fallbacks + surfaces the primary's own error and never tries the fallback deployment.""" + router = _mid_stream_opt_out_router() + primary_error = _mid_stream_opt_out_primary_error() + source = _MidStreamOptOutChatStream(_mid_stream_opt_out_trigger(primary_error)) + + with patch.object(router, "_acompletion", new=AsyncMock(return_value=_AsyncList([]))) as fallback_attempt: + wrapped = await router._acompletion_streaming_iterator( + model_response=source, + messages=[{"role": "user", "content": "Hi"}], + initial_kwargs={"model": "primary", "stream": True, **copy.deepcopy(opt_out)}, + ) + with pytest.raises(litellm.InternalServerError) as raised: + [chunk async for chunk in wrapped] + + assert raised.value is primary_error + fallback_attempt.assert_not_awaited() + + +@pytest.mark.parametrize("opt_out", _MID_STREAM_OPT_OUT_SHAPES) +def test_completion_streaming_iterator_honors_disable_fallbacks(opt_out): + """Sync counterpart of test_acompletion_streaming_iterator_honors_disable_fallbacks.""" + router = _mid_stream_opt_out_router() + primary_error = _mid_stream_opt_out_primary_error() + source = _MidStreamOptOutChatStream(_mid_stream_opt_out_trigger(primary_error)) + + with patch.object(router, "_completion", new=MagicMock(return_value=iter([]))) as fallback_attempt: + wrapped = router._completion_streaming_iterator( + model_response=source, + messages=[{"role": "user", "content": "Hi"}], + initial_kwargs={"model": "primary", "stream": True, **copy.deepcopy(opt_out)}, + ) + with pytest.raises(litellm.InternalServerError) as raised: + list(wrapped) + + assert raised.value is primary_error + fallback_attempt.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("opt_out", _MID_STREAM_OPT_OUT_SHAPES) +async def test_aresponses_streaming_iterator_honors_disable_fallbacks(opt_out): + """Same opt-out contract on the Responses API mid-stream fallback path.""" + router = _mid_stream_opt_out_router() + primary_error = _mid_stream_opt_out_primary_error() + source = _make_responses_iterator(error=_mid_stream_opt_out_trigger(primary_error), model="primary") + + with patch.object( + router, + "_ageneric_api_call_with_fallbacks_responses_attempt", + new=AsyncMock(return_value=_AsyncList([])), + ) as fallback_attempt: + wrapped = await router._aresponses_streaming_iterator( + response=source, + initial_kwargs={ + "model": "primary", + "stream": True, + "input": "Hi", + "original_generic_function": litellm.aresponses, + **copy.deepcopy(opt_out), + }, + ) + with pytest.raises(litellm.InternalServerError) as raised: + [chunk async for chunk in wrapped] + + assert raised.value is primary_error + fallback_attempt.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("opt_out", _MID_STREAM_OPT_OUT_SHAPES) +async def test_anthropic_messages_streaming_iterator_honors_disable_fallbacks(opt_out): + """Same opt-out contract on the Anthropic Messages mid-stream fallback path.""" + router = _mid_stream_opt_out_router() + primary_error = _mid_stream_opt_out_primary_error() + source = _AnthropicMessagesRaisingByteStream([], _mid_stream_opt_out_trigger(primary_error)) + + with patch.object( + router, + "_ageneric_api_call_with_fallbacks_anthropic_messages_attempt", + new=AsyncMock(return_value=_AnthropicMessagesFallbackByteStream([])), + ) as fallback_attempt: + wrapped = await router._aanthropic_messages_streaming_iterator( + response=source, + initial_kwargs={"model": "primary", "stream": True, **copy.deepcopy(opt_out)}, + ) + with pytest.raises(litellm.InternalServerError) as raised: + [chunk async for chunk in wrapped] + + assert raised.value is primary_error + fallback_attempt.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_acompletion_disable_fallbacks_reaches_the_mid_stream_hop(): + """`disable_fallbacks=True` sent to the public entrypoint survives the fallback wrapper's handoff + into the stream: the primary's own error surfaces and no fallback deployment is ever called.""" + router = _mid_stream_opt_out_router() + primary_error = _mid_stream_opt_out_primary_error() + + async def primary_stream(**kwargs): + return _MidStreamOptOutChatStream(_mid_stream_opt_out_trigger(primary_error), model=kwargs["model"]) + + with patch("litellm.acompletion", side_effect=primary_stream) as provider_calls: + response = await router.acompletion( + model="primary", messages=[{"role": "user", "content": "Hi"}], stream=True, disable_fallbacks=True + ) + with pytest.raises(litellm.InternalServerError) as raised: + [chunk async for chunk in response] + + assert raised.value is primary_error + assert [call.kwargs["metadata"]["model_group"] for call in provider_calls.call_args_list] == ["primary"] + + +def test_completion_disable_fallbacks_reaches_the_mid_stream_hop(): + """Sync counterpart of test_acompletion_disable_fallbacks_reaches_the_mid_stream_hop.""" + router = _mid_stream_opt_out_router() + primary_error = _mid_stream_opt_out_primary_error() + + def primary_stream(**kwargs): + return _MidStreamOptOutChatStream(_mid_stream_opt_out_trigger(primary_error), model=kwargs["model"]) + + with patch("litellm.completion", side_effect=primary_stream) as provider_calls: + response = router.completion( + model="primary", messages=[{"role": "user", "content": "Hi"}], stream=True, disable_fallbacks=True + ) + with pytest.raises(litellm.InternalServerError) as raised: + list(response) + + assert raised.value is primary_error + assert [call.kwargs["metadata"]["model_group"] for call in provider_calls.call_args_list] == ["primary"] + + @pytest.mark.asyncio @pytest.mark.parametrize( "raised_error", diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 285188c9c09..768d8955b8e 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -3730,6 +3730,23 @@ def test_scoped_weights_are_excluded_from_provider_params(filter_name: str) -> N assert filtered == {"provider_option": "kept"} +@pytest.mark.parametrize( + "provider_filter", + [ + litellm.utils.get_non_default_completion_params, + litellm.utils.get_non_default_transcription_params, + litellm.utils.filter_out_litellm_params, + ], +) +@pytest.mark.parametrize("setting", [("tag_regex", ["^team-a$"]), ("max_file_size_mb", 5)]) +def test_deployment_only_settings_copied_by_the_router_stay_out_of_provider_params( + provider_filter: Callable[[dict[str, object]], Mapping[str, object]], setting: tuple[str, object] +) -> None: + name, value = setting + filtered: Final = provider_filter({"provider_option": "kept", name: value}) + assert filtered == {"provider_option": "kept"}, filtered + + class TestGetOptionalParamsTencent: """Tests that tencent provider uses TencentChatConfig for parameter mapping.""" diff --git a/tests/unit/llms/vertex_ai/vertex_ai_partner_models/llama3/test_vertex_ai_partner_models_llama3_transformation.py b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/llama3/test_vertex_ai_partner_models_llama3_transformation.py index 3bca51ec6b3..05e4e36edd7 100644 --- a/tests/unit/llms/vertex_ai/vertex_ai_partner_models/llama3/test_vertex_ai_partner_models_llama3_transformation.py +++ b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/llama3/test_vertex_ai_partner_models_llama3_transformation.py @@ -11,6 +11,39 @@ from litellm.llms.vertex_ai.vertex_ai_partner_models.llama3.transformation impor ) +OPENAI_PLATFORM_PARAMS = ( + "prompt_cache_key", + "prompt_cache_retention", + "safety_identifier", + "service_tier", + "store", + "web_search_options", + "modalities", + "prediction", + "audio", +) + +SELF_DEPLOYED_ENDPOINT_MODELS = ( + "gemma/gemma-2-2b-it", + "vertex_ai/gemma/gemma-2-2b-it", + "openai/mg-endpoint-lit8592", + "vertex_ai/openai/mg-endpoint-lit8592", + "openai/5464397967697903616", +) + +MAAS_MODELS = ( + "meta/llama-4-maverick-17b-128e-instruct-maas", + "vertex_ai/meta/llama-4-maverick-17b-128e-instruct-maas", + "moonshotai/kimi-k2-thinking-maas", + "qwen/qwen3-next-80b-a3b-instruct-maas", + "google/gemma-4-26b-a4b-it-maas", + "xai/grok-4.1-fast-non-reasoning", + "openai/xai/grok-4.1-fast-reasoning", + "1984786713414729728", + "llama3", +) + + class TestVertexAILlama3Config: def test_transform_choices(self): """ @@ -56,6 +89,52 @@ class TestVertexAILlama3Config: assert response[0].message.tool_calls is not None assert response[0].finish_reason == "tool_calls" + @pytest.mark.parametrize("model", SELF_DEPLOYED_ENDPOINT_MODELS) + @pytest.mark.parametrize("param", OPENAI_PLATFORM_PARAMS) + def test_get_supported_openai_params_omits_platform_params_for_self_deployed_endpoints( + self, model: str, param: str + ): + assert param not in VertexAILlama3Config().get_supported_openai_params(model=model) + + @pytest.mark.parametrize("model", MAAS_MODELS) + @pytest.mark.parametrize("param", OPENAI_PLATFORM_PARAMS) + def test_get_supported_openai_params_keeps_platform_params_for_maas_models(self, model: str, param: str): + assert param in VertexAILlama3Config().get_supported_openai_params(model=model) + + @pytest.mark.parametrize("model", [*SELF_DEPLOYED_ENDPOINT_MODELS, *MAAS_MODELS]) + def test_get_supported_openai_params_never_lists_max_retries(self, model: str): + assert "max_retries" not in VertexAILlama3Config().get_supported_openai_params(model=model) + + @pytest.mark.parametrize("model", [*SELF_DEPLOYED_ENDPOINT_MODELS, *MAAS_MODELS]) + @pytest.mark.parametrize( + "param", + ["max_completion_tokens", "tools", "tool_choice", "response_format", "seed", "logprobs", "parallel_tool_calls"], + ) + def test_get_supported_openai_params_keeps_params_every_vertex_openai_endpoint_accepts( + self, model: str, param: str + ): + assert param in VertexAILlama3Config().get_supported_openai_params(model=model) + + @pytest.mark.parametrize("model", SELF_DEPLOYED_ENDPOINT_MODELS) + def test_map_openai_params_drops_prompt_cache_key_for_self_deployed_endpoints(self, model: str): + mapped = VertexAILlama3Config().map_openai_params( + {"prompt_cache_key": "session-lit8592", "max_completion_tokens": 10}, + {}, + model, + drop_params=True, + ) + assert mapped == {"max_tokens": 10} + + @pytest.mark.parametrize("model", MAAS_MODELS) + def test_map_openai_params_forwards_prompt_cache_key_for_maas_models(self, model: str): + mapped = VertexAILlama3Config().map_openai_params( + {"prompt_cache_key": "session-lit8592", "max_completion_tokens": 10}, + {}, + model, + drop_params=True, + ) + assert mapped == {"prompt_cache_key": "session-lit8592", "max_tokens": 10} + class TestVertexAILlama3StreamingHandler: def test_first_chunk_has_role_assistant_when_missing(self): diff --git a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py index e9ae5234094..e5ca31833ce 100644 --- a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py +++ b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py @@ -694,6 +694,89 @@ class TestVertexGemmaCompletion: assert instance["@requestFormat"] == "chatCompletions" assert "messages" in instance + @pytest.mark.parametrize( + "param", + [ + "prompt_cache_key", + "prompt_cache_retention", + "safety_identifier", + "service_tier", + "store", + "web_search_options", + "modalities", + "prediction", + "audio", + "max_retries", + ], + ) + def test_get_supported_openai_params_omits_params_the_predict_endpoint_rejects(self, param: str): + from litellm.llms.vertex_ai.vertex_gemma_models.transformation import ( + VertexGemmaConfig, + ) + + assert param not in VertexGemmaConfig().get_supported_openai_params(model="gemma-2-2b-it") + + @pytest.mark.asyncio + async def test_acompletion_drops_prompt_cache_key_when_drop_params_is_set(self): + with ( + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client") as mock_get_client, + patch( + "litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token", + return_value=("fake-access-token", "PROJECT_ID"), + ), + ): + mock_client = Mock() + mock_response = Mock() + mock_response.status_code = 200 + mock_response.json.return_value = _make_gemma_vertex_response() + mock_client.post = AsyncMock(return_value=mock_response) + mock_get_client.return_value = mock_client + + await litellm.acompletion( + model="vertex_ai/gemma/gemma-2-2b-it", + messages=[{"role": "user", "content": "Test"}], + prompt_cache_key="session-lit8592", + service_tier="default", + max_completion_tokens=16, + drop_params=True, + api_base="https://test.us-central1-project.prediction.vertexai.goog/v1/projects/PROJECT_ID/locations/us-central1/endpoints/ENDPOINT_ID:predict", + vertex_project="PROJECT_ID", + vertex_location="us-central1", + ) + + instance = mock_client.post.call_args.kwargs["json"]["instances"][0] + assert "prompt_cache_key" not in instance + assert "service_tier" not in instance + assert instance["max_tokens"] == 16 + assert instance["messages"] == [{"role": "user", "content": "Test"}] + + @pytest.mark.asyncio + async def test_acompletion_rejects_prompt_cache_key_before_calling_vertex(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(litellm, "drop_params", False) + with ( + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client") as mock_get_client, + patch( + "litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token", + return_value=("fake-access-token", "PROJECT_ID"), + ), + ): + mock_client = Mock() + mock_client.post = AsyncMock() + mock_get_client.return_value = mock_client + + with pytest.raises(litellm.UnsupportedParamsError, match="prompt_cache_key"): + await litellm.acompletion( + model="vertex_ai/gemma/gemma-2-2b-it", + messages=[{"role": "user", "content": "Test"}], + prompt_cache_key="session-lit8592", + drop_params=False, + api_base="https://test.us-central1-project.prediction.vertexai.goog/v1/projects/PROJECT_ID/locations/us-central1/endpoints/ENDPOINT_ID:predict", + vertex_project="PROJECT_ID", + vertex_location="us-central1", + ) + + mock_client.post.assert_not_called() + def test_transform_request_strips_context_management(self): """ Direct unit test for VertexGemmaConfig.transform_request: verify that diff --git a/tests/unit/types/test_litellm_params.py b/tests/unit/types/test_litellm_params.py new file mode 100644 index 00000000000..e421321aaaa --- /dev/null +++ b/tests/unit/types/test_litellm_params.py @@ -0,0 +1,655 @@ +import inspect +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass, field, fields +from operator import attrgetter +from types import MappingProxyType +from typing import Final, TypeAlias, cast, get_type_hints + +import httpx +import pytest +from aiohttp import ClientSession +from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI +from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError + +import litellm +from litellm.caching.caching import Cache +from litellm.litellm_core_utils.get_litellm_params import ( + get_litellm_params, # pyright: ignore[reportUnknownVariableType] # untyped legacy carrier +) +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.router_strategy.complexity_router.context_compaction import CompactionState +from litellm.router_utils.fallback_event_handlers import AttemptedFallbackTargets +from litellm.types import litellm_params +from litellm.types import utils as types_utils +from litellm.types.caching import DynamicCacheControl +from litellm.types.litellm_params import ( + ADDRESSED_RESPONSE_ID_FIELD, + LITELLM_OWNED_ROOTS, + TRUSTED_CALLBACK_VARS_FIELD, + CachingOptions, + owned_wire_names, + wire, + wire_names, +) +from litellm.types.llms.openai import ChatCompletionAssistantMessage, ChatCompletionUserMessage +from litellm.types.proxy.litellm_pre_call_utils import SecretFields +from litellm.types.router import ( + ConfigurableClientsideParamsCustomAuth, + CredentialLiteLLMParams, + DeploymentTypedDict, + RetryPolicy, + RouterConfig, + UpdateRouterConfig, +) +from litellm.types.router_weights import RouterWeights +from litellm.types.utils import ( + CustomPricingLiteLLMParams, + ModelResponse, + ModelResponseStream, + ProviderSpecificHeader, + StandardCallbackDynamicParams, + agentic_loop_internal_litellm_params, + all_litellm_params, + bedrock_batch_litellm_params, +) +from litellm.utils import ( + filter_out_litellm_params, # pyright: ignore[reportUnknownVariableType] # untyped legacy classifier + get_non_default_completion_params, # pyright: ignore[reportUnknownVariableType] # untyped legacy classifier + get_non_default_transcription_params, # pyright: ignore[reportUnknownVariableType] # untyped legacy classifier +) + +PROVIDER_KNOB: Final = "registry_test_provider_only_knob" + +CONNECTION_NAMES: Final = ( + "api_key", + "api_base", + "api_version", + "region_name", + "headers", + "provider_specific_header", + "client", + "shared_session", + "ssl_verify", + "request_timeout", + "force_timeout", + "stream_timeout", + "max_retries", + "tenant_id", + "client_id", + "client_secret", + "azure_username", + "azure_password", + "azure_scope", + "azure_ad_token_provider", + "litellm_credential_name", + "configurable_clientside_auth_params", + "use_xai_oauth", + "aws_batch_role_arn", + "s3_bucket_name", + "s3_region_name", + "s3_endpoint_url", + "s3_output_bucket_name", + "s3_bucket_owner", + "s3_access_key_id", + "s3_secret_access_key", + "s3_encryption_key_id", + "bedrock_tags", +) + +OPTION_NAMES: Final = ( + "custom_llm_provider", + "azure", + "use_litellm_proxy", + "use_chat_completions_api", + "use_in_pass_through", + "allowed_openai_params", + "fallbacks", + "context_window_fallback_dict", + "num_retries", + "retry_policy", + "retry_strategy", + "routing_strategy", + "cooldown_time", + "allowed_model_region", + "enable_tag_filtering", + "fastest_response", + "provider_affinity_header", + "search_tool_name", + "model_list", + "model_info", + "rpm", + "tpm", + "itpm", + "otpm", + "default_api_key_rpm_limit", + "default_api_key_tpm_limit", + "max_parallel_requests", + "weight", + "order", + "tag_regex", + "max_file_size_mb", + "auto_router_config_path", + "auto_router_config", + "auto_router_default_model", + "auto_router_embedding_model", + "auto_router_max_input_chars", + "auto_router_routing_compression", + "auto_router_model_compression", + "complexity_router_config", + "complexity_router_default_model", + "adaptive_router_config", + "adaptive_router_default_model", + "quality_router_config", + "quality_router_default_model", + "caching", + "cache", + "ttl", + "enable_prompt_caching", + "caching_groups", + "cost_per_query", + "base_model", + "max_budget", + "budget_duration", + "id", + "metadata", + "litellm_metadata", + "tags", + "litellm_trace_id", + "litellm_session_id", + "litellm_request_debug", + "logger_fn", + "verbose", + "no-log", + "max_agentic_loops", + "guardrails", + "prompt_id", + "prompt_variables", + "prompt_version", + "prompt_environment", + "prompt_label", + "litellm_system_prompt", + "custom_prompt_dict", + "roles", + "final_prompt_value", + "bos_token", + "eos_token", + "hf_model_name", + "supports_system_message", + "ensure_alternating_roles", + "user_continue_message", + "assistant_continue_message", + "disable_add_transform_inline_image_block", + "merge_reasoning_content_in_choices", + "enable_json_schema_validation", + "complete_response", + "stream_chunk_size", + "keepalive_seconds", + "allow_client_keepalive_override", + "mock_response", + "mock_timeout", +) + +AGENTIC_LOOP_STATE_NAMES: Final = ( + "_agentic_loop_depth", + "_agentic_loop_fingerprints", + "_agentic_loop_api_surface", + "_code_interpreter_interception_active", + "_code_interpreter_interception_sandbox_key", + "_code_interpreter_interception_session_scoped", + "_code_interpreter_interception_converted_stream", + "_websearch_interception_emit_native_blocks", + "_websearch_interception_converted_stream", + "_headroom_interception_converted_stream", +) + +INTERNAL_STATE_NAMES: Final = ( + "litellm_call_id", + "completion_call_id", + "model_alias_map", + "data_residency", + "litellm_logging_obj", + "preset_cache_key", + "cache_key", + "stream_response", + "_context_compaction_state", + *AGENTIC_LOOP_STATE_NAMES, + "_router_weights", + "fallback_depth", + "max_fallbacks", + "attempted_targets", + "proxy_server_request", + "secret_fields", + "litellm_trusted_callback_vars", + "_litellm_addressed_response_id", + "_litellm_strip_stream_usage", + "client_side_timeout", + "model_file_id_mapping", + "acompletion", + "aembedding", + "aimg_generation", + "atext_completion", + "text_completion", + "allm_passthrough_route", + "async_call", +) + +BEDROCK_BATCH_NAMES: Final = ( + "aws_batch_role_arn", + "s3_bucket_name", + "s3_region_name", + "s3_endpoint_url", + "s3_output_bucket_name", + "s3_bucket_owner", + "s3_access_key_id", + "s3_secret_access_key", + "s3_encryption_key_id", + "bedrock_tags", +) + +ARTIFACT_NAMES: Final = ("self", "use_client", "model_config", "rust") + +CALLBACK_VAR_NAMES: Final = tuple(StandardCallbackDynamicParams.__annotations__) + +PRICING_NAMES: Final = tuple(CustomPricingLiteLLMParams.model_fields) + +OWNED_NAMES: Final = ( + *CONNECTION_NAMES, + *OPTION_NAMES, + *INTERNAL_STATE_NAMES, + *ARTIFACT_NAMES, + *CALLBACK_VAR_NAMES, + *PRICING_NAMES, +) + +Classifier: TypeAlias = Callable[[dict[str, object]], dict[str, object]] # mutable-ok: classifiers use dict + +CLASSIFIERS: Final[Mapping[str, Classifier]] = MappingProxyType( + { # pyright: ignore[reportUnknownArgumentType] # untyped legacy classifiers + "completion": get_non_default_completion_params, + "transcription": get_non_default_transcription_params, + "filter_out": filter_out_litellm_params, + } +) + + +@pytest.mark.parametrize("classifier_name", CLASSIFIERS) +@pytest.mark.parametrize("name", OWNED_NAMES) +def test_owned_name_is_kept_out_of_provider_params(name: str, classifier_name: str) -> None: + provider_value: Final = object() + classify: Final = CLASSIFIERS[classifier_name] + + result: Final = classify({name: object(), PROVIDER_KNOB: provider_value}) # mutable-ok: classifiers take a dict + + assert result == MappingProxyType({PROVIDER_KNOB: provider_value}) + assert result[PROVIDER_KNOB] is provider_value + + +def test_a_name_no_object_declares_reaches_the_provider() -> None: + result: Final = CLASSIFIERS["completion"]({PROVIDER_KNOB: 1}) # mutable-ok: classifier input type + + assert result == MappingProxyType({PROVIDER_KNOB: 1}) + + +def _cache_key_for_model_group(cache: Cache, model_group: str, options: CachingOptions) -> str: + return cache.get_cache_key( # pyright: ignore[reportUnknownMemberType] # untyped legacy key builder + model=model_group, + messages=(MappingProxyType({"role": "user", "content": "shared prompt"}),), + metadata=MappingProxyType({"caching_groups": options.caching_groups, "model_group": model_group}), + ) + + +def test_caching_groups_is_a_flat_sequence_of_model_groups_that_share_one_cache_key( + monkeypatch: pytest.MonkeyPatch, +) -> None: + for callback_list in ("input_callback", "success_callback", "_async_success_callback"): + monkeypatch.setattr(litellm, callback_list, []) # mutable-ok: Cache() appends "cache" to these lists + options: Final = CachingOptions(caching_groups=(("gpt-4", "gpt-4o"), ("claude-3",))) + cache: Final = Cache() + + keys: Final = tuple(_cache_key_for_model_group(cache, group, options) for group in ("gpt-4", "gpt-4o", "claude-3")) + + assert (keys[0] == keys[1], keys[0] == keys[2]) == (True, False) + + +def test_all_litellm_params_is_exactly_the_owned_inventory() -> None: + assert frozenset(all_litellm_params) == frozenset(OWNED_NAMES) + assert frozenset(ARTIFACT_NAMES).isdisjoint(DECLARED_NAMES) + + +def test_every_owned_name_has_exactly_one_owner() -> None: + duplicated: Final = tuple(name for name in dict.fromkeys(all_litellm_params) if all_litellm_params.count(name) > 1) + + assert duplicated == () + + +@pytest.mark.parametrize( + ("exported", "declared"), + ( + pytest.param( + types_utils.TRUSTED_CALLBACK_VARS_FIELD, + litellm_params.TRUSTED_CALLBACK_VARS_FIELD, + id="TRUSTED_CALLBACK_VARS_FIELD", + ), + pytest.param( + types_utils.ADDRESSED_RESPONSE_ID_FIELD, + litellm_params.ADDRESSED_RESPONSE_ID_FIELD, + id="ADDRESSED_RESPONSE_ID_FIELD", + ), + ), +) +def test_types_utils_still_exports_the_field_constant(exported: str, declared: str) -> None: + assert exported == declared + + +@dataclass(frozen=True, slots=True, kw_only=True) +class _Leaf: + plain: int | None = None + renamed: int | None = field(default=None, metadata=wire("wire-name")) + + +@dataclass(frozen=True, slots=True, kw_only=True) +class _OtherLeaf: + plain: int | None = None + trailing: int | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class _Root: + first: _Leaf + second: _OtherLeaf + + +@dataclass(frozen=True, slots=True, kw_only=True) +class _RootDeclaringAKwargDirectly: + first: _Leaf + stray: int | None = None + + +def test_wire_names_are_the_field_names_in_declaration_order_unless_wire_renames_them() -> None: + assert wire_names(_Leaf) == ("plain", "wire-name") + + +def test_owned_wire_names_walk_leaves_in_declaration_order_and_keep_every_occurrence() -> None: + assert owned_wire_names(_Root) == ("plain", "wire-name", "plain", "trailing") + + +def test_owned_wire_names_refuse_a_root_that_declares_a_kwarg_outside_a_leaf() -> None: + with pytest.raises(TypeError): + owned_wire_names(_RootDeclaringAKwargDirectly) + + +def test_agentic_loop_names_concatenate_as_a_list() -> None: + extended: Final = agentic_loop_internal_litellm_params + ["caller_added"] # mutable-ok: list contract under test + + assert (type(extended), len(extended), frozenset(extended)) == ( + list, + len(AGENTIC_LOOP_STATE_NAMES) + 2, + frozenset((*AGENTIC_LOOP_STATE_NAMES, "max_agentic_loops", "caller_added")), + ) + + +def test_bedrock_batch_names_concatenate_as_a_tuple() -> None: + extended: Final = bedrock_batch_litellm_params + ("caller_added",) + + assert extended == (*BEDROCK_BATCH_NAMES, "caller_added") + + +def test_proxy_stamped_fields_keep_their_wire_names() -> None: + assert (TRUSTED_CALLBACK_VARS_FIELD, ADDRESSED_RESPONSE_ID_FIELD) == ( + "litellm_trusted_callback_vars", + "_litellm_addressed_response_id", + ) + + +def test_all_litellm_params_concatenates_with_a_list_like_the_completion_entrypoint_does() -> None: + extended: Final = ["aembedding", "extra_headers"] + all_litellm_params # mutable-ok: list contract under test + + assert (type(extended), frozenset(extended)) == (list, frozenset(("aembedding", "extra_headers", *OWNED_NAMES))) + + +CARRIED_AND_FORWARDED: Final = frozenset(("drop_params", "hugging_face", "no_log", "replicate", "together_ai")) + +CARRIER_SIGNATURE: Final = inspect.signature(get_litellm_params) # pyright: ignore[reportUnknownArgumentType] # legacy + +CARRIED_PARAMS: Final = tuple( + name for name in CARRIER_SIGNATURE.parameters if name != "kwargs" and name not in CARRIED_AND_FORWARDED +) + + +@pytest.mark.parametrize("name", CARRIED_PARAMS) +def test_every_param_get_litellm_params_carries_is_kept_out_of_provider_params(name: str) -> None: + provider_value: Final = object() + + result: Final = CLASSIFIERS["completion"]( + {name: object(), PROVIDER_KNOB: provider_value} # mutable-ok: classifier input type + ) + + assert result == MappingProxyType({PROVIDER_KNOB: provider_value}) + + +TYPED_CONFIG_MODELS: Final[Mapping[str, tuple[type[BaseModel], ...]]] = MappingProxyType( + { + "credentials": (CredentialLiteLLMParams,), + "router": (RouterConfig, UpdateRouterConfig), + } +) + +DECLARED_NAMES: Final = frozenset(name for root in LITELLM_OWNED_ROOTS for name in owned_wire_names(root)) + +ProviderClient: TypeAlias = ( + OpenAI + | AsyncOpenAI + | AzureOpenAI + | AsyncAzureOpenAI + | HTTPHandler + | AsyncHTTPHandler + | httpx.Client + | httpx.AsyncClient +) +MockResponse: TypeAlias = str | Exception | Mapping[str, object] | Sequence[float] | ModelResponse | ModelResponseStream + +TYPE_HINT_NAMESPACE: Final[Mapping[str, object]] = { + "ProviderClient": ProviderClient, + "ProviderSpecificHeader": ProviderSpecificHeader, + "ClientSession": ClientSession, + "AsyncAzureOpenAI": AsyncAzureOpenAI, + "AsyncOpenAI": AsyncOpenAI, + "AzureOpenAI": AzureOpenAI, + "OpenAI": OpenAI, + "AsyncHTTPHandler": AsyncHTTPHandler, + "HTTPHandler": HTTPHandler, + "ConfigurableClientsideParamsCustomAuth": ConfigurableClientsideParamsCustomAuth, + "RetryPolicy": RetryPolicy, + "DeploymentTypedDict": DeploymentTypedDict, + "DynamicCacheControl": DynamicCacheControl, + "ChatCompletionUserMessage": ChatCompletionUserMessage, + "ChatCompletionAssistantMessage": ChatCompletionAssistantMessage, + "MockResponse": MockResponse, + "ModelResponse": ModelResponse, + "ModelResponseStream": ModelResponseStream, + "Logging": Logging, + "SecretFields": SecretFields, + "CompactionState": CompactionState, + "RouterWeights": RouterWeights, + "AttemptedFallbackTargets": AttemptedFallbackTargets, +} + +LEAF_SAMPLES: Final[Mapping[type, Mapping[str, object]]] = { + litellm_params.ProviderConnection: {"api_key": "k", "request_timeout": 1.5}, + litellm_params.BedrockBatchConnection: {"aws_batch_role_arn": "arn", "bedrock_tags": ({"k": "v"},)}, + litellm_params.DispatchOptions: {"custom_llm_provider": "openai"}, + litellm_params.RoutingOptions: { + "fallbacks": [{"model": "gpt-4o", "api_key": "k", "temperature": 0}], + "num_retries": 2, + "retry_strategy": "constant_retry", + "routing_strategy": "simple-shuffle", + }, + litellm_params.DeploymentOptions: {"model_info": {"region": "us"}, "rpm": 2}, + litellm_params.SpecializedRouterOptions: {"adaptive_router_default_model": "gpt-4o"}, + litellm_params.CachingOptions: {"ttl": 30.0, "caching_groups": (("gpt-4o", "gpt-4o-mini"),)}, + litellm_params.CostOptions: {"max_budget": 10.0}, + litellm_params.ObservabilityOptions: {"metadata": {"request": "test"}, "no_log": True}, + litellm_params.AgenticLoopOptions: {"max_agentic_loops": 2}, + litellm_params.GuardrailOptions: {"guardrails": ("default",)}, + litellm_params.PromptOptions: {"prompt_id": "prompt", "prompt_variables": {"name": "value"}}, + litellm_params.ResponseOptions: {"stream_chunk_size": 64}, + litellm_params.MockOptions: {"mock_timeout": True}, + litellm_params.CallState: { + "completion_call_id": "call", + "model_alias_map": {"alias": "gpt-4o"}, + "data_residency": "us", + }, + litellm_params.AgenticLoopState: {"api_surface": "chat_completions", "depth": 1}, + litellm_params.RouterState: {"fallback_depth": 1}, + litellm_params.ProxyRequestState: { + "proxy_server_request": {"path": "/chat/completions"}, + "trusted_callback_vars": {"dd_api_key": "k"}, + }, + litellm_params.EntrypointState: {"acompletion": True}, +} + +LEAF_BAD_SAMPLES: Final[Mapping[type, Mapping[str, object]]] = { + litellm_params.ProviderConnection: {"api_key": 1}, + litellm_params.BedrockBatchConnection: {"aws_batch_role_arn": 1}, + litellm_params.DispatchOptions: {"custom_llm_provider": 1}, + litellm_params.RoutingOptions: {"num_retries": "2"}, + litellm_params.DeploymentOptions: {"rpm": "2"}, + litellm_params.SpecializedRouterOptions: {"auto_router_max_input_chars": "2"}, + litellm_params.CachingOptions: {"ttl": "30"}, + litellm_params.CostOptions: {"max_budget": "10"}, + litellm_params.ObservabilityOptions: {"verbose": "true"}, + litellm_params.AgenticLoopOptions: {"max_agentic_loops": "2"}, + litellm_params.GuardrailOptions: {"guardrails": (1,)}, + litellm_params.PromptOptions: {"prompt_id": 1}, + litellm_params.ResponseOptions: {"stream_chunk_size": "64"}, + litellm_params.MockOptions: {"mock_timeout": "true"}, + litellm_params.CallState: {"completion_call_id": 1}, + litellm_params.AgenticLoopState: {"depth": "1"}, + litellm_params.RouterState: {"fallback_depth": "1"}, + litellm_params.ProxyRequestState: {"proxy_server_request": "request"}, + litellm_params.EntrypointState: {"acompletion": "true"}, +} + +INVALID_LITERAL_SAMPLES: Final[tuple[tuple[type, Mapping[str, object]], ...]] = ( + (litellm_params.RoutingOptions, {"retry_strategy": "linear"}), + (litellm_params.RoutingOptions, {"routing_strategy": "random"}), + (litellm_params.AgenticLoopState, {"api_surface": "batches"}), +) + + +def _leaf_id(value: object) -> str: + return value.__name__ if isinstance(value, type) else "" + + +def _leaf_instance(leaf: type, sample: Mapping[str, object]) -> object: + constructor: Final = cast(Callable[..., object], leaf) + return constructor(**sample) + + +def _strict_leaf_validation(leaf: type, instance: object) -> object: + hints: Final[Mapping[str, object]] = cast( + Mapping[str, object], get_type_hints(type(instance), localns=TYPE_HINT_NAMESPACE) + ) + for field_info in fields(leaf): + value = cast(Callable[[object], object], attrgetter(field_info.name))(instance) + field_adapter: TypeAdapter[object] = TypeAdapter[object]( + hints[field_info.name], + config=ConfigDict(arbitrary_types_allowed=True), + ) + field_adapter.validate_python(value, strict=True) + return instance + + +@pytest.mark.parametrize("leaf,sample", LEAF_SAMPLES.items(), ids=_leaf_id) +def test_every_owned_leaf_accepts_a_strict_reader_shaped_sample(leaf: type, sample: Mapping[str, object]) -> None: + instance: Final = _leaf_instance(leaf, sample) + result: Final = _strict_leaf_validation(leaf, instance) + + assert result == instance + assert frozenset(sample) <= frozenset(field.name for field in fields(leaf)) + + +@pytest.mark.parametrize("leaf,sample", LEAF_BAD_SAMPLES.items(), ids=_leaf_id) +def test_every_owned_leaf_rejects_a_strict_wrong_typed_sample(leaf: type, sample: Mapping[str, object]) -> None: + instance: Final = _leaf_instance(leaf, sample) + + with pytest.raises(ValidationError): + _strict_leaf_validation(leaf, instance) + + +@pytest.mark.parametrize("leaf,sample", INVALID_LITERAL_SAMPLES, ids=_leaf_id) +def test_owned_leaf_literals_reject_unknown_values(leaf: type, sample: Mapping[str, object]) -> None: + instance: Final = _leaf_instance(leaf, sample) + + with pytest.raises(ValidationError): + _strict_leaf_validation(leaf, instance) + + +@pytest.mark.parametrize( + "strategy", + [ + "simple-shuffle", + "least-busy", + "usage-based-routing", + "latency-based-routing", + "cost-based-routing", + "usage-based-routing-v2", + "lar1", + ], +) +def test_routing_options_accept_every_strategy_the_router_accepts(strategy: str) -> None: + instance: Final = _leaf_instance(litellm_params.RoutingOptions, {"routing_strategy": strategy}) + + assert _strict_leaf_validation(litellm_params.RoutingOptions, instance) is instance + + +NAMES_SHARED_WITH_TYPED_MODELS: Final[Mapping[str, tuple[str, ...]]] = MappingProxyType( + { + "credentials": ( + "api_base", + "api_key", + "api_version", + "aws_batch_role_arn", + "azure_password", + "azure_scope", + "azure_username", + "bedrock_tags", + "client_id", + "client_secret", + "region_name", + "s3_access_key_id", + "s3_bucket_name", + "s3_bucket_owner", + "s3_encryption_key_id", + "s3_endpoint_url", + "s3_output_bucket_name", + "s3_region_name", + "s3_secret_access_key", + "tenant_id", + ), + "router": ( + "caching_groups", + "cooldown_time", + "enable_tag_filtering", + "fallbacks", + "max_retries", + "model_list", + "num_retries", + "retry_policy", + "routing_strategy", + ), + } +) + + +@pytest.mark.parametrize("source", TYPED_CONFIG_MODELS) +def test_names_a_typed_config_model_shares_with_the_owned_inventory_are_exactly_these(source: str) -> None: + model_names: Final = frozenset(name for model in TYPED_CONFIG_MODELS[source] for name in model.model_fields) + + assert DECLARED_NAMES & model_names == frozenset(NAMES_SHARED_WITH_TYPED_MODELS[source]) + + +@pytest.mark.parametrize("name", PRICING_NAMES) +def test_pricing_name_is_owned_by_the_pricing_model_alone(name: str) -> None: + assert name not in DECLARED_NAMES diff --git a/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.test.tsx b/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.test.tsx index d18e1993388..2b01ef6fce7 100644 --- a/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.test.tsx @@ -134,9 +134,10 @@ describe("MCPToolPermissions", () => { const selectAllButton = screen.getByRole("button", { name: "Select All" }); await userEvent.click(selectAllButton); - // Verify onChange was called with all tools selected + // Selecting every displayed tool writes the wildcard, which also covers tools the + // server adds later. expect(mockOnChange).toHaveBeenCalledWith({ - [mockServerId]: ["read_wiki_structure", "read_wiki_contents", "ask_question"], + [mockServerId]: ["*"], }); }); @@ -191,6 +192,77 @@ describe("MCPToolPermissions", () => { }); }); + describe("wildcard all-tools grant", () => { + const wildcardServerId = "server-1"; + const wildcardServer = { server_id: wildcardServerId, server_name: "Wildcard Server", alias: "Wildcard Server" }; + const wildcardTools = [ + { name: "read_wiki_structure", description: "Get documentation topics" }, + { name: "read_wiki_contents", description: "View documentation" }, + { name: "ask_question", description: "Ask questions" }, + ]; + + beforeEach(() => { + vi.mocked(networking.fetchMCPServers).mockResolvedValue([wildcardServer]); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: wildcardTools, error: false }); + }); + + it("renders every tool checked with the future-tools note when the entry is the wildcard", async () => { + renderWithProviders( + , + ); + + expect(await screen.findByText("Wildcard Server")).toBeInTheDocument(); + expect(screen.getByText("All tools allowed, including tools added to this server later")).toBeInTheDocument(); + + await userEvent.click(screen.getByText("Flat List")); + for (const checkbox of screen.getAllByRole("checkbox")) { + expect(checkbox).toBeChecked(); + } + }); + + it("writes the wildcard when Select All covers every displayed tool", async () => { + const mockOnChange = vi.fn(); + renderWithProviders( + , + ); + + expect(await screen.findByText("read_wiki_structure")).toBeInTheDocument(); + await userEvent.click(screen.getByRole("button", { name: "Select All" })); + + expect(mockOnChange).toHaveBeenCalledWith({ [wildcardServerId]: ["*"] }); + }); + + it("converts back to an enumerated list when one tool is unchecked from a wildcard grant", async () => { + const mockOnChange = vi.fn(); + renderWithProviders( + , + ); + + expect(await screen.findByText("read_wiki_structure")).toBeInTheDocument(); + await userEvent.click(screen.getByText("Flat List")); + await userEvent.click(screen.getByRole("checkbox", { name: "ask_question" })); + + expect(mockOnChange).toHaveBeenCalledWith({ + [wildcardServerId]: ["read_wiki_structure", "read_wiki_contents"], + }); + }); + }); + describe("servers reached indirectly", () => { const groupServer = { server_id: "srv-group-1", @@ -432,6 +504,8 @@ describe("MCPToolPermissions", () => { expect(await screen.findByText("list_issues")).toBeInTheDocument(); await userEvent.click(screen.getByText("Select All")); + // A toolset-sourced server never writes the wildcard: that would create a standing direct + // grant outliving the toolset. The write keeps only the tools this level grants itself. expect(mockOnChange).toHaveBeenCalledWith({ [toolsetServer.server_id]: ["delete_issue"] }); }); @@ -853,7 +927,7 @@ describe("MCPToolPermissions", () => { const written = mockOnChange.mock.calls.at(-1)?.[0] as Record; expect(written["github_mcp"]).toEqual(["list_issues"]); - expect(written[twin.server_id]).toEqual(["list_issues", "create_issue", "delete_issue"]); + expect(written[twin.server_id]).toEqual(["*"]); }); it("says nothing about shared names when every key names one server", async () => { diff --git a/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.tsx b/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.tsx index 76806ab7fd6..dabc8700e4f 100644 --- a/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.tsx +++ b/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.tsx @@ -8,7 +8,7 @@ import { useMCPAccessGroups } from "../../app/(dashboard)/hooks/mcpServers/useMC import { useMCPToolsets } from "../../app/(dashboard)/hooks/mcpServers/useMCPToolsets"; import McpCrudPermissionPanel from "../mcp_tools/McpCrudPermissionPanel"; import { classifyToolOp } from "../../utils/mcpToolCrudClassification"; -import { NO_MCP_SERVERS_SENTINEL } from "../mcp_tools/constants"; +import { MCP_ALL_TOOLS_WILDCARD, NO_MCP_SERVERS_SENTINEL } from "../mcp_tools/constants"; import { EffectiveMcpServer, McpGrantSource, @@ -18,6 +18,8 @@ import { applyToolPermissionWrite, emptyMcpAccessGroups, isConventionServer, + mcpAllowedToolsFor, + mcpGrantsAllTools, mcpToolState, resolveEffectiveMcpServers, } from "./effectiveMcpServers"; @@ -135,7 +137,12 @@ const MCPToolPermissions: React.FC = ({ }, [servers, accessToken, toolsetsLoading]); const writeAllowedTools = (entry: EffectiveMcpServer, allowed: string[]) => { - onChange(applyToolPermissionWrite({ toolPermissions, entry, allowed })); + const names = (serverTools[entry.server.server_id] ?? []).map((t) => t.name); + const next = + entry.source.kind !== "toolset" && names.length > 0 && names.every((n) => allowed.includes(n)) + ? [MCP_ALL_TOOLS_WILDCARD] + : allowed; + onChange(applyToolPermissionWrite({ toolPermissions, entry, allowed: next })); }; const isDelete = (tool: MCPTool) => classifyToolOp(tool.name, tool.description || "") === "delete"; @@ -152,7 +159,9 @@ const MCPToolPermissions: React.FC = ({ onOverridesChange?.(applyToolOverrideWrite(write)); return; } - const current = entry.allowedTools ?? (serverTools[entry.server.server_id] || []).map((t) => t.name); + const current = mcpGrantsAllTools(entry.keyedTools) + ? (serverTools[entry.server.server_id] || []).map((t) => t.name) + : entry.allowedTools ?? (serverTools[entry.server.server_id] || []).map((t) => t.name); writeAllowedTools(entry, checked ? [...current, tool.name] : current.filter((name) => name !== tool.name)); }; @@ -238,9 +247,12 @@ const MCPToolPermissions: React.FC = ({ const serverId = server.server_id; const serverName = server.server_name || server.alias || serverId; const tools = serverTools[serverId] || []; + const grantsAll = mcpGrantsAllTools(entry.keyedTools); const stateFor = (tool: MCPTool) => mcpToolState(entry, tool.name, classifyToolOp(tool.name, tool.description || "") === "delete"); - const selectedTools = tools.filter((tool) => stateFor(tool).checked).map((tool) => tool.name); + const selectedTools = grantsAll + ? tools.map((tool) => tool.name) + : tools.filter((tool) => stateFor(tool).checked).map((tool) => tool.name); const isLoading = loadingTools[serverId]; const error = toolErrors[serverId]; const viewMode = viewModes[serverId] ?? "crud"; @@ -265,6 +277,11 @@ const MCPToolPermissions: React.FC = ({ )} {server.description &&

{server.description}

} + {grantsAll && ( +

+ All tools allowed, including tools added to this server later +

+ )} {entry.ambiguousKeys.length > 0 && (

{`Also granted by ${entry.ambiguousKeys.map((key) => `"${key}"`).join(", ")}, which names another server too. Those tools stay allowed here until the servers no longer share that name`} @@ -376,10 +393,10 @@ const MCPToolPermissions: React.FC = ({ { if (disabled || state.locked) return; - writeToolToggle(entry, tool, !state.checked); + writeToolToggle(entry, tool, !(grantsAll || state.checked)); }} disabled={disabled || state.locked} className="mt-0.5" diff --git a/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.test.ts b/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.test.ts index f262b95ff50..d42bf8f736f 100644 --- a/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.test.ts +++ b/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.test.ts @@ -7,6 +7,7 @@ import { emptyMcpAccessGroups, isConventionServer, mcpAllowedToolsFor, + mcpGrantsAllTools, mcpServersForIdentifier, mcpToolOverridesFor, mcpToolPermissionKeyFor, @@ -72,6 +73,16 @@ describe("mcpServersForIdentifier", () => { }); }); +describe("mcpGrantsAllTools", () => { + it("is true only when the union carries the wildcard, never for an absent grant", () => { + expect(mcpGrantsAllTools(["*"])).toBe(true); + expect(mcpGrantsAllTools(["read_file", "*"])).toBe(true); + expect(mcpGrantsAllTools(["read_file"])).toBe(false); + expect(mcpGrantsAllTools([])).toBe(false); + expect(mcpGrantsAllTools(undefined)).toBe(false); + }); +}); + describe("mcpToolPermissionKeyFor", () => { const target = server({ server_id: "uuid-1", server_name: "github_mcp", alias: "GitHub" }); diff --git a/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.ts b/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.ts index d74705be4c7..5d8dd336479 100644 --- a/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.ts +++ b/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.ts @@ -1,5 +1,6 @@ import { z } from "zod/v4"; import { MCPServer, MCPToolset } from "../mcp_tools/types"; +import { MCP_ALL_TOOLS_WILDCARD } from "../mcp_tools/constants"; // Mirrors the backend resolver's union (direct + access_group + tool_perm + toolset), so the // editor shows exactly the servers this permission level entitles. @@ -156,6 +157,12 @@ export const mcpToolOverridesFor = ( }; }; +// An allowed-tools union carrying the wildcard grants every current and future tool on the +// server; `undefined` (no entry at all) is unrestricted for a different reason and is not a +// wildcard grant the editor should expand. +export const mcpGrantsAllTools = (allowed: readonly string[] | undefined): boolean => + allowed !== undefined && allowed.includes(MCP_ALL_TOOLS_WILDCARD); + // Tool names the given toolsets grant on this server, `undefined` when they grant none. const mcpToolsetToolsFor = ( server: MCPServer, diff --git a/ui/litellm-dashboard/src/components/mcp_tools/constants.ts b/ui/litellm-dashboard/src/components/mcp_tools/constants.ts index 66ab1a352f4..eef98383d2a 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/constants.ts +++ b/ui/litellm-dashboard/src/components/mcp_tools/constants.ts @@ -3,5 +3,8 @@ export const NO_MCP_SERVERS_SENTINEL = "no-mcp-servers"; export const ALL_PROXY_MCP_SERVERS_SENTINEL = "all-proxy-mcpservers"; +// Must match the backend MCP_ALL_TOOLS_WILDCARD constant in litellm/constants.py. +export const MCP_ALL_TOOLS_WILDCARD = "*"; + export const MCP_TOOLS_PREVIEW_FORBIDDEN_MESSAGE = "Tool preview is not available for submissions. Tools will be verified by an admin during review."; diff --git a/uv.lock b/uv.lock index 85e2b6d4e52..8f63ca2b564 100644 --- a/uv.lock +++ b/uv.lock @@ -4959,12 +4959,12 @@ proxy-dev = [ [[package]] name = "litellm-enterprise" -version = "0.1.70" +version = "0.1.71" source = { editable = "enterprise" } [[package]] name = "litellm-proxy-extras" -version = "0.4.101" +version = "0.4.102" source = { editable = "litellm-proxy-extras" } [[package]]