litellm/tests/rust-python-harness/strategies/trace_parity/runner.py
yujonglee ee08c36fc0
refactor(tests): restructure rust python harness around strategy definitions (#39628)
* wip

* refactor(tests): move sdk function tracing into rust python harness

* dead code

* fix: handle harness keyboard interrupts

* refactor(tests): deduplicate rust python harness helpers

* fix(harness): expose validated strategy choices

* wip

* refactor(harness): let strategies own parity reports

* docs(harness): update strategy structure

* refactor(harness): localize strategy report views

* wip

* fix(harness): satisfy mapping runner type checks

* fix(harness): clarify trace parity output

* wip

* fix(harness): clarify unit mapping report

* fix(harness): finalize trace parity contracts

* refactor(harness): structure parity contracts

* feat: derive unit test mapping from traces

* feat(harness): map rstest test families

* feat(ocr): port Azure document intelligence tests

* feat(harness): enforce complete unit mappings

* feat(ocr): add reducto core transforms

* feat(harness): classify host-only unit tests

* fix(ocr): complete Rust provider plumbing

* fix(harness): reuse OCR parity workers
2026-09-03 21:15:01 -07:00

182 lines
7.2 KiB
Python

from __future__ import annotations
import importlib
from collections.abc import Sequence
from pathlib import Path
from time import monotonic
from typing import Final
from ...shared.reporting.models import CaseResult, HarnessCase, HarnessRun, ResultArtifact, RunStatus, Surface
from ...shared.reporting.strategy import ModuleCaseSpec, UpdateCallback
from ...shared.native_build import ensure_trace_bridge
from .models import GatewayRouteSpec, RouteSpec, TraceExecutionFailure, TraceMode, TraceScenario, TraceSuite
from .reporting import TRACE_COMPARISON_ARTIFACT, TraceComparisonArtifact
from .sdk.execution import execute_trace
def _load_case(reference: str, harness_case: HarnessCase) -> TraceSuite | TraceExecutionFailure:
try:
module: Final = importlib.import_module(reference)
except Exception as error:
return TraceExecutionFailure("harness", f"cannot import {reference}: {type(error).__name__}: {error}")
suite: Final = getattr(module, "TRACE_SUITE", None)
if not isinstance(suite, TraceSuite):
return TraceExecutionFailure("harness", f"{reference} must export TRACE_SUITE: TraceSuite")
validation_error: Final = validate_trace_suite(suite, harness_case)
if validation_error is not None:
return TraceExecutionFailure("harness", f"{reference} {validation_error}")
return suite
def validate_trace_suite(suite: TraceSuite, harness_case: HarnessCase) -> str | None:
names: Final = tuple(scenario.name for scenario in suite.scenarios)
if not names or len(names) != len(set(names)) or any(not name or ":" in name for name in names):
return "scenario names must be non-empty, unique, and colon-free"
invalid_modes: Final = tuple(
scenario.name
for scenario in suite.scenarios
if not scenario.modes
or len(scenario.modes) != len(set(scenario.modes))
or any(mode not in {"sync", "async"} for mode in scenario.modes)
)
if invalid_modes:
return f"scenarios must use non-empty, unique sync/async modes: {', '.join(invalid_modes)}"
surface: Final = harness_case.surface
if surface == "sdk" and not isinstance(suite.route, RouteSpec):
return "must use RouteSpec for the sdk surface"
if surface == "gateway" and not isinstance(suite.route, GatewayRouteSpec):
return "must use GatewayRouteSpec for the gateway surface"
if surface is None:
return "requires an sdk or gateway surface"
if suite.route.route != harness_case.sdk_function:
return f"route {suite.route.route} does not match case function {harness_case.sdk_function}"
return None
def scenario_nodeids(
trace_suite: TraceSuite,
harness_case: HarnessCase,
selected_scenarios: frozenset[str] = frozenset(),
) -> tuple[tuple[TraceScenario, TraceMode, str], ...]:
surface: Final = harness_case.surface
if surface is None:
return ()
return tuple(
(scenario, mode, f"trace:{surface}:{harness_case.sdk_function}:{scenario.name}:{mode}")
for scenario in trace_suite.scenarios
if not selected_scenarios or scenario.name in selected_scenarios
for mode in scenario.modes
)
def _record_setup_failure(run: HarnessRun, case: HarnessCase, message: str, stage: str) -> None:
result: Final = run.results[case.key]
nodeid: Final = f"trace:{case.surface}:{case.sdk_function}:{stage}"
result.collected.add(nodeid)
result.record(nodeid, RunStatus.ERROR)
run.failures.append((nodeid, message))
def run_trace_mode(
run: HarnessRun,
result: CaseResult,
trace_suite: TraceSuite,
scenario: TraceScenario,
mode: TraceMode,
surface: Surface,
nodeid: str,
on_update: UpdateCallback,
) -> None:
started_at: Final = monotonic()
comparison: Final = _execute_mode(trace_suite, scenario, mode, surface)
duration: Final = monotonic() - started_at
if isinstance(comparison, TraceExecutionFailure):
result.record(nodeid, RunStatus.ERROR, duration)
run.failures.append((nodeid, comparison.message))
on_update(run)
return
artifact: Final = ResultArtifact(TRACE_COMPARISON_ARTIFACT, comparison.model_dump_json())
if comparison.has_errors():
result.record(nodeid, RunStatus.ERROR, duration, (artifact,))
run.failures.append(
(nodeid, "\n".join(error for error in (comparison.python_error, comparison.rust_error) if error))
)
else:
status: Final = RunStatus.PASSED if comparison.contract_matches() else RunStatus.FAILED
result.record(nodeid, status, duration, (artifact,))
if status is RunStatus.FAILED:
run.failures.append((nodeid, "trace contract mismatch; see the rendered comparison"))
on_update(run)
def _execute_mode(
trace_suite: TraceSuite,
scenario: TraceScenario,
mode: TraceMode,
surface: Surface,
) -> TraceComparisonArtifact | TraceExecutionFailure:
route: Final = trace_suite.route
if isinstance(route, GatewayRouteSpec):
if surface != "gateway":
return TraceExecutionFailure("harness", "gateway route cannot run on the sdk surface")
from .gateway.execution import execute_gateway_trace
return execute_gateway_trace(route, scenario, mode)
if surface != "sdk":
return TraceExecutionFailure("harness", "sdk route cannot run on the gateway surface")
return execute_trace(route, scenario, mode, surface)
def _run_case(
run: HarnessRun,
harness_case: HarnessCase,
selected_scenarios: frozenset[str],
on_update: UpdateCallback,
) -> None:
result: Final = run.results[harness_case.key]
spec: Final = harness_case.spec
if not isinstance(spec, ModuleCaseSpec):
return
surface: Final = harness_case.surface
if surface is None:
return
trace_suite: Final = _load_case(spec.module, harness_case)
if isinstance(trace_suite, TraceExecutionFailure):
_record_setup_failure(run, harness_case, trace_suite.message, "load")
on_update(run)
return
nodeids: Final = scenario_nodeids(trace_suite, harness_case, selected_scenarios)
result.collected.update(nodeid for _, _, nodeid in nodeids)
if not nodeids:
result.status = RunStatus.SKIPPED
on_update(run)
return
result.status = RunStatus.RUNNING
on_update(run)
for scenario, mode, nodeid in nodeids:
run_trace_mode(run, result, trace_suite, scenario, mode, surface, nodeid, on_update)
def run_trace_cases(
cases: Sequence[HarnessCase],
repo_root: Path,
on_update: UpdateCallback,
runner_args: Sequence[str] = (),
) -> tuple[int, HarnessRun]:
selected_scenarios: Final = frozenset(runner_args)
run: Final = HarnessRun.from_cases(cases)
bridge_error: Final = ensure_trace_bridge(repo_root)
if bridge_error is not None:
for harness_case in cases:
_record_setup_failure(run, harness_case, bridge_error, "bridge")
run.finished_at = monotonic()
on_update(run)
return 1, run
for harness_case in cases:
_run_case(run, harness_case, selected_scenarios, on_update)
run.finished_at = monotonic()
on_update(run)
failed: Final = any(
result.status in {RunStatus.ERROR, RunStatus.FAILED, RunStatus.MISSING} for result in run.results.values()
)
return int(failed), run