refactor(tests): discover OCR fixture targets

This commit is contained in:
Yujong Lee 2026-08-29 16:37:16 -07:00 committed by GitHub
parent ab1f7b1fb8
commit ce0dda405c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 160 additions and 33 deletions

View file

@ -3,7 +3,7 @@ from __future__ import annotations
import argparse
import logging
import os
from collections.abc import Callable
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from pathlib import Path
from typing import Final, cast
@ -27,6 +27,7 @@ from tests.test_litellm.ocr.fixture_models import (
MistralImageUrlDocument,
MistralOcrSdkInput,
OcrParityCase,
OcrSdkInputBase,
ReductoChunking,
ReductoDocumentUrlDocument,
ReductoFormatting,
@ -41,6 +42,7 @@ LOGGER: Final = logging.getLogger(__name__)
_TEXT: Final = st.just("invoice 123")
_VALUE_TEXT: Final = st.just("case-1")
_FONT_SIZE: Final = st.just(24)
_MISTRAL_MODEL: Final = "mistral/mistral-ocr-latest"
@dataclass(frozen=True, slots=True)
@ -48,7 +50,14 @@ class GeneratorArgs:
concurrency: int
examples: int
fixture_dir: Path | None
model: str
@dataclass(frozen=True, slots=True)
class OcrFixtureTarget:
name: str
provider_spec: ProviderSpec
strategy: SearchStrategy[OcrSdkInputBase]
invoke: Callable[[str, OcrSdkInputBase], object]
def _image_document(text: str, font_size: int) -> MistralImageUrlDocument:
@ -235,67 +244,94 @@ def reducto_legacy_input_strategy() -> SearchStrategy[ReductoParseLegacySdkInput
def _generate_examples(
spec: ProviderSpec,
target: OcrFixtureTarget,
root: Path,
model: str,
api_key: str,
examples: int,
concurrency: int,
sdk_call: Callable[..., object],
) -> None:
case_inputs: Final = generate_case_inputs(mistral_input_strategy(model), examples)
def invoke(api_base: str, case_input: MistralOcrSdkInput) -> object:
return sdk_call(api_base=api_base, api_key=api_key, **case_input.as_sdk_kwargs())
results: Final = record_cases(spec, root, case_inputs, invoke, OcrParityCase, concurrency)
case_inputs: Final = generate_case_inputs(target.strategy, examples)
results: Final = record_cases(
target.provider_spec,
root,
case_inputs,
target.invoke,
OcrParityCase,
concurrency,
)
for result in results:
LOGGER.info("%s %s", "cached" if result.cache_hit else "recorded", result.case.litellm_input.model)
LOGGER.info(
"%s %s %s",
"cached" if result.cache_hit else "recorded",
target.name,
result.case.litellm_input.model,
)
def _mistral_upstream_base() -> str:
configured: Final = os.environ.get("MISTRAL_API_BASE", "https://api.mistral.ai").rstrip("/")
def _mistral_upstream_base(environ: Mapping[str, str]) -> str:
configured: Final = environ.get("MISTRAL_API_BASE", "https://api.mistral.ai").rstrip("/")
return configured.removesuffix("/v1")
def _parse_args() -> GeneratorArgs:
def _mistral_target(
environ: Mapping[str, str],
sdk_call: Callable[..., object],
) -> OcrFixtureTarget | None:
api_key: Final = environ.get("MISTRAL_API_KEY")
if not api_key:
return None
def invoke(api_base: str, case_input: OcrSdkInputBase) -> object:
return sdk_call(api_base=api_base, api_key=api_key, **case_input.as_sdk_kwargs())
return OcrFixtureTarget(
name="mistral-ocr",
provider_spec=ProviderSpec(upstream_base=_mistral_upstream_base(environ)),
strategy=cast(SearchStrategy[OcrSdkInputBase], mistral_input_strategy(_MISTRAL_MODEL)),
invoke=invoke,
)
def discover_targets(
environ: Mapping[str, str],
sdk_call: Callable[..., object],
) -> tuple[OcrFixtureTarget, ...]:
candidates: Final = (_mistral_target(environ, sdk_call),)
return tuple(target for target in candidates if target is not None)
def require_targets(targets: tuple[OcrFixtureTarget, ...]) -> tuple[OcrFixtureTarget, ...]:
if targets:
return targets
raise SystemExit("No OCR fixture providers are configured. Set MISTRAL_API_KEY")
def parse_generator_args(argv: Sequence[str] | None = None) -> GeneratorArgs:
parser: Final = argparse.ArgumentParser()
parser.add_argument("--concurrency", type=int, default=4)
parser.add_argument("--examples", type=int, default=4)
parser.add_argument("--fixture-dir", type=Path)
parser.add_argument("--model", required=True)
namespace: Final = parser.parse_args()
namespace: Final = parser.parse_args(argv)
return GeneratorArgs(
concurrency=cast(int, namespace.concurrency),
examples=cast(int, namespace.examples),
fixture_dir=cast(Path | None, namespace.fixture_dir),
model=cast(str, namespace.model),
)
def main() -> None:
logging.basicConfig(level=logging.INFO, format="%(message)s")
load_dotenv()
args: Final = _parse_args()
api_key: Final = os.environ.get("MISTRAL_API_KEY") or os.environ.get("LITELLM_API_KEY")
if api_key is None:
raise SystemExit("MISTRAL_API_KEY is required")
args: Final = parse_generator_args()
sdk_call: Final = cast(Callable[..., object], litellm.ocr)
targets: Final = require_targets(discover_targets(os.environ, sdk_call))
root: Final = fixture_directory(
args.fixture_dir,
os.environ.get(FIXTURE_DIR_ENV),
Path(__file__).with_name(".fixtures"),
)
spec: Final = ProviderSpec(upstream_base=_mistral_upstream_base())
use_litellm_rust(False, ocr=None, aocr=None)
_generate_examples(
spec,
root,
args.model,
api_key,
args.examples,
args.concurrency,
cast(Callable[..., object], litellm.ocr),
)
for target in targets:
_generate_examples(target, root, args.examples, args.concurrency)
if __name__ == "__main__":

View file

@ -0,0 +1,91 @@
from __future__ import annotations
import queue
from pathlib import Path
from typing import Final
import pytest
from tests.test_litellm._fixture_recorder import generate_case_inputs
from tests.test_litellm.ocr.generate_fixtures import (
discover_targets,
parse_generator_args,
require_targets,
)
def _unused_sdk_call(**kwargs: object) -> object:
raise AssertionError(f"unexpected SDK call with {tuple(kwargs)}")
def test_parse_args_has_no_model_selection() -> None:
args: Final = parse_generator_args(["--examples", "2", "--concurrency", "3", "--fixture-dir", "/tmp/ocr"])
assert args.examples == 2
assert args.concurrency == 3
assert args.fixture_dir == Path("/tmp/ocr")
with pytest.raises(SystemExit):
parse_generator_args(["--model", "mistral/mistral-ocr-latest"])
@pytest.mark.parametrize(
"environ",
(
{},
{"MISTRAL_API_KEY": ""},
{"LITELLM_API_KEY": "generic-key"},
),
)
def test_discovery_requires_provider_specific_key(environ: dict[str, str]) -> None:
assert discover_targets(environ, _unused_sdk_call) == ()
def test_no_discovered_targets_has_actionable_error() -> None:
with pytest.raises(SystemExit, match="Set MISTRAL_API_KEY"):
require_targets(())
@pytest.mark.parametrize(
("configured", "expected"),
(
(None, "https://api.mistral.ai"),
("https://mistral.example/v1", "https://mistral.example"),
("https://mistral.example/", "https://mistral.example"),
),
)
def test_mistral_target_uses_canonical_model_and_normalized_base(
configured: str | None,
expected: str,
) -> None:
environ: Final = {
"MISTRAL_API_KEY": "mistral-secret",
**({"MISTRAL_API_BASE": configured} if configured is not None else {}),
}
targets: Final = discover_targets(environ, _unused_sdk_call)
assert len(targets) == 1
target: Final = targets[0]
assert target.name == "mistral-ocr"
assert target.provider_spec.upstream_base == expected
assert "mistral-secret" not in repr(target)
case_inputs: Final = generate_case_inputs(target.strategy, examples=1)
assert len(case_inputs) == 1
assert case_inputs[0].canonical_input()["model"] == "mistral/mistral-ocr-latest"
def test_mistral_target_invocation_forwards_discovered_credentials() -> None:
calls: Final[queue.SimpleQueue[dict[str, object]]] = queue.SimpleQueue()
def sdk_call(**kwargs: object) -> object:
calls.put(kwargs)
return object()
target: Final = discover_targets({"MISTRAL_API_KEY": "mistral-secret"}, sdk_call)[0]
case_input: Final = generate_case_inputs(target.strategy, examples=1)[0]
target.invoke("http://127.0.0.1:1234", case_input)
kwargs: Final = calls.get_nowait()
assert kwargs["api_base"] == "http://127.0.0.1:1234"
assert kwargs["api_key"] == "mistral-secret"
assert kwargs["model"] == "mistral/mistral-ocr-latest"