mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
81 lines
2.9 KiB
Python
81 lines
2.9 KiB
Python
import os
|
|
import inspect
|
|
from collections.abc import Iterator
|
|
from typing import Final
|
|
|
|
import pytest
|
|
|
|
import litellm
|
|
from litellm.rust_bridge import ocr as native_ocr
|
|
from litellm.rust_bridge.configuration import reset_rust_configuration
|
|
from litellm.rust_bridge.configuration import rust_enabled
|
|
from tests.test_litellm_rust.ocr_test_server import ocr_server # noqa: F401 # pytest fixture export
|
|
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def isolate_rust_ocr_state(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
|
|
callback_attributes: Final = (
|
|
"callbacks",
|
|
"input_callback",
|
|
"success_callback",
|
|
"failure_callback",
|
|
"_async_input_callback",
|
|
"_async_success_callback",
|
|
"_async_failure_callback",
|
|
)
|
|
original_callbacks: Final = {attribute: list(getattr(litellm, attribute)) for attribute in callback_attributes}
|
|
original_cache: Final = litellm.cache
|
|
for attribute in callback_attributes:
|
|
getattr(litellm, attribute).clear()
|
|
litellm.cache = None
|
|
reset_rust_configuration()
|
|
litellm.rust(True)
|
|
python_ocr: Final = litellm.ocr
|
|
python_aocr: Final = litellm.aocr
|
|
signature: Final = inspect.signature(python_ocr)
|
|
|
|
def arguments(args: tuple[object, ...], kwargs: dict[str, object]) -> dict[str, object]:
|
|
bound: Final = signature.bind(*args, **kwargs)
|
|
bound.apply_defaults()
|
|
extra: Final = bound.arguments.pop("kwargs")
|
|
return {**extra, **bound.arguments}
|
|
|
|
def ocr(*args: object, **kwargs: object) -> object:
|
|
if not rust_enabled():
|
|
return python_ocr(*args, **kwargs)
|
|
values: Final = arguments(args, kwargs)
|
|
return native_ocr.aocr(values) if values.get("aocr") is True else native_ocr.ocr(values)
|
|
|
|
async def aocr(*args: object, **kwargs: object) -> object:
|
|
if not rust_enabled():
|
|
return await python_aocr(*args, **kwargs)
|
|
return await native_ocr.aocr(arguments(args, kwargs))
|
|
|
|
monkeypatch.setattr(litellm, "ocr", ocr)
|
|
monkeypatch.setattr(litellm, "aocr", aocr)
|
|
yield
|
|
for attribute, callbacks in original_callbacks.items():
|
|
target = getattr(litellm, attribute)
|
|
target.clear()
|
|
target.extend(callbacks)
|
|
litellm.cache = original_cache
|
|
reset_rust_configuration()
|
|
|
|
|
|
def pytest_collection_modifyitems(items):
|
|
rust_enabled = os.environ.get("LITELLM_RUST", "").strip().lower() in {
|
|
"1",
|
|
"true",
|
|
"yes",
|
|
"on",
|
|
}
|
|
if not rust_enabled:
|
|
skip = pytest.mark.skip(reason="requires LITELLM_RUST=1 and a compiled Rust extension")
|
|
for item in items:
|
|
item.add_marker(skip)
|
|
return
|
|
|
|
try:
|
|
from litellm.rust_bridge import _native # noqa: F401 # validates the installed extension
|
|
except ImportError as error:
|
|
raise pytest.UsageError("LITELLM_RUST=1 requires a compiled litellm.rust_bridge._native extension") from error
|