mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
refactor(rust): remove delivery routing abstraction (#43514)
* wip * refactor: finish removing delivery routing abstraction * refactor(rust): remove chat completion decline admission --------- Co-authored-by: Yujong Lee <yujong@berri.ai>
This commit is contained in:
parent
f184ace25b
commit
74cad08997
21 changed files with 125 additions and 530 deletions
|
|
@ -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<dyn SecretSource>,
|
||||
}
|
||||
|
||||
/// 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<String, Value>,
|
||||
) -> 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,
|
||||
|
|
|
|||
|
|
@ -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}"));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -41,14 +41,6 @@ pub enum RouteError {
|
|||
PostCallHook(#[source] Arc<RouteError>),
|
||||
}
|
||||
|
||||
/// 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<litellm_host::machine::MachineFault> 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::<serde_json::Value>("{").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!(
|
||||
|
|
|
|||
|
|
@ -11,4 +11,4 @@ mod provider;
|
|||
pub mod resources;
|
||||
pub mod responses;
|
||||
|
||||
pub use error::{Phase, RouteError};
|
||||
pub use error::RouteError;
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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::<RustBridgeDeclined>(py));
|
||||
let network =
|
||||
chat_completions_error_to_pyerr(TransportError::Network("timed out".into()).into());
|
||||
assert!(network.is_instance_of::<RustUpstreamError>(py));
|
||||
let upstream = chat_completions_error_to_pyerr(
|
||||
let failure = route_error_to_pyerr(error);
|
||||
assert!(!failure.is_instance_of::<RustBridgeDeclined>(py));
|
||||
assert_eq!(failure.is_instance_of::<PyValueError>(py), is_request);
|
||||
assert_eq!(failure.is_instance_of::<PyRuntimeError>(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::<RustUpstreamError>(py));
|
||||
assert_eq!(
|
||||
upstream
|
||||
.value(py)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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<Map<String, Value>>,
|
||||
custom_llm_provider: Option<String>,
|
||||
) -> Option<String> {
|
||||
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::<Option<String>>()?
|
||||
{
|
||||
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<Py<PyAny>> {
|
||||
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<String> = 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<String> = 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<String> = 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")
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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})
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue