Merge branch 'litellm_internal_staging' of https://github.com/BerriAI/litellm into litellm_internal_copy_38013

# Conflicts:
#	tests/test_litellm/proxy/test_proxy_server.py
This commit is contained in:
mateo-berri 2026-09-02 18:53:17 -07:00
commit 4a785c8a1a
147 changed files with 12402 additions and 977 deletions

View file

@ -9,6 +9,7 @@ on:
- "rust-toolchain.toml"
- ".github/scripts/smoke_test_native_wheel.py"
- ".github/scripts/verify_linux_native_wheel.py"
- "tests/test_litellm/rust_bridge/native_route_wheel_test.py"
- ".github/workflows/test-rust.yml"
pull_request:
branches:
@ -23,6 +24,7 @@ on:
- "rust-toolchain.toml"
- ".github/scripts/smoke_test_native_wheel.py"
- ".github/scripts/verify_linux_native_wheel.py"
- "tests/test_litellm/rust_bridge/native_route_wheel_test.py"
- ".github/workflows/test-rust.yml"
permissions:
@ -121,3 +123,6 @@ jobs:
env:
RELEASE_WHEEL_COMMIT_SHA: ${{ github.event.pull_request.head.sha || github.sha }}
run: python .github/scripts/verify_linux_native_wheel.py dist/*.whl
- name: Test native route wheel
run: python tests/test_litellm/rust_bridge/native_route_wheel_test.py dist/*.whl

View file

@ -96,6 +96,7 @@ jobs:
- shard: misc
artifact-name: misc
test-path: >-
tests/sdk_function_trace
tests/test_litellm/batches
tests/test_litellm/secret_managers
tests/test_litellm/a2a_protocol

View file

@ -3,7 +3,7 @@
"limit": 14074
},
"reportArgumentType": {
"limit": 2216
"limit": 2215
},
"reportAssignmentType": {
"limit": 319

View file

@ -1392,6 +1392,12 @@ dependencies = [
"wasm-bindgen",
]
[[package]]
name = "lazy_static"
version = "1.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "bbd2bcb4c963f2ddae06a2efc7e9f3591312473c50c6685e1f298068316e66fe"
[[package]]
name = "libc"
version = "0.2.186"
@ -1416,6 +1422,7 @@ dependencies = [
"tokio",
"tokio-tungstenite",
"tower",
"tracing",
]
[[package]]
@ -1435,6 +1442,7 @@ dependencies = [
"sha2 0.10.9",
"thiserror 2.0.19",
"tokio",
"tracing",
]
[[package]]
@ -1442,13 +1450,18 @@ name = "litellm-python-bridge"
version = "0.1.0"
dependencies = [
"criterion",
"futures-util",
"litellm-ai-gateway",
"litellm-core",
"litellm-python-interop",
"pyo3",
"pyo3-async-runtimes",
"serde",
"serde_json",
"tokio",
"tokio-tungstenite",
"tracing",
"tracing-subscriber",
]
[[package]]
@ -1669,9 +1682,9 @@ dependencies = [
[[package]]
name = "pyo3"
version = "0.29.0"
version = "0.29.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "cd274650b21d4bfc26a0a47587962c1edb425f69287324355cd040c3ea66071c"
checksum = "4688ddedf473e32662b9b067670129a8afb8c18e351482c70d62ba4a88171e8b"
dependencies = [
"libc",
"once_cell",
@ -1697,18 +1710,18 @@ dependencies = [
[[package]]
name = "pyo3-build-config"
version = "0.29.0"
version = "0.29.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c5e2a7d2f0d013342f295c048ad19237add5154a55b1c5a254c0ec93d4109078"
checksum = "f41027e41b4bd03f6e60f9f417fe24a6341a6bb744edd62b6f709f2a52ea30e9"
dependencies = [
"target-lexicon",
]
[[package]]
name = "pyo3-ffi"
version = "0.29.0"
version = "0.29.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "ca85c467da1bbc8d866eea5deff9cf29ea5f7785054a17da36e65bda9c05845b"
checksum = "e591a95526fead067432c3b3a33fc74770b87b1e04e73671090d9c2055a2b327"
dependencies = [
"libc",
"pyo3-build-config",
@ -1716,9 +1729,9 @@ dependencies = [
[[package]]
name = "pyo3-macros"
version = "0.29.0"
version = "0.29.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9ac53762fd065daa3194dd09337a38bd793a188100fd1a9304c4ab312d901771"
checksum = "73225868fc1cd84eef2c3c230ddb91273bf1de46aeb8a4248da76d32a0924a1c"
dependencies = [
"proc-macro2",
"pyo3-macros-backend",
@ -1728,9 +1741,9 @@ dependencies = [
[[package]]
name = "pyo3-macros-backend"
version = "0.29.0"
version = "0.29.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "4ca3a1557399783172dc5bf39cfca835157732532cba56b71d2292161e53b362"
checksum = "571575aa3749fa6216757dd47d2a3e7ef360f329a40f0666a9fbd14889024952"
dependencies = [
"heck",
"proc-macro2",
@ -2276,6 +2289,15 @@ dependencies = [
"digest 0.11.3",
]
[[package]]
name = "sharded-slab"
version = "0.1.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f40ca3c46823713e0d4209592e8d6e826aa57e928f09752619fc696c499637f6"
dependencies = [
"lazy_static",
]
[[package]]
name = "shlex"
version = "2.0.1"
@ -2414,6 +2436,15 @@ dependencies = [
"syn 3.0.0",
]
[[package]]
name = "thread_local"
version = "1.1.10"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1ad99c4c6d32803332c548b1af0540b357b3f5fc0be8f6c6bfe8b2e6ae784070"
dependencies = [
"cfg-if",
]
[[package]]
name = "time"
version = "0.3.53"
@ -2662,6 +2693,17 @@ dependencies = [
"once_cell",
]
[[package]]
name = "tracing-subscriber"
version = "0.3.23"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "cb7f578e5945fb242538965c2d0b04418d38ec25c79d160cd279bf0731c8d319"
dependencies = [
"sharded-slab",
"thread_local",
"tracing-core",
]
[[package]]
name = "try-lock"
version = "0.2.5"

View file

@ -14,11 +14,13 @@ license = "MIT"
repository = "https://github.com/BerriAI/litellm"
[workspace.dependencies]
tracing = "0.1"
tracing-subscriber = { version = "0.3", default-features = false, features = ["registry", "std"] }
litellm-core = { path = "crates/core" }
litellm-ai-gateway = { path = "crates/ai-gateway", default-features = false }
litellm-python-interop = { path = "crates/python-interop" }
axum = "0.7"
pyo3 = "0.29.0"
pyo3 = "0.29.2"
pyo3-async-runtimes = { version = "0.29.0", features = ["tokio-runtime"] }
pythonize = "0.29.0"
rand = "0.8"

View file

@ -14,6 +14,7 @@ path = "src/main.rs"
required-features = ["server"]
[dependencies]
tracing.workspace = true
litellm-core = { workspace = true, features = ["bedrock-auth"] }
# reqwest (rustls + json) is used by io/ocr and ships realtime logs to the
# Python proxy callbacks API.

View file

@ -32,6 +32,7 @@ pub(super) fn truncate_error_body(body: &str) -> String {
format!("{truncated}... (truncated)")
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(super) fn ocr_provider_config(
provider: &str,
model: &str,
@ -73,12 +74,6 @@ pub(super) fn string_headers(
.collect()
}
pub(super) fn has_header(headers: &[(String, String)], name: &str) -> bool {
headers
.iter()
.any(|(key, _)| key.eq_ignore_ascii_case(name))
}
fn document_url_field(document: &Value) -> Result<Option<(&str, &str)>, Error> {
let Some(object) = document.as_object() else {
return Ok(None);

View file

@ -1,12 +1,19 @@
use litellm_core::error::Error;
use litellm_core::http_utils::http_request;
use litellm_core::ocr::transformation::OcrResponseHandling;
use serde_json::Value;
use super::common_utils::{poll_document_intelligence, truncate_error_body};
use super::types::ProviderOcrRequest;
use super::hooks::OcrLifecycleHooks;
use super::types::PreparedOcrRequest;
use crate::client::http_client;
pub(crate) async fn execute_ocr_provider_call(request: ProviderOcrRequest) -> Result<Value, Error> {
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(crate) async fn execute_ocr_provider_call(
request: PreparedOcrRequest,
hooks: &OcrLifecycleHooks,
) -> Result<Value, Error> {
let request = hooks.prepare_provider_request(request).await?;
let mut request_builder = http_client().post(&request.url).json(&request.body);
for (key, value) in &request.upstream_headers {
request_builder = request_builder.header(key, value);
@ -15,8 +22,7 @@ pub(crate) async fn execute_ocr_provider_call(request: ProviderOcrRequest) -> Re
request_builder = request_builder.timeout(duration);
}
let response = request_builder
.send()
let response = http_request(request_builder)
.await
.map_err(|err| Error::Network(err.to_string()))?;

View file

@ -1,13 +1,10 @@
use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming};
use litellm_core::error::Error;
use litellm_core::ocr::transformation::OcrAuthStrategy;
use serde_json::{Map, Value, json};
use std::future::Future;
use std::pin::Pin;
use super::common_utils::{
convert_document_url_to_data_uri, has_header, ocr_provider_config, string_headers,
};
use super::common_utils::{convert_document_url_to_data_uri, string_headers};
use super::types::{PreparedOcrRequest, ProviderOcrRequest};
use crate::integrations::custom_guardrail::{
CustomGuardrailRunner, GuardrailContext, GuardrailError, GuardrailRequest,
@ -62,6 +59,10 @@ impl OcrLifecycleHooks {
.await
.map_err(guardrail_error_to_core_error)?;
let (document, optional_params) = parse_ocr_pre_call_guardrail_request(guardrail_request)?;
let optional_params = match &request.config {
Ok(config) => config.map_ocr_params(&optional_params),
Err(_) => optional_params,
};
Ok(PreparedOcrRequest {
document,
optional_params,
@ -69,25 +70,23 @@ impl OcrLifecycleHooks {
})
}
async fn prepare_provider_request(
pub(crate) async fn prepare_provider_request(
&self,
request: PreparedOcrRequest,
) -> Result<ProviderOcrRequest, Error> {
let config = ocr_provider_config(&request.custom_llm_provider, &request.model)
.ok_or_else(|| Error::InvalidProvider(request.custom_llm_provider.clone()))?;
let config = request.config?;
let env_lookup = |key: &str| std::env::var(key).ok();
let headers = string_headers(request.extra_headers)?;
let auth_strategy = config.auth_strategy();
let api_key = (!has_header(&headers, auth_strategy.header_name()))
.then(|| config.resolve_api_key(request.api_key.as_deref(), &env_lookup))
.transpose()?;
let upstream_headers = config.validate_environment(
string_headers(request.extra_headers)?,
request.api_key.as_deref(),
&env_lookup,
)?;
let url = config.complete_url(
request.api_base.as_deref(),
&request.model,
&request.optional_params,
&env_lookup,
)?;
let filtered_params = config.map_ocr_params(&request.optional_params);
let model = request.model.clone();
let custom_llm_provider = request.custom_llm_provider.clone();
let document = if config.requires_data_uri_document() {
@ -96,9 +95,8 @@ impl OcrLifecycleHooks {
request.document
};
let body = config
.transform_ocr_request(&request.model, document, filtered_params)?
.transform_ocr_request(&request.model, document, request.optional_params)?
.data;
let upstream_headers = upstream_headers(&headers, auth_strategy, api_key.as_deref());
let body = self
.run_during_call_guardrails(&model, &custom_llm_provider, &url, body)
.await?;
@ -167,9 +165,9 @@ impl OcrLifecycleHooks {
}
}
impl CallLifecycleHooks<PreparedOcrRequest, ProviderOcrRequest, Value> for OcrLifecycleHooks {
impl CallLifecycleHooks<PreparedOcrRequest, PreparedOcrRequest, Value> for OcrLifecycleHooks {
type PreCallFuture<'a> = OcrFuture<'a, PreparedOcrRequest>;
type DuringCallFuture<'a> = OcrFuture<'a, ProviderOcrRequest>;
type DuringCallFuture<'a> = OcrFuture<'a, PreparedOcrRequest>;
type SuccessFuture<'a> = OcrLogFuture<'a>;
type FailureFuture<'a> = OcrLogFuture<'a>;
@ -186,7 +184,7 @@ impl CallLifecycleHooks<PreparedOcrRequest, ProviderOcrRequest, Value> for OcrLi
_context: &'a CallLifecycleContext,
request: PreparedOcrRequest,
) -> Self::DuringCallFuture<'a> {
Box::pin(async move { self.prepare_provider_request(request).await })
Box::pin(async move { Ok(request) })
}
fn async_log_success_event<'a>(
@ -247,21 +245,6 @@ impl CallLifecycleHooks<PreparedOcrRequest, ProviderOcrRequest, Value> for OcrLi
}
}
fn upstream_headers(
headers: &[(String, String)],
auth_strategy: OcrAuthStrategy,
api_key: Option<&str>,
) -> Vec<(String, String)> {
api_key
.map(|api_key| match auth_strategy {
OcrAuthStrategy::Bearer => ("Authorization".to_string(), format!("Bearer {api_key}")),
OcrAuthStrategy::Header(header_name) => (header_name.to_string(), api_key.to_string()),
})
.into_iter()
.chain(headers.iter().cloned())
.collect()
}
fn guardrail_context(metadata: &RequestMetadata) -> GuardrailContext {
GuardrailContext {
call_type: CallType::Ocr,

View file

@ -13,10 +13,13 @@ pub use types::OcrRequest;
use handler::execute_ocr_provider_call;
use prepare::{PreparedOcrCall, prepare_ocr_call};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn ocr(request: OcrRequest<'_>) -> Result<Value, Error> {
let PreparedOcrCall { request, hooks } = prepare_ocr_call(request);
CallLifecycle::default()
.run_request(request, &hooks, execute_ocr_provider_call)
.run_request(request, &hooks, |request| {
execute_ocr_provider_call(request, &hooks)
})
.await
}

View file

@ -3,6 +3,7 @@ use std::time::{SystemTime, UNIX_EPOCH};
use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use super::common_utils::ocr_provider_config;
use super::hooks::OcrLifecycleHooks;
use super::types::{OcrRequest, PreparedOcrRequest};
use crate::integrations::custom_guardrail::CustomGuardrailRunner;
@ -13,6 +14,7 @@ pub(crate) struct PreparedOcrCall {
pub(crate) hooks: OcrLifecycleHooks,
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(crate) fn prepare_ocr_call(request: OcrRequest<'_>) -> PreparedOcrCall {
let call_id = request
.litellm_call_id
@ -25,9 +27,25 @@ pub(crate) fn prepare_ocr_call(request: OcrRequest<'_>) -> PreparedOcrCall {
});
let model = provider_info.model.to_string();
let custom_llm_provider = provider_info.custom_llm_provider.to_string();
let config = ocr_provider_config(&custom_llm_provider, &model)
.ok_or_else(|| litellm_core::Error::InvalidProvider(custom_llm_provider.clone()));
let optional_params = match &config {
Ok(config) => {
let supported = config.supported_ocr_params();
config.map_ocr_params(
&request
.optional_params
.into_iter()
.filter(|(name, _)| supported.contains(&name.as_str()))
.collect(),
)
}
Err(_) => request.optional_params,
};
PreparedOcrCall {
request: PreparedOcrRequest {
config,
model,
custom_llm_provider,
litellm_call_id: call_id,
@ -35,7 +53,7 @@ pub(crate) fn prepare_ocr_call(request: OcrRequest<'_>) -> PreparedOcrCall {
api_key: request.api_key.map(str::to_string),
api_base: request.api_base.map(str::to_string),
extra_headers: request.extra_headers,
optional_params: request.optional_params,
optional_params,
timeout: request.timeout,
},
hooks: OcrLifecycleHooks::new(

View file

@ -2,12 +2,13 @@ use std::sync::{Arc, Mutex};
use std::time::Duration;
use litellm_core::error::Error;
use litellm_core::http_utils::has_header;
use litellm_core::ocr::transformation::OcrResponseHandling;
use serde_json::{Map, Value, json};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use super::common_utils::{has_header, ocr_provider_config, string_headers, truncate_error_body};
use super::common_utils::{ocr_provider_config, string_headers, truncate_error_body};
use super::{OcrRequest, ocr};
use crate::integrations::custom_guardrail::{
CustomGuardrail, GuardrailContext, GuardrailDecision, GuardrailError, GuardrailEventHook,

View file

@ -25,6 +25,7 @@ pub struct OcrRequest<'a> {
}
pub(crate) struct PreparedOcrRequest {
pub(crate) config: Result<&'static dyn OcrProviderConfig, litellm_core::Error>,
pub(crate) model: String,
pub(crate) custom_llm_provider: String,
pub(crate) litellm_call_id: String,

View file

@ -11,6 +11,7 @@ reqwest.workspace = true
serde.workspace = true
serde_json.workspace = true
thiserror.workspace = true
tracing.workspace = true
sha2.workspace = true
aws-config = { version = "1.9.0", default-features = false, features = ["rustls", "rt-tokio"], optional = true }
aws-credential-types = { version = "1.3.0", features = ["hardcoded-credentials"], optional = true }

View file

@ -1,11 +1,12 @@
use serde_json::Value;
use crate::error::Error;
use crate::http_utils::truncate_error_body;
use crate::http_utils::{http_request, truncate_error_body};
use super::client::http_client;
use super::types::ProviderAudioTranscriptionRequest;
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn execute_audio_transcription_provider_call(
request: ProviderAudioTranscriptionRequest,
) -> Result<Value, Error> {
@ -19,8 +20,7 @@ pub async fn execute_audio_transcription_provider_call(
if let Some(duration) = request.timeout {
request_builder = request_builder.timeout(duration);
}
let response = request_builder
.send()
let response = http_request(request_builder)
.await
.map_err(|error| Error::Network(error.to_string()))?;
let status = response.status();

View file

@ -11,6 +11,7 @@ pub use handler::execute_audio_transcription_provider_call;
pub use prepare::prepare_audio_transcription_provider_call;
pub use types::{AudioTranscriptionRequest, ProviderAudioTranscriptionRequest};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
execute_audio_transcription_provider_call(prepare_audio_transcription_provider_call(request)?)
.await

View file

@ -7,6 +7,7 @@ use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}
use super::transformation::{AudioTranscriptionAuth, AudioTranscriptionProviderConfig};
use super::types::{AudioTranscriptionRequest, ProviderAudioTranscriptionRequest};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn provider_config(provider: &str) -> Option<&'static dyn AudioTranscriptionProviderConfig> {
#[cfg(feature = "bedrock-auth")]
if provider == "bedrock" {
@ -16,6 +17,7 @@ fn provider_config(provider: &str) -> Option<&'static dyn AudioTranscriptionProv
None
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub fn prepare_audio_transcription_provider_call(
request: AudioTranscriptionRequest<'_>,
) -> Result<ProviderAudioTranscriptionRequest, Error> {

View file

@ -15,6 +15,7 @@ pub enum AudioTranscriptionAuth {
pub trait AudioTranscriptionProviderConfig: Sync {
fn supported_transcription_params(&self) -> &'static [&'static str];
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn map_transcription_params(&self, params: &Map<String, Value>) -> Map<String, Value> {
params
.iter()

View file

@ -7,6 +7,7 @@ use super::transformation::ChatCompletionsProviderConfig;
const HEADER_CONTEXT: &str = "chat completions";
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(super) fn chat_completions_provider_config(
provider: &str,
) -> Option<&'static dyn ChatCompletionsProviderConfig> {

View file

@ -1,17 +1,21 @@
use serde_json::Value;
use crate::error::Error;
use crate::http_utils::truncate_error_body;
use crate::http_utils::{http_request, truncate_error_body};
use super::client::http_client;
use super::prepare::prepare_provider_request;
use super::transformation::ChatCompletionsAuth;
use super::types::{
ChatCompletionsResponse, ProviderChatCompletionsRequest, ProviderChatResponseData,
ResolvedChatCompletionsRequest,
};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(super) async fn execute_chat_completions_provider_call(
request: ProviderChatCompletionsRequest,
request: ResolvedChatCompletionsRequest<'_>,
) -> Result<ChatCompletionsResponse, Error> {
let request = prepare_provider_request(request)?;
let body = serde_json::to_vec(&request.body).map_err(|err| {
Error::InvalidRequest(format!(
"failed to serialize chat completions request: {err}"
@ -27,7 +31,7 @@ pub(super) async fn execute_chat_completions_provider_call(
request_builder = request_builder.timeout(duration);
}
let response = request_builder.send().await.map_err(|err| {
let response = http_request(request_builder).await.map_err(|err| {
// Failing to establish the connection means the request never went out,
// so the host can still serve it. Everything else here, a timeout
// above all, may have reached the provider and been answered.

View file

@ -19,13 +19,14 @@ pub mod types;
use serde_json::{Map, Value};
use handler::execute_chat_completions_provider_call;
use prepare::{parse_messages, prepare_chat_completions_call, resolve_provider_config};
use prepare::{parse_messages, resolve_provider_config, resolve_request};
use types::{ChatCompletionsRequest, ChatCompletionsResponse};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn chat_completions(
request: ChatCompletionsRequest<'_>,
) -> Result<ChatCompletionsResponse, Error> {
execute_chat_completions_provider_call(prepare_chat_completions_call(request)?).await
execute_chat_completions_provider_call(resolve_request(request)?).await
}
/// Whether the core would accept this request, without resolving credentials or

View file

@ -6,7 +6,10 @@ use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}
use super::common_utils::{chat_completions_provider_config, string_headers};
use super::transformation::{ChatCompletionsAuth, ChatCompletionsProviderConfig};
use super::types::{ChatCompletionsRequest, ChatMessage, ProviderChatCompletionsRequest};
use super::types::{
ChatCompletionsRequest, ChatMessage, ProviderChatCompletionsRequest,
ResolvedChatCompletionsRequest,
};
pub(super) fn resolve_provider_config<'a>(
model: &'a str,
@ -34,12 +37,10 @@ pub(super) fn parse_messages(messages: Value) -> Result<Vec<ChatMessage>, Error>
.map_err(|err| Error::InvalidRequest(format!("invalid chat completions messages: {err}")))
}
pub(super) fn prepare_chat_completions_call(
pub(super) fn resolve_request(
request: ChatCompletionsRequest<'_>,
) -> Result<ProviderChatCompletionsRequest, Error> {
) -> Result<ResolvedChatCompletionsRequest<'_>, Error> {
let (model, config) = resolve_provider_config(request.model, request.custom_llm_provider)?;
let env_lookup = |key: &str| std::env::var(key).ok();
let messages = parse_messages(request.messages)?;
if messages.is_empty() {
return Err(Error::InvalidRequest(
@ -49,11 +50,29 @@ pub(super) fn prepare_chat_completions_call(
if let Some(reason) = config.unsupported_reason(&messages, &request.optional_params) {
return Err(Error::Unsupported(reason.0));
}
Ok(ResolvedChatCompletionsRequest {
model,
config,
messages,
optional_params: request.optional_params,
api_key: request.api_key,
api_base: request.api_base,
extra_headers: request.extra_headers,
timeout: request.timeout,
})
}
let mut headers = string_headers(request.extra_headers)?;
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn validate_environment(
request: &ResolvedChatCompletionsRequest<'_>,
model: &str,
config: &dyn ChatCompletionsProviderConfig,
) -> Result<(Vec<(String, String)>, ChatCompletionsAuth), Error> {
let env_lookup = |key: &str| std::env::var(key).ok();
let mut headers = string_headers(request.extra_headers.clone())?;
let auth = config.auth(
request.api_key,
&model,
model,
&request.optional_params,
&env_lookup,
)?;
@ -94,7 +113,16 @@ pub(super) fn prepare_chat_completions_call(
headers.push(((*name).to_string(), (*value).to_string()));
}
}
Ok((headers, auth))
}
pub(super) fn prepare_provider_request(
request: ResolvedChatCompletionsRequest<'_>,
) -> Result<ProviderChatCompletionsRequest, Error> {
let (headers, auth) = validate_environment(&request, &request.model, request.config)?;
let model = request.model;
let config = request.config;
let env_lookup = |key: &str| std::env::var(key).ok();
let url = config.complete_url(
request.api_base,
&model,
@ -102,7 +130,7 @@ pub(super) fn prepare_chat_completions_call(
&env_lookup,
)?;
let transformed =
config.transform_request(&model, messages, request.optional_params.clone())?;
config.transform_request(&model, request.messages, request.optional_params.clone())?;
Ok(ProviderChatCompletionsRequest {
model,

View file

@ -2,9 +2,15 @@ use serde_json::{Map, Value, json};
use crate::error::Error;
use super::prepare::prepare_chat_completions_call;
use super::prepare::{prepare_provider_request, resolve_request};
use super::transformation::ChatCompletionsAuth;
use super::types::ChatCompletionsRequest;
use super::types::{ChatCompletionsRequest, ProviderChatCompletionsRequest};
fn prepare_chat_completions_call(
request: ChatCompletionsRequest<'_>,
) -> Result<ProviderChatCompletionsRequest, Error> {
prepare_provider_request(resolve_request(request)?)
}
fn request<'a>(
model: &'a str,

View file

@ -62,9 +62,8 @@ pub trait ChatCompletionsProviderConfig: Sync {
false
}
/// Provider parameter names (post-mapping) the Rust path knows how to place
/// in the upstream body. Anything outside this set declines the request.
fn supported_params(&self) -> &'static [&'static str];
/// Supported OpenAI parameter names paired with their provider names.
fn supported_openai_params(&self) -> &'static [(&'static str, &'static str)];
/// Parameters consumed as call configuration (credentials, endpoints)
/// rather than placed in the body. Accepted, never serialized.
@ -78,7 +77,7 @@ pub trait ChatCompletionsProviderConfig: Sync {
optional_params: &Map<String, Value>,
) -> Option<Unsupported> {
unsupported_param(
self.supported_params(),
self.supported_openai_params(),
self.config_params(),
optional_params,
)
@ -100,7 +99,7 @@ pub trait ChatCompletionsProviderConfig: Sync {
}
pub fn unsupported_param(
supported: &'static [&'static str],
supported: &'static [(&'static str, &'static str)],
config: &'static [&'static str],
optional_params: &Map<String, Value>,
) -> Option<Unsupported> {
@ -115,7 +114,9 @@ pub fn unsupported_param(
.keys()
.any(|key| {
key != STREAM_PARAM
&& !supported.contains(&key.as_str())
&& !supported
.iter()
.any(|(_, provider_name)| *provider_name == key)
&& !config.contains(&key.as_str())
})
.then_some(Unsupported("unrecognized request parameter"))

View file

@ -22,6 +22,17 @@ pub struct ChatCompletionsRequest<'a> {
pub timeout: Option<Duration>,
}
pub(super) struct ResolvedChatCompletionsRequest<'a> {
pub(super) model: String,
pub(super) config: &'static dyn ChatCompletionsProviderConfig,
pub(super) messages: Vec<ChatMessage>,
pub(super) optional_params: Map<String, Value>,
pub(super) api_key: Option<&'a str>,
pub(super) api_base: Option<&'a str>,
pub(super) extra_headers: Option<Map<String, Value>>,
pub(super) timeout: Option<Duration>,
}
pub(super) struct ProviderChatCompletionsRequest {
pub(super) model: String,
pub(super) config: &'static dyn ChatCompletionsProviderConfig,

View file

@ -5,6 +5,13 @@ use serde_json::{Map, Value};
use crate::constants::UPSTREAM_ERROR_BODY_MAX_CHARS;
use crate::error::{Error, json_type_name};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn http_request(
request: reqwest::RequestBuilder,
) -> Result<reqwest::Response, reqwest::Error> {
request.send().await
}
/// Bound an upstream error body before it crosses a host boundary, so provider
/// bodies stay data-minimized.
pub fn truncate_error_body(body: &str) -> String {

View file

@ -10,6 +10,7 @@ pub(super) use crate::http_utils::{has_bearer_auth, has_header, truncate_error_b
const HEADER_CONTEXT: &str = "messages";
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(super) fn messages_provider_config(
provider: &str,
) -> Option<&'static dyn AnthropicMessagesProviderConfig> {

View file

@ -1,13 +1,17 @@
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
use crate::error::Error;
use crate::http_utils::http_request;
use super::client::http_client;
use super::common_utils::truncate_error_body;
use super::types::{AnthropicMessagesResponse, ProviderMessagesRequest};
use super::prepare::prepare_provider_request;
use super::types::{AnthropicMessagesResponse, MessagesRequest};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(super) async fn execute_messages_provider_call(
request: ProviderMessagesRequest,
request: MessagesRequest<'_>,
) -> Result<AnthropicMessagesResponse, Error> {
let request = prepare_provider_request(request)?;
let mut request_builder = http_client().post(&request.url).json(&request.body);
for (key, value) in &request.upstream_headers {
request_builder = request_builder.header(key, value);
@ -16,8 +20,7 @@ pub(super) async fn execute_messages_provider_call(
request_builder = request_builder.timeout(duration);
}
let response = request_builder
.send()
let response = http_request(request_builder)
.await
.map_err(|err| Error::Network(err.to_string()))?;
@ -40,8 +43,9 @@ pub(super) async fn execute_messages_provider_call(
}
pub(super) async fn execute_messages_provider_stream(
request: ProviderMessagesRequest,
request: MessagesRequest<'_>,
) -> Result<reqwest::Response, Error> {
let request = prepare_provider_request(request)?;
if request.provider != ANTHROPIC_MESSAGES_PROVIDER {
return Err(Error::InvalidRequest(
"streaming messages is not supported for this provider".to_string(),
@ -56,8 +60,7 @@ pub(super) async fn execute_messages_provider_stream(
request_builder = request_builder.timeout(duration);
}
let response = request_builder
.send()
let response = http_request(request_builder)
.await
.map_err(|err| Error::Network(err.to_string()))?;
let status = response.status();

View file

@ -16,15 +16,15 @@ pub mod transformation;
pub mod types;
use handler::{execute_messages_provider_call, execute_messages_provider_stream};
use prepare::prepare_messages_call;
use types::{AnthropicMessagesResponse, MessagesRequest};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesResponse, Error> {
execute_messages_provider_call(prepare_messages_call(request)?).await
execute_messages_provider_call(request).await
}
pub async fn messages_stream(request: MessagesRequest<'_>) -> Result<reqwest::Response, Error> {
execute_messages_provider_stream(prepare_messages_call(request)?).await
execute_messages_provider_stream(request).await
}
#[cfg(test)]

View file

@ -2,10 +2,11 @@ use crate::error::Error;
use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use super::common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers};
use super::transformation::MessagesAuthStrategy;
use super::transformation::{AnthropicMessagesProviderConfig, MessagesAuthStrategy};
use super::types::{MessagesRequest, ProviderMessagesRequest};
use serde_json::{Map, Value};
pub(super) fn prepare_messages_call(
pub(super) fn prepare_provider_request(
request: MessagesRequest<'_>,
) -> Result<ProviderMessagesRequest, Error> {
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
@ -29,13 +30,46 @@ pub(super) fn prepare_messages_call(
.ok_or_else(|| Error::InvalidProvider(provider.to_string()))?;
let env_lookup = |key: &str| std::env::var(key).ok();
let mut headers = string_headers(request.extra_headers)?;
let headers =
validate_environment(config, request.extra_headers, request.api_key, &env_lookup)?;
let typed_request = serde_json::from_value(request.body).map_err(|err| {
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
})?;
let transformed = config.transform_request(typed_request)?;
let body = serde_json::to_value(transformed).map_err(|err| {
Error::InvalidRequest(format!(
"failed to serialize Anthropic messages request: {err}"
))
})?;
let url = config.complete_url(request.api_base, &model, &env_lookup)?;
Ok(ProviderMessagesRequest {
provider: provider.to_string(),
model,
config,
url,
body,
upstream_headers: headers,
timeout: request.timeout,
})
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn validate_environment(
config: &dyn AnthropicMessagesProviderConfig,
extra_headers: Option<Map<String, Value>>,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<Vec<(String, String)>, Error> {
let mut headers = string_headers(extra_headers)?;
let auth_strategy = config.auth_strategy();
let already_authorized = has_header(&headers, auth_strategy.header_name())
|| (config.accepts_bearer_auth() && has_bearer_auth(&headers));
if !already_authorized {
let api_key = config.resolve_api_key(request.api_key, &env_lookup)?;
let api_key = config.resolve_api_key(api_key, env_lookup)?;
let auth_header = match auth_strategy {
MessagesAuthStrategy::Bearer => {
("authorization".to_string(), format!("Bearer {api_key}"))
@ -51,24 +85,5 @@ pub(super) fn prepare_messages_call(
}
}
let url = config.complete_url(request.api_base, &model, &env_lookup)?;
let typed_request = serde_json::from_value(request.body).map_err(|err| {
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
})?;
let transformed = config.transform_request(typed_request)?;
let body = serde_json::to_value(transformed).map_err(|err| {
Error::InvalidRequest(format!(
"failed to serialize Anthropic messages request: {err}"
))
})?;
Ok(ProviderMessagesRequest {
provider: provider.to_string(),
model,
config,
url,
body,
upstream_headers: headers,
timeout: request.timeout,
})
Ok(headers)
}

View file

@ -45,6 +45,7 @@ pub trait AnthropicMessagesProviderConfig: Sync {
]
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn transform_request(
&self,
request: AnthropicMessagesRequest,
@ -52,6 +53,7 @@ pub trait AnthropicMessagesProviderConfig: Sync {
Ok(request)
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn transform_response(
&self,
_model: &str,

View file

@ -27,6 +27,7 @@ pub enum OcrResponseHandling {
pub trait OcrProviderConfig: Sync {
fn supported_ocr_params(&self) -> &'static [&'static str];
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn map_ocr_params(&self, non_default_params: &Map<String, Value>) -> Map<String, Value> {
let mut mapped_params = Map::new();
for (param, value) in non_default_params {
@ -64,6 +65,25 @@ pub trait OcrProviderConfig: Sync {
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error>;
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn validate_environment(
&self,
headers: Vec<(String, String)>,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<Vec<(String, String)>, Error> {
let strategy = self.auth_strategy();
if crate::http_utils::has_header(&headers, strategy.header_name()) {
return Ok(headers);
}
let api_key = self.resolve_api_key(api_key, env_lookup)?;
let auth_header = match strategy {
OcrAuthStrategy::Bearer => ("Authorization".to_string(), format!("Bearer {api_key}")),
OcrAuthStrategy::Header(name) => (name.to_string(), api_key),
};
Ok(std::iter::once(auth_header).chain(headers).collect())
}
fn auth_strategy(&self) -> OcrAuthStrategy {
OcrAuthStrategy::Bearer
}

View file

@ -27,7 +27,12 @@ use crate::chat_completions::response_utils::{finish_reason_for, unix_now, usage
/// per-model gate inside `transform_request`, the function this route replaces.
/// Forwarding it would send `top_k` to a model that removed sampling params and
/// take a 400 after the call, where Python drops it and succeeds.
const SUPPORTED_PARAMS: &[&str] = &["max_tokens", "temperature", "top_p", "stop_sequences"];
const SUPPORTED_PARAMS: &[(&str, &str)] = &[
("max_tokens", "max_tokens"),
("temperature", "temperature"),
("top_p", "top_p"),
("stop", "stop_sequences"),
];
pub struct AnthropicChatCompletionsConfig;
@ -112,7 +117,8 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig {
})
}
fn supported_params(&self) -> &'static [&'static str] {
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn supported_openai_params(&self) -> &'static [(&'static str, &'static str)] {
SUPPORTED_PARAMS
}
@ -121,7 +127,7 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig {
messages: &[ChatMessage],
optional_params: &Map<String, Value>,
) -> Option<Unsupported> {
unsupported_param(SUPPORTED_PARAMS, &[], optional_params)
unsupported_param(self.supported_openai_params(), &[], optional_params)
.or_else(|| messages.iter().find_map(unsupported_message))
// Anthropic rejects a request whose first turn is not a user turn.
// Python only repairs that under `litellm.modify_params`, which the
@ -132,6 +138,7 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig {
})
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn transform_request(
&self,
model: &str,
@ -143,6 +150,7 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig {
})
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn transform_response(
&self,
_model: &str,

View file

@ -47,6 +47,7 @@ pub fn complete_anthropic_url(
}
impl AnthropicMessagesProviderConfig for AnthropicMessagesConfig {
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn complete_url(
&self,
api_base: Option<&str>,

View file

@ -142,6 +142,7 @@ fn fold_system_role_messages(request: AnthropicMessagesRequest) -> AnthropicMess
}
impl AnthropicMessagesProviderConfig for AzureAnthropicMessagesConfig {
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn complete_url(
&self,
api_base: Option<&str>,

View file

@ -46,10 +46,12 @@ fn optional_string<'a>(params: &'a Map<String, Value>, key: &str) -> Option<&'a
}
impl AudioTranscriptionProviderConfig for BedrockAudioTranscriptionConfig {
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn supported_transcription_params(&self) -> &'static [&'static str] {
SUPPORTED_PARAMS
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn transform_transcription_request(
&self,
_model: &str,
@ -83,6 +85,7 @@ impl AudioTranscriptionProviderConfig for BedrockAudioTranscriptionConfig {
})
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn transform_transcription_response(
&self,
_model: &str,

View file

@ -23,11 +23,12 @@ use super::super::constants::{AWS_BEARER_TOKEN_BEDROCK, BEDROCK_RUNTIME_ENDPOINT
/// `additionalModelRequestFields` for Anthropic base models and to
/// `inferenceConfig` otherwise, and that branch reads the model catalog the
/// core cannot see.
const SUPPORTED_PARAMS: &[&str] = &["maxTokens", "temperature", "topP", "stopSequences"];
/// Params that belong in `inferenceConfig`, in the order Python's
/// `AmazonConverseConfig` declares them, so bodies compare cleanly.
const INFERENCE_CONFIG_PARAMS: &[&str] = SUPPORTED_PARAMS;
const SUPPORTED_PARAMS: &[(&str, &str)] = &[
("max_tokens", "maxTokens"),
("temperature", "temperature"),
("top_p", "topP"),
("stop", "stopSequences"),
];
const AWS_BEDROCK_RUNTIME_ENDPOINT: &str = "aws_bedrock_runtime_endpoint";
@ -66,7 +67,7 @@ fn converse_body(conversation: &Conversation, params: &Map<String, Value>) -> Va
})
.collect();
let inference_config = Map::from_iter(INFERENCE_CONFIG_PARAMS.iter().filter_map(|name| {
let inference_config = Map::from_iter(SUPPORTED_PARAMS.iter().filter_map(|(_, name)| {
params
.get(*name)
.map(|value| ((*name).to_string(), value.clone()))
@ -162,7 +163,8 @@ impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig {
&[("Content-Type", "application/json")]
}
fn supported_params(&self) -> &'static [&'static str] {
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn supported_openai_params(&self) -> &'static [(&'static str, &'static str)] {
SUPPORTED_PARAMS
}
@ -175,32 +177,36 @@ impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig {
messages: &[ChatMessage],
optional_params: &Map<String, Value>,
) -> Option<Unsupported> {
unsupported_param(SUPPORTED_PARAMS, CONFIG_PARAMS, optional_params)
.or_else(|| messages.iter().find_map(unsupported_message))
// Python's Converse translation drops blank text blocks instead of
// substituting the placeholder the shared conversation builder
// applies, so decline blank text rather than diverge.
.or_else(|| {
messages
.iter()
.any(has_blank_text)
.then_some(Unsupported("blank message text"))
})
// Converse has no assistant prefill: Python inserts a continue turn
// when a conversation opens or closes on an assistant message, and
// only under `litellm.modify_params`, which the core cannot see.
// Declining both ends also keeps the shared builder's final
// assistant right-strip (an Anthropic rule) unreachable here.
.or_else(|| {
let conversation = build_conversation(messages);
let ends_on_assistant = conversation
.turns
.last()
.is_some_and(|turn| turn.role == TurnRole::Assistant);
(!conversation.opens_on_user_turn() || ends_on_assistant).then_some(Unsupported(
"conversation does not run user turn to user turn",
))
})
unsupported_param(
self.supported_openai_params(),
CONFIG_PARAMS,
optional_params,
)
.or_else(|| messages.iter().find_map(unsupported_message))
// Python's Converse translation drops blank text blocks instead of
// substituting the placeholder the shared conversation builder
// applies, so decline blank text rather than diverge.
.or_else(|| {
messages
.iter()
.any(has_blank_text)
.then_some(Unsupported("blank message text"))
})
// Converse has no assistant prefill: Python inserts a continue turn
// when a conversation opens or closes on an assistant message, and
// only under `litellm.modify_params`, which the core cannot see.
// Declining both ends also keeps the shared builder's final
// assistant right-strip (an Anthropic rule) unreachable here.
.or_else(|| {
let conversation = build_conversation(messages);
let ends_on_assistant = conversation
.turns
.last()
.is_some_and(|turn| turn.role == TurnRole::Assistant);
(!conversation.opens_on_user_turn() || ends_on_assistant).then_some(Unsupported(
"conversation does not run user turn to user turn",
))
})
}
fn transform_request(

View file

@ -70,10 +70,12 @@ pub struct MistralOcrConfig;
pub const MISTRAL_OCR_CONFIG: MistralOcrConfig = MistralOcrConfig;
impl OcrProviderConfig for MistralOcrConfig {
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn supported_ocr_params(&self) -> &'static [&'static str] {
SUPPORTED_OCR_PARAMS
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn transform_ocr_request(
&self,
model: &str,
@ -100,6 +102,7 @@ impl OcrProviderConfig for MistralOcrConfig {
})
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn transform_ocr_response(
&self,
model: &str,
@ -134,6 +137,7 @@ impl OcrProviderConfig for MistralOcrConfig {
})
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn complete_url(
&self,
api_base: Option<&str>,
@ -153,6 +157,7 @@ impl OcrProviderConfig for MistralOcrConfig {
}
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub fn supported_ocr_params() -> &'static [&'static str] {
MISTRAL_OCR_CONFIG.supported_ocr_params()
}
@ -161,6 +166,7 @@ pub fn map_ocr_params(non_default_params: &Map<String, Value>) -> Map<String, Va
MISTRAL_OCR_CONFIG.map_ocr_params(non_default_params)
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub fn transform_ocr_request(
model: &str,
document: Value,
@ -169,6 +175,7 @@ pub fn transform_ocr_request(
MISTRAL_OCR_CONFIG.transform_ocr_request(model, document, optional_params)
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub fn transform_ocr_response(model: &str, response_json: Value) -> Result<OcrResponseData, Error> {
MISTRAL_OCR_CONFIG.transform_ocr_response(model, response_json)
}

View file

@ -16,16 +16,21 @@ extension-module = ["pyo3/extension-module"]
panic-test = []
[dependencies]
futures-util.workspace = true
tracing.workspace = true
tracing-subscriber.workspace = true
litellm-core = { workspace = true, features = ["bedrock-auth"] }
litellm-ai-gateway = { workspace = true, default-features = false }
litellm-python-interop.workspace = true
pyo3.workspace = true
pyo3-async-runtimes.workspace = true
serde.workspace = true
serde_json.workspace = true
tokio.workspace = true
[dev-dependencies]
criterion = "0.8.2"
tokio-tungstenite.workspace = true
[[bench]]
name = "serialization"

View file

@ -0,0 +1 @@
pub(crate) const FUNCTION_TRACE_TARGET: &str = "litellm::function_trace";

View file

@ -0,0 +1,23 @@
use litellm_python_interop::release_count;
use pyo3::prelude::*;
use pyo3::types::PyDict;
#[pyfunction]
fn gil_stats(py: Python<'_>) -> PyResult<Py<PyAny>> {
let stats = PyDict::new(py);
stats.set_item("releases", release_count())?;
Ok(stats.into_any().unbind())
}
#[cfg(feature = "panic-test")]
#[pyfunction]
fn _panic_for_test() {
panic!("intentional PyO3 panic smoke test");
}
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
module.add_function(wrap_pyfunction!(gil_stats, module)?)?;
#[cfg(feature = "panic-test")]
module.add_function(wrap_pyfunction!(_panic_for_test, module)?)?;
Ok(())
}

View file

@ -0,0 +1,61 @@
use litellm_core::error::Error;
use pyo3::exceptions::{PyRuntimeError, PyValueError};
use pyo3::prelude::*;
pyo3::create_exception!(
_native,
RustBridgeDeclined,
pyo3::exceptions::PyException,
"The route declined before calling the provider, so the host may retry on its own path."
);
pyo3::create_exception!(
_native,
RustUpstreamError,
pyo3::exceptions::PyException,
"The provider call was already issued and failed. Args are (status, message); status is 0 when there was no HTTP response."
);
pub(crate) fn core_error_to_pyerr(err: Error) -> PyErr {
match err {
Error::Auth(message) => PyValueError::new_err(message),
Error::InvalidProvider(_)
| Error::InvalidRequest(_)
| Error::InvalidType { .. }
| Error::MissingField(_) => PyValueError::new_err(err.to_string()),
other => PyRuntimeError::new_err(other.to_string()),
}
}
/// Map a core 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(err: Error) -> PyErr {
match err {
Error::Unsupported(_)
| Error::Auth(_)
| Error::InvalidProvider(_)
| Error::InvalidRequest(_)
| Error::InvalidType { .. }
| Error::MissingField(_)
| Error::Routing(_)
// Nothing reached the provider, so serving it on Python cannot double
// bill and is the only way the caller gets an answer at all.
| Error::Connect(_) => RustBridgeDeclined::new_err(err.to_string()),
Error::Http { status, body } => {
RustUpstreamError::new_err((status, format!("{status}: {body}")))
}
Error::Network(message) | Error::InvalidResponse(message) => {
RustUpstreamError::new_err((0u16, message))
}
}
}
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
let py = module.py();
module.add("RustBridgeDeclined", py.get_type::<RustBridgeDeclined>())?;
module.add("RustUpstreamError", py.get_type::<RustUpstreamError>())
}

View file

@ -0,0 +1,423 @@
use std::future::Future;
use std::panic::AssertUnwindSafe;
use std::time::Duration;
use futures_util::FutureExt;
use litellm_core::error::Error;
use litellm_python_interop::{Pythonized, panic_to_pyerr, release_gil};
use pyo3::exceptions::PyRuntimeError;
use pyo3::prelude::*;
use serde::Serialize;
use tokio::runtime::{Handle, Runtime};
use tokio::time::{self, MissedTickBehavior};
pub(crate) fn run_sync<T, F>(
py: Python<'_>,
future: F,
map_error: fn(Error) -> PyErr,
) -> PyResult<Py<PyAny>>
where
T: Serialize + Send + 'static,
F: Future<Output = Result<T, Error>> + Send + 'static,
{
run_sync_on(
py,
pyo3_async_runtimes::tokio::get_runtime(),
future,
map_error,
)
}
fn run_sync_on<T, F>(
py: Python<'_>,
runtime: &Runtime,
future: F,
map_error: fn(Error) -> PyErr,
) -> PyResult<Py<PyAny>>
where
T: Serialize + Send + 'static,
F: Future<Output = Result<T, Error>> + Send + 'static,
{
if Handle::try_current().is_ok() {
return Err(PyRuntimeError::new_err(
"synchronous native routes cannot run from a Tokio context; use the async route",
));
}
let result = release_gil(py, move || runtime.block_on(wait_for_sync_result(future)))?;
let result = map_core_result(result, map_error)?;
Pythonized(result).into_pyobject(py).map(Bound::unbind)
}
pub(crate) fn run_async<T, F>(
py: Python<'_>,
future: F,
map_error: fn(Error) -> PyErr,
) -> PyResult<Bound<'_, PyAny>>
where
T: Serialize + Send + 'static,
F: Future<Output = Result<T, Error>> + Send + 'static,
{
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let result = catch_future_panic(future).await?;
let result = map_core_result(result, map_error)?;
Ok(Pythonized(result))
})
}
fn map_core_result<T>(result: Result<T, Error>, map_error: fn(Error) -> PyErr) -> PyResult<T> {
match result {
Ok(value) => Ok(value),
Err(error) => Err(
std::panic::catch_unwind(AssertUnwindSafe(|| map_error(error)))
.map_err(panic_to_pyerr)?,
),
}
}
async fn catch_future_panic<T, F>(future: F) -> PyResult<Result<T, Error>>
where
F: Future<Output = Result<T, Error>>,
{
AssertUnwindSafe(future)
.catch_unwind()
.await
.map_err(panic_to_pyerr)
}
async fn wait_for_sync_result<T, F>(future: F) -> PyResult<Result<T, Error>>
where
F: Future<Output = Result<T, Error>>,
{
let future = catch_future_panic(future);
tokio::pin!(future);
let signal_interval = Duration::from_millis(50);
let mut signal_checks =
time::interval_at(time::Instant::now() + signal_interval, signal_interval);
signal_checks.set_missed_tick_behavior(MissedTickBehavior::Delay);
loop {
tokio::select! {
result = &mut future => return result,
_ = signal_checks.tick() => Python::attach(|py| py.check_signals())?,
}
}
}
#[cfg(test)]
mod tests {
use std::ffi::CString;
use std::future::poll_fn;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, mpsc};
use std::task::Poll;
use std::thread;
use std::time::Instant;
use pyo3::panic::PanicException;
use pyo3::types::{PyDict, PyModule};
use serde::Serializer;
use tokio::runtime::Builder;
use super::*;
fn runtime_error(error: Error) -> PyErr {
PyRuntimeError::new_err(error.to_string())
}
fn panicking_error_mapper(_error: Error) -> PyErr {
panic!("error mapper panicked")
}
struct PanickingOutput;
static ASYNC_PROBE_COMPLETED: AtomicUsize = AtomicUsize::new(0);
impl Serialize for PanickingOutput {
fn serialize<S>(&self, _serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
panic!("serializer panicked")
}
}
#[pyfunction]
fn async_serialization_panic(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
run_async(py, async { Ok(PanickingOutput) }, runtime_error)
}
#[pyfunction]
fn async_runtime_probe(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
run_async(
py,
async {
ASYNC_PROBE_COMPLETED.fetch_add(1, Ordering::SeqCst);
Ok(true)
},
runtime_error,
)
}
#[pyfunction]
fn runtime_worker_count() -> usize {
pyo3_async_runtimes::tokio::get_runtime()
.metrics()
.num_workers()
}
#[pyfunction]
fn runtime_is_responsive(_py: Python<'_>, expected_completions: usize) -> bool {
let completion_deadline = Instant::now() + Duration::from_secs(2);
while ASYNC_PROBE_COMPLETED.load(Ordering::SeqCst) < expected_completions {
if Instant::now() >= completion_deadline {
return false;
}
thread::sleep(Duration::from_millis(1));
}
let (heartbeat_tx, heartbeat_rx) = mpsc::sync_channel(1);
pyo3_async_runtimes::tokio::get_runtime().spawn(async move {
let _ = heartbeat_tx.send(());
});
heartbeat_rx.recv_timeout(Duration::from_secs(2)).is_ok()
}
fn extract_bool(py: Python<'_>, result: PyResult<Py<PyAny>>) -> bool {
result
.expect("route should complete")
.bind(py)
.extract()
.expect("result should convert")
}
#[test]
fn sync_runner_polls_future_on_the_caller_thread() {
Python::initialize();
Python::attach(|py| {
let caller_thread = std::thread::current().id();
let result = run_sync(
py,
async move { Ok(std::thread::current().id() == caller_thread) },
runtime_error,
);
assert!(extract_bool(py, result));
});
}
#[test]
fn sync_runner_releases_gil_while_waiting() {
Python::initialize();
Python::attach(|py| {
let result = run_sync(
py,
async {
let gil_acquired = tokio::time::timeout(
Duration::from_secs(2),
tokio::task::spawn_blocking(|| Python::attach(|_| true)),
)
.await;
Ok(matches!(gil_acquired, Ok(Ok(true))))
},
runtime_error,
);
assert!(extract_bool(py, result));
});
}
#[test]
fn sync_runner_rejects_calls_from_a_tokio_context() {
Python::initialize();
let runtime = Builder::new_current_thread()
.enable_all()
.build()
.expect("runtime should build");
let error = runtime.block_on(async {
Python::attach(|py| {
run_sync::<bool, _>(py, async { Ok(true) }, runtime_error)
.expect_err("sync route should reject a nested Tokio runtime")
})
});
assert_eq!(
error.to_string(),
"RuntimeError: synchronous native routes cannot run from a Tokio context; use the async route"
);
}
#[test]
fn sync_runner_can_drive_a_current_thread_runtime() {
Python::initialize();
let runtime = Builder::new_current_thread()
.enable_all()
.build()
.expect("runtime should build");
Python::attach(|py| {
let result = run_sync_on(
py,
&runtime,
async {
tokio::task::yield_now().await;
Ok(true)
},
runtime_error,
);
assert!(extract_bool(py, result));
});
}
#[test]
fn sync_runner_maps_a_panicked_future() {
Python::initialize();
Python::attach(|py| {
let error = run_sync::<bool, _>(
py,
poll_fn(|_| -> Poll<Result<bool, Error>> { panic!("route future panicked") }),
runtime_error,
)
.expect_err("panicked route should become a Python exception");
assert!(error.is_instance_of::<PanicException>(py));
assert_eq!(error.to_string(), "PanicException: route future panicked");
});
}
#[test]
fn sync_runner_maps_a_panicked_error_mapper() {
Python::initialize();
Python::attach(|py| {
let error = run_sync::<bool, _>(
py,
async { Err(Error::InvalidRequest("invalid".to_string())) },
panicking_error_mapper,
)
.expect_err("panicked mapper should become a Python exception");
assert!(error.is_instance_of::<PanicException>(py));
assert_eq!(error.to_string(), "PanicException: error mapper panicked");
});
}
#[test]
fn sync_runner_surfaces_serializer_panics() {
Python::initialize();
Python::attach(|py| {
let error = run_sync(py, async { Ok(PanickingOutput) }, runtime_error)
.expect_err("serializer panic should become a Python exception");
assert!(error.is_instance_of::<PanicException>(py));
assert_eq!(error.to_string(), "PanicException: serializer panicked");
});
}
#[test]
fn sync_runner_supports_concurrent_callers_on_the_shared_runtime() {
Python::initialize();
let barrier = Arc::new(tokio::sync::Barrier::new(2));
let callers: Vec<_> = (0..2)
.map(|_| {
let barrier = Arc::clone(&barrier);
thread::spawn(move || {
Python::attach(|py| {
extract_bool(
py,
run_sync(
py,
async move {
Ok(tokio::time::timeout(Duration::from_secs(2), barrier.wait())
.await
.is_ok())
},
runtime_error,
),
)
})
})
})
.collect();
let results: Vec<_> = callers
.into_iter()
.map(|caller| caller.join().expect("caller should not panic"))
.collect();
assert_eq!(results, vec![true, true]);
}
#[test]
fn async_runner_surfaces_serializer_panics() {
Python::initialize();
Python::attach(|py| {
let module = PyModule::new(py, "runtime").expect("module should be created");
module
.add_function(
wrap_pyfunction!(async_serialization_panic, &module)
.expect("function should wrap"),
)
.expect("function should register");
let locals = PyDict::new(py);
locals
.set_item("runtime", &module)
.expect("module should enter Python locals");
let code = CString::new(
r#"
import asyncio
async def exercise():
try:
await runtime.async_serialization_panic()
except BaseException as error:
assert type(error).__name__ == "PanicException"
assert str(error) == "serializer panicked"
else:
raise AssertionError("serializer panic was not raised")
asyncio.run(exercise())
"#,
)
.expect("Python source should not contain null bytes");
py.run(&code, Some(&locals), Some(&locals))
.expect("serializer panic should reach the Python awaiter");
});
}
#[test]
fn async_result_delivery_does_not_stall_tokio_workers() {
Python::initialize();
ASYNC_PROBE_COMPLETED.store(0, Ordering::SeqCst);
Python::attach(|py| {
let module = PyModule::new(py, "runtime").expect("module should be created");
for function in [
wrap_pyfunction!(async_runtime_probe, &module).expect("function should wrap"),
wrap_pyfunction!(runtime_worker_count, &module).expect("function should wrap"),
wrap_pyfunction!(runtime_is_responsive, &module).expect("function should wrap"),
] {
module
.add_function(function)
.expect("function should register");
}
let locals = PyDict::new(py);
locals
.set_item("runtime", &module)
.expect("module should enter Python locals");
let code = CString::new(
r#"
import asyncio
async def exercise():
worker_count = runtime.runtime_worker_count()
awaitables = [runtime.async_runtime_probe() for _ in range(worker_count)]
assert runtime.runtime_is_responsive(worker_count)
assert await asyncio.gather(*awaitables) == [True] * worker_count
asyncio.run(exercise())
"#,
)
.expect("Python source should not contain null bytes");
py.run(&code, Some(&locals), Some(&locals))
.expect("result delivery should leave Tokio workers responsive");
});
}
}

View file

@ -0,0 +1,216 @@
use std::future::Future;
use std::sync::{Arc, Mutex};
use serde::Serialize;
use tracing::instrument::WithSubscriber;
use tracing::span::{Attributes, Id};
use tracing::{Dispatch, Level, Subscriber};
use tracing_subscriber::filter::{LevelFilter, filter_fn};
use tracing_subscriber::layer::Context;
use tracing_subscriber::prelude::*;
use tracing_subscriber::registry::LookupSpan;
use tracing_subscriber::{Layer, Registry};
use crate::constants::FUNCTION_TRACE_TARGET;
#[derive(Serialize)]
#[serde(untagged)]
pub(crate) enum TraceResponse<T> {
Plain(T),
Traced {
response: T,
trace: Vec<FunctionTraceEvent>,
},
}
pub(crate) async fn trace_call<T, E>(
future: impl Future<Output = Result<T, E>>,
enabled: bool,
) -> Result<TraceResponse<T>, E> {
if !enabled {
return future.await.map(TraceResponse::Plain);
}
let trace = FunctionTrace::default();
let response = future.with_subscriber(trace.dispatcher()).await?;
Ok(TraceResponse::Traced {
response,
trace: trace.events(),
})
}
#[derive(Clone, Debug, PartialEq, Serialize)]
pub struct FunctionTraceEvent {
pub function: &'static str,
pub depth: usize,
}
#[derive(Clone, Default)]
pub struct FunctionTrace {
events: Arc<Mutex<Vec<FunctionTraceEvent>>>,
}
impl FunctionTrace {
pub fn dispatcher(&self) -> Dispatch {
let filter = filter_fn(|metadata| {
metadata.is_span()
&& metadata.target() == FUNCTION_TRACE_TARGET
&& *metadata.level() == Level::TRACE
})
.with_max_level_hint(LevelFilter::TRACE);
Dispatch::new(
Registry::default().with(
FunctionTraceLayer {
trace: self.clone(),
}
.with_filter(filter),
),
)
}
pub fn events(&self) -> Vec<FunctionTraceEvent> {
self.events
.lock()
.unwrap_or_else(|error| error.into_inner())
.clone()
}
}
struct FunctionTraceLayer {
trace: FunctionTrace,
}
impl<S> Layer<S> for FunctionTraceLayer
where
S: Subscriber + for<'lookup> LookupSpan<'lookup>,
{
fn on_new_span(&self, attributes: &Attributes<'_>, id: &Id, context: Context<'_, S>) {
let depth = context
.span(id)
.map(|span| span.scope().skip(1).count())
.unwrap_or_default();
self.trace
.events
.lock()
.unwrap_or_else(|error| error.into_inner())
.push(FunctionTraceEvent {
function: attributes.metadata().name(),
depth,
});
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
async fn outer() {
tokio::task::yield_now().await;
inner().await;
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
async fn inner() {
tokio::task::yield_now().await;
}
#[tokio::test]
async fn concurrent_futures_keep_separate_traces_across_yields() {
use tracing::instrument::WithSubscriber;
let first = FunctionTrace::default();
let second = FunctionTrace::default();
let outside = FunctionTrace::default();
async {
tokio::join!(
outer().with_subscriber(first.dispatcher()),
inner().with_subscriber(second.dispatcher()),
);
inner().await;
}
.with_subscriber(outside.dispatcher())
.await;
assert_eq!(
first.events(),
vec![
FunctionTraceEvent {
function: "outer",
depth: 0
},
FunctionTraceEvent {
function: "inner",
depth: 1
},
],
);
assert_eq!(
second.events(),
vec![FunctionTraceEvent {
function: "inner",
depth: 0
}],
);
assert_eq!(
outside.events(),
vec![FunctionTraceEvent {
function: "inner",
depth: 0
}],
);
}
#[test]
fn records_matching_spans_in_creation_order() {
let trace = FunctionTrace::default();
let dispatch = trace.dispatcher();
tracing::dispatcher::with_default(&dispatch, || {
let _ignored = tracing::trace_span!(target: "other", "ignored");
let _first = tracing::trace_span!(target: FUNCTION_TRACE_TARGET, "same_name");
let _wrong_level = tracing::debug_span!(target: FUNCTION_TRACE_TARGET, "wrong_level");
let _second = tracing::trace_span!(target: FUNCTION_TRACE_TARGET, "same_name");
});
assert_eq!(
trace.events(),
vec![
FunctionTraceEvent {
function: "same_name",
depth: 0,
},
FunctionTraceEvent {
function: "same_name",
depth: 0,
},
]
);
}
#[test]
fn records_matching_span_nesting_depth() {
let trace = FunctionTrace::default();
let dispatch = trace.dispatcher();
tracing::dispatcher::with_default(&dispatch, || {
let outer = tracing::trace_span!(target: FUNCTION_TRACE_TARGET, "outer");
let _outer_guard = outer.enter();
let _inner = tracing::trace_span!(target: FUNCTION_TRACE_TARGET, "inner");
});
assert_eq!(
trace.events(),
vec![
FunctionTraceEvent {
function: "outer",
depth: 0,
},
FunctionTraceEvent {
function: "inner",
depth: 1,
},
]
);
}
}

View file

@ -1,142 +1,18 @@
use std::collections::HashMap;
use std::time::Duration;
mod constants;
mod diagnostics;
mod errors;
mod execution;
pub mod function_trace;
mod marshal;
mod routes;
use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr};
use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
use litellm_core::audio_transcription::{
AudioTranscriptionRequest, audio_transcription as run_audio_transcription,
};
use litellm_core::chat_completions::types::{ChatCompletionsRequest, ChatCompletionsResponse};
use litellm_core::chat_completions::{
chat_completions as run_chat_completions, chat_completions_decline_reason,
};
use litellm_core::error::Error;
use litellm_core::messages::messages as run_messages;
use litellm_core::messages::types::{AnthropicMessagesResponse, MessagesRequest};
use litellm_python_interop::{from_py, release_count, release_gil, to_py};
use pyo3::exceptions::{PyRuntimeError, PyValueError};
use pyo3::prelude::*;
use pyo3::types::{PyAny, PyDict};
use serde_json::{Map, Value};
use pyo3::types::PyAny;
use serde_json::Value;
pyo3::create_exception!(
_native,
RustBridgeDeclined,
pyo3::exceptions::PyException,
"The route declined before calling the provider, so the host may retry on its own path."
);
pyo3::create_exception!(
_native,
RustUpstreamError,
pyo3::exceptions::PyException,
"The provider call was already issued and failed. Args are (status, message); status is 0 when there was no HTTP response."
);
type MarshaledOcrInputs = (
Value,
Option<Map<String, Value>>,
Map<String, Value>,
Option<Duration>,
);
fn messages_response_to_py(
py: Python<'_>,
response: AnthropicMessagesResponse,
) -> PyResult<Py<PyAny>> {
to_py(py, &response)
}
fn chat_completions_response_to_py(
py: Python<'_>,
response: ChatCompletionsResponse,
) -> PyResult<Py<PyAny>> {
to_py(py, &response)
}
fn core_error_to_pyerr(err: Error) -> PyErr {
match err {
Error::Auth(message) => PyValueError::new_err(message),
Error::InvalidProvider(_)
| Error::InvalidRequest(_)
| Error::InvalidType { .. }
| Error::MissingField(_) => PyValueError::new_err(err.to_string()),
other => PyRuntimeError::new_err(other.to_string()),
}
}
/// Map a core 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.
fn chat_completions_error_to_pyerr(err: Error) -> PyErr {
match err {
Error::Unsupported(_)
| Error::Auth(_)
| Error::InvalidProvider(_)
| Error::InvalidRequest(_)
| Error::InvalidType { .. }
| Error::MissingField(_)
| Error::Routing(_)
// Nothing reached the provider, so serving it on Python cannot double
// bill and is the only way the caller gets an answer at all.
| Error::Connect(_) => RustBridgeDeclined::new_err(err.to_string()),
Error::Http { status, body } => {
RustUpstreamError::new_err((status, format!("{status}: {body}")))
}
Error::Network(message) | Error::InvalidResponse(message) => {
RustUpstreamError::new_err((0u16, message))
}
}
}
fn optional_object_to_map(
py: Python<'_>,
name: &'static str,
value: Option<Py<PyAny>>,
) -> PyResult<Map<String, Value>> {
match value {
Some(value) => match from_py(value.bind(py))? {
Value::Object(map) => Ok(map),
_ => Err(PyValueError::new_err(format!("{name} must be a dict"))),
},
None => Ok(Map::new()),
}
}
fn optional_timeout(timeout_seconds: Option<f64>) -> Option<Duration> {
timeout_seconds.and_then(|secs| {
if secs.is_finite() && secs > 0.0 {
Some(Duration::from_secs_f64(secs))
} else {
None
}
})
}
fn marshal_headers(
py: Python<'_>,
headers: Option<Py<PyAny>>,
) -> PyResult<HashMap<String, String>> {
let value = match headers {
Some(headers) => from_py(headers.bind(py))?,
None => Value::Object(Map::new()),
};
let Value::Object(headers) = value else {
return Err(PyValueError::new_err("headers must be a dict"));
};
headers
.into_iter()
.map(|(name, value)| {
value
.as_str()
.map(|value| (name, value.to_string()))
.ok_or_else(|| PyValueError::new_err("header values must be strings"))
})
.collect()
}
use crate::errors::core_error_to_pyerr;
use crate::marshal::{marshal_headers, optional_timeout};
#[pyclass]
struct ResponsesWebSocketConnection {
@ -151,16 +27,16 @@ impl ResponsesWebSocketConnection {
_cls: &Bound<'py, pyo3::types::PyType>,
py: Python<'py>,
url: String,
headers: Option<Py<PyAny>>,
#[pyo3(from_py_with = litellm_python_interop::from_py)] headers: Option<Value>,
timeout_seconds: Option<f64>,
) -> PyResult<Bound<'py, PyAny>> {
let headers = marshal_headers(py, headers)?;
let headers = marshal_headers(headers)?;
let timeout = optional_timeout(timeout_seconds);
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let inner = RustResponsesWebSocketConnection::connect_url(&url, &headers, timeout)
.await
.map_err(core_error_to_pyerr)?;
Python::attach(|py| Py::new(py, ResponsesWebSocketConnection { inner }))
Ok(ResponsesWebSocketConnection { inner })
})
}
@ -186,445 +62,126 @@ impl ResponsesWebSocketConnection {
}
}
fn marshal_inputs(
py: Python<'_>,
document: Py<PyAny>,
extra_headers: Option<Py<PyAny>>,
optional_params: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<MarshaledOcrInputs> {
let document = from_py(document.bind(py))?;
let extra_headers = match extra_headers {
Some(headers) => Some(optional_object_to_map(py, "extra_headers", Some(headers))?),
None => None,
};
let optional_params = optional_object_to_map(py, "optional_params", optional_params)?;
let timeout = optional_timeout(timeout_seconds);
#[pymodule(gil_used = false)]
mod _native {
use pyo3::prelude::*;
Ok((document, extra_headers, optional_params, timeout))
}
#[pyfunction]
#[pyo3(signature = (model, document, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
fn ocr(
py: Python<'_>,
model: String,
document: Py<PyAny>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
extra_headers: Option<Py<PyAny>>,
optional_params: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Py<PyAny>> {
let (document, extra_headers, optional_params, timeout) = marshal_inputs(
py,
document,
extra_headers,
optional_params,
timeout_seconds,
)?;
let result = release_gil(py, || {
pyo3_async_runtimes::tokio::get_runtime().block_on(run_ocr(OcrRequest {
model: &model,
document,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
optional_params,
timeout,
callbacks: Vec::new(),
guardrails: Vec::new(),
request_metadata: Default::default(),
litellm_call_id: None,
}))
});
match result {
Ok(value) => to_py(py, &value),
Err(err) => Err(core_error_to_pyerr(err)),
#[pymodule_init]
fn init(module: &Bound<'_, PyModule>) -> PyResult<()> {
super::errors::register(module)?;
super::routes::register(module)?;
module.add_class::<super::ResponsesWebSocketConnection>()?;
super::diagnostics::register(module)
}
}
#[pyfunction]
#[pyo3(signature = (model, document, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
fn aocr(
py: Python<'_>,
model: String,
document: Py<PyAny>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
extra_headers: Option<Py<PyAny>>,
optional_params: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Bound<'_, PyAny>> {
let (document, extra_headers, optional_params, timeout) = marshal_inputs(
py,
document,
extra_headers,
optional_params,
timeout_seconds,
)?;
#[cfg(test)]
mod tests {
use std::ffi::CString;
use std::time::Duration;
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let value = run_ocr(OcrRequest {
model: &model,
document,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
optional_params,
timeout,
callbacks: Vec::new(),
guardrails: Vec::new(),
request_metadata: Default::default(),
litellm_call_id: None,
})
.await
.map_err(core_error_to_pyerr)?;
use futures_util::{SinkExt, StreamExt};
use pyo3::types::PyDict;
use tokio::net::TcpListener;
use tokio_tungstenite::{accept_async, tungstenite::Message};
Python::attach(|py| to_py(py, &value))
})
}
use super::*;
#[pyfunction]
#[pyo3(signature = (model, audio, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
fn transcription(
py: Python<'_>,
model: String,
audio: Py<PyAny>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
extra_headers: Option<Py<PyAny>>,
optional_params: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Py<PyAny>> {
let audio = from_py(audio.bind(py))?;
let extra_headers = match extra_headers {
Some(headers) => Some(optional_object_to_map(py, "extra_headers", Some(headers))?),
None => None,
};
let optional_params = optional_object_to_map(py, "optional_params", optional_params)?;
let timeout = optional_timeout(timeout_seconds);
let result = release_gil(py, || {
pyo3_async_runtimes::tokio::get_runtime().block_on(run_audio_transcription(
AudioTranscriptionRequest {
model: &model,
audio,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
optional_params,
timeout,
},
))
});
match result {
Ok(value) => to_py(py, &value),
Err(err) => Err(core_error_to_pyerr(err)),
#[test]
fn module_registration_preserves_the_public_surface() {
Python::initialize();
Python::attach(|py| {
let module = pyo3::wrap_pymodule!(_native)(py).into_bound(py);
let expected = [
"RustBridgeDeclined",
"RustUpstreamError",
"ocr",
"aocr",
"transcription",
"atranscription",
"messages",
"amessages",
"chat_completions_decline",
"chat_completions",
"achat_completions",
"ResponsesWebSocketConnection",
"gil_stats",
];
let public_names: Vec<String> = module
.dict()
.keys()
.extract::<Vec<String>>()
.expect("module names should be strings")
.into_iter()
.filter(|name| !name.starts_with("__"))
.collect();
assert_eq!(public_names, expected);
});
}
#[test]
fn responses_websocket_connection_round_trips_through_python() {
Python::initialize();
let runtime = pyo3_async_runtimes::tokio::get_runtime();
let listener = runtime
.block_on(TcpListener::bind("127.0.0.1:0"))
.expect("listener should bind");
let address = listener
.local_addr()
.expect("listener should have an address");
let server = runtime.spawn(async move {
let (stream, _) = listener.accept().await.expect("server should accept");
let mut socket = accept_async(stream)
.await
.expect("handshake should succeed");
let message = socket
.next()
.await
.expect("client should send a frame")
.expect("client frame should be valid");
assert_eq!(message, Message::Text("from-python".into()));
socket
.send(Message::Text("from-server".into()))
.await
.expect("server should reply");
assert!(matches!(socket.next().await, Some(Ok(Message::Close(_)))));
});
Python::attach(|py| {
let module = pyo3::wrap_pymodule!(_native)(py).into_bound(py);
let locals = PyDict::new(py);
locals
.set_item("native", &module)
.expect("module should enter Python locals");
locals
.set_item("url", format!("ws://{address}"))
.expect("URL should enter Python locals");
let code = CString::new(
r#"
import asyncio
async def exercise():
connection = await native.ResponsesWebSocketConnection.connect(url)
assert type(connection) is native.ResponsesWebSocketConnection
await connection.send_text("from-python")
assert await connection.recv_text() == "from-server"
await connection.close()
assert await connection.recv_text() is None
asyncio.run(asyncio.wait_for(exercise(), timeout=5))
"#,
)
.expect("Python source should not contain null bytes");
py.run(&code, Some(&locals), Some(&locals))
.expect("Python WebSocket methods should round trip");
});
runtime
.block_on(async { tokio::time::timeout(Duration::from_secs(5), server).await })
.expect("server should finish")
.expect("server task should not panic");
}
}
#[pyfunction]
#[pyo3(signature = (model, audio, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
fn atranscription(
py: Python<'_>,
model: String,
audio: Py<PyAny>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
extra_headers: Option<Py<PyAny>>,
optional_params: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Bound<'_, PyAny>> {
let audio = from_py(audio.bind(py))?;
let extra_headers = match extra_headers {
Some(headers) => Some(optional_object_to_map(py, "extra_headers", Some(headers))?),
None => None,
};
let optional_params = optional_object_to_map(py, "optional_params", optional_params)?;
let timeout = optional_timeout(timeout_seconds);
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let value = run_audio_transcription(AudioTranscriptionRequest {
model: &model,
audio,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
optional_params,
timeout,
})
.await
.map_err(core_error_to_pyerr)?;
Python::attach(|py| to_py(py, &value))
})
}
type MarshaledMessagesInputs = (Value, Option<Map<String, Value>>, Option<Duration>);
fn marshal_messages_inputs(
py: Python<'_>,
body: Py<PyAny>,
extra_headers: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<MarshaledMessagesInputs> {
let body: Value = from_py(body.bind(py))?;
if !body.is_object() {
return Err(PyValueError::new_err("body must be a dict"));
}
let extra_headers = match extra_headers {
Some(headers) => Some(optional_object_to_map(py, "extra_headers", Some(headers))?),
None => None,
};
Ok((body, extra_headers, optional_timeout(timeout_seconds)))
}
#[pyfunction]
#[pyo3(signature = (model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
fn messages(
py: Python<'_>,
model: String,
body: Py<PyAny>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
extra_headers: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Py<PyAny>> {
let (body, extra_headers, timeout) =
marshal_messages_inputs(py, body, extra_headers, timeout_seconds)?;
let result = release_gil(py, || {
pyo3_async_runtimes::tokio::get_runtime().block_on(run_messages(MessagesRequest {
model: &model,
body,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
timeout,
}))
});
match result {
Ok(response) => messages_response_to_py(py, response),
Err(err) => Err(core_error_to_pyerr(err)),
}
}
#[pyfunction]
#[pyo3(signature = (model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
fn amessages(
py: Python<'_>,
model: String,
body: Py<PyAny>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
extra_headers: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Bound<'_, PyAny>> {
let (body, extra_headers, timeout) =
marshal_messages_inputs(py, body, extra_headers, timeout_seconds)?;
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let response = run_messages(MessagesRequest {
model: &model,
body,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
timeout,
})
.await
.map_err(core_error_to_pyerr)?;
Python::attach(|py| messages_response_to_py(py, response))
})
}
type MarshaledChatCompletionsInputs = (
Value,
Map<String, Value>,
Option<Map<String, Value>>,
Option<Duration>,
);
fn marshal_chat_completions_inputs(
py: Python<'_>,
messages: Py<PyAny>,
optional_params: Option<Py<PyAny>>,
extra_headers: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<MarshaledChatCompletionsInputs> {
let messages: Value = from_py(messages.bind(py))?;
if !messages.is_array() {
return Err(PyValueError::new_err("messages must be a list"));
}
let optional_params = optional_object_to_map(py, "optional_params", optional_params)?;
let extra_headers = match extra_headers {
Some(headers) => Some(optional_object_to_map(py, "extra_headers", Some(headers))?),
None => None,
};
Ok((
messages,
optional_params,
extra_headers,
optional_timeout(timeout_seconds),
))
}
/// The decline reason for this request, or `None` when the Rust path accepts
/// it. Resolves no credentials and performs no I/O, so a host can ask before
/// committing to either path.
#[pyfunction]
#[pyo3(signature = (model, messages, optional_params=None, custom_llm_provider=None))]
fn chat_completions_decline(
py: Python<'_>,
model: String,
messages: Py<PyAny>,
optional_params: Option<Py<PyAny>>,
custom_llm_provider: Option<String>,
) -> PyResult<Option<String>> {
let messages = from_py(messages.bind(py))?;
let optional_params = optional_object_to_map(py, "optional_params", optional_params)?;
Ok(chat_completions_decline_reason(
&model,
custom_llm_provider.as_deref(),
messages,
&optional_params,
)
.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))]
#[allow(clippy::too_many_arguments)]
fn chat_completions(
py: Python<'_>,
model: String,
messages: Py<PyAny>,
optional_params: Option<Py<PyAny>>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
extra_headers: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Py<PyAny>> {
let (messages, optional_params, extra_headers, timeout) = marshal_chat_completions_inputs(
py,
messages,
optional_params,
extra_headers,
timeout_seconds,
)?;
let result = release_gil(py, || {
pyo3_async_runtimes::tokio::get_runtime().block_on(run_chat_completions(
ChatCompletionsRequest {
model: &model,
messages,
optional_params,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
timeout,
},
))
});
match result {
Ok(response) => chat_completions_response_to_py(py, response),
Err(err) => Err(chat_completions_error_to_pyerr(err)),
}
}
#[pyfunction]
#[pyo3(signature = (model, messages, optional_params=None, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
fn achat_completions(
py: Python<'_>,
model: String,
messages: Py<PyAny>,
optional_params: Option<Py<PyAny>>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
extra_headers: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Bound<'_, PyAny>> {
let (messages, optional_params, extra_headers, timeout) = marshal_chat_completions_inputs(
py,
messages,
optional_params,
extra_headers,
timeout_seconds,
)?;
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let response = run_chat_completions(ChatCompletionsRequest {
model: &model,
messages,
optional_params,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
timeout,
})
.await
.map_err(chat_completions_error_to_pyerr)?;
Python::attach(|py| chat_completions_response_to_py(py, response))
})
}
#[pyfunction]
fn gil_stats(py: Python<'_>) -> PyResult<Py<PyAny>> {
let stats = PyDict::new(py);
stats.set_item("releases", release_count())?;
Ok(stats.into_any().unbind())
}
#[cfg(feature = "panic-test")]
#[pyfunction]
fn _panic_for_test() {
panic!("intentional PyO3 panic smoke test");
}
#[pymodule]
fn _native(module: &Bound<'_, PyModule>) -> PyResult<()> {
let py = module.py();
module.add_function(wrap_pyfunction!(ocr, module)?)?;
module.add_function(wrap_pyfunction!(aocr, module)?)?;
module.add_function(wrap_pyfunction!(transcription, module)?)?;
module.add_function(wrap_pyfunction!(atranscription, module)?)?;
module.add_function(wrap_pyfunction!(messages, module)?)?;
module.add_function(wrap_pyfunction!(amessages, module)?)?;
module.add("RustBridgeDeclined", py.get_type::<RustBridgeDeclined>())?;
module.add("RustUpstreamError", py.get_type::<RustUpstreamError>())?;
module.add_function(wrap_pyfunction!(chat_completions_decline, module)?)?;
module.add_function(wrap_pyfunction!(chat_completions, module)?)?;
module.add_function(wrap_pyfunction!(achat_completions, module)?)?;
module.add_class::<ResponsesWebSocketConnection>()?;
module.add_function(wrap_pyfunction!(gil_stats, module)?)?;
#[cfg(feature = "panic-test")]
module.add_function(wrap_pyfunction!(_panic_for_test, module)?)?;
Ok(())
}

View file

@ -0,0 +1,104 @@
use std::collections::HashMap;
use std::time::Duration;
use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
use serde_json::{Map, Value};
pub(crate) struct RouteOptions {
pub(crate) model: String,
pub(crate) api_key: Option<String>,
pub(crate) api_base: Option<String>,
pub(crate) custom_llm_provider: Option<String>,
pub(crate) extra_headers: Option<Map<String, Value>>,
pub(crate) timeout: Option<Duration>,
}
pub(crate) struct RouteOptionsInputs {
pub(crate) model: String,
pub(crate) api_key: Option<String>,
pub(crate) api_base: Option<String>,
pub(crate) custom_llm_provider: Option<String>,
pub(crate) extra_headers: Option<Value>,
pub(crate) timeout_seconds: Option<f64>,
}
impl RouteOptions {
pub(crate) fn from_python(inputs: RouteOptionsInputs) -> PyResult<Self> {
Ok(Self {
model: inputs.model,
api_key: inputs.api_key,
api_base: inputs.api_base,
custom_llm_provider: inputs.custom_llm_provider,
extra_headers: optional_object("extra_headers", inputs.extra_headers)?,
timeout: optional_timeout(inputs.timeout_seconds),
})
}
}
pub(crate) fn required_value(
name: &'static str,
value: Value,
expected: fn(&Value) -> bool,
expected_name: &'static str,
) -> PyResult<Value> {
if expected(&value) {
return Ok(value);
}
Err(PyValueError::new_err(format!(
"{name} must be a {expected_name}"
)))
}
pub(crate) fn object_or_empty(
name: &'static str,
value: Option<Value>,
) -> PyResult<Map<String, Value>> {
match value {
Some(value) => object(name, value),
None => Ok(Map::new()),
}
}
fn optional_object(
name: &'static str,
value: Option<Value>,
) -> PyResult<Option<Map<String, Value>>> {
value.map(|value| object(name, value)).transpose()
}
fn object(name: &'static str, value: Value) -> PyResult<Map<String, Value>> {
match value {
Value::Object(map) => Ok(map),
_ => Err(PyValueError::new_err(format!("{name} must be a dict"))),
}
}
pub(crate) fn optional_timeout(timeout_seconds: Option<f64>) -> Option<Duration> {
timeout_seconds.and_then(|secs| {
if secs.is_finite() && secs > 0.0 {
Some(Duration::from_secs_f64(secs))
} else {
None
}
})
}
pub(crate) fn marshal_headers(headers: Option<Value>) -> PyResult<HashMap<String, String>> {
let value = match headers {
Some(headers) => headers,
None => Value::Object(Map::new()),
};
let Value::Object(headers) = value else {
return Err(PyValueError::new_err("headers must be a dict"));
};
headers
.into_iter()
.map(|(name, value)| {
value
.as_str()
.map(|value| (name, value.to_string()))
.ok_or_else(|| PyValueError::new_err("header values must be strings"))
})
.collect()
}

View file

@ -0,0 +1,71 @@
use litellm_core::Error;
use std::future::Future;
use litellm_core::audio_transcription::{
AudioTranscriptionRequest, audio_transcription as run_audio_transcription,
};
use pyo3::prelude::*;
use serde_json::Value;
use crate::errors::core_error_to_pyerr;
use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty};
fn prepare_transcription(
inputs: AudioTranscriptionInputs,
) -> PyResult<impl Future<Output = Result<Value, Error>> + Send + 'static> {
let audio = inputs.audio;
let options = RouteOptions::from_python(RouteOptionsInputs {
model: inputs.model,
api_key: inputs.api_key,
api_base: inputs.api_base,
custom_llm_provider: inputs.custom_llm_provider,
extra_headers: inputs.extra_headers,
timeout_seconds: inputs.timeout_seconds,
})?;
let optional_params = object_or_empty("optional_params", inputs.optional_params)?;
Ok(async move {
let RouteOptions {
model,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout,
} = options;
run_audio_transcription(AudioTranscriptionRequest {
model: &model,
audio,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
optional_params,
timeout,
})
.await
})
}
bridge_route! {
sync = transcription,
asynchronous = atranscription,
inputs = AudioTranscriptionInputs,
required = {
model: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
audio: Value,
},
optional = {
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
extra_headers: Option<Value>,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
optional_params: Option<Value>,
timeout_seconds: Option<f64>,
},
prepare = prepare_transcription,
errors = core_error_to_pyerr,
}

View file

@ -0,0 +1,91 @@
use litellm_core::Error;
use std::future::Future;
use litellm_core::chat_completions::types::{ChatCompletionsRequest, ChatCompletionsResponse};
use litellm_core::chat_completions::{
chat_completions as run_chat_completions, chat_completions_decline_reason,
};
use pyo3::prelude::*;
use serde_json::Value;
use crate::errors::chat_completions_error_to_pyerr;
use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty, required_value};
fn prepare_chat_completions(
inputs: ChatCompletionsInputs,
) -> PyResult<impl Future<Output = Result<ChatCompletionsResponse, Error>> + Send + 'static> {
let messages = required_value("messages", inputs.messages, Value::is_array, "list")?;
let optional_params = object_or_empty("optional_params", inputs.optional_params)?;
let options = RouteOptions::from_python(RouteOptionsInputs {
model: inputs.model,
api_key: inputs.api_key,
api_base: inputs.api_base,
custom_llm_provider: inputs.custom_llm_provider,
extra_headers: inputs.extra_headers,
timeout_seconds: inputs.timeout_seconds,
})?;
Ok(async move {
let RouteOptions {
model,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout,
} = options;
run_chat_completions(ChatCompletionsRequest {
model: &model,
messages,
optional_params,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
timeout,
})
.await
})
}
#[pyfunction]
#[pyo3(signature = (model, messages, optional_params=None, custom_llm_provider=None))]
fn chat_completions_decline(
model: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)] messages: Value,
#[pyo3(from_py_with = litellm_python_interop::from_py)] optional_params: Option<Value>,
custom_llm_provider: Option<String>,
) -> PyResult<Option<String>> {
let optional_params = object_or_empty("optional_params", optional_params)?;
Ok(chat_completions_decline_reason(
&model,
custom_llm_provider.as_deref(),
messages,
&optional_params,
)
.map(str::to_string))
}
bridge_route! {
sync = chat_completions,
asynchronous = achat_completions,
inputs = ChatCompletionsInputs,
required = {
model: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
messages: Value,
},
optional = {
#[pyo3(from_py_with = litellm_python_interop::from_py)]
optional_params: Option<Value>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
extra_headers: Option<Value>,
timeout_seconds: Option<f64>,
},
prepare = prepare_chat_completions,
errors = chat_completions_error_to_pyerr,
extra = [chat_completions_decline],
}

View file

@ -0,0 +1,429 @@
use pyo3::exceptions::PyRuntimeError;
use pyo3::prelude::*;
use pyo3::types::PyCFunction;
macro_rules! bridge_route {
(
sync = $sync_name:ident,
asynchronous = $async_name:ident,
inputs = $inputs:ident,
required = { $($(#[$required_attr:meta])* $required_name:ident: $required_type:ty),+ $(,)? },
optional = { $($(#[$optional_attr:meta])* $optional_name:ident: $optional_type:ty),* $(,)? },
prepare = $prepare:path,
errors = $map_error:path
$(, extra = [$($extra:ident),* $(,)?])?
$(,)?
) => {
struct $inputs {
$($required_name: $required_type,)*
$($optional_name: $optional_type),*
}
#[pyfunction]
#[pyo3(signature = ($($required_name),*, $($optional_name=None,)* trace=false))]
#[allow(clippy::too_many_arguments)]
fn $sync_name(
py: pyo3::Python<'_>,
$($(#[$required_attr])* $required_name: $required_type,)*
$($(#[$optional_attr])* $optional_name: $optional_type,)*
trace: bool,
) -> pyo3::PyResult<pyo3::Py<pyo3::PyAny>> {
let future = $prepare($inputs {
$($required_name,)*
$($optional_name),*
})?;
$crate::execution::run_sync(
py,
$crate::function_trace::trace_call(future, trace),
$map_error,
)
}
#[pyfunction]
#[pyo3(signature = ($($required_name),*, $($optional_name=None,)* trace=false))]
#[allow(clippy::too_many_arguments)]
fn $async_name(
py: pyo3::Python<'_>,
$($(#[$required_attr])* $required_name: $required_type,)*
$($(#[$optional_attr])* $optional_name: $optional_type,)*
trace: bool,
) -> pyo3::PyResult<pyo3::Bound<'_, pyo3::PyAny>> {
let future = $prepare($inputs {
$($required_name,)*
$($optional_name),*
})?;
$crate::execution::run_async(
py,
$crate::function_trace::trace_call(future, trace),
$map_error,
)
}
pub(super) fn register(
module: &pyo3::Bound<'_, pyo3::types::PyModule>,
) -> pyo3::PyResult<()> {
$($($crate::routes::definition::add_function(module, pyo3::wrap_pyfunction!($extra, module)?)?;)*)?
$crate::routes::definition::add_function(module, pyo3::wrap_pyfunction!($sync_name, module)?)?;
$crate::routes::definition::add_function(module, pyo3::wrap_pyfunction!($async_name, module)?)?;
Ok(())
}
};
}
pub(super) fn add_function(
module: &Bound<'_, PyModule>,
function: Bound<'_, PyCFunction>,
) -> PyResult<()> {
let name: String = function.getattr("__name__")?.extract()?;
if module.hasattr(&name)? {
return Err(PyRuntimeError::new_err(format!(
"duplicate native route: {name}"
)));
}
module.add_function(function)
}
#[cfg(test)]
mod tests {
use std::ffi::CString;
use std::sync::atomic::{AtomicBool, Ordering};
use litellm_core::error::Error;
use pyo3::exceptions::PyLookupError;
use pyo3::types::{PyDict, PyList};
use super::*;
mod synthetic {
use std::future::{Future, pending};
use super::*;
static FUTURE_DROPPED: AtomicBool = AtomicBool::new(false);
struct DropGuard;
impl Drop for DropGuard {
fn drop(&mut self) {
FUTURE_DROPPED.store(true, Ordering::SeqCst);
}
}
#[pyfunction]
fn future_dropped() -> bool {
FUTURE_DROPPED.load(Ordering::SeqCst)
}
bridge_route! {
sync = echo,
asynchronous = aecho,
inputs = EchoInputs,
required = { value: String },
optional = {},
prepare = prepare_echo,
errors = map_error,
extra = [future_dropped],
}
fn prepare_echo(
inputs: EchoInputs,
) -> PyResult<impl Future<Output = Result<String, Error>> + Send + 'static> {
FUTURE_DROPPED.store(false, Ordering::SeqCst);
let drop_guard = (inputs.value == "pending").then_some(DropGuard);
Ok(async move {
let _drop_guard = drop_guard;
tokio::task::yield_now().await;
match inputs.value.as_str() {
"error" => Err(Error::InvalidRequest("synthetic error".to_string())),
"map_panic" => Err(Error::InvalidRequest("panic in mapper".to_string())),
"panic" => panic!("synthetic panic"),
"pending" => {
pending::<()>().await;
unreachable!()
}
_ => Ok(inputs.value),
}
})
}
fn map_error(error: Error) -> PyErr {
if matches!(&error, Error::InvalidRequest(message) if message == "panic in mapper") {
panic!("synthetic mapper panic")
}
PyLookupError::new_err(error.to_string())
}
}
#[test]
fn sync_and_async_route_signatures_match_the_python_contract() {
Python::initialize();
Python::attach(|py| {
let module = PyModule::new(py, "routes").expect("module should be created");
crate::routes::register(&module).expect("routes should register");
let routes = [
(
"ocr",
"aocr",
"(model, document, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None, trace=False)",
),
(
"transcription",
"atranscription",
"(model, audio, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None, trace=False)",
),
(
"messages",
"amessages",
"(model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None, trace=False)",
),
(
"chat_completions",
"achat_completions",
"(model, messages, optional_params=None, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None, trace=False)",
),
];
for (sync_name, async_name, expected) in routes {
let sync_signature: String = module
.getattr(sync_name)
.and_then(|function| function.getattr("__text_signature__"))
.and_then(|signature| signature.extract())
.expect("sync signature should be available");
let async_signature: String = module
.getattr(async_name)
.and_then(|function| function.getattr("__text_signature__"))
.and_then(|signature| signature.extract())
.expect("async signature should be available");
assert_eq!(sync_signature, expected);
assert_eq!(async_signature, expected);
}
});
}
#[test]
fn sync_and_async_routes_apply_the_same_input_validation() {
Python::initialize();
Python::attach(|py| {
let module = PyModule::new(py, "routes").expect("module should be created");
crate::routes::register(&module).expect("routes should register");
let invalid_messages = PyDict::new(py);
let sync_chat_error = module
.getattr("chat_completions")
.and_then(|function| function.call1(("model", &invalid_messages)))
.expect_err("sync chat should reject a non-list messages value");
let async_chat_error = module
.getattr("achat_completions")
.and_then(|function| function.call1(("model", &invalid_messages)))
.expect_err("async chat should reject a non-list messages value");
assert_eq!(
sync_chat_error.to_string(),
"ValueError: messages must be a list"
);
assert_eq!(async_chat_error.to_string(), sync_chat_error.to_string());
let invalid_body = PyList::empty(py);
let sync_messages_error = module
.getattr("messages")
.and_then(|function| function.call1(("model", &invalid_body)))
.expect_err("sync Messages should reject a non-dict body");
let async_messages_error = module
.getattr("amessages")
.and_then(|function| function.call1(("model", &invalid_body)))
.expect_err("async Messages should reject a non-dict body");
assert_eq!(
sync_messages_error.to_string(),
"ValueError: body must be a dict"
);
assert_eq!(
async_messages_error.to_string(),
sync_messages_error.to_string()
);
let invalid_headers = PyList::empty(py);
let kwargs = PyDict::new(py);
kwargs
.set_item("extra_headers", &invalid_headers)
.expect("kwargs should accept extra_headers");
let document = PyDict::new(py);
for (sync_name, async_name) in [("ocr", "aocr"), ("transcription", "atranscription")] {
let sync_error = module
.getattr(sync_name)
.and_then(|function| function.call(("model", &document), Some(&kwargs)))
.expect_err("sync route should reject non-dict extra_headers");
let async_error = module
.getattr(async_name)
.and_then(|function| function.call(("model", &document), Some(&kwargs)))
.expect_err("async route should reject non-dict extra_headers");
assert_eq!(
sync_error.to_string(),
"ValueError: extra_headers must be a dict"
);
assert_eq!(async_error.to_string(), sync_error.to_string());
}
});
}
#[test]
fn route_input_validation_preserves_left_to_right_order() {
Python::initialize();
Python::attach(|py| {
let module = PyModule::new(py, "routes").expect("module should be created");
crate::routes::register(&module).expect("routes should register");
let invalid = PyList::empty(py);
let chat_kwargs = PyDict::new(py);
chat_kwargs
.set_item("optional_params", &invalid)
.expect("kwargs should accept optional_params");
chat_kwargs
.set_item("extra_headers", &invalid)
.expect("kwargs should accept extra_headers");
let invalid_messages = PyDict::new(py);
let error = module
.getattr("chat_completions")
.and_then(|function| {
function.call(("model", &invalid_messages), Some(&chat_kwargs))
})
.expect_err("messages should be validated first");
assert_eq!(error.to_string(), "ValueError: messages must be a list");
let valid_messages = PyList::empty(py);
let error = module
.getattr("chat_completions")
.and_then(|function| function.call(("model", &valid_messages), Some(&chat_kwargs)))
.expect_err("optional_params should be validated before headers");
assert_eq!(
error.to_string(),
"ValueError: optional_params must be a dict"
);
let headers_kwargs = PyDict::new(py);
headers_kwargs
.set_item("extra_headers", &invalid)
.expect("kwargs should accept extra_headers");
let invalid_body = PyList::empty(py);
let error = module
.getattr("messages")
.and_then(|function| function.call(("model", &invalid_body), Some(&headers_kwargs)))
.expect_err("body should be validated before headers");
assert_eq!(error.to_string(), "ValueError: body must be a dict");
let invalid_payload =
PyModule::new(py, "invalid_payload").expect("invalid payload should be created");
for name in ["ocr", "transcription"] {
let error = module
.getattr(name)
.and_then(|function| {
function.call(("model", &invalid_payload), Some(&headers_kwargs))
})
.expect_err("payload should be validated before headers");
assert!(!error.to_string().contains("extra_headers"));
}
});
}
#[test]
fn generated_routes_execute_sync_and_async_contracts() {
Python::initialize();
Python::attach(|py| {
let module = PyModule::new(py, "synthetic").expect("module should be created");
synthetic::register(&module).expect("routes should register");
let sync_value: String = module
.getattr("echo")
.and_then(|function| function.call1(("sync",)))
.and_then(|value| value.extract())
.expect("sync route should return its value");
assert_eq!(sync_value, "sync");
let sync_error = module
.getattr("echo")
.and_then(|function| function.call1(("error",)))
.expect_err("sync route should map its error");
assert!(sync_error.is_instance_of::<PyLookupError>(py));
assert_eq!(
sync_error.to_string(),
"LookupError: invalid request: synthetic error"
);
let locals = PyDict::new(py);
locals
.set_item("routes", &module)
.expect("module should enter Python locals");
let code = CString::new(
r#"
import asyncio
async def exercise():
assert await routes.aecho("async") == "async"
try:
await routes.aecho("error")
except LookupError as error:
assert str(error) == "invalid request: synthetic error"
else:
raise AssertionError("mapped error was not raised")
try:
await routes.aecho("panic")
except BaseException as error:
assert type(error).__name__ == "PanicException"
assert str(error) == "synthetic panic"
else:
raise AssertionError("panic was not raised")
try:
await routes.aecho("map_panic")
except BaseException as error:
assert type(error).__name__ == "PanicException"
assert str(error) == "synthetic mapper panic"
else:
raise AssertionError("mapper panic was not raised")
task = asyncio.ensure_future(routes.aecho("pending"))
await asyncio.sleep(0)
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
else:
raise AssertionError("cancelled route completed")
for _ in range(100):
if routes.future_dropped():
break
await asyncio.sleep(0.001)
assert routes.future_dropped()
asyncio.run(exercise())
"#,
)
.expect("Python source should not contain null bytes");
py.run(&code, Some(&locals), Some(&locals))
.expect("async route contract should hold");
});
}
#[test]
fn route_registration_rejects_duplicate_python_names() {
Python::initialize();
Python::attach(|py| {
let module = PyModule::new(py, "synthetic").expect("module should be created");
synthetic::register(&module).expect("first registration should succeed");
let error = synthetic::register(&module)
.expect_err("duplicate registration should be rejected");
assert_eq!(
error.to_string(),
"RuntimeError: duplicate native route: future_dropped"
);
});
}
}

View file

@ -0,0 +1,65 @@
use litellm_core::Error;
use litellm_core::messages::messages as run_messages;
use litellm_core::messages::types::{AnthropicMessagesResponse, MessagesRequest};
use pyo3::prelude::*;
use serde_json::Value;
use std::future::Future;
use crate::errors::core_error_to_pyerr;
use crate::marshal::{RouteOptions, RouteOptionsInputs, required_value};
fn prepare_messages(
inputs: MessagesInputs,
) -> PyResult<impl Future<Output = Result<AnthropicMessagesResponse, Error>> + Send + 'static> {
let body = required_value("body", inputs.body, Value::is_object, "dict")?;
let options = RouteOptions::from_python(RouteOptionsInputs {
model: inputs.model,
api_key: inputs.api_key,
api_base: inputs.api_base,
custom_llm_provider: inputs.custom_llm_provider,
extra_headers: inputs.extra_headers,
timeout_seconds: inputs.timeout_seconds,
})?;
Ok(async move {
let RouteOptions {
model,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout,
} = options;
run_messages(MessagesRequest {
model: &model,
body,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
timeout,
})
.await
})
}
bridge_route! {
sync = messages,
asynchronous = amessages,
inputs = MessagesInputs,
required = {
model: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
body: Value,
},
optional = {
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
extra_headers: Option<Value>,
timeout_seconds: Option<f64>,
},
prepare = prepare_messages,
errors = core_error_to_pyerr,
}

View file

@ -0,0 +1,16 @@
use pyo3::prelude::*;
#[macro_use]
mod definition;
mod audio_transcription;
mod chat_completions;
mod messages;
mod ocr;
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
ocr::register(module)?;
audio_transcription::register(module)?;
messages::register(module)?;
chat_completions::register(module)
}

View file

@ -0,0 +1,73 @@
use litellm_core::Error;
use std::future::Future;
use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr};
use pyo3::prelude::*;
use serde_json::Value;
use crate::errors::core_error_to_pyerr;
use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty};
fn prepare_ocr(
inputs: OcrInputs,
) -> PyResult<impl Future<Output = Result<Value, Error>> + Send + 'static> {
let document = inputs.document;
let options = RouteOptions::from_python(RouteOptionsInputs {
model: inputs.model,
api_key: inputs.api_key,
api_base: inputs.api_base,
custom_llm_provider: inputs.custom_llm_provider,
extra_headers: inputs.extra_headers,
timeout_seconds: inputs.timeout_seconds,
})?;
let optional_params = object_or_empty("optional_params", inputs.optional_params)?;
Ok(async move {
let RouteOptions {
model,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout,
} = options;
run_ocr(OcrRequest {
model: &model,
document,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
optional_params,
timeout,
callbacks: Vec::new(),
guardrails: Vec::new(),
request_metadata: Default::default(),
litellm_call_id: None,
})
.await
})
}
bridge_route! {
sync = ocr,
asynchronous = aocr,
inputs = OcrInputs,
required = {
model: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
document: Value,
},
optional = {
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
extra_headers: Option<Value>,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
optional_params: Option<Value>,
timeout_seconds: Option<f64>,
},
prepare = prepare_ocr,
errors = core_error_to_pyerr,
}

View file

@ -0,0 +1,423 @@
use std::future::Future;
use std::panic::AssertUnwindSafe;
use std::time::Duration;
use futures_util::FutureExt;
use litellm_core::error::Error;
use litellm_python_interop::{Pythonized, panic_to_pyerr, release_gil};
use pyo3::exceptions::PyRuntimeError;
use pyo3::prelude::*;
use serde::Serialize;
use tokio::runtime::{Handle, Runtime};
use tokio::time::{self, MissedTickBehavior};
pub(super) fn run_sync<T, F>(
py: Python<'_>,
future: F,
map_error: fn(Error) -> PyErr,
) -> PyResult<Py<PyAny>>
where
T: Serialize + Send + 'static,
F: Future<Output = Result<T, Error>> + Send + 'static,
{
run_sync_on(
py,
pyo3_async_runtimes::tokio::get_runtime(),
future,
map_error,
)
}
fn run_sync_on<T, F>(
py: Python<'_>,
runtime: &Runtime,
future: F,
map_error: fn(Error) -> PyErr,
) -> PyResult<Py<PyAny>>
where
T: Serialize + Send + 'static,
F: Future<Output = Result<T, Error>> + Send + 'static,
{
if Handle::try_current().is_ok() {
return Err(PyRuntimeError::new_err(
"synchronous native routes cannot run from a Tokio context; use the async route",
));
}
let result = release_gil(py, move || runtime.block_on(wait_for_sync_result(future)))?;
let result = map_core_result(result, map_error)?;
Pythonized(result).into_pyobject(py).map(Bound::unbind)
}
pub(super) fn run_async<T, F>(
py: Python<'_>,
future: F,
map_error: fn(Error) -> PyErr,
) -> PyResult<Bound<'_, PyAny>>
where
T: Serialize + Send + 'static,
F: Future<Output = Result<T, Error>> + Send + 'static,
{
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let result = catch_route_panic(future).await?;
let result = map_core_result(result, map_error)?;
Ok(Pythonized(result))
})
}
fn map_core_result<T>(result: Result<T, Error>, map_error: fn(Error) -> PyErr) -> PyResult<T> {
match result {
Ok(value) => Ok(value),
Err(error) => Err(
std::panic::catch_unwind(AssertUnwindSafe(|| map_error(error)))
.map_err(panic_to_pyerr)?,
),
}
}
async fn catch_route_panic<T, F>(future: F) -> PyResult<Result<T, Error>>
where
F: Future<Output = Result<T, Error>>,
{
AssertUnwindSafe(future)
.catch_unwind()
.await
.map_err(panic_to_pyerr)
}
async fn wait_for_sync_result<T, F>(future: F) -> PyResult<Result<T, Error>>
where
F: Future<Output = Result<T, Error>>,
{
let future = catch_route_panic(future);
tokio::pin!(future);
let signal_interval = Duration::from_millis(50);
let mut signal_checks =
time::interval_at(time::Instant::now() + signal_interval, signal_interval);
signal_checks.set_missed_tick_behavior(MissedTickBehavior::Delay);
loop {
tokio::select! {
result = &mut future => return result,
_ = signal_checks.tick() => Python::attach(|py| py.check_signals())?,
}
}
}
#[cfg(test)]
mod tests {
use std::ffi::CString;
use std::future::poll_fn;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, mpsc};
use std::task::Poll;
use std::thread;
use std::time::Instant;
use pyo3::panic::PanicException;
use pyo3::types::{PyDict, PyModule};
use serde::Serializer;
use tokio::runtime::Builder;
use super::*;
fn runtime_error(error: Error) -> PyErr {
PyRuntimeError::new_err(error.to_string())
}
fn panicking_error_mapper(_error: Error) -> PyErr {
panic!("error mapper panicked")
}
struct PanickingOutput;
static ASYNC_PROBE_COMPLETED: AtomicUsize = AtomicUsize::new(0);
impl Serialize for PanickingOutput {
fn serialize<S>(&self, _serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
panic!("serializer panicked")
}
}
#[pyfunction]
fn async_serialization_panic(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
run_async(py, async { Ok(PanickingOutput) }, runtime_error)
}
#[pyfunction]
fn async_runtime_probe(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
run_async(
py,
async {
ASYNC_PROBE_COMPLETED.fetch_add(1, Ordering::SeqCst);
Ok(true)
},
runtime_error,
)
}
#[pyfunction]
fn runtime_worker_count() -> usize {
pyo3_async_runtimes::tokio::get_runtime()
.metrics()
.num_workers()
}
#[pyfunction]
fn runtime_is_responsive(_py: Python<'_>, expected_completions: usize) -> bool {
let completion_deadline = Instant::now() + Duration::from_secs(2);
while ASYNC_PROBE_COMPLETED.load(Ordering::SeqCst) < expected_completions {
if Instant::now() >= completion_deadline {
return false;
}
thread::sleep(Duration::from_millis(1));
}
let (heartbeat_tx, heartbeat_rx) = mpsc::sync_channel(1);
pyo3_async_runtimes::tokio::get_runtime().spawn(async move {
let _ = heartbeat_tx.send(());
});
heartbeat_rx.recv_timeout(Duration::from_secs(2)).is_ok()
}
fn extract_bool(py: Python<'_>, result: PyResult<Py<PyAny>>) -> bool {
result
.expect("route should complete")
.bind(py)
.extract()
.expect("result should convert")
}
#[test]
fn sync_runner_polls_future_on_the_caller_thread() {
Python::initialize();
Python::attach(|py| {
let caller_thread = std::thread::current().id();
let result = run_sync(
py,
async move { Ok(std::thread::current().id() == caller_thread) },
runtime_error,
);
assert!(extract_bool(py, result));
});
}
#[test]
fn sync_runner_releases_gil_while_waiting() {
Python::initialize();
Python::attach(|py| {
let result = run_sync(
py,
async {
let gil_acquired = tokio::time::timeout(
Duration::from_secs(2),
tokio::task::spawn_blocking(|| Python::attach(|_| true)),
)
.await;
Ok(matches!(gil_acquired, Ok(Ok(true))))
},
runtime_error,
);
assert!(extract_bool(py, result));
});
}
#[test]
fn sync_runner_rejects_calls_from_a_tokio_context() {
Python::initialize();
let runtime = Builder::new_current_thread()
.enable_all()
.build()
.expect("runtime should build");
let error = runtime.block_on(async {
Python::attach(|py| {
run_sync::<bool, _>(py, async { Ok(true) }, runtime_error)
.expect_err("sync route should reject a nested Tokio runtime")
})
});
assert_eq!(
error.to_string(),
"RuntimeError: synchronous native routes cannot run from a Tokio context; use the async route"
);
}
#[test]
fn sync_runner_can_drive_a_current_thread_runtime() {
Python::initialize();
let runtime = Builder::new_current_thread()
.enable_all()
.build()
.expect("runtime should build");
Python::attach(|py| {
let result = run_sync_on(
py,
&runtime,
async {
tokio::task::yield_now().await;
Ok(true)
},
runtime_error,
);
assert!(extract_bool(py, result));
});
}
#[test]
fn sync_runner_maps_a_panicked_future() {
Python::initialize();
Python::attach(|py| {
let error = run_sync::<bool, _>(
py,
poll_fn(|_| -> Poll<Result<bool, Error>> { panic!("route future panicked") }),
runtime_error,
)
.expect_err("panicked route should become a Python exception");
assert!(error.is_instance_of::<PanicException>(py));
assert_eq!(error.to_string(), "PanicException: route future panicked");
});
}
#[test]
fn sync_runner_maps_a_panicked_error_mapper() {
Python::initialize();
Python::attach(|py| {
let error = run_sync::<bool, _>(
py,
async { Err(Error::InvalidRequest("invalid".to_string())) },
panicking_error_mapper,
)
.expect_err("panicked mapper should become a Python exception");
assert!(error.is_instance_of::<PanicException>(py));
assert_eq!(error.to_string(), "PanicException: error mapper panicked");
});
}
#[test]
fn sync_runner_surfaces_serializer_panics() {
Python::initialize();
Python::attach(|py| {
let error = run_sync(py, async { Ok(PanickingOutput) }, runtime_error)
.expect_err("serializer panic should become a Python exception");
assert!(error.is_instance_of::<PanicException>(py));
assert_eq!(error.to_string(), "PanicException: serializer panicked");
});
}
#[test]
fn sync_runner_supports_concurrent_callers_on_the_shared_runtime() {
Python::initialize();
let barrier = Arc::new(tokio::sync::Barrier::new(2));
let callers: Vec<_> = (0..2)
.map(|_| {
let barrier = Arc::clone(&barrier);
thread::spawn(move || {
Python::attach(|py| {
extract_bool(
py,
run_sync(
py,
async move {
Ok(tokio::time::timeout(Duration::from_secs(2), barrier.wait())
.await
.is_ok())
},
runtime_error,
),
)
})
})
})
.collect();
let results: Vec<_> = callers
.into_iter()
.map(|caller| caller.join().expect("caller should not panic"))
.collect();
assert_eq!(results, vec![true, true]);
}
#[test]
fn async_runner_surfaces_serializer_panics() {
Python::initialize();
Python::attach(|py| {
let module = PyModule::new(py, "runtime").expect("module should be created");
module
.add_function(
wrap_pyfunction!(async_serialization_panic, &module)
.expect("function should wrap"),
)
.expect("function should register");
let locals = PyDict::new(py);
locals
.set_item("runtime", &module)
.expect("module should enter Python locals");
let code = CString::new(
r#"
import asyncio
async def exercise():
try:
await runtime.async_serialization_panic()
except BaseException as error:
assert type(error).__name__ == "PanicException"
assert str(error) == "serializer panicked"
else:
raise AssertionError("serializer panic was not raised")
asyncio.run(exercise())
"#,
)
.expect("Python source should not contain null bytes");
py.run(&code, Some(&locals), Some(&locals))
.expect("serializer panic should reach the Python awaiter");
});
}
#[test]
fn async_result_delivery_does_not_stall_tokio_workers() {
Python::initialize();
ASYNC_PROBE_COMPLETED.store(0, Ordering::SeqCst);
Python::attach(|py| {
let module = PyModule::new(py, "runtime").expect("module should be created");
for function in [
wrap_pyfunction!(async_runtime_probe, &module).expect("function should wrap"),
wrap_pyfunction!(runtime_worker_count, &module).expect("function should wrap"),
wrap_pyfunction!(runtime_is_responsive, &module).expect("function should wrap"),
] {
module
.add_function(function)
.expect("function should register");
}
let locals = PyDict::new(py);
locals
.set_item("runtime", &module)
.expect("module should enter Python locals");
let code = CString::new(
r#"
import asyncio
async def exercise():
worker_count = runtime.runtime_worker_count()
awaitables = [runtime.async_runtime_probe() for _ in range(worker_count)]
assert runtime.runtime_is_responsive(worker_count)
assert await asyncio.gather(*awaitables) == [True] * worker_count
asyncio.run(exercise())
"#,
)
.expect("Python source should not contain null bytes");
py.run(&code, Some(&locals), Some(&locals))
.expect("result delivery should leave Tokio workers responsive");
});
}
}

View file

@ -2,4 +2,4 @@ mod gil;
mod marshal;
pub use gil::{release_count, release_gil};
pub use marshal::{from_py, to_py};
pub use marshal::{Pythonized, from_py, panic_to_pyerr, to_py};

View file

@ -1,4 +1,8 @@
use std::any::Any;
use std::panic::{AssertUnwindSafe, catch_unwind};
use pyo3::exceptions::PyValueError;
use pyo3::panic::PanicException;
use pyo3::prelude::*;
use serde::Serialize;
use serde::de::DeserializeOwned;
@ -18,3 +22,71 @@ where
.map(Bound::unbind)
.map_err(|error| PyValueError::new_err(error.to_string()))
}
pub struct Pythonized<T>(pub T);
impl<'py, T> IntoPyObject<'py> for Pythonized<T>
where
T: Serialize,
{
type Target = PyAny;
type Output = Bound<'py, PyAny>;
type Error = PyErr;
fn into_pyobject(self, py: Python<'py>) -> PyResult<Self::Output> {
catch_unwind(AssertUnwindSafe(|| pythonize::pythonize(py, &self.0)))
.map_err(panic_to_pyerr)?
.map_err(|error| PyValueError::new_err(error.to_string()))
}
}
pub fn panic_to_pyerr(payload: Box<dyn Any + Send>) -> PyErr {
let message = payload
.downcast_ref::<String>()
.map(String::as_str)
.or_else(|| payload.downcast_ref::<&str>().copied())
.unwrap_or("panic from Rust code");
PanicException::new_err(message.to_string())
}
#[cfg(test)]
mod tests {
use serde::Serializer;
use super::*;
struct PanickingSerializer;
impl Serialize for PanickingSerializer {
fn serialize<S>(&self, _serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
panic!("serializer panicked")
}
}
#[test]
fn pythonized_converts_on_the_attached_thread() {
Python::initialize();
Python::attach(|py| {
let value: Vec<i32> = Pythonized(vec![1, 2, 3])
.into_pyobject(py)
.and_then(|value| value.extract())
.expect("value should convert");
assert_eq!(value, vec![1, 2, 3]);
});
}
#[test]
fn pythonized_maps_serializer_panics_to_a_base_exception() {
Python::initialize();
Python::attach(|py| {
let error = Pythonized(PanickingSerializer)
.into_pyobject(py)
.expect_err("serializer panic should become a Python exception");
assert!(error.is_instance_of::<PanicException>(py));
assert_eq!(error.to_string(), "PanicException: serializer panicked");
});
}
}

View file

@ -1417,7 +1417,7 @@ from .skills.main import (
)
from .containers.main import *
from .ocr.main import *
from .rust_bridge.ocr import use_litellm_rust
from .rust_bridge import use_litellm_rust
from .rag.main import *
from .sandbox.main import *
from .search.main import *

View file

@ -1567,6 +1567,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
optional_params.pop("thinking", None)
else:
optional_params["thinking"] = value
AnthropicModelInfo.translate_legacy_thinking_for_adaptive_model(
model=model, optional_params=optional_params, custom_llm_provider=self._resolved_provider
)
elif param == "reasoning_effort":
# Accept both string ("low") and dict ({"effort": "low",
# "summary": "concise"}). The Responses->Chat parser keeps the

View file

@ -14,7 +14,12 @@ import httpx
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
import litellm
from litellm.constants import DEFAULT_MODEL_CREATED_AT_TIME
from litellm.constants import (
DEFAULT_MODEL_CREATED_AT_TIME,
DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET,
)
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_file_ids_from_messages,
)
@ -648,6 +653,51 @@ class AnthropicModelInfo(BaseLLMModelInfo):
)
optional_params.pop("thinking", None)
@staticmethod
def translate_legacy_thinking_for_adaptive_model(
model: str,
optional_params: MutableMapping[str, object], # mutable-ok: in-place out-param like the sibling helpers
custom_llm_provider: str,
) -> None:
"""Translate legacy ``thinking.type=enabled`` to adaptive for the
adaptive-thinking models that reject it (4.7+ and the 5 families).
Models flagged ``supports_legacy_thinking`` (the 4.6 family) accept the
legacy shape natively, so it is forwarded verbatim and the caller's
``budget_tokens`` cap keeps applying. Caller-provided
``output_config.effort`` is never overridden.
"""
if not AnthropicModelInfo._is_adaptive_thinking_model(model, custom_llm_provider):
return
if AnthropicModelInfo._supports_legacy_thinking(model, custom_llm_provider):
return
thinking: Final = optional_params.get("thinking")
if not isinstance(thinking, dict) or thinking.get("type") != "enabled":
return
effort: Final = AnthropicModelInfo._legacy_budget_to_effort(
model=model,
budget_tokens=int(thinking.get("budget_tokens") or 0),
custom_llm_provider=custom_llm_provider,
)
existing_output_config: Final = optional_params.get("output_config")
optional_params["thinking"] = {"type": "adaptive"}
optional_params["output_config"] = {
"effort": effort,
**(existing_output_config if isinstance(existing_output_config, dict) else MappingProxyType({})),
}
@staticmethod
def _legacy_budget_to_effort(model: str, budget_tokens: int, custom_llm_provider: str) -> str:
if budget_tokens >= DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET and (
AnthropicModelInfo._supports_model_capability(model, "supports_xhigh_reasoning_effort", custom_llm_provider)
):
return "xhigh"
if budget_tokens >= DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET:
return "high"
if budget_tokens >= DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET:
return "medium"
return "low"
def is_effort_used(
self,
optional_params: dict | None,

View file

@ -3,11 +3,6 @@ from typing import Any, ClassVar, Final
import httpx
from litellm.constants import (
DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET,
)
from litellm.exceptions import AuthenticationError
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.litellm_logging import verbose_logger
@ -485,46 +480,6 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
existing_output_config.setdefault("effort", mapped_effort)
optional_params["output_config"] = existing_output_config
@staticmethod
def _translate_legacy_thinking_for_adaptive_model(
model: str, optional_params: dict, custom_llm_provider: str
) -> None:
"""Translate legacy ``thinking.type=enabled`` to adaptive for the
adaptive-thinking models that reject it (4.7+ and the 5 families).
Models flagged ``supports_legacy_thinking`` (the 4.6 family) accept the
legacy shape natively, so it is forwarded verbatim and the caller's
``budget_tokens`` cap keeps applying. Caller-provided
``output_config.effort`` is never overridden.
"""
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
if not AnthropicModelInfo._is_adaptive_thinking_model(model, custom_llm_provider):
return
if AnthropicModelInfo._supports_legacy_thinking(model, custom_llm_provider):
return
thinking: Final = optional_params.get("thinking")
if not isinstance(thinking, dict) or thinking.get("type") != "enabled":
return
budget: Final = int(thinking.get("budget_tokens") or 0)
if budget >= DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET and (
AnthropicConfig._supports_effort_level(model, "xhigh", custom_llm_provider)
):
effort = "xhigh"
elif budget >= DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET:
effort = "high"
elif budget >= DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET:
effort = "medium"
else:
effort = "low"
optional_params["thinking"] = {"type": "adaptive"}
existing_output_config = optional_params.get("output_config")
if not isinstance(existing_output_config, dict):
existing_output_config = {}
existing_output_config.setdefault("effort", effort)
optional_params["output_config"] = existing_output_config
@staticmethod
def _translate_adaptive_effort_for_non_adaptive_model(
model: str, optional_params: dict, max_tokens: int | None, custom_llm_provider: str
@ -691,7 +646,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
custom_llm_provider=self._resolved_provider,
)
self._translate_legacy_thinking_for_adaptive_model(
AnthropicModelInfo.translate_legacy_thinking_for_adaptive_model(
model=model,
optional_params=anthropic_messages_optional_request_params,
custom_llm_provider=self._resolved_provider,

View file

@ -17,13 +17,14 @@ def _promote_extra_body_to_optional_params(optional_params: dict) -> None:
``output_config`` get auto-routed into ``extra_body`` by
``add_provider_specific_params_to_optional_params``. For the Azure→Anthropic
route those keys must reach the request body and be validated, so promote
them. ``setdefault`` keeps explicit top-level values authoritative.
them. The caller's values overwrite mapped top-level duplicates, matching
the native ``anthropic`` provider, where the same passthrough lands on
top-level ``optional_params`` after mapping.
"""
extra_body: Final = optional_params.get("extra_body")
if not isinstance(extra_body, dict) or not extra_body:
return
for k, v in extra_body.items():
optional_params.setdefault(k, v)
optional_params.update(extra_body)
optional_params.pop("extra_body", None)

View file

@ -943,6 +943,9 @@ class AmazonConverseConfig(BaseConfig):
litellm.verbose_logger.warning(DROP_UNSUPPORTED_ADAPTIVE_THINKING_WARNING, model)
else:
optional_params["thinking"] = value
AnthropicModelInfo.translate_legacy_thinking_for_adaptive_model(
model=model, optional_params=optional_params, custom_llm_provider="bedrock"
)
elif param == "reasoning_effort" and isinstance(value, str):
self._handle_reasoning_effort_parameter(
model=model, reasoning_effort=value, optional_params=optional_params

View file

@ -107,6 +107,10 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
# Restore original model name
model = original_model
AnthropicModelInfo.translate_legacy_thinking_for_adaptive_model(
model=original_model, optional_params=optional_params, custom_llm_provider="bedrock"
)
# The stub model hides the original model from the parent's forced-tool-use backstop
response_format_tool_choice: Final = optional_params.get("tool_choice")
if (

View file

@ -1,7 +1,6 @@
import asyncio
import inspect
import json
import os
import ssl
from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping, Sequence
from contextlib import asynccontextmanager
@ -173,7 +172,11 @@ def _rust_responses_websocket_enabled(
custom_llm_provider: str | None,
litellm_params: GenericLiteLLMParams,
) -> bool:
return custom_llm_provider == "openai" and litellm_params.get("rust") is True
from litellm.rust_bridge.configuration import rust_enabled
raw_request_override: Final = litellm_params.get("rust")
request_override: Final = raw_request_override if isinstance(raw_request_override, bool) else None
return custom_llm_provider == "openai" and rust_enabled(request_override=request_override)
from .http_handler import get_shared_realtime_ssl_context
@ -2427,10 +2430,6 @@ class BaseLLMHTTPHandler:
"anthropic_messages",
)
@staticmethod
def _rust_env_enabled() -> bool:
return os.getenv("LITELLM_RUST", "").strip().lower() in {"1", "true", "yes", "on"}
@staticmethod
async def _maybe_rust_anthropic_messages(
*,
@ -2446,7 +2445,11 @@ class BaseLLMHTTPHandler:
) -> AnthropicMessagesResponse | None:
if custom_llm_provider not in ("azure_ai", "anthropic"):
return None
if litellm_params.get("rust") is not True and not BaseLLMHTTPHandler._rust_env_enabled():
from litellm.rust_bridge.configuration import rust_enabled
raw_request_override: Final = litellm_params.get("rust")
request_override: Final = raw_request_override if isinstance(raw_request_override, bool) else None
if not rust_enabled(request_override=request_override):
return None
if has_agentic_hook:
return None

View file

@ -330,6 +330,10 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
) -> dict:
is_thinking_enabled: Final = self.is_thinking_enabled(non_default_params)
mapped_params: Final = super().map_openai_params(non_default_params, optional_params, model, drop_params)
if "claude" in model:
AnthropicConfig.translate_legacy_thinking_for_adaptive_model(
model=model, optional_params=mapped_params, custom_llm_provider="databricks"
)
if "tools" in mapped_params:
mapped_params["tools"] = self._map_openai_to_dbrx_tool(model=model, tools=mapped_params["tools"])
if "max_completion_tokens" in non_default_params and replace_max_completion_tokens_with_max_tokens:

View file

@ -177,6 +177,10 @@ class VertexAIAnthropicConfig(AnthropicConfig):
# Restore original model name for any other processing
model = original_model
AnthropicModelInfo.translate_legacy_thinking_for_adaptive_model(
model=original_model, optional_params=optional_params, custom_llm_provider="vertex_ai"
)
return optional_params
def transform_response(

View file

@ -194,6 +194,12 @@ def _rust_ocr_supported(prepared_request: _PreparedOCRRequest) -> bool:
return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS
def _rust_ocr_enabled(prepared_request: _PreparedOCRRequest) -> bool:
raw_request_override: Final = prepared_request.litellm_params.get("rust")
request_override: Final = raw_request_override if isinstance(raw_request_override, bool) else None
return rust_ocr_bridge.rust_ocr_enabled(request_override=request_override)
def _rust_bridge_optional_params(
prepared_request: _PreparedOCRRequest,
resolve_secret: Callable[[str], str | None],
@ -422,7 +428,7 @@ async def aocr(
custom_llm_provider = prepared.custom_llm_provider
completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider})
if _rust_ocr_supported(prepared) and rust_ocr_bridge.rust_ocr_enabled():
if _rust_ocr_supported(prepared) and _rust_ocr_enabled(prepared):
from litellm.secret_managers.main import get_secret_str
rust_response: Final = await _run_rust_aocr(
@ -694,7 +700,7 @@ def ocr(
custom_llm_provider = prepared.custom_llm_provider
completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider})
if _rust_ocr_supported(prepared) and rust_ocr_bridge.rust_ocr_enabled():
if _rust_ocr_supported(prepared) and _rust_ocr_enabled(prepared):
from litellm.secret_managers.main import get_secret_str
rust_response: Final = _run_rust_ocr(

View file

@ -6,8 +6,11 @@ External callers (public IPs) only see servers with available_on_public_internet
"""
import ipaddress
import os
from collections.abc import Mapping
from dataclasses import dataclass
from typing import Any, Final
from urllib.parse import urlparse
from fastapi import Request
from pydantic import TypeAdapter, ValidationError
@ -137,7 +140,7 @@ class IPAddressUtils:
@staticmethod
def is_request_from_trusted_proxy(
request: Request,
general_settings: dict[str, Any] | None = None,
general_settings: Mapping[str, Any] | None = None,
) -> bool:
"""
Return True if X-Forwarded-* headers on this request should be trusted.
@ -190,6 +193,36 @@ class IPAddressUtils:
trusted_networks: Final = IPAddressUtils.parse_trusted_proxy_networks(trusted_ranges)
return IPAddressUtils.is_trusted_proxy(direct_ip, trusted_networks)
@staticmethod
def is_request_https(
request: Request,
general_settings: Mapping[str, Any] | None = None,
) -> bool:
"""
Whether this request's PUBLIC-facing origin is HTTPS, for deciding
whether a cookie set on the response should be marked ``Secure``.
litellm only sees a plain-HTTP hop whenever TLS terminates at a
reverse proxy, so ``request.url.scheme`` alone cannot answer this in
that deployment shape. Resolved from the first trusted signal:
1. ``PROXY_BASE_URL`` (operator-declared public origin).
2. ``X-Forwarded-Proto``, only when the request's direct peer is a
configured trusted proxy -- see ``is_request_from_trusted_proxy``.
An untrusted caller cannot spoof this header to strip Secure.
3. The request's own literal scheme (direct TLS termination, or no
reverse proxy in front of litellm).
"""
configured_base_url: Final = os.environ.get("PROXY_BASE_URL", "").strip()
if configured_base_url:
return urlparse(configured_base_url).scheme == "https"
if IPAddressUtils.is_request_from_trusted_proxy(request, general_settings=general_settings):
forwarded_proto: Final = request.headers.get("X-Forwarded-Proto")
if forwarded_proto:
return forwarded_proto.split(",")[0].strip().lower() == "https"
return request.url.scheme == "https"
@staticmethod
def extract_client_ip_from_xff_hops(
xff_header: str,

View file

@ -36,6 +36,7 @@ from pydantic import ValidationError
from litellm._logging import verbose_proxy_logger
from litellm.caching.dual_cache import DualCache
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
from litellm.proxy.management_endpoints.types import CustomOpenID, get_litellm_user_role
from litellm.proxy.utils import get_custom_url
@ -131,7 +132,7 @@ class SAMLAuthHandler:
@staticmethod
def _is_https(request: Request) -> bool:
return SAMLAuthHandler._base_url(request).startswith("https")
return IPAddressUtils.is_request_https(request)
@staticmethod
def _acs_url(request: Request) -> str:

View file

@ -92,6 +92,7 @@ from litellm.proxy.auth.auth_utils import (
has_user_setup_sso,
)
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.admin_ui_utils import (
admin_ui_disabled,
@ -1118,7 +1119,7 @@ async def google_login(
request=request,
)
if sso_redirect is not None:
_persist_return_to_cookie(sso_redirect, return_to)
_persist_return_to_cookie(sso_redirect, return_to, request)
return sso_redirect
from fastapi.responses import HTMLResponse
@ -1138,7 +1139,7 @@ async def google_login(
# helper the SSO branch uses, so /login can resume the connect flow instead of dead-ending at the
# dashboard. One implementation → the two sign-in branches cannot diverge (and the login form always
# renders, since the helper never raises on a bad return_to).
_persist_return_to_cookie(form_response, return_to)
_persist_return_to_cookie(form_response, return_to, request)
return form_response
@ -2741,6 +2742,7 @@ async def _sso_return_to_redirect(
jwt_token: str,
redis_usage_cache,
user_api_key_cache,
request: Request,
) -> RedirectResponse | None:
"""Resolve the post-SSO redirect for a ``return_to``, or None to fall through to the dashboard.
@ -2759,7 +2761,7 @@ async def _sso_return_to_redirect(
if _is_same_origin_return_path(return_to):
redirect_response = RedirectResponse(url=return_to, status_code=303)
redirect_response.set_cookie(key="token", value=jwt_token)
set_session_token_cookie(redirect_response, request, jwt_token)
redirect_response.delete_cookie("litellm_cp_return_to")
return redirect_response
@ -2782,7 +2784,25 @@ async def _sso_return_to_redirect(
return None
def _persist_return_to_cookie(response: Response, return_to: str | None) -> None:
def set_session_token_cookie(response: Response, request: Request, jwt_token: str) -> None:
"""Set the ``token`` session cookie shared by every sign-in path.
Not HttpOnly: the dashboard reads this cookie via ``document.cookie`` to
populate its own Authorization headers (see
``ui/litellm-dashboard/src/utils/cookieUtils.ts``), so marking it
HttpOnly would break login. Secure is still required whenever the public
origin is HTTPS, resolved the same trust-aware way as every other
litellm cookie."""
response.set_cookie(
key="token",
value=jwt_token,
secure=IPAddressUtils.is_request_https(request),
httponly=False,
samesite="lax",
)
def _persist_return_to_cookie(response: Response, return_to: str | None, request: Request) -> None:
"""Best-effort: persist a SAFE ``return_to`` on ``response`` as the one-shot ``litellm_cp_return_to``
cookie so ANY sign-in path — SSO / Okta / generic OR the username/password form — can resume there
afterwards. THIS is the single source of truth, called by every sign-in branch so they cannot
@ -2803,6 +2823,7 @@ def _persist_return_to_cookie(response: Response, return_to: str | None) -> None
max_age=600,
httponly=True,
samesite="lax",
secure=IPAddressUtils.is_request_https(request),
)
@ -3079,8 +3100,11 @@ class SSOAuthenticationHandler:
# incoming request is HTTP (local dev). Without
# ``Secure`` the cookie is sent over plain HTTP,
# letting a network observer read and replay the
# state value and bypass this protection.
secure_flag: Final = request is None or request.url.scheme == "https"
# state value and bypass this protection. Trust-aware:
# honors PROXY_BASE_URL / a trusted reverse proxy's
# X-Forwarded-Proto instead of only the literal scheme
# litellm sees on the wire.
secure_flag: Final = request is None or IPAddressUtils.is_request_https(request)
redirect_response.set_cookie(
key="litellm_oauth_state",
value=state_value,
@ -3628,6 +3652,7 @@ class SSOAuthenticationHandler:
jwt_token=jwt_token,
redis_usage_cache=redis_usage_cache,
user_api_key_cache=user_api_key_cache,
request=request,
)
if return_to_redirect is not None:
return return_to_redirect
@ -3636,7 +3661,7 @@ class SSOAuthenticationHandler:
litellm_dashboard_ui += "?login=success"
verbose_proxy_logger.info("Redirecting to %s", litellm_dashboard_ui)
redirect_response: Final = RedirectResponse(url=litellm_dashboard_ui, status_code=303)
redirect_response.set_cookie(key="token", value=jwt_token)
set_session_token_cookie(redirect_response, request, jwt_token)
return redirect_response
@staticmethod

View file

@ -15334,7 +15334,10 @@ async def login(request: Request):
# authorize round-trip), mirroring the SSO callback; otherwise land on the dashboard. Gated by
# _is_same_origin_return_path (strictly relative path) so it can never be an open redirect, and the
# one-shot cookie is cleared after use.
from litellm.proxy.management_endpoints.ui_sso import _sso_return_to_redirect
from litellm.proxy.management_endpoints.ui_sso import (
_sso_return_to_redirect,
set_session_token_cookie,
)
# Resume through the SAME resumer the SSO callback uses, rather than a second, narrower arm.
# _persist_return_to_cookie stores both shapes it accepts (a relative same-origin path AND a
@ -15351,6 +15354,7 @@ async def login(request: Request):
jwt_token=jwt_token,
redis_usage_cache=redis_usage_cache,
user_api_key_cache=user_api_key_cache,
request=request,
)
except Exception: # noqa: BLE001 # resuming must NEVER block a completed sign-in
# The symmetric half of _persist_return_to_cookie's "never raises" contract. The resumer
@ -15365,7 +15369,7 @@ async def login(request: Request):
# Create redirect response with cookie
redirect_response: Final = RedirectResponse(url=litellm_dashboard_ui, status_code=303)
redirect_response.set_cookie(key="token", value=jwt_token)
set_session_token_cookie(redirect_response, request, jwt_token)
if cp_return_to:
redirect_response.delete_cookie(key="litellm_cp_return_to")
return redirect_response
@ -15375,6 +15379,7 @@ async def login(request: Request):
async def login_v2(request: Request):
global premium_user, general_settings, master_key
from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object, encode_ui_session_jwt
from litellm.proxy.management_endpoints.ui_sso import set_session_token_cookie
from litellm.proxy.utils import get_custom_url
try:
@ -15409,7 +15414,7 @@ async def login_v2(request: Request):
content={"redirect_url": litellm_dashboard_ui, "token": jwt_token},
status_code=status.HTTP_200_OK,
)
json_response.set_cookie(key="token", value=jwt_token)
set_session_token_cookie(json_response, request, jwt_token)
return json_response
except Exception as e:
verbose_proxy_logger.exception("litellm.proxy.proxy_server.login_v2(): Exception occurred - %s", e)
@ -15509,6 +15514,8 @@ async def login_v3(request: Request):
@router.post("/v3/login/exchange", include_in_schema=False) # exchange single-use opaque code for JWT
async def login_v3_exchange(request: Request):
from litellm.proxy.management_endpoints.ui_sso import set_session_token_cookie
try:
if not general_settings.get("control_plane_url"):
raise ProxyException(
@ -15555,7 +15562,7 @@ async def login_v3_exchange(request: Request):
},
status_code=status.HTTP_200_OK,
)
json_response.set_cookie(key="token", value=cached_data["token"])
set_session_token_cookie(json_response, request, cached_data["token"])
return json_response
except ProxyException:
raise

View file

@ -68,6 +68,36 @@ still resolve to a deployment in `model_list`; this configuration does not creat
- abc
```
### Heuristic v2
Set `classifier_type: heuristic_v2` to classify with the bundled calibrated
success-probability model instead of the hand-written weighted scorer
```yaml
model_list:
- model_name: smart-router
litellm_params:
model: auto_router/complexity_router
complexity_router_config:
classifier_type: heuristic_v2
tiers:
SIMPLE: luna
MEDIUM: terra
COMPLEX: sol
REASONING: sol-ultra
```
No classifier model call or per-model training data is required. The classifier
uses global tier quality, request-type quality, and similar-request cohorts from
the bundled UltraFeedback artifact. It estimates success at every tier, enforces
monotonic probabilities, and returns the first tier meeting the trained 0.75
threshold. The existing complexity-router tier pool then selects and dispatches
a model from that tier
Spend logs record `routing_decision.cause: heuristic_v2`, the detected request
type, and all four predicted probabilities. Existing `classifier_type: heuristic`
configurations keep the original weighted scorer unchanged
### Renaming the tiers
`tier_labels` puts your own vocabulary on the four tiers:

File diff suppressed because it is too large Load diff

View file

@ -33,6 +33,11 @@ from litellm.litellm_core_utils.internal_call_metadata import forwarded_internal
from litellm.litellm_core_utils.prompt_templates.common_utils import request_contains_image_content
from litellm.litellm_core_utils.sensitive_data_masker import mask_credentials_in_payload
from litellm.llms.base_llm.base_utils import type_to_response_format_param
from litellm.router_strategy.adaptive_router.classifier import classify_prompt
from litellm.router_strategy.complexity_router.tier_predictor import (
TierSuccessPredictor,
resolve_tier_artifact,
)
from litellm.types.utils import (
AUTOROUTER_CLASSIFIER_CALL_ORIGIN,
ModelResponse,
@ -790,6 +795,7 @@ class ClassificationOutcome(NamedTuple):
signals: tuple[str, ...]
cause: Literal[
"heuristic_scorer",
"heuristic_v2",
"reasoning_override",
"llm_classifier",
"heuristic_first_short_circuit",
@ -978,6 +984,11 @@ class ComplexityRouter(CustomLogger):
if llm_classifier_configured
else None
)
self._tier_success_predictor: TierSuccessPredictor | None = (
TierSuccessPredictor(resolve_tier_artifact(self.config.heuristic_v2_artifact))
if self.config.classifier_type == "heuristic_v2"
else None
)
verbose_router_logger.debug("ComplexityRouter initialized for %s with tiers: %s", model_name, self.config.tiers)
@ -1350,6 +1361,8 @@ class ComplexityRouter(CustomLogger):
custom tier set, and classifier_fallback otherwise decides between the heuristic scorer and
default_model. The outcome's `cause` reports which path actually ran.
"""
if self.config.classifier_type == "heuristic_v2":
return self._classify_with_heuristic_v2(prompt)
if self.config.classifier_type == "custom":
return await self._classify_with_plugin(prompt, system_prompt, request_kwargs, raw_messages)
if self.config.classifier_type == "heuristic_first" and self.config.classifier_llm_config is not None:
@ -1359,6 +1372,24 @@ class ComplexityRouter(CustomLogger):
return ClassificationOutcome(tier=tier, score=score, signals=signals, cause=cause)
return await self._llm_classifier_outcome(prompt, system_prompt, request_kwargs, messages)
def _classify_with_heuristic_v2(self, prompt: str) -> ClassificationOutcome:
predictor: Final = self._tier_success_predictor
if predictor is None:
raise ValueError("heuristic v2 predictor is not configured")
request_type: Final = classify_prompt(prompt)
prediction: Final = predictor.predict(prompt, request_type)
tier: Final = TIER_SEVERITY_ORDER[prediction.required_tier - 1]
probability_signals: Final = tuple(
f"tier-probability:{candidate.value.lower()}={prediction.probabilities[index]:.6f}"
for index, candidate in enumerate(TIER_SEVERITY_ORDER, start=1)
)
return ClassificationOutcome(
tier=tier,
score=None,
signals=(f"request-type:{request_type.value}", *probability_signals),
cause="heuristic_v2",
)
async def _classify_heuristic_first(
self,
prompt: str,

View file

@ -14,6 +14,8 @@ from pydantic import BaseModel, ConfigDict, Field, SkipValidation, field_seriali
from litellm.types.router import AdaptiveRouterWeights, ClassifierPlugin, RoutingPlugin
from .tier_predictor import TrainedTierArtifact
class ComplexityTier(str, Enum):
"""Complexity tiers for routing decisions."""
@ -625,12 +627,19 @@ class ComplexityRouterConfig(BaseModel):
)
# Classifier strategy
classifier_type: Literal["heuristic", "llm", "custom", "heuristic_first"] = Field(
classifier_type: Literal["heuristic", "heuristic_v2", "llm", "custom", "heuristic_first"] = Field(
default="heuristic",
description=(
"Classification strategy: local regex/keyword scoring, an LLM call, a custom classifier "
"plugin, or 'heuristic_first', which scores locally and only pays for the LLM classifier "
"when the local scorer does not confidently land a cheap tier"
"Classification strategy: local regex/keyword scoring, the bundled trained four-tier heuristic, "
"an LLM call, a custom classifier plugin, or 'heuristic_first', which scores locally and only pays "
"for the LLM classifier when the local scorer does not confidently land a cheap tier"
),
)
heuristic_v2_artifact: TrainedTierArtifact | Literal["ultrafeedback"] = Field(
default="ultrafeedback",
description=(
"Success-probability artifact used by classifier_type 'heuristic_v2'. The bundled "
"UltraFeedback artifact is selected by default; an inline trained artifact may replace it"
),
)
classifier_llm_config: ClassifierLLMConfig | None = Field(
@ -1248,10 +1257,10 @@ class ComplexityRouterConfig(BaseModel):
)
if duplicated:
raise ValueError(f"tier_definitions names must be unique (case-insensitive): {', '.join(duplicated)}")
if self.classifier_type in ("heuristic", "heuristic_first"):
if self.classifier_type in ("heuristic", "heuristic_v2", "heuristic_first"):
raise ValueError(
"tier_definitions requires classifier_type 'llm' or 'custom': the heuristic scorer only "
"produces the built-in tiers"
"produces the four built-in tiers, as does heuristic_v2"
)
conflicts: Final = self._tier_definition_conflicts()
if conflicts:

View file

@ -0,0 +1,156 @@
from __future__ import annotations
import re
from collections.abc import Mapping
from dataclasses import dataclass
from pathlib import Path
from types import MappingProxyType
from typing import Final, Literal
from pydantic import BaseModel, Field, model_validator
from litellm.types.router import RequestType
class TierGlobalStatistic(BaseModel):
tier: int = Field(ge=1, le=4)
successes: float = Field(ge=0.0)
observations: float = Field(gt=0.0)
@model_validator(mode="after")
def _successes_do_not_exceed_observations(self) -> TierGlobalStatistic:
if self.successes > self.observations:
raise ValueError("successes cannot exceed observations")
return self
class TierDomainStatistic(TierGlobalStatistic):
request_type: RequestType
class TierCohortStatistic(TierGlobalStatistic):
cohort: str = Field(min_length=1)
class TierDataset(BaseModel):
name: str = Field(min_length=1)
url: str = Field(min_length=1)
license: str = Field(min_length=1)
rows: int = Field(gt=0)
success_definition: str = Field(default="quality score meets the dataset success threshold", min_length=1)
class TrainedTierArtifact(BaseModel):
schema_version: Literal[1] = 1
global_statistics: tuple[TierGlobalStatistic, ...]
domain_statistics: tuple[TierDomainStatistic, ...] = ()
cohort_statistics: tuple[TierCohortStatistic, ...] = ()
domain_prior_mass: float = Field(default=200.0, gt=0.0)
cohort_prior_mass: float = Field(default=20.0, gt=0.0)
routing_threshold: float = Field(default=0.75, ge=0.0, le=1.0)
datasets: tuple[TierDataset, ...] = ()
success_definition: str = Field(default="quality score meets the dataset success threshold", min_length=1)
split_method: str = Field(default="sha256(prompt): 70% train, 15% validation, 15% test", min_length=1)
@model_validator(mode="after")
def _statistics_are_unique(self) -> TrainedTierArtifact:
global_tiers: Final = tuple(stat.tier for stat in self.global_statistics)
if frozenset(global_tiers) != frozenset((1, 2, 3, 4)) or len(global_tiers) != 4:
raise ValueError("global statistics must contain each tier exactly once")
domain_keys: Final = tuple((stat.request_type, stat.tier) for stat in self.domain_statistics)
if len(domain_keys) != len(frozenset(domain_keys)):
raise ValueError("domain statistics must contain unique request_type and tier pairs")
cohort_keys: Final = tuple((stat.cohort, stat.tier) for stat in self.cohort_statistics)
if len(cohort_keys) != len(frozenset(cohort_keys)):
raise ValueError("cohort statistics must contain unique cohort and tier pairs")
return self
_CODE_PATTERN: Final = re.compile(
r"```|\b(def|class|function|python|javascript|typescript|sql|code)\b",
re.IGNORECASE,
)
_MATH_PATTERN: Final = re.compile(
r"\b(solve|calculate|equation|probability|theorem|proof|integral)\b|[$=]",
re.IGNORECASE,
)
_MULTIPLE_CHOICE_PATTERN: Final = re.compile(r"(?:^|\s)[A-D][.)]\s")
_TIERS: Final = (1, 2, 3, 4)
_BUILTIN_ARTIFACTS: Final = MappingProxyType({"ultrafeedback": "ultrafeedback_tiers.json"})
def resolve_tier_artifact(artifact: TrainedTierArtifact | str) -> TrainedTierArtifact:
if isinstance(artifact, TrainedTierArtifact):
return artifact
filename: Final = _BUILTIN_ARTIFACTS.get(artifact)
if filename is None:
raise ValueError(f"unknown complexity router tier artifact: {artifact}")
path: Final = Path(__file__).with_name("artifacts") / filename
return TrainedTierArtifact.model_validate_json(path.read_text())
def similarity_cohort(prompt: str, request_type: RequestType) -> str:
length: Final = len(prompt)
length_bucket: Final = (
"short" if length < 200 else "medium" if length < 800 else "long" if length < 2000 else "very_long"
)
code: Final = int(bool(_CODE_PATTERN.search(prompt)))
math: Final = int(bool(_MATH_PATTERN.search(prompt)))
multiple_choice: Final = int(bool(_MULTIPLE_CHOICE_PATTERN.search(prompt)))
non_ascii: Final = int(sum(ord(character) > 127 for character in prompt) / max(1, length) > 0.1)
return f"{request_type.value}|{length_bucket}|code={code}|math={math}|mc={multiple_choice}|intl={non_ascii}"
@dataclass(frozen=True, slots=True)
class TierPrediction:
probabilities: Mapping[int, float]
required_tier: int
class TierSuccessPredictor:
def __init__(self, artifact: TrainedTierArtifact) -> None:
self._artifact = artifact
self._global: Mapping[int, TierGlobalStatistic] = MappingProxyType(
{stat.tier: stat for stat in artifact.global_statistics}
)
self._domain: Mapping[tuple[RequestType, int], TierDomainStatistic] = MappingProxyType(
{(stat.request_type, stat.tier): stat for stat in artifact.domain_statistics}
)
self._cohort: Mapping[tuple[str, int], TierCohortStatistic] = MappingProxyType(
{(stat.cohort, stat.tier): stat for stat in artifact.cohort_statistics}
)
@property
def routing_threshold(self) -> float:
return self._artifact.routing_threshold
def predict(self, prompt: str, request_type: RequestType) -> TierPrediction:
cohort: Final = similarity_cohort(prompt, request_type)
raw: Final = tuple(self._probability(tier, request_type, cohort) for tier in _TIERS)
monotonic: Final = tuple(max(raw[:index]) for index in range(1, len(raw) + 1))
probabilities: Final[Mapping[int, float]] = MappingProxyType(
{int(tier): probability for tier, probability in zip(_TIERS, monotonic)}
)
required_tier: Final = next(
(tier for tier in _TIERS if probabilities[tier] >= self._artifact.routing_threshold),
4,
)
return TierPrediction(probabilities=probabilities, required_tier=required_tier)
def _probability(self, tier: int, request_type: RequestType, cohort: str) -> float:
global_stat: Final = self._global[tier]
global_mean: Final = (global_stat.successes + 1.0) / (global_stat.observations + 2.0)
domain_stat: Final = self._domain.get((request_type, tier))
domain_mean: Final = self._posterior_mean(domain_stat, self._artifact.domain_prior_mass, global_mean)
cohort_stat: Final = self._cohort.get((cohort, tier))
return self._posterior_mean(cohort_stat, self._artifact.cohort_prior_mass, domain_mean)
@staticmethod
def _posterior_mean(
statistic: TierGlobalStatistic | None,
prior_mass: float,
prior_mean: float,
) -> float:
if statistic is None:
return prior_mean
return (statistic.successes + prior_mass * prior_mean) / (statistic.observations + prior_mass)

View file

@ -1,9 +1,9 @@
"""LiteLLM Rust bridge package."""
from litellm.rust_bridge.configuration import use_litellm_rust
from litellm.rust_bridge.loader import (
get_native_bridge,
native_bridge_available,
)
from litellm.rust_bridge.ocr import use_litellm_rust
__all__ = ["get_native_bridge", "native_bridge_available", "use_litellm_rust"]

View file

@ -0,0 +1,49 @@
from __future__ import annotations
from collections.abc import Callable
from typing import Final, Generic, TypeVar
from litellm.rust_bridge.loader import get_native_bridge
BindingT = TypeVar("BindingT")
class _Unset:
pass
_UNSET: Final = _Unset()
class NativeBinding(Generic[BindingT]):
"""Resolve one native attribute with an explicit, resettable test override."""
def __init__(self, attribute: str, *, validate: Callable[[object], BindingT | None]) -> None:
self._attribute: Final = attribute
self._validate: Final = validate
self._override: BindingT | None | _Unset = _UNSET
def load(self) -> BindingT | None:
if not isinstance(self._override, _Unset):
return self._override
native: Final = get_native_bridge()
if native is None:
return None
return self._validate(getattr(native, self._attribute, None))
def override(self, value: BindingT | None) -> None:
self._override = value
def reset(self) -> None:
self._override = _UNSET
def native_exception_types() -> tuple[type[BaseException], type[BaseException]] | None:
native: Final = get_native_bridge()
if native is None:
return None
declined: Final = getattr(native, "RustBridgeDeclined", None)
upstream: Final = getattr(native, "RustUpstreamError", None)
if not isinstance(declined, type) or not isinstance(upstream, type):
return None
return declined, upstream

View file

@ -13,7 +13,6 @@ retrying it there would bill the customer for the same work twice.
from __future__ import annotations
import json
import os
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import dataclass
from typing import TYPE_CHECKING, Final, Protocol
@ -27,6 +26,7 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
convert_to_model_response_object,
)
from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned
from litellm.rust_bridge.configuration import rust_enabled
from litellm.rust_bridge.loader import get_native_bridge
from litellm.rust_bridge.timeouts import timeout_to_seconds
from litellm.types.utils import ModelResponse
@ -44,8 +44,6 @@ _LITELLM_METADATA_ADAPTER: Final = TypeAdapter(Mapping[str, object])
RUST_RESPONSE_HEADER: Final = "x-litellm-rust"
_TRUTHY_ENV_VALUES: Final = frozenset({"1", "true", "yes", "on"})
class RustChatCompletions(Protocol):
def __call__(
@ -181,10 +179,6 @@ def load_rust_achat_completions() -> RustAchatCompletions | None:
return loaded
def _env_enables_rust() -> bool:
return os.getenv("LITELLM_RUST", "").strip().lower() in _TRUTHY_ENV_VALUES
def _load_rust_decline() -> RustChatCompletionsDecline | None:
if _STATE.decline is not None:
return _STATE.decline
@ -253,8 +247,8 @@ def rust_chat_completions_accepts(
return False
if stream:
return False
opted_in: Final = litellm_params is not None and litellm_params.get("rust") is True
if not opted_in and not _env_enables_rust():
request_override: Final = litellm_params.get("rust") if litellm_params is not None else None
if not rust_enabled(request_override=request_override if isinstance(request_override, bool) else None):
return False
if _litellm_metadata_reaches_the_provider(custom_llm_provider, litellm_params):
verbose_logger.debug("Rust chat completions declined (litellm metadata user_id); using the Python path")

View file

@ -0,0 +1,148 @@
from __future__ import annotations
import os
import warnings
from typing import TYPE_CHECKING, Final
if TYPE_CHECKING:
from litellm.rust_bridge.messages import RustAmessages, RustMessages
from litellm.rust_bridge.ocr import RustAocr, RustOcr
from litellm.rust_bridge.responses_websocket import RustResponsesWebSocketConnection
from litellm.rust_bridge.transcription import RustAtranscription, RustTranscription
DEFAULT_RUST_ENABLED: Final = False
_TRUE_ENV_VALUES: Final = frozenset({"1", "true", "yes", "on"})
_GLOBAL_ENV_NAME: Final = "LITELLM_RUST"
_LEGACY_OCR_ENV_NAME: Final = "LITELLM_USE_RUST_OCR"
class _Unset:
pass
_UNSET: Final = _Unset()
class _RustConfiguration:
def __init__(self) -> None:
self.override: bool | None = None
_CONFIGURATION: Final = _RustConfiguration()
def _parse_env_bool(value: str | None) -> bool | None:
if value is None:
return None
return value.strip().lower() in _TRUE_ENV_VALUES
def resolve_rust_enabled(
*,
request_override: bool | None,
process_override: bool | None,
environment_override: bool | None,
legacy_ocr_override: bool | None = None,
release_default: bool = DEFAULT_RUST_ENABLED,
) -> bool:
if request_override is not None:
return request_override
if process_override is not None:
return process_override
if environment_override is not None:
return environment_override
if legacy_ocr_override is not None:
return legacy_ocr_override
return release_default
def rust_enabled(*, request_override: bool | None = None) -> bool:
if request_override is not None:
return request_override
process_override: Final = _CONFIGURATION.override
if process_override is not None:
return process_override
return resolve_rust_enabled(
request_override=None,
process_override=None,
environment_override=_parse_env_bool(os.getenv(_GLOBAL_ENV_NAME)),
)
def rust_ocr_enabled(*, request_override: bool | None = None) -> bool:
if request_override is not None:
return request_override
process_override: Final = _CONFIGURATION.override
if process_override is not None:
return process_override
global_override: Final = _parse_env_bool(os.getenv(_GLOBAL_ENV_NAME))
legacy_override: Final = None if global_override is not None else _parse_env_bool(os.getenv(_LEGACY_OCR_ENV_NAME))
if legacy_override is not None:
warnings.warn(
f"{_LEGACY_OCR_ENV_NAME} is deprecated; use {_GLOBAL_ENV_NAME} instead",
DeprecationWarning,
stacklevel=2,
)
return resolve_rust_enabled(
request_override=None,
process_override=None,
environment_override=global_override,
legacy_ocr_override=legacy_override,
)
def reset_rust_configuration() -> None:
_CONFIGURATION.override = None
def use_litellm_rust(
enabled: bool = True,
*,
ocr: RustOcr | None | _Unset = _UNSET,
aocr: RustAocr | None | _Unset = _UNSET,
messages: RustMessages | None | _Unset = _UNSET,
amessages: RustAmessages | None | _Unset = _UNSET,
responses_websocket: type[RustResponsesWebSocketConnection] | None | _Unset = _UNSET,
transcription: RustTranscription | None | _Unset = _UNSET,
atranscription: RustAtranscription | None | _Unset = _UNSET,
) -> None:
"""Set the process override for optional Rust paths.
Rust-only paths, including Bedrock transcription, are not controlled by this switch.
"""
_CONFIGURATION.override = enabled
bindings: Final = (ocr, aocr, messages, amessages, responses_websocket, transcription, atranscription)
if all(isinstance(binding, _Unset) for binding in bindings):
return
warnings.warn(
"Injecting Rust bridge implementations through use_litellm_rust() is deprecated; "
"use the internal bridge setters in tests",
DeprecationWarning,
stacklevel=2,
)
if not isinstance(ocr, _Unset) or not isinstance(aocr, _Unset):
from litellm.rust_bridge.ocr import set_rust_ocr
if not isinstance(ocr, _Unset):
set_rust_ocr(ocr=ocr)
if not isinstance(aocr, _Unset):
set_rust_ocr(aocr=aocr)
if not isinstance(messages, _Unset) or not isinstance(amessages, _Unset):
from litellm.rust_bridge.messages import set_rust_messages
if not isinstance(messages, _Unset):
set_rust_messages(messages=messages)
if not isinstance(amessages, _Unset):
set_rust_messages(amessages=amessages)
if not isinstance(responses_websocket, _Unset):
from litellm.rust_bridge.responses_websocket import set_rust_responses_websocket
set_rust_responses_websocket(connection=responses_websocket)
if not isinstance(transcription, _Unset) or not isinstance(atranscription, _Unset):
from litellm.rust_bridge.transcription import configure_rust_transcription
if not isinstance(transcription, _Unset):
configure_rust_transcription(transcription=transcription)
if not isinstance(atranscription, _Unset):
configure_rust_transcription(atranscription=atranscription)

View file

@ -2,16 +2,16 @@
from __future__ import annotations
import os
from collections.abc import Awaitable
from typing import TYPE_CHECKING, Any, Final, Protocol, cast
from typing import Final, Protocol, cast # noqa: TID251 # native extension exposes dynamically typed callables
import httpx
from litellm.rust_bridge import configuration as _configuration
from litellm.rust_bridge.timeouts import timeout_to_seconds as _timeout_to_seconds
if TYPE_CHECKING:
from litellm.rust_bridge.messages import RustAmessages, RustMessages
rust_ocr_enabled = _configuration.rust_ocr_enabled
use_litellm_rust = _configuration.use_litellm_rust
class RustOcr(Protocol):
@ -51,69 +51,20 @@ class _Unset:
_UNSET: Final[_Unset] = _Unset()
def _env_enables_rust_ocr() -> bool:
return os.getenv("LITELLM_USE_RUST_OCR", "").strip().lower() in {
"1",
"true",
"yes",
"on",
}
_rust_ocr_enabled = _env_enables_rust_ocr()
_rust_ocr_impl: RustOcr | None = None
_rust_aocr_impl: RustAocr | None = None
def use_litellm_rust(
enabled: bool = True,
def set_rust_ocr(
*,
ocr: RustOcr | None | _Unset = _UNSET,
aocr: RustAocr | None | _Unset = _UNSET,
messages: RustMessages | None | _Unset = _UNSET,
amessages: RustAmessages | None | _Unset = _UNSET,
responses_websocket: Any | None | _Unset = _UNSET,
transcription: Any | None | _Unset = _UNSET,
atranscription: Any | None | _Unset = _UNSET,
) -> None:
global _rust_ocr_enabled, _rust_ocr_impl, _rust_aocr_impl
configuring_ocr: Final = not isinstance(ocr, _Unset) or not isinstance(aocr, _Unset)
configuring_messages: Final = not isinstance(messages, _Unset) or not isinstance(amessages, _Unset)
configuring_responses_websocket: Final = not isinstance(responses_websocket, _Unset)
configuring_transcription: Final = not isinstance(transcription, _Unset) or not isinstance(atranscription, _Unset)
if configuring_ocr or (not configuring_messages and not configuring_responses_websocket):
_rust_ocr_enabled = enabled
global _rust_ocr_impl, _rust_aocr_impl
if not isinstance(ocr, _Unset):
_rust_ocr_impl = ocr
if not isinstance(aocr, _Unset):
_rust_aocr_impl = aocr
if configuring_transcription:
from litellm.rust_bridge.transcription import configure_rust_transcription
configure_rust_transcription(
enabled=enabled,
transcription=transcription,
atranscription=atranscription,
)
if not configuring_messages and not configuring_responses_websocket:
return
if configuring_messages:
from litellm.rust_bridge.messages import set_rust_messages
if not isinstance(messages, _Unset) and not isinstance(amessages, _Unset):
set_rust_messages(messages=messages, amessages=amessages)
elif not isinstance(messages, _Unset):
set_rust_messages(messages=messages)
else:
set_rust_messages(amessages=amessages)
if configuring_responses_websocket:
from litellm.rust_bridge.responses_websocket import set_rust_responses_websocket
set_rust_responses_websocket(connection=responses_websocket)
def rust_ocr_enabled() -> bool:
return _rust_ocr_enabled
def load_rust_ocr() -> RustOcr | None:

View file

@ -0,0 +1,180 @@
from __future__ import annotations
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from enum import Enum
from typing import Final, Generic, NoReturn, TypeAlias, TypeVar
from litellm.exceptions import APIError
from litellm.rust_bridge.bindings import native_exception_types
NativeT = TypeVar("NativeT")
ResultT = TypeVar("ResultT")
class FallbackMode(Enum):
PYTHON = "python"
RUST_REQUIRED = "rust_required"
@dataclass(frozen=True, slots=True)
class RustHandled(Generic[ResultT]):
value: ResultT
@dataclass(frozen=True, slots=True)
class RustDeclined:
reason: str
@dataclass(frozen=True, slots=True)
class RustUnavailable:
pass
RustAttempt: TypeAlias = RustHandled[ResultT] | RustDeclined | RustUnavailable
@dataclass(frozen=True, slots=True)
class BridgeErrorContext:
route: str
provider: str
model: str
def invoke(
*,
native_call: Callable[[], NativeT] | None,
fallback: Callable[[], ResultT],
adapt: Callable[[NativeT], ResultT],
mode: FallbackMode,
context: BridgeErrorContext,
) -> ResultT:
result: Final = attempt(native_call=native_call, adapt=adapt, context=context)
if isinstance(result, RustHandled):
return result.value
if mode is FallbackMode.PYTHON:
return fallback()
_raise_required(result, context)
async def ainvoke(
*,
native_call: Callable[[], Awaitable[NativeT]] | None,
fallback: Callable[[], Awaitable[ResultT]],
adapt: Callable[[NativeT], ResultT],
mode: FallbackMode,
context: BridgeErrorContext,
) -> ResultT:
result: Final = await aattempt(native_call=native_call, adapt=adapt, context=context)
if isinstance(result, RustHandled):
return result.value
if mode is FallbackMode.PYTHON:
return await fallback()
_raise_required(result, context)
def attempt(
*,
native_call: Callable[[], NativeT] | None,
adapt: Callable[[NativeT], ResultT],
context: BridgeErrorContext,
) -> RustAttempt[ResultT]:
if native_call is None:
return RustUnavailable()
exceptions: Final = native_exception_types()
if exceptions is None:
return RustHandled(adapt(native_call()))
declined, upstream = exceptions
try:
value: Final = native_call()
except declined as error:
return RustDeclined(reason=_decline_reason(error))
except upstream as error:
_raise_upstream(error, context)
return RustHandled(adapt(value))
async def aattempt(
*,
native_call: Callable[[], Awaitable[NativeT]] | None,
adapt: Callable[[NativeT], ResultT],
context: BridgeErrorContext,
) -> RustAttempt[ResultT]:
if native_call is None:
return RustUnavailable()
exceptions: Final = native_exception_types()
if exceptions is None:
return RustHandled(adapt(await native_call()))
declined, upstream = exceptions
try:
value: Final = await native_call()
except declined as error:
return RustDeclined(reason=_decline_reason(error))
except upstream as error:
_raise_upstream(error, context)
return RustHandled(adapt(value))
def call(operation: Callable[[], ResultT], context: BridgeErrorContext) -> ResultT:
exceptions: Final = native_exception_types()
if exceptions is None:
return operation()
upstream: Final = exceptions[1]
try:
return operation()
except upstream as error:
_raise_upstream(error, context)
async def acall(operation: Callable[[], Awaitable[ResultT]], context: BridgeErrorContext) -> ResultT:
exceptions: Final = native_exception_types()
if exceptions is None:
return await operation()
upstream: Final = exceptions[1]
try:
return await operation()
except upstream as error:
_raise_upstream(error, context)
def _decline_reason(error: BaseException) -> str:
reason: Final[object] = error.args[0] if error.args else str(error)
return reason if isinstance(reason, str) else str(reason)
def _raise_required(
result: RustDeclined | RustUnavailable,
context: BridgeErrorContext,
) -> NoReturn:
raise RuntimeError(f"Rust {context.route} bridge {_required_reason(result)}")
def _required_reason(result: RustDeclined | RustUnavailable) -> str:
match result:
case RustUnavailable():
return "is unavailable"
case RustDeclined(reason=reason):
return f"declined the request: {reason}"
def _raise_upstream(error: BaseException, context: BridgeErrorContext) -> NoReturn:
args: Final[tuple[object, ...]] = error.args
status_value: Final = args[0] if args else 0
message_value: Final = args[1] if len(args) > 1 else str(error)
status: Final = status_value if isinstance(status_value, int) else 0
message: Final = message_value if isinstance(message_value, str) else str(message_value)
raise APIError(
status_code=status or 500,
message=f"litellm rust {context.route}: {message}",
llm_provider=context.provider,
model=context.model,
) from error
def identity(value: ResultT) -> ResultT:
return value
async def async_none() -> None:
return None

View file

@ -369,6 +369,7 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
"""
custom_llm_provider: str | None = None
rust: bool | None = None
tpm: int | None = None
rpm: int | None = None
itpm: int | None = None

View file

@ -2839,6 +2839,7 @@ class StandardLoggingRoutingDecisionTierBoundaries(TypedDict):
RoutingDecisionCause = Literal[
"heuristic_scorer",
"heuristic_v2",
# The scorer found 2+ reasoning markers and forced REASONING regardless of score.
# A distinct cause rather than a marker inside `signals`, because it is the fact
# that tells a reader the score did NOT choose the tier; encoding it as free text

View file

@ -278,7 +278,10 @@ bindings = "pyo3"
features = ["extension-module"]
profile = "release"
editable-profile = "dev"
include = ["litellm/proxy/_experimental/out/**"]
include = [
"litellm/proxy/_experimental/out/**",
"litellm/router_strategy/complexity_router/artifacts/*.json",
]
exclude = [
"litellm/proxy/enterprise",
"litellm/proxy/enterprise/**",

View file

@ -24,7 +24,7 @@
"limit": 133
},
"ANN401": {
"limit": 307
"limit": 304
},
"ASYNC230": {
"limit": 11
@ -156,7 +156,7 @@
"limit": 215
},
"PLW0603": {
"limit": 191
"limit": 190
},
"PLW1508": {
"limit": 190
@ -231,7 +231,7 @@
"limit": 5
},
"TID251": {
"limit": 1073
"limit": 1071
},
"TRY002": {
"limit": 524

View file

@ -333,6 +333,7 @@ async def test_bedrock_kb_request_body_has_transformed_filters(
timeout=None,
client=None,
_is_async=False,
router: "litellm.Router | None" = None,
):
litellm_params_dict = (
litellm_params.model_dump(exclude_none=False)

View file

@ -0,0 +1,145 @@
# Rust ↔ Python SDK parity harness
This folder is the operator-facing harness for the Rust migration test plan. It runs pytest normally, listens to test events in-process, and redraws a live matrix grouped by testing strategy and SDK-level function.
The matrix always has these SDK columns:
- `ocr / aocr`
- `messages / amessages`
- `responses / aresponses`
- `count_tokens`
The harness has three deliberately broad test-strategy folders:
| Strategy | Folder |
| --- | --- |
| Public SDK parity over generated and recorded inputs | [`e2e_fuzz_tests/`](e2e_fuzz_tests/) |
| Focused tests of Rust-owned behavior | [`unit_tests_rust/`](unit_tests_rust/) |
| Isolated transform and Python-to-Rust helper coverage | [`validate_sub_methods/`](validate_sub_methods/) |
## Run it
From the repository root:
```bash
poetry run python -m tests.rust-python-harness
```
The default runs every configured test once and updates all matching cells in real time. Narrow a run by strategy, SDK function, or both:
```bash
poetry run python -m tests.rust-python-harness --strategy e2e_fuzz_tests
poetry run python -m tests.rust-python-harness --function messages
poetry run python -m tests.rust-python-harness --strategy validate_sub_methods --function ocr
```
For a guided run, use the interactive picker. It asks which strategy rows and SDK
function columns to include, then hands the terminal to the live dashboard. It never
captures keys while tests are running, so Ctrl-C and pytest debugging remain safe.
```bash
poetry run python -m tests.rust-python-harness --interactive
```
Useful operator options:
```bash
# Inspect coverage and pytest selectors without running anything.
poetry run python -m tests.rust-python-harness --list
# Stable line-oriented output for CI logs or redirected output.
poetry run python -m tests.rust-python-harness --plain
# Measure Python reference lines exercised by this parity run and build an HTML heatmap.
poetry run python -m tests.rust-python-harness --coverage
# Forward pytest options. Use the equals form when the value begins with a dash.
poetry run python -m tests.rust-python-harness --pytest-arg=-x
```
The process returns pytest's exit code. A configured selector that collects no test is also a failure. A planned cell has no selector yet and does not fail the run.
The dashboard adapts to narrow terminals, shows elapsed time and unique-test progress,
and prints the three slowest tests when the run ends. Each failure includes a focused
`poetry run pytest ... -q` command. Redirected output and CI automatically use the
line-oriented plain renderer; `--plain` lets you opt into it locally.
The final screen includes a confidence score for every SDK section. It is the direct
ratio of required strategy rows with passing evidence, such as `1/3 = 33%`; High means
all required strategies passed, Medium means some passed, and Low means none passed.
This behavioral score is intentionally shown separately from Python and Rust LOC.
Coverage reports are written outside the three strategy folders at
`target/rust-python-harness/`. Open `python-html/index.html` to inspect executed and
missing Python lines; `python.json` and `python.xml` are available for automation.
Coverage is finalized after pytest exits, because worker processes must flush their
data first.
## Port coverage and confidence
Treat these as separate signals instead of one ambiguous coverage percentage:
| Signal | Tool | What it proves |
| --- | --- | --- |
| Python reference LOC | `coverage.py` / `pytest-cov` via `--coverage` | The mapped Python behavior ran |
| Rust port LOC | `cargo-llvm-cov` | The mapped Rust implementation ran |
| Parity contracts | This harness matrix | Python and Rust had the same observable behavior |
`validate_sub_methods/` owns the future source-section inventory that maps a stable
Python qualified symbol to its Rust symbol. That inventory is the denominator for
per-function rollups; raw coverage for the entire LiteLLM repository would obscure
the port's real gaps. `unit_tests_rust/` owns direct `cargo-llvm-cov` runs, while
`e2e_fuzz_tests/` owns behavioral parity and fuzz-case counts. Keep Python, Rust, and
parity percentages visible side by side and label section confidence High only when
the mapped implementation exists, every required strategy passes, and both sides meet
their LOC thresholds. Generated Rust LCOV/HTML and the combined index also belong in
`target/rust-python-harness/`, not in a fourth strategy folder.
## Read the matrix
| Mark | Meaning |
| --- | --- |
| `✓` | All collected tests passed |
| `✗` | At least one test failed |
| `!` | Test setup or teardown failed |
| `↷` | All collected tests skipped |
| `?` | A configured selector did not collect a test |
| `—` | Strategy is planned but has no test yet |
| `n/a` | Strategy does not apply to this SDK function |
| `◐` | The configured tests cover only part of the TDD's parity contract |
The initial end-to-end entries deliberately show `◐`: the repository has Rust bridge tests for OCR, Messages, and Responses websocket plumbing, but those are not yet frozen-Python-oracle comparisons. The remaining TDD cells stay visible as planned work instead of disappearing from a green summary.
## Attach parity tests
Each of the three folders contains a concise `README.md` and a `strategy.json`. Add a pytest file or node ID to the appropriate SDK function's `selectors` list:
```json
{
"coverage": "complete",
"selectors": [
"tests/rust-python-harness/validate_sub_methods/test_messages.py"
]
}
```
Selectors use the same syntax as pytest. A file selector aggregates every test in the file; a node selector can target one test or parametrized family. The runner deduplicates selectors, so one test may intentionally prove more than one cell without executing twice.
Use these coverage values:
- `complete`: implements the full strategy contract for that SDK function.
- `partial`: useful coverage exists, but the TDD contract is not fully proven.
- `planned`: no runnable parity test exists yet.
- `not_applicable`: the strategy cannot apply, such as streaming for OCR.
Keep comparison mechanics in shared harness modules and provider/function facts in the owning strategy folder. A Python/Rust mismatch is a test failure; do not normalize away observable return types, exception classes, private response fields, chunk ordering, or callback payload differences merely to make a cell green.
## Architecture
- `catalog.py` validates and loads every strategy manifest.
- `models.py` owns typed strategy, case, coverage, and run-state models.
- `runner.py` maps live pytest events back to one or more matrix cells.
- `ui.py` renders the interactive Rich dashboard and a dependency-free plain fallback.
- `cli.py` handles filtering and preserves pytest exit semantics.
The harness is driven from Python, matching the SDK surface and existing test tooling. Rust remains responsible for the implementation under comparison; the harness does not move provider semantics into the PyO3 bridge.

View file

@ -0,0 +1,5 @@
"""Interactive Rust/Python SDK parity test harness."""
from .catalog import load_catalog
__all__ = ["load_catalog"]

View file

@ -0,0 +1,4 @@
from .cli import main
if __name__ == "__main__":
raise SystemExit(main())

View file

@ -0,0 +1,93 @@
from __future__ import annotations
import json
from pathlib import Path
from typing import Any
from .models import Coverage, HarnessCase, SDK_FUNCTIONS, Strategy
STRATEGIES_ROOT = Path(__file__).parent
def _require_string(value: Any, field: str, source: Path) -> str:
if not isinstance(value, str) or not value.strip():
raise ValueError(f"{source}: {field} must be a non-empty string")
return value
def _load_strategy(source: Path) -> Strategy:
with source.open(encoding="utf-8") as stream:
data = json.load(stream)
strategy_id = _require_string(data.get("id"), "id", source)
label = _require_string(data.get("label"), "label", source)
description = _require_string(data.get("description"), "description", source)
order = data.get("order")
if not isinstance(order, int):
raise ValueError(f"{source}: order must be an integer")
function_data = data.get("functions")
if not isinstance(function_data, dict):
raise ValueError(f"{source}: functions must be an object")
missing = set(SDK_FUNCTIONS) - set(function_data)
extra = set(function_data) - set(SDK_FUNCTIONS)
if missing or extra:
raise ValueError(
f"{source}: functions must exactly match {SDK_FUNCTIONS}; missing={missing}, extra={extra}"
)
cases: list[HarnessCase] = []
for sdk_function in SDK_FUNCTIONS:
case_data = function_data[sdk_function]
if not isinstance(case_data, dict):
raise ValueError(f"{source}: functions.{sdk_function} must be an object")
try:
coverage = Coverage(case_data.get("coverage"))
except ValueError as exc:
raise ValueError(f"{source}: invalid coverage for {sdk_function}") from exc
selectors = case_data.get("selectors", [])
if not isinstance(selectors, list) or not all(
isinstance(item, str) and item for item in selectors
):
raise ValueError(
f"{source}: selectors for {sdk_function} must be a list of strings"
)
if coverage is Coverage.NOT_APPLICABLE and selectors:
raise ValueError(
f"{source}: not_applicable case {sdk_function} cannot have selectors"
)
cases.append(
HarnessCase(
strategy_id=strategy_id,
strategy_label=label,
sdk_function=sdk_function,
coverage=coverage,
selectors=tuple(selectors),
note=str(case_data.get("note", "")),
)
)
return Strategy(
order=order,
id=strategy_id,
label=label,
description=description,
directory=source.parent,
cases=tuple(cases),
)
def load_catalog(root: Path = STRATEGIES_ROOT) -> tuple[Strategy, ...]:
sources = sorted(root.glob("*/strategy.json"))
if not sources:
raise ValueError(f"No strategy manifests found below {root}")
strategies = tuple(
sorted(
(_load_strategy(source) for source in sources),
key=lambda strategy: strategy.order,
)
)
ids = [strategy.id for strategy in strategies]
if len(ids) != len(set(ids)):
raise ValueError(f"Duplicate strategy id in {root}")
return strategies

View file

@ -0,0 +1,180 @@
from __future__ import annotations
import argparse
import importlib.util
from collections.abc import Sequence
from pathlib import Path
from .catalog import load_catalog
from .models import HarnessCase, Strategy
from .runner import run_pytest
from .ui import make_dashboard
REPO_ROOT = Path(__file__).resolve().parents[2]
COVERAGE_ROOT = REPO_ROOT / "target" / "rust-python-harness"
def _parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
prog="rust-python-harness",
description="Run Rust/Python parity tests with a live strategy-by-SDK-function dashboard.",
)
parser.add_argument(
"-i",
"--interactive",
action="store_true",
help="pick strategies and SDK functions in a guided terminal menu",
)
parser.add_argument(
"--list", action="store_true", help="show the catalog without running tests"
)
parser.add_argument(
"--strategy",
action="append",
default=[],
metavar="ID",
help="run only this strategy",
)
parser.add_argument(
"--function",
action="append",
default=[],
dest="sdk_functions",
choices=("ocr", "messages", "responses", "count_tokens"),
help="run only this SDK function",
)
parser.add_argument(
"--plain",
action="store_true",
help="disable the interactive terminal dashboard",
)
parser.add_argument(
"--coverage",
action="store_true",
help="write Python reference LOC reports (HTML, JSON, and XML)",
)
parser.add_argument(
"--pytest-arg",
action="append",
default=[],
metavar="ARG",
help="append an argument to pytest (repeatable, for example --pytest-arg=-x)",
)
return parser
def _coverage_pytest_args(output_root: Path = COVERAGE_ROOT) -> tuple[str, ...]:
output_root.mkdir(parents=True, exist_ok=True)
return (
"--cov=litellm",
"--cov-context=test",
f"--cov-report=json:{output_root / 'python.json'}",
f"--cov-report=xml:{output_root / 'python.xml'}",
f"--cov-report=html:{output_root / 'python-html'}",
)
def _pick_values(
title: str, options: Sequence[tuple[str, str]], input_fn=input
) -> set[str]:
print(f"\n{title} (Enter = all)")
for index, (value, label) in enumerate(options, start=1):
print(f" {index:>2}. {label} [{value}]")
while True:
answer = input_fn("Choose numbers, comma-separated: ").strip()
if not answer:
return set()
try:
indexes = {int(part.strip()) for part in answer.split(",")}
except ValueError:
print("Please enter numbers separated by commas.")
continue
if indexes and all(1 <= index <= len(options) for index in indexes):
return {options[index - 1][0] for index in indexes}
print(f"Choose values from 1 to {len(options)}.")
def _interactive_filters(strategies: Sequence[Strategy]) -> tuple[set[str], set[str]]:
strategy_ids = _pick_values(
"Testing strategies", [(strategy.id, strategy.label) for strategy in strategies]
)
sdk_functions = _pick_values(
"SDK functions",
[(name, name) for name in ("ocr", "messages", "responses", "count_tokens")],
)
return strategy_ids, sdk_functions
def _select(
strategies: Sequence[Strategy], strategy_ids: set[str], sdk_functions: set[str]
) -> tuple[HarnessCase, ...]:
known_ids = {strategy.id for strategy in strategies}
unknown = strategy_ids - known_ids
if unknown:
raise ValueError(f"Unknown strategy: {', '.join(sorted(unknown))}")
return tuple(
case
for strategy in strategies
if not strategy_ids or strategy.id in strategy_ids
for case in strategy.cases
if not sdk_functions or case.sdk_function in sdk_functions
)
def _print_catalog(strategies: Sequence[Strategy]) -> None:
for strategy in strategies:
print(f"{strategy.id:20} {strategy.label}")
for case in strategy.cases:
selectors = (
", ".join(case.selectors) if case.selectors else "no test configured"
)
print(f" {case.sdk_function:12} {case.coverage.value:14} {selectors}")
def main(argv: Sequence[str] | None = None) -> int:
args = _parser().parse_args(argv)
if args.coverage and importlib.util.find_spec("pytest_cov") is None:
_parser().error(
"--coverage requires the project's pytest-cov dependency; run with "
"`poetry run python -m tests.rust-python-harness --coverage`"
)
strategies = load_catalog()
if args.list:
_print_catalog(strategies)
return 0
strategy_ids = set(args.strategy)
sdk_functions = set(args.sdk_functions)
if args.interactive:
picked_strategies, picked_functions = _interactive_filters(strategies)
strategy_ids = strategy_ids or picked_strategies
sdk_functions = sdk_functions or picked_functions
try:
cases = _select(strategies, strategy_ids, sdk_functions)
except ValueError as exc:
_parser().error(str(exc))
selected_strategy_ids = {case.strategy_id for case in cases}
visible_strategies = tuple(
strategy for strategy in strategies if strategy.id in selected_strategy_ids
)
dashboard = make_dashboard(
visible_strategies,
plain=args.plain,
confidence_strategies=strategies,
)
pytest_args = [*args.pytest_arg]
if args.coverage:
pytest_args.extend(_coverage_pytest_args())
with dashboard:
exit_code, run = run_pytest(
cases=cases,
repo_root=REPO_ROOT,
on_update=dashboard.update,
pytest_args=pytest_args,
)
dashboard.finish(run, exit_code)
if args.coverage and (COVERAGE_ROOT / "python.json").exists():
print(f"Python LOC heatmap: {COVERAGE_ROOT / 'python-html' / 'index.html'}")
print(f"Machine-readable coverage: {COVERAGE_ROOT / 'python.json'}")
return exit_code

View file

@ -0,0 +1,3 @@
# End-to-end fuzz tests
Runs the same SDK call through the Python and Rust paths using generated inputs and recorded provider responses. It compares public results, streams, callbacks, and exceptions to catch behavior differences a unit test can miss.

View file

@ -0,0 +1,12 @@
{
"order": 10,
"id": "e2e_fuzz_tests",
"label": "End-to-end fuzz tests",
"description": "Compare observable Python and Rust SDK behavior over generated and recorded inputs.",
"functions": {
"ocr": {"coverage": "partial", "selectors": ["tests/test_litellm/ocr/test_rust_bridge.py"], "note": "Bridge coverage exists; frozen-oracle fuzz parity is still being added."},
"messages": {"coverage": "partial", "selectors": ["tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py"], "note": "Bridge coverage exists; frozen-oracle fuzz parity is still being added."},
"responses": {"coverage": "partial", "selectors": ["tests/test_litellm/responses/test_rust_bridge_websocket.py"], "note": "Covers the websocket bridge; full responses parity is still being added."},
"count_tokens": {"coverage": "planned", "selectors": [], "note": "No Rust count_tokens parity test is present yet."}
}
}

View file

@ -0,0 +1,219 @@
from __future__ import annotations
from dataclasses import dataclass, field
from enum import Enum
from pathlib import Path
from time import monotonic
from typing import Iterable
class Coverage(str, Enum):
COMPLETE = "complete"
PARTIAL = "partial"
PLANNED = "planned"
NOT_APPLICABLE = "not_applicable"
class RunStatus(str, Enum):
NOT_RUN = "not_run"
QUEUED = "queued"
RUNNING = "running"
PASSED = "passed"
FAILED = "failed"
SKIPPED = "skipped"
ERROR = "error"
MISSING = "missing"
PLANNED = "planned"
NOT_APPLICABLE = "not_applicable"
class ConfidenceLevel(str, Enum):
HIGH = "HIGH"
MEDIUM = "MEDIUM"
LOW = "LOW"
SDK_FUNCTIONS = ("ocr", "messages", "responses", "count_tokens")
@dataclass(frozen=True)
class HarnessCase:
strategy_id: str
strategy_label: str
sdk_function: str
coverage: Coverage
selectors: tuple[str, ...]
note: str = ""
@property
def key(self) -> str:
return f"{self.strategy_id}:{self.sdk_function}"
@dataclass(frozen=True)
class Strategy:
order: int
id: str
label: str
description: str
directory: Path
cases: tuple[HarnessCase, ...]
@dataclass
class CaseResult:
case: HarnessCase
status: RunStatus = RunStatus.NOT_RUN
collected: set[str] = field(default_factory=set)
completed: set[str] = field(default_factory=set)
passed: int = 0
failed: int = 0
skipped: int = 0
errors: int = 0
outcomes: dict[str, RunStatus] = field(default_factory=dict)
durations: dict[str, float] = field(default_factory=dict)
@property
def total(self) -> int:
return len(self.collected)
@property
def duration(self) -> float:
return sum(self.durations.values())
def record(self, nodeid: str, status: RunStatus, duration: float = 0.0) -> None:
"""Record a terminal outcome, allowing teardown errors to replace a pass."""
self.outcomes[nodeid] = status
self.durations[nodeid] = self.durations.get(nodeid, 0.0) + duration
self.completed = set(self.outcomes)
values = tuple(self.outcomes.values())
self.passed = values.count(RunStatus.PASSED)
self.failed = values.count(RunStatus.FAILED)
self.skipped = values.count(RunStatus.SKIPPED)
self.errors = values.count(RunStatus.ERROR)
self.finalize()
def set_initial_status(self) -> None:
if self.case.coverage is Coverage.NOT_APPLICABLE:
self.status = RunStatus.NOT_APPLICABLE
elif not self.case.selectors:
self.status = RunStatus.PLANNED
else:
self.status = RunStatus.QUEUED
def finalize(self) -> None:
if self.status in {RunStatus.NOT_APPLICABLE, RunStatus.PLANNED}:
return
if not self.collected:
self.status = RunStatus.MISSING
elif self.errors:
self.status = RunStatus.ERROR
elif self.failed:
self.status = RunStatus.FAILED
elif self.passed and len(self.completed) == len(self.collected):
self.status = RunStatus.PASSED
elif self.skipped and len(self.completed) == len(self.collected):
self.status = RunStatus.SKIPPED
@dataclass
class HarnessRun:
results: dict[str, CaseResult]
current_nodeid: str | None = None
failures: list[tuple[str, str]] = field(default_factory=list)
started_at: float = field(default_factory=monotonic)
finished_at: float | None = None
@property
def duration(self) -> float:
return (self.finished_at or monotonic()) - self.started_at
@property
def unique_tests(self) -> int:
return len(
{nodeid for result in self.results.values() for nodeid in result.collected}
)
@property
def completed_tests(self) -> int:
return len(
{nodeid for result in self.results.values() for nodeid in result.completed}
)
@classmethod
def from_cases(cls, cases: Iterable[HarnessCase]) -> "HarnessRun":
results = {case.key: CaseResult(case=case) for case in cases}
for result in results.values():
result.set_initial_status()
return cls(results=results)
@dataclass(frozen=True)
class SectionConfidence:
sdk_function: str
verified_strategies: int
required_strategies: int
level: ConfidenceLevel
details: tuple[str, ...]
@property
def percentage(self) -> int:
if not self.required_strategies:
return 0
return round(100 * self.verified_strategies / self.required_strategies)
def section_confidence(
run: HarnessRun, strategies: Iterable[Strategy]
) -> tuple[SectionConfidence, ...]:
strategy_list = tuple(strategies)
scores: list[SectionConfidence] = []
for sdk_function in SDK_FUNCTIONS:
cases = tuple(
case
for strategy in strategy_list
for case in strategy.cases
if case.sdk_function == sdk_function
and case.coverage is not Coverage.NOT_APPLICABLE
)
verified = 0
details: list[str] = []
for case in cases:
result = run.results.get(case.key)
status = result.status if result is not None else RunStatus.NOT_RUN
if status is RunStatus.PASSED:
verified += 1
details.append(
f"{STATUS_LABELS[status]} {case.strategy_id} ({case.coverage.value})"
)
required = len(cases)
if required and verified == required:
level = ConfidenceLevel.HIGH
elif verified:
level = ConfidenceLevel.MEDIUM
else:
level = ConfidenceLevel.LOW
scores.append(
SectionConfidence(
sdk_function=sdk_function,
verified_strategies=verified,
required_strategies=required,
level=level,
details=tuple(details),
)
)
return tuple(scores)
STATUS_LABELS = {
RunStatus.NOT_RUN: "·",
RunStatus.QUEUED: "○",
RunStatus.RUNNING: "◉",
RunStatus.PASSED: "✓",
RunStatus.FAILED: "✗",
RunStatus.SKIPPED: "↷",
RunStatus.ERROR: "!",
RunStatus.MISSING: "?",
RunStatus.PLANNED: "—",
RunStatus.NOT_APPLICABLE: "n/a",
}

View file

@ -0,0 +1,160 @@
from __future__ import annotations
import os
from collections.abc import Callable, Sequence
from pathlib import Path
from time import monotonic
import pytest
from .models import CaseResult, HarnessCase, HarnessRun, RunStatus
UpdateCallback = Callable[[HarnessRun], None]
def selector_matches_node(selector: str, nodeid: str) -> bool:
normalized_selector = selector.replace("\\", "/")
normalized_nodeid = nodeid.replace("\\", "/")
if "::" in normalized_selector:
return normalized_nodeid == normalized_selector or normalized_nodeid.startswith(
f"{normalized_selector}["
)
return normalized_nodeid == normalized_selector or normalized_nodeid.startswith(
f"{normalized_selector}::"
)
def selector_path(selector: str) -> Path:
return Path(selector.split("::", 1)[0])
def runnable_selectors(
cases: Sequence[HarnessCase], repo_root: Path
) -> tuple[str, ...]:
selectors = {
selector
for case in cases
for selector in case.selectors
if (repo_root / selector_path(selector)).exists()
}
return tuple(sorted(selectors))
class HarnessPytestPlugin:
def __init__(self, run: HarnessRun, on_update: UpdateCallback) -> None:
self.run = run
self.on_update = on_update
self.node_to_results: dict[str, list[CaseResult]] = {}
def _notify(self) -> None:
self.on_update(self.run)
def pytest_collection_modifyitems(self, items: list[pytest.Item]) -> None:
for item in items:
matched_results: list[CaseResult] = []
for result in self.run.results.values():
if any(
selector_matches_node(selector, item.nodeid)
for selector in result.case.selectors
):
result.collected.add(item.nodeid)
matched_results.append(result)
if matched_results:
self.node_to_results[item.nodeid] = matched_results
for result in self.run.results.values():
if result.status is RunStatus.QUEUED and not result.collected:
result.status = RunStatus.MISSING
self._notify()
def pytest_runtest_logstart(
self, nodeid: str, location: tuple[str, int | None, str]
) -> None:
del location
self.run.current_nodeid = nodeid
for result in self.node_to_results.get(nodeid, []):
if result.status not in {RunStatus.FAILED, RunStatus.ERROR}:
result.status = RunStatus.RUNNING
self._notify()
def pytest_runtest_logreport(self, report: pytest.TestReport) -> None:
if report.when not in {"setup", "call", "teardown"}:
return
results = self.node_to_results.get(report.nodeid, [])
if not results:
return
terminal = report.when == "call" or report.failed or report.skipped
if not terminal:
for result in results:
result.durations[report.nodeid] = (
result.durations.get(report.nodeid, 0.0) + report.duration
)
return
for result in results:
if report.when == "teardown" and not report.failed:
result.durations[report.nodeid] = (
result.durations.get(report.nodeid, 0.0) + report.duration
)
continue
if report.skipped:
status = RunStatus.SKIPPED
elif report.failed and report.when in {"setup", "teardown"}:
status = RunStatus.ERROR
elif report.failed:
status = RunStatus.FAILED
else:
status = RunStatus.PASSED
result.record(report.nodeid, status, report.duration)
if report.failed:
failure = (report.nodeid, str(report.longrepr))
if failure not in self.run.failures:
self.run.failures.append(failure)
self._notify()
def pytest_sessionfinish(
self, session: pytest.Session, exitstatus: int | pytest.ExitCode
) -> None:
del session, exitstatus
self.run.current_nodeid = None
self.run.finished_at = monotonic()
for result in self.run.results.values():
result.finalize()
self._notify()
def run_pytest(
cases: Sequence[HarnessCase],
repo_root: Path,
on_update: UpdateCallback,
pytest_args: Sequence[str] = (),
) -> tuple[int, HarnessRun]:
run = HarnessRun.from_cases(cases)
selectors = runnable_selectors(cases, repo_root)
if not selectors:
for result in run.results.values():
result.finalize()
run.finished_at = monotonic()
on_update(run)
has_missing_test = any(
result.status is RunStatus.MISSING for result in run.results.values()
)
exit_code = (
int(pytest.ExitCode.TESTS_FAILED)
if has_missing_test
else int(pytest.ExitCode.OK)
)
return exit_code, run
plugin = HarnessPytestPlugin(run=run, on_update=on_update)
args = [*selectors, "-p", "no:terminal", *pytest_args]
previous_directory = Path.cwd()
try:
os.chdir(repo_root)
exit_code = int(pytest.main(args, plugins=[plugin]))
finally:
os.chdir(previous_directory)
if exit_code == 0 and any(
result.status is RunStatus.MISSING for result in run.results.values()
):
exit_code = int(pytest.ExitCode.TESTS_FAILED)
return exit_code, run

View file

@ -0,0 +1,295 @@
from __future__ import annotations
import os
import shlex
import sys
from collections.abc import Sequence
from contextlib import AbstractContextManager
from pathlib import Path
from typing import Any
from .models import (
Coverage,
HarnessRun,
RunStatus,
SDK_FUNCTIONS,
Strategy,
section_confidence,
)
STATUS_GLYPHS = {
RunStatus.NOT_RUN: "·",
RunStatus.QUEUED: "○",
RunStatus.RUNNING: "◉",
RunStatus.PASSED: "✓",
RunStatus.FAILED: "✗",
RunStatus.SKIPPED: "↷",
RunStatus.ERROR: "!",
RunStatus.MISSING: "?",
RunStatus.PLANNED: "—",
RunStatus.NOT_APPLICABLE: "n/a",
}
STATUS_STYLES = {
RunStatus.QUEUED: "dim",
RunStatus.RUNNING: "bold cyan",
RunStatus.PASSED: "bold green",
RunStatus.FAILED: "bold red",
RunStatus.SKIPPED: "yellow",
RunStatus.ERROR: "bold red",
RunStatus.MISSING: "magenta",
RunStatus.PLANNED: "dim",
RunStatus.NOT_APPLICABLE: "dim",
}
def _format_duration(seconds: float) -> str:
if seconds < 1:
return f"{seconds * 1000:.0f}ms"
if seconds < 60:
return f"{seconds:.1f}s"
return f"{int(seconds // 60)}m {seconds % 60:.0f}s"
def _rerun_command(nodeid: str) -> str:
return f"poetry run pytest {shlex.quote(nodeid)} -q"
def _summary(run: HarnessRun) -> tuple[int, int, int, int]:
outcomes: dict[str, RunStatus] = {}
for result in run.results.values():
outcomes.update(result.outcomes)
return (
list(outcomes.values()).count(RunStatus.PASSED),
list(outcomes.values()).count(RunStatus.FAILED),
list(outcomes.values()).count(RunStatus.ERROR),
list(outcomes.values()).count(RunStatus.SKIPPED),
)
def _cell_text(run: HarnessRun, strategy_id: str, sdk_function: str) -> tuple[str, str]:
result = run.results.get(f"{strategy_id}:{sdk_function}")
if result is None:
return "", ""
counts = ""
if result.total:
counts = f" {len(result.completed)}/{result.total}"
coverage = " ◐" if result.case.coverage is Coverage.PARTIAL else ""
return f"{STATUS_GLYPHS[result.status]}{counts}{coverage}", STATUS_STYLES.get(
result.status, ""
)
class RichDashboard(AbstractContextManager["RichDashboard"]):
def __init__(
self,
strategies: Sequence[Strategy],
confidence_strategies: Sequence[Strategy],
) -> None:
from rich.console import Console
from rich.live import Live
self.strategies = strategies
self.confidence_strategies = confidence_strategies
self.console = Console()
self.live: Any = Live(
console=self.console, refresh_per_second=12, transient=False
)
def _table(self, run: HarnessRun) -> Any:
from rich import box
from rich.table import Table
from rich.text import Text
narrow = self.console.width < 96
if narrow:
table = Table(box=box.SIMPLE_HEAVY, expand=True, show_header=False)
table.add_column("Strategy", ratio=3)
table.add_column("Results", ratio=5)
for strategy in self.strategies:
values = []
for sdk_function in SDK_FUNCTIONS:
value, style = _cell_text(run, strategy.id, sdk_function)
if value:
values.append(
Text.assemble((f"{sdk_function} ", "dim"), (value, style))
)
table.add_row(strategy.label, Text(" ").join(values))
return table
table = Table(box=box.ROUNDED, expand=True, title="Strategy × SDK function")
table.add_column("Strategy", ratio=3)
for label in ("ocr/aocr", "messages", "responses", "count_tokens"):
table.add_column(label, justify="center", ratio=1)
for strategy in self.strategies:
cells = []
for sdk_function in SDK_FUNCTIONS:
value, style = _cell_text(run, strategy.id, sdk_function)
cells.append(Text(value, style=style))
table.add_row(strategy.label, *cells)
return table
def __enter__(self) -> "RichDashboard":
self.live.__enter__()
return self
def __exit__(self, *args: object) -> None:
self.live.__exit__(*args)
def update(self, run: HarnessRun) -> None:
from rich.markup import escape
from rich.panel import Panel
active = run.current_nodeid or "Waiting for test events…"
if len(active) > max(40, self.console.width - 16):
active = f"…{active[-(self.console.width - 17):]}"
passed, failed, errors, skipped = _summary(run)
progress = (
f"[bold]{run.completed_tests}/{run.unique_tests}[/bold] tests "
f"[green]{passed} passed[/green] [red]{failed + errors} failed[/red] "
f"[yellow]{skipped} skipped[/yellow] [dim]{_format_duration(run.duration)}[/dim]"
)
legend = "✓ pass ✗ fail ! error ↷ skip\n? configured test missing — planned ◐ partial coverage"
self.live.update(
Panel(
self._table(run),
title="⚡ Rust ↔ Python parity lab",
subtitle=f"{progress}\n[dim]{escape(active)}[/dim]\n{legend}",
border_style="cyan",
)
)
def finish(self, run: HarnessRun, exit_code: int) -> None:
self.update(run)
if run.failures:
from rich.markup import escape
from rich.panel import Panel
for nodeid, detail in run.failures[:5]:
rerun = _rerun_command(nodeid)
self.console.print(
Panel(
f"{escape(detail)}\n\n[bold]Rerun just this test[/bold]\n"
f"[cyan]{escape(rerun)}[/cyan]",
title=f"✗ {escape(nodeid)}",
border_style="red",
)
)
durations: dict[str, float] = {}
for result in run.results.values():
for nodeid, duration in result.durations.items():
durations[nodeid] = max(duration, durations.get(nodeid, 0.0))
if durations:
slow = sorted(durations.items(), key=lambda item: item[1], reverse=True)[:3]
self.console.print(
"[bold]Slowest tests[/bold] "
+ " • ".join(
f"{Path(nodeid).name} [dim]{_format_duration(duration)}[/dim]"
for nodeid, duration in slow
)
)
from rich import box
from rich.table import Table
confidence_table = Table(
title="Port confidence by SDK section", box=box.ROUNDED, expand=True
)
confidence_table.add_column("SDK section")
confidence_table.add_column("Score", justify="right")
confidence_table.add_column("Confidence")
confidence_table.add_column("Strategy evidence", ratio=4)
confidence_styles = {"HIGH": "green", "MEDIUM": "yellow", "LOW": "red"}
for score in section_confidence(run, self.confidence_strategies):
confidence_table.add_row(
score.sdk_function,
f"{score.verified_strategies}/{score.required_strategies} {score.percentage}%",
f"[{confidence_styles[score.level.value]}]{score.level.value}[/]",
" ".join(score.details),
)
self.console.print(confidence_table)
self.console.print(
"[dim]Score = required strategies with passing evidence. "
"LOC coverage remains a separate report.[/dim]"
)
style = "green" if exit_code == 0 else "red"
self.console.print(
f"[{style}]Harness finished in {_format_duration(run.duration)} "
f"(exit {exit_code})[/{style}]"
)
class PlainDashboard(AbstractContextManager["PlainDashboard"]):
def __init__(
self,
strategies: Sequence[Strategy],
confidence_strategies: Sequence[Strategy],
) -> None:
self.strategies = strategies
self.confidence_strategies = confidence_strategies
self._seen: dict[str, tuple[RunStatus, int]] = {}
def __enter__(self) -> "PlainDashboard":
print("Rust <-> Python SDK parity harness", flush=True)
return self
def __exit__(self, *args: object) -> None:
return None
def update(self, run: HarnessRun) -> None:
for key, result in run.results.items():
state = (result.status, len(result.completed))
if self._seen.get(key) != state:
self._seen[key] = state
progress = (
f" {len(result.completed)}/{result.total}" if result.total else ""
)
print(
f"{STATUS_GLYPHS[result.status]} {key}: {result.status.value}{progress}",
flush=True,
)
def finish(self, run: HarnessRun, exit_code: int) -> None:
self.update(run)
passed, failed, errors, skipped = _summary(run)
print(
f"Summary: {passed} passed, {failed} failed, {errors} errors, "
f"{skipped} skipped in {_format_duration(run.duration)}",
flush=True,
)
for nodeid, _ in run.failures[:5]:
print(f"Rerun: {_rerun_command(nodeid)}", flush=True)
print("Port confidence by SDK section", flush=True)
for score in section_confidence(run, self.confidence_strategies):
print(
f" {score.sdk_function:12} "
f"{score.verified_strategies}/{score.required_strategies} "
f"{score.percentage:3}% {score.level.value:6} "
f"{' | '.join(score.details)}",
flush=True,
)
print(
" Score = required strategies with passing evidence; LOC is reported separately.",
flush=True,
)
print(f"Harness finished with exit code {exit_code}", flush=True)
def make_dashboard(
strategies: Sequence[Strategy],
plain: bool = False,
confidence_strategies: Sequence[Strategy] | None = None,
) -> RichDashboard | PlainDashboard:
confidence_strategies = confidence_strategies or strategies
interactive_terminal = (
sys.stdout.isatty()
and not os.environ.get("CI")
and os.environ.get("TERM") != "dumb"
)
if not plain and interactive_terminal:
try:
import rich # noqa: F401
return RichDashboard(strategies, confidence_strategies)
except ImportError:
pass
return PlainDashboard(strategies, confidence_strategies)

View file

@ -0,0 +1,3 @@
# Rust unit tests
Holds focused Cargo tests for Rust-owned parsing, transforms, errors, and streaming behavior. These tests make failures fast to diagnose before the Python bridge or full SDK path is involved.

View file

@ -0,0 +1,12 @@
{
"order": 20,
"id": "unit_tests_rust",
"label": "Rust unit tests",
"description": "Exercise Rust-owned behavior directly with focused unit tests.",
"functions": {
"ocr": {"coverage": "planned", "selectors": []},
"messages": {"coverage": "planned", "selectors": []},
"responses": {"coverage": "planned", "selectors": []},
"count_tokens": {"coverage": "planned", "selectors": []}
}
}

View file

@ -0,0 +1,3 @@
# Validate sub-methods
Checks each request, response, stream, and error-mapping sub-method independently across Python and Rust. It also validates that traced Python helpers have an explicit Rust implementation and parity test.

View file

@ -0,0 +1,12 @@
{
"order": 30,
"id": "validate_sub_methods",
"label": "Validate sub-methods",
"description": "Compare isolated transforms and verify Python-to-Rust helper coverage.",
"functions": {
"ocr": {"coverage": "planned", "selectors": []},
"messages": {"coverage": "planned", "selectors": []},
"responses": {"coverage": "planned", "selectors": []},
"count_tokens": {"coverage": "planned", "selectors": []}
}
}

Some files were not shown because too many files have changed in this diff Show more