refactor(rust_bridge): route Bedrock transcription through the shared runtime

Replace the stateful transcription loader with NativeBinding pairs and call
runtime.run/arun from the Bedrock dispatch class so the RUST_REQUIRED catalog
row is load-bearing: missing native and admission declines are terminal, and
there is no Python replay. Cover the remaining runtime, OCR lifecycle and
configuration branches, and make decide() exhaustive.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Yujong Lee 2026-09-16 19:30:32 +00:00
parent 735ac9fc0f
commit 358d767c9e
4 changed files with 190 additions and 220 deletions

View file

@ -1,13 +1,29 @@
import base64
from typing import Final
from typing import Final, NoReturn
import httpx
from litellm.litellm_core_utils.audio_utils.utils import process_audio_file
from litellm.rust_bridge import transcription as rust_transcription_bridge
from litellm.rust_bridge import runtime
from litellm.rust_bridge.catalog import Context, Route
from litellm.rust_bridge.timeouts import timeout_to_seconds
from litellm.rust_bridge.transcription import (
NATIVE_ATRANSCRIPTION,
NATIVE_TRANSCRIPTION,
RustAtranscription,
RustTranscription,
)
from litellm.types.utils import FileTypes, TranscriptionResponse
def _no_python_implementation() -> NoReturn:
raise NotImplementedError("Bedrock audio transcription is implemented in Rust only")
async def _no_async_python_implementation() -> NoReturn:
_no_python_implementation()
class BedrockAudioTranscriptionRustDispatch:
@staticmethod
def _audio_payload(audio_file: FileTypes) -> dict[str, object]:
@ -43,19 +59,26 @@ class BedrockAudioTranscriptionRustDispatch:
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
) -> TranscriptionResponse:
rust_response: Final = rust_transcription_bridge.transcription(
model=model,
audio=self._audio_payload(audio_file),
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout=timeout,
def native(rust: RustTranscription) -> TranscriptionResponse:
return TranscriptionResponse(
**rust(
model=model,
audio=self._audio_payload(audio_file),
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout_seconds=timeout_to_seconds(timeout),
)
)
return runtime.run(
Context(Route.TRANSCRIPTION, provider=custom_llm_provider, model=model),
binding=NATIVE_TRANSCRIPTION,
native=native,
python=_no_python_implementation,
)
if rust_response is None:
raise RuntimeError("Rust audio transcription bridge is unavailable")
return TranscriptionResponse(**rust_response)
async def async_audio_transcriptions(
self,
@ -69,16 +92,23 @@ class BedrockAudioTranscriptionRustDispatch:
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
) -> TranscriptionResponse:
rust_response: Final = await rust_transcription_bridge.atranscription(
model=model,
audio=self._audio_payload(audio_file),
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout=timeout,
async def native(rust: RustAtranscription) -> TranscriptionResponse:
return TranscriptionResponse(
**await rust(
model=model,
audio=self._audio_payload(audio_file),
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout_seconds=timeout_to_seconds(timeout),
)
)
return await runtime.arun(
Context(Route.TRANSCRIPTION, provider=custom_llm_provider, model=model),
binding=NATIVE_ATRANSCRIPTION,
native=native,
python=_no_async_python_implementation,
)
if rust_response is None:
raise RuntimeError("Rust audio transcription bridge is unavailable")
return TranscriptionResponse(**rust_response)

View file

@ -4,6 +4,8 @@ import os
from enum import Enum, auto
from typing import Final
from typing_extensions import assert_never
_TRUE_ENV_VALUES: Final = frozenset({"1", "true", "yes", "on"})
_GLOBAL_ENV_NAME: Final = "LITELLM_RUST"
@ -55,6 +57,8 @@ def decide(
else rollout is Rollout.RUST_OPT_OUT
)
return Decision.RUST_WITH_FALLBACK if switch else Decision.PYTHON
case _:
assert_never(rollout)
def decision(rollout: Rollout) -> Decision:

View file

@ -1,12 +1,9 @@
from __future__ import annotations
from collections.abc import Awaitable
from dataclasses import dataclass
from typing import Final, Protocol, cast
from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables
import httpx
from litellm.rust_bridge.timeouts import timeout_to_seconds
from litellm.rust_bridge.bindings import NativeBinding
class RustTranscription(Protocol):
@ -39,110 +36,17 @@ class RustAtranscription(Protocol):
raise NotImplementedError
class _Unset:
pass
_UNSET: Final[_Unset] = _Unset()
@dataclass
class _RustTranscriptionState:
transcription: RustTranscription | None = None
atranscription: RustAtranscription | None = None
_STATE: Final = _RustTranscriptionState()
def configure_rust_transcription(
*,
transcription: RustTranscription | None | _Unset = _UNSET,
atranscription: RustAtranscription | None | _Unset = _UNSET,
) -> None:
if not isinstance(transcription, _Unset):
_STATE.transcription = transcription
if not isinstance(atranscription, _Unset):
_STATE.atranscription = atranscription
def load_rust_transcription() -> RustTranscription | None:
if _STATE.transcription is not None:
return _STATE.transcription
from litellm.rust_bridge import get_native_bridge
native_bridge: Final = get_native_bridge()
return (
None
if native_bridge is None
else cast( # cast-ok: native extension protocol is runtime-defined
RustTranscription, getattr(native_bridge, "transcription", None)
)
)
def load_rust_atranscription() -> RustAtranscription | None:
if _STATE.atranscription is not None:
return _STATE.atranscription
from litellm.rust_bridge import get_native_bridge
native_bridge: Final = get_native_bridge()
return (
None
if native_bridge is None
else cast( # cast-ok: native extension protocol is runtime-defined
RustAtranscription, getattr(native_bridge, "atranscription", None)
)
)
def transcription(
*,
model: str,
audio: 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,
) -> dict[str, object] | None:
rust_transcription: Final = load_rust_transcription()
if rust_transcription is None:
def _sync_binding(value: object) -> RustTranscription | None:
if not callable(value):
return None
return rust_transcription(
model=model,
audio=audio,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout_seconds=timeout_to_seconds(timeout),
)
return cast("RustTranscription", value) # cast-ok: callable validated at the native binding boundary
async def atranscription(
*,
model: str,
audio: 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,
) -> dict[str, object] | None:
rust_atranscription: Final = load_rust_atranscription()
if rust_atranscription is None:
def _async_binding(value: object) -> RustAtranscription | None:
if not callable(value):
return None
return await rust_atranscription(
model=model,
audio=audio,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout_seconds=timeout_to_seconds(timeout),
)
return cast("RustAtranscription", value) # cast-ok: callable validated at the native binding boundary
NATIVE_TRANSCRIPTION: Final = NativeBinding("transcription", validate=_sync_binding)
NATIVE_ATRANSCRIPTION: Final = NativeBinding("atranscription", validate=_async_binding)

View file

@ -1,16 +1,44 @@
import importlib
from __future__ import annotations
from collections.abc import Generator
from types import SimpleNamespace
from typing import Final
import pytest
import litellm
from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch
from litellm.rust_bridge import bindings, configuration
from litellm.rust_bridge.transcription import NATIVE_ATRANSCRIPTION, NATIVE_TRANSCRIPTION
rust_bridge = importlib.import_module("litellm.rust_bridge.transcription")
MODEL: Final = "bedrock/mistral.voxtral-mini-3b-2507"
AUDIO_FILE: Final = ("audio.wav", b"audio", "audio/wav")
class RustBridgeDeclined(Exception):
pass
class RustUpstreamError(Exception):
pass
@pytest.fixture(autouse=True)
def isolated_bridge(monkeypatch: pytest.MonkeyPatch) -> Generator[None]:
native: Final = SimpleNamespace(RustBridgeDeclined=RustBridgeDeclined, RustUpstreamError=RustUpstreamError)
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
monkeypatch.delenv("LITELLM_RUST", raising=False)
configuration.reset_rust_configuration()
yield
NATIVE_TRANSCRIPTION.reset()
NATIVE_ATRANSCRIPTION.reset()
configuration.reset_rust_configuration()
class SyncBridge:
def __init__(self) -> None:
self.calls: list[dict[str, object]] = []
def __init__(self, effect: BaseException | None = None) -> None:
self._effect: Final = effect
self.calls: tuple[dict[str, object], ...] = ()
def __call__(
self,
@ -23,11 +51,19 @@ class SyncBridge:
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> dict[str, object]:
self.calls.append({"model": model, "audio": audio, "optional_params": optional_params})
return {"text": "hello"}
self.calls = (
*self.calls,
{"model": model, "audio": audio, "provider": custom_llm_provider, "timeout": timeout_seconds},
)
if self._effect is not None:
raise self._effect
return {"text": "rust"}
class AsyncBridge:
def __init__(self) -> None:
self.calls: tuple[str, ...] = ()
async def __call__(
self,
model: str,
@ -39,113 +75,109 @@ class AsyncBridge:
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> dict[str, object]:
return {"text": "async"}
self.calls = (*self.calls, model)
return {"text": "async rust"}
def test_enabled_sync_bridge_receives_audio() -> None:
bridge = SyncBridge()
rust_bridge.configure_rust_transcription(transcription=bridge)
result = rust_bridge.transcription(
model="mistral.voxtral-mini-3b-2507",
audio={"data": "AQI=", "format": "wav", "filename": "audio.wav"},
def dispatch_sync() -> litellm.TranscriptionResponse:
return BedrockAudioTranscriptionRustDispatch().audio_transcriptions(
model=MODEL,
audio_file=AUDIO_FILE,
api_key=None,
api_base=None,
custom_llm_provider="bedrock",
extra_headers=None,
optional_params={"temperature": 0},
timeout=5.0,
timeout=5,
)
assert result == {"text": "hello"}
assert bridge.calls[0]["audio"] == {"data": "AQI=", "format": "wav", "filename": "audio.wav"}
@pytest.mark.asyncio
async def test_enabled_async_bridge() -> None:
rust_bridge.configure_rust_transcription(atranscription=AsyncBridge())
result = await rust_bridge.atranscription(
model="mistral.voxtral-mini-3b-2507",
audio={"data": "AQI=", "format": "wav", "filename": "audio.wav"},
api_key=None,
api_base=None,
custom_llm_provider="bedrock",
extra_headers=None,
optional_params={},
timeout=None,
def test_dispatch_marshals_audio_into_rust_call() -> None:
bridge: Final = SyncBridge()
NATIVE_TRANSCRIPTION.override(bridge)
response: Final = dispatch_sync()
assert response.text == "rust"
assert bridge.calls == (
{
"model": MODEL,
"audio": {"data": "YXVkaW8=", "format": "wav", "filename": "audio.wav"},
"provider": "bedrock",
"timeout": 5.0,
},
)
assert result == {"text": "async"}
def test_loader_returns_none_without_native_extension(monkeypatch: pytest.MonkeyPatch) -> None:
rust_bridge.configure_rust_transcription(transcription=None, atranscription=None)
monkeypatch.setattr("litellm.rust_bridge.get_native_bridge", lambda: None)
assert rust_bridge.load_rust_transcription() is None
assert rust_bridge.load_rust_atranscription() is None
@pytest.mark.parametrize("disable", ("process", "environment"))
def test_bedrock_transcription_ignores_optional_rust_switches(disable: str, monkeypatch: pytest.MonkeyPatch) -> None:
if disable == "process":
litellm.rust(False)
else:
monkeypatch.setenv("LITELLM_RUST", "0")
bridge: Final = SyncBridge()
NATIVE_TRANSCRIPTION.override(bridge)
assert dispatch_sync().text == "rust"
assert len(bridge.calls) == 1
def test_dispatch_sync_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(rust_bridge, "transcription", lambda **_: None)
def test_missing_native_binding_raises_without_python_fallback() -> None:
NATIVE_TRANSCRIPTION.override(None)
with pytest.raises(RuntimeError, match="bridge is unavailable"):
BedrockAudioTranscriptionRustDispatch().audio_transcriptions(
model="bedrock/mistral.voxtral-mini-3b-2507",
audio_file=("audio.wav", b"audio", "audio/wav"),
api_key=None,
api_base=None,
custom_llm_provider="bedrock",
extra_headers=None,
optional_params={},
timeout=5,
)
dispatch_sync()
def test_admission_decline_raises_for_required_route() -> None:
NATIVE_TRANSCRIPTION.override(SyncBridge(RustBridgeDeclined("unsupported format")))
with pytest.raises(RuntimeError, match="declined the request: unsupported format"):
dispatch_sync()
def test_upstream_error_maps_to_api_error() -> None:
NATIVE_TRANSCRIPTION.override(SyncBridge(RustUpstreamError(503, "bedrock down")))
with pytest.raises(litellm.APIError, match="bedrock down") as raised:
dispatch_sync()
assert raised.value.status_code == 503
def test_bedrock_transcription_dispatches_to_rust_from_sdk_entrypoint() -> None:
bridge: Final = SyncBridge()
NATIVE_TRANSCRIPTION.override(bridge)
response: Final = litellm.transcription(model=MODEL, file=AUDIO_FILE)
assert isinstance(response, litellm.TranscriptionResponse)
assert response.text == "rust"
assert bridge.calls[0]["model"] == MODEL.removeprefix("bedrock/")
@pytest.mark.asyncio
async def test_dispatch_async_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None:
async def unavailable(**_: object) -> None:
return None
async def test_bedrock_atranscription_dispatches_to_rust_from_sdk_entrypoint() -> None:
bridge: Final = AsyncBridge()
NATIVE_ATRANSCRIPTION.override(bridge)
monkeypatch.setattr(rust_bridge, "atranscription", unavailable)
response: Final = await litellm.atranscription(model=MODEL, file=AUDIO_FILE)
assert response.text == "async rust"
assert bridge.calls == (MODEL.removeprefix("bedrock/"),)
@pytest.mark.asyncio
async def test_async_missing_native_binding_raises_without_python_fallback() -> None:
NATIVE_ATRANSCRIPTION.override(None)
with pytest.raises(RuntimeError, match="bridge is unavailable"):
await BedrockAudioTranscriptionRustDispatch().async_audio_transcriptions(
model="bedrock/mistral.voxtral-mini-3b-2507",
audio_file=("audio.wav", b"audio", "audio/wav"),
model=MODEL,
audio_file=AUDIO_FILE,
api_key=None,
api_base=None,
custom_llm_provider="bedrock",
extra_headers=None,
optional_params={},
timeout=5,
timeout=None,
)
def test_bedrock_transcription_uses_rust_only_path() -> None:
rust_bridge.configure_rust_transcription(
transcription=lambda **_: {"text": "rust"},
atranscription=None,
)
try:
response = litellm.transcription(
model="bedrock/mistral.voxtral-mini-3b-2507",
file=("audio.wav", b"audio", "audio/wav"),
)
finally:
rust_bridge.configure_rust_transcription(transcription=None, atranscription=None)
assert response.text == "rust"
@pytest.mark.asyncio
async def test_bedrock_atranscription_uses_rust_only_path() -> None:
async def rust_response(**_: object) -> dict[str, object]:
return {"text": "rust"}
rust_bridge.configure_rust_transcription(transcription=None, atranscription=rust_response)
try:
response = await litellm.atranscription(
model="bedrock/mistral.voxtral-mini-3b-2507",
file=("audio.wav", b"audio", "audio/wav"),
)
finally:
rust_bridge.configure_rust_transcription(transcription=None, atranscription=None)
assert response.text == "rust"