litellm/tests/test_litellm_rust/conftest.py

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