mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix: preserve rust setting in trace parity
This commit is contained in:
parent
18b6049808
commit
27981c7d20
4 changed files with 6 additions and 17 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 ()
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue