mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-20 00:11:50 +00:00
fix(rust_bridge): qualify runtime calls in dispatch and drop OCR transport rows from wheel matrix
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
3de23e7f18
commit
13cb739089
2 changed files with 8 additions and 77 deletions
|
|
@ -4,12 +4,11 @@ from collections.abc import Awaitable, Callable, Mapping
|
|||
from dataclasses import dataclass
|
||||
from typing import Final, Generic, TypeVar
|
||||
|
||||
from litellm.rust_bridge import catalog
|
||||
from litellm.rust_bridge import catalog, runtime
|
||||
from litellm.rust_bridge.bindings import NativeBinding
|
||||
from litellm.rust_bridge.catalog import Context, Route, Rules
|
||||
from litellm.rust_bridge.configuration import Decision
|
||||
from litellm.rust_bridge.configuration import decision as rollout_decision
|
||||
from litellm.rust_bridge.runtime import arun, run
|
||||
|
||||
RequestT = TypeVar("RequestT")
|
||||
NativeT = TypeVar("NativeT")
|
||||
|
|
@ -50,7 +49,7 @@ class PublicDispatch(Generic[RequestT]):
|
|||
request: Final = self.request(args, kwargs)
|
||||
if request is None or (self.bypass is not None and self.bypass(request)):
|
||||
return python(*args, **kwargs)
|
||||
return run(
|
||||
return runtime.run(
|
||||
self.context(request),
|
||||
binding=binding,
|
||||
native=lambda hook: native(hook, request, args, kwargs),
|
||||
|
|
@ -74,7 +73,7 @@ class PublicDispatch(Generic[RequestT]):
|
|||
request: Final = self.request(args, kwargs)
|
||||
if request is None or (self.bypass is not None and self.bypass(request)):
|
||||
return await python(*args, **kwargs)
|
||||
return await arun(
|
||||
return await runtime.arun(
|
||||
self.context(request),
|
||||
binding=binding,
|
||||
native=lambda hook: native(hook, request, args, kwargs),
|
||||
|
|
|
|||
|
|
@ -73,32 +73,12 @@ def assert_native_request(
|
|||
headers: HTTPMessage,
|
||||
body: object,
|
||||
) -> None:
|
||||
if route not in {"ocr", "azure_ocr", "azure_di", "transcription", "messages", "chat_completions"}:
|
||||
if route not in {"transcription", "messages", "chat_completions"}:
|
||||
raise AssertionError(f"unexpected route marker: {route!r}")
|
||||
if outcome not in {"success", "429", "hang"}:
|
||||
raise AssertionError(f"unexpected outcome marker: {outcome!r}")
|
||||
if not isinstance(body, dict):
|
||||
raise TypeError(f"{route} sent {type(body).__name__}, expected a JSON object")
|
||||
if route == "ocr":
|
||||
assert path == "/v1/ocr"
|
||||
assert headers.get("authorization") == "Bearer sk-native"
|
||||
assert body["model"] == "mistral-ocr-latest"
|
||||
assert body["document"]["document_url"] == "https://example.com/document.pdf"
|
||||
assert body["include_image_base64"] is True
|
||||
return
|
||||
if route == "azure_ocr":
|
||||
assert path == "/providers/mistral/azure/ocr"
|
||||
assert headers.get("authorization") == "Bearer prepared-azure-token"
|
||||
assert body["model"] == "mistral-ocr-2505"
|
||||
assert body["document"]["document_url"] == "data:application/pdf;base64,YWJj"
|
||||
return
|
||||
if route == "azure_di":
|
||||
assert path.startswith("/documentintelligence/documentModels/prebuilt-read:analyze?")
|
||||
assert "api-version=2024-11-30" in path
|
||||
assert "pages=1%2C3" in path
|
||||
assert headers.get("ocp-apim-subscription-key") == "di-key"
|
||||
assert body == {"base64Source": "YWJj"}
|
||||
return
|
||||
if route == "transcription":
|
||||
assert path == "/model/mistral.voxtral-mini-3b-2507/converse"
|
||||
assert headers.get("authorization", "").startswith("AWS4-HMAC-SHA256 ")
|
||||
|
|
@ -120,10 +100,6 @@ def assert_native_request(
|
|||
def native_response(status: int, route: str | None) -> bytes:
|
||||
if status == 429:
|
||||
return b'{"error":"native-rate-limit"}'
|
||||
if route in {"ocr", "azure_ocr"}:
|
||||
return b'{"pages":[{"index":0,"markdown":"native-ocr"}]}'
|
||||
if route == "azure_di":
|
||||
return b'{"status":"succeeded","analyzeResult":{"pages":[]}}'
|
||||
if route == "transcription":
|
||||
return b'{"output":{"message":{"content":[{"text":"native-transcription"}]}}}'
|
||||
return ANTHROPIC_RESPONSE
|
||||
|
|
@ -144,14 +120,6 @@ def route_kwargs(route: str, api_base: str, outcome: str) -> dict[str, object]:
|
|||
"extra_headers": {"x-test-outcome": outcome, "x-test-route": route},
|
||||
"timeout_seconds": 3.0,
|
||||
}
|
||||
if route == "ocr":
|
||||
return common | {
|
||||
"model": "mistral-ocr-latest",
|
||||
"document": {"type": "document_url", "document_url": "https://example.com/document.pdf"},
|
||||
"api_key": "sk-native",
|
||||
"custom_llm_provider": "mistral",
|
||||
"optional_params": {"include_image_base64": True},
|
||||
}
|
||||
if route == "transcription":
|
||||
return common | {
|
||||
"model": "mistral.voxtral-mini-3b-2507",
|
||||
|
|
@ -189,42 +157,12 @@ def assert_success(route: str, response: object) -> None:
|
|||
if not isinstance(response, dict):
|
||||
raise TypeError(f"{route} returned {type(response).__name__}, expected dict")
|
||||
actual: Final = success_value(route, response)
|
||||
expected: Final = (
|
||||
"native-ocr" if route == "ocr" else "native-transcription" if route == "transcription" else "native-message"
|
||||
)
|
||||
expected: Final = "native-transcription" if route == "transcription" else "native-message"
|
||||
if actual != expected:
|
||||
raise AssertionError(f"{route} returned {actual!r}, expected {expected!r}")
|
||||
|
||||
|
||||
def azure_ocr_kwargs(api_base: str) -> dict[str, object]:
|
||||
return {
|
||||
"model": "mistral-ocr-2505",
|
||||
"document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
|
||||
"api_base": api_base,
|
||||
"custom_llm_provider": "azure_ai",
|
||||
"extra_headers": {
|
||||
"x-test-outcome": "success",
|
||||
"x-test-route": "azure_ocr",
|
||||
},
|
||||
"optional_params": {"azure_ad_token": "prepared-azure-token"},
|
||||
}
|
||||
|
||||
|
||||
def azure_di_kwargs(api_base: str) -> dict[str, object]:
|
||||
return {
|
||||
"model": "doc-intelligence/prebuilt-read",
|
||||
"document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
|
||||
"api_key": "di-key",
|
||||
"api_base": api_base,
|
||||
"custom_llm_provider": "azure_ai",
|
||||
"extra_headers": {"x-test-outcome": "success", "x-test-route": "azure_di"},
|
||||
"optional_params": {"req_format": "native", "pages": [0, 2]},
|
||||
}
|
||||
|
||||
|
||||
def success_value(route: str, response: dict[object, object]) -> object:
|
||||
if route == "ocr":
|
||||
return response["pages"][0]["markdown"]
|
||||
if route == "transcription":
|
||||
return response["text"]
|
||||
if route == "messages":
|
||||
|
|
@ -233,7 +171,7 @@ def success_value(route: str, response: dict[object, object]) -> object:
|
|||
|
||||
|
||||
def assert_rate_limit(native: object, route: str, error: BaseException) -> None:
|
||||
if route in {"ocr", "chat_completions"}:
|
||||
if route == "chat_completions":
|
||||
upstream_error: Final = native.RustUpstreamError
|
||||
if not isinstance(error, upstream_error) or error.args[0] != 429:
|
||||
raise AssertionError(f"{route} returned the wrong 429 error: {error!r}")
|
||||
|
|
@ -243,7 +181,7 @@ def assert_rate_limit(native: object, route: str, error: BaseException) -> None:
|
|||
|
||||
|
||||
def exercise_sync(native: object, api_base: str) -> None:
|
||||
for route in ("ocr", "transcription", "messages", "chat_completions"):
|
||||
for route in ("transcription", "messages", "chat_completions"):
|
||||
function: Final = getattr(native, route)
|
||||
assert_success(route, function(**route_kwargs(route, api_base, "success")))
|
||||
try:
|
||||
|
|
@ -252,13 +190,10 @@ def exercise_sync(native: object, api_base: str) -> None:
|
|||
assert_rate_limit(native, route, error)
|
||||
else:
|
||||
raise AssertionError(f"{route} accepted a 429 response")
|
||||
assert_success("ocr", native.ocr(**azure_ocr_kwargs(api_base)))
|
||||
di_response: Final = native.ocr(**azure_di_kwargs(api_base))
|
||||
assert di_response["provider_native_response"]["status"] == "succeeded"
|
||||
|
||||
|
||||
async def exercise_async(native: object, api_base: str) -> None:
|
||||
for route in ("ocr", "transcription", "messages", "chat_completions"):
|
||||
for route in ("transcription", "messages", "chat_completions"):
|
||||
function: Final = getattr(native, f"a{route}")
|
||||
assert_success(route, await function(**route_kwargs(route, api_base, "success")))
|
||||
try:
|
||||
|
|
@ -267,9 +202,6 @@ async def exercise_async(native: object, api_base: str) -> None:
|
|||
assert_rate_limit(native, route, error)
|
||||
else:
|
||||
raise AssertionError(f"a{route} accepted a 429 response")
|
||||
assert_success("ocr", await native.aocr(**azure_ocr_kwargs(api_base)))
|
||||
di_response: Final = await native.aocr(**azure_di_kwargs(api_base))
|
||||
assert di_response["provider_native_response"]["status"] == "succeeded"
|
||||
|
||||
|
||||
async def exercise_async_concurrency(native: object, api_base: str) -> None:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue