diff --git a/litellm-rust/crates/core/src/audio_transcription/lifecycle.rs b/litellm-rust/crates/core/src/audio_transcription/lifecycle.rs index 6e744054dd5..87f56caa486 100644 --- a/litellm-rust/crates/core/src/audio_transcription/lifecycle.rs +++ b/litellm-rust/crates/core/src/audio_transcription/lifecycle.rs @@ -36,9 +36,17 @@ pub struct AudioTranscriptionRoute; pub type AudioTranscriptionCall = CompletedCall; impl CompletedRoute for AudioTranscriptionRoute { + type Admission = + crate::call_lifecycle::admission::Inspection; type Request = OwnedAudioTranscriptionRequest; type Response = Value; + fn admit( + admission: Self::Admission, + ) -> Result<(), crate::call_lifecycle::admission::AdmissionDecline> { + super::admit(admission) + } + fn run( request: Self::Request, hooks: Arc, diff --git a/litellm-rust/crates/core/src/audio_transcription/mod.rs b/litellm-rust/crates/core/src/audio_transcription/mod.rs index eb4335c1d96..a9d1804b234 100644 --- a/litellm-rust/crates/core/src/audio_transcription/mod.rs +++ b/litellm-rust/crates/core/src/audio_transcription/mod.rs @@ -12,6 +12,12 @@ pub use handler::execute_audio_transcription_provider_call; pub use prepare::prepare_audio_transcription_provider_call; pub use types::{AudioTranscriptionRequest, ProviderAudioTranscriptionRequest}; +pub struct AudioTranscriptionAdmission { + pub model: String, + pub provider: Option, + pub audio: Value, +} + #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Result { crate::call_lifecycle::provider::run_completed::( @@ -21,17 +27,24 @@ pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Resu } pub fn admit( - model: &str, - provider: Option<&str>, - audio: &Value, + inspection: crate::call_lifecycle::admission::Inspection, ) -> Result<(), crate::call_lifecycle::admission::AdmissionDecline> { - use crate::call_lifecycle::admission::AdmissionDecline; - let resolved = crate::routing_utils::provider::get_custom_llm_provider(model, provider); - let provider = provider.or_else(|| resolved.as_ref().map(|value| value.custom_llm_provider)); + use crate::call_lifecycle::admission::{AdmissionDecline, Inspection}; + let Inspection::Inspectable(admission) = inspection else { + return Err(AdmissionDecline::Uninspectable); + }; + let resolved = crate::routing_utils::provider::get_custom_llm_provider( + &admission.model, + admission.provider.as_deref(), + ); + let provider = admission + .provider + .as_deref() + .or_else(|| resolved.as_ref().map(|value| value.custom_llm_provider)); if provider.and_then(prepare::provider_config).is_none() { return Err(AdmissionDecline::Provider); } - if let Some(format) = audio.get("format").and_then(Value::as_str) + if let Some(format) = admission.audio.get("format").and_then(Value::as_str) && !matches!(format, "wav" | "mp3" | "flac" | "ogg") { return Err(AdmissionDecline::Feature("unsupported audio format")); diff --git a/litellm-rust/crates/core/src/audio_transcription/tests.rs b/litellm-rust/crates/core/src/audio_transcription/tests.rs index 263d63337b0..98c48f52a51 100644 --- a/litellm-rust/crates/core/src/audio_transcription/tests.rs +++ b/litellm-rust/crates/core/src/audio_transcription/tests.rs @@ -4,8 +4,16 @@ use std::thread; use serde_json::{Map, json}; -use super::audio_transcription; use super::types::AudioTranscriptionRequest; +use super::{admit, audio_transcription}; + +#[test] +fn uninspectable_request_declines_in_core() { + assert_eq!( + admit(crate::call_lifecycle::admission::Inspection::Uninspectable), + Err(crate::call_lifecycle::admission::AdmissionDecline::Uninspectable) + ); +} #[tokio::test] async fn bedrock_request_is_signed_and_contains_audio() { diff --git a/litellm-rust/crates/core/src/call_lifecycle/admission.rs b/litellm-rust/crates/core/src/call_lifecycle/admission.rs index 138fb83a1ed..a064972fe95 100644 --- a/litellm-rust/crates/core/src/call_lifecycle/admission.rs +++ b/litellm-rust/crates/core/src/call_lifecycle/admission.rs @@ -30,3 +30,8 @@ pub enum AdmissionDecline { #[strum(to_string = "{0}")] Feature(&'static str), } + +pub enum Inspection { + Inspectable(T), + Uninspectable, +} diff --git a/litellm-rust/crates/core/src/call_lifecycle/provider.rs b/litellm-rust/crates/core/src/call_lifecycle/provider.rs index 7a56851787e..5173ca0bf2b 100644 --- a/litellm-rust/crates/core/src/call_lifecycle/provider.rs +++ b/litellm-rust/crates/core/src/call_lifecycle/provider.rs @@ -146,9 +146,13 @@ pub enum CompletedReply { } pub trait CompletedRoute: Send + Sync + 'static { + type Admission; type Request: Send + Sync + 'static; type Response: Clone + Send + Sync + serde::de::DeserializeOwned + 'static; + fn admit( + admission: Self::Admission, + ) -> Result<(), crate::call_lifecycle::admission::AdmissionDecline>; fn run(request: Self::Request, hooks: Arc) -> WorkflowFuture; fn context(request: &Self::Request) -> CallLifecycleContext; @@ -473,9 +477,16 @@ mod tests { struct TestRoute; impl CompletedRoute for TestRoute { + type Admission = (); type Request = (); type Response = TestResponse; + fn admit( + (): Self::Admission, + ) -> Result<(), crate::call_lifecycle::admission::AdmissionDecline> { + Ok(()) + } + fn run((): Self::Request, hooks: Arc) -> WorkflowFuture { Box::pin(async move { let request = hooks diff --git a/litellm-rust/crates/core/src/chat_completions/lifecycle.rs b/litellm-rust/crates/core/src/chat_completions/lifecycle.rs index ca6bd37fe47..904ac661139 100644 --- a/litellm-rust/crates/core/src/chat_completions/lifecycle.rs +++ b/litellm-rust/crates/core/src/chat_completions/lifecycle.rs @@ -36,9 +36,16 @@ pub struct ChatCompletionsRoute; pub type ChatCompletionsCall = CompletedCall; impl CompletedRoute for ChatCompletionsRoute { + type Admission = crate::call_lifecycle::admission::Inspection; type Request = OwnedChatCompletionsRequest; type Response = ChatCompletionsResponse; + fn admit( + admission: Self::Admission, + ) -> Result<(), crate::call_lifecycle::admission::AdmissionDecline> { + super::admit(admission) + } + fn run( request: Self::Request, hooks: Arc, diff --git a/litellm-rust/crates/core/src/chat_completions/mod.rs b/litellm-rust/crates/core/src/chat_completions/mod.rs index c0f325712c3..901ca8ea7b0 100644 --- a/litellm-rust/crates/core/src/chat_completions/mod.rs +++ b/litellm-rust/crates/core/src/chat_completions/mod.rs @@ -69,28 +69,41 @@ pub struct AdmissionContext { pub bedrock_metadata_owned: bool, } +pub struct ChatCompletionsAdmission { + pub model: String, + pub provider: Option, + pub messages: Value, + pub params: Map, + pub headers: Option>, + pub context: AdmissionContext, +} + pub fn admit( - model: &str, - provider: Option<&str>, - messages: Value, - params: &Map, - headers: Option<&Map>, - context: AdmissionContext, + inspection: crate::call_lifecycle::admission::Inspection, ) -> Result<(), crate::call_lifecycle::admission::AdmissionDecline> { - use crate::call_lifecycle::admission::AdmissionDecline; - let resolved = crate::routing_utils::provider::get_custom_llm_provider(model, provider); - let provider = provider.or_else(|| resolved.as_ref().map(|value| value.custom_llm_provider)); - if context.stream { + use crate::call_lifecycle::admission::{AdmissionDecline, Inspection}; + let Inspection::Inspectable(admission) = inspection else { + return Err(AdmissionDecline::Uninspectable); + }; + let resolved = crate::routing_utils::provider::get_custom_llm_provider( + &admission.model, + admission.provider.as_deref(), + ); + let provider = admission + .provider + .as_deref() + .or_else(|| resolved.as_ref().map(|value| value.custom_llm_provider)); + if admission.context.stream { return Err(AdmissionDecline::Feature("streaming")); } - if (provider == Some("anthropic") && context.anthropic_user_id) - || (provider == Some("bedrock") && context.bedrock_metadata_owned) + if (provider == Some("anthropic") && admission.context.anthropic_user_id) + || (provider == Some("bedrock") && admission.context.bedrock_metadata_owned) { return Err(AdmissionDecline::HostOperations); } #[cfg(feature = "bedrock-auth")] if provider == Some("bedrock") - && headers.is_some_and(|headers| { + && admission.headers.as_ref().is_some_and(|headers| { headers .keys() .any(|name| crate::providers::bedrock::aws_base::is_sigv4_computed_header(name)) @@ -100,8 +113,13 @@ pub fn admit( "request forwards a header AWS SigV4 computes", )); } - let _ = headers; - match chat_completions_decline_reason(model, provider, messages, params) { + let _ = admission.headers; + match chat_completions_decline_reason( + &admission.model, + provider, + admission.messages, + &admission.params, + ) { Some(reason) => Err(AdmissionDecline::Feature(reason)), None => Ok(()), } diff --git a/litellm-rust/crates/core/src/chat_completions/tests.rs b/litellm-rust/crates/core/src/chat_completions/tests.rs index f8594dee447..a384a5e549b 100644 --- a/litellm-rust/crates/core/src/chat_completions/tests.rs +++ b/litellm-rust/crates/core/src/chat_completions/tests.rs @@ -2,10 +2,19 @@ use serde_json::{Map, Value, json}; use crate::error::Error; +use super::admit; use super::prepare::{prepare_provider_request, resolve_request}; use super::transformation::ChatCompletionsAuth; use super::types::{ChatCompletionsRequest, ProviderChatCompletionsRequest}; +#[test] +fn uninspectable_request_declines_in_core() { + assert_eq!( + admit(crate::call_lifecycle::admission::Inspection::Uninspectable), + Err(crate::call_lifecycle::admission::AdmissionDecline::Uninspectable) + ); +} + fn prepare_chat_completions_call( request: ChatCompletionsRequest<'_>, ) -> Result { diff --git a/litellm-rust/crates/core/src/messages/lifecycle.rs b/litellm-rust/crates/core/src/messages/lifecycle.rs index 07dedb79cf0..fab3e224928 100644 --- a/litellm-rust/crates/core/src/messages/lifecycle.rs +++ b/litellm-rust/crates/core/src/messages/lifecycle.rs @@ -34,9 +34,16 @@ pub struct MessagesRoute; pub type MessagesCall = CompletedCall; impl CompletedRoute for MessagesRoute { + type Admission = crate::call_lifecycle::admission::Inspection; type Request = OwnedMessagesRequest; type Response = AnthropicMessagesResponse; + fn admit( + admission: Self::Admission, + ) -> Result<(), crate::call_lifecycle::admission::AdmissionDecline> { + super::admit(admission) + } + fn run( request: Self::Request, hooks: Arc, diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index 9b817b4b69a..619fc6ce3b4 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -19,6 +19,12 @@ pub mod types; use handler::execute_messages_provider_stream; use types::{AnthropicMessagesResponse, MessagesRequest}; +pub struct MessagesAdmission { + pub model: String, + pub provider: Option, + pub has_agentic_hook: bool, +} + #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] pub async fn messages(request: MessagesRequest<'_>) -> Result { crate::call_lifecycle::provider::run_completed::(request.into()).await @@ -29,20 +35,27 @@ pub async fn messages_stream(request: MessagesRequest<'_>) -> Result, - has_agentic_hook: bool, + inspection: crate::call_lifecycle::admission::Inspection, ) -> Result<(), crate::call_lifecycle::admission::AdmissionDecline> { - use crate::call_lifecycle::admission::AdmissionDecline; - let resolved = crate::routing_utils::provider::get_custom_llm_provider(model, provider); - let provider = provider.or_else(|| resolved.as_ref().map(|value| value.custom_llm_provider)); + use crate::call_lifecycle::admission::{AdmissionDecline, Inspection}; + let Inspection::Inspectable(admission) = inspection else { + return Err(AdmissionDecline::Uninspectable); + }; + let resolved = crate::routing_utils::provider::get_custom_llm_provider( + &admission.model, + admission.provider.as_deref(), + ); + let provider = admission + .provider + .as_deref() + .or_else(|| resolved.as_ref().map(|value| value.custom_llm_provider)); if provider .and_then(common_utils::messages_provider_config) .is_none() { return Err(AdmissionDecline::Provider); } - if has_agentic_hook { + if admission.has_agentic_hook { return Err(AdmissionDecline::HostOperations); } Ok(()) diff --git a/litellm-rust/crates/core/src/messages/tests.rs b/litellm-rust/crates/core/src/messages/tests.rs index df9f7051011..e0464899947 100644 --- a/litellm-rust/crates/core/src/messages/tests.rs +++ b/litellm-rust/crates/core/src/messages/tests.rs @@ -9,8 +9,16 @@ use crate::error::Error; use super::common_utils::{ has_bearer_auth, has_header, messages_provider_config, string_headers, truncate_error_body, }; -use super::messages; use super::types::MessagesRequest; +use super::{admit, messages}; + +#[test] +fn uninspectable_request_declines_in_core() { + assert_eq!( + admit(crate::call_lifecycle::admission::Inspection::Uninspectable), + Err(crate::call_lifecycle::admission::AdmissionDecline::Uninspectable) + ); +} async fn read_http_request(socket: &mut TcpStream) -> String { let mut request = Vec::new(); diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/completed.rs b/litellm-rust/crates/python-bridge/src/lifecycle/completed.rs index 49329bcdef6..69f05d4d818 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle/completed.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/completed.rs @@ -23,7 +23,7 @@ pub(crate) trait PythonCompletedRoute: CompletedRoute { const SYNC_CALL_TYPE: PythonCallType; const ASYNC_CALL_TYPE: PythonCallType; - fn admit(request: &Bound<'_, PyDict>) -> PyResult<()>; + fn project_admission(request: &Bound<'_, PyDict>) -> PyResult; fn project(request: &Bound<'_, PyDict>) -> PyResult; } @@ -242,7 +242,7 @@ pub(crate) fn run( where R::Response: Serialize, { - R::admit(&request)?; + crate::errors::admit(R::admit(R::project_admission(&request)?))?; let controls = crate::cache::snapshot( py, if asynchronous { diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions/lifecycle.rs index 588b1a06309..fa806bfae8d 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions/lifecycle.rs @@ -1,3 +1,5 @@ +use litellm_core::call_lifecycle::admission::Inspection; +use litellm_core::chat_completions::ChatCompletionsAdmission; use litellm_core::chat_completions::lifecycle::{ ChatCompletionsRoute, OwnedChatCompletionsRequest, }; @@ -15,7 +17,7 @@ impl PythonCompletedRoute for ChatCompletionsRoute { const SYNC_CALL_TYPE: PythonCallType = PythonCallType::Completion; const ASYNC_CALL_TYPE: PythonCallType = PythonCallType::AsyncCompletion; - fn admit(request: &Bound<'_, PyDict>) -> PyResult<()> { + fn project_admission(request: &Bound<'_, PyDict>) -> PyResult { let model = required(request, RequestField::Model)?; let provider = request.get_item(RequestField::CustomLlmProvider.key(request.py()))?; let messages = required(request, RequestField::Messages)?; @@ -29,26 +31,24 @@ impl PythonCompletedRoute for ChatCompletionsRoute { || !exact_optional_object(headers.as_ref()) || !exact_optional_object(facts.as_ref()) { - return crate::errors::admit(Err( - litellm_core::call_lifecycle::admission::AdmissionDecline::Uninspectable, - )); + return Ok(Inspection::Uninspectable); } let provider: Option = provider .as_ref() .map(|value| value.extract::>()) .transpose()? .flatten(); - crate::errors::admit(litellm_core::chat_completions::admit( - &model.extract::()?, - provider.as_deref(), - from_py(&messages)?, - &object(request, RequestField::OptionalParams)?, - Some(&object(request, RequestField::ExtraHeaders)?), - facts + Ok(Inspection::Inspectable(ChatCompletionsAdmission { + model: model.extract()?, + provider, + messages: from_py(&messages)?, + params: object(request, RequestField::OptionalParams)?, + headers: Some(object(request, RequestField::ExtraHeaders)?), + context: facts .map(|value| from_py(&value)) .transpose()? .unwrap_or_default(), - )) + })) } fn project(request: &Bound<'_, PyDict>) -> PyResult { diff --git a/litellm-rust/crates/python-bridge/src/routes/definition.rs b/litellm-rust/crates/python-bridge/src/routes/definition.rs index 94026410e3c..18402e2960c 100644 --- a/litellm-rust/crates/python-bridge/src/routes/definition.rs +++ b/litellm-rust/crates/python-bridge/src/routes/definition.rs @@ -347,6 +347,27 @@ mod tests { async_messages_error.to_string(), sync_messages_error.to_string() ); + + let invalid_audio = PyList::empty(py); + let request = PyDict::new(py); + request.set_item("model", "bedrock/model").unwrap(); + request.set_item("audio", &invalid_audio).unwrap(); + let sync_transcription_error = module + .getattr("transcription") + .and_then(|function| function.call1((&request, (), PyDict::new(py), py.None()))) + .expect_err("sync transcription should reject a non-dict audio value"); + let async_transcription_error = module + .getattr("atranscription") + .and_then(|function| function.call1((&request, (), PyDict::new(py), py.None()))) + .expect_err("async transcription should reject a non-dict audio value"); + + assert!( + sync_transcription_error.is_instance_of::(py) + ); + assert_eq!( + async_transcription_error.to_string(), + sync_transcription_error.to_string() + ); }); } diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/messages/lifecycle.rs index 47421136f33..bfe0d076b59 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/lifecycle.rs @@ -1,6 +1,8 @@ use pyo3::prelude::*; use pyo3::types::{PyDict, PyTuple}; +use litellm_core::call_lifecycle::admission::Inspection; +use litellm_core::messages::MessagesAdmission; use litellm_core::messages::lifecycle::{MessagesRoute, OwnedMessagesRequest}; use litellm_python_interop::from_py_preserving_errors as from_py; @@ -14,7 +16,7 @@ impl PythonCompletedRoute for MessagesRoute { const SYNC_CALL_TYPE: PythonCallType = PythonCallType::AnthropicMessages; const ASYNC_CALL_TYPE: PythonCallType = PythonCallType::AnthropicMessages; - fn admit(request: &Bound<'_, PyDict>) -> PyResult<()> { + fn project_admission(request: &Bound<'_, PyDict>) -> PyResult { let model = required(request, RequestField::Model)?; let provider = request.get_item(RequestField::CustomLlmProvider.key(request.py()))?; let body = request.get_item(RequestField::Body.key(request.py()))?; @@ -24,23 +26,21 @@ impl PythonCompletedRoute for MessagesRoute { || !exact_optional_object(body.as_ref()) || !exact_optional_bool(host_hook.as_ref()) { - return crate::errors::admit(Err( - litellm_core::call_lifecycle::admission::AdmissionDecline::Uninspectable, - )); + return Ok(Inspection::Uninspectable); } let provider: Option = provider .as_ref() .map(|value| value.extract::>()) .transpose()? .flatten(); - crate::errors::admit(litellm_core::messages::admit( - &model.extract::()?, - provider.as_deref(), - host_hook + Ok(Inspection::Inspectable(MessagesAdmission { + model: model.extract()?, + provider, + has_agentic_hook: host_hook .map(|value| value.extract()) .transpose()? .unwrap_or(false), - )) + })) } fn project(request: &Bound<'_, PyDict>) -> PyResult { diff --git a/litellm-rust/crates/python-bridge/src/routes/transcription/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/transcription/lifecycle.rs index 917a25d5ffc..4a300d3d65e 100644 --- a/litellm-rust/crates/python-bridge/src/routes/transcription/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/transcription/lifecycle.rs @@ -1,9 +1,11 @@ use pyo3::prelude::*; use pyo3::types::{PyDict, PyTuple}; +use litellm_core::audio_transcription::AudioTranscriptionAdmission; use litellm_core::audio_transcription::lifecycle::{ AudioTranscriptionRoute, OwnedAudioTranscriptionRequest, }; +use litellm_core::call_lifecycle::admission::Inspection; use litellm_python_interop::from_py_preserving_errors as from_py; use crate::lifecycle::completed::{self, PythonCompletedRoute}; @@ -16,7 +18,7 @@ impl PythonCompletedRoute for AudioTranscriptionRoute { const SYNC_CALL_TYPE: PythonCallType = PythonCallType::Transcription; const ASYNC_CALL_TYPE: PythonCallType = PythonCallType::AsyncTranscription; - fn admit(request: &Bound<'_, PyDict>) -> PyResult<()> { + fn project_admission(request: &Bound<'_, PyDict>) -> PyResult { let model = required(request, RequestField::Model)?; let provider = request.get_item(RequestField::CustomLlmProvider.key(request.py()))?; let audio_value = required(request, RequestField::Audio)?; @@ -26,16 +28,14 @@ impl PythonCompletedRoute for AudioTranscriptionRoute { || !exact_optional_object(Some(&audio_value)) || !exact_optional_object(optional_params.as_ref()) { - return crate::errors::admit(Err( - litellm_core::call_lifecycle::admission::AdmissionDecline::Uninspectable, - )); + return Ok(Inspection::Uninspectable); } let audio = from_py(&audio_value)?; - crate::errors::admit(litellm_core::audio_transcription::admit( - &model.extract::()?, - optional_string(request, RequestField::CustomLlmProvider)?.as_deref(), - &audio, - )) + Ok(Inspection::Inspectable(AudioTranscriptionAdmission { + model: model.extract()?, + provider: optional_string(request, RequestField::CustomLlmProvider)?, + audio, + })) } fn project(request: &Bound<'_, PyDict>) -> PyResult {