This commit is contained in:
Yujong Lee 2026-09-14 17:31:07 -07:00
parent 592fd00504
commit 6aeae9354a
39 changed files with 919 additions and 902 deletions

View file

@ -17,7 +17,6 @@ pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Resu
.await
}
pub fn admit(
model: &str,
provider: Option<&str>,

View file

@ -55,7 +55,6 @@ pub fn chat_completions_decline_reason(
.map(|reason| reason.0)
}
#[derive(Clone, Copy, Debug, Default, serde::Deserialize)]
#[serde(deny_unknown_fields)]
pub struct AdmissionContext {

View file

@ -27,7 +27,6 @@ pub async fn messages_stream(request: MessagesRequest<'_>) -> Result<reqwest::Re
execute_messages_provider_stream(request).await
}
pub fn admit(
model: &str,
provider: Option<&str>,

View file

@ -112,7 +112,11 @@ mod tests {
Python::initialize();
Python::attach(|py| {
for (error, status, message) in [
(Error::Connect("connection refused".into()), 0, "connection refused"),
(
Error::Connect("connection refused".into()),
0,
"connection refused",
),
(
Error::Http {
status: 429,

View file

@ -269,7 +269,7 @@ mod tests {
(
"messages",
"amessages",
"(model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None, has_agentic_hook=None)",
"(model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None, has_agentic_hook=None, on_request=None)",
),
(
"chat_completions",

View file

@ -26,6 +26,14 @@ fn prepare_messages(
options.custom_llm_provider.as_deref(),
inputs.has_agentic_hook.unwrap_or(false),
))?;
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 {
@ -66,6 +74,7 @@ bridge_route! {
extra_headers: Option<serde_json::Value>,
timeout_seconds: Option<f64>,
has_agentic_hook: Option<bool>,
on_request: Option<Py<PyAny>>,
},
prepare = prepare_messages,
errors = execution_error_to_pyerr,

View file

@ -34,8 +34,9 @@ fn count_input_tokens<'py>(
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)?;
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();
@ -72,7 +73,8 @@ fn cached_counter(
.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 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())),
@ -109,7 +111,9 @@ fn admission_error_to_pyerr(error: Error) -> PyErr {
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::Load(_) | Error::Ranks(_) | Error::UnicodeClasses => {
RustBridgeUnavailable::new_err(message)
}
Error::UnsupportedTokenizer
| Error::RequestParse(_)
| Error::MissingInput

View file

@ -449,6 +449,83 @@ class AnthropicChatCompletion(BaseLLM):
timeout=timeout,
)
def completion_dispatch() -> "ModelResponse | CustomStreamWrapper":
request_headers, data = finish_request(
config.transform_request(
model=model,
messages=messages,
optional_params=transform_params,
litellm_params=litellm_params,
headers=headers,
)
)
if stream is True:
data["stream"] = stream
completion_stream, response_headers = make_sync_call(
client=client,
api_base=api_base,
headers=request_headers,
data=json.dumps(data),
model=model,
messages=messages,
logging_obj=logging_obj,
timeout=timeout,
json_mode=json_mode,
speed=optional_params.get("speed") if optional_params else None,
tool_name_reverse_map=(
litellm_params.get(ANTHROPIC_TOOL_NAME_REVERSE_MAP_KEY)
if isinstance(litellm_params, dict)
else None
),
)
return CustomStreamWrapper(
completion_stream=completion_stream,
model=model,
custom_llm_provider="anthropic",
logging_obj=logging_obj,
_response_headers=process_anthropic_headers(response_headers),
)
sync_client: Final = (
client if isinstance(client, HTTPHandler) else _get_httpx_client(params={"timeout": timeout})
)
try:
response: Final = sync_client.post(
api_base,
headers=request_headers,
data=json.dumps(data),
timeout=timeout,
logging_obj=logging_obj,
)
except Exception as e:
status_code: Final = getattr(e, "status_code", 500)
error_headers = getattr(e, "headers", None)
error_text = getattr(e, "text", str(e))
error_response: Final[object] = getattr(e, "response", None)
if error_headers is None and error_response:
error_headers = getattr(error_response, "headers", None)
if error_response and hasattr(error_response, "text"):
error_text = getattr(error_response, "text", error_text)
raise AnthropicError(
message=error_text,
status_code=status_code,
headers=error_headers,
)
return config.transform_response(
model=model,
raw_response=response,
model_response=model_response,
logging_obj=logging_obj,
api_key=api_key,
request_data=data,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
encoding=encoding,
json_mode=json_mode,
)
rust_optional_params: Final = { # mutable-ok: json.dumps in the bridge rejects a mappingproxy
**AnthropicConfig.get_config(model=model),
**optional_params,
@ -462,9 +539,10 @@ class AnthropicChatCompletion(BaseLLM):
"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
)
def log_rust_pre_call() -> None:
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,
@ -488,7 +566,7 @@ class AnthropicChatCompletion(BaseLLM):
on_response=log_rust_post_call,
python_fallback=acompletion_dispatch,
)
rust_response: Final = rust_chat_completions_bridge.chat_completions(
return rust_chat_completions_bridge.chat_completions(
model=model,
messages=messages,
optional_params=rust_optional_params,
@ -502,94 +580,7 @@ class AnthropicChatCompletion(BaseLLM):
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()
else:
headers, data = finish_request(
config.transform_request(
model=model,
messages=messages,
optional_params=transform_params,
litellm_params=litellm_params,
headers=headers,
)
)
## COMPLETION CALL
if (
stream is True
): # if function call - fake the streaming (need complete blocks for output parsing in openai format)
data["stream"] = stream
completion_stream, headers = make_sync_call(
client=client,
api_base=api_base,
headers=headers,
data=json.dumps(data),
model=model,
messages=messages,
logging_obj=logging_obj,
timeout=timeout,
json_mode=json_mode,
speed=optional_params.get("speed") if optional_params else None,
tool_name_reverse_map=(
litellm_params.get(ANTHROPIC_TOOL_NAME_REVERSE_MAP_KEY)
if isinstance(litellm_params, dict)
else None
),
)
return CustomStreamWrapper(
completion_stream=completion_stream,
model=model,
custom_llm_provider="anthropic",
logging_obj=logging_obj,
_response_headers=process_anthropic_headers(headers),
)
else:
if client is None or not isinstance(client, HTTPHandler):
client = _get_httpx_client(params={"timeout": timeout})
else:
client = client
try:
response: Final = client.post(
api_base,
headers=headers,
data=json.dumps(data),
timeout=timeout,
logging_obj=logging_obj,
)
except Exception as e:
status_code: Final = getattr(e, "status_code", 500)
error_headers = getattr(e, "headers", None)
error_text = getattr(e, "text", str(e))
error_response: Final[object] = getattr(e, "response", None)
if error_headers is None and error_response:
error_headers = getattr(error_response, "headers", None)
if error_response and hasattr(error_response, "text"):
error_text = getattr(error_response, "text", error_text)
raise AnthropicError(
message=error_text,
status_code=status_code,
headers=error_headers,
)
return config.transform_response(
model=model,
raw_response=response,
model_response=model_response,
logging_obj=logging_obj,
api_key=api_key,
request_data=data,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
encoding=encoding,
json_mode=json_mode,
python_fallback=completion_dispatch,
)
def embedding(self):

View file

@ -50,9 +50,8 @@ class BedrockAudioTranscriptionRustDispatch:
extra_headers=extra_headers,
optional_params=optional_params,
timeout=timeout,
python_fallback=None,
)
if rust_response is None:
raise RuntimeError("Rust audio transcription bridge is unavailable")
return TranscriptionResponse(**rust_response)
async def async_audio_transcriptions(
@ -76,7 +75,6 @@ class BedrockAudioTranscriptionRustDispatch:
extra_headers=extra_headers,
optional_params=optional_params,
timeout=timeout,
python_fallback=None,
)
if rust_response is None:
raise RuntimeError("Rust audio transcription bridge is unavailable")
return TranscriptionResponse(**rust_response)

View file

@ -411,49 +411,139 @@ class BedrockConverseLLM(BaseAWSLLM):
"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
)
def log_rust_pre_call() -> None:
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(
def completion_dispatch() -> ModelResponse | CustomStreamWrapper:
request_data: Final = litellm.AmazonConverseConfig()._transform_request(
model=model,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
headers=extra_headers,
)
data: Final = json.dumps(request_data)
prepped: Final = self.get_request_headers(
credentials=credentials,
aws_region_name=aws_region_name,
extra_headers=extra_headers,
endpoint_url=proxy_endpoint_url,
data=data,
headers=headers,
api_key=api_key,
)
logging_obj.pre_call(
input=messages,
api_key="",
additional_args={
"complete_input_dict": data,
"api_base": proxy_endpoint_url,
"headers": prepped.headers,
},
)
client_timeout: Final = httpx.Timeout(timeout) if isinstance(timeout, (float, int)) else timeout
sync_client: Final = (
client
if isinstance(client, HTTPHandler)
else _get_httpx_client({} if client_timeout is None else {"timeout": client_timeout})
)
if stream is True:
completion_stream, response_headers = make_sync_call(
client=sync_client,
api_base=proxy_endpoint_url,
headers=prepped.headers,
data=data,
model=model,
messages=messages,
optional_params=rust_optional_params,
model_response=model_response,
api_key=api_key,
api_base=proxy_endpoint_url,
logging_obj=logging_obj,
json_mode=json_mode,
fake_stream=fake_stream,
stream_chunk_size=stream_chunk_size,
)
return CustomStreamWrapper(
completion_stream=completion_stream,
model=model,
custom_llm_provider="bedrock",
extra_headers=headers,
timeout=timeout,
logging_obj=logging_obj,
_response_headers=response_headers,
)
try:
response: Final = sync_client.post(
url=proxy_endpoint_url,
headers=prepped.headers,
data=data,
logging_obj=logging_obj,
)
response.raise_for_status()
except httpx.HTTPStatusError as err:
error_code: Final = err.response.status_code
raise BedrockError(
status_code=error_code,
message=error_response_text(err.response),
headers=err.response.headers,
response=err.response,
)
except httpx.TimeoutException:
raise BedrockError(status_code=408, message="Timeout error occurred.")
transformed_response: Final = litellm.AmazonConverseConfig()._transform_response(
model=model,
response=response,
model_response=model_response,
stream=stream if isinstance(stream, bool) else False,
logging_obj=logging_obj,
api_key="",
data=data,
messages=messages,
optional_params=optional_params,
encoding=encoding,
)
transformed_response.set_provider_response_headers(response.headers)
return transformed_response
if acompletion:
return rust_chat_completions_bridge.achat_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: self.async_completion(
model=model,
messages=messages,
api_base=proxy_endpoint_url,
model_response=model_response,
encoding=encoding,
logging_obj=logging_obj,
optional_params=optional_params,
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,
messages=messages,
api_base=proxy_endpoint_url,
model_response=model_response,
encoding=encoding,
logging_obj=logging_obj,
optional_params=optional_params,
stream=stream,
litellm_params=litellm_params,
logger_fn=logger_fn,
headers=headers,
timeout=timeout,
client=client,
credentials=credentials,
api_key=api_key,
),
)
rust_response: Final = rust_chat_completions_bridge.chat_completions(
logger_fn=logger_fn,
headers=headers,
timeout=timeout,
client=client,
credentials=credentials,
api_key=api_key,
),
)
return rust_chat_completions_bridge.chat_completions(
model=model,
messages=messages,
optional_params=rust_optional_params,
@ -467,151 +557,5 @@ class BedrockConverseLLM(BaseAWSLLM):
litellm_params=litellm_params,
on_request=log_rust_pre_call,
on_response=log_rust_post_call,
python_fallback=lambda: None,
python_fallback=completion_dispatch,
)
if rust_response is not None:
return rust_response
### ROUTING (ASYNC, STREAMING, SYNC)
if acompletion:
if isinstance(client, HTTPHandler):
client = None
if stream is True:
return self.async_streaming(
model=model,
messages=messages,
api_base=proxy_endpoint_url,
model_response=model_response,
encoding=encoding,
logging_obj=logging_obj,
optional_params=optional_params,
stream=True,
litellm_params=litellm_params,
logger_fn=logger_fn,
headers=headers,
timeout=timeout,
client=client,
json_mode=json_mode,
fake_stream=fake_stream,
credentials=credentials,
api_key=api_key,
stream_chunk_size=stream_chunk_size,
)
### ASYNC COMPLETION
return self.async_completion(
model=model,
messages=messages,
api_base=proxy_endpoint_url,
model_response=model_response,
encoding=encoding,
logging_obj=logging_obj,
optional_params=optional_params,
stream=stream,
litellm_params=litellm_params,
logger_fn=logger_fn,
headers=headers,
timeout=timeout,
client=client,
credentials=credentials,
api_key=api_key,
)
## TRANSFORMATION ##
_data: Final = litellm.AmazonConverseConfig()._transform_request(
model=model,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
headers=extra_headers,
)
data: Final = json.dumps(_data)
prepped: Final = self.get_request_headers(
credentials=credentials,
aws_region_name=aws_region_name,
extra_headers=extra_headers,
endpoint_url=proxy_endpoint_url,
data=data,
headers=headers,
api_key=api_key,
)
## LOGGING
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:
if isinstance(timeout, float) or isinstance(timeout, int):
timeout = httpx.Timeout(timeout)
_params["timeout"] = timeout
client = _get_httpx_client(_params)
else:
client = client
if stream is not None and stream is True:
completion_stream, response_headers = make_sync_call(
client=(client if client is not None and isinstance(client, HTTPHandler) else None),
api_base=proxy_endpoint_url,
headers=prepped.headers,
data=data,
model=model,
messages=messages,
logging_obj=logging_obj,
json_mode=json_mode,
fake_stream=fake_stream,
stream_chunk_size=stream_chunk_size,
)
streaming_response: Final = CustomStreamWrapper(
completion_stream=completion_stream,
model=model,
custom_llm_provider="bedrock",
logging_obj=logging_obj,
_response_headers=response_headers,
)
return streaming_response
### COMPLETION
try:
response: Final = client.post(
url=proxy_endpoint_url,
headers=prepped.headers,
data=data,
logging_obj=logging_obj,
)
response.raise_for_status()
except httpx.HTTPStatusError as err:
error_code: Final = err.response.status_code
raise BedrockError(
status_code=error_code,
message=error_response_text(err.response),
headers=err.response.headers,
response=err.response,
)
except httpx.TimeoutException:
raise BedrockError(status_code=408, message="Timeout error occurred.")
sync_transformed_response: Final = litellm.AmazonConverseConfig()._transform_response(
model=model,
response=response,
model_response=model_response,
stream=stream if isinstance(stream, bool) else False,
logging_obj=logging_obj,
api_key="",
data=data,
messages=messages,
optional_params=optional_params,
encoding=encoding,
)
sync_transformed_response.set_provider_response_headers(response.headers)
return sync_transformed_response

View file

@ -2264,36 +2264,90 @@ class BaseLLMHTTPHandler:
# internally -- this only deduplicates the success path.
request_body_json: Final = json.dumps(request_body)
logging_obj.pre_call(
input=[{"role": "user", "content": request_body_json}],
api_key="",
additional_args={
"complete_input_dict": request_body,
"api_base": str(request_url),
"headers": headers,
},
)
def log_pre_call() -> None:
logging_obj.pre_call(
input=[{"role": "user", "content": request_body_json}],
api_key="",
additional_args={
"complete_input_dict": request_body,
"api_base": str(request_url),
"headers": headers,
},
)
rust_messages_response: Final = await self._maybe_rust_anthropic_messages(
custom_llm_provider=custom_llm_provider,
litellm_params=litellm_params,
has_agentic_hook=self._has_agentic_completion_hook(logging_obj),
model=model,
api_key=api_key,
api_base=api_base,
headers=headers,
request_body=request_body,
timeout=self._resolve_anthropic_messages_timeout(
litellm_params=litellm_params,
async def python_fallback() -> AnthropicMessagesResponse | AsyncIterator:
log_pre_call()
response: Final = await self._async_post_anthropic_messages_with_http_error_retry(
async_httpx_client=async_httpx_client,
request_url=request_url,
headers=headers,
signed_json_body=(signed_json_body if signed_json_body is not None else request_body_json),
request_body=request_body,
stream=stream or False,
custom_llm_provider=custom_llm_provider,
),
)
if rust_messages_response is not None:
logging_obj=logging_obj,
provider_config=anthropic_messages_provider_config,
litellm_params=litellm_params,
api_key=api_key,
model=model,
timeout=self._resolve_anthropic_messages_timeout(
litellm_params=litellm_params,
stream=stream or False,
custom_llm_provider=custom_llm_provider,
),
)
logging_obj.model_call_details["httpx_response"] = response
if stream:
return self._rust_anthropic_messages_fake_stream(rust_messages_response)
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
AnthropicMessagesStreamingResponse,
anthropic_messages_stream_hidden_params,
)
completion_stream: Final = anthropic_messages_provider_config.get_async_streaming_response_iterator(
model=model,
httpx_response=response,
request_body=request_body,
litellm_logging_obj=logging_obj,
)
stream_hidden_params: Final = anthropic_messages_stream_hidden_params(response.headers)
if not self._has_agentic_completion_hook(logging_obj):
return AnthropicMessagesStreamingResponse(
completion_stream=completion_stream,
hidden_params=stream_hidden_params,
)
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
AgenticAnthropicStreamingIterator,
)
held_back_tool_names: Final = self._server_fulfilled_tools_in_request(
logging_obj=logging_obj,
tools=anthropic_messages_optional_request_params.get("tools"),
)
agentic_stream: Final = AgenticAnthropicStreamingIterator(
completion_stream=completion_stream,
http_handler=self,
model=model,
messages=messages,
anthropic_messages_provider_config=anthropic_messages_provider_config,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
kwargs={**kwargs, "api_key": api_key} if api_key else kwargs,
hold_back=bool(held_back_tool_names),
server_fulfilled_tool_names=held_back_tool_names,
)
return AnthropicMessagesStreamingResponse(
completion_stream=agentic_stream,
hidden_params=stream_hidden_params,
)
initial_response: Final = anthropic_messages_provider_config.transform_anthropic_messages_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
)
return await self._finalize_anthropic_messages_response(
initial_response=rust_messages_response,
initial_response=initial_response,
model=model,
messages=messages,
anthropic_messages_provider_config=anthropic_messages_provider_config,
@ -2304,96 +2358,42 @@ class BaseLLMHTTPHandler:
kwargs=kwargs,
)
response: Final = await self._async_post_anthropic_messages_with_http_error_retry(
async_httpx_client=async_httpx_client,
request_url=request_url,
headers=headers,
signed_json_body=(signed_json_body if signed_json_body is not None else request_body_json),
request_body=request_body,
stream=stream or False,
logging_obj=logging_obj,
provider_config=anthropic_messages_provider_config,
litellm_params=litellm_params,
api_key=api_key,
model=model,
timeout=self._resolve_anthropic_messages_timeout(
litellm_params=litellm_params,
stream=stream or False,
custom_llm_provider=custom_llm_provider,
),
)
# used for logging + cost tracking
logging_obj.model_call_details["httpx_response"] = response
initial_response: AsyncIterator | AnthropicMessagesResponse
if stream:
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
AnthropicMessagesStreamingResponse,
anthropic_messages_stream_hidden_params,
)
completion_stream: Final = anthropic_messages_provider_config.get_async_streaming_response_iterator(
model=model,
httpx_response=response,
request_body=request_body,
litellm_logging_obj=logging_obj,
)
stream_hidden_params: Final = anthropic_messages_stream_hidden_params(response.headers)
if not self._has_agentic_completion_hook(logging_obj):
# No callback overrides async_should_run_agentic_loop, so the
# agentic wrapper's only effect would be buffering every chunk
# and rebuilding the response from SSE at end-of-stream to call
# hooks that all return (False, {}). Stream through directly and
# skip that per-chunk + end-of-stream overhead.
return AnthropicMessagesStreamingResponse(
completion_stream=completion_stream,
hidden_params=stream_hidden_params,
)
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
AgenticAnthropicStreamingIterator,
)
held_back_tool_names: Final = self._server_fulfilled_tools_in_request(
logging_obj=logging_obj,
tools=anthropic_messages_optional_request_params.get("tools"),
)
initial_response = AgenticAnthropicStreamingIterator(
completion_stream=completion_stream,
http_handler=self,
async def adapt_rust_response(response: dict[str, object]) -> AnthropicMessagesResponse | AsyncIterator:
response_obj: Final = cast(AnthropicMessagesResponse, dict(response))
response_obj["_hidden_params"] = {"additional_headers": {"x-litellm-rust": "true"}}
if stream:
return self._rust_anthropic_messages_fake_stream(response_obj)
return await self._finalize_anthropic_messages_response(
initial_response=response_obj,
model=model,
messages=messages,
anthropic_messages_provider_config=anthropic_messages_provider_config,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
kwargs={**kwargs, "api_key": api_key} if api_key else kwargs,
hold_back=bool(held_back_tool_names),
server_fulfilled_tool_names=held_back_tool_names,
)
return AnthropicMessagesStreamingResponse(
completion_stream=initial_response,
hidden_params=stream_hidden_params,
)
else:
initial_response = anthropic_messages_provider_config.transform_anthropic_messages_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
api_key=api_key,
kwargs=kwargs,
)
return await self._finalize_anthropic_messages_response(
initial_response=initial_response,
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"}
return await rust_messages_bridge.amessages(
model=model,
messages=messages,
anthropic_messages_provider_config=anthropic_messages_provider_config,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
body=upstream_body,
has_agentic_hook=self._has_agentic_completion_hook(logging_obj),
api_key=api_key,
kwargs=kwargs,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=headers,
timeout=self._resolve_anthropic_messages_timeout(
litellm_params=litellm_params,
stream=stream or False,
custom_llm_provider=custom_llm_provider,
),
on_request=log_pre_call,
python_fallback=python_fallback,
adapt=adapt_rust_response,
)
async def _finalize_anthropic_messages_response(
@ -2432,39 +2432,6 @@ class BaseLLMHTTPHandler:
"anthropic_messages",
)
@staticmethod
async def _maybe_rust_anthropic_messages(
*,
custom_llm_provider: str,
litellm_params: GenericLiteLLMParams,
has_agentic_hook: bool,
model: str,
api_key: str | None,
api_base: str | None,
headers: dict,
request_body: dict,
timeout: float | httpx.Timeout | None,
) -> AnthropicMessagesResponse | 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"}
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
response_obj: Final = cast(AnthropicMessagesResponse, dict(rust_response))
response_obj["_hidden_params"] = {"additional_headers": {"x-litellm-rust": "true"}}
return response_obj
@staticmethod
def _rust_anthropic_messages_fake_stream(
rust_response: AnthropicMessagesResponse,
@ -6635,27 +6602,26 @@ class BaseLLMHTTPHandler:
async def _backend_connection():
from litellm.rust_bridge.responses import websocket as rust_responses_websocket
rust_backend: Final = await rust_responses_websocket.connect(
async def python_fallback() -> ClientConnection:
return await websockets.connect(
ws_url,
additional_headers=headers,
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
ssl=ssl_context,
)
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,
python_fallback=python_fallback,
)
if rust_backend is not None:
try:
yield rust_backend
finally:
await rust_backend.close()
return
async with websockets.connect(
ws_url,
additional_headers=headers,
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
ssl=ssl_context,
) as backend:
try:
yield backend
finally:
await backend.close()
async with _backend_connection() as backend_ws:
_request_data: Final[dict[str, object]] = {}

View file

@ -80,14 +80,18 @@ async def aocr(*args: object, **kwargs: object) -> OCRResponse: # kwargs-ok: pr
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)
async def adapt(value: OCRResponse) -> OCRResponse:
return value
return await ainvoke(
execution=execution,
native_call=native_call,
python_fallback=python_fallback,
adapt=lambda value: value,
adapt=adapt,
context=BridgeErrorContext(
route=COMPONENT.name.value,
provider=request.custom_llm_provider or "",

View file

@ -1392,50 +1392,34 @@ async def count_request_input_tokens(
Tokenizing is the reservation path's dominant CPU cost and is O(prompt), so
counting a large prompt inline stalls every other request on the worker.
Models whose tokenizer the Rust bridge ports (Anthropic, tiktoken cl100k_base
and o200k_base) are counted from the raw body by the bridge when it is enabled, once per
distinct tokenizer, which parses and tokenizes with the GIL released.
Everything it declines is counted in Python, large prompts in a worker
thread. The counts are reused by both the max-cost and the input-cost
estimate.
The Rust bridge counts the complete request once per distinct tokenizer.
If that operation is unavailable or declines any tokenizer, the complete
request is counted in Python, with large prompts moved to a worker thread.
The counts are reused by both the max-cost and input-cost estimates.
"""
models: Final = _get_request_models(request_body=request_body, route=route, llm_router=llm_router)
if not models:
return MappingProxyType({})
tokenizers: Final[Mapping[str, RustTokenizer | None]] = MappingProxyType(
tokenizers: Final[Mapping[str, RustTokenizer]] = MappingProxyType(
{model: rust_tokenizer(model) for model in models}
)
distinct_tokenizers: Final[tuple[RustTokenizer, ...]] = tuple(
dict.fromkeys(tokenizer for tokenizer in tokenizers.values() if tokenizer is not None)
)
rust_counts_by_tokenizer: Final[Mapping[RustTokenizer, int]] = MappingProxyType(
{
tokenizer: count.input_tokens
for tokenizer in distinct_tokenizers
if raw_body is not None and (count := await count_input_tokens(raw_body, tokenizer)) is not None
}
)
rust_counts: Final = MappingProxyType(
{
model: rust_counts_by_tokenizer[tokenizer]
for model, tokenizer in tokenizers.items()
if tokenizer is not None and tokenizer in rust_counts_by_tokenizer
}
)
python_models: Final = tuple(model for model in models if model not in rust_counts)
python_counts: Final = (
MappingProxyType({})
if not python_models
else _count_input_tokens_for_models(request_body=request_body, models=python_models)
if _approximate_input_size(request_body) < TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS
else await asyncio.to_thread(
async def python_fallback() -> Mapping[str, int]:
if _approximate_input_size(request_body) < TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS:
return _count_input_tokens_for_models(request_body=request_body, models=models)
return await asyncio.to_thread(
_count_input_tokens_for_models,
request_body=request_body,
models=python_models,
models=models,
)
counts: Final = await count_input_tokens(
body=raw_body,
tokenizers=tokenizers,
python_fallback=python_fallback,
)
verbose_proxy_logger.debug("input token counts: rust=%s python=%s", dict(rust_counts), dict(python_counts))
return MappingProxyType({**rust_counts, **python_counts})
verbose_proxy_logger.debug("input token counts: %s", dict(counts))
return counts
def _count_input_tokens_for_models(

View file

@ -22,7 +22,7 @@ Optional capabilities use the `litellm.rust(bool)` process override first, `LITE
## Fallback contract
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
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` permits fallback only when the resolved capability declares Python availability and the caller supplies a real Python callback. Python-capable decisions without a callback, and Rust-required decisions with one, are contract errors
Provider failures, host callback failures, cancellation, conversion failures, and response adaptation failures propagate without replay. Adaptation runs outside the decline-catching boundary

View file

@ -109,6 +109,7 @@ def messages(
extra_headers: object = None,
timeout_seconds: float | None = None,
has_agentic_hook: bool | None = None,
on_request: Callable[[], None] | None = None,
) -> dict[str, object]: ...
def amessages(
model: str,
@ -119,6 +120,7 @@ def amessages(
extra_headers: object = None,
timeout_seconds: float | None = None,
has_agentic_hook: bool | None = None,
on_request: Callable[[], None] | None = None,
) -> Future[dict[str, object]]: ...
def chat_completions(
model: str,
@ -175,38 +177,37 @@ def count_input_tokens(
legacy_accounting: bool,
resource_loader: Callable[[str], str],
) -> Future[dict[str, object]]: ...
def gil_stats() -> dict[str, int]: ...
__all__ = [
"RustBridgeUnavailable",
"_OCR_MAX_FILE_BYTES",
"ResponsesWebSocketConnection",
"RustBridgeDeclined",
"RustBridgeUnavailable",
"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",
"_messages_lifecycle",
"_moderation_lifecycle",
"_ocr_file_document",
"_ocr_lifecycle",
"_ocr_mime_type",
"_ocr_upload_document",
"_rerank_lifecycle",
"ResponsesWebSocketConnection",
"_responses_lifecycle",
"_speech_lifecycle",
"_transcription_lifecycle",
"achat_completions",
"amessages",
"aocr",
"atranscription",
"chat_completions",
"count_input_tokens",
"gil_stats",
"messages",
"ocr",
"transcription",
]

View file

@ -1,3 +1,4 @@
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final
@ -90,7 +91,7 @@ def _component(
return NativeComponent(name=name, capability=capability, exports=exports)
COMPONENTS: Final = MappingProxyType(
COMPONENTS: Final[Mapping[ComponentName, NativeComponent]] = MappingProxyType(
{
ComponentName.OCR: _component(
ComponentName.OCR,

View file

@ -8,7 +8,7 @@ import httpx
from pydantic import TypeAdapter, ValidationError
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
convert_to_model_response_object,
convert_to_model_response_object, # pyright: ignore[reportUnknownVariableType] # legacy converter lacks complete annotations
)
from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned
from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset
@ -112,8 +112,9 @@ def chat_completions(
on_response(rust_response)
return _build_model_response(rust_response, model_response)
native_call: Final[Callable[[], Mapping[str, object]] | None] = (
lambda: rust_chat_completions(
def native_call() -> Mapping[str, object]:
assert rust_chat_completions is not None
return rust_chat_completions(
model=model,
messages=messages,
optional_params=optional_params,
@ -125,12 +126,10 @@ def chat_completions(
host_facts=_host_facts(stream, litellm_params),
on_request=on_request,
)
if rust_chat_completions is not None
else None
)
return invoke(
execution=execution,
native_call=native_call,
native_call=native_call if rust_chat_completions is not None else None,
python_fallback=python_fallback,
adapt=adapt,
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
@ -163,12 +162,13 @@ async def achat_completions(
)
rust_achat_completions: Final = execution.select(_ACHAT)
def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
async 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(
def native_call() -> Awaitable[Mapping[str, object]]:
assert rust_achat_completions is not None
return rust_achat_completions(
model=model,
messages=messages,
optional_params=optional_params,
@ -180,12 +180,10 @@ async def achat_completions(
host_facts=_host_facts(stream, litellm_params),
on_request=on_request,
)
if rust_achat_completions is not None
else None
)
return await ainvoke(
execution=execution,
native_call=native_call,
native_call=native_call if rust_achat_completions is not None else None,
python_fallback=python_fallback,
adapt=adapt,
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),

View file

@ -1,6 +1,6 @@
from __future__ import annotations
from collections.abc import Awaitable, Mapping
from collections.abc import Awaitable, Callable, Mapping
from typing import Protocol
@ -15,6 +15,7 @@ class RustMessages(Protocol):
extra_headers: Mapping[str, object] | None,
timeout_seconds: float | None,
has_agentic_hook: bool = False,
on_request: Callable[[], None] | None = None,
) -> dict[str, object]:
raise NotImplementedError
@ -30,5 +31,6 @@ class RustAmessages(Protocol):
extra_headers: Mapping[str, object] | None,
timeout_seconds: float | None,
has_agentic_hook: bool = False,
on_request: Callable[[], None] | None = None,
) -> Awaitable[dict[str, object]]:
raise NotImplementedError

View file

@ -5,6 +5,7 @@ from __future__ import annotations
from collections.abc import Awaitable, Callable, Mapping
from typing import (
Final,
TypeVar,
cast, # noqa: TID251 # native callable signatures are checked by bridge contract tests
)
@ -28,6 +29,7 @@ def _as_amessages(value: object) -> RustAmessages | None:
_MESSAGES: Final = COMPONENT.bind("messages", validate=_as_messages)
_AMESSAGES: Final = COMPONENT.bind("amessages", validate=_as_amessages)
ResultT = TypeVar("ResultT")
def set_rust_messages(
@ -57,7 +59,10 @@ def messages(
extra_headers: Mapping[str, object] | None,
timeout: float | httpx.Timeout | None,
has_agentic_hook: bool = False,
) -> dict[str, object] | None:
on_request: Callable[[], None] = lambda: None,
python_fallback: Callable[[], ResultT],
adapt: Callable[[dict[str, object]], ResultT],
) -> ResultT:
execution: Final = COMPONENT.resolve(
CapabilityContext(
provider=custom_llm_provider or "",
@ -77,6 +82,7 @@ def messages(
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
on_request=on_request,
)
)
if rust_messages is not None
@ -85,8 +91,8 @@ def messages(
return invoke(
execution=execution,
native_call=native_call,
python_fallback=lambda: None,
adapt=lambda value: value,
python_fallback=python_fallback,
adapt=adapt,
context=BridgeErrorContext(
route=COMPONENT.name.value,
provider=custom_llm_provider or "",
@ -105,7 +111,10 @@ async def amessages(
extra_headers: Mapping[str, object] | None,
timeout: float | httpx.Timeout | None,
has_agentic_hook: bool = False,
) -> dict[str, object] | None:
on_request: Callable[[], None] = lambda: None,
python_fallback: Callable[[], Awaitable[ResultT]],
adapt: Callable[[dict[str, object]], Awaitable[ResultT]],
) -> ResultT:
execution: Final = COMPONENT.resolve(
CapabilityContext(
provider=custom_llm_provider or "",
@ -125,20 +134,18 @@ async def amessages(
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
on_request=on_request,
)
)
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,
adapt=adapt,
context=BridgeErrorContext(
route=COMPONENT.name.value,
provider=custom_llm_provider or "",

View file

@ -1,21 +1,9 @@
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 (
aocr,
load_rust_aocr,
load_rust_ocr,
ocr,
)
from litellm.rust_bridge.ocr.types import LiteLLMOcrRequest
__all__: Final = (
"COMPONENT",
"LiteLLMOcrRequest",
"RustAocr",
"RustOcr",
"aocr",
"load_rust_aocr",
"load_rust_ocr",
"ocr",
)

View file

@ -1,7 +1,7 @@
from __future__ import annotations
from collections.abc import Mapping
from typing import Final, Protocol, cast
from typing import Final, Protocol
import litellm
from litellm.llms.base_llm.ocr.transformation import OCRResponse
@ -35,7 +35,7 @@ class OcrLifecycleHost:
request: LiteLLMOcrRequest,
request_provider: str,
) -> Exception:
mapper: Final = cast(ExceptionMapper, litellm.exception_type) # cast-ok: legacy public exception mapper
mapper: Final[ExceptionMapper] = litellm.exception_type # pyright: ignore[reportAssignmentType] # legacy public mapper is callable
try:
return mapper(
model=request.model.removeprefix(f"{request_provider}/"),

View file

@ -1,8 +1,7 @@
from __future__ import annotations
from collections.abc import Awaitable, Mapping
from collections.abc import Mapping
from dataclasses import dataclass
from typing import Protocol
import httpx
@ -18,35 +17,3 @@ class LiteLLMOcrRequest:
extra_headers: dict[str, object] | None
kwargs: Mapping[str, object]
input_sources: Mapping[str, str] | None = None
class RustOcr(Protocol):
def __call__(
self,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
input_sources: dict[str, str],
timeout_seconds: float | None,
) -> dict[str, object]:
raise NotImplementedError
class RustAocr(Protocol):
def __call__(
self,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
input_sources: dict[str, str],
timeout_seconds: float | None,
) -> Awaitable[dict[str, object]]:
raise NotImplementedError

View file

@ -4,35 +4,9 @@ from __future__ import annotations
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final, cast # noqa: TID251 # native extension exposes dynamically typed callables
import httpx
from typing import Final
from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KEY, OCRResponse
from litellm.rust_bridge.ocr.definition import COMPONENT
from litellm.rust_bridge.ocr.types import RustAocr, RustOcr
from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke, invoke
from litellm.rust_bridge.timeouts import timeout_to_seconds as _timeout_to_seconds
def _as_ocr(value: object) -> RustOcr | None:
return cast(RustOcr, value) if callable(value) else None
def _as_aocr(value: object) -> RustAocr | None:
return cast(RustAocr, value) if callable(value) else None
_OCR: Final = COMPONENT.bind("ocr", validate=_as_ocr)
_AOCR: Final = COMPONENT.bind("aocr", validate=_as_aocr)
def load_rust_ocr() -> RustOcr | None:
return COMPONENT.resolve().select(_OCR)
def load_rust_aocr() -> RustAocr | None:
return COMPONENT.resolve().select(_AOCR)
def adapt_response(response: Mapping[str, object]) -> OCRResponse:
@ -43,80 +17,3 @@ def adapt_response(response: Mapping[str, object]) -> OCRResponse:
if isinstance(provider_native_response, Mapping):
normalized.set_provider_native_response(provider_native_response)
return normalized
def ocr(
*,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
input_sources: Mapping[str, str] | None = None,
) -> dict[str, object] | None:
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),
)
async def aocr(
*,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
input_sources: Mapping[str, str] | None = None,
) -> dict[str, object] | None:
execution: Final = COMPONENT.resolve()
rust_aocr: Final = execution.select(_AOCR)
async def python_fallback() -> None:
return None
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),
)

View file

@ -3,7 +3,7 @@
from __future__ import annotations
from collections.abc import Awaitable, Callable
from typing import Final, Protocol, cast # noqa: TID251 # native class is validated at load time
from typing import Final, Protocol, TypeVar, cast # noqa: TID251 # native class is validated at load time
import httpx
from websockets.exceptions import ConnectionClosedOK
@ -23,6 +23,14 @@ class RustResponsesWebSocket(Protocol):
def close(self) -> Awaitable[None]: ...
class ResponsesWebSocket(Protocol):
def send(self, text: str) -> Awaitable[None]: ...
def recv(self) -> Awaitable[str | bytes]: ...
def close(self) -> Awaitable[None]: ...
class RustResponsesWebSocketConnection(Protocol):
@classmethod
def connect(
@ -71,6 +79,9 @@ class _ConnectionAdapter:
await self._connection.close()
ConnectionT = TypeVar("ConnectionT", bound=ResponsesWebSocket)
async def connect(
*,
url: str,
@ -78,7 +89,8 @@ async def connect(
timeout: float | httpx.Timeout | None,
custom_llm_provider: str | None,
model: str,
) -> _ConnectionAdapter | None:
python_fallback: Callable[[], Awaitable[ConnectionT]],
) -> ResponsesWebSocket:
execution: Final = COMPONENT.resolve(
CapabilityContext(
provider=custom_llm_provider or "",
@ -100,16 +112,14 @@ async def connect(
if connection_type is not None
else None
)
async def python_fallback() -> None:
return None
connection: Final = await ainvoke(
async def adapt(connection: RustResponsesWebSocket) -> ResponsesWebSocket:
return _ConnectionAdapter(connection)
return await ainvoke(
execution=execution,
native_call=native_call,
python_fallback=python_fallback,
adapt=lambda value: value,
adapt=adapt,
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
)
if connection is None:
return None
return _ConnectionAdapter(connection)

View file

@ -3,7 +3,7 @@ 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, overload
from typing import Final, Literal, Protocol, TypeVar, overload
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.configuration import (
@ -55,7 +55,7 @@ class NativeLifecycle(Protocol[RequestT, ResponseT]):
def _lifecycle(value: object) -> NativeLifecycle[object, object] | None:
if not callable(value):
return None
return cast(NativeLifecycle[object, object], value) # cast-ok: callable native entrypoint validated above
return value # pyright: ignore[reportReturnType] # callable shape is validated by native contract tests
@dataclass(frozen=True, slots=True)
@ -65,9 +65,7 @@ class ComponentExecution:
def require_supported(self) -> None:
if self.decision is ExecutionDecision.UNSUPPORTED:
raise RustRouteUnsupportedError(
f"No Python or Rust implementation for {self.route_name.value}"
)
raise RustRouteUnsupportedError(f"No Python or Rust implementation for {self.route_name.value}")
def select(self, binding: NativeBinding[BindingT]) -> BindingT | None:
self.require_supported()
@ -75,9 +73,7 @@ class ComponentExecution:
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"
)
raise RustRouteUnavailableError(f"Rust {self.route_name.value} bridge is unavailable")
return selected
@ -87,10 +83,11 @@ class NativeComponent:
capability: CapabilitySpec
exports: tuple[str, ...]
def resolve(self, context: CapabilityContext = CapabilityContext()) -> ComponentExecution:
def resolve(self, context: CapabilityContext | None = None) -> ComponentExecution:
resolved_context: Final = context if context is not None else CapabilityContext()
return ComponentExecution(
route_name=self.name,
decision=capability_decision(self.capability, context=context),
decision=capability_decision(self.capability, context=resolved_context),
)
def bind(

View file

@ -28,13 +28,15 @@ class BridgeErrorContext:
def invoke(
*,
native_call: Callable[[], NativeT] | None,
python_fallback: Callable[[], ResultT],
python_fallback: Callable[[], ResultT] | None,
adapt: Callable[[NativeT], ResultT],
execution: ComponentExecution,
context: BridgeErrorContext,
) -> ResultT:
execution.require_supported()
_validate_fallback(execution, python_fallback)
if execution.decision is ExecutionDecision.PYTHON:
assert python_fallback is not None
return python_fallback()
if native_call is None:
return _unavailable_or_fallback(execution, python_fallback)
@ -61,20 +63,22 @@ def invoke(
async def ainvoke(
*,
native_call: Callable[[], Awaitable[NativeT]] | None,
python_fallback: Callable[[], Awaitable[ResultT]],
adapt: Callable[[NativeT], ResultT],
python_fallback: Callable[[], Awaitable[ResultT]] | None,
adapt: Callable[[NativeT], Awaitable[ResultT]],
execution: ComponentExecution,
context: BridgeErrorContext,
) -> ResultT:
execution.require_supported()
_validate_fallback(execution, python_fallback)
if execution.decision is ExecutionDecision.PYTHON:
assert python_fallback is not None
return await python_fallback()
if native_call is None:
return await _aunavailable_or_fallback(execution, python_fallback)
exceptions: Final = native_exception_types()
if exceptions is None:
return adapt(await native_call())
return await adapt(await native_call())
declined, upstream = exceptions
unavailable: Final = native_unavailable_exception()
host_callback: Final = native_host_callback_exception()
@ -88,53 +92,65 @@ async def ainvoke(
return await _adeclined_or_fallback(execution, python_fallback, error)
except upstream as error:
_raise_upstream(error, context)
return adapt(value)
return await adapt(value)
def _unavailable_or_fallback(
execution: ComponentExecution,
python_fallback: Callable[[], ResultT],
python_fallback: Callable[[], ResultT] | None,
) -> ResultT:
if execution.decision is ExecutionDecision.RUST_WITH_FALLBACK:
assert python_fallback is not None
return python_fallback()
raise RustRouteUnavailableError(f"Rust {execution.route_name.value} bridge is unavailable")
async def _aunavailable_or_fallback(
execution: ComponentExecution,
python_fallback: Callable[[], Awaitable[ResultT]],
python_fallback: Callable[[], Awaitable[ResultT]] | None,
) -> ResultT:
if execution.decision is ExecutionDecision.RUST_WITH_FALLBACK:
assert python_fallback is not None
return await python_fallback()
raise RustRouteUnavailableError(f"Rust {execution.route_name.value} bridge is unavailable")
def _declined_or_fallback(
execution: ComponentExecution,
python_fallback: Callable[[], ResultT],
python_fallback: Callable[[], ResultT] | None,
error: BaseException,
) -> ResultT:
if execution.decision is ExecutionDecision.RUST_WITH_FALLBACK:
assert python_fallback is not None
return python_fallback()
_raise_declined(execution, error)
async def _adeclined_or_fallback(
execution: ComponentExecution,
python_fallback: Callable[[], Awaitable[ResultT]],
python_fallback: Callable[[], Awaitable[ResultT]] | None,
error: BaseException,
) -> ResultT:
if execution.decision is ExecutionDecision.RUST_WITH_FALLBACK:
assert python_fallback is not None
return await python_fallback()
_raise_declined(execution, error)
def _validate_fallback(execution: ComponentExecution, python_fallback: object | None) -> None:
has_python: Final = execution.decision in (ExecutionDecision.PYTHON, ExecutionDecision.RUST_WITH_FALLBACK)
if has_python and not callable(python_fallback):
raise ValueError(
f"{execution.route_name.value} declares a Python implementation but no callable fallback was supplied"
)
if not has_python and python_fallback is not None:
raise ValueError(f"{execution.route_name.value} declares no Python implementation but a fallback was supplied")
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
raise RustRouteDeclinedError(f"Rust {execution.route_name.value} bridge declined the request: {reason}") from error
def _raise_host_callback(error: BaseException) -> NoReturn:

View file

@ -1,5 +1,7 @@
from __future__ import annotations
from collections.abc import Awaitable, Callable, Mapping
from types import MappingProxyType
from typing import Final, cast # noqa: TID251 # native extension exposes untyped callables
from pydantic import TypeAdapter
@ -47,27 +49,43 @@ def _tokenizer_resource(tokenizer: str) -> str:
raise ValueError(f"unsupported Rust tokenizer resource: {tokenizer}")
async def count_input_tokens(body: bytes, tokenizer: RustTokenizer) -> InputTokenCount | None:
async def count_input_tokens(
*,
body: bytes | None,
tokenizers: Mapping[str, RustTokenizer],
python_fallback: Callable[[], Awaitable[Mapping[str, int]]],
) -> Mapping[str, int]:
counter: Final = TOKEN_COUNTER.load()
distinct_tokenizers: Final = tuple(dict.fromkeys(tokenizers.values()))
async def python_fallback() -> None:
return None
async def native_call() -> Mapping[str, int]:
assert counter is not None
assert body is not None
counts: Final = tuple(
[
_INPUT_TOKEN_COUNT.validate_python(
await counter(
body,
tokenizer.kind,
tokenizer.encoding,
tokenizer.disabled,
tokenizer.legacy_accounting,
_tokenizer_resource,
)
).input_tokens
for tokenizer in distinct_tokenizers
]
)
counts_by_tokenizer: Final = MappingProxyType(dict(zip(distinct_tokenizers, counts)))
return MappingProxyType({model: counts_by_tokenizer[tokenizer] for model, tokenizer in tokenizers.items()})
async def adapt(counts: Mapping[str, int]) -> Mapping[str, int]:
return counts
return await ainvoke(
execution=ComponentExecution(COMPONENT.name, ExecutionDecision.RUST_WITH_FALLBACK),
native_call=(
lambda: counter(
body,
tokenizer.kind,
tokenizer.encoding,
tokenizer.disabled,
tokenizer.legacy_accounting,
_tokenizer_resource,
)
)
if counter is not None
else None,
native_call=native_call if counter is not None and body is not None else None,
python_fallback=python_fallback,
adapt=_INPUT_TOKEN_COUNT.validate_python,
adapt=adapt,
context=BridgeErrorContext(route=COMPONENT.name.value, provider="", model=""),
)

View file

@ -1,5 +1,6 @@
from __future__ import annotations
from collections.abc import Awaitable, Callable
from typing import (
Final,
cast, # noqa: TID251 # native callable signatures are checked by bridge contract tests
@ -54,7 +55,8 @@ def transcription(
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
) -> dict[str, object] | None:
python_fallback: Callable[[], dict[str, object]] | None,
) -> dict[str, object]:
execution: Final = COMPONENT.resolve(CapabilityContext(provider=custom_llm_provider or "", model=model))
rust_transcription: Final = execution.select(_TRANSCRIPTION)
return invoke(
@ -73,7 +75,7 @@ def transcription(
)
if rust_transcription is not None
else None,
python_fallback=lambda: None,
python_fallback=python_fallback,
adapt=lambda response: response,
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
)
@ -89,11 +91,13 @@ async def atranscription(
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
) -> dict[str, object] | None:
python_fallback: Callable[[], Awaitable[dict[str, object]]] | None,
) -> dict[str, object]:
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
async def adapt(response: dict[str, object]) -> dict[str, object]:
return response
return await ainvoke(
execution=execution,
@ -112,6 +116,6 @@ async def atranscription(
if rust_atranscription is not None
else None,
python_fallback=python_fallback,
adapt=lambda response: response,
adapt=adapt,
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
)

View file

@ -12,7 +12,6 @@ from litellm.rust_bridge import configuration
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
from litellm.types.router import GenericLiteLLMParams
rust_messages = importlib.import_module("litellm.rust_bridge.messages")
rust_bridge_loader = importlib.import_module("litellm.rust_bridge.loader")
@ -32,6 +31,7 @@ REQUEST_BODY: dict[str, object] = {
"max_tokens": 64,
"messages": [{"role": "user", "content": "hi"}],
}
PYTHON_MESSAGES_RESPONSE: dict[str, object] = {"id": "python_fallback"}
class RecordingMessages:
@ -48,6 +48,7 @@ class RecordingMessages:
extra_headers: dict[str, object] | None,
timeout_seconds: float | None,
has_agentic_hook: bool = False,
on_request=None,
) -> dict[str, object]:
self.calls.append(
{
@ -60,6 +61,8 @@ class RecordingMessages:
"timeout_seconds": timeout_seconds,
}
)
if on_request is not None:
on_request()
return dict(FAKE_MESSAGES_RESPONSE)
@ -77,6 +80,7 @@ class RecordingAsyncMessages:
extra_headers: dict[str, object] | None,
timeout_seconds: float | None,
has_agentic_hook: bool = False,
on_request=None,
) -> dict[str, object]:
self.calls.append(
{
@ -89,6 +93,8 @@ class RecordingAsyncMessages:
"timeout_seconds": timeout_seconds,
}
)
if on_request is not None:
on_request()
return dict(FAKE_MESSAGES_RESPONSE)
@ -110,6 +116,12 @@ class RaisingAsyncMessages:
raise RuntimeError("upstream request failed with status 400: bad request")
class DecliningAsyncMessages:
async def __call__(self, **kwargs: object) -> dict[str, object]:
native = pytest.importorskip("litellm.rust_bridge._native")
raise native.RustBridgeDeclined("unsupported request")
@pytest.fixture(autouse=True)
def _reset_rust_flag():
rust_messages.set_rust_messages(messages=None, amessages=None)
@ -135,7 +147,7 @@ def test_load_rust_amessages_returns_injected_impl():
assert rust_messages.load_rust_amessages() is bridge
def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch):
def test_messages_wrapper_returns_fallback_when_bridge_absent(monkeypatch):
monkeypatch.setattr(
importlib.import_module("litellm.rust_bridge.bindings"),
"get_native_bridge",
@ -151,8 +163,10 @@ def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch):
custom_llm_provider="azure_ai",
extra_headers={},
timeout=30.0,
python_fallback=lambda: dict(PYTHON_MESSAGES_RESPONSE),
adapt=lambda response: response,
)
assert result is None
assert result == PYTHON_MESSAGES_RESPONSE
def test_messages_wrapper_forwards_args_and_converts_timeout():
@ -168,6 +182,8 @@ def test_messages_wrapper_forwards_args_and_converts_timeout():
custom_llm_provider="azure_ai",
extra_headers={"anthropic-beta": "token-efficient-tools-2025-02-19"},
timeout=httpx.Timeout(600.0, read=42.0),
python_fallback=lambda: pytest.fail("native request should not fall back"),
adapt=lambda response: response,
)
assert response == FAKE_MESSAGES_RESPONSE
@ -188,6 +204,12 @@ async def test_amessages_wrapper_forwards_args():
litellm.rust(True)
rust_messages.set_rust_messages(amessages=bridge)
async def python_fallback() -> dict[str, object]:
pytest.fail("native request should not fall back")
async def adapt(response: dict[str, object]) -> dict[str, object]:
return response
response = await rust_messages.amessages(
model="claude-sonnet-4-5",
body=REQUEST_BODY,
@ -196,6 +218,8 @@ async def test_amessages_wrapper_forwards_args():
custom_llm_provider="azure_ai",
extra_headers=None,
timeout=12.5,
python_fallback=python_fallback,
adapt=adapt,
)
assert response == FAKE_MESSAGES_RESPONSE
@ -203,10 +227,9 @@ async def test_amessages_wrapper_forwards_args():
assert bridge.calls[0]["timeout_seconds"] == 12.5
def _gate(**overrides):
async def _gate(**overrides):
kwargs = {
"custom_llm_provider": "azure_ai",
"litellm_params": GenericLiteLLMParams(api_key="sk-azure"),
"has_agentic_hook": False,
"model": "claude-sonnet-4-5",
"api_key": "sk-azure",
@ -216,7 +239,28 @@ def _gate(**overrides):
"timeout": 30.0,
}
kwargs.update(overrides)
return BaseLLMHTTPHandler._maybe_rust_anthropic_messages(**kwargs)
request_body = kwargs.pop("request_body")
async def python_fallback() -> dict[str, object]:
return dict(PYTHON_MESSAGES_RESPONSE)
async def adapt(response: dict[str, object]) -> dict[str, object]:
adapted = dict(response)
adapted["_hidden_params"] = {"additional_headers": {"x-litellm-rust": "true"}}
return adapted
return await rust_messages.amessages(
model=kwargs["model"],
body={key: value for key, value in request_body.items() if key != "stream"},
has_agentic_hook=kwargs["has_agentic_hook"],
api_key=kwargs["api_key"],
api_base=kwargs["api_base"],
custom_llm_provider=kwargs["custom_llm_provider"],
extra_headers=kwargs["headers"],
timeout=kwargs["timeout"],
python_fallback=python_fallback,
adapt=adapt,
)
@pytest.mark.asyncio
@ -256,9 +300,9 @@ async def test_gate_skips_rust_when_flag_absent(monkeypatch):
bridge = ExplodingAsyncMessages()
rust_messages.set_rust_messages(amessages=bridge)
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure"))
response = await _gate()
assert response is None
assert response == PYTHON_MESSAGES_RESPONSE
assert bridge.calls == 0
@ -268,7 +312,7 @@ async def test_gate_uses_process_enable_without_request_override():
rust_messages.set_rust_messages(amessages=bridge)
litellm.rust(True)
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure"))
response = await _gate()
assert response is not None
assert bridge.calls[0]["custom_llm_provider"] == "azure_ai"
@ -282,7 +326,6 @@ async def test_gate_invokes_rust_for_native_anthropic_provider():
response = await _gate(
custom_llm_provider="anthropic",
litellm_params=GenericLiteLLMParams(api_key="sk-ant"),
api_key="sk-ant",
api_base="https://api.anthropic.com",
headers={"x-api-key": "sk-ant", "anthropic-version": "2023-06-01"},
@ -302,7 +345,6 @@ async def test_gate_invokes_rust_when_env_var_set(monkeypatch):
response = await _gate(
custom_llm_provider="anthropic",
litellm_params=GenericLiteLLMParams(api_key="sk-ant"),
)
assert response is not None
@ -317,29 +359,26 @@ async def test_gate_env_var_falsey_does_not_enable(monkeypatch):
response = await _gate(
custom_llm_provider="anthropic",
litellm_params=GenericLiteLLMParams(api_key="sk-ant"),
)
assert response is None
assert response == PYTHON_MESSAGES_RESPONSE
assert bridge.calls == 0
@pytest.mark.asyncio
async def test_gate_skips_rust_for_unsupported_provider():
native = pytest.importorskip("litellm.rust_bridge._native")
async def test_gate_falls_back_for_unsupported_provider():
litellm.rust(True)
rust_messages.set_rust_messages(amessages=native.amessages)
rust_messages.set_rust_messages(amessages=DecliningAsyncMessages())
response = await _gate(custom_llm_provider="openai", api_base="http://127.0.0.1:1")
assert response is None
assert response == PYTHON_MESSAGES_RESPONSE
@pytest.mark.asyncio
async def test_gate_skips_rust_for_agentic_hook():
native = pytest.importorskip("litellm.rust_bridge._native")
async def test_gate_falls_back_for_agentic_hook():
litellm.rust(True)
rust_messages.set_rust_messages(amessages=native.amessages)
rust_messages.set_rust_messages(amessages=DecliningAsyncMessages())
response = await _gate(has_agentic_hook=True, api_base="http://127.0.0.1:1")
assert response is None
assert response == PYTHON_MESSAGES_RESPONSE
@pytest.mark.asyncio
@ -387,4 +426,4 @@ async def test_gate_falls_back_when_bridge_unavailable(monkeypatch):
response = await _gate()
assert response is None
assert response == PYTHON_MESSAGES_RESPONSE

View file

@ -7,7 +7,7 @@ from litellm.rust_bridge.ocr import value as rust_ocr_bridge
def test_rust_ocr_response_retains_provider_native_response():
provider_response = {"status": "succeeded", "analyzeResult": {"content": "native"}}
response = rust_ocr_bridge._response(
response = rust_ocr_bridge.adapt_response(
{
"pages": [],
"model": "prebuilt-layout",

View file

@ -4,6 +4,7 @@ import json
import math
from types import MappingProxyType
from typing import Final
from unittest.mock import AsyncMock
import pytest
@ -12,6 +13,7 @@ from litellm.caching import DualCache
from litellm.proxy import proxy_server
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.spend_tracking import budget_reservation
from litellm.proxy.spend_tracking.budget_reservation import (
count_request_input_tokens,
estimate_request_max_cost,
@ -232,58 +234,37 @@ class _FakeNative:
RustUpstreamError = _FakeUpstream
class _RecordingCounter:
"""Stands in for one native counter; records `(tokenizer, body)` on the shared factory."""
def __init__(self, factory: _RecordingFactory, tokenizer: rust_token_counter.RustTokenizer) -> None:
self.factory = factory
self.tokenizer = tokenizer
async def acount_request(self, body: bytes) -> object:
self.factory.calls.append((self.tokenizer, body))
return {"model": "", "input_tokens": RUST_INPUT_TOKENS_BY_TOKENIZER[self.tokenizer]}
class _RecordingFactory:
"""Stands in for the native `TokenCounter` class: called with tokenizer JSON, or `from_*_ranks`."""
def __init__(self) -> None:
self.calls: list[tuple[rust_token_counter.RustTokenizer, bytes]] = []
self.calls: list[tuple[str, bytes]] = []
def __call__(self, tokenizer_json: str) -> _RecordingCounter:
return _RecordingCounter(self, "anthropic")
def from_cl100k_ranks(self, rank_file: str) -> _RecordingCounter:
return _RecordingCounter(self, "cl100k_base")
def from_o200k_ranks(self, rank_file: str) -> _RecordingCounter:
return _RecordingCounter(self, "o200k_base")
class _DecliningCounter:
async def acount_request(self, body: bytes) -> object:
raise _FakeDeclined("unsupported content block")
async def __call__(
self,
body: bytes,
kind: str | None,
encoding: str,
disabled: bool,
legacy_accounting: bool,
resource_loader,
) -> object:
tokenizer: Final = kind or encoding
self.calls.append((tokenizer, body))
if disabled or legacy_accounting or tokenizer not in RUST_INPUT_TOKENS_BY_TOKENIZER:
raise _FakeDeclined("unsupported tokenizer configuration")
return {"model": "", "input_tokens": RUST_INPUT_TOKENS_BY_TOKENIZER[tokenizer]}
class _DecliningFactory:
def __call__(self, tokenizer_json: str) -> _DecliningCounter:
return _DecliningCounter()
def from_cl100k_ranks(self, rank_file: str) -> _DecliningCounter:
return _DecliningCounter()
def from_o200k_ranks(self, rank_file: str) -> _DecliningCounter:
return _DecliningCounter()
async def __call__(self, *args, **kwargs) -> object:
raise _FakeDeclined("unsupported content block")
@pytest.fixture
def rust_counter(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(bindings, "get_native_bridge", lambda: _FakeNative())
rust_token_counter._counter.cache_clear()
configuration.reset_rust_configuration()
yield
rust_token_counter.TOKEN_COUNTER.reset()
rust_token_counter._counter.cache_clear()
configuration.reset_rust_configuration()
@ -350,7 +331,7 @@ async def test_tiktoken_o200k_models_are_counted_by_rust(rust_counter: None, mod
@pytest.mark.asyncio
async def test_multi_model_request_counts_once_per_tokenizer_and_python_for_the_rest(rust_counter: None) -> None:
async def test_multi_model_decline_discards_every_native_count(rust_counter: None) -> None:
factory: Final = _RecordingFactory()
litellm.rust(True)
rust_token_counter.TOKEN_COUNTER.override(factory)
@ -372,16 +353,14 @@ async def test_multi_model_request_counts_once_per_tokenizer_and_python_for_the_
request_body=body, route="/v1/chat/completions", llm_router=None, raw_body=raw_body
)
assert factory.calls == [("cl100k_base", raw_body), ("anthropic", raw_body), ("o200k_base", raw_body)]
assert dict(counts) == {
CL100K_MODEL: RUST_INPUT_TOKENS_BY_TOKENIZER["cl100k_base"],
"gemini/gemini-2.5-pro": RUST_INPUT_TOKENS_BY_TOKENIZER["cl100k_base"],
ANTHROPIC_TOKENIZER_MODEL: RUST_INPUT_TOKENS,
O200K_MODEL: RUST_INPUT_TOKENS_BY_TOKENIZER["o200k_base"],
"gpt-5": RUST_INPUT_TOKENS_BY_TOKENIZER["o200k_base"],
"replicate/meta/llama-2-70b-chat": python_counts["replicate/meta/llama-2-70b-chat"],
}
assert counts["replicate/meta/llama-2-70b-chat"] not in RUST_INPUT_TOKENS_BY_TOKENIZER.values()
assert factory.calls == [
("cl100k_base", raw_body),
("anthropic", raw_body),
("o200k_base", raw_body),
("llama2", raw_body),
]
assert dict(counts) == dict(python_counts)
assert not set(counts.values()) & set(RUST_INPUT_TOKENS_BY_TOKENIZER.values())
@pytest.mark.asyncio
@ -431,7 +410,7 @@ async def test_direct_budget_counter_ignores_disabled_public_rollout(rust_counte
@pytest.mark.asyncio
@pytest.mark.parametrize("model", ("replicate/meta/llama-2-70b-chat", "meta-llama/Llama-3-8b", "text-davinci-003"))
async def test_models_without_a_rust_tokenizer_stay_in_python(
async def test_models_with_unsupported_native_configuration_fall_back_to_python(
rust_counter: None, monkeypatch: pytest.MonkeyPatch, model: str
) -> None:
monkeypatch.setattr(
@ -449,6 +428,28 @@ async def test_models_without_a_rust_tokenizer_stay_in_python(
request_body=body, route="/v1/chat/completions", llm_router=None, raw_body=json.dumps(body).encode()
)
assert factory.calls == []
tokenizer: Final = rust_token_counter.rust_tokenizer(model)
assert factory.calls == [(tokenizer.kind or tokenizer.encoding, json.dumps(body).encode())]
assert dict(counts) == dict(python_counts)
assert counts[model] not in RUST_INPUT_TOKENS_BY_TOKENIZER.values()
@pytest.mark.asyncio
async def test_large_python_fallback_runs_off_event_loop(
rust_counter: None,
monkeypatch: pytest.MonkeyPatch,
) -> None:
litellm.rust(False)
body: Final = {"model": O200K_MODEL, "input": "x" * budget_reservation.TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS}
offload: Final = AsyncMock(side_effect=lambda function, **kwargs: function(**kwargs))
monkeypatch.setattr(budget_reservation.asyncio, "to_thread", offload)
counts: Final = await count_request_input_tokens(
request_body=body,
route="/v1/responses",
llm_router=None,
raw_body=json.dumps(body).encode(),
)
assert counts[O200K_MODEL] > 0
offload.assert_awaited_once()

View file

@ -14,9 +14,15 @@ class _FakeNativeConnection:
async def send_text(self, text: str) -> None:
self.sent.append(text)
async def send(self, text: str) -> None:
self.sent.append(text)
async def recv_text(self) -> str:
return "response.completed"
async def recv(self) -> str:
return "response.completed"
async def close(self) -> None:
self.closed = True
@ -39,6 +45,10 @@ class _FakeNativeBridge:
return _FakeNativeConnection()
async def _no_connection() -> _FakeNativeConnection:
raise AssertionError("Python fallback must not connect")
@pytest.fixture(autouse=True)
def reset_responses_websocket():
responses_websocket.set_rust_responses_websocket(connection=None)
@ -57,20 +67,22 @@ async def test_adapter_raises_clean_close_when_rust_connection_ends() -> None:
@pytest.mark.asyncio
async def test_bridge_unavailable_returns_none(monkeypatch: pytest.MonkeyPatch) -> None:
async def test_bridge_unavailable_executes_fallback(monkeypatch: pytest.MonkeyPatch) -> None:
configuration.rust(True)
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None)
fallback = _FakeNativeConnection()
assert (
await responses_websocket.connect(
url="wss://example.test/responses",
custom_llm_provider="openai",
model="test",
headers={},
timeout=None,
)
is None
)
async def python_fallback() -> _FakeNativeConnection:
return fallback
assert await responses_websocket.connect(
url="wss://example.test/responses",
custom_llm_provider="openai",
model="test",
headers={},
timeout=None,
python_fallback=python_fallback,
) is fallback
@pytest.mark.asyncio
@ -86,6 +98,7 @@ async def test_enabled_bridge_connects_and_adapts_socket(
model="test",
headers={"Authorization": "Bearer key"},
timeout=1.0,
python_fallback=_no_connection,
)
assert connection is not None
@ -103,16 +116,19 @@ async def test_disabled_websocket_does_not_connect() -> None:
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
)
fallback = _FakeNativeConnection()
async def python_fallback() -> _FakeNativeConnection:
return fallback
assert await responses_websocket.connect(
url="ws://127.0.0.1:1",
headers={},
timeout=0.1,
custom_llm_provider="openai",
model="test",
python_fallback=python_fallback,
) is fallback
@pytest.mark.asyncio
@ -122,16 +138,19 @@ async def test_native_websocket_decline_falls_back_but_connection_failure_does_n
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
)
fallback = _FakeNativeConnection()
async def python_fallback() -> _FakeNativeConnection:
return fallback
assert await responses_websocket.connect(
url="ws://127.0.0.1:1",
headers={},
timeout=0.1,
custom_llm_provider="azure",
model="test",
python_fallback=python_fallback,
) is fallback
with pytest.raises(APIError):
await responses_websocket.connect(
url="ws://127.0.0.1:1",
@ -139,4 +158,5 @@ async def test_native_websocket_decline_falls_back_but_connection_failure_does_n
timeout=0.1,
custom_llm_provider="openai",
model="test",
python_fallback=python_fallback,
)

View file

@ -21,7 +21,8 @@ from litellm.rust_bridge.configuration import (
)
from litellm.rust_bridge.errors import RustRouteUnsupportedError
from litellm.rust_bridge.route import NativeComponent
from litellm.rust_bridge.token_counter import COMPONENT, definition
from litellm.rust_bridge.token_counter import definition
from litellm.rust_bridge.token_counter import COMPONENT
@pytest.mark.parametrize(("value", "expected"), (("1", True), ("0", False), (" 1 ", True), (" 0 ", False)))

View file

@ -12,10 +12,10 @@ from litellm.rust_bridge.catalog import COMPONENTS
from litellm.rust_bridge.configuration import (
CapabilityContext,
CapabilityDefinition,
ComponentName,
DeliveryMode,
ExecutionDecision,
RolloutPolicy,
ComponentName,
RustImplementationState,
)
from litellm.rust_bridge.errors import RustRouteUnavailableError, RustRouteUnsupportedError

View file

@ -9,7 +9,7 @@ import pytest
from litellm.exceptions import APIError
from litellm.rust_bridge import bindings, runtime
from litellm.rust_bridge.configuration import ComponentName, ExecutionDecision
from litellm.rust_bridge.configuration import ExecutionDecision, ComponentName
from litellm.rust_bridge.errors import RustRouteDeclinedError, RustRouteUnavailableError, RustRouteUnsupportedError
from litellm.rust_bridge.route import ComponentExecution
@ -46,7 +46,7 @@ async def _invoke(
decision: ExecutionDecision,
native_call: Callable[[], object] | None,
adapt: Callable[[object], object],
fallback: Callable[[], object] = lambda: "python",
fallback: Callable[[], object] | None = lambda: "python",
) -> object:
execution: Final = ComponentExecution(route_name=ComponentName.MESSAGES, decision=decision)
context: Final = runtime.BridgeErrorContext(route="messages", provider="anthropic", model="model")
@ -64,13 +64,17 @@ async def _invoke(
return native_call()
async def afallback() -> object:
assert fallback is not None
return fallback()
async def aadapt(value: object) -> object:
return adapt(value)
return await runtime.ainvoke(
execution=execution,
native_call=call if native_call is not None else None,
python_fallback=afallback,
adapt=adapt,
python_fallback=afallback if fallback is not None else None,
adapt=aadapt,
context=context,
)
@ -111,11 +115,32 @@ async def test_unavailable_and_declined_follow_policy(
if required:
expected: Final = RustRouteDeclinedError if isinstance(error, RustBridgeDeclined) else RustRouteUnavailableError
with pytest.raises(expected):
await _invoke(asynchronous, decision, native_call, str)
await _invoke(asynchronous, decision, native_call, str, None)
return
assert await _invoke(asynchronous, decision, native_call, str) == "python"
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
@pytest.mark.parametrize("native_call", (None, lambda: (_ for _ in ()).throw(RustBridgeDeclined("declined"))))
async def test_optional_failure_runs_python_exactly_once(
asynchronous: bool,
native_call: Callable[[], object] | None,
) -> None:
fallback_calls: Final[list[None]] = []
result: Final = await _invoke(
asynchronous,
ExecutionDecision.RUST_WITH_FALLBACK,
native_call,
str,
lambda: fallback_calls.append(None) or "python",
)
assert result == "python"
assert fallback_calls == [None]
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
@pytest.mark.parametrize("decision", (ExecutionDecision.RUST_WITH_FALLBACK, ExecutionDecision.RUST_REQUIRED))
@ -129,12 +154,37 @@ async def test_native_success_does_not_run_fallback(
decision,
lambda: value,
lambda native: native,
lambda: calls.append("python") or "python",
None if decision is ExecutionDecision.RUST_REQUIRED else lambda: calls.append("python") or "python",
)
assert result is value
assert calls == []
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
@pytest.mark.parametrize("decision", (ExecutionDecision.PYTHON, ExecutionDecision.RUST_WITH_FALLBACK))
async def test_python_capability_requires_fallback(asynchronous: bool, decision: ExecutionDecision) -> None:
native_calls: Final[list[bool]] = []
with pytest.raises(ValueError, match="declares a Python implementation"):
await _invoke(asynchronous, decision, lambda: native_calls.append(True), str, None)
assert native_calls == []
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
async def test_rust_required_rejects_python_fallback(asynchronous: bool) -> None:
native_calls: Final[list[bool]] = []
with pytest.raises(ValueError, match="declares no Python implementation"):
await _invoke(
asynchronous,
ExecutionDecision.RUST_REQUIRED,
lambda: native_calls.append(True),
str,
lambda: "python",
)
assert native_calls == []
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
async def test_upstream_failure_never_falls_back(asynchronous: bool) -> None:
@ -170,6 +220,31 @@ async def test_adaptation_failure_never_falls_back(asynchronous: bool, error: Ex
assert caught.value is error
@pytest.mark.asyncio
async def test_async_adaptation_is_awaited_and_failure_never_falls_back() -> None:
fallback_calls: Final[list[None]] = []
async def adapt(_value: object) -> object:
await asyncio.sleep(0)
raise RuntimeError("async adapt failed")
execution: Final = ComponentExecution(
route_name=ComponentName.MESSAGES,
decision=ExecutionDecision.RUST_WITH_FALLBACK,
)
with pytest.raises(RuntimeError, match="async adapt failed"):
await runtime.ainvoke(
execution=execution,
native_call=lambda: asyncio.sleep(0, result="native"),
python_fallback=lambda: asyncio.sleep(0, result=fallback_calls.append(None)),
adapt=adapt,
context=runtime.BridgeErrorContext(route="messages", provider="anthropic", model="model"),
)
assert fallback_calls == []
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
async def test_host_callback_failure_preserves_its_cause(asynchronous: bool) -> None:

View file

@ -1,7 +1,7 @@
from __future__ import annotations
import json
from collections.abc import Callable, Iterator
from collections.abc import Awaitable, Callable, Iterator, Mapping
from typing import Final
import pytest
@ -20,7 +20,6 @@ BODY: Final = json.dumps({"model": MODEL, "messages": [{"role": "user", "content
ANTHROPIC: Final = bridge.RustTokenizer(kind="anthropic", encoding="", disabled=False, legacy_accounting=False)
CL100K: Final = bridge.RustTokenizer(kind=None, encoding="cl100k_base", disabled=False, legacy_accounting=False)
O200K: Final = bridge.RustTokenizer(kind=None, encoding="o200k_base", disabled=False, legacy_accounting=False)
TOKENIZERS: Final = (ANTHROPIC, CL100K, O200K)
_FakeDeclined = native.RustBridgeDeclined
_FakeUnavailable = native.RustBridgeUnavailable
@ -34,8 +33,8 @@ class _FakeNative:
class _RecordingCounter:
def __init__(self, error: Exception | None = None) -> None:
self.error: Final = error
def __init__(self, errors: Mapping[str, Exception] | None = None) -> None:
self.errors: Final = errors or {}
self.calls: Final[list[tuple[bytes, bridge.RustTokenizer, str]]] = []
async def __call__(
@ -50,10 +49,10 @@ class _RecordingCounter:
tokenizer: Final = bridge.RustTokenizer(kind, encoding, disabled, legacy_accounting)
resource_name: Final = kind or encoding
self.calls.append((body, tokenizer, resource_name))
if self.error is not None:
raise self.error
if error := self.errors.get(resource_name):
raise error
resource_loader(resource_name)
return {"model": MODEL, "input_tokens": 42}
return {"model": MODEL, "input_tokens": {"anthropic": 42, "cl100k_base": 17}[resource_name]}
@pytest.fixture(autouse=True)
@ -67,61 +66,118 @@ def reset_bridge(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
configuration.reset_rust_configuration()
def _fallback(result: Mapping[str, int], calls: list[None]) -> Callable[[], Awaitable[Mapping[str, int]]]:
async def call() -> Mapping[str, int]:
calls.append(None)
return result
return call
@pytest.mark.asyncio
@pytest.mark.parametrize("tokenizer", TOKENIZERS)
async def test_direct_bridge_bypasses_disabled_public_rollout(tokenizer: bridge.RustTokenizer) -> None:
async def test_direct_bridge_bypasses_disabled_public_rollout() -> None:
counter: Final = _RecordingCounter()
bridge.TOKEN_COUNTER.override(counter)
fallback_calls: list[None] = []
litellm.rust(False)
assert await bridge.count_input_tokens(BODY, tokenizer) == bridge.InputTokenCount(model=MODEL, input_tokens=42)
assert counter.calls == [(BODY, tokenizer, tokenizer.kind or tokenizer.encoding)]
result: Final = await bridge.count_input_tokens(
body=BODY,
tokenizers={MODEL: ANTHROPIC},
python_fallback=_fallback({MODEL: 9}, fallback_calls),
)
assert result == {MODEL: 42}
assert fallback_calls == []
assert counter.calls == [(BODY, ANTHROPIC, "anthropic")]
@pytest.mark.asyncio
@pytest.mark.parametrize("tokenizer", TOKENIZERS)
async def test_one_native_count_entrypoint_receives_configuration_and_body(tokenizer: bridge.RustTokenizer) -> None:
async def test_native_count_deduplicates_tokenizers() -> None:
counter: Final = _RecordingCounter()
bridge.TOKEN_COUNTER.override(counter)
result: Final = await bridge.count_input_tokens(BODY, tokenizer)
assert result == bridge.InputTokenCount(model=MODEL, input_tokens=42)
assert counter.calls == [(BODY, tokenizer, tokenizer.kind or tokenizer.encoding)]
fallback_calls: list[None] = []
result: Final = await bridge.count_input_tokens(
body=BODY,
tokenizers={MODEL: ANTHROPIC, "other-claude": ANTHROPIC, CL100K_MODEL: CL100K},
python_fallback=_fallback({}, fallback_calls),
)
assert result == {MODEL: 42, "other-claude": 42, CL100K_MODEL: 17}
assert fallback_calls == []
assert counter.calls == [
(BODY, ANTHROPIC, "anthropic"),
(BODY, CL100K, "cl100k_base"),
]
@pytest.mark.asyncio
@pytest.mark.parametrize("error", (_FakeDeclined("unsupported"), _FakeUnavailable("resource")))
async def test_decline_and_resource_unavailability_fall_back(error: Exception) -> None:
bridge.TOKEN_COUNTER.override(_RecordingCounter(error))
assert await bridge.count_input_tokens(BODY, ANTHROPIC) is None
async def test_any_decline_or_unavailability_discards_partial_native_counts(error: Exception) -> None:
counter: Final = _RecordingCounter({"cl100k_base": error})
bridge.TOKEN_COUNTER.override(counter)
fallback_calls: list[None] = []
result: Final = await bridge.count_input_tokens(
body=BODY,
tokenizers={MODEL: ANTHROPIC, CL100K_MODEL: CL100K},
python_fallback=_fallback({MODEL: 9, CL100K_MODEL: 8}, fallback_calls),
)
assert result == {MODEL: 9, CL100K_MODEL: 8}
assert fallback_calls == [None]
assert [call[2] for call in counter.calls] == ["anthropic", "cl100k_base"]
@pytest.mark.asyncio
async def test_unexpected_counting_failure_propagates() -> None:
bridge.TOKEN_COUNTER.override(_RecordingCounter(RuntimeError("encode failed")))
@pytest.mark.parametrize("body", (None, BODY))
async def test_missing_body_or_binding_runs_complete_python_fallback(body: bytes | None) -> None:
bridge.TOKEN_COUNTER.override(None)
fallback_calls: list[None] = []
result: Final = await bridge.count_input_tokens(
body=body,
tokenizers={MODEL: ANTHROPIC},
python_fallback=_fallback({MODEL: 9}, fallback_calls),
)
assert result == {MODEL: 9}
assert fallback_calls == [None]
@pytest.mark.asyncio
async def test_unexpected_counting_failure_propagates_without_python_replay() -> None:
counter: Final = _RecordingCounter({"anthropic": RuntimeError("encode failed")})
bridge.TOKEN_COUNTER.override(counter)
fallback_calls: list[None] = []
with pytest.raises(RuntimeError, match="encode failed"):
await bridge.count_input_tokens(BODY, ANTHROPIC)
await bridge.count_input_tokens(
body=BODY,
tokenizers={MODEL: ANTHROPIC},
python_fallback=_fallback({MODEL: 9}, fallback_calls),
)
assert fallback_calls == []
@pytest.mark.parametrize(
("model", "expected"),
(
(MODEL, ANTHROPIC),
("gpt-4", CL100K),
("gpt-4o", O200K),
(CL100K_MODEL, CL100K),
(O200K_MODEL, O200K),
("replicate/meta/llama-2-70b-chat", bridge.RustTokenizer("llama2", "", False, False)),
),
)
def test_tokenizer_configuration_matches_python_selection(
model: str, expected: bridge.RustTokenizer
) -> None:
bridge.TOKEN_COUNTER.override(_RecordingCounter())
def test_tokenizer_configuration_matches_python_selection(model: str, expected: bridge.RustTokenizer) -> None:
assert bridge.rust_tokenizer(model) == expected
def test_unsupported_configuration_is_left_for_native_admission(monkeypatch: pytest.MonkeyPatch) -> None:
bridge.TOKEN_COUNTER.override(_RecordingCounter())
monkeypatch.setattr(litellm, "disable_token_counter", True)
tokenizer: Final = bridge.rust_tokenizer(MODEL)
assert tokenizer == bridge.RustTokenizer("anthropic", "", True, False)
assert bridge.rust_tokenizer(MODEL) == bridge.RustTokenizer("anthropic", "", True, False)
PARITY_REQUESTS: Final[tuple[dict[str, object], ...]] = (
@ -149,17 +205,31 @@ async def test_native_count_matches_python_budget_counter(
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
bridge.TOKEN_COUNTER.reset()
body: Final = json.dumps(request_body).replace(MODEL, model)
rust_count: Final = await bridge.count_input_tokens(body.encode(), tokenizer)
python_count: Final = _count_input_tokens(request_body=json.loads(body), model=model)
assert rust_count is not None
assert rust_count.input_tokens == python_count
result: Final = await bridge.count_input_tokens(
body=body.encode(),
tokenizers={model: tokenizer},
python_fallback=_fallback({}, []),
)
assert result == {model: python_count}
@pytest.mark.requires_rust_extension
@pytest.mark.asyncio
async def test_native_declines_unsupported_request_without_loading_resources(
async def test_native_declines_unsupported_request_and_runs_complete_fallback(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
bridge.TOKEN_COUNTER.reset()
assert await bridge.count_input_tokens(b'{"input":1.5}', ANTHROPIC) is None
fallback_calls: list[None] = []
result: Final = await bridge.count_input_tokens(
body=b'{"input":1.5}',
tokenizers={MODEL: ANTHROPIC},
python_fallback=_fallback({MODEL: 7}, fallback_calls),
)
assert result == {MODEL: 7}
assert fallback_calls == [None]

View file

@ -71,6 +71,7 @@ def test_enabled_sync_bridge_receives_audio(enabled: bool) -> None:
extra_headers=None,
optional_params={"temperature": 0},
timeout=5.0,
python_fallback=None,
)
assert result == {"text": "hello"}
assert bridge.calls[0]["audio"] == {"data": "AQI=", "format": "wav", "filename": "audio.wav"}
@ -90,6 +91,7 @@ async def test_enabled_async_bridge(enabled: bool) -> None:
extra_headers=None,
optional_params={},
timeout=None,
python_fallback=None,
)
assert result == {"text": "async"}
@ -214,6 +216,7 @@ async def test_bedrock_transcription_errors_never_fall_back(
extra_headers=None,
optional_params={},
timeout=None,
python_fallback=None,
)
with pytest.raises(expected, match=message):
await rust_bridge.atranscription(
@ -225,6 +228,7 @@ async def test_bedrock_transcription_errors_never_fall_back(
extra_headers=None,
optional_params={},
timeout=None,
python_fallback=None,
)
@ -239,29 +243,29 @@ async def test_python_transcription_skips_rust_when_enabled() -> None:
pytest.fail("Python provider must not call Rust")
rust_bridge.configure_rust_transcription(transcription=unexpected, atranscription=aunexpected)
assert (
rust_bridge.transcription(
model="model",
audio={},
api_key=None,
api_base=None,
custom_llm_provider="openai",
extra_headers=None,
optional_params={},
timeout=None,
)
is None
)
assert (
await rust_bridge.atranscription(
model="model",
audio={},
api_key=None,
api_base=None,
custom_llm_provider="openai",
extra_headers=None,
optional_params={},
timeout=None,
)
is None
)
assert rust_bridge.transcription(
model="model",
audio={},
api_key=None,
api_base=None,
custom_llm_provider="openai",
extra_headers=None,
optional_params={},
timeout=None,
python_fallback=lambda: {"text": "python"},
) == {"text": "python"}
async def python_fallback() -> dict[str, object]:
return {"text": "python"}
assert await rust_bridge.atranscription(
model="model",
audio={},
api_key=None,
api_base=None,
custom_llm_provider="openai",
extra_headers=None,
optional_params={},
timeout=None,
python_fallback=python_fallback,
) == {"text": "python"}

View file

@ -8,7 +8,7 @@ from typing import Final
import pytest
import litellm
from litellm.rust_bridge import ocr as rust_ocr_bridge
from litellm.rust_bridge import _native
pytestmark = pytest.mark.requires_rust_extension
@ -79,15 +79,15 @@ def test_native_ocr_with_compiled_rust_extension(
host: Final = str(address[0])
port: Final = int(address[1])
response: Final = rust_ocr_bridge.ocr(
model="mistral-ocr-latest",
document={"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
response: Final = _native.ocr(
"mistral-ocr-latest",
{"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
api_key="test-key",
api_base=f"http://{host}:{port}",
custom_llm_provider="mistral",
extra_headers=None,
optional_params={},
timeout=None,
timeout_seconds=None,
)
assert response is not None

View file

@ -8,7 +8,7 @@ from litellm.rust_bridge import _native
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.catalog import NATIVE_EXPORTS
from litellm.rust_bridge.chat_completions.lifecycle import LIFECYCLE as CHAT_COMPLETIONS
from litellm.rust_bridge.configuration import ComponentName, ExecutionDecision
from litellm.rust_bridge.configuration import ExecutionDecision, ComponentName
from litellm.rust_bridge.embeddings.lifecycle import LIFECYCLE as EMBEDDINGS
from litellm.rust_bridge.image_edit.lifecycle import LIFECYCLE as IMAGE_EDIT
from litellm.rust_bridge.image_generation.lifecycle import LIFECYCLE as IMAGE_GENERATION