Merge pull request #31488 from BerriAI/litellm_rust_ocr_e2e_llm_translation
Some checks are pending
GitHub Actions Security Analysis / zizmor (push) Waiting to run

This commit is contained in:
mubashir1osmani 2026-06-28 08:14:13 -07:00 • committed by GitHub
commit b443037783
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 437 additions and 161 deletions

148
tests/e2e/CLAUDE.md Normal file
View file

@ -0,0 +1,148 @@
# e2e harness conventions
Code-style rules for writing tests under `tests/e2e/`. The harness already encodes the plumbing; your job is the feature-specific behavior, not reinventing it. For what a complete test must do (the lifecycle contract, asserting both recorded state and enforced behavior) and how to run a suite, see `CONTRIBUTING.md` in this directory. Repo-wide conventions live in the root `CLAUDE.md`
## Suite folders
Each subdirectory under `tests/e2e/` is one suite, scoped to an endpoint family or behavior area. If you add a new folder, you must add a line here describing what kind of tests belong in it, so the layout stays self-describing. `gateway/` is the exception: it holds proxy configuration only and never tests
- `llm_translation/` - LLM endpoint and provider-translation behavior: passthrough, custom pricing, OCR
- `embeddings/` - the `/embeddings` endpoint across providers
- `batches/` - the `/batches` endpoint (placeholder until the first test lands)
- `realtime/` - realtime websocket sessions, including the pipecat audio path
- `budgets/` - budget definition, enforcement, and reset windows (key, team, tag, soft, multi-window)
- `spend_tracking/` - spend logging and cost attribution on `/spend/*`
- `models_mgmt/` - model-management routes (add/update, tpm persistence)
- `logging/` - logging-integration delivery (datadog and friends)
- `security/` - secret handling and log-leak protection
- `router/` - routing and reliability behavior (rate limits, fallbacks, cooldowns)
- `gateway/` - proxy configuration only (`litellm-config.yml`); no tests
## Lay the pattern down in a class
Keep the cases for one feature inside a class so the file reads as a spec for how that feature behaves in production. The class name says what is under test; each method is one behavior. Think of it as documenting the contract, with the rough intent being
```python
# pseudo-code to convey intent, not the real API
class TestPromptCompression:
def test_prompt_compression_add_to_virtual_key(self):
new_key = self.resources.create_key(user_id, compression=True) # turn the feature on
resources._defer(new_key) # queue key deletion
def test_prompt_compression_accumulate_spend(self, key_id, user_id):
for _ in range(10):
response = self.resources.gateway.post("gemini-2.5-flash", key_id, user_id)
compressed_value = ...
assert response.cost == compressed_value # the cost was actually reduced
```
That snippet only conveys intent. What you actually write uses the real harness: the `client` fixture for your suite, the `scoped_key` fixture for an auto-deleted key, typed pydantic bodies from `models.py`, and `unwrap(...)` on the tagged-union result. `tests/e2e/llm_translation/test_custom_pricing_e2e.py` is the reference to copy from; it creates a scoped key, drives a real gemini call, polls `/spend/logs` to a deadline for the cost-breakdown row, then asserts the input and output costs match the configured custom rates and that a sibling deployment kept its own price. Read it before writing yours
## Use the shared transport; never touch requests directly
Every HTTP call goes through the shared transport, never through `requests.*` in a test. `e2e_http.py` is the only module permitted to call `requests.*`, and that is enforced in CI by `tests/code_coverage_tests/check_e2e_no_raw_requests.py`. A test that imports requests will fail the check
The shape is layered so tests stay declarative
`transport.py` exposes a `Transport` Protocol with `post`, `get`, `delete`, `send`, `stream`, `probe`, plus `bearer(key)` and the `master` header. `HttpTransport` fulfils it, and `SplitTransport` routes each call by path to the data plane or the control plane so a split control-plane/data-plane deployment works without any change in the test
`e2e_gateway.py` holds `Gateway`, a frozen dataclass that wraps a `Transport` and adds the operations tests reuse: `generate_key` / `delete_key` / `key_info`, `model_info`, the LLM calls `chat` / `chat_stream` / `embed` / `ocr`, the spend read-back `spend_logs`, and the poll helpers `poll_logs_for_key` / `poll_logs_for_request_id` that loop to `poll_timeout` instead of sleeping once. Add a new route as a method here so other suites get it for free
Each suite provides its own `client` fixture (see `llm_translation/passthrough_client.py`), a frozen dataclass that holds the shared `Gateway` and adds suite-specific routes. Cleanup runs through that same `Gateway`, so whatever keys or customers your test creates get torn down by the `resources` fixture
Request and response bodies are typed pydantic models in `models.py`; only the fields a test reads are modelled, and nothing passes raw dicts. Outcomes come back as a `Result[R]` tagged union (`Success`, `NetworkError`, `UnauthorizedError`, `RateLimitedError`, `ValidationError`, `UnknownApiError`). Handle them with `match`, or call `unwrap(...)` when a non-success should fail the test. The skip-vs-fail split is deliberate: a test marked `e2e` skips when no proxy answers its liveness probe, but once a request reaches the proxy any wrong behavior is a hard failure, never a skip
Mark live tests with `@pytest.mark.e2e` (on the class or the module). Pure coverage of the harness itself carries no marker and runs regardless. Use `scoped_key` for a fresh all-models key that auto-deletes, `resources` when you need to create and tear down more than a key, and `unique_marker()` from `e2e_config` to keep prompts, tags, and customer ids from colliding across concurrent runs and the shared response cache
## Typing
The harness is fully typed and new code must not add `Any` or widen the basedpyright budgets. When a response field is untyped, model it in `models.py` (just the fields you read) and let pydantic validate it, rather than threading a `dict` or `Any` through the test
## Coverage registry
The set of tests we want is a registry checked into this repo, one row per behavior; that file is the definition of done and the denominator. Each e2e test declares what it covers with `@pytest.mark.covers("...")`, and a small collector diffs the registry against the tests and ships coverage to the existing Grafana. No Allure, no new dependencies
Coverage is organized as module > feature > test. There are six modules: LLMs, MCPs, Management/UI, Reliability & Performance, Logging & Guardrails, and Other. A feature is either an endpoint (`/chat/completions`) or a behavior (fallbacks, rate limits; config-driven, with no route of its own). A cell reads like `llm.chat_completions.bedrock_converse.tool_use.stream.works`
The metric is coverage: the share of registry rows that have a passing covering test, reported to Grafana per module so a gap surfaces as an uncovered row rather than a silent absence
### Naming grammar per module
LLMs - endpoint features (subject = the route), seeded from the Claude Code compat matrix
```
llm.<endpoint>.<route>.<capability>.<streaming>.<assertion>
endpoint : chat_completions | messages | responses | embeddings | batches | files
| rerank | images_generations | audio_speech | audio_transcriptions | moderations
route : openai | azure_openai | anthropic | bedrock_invoke | bedrock_converse | vertex | azure_foundry
(vocab varies per endpoint; messages is anthropic-format only)
capability : basic | tool_use | prompt_cache_5m | prompt_cache_1h | vision | thinking
| thinking_tool_use | pdf_input | web_search | structured_output | count_tokens
| tool_search | long_context_1m
streaming : stream | nonstream (omit where n/a)
assertion : works | cost_logged
label (not in id): model = haiku-4.5 | sonnet-4.6 | opus-4.7 | gpt-*
e.g. llm.chat_completions.bedrock_converse.tool_use.stream.works
llm.messages.anthropic.prompt_cache_1h.nonstream.cache_hit
```
Management / UI - endpoint features (surface tag: api | ui)
```
mgmt.<endpoint>.<assertion>
endpoint : key.generate | key.update | key.delete | team.new | user.new
| budget.new | model.add | ... (one per management route)
assertion : persists | member_forbidden | admin_only | happy_path
e.g. mgmt.key.generate.persists (surface=api)
mgmt.key.generate.happy_path (surface=ui)
```
MCPs - endpoint features with the protocol op as the variant
```
mcp.<operation>.<auth_family>.<assertion>
operation : list_tools | call_tool | list_resources | read_resource | list_prompts | get_prompt
auth_family : none | api_key | bearer | oauth
assertion : succeeds | denied_without_permission
e.g. mcp.call_tool.oauth.succeeds
```
Reliability & Performance - behavior features (no route; endpoint is exercised_on)
```
reliability.<behavior>.<variant>.<assertion>
behavior : fallback | retry | cooldown | timeout | ratelimit | routing | cache | circuit_breaker | perf
variant : <trigger> 5xx | context_window | content_policy | 429 | timeout
<strategy> simple_shuffle | usage_based | latency_based | cost_based | least_busy
<dimension> latency | throughput (perf only; SLO/threshold assertion, not binary)
assertion : routes_to_fallback | succeeds_within_retries | picks_under_tpm | returns_cached
| trips_then_recovers | under_slo
e.g. reliability.fallback.context_window.routes_to_fallback exercised_on=[chat_completions]
reliability.ratelimit.rpm.blocks_over_limit exercised_on=[chat_completions, messages]
```
Logging & Guardrails - behavior features (config-driven; endpoint is exercised_on)
```
logging.<integration>.<event>.<assertion>
integration : langfuse | s3 | otel | prometheus | datadog | ...
event : success | failure | stream
assertion : logs_spend | writes_object | exports_metric
e.g. logging.langfuse.success.logs_spend exercised_on=[chat_completions]
guardrail.<provider>.<hook_point>.<assertion>
provider : presidio | lakera | bedrock | aporia | ...
hook_point : pre_call | post_call | during | logging_only
assertion : blocks | masks | allows
e.g. guardrail.presidio.pre_call.masks exercised_on=[chat_completions]
```
Other - holding pen (endpoint or behavior)
```
other.<area>.<case>.<assertion>
area : auth | lifecycle | config | ...
rule : audited periodically; a cluster here promotes to a new component
e.g. other.auth.jwt.valid_token_allows
other.lifecycle.readiness.reports_db
```

129
tests/e2e/CONTRIBUTING.md Normal file
View file

@ -0,0 +1,129 @@
# Contributors Guide
This directory holds the live end-to-end suites that prove product correctness against a real running proxy and real provider APIs. The goal of this guide is simple: when you ship a feature, you add e2e coverage that walks that feature the way production does, across every route and edge case it touches, so a later change that breaks it fails here first
Read this before adding a test and i recommend reading through CLAUDE.md
When contributing to this directory, please first discuss the change you wish to make via issue or pull request. We require screenshots and proof of your tests working on a live proxy.
## Setup
The suites run against a live proxy, so bring one up first. `docker-compose.yml` here starts that proxy with its Postgres and Redis, serving `gateway/litellm-config.yml`; add any model, pricing override, or guardrail your test needs to that file and read it back in the test rather than hardcoding values. `gateway/` holds proxy configuration only, so never put tests there
## Running the tests locally
1. Create a .env file and add provider keys:
```bash
OPENAI_API_KEY="sk-..."
ANTHROPIC_API_KEY="sk-..."
2. Bring the stack up from this directory:
```bash
docker compose up -d
curl -fs http://localhost:4000/health/liveliness
```
3. Run a suite against it; the harness reads `LITELLM_PROXY_URL` (default `http://localhost:4000`):
```bash
uv run pytest tests/e2e/llm_translation/ -v
```
4. Tear it down when you're done:
```bash
docker compose down -v
```
Tests marked `@pytest.mark.e2e` skip when no proxy answers `/health/liveliness`, so a run that reports everything skipped means the stack isn't up, not that anything passed
## What a complete test looks like
A feature test is complete only when it walks the feature end to end, in this order
1. CREATE the resource (key / team / budget / ...) and immediately queue its deletion
2. CONFIGURE the feature's setting on it (assign the budget, turn on compression, set the limit)
3. ACT; drive real traffic through the gateway exactly like prod does (right model, real auth headers, enough calls to actually trigger the behavior)
4. SETTLE; poll the DB / spend logs until the write lands. Writes are eventually consistent (spend flushes on proxy_batch_write_at, ~60s), so poll to a deadline. Never sleep once
5. ASSERT the recorded state the feature promises (spend > budget, cost reduced, tag attributed, ...)
6. ASSERT the enforced behavior the gateway returns (429 budget_exceeded, block, refusal, ...)
7. TEARDOWN; every resource you created is deleted
### The one rule that makes it complete
It must assert BOTH sides: the recorded state (step 5) AND the enforced behavior (step 6)
A test that only checks "the call went through", or only checks spend without checking the 429, is not complete; it is checking plumbing, not the product promise
### Example: budget enforcement
```
create a key -> (1)
assign a budget -> (2)
send a bunch of calls -> (3)
poll for db spend -> (4)
assert spend > budget -> (5)
assert status_code == 429 -> (6) ("budget_exceeded")
key auto-deleted on teardown -> (7)
```
### The skeleton every test fills in
```
setup -> create the resource + queue cleanup
configure -> apply the feature's knob
act -> send real calls like production
settle -> poll the DB until the write lands
assert -> recorded state is correct (the feature happened)
assert -> gateway enforced it (the product promise held)
teardown -> delete everything you created
```
If a step is missing, the test is not done. That is the whole pattern
## Style: lay the pattern down in a class
Keep the cases for one feature inside a class so the file reads as a spec for how that feature behaves in production. The class name says what is under test; each method is one behavior. Think of it as documenting the contract, with the rough intent being
```python
# pseudo-code to convey intent
class TestPromptCompression:
def test_prompt_compression_add_to_virtual_key(self):
new_key = self.resources.create_key(user_id, compression=True) # turn the feature on
resources._defer(new_key) # queue key deletion
def test_prompt_compression_accumulate_spend(self, key_id, user_id):
for _ in range(10):
response = self.resources.gateway.post("gemini-2.5-flash", key_id, user_id)
compressed_value = ...
assert response.cost == compressed_value # the cost was actually reduced
```
That snippet only conveys intent. What you actually write uses the real harness: the `client` fixture for your suite, the `scoped_key` fixture for an auto-deleted key, typed pydantic bodies from `models.py`, and `unwrap(...)` on the tagged-union result. `tests/e2e/llm_translation/test_custom_pricing_e2e.py` is the reference to copy from; it creates a scoped key, drives a real gemini call, polls `/spend/logs` to a deadline for the cost-breakdown row, then asserts the input and output costs match the configured custom rates and that a sibling deployment kept its own price. Read it before writing yours
## Use the shared transport; never touch requests directly
Every HTTP call goes through the shared transport, never through `requests.*` in a test. `e2e_http.py` is the only module permitted to call `requests.*`, and that is enforced in CI by `tests/code_coverage_tests/check_e2e_no_raw_requests.py`. A test that imports requests will fail the check
The shape is layered so tests stay declarative
`transport.py` exposes a `Transport` Protocol with `post`, `get`, `delete`, `send`, `stream`, `probe`, plus `bearer(key)` and the `master` header. `HttpTransport` fulfils it, and `SplitTransport` routes each call by path to the data plane or the control plane so a split control-plane/data-plane deployment works without any change in the test
`e2e_gateway.py` holds `Gateway`, a frozen dataclass that wraps a `Transport` and adds the operations tests reuse: `generate_key` / `delete_key` / `key_info`, `model_info`, the LLM calls `chat` / `chat_stream` / `embed` / `ocr`, the spend read-back `spend_logs`, and the poll helpers `poll_logs_for_key` / `poll_logs_for_request_id` that loop to `poll_timeout` instead of sleeping once. Add a new route as a method here so other suites get it for free
Each suite provides its own `client` fixture (see `llm_translation/passthrough_client.py`), a frozen dataclass that holds the shared `Gateway` and adds suite-specific routes. Cleanup runs through that same `Gateway`, so whatever keys or customers your test creates get torn down by the `resources` fixture
Request and response bodies are typed pydantic models in `models.py`; only the fields a test reads are modelled, and nothing passes raw dicts. Outcomes come back as a `Result[R]` tagged union (`Success`, `NetworkError`, `UnauthorizedError`, `RateLimitedError`, `ValidationError`, `UnknownApiError`). Handle them with `match`, or call `unwrap(...)` when a non-success should fail the test. The skip-vs-fail split is deliberate: a test marked `e2e` skips when no proxy answers its liveness probe, but once a request reaches the proxy any wrong behavior is a hard failure, never a skip
Mark live tests with `@pytest.mark.e2e` (on the class or the module). Pure coverage of the harness itself carries no marker and runs regardless. Use `scoped_key` for a fresh all-models key that auto-deletes, `resources` when you need to create and tear down more than a key, and `unique_marker()` from `e2e_config` to keep prompts, tags, and customer ids from colliding across concurrent runs and the shared response cache
## Pre-commit steps
Before you push
- Run basedpyright over your changes; the harness is fully typed and new code must not add `Any` or widen the budgets
- Bring the stack up with docker-compose from this directory and run your suite locally against it, so you exercise the same skip-vs-fail path CI does
- Use the config at `tests/e2e/gateway/litellm-config.yml` if your feature needs a model, pricing override, guardrail, or other proxy setting declared up front; add the deployment there and read it back in the test rather than hardcoding values
- Capture screenshots of the tests passing and attach them to the PR as proof of fix

View file

@ -34,6 +34,8 @@ from models import (
KeyInfoResponse,
ModelInfoEntry,
ModelInfoResponse,
OcrBody,
OcrResponse,
SpendLogRow,
SpendLogs,
SpendLogsParams,
@ -120,9 +122,7 @@ class Gateway:
)
def chat_stream(self, key: str, body: ChatBody) -> StreamingResponse:
return self.transport.stream(
"/chat/completions", headers=self.transport.bearer(key), json=body
)
return self.transport.stream("/chat/completions", headers=self.transport.bearer(key), json=body)
def embed(self, key: str, body: EmbedBody) -> Result[EmbedResponse]:
return self.transport.post(
@ -132,6 +132,14 @@ class Gateway:
response_type=EmbedResponse,
)
def ocr(self, key: str, body: OcrBody) -> Result[OcrResponse]:
return self.transport.post(
"/v1/ocr",
headers=self.transport.bearer(key),
json=body,
response_type=OcrResponse,
)
# ---- spend read-back ------------------------------------------------
def spend_logs(self, params: SpendLogsParams) -> list[SpendLogRow]:
@ -150,9 +158,7 @@ class Gateway:
def poll_logs_for_key(
self, key: str, *, min_rows: int = 1, predicate: RowsPredicate | None = None
) -> list[SpendLogRow]:
return self._poll(
lambda: self.spend_logs(SpendLogsParams(api_key=key)), min_rows, predicate
)
return self._poll(lambda: self.spend_logs(SpendLogsParams(api_key=key)), min_rows, predicate)
def poll_logs_for_request_id(
self,

View file

@ -1,148 +0,0 @@
"""
Gateway E2E smoke for Rust-backed OCR.
Start the proxy with:
LITELLM_USE_RUST_OCR=1 litellm --config tests/e2e/gateway/litellm-config.yml --port 4000
"""
from __future__ import annotations
import os
from dataclasses import dataclass
from pathlib import Path
from typing import Any
import httpx
import pytest
import yaml
TEST_PDF_URL = (
"https://cdn.jsdelivr.net/gh/BerriAI/litellm"
"@d769e81c90d453240c61fc572cdb27fae06a89d0"
"/tests/llm_translation/fixtures/dummy.pdf"
)
TEST_IMAGE_URL = (
"https://cdn.jsdelivr.net/gh/BerriAI/litellm"
"@d769e81c90d453240c61fc572cdb27fae06a89d0"
"/tests/image_gen_tests/test_image.png"
)
RUST_OCR_GATEWAY_CASES = [
pytest.param(
"rust-ocr-mistral",
{"type": "document_url", "document_url": TEST_PDF_URL},
id="mistral",
),
pytest.param(
"rust-ocr-azure-ai",
{"type": "document_url", "document_url": TEST_PDF_URL},
id="azure_ai",
),
pytest.param(
"rust-ocr-azure-document-intelligence",
{"type": "document_url", "document_url": TEST_PDF_URL},
id="azure_document_intelligence",
),
pytest.param(
"rust-ocr-vertex-mistral",
{"type": "document_url", "document_url": TEST_PDF_URL},
id="vertex_mistral",
),
pytest.param(
"rust-ocr-vertex-deepseek",
{
"type": "image_url",
"image_url": os.getenv("RUST_OCR_IMAGE_URL", TEST_IMAGE_URL),
},
id="vertex_deepseek",
),
]
CONFIG_PATH = Path(__file__).with_name("litellm-config.yml")
@dataclass(frozen=True)
class OcrGateway:
base_url: str
master_key: str
def model_names(self) -> set[str]:
with httpx.Client(
timeout=float(os.getenv("E2E_REQUEST_TIMEOUT", "120"))
) as client:
response = client.get(
f"{self.base_url.rstrip('/')}/model/info",
headers={"Authorization": f"Bearer {self.master_key}"},
)
assert response.status_code == 200, response.text
return {
model["model_name"]
for model in response.json().get("data", [])
if "model_name" in model
}
def ocr(self, model: str, document: dict[str, str]) -> httpx.Response:
with httpx.Client(
timeout=float(os.getenv("E2E_REQUEST_TIMEOUT", "120"))
) as client:
return client.post(
f"{self.base_url.rstrip('/')}/v1/ocr",
headers={"Authorization": f"Bearer {self.master_key}"},
json={"model": model, "document": document},
)
@dataclass(frozen=True)
class OcrResources:
gateway: OcrGateway
@pytest.fixture
def resources() -> OcrResources:
proxy_url = os.getenv("LITELLM_PROXY_URL")
if not proxy_url:
pytest.skip(
"Start a Rust OCR proxy and set LITELLM_PROXY_URL, e.g. http://localhost:4000"
)
return OcrResources(
gateway=OcrGateway(
base_url=proxy_url,
master_key=os.getenv("LITELLM_MASTER_KEY", "sk-1234"),
)
)
def _assert_ocr_response_shape(response_json: dict[str, Any]) -> None:
assert response_json["object"] == "ocr"
assert response_json["model"]
assert isinstance(response_json["pages"], list)
assert len(response_json["pages"]) > 0
assert "index" in response_json["pages"][0]
assert "markdown" in response_json["pages"][0]
class TestRustOcrGateway:
def test_rust_ocr_models_are_on_gateway_config(self) -> None:
config = yaml.safe_load(CONFIG_PATH.read_text())
configured_models = {
model_config["model_name"] for model_config in config["model_list"]
}
expected_models = {case.values[0] for case in RUST_OCR_GATEWAY_CASES}
assert expected_models.issubset(configured_models)
def test_running_gateway_loaded_rust_ocr_models(
self, resources: OcrResources
) -> None:
expected_models = {case.values[0] for case in RUST_OCR_GATEWAY_CASES}
assert expected_models.issubset(resources.gateway.model_names())
@pytest.mark.parametrize(("model", "document"), RUST_OCR_GATEWAY_CASES)
def test_rust_ocr_model_gateway_response(
self, resources: OcrResources, model: str, document: dict[str, str]
) -> None:
response = resources.gateway.ocr(model, document)
assert response.status_code == 200, response.text
_assert_ocr_response_shape(response.json())

View file

@ -0,0 +1,117 @@
"""Live e2e: Rust-backed OCR is reachable through the gateway across providers.
The gateway config declares one rust-ocr deployment per provider (mistral,
azure_ai, azure document intelligence, vertex mistral, vertex deepseek). Start the
proxy with the Rust OCR path enabled:
LITELLM_USE_RUST_OCR=1 litellm --config tests/e2e/gateway/litellm-config.yml
Three behaviors are checked: the config declares every provider's deployment (a
pure config read, no proxy needed); the running proxy loaded them onto /model/info;
and each one returns a well-formed OCR document over /v1/ocr. Per the e2e
"skip on environment, fail on behavior" rule, the proxy-backed cases skip when no
proxy answers but fail (never skip) once a request reaches it, so a provider whose
credentials are missing surfaces as a hard failure rather than silent green.
"""
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
import pytest
import yaml
from pydantic import BaseModel
from e2e_http import unwrap
from models import OcrBody, OcrDocument, OcrResponse
from passthrough_client import PassthroughClient
# Tiny in-repo fixtures served via jsdelivr (sha-pinned, immutable) so the request
# bodies stay stable across runs.
TEST_PDF_URL = (
"https://cdn.jsdelivr.net/gh/BerriAI/litellm"
"@d769e81c90d453240c61fc572cdb27fae06a89d0"
"/tests/llm_translation/fixtures/dummy.pdf"
)
TEST_IMAGE_URL = (
"https://cdn.jsdelivr.net/gh/BerriAI/litellm"
"@d769e81c90d453240c61fc572cdb27fae06a89d0"
"/tests/image_gen_tests/test_image.png"
)
CONFIG_PATH = Path(__file__).resolve().parents[1] / "gateway" / "litellm-config.yml"
@dataclass(frozen=True, slots=True)
class _OcrCase:
model: str
document: OcrDocument
RUST_OCR_CASES: tuple[_OcrCase, ...] = (
_OcrCase(
"rust-ocr-mistral",
OcrDocument(type="document_url", document_url=TEST_PDF_URL),
),
_OcrCase(
"rust-ocr-azure-ai",
OcrDocument(type="document_url", document_url=TEST_PDF_URL),
),
_OcrCase(
"rust-ocr-azure-document-intelligence",
OcrDocument(type="document_url", document_url=TEST_PDF_URL),
),
_OcrCase(
"rust-ocr-vertex-mistral",
OcrDocument(type="document_url", document_url=TEST_PDF_URL),
),
_OcrCase(
"rust-ocr-vertex-deepseek",
OcrDocument(type="image_url", image_url=TEST_IMAGE_URL),
),
)
_EXPECTED_MODELS = frozenset(case.model for case in RUST_OCR_CASES)
_CASE_IDS = tuple(case.model.removeprefix("rust-ocr-") for case in RUST_OCR_CASES)
class _ConfiguredModel(BaseModel):
model_name: str
class _GatewayConfig(BaseModel):
model_list: list[_ConfiguredModel]
def _configured_model_names() -> frozenset[str]:
config = _GatewayConfig.model_validate(yaml.safe_load(CONFIG_PATH.read_text()))
return frozenset(entry.model_name for entry in config.model_list)
def _assert_ocr_document(response: OcrResponse) -> None:
assert response.object == "ocr", f"expected object='ocr', got {response.object!r}"
assert response.model, "response missing the resolved model name"
assert response.pages, "OCR returned no pages"
assert response.pages[0].markdown is not None, "first page has no markdown"
def test_rust_ocr_models_declared_in_gateway_config() -> None:
"""Pure config read (no proxy): every provider's rust-ocr deployment the suite
exercises is declared in the gateway config the proxy runs with. A case added
here without a matching deployment fails before any live call is attempted."""
missing = _EXPECTED_MODELS - _configured_model_names()
assert not missing, f"rust-ocr models absent from {CONFIG_PATH.name}: {missing}"
@pytest.mark.e2e
class TestRustOcrGateway:
def test_gateway_loaded_rust_ocr_models(self, client: PassthroughClient) -> None:
loaded = frozenset(entry.model_name for entry in client.gateway.model_info())
missing = _EXPECTED_MODELS - loaded
assert not missing, f"proxy did not load rust-ocr models: {missing}"
@pytest.mark.parametrize("case", RUST_OCR_CASES, ids=_CASE_IDS)
def test_rust_ocr_response(self, client: PassthroughClient, scoped_key: str, case: _OcrCase) -> None:
response = unwrap(client.gateway.ocr(scoped_key, OcrBody(model=case.model, document=case.document)))
_assert_ocr_document(response)

View file

@ -125,6 +125,34 @@ class EmbedResponse(BaseModel):
model: str | None = None
# ---------- ocr ----------
class OcrDocument(BaseModel):
"""A document for /v1/ocr in Mistral OCR format: a document_url for PDFs/docs
or an image_url for images. exclude_none on serialize drops the unset one."""
type: str
document_url: str | None = None
image_url: str | None = None
class OcrBody(BaseModel):
model: str
document: OcrDocument
class OcrPage(BaseModel):
index: int
markdown: str
class OcrResponse(BaseModel):
object: str | None = None
model: str | None = None
pages: list[OcrPage] = []
# ---------- spend logs ----------
@ -216,14 +244,10 @@ class CustomPricing(BaseModel):
def token_cost(self, prompt_tokens: int, completion_tokens: int) -> float:
"""Spend for a fresh (uncached) call under these rates: the proxy's
custom-pricing formula (prompt * input + completion * output)."""
assert (
self.input_cost_per_token is not None
and self.output_cost_per_token is not None
), "custom pricing has no per-token rates"
return (
prompt_tokens * self.input_cost_per_token
+ completion_tokens * self.output_cost_per_token
assert self.input_cost_per_token is not None and self.output_cost_per_token is not None, (
"custom pricing has no per-token rates"
)
return prompt_tokens * self.input_cost_per_token + completion_tokens * self.output_cost_per_token
class ModelInfoEntry(BaseModel):