diff --git a/tests/rust-python-harness/strategies/trace_parity/models.py b/tests/rust-python-harness/strategies/trace_parity/models.py index fd4585d85c1..3701fb41e14 100644 --- a/tests/rust-python-harness/strategies/trace_parity/models.py +++ b/tests/rust-python-harness/strategies/trace_parity/models.py @@ -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) diff --git a/tests/rust-python-harness/strategies/trace_parity/sdk/execution.py b/tests/rust-python-harness/strategies/trace_parity/sdk/execution.py index 2ef8a0584bc..ffdbe842e85 100644 --- a/tests/rust-python-harness/strategies/trace_parity/sdk/execution.py +++ b/tests/rust-python-harness/strategies/trace_parity/sdk/execution.py @@ -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 () diff --git a/tests/rust-python-harness/strategies/trace_parity/sdk/ocr/case.py b/tests/rust-python-harness/strategies/trace_parity/sdk/ocr/case.py index 58b16c68383..bb21e8ab0c5 100644 --- a/tests/rust-python-harness/strategies/trace_parity/sdk/ocr/case.py +++ b/tests/rust-python-harness/strategies/trace_parity/sdk/ocr/case.py @@ -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, ), ), ) diff --git a/tests/rust-python-harness/strategies/trace_parity/test_runner.py b/tests/rust-python-harness/strategies/trace_parity/test_runner.py index b940221bb7a..57364057788 100644 --- a/tests/rust-python-harness/strategies/trace_parity/test_runner.py +++ b/tests/rust-python-harness/strategies/trace_parity/test_runner.py @@ -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"