mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
wip
This commit is contained in:
parent
c8c916e549
commit
6276022f34
105 changed files with 2955 additions and 2672 deletions
3
Makefile
3
Makefile
|
|
@ -299,6 +299,9 @@ test-rust-extension:
|
|||
[ "$$#" -eq 1 ] && \
|
||||
UV_PROJECT_ENVIRONMENT="$$temporary/venv" $(UV) sync --python 3.12 --frozen --no-install-project --all-groups --all-extras && \
|
||||
$(UV) pip install --python "$$temporary/venv/bin/python" --no-deps "$$1" && \
|
||||
"$$temporary/venv/bin/python" -I -m mypy.stubtest \
|
||||
--mypy-config-file tests/test_litellm/rust_bridge/stubtest.ini \
|
||||
litellm.rust_bridge._native && \
|
||||
LITELLM_RUST=1 LITELLM_LOCAL_MODEL_COST_MAP=True \
|
||||
"$$temporary/venv/bin/python" -I -m pytest --import-mode=importlib -m requires_rust_extension tests/test_litellm_rust
|
||||
|
||||
|
|
|
|||
|
|
@ -17,5 +17,25 @@ pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Resu
|
|||
.await
|
||||
}
|
||||
|
||||
|
||||
pub fn admit(
|
||||
model: &str,
|
||||
provider: Option<&str>,
|
||||
audio: &Value,
|
||||
) -> Result<(), crate::call_lifecycle::admission::AdmissionDecline> {
|
||||
use crate::call_lifecycle::admission::AdmissionDecline;
|
||||
let resolved = crate::routing_utils::provider::get_custom_llm_provider(model, provider);
|
||||
let provider = provider.or_else(|| resolved.as_ref().map(|value| value.custom_llm_provider));
|
||||
if provider.and_then(prepare::provider_config).is_none() {
|
||||
return Err(AdmissionDecline::Provider);
|
||||
}
|
||||
if let Some(format) = audio.get("format").and_then(Value::as_str)
|
||||
&& !matches!(format, "wav" | "mp3" | "flac" | "ogg")
|
||||
{
|
||||
return Err(AdmissionDecline::Feature("unsupported audio format"));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
|
|
|||
|
|
@ -8,7 +8,9 @@ use super::transformation::{AudioTranscriptionAuth, AudioTranscriptionProviderCo
|
|||
use super::types::{AudioTranscriptionRequest, ProviderAudioTranscriptionRequest};
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn provider_config(provider: &str) -> Option<&'static dyn AudioTranscriptionProviderConfig> {
|
||||
pub(super) fn provider_config(
|
||||
provider: &str,
|
||||
) -> Option<&'static dyn AudioTranscriptionProviderConfig> {
|
||||
#[cfg(feature = "bedrock-auth")]
|
||||
if provider == "bedrock" {
|
||||
return Some(&BEDROCK_AUDIO_TRANSCRIPTION_CONFIG);
|
||||
|
|
|
|||
|
|
@ -18,3 +18,13 @@ pub enum UnimplementedRoute {
|
|||
pub fn admit_unimplemented(route: UnimplementedRoute) -> Result<Infallible, UnimplementedRoute> {
|
||||
Err(route)
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq, strum::Display)]
|
||||
pub enum AdmissionDecline {
|
||||
#[strum(to_string = "provider is not supported by this native route")]
|
||||
Provider,
|
||||
#[strum(to_string = "required host operations are not supported")]
|
||||
HostOperations,
|
||||
#[strum(to_string = "{0}")]
|
||||
Feature(&'static str),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -55,5 +55,55 @@ pub fn chat_completions_decline_reason(
|
|||
.map(|reason| reason.0)
|
||||
}
|
||||
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, serde::Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct AdmissionContext {
|
||||
#[serde(default)]
|
||||
pub stream: bool,
|
||||
#[serde(default)]
|
||||
pub anthropic_user_id: bool,
|
||||
#[serde(default)]
|
||||
pub bedrock_metadata_owned: bool,
|
||||
}
|
||||
|
||||
pub fn admit(
|
||||
model: &str,
|
||||
provider: Option<&str>,
|
||||
messages: Value,
|
||||
params: &Map<String, Value>,
|
||||
headers: Option<&Map<String, Value>>,
|
||||
context: AdmissionContext,
|
||||
) -> Result<(), crate::call_lifecycle::admission::AdmissionDecline> {
|
||||
use crate::call_lifecycle::admission::AdmissionDecline;
|
||||
let resolved = crate::routing_utils::provider::get_custom_llm_provider(model, provider);
|
||||
let provider = provider.or_else(|| resolved.as_ref().map(|value| value.custom_llm_provider));
|
||||
if context.stream {
|
||||
return Err(AdmissionDecline::Feature("streaming"));
|
||||
}
|
||||
if (provider == Some("anthropic") && context.anthropic_user_id)
|
||||
|| (provider == Some("bedrock") && context.bedrock_metadata_owned)
|
||||
{
|
||||
return Err(AdmissionDecline::HostOperations);
|
||||
}
|
||||
#[cfg(feature = "bedrock-auth")]
|
||||
if provider == Some("bedrock")
|
||||
&& headers.is_some_and(|headers| {
|
||||
headers
|
||||
.keys()
|
||||
.any(|name| crate::providers::bedrock::aws_base::is_sigv4_computed_header(name))
|
||||
})
|
||||
{
|
||||
return Err(AdmissionDecline::Feature(
|
||||
"request forwards a header AWS SigV4 computes",
|
||||
));
|
||||
}
|
||||
let _ = headers;
|
||||
match chat_completions_decline_reason(model, provider, messages, params) {
|
||||
Some(reason) => Err(AdmissionDecline::Feature(reason)),
|
||||
None => Ok(()),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
|
|
|||
|
|
@ -27,5 +27,26 @@ pub async fn messages_stream(request: MessagesRequest<'_>) -> Result<reqwest::Re
|
|||
execute_messages_provider_stream(request).await
|
||||
}
|
||||
|
||||
|
||||
pub fn admit(
|
||||
model: &str,
|
||||
provider: Option<&str>,
|
||||
has_agentic_hook: bool,
|
||||
) -> Result<(), crate::call_lifecycle::admission::AdmissionDecline> {
|
||||
use crate::call_lifecycle::admission::AdmissionDecline;
|
||||
let resolved = crate::routing_utils::provider::get_custom_llm_provider(model, provider);
|
||||
let provider = provider.or_else(|| resolved.as_ref().map(|value| value.custom_llm_provider));
|
||||
if provider
|
||||
.and_then(common_utils::messages_provider_config)
|
||||
.is_none()
|
||||
{
|
||||
return Err(AdmissionDecline::Provider);
|
||||
}
|
||||
if has_agentic_hook {
|
||||
return Err(AdmissionDecline::HostOperations);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
|
|
|||
|
|
@ -19,6 +19,15 @@ pub use lifecycle::{
|
|||
};
|
||||
pub use types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument};
|
||||
|
||||
pub fn admit_value(
|
||||
model: &str,
|
||||
provider: Option<&str>,
|
||||
) -> Result<(), crate::call_lifecycle::admission::AdmissionDecline> {
|
||||
registry::resolve_wire_adapter(model, provider)
|
||||
.map(|_| ())
|
||||
.map_err(|_| crate::call_lifecycle::admission::AdmissionDecline::Provider)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
#[path = "../../tests/azure_ai_ocr.rs"]
|
||||
mod azure_ai_tests;
|
||||
|
|
|
|||
|
|
@ -335,3 +335,12 @@ mod tests {
|
|||
assert!(!nested_without_flat.data.contains_key("model"));
|
||||
}
|
||||
}
|
||||
|
||||
pub fn admit(
|
||||
provider: Option<&str>,
|
||||
) -> Result<(), crate::call_lifecycle::admission::AdmissionDecline> {
|
||||
match provider {
|
||||
Some("openai") => Ok(()),
|
||||
_ => Err(crate::call_lifecycle::admission::AdmissionDecline::Provider),
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -9,6 +9,18 @@ pyo3::create_exception!(
|
|||
"Core admission declined without effects, so the host may select its legacy path once."
|
||||
);
|
||||
|
||||
pyo3::create_exception!(
|
||||
_native,
|
||||
RustHostCallbackError,
|
||||
pyo3::exceptions::PyException
|
||||
);
|
||||
|
||||
pyo3::create_exception!(
|
||||
_native,
|
||||
RustBridgeUnavailable,
|
||||
pyo3::exceptions::PyException
|
||||
);
|
||||
|
||||
pyo3::create_exception!(
|
||||
_native,
|
||||
RustUpstreamError,
|
||||
|
|
@ -31,7 +43,7 @@ pub(crate) fn core_error_to_pyerr(err: Error) -> PyErr {
|
|||
pub(crate) fn execution_error_to_pyerr(error: Error) -> PyErr {
|
||||
match error {
|
||||
Error::Http { status, body } => RustUpstreamError::new_err((status, body)),
|
||||
Error::Network(message) | Error::InvalidResponse(message) => {
|
||||
Error::Connect(message) | Error::Network(message) | Error::InvalidResponse(message) => {
|
||||
RustUpstreamError::new_err((0u16, message))
|
||||
}
|
||||
other => core_error_to_pyerr(other),
|
||||
|
|
@ -40,10 +52,25 @@ pub(crate) fn execution_error_to_pyerr(error: Error) -> PyErr {
|
|||
|
||||
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
let py = module.py();
|
||||
module.add(
|
||||
"RustBridgeUnavailable",
|
||||
py.get_type::<RustBridgeUnavailable>(),
|
||||
)?;
|
||||
module.add("RustBridgeDeclined", py.get_type::<RustBridgeDeclined>())?;
|
||||
module.add(
|
||||
"RustHostCallbackError",
|
||||
py.get_type::<RustHostCallbackError>(),
|
||||
)?;
|
||||
module.add("RustUpstreamError", py.get_type::<RustUpstreamError>())
|
||||
}
|
||||
|
||||
pub(crate) fn host_callback_error(py: Python<'_>, error: PyErr) -> PyErr {
|
||||
let wrapped = RustHostCallbackError::new_err(error.to_string());
|
||||
wrapped.set_context(py, Some(error.clone_ref(py)));
|
||||
wrapped.set_cause(py, Some(error));
|
||||
wrapped
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
|
@ -70,7 +97,6 @@ mod tests {
|
|||
Error::MissingAzureDocumentIntelligenceCredentials,
|
||||
Error::MissingReductoApiKey,
|
||||
Error::Routing("routing failed".into()),
|
||||
Error::Connect("connection refused".into()),
|
||||
] {
|
||||
let expected = core_error_to_pyerr(error.clone());
|
||||
let actual = execution_error_to_pyerr(error);
|
||||
|
|
@ -86,6 +112,7 @@ mod tests {
|
|||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
for (error, status, message) in [
|
||||
(Error::Connect("connection refused".into()), 0, "connection refused"),
|
||||
(
|
||||
Error::Http {
|
||||
status: 429,
|
||||
|
|
@ -114,3 +141,9 @@ mod tests {
|
|||
});
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn admit(
|
||||
result: Result<(), litellm_core::call_lifecycle::admission::AdmissionDecline>,
|
||||
) -> PyResult<()> {
|
||||
result.map_err(|reason| RustBridgeDeclined::new_err(reason.to_string()))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -8,7 +8,6 @@ mod function_trace;
|
|||
mod lifecycle;
|
||||
mod marshal;
|
||||
mod routes;
|
||||
mod token_counter;
|
||||
|
||||
use pyo3::prelude::*;
|
||||
|
||||
|
|
@ -20,7 +19,6 @@ mod _native {
|
|||
fn init(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
super::errors::register(module)?;
|
||||
super::routes::register(module)?;
|
||||
super::token_counter::register(module)?;
|
||||
super::diagnostics::register(module)
|
||||
}
|
||||
}
|
||||
|
|
@ -44,7 +42,9 @@ mod tests {
|
|||
let module = pyo3::wrap_pymodule!(_native)(py).into_bound(py);
|
||||
|
||||
let expected = [
|
||||
"RustBridgeUnavailable",
|
||||
"RustBridgeDeclined",
|
||||
"RustHostCallbackError",
|
||||
"RustUpstreamError",
|
||||
"ocr",
|
||||
"aocr",
|
||||
|
|
@ -52,11 +52,10 @@ mod tests {
|
|||
"atranscription",
|
||||
"messages",
|
||||
"amessages",
|
||||
"chat_completions_decline",
|
||||
"chat_completions",
|
||||
"achat_completions",
|
||||
"ResponsesWebSocketConnection",
|
||||
"TokenCounter",
|
||||
"count_input_tokens",
|
||||
"gil_stats",
|
||||
];
|
||||
|
||||
|
|
@ -148,7 +147,7 @@ mod tests {
|
|||
import asyncio
|
||||
|
||||
async def exercise():
|
||||
connection = await native.ResponsesWebSocketConnection.connect(url)
|
||||
connection = await native.ResponsesWebSocketConnection.connect(url, custom_llm_provider="openai")
|
||||
assert type(connection) is native.ResponsesWebSocketConnection
|
||||
await connection.send_text("from-python")
|
||||
assert await connection.recv_text() == "from-server"
|
||||
|
|
|
|||
|
|
@ -2,9 +2,7 @@ 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 litellm_core::chat_completions::{AdmissionContext, chat_completions as run_chat_completions};
|
||||
use pyo3::prelude::*;
|
||||
use serde_json::Value;
|
||||
|
||||
|
|
@ -16,6 +14,12 @@ fn prepare_chat_completions(
|
|||
) -> PyResult<impl Future<Output = Result<ChatCompletionsResponse, Error>> + Send + 'static> {
|
||||
let messages = required_array("messages", inputs.messages)?;
|
||||
let optional_params = object_or_empty("optional_params", inputs.optional_params)?;
|
||||
let context: AdmissionContext = inputs
|
||||
.host_facts
|
||||
.map(serde_json::from_value)
|
||||
.transpose()
|
||||
.map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string()))?
|
||||
.unwrap_or_default();
|
||||
let options = RouteOptions::from_python(RouteOptionsInputs {
|
||||
model: inputs.model,
|
||||
api_key: inputs.api_key,
|
||||
|
|
@ -25,6 +29,23 @@ fn prepare_chat_completions(
|
|||
timeout_seconds: inputs.timeout_seconds,
|
||||
})?;
|
||||
|
||||
crate::errors::admit(litellm_core::chat_completions::admit(
|
||||
&options.model,
|
||||
options.custom_llm_provider.as_deref(),
|
||||
Value::Array(messages.clone()),
|
||||
&optional_params,
|
||||
options.extra_headers.as_ref(),
|
||||
context,
|
||||
))?;
|
||||
if let Some(on_request) = inputs.on_request {
|
||||
Python::attach(|py| {
|
||||
on_request
|
||||
.call0(py)
|
||||
.map(|_| ())
|
||||
.map_err(|error| crate::errors::host_callback_error(py, error))
|
||||
})?;
|
||||
}
|
||||
|
||||
Ok(async move {
|
||||
let RouteOptions {
|
||||
model,
|
||||
|
|
@ -48,24 +69,6 @@ fn prepare_chat_completions(
|
|||
})
|
||||
}
|
||||
|
||||
#[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,
|
||||
|
|
@ -84,8 +87,10 @@ bridge_route! {
|
|||
#[pyo3(from_py_with = litellm_python_interop::from_py)]
|
||||
extra_headers: Option<serde_json::Value>,
|
||||
timeout_seconds: Option<f64>,
|
||||
#[pyo3(from_py_with = litellm_python_interop::from_py)]
|
||||
host_facts: Option<Value>,
|
||||
on_request: Option<Py<PyAny>>,
|
||||
},
|
||||
prepare = prepare_chat_completions,
|
||||
errors = execution_error_to_pyerr,
|
||||
extra = [chat_completions_decline],
|
||||
}
|
||||
|
|
|
|||
|
|
@ -5,17 +5,18 @@ use pyo3::types::PyCFunction;
|
|||
macro_rules! unimplemented_lifecycle_route {
|
||||
($route:ident, $entrypoint:ident) => {
|
||||
#[pyo3::pyfunction]
|
||||
#[pyo3(signature = (request, args, kwargs, asynchronous))]
|
||||
#[pyo3(signature = (request, args, kwargs, asynchronous, host))]
|
||||
fn $entrypoint(
|
||||
request: pyo3::Bound<'_, pyo3::PyAny>,
|
||||
args: pyo3::Bound<'_, pyo3::types::PyTuple>,
|
||||
kwargs: pyo3::Bound<'_, pyo3::types::PyDict>,
|
||||
asynchronous: bool,
|
||||
host: pyo3::Bound<'_, pyo3::PyAny>,
|
||||
) -> pyo3::PyResult<pyo3::Py<pyo3::PyAny>> {
|
||||
use litellm_core::call_lifecycle::admission::{
|
||||
UnimplementedRoute, admit_unimplemented,
|
||||
};
|
||||
let _ = (request, args, kwargs, asynchronous);
|
||||
let _ = (request, args, kwargs, asynchronous, host);
|
||||
match admit_unimplemented(UnimplementedRoute::$route) {
|
||||
Ok(never) => match never {},
|
||||
Err(route) => Err($crate::errors::RustBridgeDeclined::new_err(format!(
|
||||
|
|
@ -268,12 +269,12 @@ mod tests {
|
|||
(
|
||||
"messages",
|
||||
"amessages",
|
||||
"(model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None)",
|
||||
"(model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None, has_agentic_hook=None)",
|
||||
),
|
||||
(
|
||||
"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)",
|
||||
"(model, messages, optional_params=None, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None, host_facts=None, on_request=None)",
|
||||
),
|
||||
];
|
||||
|
||||
|
|
@ -458,46 +459,6 @@ mod tests {
|
|||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn chat_completions_decline_keeps_existing_reasons() {
|
||||
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 decline = module
|
||||
.getattr("chat_completions_decline")
|
||||
.expect("decline helper should be registered");
|
||||
let empty = PyList::empty(py);
|
||||
let unreadable = py
|
||||
.eval(c"'nope'", None, None)
|
||||
.expect("string messages should convert");
|
||||
|
||||
let unknown: Option<String> = decline
|
||||
.call1(("unknown-model", &empty))
|
||||
.and_then(|value| value.extract())
|
||||
.expect("unknown providers should decline");
|
||||
assert_eq!(
|
||||
unknown.as_deref(),
|
||||
Some("provider is not on the rust chat completions path")
|
||||
);
|
||||
|
||||
let empty_reason: Option<String> = decline
|
||||
.call1(("anthropic/claude-sonnet-4-5", &empty))
|
||||
.and_then(|value| value.extract())
|
||||
.expect("empty lists should decline");
|
||||
assert_eq!(empty_reason.as_deref(), Some("empty message list"));
|
||||
|
||||
let unreadable_reason: Option<String> = decline
|
||||
.call1(("anthropic/claude-sonnet-4-5", unreadable))
|
||||
.and_then(|value| value.extract())
|
||||
.expect("non-list messages should decline");
|
||||
assert_eq!(
|
||||
unreadable_reason.as_deref(),
|
||||
Some("unreadable message list")
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generated_routes_execute_sync_and_async_contracts() {
|
||||
Python::initialize();
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ use pyo3::prelude::*;
|
|||
use serde_json::Value;
|
||||
use std::future::Future;
|
||||
|
||||
use crate::errors::core_error_to_pyerr;
|
||||
use crate::errors::{admit, execution_error_to_pyerr};
|
||||
use crate::marshal::{RouteOptions, RouteOptionsInputs, required_object};
|
||||
|
||||
fn prepare_messages(
|
||||
|
|
@ -21,6 +21,12 @@ fn prepare_messages(
|
|||
timeout_seconds: inputs.timeout_seconds,
|
||||
})?;
|
||||
|
||||
admit(litellm_core::messages::admit(
|
||||
&options.model,
|
||||
options.custom_llm_provider.as_deref(),
|
||||
inputs.has_agentic_hook.unwrap_or(false),
|
||||
))?;
|
||||
|
||||
Ok(async move {
|
||||
let RouteOptions {
|
||||
model,
|
||||
|
|
@ -59,7 +65,8 @@ bridge_route! {
|
|||
#[pyo3(from_py_with = litellm_python_interop::from_py)]
|
||||
extra_headers: Option<serde_json::Value>,
|
||||
timeout_seconds: Option<f64>,
|
||||
has_agentic_hook: Option<bool>,
|
||||
},
|
||||
prepare = prepare_messages,
|
||||
errors = core_error_to_pyerr,
|
||||
errors = execution_error_to_pyerr,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ mod ocr;
|
|||
mod rerank;
|
||||
mod responses;
|
||||
mod speech;
|
||||
mod token_counter;
|
||||
mod transcription;
|
||||
|
||||
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
|
|
@ -27,6 +28,7 @@ pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
|||
rerank::register(module)?;
|
||||
responses::register(module)?;
|
||||
speech::register(module)?;
|
||||
token_counter::register(module)?;
|
||||
|
||||
#[cfg(feature = "trace-parity")]
|
||||
{
|
||||
|
|
|
|||
|
|
@ -29,6 +29,7 @@ impl PythonLogger {
|
|||
pub(super) fn update_ocr(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
host: &Py<PyAny>,
|
||||
kwargs: &Py<PyDict>,
|
||||
pre_call: &OcrLoggingFields,
|
||||
secret_fields: &[&str],
|
||||
|
|
@ -58,7 +59,7 @@ impl PythonLogger {
|
|||
params.set_item(name, value)?;
|
||||
}
|
||||
}
|
||||
for name in custom_pricing_fields(py)? {
|
||||
for name in custom_pricing_fields(py, host)? {
|
||||
if let Some(value) = kwargs.bind(py).get_item(&name)?
|
||||
&& !value.is_none()
|
||||
{
|
||||
|
|
@ -127,15 +128,10 @@ impl PythonLogger {
|
|||
}
|
||||
}
|
||||
|
||||
fn custom_pricing_fields(py: Python<'_>) -> PyResult<Vec<String>> {
|
||||
py.import("litellm.types.utils")?
|
||||
.getattr("CustomPricingLiteLLMParams")?
|
||||
.getattr("model_fields")?
|
||||
.cast_into::<PyDict>()?
|
||||
.keys()
|
||||
.iter()
|
||||
.map(|name| name.extract::<String>())
|
||||
.collect()
|
||||
fn custom_pricing_fields(py: Python<'_>, host: &Py<PyAny>) -> PyResult<Vec<String>> {
|
||||
host.bind(py)
|
||||
.call_method0("custom_pricing_fields")?
|
||||
.extract()
|
||||
}
|
||||
|
||||
fn redact(
|
||||
|
|
@ -158,21 +154,26 @@ fn redact(
|
|||
Ok(redacted.unbind())
|
||||
}
|
||||
|
||||
pub(super) fn response(py: Python<'_>, response: &LiteLLMOcrResponse) -> PyResult<Py<PyAny>> {
|
||||
py.import("litellm.rust_bridge.ocr.value")?
|
||||
.getattr("_response")?
|
||||
pub(super) fn response(
|
||||
py: Python<'_>,
|
||||
host: &Py<PyAny>,
|
||||
response: &LiteLLMOcrResponse,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
host.bind(py)
|
||||
.getattr("response")?
|
||||
.call1((to_py(py, response)?,))
|
||||
.map(Bound::unbind)
|
||||
}
|
||||
|
||||
pub(super) fn map_failure(
|
||||
py: Python<'_>,
|
||||
host: &Py<PyAny>,
|
||||
error: &Py<PyBaseException>,
|
||||
request: &Bound<'_, PyAny>,
|
||||
provider: &str,
|
||||
) -> PyResult<Py<PyBaseException>> {
|
||||
Ok(py
|
||||
.import("litellm.rust_bridge.ocr.lifecycle")?
|
||||
Ok(host
|
||||
.bind(py)
|
||||
.getattr("map_failure")?
|
||||
.call1((error, request, provider))?
|
||||
.extract()?)
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ use crate::lifecycle::{
|
|||
|
||||
struct PythonOcrHost {
|
||||
state: PythonCallState,
|
||||
adapter: Py<PyAny>,
|
||||
data: OcrHostData,
|
||||
}
|
||||
|
||||
|
|
@ -92,6 +93,7 @@ impl PythonOcrHost {
|
|||
let pre_call = projected.pre_call.as_ref().ok_or_else(missing_state)?;
|
||||
self.state.logger()?.update_ocr(
|
||||
py,
|
||||
&self.adapter,
|
||||
&self.state.kwargs,
|
||||
pre_call,
|
||||
&projected.fields.secret_fields,
|
||||
|
|
@ -215,7 +217,8 @@ impl PythonRoute for PythonOcrHost {
|
|||
}
|
||||
OcrHostOperation::ConstructResponse(response) => {
|
||||
self.state.end = Some(now(py)?);
|
||||
self.state.response = Some(callbacks::response(py, response.as_ref())?);
|
||||
self.state.response =
|
||||
Some(callbacks::response(py, &self.adapter, response.as_ref())?);
|
||||
OcrHostResult::Lifecycle(Ok(()))
|
||||
}
|
||||
OcrHostOperation::MapFailure(error) => {
|
||||
|
|
@ -234,7 +237,7 @@ impl PythonRoute for PythonOcrHost {
|
|||
),
|
||||
OcrHostData::Released => return Err(missing_state()),
|
||||
};
|
||||
let mapped = callbacks::map_failure(py, error, request, provider)?;
|
||||
let mapped = callbacks::map_failure(py, &self.adapter, error, request, provider)?;
|
||||
self.state
|
||||
.retain_error(py, PyErr::from_value(mapped.into_bound(py).into_any()));
|
||||
OcrHostResult::Lifecycle(Ok(()))
|
||||
|
|
@ -249,6 +252,7 @@ impl PythonRoute for PythonOcrHost {
|
|||
self.data = OcrHostData::Released;
|
||||
}
|
||||
fn traverse(&self, visit: &pyo3::gc::PyVisit<'_>) -> Result<(), pyo3::gc::PyTraverseError> {
|
||||
visit.call(&self.adapter)?;
|
||||
match &self.data {
|
||||
OcrHostData::Unprojected { request } => visit.call(request),
|
||||
OcrHostData::Projected(projected) => {
|
||||
|
|
@ -282,6 +286,7 @@ fn _ocr_lifecycle(
|
|||
args: Bound<'_, PyTuple>,
|
||||
kwargs: Bound<'_, PyDict>,
|
||||
asynchronous: bool,
|
||||
host: Bound<'_, PyAny>,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
let client = OcrClient::shared().map_err(ocr_error_to_pyerr)?;
|
||||
let call = admitted_call(OcrCall::admit(
|
||||
|
|
@ -299,6 +304,7 @@ fn _ocr_lifecycle(
|
|||
asynchronous,
|
||||
if asynchronous { "aocr" } else { "ocr" },
|
||||
)?,
|
||||
adapter: host.unbind(),
|
||||
data: OcrHostData::Unprojected {
|
||||
request: request.unbind(),
|
||||
},
|
||||
|
|
|
|||
|
|
@ -28,6 +28,11 @@ fn prepare_ocr(
|
|||
.map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string()))?
|
||||
.unwrap_or_default();
|
||||
|
||||
crate::errors::admit(litellm_core::ocr::admit_value(
|
||||
&options.model,
|
||||
options.custom_llm_provider.as_deref(),
|
||||
))?;
|
||||
|
||||
Ok(async move {
|
||||
let RouteOptions {
|
||||
model,
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ use pyo3::prelude::*;
|
|||
use pyo3::types::PyAny;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::errors::core_error_to_pyerr;
|
||||
use crate::errors::{admit, execution_error_to_pyerr};
|
||||
use crate::marshal::{marshal_headers, optional_timeout};
|
||||
|
||||
#[pyclass]
|
||||
|
|
@ -14,20 +14,24 @@ struct ResponsesWebSocketConnection {
|
|||
#[pymethods]
|
||||
impl ResponsesWebSocketConnection {
|
||||
#[classmethod]
|
||||
#[pyo3(signature = (url, headers=None, timeout_seconds=None))]
|
||||
#[pyo3(signature = (url, headers=None, timeout_seconds=None, custom_llm_provider=None))]
|
||||
fn connect<'py>(
|
||||
_cls: &Bound<'py, pyo3::types::PyType>,
|
||||
py: Python<'py>,
|
||||
url: String,
|
||||
#[pyo3(from_py_with = litellm_python_interop::from_py)] headers: Option<Value>,
|
||||
timeout_seconds: Option<f64>,
|
||||
custom_llm_provider: Option<String>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
admit(litellm_core::responses::websocket::admit(
|
||||
custom_llm_provider.as_deref(),
|
||||
))?;
|
||||
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)?;
|
||||
.map_err(execution_error_to_pyerr)?;
|
||||
Ok(ResponsesWebSocketConnection { inner })
|
||||
})
|
||||
}
|
||||
|
|
@ -35,21 +39,24 @@ impl ResponsesWebSocketConnection {
|
|||
fn send_text<'py>(&self, py: Python<'py>, text: String) -> PyResult<Bound<'py, PyAny>> {
|
||||
let inner = self.inner.clone();
|
||||
pyo3_async_runtimes::tokio::future_into_py(py, async move {
|
||||
inner.send_text(text).await.map_err(core_error_to_pyerr)
|
||||
inner
|
||||
.send_text(text)
|
||||
.await
|
||||
.map_err(execution_error_to_pyerr)
|
||||
})
|
||||
}
|
||||
|
||||
fn recv_text<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
let inner = self.inner.clone();
|
||||
pyo3_async_runtimes::tokio::future_into_py(py, async move {
|
||||
inner.recv_text().await.map_err(core_error_to_pyerr)
|
||||
inner.recv_text().await.map_err(execution_error_to_pyerr)
|
||||
})
|
||||
}
|
||||
|
||||
fn close<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
let inner = self.inner.clone();
|
||||
pyo3_async_runtimes::tokio::future_into_py(py, async move {
|
||||
inner.close().await.map_err(core_error_to_pyerr)
|
||||
inner.close().await.map_err(execution_error_to_pyerr)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,127 @@
|
|||
use std::collections::HashMap;
|
||||
use std::num::NonZero;
|
||||
use std::sync::{Arc, Mutex, OnceLock};
|
||||
use std::thread::available_parallelism;
|
||||
|
||||
use litellm_python_interop::release_gil;
|
||||
use litellm_token_counter::{
|
||||
CountableRequest, Error, InputTokenCount, TokenCounter as CoreTokenCounter,
|
||||
};
|
||||
use pyo3::exceptions::PyRuntimeError;
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::PyAny;
|
||||
use tokio::sync::Semaphore;
|
||||
|
||||
use crate::constants::TOKEN_COUNT_FALLBACK_PARALLELISM;
|
||||
use crate::errors::{RustBridgeDeclined, RustBridgeUnavailable};
|
||||
use crate::execution::run_async;
|
||||
|
||||
struct CachedCounter {
|
||||
counter: Arc<CoreTokenCounter>,
|
||||
encode_slots: Arc<Semaphore>,
|
||||
}
|
||||
|
||||
static COUNTERS: OnceLock<Mutex<HashMap<&'static str, Arc<CachedCounter>>>> = OnceLock::new();
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (body, kind, encoding, disabled, legacy_accounting, resource_loader))]
|
||||
fn count_input_tokens<'py>(
|
||||
py: Python<'py>,
|
||||
body: &[u8],
|
||||
kind: Option<&str>,
|
||||
encoding: &str,
|
||||
disabled: bool,
|
||||
legacy_accounting: bool,
|
||||
resource_loader: Py<PyAny>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let tokenizer = litellm_token_counter::admit_tokenizer(kind, encoding, disabled, legacy_accounting)
|
||||
.map_err(admission_error_to_pyerr)?;
|
||||
CoreTokenCounter::admit_request(body).map_err(admission_error_to_pyerr)?;
|
||||
let cached = cached_counter(py, tokenizer, resource_loader)?;
|
||||
let body = body.to_vec();
|
||||
run_async(
|
||||
py,
|
||||
async move {
|
||||
let _slot = Arc::clone(&cached.encode_slots)
|
||||
.acquire_owned()
|
||||
.await
|
||||
.map_err(|error| Error::Task(error.to_string()))?;
|
||||
tokio::task::spawn_blocking(move || count_body(&cached.counter, &body))
|
||||
.await
|
||||
.map_err(|error| Error::Task(error.to_string()))?
|
||||
},
|
||||
token_count_error_to_pyerr,
|
||||
)
|
||||
}
|
||||
|
||||
fn cached_counter(
|
||||
py: Python<'_>,
|
||||
tokenizer: &'static str,
|
||||
resource_loader: Py<PyAny>,
|
||||
) -> PyResult<Arc<CachedCounter>> {
|
||||
let counters = COUNTERS.get_or_init(|| Mutex::new(HashMap::new()));
|
||||
if let Some(counter) = counters
|
||||
.lock()
|
||||
.map_err(|error| PyRuntimeError::new_err(error.to_string()))?
|
||||
.get(tokenizer)
|
||||
.cloned()
|
||||
{
|
||||
return Ok(counter);
|
||||
}
|
||||
let resource: String = resource_loader
|
||||
.call1(py, (tokenizer,))
|
||||
.and_then(|value| value.extract(py))
|
||||
.map_err(|error| RustBridgeUnavailable::new_err(error.to_string()))?;
|
||||
let counter = release_gil(py, move || load_counter(tokenizer, &resource)).map_err(token_count_error_to_pyerr)?;
|
||||
let cached = Arc::new(CachedCounter {
|
||||
counter: Arc::new(counter),
|
||||
encode_slots: Arc::new(Semaphore::new(encode_parallelism())),
|
||||
});
|
||||
counters
|
||||
.lock()
|
||||
.map_err(|error| PyRuntimeError::new_err(error.to_string()))?
|
||||
.insert(tokenizer, Arc::clone(&cached));
|
||||
Ok(cached)
|
||||
}
|
||||
|
||||
fn load_counter(tokenizer: &str, resource: &str) -> Result<CoreTokenCounter, Error> {
|
||||
match tokenizer {
|
||||
"anthropic" => CoreTokenCounter::from_json(resource),
|
||||
"cl100k_base" => CoreTokenCounter::from_cl100k_ranks(resource),
|
||||
"o200k_base" => CoreTokenCounter::from_o200k_ranks(resource),
|
||||
_ => Err(Error::UnsupportedTokenizer),
|
||||
}
|
||||
}
|
||||
|
||||
fn encode_parallelism() -> usize {
|
||||
available_parallelism().map_or(TOKEN_COUNT_FALLBACK_PARALLELISM, NonZero::get)
|
||||
}
|
||||
|
||||
fn count_body(counter: &CoreTokenCounter, body: &[u8]) -> Result<InputTokenCount, Error> {
|
||||
let request = CountableRequest::parse(body)?;
|
||||
counter.count_request(&request)
|
||||
}
|
||||
|
||||
fn admission_error_to_pyerr(error: Error) -> PyErr {
|
||||
RustBridgeDeclined::new_err(error.to_string())
|
||||
}
|
||||
|
||||
fn token_count_error_to_pyerr(error: Error) -> PyErr {
|
||||
let message = error.to_string();
|
||||
match error {
|
||||
Error::Load(_) | Error::Ranks(_) | Error::UnicodeClasses => RustBridgeUnavailable::new_err(message),
|
||||
Error::UnsupportedTokenizer
|
||||
| Error::RequestParse(_)
|
||||
| Error::MissingInput
|
||||
| Error::FloatText
|
||||
| Error::ContentBlock
|
||||
| Error::ArrayItems
|
||||
| Error::JsonSerialization(_)
|
||||
| Error::JsonUtf8(_) => RustBridgeDeclined::new_err(message),
|
||||
Error::Encode(_) | Error::Task(_) => PyRuntimeError::new_err(message),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
module.add_function(wrap_pyfunction!(count_input_tokens, module)?)
|
||||
}
|
||||
|
|
@ -7,7 +7,7 @@ use litellm_core::audio_transcription::{
|
|||
use pyo3::prelude::*;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::errors::core_error_to_pyerr;
|
||||
use crate::errors::{admit, execution_error_to_pyerr};
|
||||
use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty};
|
||||
|
||||
fn prepare_transcription(
|
||||
|
|
@ -24,6 +24,12 @@ fn prepare_transcription(
|
|||
})?;
|
||||
let optional_params = object_or_empty("optional_params", inputs.optional_params)?;
|
||||
|
||||
admit(litellm_core::audio_transcription::admit(
|
||||
&options.model,
|
||||
options.custom_llm_provider.as_deref(),
|
||||
&audio,
|
||||
))?;
|
||||
|
||||
Ok(async move {
|
||||
let RouteOptions {
|
||||
model,
|
||||
|
|
@ -67,5 +73,5 @@ bridge_route! {
|
|||
timeout_seconds: Option<f64>,
|
||||
},
|
||||
prepare = prepare_transcription,
|
||||
errors = core_error_to_pyerr,
|
||||
errors = execution_error_to_pyerr,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,105 +0,0 @@
|
|||
use std::num::NonZero;
|
||||
use std::sync::Arc;
|
||||
use std::thread::available_parallelism;
|
||||
|
||||
use litellm_python_interop::release_gil;
|
||||
use litellm_token_counter::{
|
||||
CountableRequest, Error, InputTokenCount, TokenCounter as CoreTokenCounter,
|
||||
};
|
||||
use pyo3::exceptions::{PyRuntimeError, PyValueError};
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::PyAny;
|
||||
use tokio::sync::Semaphore;
|
||||
|
||||
use crate::constants::TOKEN_COUNT_FALLBACK_PARALLELISM;
|
||||
use crate::errors::RustBridgeDeclined;
|
||||
use crate::execution::run_async;
|
||||
|
||||
/// Counts the input tokens of a raw request body off the Python event loop with
|
||||
/// the GIL released. Python owns which requests get here and what to do with
|
||||
/// the count. At most one encode per core runs at a time; the rest wait in the
|
||||
/// async task, where a cancelled Python awaiter drops them before any blocking
|
||||
/// work is scheduled.
|
||||
#[pyclass(frozen)]
|
||||
struct TokenCounter {
|
||||
inner: Arc<CoreTokenCounter>,
|
||||
encode_slots: Arc<Semaphore>,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl TokenCounter {
|
||||
#[new]
|
||||
fn new(py: Python<'_>, tokenizer_json: &str) -> PyResult<Self> {
|
||||
Self::load(py, || CoreTokenCounter::from_json(tokenizer_json))
|
||||
}
|
||||
|
||||
#[staticmethod]
|
||||
fn from_cl100k_ranks(py: Python<'_>, rank_file: &str) -> PyResult<Self> {
|
||||
Self::load(py, || CoreTokenCounter::from_cl100k_ranks(rank_file))
|
||||
}
|
||||
|
||||
#[staticmethod]
|
||||
fn from_o200k_ranks(py: Python<'_>, rank_file: &str) -> PyResult<Self> {
|
||||
Self::load(py, || CoreTokenCounter::from_o200k_ranks(rank_file))
|
||||
}
|
||||
|
||||
fn acount_request<'py>(&self, py: Python<'py>, body: &[u8]) -> PyResult<Bound<'py, PyAny>> {
|
||||
let counter = Arc::clone(&self.inner);
|
||||
let encode_slots = Arc::clone(&self.encode_slots);
|
||||
let body = body.to_vec();
|
||||
run_async(
|
||||
py,
|
||||
async move {
|
||||
let _slot = encode_slots
|
||||
.acquire_owned()
|
||||
.await
|
||||
.map_err(|error| Error::Task(error.to_string()))?;
|
||||
tokio::task::spawn_blocking(move || count_body(&counter, &body))
|
||||
.await
|
||||
.map_err(|error| Error::Task(error.to_string()))?
|
||||
},
|
||||
token_count_error_to_pyerr,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
impl TokenCounter {
|
||||
fn load(
|
||||
py: Python<'_>,
|
||||
load: impl FnOnce() -> Result<CoreTokenCounter, Error> + Send,
|
||||
) -> PyResult<Self> {
|
||||
let inner = release_gil(py, load).map_err(token_count_error_to_pyerr)?;
|
||||
Ok(Self {
|
||||
inner: Arc::new(inner),
|
||||
encode_slots: Arc::new(Semaphore::new(encode_parallelism())),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn encode_parallelism() -> usize {
|
||||
available_parallelism().map_or(TOKEN_COUNT_FALLBACK_PARALLELISM, NonZero::get)
|
||||
}
|
||||
|
||||
fn count_body(counter: &CoreTokenCounter, body: &[u8]) -> Result<InputTokenCount, Error> {
|
||||
let request = CountableRequest::parse(body)?;
|
||||
counter.count_request(&request)
|
||||
}
|
||||
|
||||
fn token_count_error_to_pyerr(error: Error) -> PyErr {
|
||||
let message = error.to_string();
|
||||
match error {
|
||||
Error::Load(_) | Error::Ranks(_) | Error::UnicodeClasses => PyValueError::new_err(message),
|
||||
Error::RequestParse(_)
|
||||
| Error::MissingInput
|
||||
| Error::FloatText
|
||||
| Error::ContentBlock
|
||||
| Error::ArrayItems
|
||||
| Error::JsonSerialization(_)
|
||||
| Error::JsonUtf8(_) => RustBridgeDeclined::new_err(message),
|
||||
Error::Encode(_) | Error::Task(_) => PyRuntimeError::new_err(message),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
module.add_class::<TokenCounter>()
|
||||
}
|
||||
|
|
@ -25,6 +25,7 @@ pub struct InputTokenCount {
|
|||
}
|
||||
|
||||
enum Encoder {
|
||||
Admission,
|
||||
HuggingFace {
|
||||
tokenizer: Box<tokenizers::Tokenizer>,
|
||||
byte_level: Option<ByteLevelCounter>,
|
||||
|
|
@ -40,6 +41,15 @@ pub struct TokenCounter {
|
|||
}
|
||||
|
||||
impl TokenCounter {
|
||||
pub fn admit_request(body: &[u8]) -> Result<(), Error> {
|
||||
let request = CountableRequest::parse(body)?;
|
||||
Self {
|
||||
encoder: Encoder::Admission,
|
||||
}
|
||||
.count_request(&request)
|
||||
.map(|_| ())
|
||||
}
|
||||
|
||||
/// Load a HuggingFace `tokenizer.json` document. The host reads the file.
|
||||
pub fn from_json(tokenizer_json: &str) -> Result<Self, Error> {
|
||||
let tokenizer = tokenizer_json
|
||||
|
|
@ -74,6 +84,7 @@ impl TokenCounter {
|
|||
|
||||
pub fn count_text(&self, text: &str) -> Result<usize, Error> {
|
||||
match &self.encoder {
|
||||
Encoder::Admission => Ok(0),
|
||||
Encoder::Tiktoken(counter) => Ok(counter.count(text)),
|
||||
Encoder::HuggingFace {
|
||||
tokenizer,
|
||||
|
|
|
|||
|
|
@ -4,6 +4,8 @@ use thiserror::Error as ThisError;
|
|||
|
||||
#[derive(Debug, ThisError)]
|
||||
pub enum Error {
|
||||
#[error("tokenizer or accounting configuration is not supported")]
|
||||
UnsupportedTokenizer,
|
||||
#[error("failed to load tokenizer: {0}")]
|
||||
Load(#[source] tokenizers::Error),
|
||||
#[error("failed to load tokenizer: tiktoken rank file: {0}")]
|
||||
|
|
|
|||
|
|
@ -19,3 +19,20 @@ mod unicode_classes;
|
|||
pub use counter::{InputTokenCount, TokenCounter};
|
||||
pub use error::Error;
|
||||
pub use types::CountableRequest;
|
||||
|
||||
pub fn admit_tokenizer(
|
||||
kind: Option<&str>,
|
||||
encoding: &str,
|
||||
disabled: bool,
|
||||
legacy_accounting: bool,
|
||||
) -> Result<&'static str, Error> {
|
||||
if disabled {
|
||||
return Err(Error::UnsupportedTokenizer);
|
||||
}
|
||||
match (kind, encoding, legacy_accounting) {
|
||||
(Some("anthropic"), _, _) => Ok("anthropic"),
|
||||
(None, "cl100k_base", false) => Ok("cl100k_base"),
|
||||
(None, "o200k_base", false) => Ok("o200k_base"),
|
||||
_ => Err(Error::UnsupportedTokenizer),
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -911,7 +911,7 @@ openai_compatible_endpoints: Final[list] = [
|
|||
]
|
||||
|
||||
|
||||
openai_compatible_providers: Final[list] = [
|
||||
openai_compatible_providers: Final[list[str]] = [
|
||||
"anyscale",
|
||||
"groq",
|
||||
"nvidia_nim",
|
||||
|
|
|
|||
|
|
@ -26,7 +26,6 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
get_async_httpx_client,
|
||||
)
|
||||
from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge
|
||||
from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts
|
||||
from litellm.types.llms.anthropic import (
|
||||
ContentBlockDelta,
|
||||
ContentBlockStart,
|
||||
|
|
@ -372,27 +371,21 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
transform_params: Final = {**optional_params, "is_vertex_request": is_vertex_request}
|
||||
|
||||
def finish_request(request_data: dict) -> tuple[dict, dict]: # mutable-ok: rewritten in place downstream
|
||||
"""Filter beta headers and emit pre_call, returning `(headers, data)`.
|
||||
|
||||
The pair stays mutable because the streaming path rewrites it in
|
||||
place (`data["stream"] = True`) before sending. A Rust attempt that
|
||||
declined already emitted pre_call for this request, so skip it there.
|
||||
"""
|
||||
"""Filter beta headers and emit pre_call, returning `(headers, data)`."""
|
||||
request_headers, data = update_request_with_filtered_beta(
|
||||
headers=headers,
|
||||
request_data=request_data,
|
||||
provider=custom_llm_provider,
|
||||
)
|
||||
if not serves_via_rust:
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key=api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": api_base,
|
||||
"headers": request_headers,
|
||||
},
|
||||
)
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key=api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": api_base,
|
||||
"headers": request_headers,
|
||||
},
|
||||
)
|
||||
print_verbose(f"_is_function_call: {_is_function_call}")
|
||||
return request_headers, data
|
||||
|
||||
|
|
@ -456,54 +449,30 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
timeout=timeout,
|
||||
)
|
||||
|
||||
# The Rust core owns the whole call for the subset it accepts, so ask
|
||||
# before transforming: whichever path runs emits pre_call exactly once.
|
||||
# `get_config` merges the class-level defaults (Anthropic's required
|
||||
# `max_tokens` among them) that `transform_request` would have applied.
|
||||
rust_optional_params: Final = { # mutable-ok: json.dumps in the bridge rejects a mappingproxy
|
||||
**AnthropicConfig.get_config(model=model),
|
||||
**optional_params,
|
||||
}
|
||||
serves_via_rust: Final = rust_chat_completions_accepts(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
stream=stream,
|
||||
rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict
|
||||
"complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
**rust_optional_params,
|
||||
},
|
||||
"api_base": api_base,
|
||||
"headers": headers,
|
||||
}
|
||||
log_rust_pre_call: Final = lambda: logging_obj.pre_call(
|
||||
input=messages, api_key=api_key, additional_args=rust_logging_args
|
||||
)
|
||||
if serves_via_rust:
|
||||
rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict
|
||||
"complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
**rust_optional_params,
|
||||
},
|
||||
"api_base": api_base,
|
||||
"headers": headers,
|
||||
}
|
||||
logging_obj.pre_call(input=messages, api_key=api_key, additional_args=rust_logging_args)
|
||||
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
|
||||
logging_obj=logging_obj,
|
||||
messages=messages,
|
||||
api_key=api_key,
|
||||
additional_args=rust_logging_args,
|
||||
)
|
||||
if acompletion is True:
|
||||
return rust_chat_completions_bridge.achat_completions_or_fallback(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
on_response=log_rust_post_call,
|
||||
python_fallback=acompletion_dispatch,
|
||||
)
|
||||
rust_response: Final = rust_chat_completions_bridge.chat_completions(
|
||||
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
|
||||
logging_obj=logging_obj,
|
||||
messages=messages,
|
||||
api_key=api_key,
|
||||
additional_args=rust_logging_args,
|
||||
)
|
||||
if acompletion is True:
|
||||
return rust_chat_completions_bridge.achat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
|
|
@ -513,10 +482,30 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
stream=stream,
|
||||
litellm_params=litellm_params,
|
||||
on_request=log_rust_pre_call,
|
||||
on_response=log_rust_post_call,
|
||||
python_fallback=acompletion_dispatch,
|
||||
)
|
||||
if rust_response is not None:
|
||||
return rust_response
|
||||
rust_response: Final = rust_chat_completions_bridge.chat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
stream=stream,
|
||||
litellm_params=litellm_params,
|
||||
on_request=log_rust_pre_call,
|
||||
on_response=log_rust_post_call,
|
||||
python_fallback=lambda: None,
|
||||
)
|
||||
if rust_response is not None:
|
||||
return rust_response
|
||||
|
||||
if acompletion is True:
|
||||
return acompletion_dispatch()
|
||||
|
|
|
|||
|
|
@ -23,8 +23,6 @@ class BedrockAudioTranscriptionRustDispatch:
|
|||
audio_format: Final = formats.get(processed_audio.content_type) or (
|
||||
processed_audio.filename.rsplit(".", 1)[-1].lower() if "." in processed_audio.filename else ""
|
||||
)
|
||||
if audio_format not in {"wav", "mp3", "flac", "ogg"}:
|
||||
raise ValueError(f"Unsupported Bedrock audio format for file {processed_audio.filename!r}")
|
||||
return {
|
||||
"data": base64.b64encode(processed_audio.file_content).decode("ascii"),
|
||||
"format": audio_format,
|
||||
|
|
|
|||
|
|
@ -17,7 +17,6 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
get_async_httpx_client,
|
||||
)
|
||||
from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge
|
||||
from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts
|
||||
from litellm.types.utils import ModelResponse
|
||||
from litellm.utils import CustomStreamWrapper
|
||||
|
||||
|
|
@ -196,7 +195,6 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
headers: dict = {},
|
||||
client: AsyncHTTPHandler | None = None,
|
||||
api_key: str | None = None,
|
||||
skip_pre_call_logging: bool = False,
|
||||
) -> ModelResponse | CustomStreamWrapper:
|
||||
request_data: Final = await litellm.AmazonConverseConfig()._async_transform_request(
|
||||
model=model,
|
||||
|
|
@ -222,16 +220,15 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
# The Rust path already logged this request's pre_call before handing
|
||||
# it here, and it only declines before the provider is called, so this
|
||||
# is the same attempt continuing rather than a second one.
|
||||
if not skip_pre_call_logging:
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": api_base,
|
||||
"headers": prepped.headers,
|
||||
},
|
||||
)
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": api_base,
|
||||
"headers": prepped.headers,
|
||||
},
|
||||
)
|
||||
|
||||
headers = dict(prepped.headers)
|
||||
if client is None or not isinstance(client, AsyncHTTPHandler):
|
||||
|
|
@ -401,43 +398,30 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
# Filter beta headers in HTTP headers before making the request
|
||||
headers = update_headers_with_filtered_beta(headers=headers, provider="bedrock_converse")
|
||||
|
||||
# The Rust core owns the whole call for the subset it accepts. Ask
|
||||
# before transforming so whichever path runs emits pre_call once, and
|
||||
# hand down the credentials, region and endpoint this handler already
|
||||
# resolved so both paths sign as the same principal. Bearer-token auth
|
||||
# resolves no SigV4 principal at all, and each path reads that token
|
||||
# itself.
|
||||
rust_optional_params: Final = { # mutable-ok: json.dumps in the bridge rejects a mappingproxy
|
||||
**optional_params,
|
||||
**_sigv4_principal(credentials),
|
||||
"aws_region_name": aws_region_name,
|
||||
}
|
||||
serves_via_rust: Final = rust_chat_completions_accepts(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
custom_llm_provider="bedrock",
|
||||
litellm_params=litellm_params,
|
||||
stream=stream,
|
||||
rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict
|
||||
"complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent
|
||||
"messages": messages,
|
||||
**optional_params,
|
||||
},
|
||||
"api_base": proxy_endpoint_url,
|
||||
"headers": headers,
|
||||
}
|
||||
log_rust_pre_call: Final = lambda: logging_obj.pre_call(
|
||||
input=messages, api_key="", additional_args=rust_logging_args
|
||||
)
|
||||
if serves_via_rust:
|
||||
rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict
|
||||
"complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent
|
||||
"messages": messages,
|
||||
**optional_params,
|
||||
},
|
||||
"api_base": proxy_endpoint_url,
|
||||
"headers": headers,
|
||||
}
|
||||
logging_obj.pre_call(input=messages, api_key="", additional_args=rust_logging_args)
|
||||
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
|
||||
logging_obj=logging_obj,
|
||||
messages=messages,
|
||||
api_key="",
|
||||
additional_args=rust_logging_args,
|
||||
)
|
||||
if acompletion:
|
||||
return rust_chat_completions_bridge.achat_completions_or_fallback(
|
||||
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
|
||||
logging_obj=logging_obj,
|
||||
messages=messages,
|
||||
api_key="",
|
||||
additional_args=rust_logging_args,
|
||||
)
|
||||
if acompletion:
|
||||
return rust_chat_completions_bridge.achat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
|
|
@ -447,6 +431,9 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
custom_llm_provider="bedrock",
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
stream=stream,
|
||||
litellm_params=litellm_params,
|
||||
on_request=log_rust_pre_call,
|
||||
on_response=log_rust_post_call,
|
||||
python_fallback=lambda: self.async_completion(
|
||||
model=model,
|
||||
|
|
@ -464,23 +451,26 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
client=client,
|
||||
credentials=credentials,
|
||||
api_key=api_key,
|
||||
skip_pre_call_logging=True,
|
||||
),
|
||||
)
|
||||
rust_response: Final = rust_chat_completions_bridge.chat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
api_base=proxy_endpoint_url,
|
||||
custom_llm_provider="bedrock",
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
on_response=log_rust_post_call,
|
||||
)
|
||||
if rust_response is not None:
|
||||
return rust_response
|
||||
rust_response: Final = rust_chat_completions_bridge.chat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
api_base=proxy_endpoint_url,
|
||||
custom_llm_provider="bedrock",
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
stream=stream,
|
||||
litellm_params=litellm_params,
|
||||
on_request=log_rust_pre_call,
|
||||
on_response=log_rust_post_call,
|
||||
python_fallback=lambda: None,
|
||||
)
|
||||
if rust_response is not None:
|
||||
return rust_response
|
||||
|
||||
### ROUTING (ASYNC, STREAMING, SYNC)
|
||||
if acompletion:
|
||||
|
|
@ -548,21 +538,15 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
)
|
||||
|
||||
## LOGGING
|
||||
# Reaching here with `serves_via_rust` set means the synchronous Rust
|
||||
# attempt declined at call time, before the provider was called, and
|
||||
# already logged this request. That is the same attempt continuing.
|
||||
# The asynchronous branch above returns before this point, and hands
|
||||
# its own fallback `skip_pre_call_logging=True` for the same reason.
|
||||
if not serves_via_rust:
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": proxy_endpoint_url,
|
||||
"headers": prepped.headers,
|
||||
},
|
||||
)
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": proxy_endpoint_url,
|
||||
"headers": prepped.headers,
|
||||
},
|
||||
)
|
||||
if client is None or isinstance(client, AsyncHTTPHandler):
|
||||
_params: Final = {}
|
||||
if timeout is not None:
|
||||
|
|
|
|||
|
|
@ -162,15 +162,6 @@ from litellm.utils import (
|
|||
async_pre_call_deployment_hook,
|
||||
)
|
||||
|
||||
|
||||
def _rust_responses_websocket_enabled(
|
||||
custom_llm_provider: str | None,
|
||||
) -> bool:
|
||||
from litellm.rust_bridge.configuration import RouteName, rust_enabled
|
||||
|
||||
return custom_llm_provider == "openai" and rust_enabled(RouteName.RESPONSES)
|
||||
|
||||
|
||||
from .http_handler import get_shared_realtime_ssl_context
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -2454,34 +2445,19 @@ class BaseLLMHTTPHandler:
|
|||
request_body: dict,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> AnthropicMessagesResponse | None:
|
||||
if custom_llm_provider not in ("azure_ai", "anthropic"):
|
||||
return None
|
||||
from litellm.rust_bridge.configuration import RouteName, rust_enabled
|
||||
|
||||
if not rust_enabled(RouteName.MESSAGES):
|
||||
return None
|
||||
if has_agentic_hook:
|
||||
return None
|
||||
|
||||
from litellm.rust_bridge import messages as rust_messages_bridge
|
||||
|
||||
upstream_body: Final = {key: value for key, value in request_body.items() if key != "stream"}
|
||||
try:
|
||||
rust_response: Final = await rust_messages_bridge.amessages(
|
||||
model=model,
|
||||
body=upstream_body,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
)
|
||||
except Exception as rust_error: # noqa: BLE001 # rollout-safety fallback: any Rust bridge failure must fall back to the Python path
|
||||
verbose_logger.debug(
|
||||
"Rust Anthropic messages bridge raised %s; falling back to Python path",
|
||||
type(rust_error).__name__,
|
||||
)
|
||||
return None
|
||||
rust_response: Final = await rust_messages_bridge.amessages(
|
||||
model=model,
|
||||
body=upstream_body,
|
||||
has_agentic_hook=has_agentic_hook,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
)
|
||||
if rust_response is None:
|
||||
return None
|
||||
|
||||
|
|
@ -6657,17 +6633,21 @@ class BaseLLMHTTPHandler:
|
|||
|
||||
@asynccontextmanager
|
||||
async def _backend_connection():
|
||||
if _rust_responses_websocket_enabled(custom_llm_provider):
|
||||
from litellm.rust_bridge.responses import websocket as rust_responses_websocket
|
||||
from litellm.rust_bridge.responses import websocket as rust_responses_websocket
|
||||
|
||||
rust_backend: Final = await rust_responses_websocket.connect(
|
||||
url=ws_url,
|
||||
headers={str(key): str(value) for key, value in headers.items()},
|
||||
timeout=timeout,
|
||||
)
|
||||
if rust_backend is not None:
|
||||
rust_backend: Final = await rust_responses_websocket.connect(
|
||||
url=ws_url,
|
||||
headers={str(key): str(value) for key, value in headers.items()},
|
||||
timeout=timeout,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model=model,
|
||||
)
|
||||
if rust_backend is not None:
|
||||
try:
|
||||
yield rust_backend
|
||||
return
|
||||
finally:
|
||||
await rust_backend.close()
|
||||
return
|
||||
|
||||
async with websockets.connect(
|
||||
ws_url,
|
||||
|
|
|
|||
|
|
@ -4,8 +4,7 @@ from typing import Final, Literal, Protocol, cast # noqa: TID251 # native call
|
|||
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
from litellm.rust_bridge.bindings import NativeBinding
|
||||
from litellm.rust_bridge.configuration import rust_ocr_enabled
|
||||
from litellm.rust_bridge.ocr.definition import COMPONENT
|
||||
|
||||
|
||||
class FileReader(Protocol):
|
||||
|
|
@ -30,7 +29,7 @@ class NativeMimeType(Protocol):
|
|||
def __call__(self, file_name: str) -> str: ...
|
||||
|
||||
|
||||
_FILE_DOCUMENT: Final = NativeBinding(
|
||||
_FILE_DOCUMENT: Final = COMPONENT.bind(
|
||||
"_ocr_file_document",
|
||||
validate=lambda value: (
|
||||
cast( # cast-ok: native export owns the callable signature
|
||||
|
|
@ -40,7 +39,7 @@ _FILE_DOCUMENT: Final = NativeBinding(
|
|||
else None
|
||||
),
|
||||
)
|
||||
_UPLOAD_DOCUMENT: Final = NativeBinding(
|
||||
_UPLOAD_DOCUMENT: Final = COMPONENT.bind(
|
||||
"_ocr_upload_document",
|
||||
validate=lambda value: (
|
||||
cast( # cast-ok: native export owns the callable signature
|
||||
|
|
@ -50,10 +49,10 @@ _UPLOAD_DOCUMENT: Final = NativeBinding(
|
|||
else None
|
||||
),
|
||||
)
|
||||
_MAX_FILE_BYTES: Final = NativeBinding(
|
||||
_MAX_FILE_BYTES: Final = COMPONENT.bind(
|
||||
"_OCR_MAX_FILE_BYTES", validate=lambda value: value if isinstance(value, int) and value > 0 else None
|
||||
)
|
||||
_MIME_TYPE: Final = NativeBinding(
|
||||
_MIME_TYPE: Final = COMPONENT.bind(
|
||||
"_ocr_mime_type",
|
||||
validate=lambda value: (
|
||||
cast( # cast-ok: native export owns the callable signature
|
||||
|
|
@ -67,7 +66,7 @@ _PYTHON_MAX_FILE_BYTES: Final = 50 * 1024 * 1024
|
|||
|
||||
|
||||
def get_mime_type(file_path: str) -> str:
|
||||
native: Final = _MIME_TYPE.load() if rust_ocr_enabled() else None
|
||||
native: Final = COMPONENT.resolve().select(_MIME_TYPE)
|
||||
if native is None:
|
||||
from litellm.ocr import legacy
|
||||
|
||||
|
|
@ -76,14 +75,14 @@ def get_mime_type(file_path: str) -> str:
|
|||
|
||||
|
||||
def get_max_file_bytes() -> int:
|
||||
limit: Final = _MAX_FILE_BYTES.load() if rust_ocr_enabled() else None
|
||||
limit: Final = COMPONENT.resolve().select(_MAX_FILE_BYTES)
|
||||
if limit is None:
|
||||
return _PYTHON_MAX_FILE_BYTES
|
||||
return limit
|
||||
|
||||
|
||||
def convert_file_document_to_url_document(document: FileDocument) -> dict[str, str]:
|
||||
native: Final = _FILE_DOCUMENT.load() if rust_ocr_enabled() else None
|
||||
native: Final = COMPONENT.resolve().select(_FILE_DOCUMENT)
|
||||
if native is None:
|
||||
from litellm.ocr import legacy
|
||||
|
||||
|
|
@ -94,7 +93,7 @@ def convert_file_document_to_url_document(document: FileDocument) -> dict[str, s
|
|||
def convert_upload_to_url_document(
|
||||
file_content: bytes, filename: str | None, content_type: str | None
|
||||
) -> dict[str, str]:
|
||||
native: Final = _UPLOAD_DOCUMENT.load() if rust_ocr_enabled() else None
|
||||
native: Final = COMPONENT.resolve().select(_UPLOAD_DOCUMENT)
|
||||
if native is None:
|
||||
from litellm.ocr import legacy
|
||||
|
||||
|
|
|
|||
|
|
@ -6,10 +6,11 @@ import httpx
|
|||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.ocr import legacy
|
||||
from litellm.ocr.input import convert_file_document_to_url_document, get_mime_type
|
||||
from litellm.rust_bridge.bindings import native_exception_types
|
||||
from litellm.rust_bridge.configuration import rust_ocr_enabled
|
||||
from litellm.rust_bridge.ocr import LiteLLMOcrRequest
|
||||
from litellm.rust_bridge.ocr.definition import COMPONENT
|
||||
from litellm.rust_bridge.ocr.host import HOST
|
||||
from litellm.rust_bridge.ocr.lifecycle import select
|
||||
from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke, invoke
|
||||
|
||||
__all__ = ("aocr", "convert_file_document_to_url_document", "get_mime_type", "ocr")
|
||||
|
||||
|
|
@ -48,32 +49,48 @@ def ocr(
|
|||
**kwargs: object, # kwargs-ok: preserve the public OCR call shape
|
||||
) -> OCRResponse | Coroutine[object, object, OCRResponse]:
|
||||
request: Final = _public_request("ocr", args, kwargs)
|
||||
native: Final = select(request) if rust_ocr_enabled() else None
|
||||
if native is not None:
|
||||
try:
|
||||
return native(request, args, kwargs, False)
|
||||
except _decline_types():
|
||||
pass
|
||||
execution: Final = COMPONENT.resolve()
|
||||
native: Final = select(request, execution)
|
||||
fallback: Final = cast( # cast-ok: forward the original call shape through the legacy @client decorator
|
||||
Callable[..., OCRResponse | Coroutine[object, object, OCRResponse]], legacy.ocr
|
||||
)
|
||||
return fallback(*args, **kwargs)
|
||||
native_call: Final[Callable[[], OCRResponse] | None] = (
|
||||
(lambda: native(request, args, kwargs, False, HOST)) if native is not None else None
|
||||
)
|
||||
return invoke(
|
||||
execution=execution,
|
||||
native_call=native_call,
|
||||
python_fallback=lambda: fallback(*args, **kwargs),
|
||||
adapt=lambda value: value,
|
||||
context=BridgeErrorContext(
|
||||
route=COMPONENT.name.value,
|
||||
provider=request.custom_llm_provider or "",
|
||||
model=request.model,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
async def aocr(*args: object, **kwargs: object) -> OCRResponse: # kwargs-ok: preserve the public OCR call shape
|
||||
request: Final = _public_request("aocr", args, kwargs)
|
||||
native: Final = select(request) if rust_ocr_enabled() else None
|
||||
if native is not None:
|
||||
try:
|
||||
return await native(request, args, kwargs, True)
|
||||
except _decline_types():
|
||||
pass
|
||||
execution: Final = COMPONENT.resolve()
|
||||
native: Final = select(request, execution)
|
||||
fallback: Final = cast( # cast-ok: forward the original call shape through the legacy @client decorator
|
||||
Callable[..., Awaitable[OCRResponse]], legacy.aocr
|
||||
)
|
||||
return await fallback(*args, **kwargs)
|
||||
native_call: Final[Callable[[], Awaitable[OCRResponse]] | None] = (
|
||||
(lambda: native(request, args, kwargs, True, HOST)) if native is not None else None
|
||||
)
|
||||
async def python_fallback() -> OCRResponse:
|
||||
return await fallback(*args, **kwargs)
|
||||
|
||||
|
||||
def _decline_types() -> tuple[type[BaseException], ...]:
|
||||
exception_types: Final = native_exception_types()
|
||||
return (exception_types[0],) if exception_types is not None else ()
|
||||
return await ainvoke(
|
||||
execution=execution,
|
||||
native_call=native_call,
|
||||
python_fallback=python_fallback,
|
||||
adapt=lambda value: value,
|
||||
context=BridgeErrorContext(
|
||||
route=COMPONENT.name.value,
|
||||
provider=request.custom_llm_provider or "",
|
||||
model=request.model,
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,25 +1,35 @@
|
|||
# Native route foundation
|
||||
# Native bridge catalog
|
||||
|
||||
Each route has a Python package and a matching Rust module under `crates/python-bridge/src/routes/`. Python `__init__.py` files are thin entrypoints exporting `ROUTE: NativeRoute` and public adapters. Request/response protocols live in `types.py`, value adapters in `value.py`, callback adapters in `callbacks.py` where needed, and full-call bindings in `lifecycle.py`. Unimplemented routes add these files when they gain an implementation
|
||||
Every SDK API has one `NativeComponent` in the immutable `COMPONENTS` catalog. A component declares its native exports and one `CapabilitySpec`, which resolves implementation availability and rollout from `CapabilityContext(provider, model, delivery)`
|
||||
|
||||
WebSocket is a Responses transport: its adapters live in `responses/websocket.py` and Rust `routes/responses/websocket.rs`, using the Responses route policy. Token counting is a utility outside the route registry, in `token_counter.py` and Rust `src/token_counter.rs`
|
||||
`DeliveryMode` contains `COMPLETED`, `STREAMING`, and `WEBSOCKET`. Lifecycle is an implementation detail, so lifecycle and value entrypoints for the same API share the same completed-delivery policy
|
||||
|
||||
`configuration.py` owns release policy. OCR is default-on, Messages and other optional routes are default-off, and transcription is required-native. The process override takes precedence over the environment except for OCR's existing environment opt-out. Required-native execution ignores optional rollout switches. A default is an enablement choice, not a claim that a lifecycle implementation exists
|
||||
`RustImplementationState` records whether Rust is unimplemented, experimental, or ready. `RolloutPolicy` independently selects unsupported, Python-only, Rust opt-in, Rust opt-out, or Rust-required execution. Optional Rust execution can fall back to Python. Rust-required execution cannot
|
||||
|
||||
`NativeRoute.select(binding)` checks policy before discovering the native module. `NativeBinding` handles validation and resettable overrides, with injectable discovery for tests. Native exports keep their existing names; moving a Python module into a package does not change its import path
|
||||
OCR completed delivery is ready and default-on. Messages, chat completions, token counting, and Responses WebSocket transport are experimental and opt-in. Other completed APIs remain Python-only. Bedrock transcription requires Rust because it has no Python implementation; Python-backed transcription providers remain on Python
|
||||
|
||||
## Lifecycle contract
|
||||
```python
|
||||
execution = COMPONENT.resolve(
|
||||
CapabilityContext(
|
||||
provider=provider,
|
||||
model=model,
|
||||
delivery=DeliveryMode.COMPLETED,
|
||||
)
|
||||
)
|
||||
```
|
||||
|
||||
Full-call bindings implement `NativeLifecycle[Request, Response]`: `(request, args, kwargs, asynchronous)`. The synchronous form returns a response, while the asynchronous form returns an inline-driven coroutine. Original positional arguments, keyword arguments and Python object identities stay available to the host
|
||||
Optional capabilities use the `litellm.rust(bool)` process override first, `LITELLM_RUST=1` or `LITELLM_RUST=0` second, then their catalog default. Python-only, Rust-required, and unsupported capabilities ignore overrides
|
||||
|
||||
Core owns effect-free admission and callback sequencing. The PyO3 `PythonRoute` implementation retains Python objects, projects consumed fields and executes core-selected hooks. The shared native handle and Python `lifecycle.py` driver preserve caller task/context, error identity, cancellation and cleanup. Python logging continues to select registered integrations and their dispatch modes
|
||||
## Fallback contract
|
||||
|
||||
Only disabled/unavailable native execution or a typed pre-effect admission decline permits Python fallback. Callback failures, projection errors and post-admission failures must not replay the request. Keep success/failure dispatch after fallible response finalization
|
||||
Each API calls its native entrypoint at most once. Rust performs request admission inside that entrypoint before provider calls or host callbacks. An unavailable binding or `RustBridgeDeclined` selects the supplied Python fallback only when the policy allows it
|
||||
|
||||
OCR implements this contract today. Messages, chat completions and transcription retain their existing value-based execution while their new full-call lifecycle slots are unfinished. Embeddings, rerank, image generation/edit, speech, moderation and Responses have lifecycle slots but no public SDK wiring here. The Rust `unimplemented_lifecycle_route!` macro registers each slot and maps a pure core decline to `RustBridgeDeclined`, without inspecting the request. Deliberately avoid `todo!()` in Python-callable paths because it panics instead of providing safe admission fallback
|
||||
Provider failures, host callback failures, cancellation, conversion failures, and response adaptation failures propagate without replay. Adaptation runs outside the decline-catching boundary
|
||||
|
||||
## Extending a route
|
||||
`invoke` and `ainvoke` return the native result or execute the supplied fallback directly. There is no public admission, prepare, accepts, or can-handle API
|
||||
|
||||
Replace the route's Rust lifecycle stub with a typed core call and a `PythonRoute` host, following OCR's `project`, `callbacks` and `lifecycle` split. Give its Python binding concrete request/response types, wire the public entrypoint through admission-only fallback, and prove positive native execution and callback parity before changing its release default
|
||||
Token counting follows the same component policy. Its one native counting entrypoint validates the tokenizer configuration and request body, obtains and caches the required tokenizer resource, then counts. Unsupported inputs decline, known resource loading failures report native unavailability, and unexpected counting failures propagate
|
||||
|
||||
Streaming and WebSocket sessions do not yet use the full-call lifecycle contract. Their follow-up needs explicit chunk delivery, backpressure, final response aggregation, consumer close, cancellation acknowledgement, deferred terminal dispatch and exactly-once cleanup. Returning an iterator or opening a socket is not terminal success. WebSocket uses the Responses policy; token counting keeps the global optional switch. Both use shared loading and retain their own session/utility protocols
|
||||
## Package layout
|
||||
|
||||
Python component packages keep their descriptor in `definition.py`, dynamic call protocols in `types.py`, and entrypoint adapters in `value.py`, `lifecycle.py`, or transport modules. Rust mirrors those APIs below `crates/python-bridge/src/routes/`
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
from asyncio import Future
|
||||
from collections.abc import Coroutine
|
||||
from typing import Literal, overload
|
||||
from collections.abc import Callable, Coroutine
|
||||
from typing import Literal, final, overload
|
||||
|
||||
from typing_extensions import Never
|
||||
|
||||
|
|
@ -8,45 +8,55 @@ from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
|||
from litellm.rust_bridge.ocr import LiteLLMOcrRequest
|
||||
|
||||
class RustBridgeDeclined(Exception): ...
|
||||
class RustBridgeUnavailable(Exception): ...
|
||||
class RustHostCallbackError(Exception): ...
|
||||
class RustUpstreamError(Exception): ...
|
||||
|
||||
@overload
|
||||
def _ocr_lifecycle(
|
||||
request: LiteLLMOcrRequest, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: Literal[False]
|
||||
request: LiteLLMOcrRequest,
|
||||
args: tuple[object, ...],
|
||||
kwargs: dict[str, object],
|
||||
asynchronous: Literal[False],
|
||||
host: object,
|
||||
) -> OCRResponse: ...
|
||||
@overload
|
||||
def _ocr_lifecycle(
|
||||
request: LiteLLMOcrRequest, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: Literal[True]
|
||||
request: LiteLLMOcrRequest,
|
||||
args: tuple[object, ...],
|
||||
kwargs: dict[str, object],
|
||||
asynchronous: Literal[True],
|
||||
host: object,
|
||||
) -> Coroutine[object, object, OCRResponse]: ...
|
||||
def _messages_lifecycle(
|
||||
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool
|
||||
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object
|
||||
) -> Never: ...
|
||||
def _chat_completions_lifecycle(
|
||||
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool
|
||||
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object
|
||||
) -> Never: ...
|
||||
def _transcription_lifecycle(
|
||||
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool
|
||||
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object
|
||||
) -> Never: ...
|
||||
def _embeddings_lifecycle(
|
||||
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool
|
||||
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object
|
||||
) -> Never: ...
|
||||
def _rerank_lifecycle(
|
||||
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool
|
||||
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object
|
||||
) -> Never: ...
|
||||
def _image_generation_lifecycle(
|
||||
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool
|
||||
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object
|
||||
) -> Never: ...
|
||||
def _image_edit_lifecycle(
|
||||
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool
|
||||
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object
|
||||
) -> Never: ...
|
||||
def _speech_lifecycle(
|
||||
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool
|
||||
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object
|
||||
) -> Never: ...
|
||||
def _moderation_lifecycle(
|
||||
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool
|
||||
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object
|
||||
) -> Never: ...
|
||||
def _responses_lifecycle(
|
||||
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool
|
||||
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object
|
||||
) -> Never: ...
|
||||
def ocr(
|
||||
model: str,
|
||||
|
|
@ -98,6 +108,7 @@ def messages(
|
|||
custom_llm_provider: str | None = None,
|
||||
extra_headers: object = None,
|
||||
timeout_seconds: float | None = None,
|
||||
has_agentic_hook: bool | None = None,
|
||||
) -> dict[str, object]: ...
|
||||
def amessages(
|
||||
model: str,
|
||||
|
|
@ -107,6 +118,7 @@ def amessages(
|
|||
custom_llm_provider: str | None = None,
|
||||
extra_headers: object = None,
|
||||
timeout_seconds: float | None = None,
|
||||
has_agentic_hook: bool | None = None,
|
||||
) -> Future[dict[str, object]]: ...
|
||||
def chat_completions(
|
||||
model: str,
|
||||
|
|
@ -117,6 +129,8 @@ def chat_completions(
|
|||
custom_llm_provider: str | None = None,
|
||||
extra_headers: object = None,
|
||||
timeout_seconds: float | None = None,
|
||||
host_facts: object = None,
|
||||
on_request: Callable[[], None] | None = None,
|
||||
) -> dict[str, object]: ...
|
||||
def achat_completions(
|
||||
model: str,
|
||||
|
|
@ -127,10 +141,9 @@ def achat_completions(
|
|||
custom_llm_provider: str | None = None,
|
||||
extra_headers: object = None,
|
||||
timeout_seconds: float | None = None,
|
||||
host_facts: object = None,
|
||||
on_request: Callable[[], None] | None = None,
|
||||
) -> Future[dict[str, object]]: ...
|
||||
def chat_completions_decline(
|
||||
model: str, messages: object, optional_params: object = None, custom_llm_provider: str | None = None
|
||||
) -> str | None: ...
|
||||
|
||||
_OCR_MAX_FILE_BYTES: int
|
||||
|
||||
|
|
@ -140,21 +153,60 @@ def _ocr_upload_document(
|
|||
file_content: bytes, file_name: str | None = None, content_type: str | None = None
|
||||
) -> dict[str, object]: ...
|
||||
|
||||
@final
|
||||
class ResponsesWebSocketConnection:
|
||||
@classmethod
|
||||
def connect(
|
||||
cls, url: str, headers: object = None, timeout_seconds: float | None = None
|
||||
cls,
|
||||
url: str,
|
||||
headers: object = None,
|
||||
timeout_seconds: float | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
) -> Future[ResponsesWebSocketConnection]: ...
|
||||
def send_text(self, text: str) -> Future[None]: ...
|
||||
def recv_text(self) -> Future[str | None]: ...
|
||||
def close(self) -> Future[None]: ...
|
||||
|
||||
class TokenCounter:
|
||||
def __init__(self, tokenizer_json: str) -> None: ...
|
||||
@staticmethod
|
||||
def from_cl100k_ranks(rank_file: str) -> TokenCounter: ...
|
||||
@staticmethod
|
||||
def from_o200k_ranks(rank_file: str) -> TokenCounter: ...
|
||||
def acount_request(self, body: bytes) -> Future[dict[str, object]]: ...
|
||||
def count_input_tokens(
|
||||
body: bytes,
|
||||
kind: str | None,
|
||||
encoding: str,
|
||||
disabled: bool,
|
||||
legacy_accounting: bool,
|
||||
resource_loader: Callable[[str], str],
|
||||
) -> Future[dict[str, object]]: ...
|
||||
|
||||
def gil_stats() -> dict[str, int]: ...
|
||||
|
||||
__all__ = [
|
||||
"RustBridgeUnavailable",
|
||||
"RustBridgeDeclined",
|
||||
"RustHostCallbackError",
|
||||
"RustUpstreamError",
|
||||
"ocr",
|
||||
"aocr",
|
||||
"_OCR_MAX_FILE_BYTES",
|
||||
"_ocr_upload_document",
|
||||
"_ocr_file_document",
|
||||
"_ocr_mime_type",
|
||||
"_ocr_lifecycle",
|
||||
"_transcription_lifecycle",
|
||||
"transcription",
|
||||
"atranscription",
|
||||
"_messages_lifecycle",
|
||||
"messages",
|
||||
"amessages",
|
||||
"_chat_completions_lifecycle",
|
||||
"chat_completions",
|
||||
"achat_completions",
|
||||
"_embeddings_lifecycle",
|
||||
"_image_edit_lifecycle",
|
||||
"_image_generation_lifecycle",
|
||||
"_moderation_lifecycle",
|
||||
"_rerank_lifecycle",
|
||||
"ResponsesWebSocketConnection",
|
||||
"_responses_lifecycle",
|
||||
"_speech_lifecycle",
|
||||
"count_input_tokens",
|
||||
"gil_stats",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -63,3 +63,15 @@ def native_exception_types() -> tuple[type[BaseException], type[BaseException]]
|
|||
if not isinstance(declined, type) or not isinstance(upstream, type):
|
||||
return None
|
||||
return declined, upstream
|
||||
|
||||
|
||||
def native_unavailable_exception() -> tuple[type[BaseException], ...]:
|
||||
native: Final = get_native_bridge()
|
||||
unavailable: Final = getattr(native, "RustBridgeUnavailable", None)
|
||||
return (unavailable,) if isinstance(unavailable, type) and issubclass(unavailable, BaseException) else ()
|
||||
|
||||
|
||||
def native_host_callback_exception() -> tuple[type[BaseException], ...]:
|
||||
native: Final = get_native_bridge()
|
||||
callback: Final = getattr(native, "RustHostCallbackError", None)
|
||||
return (callback,) if isinstance(callback, type) and issubclass(callback, BaseException) else ()
|
||||
|
|
|
|||
145
litellm/rust_bridge/catalog.py
Normal file
145
litellm/rust_bridge/catalog.py
Normal file
|
|
@ -0,0 +1,145 @@
|
|||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.configuration import (
|
||||
CapabilityContext,
|
||||
CapabilityDefinition,
|
||||
CapabilitySpec,
|
||||
DeliveryMode,
|
||||
RolloutPolicy,
|
||||
RouteName,
|
||||
RustImplementationState,
|
||||
UtilityName,
|
||||
)
|
||||
from litellm.rust_bridge.route import NativeComponent
|
||||
|
||||
|
||||
def _unimplemented(*, python_available: bool = True) -> CapabilityDefinition:
|
||||
return CapabilityDefinition(
|
||||
rust=RustImplementationState.UNIMPLEMENTED,
|
||||
python_available=python_available,
|
||||
rollout=RolloutPolicy.PYTHON_ONLY if python_available else RolloutPolicy.UNSUPPORTED,
|
||||
)
|
||||
|
||||
|
||||
def _experimental() -> CapabilityDefinition:
|
||||
return CapabilityDefinition(
|
||||
rust=RustImplementationState.EXPERIMENTAL,
|
||||
python_available=True,
|
||||
rollout=RolloutPolicy.RUST_OPT_IN,
|
||||
)
|
||||
|
||||
|
||||
def _ready_default() -> CapabilityDefinition:
|
||||
return CapabilityDefinition(
|
||||
rust=RustImplementationState.READY,
|
||||
python_available=True,
|
||||
rollout=RolloutPolicy.RUST_OPT_OUT,
|
||||
)
|
||||
|
||||
|
||||
def _completed_only(context: CapabilityContext, completed: CapabilityDefinition) -> CapabilityDefinition:
|
||||
return completed if context.delivery is DeliveryMode.COMPLETED else _unimplemented()
|
||||
|
||||
|
||||
def _ocr_capability(context: CapabilityContext) -> CapabilityDefinition:
|
||||
return _completed_only(context, _ready_default())
|
||||
|
||||
|
||||
def _experimental_completed(context: CapabilityContext) -> CapabilityDefinition:
|
||||
return _completed_only(context, _experimental())
|
||||
|
||||
|
||||
def _python_completed(context: CapabilityContext) -> CapabilityDefinition:
|
||||
return _completed_only(context, _unimplemented())
|
||||
|
||||
|
||||
def _responses_capability(context: CapabilityContext) -> CapabilityDefinition:
|
||||
return _experimental() if context.delivery is DeliveryMode.WEBSOCKET else _unimplemented()
|
||||
|
||||
|
||||
def _transcription_capability(context: CapabilityContext) -> CapabilityDefinition:
|
||||
if context.delivery is not DeliveryMode.COMPLETED:
|
||||
return _unimplemented()
|
||||
if context.provider == "bedrock":
|
||||
return CapabilityDefinition(
|
||||
rust=RustImplementationState.EXPERIMENTAL,
|
||||
python_available=False,
|
||||
rollout=RolloutPolicy.RUST_REQUIRED,
|
||||
)
|
||||
|
||||
import litellm
|
||||
from litellm.constants import AZURE_OPENAI_AUDIO_PROVIDERS
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
provider: Final = next((provider for provider in LlmProviders if provider.value == context.provider), None)
|
||||
python_available: Final = provider is not None and (
|
||||
context.provider in AZURE_OPENAI_AUDIO_PROVIDERS
|
||||
or context.provider in litellm.openai_compatible_providers
|
||||
or ProviderConfigManager.get_provider_audio_transcription_config(model=context.model, provider=provider)
|
||||
is not None
|
||||
)
|
||||
return _unimplemented(python_available=python_available)
|
||||
|
||||
|
||||
def _component(
|
||||
name: RouteName | UtilityName,
|
||||
capability: CapabilitySpec,
|
||||
exports: tuple[str, ...],
|
||||
) -> NativeComponent:
|
||||
return NativeComponent(name=name, capability=capability, exports=exports)
|
||||
|
||||
|
||||
COMPONENTS: Final = MappingProxyType(
|
||||
{
|
||||
RouteName.OCR: _component(
|
||||
RouteName.OCR,
|
||||
_ocr_capability,
|
||||
(
|
||||
"ocr",
|
||||
"aocr",
|
||||
"_ocr_file_document",
|
||||
"_ocr_upload_document",
|
||||
"_OCR_MAX_FILE_BYTES",
|
||||
"_ocr_mime_type",
|
||||
"_ocr_lifecycle",
|
||||
),
|
||||
),
|
||||
RouteName.MESSAGES: _component(
|
||||
RouteName.MESSAGES,
|
||||
_experimental_completed,
|
||||
("messages", "amessages", "_messages_lifecycle"),
|
||||
),
|
||||
RouteName.CHAT_COMPLETIONS: _component(
|
||||
RouteName.CHAT_COMPLETIONS,
|
||||
_experimental_completed,
|
||||
("chat_completions", "achat_completions", "_chat_completions_lifecycle"),
|
||||
),
|
||||
RouteName.TRANSCRIPTION: _component(
|
||||
RouteName.TRANSCRIPTION,
|
||||
_transcription_capability,
|
||||
("transcription", "atranscription", "_transcription_lifecycle"),
|
||||
),
|
||||
RouteName.EMBEDDINGS: _component(RouteName.EMBEDDINGS, _python_completed, ("_embeddings_lifecycle",)),
|
||||
RouteName.RERANK: _component(RouteName.RERANK, _python_completed, ("_rerank_lifecycle",)),
|
||||
RouteName.IMAGE_GENERATION: _component(
|
||||
RouteName.IMAGE_GENERATION, _python_completed, ("_image_generation_lifecycle",)
|
||||
),
|
||||
RouteName.IMAGE_EDIT: _component(RouteName.IMAGE_EDIT, _python_completed, ("_image_edit_lifecycle",)),
|
||||
RouteName.SPEECH: _component(RouteName.SPEECH, _python_completed, ("_speech_lifecycle",)),
|
||||
RouteName.MODERATION: _component(RouteName.MODERATION, _python_completed, ("_moderation_lifecycle",)),
|
||||
RouteName.RESPONSES: _component(
|
||||
RouteName.RESPONSES,
|
||||
_responses_capability,
|
||||
("ResponsesWebSocketConnection", "_responses_lifecycle"),
|
||||
),
|
||||
UtilityName.TOKEN_COUNTER: _component(
|
||||
UtilityName.TOKEN_COUNTER,
|
||||
_experimental_completed,
|
||||
("count_input_tokens",),
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
NATIVE_EXPORTS: Final = frozenset(export for component in COMPONENTS.values() for export in component.exports)
|
||||
|
|
@ -1,39 +1,31 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.chat_completions.callbacks import response_logger
|
||||
from litellm.rust_bridge.chat_completions.definition import COMPONENT
|
||||
from litellm.rust_bridge.chat_completions.types import (
|
||||
ResponseObserver,
|
||||
RustAchatCompletions,
|
||||
RustChatCompletions,
|
||||
RustChatCompletionsDecline,
|
||||
)
|
||||
from litellm.rust_bridge.chat_completions.value import (
|
||||
ROUTE,
|
||||
RUST_CHAT_COMPLETIONS_PROVIDERS,
|
||||
RUST_RESPONSE_HEADER,
|
||||
achat_completions,
|
||||
achat_completions_or_fallback,
|
||||
chat_completions,
|
||||
load_rust_achat_completions,
|
||||
load_rust_chat_completions,
|
||||
rust_chat_completions_accepts,
|
||||
set_rust_chat_completions,
|
||||
)
|
||||
|
||||
__all__: Final = (
|
||||
"ROUTE",
|
||||
"RUST_CHAT_COMPLETIONS_PROVIDERS",
|
||||
"COMPONENT",
|
||||
"RUST_RESPONSE_HEADER",
|
||||
"ResponseObserver",
|
||||
"RustAchatCompletions",
|
||||
"RustChatCompletions",
|
||||
"RustChatCompletionsDecline",
|
||||
"achat_completions",
|
||||
"achat_completions_or_fallback",
|
||||
"chat_completions",
|
||||
"load_rust_achat_completions",
|
||||
"load_rust_chat_completions",
|
||||
"response_logger",
|
||||
"rust_chat_completions_accepts",
|
||||
"set_rust_chat_completions",
|
||||
)
|
||||
|
|
|
|||
6
litellm/rust_bridge/chat_completions/definition.py
Normal file
6
litellm/rust_bridge/chat_completions/definition.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.catalog import COMPONENTS
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
|
||||
COMPONENT: Final = COMPONENTS[RouteName.CHAT_COMPLETIONS]
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.chat_completions import ROUTE
|
||||
from litellm.rust_bridge.chat_completions.definition import COMPONENT
|
||||
|
||||
LIFECYCLE: Final = ROUTE.lifecycle()
|
||||
LIFECYCLE: Final = COMPONENT.lifecycle()
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Mapping, Sequence
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from typing import Protocol
|
||||
|
||||
|
||||
|
|
@ -15,6 +15,8 @@ class RustChatCompletions(Protocol):
|
|||
custom_llm_provider: str | None,
|
||||
extra_headers: Mapping[str, object] | None,
|
||||
timeout_seconds: float | None,
|
||||
host_facts: Mapping[str, bool] | None = None,
|
||||
on_request: Callable[[], None] | None = None,
|
||||
) -> Mapping[str, object]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
|
@ -30,21 +32,12 @@ class RustAchatCompletions(Protocol):
|
|||
custom_llm_provider: str | None,
|
||||
extra_headers: Mapping[str, object] | None,
|
||||
timeout_seconds: float | None,
|
||||
host_facts: Mapping[str, bool] | None = None,
|
||||
on_request: Callable[[], None] | None = None,
|
||||
) -> Awaitable[Mapping[str, object]]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class RustChatCompletionsDecline(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
model: str,
|
||||
messages: Sequence[object],
|
||||
optional_params: Mapping[str, object] | None,
|
||||
custom_llm_provider: str | None,
|
||||
) -> str | None:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class ResponseObserver(Protocol):
|
||||
"""Invoked with the payload the core returned, on success only.
|
||||
|
||||
|
|
|
|||
|
|
@ -1,55 +1,28 @@
|
|||
"""Native chat completions bindings.
|
||||
|
||||
The Rust core owns the conversation translation, the provider call, and the
|
||||
response normalization for the subset of `/chat/completions` requests it
|
||||
accepts. This module only marshals inputs and hands the normalized result to
|
||||
LiteLLM's existing `ModelResponse` builder.
|
||||
|
||||
``None`` means the provider was never called, so the caller is free to serve the
|
||||
request on the Python path. A failure after the call was issued raises instead:
|
||||
retrying it there would bill the customer for the same work twice.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import Final, cast # noqa: TID251 # native callables are validated at load time
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.exceptions import APIError
|
||||
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
|
||||
convert_to_model_response_object,
|
||||
)
|
||||
from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned
|
||||
from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset, native_exception_types
|
||||
from litellm.rust_bridge.chat_completions.types import (
|
||||
ResponseObserver,
|
||||
RustAchatCompletions,
|
||||
RustChatCompletions,
|
||||
RustChatCompletionsDecline,
|
||||
)
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
from litellm.rust_bridge.route import NativeRoute
|
||||
from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset
|
||||
from litellm.rust_bridge.chat_completions.definition import COMPONENT
|
||||
from litellm.rust_bridge.chat_completions.types import ResponseObserver, RustAchatCompletions, RustChatCompletions
|
||||
from litellm.rust_bridge.configuration import CapabilityContext, DeliveryMode
|
||||
from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke, invoke
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
# Providers whose `/chat/completions` deployments the Rust core can serve. A
|
||||
# provider outside this set never reaches the bridge.
|
||||
RUST_CHAT_COMPLETIONS_PROVIDERS: Final = frozenset({"anthropic", "bedrock"})
|
||||
|
||||
# `litellm_params` values are `object`, so validate the one this module reads
|
||||
# rather than narrowing an unparameterized `Mapping` and typing the result Any.
|
||||
_LITELLM_METADATA_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
||||
|
||||
RUST_RESPONSE_HEADER: Final = "x-litellm-rust"
|
||||
|
||||
|
||||
ROUTE: Final = NativeRoute(RouteName.CHAT_COMPLETIONS)
|
||||
|
||||
|
||||
def _as_chat(value: object) -> RustChatCompletions | None:
|
||||
return cast(RustChatCompletions, value) if callable(value) else None
|
||||
|
||||
|
|
@ -58,38 +31,25 @@ def _as_achat(value: object) -> RustAchatCompletions | None:
|
|||
return cast(RustAchatCompletions, value) if callable(value) else None
|
||||
|
||||
|
||||
def _as_decline(value: object) -> RustChatCompletionsDecline | None:
|
||||
return cast(RustChatCompletionsDecline, value) if callable(value) else None
|
||||
|
||||
|
||||
_CHAT: Final = ROUTE.bind("chat_completions", validate=_as_chat)
|
||||
_ACHAT: Final = ROUTE.bind("achat_completions", validate=_as_achat)
|
||||
_DECLINE: Final = ROUTE.bind("chat_completions_decline", validate=_as_decline)
|
||||
_CHAT: Final = COMPONENT.bind("chat_completions", validate=_as_chat)
|
||||
_ACHAT: Final = COMPONENT.bind("achat_completions", validate=_as_achat)
|
||||
|
||||
|
||||
def set_rust_chat_completions(
|
||||
*,
|
||||
chat_completions: RustChatCompletions | None | BindingUnset = BINDING_UNSET,
|
||||
achat_completions: RustAchatCompletions | None | BindingUnset = BINDING_UNSET,
|
||||
decline: RustChatCompletionsDecline | None | BindingUnset = BINDING_UNSET,
|
||||
) -> None:
|
||||
"""Inject the native callables, so tests can supply a double instead of
|
||||
patching module attributes."""
|
||||
_CHAT.configure(chat_completions)
|
||||
_ACHAT.configure(achat_completions)
|
||||
_DECLINE.configure(decline)
|
||||
|
||||
|
||||
def load_rust_chat_completions() -> RustChatCompletions | None:
|
||||
return ROUTE.select(_CHAT)
|
||||
return COMPONENT.resolve().select(_CHAT)
|
||||
|
||||
|
||||
def load_rust_achat_completions() -> RustAchatCompletions | None:
|
||||
return ROUTE.select(_ACHAT)
|
||||
|
||||
|
||||
def _load_rust_decline() -> RustChatCompletionsDecline | None:
|
||||
return ROUTE.select(_DECLINE)
|
||||
return COMPONENT.resolve().select(_ACHAT)
|
||||
|
||||
|
||||
def _anthropic_user_id_reaches_the_body(litellm_params: Mapping[str, object] | None) -> bool:
|
||||
|
|
@ -101,134 +61,21 @@ def _anthropic_user_id_reaches_the_body(litellm_params: Mapping[str, object] | N
|
|||
return entries.get("user_id") is not None
|
||||
|
||||
|
||||
def _litellm_metadata_reaches_the_provider(
|
||||
custom_llm_provider: str | None, litellm_params: Mapping[str, object] | None
|
||||
) -> bool:
|
||||
"""Whether the Python transform would promote proxy-owned attribution into the
|
||||
provider request, below this gate and inside the function the Rust route replaces.
|
||||
|
||||
`AnthropicConfig.transform_request` promotes a valid `metadata["user_id"]`
|
||||
into the Messages body, so the core never sees the key and would send the
|
||||
request to Anthropic with the abuse-detection attribution missing.
|
||||
|
||||
`AmazonConverseConfig` resolves proxy-owned `requestMetadata` onto the
|
||||
Converse body whenever the operator armed `bedrock_request_metadata_fields`.
|
||||
Owning that field also means evicting a caller-supplied one, which the core
|
||||
cannot do either, so ownership alone is the condition rather than whether
|
||||
anything resolved.
|
||||
|
||||
Deliberately a superset of Python's condition in both cases: declining a
|
||||
request Python would not have attributed anyway costs only the Rust path,
|
||||
while missing one loses the attribution silently.
|
||||
"""
|
||||
match custom_llm_provider:
|
||||
case "anthropic":
|
||||
return _anthropic_user_id_reaches_the_body(litellm_params)
|
||||
case "bedrock":
|
||||
return bedrock_request_metadata_is_owned()
|
||||
case _:
|
||||
return False
|
||||
|
||||
|
||||
def rust_chat_completions_accepts(
|
||||
*,
|
||||
model: str,
|
||||
messages: Sequence[object],
|
||||
optional_params: Mapping[str, object],
|
||||
custom_llm_provider: str | None,
|
||||
litellm_params: Mapping[str, object] | None,
|
||||
stream: object,
|
||||
) -> bool:
|
||||
"""Whether the Rust path will serve this request.
|
||||
|
||||
Asked before the caller commits to either path, so pre-call logging is
|
||||
emitted exactly once, on whichever path actually runs. The core's own
|
||||
capability gate answers the second half; it resolves no credentials and
|
||||
performs no I/O.
|
||||
"""
|
||||
if custom_llm_provider not in RUST_CHAT_COMPLETIONS_PROVIDERS:
|
||||
return False
|
||||
if stream:
|
||||
return False
|
||||
if not ROUTE.enabled():
|
||||
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")
|
||||
return False
|
||||
decline: Final = _load_rust_decline()
|
||||
if decline is None:
|
||||
return False
|
||||
try:
|
||||
reason: Final = decline(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
except Exception as rust_error: # noqa: BLE001 # rollout-safety fallback: any Rust bridge failure must fall back to the Python path
|
||||
verbose_logger.debug(
|
||||
"Rust chat completions gate raised %s; staying on the Python path",
|
||||
type(rust_error).__name__,
|
||||
)
|
||||
return False
|
||||
if reason is not None:
|
||||
verbose_logger.debug("Rust chat completions declined (%s); using the Python path", reason)
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _rust_bridge_exceptions() -> tuple[type[BaseException], type[BaseException]] | None:
|
||||
"""`(declined, upstream_failed)` from the native module, or None when absent."""
|
||||
return native_exception_types()
|
||||
|
||||
|
||||
def _reraise_or_decline(
|
||||
rust_error: BaseException,
|
||||
*,
|
||||
model: str,
|
||||
custom_llm_provider: str | None,
|
||||
) -> None:
|
||||
"""Re-raise a failure the provider already saw, or return so the caller declines.
|
||||
|
||||
A request that never reached the provider is safe to serve on the Python
|
||||
path. One that did is not: the provider has already done the work, so a
|
||||
second attempt bills for it twice. Those surface as an `APIError` carrying
|
||||
the upstream status, which LiteLLM's exception mapping already understands.
|
||||
"""
|
||||
exceptions: Final = _rust_bridge_exceptions()
|
||||
if exceptions is None:
|
||||
verbose_logger.debug(
|
||||
"Rust chat completions bridge raised %s; falling back to Python path",
|
||||
type(rust_error).__name__,
|
||||
)
|
||||
return
|
||||
declined, upstream_failed = exceptions
|
||||
if isinstance(rust_error, upstream_failed):
|
||||
args: Final = rust_error.args
|
||||
status: Final = args[0] if args else 0
|
||||
message: Final = args[1] if len(args) > 1 else ""
|
||||
raise APIError(
|
||||
status_code=int(status) or 500,
|
||||
message=f"litellm rust chat completions: {message}",
|
||||
llm_provider=custom_llm_provider or "",
|
||||
model=model,
|
||||
)
|
||||
if not isinstance(rust_error, declined):
|
||||
raise rust_error
|
||||
verbose_logger.debug(
|
||||
"Rust chat completions declined before calling the provider (%s); using the Python path",
|
||||
rust_error,
|
||||
def _host_facts(stream: object, litellm_params: Mapping[str, object] | None) -> Mapping[str, bool]:
|
||||
return MappingProxyType(
|
||||
{
|
||||
"stream": bool(stream),
|
||||
"anthropic_user_id": _anthropic_user_id_reaches_the_body(litellm_params),
|
||||
"bedrock_metadata_owned": bedrock_request_metadata_is_owned(),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _build_model_response(
|
||||
rust_response: Mapping[str, object],
|
||||
model_response: ModelResponse,
|
||||
) -> ModelResponse:
|
||||
def _build_model_response(rust_response: Mapping[str, object], model_response: ModelResponse) -> ModelResponse:
|
||||
built: Final = convert_to_model_response_object(
|
||||
response_object=dict(rust_response), # mutable-ok: the converter takes a real dict and rewrites it
|
||||
model_response_object=model_response,
|
||||
hidden_params={"additional_headers": {RUST_RESPONSE_HEADER: "true"}}, # mutable-ok: rewritten by the converter
|
||||
hidden_params={"additional_headers": {RUST_RESPONSE_HEADER: "true"}}, # mutable-ok: converter rewrites it
|
||||
)
|
||||
if not isinstance(built, ModelResponse):
|
||||
raise TypeError(f"expected a ModelResponse from the rust path, got {type(built).__name__}")
|
||||
|
|
@ -246,13 +93,27 @@ def chat_completions(
|
|||
custom_llm_provider: str | None,
|
||||
extra_headers: Mapping[str, object] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
on_response: ResponseObserver,
|
||||
) -> ModelResponse | None:
|
||||
rust_chat_completions: Final = load_rust_chat_completions()
|
||||
if rust_chat_completions is None:
|
||||
return None
|
||||
try:
|
||||
rust_response: Final = rust_chat_completions(
|
||||
python_fallback: Callable[[], object],
|
||||
stream: object = False,
|
||||
litellm_params: Mapping[str, object] | None = None,
|
||||
on_request: Callable[[], None] = lambda: None,
|
||||
on_response: ResponseObserver = lambda _response: None,
|
||||
) -> object:
|
||||
execution: Final = COMPONENT.resolve(
|
||||
CapabilityContext(
|
||||
provider=custom_llm_provider or "",
|
||||
model=model,
|
||||
delivery=DeliveryMode.STREAMING if bool(stream) else DeliveryMode.COMPLETED,
|
||||
)
|
||||
)
|
||||
rust_chat_completions: Final = execution.select(_CHAT)
|
||||
|
||||
def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
|
||||
on_response(rust_response)
|
||||
return _build_model_response(rust_response, model_response)
|
||||
|
||||
native_call: Final[Callable[[], Mapping[str, object]] | None] = (
|
||||
lambda: rust_chat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
|
|
@ -261,12 +122,19 @@ def chat_completions(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
host_facts=_host_facts(stream, litellm_params),
|
||||
on_request=on_request,
|
||||
)
|
||||
except Exception as rust_error: # noqa: BLE001 # rollout safety: the helper re-raises anything the provider already saw
|
||||
_reraise_or_decline(rust_error, model=model, custom_llm_provider=custom_llm_provider)
|
||||
return None
|
||||
on_response(rust_response)
|
||||
return _build_model_response(rust_response, model_response)
|
||||
if rust_chat_completions is not None
|
||||
else None
|
||||
)
|
||||
return invoke(
|
||||
execution=execution,
|
||||
native_call=native_call,
|
||||
python_fallback=python_fallback,
|
||||
adapt=adapt,
|
||||
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
||||
|
||||
async def achat_completions(
|
||||
|
|
@ -280,13 +148,27 @@ async def achat_completions(
|
|||
custom_llm_provider: str | None,
|
||||
extra_headers: Mapping[str, object] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
on_response: ResponseObserver,
|
||||
) -> ModelResponse | None:
|
||||
rust_achat_completions: Final = load_rust_achat_completions()
|
||||
if rust_achat_completions is None:
|
||||
return None
|
||||
try:
|
||||
rust_response: Final = await rust_achat_completions(
|
||||
python_fallback: Callable[[], Awaitable[object]],
|
||||
stream: object = False,
|
||||
litellm_params: Mapping[str, object] | None = None,
|
||||
on_request: Callable[[], None] = lambda: None,
|
||||
on_response: ResponseObserver = lambda _response: None,
|
||||
) -> object:
|
||||
execution: Final = COMPONENT.resolve(
|
||||
CapabilityContext(
|
||||
provider=custom_llm_provider or "",
|
||||
model=model,
|
||||
delivery=DeliveryMode.STREAMING if bool(stream) else DeliveryMode.COMPLETED,
|
||||
)
|
||||
)
|
||||
rust_achat_completions: Final = execution.select(_ACHAT)
|
||||
|
||||
def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
|
||||
on_response(rust_response)
|
||||
return _build_model_response(rust_response, model_response)
|
||||
|
||||
native_call: Final[Callable[[], Awaitable[Mapping[str, object]]] | None] = (
|
||||
lambda: rust_achat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
|
|
@ -295,48 +177,16 @@ async def achat_completions(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
host_facts=_host_facts(stream, litellm_params),
|
||||
on_request=on_request,
|
||||
)
|
||||
except Exception as rust_error: # noqa: BLE001 # rollout safety: the helper re-raises anything the provider already saw
|
||||
_reraise_or_decline(rust_error, model=model, custom_llm_provider=custom_llm_provider)
|
||||
return None
|
||||
on_response(rust_response)
|
||||
return _build_model_response(rust_response, model_response)
|
||||
|
||||
|
||||
async def achat_completions_or_fallback(
|
||||
*,
|
||||
model: str,
|
||||
messages: Sequence[object],
|
||||
optional_params: Mapping[str, object],
|
||||
model_response: ModelResponse,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: Mapping[str, object] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
on_response: ResponseObserver,
|
||||
python_fallback: Callable[[], Awaitable[object]],
|
||||
) -> object:
|
||||
"""Await the Rust path, falling back to the caller's own Python path when
|
||||
the bridge is unavailable or the call fails.
|
||||
|
||||
The caller supplies the fallback, so the bridge stays free of provider
|
||||
dispatch. This exists because a caller that dispatches asynchronously has
|
||||
already returned a coroutine by the time a Rust failure surfaces, and so
|
||||
cannot fall back on its own.
|
||||
"""
|
||||
response: Final = await achat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout=timeout,
|
||||
on_response=on_response,
|
||||
if rust_achat_completions is not None
|
||||
else None
|
||||
)
|
||||
return await ainvoke(
|
||||
execution=execution,
|
||||
native_call=native_call,
|
||||
python_fallback=python_fallback,
|
||||
adapt=adapt,
|
||||
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
if response is not None:
|
||||
return response
|
||||
return await python_fallback()
|
||||
|
|
|
|||
|
|
@ -3,14 +3,10 @@ from __future__ import annotations
|
|||
import os
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
from typing import Final, Protocol, TypeAlias, assert_never
|
||||
|
||||
DEFAULT_RUST_ENABLED: Final = False
|
||||
_GLOBAL_ENV_NAME: Final = "LITELLM_RUST"
|
||||
_ENV_BOOL: Final = TypeAdapter(bool)
|
||||
|
||||
|
||||
class RouteName(str, Enum):
|
||||
|
|
@ -27,33 +23,77 @@ class RouteName(str, Enum):
|
|||
RESPONSES = "responses"
|
||||
|
||||
|
||||
class RouteMode(str, Enum):
|
||||
OPTIONAL = "optional"
|
||||
REQUIRED = "required"
|
||||
class UtilityName(str, Enum):
|
||||
TOKEN_COUNTER = "token_counter"
|
||||
|
||||
|
||||
class RustImplementationState(str, Enum):
|
||||
UNIMPLEMENTED = "unimplemented"
|
||||
EXPERIMENTAL = "experimental"
|
||||
READY = "ready"
|
||||
|
||||
|
||||
class RolloutPolicy(str, Enum):
|
||||
UNSUPPORTED = "unsupported"
|
||||
PYTHON_ONLY = "python_only"
|
||||
RUST_OPT_IN = "rust_opt_in"
|
||||
RUST_OPT_OUT = "rust_opt_out"
|
||||
RUST_REQUIRED = "rust_required"
|
||||
|
||||
|
||||
class ExecutionDecision(str, Enum):
|
||||
UNSUPPORTED = "unsupported"
|
||||
PYTHON = "python"
|
||||
RUST_WITH_FALLBACK = "rust_with_fallback"
|
||||
RUST_REQUIRED = "rust_required"
|
||||
|
||||
|
||||
class DeliveryMode(str, Enum):
|
||||
COMPLETED = "completed"
|
||||
STREAMING = "streaming"
|
||||
WEBSOCKET = "websocket"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RoutePolicy:
|
||||
default_enabled: bool = False
|
||||
mode: RouteMode = RouteMode.OPTIONAL
|
||||
environment_opt_out: bool = False
|
||||
class CapabilityContext:
|
||||
provider: str = ""
|
||||
model: str = ""
|
||||
delivery: DeliveryMode = DeliveryMode.COMPLETED
|
||||
|
||||
|
||||
ROUTE_POLICIES: Final = MappingProxyType(
|
||||
{
|
||||
RouteName.OCR: RoutePolicy(default_enabled=True, environment_opt_out=True),
|
||||
RouteName.MESSAGES: RoutePolicy(),
|
||||
RouteName.CHAT_COMPLETIONS: RoutePolicy(),
|
||||
RouteName.TRANSCRIPTION: RoutePolicy(mode=RouteMode.REQUIRED),
|
||||
RouteName.EMBEDDINGS: RoutePolicy(),
|
||||
RouteName.RERANK: RoutePolicy(),
|
||||
RouteName.IMAGE_GENERATION: RoutePolicy(),
|
||||
RouteName.IMAGE_EDIT: RoutePolicy(),
|
||||
RouteName.SPEECH: RoutePolicy(),
|
||||
RouteName.MODERATION: RoutePolicy(),
|
||||
RouteName.RESPONSES: RoutePolicy(),
|
||||
}
|
||||
)
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CapabilityDefinition:
|
||||
rust: RustImplementationState
|
||||
python_available: bool
|
||||
rollout: RolloutPolicy
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self.rollout is RolloutPolicy.UNSUPPORTED:
|
||||
if self.rust is not RustImplementationState.UNIMPLEMENTED or self.python_available:
|
||||
raise ValueError("an unsupported capability must have neither implementation")
|
||||
return
|
||||
if self.rust is RustImplementationState.UNIMPLEMENTED:
|
||||
if self.rollout is not RolloutPolicy.PYTHON_ONLY or not self.python_available:
|
||||
raise ValueError("an unimplemented Rust capability must use its Python implementation")
|
||||
return
|
||||
if self.rollout is RolloutPolicy.PYTHON_ONLY:
|
||||
raise ValueError("an implemented Rust capability must declare a Rust rollout")
|
||||
if self.rollout is RolloutPolicy.RUST_OPT_IN or self.rollout is RolloutPolicy.RUST_OPT_OUT:
|
||||
if not self.python_available:
|
||||
raise ValueError("an optional Rust capability requires a Python fallback")
|
||||
return
|
||||
if self.rollout is RolloutPolicy.RUST_REQUIRED:
|
||||
if self.python_available:
|
||||
raise ValueError("a required Rust capability must have no Python implementation")
|
||||
return
|
||||
assert_never(self.rollout)
|
||||
|
||||
|
||||
class CapabilityResolver(Protocol):
|
||||
def __call__(self, context: CapabilityContext, /) -> CapabilityDefinition: ...
|
||||
|
||||
|
||||
CapabilitySpec: TypeAlias = CapabilityDefinition | CapabilityResolver
|
||||
|
||||
|
||||
class _RustConfiguration:
|
||||
|
|
@ -67,38 +107,72 @@ _CONFIGURATION: Final = _RustConfiguration()
|
|||
def _parse_env_bool(value: str | None) -> bool | None:
|
||||
if value is None:
|
||||
return None
|
||||
return _ENV_BOOL.validate_python(value.strip())
|
||||
match value.strip():
|
||||
case "1":
|
||||
return True
|
||||
case "0":
|
||||
return False
|
||||
case invalid:
|
||||
raise ValueError(f"{_GLOBAL_ENV_NAME} must be '1' or '0', got {invalid!r}")
|
||||
|
||||
|
||||
def resolve_rust_enabled(
|
||||
def resolve_capability(
|
||||
capability: CapabilityDefinition,
|
||||
*,
|
||||
process_override: bool | None,
|
||||
environment_override: bool | None,
|
||||
release_default: bool = DEFAULT_RUST_ENABLED,
|
||||
) -> bool:
|
||||
if process_override is not None:
|
||||
return process_override
|
||||
if environment_override is not None:
|
||||
return environment_override
|
||||
return release_default
|
||||
) -> ExecutionDecision:
|
||||
match capability.rollout:
|
||||
case RolloutPolicy.UNSUPPORTED:
|
||||
return ExecutionDecision.UNSUPPORTED
|
||||
case RolloutPolicy.PYTHON_ONLY:
|
||||
return ExecutionDecision.PYTHON
|
||||
case RolloutPolicy.RUST_REQUIRED:
|
||||
return ExecutionDecision.RUST_REQUIRED
|
||||
case RolloutPolicy.RUST_OPT_IN | RolloutPolicy.RUST_OPT_OUT:
|
||||
enabled: Final = (
|
||||
process_override
|
||||
if process_override is not None
|
||||
else environment_override
|
||||
if environment_override is not None
|
||||
else capability.rollout is RolloutPolicy.RUST_OPT_OUT
|
||||
)
|
||||
return ExecutionDecision.RUST_WITH_FALLBACK if enabled else ExecutionDecision.PYTHON
|
||||
assert_never(capability.rollout)
|
||||
|
||||
|
||||
def rust_enabled(route: RouteName | None = None) -> bool:
|
||||
policy: Final = ROUTE_POLICIES[route] if route is not None else RoutePolicy()
|
||||
if policy.mode is RouteMode.REQUIRED:
|
||||
return True
|
||||
environment: Final = _parse_env_bool(os.getenv(_GLOBAL_ENV_NAME))
|
||||
if policy.environment_opt_out and environment is False:
|
||||
return False
|
||||
return resolve_rust_enabled(
|
||||
def _definition(spec: CapabilitySpec, context: CapabilityContext) -> CapabilityDefinition:
|
||||
if isinstance(spec, CapabilityDefinition):
|
||||
return spec
|
||||
return spec(context)
|
||||
|
||||
|
||||
def capability_decision(spec: CapabilitySpec, *, context: CapabilityContext) -> ExecutionDecision:
|
||||
capability: Final = _definition(spec, context)
|
||||
environment_override: Final = (
|
||||
_parse_env_bool(os.getenv(_GLOBAL_ENV_NAME))
|
||||
if _CONFIGURATION.override is None
|
||||
and capability.rollout in (RolloutPolicy.RUST_OPT_IN, RolloutPolicy.RUST_OPT_OUT)
|
||||
else None
|
||||
)
|
||||
return resolve_capability(
|
||||
capability,
|
||||
process_override=_CONFIGURATION.override,
|
||||
environment_override=environment,
|
||||
release_default=policy.default_enabled,
|
||||
environment_override=environment_override,
|
||||
)
|
||||
|
||||
|
||||
def rust_ocr_enabled() -> bool:
|
||||
return rust_enabled(RouteName.OCR)
|
||||
def rust_enabled() -> bool:
|
||||
environment_override: Final = (
|
||||
_parse_env_bool(os.getenv(_GLOBAL_ENV_NAME)) if _CONFIGURATION.override is None else None
|
||||
)
|
||||
return (
|
||||
_CONFIGURATION.override
|
||||
if _CONFIGURATION.override is not None
|
||||
else environment_override
|
||||
if environment_override is not None
|
||||
else DEFAULT_RUST_ENABLED
|
||||
)
|
||||
|
||||
|
||||
def reset_rust_configuration() -> None:
|
||||
|
|
@ -106,8 +180,5 @@ def reset_rust_configuration() -> None:
|
|||
|
||||
|
||||
def rust(enabled: bool) -> None:
|
||||
"""Set the process override for optional Rust paths.
|
||||
|
||||
Rust-only paths, including Bedrock transcription, are not controlled by this switch.
|
||||
"""
|
||||
"""Set the process override for optional Rust capabilities."""
|
||||
_CONFIGURATION.override = enabled
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
from litellm.rust_bridge.route import NativeRoute
|
||||
from litellm.rust_bridge.embeddings.definition import COMPONENT
|
||||
|
||||
ROUTE: Final = NativeRoute(RouteName.EMBEDDINGS)
|
||||
__all__: Final = ("COMPONENT",)
|
||||
|
|
|
|||
6
litellm/rust_bridge/embeddings/definition.py
Normal file
6
litellm/rust_bridge/embeddings/definition.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.catalog import COMPONENTS
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
|
||||
COMPONENT: Final = COMPONENTS[RouteName.EMBEDDINGS]
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.embeddings import ROUTE
|
||||
from litellm.rust_bridge.embeddings.definition import COMPONENT
|
||||
|
||||
LIFECYCLE: Final = ROUTE.lifecycle()
|
||||
LIFECYCLE: Final = COMPONENT.lifecycle()
|
||||
|
|
|
|||
10
litellm/rust_bridge/errors.py
Normal file
10
litellm/rust_bridge/errors.py
Normal file
|
|
@ -0,0 +1,10 @@
|
|||
class RustRouteUnavailableError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
class RustRouteDeclinedError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
class RustRouteUnsupportedError(NotImplementedError):
|
||||
pass
|
||||
|
|
@ -1,6 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
from litellm.rust_bridge.route import NativeRoute
|
||||
from litellm.rust_bridge.image_edit.definition import COMPONENT
|
||||
|
||||
ROUTE: Final = NativeRoute(RouteName.IMAGE_EDIT)
|
||||
__all__: Final = ("COMPONENT",)
|
||||
|
|
|
|||
6
litellm/rust_bridge/image_edit/definition.py
Normal file
6
litellm/rust_bridge/image_edit/definition.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.catalog import COMPONENTS
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
|
||||
COMPONENT: Final = COMPONENTS[RouteName.IMAGE_EDIT]
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.image_edit import ROUTE
|
||||
from litellm.rust_bridge.image_edit.definition import COMPONENT
|
||||
|
||||
LIFECYCLE: Final = ROUTE.lifecycle()
|
||||
LIFECYCLE: Final = COMPONENT.lifecycle()
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
from litellm.rust_bridge.route import NativeRoute
|
||||
from litellm.rust_bridge.image_generation.definition import COMPONENT
|
||||
|
||||
ROUTE: Final = NativeRoute(RouteName.IMAGE_GENERATION)
|
||||
__all__: Final = ("COMPONENT",)
|
||||
|
|
|
|||
6
litellm/rust_bridge/image_generation/definition.py
Normal file
6
litellm/rust_bridge/image_generation/definition.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.catalog import COMPONENTS
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
|
||||
COMPONENT: Final = COMPONENTS[RouteName.IMAGE_GENERATION]
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.image_generation import ROUTE
|
||||
from litellm.rust_bridge.image_generation.definition import COMPONENT
|
||||
|
||||
LIFECYCLE: Final = ROUTE.lifecycle()
|
||||
LIFECYCLE: Final = COMPONENT.lifecycle()
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.messages.definition import COMPONENT
|
||||
from litellm.rust_bridge.messages.types import RustAmessages, RustMessages
|
||||
from litellm.rust_bridge.messages.value import (
|
||||
ROUTE,
|
||||
amessages,
|
||||
load_rust_amessages,
|
||||
load_rust_messages,
|
||||
|
|
@ -11,7 +11,7 @@ from litellm.rust_bridge.messages.value import (
|
|||
)
|
||||
|
||||
__all__: Final = (
|
||||
"ROUTE",
|
||||
"COMPONENT",
|
||||
"RustAmessages",
|
||||
"RustMessages",
|
||||
"amessages",
|
||||
|
|
|
|||
6
litellm/rust_bridge/messages/definition.py
Normal file
6
litellm/rust_bridge/messages/definition.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.catalog import COMPONENTS
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
|
||||
COMPONENT: Final = COMPONENTS[RouteName.MESSAGES]
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.messages import ROUTE
|
||||
from litellm.rust_bridge.messages.definition import COMPONENT
|
||||
|
||||
LIFECYCLE: Final = ROUTE.lifecycle()
|
||||
LIFECYCLE: Final = COMPONENT.lifecycle()
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable
|
||||
from collections.abc import Awaitable, Mapping
|
||||
from typing import Protocol
|
||||
|
||||
|
||||
|
|
@ -8,12 +8,13 @@ class RustMessages(Protocol):
|
|||
def __call__(
|
||||
self,
|
||||
model: str,
|
||||
body: dict[str, object],
|
||||
body: Mapping[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
extra_headers: Mapping[str, object] | None,
|
||||
timeout_seconds: float | None,
|
||||
has_agentic_hook: bool = False,
|
||||
) -> dict[str, object]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
|
@ -22,11 +23,12 @@ class RustAmessages(Protocol):
|
|||
def __call__(
|
||||
self,
|
||||
model: str,
|
||||
body: dict[str, object],
|
||||
body: Mapping[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
extra_headers: Mapping[str, object] | None,
|
||||
timeout_seconds: float | None,
|
||||
has_agentic_hook: bool = False,
|
||||
) -> Awaitable[dict[str, object]]:
|
||||
raise NotImplementedError
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from typing import (
|
||||
Final,
|
||||
cast, # noqa: TID251 # native callable signatures are checked by bridge contract tests
|
||||
|
|
@ -10,13 +11,12 @@ from typing import (
|
|||
import httpx
|
||||
|
||||
from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
from litellm.rust_bridge.configuration import CapabilityContext, DeliveryMode
|
||||
from litellm.rust_bridge.messages.definition import COMPONENT
|
||||
from litellm.rust_bridge.messages.types import RustAmessages, RustMessages
|
||||
from litellm.rust_bridge.route import NativeRoute
|
||||
from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke, invoke
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
|
||||
ROUTE: Final = NativeRoute(RouteName.MESSAGES)
|
||||
|
||||
|
||||
def _as_messages(value: object) -> RustMessages | None:
|
||||
return cast(RustMessages, value) if callable(value) else None # cast-ok: validated callable native binding
|
||||
|
|
@ -26,8 +26,8 @@ def _as_amessages(value: object) -> RustAmessages | None:
|
|||
return cast(RustAmessages, value) if callable(value) else None # cast-ok: validated callable native binding
|
||||
|
||||
|
||||
_MESSAGES: Final = ROUTE.bind("messages", validate=_as_messages)
|
||||
_AMESSAGES: Final = ROUTE.bind("amessages", validate=_as_amessages)
|
||||
_MESSAGES: Final = COMPONENT.bind("messages", validate=_as_messages)
|
||||
_AMESSAGES: Final = COMPONENT.bind("amessages", validate=_as_amessages)
|
||||
|
||||
|
||||
def set_rust_messages(
|
||||
|
|
@ -40,56 +40,108 @@ def set_rust_messages(
|
|||
|
||||
|
||||
def load_rust_messages() -> RustMessages | None:
|
||||
return ROUTE.select(_MESSAGES)
|
||||
return COMPONENT.resolve().select(_MESSAGES)
|
||||
|
||||
|
||||
def load_rust_amessages() -> RustAmessages | None:
|
||||
return ROUTE.select(_AMESSAGES)
|
||||
return COMPONENT.resolve().select(_AMESSAGES)
|
||||
|
||||
|
||||
def messages(
|
||||
*,
|
||||
model: str,
|
||||
body: dict[str, object],
|
||||
body: Mapping[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
extra_headers: Mapping[str, object] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
has_agentic_hook: bool = False,
|
||||
) -> dict[str, object] | None:
|
||||
rust_messages: Final = load_rust_messages()
|
||||
if rust_messages is None:
|
||||
return None
|
||||
return rust_messages(
|
||||
model=model,
|
||||
body=body,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
execution: Final = COMPONENT.resolve(
|
||||
CapabilityContext(
|
||||
provider=custom_llm_provider or "",
|
||||
model=model,
|
||||
delivery=DeliveryMode.STREAMING if body.get("stream") is True else DeliveryMode.COMPLETED,
|
||||
)
|
||||
)
|
||||
rust_messages: Final = execution.select(_MESSAGES)
|
||||
native_call: Final[Callable[[], dict[str, object]] | None] = (
|
||||
(
|
||||
lambda: rust_messages(
|
||||
model=model,
|
||||
body=body,
|
||||
has_agentic_hook=has_agentic_hook,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
)
|
||||
)
|
||||
if rust_messages is not None
|
||||
else None
|
||||
)
|
||||
return invoke(
|
||||
execution=execution,
|
||||
native_call=native_call,
|
||||
python_fallback=lambda: None,
|
||||
adapt=lambda value: value,
|
||||
context=BridgeErrorContext(
|
||||
route=COMPONENT.name.value,
|
||||
provider=custom_llm_provider or "",
|
||||
model=model,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
async def amessages(
|
||||
*,
|
||||
model: str,
|
||||
body: dict[str, object],
|
||||
body: Mapping[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
extra_headers: Mapping[str, object] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
has_agentic_hook: bool = False,
|
||||
) -> dict[str, object] | None:
|
||||
rust_amessages: Final = load_rust_amessages()
|
||||
if rust_amessages is None:
|
||||
return None
|
||||
return await rust_amessages(
|
||||
model=model,
|
||||
body=body,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
execution: Final = COMPONENT.resolve(
|
||||
CapabilityContext(
|
||||
provider=custom_llm_provider or "",
|
||||
model=model,
|
||||
delivery=DeliveryMode.STREAMING if body.get("stream") is True else DeliveryMode.COMPLETED,
|
||||
)
|
||||
)
|
||||
rust_amessages: Final = execution.select(_AMESSAGES)
|
||||
native_call: Final[Callable[[], Awaitable[dict[str, object]]] | None] = (
|
||||
(
|
||||
lambda: rust_amessages(
|
||||
model=model,
|
||||
body=body,
|
||||
has_agentic_hook=has_agentic_hook,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
)
|
||||
)
|
||||
if rust_amessages is not None
|
||||
else None
|
||||
)
|
||||
|
||||
async def python_fallback() -> None:
|
||||
return None
|
||||
|
||||
return await ainvoke(
|
||||
execution=execution,
|
||||
native_call=native_call,
|
||||
python_fallback=python_fallback,
|
||||
adapt=lambda value: value,
|
||||
context=BridgeErrorContext(
|
||||
route=COMPONENT.name.value,
|
||||
provider=custom_llm_provider or "",
|
||||
model=model,
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
from litellm.rust_bridge.route import NativeRoute
|
||||
from litellm.rust_bridge.moderation.definition import COMPONENT
|
||||
|
||||
ROUTE: Final = NativeRoute(RouteName.MODERATION)
|
||||
__all__: Final = ("COMPONENT",)
|
||||
|
|
|
|||
6
litellm/rust_bridge/moderation/definition.py
Normal file
6
litellm/rust_bridge/moderation/definition.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.catalog import COMPONENTS
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
|
||||
COMPONENT: Final = COMPONENTS[RouteName.MODERATION]
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.moderation import ROUTE
|
||||
from litellm.rust_bridge.moderation.definition import COMPONENT
|
||||
|
||||
LIFECYCLE: Final = ROUTE.lifecycle()
|
||||
LIFECYCLE: Final = COMPONENT.lifecycle()
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.ocr.definition import COMPONENT
|
||||
from litellm.rust_bridge.ocr.types import LiteLLMOcrRequest, RustAocr, RustOcr
|
||||
from litellm.rust_bridge.ocr.value import (
|
||||
ROUTE,
|
||||
aocr,
|
||||
load_rust_aocr,
|
||||
load_rust_ocr,
|
||||
|
|
@ -10,7 +10,7 @@ from litellm.rust_bridge.ocr.value import (
|
|||
)
|
||||
|
||||
__all__: Final = (
|
||||
"ROUTE",
|
||||
"COMPONENT",
|
||||
"LiteLLMOcrRequest",
|
||||
"RustAocr",
|
||||
"RustOcr",
|
||||
|
|
|
|||
6
litellm/rust_bridge/ocr/definition.py
Normal file
6
litellm/rust_bridge/ocr/definition.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.catalog import COMPONENTS
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
|
||||
COMPONENT: Final = COMPONENTS[RouteName.OCR]
|
||||
52
litellm/rust_bridge/ocr/host.py
Normal file
52
litellm/rust_bridge/ocr/host.py
Normal file
|
|
@ -0,0 +1,52 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Final, Protocol, cast
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.rust_bridge.ocr.types import LiteLLMOcrRequest
|
||||
from litellm.rust_bridge.ocr.value import adapt_response
|
||||
from litellm.types.utils import CustomPricingLiteLLMParams
|
||||
|
||||
|
||||
class ExceptionMapper(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
*,
|
||||
model: str,
|
||||
custom_llm_provider: str | None,
|
||||
original_exception: Exception,
|
||||
completion_kwargs: dict[str, object],
|
||||
extra_kwargs: dict[str, object],
|
||||
) -> Exception: ...
|
||||
|
||||
|
||||
class OcrLifecycleHost:
|
||||
def response(self, response: Mapping[str, object]) -> OCRResponse:
|
||||
return adapt_response(response)
|
||||
|
||||
def custom_pricing_fields(self) -> tuple[str, ...]:
|
||||
return tuple(CustomPricingLiteLLMParams.model_fields)
|
||||
|
||||
def map_failure(
|
||||
self,
|
||||
error: Exception,
|
||||
request: LiteLLMOcrRequest,
|
||||
request_provider: str,
|
||||
) -> Exception:
|
||||
mapper: Final = cast(ExceptionMapper, litellm.exception_type) # cast-ok: legacy public exception mapper
|
||||
try:
|
||||
return mapper(
|
||||
model=request.model.removeprefix(f"{request_provider}/"),
|
||||
custom_llm_provider=request_provider,
|
||||
original_exception=error,
|
||||
completion_kwargs=dict(request.kwargs), # mutable-ok: exception mapper requires owned kwargs
|
||||
extra_kwargs=dict(request.kwargs), # mutable-ok: exception mapper requires owned kwargs
|
||||
)
|
||||
except Exception as public_error:
|
||||
public_error.__context__ = error
|
||||
return public_error
|
||||
|
||||
|
||||
HOST: Final = OcrLifecycleHost()
|
||||
|
|
@ -1,60 +1,31 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables
|
||||
from typing import Final, cast # noqa: TID251 # validates dynamically loaded native callables
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.rust_bridge.ocr import ROUTE, LiteLLMOcrRequest
|
||||
from litellm.rust_bridge.route import NativeLifecycle
|
||||
from litellm.rust_bridge.ocr.definition import COMPONENT
|
||||
from litellm.rust_bridge.ocr.host import HOST
|
||||
from litellm.rust_bridge.ocr.types import LiteLLMOcrRequest
|
||||
from litellm.rust_bridge.route import ComponentExecution, NativeLifecycle
|
||||
|
||||
NativeOcrLifecycle = NativeLifecycle[LiteLLMOcrRequest, OCRResponse]
|
||||
|
||||
|
||||
class ExceptionMapper(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
*,
|
||||
model: str,
|
||||
custom_llm_provider: str | None,
|
||||
original_exception: Exception,
|
||||
completion_kwargs: dict[str, object],
|
||||
extra_kwargs: dict[str, object],
|
||||
) -> Exception: ...
|
||||
|
||||
|
||||
def _binding(value: object) -> NativeOcrLifecycle | None:
|
||||
if not callable(value):
|
||||
return None
|
||||
return cast("NativeOcrLifecycle", value) # cast-ok: callable validated at the native binding boundary
|
||||
|
||||
|
||||
LIFECYCLE: Final = ROUTE.bind("_ocr_lifecycle", validate=_binding)
|
||||
LIFECYCLE: Final = COMPONENT.bind("_ocr_lifecycle", validate=_binding)
|
||||
NATIVE_OCR_LIFECYCLE: Final = LIFECYCLE
|
||||
|
||||
|
||||
def select(request: LiteLLMOcrRequest) -> NativeOcrLifecycle | None:
|
||||
def select(request: LiteLLMOcrRequest, execution: ComponentExecution) -> NativeOcrLifecycle | None:
|
||||
if request.kwargs.get("aocr"):
|
||||
return None
|
||||
return ROUTE.select(NATIVE_OCR_LIFECYCLE)
|
||||
|
||||
|
||||
def arguments(request: LiteLLMOcrRequest) -> Mapping[str, object]:
|
||||
return request.kwargs
|
||||
return execution.select(NATIVE_OCR_LIFECYCLE)
|
||||
|
||||
|
||||
def map_failure(error: Exception, request: LiteLLMOcrRequest, request_provider: str) -> Exception:
|
||||
mapper: Final = cast( # cast-ok: bounded adapter for the legacy public exception mapper
|
||||
ExceptionMapper, litellm.exception_type
|
||||
)
|
||||
try:
|
||||
return mapper(
|
||||
model=request.model.removeprefix(f"{request_provider}/"),
|
||||
custom_llm_provider=request_provider,
|
||||
original_exception=error,
|
||||
completion_kwargs=dict(arguments(request)), # mutable-ok: exception mapper requires owned kwargs
|
||||
extra_kwargs=dict(request.kwargs), # mutable-ok: exception mapper requires owned kwargs
|
||||
)
|
||||
except Exception as public_error:
|
||||
public_error.__context__ = error
|
||||
return public_error
|
||||
return HOST.map_failure(error, request, request_provider)
|
||||
|
|
|
|||
|
|
@ -9,9 +9,9 @@ from typing import Final, cast # noqa: TID251 # native extension exposes dynam
|
|||
import httpx
|
||||
|
||||
from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KEY, OCRResponse
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
from litellm.rust_bridge.ocr.definition import COMPONENT
|
||||
from litellm.rust_bridge.ocr.types import RustAocr, RustOcr
|
||||
from litellm.rust_bridge.route import NativeRoute
|
||||
from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke, invoke
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds as _timeout_to_seconds
|
||||
|
||||
|
||||
|
|
@ -23,20 +23,19 @@ def _as_aocr(value: object) -> RustAocr | None:
|
|||
return cast(RustAocr, value) if callable(value) else None
|
||||
|
||||
|
||||
ROUTE: Final = NativeRoute(RouteName.OCR)
|
||||
_OCR: Final = ROUTE.bind("ocr", validate=_as_ocr)
|
||||
_AOCR: Final = ROUTE.bind("aocr", validate=_as_aocr)
|
||||
_OCR: Final = COMPONENT.bind("ocr", validate=_as_ocr)
|
||||
_AOCR: Final = COMPONENT.bind("aocr", validate=_as_aocr)
|
||||
|
||||
|
||||
def load_rust_ocr() -> RustOcr | None:
|
||||
return ROUTE.select(_OCR)
|
||||
return COMPONENT.resolve().select(_OCR)
|
||||
|
||||
|
||||
def load_rust_aocr() -> RustAocr | None:
|
||||
return ROUTE.select(_AOCR)
|
||||
return COMPONENT.resolve().select(_AOCR)
|
||||
|
||||
|
||||
def _response(response: Mapping[str, object]) -> OCRResponse:
|
||||
def adapt_response(response: Mapping[str, object]) -> OCRResponse:
|
||||
provider_native_response: Final = response.get(PROVIDER_NATIVE_RESPONSE_KEY)
|
||||
normalized: Final = OCRResponse.model_validate(
|
||||
MappingProxyType({key: value for key, value in response.items() if key != PROVIDER_NATIVE_RESPONSE_KEY})
|
||||
|
|
@ -58,19 +57,28 @@ def ocr(
|
|||
timeout: float | httpx.Timeout | None,
|
||||
input_sources: Mapping[str, str] | None = None,
|
||||
) -> dict[str, object] | None:
|
||||
rust_ocr: Final = load_rust_ocr()
|
||||
if rust_ocr is None:
|
||||
return None
|
||||
return rust_ocr(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
input_sources=dict(input_sources or {}), # mutable-ok: native boundary requires a concrete dict
|
||||
timeout_seconds=_timeout_to_seconds(timeout),
|
||||
execution: Final = COMPONENT.resolve()
|
||||
rust_ocr: Final = execution.select(_OCR)
|
||||
return invoke(
|
||||
execution=execution,
|
||||
native_call=(
|
||||
lambda: rust_ocr(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
input_sources=dict(input_sources or {}), # mutable-ok: native boundary requires a concrete dict
|
||||
timeout_seconds=_timeout_to_seconds(timeout),
|
||||
)
|
||||
)
|
||||
if rust_ocr is not None
|
||||
else None,
|
||||
python_fallback=lambda: None,
|
||||
adapt=lambda value: value,
|
||||
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -86,17 +94,29 @@ async def aocr(
|
|||
timeout: float | httpx.Timeout | None,
|
||||
input_sources: Mapping[str, str] | None = None,
|
||||
) -> dict[str, object] | None:
|
||||
rust_aocr: Final = load_rust_aocr()
|
||||
if rust_aocr is None:
|
||||
execution: Final = COMPONENT.resolve()
|
||||
rust_aocr: Final = execution.select(_AOCR)
|
||||
async def python_fallback() -> None:
|
||||
return None
|
||||
return await rust_aocr(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
input_sources=dict(input_sources or {}), # mutable-ok: native boundary requires a concrete dict
|
||||
timeout_seconds=_timeout_to_seconds(timeout),
|
||||
|
||||
return await ainvoke(
|
||||
execution=execution,
|
||||
native_call=(
|
||||
lambda: rust_aocr(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
input_sources=dict(input_sources or {}), # mutable-ok: native boundary requires a concrete dict
|
||||
timeout_seconds=_timeout_to_seconds(timeout),
|
||||
)
|
||||
)
|
||||
if rust_aocr is not None
|
||||
else None,
|
||||
python_fallback=python_fallback,
|
||||
adapt=lambda value: value,
|
||||
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
from litellm.rust_bridge.route import NativeRoute
|
||||
from litellm.rust_bridge.rerank.definition import COMPONENT
|
||||
|
||||
ROUTE: Final = NativeRoute(RouteName.RERANK)
|
||||
__all__: Final = ("COMPONENT",)
|
||||
|
|
|
|||
6
litellm/rust_bridge/rerank/definition.py
Normal file
6
litellm/rust_bridge/rerank/definition.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.catalog import COMPONENTS
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
|
||||
COMPONENT: Final = COMPONENTS[RouteName.RERANK]
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.rerank import ROUTE
|
||||
from litellm.rust_bridge.rerank.definition import COMPONENT
|
||||
|
||||
LIFECYCLE: Final = ROUTE.lifecycle()
|
||||
LIFECYCLE: Final = COMPONENT.lifecycle()
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
from litellm.rust_bridge.route import NativeRoute
|
||||
from litellm.rust_bridge.responses.definition import COMPONENT
|
||||
|
||||
ROUTE: Final = NativeRoute(RouteName.RESPONSES)
|
||||
__all__: Final = ("COMPONENT",)
|
||||
|
|
|
|||
6
litellm/rust_bridge/responses/definition.py
Normal file
6
litellm/rust_bridge/responses/definition.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.catalog import COMPONENTS
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
|
||||
COMPONENT: Final = COMPONENTS[RouteName.RESPONSES]
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.responses import ROUTE
|
||||
from litellm.rust_bridge.responses.definition import COMPONENT
|
||||
|
||||
LIFECYCLE: Final = ROUTE.lifecycle()
|
||||
LIFECYCLE: Final = COMPONENT.lifecycle()
|
||||
|
|
|
|||
|
|
@ -2,14 +2,16 @@
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Final, Protocol, cast # noqa: TID251 # native class is validated at load time
|
||||
|
||||
import httpx
|
||||
from websockets.exceptions import ConnectionClosedOK
|
||||
|
||||
from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset
|
||||
from litellm.rust_bridge.responses import ROUTE
|
||||
from litellm.rust_bridge.configuration import CapabilityContext, DeliveryMode
|
||||
from litellm.rust_bridge.responses.definition import COMPONENT
|
||||
from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
|
||||
|
||||
|
|
@ -28,6 +30,7 @@ class RustResponsesWebSocketConnection(Protocol):
|
|||
url: str,
|
||||
headers: dict[str, str],
|
||||
timeout_seconds: float | None,
|
||||
custom_llm_provider: str | None,
|
||||
) -> Awaitable[RustResponsesWebSocket]: ...
|
||||
|
||||
|
||||
|
|
@ -37,7 +40,7 @@ def _as_connection(value: object) -> RustResponsesWebSocketConnection | None:
|
|||
return cast(RustResponsesWebSocketConnection, value) # cast-ok: native connection factory validated above
|
||||
|
||||
|
||||
_CONNECTION: Final = ROUTE.bind("ResponsesWebSocketConnection", validate=_as_connection)
|
||||
_CONNECTION: Final = COMPONENT.bind("ResponsesWebSocketConnection", validate=_as_connection)
|
||||
|
||||
|
||||
def set_rust_responses_websocket(
|
||||
|
|
@ -48,7 +51,7 @@ def set_rust_responses_websocket(
|
|||
|
||||
|
||||
def load_rust_responses_websocket() -> RustResponsesWebSocketConnection | None:
|
||||
return ROUTE.select(_CONNECTION)
|
||||
return COMPONENT.resolve(CapabilityContext(delivery=DeliveryMode.WEBSOCKET)).select(_CONNECTION)
|
||||
|
||||
|
||||
class _ConnectionAdapter:
|
||||
|
|
@ -73,16 +76,40 @@ async def connect(
|
|||
url: str,
|
||||
headers: dict[str, str],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
custom_llm_provider: str | None,
|
||||
model: str,
|
||||
) -> _ConnectionAdapter | None:
|
||||
connection_type: Final = load_rust_responses_websocket()
|
||||
if connection_type is None:
|
||||
return None
|
||||
try:
|
||||
connection: Final = await connection_type.connect(
|
||||
url=url,
|
||||
headers=headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
execution: Final = COMPONENT.resolve(
|
||||
CapabilityContext(
|
||||
provider=custom_llm_provider or "",
|
||||
model=model,
|
||||
delivery=DeliveryMode.WEBSOCKET,
|
||||
)
|
||||
except Exception: # noqa: BLE001 # bridge failures must fall back to Python
|
||||
)
|
||||
connection_type: Final = execution.select(_CONNECTION)
|
||||
|
||||
native_call: Final[Callable[[], Awaitable[RustResponsesWebSocket]] | None] = (
|
||||
(
|
||||
lambda: connection_type.connect(
|
||||
url=url,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
headers=headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
)
|
||||
)
|
||||
if connection_type is not None
|
||||
else None
|
||||
)
|
||||
async def python_fallback() -> None:
|
||||
return None
|
||||
|
||||
connection: Final = await ainvoke(
|
||||
execution=execution,
|
||||
native_call=native_call,
|
||||
python_fallback=python_fallback,
|
||||
adapt=lambda value: value,
|
||||
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
if connection is None:
|
||||
return None
|
||||
return _ConnectionAdapter(connection)
|
||||
|
|
|
|||
|
|
@ -3,17 +3,18 @@ from __future__ import annotations
|
|||
from collections.abc import Callable, Coroutine
|
||||
from dataclasses import dataclass
|
||||
from types import ModuleType
|
||||
from typing import (
|
||||
Final,
|
||||
Literal,
|
||||
Protocol,
|
||||
TypeVar,
|
||||
cast, # noqa: TID251 # validate callability at the native boundary
|
||||
overload,
|
||||
)
|
||||
from typing import Final, Literal, Protocol, TypeVar, cast, overload
|
||||
|
||||
from litellm.rust_bridge.bindings import NativeBinding
|
||||
from litellm.rust_bridge.configuration import ROUTE_POLICIES, RouteName, RoutePolicy, rust_enabled
|
||||
from litellm.rust_bridge.configuration import (
|
||||
CapabilityContext,
|
||||
CapabilitySpec,
|
||||
ExecutionDecision,
|
||||
RouteName,
|
||||
UtilityName,
|
||||
capability_decision,
|
||||
)
|
||||
from litellm.rust_bridge.errors import RustRouteUnavailableError, RustRouteUnsupportedError
|
||||
|
||||
BindingT = TypeVar("BindingT")
|
||||
RequestT = TypeVar("RequestT", contravariant=True)
|
||||
|
|
@ -28,6 +29,7 @@ class NativeLifecycle(Protocol[RequestT, ResponseT]):
|
|||
args: tuple[object, ...],
|
||||
kwargs: dict[str, object], # mutable-ok: PyO3 requires the original concrete dict
|
||||
asynchronous: Literal[False],
|
||||
host: object,
|
||||
) -> ResponseT: ...
|
||||
|
||||
@overload
|
||||
|
|
@ -37,6 +39,7 @@ class NativeLifecycle(Protocol[RequestT, ResponseT]):
|
|||
args: tuple[object, ...],
|
||||
kwargs: dict[str, object], # mutable-ok: PyO3 requires the original concrete dict
|
||||
asynchronous: Literal[True],
|
||||
host: object,
|
||||
) -> Coroutine[object, object, ResponseT]: ...
|
||||
|
||||
@overload
|
||||
|
|
@ -46,6 +49,7 @@ class NativeLifecycle(Protocol[RequestT, ResponseT]):
|
|||
args: tuple[object, ...],
|
||||
kwargs: dict[str, object], # mutable-ok: PyO3 requires the original concrete dict
|
||||
asynchronous: bool,
|
||||
host: object,
|
||||
) -> ResponseT | Coroutine[object, object, ResponseT]: ...
|
||||
|
||||
|
||||
|
|
@ -56,15 +60,39 @@ def _lifecycle(value: object) -> NativeLifecycle[object, object] | None:
|
|||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class NativeRoute:
|
||||
name: RouteName
|
||||
class ComponentExecution:
|
||||
route_name: RouteName | UtilityName
|
||||
decision: ExecutionDecision
|
||||
|
||||
@property
|
||||
def policy(self) -> RoutePolicy:
|
||||
return ROUTE_POLICIES[self.name]
|
||||
def require_supported(self) -> None:
|
||||
if self.decision is ExecutionDecision.UNSUPPORTED:
|
||||
raise RustRouteUnsupportedError(
|
||||
f"No Python or Rust implementation for {self.route_name.value}"
|
||||
)
|
||||
|
||||
def enabled(self) -> bool:
|
||||
return rust_enabled(self.name)
|
||||
def select(self, binding: NativeBinding[BindingT]) -> BindingT | None:
|
||||
self.require_supported()
|
||||
if self.decision is ExecutionDecision.PYTHON:
|
||||
return None
|
||||
selected: Final = binding.load()
|
||||
if selected is None and self.decision is ExecutionDecision.RUST_REQUIRED:
|
||||
raise RustRouteUnavailableError(
|
||||
f"Rust {self.route_name.value} bridge is unavailable"
|
||||
)
|
||||
return selected
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class NativeComponent:
|
||||
name: RouteName | UtilityName
|
||||
capability: CapabilitySpec
|
||||
exports: tuple[str, ...]
|
||||
|
||||
def resolve(self, context: CapabilityContext = CapabilityContext()) -> ComponentExecution:
|
||||
return ComponentExecution(
|
||||
route_name=self.name,
|
||||
decision=capability_decision(self.capability, context=context),
|
||||
)
|
||||
|
||||
def bind(
|
||||
self,
|
||||
|
|
@ -73,11 +101,10 @@ class NativeRoute:
|
|||
validate: Callable[[object], BindingT | None],
|
||||
module_loader: Callable[[], ModuleType | None] | None = None,
|
||||
) -> NativeBinding[BindingT]:
|
||||
if export not in self.exports:
|
||||
raise ValueError(f"native export {export!r} is not declared for {self.name.value}")
|
||||
return NativeBinding(export, validate=validate, module_loader=module_loader)
|
||||
|
||||
def select(self, binding: NativeBinding[BindingT]) -> BindingT | None:
|
||||
return binding.load() if self.enabled() else None
|
||||
|
||||
def lifecycle(self) -> NativeBinding[NativeLifecycle[object, object]]:
|
||||
export: Final = f"_{self.name.value}_lifecycle"
|
||||
return self.bind(export, validate=_lifecycle)
|
||||
|
|
|
|||
|
|
@ -2,39 +2,22 @@ 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 typing import Final, NoReturn, TypeVar
|
||||
|
||||
from litellm.exceptions import APIError
|
||||
from litellm.rust_bridge.bindings import native_exception_types
|
||||
from litellm.rust_bridge.bindings import (
|
||||
native_exception_types,
|
||||
native_host_callback_exception,
|
||||
native_unavailable_exception,
|
||||
)
|
||||
from litellm.rust_bridge.configuration import ExecutionDecision
|
||||
from litellm.rust_bridge.errors import RustRouteDeclinedError, RustRouteUnavailableError
|
||||
from litellm.rust_bridge.route import ComponentExecution
|
||||
|
||||
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
|
||||
|
|
@ -45,95 +28,120 @@ class BridgeErrorContext:
|
|||
def invoke(
|
||||
*,
|
||||
native_call: Callable[[], NativeT] | None,
|
||||
fallback: Callable[[], ResultT],
|
||||
python_fallback: Callable[[], ResultT],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
mode: FallbackMode,
|
||||
execution: ComponentExecution,
|
||||
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)
|
||||
execution.require_supported()
|
||||
if execution.decision is ExecutionDecision.PYTHON:
|
||||
return python_fallback()
|
||||
if native_call is None:
|
||||
return _unavailable_or_fallback(execution, python_fallback)
|
||||
|
||||
exceptions: Final = native_exception_types()
|
||||
if exceptions is None:
|
||||
return adapt(native_call())
|
||||
declined, upstream = exceptions
|
||||
unavailable: Final = native_unavailable_exception()
|
||||
host_callback: Final = native_host_callback_exception()
|
||||
try:
|
||||
value: Final = native_call()
|
||||
except host_callback as error:
|
||||
_raise_host_callback(error)
|
||||
except unavailable:
|
||||
return _unavailable_or_fallback(execution, python_fallback)
|
||||
except declined as error:
|
||||
return _declined_or_fallback(execution, python_fallback, error)
|
||||
except upstream as error:
|
||||
_raise_upstream(error, context)
|
||||
return adapt(value)
|
||||
|
||||
|
||||
async def ainvoke(
|
||||
*,
|
||||
native_call: Callable[[], Awaitable[NativeT]] | None,
|
||||
fallback: Callable[[], Awaitable[ResultT]],
|
||||
python_fallback: Callable[[], Awaitable[ResultT]],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
mode: FallbackMode,
|
||||
execution: ComponentExecution,
|
||||
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]:
|
||||
execution.require_supported()
|
||||
if execution.decision is ExecutionDecision.PYTHON:
|
||||
return await python_fallback()
|
||||
if native_call is None:
|
||||
return RustUnavailable()
|
||||
return await _aunavailable_or_fallback(execution, python_fallback)
|
||||
|
||||
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()))
|
||||
return adapt(await native_call())
|
||||
declined, upstream = exceptions
|
||||
unavailable: Final = native_unavailable_exception()
|
||||
host_callback: Final = native_host_callback_exception()
|
||||
try:
|
||||
value: Final = await native_call()
|
||||
except host_callback as error:
|
||||
_raise_host_callback(error)
|
||||
except unavailable:
|
||||
return await _aunavailable_or_fallback(execution, python_fallback)
|
||||
except declined as error:
|
||||
return RustDeclined(reason=_decline_reason(error))
|
||||
return await _adeclined_or_fallback(execution, python_fallback, error)
|
||||
except upstream as error:
|
||||
_raise_upstream(error, context)
|
||||
return RustHandled(adapt(value))
|
||||
return adapt(value)
|
||||
|
||||
|
||||
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 _unavailable_or_fallback(
|
||||
execution: ComponentExecution,
|
||||
python_fallback: Callable[[], ResultT],
|
||||
) -> ResultT:
|
||||
if execution.decision is ExecutionDecision.RUST_WITH_FALLBACK:
|
||||
return python_fallback()
|
||||
raise RustRouteUnavailableError(f"Rust {execution.route_name.value} bridge is unavailable")
|
||||
|
||||
|
||||
def _raise_required(
|
||||
result: RustDeclined | RustUnavailable,
|
||||
context: BridgeErrorContext,
|
||||
) -> NoReturn:
|
||||
raise RuntimeError(f"Rust {context.route} bridge {_required_reason(result)}")
|
||||
async def _aunavailable_or_fallback(
|
||||
execution: ComponentExecution,
|
||||
python_fallback: Callable[[], Awaitable[ResultT]],
|
||||
) -> ResultT:
|
||||
if execution.decision is ExecutionDecision.RUST_WITH_FALLBACK:
|
||||
return await python_fallback()
|
||||
raise RustRouteUnavailableError(f"Rust {execution.route_name.value} bridge is unavailable")
|
||||
|
||||
|
||||
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 _declined_or_fallback(
|
||||
execution: ComponentExecution,
|
||||
python_fallback: Callable[[], ResultT],
|
||||
error: BaseException,
|
||||
) -> ResultT:
|
||||
if execution.decision is ExecutionDecision.RUST_WITH_FALLBACK:
|
||||
return python_fallback()
|
||||
_raise_declined(execution, error)
|
||||
|
||||
|
||||
async def _adeclined_or_fallback(
|
||||
execution: ComponentExecution,
|
||||
python_fallback: Callable[[], Awaitable[ResultT]],
|
||||
error: BaseException,
|
||||
) -> ResultT:
|
||||
if execution.decision is ExecutionDecision.RUST_WITH_FALLBACK:
|
||||
return await python_fallback()
|
||||
_raise_declined(execution, error)
|
||||
|
||||
|
||||
def _raise_declined(execution: ComponentExecution, error: BaseException) -> NoReturn:
|
||||
reason_value: Final[object] = error.args[0] if error.args else str(error)
|
||||
reason: Final = reason_value if isinstance(reason_value, str) else str(reason_value)
|
||||
raise RustRouteDeclinedError(
|
||||
f"Rust {execution.route_name.value} bridge declined the request: {reason}"
|
||||
) from error
|
||||
|
||||
|
||||
def _raise_host_callback(error: BaseException) -> NoReturn:
|
||||
cause: Final = error.__cause__
|
||||
if cause is not None:
|
||||
raise cause
|
||||
raise error
|
||||
|
||||
|
||||
def _raise_upstream(error: BaseException, context: BridgeErrorContext) -> NoReturn:
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
from litellm.rust_bridge.route import NativeRoute
|
||||
from litellm.rust_bridge.speech.definition import COMPONENT
|
||||
|
||||
ROUTE: Final = NativeRoute(RouteName.SPEECH)
|
||||
__all__: Final = ("COMPONENT",)
|
||||
|
|
|
|||
6
litellm/rust_bridge/speech/definition.py
Normal file
6
litellm/rust_bridge/speech/definition.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.catalog import COMPONENTS
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
|
||||
COMPONENT: Final = COMPONENTS[RouteName.SPEECH]
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.speech import ROUTE
|
||||
from litellm.rust_bridge.speech.definition import COMPONENT
|
||||
|
||||
LIFECYCLE: Final = ROUTE.lifecycle()
|
||||
LIFECYCLE: Final = COMPONENT.lifecycle()
|
||||
|
|
|
|||
|
|
@ -1,114 +0,0 @@
|
|||
"""Thin Python wrapper for the native Rust input token counter."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable
|
||||
from dataclasses import dataclass
|
||||
from functools import lru_cache
|
||||
from typing import Final, Literal, Protocol, cast # noqa: TID251 # native extension exposes untyped callables
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.default_encoding import cl100k_base_rank_file, o200k_base_rank_file
|
||||
from litellm.litellm_core_utils.token_counter import openai_tokenizer_encoding, uses_legacy_message_accounting
|
||||
from litellm.rust_bridge.bindings import NativeBinding
|
||||
from litellm.rust_bridge.configuration import rust_enabled
|
||||
from litellm.rust_bridge.runtime import BridgeErrorContext, RustHandled, aattempt
|
||||
from litellm.utils import claude_json_str, huggingface_tokenizer_kind
|
||||
|
||||
RustTokenizer = Literal["anthropic", "cl100k_base", "o200k_base"]
|
||||
|
||||
|
||||
class RustTokenCounter(Protocol):
|
||||
def acount_request(self, body: bytes) -> Awaitable[object]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class RustTokenCounterFactory(Protocol):
|
||||
def __call__(self, tokenizer_json: str) -> RustTokenCounter:
|
||||
raise NotImplementedError
|
||||
|
||||
def from_cl100k_ranks(self, rank_file: str) -> RustTokenCounter:
|
||||
raise NotImplementedError
|
||||
|
||||
def from_o200k_ranks(self, rank_file: str) -> RustTokenCounter:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class InputTokenCount:
|
||||
model: str | None
|
||||
input_tokens: int
|
||||
|
||||
|
||||
_INPUT_TOKEN_COUNT: Final = TypeAdapter(InputTokenCount)
|
||||
|
||||
|
||||
def _as_factory(value: object) -> RustTokenCounterFactory | None:
|
||||
return (
|
||||
cast( # cast-ok: native extension protocol is runtime-defined
|
||||
RustTokenCounterFactory, value
|
||||
)
|
||||
if callable(value)
|
||||
else None
|
||||
)
|
||||
|
||||
|
||||
TOKEN_COUNTER: Final = NativeBinding("TokenCounter", validate=_as_factory)
|
||||
|
||||
|
||||
def rust_tokenizer(model: str) -> RustTokenizer | None:
|
||||
"""The Rust counter for the tokenizer `litellm.token_counter` selects for `model`, `None` when Python must count.
|
||||
|
||||
Mirrors `_select_tokenizer_helper`: the Anthropic tokenizer has a Rust port, the other HuggingFace
|
||||
downloads do not, and of the tiktoken encodings `cl100k_base` and `o200k_base` do (p50k/r50k do not). Rust
|
||||
prices every message with the default constants, so the legacy `gpt-3.5-turbo-0301` accounting stays in
|
||||
Python."""
|
||||
if litellm.disable_token_counter is True:
|
||||
return None
|
||||
kind: Final = None if litellm.disable_hf_tokenizer_download is True else huggingface_tokenizer_kind(model)
|
||||
if kind == "anthropic":
|
||||
return "anthropic"
|
||||
if kind is not None or uses_legacy_message_accounting(model):
|
||||
return None
|
||||
match openai_tokenizer_encoding(model).name:
|
||||
case "cl100k_base":
|
||||
return "cl100k_base"
|
||||
case "o200k_base":
|
||||
return "o200k_base"
|
||||
case _:
|
||||
return None
|
||||
|
||||
|
||||
@lru_cache(maxsize=4)
|
||||
def _counter(factory: RustTokenCounterFactory, tokenizer: RustTokenizer) -> RustTokenCounter:
|
||||
match tokenizer:
|
||||
case "anthropic":
|
||||
return factory(claude_json_str)
|
||||
case "cl100k_base":
|
||||
return factory.from_cl100k_ranks(cl100k_base_rank_file())
|
||||
case "o200k_base":
|
||||
return factory.from_o200k_ranks(o200k_base_rank_file())
|
||||
|
||||
|
||||
async def count_input_tokens(body: bytes, tokenizer: RustTokenizer) -> InputTokenCount | None:
|
||||
if not rust_enabled():
|
||||
return None
|
||||
factory: Final = TOKEN_COUNTER.load()
|
||||
if factory is None:
|
||||
return None
|
||||
try:
|
||||
attempt: Final = await aattempt(
|
||||
native_call=lambda: _counter(factory, tokenizer).acount_request(body),
|
||||
adapt=_INPUT_TOKEN_COUNT.validate_python,
|
||||
context=BridgeErrorContext(route="token_counter", provider=tokenizer, model=""),
|
||||
)
|
||||
except (RuntimeError, ValueError) as error:
|
||||
verbose_logger.debug("Rust token counter (%s) failed, counting in Python: %s", tokenizer, error)
|
||||
return None
|
||||
if not isinstance(attempt, RustHandled):
|
||||
return None
|
||||
verbose_logger.debug("Rust token counter (%s) counted %d input tokens", tokenizer, attempt.value.input_tokens)
|
||||
return attempt.value
|
||||
12
litellm/rust_bridge/token_counter/__init__.py
Normal file
12
litellm/rust_bridge/token_counter/__init__.py
Normal file
|
|
@ -0,0 +1,12 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.token_counter.types import InputTokenCount, RustTokenizer
|
||||
from litellm.rust_bridge.token_counter.value import TOKEN_COUNTER, count_input_tokens, rust_tokenizer
|
||||
|
||||
__all__: Final = (
|
||||
"TOKEN_COUNTER",
|
||||
"InputTokenCount",
|
||||
"RustTokenizer",
|
||||
"count_input_tokens",
|
||||
"rust_tokenizer",
|
||||
)
|
||||
6
litellm/rust_bridge/token_counter/definition.py
Normal file
6
litellm/rust_bridge/token_counter/definition.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.catalog import COMPONENTS
|
||||
from litellm.rust_bridge.configuration import UtilityName
|
||||
|
||||
COMPONENT: Final = COMPONENTS[UtilityName.TOKEN_COUNTER]
|
||||
31
litellm/rust_bridge/token_counter/types.py
Normal file
31
litellm/rust_bridge/token_counter/types.py
Normal file
|
|
@ -0,0 +1,31 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Protocol
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RustTokenizer:
|
||||
kind: str | None
|
||||
encoding: str
|
||||
disabled: bool
|
||||
legacy_accounting: bool
|
||||
|
||||
|
||||
class RustTokenCounter(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
body: bytes,
|
||||
kind: str | None,
|
||||
encoding: str,
|
||||
disabled: bool,
|
||||
legacy_accounting: bool,
|
||||
resource_loader: Callable[[str], str],
|
||||
) -> Awaitable[object]: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class InputTokenCount:
|
||||
model: str | None
|
||||
input_tokens: int
|
||||
75
litellm/rust_bridge/token_counter/value.py
Normal file
75
litellm/rust_bridge/token_counter/value.py
Normal file
|
|
@ -0,0 +1,75 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from typing import Final, cast # noqa: TID251 # native extension exposes untyped callables
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.default_encoding import cl100k_base_rank_file, o200k_base_rank_file
|
||||
from litellm.litellm_core_utils.token_counter import openai_tokenizer_encoding, uses_legacy_message_accounting
|
||||
from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke
|
||||
from litellm.rust_bridge.token_counter.definition import COMPONENT
|
||||
from litellm.rust_bridge.token_counter.types import InputTokenCount, RustTokenCounter, RustTokenizer
|
||||
from litellm.utils import claude_json_str, huggingface_tokenizer_kind
|
||||
|
||||
_INPUT_TOKEN_COUNT: Final = TypeAdapter(InputTokenCount)
|
||||
|
||||
|
||||
def _as_counter(value: object) -> RustTokenCounter | None:
|
||||
return cast(RustTokenCounter, value) if callable(value) else None
|
||||
|
||||
|
||||
TOKEN_COUNTER: Final = COMPONENT.bind("count_input_tokens", validate=_as_counter)
|
||||
|
||||
|
||||
def rust_tokenizer(model: str) -> RustTokenizer | None:
|
||||
execution: Final = COMPONENT.resolve()
|
||||
if execution.select(TOKEN_COUNTER) is None:
|
||||
return None
|
||||
kind: Final = None if litellm.disable_hf_tokenizer_download is True else huggingface_tokenizer_kind(model)
|
||||
encoding: Final = openai_tokenizer_encoding(model).name if kind is None else ""
|
||||
return RustTokenizer(
|
||||
kind=kind,
|
||||
encoding=encoding,
|
||||
disabled=litellm.disable_token_counter is True,
|
||||
legacy_accounting=uses_legacy_message_accounting(model),
|
||||
)
|
||||
|
||||
|
||||
def _tokenizer_resource(tokenizer: str) -> str:
|
||||
match tokenizer:
|
||||
case "anthropic":
|
||||
return claude_json_str
|
||||
case "cl100k_base":
|
||||
return cl100k_base_rank_file()
|
||||
case "o200k_base":
|
||||
return o200k_base_rank_file()
|
||||
case _:
|
||||
raise ValueError(f"unsupported Rust tokenizer resource: {tokenizer}")
|
||||
|
||||
|
||||
async def count_input_tokens(body: bytes, tokenizer: RustTokenizer) -> InputTokenCount | None:
|
||||
execution: Final = COMPONENT.resolve()
|
||||
counter: Final = execution.select(TOKEN_COUNTER)
|
||||
|
||||
async def python_fallback() -> None:
|
||||
return None
|
||||
|
||||
return await ainvoke(
|
||||
execution=execution,
|
||||
native_call=(
|
||||
lambda: counter(
|
||||
body,
|
||||
tokenizer.kind,
|
||||
tokenizer.encoding,
|
||||
tokenizer.disabled,
|
||||
tokenizer.legacy_accounting,
|
||||
_tokenizer_resource,
|
||||
)
|
||||
)
|
||||
if counter is not None
|
||||
else None,
|
||||
python_fallback=python_fallback,
|
||||
adapt=_INPUT_TOKEN_COUNT.validate_python,
|
||||
context=BridgeErrorContext(route=COMPONENT.name.value, provider="", model=""),
|
||||
)
|
||||
|
|
@ -1,8 +1,8 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.transcription.definition import COMPONENT
|
||||
from litellm.rust_bridge.transcription.types import RustAtranscription, RustTranscription
|
||||
from litellm.rust_bridge.transcription.value import (
|
||||
ROUTE,
|
||||
atranscription,
|
||||
configure_rust_transcription,
|
||||
load_rust_atranscription,
|
||||
|
|
@ -11,7 +11,7 @@ from litellm.rust_bridge.transcription.value import (
|
|||
)
|
||||
|
||||
__all__: Final = (
|
||||
"ROUTE",
|
||||
"COMPONENT",
|
||||
"RustAtranscription",
|
||||
"RustTranscription",
|
||||
"atranscription",
|
||||
|
|
|
|||
6
litellm/rust_bridge/transcription/definition.py
Normal file
6
litellm/rust_bridge/transcription/definition.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.catalog import COMPONENTS
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
|
||||
COMPONENT: Final = COMPONENTS[RouteName.TRANSCRIPTION]
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.transcription import ROUTE
|
||||
from litellm.rust_bridge.transcription.definition import COMPONENT
|
||||
|
||||
LIFECYCLE: Final = ROUTE.lifecycle()
|
||||
LIFECYCLE: Final = COMPONENT.lifecycle()
|
||||
|
|
|
|||
|
|
@ -8,13 +8,12 @@ from typing import (
|
|||
import httpx
|
||||
|
||||
from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
from litellm.rust_bridge.route import NativeRoute
|
||||
from litellm.rust_bridge.configuration import CapabilityContext
|
||||
from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke, invoke
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
from litellm.rust_bridge.transcription.definition import COMPONENT
|
||||
from litellm.rust_bridge.transcription.types import RustAtranscription, RustTranscription
|
||||
|
||||
ROUTE: Final = NativeRoute(RouteName.TRANSCRIPTION)
|
||||
|
||||
|
||||
def _as_transcription(value: object) -> RustTranscription | None:
|
||||
return cast(RustTranscription, value) if callable(value) else None # cast-ok: validated callable native binding
|
||||
|
|
@ -24,8 +23,8 @@ def _as_atranscription(value: object) -> RustAtranscription | None:
|
|||
return cast(RustAtranscription, value) if callable(value) else None # cast-ok: validated callable native binding
|
||||
|
||||
|
||||
_TRANSCRIPTION: Final = ROUTE.bind("transcription", validate=_as_transcription)
|
||||
_ATRANSCRIPTION: Final = ROUTE.bind("atranscription", validate=_as_atranscription)
|
||||
_TRANSCRIPTION: Final = COMPONENT.bind("transcription", validate=_as_transcription)
|
||||
_ATRANSCRIPTION: Final = COMPONENT.bind("atranscription", validate=_as_atranscription)
|
||||
|
||||
|
||||
def configure_rust_transcription(
|
||||
|
|
@ -37,12 +36,12 @@ def configure_rust_transcription(
|
|||
_ATRANSCRIPTION.configure(atranscription)
|
||||
|
||||
|
||||
def load_rust_transcription() -> RustTranscription | None:
|
||||
return ROUTE.select(_TRANSCRIPTION)
|
||||
def load_rust_transcription(*, context: CapabilityContext) -> RustTranscription | None:
|
||||
return COMPONENT.resolve(context).select(_TRANSCRIPTION)
|
||||
|
||||
|
||||
def load_rust_atranscription() -> RustAtranscription | None:
|
||||
return ROUTE.select(_ATRANSCRIPTION)
|
||||
def load_rust_atranscription(*, context: CapabilityContext) -> RustAtranscription | None:
|
||||
return COMPONENT.resolve(context).select(_ATRANSCRIPTION)
|
||||
|
||||
|
||||
def transcription(
|
||||
|
|
@ -56,18 +55,27 @@ def transcription(
|
|||
optional_params: dict[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> dict[str, object] | None:
|
||||
rust_transcription: Final = load_rust_transcription()
|
||||
if rust_transcription is None:
|
||||
return None
|
||||
return rust_transcription(
|
||||
model=model,
|
||||
audio=audio,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
execution: Final = COMPONENT.resolve(CapabilityContext(provider=custom_llm_provider or "", model=model))
|
||||
rust_transcription: Final = execution.select(_TRANSCRIPTION)
|
||||
return invoke(
|
||||
execution=execution,
|
||||
native_call=(
|
||||
lambda: rust_transcription(
|
||||
model=model,
|
||||
audio=audio,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
)
|
||||
)
|
||||
if rust_transcription is not None
|
||||
else None,
|
||||
python_fallback=lambda: None,
|
||||
adapt=lambda response: response,
|
||||
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -82,16 +90,28 @@ async def atranscription(
|
|||
optional_params: dict[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> dict[str, object] | None:
|
||||
rust_atranscription: Final = load_rust_atranscription()
|
||||
if rust_atranscription is None:
|
||||
execution: Final = COMPONENT.resolve(CapabilityContext(provider=custom_llm_provider or "", model=model))
|
||||
rust_atranscription: Final = execution.select(_ATRANSCRIPTION)
|
||||
async def python_fallback() -> None:
|
||||
return None
|
||||
return await rust_atranscription(
|
||||
model=model,
|
||||
audio=audio,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
|
||||
return await ainvoke(
|
||||
execution=execution,
|
||||
native_call=(
|
||||
lambda: rust_atranscription(
|
||||
model=model,
|
||||
audio=audio,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
)
|
||||
)
|
||||
if rust_atranscription is not None
|
||||
else None,
|
||||
python_fallback=python_fallback,
|
||||
adapt=lambda response: response,
|
||||
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -181,6 +181,7 @@ dev = [
|
|||
"hypothesis==6.165.10",
|
||||
"reportlab==5.0.1",
|
||||
"basedpyright==1.39.7",
|
||||
"mypy==1.20.1",
|
||||
"keyring==25.7.0",
|
||||
"pytest==9.0.3",
|
||||
"tomli==2.4.1; python_version < '3.11'",
|
||||
|
|
@ -288,6 +289,7 @@ profile = "release"
|
|||
editable-profile = "dev"
|
||||
include = [
|
||||
"litellm/proxy/_experimental/out/**",
|
||||
"litellm/rust_bridge/_native.pyi",
|
||||
"litellm/router_strategy/complexity_router/artifacts/*.json",
|
||||
]
|
||||
exclude = [
|
||||
|
|
|
|||
|
|
@ -57,11 +57,11 @@ PUBLIC_RUST_DISPATCH_MAPPINGS: Final = (
|
|||
mapping(span="public_sdk_entrypoint", python_frame=r"ocr/main\.py:\d+ a?ocr$"),
|
||||
mapping(span="public_request", python_frame=r"ocr/main\.py:\d+ _public_request$"),
|
||||
mapping(span="bind_request", python_frame=r"ocr/main\.py:\d+ _bind_request$"),
|
||||
mapping(span="rust_ocr_enabled", python_frame=r"rust_bridge/configuration\.py:\d+ rust_ocr_enabled$"),
|
||||
mapping(span="rust_ocr_enabled", python_frame=r"rust_bridge/configuration\.py:\d+ capability_decision$"),
|
||||
mapping(span="select_native_ocr", python_frame=r"rust_bridge/ocr/lifecycle\.py:\d+ select$"),
|
||||
mapping(span="load_native_bridge", python_frame=r"rust_bridge/bindings\.py:\d+ NativeBinding\.load$"),
|
||||
mapping(span="native_call_setup", python_frame=r"rust_bridge/lifecycle\.py:\d+ setup$"),
|
||||
mapping(span="native_response", python_frame=r"rust_bridge/ocr\.py:\d+ _response$"),
|
||||
mapping(span="native_response", python_frame=r"rust_bridge/ocr/host\.py:\d+ OcrLifecycleHost\.response$"),
|
||||
mapping(span="native_call_finalize", python_frame=r"rust_bridge/lifecycle\.py:\d+ finalize$"),
|
||||
mapping(
|
||||
span="native_success_bookkeeping",
|
||||
|
|
|
|||
|
|
@ -47,6 +47,7 @@ class RecordingMessages:
|
|||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
timeout_seconds: float | None,
|
||||
has_agentic_hook: bool = False,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
{
|
||||
|
|
@ -75,6 +76,7 @@ class RecordingAsyncMessages:
|
|||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
timeout_seconds: float | None,
|
||||
has_agentic_hook: bool = False,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
{
|
||||
|
|
@ -238,19 +240,19 @@ async def test_gate_invokes_rust_and_marks_response_header():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_falls_back_to_python_when_bridge_raises():
|
||||
async def test_gate_propagates_unclassified_bridge_failure():
|
||||
bridge = RaisingAsyncMessages()
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
response = await _gate()
|
||||
|
||||
assert response is None
|
||||
with pytest.raises(RuntimeError, match="upstream request failed"):
|
||||
await _gate()
|
||||
assert bridge.calls == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_skips_rust_when_flag_absent():
|
||||
async def test_gate_skips_rust_when_flag_absent(monkeypatch):
|
||||
monkeypatch.delenv("LITELLM_RUST", raising=False)
|
||||
bridge = ExplodingAsyncMessages()
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
|
|
@ -324,26 +326,20 @@ async def test_gate_env_var_falsey_does_not_enable(monkeypatch):
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_skips_rust_for_unsupported_provider():
|
||||
bridge = ExplodingAsyncMessages()
|
||||
native = pytest.importorskip("litellm.rust_bridge._native")
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
response = await _gate(custom_llm_provider="openai")
|
||||
|
||||
rust_messages.set_rust_messages(amessages=native.amessages)
|
||||
response = await _gate(custom_llm_provider="openai", api_base="http://127.0.0.1:1")
|
||||
assert response is None
|
||||
assert bridge.calls == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_skips_rust_for_agentic_hook():
|
||||
bridge = ExplodingAsyncMessages()
|
||||
native = pytest.importorskip("litellm.rust_bridge._native")
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
response = await _gate(has_agentic_hook=True)
|
||||
|
||||
rust_messages.set_rust_messages(amessages=native.amessages)
|
||||
response = await _gate(has_agentic_hook=True, api_base="http://127.0.0.1:1")
|
||||
assert response is None
|
||||
assert bridge.calls == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -13,9 +13,9 @@ from unittest.mock import MagicMock, patch
|
|||
import boto3
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from botocore.credentials import Credentials
|
||||
from botocore.exceptions import ClientError
|
||||
|
||||
from litellm.llms.bedrock.chat.converse_handler import BedrockConverseLLM
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.rust_bridge import chat_completions as bridge
|
||||
|
|
@ -54,29 +54,27 @@ RESOLVED_CREDENTIALS = Credentials(
|
|||
@pytest.fixture(autouse=True)
|
||||
def reset_bridge(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
bridge.set_rust_chat_completions(
|
||||
chat_completions=None, achat_completions=None, decline=None
|
||||
)
|
||||
bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None)
|
||||
yield
|
||||
bridge.set_rust_chat_completions(
|
||||
chat_completions=None, achat_completions=None, decline=None
|
||||
)
|
||||
bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None)
|
||||
|
||||
|
||||
def _inject(*, decline_reason=None, error: Exception | None = None):
|
||||
seen: dict[str, list[dict]] = {"gate": [], "call": []}
|
||||
|
||||
def gate(**kwargs):
|
||||
seen["gate"].append(kwargs)
|
||||
return decline_reason
|
||||
|
||||
def native(**kwargs):
|
||||
seen["call"].append(kwargs)
|
||||
seen["gate"].append(kwargs)
|
||||
if decline_reason is not None:
|
||||
from litellm.rust_bridge import _native
|
||||
|
||||
raise _native.RustBridgeDeclined(decline_reason)
|
||||
if error is not None:
|
||||
raise error
|
||||
seen["call"].append(kwargs)
|
||||
kwargs["on_request"]()
|
||||
return dict(RUST_RESPONSE)
|
||||
|
||||
bridge.set_rust_chat_completions(decline=gate, chat_completions=native)
|
||||
bridge.set_rust_chat_completions(chat_completions=native)
|
||||
return seen
|
||||
|
||||
|
||||
|
|
@ -142,9 +140,7 @@ def test_the_core_receives_the_converse_url_this_handler_already_built():
|
|||
seen = _inject()
|
||||
_run()
|
||||
|
||||
assert seen["call"][0]["api_base"].endswith(
|
||||
"/model/anthropic.claude-sonnet-4-5-v1%3A0/converse"
|
||||
)
|
||||
assert seen["call"][0]["api_base"].endswith("/model/anthropic.claude-sonnet-4-5-v1%3A0/converse")
|
||||
assert "bedrock-runtime.us-east-1.amazonaws.com" in seen["call"][0]["api_base"]
|
||||
|
||||
|
||||
|
|
@ -215,9 +211,7 @@ async def test_the_async_path_falls_back_when_the_core_declines(monkeypatch):
|
|||
async def declining_native(**_kwargs):
|
||||
raise _Declined("blank message text")
|
||||
|
||||
bridge.set_rust_chat_completions(
|
||||
decline=lambda **_kwargs: None, achat_completions=declining_native
|
||||
)
|
||||
bridge.set_rust_chat_completions(achat_completions=declining_native)
|
||||
|
||||
sentinel = object()
|
||||
|
||||
|
|
@ -225,16 +219,10 @@ async def test_the_async_path_falls_back_when_the_core_declines(monkeypatch):
|
|||
return sentinel
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS
|
||||
),
|
||||
patch.object(
|
||||
BedrockConverseLLM, "async_completion", side_effect=python_path
|
||||
) as python_call,
|
||||
patch.object(BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS),
|
||||
patch.object(BedrockConverseLLM, "async_completion", side_effect=python_path) as python_call,
|
||||
):
|
||||
result = await BedrockConverseLLM().completion(
|
||||
**_completion_kwargs(acompletion=True)
|
||||
)
|
||||
result = await BedrockConverseLLM().completion(**_completion_kwargs(acompletion=True))
|
||||
|
||||
assert result is sentinel
|
||||
assert python_call.called, "a failing rust call must re-enter the python path"
|
||||
|
|
@ -242,22 +230,17 @@ async def test_the_async_path_falls_back_when_the_core_declines(monkeypatch):
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_async_path_serves_the_rust_response_without_the_fallback():
|
||||
async def native(**_kwargs):
|
||||
async def native(**kwargs):
|
||||
kwargs["on_request"]()
|
||||
return dict(RUST_RESPONSE)
|
||||
|
||||
bridge.set_rust_chat_completions(
|
||||
decline=lambda **_kwargs: None, achat_completions=native
|
||||
)
|
||||
bridge.set_rust_chat_completions(achat_completions=native)
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS
|
||||
),
|
||||
patch.object(BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS),
|
||||
patch.object(BedrockConverseLLM, "async_completion") as python_call,
|
||||
):
|
||||
result = await BedrockConverseLLM().completion(
|
||||
**_completion_kwargs(acompletion=True)
|
||||
)
|
||||
result = await BedrockConverseLLM().completion(**_completion_kwargs(acompletion=True))
|
||||
|
||||
assert result.choices[0].message.content == "hello from rust"
|
||||
assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
|
|
@ -284,26 +267,19 @@ async def test_pre_call_logging_fires_once_even_when_the_rust_path_declines():
|
|||
|
||||
async def python_path(**kwargs):
|
||||
served.append(kwargs)
|
||||
logging_obj.pre_call(input=kwargs["messages"], api_key="", additional_args={})
|
||||
return ModelResponse()
|
||||
|
||||
with (
|
||||
patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()),
|
||||
patch.object(
|
||||
BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS
|
||||
),
|
||||
patch.object(
|
||||
BedrockConverseLLM, "async_completion", side_effect=python_path
|
||||
),
|
||||
patch.object(BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS),
|
||||
patch.object(BedrockConverseLLM, "async_completion", side_effect=python_path),
|
||||
):
|
||||
bridge.set_rust_chat_completions(
|
||||
decline=lambda **_kwargs: None, achat_completions=declining_native
|
||||
)
|
||||
await BedrockConverseLLM().completion(
|
||||
**_completion_kwargs(acompletion=True, logging_obj=logging_obj)
|
||||
)
|
||||
bridge.set_rust_chat_completions(achat_completions=declining_native)
|
||||
await BedrockConverseLLM().completion(**_completion_kwargs(acompletion=True, logging_obj=logging_obj))
|
||||
|
||||
assert logging_obj.pre_call.call_count == 1
|
||||
assert served and served[0]["skip_pre_call_logging"] is True
|
||||
assert served and "skip_pre_call_logging" not in served[0]
|
||||
|
||||
|
||||
CONVERSE_RESPONSE = {
|
||||
|
|
@ -313,9 +289,7 @@ CONVERSE_RESPONSE = {
|
|||
}
|
||||
|
||||
|
||||
async def _drive_async_completion(
|
||||
*, skip_pre_call_logging: bool, logging_obj, credentials: Credentials = RESOLVED_CREDENTIALS
|
||||
):
|
||||
async def _drive_async_completion(*, logging_obj, credentials: Credentials = RESOLVED_CREDENTIALS):
|
||||
"""Run the real `async_completion` with a stubbed transport."""
|
||||
import httpx as _httpx
|
||||
|
||||
|
|
@ -345,22 +319,14 @@ async def _drive_async_completion(
|
|||
credentials=credentials,
|
||||
headers={},
|
||||
client=client,
|
||||
skip_pre_call_logging=skip_pre_call_logging,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_completion_honors_the_pre_call_suppression():
|
||||
logging_obj = MagicMock()
|
||||
await _drive_async_completion(skip_pre_call_logging=True, logging_obj=logging_obj)
|
||||
assert logging_obj.pre_call.call_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_completion_logs_pre_call_by_default():
|
||||
"""The suppression must be opt-in, so every existing caller keeps its log."""
|
||||
logging_obj = MagicMock()
|
||||
await _drive_async_completion(skip_pre_call_logging=False, logging_obj=logging_obj)
|
||||
await _drive_async_completion(logging_obj=logging_obj)
|
||||
assert logging_obj.pre_call.call_count == 1
|
||||
|
||||
|
||||
|
|
@ -373,7 +339,7 @@ async def test_async_completion_signs_off_the_event_loop(monkeypatch):
|
|||
release = asyncio.create_task(probe.release_refresh_from_the_loop())
|
||||
|
||||
response = await _drive_async_completion(
|
||||
skip_pre_call_logging=False, logging_obj=MagicMock(), credentials=probe.credentials()
|
||||
logging_obj=MagicMock(), credentials=probe.credentials()
|
||||
)
|
||||
await release
|
||||
|
||||
|
|
@ -414,9 +380,7 @@ def test_pre_call_logging_fires_once_when_the_sync_rust_path_declines():
|
|||
logging_obj = MagicMock()
|
||||
|
||||
with patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()):
|
||||
bridge.set_rust_chat_completions(
|
||||
decline=lambda **_kwargs: None, chat_completions=declining_native
|
||||
)
|
||||
bridge.set_rust_chat_completions(chat_completions=declining_native)
|
||||
response = _run(
|
||||
logging_obj=logging_obj,
|
||||
client=_sync_client_returning_converse_response(),
|
||||
|
|
@ -462,20 +426,15 @@ async def test_post_call_logging_fires_on_the_async_rust_path():
|
|||
cannot drift apart the way the pre_call suppression once did."""
|
||||
import json
|
||||
|
||||
async def native(**_kwargs):
|
||||
async def native(**kwargs):
|
||||
kwargs["on_request"]()
|
||||
return dict(RUST_RESPONSE)
|
||||
|
||||
bridge.set_rust_chat_completions(
|
||||
decline=lambda **_kwargs: None, achat_completions=native
|
||||
)
|
||||
bridge.set_rust_chat_completions(achat_completions=native)
|
||||
logging_obj = MagicMock()
|
||||
|
||||
with patch.object(
|
||||
BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS
|
||||
):
|
||||
await BedrockConverseLLM().completion(
|
||||
**_completion_kwargs(acompletion=True, logging_obj=logging_obj)
|
||||
)
|
||||
with patch.object(BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS):
|
||||
await BedrockConverseLLM().completion(**_completion_kwargs(acompletion=True, logging_obj=logging_obj))
|
||||
|
||||
assert logging_obj.post_call.call_count == 1
|
||||
logged = logging_obj.post_call.call_args.kwargs["original_response"]
|
||||
|
|
@ -500,9 +459,7 @@ def test_post_call_is_not_logged_twice_when_the_sync_rust_call_declines():
|
|||
logging_obj, calls = _recording_logging_obj()
|
||||
|
||||
with patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()):
|
||||
bridge.set_rust_chat_completions(
|
||||
decline=lambda **_kwargs: None, chat_completions=declining_native
|
||||
)
|
||||
bridge.set_rust_chat_completions(chat_completions=declining_native)
|
||||
response = _run(
|
||||
logging_obj=logging_obj,
|
||||
client=_sync_client_returning_converse_response(),
|
||||
|
|
|
|||
|
|
@ -13,29 +13,28 @@ from botocore.credentials import RefreshableCredentials
|
|||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.code_interpreter_interception.handler import (
|
||||
CodeInterpreterInterceptionLogger,
|
||||
LITELLM_CODE_EXECUTION_TOOL_NAME,
|
||||
CodeInterpreterInterceptionLogger,
|
||||
)
|
||||
from litellm.llms.azure.videos.transformation import AzureVideoConfig
|
||||
from litellm.llms.base_llm.audio_transcription.transformation import (
|
||||
AudioTranscriptionRequestData,
|
||||
BaseAudioTranscriptionConfig,
|
||||
)
|
||||
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
|
||||
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
|
||||
from litellm.llms.base_llm.search.transformation import BaseSearchConfig, SearchResponse
|
||||
from litellm.llms.bedrock.base_aws_llm import SignsRequestsWithAWS
|
||||
from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import (
|
||||
AmazonAnthropicClaudeMessagesConfig,
|
||||
)
|
||||
from litellm.llms.brave.search.transformation import BraveSearchConfig
|
||||
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import (
|
||||
BaseLLMHTTPHandler,
|
||||
_collect_ws_project_quota_callbacks,
|
||||
_google_genai_streaming_hidden_params,
|
||||
_has_pre_call_deployment_hook,
|
||||
_rust_responses_websocket_enabled,
|
||||
)
|
||||
from litellm.llms.azure.videos.transformation import AzureVideoConfig
|
||||
from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import (
|
||||
AmazonAnthropicClaudeMessagesConfig,
|
||||
)
|
||||
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
|
||||
from litellm.llms.openai.videos.transformation import OpenAIVideoConfig
|
||||
|
|
@ -1405,9 +1404,7 @@ def test_sync_delete_responses_sets_json_content_type():
|
|||
({}, True, None, None),
|
||||
],
|
||||
)
|
||||
def test_resolve_anthropic_messages_timeout(
|
||||
monkeypatch, litellm_params_kwargs, stream, global_timeout, expected
|
||||
):
|
||||
def test_resolve_anthropic_messages_timeout(monkeypatch, litellm_params_kwargs, stream, global_timeout, expected):
|
||||
from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS
|
||||
|
||||
if global_timeout is None:
|
||||
|
|
@ -1423,9 +1420,7 @@ def test_resolve_anthropic_messages_timeout(
|
|||
)
|
||||
else:
|
||||
monkeypatch.setattr("litellm.request_timeout", global_timeout, raising=False)
|
||||
monkeypatch.setattr(
|
||||
"litellm.request_timeout_explicitly_set", True, raising=False
|
||||
)
|
||||
monkeypatch.setattr("litellm.request_timeout_explicitly_set", True, raising=False)
|
||||
|
||||
resolved = BaseLLMHTTPHandler._resolve_anthropic_messages_timeout(
|
||||
litellm_params=GenericLiteLLMParams(**litellm_params_kwargs),
|
||||
|
|
@ -1450,9 +1445,7 @@ async def test_async_anthropic_messages_handler_forwards_request_timeout(monkeyp
|
|||
return_value=({"x-api-key": "k"}, "https://api.anthropic.com")
|
||||
)
|
||||
mock_config.should_filter_anthropic_beta_headers = Mock(return_value=False)
|
||||
mock_config.transform_anthropic_messages_request = Mock(
|
||||
return_value={"model": "claude", "messages": []}
|
||||
)
|
||||
mock_config.transform_anthropic_messages_request = Mock(return_value={"model": "claude", "messages": []})
|
||||
mock_config.get_complete_url = Mock(return_value="https://api.anthropic.com/v1/messages")
|
||||
mock_config.sign_request = Mock(return_value=({"x-api-key": "k"}, None))
|
||||
mock_config.max_retry_on_anthropic_messages_http_error = 1
|
||||
|
|
@ -1498,9 +1491,7 @@ async def test_async_anthropic_messages_handler_forwards_stream_timeout(monkeypa
|
|||
return_value=({"x-api-key": "k"}, "https://api.anthropic.com")
|
||||
)
|
||||
mock_config.should_filter_anthropic_beta_headers = Mock(return_value=False)
|
||||
mock_config.transform_anthropic_messages_request = Mock(
|
||||
return_value={"model": "claude", "messages": []}
|
||||
)
|
||||
mock_config.transform_anthropic_messages_request = Mock(return_value={"model": "claude", "messages": []})
|
||||
mock_config.get_complete_url = Mock(return_value="https://api.anthropic.com/v1/messages")
|
||||
mock_config.sign_request = Mock(return_value=({"x-api-key": "k"}, None))
|
||||
mock_config.max_retry_on_anthropic_messages_http_error = 1
|
||||
|
|
@ -1691,6 +1682,7 @@ def _make_responses_handler_call(signed_body):
|
|||
signing provider (e.g. Bedrock Mantle).
|
||||
"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
|
@ -1746,6 +1738,7 @@ def test_responses_handler_signs_after_fake_stream_prep_strips_stream():
|
|||
We snapshot request_data at sign time and assert "stream" is already gone.
|
||||
"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
|
|
@ -1809,6 +1802,7 @@ def _make_compact_handler_call(signed_body, is_async):
|
|||
signing provider (e.g. Bedrock Mantle SigV4 / bearer).
|
||||
"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
|
@ -1910,7 +1904,13 @@ async def test_async_anthropic_messages_handler_passes_api_key_to_agentic_hooks(
|
|||
)
|
||||
mock_config.sign_request = Mock(return_value=({}, None))
|
||||
|
||||
fake_raw_response = {"id": "msg_1", "type": "message", "role": "assistant", "content": [], "stop_reason": "end_turn"}
|
||||
fake_raw_response = {
|
||||
"id": "msg_1",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [],
|
||||
"stop_reason": "end_turn",
|
||||
}
|
||||
mock_config.transform_anthropic_messages_response = Mock(return_value=fake_raw_response)
|
||||
|
||||
mock_logging_obj = Mock()
|
||||
|
|
@ -1930,10 +1930,17 @@ async def test_async_anthropic_messages_handler_passes_api_key_to_agentic_hooks(
|
|||
mock_httpx_response.status_code = 200
|
||||
|
||||
with (
|
||||
patch.object(handler, "_async_post_anthropic_messages_with_http_error_retry", new=AsyncMock(return_value=mock_httpx_response)),
|
||||
patch.object(
|
||||
handler,
|
||||
"_async_post_anthropic_messages_with_http_error_retry",
|
||||
new=AsyncMock(return_value=mock_httpx_response),
|
||||
),
|
||||
patch.object(handler, "_call_agentic_completion_hooks", side_effect=fake_agentic_hooks),
|
||||
patch("litellm.llms.custom_httpx.llm_http_handler.get_async_httpx_client"),
|
||||
patch("litellm.litellm_core_utils.get_provider_specific_headers.ProviderSpecificHeaderUtils.get_provider_specific_headers", return_value=None),
|
||||
patch(
|
||||
"litellm.litellm_core_utils.get_provider_specific_headers.ProviderSpecificHeaderUtils.get_provider_specific_headers",
|
||||
return_value=None,
|
||||
),
|
||||
):
|
||||
result = await handler.async_anthropic_messages_handler(
|
||||
model="claude-haiku",
|
||||
|
|
@ -2175,7 +2182,9 @@ def test_audio_transcriptions_sends_dict_data_as_json_body():
|
|||
form-encodes it and silently ignores json=; JSON-body providers (e.g.
|
||||
Google Speech-to-Text) need an application/json body."""
|
||||
captured = {}
|
||||
client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_capture_json_transcription_request(captured))))
|
||||
client = HTTPHandler(
|
||||
client=httpx.Client(transport=httpx.MockTransport(_capture_json_transcription_request(captured)))
|
||||
)
|
||||
|
||||
response = BaseLLMHTTPHandler().audio_transcriptions(
|
||||
client=client,
|
||||
|
|
@ -2259,9 +2268,7 @@ def _transform_subtitle_response(payload):
|
|||
|
||||
|
||||
def test_subtitle_synthesis_fallback_without_timings_drops_words():
|
||||
response = _transform_subtitle_response(
|
||||
{"text": "hello world", "words": [{"word": "hello"}, {"word": "world"}]}
|
||||
)
|
||||
response = _transform_subtitle_response({"text": "hello world", "words": [{"word": "hello"}, {"word": "world"}]})
|
||||
|
||||
assert response.text == "hello world"
|
||||
assert "words" not in response
|
||||
|
|
@ -2481,9 +2488,7 @@ async def test_anthropic_invalid_thinking_signature_retry_resigns_bedrock_reques
|
|||
ok_response = httpx.Response(200, json={"id": "msg_1"}, request=httpx.Request("POST", request_url))
|
||||
|
||||
class FakeAsyncClient:
|
||||
async def post(
|
||||
self, url, headers, data, stream=False, logging_obj=None, timeout=None
|
||||
):
|
||||
async def post(self, url, headers, data, stream=False, logging_obj=None, timeout=None):
|
||||
posts.append({"headers": dict(headers), "data": data})
|
||||
return invalid_signature_response if len(posts) == 1 else ok_response
|
||||
|
||||
|
|
@ -2912,21 +2917,6 @@ async def test_generic_http_handler_async_streaming_forwards_provider_response_h
|
|||
assert "".join([chunk.choices[0].delta.content or "" for chunk in collected]) == "hi"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider, enabled, expected",
|
||||
[("openai", True, True), ("openai", False, False), ("azure", True, False),
|
||||
("hosted_vllm", True, False), (None, True, False)],
|
||||
)
|
||||
def test_the_rust_responses_websocket_needs_openai_and_process_enablement(
|
||||
custom_llm_provider, enabled, expected, monkeypatch
|
||||
):
|
||||
from litellm.rust_bridge import configuration
|
||||
|
||||
configuration.reset_rust_configuration()
|
||||
monkeypatch.setenv("LITELLM_RUST", "1" if enabled else "0")
|
||||
assert _rust_responses_websocket_enabled(custom_llm_provider) is expected
|
||||
|
||||
|
||||
def test_a_plain_callback_does_not_advertise_a_pre_call_deployment_hook(monkeypatch):
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
|
|
@ -3044,7 +3034,13 @@ def _capture_video_create_request(captured):
|
|||
captured["body"] = request.content
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={"id": "video_123", "object": "video", "status": "queued", "created_at": 1712697600, "model": "sora-2"},
|
||||
json={
|
||||
"id": "video_123",
|
||||
"object": "video",
|
||||
"status": "queued",
|
||||
"created_at": 1712697600,
|
||||
"model": "sora-2",
|
||||
},
|
||||
)
|
||||
|
||||
return respond
|
||||
|
|
@ -3067,7 +3063,9 @@ def test_video_generation_without_file_sends_multipart_form_data():
|
|||
captured = {}
|
||||
client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_capture_video_create_request(captured))))
|
||||
|
||||
result = BaseLLMHTTPHandler().video_generation_handler(client=client, **_video_create_call_kwargs(OpenAIVideoConfig()))
|
||||
result = BaseLLMHTTPHandler().video_generation_handler(
|
||||
client=client, **_video_create_call_kwargs(OpenAIVideoConfig())
|
||||
)
|
||||
|
||||
assert captured["content_type"].startswith("multipart/form-data")
|
||||
assert _multipart_text_fields(captured["content_type"], captured["body"]) == {
|
||||
|
|
@ -3107,7 +3105,9 @@ def test_azure_video_generation_without_file_sends_multipart_form_data():
|
|||
captured = {}
|
||||
client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_capture_video_create_request(captured))))
|
||||
|
||||
result = BaseLLMHTTPHandler().video_generation_handler(client=client, **_video_create_call_kwargs(AzureVideoConfig()))
|
||||
result = BaseLLMHTTPHandler().video_generation_handler(
|
||||
client=client, **_video_create_call_kwargs(AzureVideoConfig())
|
||||
)
|
||||
|
||||
assert captured["content_type"].startswith("multipart/form-data")
|
||||
assert _multipart_text_fields(captured["content_type"], captured["body"]) == {
|
||||
|
|
@ -3122,7 +3122,9 @@ def test_video_generation_json_provider_keeps_json_body():
|
|||
captured = {}
|
||||
client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_capture_video_create_request(captured))))
|
||||
|
||||
result = BaseLLMHTTPHandler().video_generation_handler(client=client, **_video_create_call_kwargs(_JSONBodyVideoConfig()))
|
||||
result = BaseLLMHTTPHandler().video_generation_handler(
|
||||
client=client, **_video_create_call_kwargs(_JSONBodyVideoConfig())
|
||||
)
|
||||
|
||||
assert captured["content_type"] == "application/json"
|
||||
assert json.loads(captured["body"]) == {"model": "sora-2", "prompt": "a cat surfing", "seconds": "4"}
|
||||
|
|
@ -3152,6 +3154,7 @@ def test_video_generation_with_input_reference_keeps_file_multipart():
|
|||
AZURE_AI_BASE = "https://myfoundry.services.ai.azure.com"
|
||||
AZURE_AI_CHAT_COMPLETIONS_URL = f"{AZURE_AI_BASE}/models/chat/completions"
|
||||
|
||||
|
||||
def _a_tool_with_an_unsupported_field() -> dict:
|
||||
return {
|
||||
"type": "function",
|
||||
|
|
@ -3159,14 +3162,13 @@ def _a_tool_with_an_unsupported_field() -> dict:
|
|||
"strict": True,
|
||||
}
|
||||
|
||||
|
||||
A_COMPLETION = {
|
||||
"id": "chatcmpl-1",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "grok-3",
|
||||
"choices": [
|
||||
{"index": 0, "message": {"role": "assistant", "content": "sent"}, "finish_reason": "stop"}
|
||||
],
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "sent"}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
|
||||
}
|
||||
|
||||
|
|
@ -3210,9 +3212,7 @@ def _call_azure_ai(recorder: _RecordedAzureAI, **overrides):
|
|||
|
||||
|
||||
def test_a_tool_field_the_provider_rejects_is_dropped_and_the_call_retried():
|
||||
recorder = _RecordedAzureAI(
|
||||
[_rejection(TOOL_LEVEL_REJECTION), httpx.Response(200, json=A_COMPLETION)]
|
||||
)
|
||||
recorder = _RecordedAzureAI([_rejection(TOOL_LEVEL_REJECTION), httpx.Response(200, json=A_COMPLETION)])
|
||||
|
||||
response = _call_azure_ai(recorder)
|
||||
|
||||
|
|
@ -3223,9 +3223,7 @@ def test_a_tool_field_the_provider_rejects_is_dropped_and_the_call_retried():
|
|||
|
||||
|
||||
def test_the_retry_changes_only_the_field_the_provider_named():
|
||||
recorder = _RecordedAzureAI(
|
||||
[_rejection(TOOL_LEVEL_REJECTION), httpx.Response(200, json=A_COMPLETION)]
|
||||
)
|
||||
recorder = _RecordedAzureAI([_rejection(TOOL_LEVEL_REJECTION), httpx.Response(200, json=A_COMPLETION)])
|
||||
|
||||
_call_azure_ai(recorder)
|
||||
|
||||
|
|
@ -3264,9 +3262,7 @@ def test_an_extra_input_outside_a_tool_is_not_retried_unless_dropping_params_was
|
|||
|
||||
|
||||
def test_an_extra_input_outside_a_tool_is_retried_when_dropping_params_was_asked_for():
|
||||
recorder = _RecordedAzureAI(
|
||||
[_rejection(UNRELATED_REJECTION), httpx.Response(200, json=A_COMPLETION)]
|
||||
)
|
||||
recorder = _RecordedAzureAI([_rejection(UNRELATED_REJECTION), httpx.Response(200, json=A_COMPLETION)])
|
||||
|
||||
response = _call_azure_ai(recorder, drop_params=True)
|
||||
|
||||
|
|
@ -3280,9 +3276,7 @@ async def test_a_tool_field_the_provider_rejects_is_dropped_and_retried_on_the_a
|
|||
):
|
||||
import respx
|
||||
|
||||
recorder = _RecordedAzureAI(
|
||||
[_rejection(TOOL_LEVEL_REJECTION), httpx.Response(200, json=A_COMPLETION)]
|
||||
)
|
||||
recorder = _RecordedAzureAI([_rejection(TOOL_LEVEL_REJECTION), httpx.Response(200, json=A_COMPLETION)])
|
||||
|
||||
with respx.mock(assert_all_called=True) as router:
|
||||
router.post(AZURE_AI_CHAT_COMPLETIONS_URL).mock(side_effect=recorder)
|
||||
|
|
@ -3487,7 +3481,9 @@ def _start_async_completion(config, logging_obj=None):
|
|||
custom_llm_provider="openai",
|
||||
model_response=ModelResponse(),
|
||||
encoding=None,
|
||||
logging_obj=logging_obj if logging_obj is not None else Mock(dynamic_success_callbacks=None, model_call_details={}),
|
||||
logging_obj=logging_obj
|
||||
if logging_obj is not None
|
||||
else Mock(dynamic_success_callbacks=None, model_call_details={}),
|
||||
optional_params={},
|
||||
timeout=10.0,
|
||||
litellm_params={},
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import importlib
|
||||
from collections.abc import AsyncGenerator
|
||||
from datetime import datetime
|
||||
from io import BytesIO
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
from unittest.mock import Mock
|
||||
|
||||
|
|
@ -61,8 +61,11 @@ async def test_python_request_response_and_callbacks(
|
|||
if dispatch != "disabled":
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
NATIVE_OCR_LIFECYCLE.override(Mock(side_effect=Declined()) if dispatch == "declined" else None)
|
||||
main: Final = importlib.import_module("litellm.ocr.main")
|
||||
monkeypatch.setattr(main, "native_exception_types", lambda: (Declined, RuntimeError))
|
||||
native: Final = SimpleNamespace(
|
||||
RustBridgeDeclined=Declined,
|
||||
RustUpstreamError=RuntimeError,
|
||||
)
|
||||
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
|
||||
logger: Final = Mock(spec=CustomLogger)
|
||||
monkeypatch.setattr(litellm, "input_callback", [logger])
|
||||
arguments: Final = {
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@ from __future__ import annotations
|
|||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.custom_httpx.llm_http_handler import _rust_responses_websocket_enabled
|
||||
from litellm.rust_bridge import configuration
|
||||
from litellm.rust_bridge.responses import websocket as responses_websocket
|
||||
|
||||
|
|
@ -35,6 +34,7 @@ class _FakeNativeBridge:
|
|||
url: str,
|
||||
headers: dict[str, str],
|
||||
timeout_seconds: float | None,
|
||||
custom_llm_provider: str | None,
|
||||
) -> _FakeNativeConnection:
|
||||
return _FakeNativeConnection()
|
||||
|
||||
|
|
@ -48,14 +48,6 @@ def reset_responses_websocket():
|
|||
configuration.reset_rust_configuration()
|
||||
|
||||
|
||||
def test_rust_websocket_bridge_uses_process_enablement() -> None:
|
||||
configuration.rust(False)
|
||||
assert not _rust_responses_websocket_enabled("openai")
|
||||
configuration.rust(True)
|
||||
assert _rust_responses_websocket_enabled("openai")
|
||||
assert not _rust_responses_websocket_enabled("anthropic")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_adapter_raises_clean_close_when_rust_connection_ends() -> None:
|
||||
adapter = responses_websocket._ConnectionAdapter(_ClosedNativeConnection())
|
||||
|
|
@ -72,6 +64,8 @@ async def test_bridge_unavailable_returns_none(monkeypatch: pytest.MonkeyPatch)
|
|||
assert (
|
||||
await responses_websocket.connect(
|
||||
url="wss://example.test/responses",
|
||||
custom_llm_provider="openai",
|
||||
model="test",
|
||||
headers={},
|
||||
timeout=None,
|
||||
)
|
||||
|
|
@ -88,6 +82,8 @@ async def test_enabled_bridge_connects_and_adapts_socket(
|
|||
|
||||
connection = await responses_websocket.connect(
|
||||
url="wss://example.test/responses",
|
||||
custom_llm_provider="openai",
|
||||
model="test",
|
||||
headers={"Authorization": "Bearer key"},
|
||||
timeout=1.0,
|
||||
)
|
||||
|
|
@ -96,3 +92,51 @@ async def test_enabled_bridge_connects_and_adapts_socket(
|
|||
await connection.send("response.create")
|
||||
assert await connection.recv() == "response.completed"
|
||||
await connection.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disabled_websocket_does_not_connect() -> None:
|
||||
class UnexpectedConnection:
|
||||
@classmethod
|
||||
async def connect(cls, **kwargs: object) -> None:
|
||||
raise AssertionError("disabled Rust must not connect")
|
||||
|
||||
responses_websocket.set_rust_responses_websocket(connection=UnexpectedConnection)
|
||||
configuration.rust(False)
|
||||
assert (
|
||||
await responses_websocket.connect(
|
||||
url="ws://127.0.0.1:1",
|
||||
headers={},
|
||||
timeout=0.1,
|
||||
custom_llm_provider="openai",
|
||||
model="test",
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_websocket_decline_falls_back_but_connection_failure_does_not() -> None:
|
||||
from litellm.exceptions import APIError
|
||||
|
||||
native = pytest.importorskip("litellm.rust_bridge._native")
|
||||
responses_websocket.set_rust_responses_websocket(connection=native.ResponsesWebSocketConnection)
|
||||
configuration.rust(True)
|
||||
assert (
|
||||
await responses_websocket.connect(
|
||||
url="ws://127.0.0.1:1",
|
||||
headers={},
|
||||
timeout=0.1,
|
||||
custom_llm_provider="azure",
|
||||
model="test",
|
||||
)
|
||||
is None
|
||||
)
|
||||
with pytest.raises(APIError):
|
||||
await responses_websocket.connect(
|
||||
url="ws://127.0.0.1:1",
|
||||
headers={},
|
||||
timeout=0.1,
|
||||
custom_llm_provider="openai",
|
||||
model="test",
|
||||
)
|
||||
|
|
|
|||
2
tests/test_litellm/rust_bridge/stubtest.ini
Normal file
2
tests/test_litellm/rust_bridge/stubtest.ini
Normal file
|
|
@ -0,0 +1,2 @@
|
|||
[mypy]
|
||||
follow_imports = skip
|
||||
|
|
@ -1,20 +1,17 @@
|
|||
"""Tests for the Rust chat completions bridge.
|
||||
|
||||
The native callables are dependency-injected through
|
||||
``set_rust_chat_completions`` rather than patched, so these run without the
|
||||
compiled extension present.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.rust_bridge import configuration
|
||||
from litellm.rust_bridge import chat_completions as bridge
|
||||
from litellm.rust_bridge import configuration
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
RUST_RESPONSE = {
|
||||
native = pytest.importorskip("litellm.rust_bridge._native")
|
||||
|
||||
RUST_RESPONSE: Final = {
|
||||
"created": 1_700_000_000,
|
||||
"model": "claude-sonnet-4-5-20260101",
|
||||
"choices": [
|
||||
|
|
@ -24,27 +21,12 @@ RUST_RESPONSE = {
|
|||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 11,
|
||||
"completion_tokens": 4,
|
||||
"total_tokens": 15,
|
||||
"prompt_tokens_details": {
|
||||
"cached_tokens": 0,
|
||||
"cache_creation_tokens": 0,
|
||||
"text_tokens": 11,
|
||||
},
|
||||
},
|
||||
"usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15},
|
||||
}
|
||||
MESSAGES: Final = [{"role": "user", "content": "hi"}]
|
||||
|
||||
MESSAGES = [{"role": "user", "content": "hi"}]
|
||||
|
||||
|
||||
class _FakeDeclined(Exception):
|
||||
"""Stands in for the native `RustBridgeDeclined`."""
|
||||
|
||||
|
||||
class _FakeUpstream(Exception):
|
||||
"""Stands in for the native `RustUpstreamError`; args are (status, message)."""
|
||||
_FakeDeclined = native.RustBridgeDeclined
|
||||
_FakeUpstream = native.RustUpstreamError
|
||||
|
||||
|
||||
class _FakeNative:
|
||||
|
|
@ -52,181 +34,42 @@ class _FakeNative:
|
|||
RustUpstreamError = _FakeUpstream
|
||||
|
||||
|
||||
def _fake_native_bridge(monkeypatch):
|
||||
"""Expose the bridge's exception classes without the compiled extension."""
|
||||
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative())
|
||||
|
||||
|
||||
def _hide_native_bridge(monkeypatch):
|
||||
"""Simulate a wheel built without the compiled extension.
|
||||
|
||||
There is no injection seam for "the .so is absent", so the loader itself is
|
||||
replaced; every other case here uses `set_rust_chat_completions`.
|
||||
"""
|
||||
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_bridge(monkeypatch):
|
||||
"""Every test starts with no injected callables, and leaves none behind."""
|
||||
bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None, decline=None)
|
||||
configuration.reset_rust_configuration()
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
yield
|
||||
bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None, decline=None)
|
||||
configuration.reset_rust_configuration()
|
||||
|
||||
|
||||
class _RecordingDecline:
|
||||
"""A stand-in for the native gate that records what it was asked."""
|
||||
|
||||
def __init__(self, reason: str | None = None):
|
||||
self.reason = reason
|
||||
self.calls: list[dict] = []
|
||||
|
||||
def __call__(self, **kwargs):
|
||||
self.calls.append(kwargs)
|
||||
return self.reason
|
||||
|
||||
|
||||
class _RecordingCall:
|
||||
def __init__(self, result=None, error: Exception | None = None):
|
||||
self.result = result if result is not None else dict(RUST_RESPONSE)
|
||||
self.error = error
|
||||
self.calls: list[dict] = []
|
||||
def __init__(self, result: object = RUST_RESPONSE, error: Exception | None = None) -> None:
|
||||
self.result: Final = result
|
||||
self.error: Final = error
|
||||
self.calls: Final[list[dict[str, object]]] = []
|
||||
|
||||
def __call__(self, **kwargs):
|
||||
def __call__(self, **kwargs: object) -> object:
|
||||
self.calls.append(kwargs)
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
on_request: Final = kwargs["on_request"]
|
||||
assert callable(on_request)
|
||||
on_request()
|
||||
return self.result
|
||||
|
||||
|
||||
class _RecordingAsyncCall(_RecordingCall):
|
||||
async def __call__(self, **kwargs):
|
||||
return _RecordingCall.__call__(self, **kwargs)
|
||||
async def __call__(self, **kwargs: object) -> object:
|
||||
return super().__call__(**kwargs)
|
||||
|
||||
|
||||
def _accepts(**overrides) -> bool:
|
||||
kwargs = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"messages": MESSAGES,
|
||||
"optional_params": {"max_tokens": 16},
|
||||
"custom_llm_provider": "anthropic",
|
||||
"litellm_params": {},
|
||||
"stream": None,
|
||||
}
|
||||
kwargs.update(overrides)
|
||||
return bridge.rust_chat_completions_accepts(**kwargs)
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_bridge(monkeypatch: pytest.MonkeyPatch):
|
||||
bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None)
|
||||
configuration.reset_rust_configuration()
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
yield
|
||||
bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None)
|
||||
configuration.reset_rust_configuration()
|
||||
|
||||
|
||||
class TestGate:
|
||||
def test_declines_when_the_deployment_did_not_opt_in(self, monkeypatch):
|
||||
monkeypatch.delenv("LITELLM_RUST", raising=False)
|
||||
gate = _RecordingDecline()
|
||||
bridge.set_rust_chat_completions(decline=gate)
|
||||
assert _accepts(litellm_params={}) is False
|
||||
assert _accepts(litellm_params=None) is False
|
||||
assert gate.calls == [], "the gate must not be consulted before opt-in"
|
||||
|
||||
def test_accepts_when_the_deployment_opted_in_and_the_core_agrees(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
gate = _RecordingDecline()
|
||||
bridge.set_rust_chat_completions(decline=gate)
|
||||
assert _accepts() is True
|
||||
assert gate.calls[0]["model"] == "claude-sonnet-4-5"
|
||||
assert gate.calls[0]["custom_llm_provider"] == "anthropic"
|
||||
|
||||
def test_process_enable_applies_without_request_override(self):
|
||||
bridge.set_rust_chat_completions(decline=_RecordingDecline())
|
||||
configuration.rust(True)
|
||||
|
||||
assert _accepts(litellm_params={}) is True
|
||||
|
||||
def test_the_env_var_opts_in_without_a_per_model_flag(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_RUST", "true")
|
||||
bridge.set_rust_chat_completions(decline=_RecordingDecline())
|
||||
assert _accepts(litellm_params={}) is True
|
||||
|
||||
def test_declines_streaming_and_providers_off_the_path(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
gate = _RecordingDecline()
|
||||
bridge.set_rust_chat_completions(decline=gate)
|
||||
assert _accepts(stream=True) is False
|
||||
assert _accepts(custom_llm_provider="openai") is False
|
||||
assert _accepts(custom_llm_provider=None) is False
|
||||
assert gate.calls == []
|
||||
|
||||
def test_declines_an_anthropic_request_carrying_a_litellm_metadata_user_id(self, monkeypatch):
|
||||
"""`AnthropicConfig.transform_request` copies a valid `user_id` into the Messages body.
|
||||
|
||||
It does that inside the function the Rust route replaces, and the core is
|
||||
handed `optional_params` only, so accepting here would send the request
|
||||
to Anthropic with the abuse-detection attribution silently missing.
|
||||
"""
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
gate = _RecordingDecline()
|
||||
bridge.set_rust_chat_completions(decline=gate)
|
||||
assert _accepts(litellm_params={"metadata": {"user_id": "u-123"}}) is False
|
||||
assert gate.calls == [], "the core must not be consulted for a request it cannot see the key of"
|
||||
|
||||
# Bedrock's Converse transform reads no `user_id`, and an Anthropic request
|
||||
# whose metadata carries none is one Python would not attribute either.
|
||||
assert (
|
||||
_accepts(
|
||||
custom_llm_provider="bedrock",
|
||||
model="bedrock/us-east-1/anthropic.claude-v2",
|
||||
litellm_params={"metadata": {"user_id": "u-123"}},
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert _accepts(litellm_params={"metadata": {"trace_id": "t-1"}}) is True
|
||||
assert _accepts(litellm_params={"metadata": {"user_id": None}}) is True
|
||||
assert _accepts(litellm_params={"metadata": None}) is True
|
||||
|
||||
def test_declines_a_bedrock_request_while_the_proxy_owns_request_metadata(self, monkeypatch):
|
||||
"""`AmazonConverseConfig` resolves proxy-owned `requestMetadata` onto the
|
||||
Converse body from `litellm_params`, and owning that field also means
|
||||
evicting a caller-supplied one. The core can do neither, so an operator
|
||||
who armed `bedrock_request_metadata_fields` keeps the Python path.
|
||||
"""
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
gate = _RecordingDecline()
|
||||
bridge.set_rust_chat_completions(decline=gate)
|
||||
bedrock = {
|
||||
"custom_llm_provider": "bedrock",
|
||||
"model": "bedrock/us-east-1/anthropic.claude-v2",
|
||||
}
|
||||
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ["user_api_key_team_id"])
|
||||
assert _accepts(**bedrock) is False
|
||||
assert gate.calls == [], "the core must not be consulted for a field it cannot write"
|
||||
assert _accepts() is True, "arming Bedrock attribution must not decline Anthropic"
|
||||
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", None)
|
||||
assert _accepts(**bedrock) is True, "the decline follows the operator's opt-in alone"
|
||||
|
||||
def test_declines_when_the_core_declines(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
bridge.set_rust_chat_completions(decline=_RecordingDecline("streaming"))
|
||||
assert _accepts() is False
|
||||
|
||||
def test_declines_when_the_bridge_is_unavailable(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
_hide_native_bridge(monkeypatch)
|
||||
assert _accepts() is False
|
||||
|
||||
def test_declines_when_the_gate_itself_raises(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
|
||||
def exploding(**_kwargs):
|
||||
raise RuntimeError("boom")
|
||||
|
||||
bridge.set_rust_chat_completions(decline=exploding)
|
||||
assert _accepts() is False
|
||||
def _fake_native_bridge(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative())
|
||||
|
||||
|
||||
def _call_kwargs(model_response: ModelResponse) -> dict:
|
||||
def _call_kwargs(model_response: ModelResponse, fallback: object = "python") -> dict[str, object]:
|
||||
return {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"messages": MESSAGES,
|
||||
|
|
@ -237,159 +80,92 @@ def _call_kwargs(model_response: ModelResponse) -> dict:
|
|||
"custom_llm_provider": "anthropic",
|
||||
"extra_headers": {},
|
||||
"timeout": 30.0,
|
||||
"on_response": lambda _rust_response: None,
|
||||
"python_fallback": lambda: fallback,
|
||||
}
|
||||
|
||||
|
||||
class TestSyncCall:
|
||||
def test_builds_a_model_response_and_stamps_the_rust_header(self):
|
||||
native = _RecordingCall()
|
||||
bridge.set_rust_chat_completions(chat_completions=native)
|
||||
model_response = ModelResponse()
|
||||
original_id = model_response.id
|
||||
def test_sync_native_entrypoint_runs_once_and_logs_once() -> None:
|
||||
events: Final[list[str]] = []
|
||||
native_call: Final = _RecordingCall()
|
||||
bridge.set_rust_chat_completions(chat_completions=native_call)
|
||||
kwargs: Final = _call_kwargs(ModelResponse())
|
||||
kwargs.update({"on_request": lambda: events.append("pre"), "on_response": lambda _value: events.append("post")})
|
||||
|
||||
result = bridge.chat_completions(**_call_kwargs(model_response))
|
||||
result: Final = bridge.chat_completions(**kwargs)
|
||||
|
||||
assert result is not None
|
||||
assert result.choices[0].message.content == "hello from rust"
|
||||
assert result.choices[0].finish_reason == "stop"
|
||||
assert result.model == "claude-sonnet-4-5-20260101"
|
||||
assert result.usage.prompt_tokens == 11
|
||||
assert result.usage.completion_tokens == 4
|
||||
assert result.usage.total_tokens == 15
|
||||
assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert result.id == original_id, "the rust path must keep the chatcmpl id litellm already minted"
|
||||
|
||||
def test_passes_the_timeout_through_as_seconds(self):
|
||||
native = _RecordingCall()
|
||||
bridge.set_rust_chat_completions(chat_completions=native)
|
||||
bridge.chat_completions(**_call_kwargs(ModelResponse()))
|
||||
assert native.calls[0]["timeout_seconds"] == 30.0
|
||||
|
||||
def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch):
|
||||
_hide_native_bridge(monkeypatch)
|
||||
assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None
|
||||
|
||||
def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch):
|
||||
_fake_native_bridge(monkeypatch)
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming")))
|
||||
assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None
|
||||
assert isinstance(result, ModelResponse)
|
||||
assert result.choices[0].message.content == "hello from rust"
|
||||
assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert len(native_call.calls) == 1
|
||||
assert events == ["pre", "post"]
|
||||
|
||||
|
||||
class TestAsyncCall:
|
||||
@pytest.mark.asyncio
|
||||
async def test_builds_a_model_response(self):
|
||||
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall())
|
||||
result = await bridge.achat_completions(**_call_kwargs(ModelResponse()))
|
||||
assert result is not None
|
||||
assert result.choices[0].message.content == "hello from rust"
|
||||
assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
def test_decline_has_no_logging_effect_and_runs_one_fallback(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
_fake_native_bridge(monkeypatch)
|
||||
events: Final[list[str]] = []
|
||||
native_call: Final = _RecordingCall(error=_FakeDeclined("unsupported"))
|
||||
bridge.set_rust_chat_completions(chat_completions=native_call)
|
||||
kwargs: Final = _call_kwargs(ModelResponse(), fallback="python")
|
||||
kwargs.update({"on_request": lambda: events.append("pre"), "on_response": lambda _value: events.append("post")})
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch):
|
||||
_hide_native_bridge(monkeypatch)
|
||||
assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch):
|
||||
_fake_native_bridge(monkeypatch)
|
||||
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming")))
|
||||
assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None
|
||||
assert bridge.chat_completions(**kwargs) == "python"
|
||||
assert len(native_call.calls) == 1
|
||||
assert events == []
|
||||
|
||||
|
||||
class TestAsyncFallbackWrapper:
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_the_rust_response_without_running_the_fallback(self):
|
||||
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall())
|
||||
ran = []
|
||||
|
||||
async def fallback():
|
||||
ran.append(True)
|
||||
return "python"
|
||||
|
||||
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
|
||||
assert result.choices[0].message.content == "hello from rust"
|
||||
assert ran == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runs_the_fallback_when_the_core_declines(self, monkeypatch):
|
||||
_fake_native_bridge(monkeypatch)
|
||||
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming")))
|
||||
|
||||
async def fallback():
|
||||
return "python"
|
||||
|
||||
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
|
||||
assert result == "python"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runs_the_fallback_when_the_bridge_is_unavailable(self, monkeypatch):
|
||||
_hide_native_bridge(monkeypatch)
|
||||
|
||||
async def fallback():
|
||||
return "python"
|
||||
|
||||
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
|
||||
assert result == "python"
|
||||
def test_streaming_uses_python_without_loading_native() -> None:
|
||||
native_call: Final = _RecordingCall()
|
||||
bridge.set_rust_chat_completions(chat_completions=native_call)
|
||||
kwargs: Final = _call_kwargs(ModelResponse())
|
||||
kwargs["stream"] = True
|
||||
assert bridge.chat_completions(**kwargs) == "python"
|
||||
assert native_call.calls == []
|
||||
|
||||
|
||||
class TestFailureClassification:
|
||||
"""A failure the provider already saw must not be retried on the Python
|
||||
path: it would bill the customer for the same work twice."""
|
||||
def test_host_facts_reach_the_single_native_call(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
native_call: Final = _RecordingCall()
|
||||
bridge.set_rust_chat_completions(chat_completions=native_call)
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ["team_id"])
|
||||
kwargs: Final = _call_kwargs(ModelResponse())
|
||||
kwargs["litellm_params"] = {"metadata": {"user_id": "u-1"}}
|
||||
bridge.chat_completions(**kwargs)
|
||||
assert native_call.calls[0]["host_facts"] == {
|
||||
"stream": False,
|
||||
"anthropic_user_id": True,
|
||||
"bedrock_metadata_owned": True,
|
||||
}
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _native_exceptions(self, monkeypatch):
|
||||
_fake_native_bridge(monkeypatch)
|
||||
|
||||
def test_a_decline_falls_back_because_nothing_was_sent(self):
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming")))
|
||||
assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None
|
||||
def test_upstream_and_adaptation_failures_never_fall_back(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
_fake_native_bridge(monkeypatch)
|
||||
fallback_calls: Final[list[bool]] = []
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(429, "rate limited")))
|
||||
kwargs: Final = _call_kwargs(ModelResponse())
|
||||
kwargs["python_fallback"] = lambda: fallback_calls.append(True)
|
||||
with pytest.raises(litellm.APIError, match="rate limited"):
|
||||
bridge.chat_completions(**kwargs)
|
||||
assert fallback_calls == []
|
||||
|
||||
def test_an_upstream_failure_is_surfaced_with_its_status(self):
|
||||
from litellm.exceptions import APIError
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall())
|
||||
kwargs["on_response"] = lambda _value: (_ for _ in ()).throw(RuntimeError("adapt failed"))
|
||||
with pytest.raises(RuntimeError, match="adapt failed"):
|
||||
bridge.chat_completions(**kwargs)
|
||||
assert fallback_calls == []
|
||||
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(429, "429: rate limited")))
|
||||
with pytest.raises(APIError) as raised:
|
||||
bridge.chat_completions(**_call_kwargs(ModelResponse()))
|
||||
assert raised.value.status_code == 429
|
||||
assert "rate limited" in str(raised.value)
|
||||
|
||||
def test_a_transport_failure_with_no_response_surfaces_as_a_500(self):
|
||||
from litellm.exceptions import APIError
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_native_and_fallback_paths(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
native_call: Final = _RecordingAsyncCall()
|
||||
bridge.set_rust_chat_completions(achat_completions=native_call)
|
||||
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(0, "connection reset")))
|
||||
with pytest.raises(APIError) as raised:
|
||||
bridge.chat_completions(**_call_kwargs(ModelResponse()))
|
||||
assert raised.value.status_code == 500
|
||||
async def fallback() -> str:
|
||||
return "python"
|
||||
|
||||
def test_an_unrecognized_error_is_not_swallowed(self):
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=RuntimeError("something else")))
|
||||
with pytest.raises(RuntimeError):
|
||||
bridge.chat_completions(**_call_kwargs(ModelResponse()))
|
||||
kwargs: Final = _call_kwargs(ModelResponse())
|
||||
kwargs["python_fallback"] = fallback
|
||||
result: Final = await bridge.achat_completions(**kwargs)
|
||||
assert isinstance(result, ModelResponse)
|
||||
assert len(native_call.calls) == 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_async_wrapper_does_not_fall_back_on_an_upstream_failure(self):
|
||||
from litellm.exceptions import APIError
|
||||
|
||||
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeUpstream(500, "500: boom")))
|
||||
ran = []
|
||||
|
||||
async def fallback():
|
||||
ran.append(True)
|
||||
return "python"
|
||||
|
||||
with pytest.raises(APIError):
|
||||
await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
|
||||
assert ran == [], "a request the provider already served must not be re-issued"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_async_wrapper_falls_back_on_a_decline(self):
|
||||
bridge.set_rust_chat_completions(
|
||||
achat_completions=_RecordingAsyncCall(error=_FakeDeclined("blank message text"))
|
||||
)
|
||||
|
||||
async def fallback():
|
||||
return "python"
|
||||
|
||||
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
|
||||
assert result == "python"
|
||||
configuration.rust(False)
|
||||
assert await bridge.achat_completions(**kwargs) == "python"
|
||||
|
|
|
|||
|
|
@ -1,138 +0,0 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
from collections.abc import Generator
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.rust_bridge import configuration
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _isolated_configuration( # pyright: ignore[reportUnusedFunction] # pytest discovers fixtures dynamically
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> Generator[None]:
|
||||
configuration.reset_rust_configuration()
|
||||
monkeypatch.delenv("LITELLM_RUST", raising=False)
|
||||
yield
|
||||
configuration.reset_rust_configuration()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("process", "environment", "release_default", "expected"),
|
||||
(
|
||||
(False, True, True, False),
|
||||
(True, False, False, True),
|
||||
(None, False, True, False),
|
||||
(None, True, False, True),
|
||||
(None, None, False, False),
|
||||
(None, None, True, True),
|
||||
),
|
||||
)
|
||||
def test_resolution_precedence(
|
||||
process: bool | None,
|
||||
environment: bool | None,
|
||||
release_default: bool,
|
||||
expected: bool,
|
||||
) -> None:
|
||||
assert (
|
||||
configuration.resolve_rust_enabled(
|
||||
process_override=process,
|
||||
environment_override=environment,
|
||||
release_default=release_default,
|
||||
)
|
||||
is expected
|
||||
)
|
||||
|
||||
|
||||
def test_release_default_remains_disabled() -> None:
|
||||
assert configuration.DEFAULT_RUST_ENABLED is False
|
||||
assert configuration.rust_enabled() is False
|
||||
assert configuration.rust_ocr_enabled() is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize("route", tuple(configuration.RouteName))
|
||||
def test_each_route_has_an_explicit_release_default(route: configuration.RouteName) -> None:
|
||||
assert configuration.rust_enabled(route) is (
|
||||
route in {configuration.RouteName.OCR, configuration.RouteName.TRANSCRIPTION}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("enabled", (True, False))
|
||||
def test_required_transcription_ignores_optional_rollout(monkeypatch: pytest.MonkeyPatch, enabled: bool) -> None:
|
||||
monkeypatch.setenv("LITELLM_RUST", "0")
|
||||
configuration.rust(enabled)
|
||||
assert configuration.rust_enabled(configuration.RouteName.TRANSCRIPTION) is True
|
||||
assert configuration.rust_enabled(configuration.RouteName.MESSAGES) is enabled
|
||||
assert configuration.rust_enabled(configuration.RouteName.OCR) is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize("process", [None, False, True])
|
||||
@pytest.mark.parametrize("environment", [None, "0", "1", "off"])
|
||||
def test_ocr_configuration(monkeypatch: pytest.MonkeyPatch, process: bool | None, environment: str | None) -> None:
|
||||
if environment is not None:
|
||||
monkeypatch.setenv("LITELLM_RUST", environment)
|
||||
if process is not None:
|
||||
configuration.rust(process)
|
||||
|
||||
assert configuration.rust_ocr_enabled() is (environment not in {"0", "off"} and process is not False)
|
||||
|
||||
|
||||
def test_process_override_wins_over_environment(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("LITELLM_RUST", "0")
|
||||
configuration.rust(True)
|
||||
|
||||
assert configuration.rust_enabled() is True
|
||||
|
||||
|
||||
def test_global_environment_accepts_explicit_false(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("LITELLM_RUST", "off")
|
||||
|
||||
assert configuration.rust_enabled() is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", ("", " ", "sometimes", "2"))
|
||||
def test_invalid_environment_value_disables_rust(monkeypatch: pytest.MonkeyPatch, value: str) -> None:
|
||||
monkeypatch.setenv("LITELLM_RUST", value)
|
||||
|
||||
assert configuration.rust_enabled() is False
|
||||
|
||||
|
||||
def test_process_override_and_reset_apply_to_existing_threads(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
|
||||
with ThreadPoolExecutor(max_workers=1) as executor:
|
||||
assert executor.submit(configuration.rust_enabled).result() is True
|
||||
configuration.rust(False)
|
||||
assert executor.submit(configuration.rust_enabled).result() is False
|
||||
configuration.reset_rust_configuration()
|
||||
assert executor.submit(configuration.rust_enabled).result() is True
|
||||
|
||||
|
||||
def test_explicit_override_precedes_invalid_environment(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("LITELLM_RUST", "sometimes")
|
||||
|
||||
configuration.rust(True)
|
||||
assert configuration.rust_enabled() is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("value", "expected"), (("1", "True"), ("0", "False")))
|
||||
def test_environment_controls_startup(value: str, expected: str) -> None:
|
||||
environment: Final = {**os.environ, "LITELLM_RUST": value}
|
||||
result: Final = subprocess.run(
|
||||
(
|
||||
sys.executable,
|
||||
"-c",
|
||||
"from litellm.rust_bridge.configuration import rust_enabled; print(rust_enabled())",
|
||||
),
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
env=environment,
|
||||
)
|
||||
|
||||
assert result.stdout.strip() == expected
|
||||
|
|
@ -1,27 +1,22 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.rust_bridge.configuration import (
|
||||
_parse_env_bool, # pyright: ignore[reportPrivateUsage] # directly test env parsing contract
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", ("1", "true", "t", "yes", "y", "on", "TRUE", " yes "))
|
||||
def test_parse_env_bool_accepts_standard_true_values(value: str) -> None:
|
||||
assert _parse_env_bool(value) is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", ("0", "false", "f", "no", "n", "off", "FALSE", " no "))
|
||||
def test_parse_env_bool_accepts_standard_false_values(value: str) -> None:
|
||||
assert _parse_env_bool(value) is False
|
||||
@pytest.mark.parametrize(("value", "expected"), (("1", True), ("0", False), (" 1 ", True), (" 0 ", False)))
|
||||
def test_parse_env_bool_accepts_binary_values(value: str, expected: bool) -> None:
|
||||
assert _parse_env_bool(value) is expected
|
||||
|
||||
|
||||
def test_parse_env_bool_preserves_unset_value() -> None:
|
||||
assert _parse_env_bool(None) is None
|
||||
|
||||
|
||||
def test_parse_env_bool_rejects_unknown_value() -> None:
|
||||
with pytest.raises(ValidationError):
|
||||
_parse_env_bool("enabled")
|
||||
@pytest.mark.parametrize("value", ("enabled", "true", "false", "yes", "no", "on", "off", ""))
|
||||
def test_parse_env_bool_rejects_unknown_value(value: str) -> None:
|
||||
with pytest.raises(ValueError, match="must be '1' or '0'"):
|
||||
_parse_env_bool(value)
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
from collections.abc import Generator, Mapping
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
|
||||
|
|
@ -51,7 +52,7 @@ def test_admitted_failure_is_returned_without_replay() -> None:
|
|||
assert caught.value is failure
|
||||
finally:
|
||||
NATIVE_OCR_LIFECYCLE.reset()
|
||||
litellm.rust(None)
|
||||
configuration.reset_rust_configuration()
|
||||
assert native.call_count == 1
|
||||
|
||||
|
||||
|
|
@ -64,6 +65,7 @@ def test_public_binding_keeps_positional_fields_and_defaults_out_of_native_hook_
|
|||
args: tuple[object, ...],
|
||||
kwargs: Mapping[str, object],
|
||||
asynchronous: bool,
|
||||
host: object,
|
||||
) -> OCRResponse:
|
||||
captured.append((request, args, kwargs, asynchronous))
|
||||
return OCRResponse(pages=[], model=request.model)
|
||||
|
|
@ -74,7 +76,7 @@ def test_public_binding_keeps_positional_fields_and_defaults_out_of_native_hook_
|
|||
response: Final = litellm.ocr("mistral/mistral-ocr-latest", document)
|
||||
finally:
|
||||
NATIVE_OCR_LIFECYCLE.reset()
|
||||
litellm.rust(None)
|
||||
configuration.reset_rust_configuration()
|
||||
|
||||
request, call_args, hook_kwargs, asynchronous = captured[0]
|
||||
assert response.model == "mistral/mistral-ocr-latest"
|
||||
|
|
@ -94,6 +96,7 @@ def test_public_binding_keeps_keyword_model_and_document_in_native_hook_kwargs()
|
|||
args: tuple[object, ...],
|
||||
kwargs: Mapping[str, object],
|
||||
asynchronous: bool,
|
||||
host: object,
|
||||
) -> OCRResponse:
|
||||
assert args == ()
|
||||
captured.append(kwargs)
|
||||
|
|
@ -105,7 +108,7 @@ def test_public_binding_keeps_keyword_model_and_document_in_native_hook_kwargs()
|
|||
litellm.ocr(model="mistral/mistral-ocr-latest", document=document)
|
||||
finally:
|
||||
NATIVE_OCR_LIFECYCLE.reset()
|
||||
litellm.rust(None)
|
||||
configuration.reset_rust_configuration()
|
||||
|
||||
assert captured[0]["model"] == "mistral/mistral-ocr-latest"
|
||||
assert captured[0]["document"] is document
|
||||
|
|
@ -123,7 +126,7 @@ def test_public_duplicate_argument_error_does_not_depend_on_native_selection(ena
|
|||
litellm.ocr("mistral/mistral-ocr-latest", document, model="duplicate")
|
||||
finally:
|
||||
NATIVE_OCR_LIFECYCLE.reset()
|
||||
litellm.rust(None)
|
||||
configuration.reset_rust_configuration()
|
||||
assert native.call_count == 0
|
||||
|
||||
|
||||
|
|
@ -137,14 +140,14 @@ def test_public_missing_required_argument_error_does_not_depend_on_native_select
|
|||
litellm.ocr("mistral/mistral-ocr-latest")
|
||||
finally:
|
||||
NATIVE_OCR_LIFECYCLE.reset()
|
||||
litellm.rust(None)
|
||||
configuration.reset_rust_configuration()
|
||||
assert native.call_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", [False, True])
|
||||
@pytest.mark.parametrize("enabled", [False, True, None])
|
||||
async def test_environment_opt_out_never_loads_native(
|
||||
@pytest.mark.parametrize("enabled", [False, None])
|
||||
async def test_environment_opt_out_skips_native_without_process_enable(
|
||||
monkeypatch: pytest.MonkeyPatch, asynchronous: bool, enabled: bool | None
|
||||
) -> None:
|
||||
monkeypatch.setenv("LITELLM_RUST", "0")
|
||||
|
|
@ -153,7 +156,8 @@ async def test_environment_opt_out_never_loads_native(
|
|||
monkeypatch.setattr(legacy, "aocr" if asynchronous else "ocr", fallback)
|
||||
load: Final = Mock(side_effect=AssertionError("native must not be loaded"))
|
||||
monkeypatch.setattr(bindings, "get_native_bridge", load)
|
||||
litellm.rust(enabled)
|
||||
if enabled is not None:
|
||||
litellm.rust(enabled)
|
||||
document: Final = {"type": "file", "file": b"pdf"}
|
||||
|
||||
result: Final = (
|
||||
|
|
@ -196,6 +200,10 @@ class Declined(Exception):
|
|||
pass
|
||||
|
||||
|
||||
class Upstream(Exception):
|
||||
pass
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", [False, True])
|
||||
@pytest.mark.parametrize("declined", [False, True])
|
||||
|
|
@ -205,10 +213,11 @@ async def test_only_native_declines_replay_on_legacy(
|
|||
failure: Final = Declined("unsupported") if declined else RuntimeError("provider already called")
|
||||
native: Final = AsyncMock(side_effect=failure) if asynchronous else Mock(side_effect=failure)
|
||||
NATIVE_OCR_LIFECYCLE.override(native)
|
||||
import importlib
|
||||
|
||||
main: Final = importlib.import_module("litellm.ocr.main")
|
||||
monkeypatch.setattr(main, "native_exception_types", lambda: (Declined, RuntimeError))
|
||||
monkeypatch.setattr(
|
||||
bindings,
|
||||
"get_native_bridge",
|
||||
lambda: SimpleNamespace(RustBridgeDeclined=Declined, RustUpstreamError=Upstream),
|
||||
)
|
||||
response: Final = OCRResponse(pages=[], model="mistral-ocr-latest")
|
||||
fallback: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
|
||||
monkeypatch.setattr(legacy, "aocr" if asynchronous else "ocr", fallback)
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterator
|
||||
from types import ModuleType
|
||||
from typing import Final
|
||||
|
||||
|
|
@ -7,8 +8,25 @@ import pytest
|
|||
|
||||
from litellm.rust_bridge import configuration
|
||||
from litellm.rust_bridge.bindings import BINDING_UNSET
|
||||
from litellm.rust_bridge.configuration import RouteName
|
||||
from litellm.rust_bridge.route import NativeRoute
|
||||
from litellm.rust_bridge.catalog import COMPONENTS
|
||||
from litellm.rust_bridge.configuration import (
|
||||
CapabilityContext,
|
||||
CapabilityDefinition,
|
||||
DeliveryMode,
|
||||
ExecutionDecision,
|
||||
RolloutPolicy,
|
||||
RouteName,
|
||||
RustImplementationState,
|
||||
)
|
||||
from litellm.rust_bridge.errors import RustRouteUnavailableError, RustRouteUnsupportedError
|
||||
from litellm.rust_bridge.route import NativeComponent
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_configuration() -> Iterator[None]:
|
||||
configuration.reset_rust_configuration()
|
||||
yield
|
||||
configuration.reset_rust_configuration()
|
||||
|
||||
|
||||
def _string(value: object) -> str | None:
|
||||
|
|
@ -16,22 +34,30 @@ def _string(value: object) -> str | None:
|
|||
|
||||
|
||||
def _unexpected_load() -> ModuleType:
|
||||
raise AssertionError("disabled route loaded the extension")
|
||||
raise AssertionError("Python selection loaded the native extension")
|
||||
|
||||
|
||||
def test_disabled_route_does_not_load_native(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_python_decision_does_not_load_native(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.delenv("LITELLM_RUST", raising=False)
|
||||
configuration.reset_rust_configuration()
|
||||
route: Final = NativeRoute(RouteName.MESSAGES)
|
||||
binding: Final = route.bind("messages", validate=_string, module_loader=_unexpected_load)
|
||||
assert route.select(binding) is None
|
||||
component: Final = COMPONENTS[RouteName.MESSAGES]
|
||||
binding: Final = component.bind("messages", validate=_string, module_loader=_unexpected_load)
|
||||
assert component.resolve().select(binding) is None
|
||||
|
||||
|
||||
def test_delivery_mode_uses_one_component_policy(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
component: Final = COMPONENTS[RouteName.MESSAGES]
|
||||
completed: Final = component.resolve(CapabilityContext(delivery=DeliveryMode.COMPLETED))
|
||||
streaming: Final = component.resolve(CapabilityContext(delivery=DeliveryMode.STREAMING))
|
||||
assert completed.decision is ExecutionDecision.RUST_WITH_FALLBACK
|
||||
assert streaming.decision is ExecutionDecision.PYTHON
|
||||
|
||||
|
||||
def test_binding_discovery_validation_and_override() -> None:
|
||||
module: Final = ModuleType("fake_native")
|
||||
setattr(module, "messages", "native")
|
||||
route: Final = NativeRoute(RouteName.MESSAGES)
|
||||
binding: Final = route.bind("messages", validate=_string, module_loader=lambda: module)
|
||||
component: Final = COMPONENTS[RouteName.MESSAGES]
|
||||
binding: Final = component.bind("messages", validate=_string, module_loader=lambda: module)
|
||||
assert binding.load() == "native"
|
||||
binding.configure("override")
|
||||
binding.configure(BINDING_UNSET)
|
||||
|
|
@ -44,17 +70,94 @@ def test_binding_discovery_validation_and_override() -> None:
|
|||
assert binding.load() is None
|
||||
|
||||
|
||||
def test_missing_native_is_unavailable() -> None:
|
||||
route: Final = NativeRoute(RouteName.OCR)
|
||||
binding: Final = route.bind("ocr", validate=_string, module_loader=lambda: None)
|
||||
assert binding.load() is None
|
||||
def test_component_rejects_undeclared_binding() -> None:
|
||||
with pytest.raises(ValueError, match="not declared"):
|
||||
COMPONENTS[RouteName.MESSAGES].bind("typo", validate=_string)
|
||||
|
||||
|
||||
def test_explicit_enablement_loads_optional_route(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
configuration.reset_rust_configuration()
|
||||
module: Final = ModuleType("fake_native")
|
||||
setattr(module, "messages", "native")
|
||||
route: Final = NativeRoute(RouteName.MESSAGES)
|
||||
binding: Final = route.bind("messages", validate=_string, module_loader=lambda: module)
|
||||
assert route.select(binding) == "native"
|
||||
class ModelCapability:
|
||||
def __init__(self) -> None:
|
||||
self.calls: int = 0
|
||||
|
||||
def __call__(self, context: CapabilityContext) -> CapabilityDefinition:
|
||||
self.calls += 1
|
||||
if context.provider == "provider" and context.model == "new-model":
|
||||
return CapabilityDefinition(
|
||||
rust=RustImplementationState.EXPERIMENTAL,
|
||||
python_available=False,
|
||||
rollout=RolloutPolicy.RUST_REQUIRED,
|
||||
)
|
||||
return CapabilityDefinition(
|
||||
rust=RustImplementationState.READY,
|
||||
python_available=True,
|
||||
rollout=RolloutPolicy.RUST_OPT_IN,
|
||||
)
|
||||
|
||||
|
||||
def test_dynamic_capability_resolves_once_before_binding_selection() -> None:
|
||||
resolver: Final = ModelCapability()
|
||||
component: Final = NativeComponent(
|
||||
name=RouteName.TRANSCRIPTION,
|
||||
capability=resolver,
|
||||
exports=("transcription",),
|
||||
)
|
||||
binding: Final = component.bind("transcription", validate=_string, module_loader=_unexpected_load)
|
||||
binding.override("native")
|
||||
execution: Final = component.resolve(CapabilityContext(provider="provider", model="new-model"))
|
||||
assert execution.select(binding) == "native"
|
||||
assert resolver.calls == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize("process_override", (None, False, True))
|
||||
@pytest.mark.parametrize("environment_override", ("0", "1", "invalid"))
|
||||
def test_bedrock_transcription_requires_rust_regardless_of_overrides(
|
||||
monkeypatch: pytest.MonkeyPatch, process_override: bool | None, environment_override: str
|
||||
) -> None:
|
||||
monkeypatch.setenv("LITELLM_RUST", environment_override)
|
||||
if process_override is not None:
|
||||
configuration.rust(process_override)
|
||||
component: Final = COMPONENTS[RouteName.TRANSCRIPTION]
|
||||
binding: Final = component.bind("transcription", validate=_string, module_loader=lambda: None)
|
||||
execution: Final = component.resolve(CapabilityContext(provider="bedrock", model="model"))
|
||||
assert execution.decision is ExecutionDecision.RUST_REQUIRED
|
||||
with pytest.raises(RustRouteUnavailableError, match="transcription bridge is unavailable"):
|
||||
execution.select(binding)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider", ("openai", "azure", "azure_ai", "groq", "mistral", "nvidia_riva", "soniox"))
|
||||
def test_python_transcription_providers_skip_native_discovery(provider: str) -> None:
|
||||
configuration.rust(True)
|
||||
component: Final = COMPONENTS[RouteName.TRANSCRIPTION]
|
||||
binding: Final = component.bind("transcription", validate=_string, module_loader=_unexpected_load)
|
||||
execution: Final = component.resolve(CapabilityContext(provider=provider, model="model"))
|
||||
assert execution.decision is ExecutionDecision.PYTHON
|
||||
assert execution.select(binding) is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider", ("unknown-provider", "anthropic"))
|
||||
def test_unsupported_transcription_never_selects_an_implementation(provider: str) -> None:
|
||||
component: Final = COMPONENTS[RouteName.TRANSCRIPTION]
|
||||
binding: Final = component.bind("transcription", validate=_string, module_loader=_unexpected_load)
|
||||
execution: Final = component.resolve(CapabilityContext(provider=provider, model="model"))
|
||||
assert execution.decision is ExecutionDecision.UNSUPPORTED
|
||||
with pytest.raises(RustRouteUnsupportedError, match="No Python or Rust implementation for transcription"):
|
||||
execution.select(binding)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("rust", "python_available", "rollout"),
|
||||
(
|
||||
(RustImplementationState.READY, True, RolloutPolicy.RUST_REQUIRED),
|
||||
(RustImplementationState.EXPERIMENTAL, False, RolloutPolicy.RUST_OPT_IN),
|
||||
(RustImplementationState.READY, False, RolloutPolicy.RUST_OPT_OUT),
|
||||
(RustImplementationState.UNIMPLEMENTED, False, RolloutPolicy.PYTHON_ONLY),
|
||||
(RustImplementationState.UNIMPLEMENTED, False, RolloutPolicy.RUST_REQUIRED),
|
||||
(RustImplementationState.UNIMPLEMENTED, True, RolloutPolicy.UNSUPPORTED),
|
||||
(RustImplementationState.READY, False, RolloutPolicy.UNSUPPORTED),
|
||||
),
|
||||
)
|
||||
def test_invalid_capability_definitions_rejected(
|
||||
rust: RustImplementationState, python_available: bool, rollout: RolloutPolicy
|
||||
) -> None:
|
||||
with pytest.raises(ValueError, match="capability"):
|
||||
CapabilityDefinition(rust=rust, python_available=python_available, rollout=rollout)
|
||||
|
|
|
|||
|
|
@ -1,95 +1,191 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Callable
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.exceptions import APIError
|
||||
from litellm.rust_bridge import bindings, runtime
|
||||
|
||||
from litellm.rust_bridge.configuration import ExecutionDecision, RouteName
|
||||
from litellm.rust_bridge.errors import RustRouteDeclinedError, RustRouteUnavailableError, RustRouteUnsupportedError
|
||||
from litellm.rust_bridge.route import ComponentExecution
|
||||
|
||||
class RustBridgeDeclined(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class RustBridgeUnavailable(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class RustUpstreamError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class RustHostCallbackError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def native_exceptions(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
native = SimpleNamespace(
|
||||
native: Final = SimpleNamespace(
|
||||
RustBridgeDeclined=RustBridgeDeclined,
|
||||
RustBridgeUnavailable=RustBridgeUnavailable,
|
||||
RustHostCallbackError=RustHostCallbackError,
|
||||
RustUpstreamError=RustUpstreamError,
|
||||
)
|
||||
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
|
||||
|
||||
|
||||
def context() -> runtime.BridgeErrorContext:
|
||||
return runtime.BridgeErrorContext(route="messages", provider="anthropic", model="model")
|
||||
async def _invoke(
|
||||
asynchronous: bool,
|
||||
decision: ExecutionDecision,
|
||||
native_call: Callable[[], object] | None,
|
||||
adapt: Callable[[object], object],
|
||||
fallback: Callable[[], object] = lambda: "python",
|
||||
) -> object:
|
||||
execution: Final = ComponentExecution(route_name=RouteName.MESSAGES, decision=decision)
|
||||
context: Final = runtime.BridgeErrorContext(route="messages", provider="anthropic", model="model")
|
||||
if not asynchronous:
|
||||
return runtime.invoke(
|
||||
execution=execution,
|
||||
native_call=native_call,
|
||||
python_fallback=fallback,
|
||||
adapt=adapt,
|
||||
context=context,
|
||||
)
|
||||
|
||||
async def call() -> object:
|
||||
assert native_call is not None
|
||||
return native_call()
|
||||
|
||||
def test_invoke_tags_native_decline_before_running_fallback() -> None:
|
||||
calls: list[str] = []
|
||||
async def afallback() -> object:
|
||||
return fallback()
|
||||
|
||||
def decline() -> object:
|
||||
calls.append("rust")
|
||||
raise RustBridgeDeclined("unsupported")
|
||||
|
||||
value = runtime.invoke(
|
||||
native_call=decline,
|
||||
fallback=lambda: calls.append("python") or "fallback",
|
||||
adapt=str,
|
||||
mode=runtime.FallbackMode.PYTHON,
|
||||
context=context(),
|
||||
return await runtime.ainvoke(
|
||||
execution=execution,
|
||||
native_call=call if native_call is not None else None,
|
||||
python_fallback=afallback,
|
||||
adapt=adapt,
|
||||
context=context,
|
||||
)
|
||||
|
||||
assert value == "fallback"
|
||||
assert calls == ["rust", "python"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
async def test_python_selection_runs_only_the_fallback(asynchronous: bool) -> None:
|
||||
calls: Final[list[str]] = []
|
||||
|
||||
def native() -> str:
|
||||
calls.append("native")
|
||||
return "native"
|
||||
|
||||
result: Final = await _invoke(
|
||||
asynchronous,
|
||||
ExecutionDecision.PYTHON,
|
||||
native,
|
||||
str,
|
||||
lambda: calls.append("python") or "python",
|
||||
)
|
||||
assert result == "python"
|
||||
assert calls == ["python"]
|
||||
|
||||
|
||||
def test_invoke_translates_upstream_without_fallback() -> None:
|
||||
def fail() -> object:
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
@pytest.mark.parametrize("error", (None, RustBridgeDeclined("unsupported"), RustBridgeUnavailable()))
|
||||
@pytest.mark.parametrize("required", (False, True))
|
||||
async def test_unavailable_and_declined_follow_policy(
|
||||
asynchronous: bool, error: Exception | None, required: bool
|
||||
) -> None:
|
||||
def fail() -> str:
|
||||
assert error is not None
|
||||
raise error
|
||||
|
||||
decision: Final = ExecutionDecision.RUST_REQUIRED if required else ExecutionDecision.RUST_WITH_FALLBACK
|
||||
native_call: Final = fail if error is not None else None
|
||||
if required:
|
||||
expected: Final = RustRouteDeclinedError if isinstance(error, RustBridgeDeclined) else RustRouteUnavailableError
|
||||
with pytest.raises(expected):
|
||||
await _invoke(asynchronous, decision, native_call, str)
|
||||
return
|
||||
assert await _invoke(asynchronous, decision, native_call, str) == "python"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
@pytest.mark.parametrize("decision", (ExecutionDecision.RUST_WITH_FALLBACK, ExecutionDecision.RUST_REQUIRED))
|
||||
@pytest.mark.parametrize("value", (None, False, 0, "native"))
|
||||
async def test_native_success_does_not_run_fallback(
|
||||
asynchronous: bool, decision: ExecutionDecision, value: None | bool | int | str
|
||||
) -> None:
|
||||
calls: Final[list[str]] = []
|
||||
result: Final = await _invoke(
|
||||
asynchronous,
|
||||
decision,
|
||||
lambda: value,
|
||||
lambda native: native,
|
||||
lambda: calls.append("python") or "python",
|
||||
)
|
||||
assert result is value
|
||||
assert calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
async def test_upstream_failure_never_falls_back(asynchronous: bool) -> None:
|
||||
def fail() -> str:
|
||||
raise RustUpstreamError(429, "rate limited")
|
||||
|
||||
with pytest.raises(APIError, match="rate limited") as caught:
|
||||
runtime.invoke(
|
||||
native_call=fail,
|
||||
fallback=lambda: pytest.fail("fallback must not run"),
|
||||
adapt=str,
|
||||
mode=runtime.FallbackMode.PYTHON,
|
||||
context=context(),
|
||||
)
|
||||
|
||||
await _invoke(asynchronous, ExecutionDecision.RUST_WITH_FALLBACK, fail, str)
|
||||
assert caught.value.status_code == 429
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ainvoke_handles_native_success() -> None:
|
||||
async def native() -> int:
|
||||
return 3
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
@pytest.mark.parametrize("error", (RuntimeError("execution failed"), asyncio.CancelledError()))
|
||||
async def test_execution_failure_never_falls_back(asynchronous: bool, error: BaseException) -> None:
|
||||
def fail() -> str:
|
||||
raise error
|
||||
|
||||
async def fallback() -> str:
|
||||
pytest.fail("fallback must not run")
|
||||
|
||||
assert (
|
||||
await runtime.ainvoke(
|
||||
native_call=native,
|
||||
fallback=fallback,
|
||||
adapt=str,
|
||||
mode=runtime.FallbackMode.PYTHON,
|
||||
context=context(),
|
||||
)
|
||||
== "3"
|
||||
)
|
||||
with pytest.raises(type(error)) as caught:
|
||||
await _invoke(asynchronous, ExecutionDecision.RUST_WITH_FALLBACK, fail, str)
|
||||
assert caught.value is error
|
||||
|
||||
|
||||
def test_required_mode_rejects_unavailable_bridge() -> None:
|
||||
with pytest.raises(RuntimeError, match="is unavailable"):
|
||||
runtime.invoke(
|
||||
native_call=None,
|
||||
fallback=lambda: pytest.fail("fallback must not run"),
|
||||
adapt=str,
|
||||
mode=runtime.FallbackMode.RUST_REQUIRED,
|
||||
context=context(),
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
@pytest.mark.parametrize("error", (RustBridgeDeclined("adapt"), RustBridgeUnavailable(), RuntimeError("adapt")))
|
||||
async def test_adaptation_failure_never_falls_back(asynchronous: bool, error: Exception) -> None:
|
||||
def adapt(_value: str) -> str:
|
||||
raise error
|
||||
|
||||
with pytest.raises(type(error)) as caught:
|
||||
await _invoke(asynchronous, ExecutionDecision.RUST_WITH_FALLBACK, lambda: "native", adapt)
|
||||
assert caught.value is error
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
async def test_host_callback_failure_preserves_its_cause(asynchronous: bool) -> None:
|
||||
callback_error: Final = RustBridgeDeclined("callback raised native decline type")
|
||||
|
||||
def fail() -> str:
|
||||
wrapped: Final = RustHostCallbackError("callback failed")
|
||||
wrapped.__cause__ = callback_error
|
||||
raise wrapped
|
||||
|
||||
with pytest.raises(RustBridgeDeclined) as caught:
|
||||
await _invoke(asynchronous, ExecutionDecision.RUST_WITH_FALLBACK, fail, str)
|
||||
assert caught.value is callback_error
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
async def test_unsupported_execution_runs_nothing(asynchronous: bool) -> None:
|
||||
with pytest.raises(RustRouteUnsupportedError):
|
||||
await _invoke(asynchronous, ExecutionDecision.UNSUPPORTED, lambda: "native", str)
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue