fix(native): make OCR admission provider aware

This commit is contained in:
Yujong Lee 2026-09-05 23:19:19 -07:00
parent e386004302
commit 152d041736
20 changed files with 270 additions and 191 deletions

View file

@ -35,8 +35,10 @@ pub async fn ocr(
.await
}
pub fn ocr_provider_supported(model: &str, provider: &str) -> bool {
common_utils::ocr_provider_config(provider, model).is_some()
pub fn ocr_provider_supported(model: &str, provider: &str, request_format: Option<&str>) -> bool {
common_utils::ocr_provider_config(provider, model).is_some_and(|config| {
request_format != Some("native") || config.supported_ocr_params().contains(&"req_format")
})
}
#[cfg(test)]

View file

@ -6,7 +6,6 @@ pub enum NativeRouteDecline {
Streaming,
AgenticHook,
CustomClient,
NativeResponseFormat,
}
impl NativeRouteDecline {
@ -16,7 +15,6 @@ impl NativeRouteDecline {
Self::Streaming => "native streaming is unavailable",
Self::AgenticHook => "native agentic hooks are unavailable",
Self::CustomClient => "native custom clients are unavailable",
Self::NativeResponseFormat => "native OCR response format is unavailable",
}
}
}
@ -37,8 +35,7 @@ pub fn native_route_decline(
if capabilities.has_custom_client {
return Some(NativeRouteDecline::CustomClient);
}
(capabilities.request_format.as_deref() == Some("native"))
.then_some(NativeRouteDecline::NativeResponseFormat)
None
}
#[cfg(test)]
@ -52,6 +49,7 @@ mod tests {
has_agentic_hook: true,
has_custom_client: true,
request_format: Some("native".into()),
..Default::default()
};
assert_eq!(
native_route_decline(false, &all_unsupported),
@ -84,13 +82,6 @@ mod tests {
},
NativeRouteDecline::CustomClient,
),
(
RequestCapabilities {
request_format: Some("native".into()),
..Default::default()
},
NativeRouteDecline::NativeResponseFormat,
),
];
for (capabilities, expected) in cases {
assert_eq!(native_route_decline(true, &capabilities), Some(expected));

View file

@ -332,6 +332,11 @@ litellm_context = replace(
)
assert routes.ocr_decline('model', 'mistral', context=native_context) is not None
assert routes.ocr_decline('model', 'mistral', context=litellm_context) is None
assert routes.ocr_decline(
'doc-intelligence/prebuilt-layout',
'azure_ai',
context=native_context,
) is None
",
Some(&locals),
Some(&locals),

View file

@ -23,11 +23,12 @@ fn prepare_ocr(
options: NativeRequestOptions,
context: NativeRequestContext,
) -> PyResult<impl Future<Output = Result<Value, Error>> + Send + 'static> {
let context: LiteLlmRequestContext = context.into();
let provider_supported = litellm_ai_gateway::io::ocr::ocr_provider_supported(
&input.model,
options.provider("mistral"),
context.capabilities.request_format.as_deref(),
);
let context: LiteLlmRequestContext = context.into();
if let Some(reason) = super::definition::request_decline(provider_supported, &context) {
return Err(crate::errors::RustBridgeDeclined::new_err(reason));
}
@ -58,10 +59,12 @@ fn ocr_decline(
context: NativeRequestContext,
) -> Option<String> {
let context: LiteLlmRequestContext = context.into();
super::definition::request_decline(
litellm_ai_gateway::io::ocr::ocr_provider_supported(model, custom_llm_provider),
&context,
)
let provider_supported = litellm_ai_gateway::io::ocr::ocr_provider_supported(
model,
custom_llm_provider,
context.capabilities.request_format.as_deref(),
);
super::definition::request_decline(provider_supported, &context)
}
bridge_route! {

View file

@ -482,8 +482,6 @@ class AnthropicChatCompletion(BaseLLM):
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(

View file

@ -442,12 +442,10 @@ class BedrockConverseLLM(BaseAWSLLM):
if client is None or isinstance(client, AsyncHTTPHandler):
_params: Final = {}
if timeout is not None:
if isinstance(timeout, float) or isinstance(timeout, int):
if isinstance(timeout, (float, 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(

View file

@ -6397,6 +6397,7 @@ class BaseLLMHTTPHandler:
},
)
from litellm.rust_bridge.request import request_context
from litellm.rust_bridge.responses_websocket import open_connection
async with open_connection(
@ -6405,6 +6406,11 @@ class BaseLLMHTTPHandler:
timeout=timeout,
model=model,
provider=custom_llm_provider,
context=request_context(
logging_obj=logging_obj,
request_model=logging_obj.model,
litellm_params=litellm_params.model_dump(),
),
fallback=lambda: websockets.connect(
ws_url,
additional_headers=headers,

View file

@ -7851,11 +7851,7 @@ def transcription(
litellm_params=litellm_params_dict,
model_response=model_response,
atranscription=atranscription,
client=(
client
if client is not None and (isinstance(client, HTTPHandler) or isinstance(client, AsyncHTTPHandler))
else None
),
client=(client if client is not None and isinstance(client, (HTTPHandler, AsyncHTTPHandler)) else None),
timeout=timeout,
max_retries=max_retries,
logging_obj=litellm_logging_obj,
@ -7872,11 +7868,7 @@ def transcription(
litellm_params=litellm_params_dict,
model_response=model_response,
atranscription=atranscription,
client=(
client
if client is not None and (isinstance(client, HTTPHandler) or isinstance(client, AsyncHTTPHandler))
else None
),
client=(client if client is not None and isinstance(client, (HTTPHandler, AsyncHTTPHandler)) else None),
timeout=timeout,
max_retries=max_retries,
logging_obj=litellm_logging_obj,

View file

@ -31,8 +31,10 @@ from litellm.llms.base_llm.ocr.transformation import (
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.rust_bridge import ocr as rust_ocr_bridge
from litellm.rust_bridge.request import (
NativeRequestCapabilities,
NativeRequestOptions,
PreparedNativeCall,
request_context,
vertex_options,
)
from litellm.rust_bridge.timeouts import timeout_to_seconds
@ -57,6 +59,7 @@ class _PreparedOCRRequest:
litellm_params: dict[str, object]
effective_timeout: float | httpx.Timeout
litellm_logging_obj: LiteLLMLoggingObj
execution_mode: str = "sync"
def _prepare_ocr_request(
@ -68,6 +71,7 @@ def _prepare_ocr_request(
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
kwargs: dict[str, object],
execution_mode: str = "sync",
) -> _PreparedOCRRequest:
litellm_logging_obj: Final = cast(LiteLLMLoggingObj, kwargs.pop("litellm_logging_obj"))
litellm_call_id: Final = cast(str | None, kwargs.get("litellm_call_id", None))
@ -177,6 +181,7 @@ def _prepare_ocr_request(
litellm_params=dict(litellm_params),
effective_timeout=effective_timeout,
litellm_logging_obj=litellm_logging_obj,
execution_mode=execution_mode,
)
@ -272,6 +277,21 @@ def _prepare_rust_ocr_call(
),
timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout),
),
context=request_context(
logging_obj=prepared_request.litellm_logging_obj,
request_model=getattr(prepared_request.litellm_logging_obj, "model", prepared_request.model),
litellm_params=prepared_request.litellm_params,
capabilities=NativeRequestCapabilities(
execution_mode=prepared_request.execution_mode,
input_source_kind=str(prepared_request.document.get("type") or "unknown"),
request_format=(
value
if isinstance((value := prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM)), str)
else None
),
native_response_format=(prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native"),
),
),
)
@ -422,6 +442,7 @@ async def aocr(
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
kwargs=kwargs,
execution_mode="async",
)
model = prepared.model
custom_llm_provider = prepared.custom_llm_provider

View file

@ -47,6 +47,8 @@ from litellm.rust_bridge.request import (
anthropic_options,
bedrock_options,
call_native,
request_context,
with_capabilities,
)
from litellm.rust_bridge.runtime import (
BridgeErrorContext,
@ -166,12 +168,14 @@ def _provider_eligibility_options(
def _eligibility_context(
*,
execution_mode: str | None = None,
stream: bool,
has_custom_client: bool = False,
has_agentic_hook: bool = False,
) -> NativeRequestContext:
return NativeRequestContext(
capabilities=NativeRequestCapabilities(
execution_mode=execution_mode,
stream=stream,
has_custom_client=has_custom_client,
has_agentic_hook=has_agentic_hook,
@ -179,6 +183,11 @@ def _eligibility_context(
)
def _execution_context(context: NativeRequestContext | None, mode: str) -> NativeRequestContext:
current = context or NativeRequestContext()
return with_capabilities(current, replace(current.capabilities, execution_mode=mode))
def rust_chat_completions_accepts(
*,
model: str,
@ -235,6 +244,7 @@ def chat_completions(
on_response: ResponseObserver,
bedrock: NativeBedrockOptions | None = None,
anthropic: NativeAnthropicOptions | None = None,
context: NativeRequestContext | None = None,
) -> ModelResponse | None:
def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
on_response(rust_response)
@ -256,7 +266,7 @@ def chat_completions(
bedrock=bedrock,
anthropic=anthropic,
),
context=NativeRequestContext(),
context=_execution_context(context, "sync"),
),
call=call_native,
fallback=lambda: None,
@ -279,6 +289,7 @@ async def achat_completions(
on_response: ResponseObserver,
bedrock: NativeBedrockOptions | None = None,
anthropic: NativeAnthropicOptions | None = None,
context: NativeRequestContext | None = None,
) -> ModelResponse | None:
def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
on_response(rust_response)
@ -300,7 +311,7 @@ async def achat_completions(
bedrock=bedrock,
anthropic=anthropic,
),
context=NativeRequestContext(),
context=_execution_context(context, "async"),
),
call=call_native,
fallback=async_none,
@ -324,6 +335,7 @@ async def achat_completions_or_fallback(
python_fallback: Callable[[], Awaitable[object]],
bedrock: NativeBedrockOptions | None = None,
anthropic: NativeAnthropicOptions | None = None,
context: NativeRequestContext | None = None,
) -> object:
"""Await the Rust path, falling back to the caller's own Python path when
the bridge is unavailable or the call fails.
@ -354,7 +366,7 @@ async def achat_completions_or_fallback(
bedrock=bedrock,
anthropic=anthropic,
),
context=NativeRequestContext(),
context=_execution_context(context, "async"),
),
call=call_native,
fallback=python_fallback,
@ -383,6 +395,7 @@ class _ChatOperation:
custom_llm_provider=ctx.custom_llm_provider,
options=_provider_eligibility_options(ctx.custom_llm_provider, ctx.litellm_params, ctx.optional_params),
context=_eligibility_context(
execution_mode="async" if ctx.acompletion else "sync",
stream=bool(ctx.stream),
has_custom_client=ctx.client is not None or ctx.shared_session is not None,
has_agentic_hook=BaseLLMHTTPHandler.has_agentic_completion_hook(ctx.logging),
@ -450,10 +463,16 @@ class _ChatOperation:
extra_headers=headers,
timeout_seconds=timeout_to_seconds(float(ctx.timeout) if isinstance(ctx.timeout, str) else ctx.timeout),
),
context=_eligibility_context(
stream=bool(ctx.stream),
has_custom_client=ctx.client is not None or ctx.shared_session is not None,
has_agentic_hook=BaseLLMHTTPHandler.has_agentic_completion_hook(ctx.logging),
context=request_context(
logging_obj=ctx.logging,
request_model=ctx.logging.model,
litellm_params=ctx.litellm_params,
capabilities=NativeRequestCapabilities(
execution_mode="async" if ctx.acompletion else "sync",
stream=bool(ctx.stream),
has_custom_client=ctx.client is not None or ctx.shared_session is not None,
has_agentic_hook=BaseLLMHTTPHandler.has_agentic_completion_hook(ctx.logging),
),
),
)

View file

@ -25,10 +25,13 @@ from litellm.rust_bridge.protocols import RustAmessages, RustMessages, RustRoute
from litellm.rust_bridge.request import (
NativeMessagesRequest,
NativePreCallDetails,
NativeRequestCapabilities,
NativeRequestContext,
NativeRequestOptions,
PreparedNativeCall,
call_native,
request_context,
with_capabilities,
)
from litellm.rust_bridge.runtime import (
BridgeErrorContext,
@ -101,6 +104,10 @@ def messages(
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
timeout: float | httpx.Timeout | None,
stream: bool = False,
has_custom_client: bool = False,
has_agentic_hook: bool = False,
context: NativeRequestContext | None = None,
) -> dict[str, object] | None:
return _MESSAGES.invoke(
prepare=lambda: PreparedNativeCall(
@ -115,7 +122,15 @@ def messages(
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
),
context=NativeRequestContext(),
context=with_capabilities(
context or NativeRequestContext(),
NativeRequestCapabilities(
execution_mode="sync",
stream=stream,
has_custom_client=has_custom_client,
has_agentic_hook=has_agentic_hook,
),
),
),
call=call_native,
preflight=lambda: assess_route(_PREFLIGHT, model, custom_llm_provider or ""),
@ -134,6 +149,10 @@ async def amessages(
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
timeout: float | httpx.Timeout | None,
stream: bool = False,
has_custom_client: bool = False,
has_agentic_hook: bool = False,
context: NativeRequestContext | None = None,
) -> dict[str, object] | None:
return await _MESSAGES.ainvoke(
prepare=lambda: PreparedNativeCall(
@ -148,7 +167,15 @@ async def amessages(
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
),
context=NativeRequestContext(),
context=with_capabilities(
context or NativeRequestContext(),
NativeRequestCapabilities(
execution_mode="async",
stream=stream,
has_custom_client=has_custom_client,
has_agentic_hook=has_agentic_hook,
),
),
),
call=call_native,
preflight=lambda: assess_route(_PREFLIGHT, model, custom_llm_provider or ""),
@ -190,6 +217,9 @@ class _MessagesOperation:
api_key: str | None
api_base: str | None
python: Callable[[], MessagesResult]
stream: bool = False
asynchronous: bool = False
has_custom_client: bool = False
logged: bool = False
def prepare(self) -> PreparedNativeCall[NativeMessagesRequest]:
@ -275,7 +305,17 @@ class _MessagesOperation:
BaseLLMHTTPHandler.resolve_anthropic_messages_timeout(self.params, False, self.provider)
),
),
context=NativeRequestContext(),
context=request_context(
logging_obj=self.logging,
request_model=self.logging.model if self.logging is not None else self.model,
litellm_params=self.params.model_dump(),
capabilities=NativeRequestCapabilities(
execution_mode="async" if self.asynchronous else "sync",
stream=self.stream,
has_custom_client=self.has_custom_client,
has_agentic_hook=BaseLLMHTTPHandler.has_agentic_completion_hook(self.logging),
),
),
)
def fallback(self) -> MessagesResult:
@ -313,7 +353,20 @@ def dispatch_messages(
has_custom_client: bool,
fallback: Callable[[], MessagesResult],
) -> MessagesResult:
operation: Final = _MessagesOperation(model, provider, messages, body, params, logging, api_key, api_base, fallback)
operation: Final = _MessagesOperation(
model,
provider,
messages,
body,
params,
logging,
api_key,
api_base,
fallback,
stream,
asynchronous,
has_custom_client,
)
def preflight() -> PythonFallback | None:
return assess_route(

View file

@ -6,7 +6,6 @@ from collections.abc import Awaitable, Callable, Mapping
from typing import Final, TypeVar
from . import configuration as _configuration
from .bindings import UNCHANGED, Unchanged
from .protocols import RustAocr, RustOcr, RustRouteDecline
from .request import NativeOCRRequest, PreparedNativeCall, call_native
from .runtime import (
@ -16,8 +15,6 @@ from .runtime import (
assess_route,
)
rust_ocr_enabled = _configuration.rust_ocr_enabled
rust = _configuration.rust
ResultT = TypeVar("ResultT")
@ -25,40 +22,17 @@ _OCR: Final[EndpointDispatch[RustOcr, RustAocr]] = EndpointDispatch.native(
route="ocr",
sync=lambda native: native.ocr,
asynchronous=lambda native: native.aocr,
enabled=_configuration.rust_ocr_enabled,
enabled=_configuration.rust_enabled,
)
_PREFLIGHT: Final[EndpointBinding[RustRouteDecline]] = EndpointBinding.native(
route="ocr",
select=lambda native: native.ocr_decline,
enabled=_configuration.rust_ocr_enabled,
enabled=_configuration.rust_enabled,
)
def set_rust_ocr(
*,
ocr: RustOcr | None | Unchanged = UNCHANGED,
aocr: RustAocr | None | Unchanged = UNCHANGED,
decline: RustRouteDecline | None | Unchanged = UNCHANGED,
) -> None:
if not isinstance(decline, Unchanged):
if decline is None:
_PREFLIGHT.reset()
else:
_PREFLIGHT.override(decline)
if not isinstance(ocr, Unchanged):
if ocr is None:
_OCR.sync.reset()
else:
_OCR.sync.override(ocr)
if not isinstance(aocr, Unchanged):
if aocr is None:
_OCR.asynchronous.reset()
else:
_OCR.asynchronous.override(aocr)
def load_rust_ocr() -> RustOcr | None:
return _OCR.sync.load()

View file

@ -17,11 +17,13 @@ from litellm.rust_bridge.protocols import (
RustRouteDecline,
)
from litellm.rust_bridge.request import (
NativeRequestCapabilities,
NativeRequestContext,
NativeRequestOptions,
NativeResponsesWebSocketRequest,
PreparedNativeCall,
call_native,
with_capabilities,
)
from litellm.rust_bridge.runtime import (
BridgeErrorContext,
@ -93,6 +95,7 @@ async def connect(
model: str = "responses websocket",
provider: str = "openai",
fallback: Callable[[], Awaitable[Connection | None]] = async_none,
context: NativeRequestContext | None = None,
) -> Connection | None:
return await _RESPONSES_WEBSOCKET.ainvoke(
prepare=lambda: PreparedNativeCall(
@ -104,7 +107,14 @@ async def connect(
timeout_seconds=timeout_to_seconds(timeout),
custom_llm_provider=provider,
),
context=NativeRequestContext(),
context=with_capabilities(
context or NativeRequestContext(),
NativeRequestCapabilities(
execution_mode="async",
websocket_mode="native",
requires_connection=True,
),
),
),
call=lambda connection_type, request: call_native(connection_type.connect, request),
preflight=lambda: assess_route(_PREFLIGHT, model, provider),
@ -123,6 +133,7 @@ async def open_connection(
model: str,
provider: str,
fallback: Callable[[], AbstractAsyncContextManager[Connection]],
context: NativeRequestContext | None = None,
) -> AsyncGenerator[Connection]:
async with AsyncExitStack() as stack:
@ -136,6 +147,7 @@ async def open_connection(
model=model,
provider=provider,
fallback=python_connection,
context=context,
)
if backend is None:
raise RuntimeError("WebSocket connection returned no connection")

View file

@ -3,7 +3,7 @@ from __future__ import annotations
from collections.abc import Awaitable, Callable
from dataclasses import dataclass, field
from enum import Enum
from typing import Final, Generic, NoReturn, Protocol, TypeAlias, TypeVar
from typing import Final, Generic, NoReturn, Protocol, TypeAlias, TypeVar, assert_never
from litellm.exceptions import APIError, AuthenticationError, InternalServerError, RateLimitError
from litellm.rust_bridge.bindings import (
@ -180,6 +180,8 @@ class EndpointBinding(Generic[BindingT]):
return value
case PythonFallback():
return fallback()
case _ as unreachable:
assert_never(unreachable)
async def ainvoke(
self,
@ -205,6 +207,8 @@ class EndpointBinding(Generic[BindingT]):
return value
case PythonFallback():
return await fallback()
case _ as unreachable:
assert_never(unreachable)
def assess(
self,
@ -257,6 +261,8 @@ class EndpointBinding(Generic[BindingT]):
return value
case PythonFallback():
self._raise_required(result)
case _ as unreachable:
assert_never(unreachable)
async def arequire(
self,
@ -281,6 +287,8 @@ class EndpointBinding(Generic[BindingT]):
return value
case PythonFallback():
self._raise_required(result)
case _ as unreachable:
assert_never(unreachable)
def can_attempt(
self,

View file

@ -20,12 +20,15 @@ from litellm.rust_bridge.configuration import rust_enabled
from litellm.rust_bridge.protocols import RustAtranscription, RustRouteDecline, RustTranscription
from litellm.rust_bridge.request import (
NativePreCallDetails,
NativeRequestCapabilities,
NativeRequestContext,
NativeRequestOptions,
NativeTranscriptionRequest,
PreparedNativeCall,
bedrock_options,
call_native,
request_context,
with_capabilities,
)
from litellm.rust_bridge.runtime import (
BridgeErrorContext,
@ -90,13 +93,17 @@ def load_rust_atranscription() -> RustAtranscription | None:
def transcription(
*,
model: str,
audio: dict[str, object],
audio: 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,
stream: bool = False,
has_custom_client: bool = False,
input_source_kind: str | None = None,
context: NativeRequestContext | None = None,
) -> dict[str, object] | None:
return _TRANSCRIPTION.invoke(
prepare=lambda: PreparedNativeCall(
@ -113,7 +120,15 @@ def transcription(
timeout_seconds=timeout_to_seconds(timeout),
bedrock=bedrock_options(optional_params),
),
context=NativeRequestContext(),
context=with_capabilities(
context or NativeRequestContext(),
NativeRequestCapabilities(
execution_mode="sync",
stream=stream,
has_custom_client=has_custom_client,
input_source_kind=input_source_kind,
),
),
),
call=call_native,
preflight=lambda: assess_route(_PREFLIGHT, model, custom_llm_provider or ""),
@ -126,13 +141,17 @@ def transcription(
async def atranscription(
*,
model: str,
audio: dict[str, object],
audio: 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,
stream: bool = False,
has_custom_client: bool = False,
input_source_kind: str | None = None,
context: NativeRequestContext | None = None,
) -> dict[str, object] | None:
return await _TRANSCRIPTION.ainvoke(
prepare=lambda: PreparedNativeCall(
@ -149,7 +168,15 @@ async def atranscription(
timeout_seconds=timeout_to_seconds(timeout),
bedrock=bedrock_options(optional_params),
),
context=NativeRequestContext(),
context=with_capabilities(
context or NativeRequestContext(),
NativeRequestCapabilities(
execution_mode="async",
stream=stream,
has_custom_client=has_custom_client,
input_source_kind=input_source_kind,
),
),
),
call=call_native,
preflight=lambda: assess_route(_PREFLIGHT, model, custom_llm_provider or ""),
@ -162,6 +189,17 @@ async def atranscription(
TranscriptionResult = TranscriptionResponse | Coroutine[object, object, TranscriptionResponse]
def _input_source_kind(file: FileTypes) -> str:
content: Final = file[1] if isinstance(file, tuple) else file
if isinstance(content, (bytes, bytearray, memoryview)):
return "bytes"
if isinstance(content, IOBase):
return "file"
if isinstance(content, str):
return "path"
return "opaque"
@dataclass
class _TranscriptionOperation:
model: str
@ -174,6 +212,8 @@ class _TranscriptionOperation:
timeout: float | httpx.Timeout | None
logging: Logging
python: Callable[[FileTypes], TranscriptionResult]
asynchronous: bool = False
has_custom_client: bool = False
fallback_file: FileTypes | None = None
logged: bool = False
@ -228,7 +268,17 @@ class _TranscriptionOperation:
timeout_seconds=timeout_to_seconds(self.timeout),
bedrock=bedrock_options(self.optional_params),
),
context=NativeRequestContext(),
context=request_context(
logging_obj=self.logging,
request_model=self.logging.model,
litellm_params=self.logging.litellm_params,
capabilities=NativeRequestCapabilities(
execution_mode="async" if self.asynchronous else "sync",
stream=self.optional_params.get("stream") is True,
has_custom_client=self.has_custom_client,
input_source_kind=_input_source_kind(self.file),
),
),
)
def fallback(self) -> TranscriptionResult:
@ -265,7 +315,18 @@ def dispatch_transcription(
fallback: Callable[[FileTypes], TranscriptionResult],
) -> TranscriptionResult:
operation: Final = _TranscriptionOperation(
model, provider, file, api_key, api_base, headers, optional_params, timeout, logging, fallback
model,
provider,
file,
api_key,
api_base,
headers,
optional_params,
timeout,
logging,
fallback,
asynchronous,
has_custom_client,
)
def preflight() -> PythonFallback | None:

View file

@ -128,16 +128,6 @@ def test_load_rust_messages_returns_injected_impl():
assert rust_messages.load_rust_messages() is bridge
def test_bare_rust_still_toggles_ocr():
from litellm.rust_bridge.ocr import rust_ocr_enabled
litellm.rust(True)
assert rust_ocr_enabled() is True
litellm.rust(False)
assert rust_ocr_enabled() is False
def test_load_rust_amessages_returns_injected_impl():
bridge = RecordingAsyncMessages()
litellm.rust(True)

View file

@ -10,9 +10,10 @@ import pytest
import litellm
from litellm._logging import verbose_logger
from litellm.integrations.code_interpreter_interception.handler import (
CodeInterpreterInterceptionLogger,
LITELLM_CODE_EXECUTION_TOOL_NAME,
CodeInterpreterInterceptionLogger,
)
from litellm.llms.azure.videos.transformation import AzureVideoConfig
from litellm.llms.base_llm.audio_transcription.transformation import (
AudioTranscriptionRequestData,
BaseAudioTranscriptionConfig,
@ -24,9 +25,7 @@ from litellm.llms.custom_httpx.llm_http_handler import (
_collect_ws_project_quota_callbacks,
_google_genai_streaming_hidden_params,
_has_pre_call_deployment_hook,
_rust_responses_websocket_enabled,
)
from litellm.llms.azure.videos.transformation import AzureVideoConfig
from litellm.llms.openai.videos.transformation import OpenAIVideoConfig
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.router import GenericLiteLLMParams
@ -1467,6 +1466,7 @@ def _make_responses_handler_call(signed_body):
signing provider (e.g. Bedrock Mantle).
"""
from unittest.mock import MagicMock
from litellm.llms.custom_httpx.http_handler import HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.types.router import GenericLiteLLMParams
@ -1522,6 +1522,7 @@ def test_responses_handler_signs_after_fake_stream_prep_strips_stream():
We snapshot request_data at sign time and assert "stream" is already gone.
"""
from unittest.mock import MagicMock
from litellm.llms.custom_httpx.http_handler import HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.types.llms.openai import ResponsesAPIResponse
@ -1585,6 +1586,7 @@ def _make_compact_handler_call(signed_body, is_async):
signing provider (e.g. Bedrock Mantle SigV4 / bearer).
"""
from unittest.mock import MagicMock
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.types.router import GenericLiteLLMParams
@ -2688,21 +2690,6 @@ async def test_generic_http_handler_async_streaming_forwards_provider_response_h
assert "".join([chunk.choices[0].delta.content or "" for chunk in collected]) == "hi"
@pytest.mark.parametrize(
"custom_llm_provider, enabled, expected",
[("openai", True, True), ("openai", False, False), ("azure", True, False),
("hosted_vllm", True, False), (None, True, False)],
)
def test_the_rust_responses_websocket_needs_openai_and_process_enablement(
custom_llm_provider, enabled, expected, monkeypatch
):
from litellm.rust_bridge import configuration
configuration.reset_rust_configuration()
monkeypatch.setenv("LITELLM_RUST", "1" if enabled else "0")
assert _rust_responses_websocket_enabled(custom_llm_provider) is expected
def test_a_plain_callback_does_not_advertise_a_pre_call_deployment_hook(monkeypatch):
from litellm.integrations.custom_logger import CustomLogger

View file

@ -12,10 +12,8 @@ from litellm.llms.azure_ai.ocr.common_utils import (
is_azure_document_intelligence_model,
)
from litellm.ocr.main import _prepare_ocr_request
from litellm.rust_bridge.ocr import _rust_bridge_api_base
_DOC = {"type": "document_url", "document_url": "https://example.com/doc.pdf"}
_DOC_INTELLIGENCE_ENDPOINT = "https://di.cognitiveservices.azure.com"
_AZURE_AI_API_BASE = "https://generic-azure-ai.example.com"
@ -24,13 +22,6 @@ class _FakeLogging:
return None
def _resolve_secret(name: str) -> str | None:
return {
"AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT": _DOC_INTELLIGENCE_ENDPOINT,
"AZURE_AI_API_BASE": _AZURE_AI_API_BASE,
}.get(name)
def _prepare(model: str, api_base: str | None):
return _prepare_ocr_request(
model=model,
@ -57,15 +48,13 @@ class TestIsAzureDocumentIntelligenceModel:
class TestDocIntelligenceApiBaseResolution:
def test_generic_azure_ai_base_does_not_hijack_doc_intelligence(self, monkeypatch):
"""Without an explicit api_base, the AZURE_AI_API_BASE fallback must not
overwrite the endpoint, so it resolves to the Document Intelligence one."""
"""The generic Azure AI fallback must stay out of native admission."""
monkeypatch.setenv("AZURE_AI_API_BASE", _AZURE_AI_API_BASE)
monkeypatch.delenv("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT", raising=False)
prepared = _prepare("azure_ai/doc-intelligence/prebuilt-layout", None)
assert prepared.api_base is None
assert _rust_bridge_api_base(prepared, _resolve_secret) == _DOC_INTELLIGENCE_ENDPOINT
def test_explicit_api_base_is_honoured_for_doc_intelligence(self, monkeypatch):
"""A caller-supplied api_base must always win, even for doc-intelligence."""
@ -75,7 +64,6 @@ class TestDocIntelligenceApiBaseResolution:
prepared = _prepare("azure_ai/doc-intelligence/prebuilt-layout", custom)
assert prepared.api_base == custom
assert _rust_bridge_api_base(prepared, _resolve_secret) == custom
def test_generic_azure_ai_base_still_applies_to_mistral_ocr(self, monkeypatch):
"""Non doc-intelligence azure_ai models keep using AZURE_AI_API_BASE."""

View file

@ -1,54 +1,14 @@
"""
Tests for the OCR `req_format` option in the SDK request path:
providers that don't support a native response must reject it, and the Rust
bridge (which only returns the normalized shape) must not serve native requests.
Tests for the OCR `req_format` option in the SDK request path.
"""
import dataclasses
from unittest.mock import MagicMock
import pytest
import litellm
from litellm.llms.azure_ai.ocr.cohere_parse_transformation import AzureAICohereParseConfig
from litellm.llms.cohere.ocr.transformation import CohereParseConfig
from litellm.rust_bridge.ocr import PreparedOCRRequest, _rust_ocr_supported
DOCUMENT = {"type": "document_url", "document_url": "https://example.com/doc.pdf"}
def _prepared(optional_params: dict[str, object]) -> PreparedOCRRequest:
return PreparedOCRRequest(
model="doc-intelligence/prebuilt-layout",
document=dict(DOCUMENT),
api_key="fake-key",
api_base="https://example.cognitiveservices.azure.com",
custom_llm_provider="azure_ai",
extra_headers=None,
provider_config=MagicMock(),
optional_params=optional_params,
litellm_params={},
effective_timeout=60.0,
litellm_logging_obj=MagicMock(),
)
@pytest.mark.parametrize("optional_params", [{}, {"req_format": "litellm"}])
def test_rust_ocr_serves_default_format(optional_params):
assert _rust_ocr_supported(_prepared(optional_params)) is True
def test_rust_ocr_skipped_for_native_format():
assert _rust_ocr_supported(_prepared({"req_format": "native"})) is False
@pytest.mark.parametrize("provider_config", [CohereParseConfig(), AzureAICohereParseConfig()])
def test_rust_ocr_skipped_for_configs_without_bridge_support(provider_config):
prepared = dataclasses.replace(_prepared({}), provider_config=provider_config)
assert _rust_ocr_supported(prepared) is False
@pytest.mark.asyncio
async def test_native_format_rejected_for_provider_without_support_as_bad_request():
with pytest.raises(litellm.BadRequestError, match="not supported for provider") as exc_info:

View file

@ -213,19 +213,26 @@ def build_prepared_request(
@pytest.fixture(autouse=True)
def _reset_rust_flag():
"""Keep the global toggle isolated between tests."""
rust_bridge.set_rust_ocr(ocr=None, aocr=None, decline=None)
rust_bridge._OCR.sync.reset()
rust_bridge._OCR.asynchronous.reset()
rust_bridge._PREFLIGHT.reset()
configuration.reset_rust_configuration()
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
rust_bridge.set_rust_ocr(
decline=lambda model, custom_llm_provider, *, context: (
rust_bridge._PREFLIGHT.override(
lambda model, custom_llm_provider, *, context: (
"unsupported feature"
if any(getattr(context.capabilities, key) for key in ("stream", "has_agentic_hook", "has_custom_client"))
or context.capabilities.request_format == "native"
or (
context.capabilities.request_format == "native"
and not (custom_llm_provider == "azure_ai" and "doc-intelligence" in model)
)
else None
)
)
yield
rust_bridge.set_rust_ocr(ocr=None, aocr=None, decline=None)
rust_bridge._OCR.sync.reset()
rust_bridge._OCR.asynchronous.reset()
rust_bridge._PREFLIGHT.reset()
configuration.reset_rust_configuration()
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
@ -235,7 +242,7 @@ def fake_bridge():
"""Enable the Rust path with an injected recording bridge (no native wheel)."""
bridge = RecordingBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.sync.override(bridge)
return bridge
@ -244,27 +251,27 @@ def fake_async_bridge():
"""Enable the async Rust path with an injected recording bridge."""
bridge = RecordingAsyncBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(aocr=bridge)
rust_bridge._OCR.asynchronous.override(bridge)
return bridge
def test_rust_toggles_flag():
assert rust_bridge.rust_ocr_enabled() is False
assert configuration.rust_enabled() is False
litellm.rust(True)
assert rust_bridge.rust_ocr_enabled() is True
assert configuration.rust_enabled() is True
litellm.rust(False)
assert rust_bridge.rust_ocr_enabled() is False
assert configuration.rust_enabled() is False
def test_env_var_enables_rust_ocr(monkeypatch):
monkeypatch.setenv("LITELLM_RUST", "1")
assert rust_bridge.rust_ocr_enabled() is True
assert configuration.rust_enabled() is True
def test_load_rust_ocr_returns_injected_impl():
bridge = RecordingBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.sync.override(bridge)
assert rust_bridge.load_rust_ocr() is bridge
@ -328,7 +335,7 @@ def test_native_bridge_available_reflects_loader(monkeypatch):
def test_load_rust_aocr_returns_injected_impl():
bridge = RecordingAsyncBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(aocr=bridge)
rust_bridge._OCR.asynchronous.override(bridge)
assert rust_bridge.load_rust_aocr() is bridge
@ -337,7 +344,8 @@ def test_toggle_without_ocr_arg_preserves_injected_impl():
bridge = RecordingBridge()
async_bridge = RecordingAsyncBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge, aocr=async_bridge)
rust_bridge._OCR.sync.override(bridge)
rust_bridge._OCR.asynchronous.override(async_bridge)
litellm.rust(False)
assert rust_bridge.load_rust_ocr() is bridge
@ -356,9 +364,12 @@ def test_explicit_ocr_none_clears_injected_impl(monkeypatch):
bridge = RecordingBridge()
async_bridge = RecordingAsyncBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge, aocr=async_bridge)
rust_bridge._OCR.sync.override(bridge)
rust_bridge._OCR.asynchronous.override(async_bridge)
rust_bridge.set_rust_ocr(ocr=None, aocr=None, decline=None)
rust_bridge._OCR.sync.reset()
rust_bridge._OCR.asynchronous.reset()
rust_bridge._PREFLIGHT.reset()
assert rust_bridge.load_rust_ocr() is None
assert rust_bridge.load_rust_aocr() is None
@ -405,7 +416,7 @@ def test_bridge_wrapper_forwards_prepared_args_and_wraps_response():
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.sync.override(bridge)
response = rust_bridge.dispatch_ocr(
prepare=lambda: PreparedNativeCall(
request=NativeOCRRequest(
@ -452,7 +463,7 @@ async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response():
litellm.rust(True)
rust_bridge.set_rust_ocr(aocr=bridge)
rust_bridge._OCR.asynchronous.override(bridge)
async def unexpected_fallback():
pytest.fail("unexpected Python fallback")
@ -495,7 +506,7 @@ def test_run_rust_ocr_prepares_request_and_wraps_response():
bridge = RecordingBridge()
logging_obj = RecordingLogging()
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.sync.override(bridge)
response = ocr_main._run_rust_ocr(
fallback=lambda: pytest.fail("unexpected Python fallback"),
@ -530,7 +541,7 @@ def test_run_rust_ocr_prepares_request_and_wraps_response():
def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing():
bridge = RecordingBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.sync.override(bridge)
ocr_main._run_rust_ocr(
fallback=lambda: pytest.fail("unexpected Python fallback"),
@ -544,7 +555,7 @@ def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing():
def test_run_rust_ocr_prefers_explicit_key_over_resolver():
bridge = RecordingBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.sync.override(bridge)
def _resolver(name: str) -> str | None:
raise AssertionError(f"resolver should not be called for {name}")
@ -565,7 +576,7 @@ def test_run_rust_ocr_uses_provider_api_key_env_var():
bridge = RecordingBridge()
resolver_calls = []
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.sync.override(bridge)
def _resolver(name):
resolver_calls.append(name)
@ -589,7 +600,7 @@ def test_run_rust_ocr_uses_provider_api_key_env_var():
def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
bridge = RecordingBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.sync.override(bridge)
ocr_main._run_rust_ocr(
fallback=lambda: pytest.fail("unexpected Python fallback"),
@ -614,7 +625,7 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_manager():
bridge = RecordingBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.sync.override(bridge)
def _resolver(name: str) -> str | None:
return {
@ -638,7 +649,7 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana
def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager():
bridge = RecordingBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.sync.override(bridge)
ocr_main._run_rust_ocr(
fallback=lambda: pytest.fail("unexpected Python fallback"),
@ -657,7 +668,7 @@ def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager():
def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint():
bridge = RecordingBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.sync.override(bridge)
ocr_main._run_rust_ocr(
fallback=lambda: pytest.fail("unexpected Python fallback"),
@ -679,7 +690,7 @@ def test_run_rust_ocr_runs_pre_call_logging():
logging_obj = RecordingLogging()
bridge = RecordingBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.sync.override(bridge)
ocr_main._run_rust_ocr(
fallback=lambda: pytest.fail("unexpected Python fallback"),
@ -744,7 +755,7 @@ def test_ocr_exception_type_uses_resolved_provider_context(
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=RaisingBridge())
rust_bridge._OCR.sync.override(RaisingBridge())
with pytest.raises(CapturedException):
litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
@ -790,7 +801,7 @@ async def test_aocr_exception_type_uses_resolved_provider_context(
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
litellm.rust(True)
rust_bridge.set_rust_ocr(aocr=RaisingAsyncBridge())
rust_bridge._OCR.asynchronous.override(RaisingAsyncBridge())
with pytest.raises(CapturedException):
await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test")