diff --git a/litellm-rust/crates/core/src/audio_transcription/mod.rs b/litellm-rust/crates/core/src/audio_transcription/mod.rs index 795d0a72f56..84042a3629d 100644 --- a/litellm-rust/crates/core/src/audio_transcription/mod.rs +++ b/litellm-rust/crates/core/src/audio_transcription/mod.rs @@ -17,7 +17,6 @@ pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Resu .await } - pub fn admit( model: &str, provider: Option<&str>, diff --git a/litellm-rust/crates/core/src/chat_completions/mod.rs b/litellm-rust/crates/core/src/chat_completions/mod.rs index 7441f09d367..6ac42928dd2 100644 --- a/litellm-rust/crates/core/src/chat_completions/mod.rs +++ b/litellm-rust/crates/core/src/chat_completions/mod.rs @@ -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 { diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index 4af86ad43bb..02106d4d228 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -27,7 +27,6 @@ pub async fn messages_stream(request: MessagesRequest<'_>) -> Result, diff --git a/litellm-rust/crates/python-bridge/src/errors.rs b/litellm-rust/crates/python-bridge/src/errors.rs index 216838dde6c..58434eb423c 100644 --- a/litellm-rust/crates/python-bridge/src/errors.rs +++ b/litellm-rust/crates/python-bridge/src/errors.rs @@ -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, diff --git a/litellm-rust/crates/python-bridge/src/routes/definition.rs b/litellm-rust/crates/python-bridge/src/routes/definition.rs index 2a8e4d62668..01517919d39 100644 --- a/litellm-rust/crates/python-bridge/src/routes/definition.rs +++ b/litellm-rust/crates/python-bridge/src/routes/definition.rs @@ -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", diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/value.rs b/litellm-rust/crates/python-bridge/src/routes/messages/value.rs index a436905d171..4491b523d44 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/value.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/value.rs @@ -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, timeout_seconds: Option, has_agentic_hook: Option, + on_request: Option>, }, prepare = prepare_messages, errors = execution_error_to_pyerr, diff --git a/litellm-rust/crates/python-bridge/src/routes/token_counter/mod.rs b/litellm-rust/crates/python-bridge/src/routes/token_counter/mod.rs index 2a5808fecee..bd0839e7d53 100644 --- a/litellm-rust/crates/python-bridge/src/routes/token_counter/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/token_counter/mod.rs @@ -34,8 +34,9 @@ fn count_input_tokens<'py>( legacy_accounting: bool, resource_loader: Py, ) -> PyResult> { - 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 diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index b9e3a6800ab..dd6fec2fdef 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -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): diff --git a/litellm/llms/bedrock/audio_transcription/__init__.py b/litellm/llms/bedrock/audio_transcription/__init__.py index 90297b3b492..d7b328ae421 100644 --- a/litellm/llms/bedrock/audio_transcription/__init__.py +++ b/litellm/llms/bedrock/audio_transcription/__init__.py @@ -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) diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 4928a29306c..335a7ff3a1e 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -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 diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index ee468535c2a..9356dccce83 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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]] = {} diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 2548bf10a14..2f9d1830088 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -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 "", diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 6074a50a69b..5ea64e3a2bc 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -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( diff --git a/litellm/rust_bridge/README.md b/litellm/rust_bridge/README.md index bc40449dc34..2ad662e7560 100644 --- a/litellm/rust_bridge/README.md +++ b/litellm/rust_bridge/README.md @@ -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 diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index 7912d7c1a11..ab498fb10b0 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -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", ] diff --git a/litellm/rust_bridge/catalog.py b/litellm/rust_bridge/catalog.py index eb309245115..a3ce25ccfdb 100644 --- a/litellm/rust_bridge/catalog.py +++ b/litellm/rust_bridge/catalog.py @@ -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, diff --git a/litellm/rust_bridge/chat_completions/value.py b/litellm/rust_bridge/chat_completions/value.py index 710abb93353..b92efc20ea0 100644 --- a/litellm/rust_bridge/chat_completions/value.py +++ b/litellm/rust_bridge/chat_completions/value.py @@ -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), diff --git a/litellm/rust_bridge/messages/types.py b/litellm/rust_bridge/messages/types.py index 47c51a01382..adfb2810bc7 100644 --- a/litellm/rust_bridge/messages/types.py +++ b/litellm/rust_bridge/messages/types.py @@ -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 diff --git a/litellm/rust_bridge/messages/value.py b/litellm/rust_bridge/messages/value.py index 2542995a27c..6067c22d6e8 100644 --- a/litellm/rust_bridge/messages/value.py +++ b/litellm/rust_bridge/messages/value.py @@ -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 "", diff --git a/litellm/rust_bridge/ocr/__init__.py b/litellm/rust_bridge/ocr/__init__.py index 5140e733266..40b4b27bbef 100644 --- a/litellm/rust_bridge/ocr/__init__.py +++ b/litellm/rust_bridge/ocr/__init__.py @@ -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", ) diff --git a/litellm/rust_bridge/ocr/host.py b/litellm/rust_bridge/ocr/host.py index 9f4489b2d0e..b1670b8130a 100644 --- a/litellm/rust_bridge/ocr/host.py +++ b/litellm/rust_bridge/ocr/host.py @@ -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}/"), diff --git a/litellm/rust_bridge/ocr/types.py b/litellm/rust_bridge/ocr/types.py index 65a273506c6..14023a75369 100644 --- a/litellm/rust_bridge/ocr/types.py +++ b/litellm/rust_bridge/ocr/types.py @@ -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 diff --git a/litellm/rust_bridge/ocr/value.py b/litellm/rust_bridge/ocr/value.py index ae952014894..ed8c6451d94 100644 --- a/litellm/rust_bridge/ocr/value.py +++ b/litellm/rust_bridge/ocr/value.py @@ -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), - ) diff --git a/litellm/rust_bridge/responses/websocket.py b/litellm/rust_bridge/responses/websocket.py index 75bb4d10a6a..e3e6c0c0285 100644 --- a/litellm/rust_bridge/responses/websocket.py +++ b/litellm/rust_bridge/responses/websocket.py @@ -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) diff --git a/litellm/rust_bridge/route.py b/litellm/rust_bridge/route.py index f3e932b6a3f..df214e73bf1 100644 --- a/litellm/rust_bridge/route.py +++ b/litellm/rust_bridge/route.py @@ -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( diff --git a/litellm/rust_bridge/runtime.py b/litellm/rust_bridge/runtime.py index 5cd8fa0ca56..2ac704519da 100644 --- a/litellm/rust_bridge/runtime.py +++ b/litellm/rust_bridge/runtime.py @@ -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: diff --git a/litellm/rust_bridge/token_counter/value.py b/litellm/rust_bridge/token_counter/value.py index 391204cdc06..c1851399da2 100644 --- a/litellm/rust_bridge/token_counter/value.py +++ b/litellm/rust_bridge/token_counter/value.py @@ -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=""), ) diff --git a/litellm/rust_bridge/transcription/value.py b/litellm/rust_bridge/transcription/value.py index 072e69f1978..b410681fce7 100644 --- a/litellm/rust_bridge/transcription/value.py +++ b/litellm/rust_bridge/transcription/value.py @@ -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), ) diff --git a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py index 8bf10fa9414..f6e7f8f5716 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -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 diff --git a/tests/test_litellm/ocr/test_ocr_native_format.py b/tests/test_litellm/ocr/test_ocr_native_format.py index 997f65874a0..568e5712047 100644 --- a/tests/test_litellm/ocr/test_ocr_native_format.py +++ b/tests/test_litellm/ocr/test_ocr_native_format.py @@ -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", diff --git a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py index 19bd16962cf..feedd084290 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py +++ b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py @@ -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() diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py index 9f07d542b72..a45b7fbc326 100644 --- a/tests/test_litellm/responses/test_rust_bridge_websocket.py +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -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, ) diff --git a/tests/test_litellm/rust_bridge/test_configuration_env.py b/tests/test_litellm/rust_bridge/test_configuration_env.py index da2db6a5bdb..6be854004a1 100644 --- a/tests/test_litellm/rust_bridge/test_configuration_env.py +++ b/tests/test_litellm/rust_bridge/test_configuration_env.py @@ -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))) diff --git a/tests/test_litellm/rust_bridge/test_route.py b/tests/test_litellm/rust_bridge/test_route.py index 062fbc8cbf6..c9bad451e05 100644 --- a/tests/test_litellm/rust_bridge/test_route.py +++ b/tests/test_litellm/rust_bridge/test_route.py @@ -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 diff --git a/tests/test_litellm/rust_bridge/test_runtime.py b/tests/test_litellm/rust_bridge/test_runtime.py index 76a0d3c25b2..570926bb473 100644 --- a/tests/test_litellm/rust_bridge/test_runtime.py +++ b/tests/test_litellm/rust_bridge/test_runtime.py @@ -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: diff --git a/tests/test_litellm/rust_bridge/test_token_counter.py b/tests/test_litellm/rust_bridge/test_token_counter.py index c32238cec22..80859bc28d4 100644 --- a/tests/test_litellm/rust_bridge/test_token_counter.py +++ b/tests/test_litellm/rust_bridge/test_token_counter.py @@ -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] diff --git a/tests/test_litellm/test_audio_transcription_rust_bridge.py b/tests/test_litellm/test_audio_transcription_rust_bridge.py index 7cd92107ae9..98b5a5eb9ed 100644 --- a/tests/test_litellm/test_audio_transcription_rust_bridge.py +++ b/tests/test_litellm/test_audio_transcription_rust_bridge.py @@ -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"} diff --git a/tests/test_litellm_rust/test_ocr.py b/tests/test_litellm_rust/test_ocr.py index e0e06d685b8..f9b11c3c87b 100644 --- a/tests/test_litellm_rust/test_ocr.py +++ b/tests/test_litellm_rust/test_ocr.py @@ -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 diff --git a/tests/test_litellm_rust/test_route_foundation.py b/tests/test_litellm_rust/test_route_foundation.py index 900c0fa3aee..cc77bf808c1 100644 --- a/tests/test_litellm_rust/test_route_foundation.py +++ b/tests/test_litellm_rust/test_route_foundation.py @@ -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