mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
refactor(ocr): route native requests through core (#40532)
* refactor(ocr): route native Mistral through core * fix(ocr): preserve Azure API base resolution * chore(ocr): document bridge boundary casts * fix(ocr): keep Azure environment resolution in Rust * fix(ocr): centralize native execution and isolate request logging * refactor(ocr): narrow native migration to bridge routing --------- Co-authored-by: Stack Plan <stack-plan@example.invalid>
This commit is contained in:
parent
83616c0e09
commit
89f1f9567d
4 changed files with 187 additions and 32 deletions
|
|
@ -2,6 +2,7 @@ use litellm_core::Error;
|
|||
use std::future::Future;
|
||||
|
||||
use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr};
|
||||
use litellm_core::ocr::wire::{OcrWireRequest, decode_request, is_supported_request};
|
||||
use pyo3::prelude::*;
|
||||
use serde_json::Value;
|
||||
|
||||
|
|
@ -31,6 +32,21 @@ fn prepare_ocr(
|
|||
extra_headers,
|
||||
timeout,
|
||||
} = options;
|
||||
if is_supported_request(&model, custom_llm_provider.as_deref()) {
|
||||
let request = decode_request(OcrWireRequest {
|
||||
model,
|
||||
document,
|
||||
api_key,
|
||||
api_base,
|
||||
custom_llm_provider,
|
||||
extra_headers,
|
||||
optional_params,
|
||||
timeout_seconds: timeout.map(|value| value.as_secs_f64()),
|
||||
})?;
|
||||
return litellm_core::ocr::ocr(request)
|
||||
.await
|
||||
.map(|response| response.into_json());
|
||||
}
|
||||
run_ocr(OcrRequest {
|
||||
model: &model,
|
||||
document,
|
||||
|
|
|
|||
|
|
@ -1,8 +1,9 @@
|
|||
import json
|
||||
|
||||
import litellm
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
|
||||
|
||||
def _reducto_parse_response() -> dict:
|
||||
return {
|
||||
|
|
@ -68,15 +69,11 @@ def disable_aiohttp_transport():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_parse_v3_file_upload_and_response_mapping(
|
||||
disable_aiohttp_transport, respx_mock
|
||||
):
|
||||
async def test_parse_v3_file_upload_and_response_mapping(disable_aiohttp_transport, respx_mock):
|
||||
upload_route = respx_mock.post("https://platform.reducto.ai/upload").respond(
|
||||
json={"file_id": "reducto://uploaded.pdf"}
|
||||
)
|
||||
parse_route = respx_mock.post("https://platform.reducto.ai/parse").respond(
|
||||
json=_reducto_parse_response()
|
||||
)
|
||||
parse_route = respx_mock.post("https://platform.reducto.ai/parse").respond(json=_reducto_parse_response())
|
||||
|
||||
response = await litellm.aocr(
|
||||
model="reducto/parse-v3",
|
||||
|
|
@ -123,15 +120,11 @@ async def test_parse_v3_file_upload_and_response_mapping(
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_parse_v3_reducto_id_passthrough_skips_upload(
|
||||
disable_aiohttp_transport, respx_mock
|
||||
):
|
||||
async def test_parse_v3_reducto_id_passthrough_skips_upload(disable_aiohttp_transport, respx_mock):
|
||||
upload_route = respx_mock.post("https://platform.reducto.ai/upload").respond(
|
||||
json={"file_id": "reducto://should-not-upload.pdf"}
|
||||
)
|
||||
parse_route = respx_mock.post("https://platform.reducto.ai/parse").respond(
|
||||
json=_reducto_parse_response()
|
||||
)
|
||||
parse_route = respx_mock.post("https://platform.reducto.ai/parse").respond(json=_reducto_parse_response())
|
||||
|
||||
response = await litellm.aocr(
|
||||
model="reducto/parse-v3",
|
||||
|
|
|
|||
|
|
@ -227,12 +227,7 @@ async def exercise_async(native: object, api_base: str) -> None:
|
|||
|
||||
async def exercise_async_concurrency(native: object, api_base: str) -> None:
|
||||
responses: Final = await asyncio.wait_for(
|
||||
asyncio.gather(
|
||||
*(
|
||||
native.amessages(**route_kwargs("messages", api_base, "success"))
|
||||
for _ in range(32)
|
||||
)
|
||||
),
|
||||
asyncio.gather(*(native.amessages(**route_kwargs("messages", api_base, "success")) for _ in range(32))),
|
||||
timeout=15,
|
||||
)
|
||||
for response in responses:
|
||||
|
|
|
|||
|
|
@ -1,33 +1,48 @@
|
|||
import json
|
||||
import threading
|
||||
from collections.abc import Generator
|
||||
from dataclasses import dataclass
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.rust_bridge import ocr as rust_ocr_bridge
|
||||
|
||||
pytestmark = pytest.mark.requires_rust_extension
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RecordedOCRRequest:
|
||||
body: object
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def ocr_server() -> Generator[tuple[ThreadingHTTPServer, list[RecordedOCRRequest]]]:
|
||||
requests: Final[list[RecordedOCRRequest]] = []
|
||||
def ocr_server() -> Generator[tuple[ThreadingHTTPServer, list[dict[str, object]]]]:
|
||||
requests: Final[list[dict[str, object]]] = []
|
||||
|
||||
class Handler(BaseHTTPRequestHandler):
|
||||
def do_POST(self) -> None:
|
||||
requests.append(
|
||||
RecordedOCRRequest(
|
||||
body=json.loads(self.rfile.read(int(self.headers["Content-Length"]))),
|
||||
)
|
||||
{
|
||||
"headers": {name.lower(): value for name, value in self.headers.items()},
|
||||
"body": json.loads(self.rfile.read(int(self.headers["Content-Length"]))),
|
||||
}
|
||||
)
|
||||
if self.headers.get("x-test-stall") == "true":
|
||||
self.connection.settimeout(2)
|
||||
try:
|
||||
self.rfile.read(1)
|
||||
except TimeoutError:
|
||||
pass
|
||||
return
|
||||
if self.headers.get("User-Agent", "").startswith("python-httpx"):
|
||||
self.send_response(418)
|
||||
self.end_headers()
|
||||
return
|
||||
status = int(self.headers.get("x-test-status", "200"))
|
||||
if status != 200:
|
||||
body = b'{"error":"provider unavailable"}'
|
||||
self.send_response(status)
|
||||
self.send_header("Content-Length", str(len(body)))
|
||||
self.end_headers()
|
||||
self.wfile.write(body)
|
||||
return
|
||||
response: Final = json.dumps(
|
||||
{
|
||||
"pages": [{"index": 0, "markdown": "native OCR response", "images": [], "dimensions": None}],
|
||||
|
|
@ -56,7 +71,7 @@ def ocr_server() -> Generator[tuple[ThreadingHTTPServer, list[RecordedOCRRequest
|
|||
|
||||
|
||||
def test_native_ocr_with_compiled_rust_extension(
|
||||
ocr_server: tuple[ThreadingHTTPServer, list[RecordedOCRRequest]],
|
||||
ocr_server: tuple[ThreadingHTTPServer, list[dict[str, object]]],
|
||||
) -> None:
|
||||
server, requests = ocr_server
|
||||
address: Final = server.server_address
|
||||
|
|
@ -77,7 +92,143 @@ def test_native_ocr_with_compiled_rust_extension(
|
|||
assert response is not None
|
||||
assert response["pages"][0]["markdown"] == "native OCR response"
|
||||
assert len(requests) == 1
|
||||
assert requests[0].body == {
|
||||
assert not requests[0]["headers"].get("user-agent", "").startswith("python-httpx")
|
||||
assert requests[0]["body"] == {
|
||||
"model": "mistral-ocr-latest",
|
||||
"document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("asynchronous", [False, True])
|
||||
@pytest.mark.parametrize("model", ["mistral/mistral-ocr-latest", "azure_ai/doc-intelligence/prebuilt-read"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_public_ocr_matches_python(model, asynchronous):
|
||||
import json
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from threading import Thread
|
||||
from typing import Final
|
||||
from urllib.parse import parse_qsl, urlsplit
|
||||
|
||||
from litellm.rust_bridge import _native
|
||||
|
||||
assert callable(_native.ocr)
|
||||
calls: Final = []
|
||||
|
||||
class Handler(BaseHTTPRequestHandler):
|
||||
def do_POST(self):
|
||||
body: Final = json.loads(self.rfile.read(int(self.headers["Content-Length"])))
|
||||
target: Final = urlsplit(self.path)
|
||||
calls.append(
|
||||
(
|
||||
target.path,
|
||||
parse_qsl(target.query),
|
||||
self.headers.get("Authorization"),
|
||||
self.headers.get("Ocp-Apim-Subscription-Key"),
|
||||
body,
|
||||
)
|
||||
)
|
||||
payload: Final = (
|
||||
{"status": "succeeded", "analyzeResult": {"pages": []}}
|
||||
if "doc-intelligence" in model
|
||||
else {"pages": [{"index": 0, "markdown": "hello"}]}
|
||||
)
|
||||
encoded: Final = json.dumps(payload).encode()
|
||||
self.send_response(200)
|
||||
self.send_header("Content-Type", "application/json")
|
||||
self.send_header("Content-Length", str(len(encoded)))
|
||||
self.end_headers()
|
||||
self.wfile.write(encoded)
|
||||
|
||||
def log_message(self, *_args):
|
||||
pass
|
||||
|
||||
server: Final = ThreadingHTTPServer(("127.0.0.1", 0), Handler)
|
||||
thread: Final = Thread(target=server.serve_forever, daemon=True)
|
||||
thread.start()
|
||||
responses: Final = []
|
||||
try:
|
||||
for enabled in (False, True):
|
||||
litellm.rust(enabled)
|
||||
arguments: Final = {
|
||||
"model": model,
|
||||
"document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
|
||||
"api_key": "test-key",
|
||||
"api_base": f"http://127.0.0.1:{server.server_port}",
|
||||
"pages": [0, 2],
|
||||
"timeout": 3.0,
|
||||
}
|
||||
response: Final = await litellm.aocr(**arguments) if asynchronous else litellm.ocr(**arguments)
|
||||
responses.append(response.model_dump())
|
||||
assert len(calls) == 2
|
||||
assert calls[0] == calls[1]
|
||||
for key in ("model", "pages", "object"):
|
||||
assert responses[0][key] == responses[1][key]
|
||||
finally:
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
thread.join(timeout=3)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("asynchronous", [False, True])
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_ocr_failures_do_not_retry_on_python(ocr_server, asynchronous):
|
||||
server, requests = ocr_server
|
||||
arguments = {
|
||||
"model": "mistral-ocr-latest",
|
||||
"custom_llm_provider": "mistral",
|
||||
"document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
|
||||
"api_key": "test-key",
|
||||
"api_base": f"http://127.0.0.1:{server.server_port}",
|
||||
"extra_headers": {"x-test-status": "503"},
|
||||
"num_retries": 0,
|
||||
}
|
||||
litellm.rust(True)
|
||||
with pytest.raises(litellm.ServiceUnavailableError) as caught:
|
||||
await litellm.aocr(**arguments) if asynchronous else litellm.ocr(**arguments)
|
||||
assert caught.value.status_code == 503
|
||||
assert len(requests) == 1
|
||||
assert not requests[0]["headers"].get("user-agent", "").startswith("python-httpx")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("custom_provider", ["mistral", "not-a-provider"])
|
||||
def test_native_ocr_rejects_invalid_input_before_network(ocr_server, custom_provider):
|
||||
from litellm.rust_bridge import _native
|
||||
|
||||
server, requests = ocr_server
|
||||
with pytest.raises(ValueError, match=r"invalid (OCR request field|provider)|invalid request"):
|
||||
_native.ocr(
|
||||
model="mistral-ocr-latest",
|
||||
custom_llm_provider=custom_provider,
|
||||
document={"type": "document_url"},
|
||||
api_key="test-key",
|
||||
api_base=f"http://127.0.0.1:{server.server_port}",
|
||||
)
|
||||
assert requests == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize("asynchronous", [False, True])
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_ocr_enforces_request_deadline_without_fallback(ocr_server, asynchronous):
|
||||
import asyncio
|
||||
import time
|
||||
|
||||
server, requests = ocr_server
|
||||
litellm.rust(True)
|
||||
arguments = {
|
||||
"model": "mistral/mistral-ocr-latest",
|
||||
"document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
|
||||
"api_key": "test-key",
|
||||
"api_base": f"http://127.0.0.1:{server.server_port}",
|
||||
"extra_headers": {"x-test-stall": "true"},
|
||||
"timeout": 0.1,
|
||||
"num_retries": 0,
|
||||
}
|
||||
started = time.monotonic()
|
||||
with pytest.raises(litellm.APIConnectionError):
|
||||
await asyncio.wait_for(
|
||||
litellm.aocr(**arguments) if asynchronous else asyncio.to_thread(litellm.ocr, **arguments),
|
||||
timeout=3,
|
||||
)
|
||||
assert 0.09 <= time.monotonic() - started < 3
|
||||
assert len(requests) == 1
|
||||
assert not requests[0]["headers"].get("user-agent", "").startswith("python-httpx")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue