fix: preserve rust setting in trace parity

This commit is contained in:
Yujong Lee 2026-09-14 12:17:15 -07:00
parent 18b6049808
commit 27981c7d20
4 changed files with 6 additions and 17 deletions

View file

@ -66,7 +66,6 @@ class TraceScenario:
fixture: Callable[[Engine, str], RouteFixture]
mappings: tuple[TraceMapping, ...]
asynchronous: bool
python_rust_enabled: bool = False
@dataclass(frozen=True, slots=True)

View file

@ -3,7 +3,6 @@ from __future__ import annotations
import asyncio
import os
from collections.abc import AsyncIterable, Awaitable, Iterable
from contextlib import nullcontext
from dataclasses import dataclass
from pathlib import Path
from typing import Final, Protocol, cast
@ -107,9 +106,7 @@ def _collect(
return _CollectedTrace(tuple(profiler.events), error)
def collect_trace(
spec: RouteSpec, engine: Engine, *, asynchronous: bool, python_rust_enabled: bool = False
) -> tuple[FunctionTraceEvent, ...] | TraceExecutionFailure:
def collect_trace(spec: RouteSpec, engine: Engine, *, asynchronous: bool) -> tuple[FunctionTraceEvent, ...] | TraceExecutionFailure:
function: Final = _entrypoint(spec, engine, asynchronous=asynchronous)
if isinstance(function, TraceExecutionFailure):
return function
@ -130,12 +127,7 @@ def collect_trace(
consume_stream=base_fixture.consume_stream,
environment=base_fixture.environment,
)
environment: Final = (
patch.dict(os.environ, {"LITELLM_RUST": "1" if python_rust_enabled else "0"})
if engine == "python"
else nullcontext()
)
with environment, patch.dict(os.environ, fixture.environment):
with patch.dict(os.environ, fixture.environment):
collected: Final = _collect(function, fixture, engine, asynchronous=asynchronous)
provider.take_requests(len(fixture.provider_responses))
except Exception as error:
@ -172,7 +164,6 @@ def execute_trace(
scenario_route,
"python",
asynchronous=scenario.asynchronous,
python_rust_enabled=scenario.python_rust_enabled,
)
if engine != "rust"
else ()

View file

@ -498,14 +498,12 @@ TRACE_SUITE: Final = TraceSuite(
fixture=_mistral_fixture,
mappings=PUBLIC_RUST_DISPATCH_MAPPINGS,
asynchronous=False,
python_rust_enabled=True,
),
TraceScenario(
name="async-public-rust-dispatch",
fixture=_mistral_fixture,
mappings=PUBLIC_RUST_DISPATCH_MAPPINGS,
asynchronous=True,
python_rust_enabled=True,
),
),
)

View file

@ -80,7 +80,7 @@ def test_python_engine_skips_native_bridge(monkeypatch: pytest.MonkeyPatch, tmp_
assert selected == [(frozenset({"mistral"}), "python")]
def test_python_trace_controls_native_ocr_dispatch(monkeypatch: pytest.MonkeyPatch) -> None:
def test_python_trace_preserves_native_ocr_dispatch_setting(monkeypatch: pytest.MonkeyPatch) -> None:
execution: Final = importlib.import_module("tests.rust-python-harness.strategies.trace_parity.sdk.execution")
route: Final = RouteSpec("ocr", ("ocr", "aocr"), ("ocr", "aocr"), _fixture)
observed: list[str | None] = []
@ -98,11 +98,12 @@ def test_python_trace_controls_native_ocr_dispatch(monkeypatch: pytest.MonkeyPat
error=None,
)
monkeypatch.setenv("LITELLM_RUST", "1")
monkeypatch.setattr(execution, "_collect", collect)
monkeypatch.setenv("LITELLM_RUST", "0")
collect_trace(route, "python", asynchronous=False)
collect_trace(route, "python", asynchronous=True, python_rust_enabled=True)
monkeypatch.setenv("LITELLM_RUST", "1")
collect_trace(route, "python", asynchronous=True)
assert observed == ["0", "1"]
assert os.environ["LITELLM_RUST"] == "1"