mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-17 23:51:30 +00:00
refactor(tests): discover OCR fixture targets
This commit is contained in:
parent
ab1f7b1fb8
commit
ce0dda405c
2 changed files with 160 additions and 33 deletions
|
|
@ -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__":
|
||||
|
|
|
|||
91
tests/test_litellm/ocr/test_generate_fixtures.py
Normal file
91
tests/test_litellm/ocr/test_generate_fixtures.py
Normal 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"
|
||||
Loading…
Add table
Reference in a new issue