refactor(rust): move completed route admission into core
Some checks failed
ai-gateway image / ai-gateway release image (push) Has been cancelled
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled

This commit is contained in:
Yujong Lee 2026-09-14 20:42:36 -07:00
parent 5bc0591784
commit e0069cf5d6
16 changed files with 191 additions and 63 deletions

View file

@ -36,9 +36,17 @@ pub struct AudioTranscriptionRoute;
pub type AudioTranscriptionCall = CompletedCall<AudioTranscriptionRoute>;
impl CompletedRoute for AudioTranscriptionRoute {
type Admission =
crate::call_lifecycle::admission::Inspection<super::AudioTranscriptionAdmission>;
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<dyn ProviderHooks>,

View file

@ -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<String>,
pub audio: Value,
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
crate::call_lifecycle::provider::run_completed::<lifecycle::AudioTranscriptionRoute>(
@ -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<AudioTranscriptionAdmission>,
) -> 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"));

View file

@ -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() {

View file

@ -30,3 +30,8 @@ pub enum AdmissionDecline {
#[strum(to_string = "{0}")]
Feature(&'static str),
}
pub enum Inspection<T> {
Inspectable(T),
Uninspectable,
}

View file

@ -146,9 +146,13 @@ pub enum CompletedReply<Q> {
}
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<dyn ProviderHooks>)
-> WorkflowFuture<Self::Response>;
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<dyn ProviderHooks>) -> WorkflowFuture<Self::Response> {
Box::pin(async move {
let request = hooks

View file

@ -36,9 +36,16 @@ pub struct ChatCompletionsRoute;
pub type ChatCompletionsCall = CompletedCall<ChatCompletionsRoute>;
impl CompletedRoute for ChatCompletionsRoute {
type Admission = crate::call_lifecycle::admission::Inspection<super::ChatCompletionsAdmission>;
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<dyn ProviderHooks>,

View file

@ -69,28 +69,41 @@ pub struct AdmissionContext {
pub bedrock_metadata_owned: bool,
}
pub struct ChatCompletionsAdmission {
pub model: String,
pub provider: Option<String>,
pub messages: Value,
pub params: Map<String, Value>,
pub headers: Option<Map<String, Value>>,
pub context: AdmissionContext,
}
pub fn admit(
model: &str,
provider: Option<&str>,
messages: Value,
params: &Map<String, Value>,
headers: Option<&Map<String, Value>>,
context: AdmissionContext,
inspection: crate::call_lifecycle::admission::Inspection<ChatCompletionsAdmission>,
) -> 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(()),
}

View file

@ -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<ProviderChatCompletionsRequest, Error> {

View file

@ -34,9 +34,16 @@ pub struct MessagesRoute;
pub type MessagesCall = CompletedCall<MessagesRoute>;
impl CompletedRoute for MessagesRoute {
type Admission = crate::call_lifecycle::admission::Inspection<super::MessagesAdmission>;
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<dyn ProviderHooks>,

View file

@ -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<String>,
pub has_agentic_hook: bool,
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesResponse, Error> {
crate::call_lifecycle::provider::run_completed::<lifecycle::MessagesRoute>(request.into()).await
@ -29,20 +35,27 @@ pub async fn messages_stream(request: MessagesRequest<'_>) -> Result<reqwest::Re
}
pub fn admit(
model: &str,
provider: Option<&str>,
has_agentic_hook: bool,
inspection: crate::call_lifecycle::admission::Inspection<MessagesAdmission>,
) -> 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(())

View file

@ -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();

View file

@ -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<Self::Admission>;
fn project(request: &Bound<'_, PyDict>) -> PyResult<Self::Request>;
}
@ -242,7 +242,7 @@ pub(crate) fn run<R: PythonCompletedRoute>(
where
R::Response: Serialize,
{
R::admit(&request)?;
crate::errors::admit(R::admit(R::project_admission(&request)?))?;
let controls = crate::cache::snapshot(
py,
if asynchronous {

View file

@ -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<Self::Admission> {
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<String> = provider
.as_ref()
.map(|value| value.extract::<Option<String>>())
.transpose()?
.flatten();
crate::errors::admit(litellm_core::chat_completions::admit(
&model.extract::<String>()?,
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<OwnedChatCompletionsRequest> {

View file

@ -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::<crate::errors::RustBridgeDeclined>(py)
);
assert_eq!(
async_transcription_error.to_string(),
sync_transcription_error.to_string()
);
});
}

View file

@ -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<Self::Admission> {
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<String> = provider
.as_ref()
.map(|value| value.extract::<Option<String>>())
.transpose()?
.flatten();
crate::errors::admit(litellm_core::messages::admit(
&model.extract::<String>()?,
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<OwnedMessagesRequest> {

View file

@ -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<Self::Admission> {
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::<String>()?,
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<OwnedAudioTranscriptionRequest> {