mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
* 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
124 lines
4.3 KiB
Python
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)
|