diff --git a/litellm-rust/crates/core/src/chat_completions/mod.rs b/litellm-rust/crates/core/src/chat_completions/mod.rs index 8f593689988..ed3909870fa 100644 --- a/litellm-rust/crates/core/src/chat_completions/mod.rs +++ b/litellm-rust/crates/core/src/chat_completions/mod.rs @@ -5,8 +5,7 @@ mod common_utils; pub(crate) mod handler; mod prepare; use litellm_types::utils::ChatCompletionsResponse; -use prepare::{parse_messages, prepare_provider_request, resolve_provider_config, resolve_request}; -use serde_json::{Map, Value}; +use prepare::{prepare_provider_request, resolve_request}; use crate::chat_completions::types::ChatCompletionsRequest; use litellm_auth::AuthServices; @@ -20,33 +19,6 @@ pub struct ChatCompletionsRoute { secrets: Arc, } -/// Whether the core would accept this request, without resolving credentials or -/// touching the network. -/// -/// A host that keeps the Python implementation asks this first so it can emit -/// its pre-call logging exactly once, on whichever path is about to run. -/// Returns the decline reason, or `None` when the request is accepted. -pub fn chat_completions_decline_reason( - model: &str, - custom_llm_provider: Option<&str>, - messages: Value, - optional_params: &Map, -) -> Option<&'static str> { - let Ok(resolved) = resolve_provider_config(model, custom_llm_provider) else { - return Some("provider is not on the rust chat completions path"); - }; - let config = resolved.config; - let Ok(messages) = parse_messages(messages) else { - return Some("unreadable message list"); - }; - if messages.is_empty() { - return Some("empty message list"); - } - config - .unsupported_reason(&messages, optional_params) - .map(|reason| reason.0) -} - impl ChatCompletionsRoute { pub fn new( http: litellm_http::Client, diff --git a/litellm-rust/crates/core/src/chat_completions/prepare.rs b/litellm-rust/crates/core/src/chat_completions/prepare.rs index c85ed8fed75..ec832b3e59a 100644 --- a/litellm-rust/crates/core/src/chat_completions/prepare.rs +++ b/litellm-rust/crates/core/src/chat_completions/prepare.rs @@ -185,13 +185,10 @@ mod tests { } } - /// `ProviderChatCompletionsRequest` deliberately has no `Debug` (its headers - /// carry resolved credentials), so unwrap the failure case by hand. - fn decline(request: ChatCompletionsRequest<'_>) -> Error { - match prepare_chat_completions_call(request) { - Err(error) => error, - Ok(prepared) => panic!("expected a decline, prepared a call to {}", prepared.url), - } + fn preparation_error(request: ChatCompletionsRequest<'_>) -> Error { + prepare_chat_completions_call(request) + .err() + .expect("request preparation should fail") } #[test] @@ -337,24 +334,19 @@ mod tests { ); } - #[test] - fn declines_an_unsupported_request_before_resolving_credentials() { - let mut call = request( - "claude-sonnet-4-5", - Some("anthropic"), - json!([{"role": "user", "content": "hi"}]), - json!({"stream": true}), - ); - call.api_key = None; - // No api_key is set and no env is consulted: the gate must run first, so the - // error is the decline rather than a missing-credential error. - assert_eq!(decline(call), Error::Unsupported("streaming")); + #[rstest::rstest] + fn rejects_empty_messages_before_resolving_credentials() { + let call = ChatCompletionsRequest { + api_key: None, + ..request("claude-sonnet-4-5", Some("anthropic"), json!([]), json!({})) + }; + assert!(matches!(preparation_error(call), Error::InvalidRequest(_))); } #[test] fn rejects_an_unknown_provider() { assert_eq!( - decline(request( + preparation_error(request( "openai/gpt-4o", None, json!([{"role": "user", "content": "hi"}]), @@ -367,7 +359,7 @@ mod tests { #[test] fn rejects_a_model_with_no_resolvable_provider() { assert!(matches!( - decline(request( + preparation_error(request( "claude-sonnet-4-5", None, json!([{"role": "user", "content": "hi"}]), @@ -380,7 +372,7 @@ mod tests { #[test] fn rejects_an_empty_or_malformed_message_list() { assert_eq!( - decline(request( + preparation_error(request( "anthropic/claude-sonnet-4-5", None, json!([]), @@ -393,7 +385,7 @@ mod tests { ) ); assert!(matches!( - decline(request( + preparation_error(request( "anthropic/claude-sonnet-4-5", None, json!("not a list"), @@ -413,7 +405,7 @@ mod tests { ); call.extra_headers = Some(Map::from_iter([("x-trace".to_string(), json!(7))])); assert_eq!( - decline(call), + preparation_error(call), Error::Headers(litellm_http::request::HeaderError { context: "chat completions", name: "x-trace".to_string(), @@ -511,18 +503,17 @@ mod tests { ); } + #[rstest::rstest] + #[case::authorization("Authorization")] + #[case::amz_date("x-amz-date")] + #[case::security_token("x-amz-security-token")] + #[case::date("Date")] #[tokio::test] - async fn a_forwarded_header_the_signer_computes_declines_to_python() { - // Reattaching the caller's copy next to the computed one puts the name on - // the wire twice and Bedrock rejects the pair, so a request carrying one - // has to go to Python instead of being signed here. - for forwarded in [ - "Authorization", - "x-amz-date", - "x-amz-security-token", - "Date", - ] { - let mut call = request( + async fn rejects_a_forwarded_header_the_signer_computes(#[case] forwarded: &str) { + let call = ChatCompletionsRequest { + api_key: None, + extra_headers: Some(Map::from_iter([(forwarded.to_string(), json!("forged"))])), + ..request( "bedrock/us-east-1/anthropic.claude-v2", None, json!([{"role": "user", "content": "hi"}]), @@ -531,29 +522,27 @@ mod tests { "aws_access_key_id": "AKIDEXAMPLE", "aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY" }), - ); - call.api_key = None; - call.extra_headers = Some(Map::from_iter([(forwarded.to_string(), json!("forged"))])); - let prepared = prepare_chat_completions_call(call).expect("prepares"); - let authenticated = resolve_auth( - &litellm_auth::AuthServices::default(), - prepared.environment, - &|_| None, ) - .await - .expect("resolves"); - let error = crate::chat_completions::handler::outbound_request( - authenticated, - prepared.url, - &prepared.body, - prepared.timeout, - ) - .expect_err("{forwarded} should decline instead of being signed"); - assert!( - matches!(error, Error::Unsupported(_)), - "{forwarded} declined as {error:?}, which the host would not fall back on" - ); - } + }; + let prepared = prepare_chat_completions_call(call).expect("prepares"); + let authenticated = resolve_auth( + &litellm_auth::AuthServices::default(), + prepared.environment, + &|_| None, + ) + .await + .expect("resolves"); + let error = crate::chat_completions::handler::outbound_request( + authenticated, + prepared.url, + &prepared.body, + prepared.timeout, + ) + .expect_err("conflicting signing headers must fail"); + assert!( + matches!(error, Error::Unsupported(_)), + "{forwarded} returned {error:?}" + ); } #[test] @@ -646,112 +635,4 @@ mod tests { "prepare did not carry the bearer token" ); } - - fn decline_reason( - model: &str, - provider: Option<&str>, - messages: Value, - params: Value, - ) -> Option<&'static str> { - let params = match params { - Value::Object(map) => map, - other => panic!("params must be an object, got {other}"), - }; - crate::chat_completions::chat_completions_decline_reason(model, provider, messages, ¶ms) - } - - #[test] - fn the_gate_accepts_what_prepare_accepts() { - assert_eq!( - decline_reason( - "anthropic/claude-sonnet-4-5", - None, - json!([{"role": "user", "content": "hi"}]), - json!({"max_tokens": 16}), - ), - None - ); - } - - #[test] - fn the_gate_declines_without_resolving_credentials_or_calling_out() { - assert_eq!( - decline_reason( - "anthropic/claude-sonnet-4-5", - None, - json!([{"role": "user", "content": "hi"}]), - json!({"stream": true}), - ), - Some("streaming") - ); - assert_eq!( - decline_reason( - "openai/gpt-4o", - None, - json!([{"role": "user", "content": "hi"}]), - json!({}), - ), - Some("provider is not on the rust chat completions path") - ); - assert_eq!( - decline_reason( - "claude-sonnet-4-5", - None, - json!([{"role": "user", "content": "hi"}]), - json!({}), - ), - Some("provider is not on the rust chat completions path") - ); - assert_eq!( - decline_reason( - "anthropic/claude-sonnet-4-5", - None, - json!("nope"), - json!({}) - ), - Some("unreadable message list") - ); - assert_eq!( - decline_reason("anthropic/claude-sonnet-4-5", None, json!([]), json!({})), - Some("empty message list") - ); - } - - #[test] - fn the_gate_agrees_with_prepare_on_every_case_it_accepts() { - // A gate that accepts what prepare then declines would make the host emit - // its pre-call logging on a path that falls back, so pin the agreement. - for (messages, params) in [ - ( - json!([{"role": "user", "content": "hi"}]), - json!({"max_tokens": 8}), - ), - ( - json!([{"role": "system", "content": "s"}, {"role": "user", "content": "hi"}]), - json!({"temperature": 0.1}), - ), - ( - json!([{"role": "user", "content": "hi"}, {"role": "assistant", "content": "yo"}]), - json!({}), - ), - ] { - assert_eq!( - decline_reason( - "anthropic/claude-sonnet-4-5", - None, - messages.clone(), - params.clone() - ), - None, - "gate declined {messages}" - ); - prepare_chat_completions_call(request( - "anthropic/claude-sonnet-4-5", - None, - messages.clone(), - params, - )) - .unwrap_or_else(|error| panic!("prepare declined {messages}: {error}")); - } - } } diff --git a/litellm-rust/crates/core/src/error.rs b/litellm-rust/crates/core/src/error.rs index a301c85fbeb..d42b38afe75 100644 --- a/litellm-rust/crates/core/src/error.rs +++ b/litellm-rust/crates/core/src/error.rs @@ -41,14 +41,6 @@ pub enum RouteError { PostCallHook(#[source] Arc), } -/// Whether the provider had already been called when the route failed. Before the send, a -/// host may retry on another path; after it, the provider has done the work and billed for it. -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum Phase { - BeforeSend, - AfterSend, -} - impl From for RouteError { fn from(fault: litellm_host::machine::MachineFault) -> Self { use litellm_host::machine::MachineFault; @@ -64,26 +56,6 @@ impl RouteError { Self::PostCallHook(Arc::new(error)) } - pub fn phase(&self) -> Phase { - match self { - Self::InvalidResponse(_) - | Self::PostCallHook(_) - | Self::Transport(TransportError::Http { .. } | TransportError::Network(_)) => { - Phase::AfterSend - } - Self::Transport(TransportError::Connect(_)) - | Self::InvalidType { .. } - | Self::MissingField(_) - | Self::InvalidProvider(_) - | Self::InvalidRequest(_) - | Self::Unsupported(_) - | Self::Auth(_) - | Self::Headers(_) - | Self::Http(_) - | Self::Secret(_) => Phase::BeforeSend, - } - } - /// The caller's request is what is wrong, as opposed to the environment, the wire, or /// the provider's answer. pub fn is_request(&self) -> bool { @@ -143,34 +115,10 @@ impl Eq for SecretError {} #[cfg(test)] mod tests { - use super::{Phase, RouteError}; - use litellm_http::transport::Error as TransportError; + use super::RouteError; use litellm_llms::{Error as LlmError, ErrorDetail}; use rstest::rstest; - #[test] - fn only_a_provider_answer_or_a_lost_connection_counts_as_after_send() { - let after = [ - RouteError::InvalidResponse("bad json".into()), - RouteError::Transport(TransportError::Http { - status: 500, - body: "boom".into(), - }), - RouteError::Transport(TransportError::Network("reset".into())), - ]; - for error in after { - assert_eq!(error.phase(), Phase::AfterSend, "{error:?}"); - } - let before = [ - RouteError::Transport(TransportError::Connect("refused".into())), - RouteError::Unsupported("streaming"), - RouteError::Auth(litellm_auth::Error::InvalidHeader), - ]; - for error in before { - assert_eq!(error.phase(), Phase::BeforeSend, "{error:?}"); - } - } - #[test] fn a_missing_api_key_is_the_environment_not_the_request() { assert!( @@ -185,12 +133,9 @@ mod tests { assert!(!RouteError::InvalidResponse("bad json".into()).is_request()); } #[rstest] - #[case::request(true, Phase::BeforeSend)] - #[case::response(false, Phase::AfterSend)] - fn contextual_errors_preserve_sources_and_route_classification( - #[case] request: bool, - #[case] phase: Phase, - ) { + #[case::request(true)] + #[case::response(false)] + fn contextual_errors_preserve_sources_and_route_classification(#[case] request: bool) { let source = serde_json::from_str::("{").unwrap_err(); let source_message = source.to_string(); let detail = ErrorDetail::invalid("test payload", source); @@ -199,7 +144,6 @@ mod tests { } else { LlmError::InvalidResponse(detail) }); - assert_eq!(error.phase(), phase); assert_eq!(error.is_request(), request); let category = if request { "request" } else { "response" }; assert_eq!( diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index a0ae582210e..fe487f41544 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -11,4 +11,4 @@ mod provider; pub mod resources; pub mod responses; -pub use error::{Phase, RouteError}; +pub use error::RouteError; diff --git a/litellm-rust/crates/core/tests/chat_completions.rs b/litellm-rust/crates/core/tests/chat_completions.rs index 28e0a65589c..681670b56f6 100644 --- a/litellm-rust/crates/core/tests/chat_completions.rs +++ b/litellm-rust/crates/core/tests/chat_completions.rs @@ -1,8 +1,6 @@ use std::time::Duration; -use litellm_core::chat_completions::{ - Error, chat_completions_decline_reason, types::ChatCompletionsRequest, -}; +use litellm_core::chat_completions::{Error, types::ChatCompletionsRequest}; use litellm_http::transport::Error as TransportError; use litellm_types::utils::ChatCompletionsResponse; use rstest::{fixture, rstest}; @@ -156,8 +154,6 @@ async fn bedrock_round_trip_is_signed_and_normalized(request: ChatCompletionsReq assert_eq!(response.usage.total_tokens, 15); } -/// The provider already answered and billed these, so the host must not retry them on -/// its own path: they surface as `InvalidResponse`, never as a pre-send decline. #[rstest] #[case::missing_usage( r#"{"model":"m","content":[{"type":"text","text":"hi"}],"stop_reason":"end_turn"}"# @@ -209,10 +205,9 @@ async fn an_upstream_error_status_keeps_its_code_and_body( ); } -/// Nothing was sent, so nothing was billed and the host can still serve the request. #[rstest] #[tokio::test] -async fn a_connection_that_is_never_established_declines_instead_of_failing( +async fn a_connection_that_is_never_established_returns_a_connect_error( request: ChatCompletionsRequest<'static>, ) { let error = complete(ChatCompletionsRequest { @@ -230,9 +225,7 @@ async fn a_connection_that_is_never_established_declines_instead_of_failing( #[rstest] #[tokio::test] -async fn a_timeout_after_sending_is_not_a_pre_send_decline( - request: ChatCompletionsRequest<'static>, -) { +async fn a_timeout_after_sending_returns_a_network_error(request: ChatCompletionsRequest<'static>) { let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE).set_delay(Duration::from_secs(5))]).await; let base = upstream.uri(); @@ -251,79 +244,6 @@ async fn a_timeout_after_sending_is_not_a_pre_send_decline( ); } -#[rstest] -#[case::accepted("anthropic/claude-sonnet-4-5", None, hi(), json!({"max_tokens": 16}), None)] -#[case::accepted_bedrock("bedrock/anthropic.claude-sonnet-4-5", None, hi(), json!({}), None)] -#[case::unknown_provider( - "gpt-4o", - Some("openai"), - hi(), - json!({}), - Some("provider is not on the rust chat completions path") -)] -#[case::unreadable_messages( - "anthropic/claude-sonnet-4-5", - None, - json!("hi"), - json!({}), - Some("unreadable message list") -)] -#[case::empty_messages("anthropic/claude-sonnet-4-5", None, json!([]), json!({}), Some("empty message list"))] -#[case::streaming( - "anthropic/claude-sonnet-4-5", - None, - hi(), - json!({"stream": true}), - Some("streaming") -)] -#[case::unrecognized_param( - "anthropic/claude-sonnet-4-5", - None, - hi(), - json!({"not_a_param": 1}), - Some("unrecognized request parameter") -)] -#[case::opens_on_assistant_turn( - "anthropic/claude-sonnet-4-5", - None, - json!([{"role": "assistant", "content": "hi"}]), - json!({}), - Some("conversation does not open on a user turn") -)] -fn decline_reason_names_why_the_core_would_not_serve_the_request( - #[case] model: &str, - #[case] provider: Option<&str>, - #[case] messages: Value, - #[case] params: Value, - #[case] reason: Option<&str>, -) { - assert_eq!( - chat_completions_decline_reason(model, provider, messages, &object(params)), - reason - ); -} - -/// A request the decline check accepts must not be declined by the call itself. -#[rstest] -#[tokio::test] -async fn a_declined_request_fails_the_call_before_sending( - request: ChatCompletionsRequest<'static>, -) { - let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await; - let base = upstream.uri(); - - let error = complete(ChatCompletionsRequest { - optional_params: object(json!({"stream": true})), - api_base: Some(&base), - ..request - }) - .await - .expect_err("streaming is declined"); - - assert_eq!(error, Error::Unsupported("streaming")); - assert!(received(&upstream).await.is_empty()); -} - #[rstest] #[case::direct(false)] #[case::hosted(true)] @@ -427,7 +347,6 @@ async fn a_post_call_hook_failure_never_looks_safe_to_retry( ) .await .unwrap_err(); - assert_eq!(error.phase(), litellm_core::error::Phase::AfterSend); let Error::PostCallHook(source) = error else { panic!("expected retained callback error") }; diff --git a/litellm-rust/crates/core/tests/messages/response.rs b/litellm-rust/crates/core/tests/messages/response.rs index 87cfb1bb651..c9bcc210126 100644 --- a/litellm-rust/crates/core/tests/messages/response.rs +++ b/litellm-rust/crates/core/tests/messages/response.rs @@ -1,7 +1,4 @@ -use litellm_core::{ - Phase, - messages::{MessagesResponse, messages_body}, -}; +use litellm_core::messages::{MessagesResponse, messages_body}; use litellm_http::transport::Error as TransportError; use rstest::rstest; @@ -206,7 +203,7 @@ async fn an_unreadable_success_body_is_an_invalid_response( .await .expect_err("an unreadable body fails"); - assert_eq!(error.phase(), Phase::AfterSend, "{error:?}"); + assert!(matches!(error, Error::InvalidResponse(_)), "{error:?}"); } #[rstest] diff --git a/litellm-rust/crates/python-bridge/src/errors.rs b/litellm-rust/crates/python-bridge/src/errors.rs index 2e6fbd9d606..eaeef7c50ff 100644 --- a/litellm-rust/crates/python-bridge/src/errors.rs +++ b/litellm-rust/crates/python-bridge/src/errors.rs @@ -1,4 +1,4 @@ -use litellm_core::{Phase, RouteError}; +use litellm_core::RouteError; use litellm_http::transport::Error as TransportError; use pyo3::{ exceptions::{PyRuntimeError, PyValueError}, @@ -20,7 +20,12 @@ pyo3::create_exception!( ); pub(crate) fn route_error_to_pyerr(error: RouteError) -> PyErr { - by_fault(error.is_request(), error.to_string()) + match error { + RouteError::Transport(TransportError::Http { status, body }) => { + RustUpstreamError::new_err((status, body)) + } + other => by_fault(other.is_request(), other.to_string()), + } } /// A request the caller got wrong is a `ValueError`; anything else is a `RuntimeError`. @@ -32,46 +37,39 @@ pub(crate) fn by_fault(is_request: bool, message: String) -> PyErr { } } -/// Map a route error for a route whose host keeps a Python implementation. -/// -/// The distinction the host needs is whether the provider was already called. -/// Everything raised before the request goes out is safe for the host to retry -/// on its own path; anything after it is not, because the provider has already -/// done the work and billed for it. -pub(crate) fn chat_completions_error_to_pyerr(error: RouteError) -> PyErr { - match error.phase() { - Phase::BeforeSend => RustBridgeDeclined::new_err(error.to_string()), - Phase::AfterSend => RustUpstreamError::new_err(match error { - RouteError::Transport(TransportError::Http { status, body }) => (status, body), - RouteError::Transport(TransportError::Network(message)) => (0u16, message), - RouteError::InvalidResponse(detail) => (0u16, detail.to_string()), - other => (0u16, other.to_string()), - }), - } -} - #[cfg(test)] mod tests { use super::*; - #[test] - fn transport_status_and_dispatch_certainty_survive_python_mapping() { + #[rstest::rstest] + #[case::unsupported(RouteError::Unsupported("test capability"), true)] + #[case::invalid_provider(RouteError::InvalidProvider("unknown".into()), true)] + #[case::invalid_request(RouteError::InvalidRequest("empty messages".into()), true)] + #[case::connection(TransportError::Connect("unreachable".into()).into(), false)] + #[case::network(TransportError::Network("timed out".into()).into(), false)] + #[case::invalid_response(RouteError::InvalidResponse("missing usage".into()), false)] + fn route_failures_are_terminal(#[case] error: RouteError, #[case] is_request: bool) { Python::initialize(); Python::attach(|py| { - let connect = chat_completions_error_to_pyerr( - TransportError::Connect("unreachable".into()).into(), - ); - assert!(connect.is_instance_of::(py)); - let network = - chat_completions_error_to_pyerr(TransportError::Network("timed out".into()).into()); - assert!(network.is_instance_of::(py)); - let upstream = chat_completions_error_to_pyerr( + let failure = route_error_to_pyerr(error); + assert!(!failure.is_instance_of::(py)); + assert_eq!(failure.is_instance_of::(py), is_request); + assert_eq!(failure.is_instance_of::(py), !is_request); + }); + } + + #[rstest::rstest] + fn transport_status_survives_python_mapping() { + Python::initialize(); + Python::attach(|py| { + let upstream = route_error_to_pyerr( TransportError::Http { status: 429, body: "slow down".into(), } .into(), ); + assert!(upstream.is_instance_of::(py)); assert_eq!( upstream .value(py) diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 51e112fa1be..e74e3556dd7 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -28,7 +28,7 @@ mod _native { use crate::routes::audio_transcription::{atranscription, transcription}; #[pymodule_export] use crate::routes::chat_completions::{ - achat_completions, acompletion, chat_completions, chat_completions_decline, completion, + achat_completions, acompletion, chat_completions, completion, }; #[pymodule_export] use crate::routes::embeddings::{aembedding, embedding}; @@ -94,7 +94,6 @@ mod tests { "atranscription", "messages", "amessages", - "chat_completions_decline", "chat_completions", "achat_completions", "completion", diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs index 960eb1f4697..0f9db3b6305 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs @@ -2,18 +2,14 @@ mod host; use pyo3::types::{PyDict, PyTuple}; -use crate::errors::RustBridgeDeclined; use crate::logger::{run_async, run_sync}; -use litellm_core::chat_completions::{ - ChatCompletionsRoute, Error, chat_completions_decline_reason, types::ChatCompletionsRequest, -}; -use litellm_host_python::from_py_argument; +use litellm_core::chat_completions::{ChatCompletionsRoute, Error, types::ChatCompletionsRequest}; use litellm_types::utils::ChatCompletionsResponse; use pyo3::prelude::*; use serde_json::{Map, Value}; use crate::{ - errors::chat_completions_error_to_pyerr, + errors::route_error_to_pyerr, marshal::{ RouteOptions, extra_headers_argument, messages_argument, optional_params_argument, optional_timeout, @@ -52,23 +48,6 @@ async fn execute( .await } -#[pyfunction] -#[pyo3(signature = (model, messages, optional_params=None, custom_llm_provider=None))] -pub(crate) fn chat_completions_decline( - model: String, - #[pyo3(from_py_with = from_py_argument)] messages: Value, - #[pyo3(from_py_with = optional_params_argument)] optional_params: Option>, - custom_llm_provider: Option, -) -> Option { - chat_completions_decline_reason( - &model, - custom_llm_provider.as_deref(), - messages, - &optional_params.unwrap_or_default(), - ) - .map(str::to_string) -} - #[pyfunction] #[pyo3(signature = (model, messages, optional_params=None, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None))] #[expect( @@ -105,7 +84,7 @@ pub(crate) fn chat_completions( optional_params.unwrap_or_default(), options, ), - chat_completions_error_to_pyerr, + route_error_to_pyerr, ) } @@ -145,7 +124,7 @@ pub(crate) fn achat_completions<'py>( optional_params.unwrap_or_default(), options, ), - chat_completions_error_to_pyerr, + route_error_to_pyerr, ) } @@ -162,32 +141,6 @@ fn run_public( request.clone().unbind(), "litellm.rust_bridge.chat_completions.route_host", ); - if let Some(reason) = py - .import("litellm.rust_bridge.chat_completions.route_host")? - .getattr("decline_reason")? - .call1((&request,))? - .extract::>()? - { - return Err(RustBridgeDeclined::new_err(reason)); - } - let admission = host::project(&host, py, &kwargs)?; - if let Some(reason) = chat_completions_decline_reason( - &admission.model, - admission.custom_llm_provider.as_deref(), - admission.messages, - &admission.optional_params, - ) { - return Err(RustBridgeDeclined::new_err(reason)); - } - if admission - .optional_params - .get("stream") - .is_some_and(|value| value == &serde_json::Value::Bool(true)) - { - return Err(RustBridgeDeclined::new_err( - "native Python chat_completions streaming", - )); - } let route = ChatCompletionsRoute::new( crate::http::provider_client(py, &kwargs, asynchronous)? .map_err(crate::http::client_error)?, @@ -232,46 +185,3 @@ pub(crate) fn acompletion( ) -> PyResult> { run_public(py, request, args, kwargs, true) } - -#[cfg(test)] -mod tests { - use pyo3::{prelude::*, types::PyList}; - - #[test] - fn chat_completions_decline_keeps_existing_reasons() { - Python::initialize(); - Python::attach(|py| { - let decline = crate::native_module(py) - .getattr("chat_completions_decline") - .expect("decline helper should be registered"); - let empty = PyList::empty(py); - let unreadable = py - .eval(c"'nope'", None, None) - .expect("string messages should convert"); - - let unknown: Option = decline - .call1(("unknown-model", &empty)) - .and_then(|value| value.extract()) - .expect("unknown providers should decline"); - assert_eq!( - unknown.as_deref(), - Some("provider is not on the rust chat completions path") - ); - - let empty_reason: Option = decline - .call1(("anthropic/claude-sonnet-4-5", &empty)) - .and_then(|value| value.extract()) - .expect("empty lists should decline"); - assert_eq!(empty_reason.as_deref(), Some("empty message list")); - - let unreadable_reason: Option = decline - .call1(("anthropic/claude-sonnet-4-5", unreadable)) - .and_then(|value| value.extract()) - .expect("non-list messages should decline"); - assert_eq!( - unreadable_reason.as_deref(), - Some("unreadable message list") - ); - }); - } -} diff --git a/litellm/chat_completions/dispatch.py b/litellm/chat_completions/dispatch.py index e0fce35bb7a..b2b9662dd78 100644 --- a/litellm/chat_completions/dispatch.py +++ b/litellm/chat_completions/dispatch.py @@ -4,7 +4,7 @@ from types import MappingProxyType from typing import Final, TypeAlias, cast # noqa: TID251 # native binding selects a sync result or an async awaitable from litellm import main -from litellm.rust_bridge.catalog import Delivery, Route, RouteContext +from litellm.rust_bridge.catalog import Route, RouteContext from litellm.rust_bridge.chat_completions.entrypoints import ( NATIVE_ACOMPLETION, NATIVE_COMPLETION, @@ -78,7 +78,6 @@ def _context(request: LiteLLMChatCompletionsRequest) -> RouteContext: Route.CHAT_COMPLETIONS, provider=request.custom_llm_provider, model=request.model, - delivery=Delivery.STREAMING if request.stream else Delivery.COMPLETED, ) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index ce9f7a2ea54..668e8f1c573 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -188,10 +188,10 @@ from litellm.utils import ( def _rust_responses_websocket_enabled( custom_llm_provider: str | None, ) -> bool: - from litellm.rust_bridge.catalog import Delivery, Route, RouteContext, decision + from litellm.rust_bridge.catalog import Route, RouteContext, decision from litellm.rust_bridge.configuration import Decision - context: Final = RouteContext(Route.RESPONSES, provider=custom_llm_provider, delivery=Delivery.WEBSOCKET) + context: Final = RouteContext(Route.RESPONSES, provider=custom_llm_provider) return decision(context) is not Decision.PYTHON diff --git a/litellm/messages/dispatch.py b/litellm/messages/dispatch.py index 28339ac5c94..54aad495d1b 100644 --- a/litellm/messages/dispatch.py +++ b/litellm/messages/dispatch.py @@ -6,7 +6,7 @@ from typing import Final, TypeAlias, cast # noqa: TID251 # native binding sele from litellm.exceptions import BadRequestError from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider from litellm.llms.anthropic.pass_through.messages import handler as main -from litellm.rust_bridge.catalog import Delivery, Route, RouteContext +from litellm.rust_bridge.catalog import Route, RouteContext from litellm.rust_bridge.dispatch import PublicDispatch, call_hook from litellm.rust_bridge.messages.entrypoints import ( NATIVE_AMESSAGES, @@ -85,7 +85,6 @@ def _context(request: LiteLLMMessagesRequest) -> RouteContext: Route.MESSAGES, provider=_resolved_provider(request), model=request.model, - delivery=Delivery.STREAMING if request.stream else Delivery.COMPLETED, ) diff --git a/litellm/responses/dispatch.py b/litellm/responses/dispatch.py index 52f9219b3fd..6fe9451cc9a 100644 --- a/litellm/responses/dispatch.py +++ b/litellm/responses/dispatch.py @@ -5,7 +5,7 @@ from typing import Final, TypeAlias, cast # noqa: TID251 # native binding sele from litellm.responses import main from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator -from litellm.rust_bridge.catalog import Delivery, Route, RouteContext +from litellm.rust_bridge.catalog import Route, RouteContext from litellm.rust_bridge.dispatch import PublicDispatch, call_hook from litellm.rust_bridge.public_call import bind, optional_bool, optional_mapping, optional_str, signature from litellm.rust_bridge.responses.entrypoints import ( @@ -70,7 +70,6 @@ def _context(request: LiteLLMResponsesRequest) -> RouteContext: Route.RESPONSES, provider=request.custom_llm_provider, model=request.model, - delivery=Delivery.STREAMING if request.stream else Delivery.COMPLETED, ) diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index 715bc15ca8d..4fe95d040a5 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -108,12 +108,6 @@ def amessages( args: tuple[object, ...], kwargs: dict[str, object], ) -> Coroutine[object, object, AnthropicMessagesResponse | AsyncIterator[bytes]]: ... -def chat_completions_decline( - model: str, - messages: Sequence[object], - optional_params: Mapping[str, object] | None = None, - custom_llm_provider: str | None = None, -) -> str | None: ... def chat_completions( model: str, messages: Sequence[object], @@ -413,7 +407,6 @@ __all__ = [ "aresponses", "atranscription", "chat_completions", - "chat_completions_decline", "completion", "embedding", "gil_stats", diff --git a/litellm/rust_bridge/catalog.py b/litellm/rust_bridge/catalog.py index 46acb94c958..32926a8464b 100644 --- a/litellm/rust_bridge/catalog.py +++ b/litellm/rust_bridge/catalog.py @@ -7,7 +7,7 @@ admission separately decides whether the selected implementation can execute. from __future__ import annotations from dataclasses import dataclass -from enum import Enum, auto +from enum import Enum from typing import Final, TypeAlias from litellm.rust_bridge.configuration import Decision, Rollout @@ -27,18 +27,11 @@ class Route(str, Enum): TOKENIZER = "tokenizer" -class Delivery(Enum): - COMPLETED = auto() - STREAMING = auto() - WEBSOCKET = auto() - - @dataclass(frozen=True, slots=True) class RouteContext: route: Route provider: str | None = None model: str | None = None - delivery: Delivery = Delivery.COMPLETED @dataclass(frozen=True, slots=True) @@ -47,7 +40,6 @@ class RouteRule: rollout: Rollout providers: frozenset[str] | None = None models: frozenset[str] | None = None - deliveries: frozenset[Delivery] | None = None def matches(self, context: Context) -> bool: return ( @@ -55,7 +47,6 @@ class RouteRule: and context.route is self.route and (self.providers is None or context.provider in self.providers) and (self.models is None or context.model in self.models) - and (self.deliveries is None or context.delivery in self.deliveries) ) diff --git a/litellm/rust_bridge/chat_completions/route_host.py b/litellm/rust_bridge/chat_completions/route_host.py index 9902e8e9946..d6b222dcf41 100644 --- a/litellm/rust_bridge/chat_completions/route_host.py +++ b/litellm/rust_bridge/chat_completions/route_host.py @@ -7,7 +7,6 @@ import litellm from litellm.constants import OPENAI_CHAT_COMPLETION_PARAMS from litellm.rust_bridge import failures from litellm.rust_bridge.chat_completions.entrypoints import LiteLLMChatCompletionsRequest -from litellm.rust_bridge.public_call import inference_decline_reason from litellm.types.utils import ModelResponse _TRANSPORT_PARAMETERS: Final = frozenset( @@ -45,7 +44,3 @@ def arguments(request: LiteLLMChatCompletionsRequest) -> Mapping[str, object]: def map_failure(error: Exception, request: LiteLLMChatCompletionsRequest) -> Exception: provider: Final = request.custom_llm_provider or request.model.partition("/")[0] return failures.map_native_failure(error, request.model, provider, arguments(request), request.api_base) - - -def decline_reason(request: LiteLLMChatCompletionsRequest) -> str | None: - return inference_decline_reason(PARAMETERS, {**request.parameters, **request.kwargs}) diff --git a/litellm/rust_bridge/dispatch.py b/litellm/rust_bridge/dispatch.py index d0a39b18c53..c85339eba8f 100644 --- a/litellm/rust_bridge/dispatch.py +++ b/litellm/rust_bridge/dispatch.py @@ -37,7 +37,7 @@ class PublicDispatch(Generic[RequestT]): for rule in rules: if not isinstance(rule, RouteRule) or rule.route is not self.route: continue - if rule.providers is not None or rule.models is not None or rule.deliveries is not None: + if rule.providers is not None or rule.models is not None: if rollout_decision(rule.rollout) is not Decision.PYTHON: return True continue diff --git a/tests/test_litellm_rust/test_inference.py b/tests/test_litellm_rust/test_inference.py index 863741b01a6..8f071c9aeda 100644 --- a/tests/test_litellm_rust/test_inference.py +++ b/tests/test_litellm_rust/test_inference.py @@ -184,19 +184,30 @@ async def test_unstarted_native_inference_has_no_provider_or_callback_effects( {"model_list": []}, ), ) -async def test_native_inference_declines_unsupported_requests_before_callbacks( - route: Route, +async def test_native_responses_declines_unsupported_requests_before_callbacks( recording_server: RecordingServer, options: Mapping[str, object], ) -> None: recorder: Final = RecordingLogger() recording_server.expected_requests = 0 with pytest.raises(_native.RustBridgeDeclined): - native_call(route, True, recording_server, {**options, "callbacks": [recorder]}) + native_call("responses", True, recording_server, {**options, "callbacks": [recorder]}) assert not recording_server.requests assert not recorder.events +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_native_chat_validation_failure_is_terminal( + asynchronous: bool, recording_server: RecordingServer +) -> None: + recording_server.expected_requests = 0 + with pytest.raises(Exception, match="chat completions requires at least one message") as failure: + await execute("chat", asynchronous, recording_server, {"messages": []}) + assert not isinstance(failure.value, _native.RustBridgeDeclined) + assert not recording_server.requests + + @pytest.mark.asyncio async def test_native_projection_reads_positional_parameters(route: Route, recording_server: RecordingServer) -> None: from litellm.chat_completions.dispatch import ( diff --git a/tests/unit/rust_bridge/test_catalog.py b/tests/unit/rust_bridge/test_catalog.py index 2b3edac612e..d3cab8679ef 100644 --- a/tests/unit/rust_bridge/test_catalog.py +++ b/tests/unit/rust_bridge/test_catalog.py @@ -10,7 +10,6 @@ from litellm.rust_bridge.catalog import ( CacheContext, CacheRule, Context, - Delivery, LoggerContext, Route, RouteContext, @@ -34,21 +33,19 @@ def isolated_configuration(monkeypatch: pytest.MonkeyPatch) -> Generator[None]: @pytest.mark.parametrize("route", tuple(Route)) @pytest.mark.parametrize("provider", (None, "bedrock", "mistral", "anthropic", "openai", "azure_ai", "unknown")) -@pytest.mark.parametrize("delivery", tuple(Delivery)) @pytest.mark.parametrize("process", (None, False, True)) @pytest.mark.parametrize("environment", (None, "0", "1")) def test_shipped_decisions( monkeypatch: pytest.MonkeyPatch, route: Route, provider: str | None, - delivery: Delivery, process: bool | None, environment: str | None, ) -> None: configuration.rust(process) if environment is not None: monkeypatch.setenv("LITELLM_RUST", environment) - context: Final = RouteContext(route, provider=provider, model="test-model", delivery=delivery) + context: Final = RouteContext(route, provider=provider, model="test-model") if route is Route.OCR or (route is Route.TRANSCRIPTION and provider == "bedrock"): assert catalog.rollout(context) is Rollout.RUST_REQUIRED @@ -111,14 +108,12 @@ def test_response_cache_rules_select_the_whole_backend_runtime() -> None: ("context", "expected"), ( ( - RouteContext(Route.RESPONSES, provider="openai", model="m", delivery=Delivery.WEBSOCKET), + RouteContext(Route.RESPONSES, provider="openai", model="m"), Decision.RUST_REQUIRED, ), - (RouteContext(Route.RESPONSES, provider="openai", model="m"), Decision.PYTHON), - (RouteContext(Route.RESPONSES, provider="openai", model="m", delivery=Delivery.STREAMING), Decision.PYTHON), - (RouteContext(Route.RESPONSES, provider="openai", model="other", delivery=Delivery.WEBSOCKET), Decision.PYTHON), - (RouteContext(Route.RESPONSES, provider="anthropic", model="m", delivery=Delivery.WEBSOCKET), Decision.PYTHON), - (RouteContext(Route.MESSAGES, provider="openai", model="m", delivery=Delivery.WEBSOCKET), Decision.PYTHON), + (RouteContext(Route.RESPONSES, provider="openai", model="other"), Decision.PYTHON), + (RouteContext(Route.RESPONSES, provider="anthropic", model="m"), Decision.PYTHON), + (RouteContext(Route.MESSAGES, provider="openai", model="m"), Decision.PYTHON), ), ) def test_first_matching_rule_respects_every_constraint(context: RouteContext, expected: Decision) -> None: @@ -128,7 +123,6 @@ def test_first_matching_rule_respects_every_constraint(context: RouteContext, ex Rollout.RUST_REQUIRED, providers=frozenset({"openai"}), models=frozenset({"m"}), - deliveries=frozenset({Delivery.WEBSOCKET}), ), RouteRule(Route.RESPONSES, Rollout.PYTHON_ONLY), ) diff --git a/tests/unit/rust_bridge/test_dispatch.py b/tests/unit/rust_bridge/test_dispatch.py index 0dc06b2905f..9c3e35edda4 100644 --- a/tests/unit/rust_bridge/test_dispatch.py +++ b/tests/unit/rust_bridge/test_dispatch.py @@ -6,7 +6,7 @@ import pytest from litellm.rust_bridge import configuration from litellm.rust_bridge.bindings import NativeBinding -from litellm.rust_bridge.catalog import CacheRule, Delivery, Route, RouteContext, RouteRule, Rules, SecretManagerRule +from litellm.rust_bridge.catalog import CacheRule, Route, RouteContext, RouteRule, Rules, SecretManagerRule from litellm.rust_bridge.configuration import Rollout from litellm.rust_bridge.dispatch import PublicDispatch from litellm.rust_bridge.runtime import NO_PYTHON, NoPythonImplementationError @@ -99,12 +99,12 @@ def test_native_stream_result_is_not_consumed_or_wrapped() -> None: rules: Final[Rules] = ( CacheRule(Rollout.PYTHON_ONLY), SecretManagerRule(Rollout.PYTHON_ONLY), - RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED, deliveries=frozenset({Delivery.STREAMING})), + RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED), ) dispatch: Final = PublicDispatch( route=Route.CHAT_COMPLETIONS, request=lambda args, kwargs: request, - context=lambda value: RouteContext(Route.CHAT_COMPLETIONS, model=value.model, delivery=Delivery.STREAMING), + context=lambda value: RouteContext(Route.CHAT_COMPLETIONS, model=value.model), ) def native(request: Request, args: tuple[object, ...], kwargs: Mapping[str, object]) -> Iterator[int]: @@ -157,13 +157,11 @@ async def test_async_route_without_rules_preserves_async_iterator_result(rules: @pytest.mark.asyncio async def test_async_dispatch_accepts_websocket_style_none_result() -> None: request: Final = Request(model="realtime-model") - rules: Final[Rules] = ( - RouteRule(Route.RESPONSES, Rollout.RUST_REQUIRED, deliveries=frozenset({Delivery.WEBSOCKET})), - ) + rules: Final[Rules] = (RouteRule(Route.RESPONSES, Rollout.RUST_REQUIRED),) dispatch: Final = PublicDispatch( route=Route.RESPONSES, request=lambda args, kwargs: request, - context=lambda value: RouteContext(Route.RESPONSES, model=value.model, delivery=Delivery.WEBSOCKET), + context=lambda value: RouteContext(Route.RESPONSES, model=value.model), ) async def python(*args: object, **kwargs: object) -> None: # kwargs-ok: public pass-through shape diff --git a/tests/unit/rust_bridge/test_runtime.py b/tests/unit/rust_bridge/test_runtime.py index 33c8e6112a6..2f166631667 100644 --- a/tests/unit/rust_bridge/test_runtime.py +++ b/tests/unit/rust_bridge/test_runtime.py @@ -10,7 +10,7 @@ from litellm.exceptions import APIError from litellm.llms.base_llm.ocr.transformation import OCRResponse from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict from litellm.rust_bridge import bindings, configuration, runtime -from litellm.rust_bridge.catalog import Delivery, Route, RouteContext, RouteRule +from litellm.rust_bridge.catalog import Route, RouteContext, RouteRule from litellm.rust_bridge.configuration import Rollout from litellm.rust_bridge.lifecycle import Complete, Open, Stream, SyncStream, Yield @@ -162,14 +162,10 @@ def test_context_outside_rule_stays_on_python() -> None: RouteContext(Route.TRANSCRIPTION, provider="openai"), ), ) -@pytest.mark.parametrize("delivery", tuple(Delivery)) -async def test_shipped_python_routes_never_load_native( - monkeypatch: pytest.MonkeyPatch, context: RouteContext, delivery: Delivery -) -> None: +async def test_shipped_python_routes_never_load_native(monkeypatch: pytest.MonkeyPatch, context: RouteContext) -> None: monkeypatch.setenv("LITELLM_RUST", "1") configuration.rust(True) calls: Final = recorder() - request: Final = RouteContext(context.route, provider=context.provider, delivery=delivery) def reject_load(value: object) -> NativeFn | None: pytest.fail("Python-only dispatch must not load a native binding") @@ -182,8 +178,8 @@ async def test_shipped_python_routes_never_load_native( async def python() -> str: return calls.python() - assert runtime.run(request, binding=bound, native=lambda fn: fn(), python=calls.python) == PYTHON - assert await runtime.arun(request, binding=bound, native=native, python=python) == PYTHON + assert runtime.run(context, binding=bound, native=lambda fn: fn(), python=calls.python) == PYTHON + assert await runtime.arun(context, binding=bound, native=native, python=python) == PYTHON assert calls.calls == (PYTHON, PYTHON)