litellm/tests/rust-python-harness/shared/parity/replay.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

124 lines
4.3 KiB
Python

from __future__ import annotations
import base64
import queue
from contextlib import AbstractContextManager
from typing import Final
from pydantic import JsonValue, TypeAdapter
from .http import local_response_header
from .local_server import LocalHttpHandler, LocalHttpServer, serve_in_thread
from .models import CapturedRequest
from .recorded_http import RecordedHttpResponse, RecordedHttpStreamResponse, RecordedResponse
JSON_VALUE: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
EXCLUDED_REQUEST_HEADERS: Final = frozenset(
{
"host",
"content-length",
"connection",
"accept-encoding",
"user-agent",
"x-litellm-parity-route",
}
)
EXCLUDED_RESPONSE_HEADERS: Final = frozenset({"content-length", "transfer-encoding", "connection"})
def _replay_response_header(name: str, value: str, provider_url: str) -> str:
if name.lower() == "retry-after":
return "0"
return local_response_header(name, value, provider_url)
class ReplayServer(LocalHttpServer):
def __init__(self) -> None:
super().__init__(("127.0.0.1", 0), _ReplayHandler)
self.responses: queue.Queue[RecordedResponse] = queue.Queue()
self.requests: queue.Queue[CapturedRequest] = queue.Queue()
def enqueue_response(self, response: RecordedResponse) -> None:
self.responses.put(response)
def take_requests(self, expected_count: int) -> tuple[CapturedRequest, ...]:
request_count: Final = self.requests.qsize()
if request_count != expected_count:
raise AssertionError(f"expected exactly {expected_count} provider requests, received {request_count}")
return tuple(self.requests.get_nowait() for _ in range(request_count))
def reset(self) -> None:
while not self.responses.empty():
self.responses.get_nowait()
while not self.requests.empty():
self.requests.get_nowait()
class _ReplayHandler(LocalHttpHandler):
def do_POST(self) -> None:
self._replay()
def do_GET(self) -> None:
self._replay()
def do_PUT(self) -> None:
self._replay()
def do_PATCH(self) -> None:
self._replay()
def do_DELETE(self) -> None:
self._replay()
def _replay(self) -> None:
provider: Final = self.server
assert isinstance(provider, ReplayServer)
length: Final = int(self.headers.get("content-length") or "0")
raw_body: Final = self.rfile.read(length) if length else b""
content_type: Final = self.headers.get("content-type", "")
body: Final = (
JSON_VALUE.validate_json(raw_body)
if raw_body and content_type.lower().startswith("application/json")
else base64.b64encode(raw_body).decode("ascii")
if raw_body
else None
)
headers: Final = tuple(
sorted(
(name.lower(), value)
for name, value in self.headers.raw_items()
if name.lower() not in EXCLUDED_REQUEST_HEADERS
)
)
provider.requests.put(
CapturedRequest(
method=self.command,
path=self.path,
headers=headers,
body=body,
user_agent=self.headers.get("user-agent"),
)
)
try:
response: Final = provider.responses.get(timeout=5)
except queue.Empty:
self.send_error(500, "no replay response queued")
return
self.send_response_only(response.status_code)
for header in response.headers:
if header.name.lower() not in EXCLUDED_RESPONSE_HEADERS:
self.send_header(header.name, _replay_response_header(header.name, header.value, provider.url))
if isinstance(response, RecordedHttpResponse):
response_body: Final = response.body_bytes()
self.send_header("content-length", str(len(response_body)))
self.end_headers()
self.wfile.write(response_body)
return
assert isinstance(response, RecordedHttpStreamResponse)
self.send_header("transfer-encoding", "chunked")
self.end_headers()
self.write_chunked(chunk.data_bytes() for chunk in response.chunks)
def replay_server() -> AbstractContextManager[ReplayServer]:
return serve_in_thread(ReplayServer(), poll_interval=0.01)