litellm/tests/rust-python-harness/strategies/trace_parity/models.py
2026-09-14 14:01:28 -07:00

81 lines
2.5 KiB
Python

from __future__ import annotations
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from typing import Final, Literal, TypeAlias, cast
from ...shared.parity.recorded_http import RecordedResponse
from ...shared.reporting.models import SdkFunction
from ...shared.tracing.steps import Engine, TraceMapping
TraceEngine = Literal["python", "rust", "both"]
TraceFailureSource = Literal["python", "rust", "harness"]
@dataclass(frozen=True, slots=True)
class RouteFixture:
kwargs: dict[str, object]
provider_responses: tuple[RecordedResponse, ...]
expected_failure: bool = False
consume_stream: bool = False
environment: tuple[tuple[str, str], ...] = ()
def derive(
self,
*,
kwargs: Mapping[str, object] | None = None,
provider_responses: tuple[RecordedResponse, ...] | None = None,
expected_failure: bool | None = None,
consume_stream: bool | None = None,
) -> RouteFixture:
return RouteFixture(
kwargs={**self.kwargs, **(kwargs or {})},
provider_responses=self.provider_responses if provider_responses is None else provider_responses,
expected_failure=self.expected_failure if expected_failure is None else expected_failure,
consume_stream=self.consume_stream if consume_stream is None else consume_stream,
environment=self.environment,
)
def with_body(self, **updates: object) -> RouteFixture:
raw_body: Final = self.kwargs.get("body")
if not isinstance(raw_body, dict):
raise ValueError("route fixture does not contain an object body")
body: Final = cast(dict[str, object], raw_body)
return self.derive(kwargs={"body": {**body, **updates}})
@dataclass(frozen=True, slots=True)
class RouteSpec:
route: SdkFunction
python_entrypoints: tuple[str, str]
rust_entrypoints: tuple[str, str] | None
fixture: Callable[[Engine, str], RouteFixture]
@dataclass(frozen=True, slots=True)
class GatewayRouteSpec:
route: SdkFunction
rust_supported: bool = True
TraceRouteSpec: TypeAlias = RouteSpec | GatewayRouteSpec
@dataclass(frozen=True, slots=True)
class TraceScenario:
name: str
fixture: Callable[[Engine, str], RouteFixture]
mappings: tuple[TraceMapping, ...]
asynchronous: bool
@dataclass(frozen=True, slots=True)
class TraceSuite:
route: TraceRouteSpec
scenarios: tuple[TraceScenario, ...]
@dataclass(frozen=True, slots=True)
class TraceExecutionFailure:
engine: TraceFailureSource
message: str