mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
test(rust): strengthen native callback coverage
This commit is contained in:
parent
b1310a1bf8
commit
94ce534b93
9 changed files with 194 additions and 214 deletions
|
|
@ -37,6 +37,12 @@ class RecordingLogger(CustomLogger):
|
|||
|
||||
def _record(self, name: str, kwargs: object = None, response: object = None) -> None:
|
||||
details: Final = kwargs if isinstance(kwargs, dict) else {}
|
||||
try:
|
||||
snapshot: Final = copy.deepcopy(details)
|
||||
except Exception:
|
||||
snapshot = dict(details)
|
||||
if "exception" in details:
|
||||
snapshot["exception"] = details["exception"]
|
||||
try:
|
||||
loop: Final = asyncio.get_running_loop()
|
||||
has_running_loop: Final = True
|
||||
|
|
@ -50,7 +56,7 @@ class RecordingLogger(CustomLogger):
|
|||
thread=threading.current_thread(),
|
||||
loop=loop,
|
||||
has_running_loop=has_running_loop,
|
||||
kwargs=copy.deepcopy(details),
|
||||
kwargs=snapshot,
|
||||
response=response,
|
||||
)
|
||||
with self._condition:
|
||||
|
|
@ -68,7 +74,11 @@ class RecordingLogger(CustomLogger):
|
|||
return tuple(event for event in self._events if event.name == name)
|
||||
|
||||
async def wait_for_async(self, name: str, count: int = 1, timeout: float = 10) -> tuple[HookEvent, ...]:
|
||||
return await asyncio.wait_for(asyncio.to_thread(self.wait_for, name, count, timeout), timeout=timeout + 1)
|
||||
await asyncio.wait_for(asyncio.to_thread(self.wait_for, name, count, timeout), timeout=timeout + 1)
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
|
||||
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), timeout=timeout)
|
||||
return tuple(event for event in self.events if event.name == name)
|
||||
|
||||
def log_pre_api_call(self, model, messages, kwargs):
|
||||
self._record("log_pre_api_call", kwargs)
|
||||
|
|
|
|||
|
|
@ -1,19 +1,16 @@
|
|||
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.recording_server import recording_server # noqa: F401 # pytest fixture export
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def isolate_rust_state(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
|
||||
def isolate_rust_state() -> Iterator[None]:
|
||||
callback_attributes: Final = (
|
||||
"callbacks",
|
||||
"input_callback",
|
||||
|
|
@ -30,29 +27,6 @@ def isolate_rust_state(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
|
|||
litellm.cache = None # test-quality-ok: isolate the process-global cache from native extension tests
|
||||
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)
|
||||
|
|
|
|||
60
tests/test_litellm_rust/contracts.py
Normal file
60
tests/test_litellm_rust/contracts.py
Normal file
|
|
@ -0,0 +1,60 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.rust_bridge import ocr as native_ocr
|
||||
from tests.test_litellm_rust.recording_server import RecordingServer
|
||||
|
||||
OCR_DOCUMENT: Final = {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}
|
||||
OCR_MODEL: Final = "mistral/mistral-ocr-latest"
|
||||
OCR_RESPONSE: Final = {
|
||||
"pages": [{"index": 0, "markdown": "native OCR response", "images": [], "dimensions": None}],
|
||||
"model": "mistral-ocr-latest",
|
||||
"usage_info": {"pages_processed": 1, "doc_size_bytes": 3},
|
||||
}
|
||||
|
||||
MESSAGES_MODEL: Final = "anthropic/claude-sonnet-4-5-20250929"
|
||||
MESSAGES: Final = [{"role": "user", "content": "Hello"}]
|
||||
MESSAGES_RESPONSE: Final = {
|
||||
"id": "msg_native",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-sonnet-4-5-20250929",
|
||||
"content": [{"type": "text", "text": "Hello from native Messages"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 5, "output_tokens": 4},
|
||||
}
|
||||
|
||||
|
||||
def ocr_arguments(server: RecordingServer, **kwargs: object) -> dict[str, object]:
|
||||
return {
|
||||
"model": OCR_MODEL,
|
||||
"document": dict(OCR_DOCUMENT),
|
||||
"api_key": "test-key",
|
||||
"api_base": server.base_url,
|
||||
**kwargs,
|
||||
}
|
||||
|
||||
|
||||
def call_native_ocr(server: RecordingServer, **kwargs: object) -> OCRResponse:
|
||||
return native_ocr.ocr(ocr_arguments(server, **kwargs))
|
||||
|
||||
|
||||
async def call_native_aocr(server: RecordingServer, **kwargs: object) -> OCRResponse:
|
||||
return await native_ocr.aocr(ocr_arguments(server, **kwargs))
|
||||
|
||||
|
||||
def request_body(kwargs: dict[str, object]) -> dict[str, object]:
|
||||
additional_args = kwargs["additional_args"]
|
||||
assert isinstance(additional_args, dict)
|
||||
body = additional_args["complete_input_dict"]
|
||||
assert isinstance(body, dict)
|
||||
return body
|
||||
|
||||
|
||||
def request_headers(kwargs: dict[str, object]) -> dict[str, object]:
|
||||
additional_args = kwargs["additional_args"]
|
||||
assert isinstance(additional_args, dict)
|
||||
headers = additional_args["headers"]
|
||||
assert isinstance(headers, dict)
|
||||
return headers
|
||||
|
|
@ -33,6 +33,7 @@ class RecordingServer:
|
|||
requests: list[RecordedRequest]
|
||||
responses: list[ResponseSpec]
|
||||
default_response: ResponseSpec
|
||||
expected_requests: int | None = 1
|
||||
|
||||
@property
|
||||
def base_url(self) -> str:
|
||||
|
|
@ -101,3 +102,6 @@ def recording_server() -> Iterator[RecordingServer]:
|
|||
server.shutdown()
|
||||
server.server_close()
|
||||
thread.join()
|
||||
if recording_server.expected_requests is not None:
|
||||
assert len(recording_server.requests) == recording_server.expected_requests
|
||||
assert recording_server.responses == []
|
||||
|
|
|
|||
|
|
@ -5,23 +5,11 @@ from typing import Final, cast
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
from tests.test_litellm_rust.contracts import MESSAGES, MESSAGES_MODEL, MESSAGES_RESPONSE
|
||||
from tests.test_litellm_rust.recording_server import RecordingServer, ResponseSpec
|
||||
|
||||
pytestmark = pytest.mark.requires_rust_extension
|
||||
|
||||
MODEL: Final = "anthropic/claude-sonnet-4-5-20250929"
|
||||
MESSAGES: Final = [{"role": "user", "content": "Hello"}]
|
||||
MESSAGES_RESPONSE: Final = {
|
||||
"id": "msg_native",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-sonnet-4-5-20250929",
|
||||
"content": [{"type": "text", "text": "Hello from native Messages"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 5, "output_tokens": 4},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def messages_server(recording_server: RecordingServer) -> RecordingServer:
|
||||
|
|
@ -31,7 +19,7 @@ def messages_server(recording_server: RecordingServer) -> RecordingServer:
|
|||
|
||||
async def call_messages(server: RecordingServer, **kwargs: object):
|
||||
return await litellm.anthropic.messages.acreate(
|
||||
model=MODEL,
|
||||
model=MESSAGES_MODEL,
|
||||
messages=MESSAGES,
|
||||
max_tokens=64,
|
||||
api_key="test-key",
|
||||
|
|
@ -45,12 +33,17 @@ def assert_native_request(server: RecordingServer) -> None:
|
|||
assert "accept-encoding" not in server.requests[0].headers
|
||||
|
||||
|
||||
def assert_native_response(response: object) -> None:
|
||||
assert isinstance(response, dict)
|
||||
assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_messages_sends_expected_provider_request(messages_server: RecordingServer) -> None:
|
||||
response: Final = await call_messages(messages_server)
|
||||
|
||||
assert response["content"] == [{"type": "text", "text": "Hello from native Messages"}]
|
||||
assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert_native_response(response)
|
||||
assert_native_request(messages_server)
|
||||
request: Final = messages_server.requests[0]
|
||||
assert request.path == "/v1/messages"
|
||||
|
|
@ -65,8 +58,9 @@ async def test_messages_sends_expected_provider_request(messages_server: Recordi
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_messages_sends_custom_headers(messages_server: RecordingServer) -> None:
|
||||
await call_messages(messages_server, extra_headers={"x-trace-id": "trace-1"})
|
||||
response: Final = await call_messages(messages_server, extra_headers={"x-trace-id": "trace-1"})
|
||||
|
||||
assert_native_response(response)
|
||||
assert messages_server.requests[0].headers["x-trace-id"] == "trace-1"
|
||||
|
||||
|
||||
|
|
@ -76,13 +70,14 @@ async def test_messages_resolves_provider_credentials(
|
|||
) -> None:
|
||||
monkeypatch.setenv("ANTHROPIC_API_KEY", "environment-key")
|
||||
|
||||
await litellm.anthropic.messages.acreate(
|
||||
model=MODEL,
|
||||
response: Final = await litellm.anthropic.messages.acreate(
|
||||
model=MESSAGES_MODEL,
|
||||
messages=MESSAGES,
|
||||
max_tokens=64,
|
||||
api_base=messages_server.base_url,
|
||||
)
|
||||
|
||||
assert_native_response(response)
|
||||
assert messages_server.requests[0].headers["x-api-key"] == "environment-key"
|
||||
|
||||
|
||||
|
|
@ -92,14 +87,15 @@ async def test_messages_explicit_credentials_override_defaults(
|
|||
) -> None:
|
||||
monkeypatch.setenv("ANTHROPIC_API_KEY", "environment-key")
|
||||
|
||||
await call_messages(messages_server)
|
||||
response: Final = await call_messages(messages_server)
|
||||
|
||||
assert_native_response(response)
|
||||
assert messages_server.requests[0].headers["x-api-key"] == "test-key"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_messages_uses_foundry_endpoint_and_credentials(messages_server: RecordingServer) -> None:
|
||||
await litellm.anthropic.messages.acreate(
|
||||
response: Final = await litellm.anthropic.messages.acreate(
|
||||
model="azure_ai/claude-opus-4.5",
|
||||
messages=MESSAGES,
|
||||
max_tokens=64,
|
||||
|
|
@ -107,6 +103,7 @@ async def test_azure_messages_uses_foundry_endpoint_and_credentials(messages_ser
|
|||
api_base=f"{messages_server.base_url}/anthropic",
|
||||
)
|
||||
|
||||
assert_native_response(response)
|
||||
assert_native_request(messages_server)
|
||||
assert messages_server.requests[0].path == "/anthropic/v1/messages"
|
||||
assert messages_server.requests[0].headers["x-api-key"] == "azure-key"
|
||||
|
|
|
|||
|
|
@ -9,23 +9,17 @@ import pytest
|
|||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from tests.test_litellm_rust.callback_recorder import RecordingLogger
|
||||
from tests.test_litellm_rust.contracts import (
|
||||
MESSAGES,
|
||||
MESSAGES_MODEL,
|
||||
MESSAGES_RESPONSE,
|
||||
request_body,
|
||||
request_headers,
|
||||
)
|
||||
from tests.test_litellm_rust.recording_server import RecordingServer, ResponseSpec
|
||||
|
||||
pytestmark = pytest.mark.requires_rust_extension
|
||||
|
||||
MODEL: Final = "anthropic/claude-sonnet-4-5-20250929"
|
||||
MESSAGES: Final = [{"role": "user", "content": "Hello"}]
|
||||
MESSAGES_RESPONSE: Final = {
|
||||
"id": "msg_native",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-sonnet-4-5-20250929",
|
||||
"content": [{"type": "text", "text": "Hello from native Messages"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 5, "output_tokens": 4},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def messages_server(recording_server: RecordingServer) -> RecordingServer:
|
||||
|
|
@ -35,7 +29,7 @@ def messages_server(recording_server: RecordingServer) -> RecordingServer:
|
|||
|
||||
async def call_messages(server: RecordingServer, callbacks: list[CustomLogger], **kwargs: object):
|
||||
return await litellm.anthropic.messages.acreate(
|
||||
model=MODEL,
|
||||
model=MESSAGES_MODEL,
|
||||
messages=MESSAGES,
|
||||
max_tokens=64,
|
||||
api_key="test-key",
|
||||
|
|
@ -45,16 +39,7 @@ async def call_messages(server: RecordingServer, callbacks: list[CustomLogger],
|
|||
)
|
||||
|
||||
|
||||
def request_body(kwargs: dict) -> dict:
|
||||
return kwargs["additional_args"]["complete_input_dict"]
|
||||
|
||||
|
||||
def request_headers(kwargs: dict) -> dict:
|
||||
return kwargs["additional_args"]["headers"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.xfail(strict=True, reason="UC-MSG-PRECALL-VIEW: native pre-call arguments differ from legacy")
|
||||
async def test_messages_pre_call_receives_expected_provider_request(messages_server: RecordingServer) -> None:
|
||||
observations: Final = []
|
||||
|
||||
|
|
@ -90,10 +75,6 @@ async def test_messages_pre_call_receives_expected_provider_request(messages_ser
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("raise_after_edit", [False, True])
|
||||
@pytest.mark.xfail(
|
||||
strict=True,
|
||||
reason="UC-MSG-PRECALL-MUTATION: native transport snapshots body and headers before pre-call",
|
||||
)
|
||||
async def test_messages_pre_call_edits_reach_later_callbacks_and_provider(
|
||||
messages_server: RecordingServer, raise_after_edit: bool
|
||||
) -> None:
|
||||
|
|
@ -168,6 +149,26 @@ async def test_messages_callbacks_run_once(messages_server: RecordingServer) ->
|
|||
assert recorder.names.count("log_pre_api_call") == 1
|
||||
assert recorder.names.count("async_logging_hook") == 1
|
||||
assert recorder.names.count("async_log_success_event") == 1
|
||||
assert "log_failure_event" not in recorder.names
|
||||
assert "async_log_failure_event" not in recorder.names
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.xfail(strict=True, reason="accepted native Messages errors are replayed through the Python transport")
|
||||
async def test_messages_failure_callbacks_receive_original_provider_error(messages_server: RecordingServer) -> None:
|
||||
messages_server.default_response = ResponseSpec(body={"error": {"message": "provider unavailable"}}, status=500)
|
||||
messages_server.expected_requests = None
|
||||
recorder: Final = RecordingLogger()
|
||||
|
||||
with pytest.raises(litellm.InternalServerError) as caught:
|
||||
await call_messages(messages_server, [recorder])
|
||||
events: Final = await recorder.wait_for_async("async_log_failure_event")
|
||||
|
||||
assert len(events) == 1
|
||||
assert events[0].call_type == "anthropic_messages"
|
||||
assert events[0].kwargs["exception"] is caught.value
|
||||
assert len(messages_server.requests) == 1
|
||||
assert "async_log_success_event" not in recorder.names
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -206,10 +207,6 @@ async def test_messages_pre_call_runs_in_callers_execution_context(messages_serv
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.xfail(
|
||||
strict=True,
|
||||
reason="UC-MSG-STREAM-COMPLETION: native fake stream logs before assembled stream finalization",
|
||||
)
|
||||
async def test_messages_stream_logs_success_after_exhaustion(messages_server: RecordingServer) -> None:
|
||||
recorder: Final = RecordingLogger()
|
||||
stream: Final = await call_messages(messages_server, [recorder], stream=True)
|
||||
|
|
|
|||
|
|
@ -1,22 +1,20 @@
|
|||
import time
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from tests.test_litellm_rust.contracts import (
|
||||
OCR_DOCUMENT,
|
||||
OCR_MODEL,
|
||||
OCR_RESPONSE,
|
||||
call_native_ocr,
|
||||
)
|
||||
from tests.test_litellm_rust.callback_recorder import RecordingLogger
|
||||
from tests.test_litellm_rust.recording_server import RecordingServer, ResponseSpec
|
||||
|
||||
pytestmark = pytest.mark.requires_rust_extension
|
||||
|
||||
DOCUMENT: Final = {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}
|
||||
MODEL: Final = "mistral/mistral-ocr-latest"
|
||||
OCR_RESPONSE: Final = {
|
||||
"pages": [{"index": 0, "markdown": "native OCR response", "images": [], "dimensions": None}],
|
||||
"model": "mistral-ocr-latest",
|
||||
"usage_info": {"pages_processed": 1, "doc_size_bytes": 3},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def ocr_server(recording_server: RecordingServer) -> RecordingServer:
|
||||
|
|
@ -25,13 +23,7 @@ def ocr_server(recording_server: RecordingServer) -> RecordingServer:
|
|||
|
||||
|
||||
def call_ocr(server: RecordingServer, **kwargs: object) -> OCRResponse:
|
||||
return litellm.ocr(
|
||||
model=MODEL,
|
||||
document=dict(DOCUMENT),
|
||||
api_key="test-key",
|
||||
api_base=server.base_url,
|
||||
**kwargs,
|
||||
)
|
||||
return call_native_ocr(server, **kwargs)
|
||||
|
||||
|
||||
def assert_native_request(server: RecordingServer) -> None:
|
||||
|
|
@ -45,19 +37,22 @@ def test_ocr_sends_expected_provider_request(ocr_server: RecordingServer) -> Non
|
|||
assert response.pages[0].markdown == "native OCR response"
|
||||
assert_native_request(ocr_server)
|
||||
assert ocr_server.requests[0].path == "/v1/ocr"
|
||||
assert ocr_server.requests[0].body == {"model": "mistral-ocr-latest", "document": DOCUMENT}
|
||||
assert ocr_server.requests[0].body == {"model": "mistral-ocr-latest", "document": OCR_DOCUMENT}
|
||||
|
||||
|
||||
def test_ocr_rejects_unsupported_file_document_before_callbacks(ocr_server: RecordingServer) -> None:
|
||||
ocr_server.expected_requests = 0
|
||||
recorder: Final = RecordingLogger()
|
||||
|
||||
with pytest.raises(NotImplementedError, match="OCR file document preparation"):
|
||||
litellm.ocr(
|
||||
model=MODEL,
|
||||
call_native_ocr(
|
||||
ocr_server,
|
||||
document={"type": "file", "file": b"%PDF-1.4", "mime_type": "application/pdf"},
|
||||
api_key="test-key",
|
||||
api_base=ocr_server.base_url,
|
||||
callbacks=[recorder],
|
||||
)
|
||||
|
||||
assert ocr_server.requests == []
|
||||
assert recorder.events == ()
|
||||
|
||||
|
||||
def test_ocr_sends_optional_parameters(ocr_server: RecordingServer) -> None:
|
||||
|
|
@ -77,7 +72,7 @@ def test_ocr_sends_custom_headers(ocr_server: RecordingServer) -> None:
|
|||
def test_ocr_resolves_provider_credentials(ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("MISTRAL_API_KEY", "environment-key")
|
||||
|
||||
litellm.ocr(model=MODEL, document=DOCUMENT, api_base=ocr_server.base_url)
|
||||
call_native_ocr(ocr_server, api_key=None)
|
||||
|
||||
assert ocr_server.requests[0].headers["authorization"] == "Bearer environment-key"
|
||||
|
||||
|
|
@ -96,7 +91,7 @@ def test_ocr_resolves_provider_endpoint(ocr_server: RecordingServer, monkeypatch
|
|||
monkeypatch.setenv("AZURE_AI_API_KEY", "azure-key")
|
||||
monkeypatch.setenv("AZURE_AI_API_BASE", ocr_server.base_url)
|
||||
|
||||
litellm.ocr(model="azure_ai/pixtral-12b-2409", document=DOCUMENT)
|
||||
call_native_ocr(ocr_server, model="azure_ai/pixtral-12b-2409", api_key=None, api_base=None)
|
||||
|
||||
assert_native_request(ocr_server)
|
||||
assert ocr_server.requests[0].path == "/providers/mistral/azure/ocr"
|
||||
|
|
@ -104,11 +99,10 @@ def test_ocr_resolves_provider_endpoint(ocr_server: RecordingServer, monkeypatch
|
|||
|
||||
|
||||
def test_ocr_resolves_vertex_project_and_location(ocr_server: RecordingServer) -> None:
|
||||
litellm.ocr(
|
||||
call_native_ocr(
|
||||
ocr_server,
|
||||
model="vertex_ai/mistral-ocr-2505",
|
||||
document=DOCUMENT,
|
||||
api_key="vertex-token",
|
||||
api_base=ocr_server.base_url,
|
||||
vertex_project="project-1",
|
||||
vertex_location="us-central1",
|
||||
)
|
||||
|
|
@ -141,22 +135,8 @@ def test_ocr_provider_error_preserves_status_and_context(ocr_server: RecordingSe
|
|||
|
||||
def test_ocr_honors_request_timeout(ocr_server: RecordingServer) -> None:
|
||||
ocr_server.enqueue(ResponseSpec(body=OCR_RESPONSE, delay=0.2))
|
||||
started_at: Final = time.monotonic()
|
||||
|
||||
with pytest.raises(RuntimeError, match="OCR transport failed"):
|
||||
call_ocr(ocr_server, timeout=0.01)
|
||||
|
||||
assert time.monotonic() - started_at < 0.15
|
||||
assert len(ocr_server.requests) == 1
|
||||
|
||||
|
||||
def test_ocr_respects_runtime_toggle(ocr_server: RecordingServer) -> None:
|
||||
litellm.rust(False)
|
||||
call_ocr(ocr_server)
|
||||
litellm.rust(True)
|
||||
call_ocr(ocr_server)
|
||||
|
||||
assert len(ocr_server.requests) == 2
|
||||
assert ocr_server.requests[0].headers["user-agent"].startswith("litellm/")
|
||||
assert ocr_server.requests[0].headers["accept-encoding"] != "identity"
|
||||
assert ocr_server.requests[1].headers["accept-encoding"] == "identity"
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
import asyncio
|
||||
import copy
|
||||
import json
|
||||
import queue
|
||||
import threading
|
||||
from typing import Final
|
||||
|
|
@ -10,18 +9,19 @@ import pytest
|
|||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from tests.test_litellm_rust.callback_recorder import RecordingLogger
|
||||
from tests.test_litellm_rust.contracts import (
|
||||
OCR_DOCUMENT,
|
||||
OCR_MODEL,
|
||||
OCR_RESPONSE,
|
||||
call_native_aocr,
|
||||
call_native_ocr,
|
||||
request_body,
|
||||
request_headers,
|
||||
)
|
||||
from tests.test_litellm_rust.recording_server import RecordingServer, ResponseSpec
|
||||
|
||||
pytestmark = pytest.mark.requires_rust_extension
|
||||
|
||||
DOCUMENT: Final = {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}
|
||||
MODEL: Final = "mistral/mistral-ocr-latest"
|
||||
OCR_RESPONSE: Final = {
|
||||
"pages": [{"index": 0, "markdown": "native OCR response", "images": [], "dimensions": None}],
|
||||
"model": "mistral-ocr-latest",
|
||||
"usage_info": {"pages_processed": 1, "doc_size_bytes": 3},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def ocr_server(recording_server: RecordingServer) -> RecordingServer:
|
||||
|
|
@ -30,33 +30,11 @@ def ocr_server(recording_server: RecordingServer) -> RecordingServer:
|
|||
|
||||
|
||||
def call_ocr(server: RecordingServer, callbacks: list[CustomLogger], **kwargs: object):
|
||||
return litellm.ocr(
|
||||
model=MODEL,
|
||||
document=dict(DOCUMENT),
|
||||
api_key="test-key",
|
||||
api_base=server.base_url,
|
||||
callbacks=callbacks,
|
||||
**kwargs,
|
||||
)
|
||||
return call_native_ocr(server, callbacks=callbacks, **kwargs)
|
||||
|
||||
|
||||
async def call_aocr(server: RecordingServer, callbacks: list[CustomLogger], **kwargs: object):
|
||||
return await litellm.aocr(
|
||||
model=MODEL,
|
||||
document=dict(DOCUMENT),
|
||||
api_key="test-key",
|
||||
api_base=server.base_url,
|
||||
callbacks=callbacks,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
def request_body(kwargs: dict) -> dict:
|
||||
return kwargs["additional_args"]["complete_input_dict"]
|
||||
|
||||
|
||||
def request_headers(kwargs: dict) -> dict:
|
||||
return kwargs["additional_args"]["headers"]
|
||||
return await call_native_aocr(server, callbacks=callbacks, **kwargs)
|
||||
|
||||
|
||||
def test_pre_call_receives_expected_provider_request(ocr_server: RecordingServer) -> None:
|
||||
|
|
@ -75,7 +53,7 @@ def test_pre_call_receives_expected_provider_request(ocr_server: RecordingServer
|
|||
assert additional_args["api_base"] == f"{ocr_server.base_url}/v1/ocr"
|
||||
assert additional_args["complete_input_dict"] == {
|
||||
"model": "mistral-ocr-latest",
|
||||
"document": DOCUMENT,
|
||||
"document": OCR_DOCUMENT,
|
||||
"pages": [0],
|
||||
}
|
||||
|
||||
|
|
@ -120,7 +98,7 @@ def test_pre_call_header_edits_reach_later_callbacks_and_provider(ocr_server: Re
|
|||
|
||||
|
||||
def test_pre_call_nested_mutation_updates_retained_references(ocr_server: RecordingServer) -> None:
|
||||
original: Final = dict(DOCUMENT)
|
||||
original: Final = dict(OCR_DOCUMENT)
|
||||
replacement_url: Final = "data:application/pdf;base64,ZGVm"
|
||||
retained: Final = []
|
||||
|
||||
|
|
@ -132,11 +110,9 @@ def test_pre_call_nested_mutation_updates_retained_references(ocr_server: Record
|
|||
def log_pre_api_call(self, model, messages, kwargs):
|
||||
request_body(kwargs)["document"]["document_url"] = replacement_url
|
||||
|
||||
litellm.ocr(
|
||||
model=MODEL,
|
||||
call_native_ocr(
|
||||
ocr_server,
|
||||
document=original,
|
||||
api_key="test-key",
|
||||
api_base=ocr_server.base_url,
|
||||
callbacks=[Retain(), Edit()],
|
||||
)
|
||||
|
||||
|
|
@ -146,7 +122,7 @@ def test_pre_call_nested_mutation_updates_retained_references(ocr_server: Record
|
|||
|
||||
|
||||
def test_pre_call_field_replacement_preserves_original_references(ocr_server: RecordingServer) -> None:
|
||||
original: Final = dict(DOCUMENT)
|
||||
original: Final = dict(OCR_DOCUMENT)
|
||||
replacement: Final = {"type": "document_url", "document_url": "data:application/pdf;base64,ZGVm"}
|
||||
retained: Final = []
|
||||
|
||||
|
|
@ -156,16 +132,14 @@ def test_pre_call_field_replacement_preserves_original_references(ocr_server: Re
|
|||
retained.append(body["document"])
|
||||
body["document"] = replacement
|
||||
|
||||
litellm.ocr(
|
||||
model=MODEL,
|
||||
call_native_ocr(
|
||||
ocr_server,
|
||||
document=original,
|
||||
api_key="test-key",
|
||||
api_base=ocr_server.base_url,
|
||||
callbacks=[RetainAndReplace()],
|
||||
)
|
||||
|
||||
assert retained[0] is original
|
||||
assert original["document_url"] == DOCUMENT["document_url"]
|
||||
assert original["document_url"] == OCR_DOCUMENT["document_url"]
|
||||
assert ocr_server.requests[0].body["document"] == replacement
|
||||
|
||||
|
||||
|
|
@ -183,7 +157,7 @@ def test_pre_call_body_rebinding_does_not_replace_inflight_request(ocr_server: R
|
|||
call_ocr(ocr_server, [Rebind(), Observe()])
|
||||
|
||||
assert observed == [{"replacement": True}]
|
||||
assert ocr_server.requests[0].body == {"model": "mistral-ocr-latest", "document": DOCUMENT}
|
||||
assert ocr_server.requests[0].body == {"model": "mistral-ocr-latest", "document": OCR_DOCUMENT}
|
||||
|
||||
|
||||
def test_queued_payload_observes_later_callback_mutations(ocr_server: RecordingServer) -> None:
|
||||
|
|
@ -202,28 +176,6 @@ def test_queued_payload_observes_later_callback_mutations(ocr_server: RecordingS
|
|||
assert queued[0]["queued-edit"] is True
|
||||
|
||||
|
||||
def test_callback_copies_preserve_expected_sharing(ocr_server: RecordingServer) -> None:
|
||||
copies: Final = {}
|
||||
replacement_url: Final = "data:application/pdf;base64,ZGVm"
|
||||
|
||||
class CopyPayload(CustomLogger):
|
||||
def log_pre_api_call(self, model, messages, kwargs):
|
||||
body = request_body(kwargs)
|
||||
copies["shallow"] = dict(body)
|
||||
copies["deep"] = copy.deepcopy(body)
|
||||
copies["serialized"] = json.dumps(body)
|
||||
|
||||
class Edit(CustomLogger):
|
||||
def log_pre_api_call(self, model, messages, kwargs):
|
||||
request_body(kwargs)["document"]["document_url"] = replacement_url
|
||||
|
||||
call_ocr(ocr_server, [CopyPayload(), Edit()])
|
||||
|
||||
assert copies["shallow"]["document"]["document_url"] == replacement_url
|
||||
assert copies["deep"]["document"]["document_url"] == "data:application/pdf;base64,YWJj"
|
||||
assert json.loads(copies["serialized"])["document"]["document_url"] == "data:application/pdf;base64,YWJj"
|
||||
|
||||
|
||||
def test_pre_call_state_reaches_terminal_callbacks(ocr_server: RecordingServer) -> None:
|
||||
token: Final = object()
|
||||
terminal_tokens: queue.SimpleQueue[object] = queue.SimpleQueue()
|
||||
|
|
@ -283,31 +235,6 @@ async def test_failure_callbacks_receive_expected_context_and_error(ocr_server:
|
|||
assert all(observation[3] is None for observation in observations)
|
||||
|
||||
|
||||
def test_background_callback_can_mutate_retained_state_after_return(ocr_server: RecordingServer) -> None:
|
||||
release: Final = threading.Event()
|
||||
finished: Final = threading.Event()
|
||||
retained: Final = []
|
||||
|
||||
class BackgroundEdit(CustomLogger):
|
||||
def log_pre_api_call(self, model, messages, kwargs):
|
||||
body = request_body(kwargs)
|
||||
retained.append(body)
|
||||
|
||||
def edit() -> None:
|
||||
release.wait(10)
|
||||
body["background-edit"] = True
|
||||
finished.set()
|
||||
|
||||
threading.Thread(target=edit, daemon=True).start()
|
||||
|
||||
call_ocr(ocr_server, [BackgroundEdit()])
|
||||
|
||||
assert "background-edit" not in retained[0]
|
||||
release.set()
|
||||
assert finished.wait(10)
|
||||
assert retained[0]["background-edit"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_runs_in_callers_execution_context(ocr_server: RecordingServer) -> None:
|
||||
caller_loop: Final = asyncio.get_running_loop()
|
||||
|
|
@ -316,7 +243,7 @@ async def test_pre_call_runs_in_callers_execution_context(ocr_server: RecordingS
|
|||
|
||||
await call_aocr(ocr_server, [recorder])
|
||||
|
||||
events: Final = recorder.wait_for("log_pre_api_call")
|
||||
events: Final = await recorder.wait_for_async("log_pre_api_call")
|
||||
assert len(events) == 1
|
||||
assert events[0].loop is caller_loop
|
||||
assert events[0].thread is caller_thread
|
||||
|
|
|
|||
31
tests/test_litellm_rust/test_sdk_dispatch.py
Normal file
31
tests/test_litellm_rust/test_sdk_dispatch.py
Normal file
|
|
@ -0,0 +1,31 @@
|
|||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from tests.test_litellm_rust.contracts import OCR_DOCUMENT, OCR_MODEL, OCR_RESPONSE
|
||||
from tests.test_litellm_rust.recording_server import RecordingServer, ResponseSpec
|
||||
|
||||
pytestmark = pytest.mark.requires_rust_extension
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def ocr_server(recording_server: RecordingServer) -> RecordingServer:
|
||||
recording_server.default_response = ResponseSpec(body=OCR_RESPONSE)
|
||||
return recording_server
|
||||
|
||||
|
||||
@pytest.mark.xfail(
|
||||
strict=True,
|
||||
reason="public litellm.ocr does not dispatch to the native OCR bridge yet",
|
||||
)
|
||||
def test_public_ocr_entrypoint_uses_native_transport_when_enabled(ocr_server: RecordingServer) -> None:
|
||||
response: Final = litellm.ocr(
|
||||
model=OCR_MODEL,
|
||||
document=OCR_DOCUMENT,
|
||||
api_key="test-key",
|
||||
api_base=ocr_server.base_url,
|
||||
)
|
||||
|
||||
assert response.pages[0].markdown == "native OCR response"
|
||||
assert ocr_server.requests[0].headers["accept-encoding"] == "identity"
|
||||
Loading…
Add table
Reference in a new issue