mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
test(ocr): add SDK callback E2E parity (#40061)
* test(ocr): cover SDK callback parity * test(ocr): compare callback kwargs
This commit is contained in:
parent
9acc01efce
commit
5f2b4d27d7
9 changed files with 453 additions and 30 deletions
|
|
@ -29,6 +29,7 @@ from litellm.llms.base_llm.ocr.transformation import (
|
|||
)
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.rust_bridge import ocr as rust_ocr_bridge
|
||||
from litellm.rust_bridge.bindings import native_exception_types
|
||||
from litellm.rust_bridge.configuration import rust_enabled
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.utils import ProviderConfigManager, client
|
||||
|
|
@ -281,6 +282,33 @@ def _prepare_rust_ocr_call(
|
|||
)
|
||||
|
||||
|
||||
def _map_rust_ocr_error(
|
||||
error: Exception,
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
exception_types: tuple[type[BaseException], type[BaseException]] | None,
|
||||
) -> Exception:
|
||||
if exception_types is None:
|
||||
return error
|
||||
_, upstream_error = exception_types
|
||||
if not isinstance(error, upstream_error):
|
||||
return error
|
||||
error_args: Final = cast( # cast-ok: BaseException.args is typed with Any in the standard library stubs
|
||||
tuple[object, ...], error.args
|
||||
)
|
||||
status_value: Final = error_args[0] if error_args else 0
|
||||
message_value: Final = error_args[1] if len(error_args) > 1 else str(error)
|
||||
status: Final = status_value if isinstance(status_value, int) else 0
|
||||
message: Final = message_value if isinstance(message_value, str) else str(message_value)
|
||||
error_factory: Final = cast( # cast-ok: the legacy provider interface leaves callable parameters untyped
|
||||
Callable[..., Exception], prepared_request.provider_config.get_error_class
|
||||
)
|
||||
return error_factory(
|
||||
error_message=message,
|
||||
status_code=status or 500,
|
||||
headers={}, # mutable-ok: provider error factories require a concrete header dict
|
||||
)
|
||||
|
||||
|
||||
def _run_rust_ocr(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
|
|
@ -291,16 +319,19 @@ def _run_rust_ocr(
|
|||
prepared_request=prepared_request,
|
||||
resolve_api_key=resolve_api_key,
|
||||
)
|
||||
rust_response: Final = rust_ocr_bridge.ocr(
|
||||
model=prepared_request.model,
|
||||
document=prepared_request.document,
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared.api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
extra_headers=prepared.headers,
|
||||
optional_params=prepared.optional_params,
|
||||
timeout=prepared_request.effective_timeout,
|
||||
)
|
||||
try:
|
||||
rust_response: Final = rust_ocr_bridge.ocr(
|
||||
model=prepared_request.model,
|
||||
document=prepared_request.document,
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared.api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
extra_headers=prepared.headers,
|
||||
optional_params=prepared.optional_params,
|
||||
timeout=prepared_request.effective_timeout,
|
||||
)
|
||||
except Exception as error:
|
||||
raise _map_rust_ocr_error(error, prepared_request, native_exception_types()) from error
|
||||
if rust_response is None:
|
||||
return None
|
||||
return OCRResponse.model_validate(rust_response)
|
||||
|
|
@ -316,16 +347,19 @@ async def _run_rust_aocr(
|
|||
prepared_request=prepared_request,
|
||||
resolve_api_key=resolve_api_key,
|
||||
)
|
||||
rust_response: Final = await rust_ocr_bridge.aocr(
|
||||
model=prepared_request.model,
|
||||
document=prepared_request.document,
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared.api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
extra_headers=prepared.headers,
|
||||
optional_params=prepared.optional_params,
|
||||
timeout=prepared_request.effective_timeout,
|
||||
)
|
||||
try:
|
||||
rust_response: Final = await rust_ocr_bridge.aocr(
|
||||
model=prepared_request.model,
|
||||
document=prepared_request.document,
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared.api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
extra_headers=prepared.headers,
|
||||
optional_params=prepared.optional_params,
|
||||
timeout=prepared_request.effective_timeout,
|
||||
)
|
||||
except Exception as error:
|
||||
raise _map_rust_ocr_error(error, prepared_request, native_exception_types()) from error
|
||||
if rust_response is None:
|
||||
return None
|
||||
return OCRResponse.model_validate(rust_response)
|
||||
|
|
|
|||
|
|
@ -78,3 +78,4 @@ def assert_parity(baseline: Execution, candidate: Execution, baseline_user_agent
|
|||
validate_harness(baseline, candidate, baseline_user_agent)
|
||||
assert_request_parity(baseline.requests, candidate.requests)
|
||||
assert_value_parity(baseline.report, candidate.report)
|
||||
assert_value_parity(baseline.callbacks, candidate.callbacks, path="$.callbacks")
|
||||
|
|
|
|||
|
|
@ -37,6 +37,25 @@ class SDKError(BaseModel):
|
|||
llm_provider: str | None
|
||||
|
||||
|
||||
class CallbackObservation(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
hook: Literal[
|
||||
"log_success_event",
|
||||
"async_log_success_event",
|
||||
"log_failure_event",
|
||||
"async_log_failure_event",
|
||||
]
|
||||
phase: Literal["success", "failure"]
|
||||
model: str | None
|
||||
call_type: str | None
|
||||
litellm_call_id: str | None
|
||||
metadata: JsonValue
|
||||
kwargs: JsonValue
|
||||
payload: JsonValue
|
||||
error: SDKError | None
|
||||
|
||||
|
||||
class SDKJsonChunk(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
|
|
@ -119,6 +138,7 @@ class Execution(BaseModel):
|
|||
|
||||
requests: tuple[CapturedRequest, ...]
|
||||
report: SDKReport
|
||||
callbacks: tuple[CallbackObservation, ...] | None = None
|
||||
|
||||
|
||||
class SDKCommand(BaseModel):
|
||||
|
|
@ -133,6 +153,7 @@ class WorkerSuccess(BaseModel):
|
|||
|
||||
status: Literal["ok"] = "ok"
|
||||
report: SDKReport
|
||||
callbacks: tuple[CallbackObservation, ...] | None = None
|
||||
|
||||
|
||||
class WorkerFailure(BaseModel):
|
||||
|
|
|
|||
|
|
@ -39,9 +39,7 @@ class SubprocessRunner:
|
|||
return (
|
||||
sys.executable,
|
||||
"-m",
|
||||
".".join(
|
||||
self.entrypoint.resolve().relative_to(PROJECT_ROOT).with_suffix("").parts
|
||||
),
|
||||
".".join(self.entrypoint.resolve().relative_to(PROJECT_ROOT).with_suffix("").parts),
|
||||
"--parity-worker",
|
||||
provider_url,
|
||||
)
|
||||
|
|
@ -113,7 +111,11 @@ class SubprocessWorker:
|
|||
)
|
||||
assert isinstance(result, WorkerSuccess)
|
||||
try:
|
||||
return Execution(requests=self.provider.take_requests(len(responses)), report=result.report)
|
||||
return Execution(
|
||||
requests=self.provider.take_requests(len(responses)),
|
||||
report=result.report,
|
||||
callbacks=result.callbacks,
|
||||
)
|
||||
except AssertionError:
|
||||
self.provider.reset()
|
||||
raise
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ import pytest
|
|||
from pydantic import BaseModel, ConfigDict, JsonValue, PrivateAttr
|
||||
|
||||
from .compare import assert_model_parity, assert_parity
|
||||
from .models import CapturedRequest, Execution, SDKError, SDKSuccess, sdk_error_report
|
||||
from .models import CallbackObservation, CapturedRequest, Execution, SDKError, SDKSuccess, sdk_error_report
|
||||
|
||||
SENTINEL: Final = "python-parity-fallback"
|
||||
|
||||
|
|
@ -69,6 +69,56 @@ def test_parity_rejects_response_difference() -> None:
|
|||
assert_parity(python, rust, SENTINEL)
|
||||
|
||||
|
||||
def test_parity_distinguishes_unobserved_callbacks_from_zero_events() -> None:
|
||||
python: Final = _execution(user_agent=SENTINEL)
|
||||
rust: Final = _execution(user_agent="litellm-rust").model_copy(update={"callbacks": ()})
|
||||
|
||||
with pytest.raises(AssertionError, match=r"\$\.callbacks"):
|
||||
assert_parity(python, rust, SENTINEL)
|
||||
|
||||
|
||||
def test_parity_rejects_callback_payload_difference() -> None:
|
||||
observation: Final = CallbackObservation(
|
||||
hook="log_success_event",
|
||||
phase="success",
|
||||
model="test-model",
|
||||
call_type="ocr",
|
||||
litellm_call_id="test-call",
|
||||
metadata={"profile": "success"},
|
||||
kwargs={"model": "test-model"},
|
||||
payload={"model": "test-model", "pages": []},
|
||||
error=None,
|
||||
)
|
||||
python: Final = _execution(user_agent=SENTINEL).model_copy(update={"callbacks": (observation,)})
|
||||
rust: Final = _execution(user_agent="litellm-rust").model_copy(
|
||||
update={"callbacks": (observation.model_copy(update={"payload": {"model": "changed", "pages": []}}),)}
|
||||
)
|
||||
|
||||
with pytest.raises(AssertionError, match=r"\$\.callbacks"):
|
||||
assert_parity(python, rust, SENTINEL)
|
||||
|
||||
|
||||
def test_parity_rejects_callback_kwargs_difference() -> None:
|
||||
observation: Final = CallbackObservation(
|
||||
hook="log_success_event",
|
||||
phase="success",
|
||||
model="test-model",
|
||||
call_type="ocr",
|
||||
litellm_call_id="test-call",
|
||||
metadata={"profile": "success"},
|
||||
kwargs={"model": "test-model"},
|
||||
payload={"model": "test-model", "pages": []},
|
||||
error=None,
|
||||
)
|
||||
python: Final = _execution(user_agent=SENTINEL).model_copy(update={"callbacks": (observation,)})
|
||||
rust: Final = _execution(user_agent="litellm-rust").model_copy(
|
||||
update={"callbacks": (observation.model_copy(update={"kwargs": {"model": "changed"}}),)}
|
||||
)
|
||||
|
||||
with pytest.raises(AssertionError, match=r"\$\.callbacks"):
|
||||
assert_parity(python, rust, SENTINEL)
|
||||
|
||||
|
||||
def test_parity_rejects_error_difference() -> None:
|
||||
python: Final = Execution(
|
||||
requests=(),
|
||||
|
|
|
|||
|
|
@ -18,8 +18,8 @@ CASES: Final[tuple[CaseDefinition, ...]] = (
|
|||
coverage=Coverage.PARTIAL,
|
||||
module="tests.rust-python-harness.strategies.e2e_parity.sdk.ocr.test_sdk_parity",
|
||||
note=(
|
||||
"Recorded sync/async SDK parity; invalid-model provider errors differ, "
|
||||
"and Reducto lacks a Rust contract."
|
||||
"Recorded sync/async SDK parity with focused success/error callback profiles; "
|
||||
"Reducto lacks a Rust contract, and known provider parity gaps remain."
|
||||
),
|
||||
),
|
||||
surface="sdk",
|
||||
|
|
|
|||
|
|
@ -1,23 +1,31 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import datetime
|
||||
import queue
|
||||
import sys
|
||||
import tempfile
|
||||
import time
|
||||
import traceback
|
||||
from collections.abc import Callable, Coroutine, Generator
|
||||
from collections.abc import Callable, Coroutine, Generator, Mapping
|
||||
from contextlib import contextmanager
|
||||
from enum import Enum
|
||||
from functools import partial
|
||||
from pathlib import Path
|
||||
from typing import Annotated, Final, Literal, cast
|
||||
from urllib.parse import urlsplit, urlunsplit
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, JsonValue, TypeAdapter
|
||||
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
|
||||
from .....shared.parity.compare import assert_parity
|
||||
from .....shared.parity.fixtures.store import fixture_id, recorded_fixtures
|
||||
from .....shared.parity.models import (
|
||||
JSON_VALUE_ADAPTER,
|
||||
CallbackObservation,
|
||||
Execution,
|
||||
SDKCommand,
|
||||
SDKError,
|
||||
SDKReport,
|
||||
|
|
@ -39,6 +47,9 @@ from .fixtures.config import configured_fixture_directory
|
|||
from .fixtures.models import OcrParityCase, OcrSdkInput
|
||||
|
||||
API_KEY: Final = "test-key"
|
||||
CALLBACK_DELAY_SECONDS: Final = 0.05
|
||||
CALLBACK_DRAIN_TIMEOUT_SECONDS: Final = 10.0
|
||||
CALLBACK_TERMINALS: Final[tuple[Literal["success", "failure"], ...]] = ("success", "failure")
|
||||
PYTHON_HTTP_SENTINEL: Final = "python-ocr-parity-fallback"
|
||||
PYTHON_VARIANT: Final = ExecutionVariant(name="Python", environment=(("LITELLM_RUST", "0"),))
|
||||
RUST_VARIANT: Final = ExecutionVariant(name="Rust", environment=(("LITELLM_RUST", "1"),))
|
||||
|
|
@ -75,10 +86,148 @@ class InvalidOcrWorkerCase(BaseModel):
|
|||
case: InvalidOcrCase
|
||||
|
||||
|
||||
OcrWorkerCase = Annotated[RecordedOcrWorkerCase | InvalidOcrWorkerCase, Field(discriminator="kind")]
|
||||
class CallbackOcrWorkerCase(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
kind: Literal["callback"] = "callback"
|
||||
case: OcrParityCase
|
||||
terminal: Literal["success", "failure"]
|
||||
|
||||
|
||||
OcrWorkerCase = Annotated[
|
||||
RecordedOcrWorkerCase | InvalidOcrWorkerCase | CallbackOcrWorkerCase,
|
||||
Field(discriminator="kind"),
|
||||
]
|
||||
OCR_WORKER_CASE_ADAPTER: Final[TypeAdapter[OcrWorkerCase]] = TypeAdapter(OcrWorkerCase)
|
||||
|
||||
|
||||
class RecordingCallback(CustomLogger):
|
||||
def __init__(self) -> None:
|
||||
self.message_logging: Final = True
|
||||
self.turn_off_message_logging: Final = False
|
||||
self._observations: Final[queue.SimpleQueue[CallbackObservation]] = queue.SimpleQueue()
|
||||
|
||||
def _normalized_kwargs(self, value: object, key: str | None = None) -> JsonValue:
|
||||
if value is None or isinstance(value, (bool, int, float)):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
if key != "api_base":
|
||||
return value
|
||||
parsed: Final = urlsplit(value)
|
||||
return urlunsplit(("", "", parsed.path, parsed.query, parsed.fragment))
|
||||
if isinstance(value, datetime.datetime):
|
||||
return "datetime"
|
||||
if isinstance(value, Exception):
|
||||
return sdk_error_report(value).model_dump(mode="json")
|
||||
if isinstance(value, BaseModel):
|
||||
return self._normalized_kwargs(value.model_dump(mode="json"), key)
|
||||
if isinstance(value, Mapping):
|
||||
if any(not isinstance(map_key, str) for map_key in value):
|
||||
raise TypeError("callback kwarg mappings must use string keys")
|
||||
return {
|
||||
map_key: self._normalized_kwargs(map_value, map_key)
|
||||
for map_key, map_value in value.items()
|
||||
}
|
||||
if isinstance(value, (list, tuple)):
|
||||
return [self._normalized_kwargs(item) for item in value]
|
||||
raise TypeError(f"unsupported callback kwarg type: {type(value)}")
|
||||
|
||||
def _record(
|
||||
self,
|
||||
hook: Literal[
|
||||
"log_success_event",
|
||||
"async_log_success_event",
|
||||
"log_failure_event",
|
||||
"async_log_failure_event",
|
||||
],
|
||||
phase: Literal["success", "failure"],
|
||||
kwargs: dict[str, object],
|
||||
response_obj: object,
|
||||
) -> None:
|
||||
raw_litellm_params: Final = kwargs.get("litellm_params")
|
||||
litellm_params: Final[Mapping[str, object]] = (
|
||||
cast(Mapping[str, object], raw_litellm_params) if isinstance(raw_litellm_params, Mapping) else {}
|
||||
)
|
||||
raw_metadata: Final = litellm_params.get("metadata")
|
||||
metadata_mapping: Final[Mapping[str, object]] = (
|
||||
cast(Mapping[str, object], raw_metadata) if isinstance(raw_metadata, Mapping) else {}
|
||||
)
|
||||
metadata: Final = JSON_VALUE_ADAPTER.validate_python(
|
||||
{key: metadata_mapping[key] for key in ("callback_profile", "sdk_route") if key in metadata_mapping}
|
||||
)
|
||||
raw_error: Final = kwargs.get("exception")
|
||||
error: Final = sdk_error_report(raw_error) if isinstance(raw_error, Exception) else None
|
||||
normalized_kwargs: Final = self._normalized_kwargs(kwargs)
|
||||
payload_source: Final = (
|
||||
response_obj.model_dump(mode="json") if isinstance(response_obj, BaseModel) else response_obj
|
||||
)
|
||||
payload: Final = JSON_VALUE_ADAPTER.validate_python(payload_source)
|
||||
raw_model: Final = kwargs.get("model")
|
||||
raw_call_type: Final = kwargs.get("call_type")
|
||||
raw_call_id: Final = kwargs.get("litellm_call_id")
|
||||
self._observations.put(
|
||||
CallbackObservation(
|
||||
hook=hook,
|
||||
phase=phase,
|
||||
model=raw_model if isinstance(raw_model, str) else None,
|
||||
call_type=str(raw_call_type) if raw_call_type is not None else None,
|
||||
litellm_call_id=raw_call_id if isinstance(raw_call_id, str) else None,
|
||||
metadata=metadata,
|
||||
kwargs=normalized_kwargs,
|
||||
payload=payload,
|
||||
error=error,
|
||||
)
|
||||
)
|
||||
|
||||
def log_success_event(
|
||||
self,
|
||||
kwargs: dict[str, object],
|
||||
response_obj: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None:
|
||||
del start_time, end_time
|
||||
time.sleep(CALLBACK_DELAY_SECONDS)
|
||||
self._record("log_success_event", "success", kwargs, response_obj)
|
||||
|
||||
async def async_log_success_event(
|
||||
self,
|
||||
kwargs: dict[str, object],
|
||||
response_obj: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None:
|
||||
del start_time, end_time
|
||||
await asyncio.sleep(CALLBACK_DELAY_SECONDS)
|
||||
self._record("async_log_success_event", "success", kwargs, response_obj)
|
||||
|
||||
def log_failure_event(
|
||||
self,
|
||||
kwargs: dict[str, object],
|
||||
response_obj: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None:
|
||||
del start_time, end_time
|
||||
time.sleep(CALLBACK_DELAY_SECONDS)
|
||||
self._record("log_failure_event", "failure", kwargs, response_obj)
|
||||
|
||||
async def async_log_failure_event(
|
||||
self,
|
||||
kwargs: dict[str, object],
|
||||
response_obj: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None:
|
||||
del start_time, end_time
|
||||
await asyncio.sleep(CALLBACK_DELAY_SECONDS)
|
||||
self._record("async_log_failure_event", "failure", kwargs, response_obj)
|
||||
|
||||
def observations(self) -> tuple[CallbackObservation, ...]:
|
||||
observations: Final = tuple(self._observations.get_nowait() for _ in range(self._observations.qsize()))
|
||||
return tuple(sorted(observations, key=lambda observation: observation.hook))
|
||||
|
||||
|
||||
INVALID_OCR_CASES: Final = (
|
||||
InvalidOcrCase(
|
||||
name="unsupported_provider",
|
||||
|
|
@ -229,6 +378,106 @@ def _execute_invalid_sdk_case(
|
|||
return _execute_sdk_call(call_kwargs, route, event_loop)
|
||||
|
||||
|
||||
def _callback_call_id(route: SDKRoute, terminal: Literal["success", "failure"]) -> str:
|
||||
return f"ocr-callback-{route.value}-{terminal}"
|
||||
|
||||
|
||||
def _callback_metadata(route: SDKRoute, terminal: Literal["success", "failure"]) -> dict[str, str]:
|
||||
return {"callback_profile": terminal, "sdk_route": route.value}
|
||||
|
||||
|
||||
def _drain_callback_delivery(route: SDKRoute, event_loop: asyncio.AbstractEventLoop) -> None:
|
||||
if route is SDKRoute.AOCR:
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
|
||||
async def drain_async_callbacks() -> None:
|
||||
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), timeout=CALLBACK_DRAIN_TIMEOUT_SECONDS)
|
||||
await GLOBAL_LOGGING_WORKER.stop()
|
||||
|
||||
event_loop.run_until_complete(drain_async_callbacks())
|
||||
|
||||
from litellm.litellm_core_utils.thread_pool_executor import executor
|
||||
|
||||
executor.shutdown(wait=True, cancel_futures=False)
|
||||
|
||||
|
||||
def _execute_callback_sdk_case(
|
||||
case: OcrParityCase,
|
||||
route: SDKRoute,
|
||||
terminal: Literal["success", "failure"],
|
||||
mock_url: str,
|
||||
event_loop: asyncio.AbstractEventLoop,
|
||||
) -> WorkerSuccess:
|
||||
callback: Final = RecordingCallback()
|
||||
call_kwargs: Final = {
|
||||
**_call_kwargs(case.litellm_input, mock_url, route),
|
||||
"callbacks": [callback],
|
||||
"litellm_call_id": _callback_call_id(route, terminal),
|
||||
"litellm_trace_id": _callback_call_id(route, terminal),
|
||||
"metadata": _callback_metadata(route, terminal),
|
||||
}
|
||||
report: Final = _execute_sdk_call(call_kwargs, route, event_loop)
|
||||
_drain_callback_delivery(route, event_loop)
|
||||
return WorkerSuccess(report=report, callbacks=callback.observations())
|
||||
|
||||
|
||||
def _assert_callback_lifecycle(
|
||||
execution: Execution,
|
||||
route: SDKRoute,
|
||||
terminal: Literal["success", "failure"],
|
||||
) -> None:
|
||||
observations: Final = execution.callbacks
|
||||
assert observations is not None, f"{route.value} {terminal} callbacks were not observed"
|
||||
expected_hooks: Final = (
|
||||
(f"log_{terminal}_event",)
|
||||
if route is SDKRoute.OCR
|
||||
else ("async_log_success_event",)
|
||||
if terminal == "success"
|
||||
else ("async_log_failure_event", "log_failure_event")
|
||||
)
|
||||
actual_hooks: Final = tuple(observation.hook for observation in observations)
|
||||
assert actual_hooks == expected_hooks, (
|
||||
f"{route.value} {terminal} expected callback hooks {expected_hooks}, received {actual_hooks}"
|
||||
)
|
||||
expected_call_id: Final = _callback_call_id(route, terminal)
|
||||
expected_metadata: Final = _callback_metadata(route, terminal)
|
||||
for observation in observations:
|
||||
assert observation.phase == terminal
|
||||
assert observation.call_type == route.value
|
||||
assert observation.litellm_call_id == expected_call_id
|
||||
assert observation.metadata == expected_metadata
|
||||
assert observation.model
|
||||
if terminal == "success":
|
||||
assert isinstance(execution.report, SDKSuccess)
|
||||
assert observation.payload == execution.report.response
|
||||
assert observation.error is None
|
||||
else:
|
||||
assert isinstance(execution.report, SDKError)
|
||||
assert observation.payload is None
|
||||
assert observation.error is not None
|
||||
assert observation.error.exception_type
|
||||
assert observation.error.message
|
||||
assert observation.error.status_code is not None
|
||||
assert observation.error.status_code >= 400
|
||||
|
||||
|
||||
def _check_callback_ocr_sdk_parity(
|
||||
case: OcrParityCase,
|
||||
route: SDKRoute,
|
||||
terminal: Literal["success", "failure"],
|
||||
case_file: Path,
|
||||
runner: SubprocessRunner,
|
||||
) -> None:
|
||||
with execution_worker_pair(runner, PYTHON_VARIANT, RUST_VARIANT) as workers:
|
||||
python_worker, rust_worker = workers
|
||||
python: Final = python_worker.execute(case_file, route.value, case.provider_responses)
|
||||
rust: Final = rust_worker.execute(case_file, route.value, case.provider_responses)
|
||||
|
||||
_assert_callback_lifecycle(python, route, terminal)
|
||||
_assert_callback_lifecycle(rust, route, terminal)
|
||||
assert_parity(python, rust, PYTHON_HTTP_SENTINEL)
|
||||
|
||||
|
||||
def _check_recorded_ocr_sdk_parity(
|
||||
ocr_fixture: OcrParityCase,
|
||||
route: SDKRoute,
|
||||
|
|
@ -276,6 +525,25 @@ def _write_worker_case(directory: Path, index: int, case: OcrWorkerCase) -> Path
|
|||
return case_file
|
||||
|
||||
|
||||
def _callback_fixture(
|
||||
fixtures: tuple[OcrParityCase, ...],
|
||||
terminal: Literal["success", "failure"],
|
||||
) -> OcrParityCase:
|
||||
matching: Final = tuple(
|
||||
fixture
|
||||
for fixture in fixtures
|
||||
if fixture.litellm_input.contract == "mistral"
|
||||
and (
|
||||
all(response.status_code < 400 for response in fixture.provider_responses)
|
||||
if terminal == "success"
|
||||
else any(response.status_code >= 400 for response in fixture.provider_responses)
|
||||
)
|
||||
)
|
||||
if not matching:
|
||||
raise AssertionError(f"no recorded Mistral OCR {terminal} fixture is available for callback parity")
|
||||
return min(matching, key=lambda fixture: fixture_id(fixture.litellm_input, fixture.litellm_input.model))
|
||||
|
||||
|
||||
@contextmanager
|
||||
def parity_checks() -> Generator[tuple[E2ECheck, ...]]:
|
||||
fixtures: Final = tuple(
|
||||
|
|
@ -298,6 +566,17 @@ def parity_checks() -> Generator[tuple[E2ECheck, ...]]:
|
|||
_write_worker_case(directory, len(recorded_files) + index, InvalidOcrWorkerCase(case=case))
|
||||
for index, case in enumerate(INVALID_OCR_CASES)
|
||||
)
|
||||
callback_cases: Final[tuple[tuple[Literal["success", "failure"], OcrParityCase], ...]] = tuple(
|
||||
(terminal, _callback_fixture(fixtures, terminal)) for terminal in CALLBACK_TERMINALS
|
||||
)
|
||||
callback_files: Final = tuple(
|
||||
_write_worker_case(
|
||||
directory,
|
||||
len(recorded_files) + len(invalid_files) + index,
|
||||
CallbackOcrWorkerCase(case=case, terminal=terminal),
|
||||
)
|
||||
for index, (terminal, case) in enumerate(callback_cases)
|
||||
)
|
||||
with execution_worker_pair(runner, PYTHON_VARIANT, RUST_VARIANT) as workers:
|
||||
recorded: Final = tuple(
|
||||
E2ECheck(
|
||||
|
|
@ -315,7 +594,15 @@ def parity_checks() -> Generator[tuple[E2ECheck, ...]]:
|
|||
for case, case_file in zip(INVALID_OCR_CASES, invalid_files, strict=True)
|
||||
for route in SDKRoute
|
||||
)
|
||||
yield (*recorded, *invalid)
|
||||
callbacks: Final = tuple(
|
||||
E2ECheck(
|
||||
f"callback:{route.value}:{terminal}",
|
||||
partial(_check_callback_ocr_sdk_parity, case, route, terminal, case_file, runner),
|
||||
)
|
||||
for (terminal, case), case_file in zip(callback_cases, callback_files, strict=True)
|
||||
for route in SDKRoute
|
||||
)
|
||||
yield (*recorded, *invalid, *callbacks)
|
||||
|
||||
|
||||
def _execute_worker_command(
|
||||
|
|
@ -333,6 +620,8 @@ def _execute_worker_command(
|
|||
return WorkerSuccess(report=_execute_sdk_case(recorded.litellm_input, route, mock_url, event_loop))
|
||||
case InvalidOcrWorkerCase(case=invalid):
|
||||
return WorkerSuccess(report=_execute_invalid_sdk_case(invalid, route, mock_url, event_loop))
|
||||
case CallbackOcrWorkerCase(case=callback_case, terminal=terminal):
|
||||
return _execute_callback_sdk_case(callback_case, route, terminal, mock_url, event_loop)
|
||||
except Exception:
|
||||
return WorkerFailure(error=traceback.format_exc())
|
||||
|
||||
|
|
|
|||
|
|
@ -230,6 +230,10 @@ _HOST_ONLY_BRIDGE_EXCLUSIONS: Final = tuple(
|
|||
"test_ocr_exception_type_uses_resolved_provider_context",
|
||||
"Python wraps bridge exceptions into public errors.",
|
||||
),
|
||||
(
|
||||
"test_rust_upstream_error_uses_ocr_provider_error_mapping",
|
||||
"Python maps native upstream errors through the selected OCR provider config.",
|
||||
),
|
||||
("test_aocr_routes_to_async_rust_when_enabled", "Python selects and invokes the async native bridge."),
|
||||
("test_aocr_exception_type_uses_resolved_provider_context", "Python wraps async bridge exceptions."),
|
||||
("test_ocr_forwards_timeout_to_rust", "Python converts and forwards explicit timeouts."),
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ import httpx
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.rust_bridge import configuration
|
||||
|
||||
|
|
@ -39,6 +40,10 @@ class CapturedException(Exception):
|
|||
pass
|
||||
|
||||
|
||||
class RustUpstreamError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class RecordingBridge:
|
||||
"""A fake ``RustOcr`` callable that records the args it was handed."""
|
||||
|
||||
|
|
@ -183,6 +188,9 @@ class FakeOCRConfig:
|
|||
) -> str:
|
||||
return f"{api_base or 'https://api.mistral.ai/v1'}/ocr"
|
||||
|
||||
def get_error_class(self, error_message: str, status_code: int, headers: dict[str, str]) -> BaseLLMException:
|
||||
return BaseLLMException(status_code=status_code, message=error_message, headers=headers)
|
||||
|
||||
|
||||
def build_prepared_request(
|
||||
*,
|
||||
|
|
@ -486,6 +494,20 @@ def test_run_rust_ocr_prepares_request_and_wraps_response():
|
|||
}
|
||||
|
||||
|
||||
def test_rust_upstream_error_uses_ocr_provider_error_mapping():
|
||||
error = RustUpstreamError(400, '{"message":"invalid model"}')
|
||||
|
||||
mapped = ocr_main._map_rust_ocr_error(
|
||||
error,
|
||||
build_prepared_request(),
|
||||
(RuntimeError, RustUpstreamError),
|
||||
)
|
||||
|
||||
assert isinstance(mapped, BaseLLMException)
|
||||
assert mapped.status_code == 400
|
||||
assert mapped.message == '{"message":"invalid model"}'
|
||||
|
||||
|
||||
def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing():
|
||||
bridge = RecordingBridge()
|
||||
litellm.rust(True)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue