mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-24 00:52:24 +00:00
test(vcr): guard leaked cassette patches and make injected-transport embedding tests immune (#42542)
* test(vcr): guard leaked cassette patches and make injected-transport embedding tests immune * test(vcr): derive the leak guard's patch points from vcrpy's own reset list * test(vcr): share CapturingTransport and switch the encoding_format embedding test to it --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
075536eca1
commit
c37fe46534
8 changed files with 331 additions and 199 deletions
|
|
@ -8,6 +8,7 @@ from __future__ import annotations
|
|||
import ast
|
||||
import atexit
|
||||
import hashlib
|
||||
import inspect
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
|
|
@ -15,10 +16,18 @@ import socket
|
|||
import sys
|
||||
import threading
|
||||
from collections import defaultdict
|
||||
from typing import Iterable
|
||||
from collections.abc import Iterable, Iterator
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
from unittest import mock
|
||||
|
||||
import aiohttp
|
||||
import pytest
|
||||
import vcr
|
||||
import vcr.matchers as _vcr_matchers
|
||||
import vcr.patch as _vcr_patch
|
||||
|
||||
from tests._vcr_redis_persister import (
|
||||
MAX_EPISODES_PER_CASSETTE,
|
||||
|
|
@ -127,9 +136,7 @@ def emit_vcr_diagnostic_log(terminalreporter) -> None:
|
|||
with open(path, "r", encoding="utf-8") as fh:
|
||||
content = fh.read()
|
||||
except OSError as exc:
|
||||
read_errors.append(
|
||||
f" [failed to read {name}: {type(exc).__name__}: {exc}]"
|
||||
)
|
||||
read_errors.append(f" [failed to read {name}: {type(exc).__name__}: {exc}]")
|
||||
continue
|
||||
for line in content.splitlines():
|
||||
if not line.strip():
|
||||
|
|
@ -142,9 +149,7 @@ def emit_vcr_diagnostic_log(terminalreporter) -> None:
|
|||
return
|
||||
|
||||
terminalreporter.write_sep("=", "VCR DIAGNOSTIC LOG", bold=True)
|
||||
terminalreporter.write_line(
|
||||
f" source dir: {directory} (deduplicated; full log archived as a CI artifact)"
|
||||
)
|
||||
terminalreporter.write_line(f" source dir: {directory} (deduplicated; full log archived as a CI artifact)")
|
||||
for line in read_errors:
|
||||
terminalreporter.write_line(line)
|
||||
|
||||
|
|
@ -235,9 +240,7 @@ def pin_httpx_multipart_boundary(monkeypatch) -> None:
|
|||
boundary = VCR_FIXED_MULTIPART_BOUNDARY.encode("ascii")
|
||||
return _original_init(self, data=data, files=files, boundary=boundary, **kwargs)
|
||||
|
||||
monkeypatch.setattr(
|
||||
_httpx_multipart.MultipartStream, "__init__", _init_with_fixed_boundary
|
||||
)
|
||||
monkeypatch.setattr(_httpx_multipart.MultipartStream, "__init__", _init_with_fixed_boundary)
|
||||
|
||||
|
||||
@pytest.fixture(scope="session", autouse=True)
|
||||
|
|
@ -270,11 +273,7 @@ def _replace_b64_json_in_place(obj) -> bool:
|
|||
changed = False
|
||||
if isinstance(obj, dict):
|
||||
for key, value in obj.items():
|
||||
if (
|
||||
key == "b64_json"
|
||||
and isinstance(value, str)
|
||||
and len(value) > len(VCR_IMAGE_B64_PLACEHOLDER)
|
||||
):
|
||||
if key == "b64_json" and isinstance(value, str) and len(value) > len(VCR_IMAGE_B64_PLACEHOLDER):
|
||||
obj[key] = VCR_IMAGE_B64_PLACEHOLDER
|
||||
changed = True
|
||||
elif _replace_b64_json_in_place(value):
|
||||
|
|
@ -296,16 +295,12 @@ def _strip_image_b64_payloads(response):
|
|||
preserves all those checks while shrinking cassettes by ~99%.
|
||||
"""
|
||||
if not isinstance(response, dict):
|
||||
vcr_diag_write_line(
|
||||
f"[vcr-strip-b64] response is {type(response).__name__!r}, not "
|
||||
"dict; skipping b64 scrub"
|
||||
)
|
||||
vcr_diag_write_line(f"[vcr-strip-b64] response is {type(response).__name__!r}, not dict; skipping b64 scrub")
|
||||
return response
|
||||
body = response.get("body")
|
||||
if not isinstance(body, dict):
|
||||
vcr_diag_write_line(
|
||||
f"[vcr-strip-b64] response['body'] is {type(body).__name__!r}, "
|
||||
"not dict; skipping b64 scrub"
|
||||
f"[vcr-strip-b64] response['body'] is {type(body).__name__!r}, not dict; skipping b64 scrub"
|
||||
)
|
||||
return response
|
||||
raw = body.get("string")
|
||||
|
|
@ -316,10 +311,7 @@ def _strip_image_b64_payloads(response):
|
|||
try:
|
||||
text = bytes(raw).decode("utf-8")
|
||||
except UnicodeDecodeError:
|
||||
vcr_diag_write_line(
|
||||
"[vcr-strip-b64] response body bytes are not valid UTF-8; "
|
||||
"skipping b64 scrub"
|
||||
)
|
||||
vcr_diag_write_line("[vcr-strip-b64] response body bytes are not valid UTF-8; skipping b64 scrub")
|
||||
return response
|
||||
was_bytes = True
|
||||
elif isinstance(raw, str):
|
||||
|
|
@ -327,8 +319,7 @@ def _strip_image_b64_payloads(response):
|
|||
was_bytes = False
|
||||
else:
|
||||
vcr_diag_write_line(
|
||||
f"[vcr-strip-b64] response['body']['string'] is "
|
||||
f"{type(raw).__name__!r}, not bytes/str; skipping b64 scrub"
|
||||
f"[vcr-strip-b64] response['body']['string'] is {type(raw).__name__!r}, not bytes/str; skipping b64 scrub"
|
||||
)
|
||||
return response
|
||||
|
||||
|
|
@ -349,9 +340,7 @@ def _strip_image_b64_payloads(response):
|
|||
for key in list(headers):
|
||||
if str(key).lower() == "content-length":
|
||||
value = headers[key]
|
||||
headers[key] = (
|
||||
[new_len_value] if isinstance(value, list) else new_len_value
|
||||
)
|
||||
headers[key] = [new_len_value] if isinstance(value, list) else new_len_value
|
||||
return response
|
||||
|
||||
|
||||
|
|
@ -409,15 +398,11 @@ def _canonical_body(request) -> tuple[bytes, str]:
|
|||
# selected. This mirrors the existing SigV4 / multipart-boundary / b64-image
|
||||
# normalizations already in this module, and means the already-bloated
|
||||
# cassettes start replaying immediately without a flush + re-record.
|
||||
_VCR_UUID_RE = re.compile(
|
||||
rb"[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}"
|
||||
)
|
||||
_VCR_UUID_RE = re.compile(rb"[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}")
|
||||
_VCR_LITELLM_BATCH_JOB_RE = re.compile(rb"litellm-batch-[0-9a-fA-F]{8}")
|
||||
# ISO-8601 timestamps, e.g. ``2026-05-25T03:40:37.262045Z`` /
|
||||
# ``2026-05-25T03:40:37+00:00``.
|
||||
_VCR_ISO_TS_RE = re.compile(
|
||||
rb"\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}(?:\.\d+)?(?:Z|[+-]\d{2}:?\d{2})?"
|
||||
)
|
||||
_VCR_ISO_TS_RE = re.compile(rb"\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}(?:\.\d+)?(?:Z|[+-]\d{2}:?\d{2})?")
|
||||
# Unix epoch as 13-digit milliseconds, then 10-digit ``time.time()`` float,
|
||||
# then 10-digit integer seconds. Anchored to ``1`` + 9/12 digits, which keeps
|
||||
# them inside the 2001-2033 / 2001-2033 epoch windows and avoids matching
|
||||
|
|
@ -639,10 +624,7 @@ def _should_drop_telemetry_record(request) -> bool:
|
|||
return False
|
||||
if not _is_telemetry_request(request):
|
||||
return False
|
||||
if (
|
||||
_is_telemetry_export_request(request)
|
||||
and not _current_test_replays_telemetry_export()
|
||||
):
|
||||
if _is_telemetry_export_request(request) and not _current_test_replays_telemetry_export():
|
||||
return True
|
||||
return not _current_test_records_telemetry()
|
||||
|
||||
|
|
@ -767,9 +749,7 @@ def _iter_header_values(headers, name: str):
|
|||
yield value
|
||||
|
||||
|
||||
_AWS_SIGV4_CREDENTIAL_RE = re.compile(
|
||||
r"AWS4-HMAC-SHA256\s+Credential=([^/\s,]+)/", re.IGNORECASE
|
||||
)
|
||||
_AWS_SIGV4_CREDENTIAL_RE = re.compile(r"AWS4-HMAC-SHA256\s+Credential=([^/\s,]+)/", re.IGNORECASE)
|
||||
|
||||
# Google OAuth2 access tokens always start with ``ya29.`` regardless of how
|
||||
# they were minted (service account, metadata server, impersonation).
|
||||
|
|
@ -891,9 +871,7 @@ def _normalize_multipart_boundary(request) -> None:
|
|||
return
|
||||
|
||||
try:
|
||||
headers[content_type_key] = content_type_value.replace(
|
||||
match.group(0), fixed_param
|
||||
)
|
||||
headers[content_type_key] = content_type_value.replace(match.group(0), fixed_param)
|
||||
except (TypeError, AttributeError):
|
||||
return
|
||||
|
||||
|
|
@ -985,8 +963,7 @@ def _materialize_iterable_body(request) -> None:
|
|||
uri = getattr(request, "uri", getattr(request, "url", "?"))
|
||||
first_type = type(chunks[0]).__name__ if chunks else "empty"
|
||||
vcr_diag_write_line(
|
||||
f"[vcr-materialize] FALLBACK: {method} {uri} chunk type "
|
||||
f"{first_type!r} not coerced to bytes; storing b''"
|
||||
f"[vcr-materialize] FALLBACK: {method} {uri} chunk type {first_type!r} not coerced to bytes; storing b''"
|
||||
)
|
||||
out = b""
|
||||
|
||||
|
|
@ -1026,9 +1003,7 @@ def _key_fingerprint_matcher(r1, r2) -> None:
|
|||
return
|
||||
|
||||
def _fp(req):
|
||||
for value in _iter_header_values(
|
||||
getattr(req, "headers", None), KEY_FINGERPRINT_HEADER
|
||||
):
|
||||
for value in _iter_header_values(getattr(req, "headers", None), KEY_FINGERPRINT_HEADER):
|
||||
if value is None:
|
||||
continue
|
||||
return value if isinstance(value, str) else str(value)
|
||||
|
|
@ -1159,13 +1134,11 @@ def _print_atexit_banner() -> None:
|
|||
_emit("VCR CASSETTE CACHE DEGRADED")
|
||||
if save_failures:
|
||||
_emit(
|
||||
f" {save_failures} cassette save failure(s); last error: "
|
||||
f"{health.get('save_failure_last_error', '')}"
|
||||
f" {save_failures} cassette save failure(s); last error: {health.get('save_failure_last_error', '')}"
|
||||
)
|
||||
if load_failures:
|
||||
_emit(
|
||||
f" {load_failures} cassette load failure(s); last error: "
|
||||
f"{health.get('load_failure_last_error', '')}"
|
||||
f" {load_failures} cassette load failure(s); last error: {health.get('load_failure_last_error', '')}"
|
||||
)
|
||||
if snapshot:
|
||||
_emit(_format_capacity_line(snapshot))
|
||||
|
|
@ -1276,11 +1249,7 @@ class _RespxUsageVisitor(ast.NodeVisitor):
|
|||
if isinstance(dec, ast.Call):
|
||||
dec = dec.func
|
||||
if isinstance(dec, ast.Attribute):
|
||||
return (
|
||||
isinstance(dec.value, ast.Name)
|
||||
and dec.value.id == "respx"
|
||||
and dec.attr == "mock"
|
||||
)
|
||||
return isinstance(dec.value, ast.Name) and dec.value.id == "respx" and dec.attr == "mock"
|
||||
return False
|
||||
|
||||
def _is_pytest_mark_respx(self, dec: ast.expr) -> bool:
|
||||
|
|
@ -1307,9 +1276,7 @@ class _RespxUsageVisitor(ast.NodeVisitor):
|
|||
# ``def test_foo(respx_mock): ...`` — pytest supplies the fixture
|
||||
# whenever the parameter name appears, regardless of marker.
|
||||
all_args = (
|
||||
list(args.args)
|
||||
+ list(args.kwonlyargs)
|
||||
+ (list(args.posonlyargs) if hasattr(args, "posonlyargs") else [])
|
||||
list(args.args) + list(args.kwonlyargs) + (list(args.posonlyargs) if hasattr(args, "posonlyargs") else [])
|
||||
)
|
||||
for a in all_args:
|
||||
if a.arg == "respx_mock":
|
||||
|
|
@ -1566,9 +1533,7 @@ def _emit_outcome_payload(
|
|||
},
|
||||
)
|
||||
)
|
||||
node.user_properties.append(
|
||||
(_USER_PROP_RECORDED_BY, os.environ.get("PYTEST_XDIST_WORKER", ""))
|
||||
)
|
||||
node.user_properties.append((_USER_PROP_RECORDED_BY, os.environ.get("PYTEST_XDIST_WORKER", "")))
|
||||
|
||||
|
||||
def aggregate_report_outcome(report) -> None:
|
||||
|
|
@ -1616,9 +1581,7 @@ def aggregate_report_outcome(report) -> None:
|
|||
if verdict == VERDICT_MISS_OVERFLOW:
|
||||
_session_stats["overflow_tests"].append(nodeid)
|
||||
elif verdict == VERDICT_UNMARKED_LIVE_CALL:
|
||||
_session_stats["unmarked_live_call_tests"].append(
|
||||
(nodeid, list(outcome.get("live_call_hosts") or []))
|
||||
)
|
||||
_session_stats["unmarked_live_call_tests"].append((nodeid, list(outcome.get("live_call_hosts") or [])))
|
||||
|
||||
skip_reason = outcome.get("skip_reason")
|
||||
if skip_reason:
|
||||
|
|
@ -1635,9 +1598,7 @@ def session_stats_snapshot() -> dict:
|
|||
"overflow_tests": list(_session_stats["overflow_tests"]),
|
||||
"unmarked_live_call_tests": list(_session_stats["unmarked_live_call_tests"]),
|
||||
"skip_reason_counts": dict(_session_stats["skip_reason_counts"]),
|
||||
"skip_reason_examples": {
|
||||
k: list(v) for k, v in _session_stats["skip_reason_examples"].items()
|
||||
},
|
||||
"skip_reason_examples": {k: list(v) for k, v in _session_stats["skip_reason_examples"].items()},
|
||||
}
|
||||
|
||||
|
||||
|
|
@ -1810,9 +1771,7 @@ def record_vcr_outcome(request, vcr) -> None:
|
|||
# Cassette is None ⇒ test wasn't VCR-marked. Honor the skip reason
|
||||
# we tagged at collection time, and pull live-call hosts captured by
|
||||
# the socket probe (if any).
|
||||
skip_reason = getattr(
|
||||
request.node, VCR_SKIP_REASON_USER_ATTR, SKIP_REASON_FILE_OPT_OUT
|
||||
)
|
||||
skip_reason = getattr(request.node, VCR_SKIP_REASON_USER_ATTR, SKIP_REASON_FILE_OPT_OUT)
|
||||
_session_stats["skip_reason_counts"][skip_reason] += 1
|
||||
|
||||
hosts = getattr(request.node, _LIVE_CALL_BUFFER_KEY, []) or []
|
||||
|
|
@ -1837,9 +1796,7 @@ def record_vcr_outcome(request, vcr) -> None:
|
|||
live_call_hosts=hosts,
|
||||
)
|
||||
if vcr_outcome_logging_enabled():
|
||||
request.node.user_properties.append(
|
||||
(_USER_PROP_VERDICT_LINE, _format_verdict_line(verdict, None, extra))
|
||||
)
|
||||
request.node.user_properties.append((_USER_PROP_VERDICT_LINE, _format_verdict_line(verdict, None, extra)))
|
||||
|
||||
|
||||
def install_live_call_probe(request, vcr) -> None:
|
||||
|
|
@ -1858,9 +1815,7 @@ def install_live_call_probe(request, vcr) -> None:
|
|||
# Track the current test for telemetry-leak suppression (applies to every
|
||||
# test, VCR-marked or not). See ``_should_drop_telemetry_record``.
|
||||
global _current_test_nodeid
|
||||
_current_test_nodeid = str(
|
||||
getattr(getattr(request, "node", None), "nodeid", "") or ""
|
||||
)
|
||||
_current_test_nodeid = str(getattr(getattr(request, "node", None), "nodeid", "") or "")
|
||||
if vcr is not None or vcr_disabled():
|
||||
return None
|
||||
probe = _LiveCallProbe()
|
||||
|
|
@ -1876,10 +1831,7 @@ def _format_capacity_line(snapshot: dict) -> str:
|
|||
pct = float(snapshot.get("used_pct", 0.0) or 0.0)
|
||||
used_mb = used / (1024 * 1024)
|
||||
cap_mb = cap / (1024 * 1024)
|
||||
return (
|
||||
f" Cassette Redis usage: {used_mb:.1f} MiB / {cap_mb:.1f} MiB "
|
||||
f"({pct:.1f}% of maxmemory)"
|
||||
)
|
||||
return f" Cassette Redis usage: {used_mb:.1f} MiB / {cap_mb:.1f} MiB ({pct:.1f}% of maxmemory)"
|
||||
|
||||
|
||||
def emit_vcr_classification_summary(terminalreporter) -> None:
|
||||
|
|
@ -1940,14 +1892,10 @@ def emit_vcr_classification_summary(terminalreporter) -> None:
|
|||
total_leaks = sum(leak_counts.values())
|
||||
terminalreporter.write_sep("-", "VCR COST LEAK CHECK", bold=True)
|
||||
if total_leaks:
|
||||
rendered = ", ".join(
|
||||
f"{verdict}={count}" for verdict, count in leak_counts.items() if count
|
||||
)
|
||||
rendered = ", ".join(f"{verdict}={count}" for verdict, count in leak_counts.items() if count)
|
||||
terminalreporter.write_line(f" FAIL: {rendered}")
|
||||
else:
|
||||
terminalreporter.write_line(
|
||||
" PASS: no overflow, partial, not-persisted, or unmarked live-call verdicts"
|
||||
)
|
||||
terminalreporter.write_line(" PASS: no overflow, partial, not-persisted, or unmarked live-call verdicts")
|
||||
|
||||
overflow = snapshot["overflow_tests"]
|
||||
if overflow:
|
||||
|
|
@ -2007,18 +1955,14 @@ def emit_cassette_cache_session_banner(terminalreporter) -> None:
|
|||
snapshot = cassette_cache_capacity_snapshot()
|
||||
|
||||
if save_failures or load_failures:
|
||||
terminalreporter.write_sep(
|
||||
"=", "VCR CASSETTE CACHE DEGRADED", red=True, bold=True
|
||||
)
|
||||
terminalreporter.write_sep("=", "VCR CASSETTE CACHE DEGRADED", red=True, bold=True)
|
||||
if save_failures:
|
||||
terminalreporter.write_line(
|
||||
f" {save_failures} cassette save failure(s); last error: "
|
||||
f"{health.get('save_failure_last_error', '')}"
|
||||
f" {save_failures} cassette save failure(s); last error: {health.get('save_failure_last_error', '')}"
|
||||
)
|
||||
if load_failures:
|
||||
terminalreporter.write_line(
|
||||
f" {load_failures} cassette load failure(s); last error: "
|
||||
f"{health.get('load_failure_last_error', '')}"
|
||||
f" {load_failures} cassette load failure(s); last error: {health.get('load_failure_last_error', '')}"
|
||||
)
|
||||
terminalreporter.write_line(
|
||||
" Tests still passed because cassette persistence is best-effort, "
|
||||
|
|
@ -2031,9 +1975,7 @@ def emit_cassette_cache_session_banner(terminalreporter) -> None:
|
|||
return
|
||||
|
||||
if snapshot and snapshot["used_pct"] >= CASSETTE_CACHE_HIGH_WATER_FRACTION * 100:
|
||||
terminalreporter.write_sep(
|
||||
"=", "VCR CASSETTE CACHE NEAR CAPACITY", yellow=True, bold=True
|
||||
)
|
||||
terminalreporter.write_sep("=", "VCR CASSETTE CACHE NEAR CAPACITY", yellow=True, bold=True)
|
||||
terminalreporter.write_line(_format_capacity_line(snapshot))
|
||||
terminalreporter.write_line(
|
||||
" No save failures yet, but Redis is approaching maxmemory. "
|
||||
|
|
@ -2082,13 +2024,104 @@ class VerboseReporterState:
|
|||
if reporter is None:
|
||||
return
|
||||
verdict = next(
|
||||
(
|
||||
v
|
||||
for k, v in (report.user_properties or [])
|
||||
if k == _USER_PROP_VERDICT_LINE
|
||||
),
|
||||
(v for k, v in (report.user_properties or []) if k == _USER_PROP_VERDICT_LINE),
|
||||
None,
|
||||
)
|
||||
if not verdict:
|
||||
return
|
||||
reporter.write_line(f"{verdict} :: {report.nodeid}")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class VcrPatchPoint:
|
||||
owner: object
|
||||
attribute: str
|
||||
original: object
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
return f"{_patch_owner_name(self.owner)}.{self.attribute}"
|
||||
|
||||
def current(self) -> object:
|
||||
current: Final[object] = getattr(self.owner, self.attribute)
|
||||
return current
|
||||
|
||||
def is_patched(self) -> bool:
|
||||
return self.current() is not self.original
|
||||
|
||||
def restore(self) -> None:
|
||||
setattr(self.owner, self.attribute, self.original)
|
||||
|
||||
|
||||
def _patch_owner_name(owner: object) -> str:
|
||||
if inspect.isclass(owner):
|
||||
return f"{owner.__module__}.{owner.__qualname__}"
|
||||
if inspect.ismodule(owner):
|
||||
return owner.__name__
|
||||
return repr(owner)
|
||||
|
||||
|
||||
def _vcr_patch_point(patcher: mock._patch[object]) -> VcrPatchPoint:
|
||||
owner: Final[object] = patcher.getter()
|
||||
return VcrPatchPoint(owner=owner, attribute=patcher.attribute, original=patcher.new)
|
||||
|
||||
|
||||
_VCR_PATCH_POINTS: Final = (
|
||||
*(_vcr_patch_point(patcher) for patcher in _vcr_patch.reset_patchers()),
|
||||
VcrPatchPoint(aiohttp.ClientSession, "_request", _vcr_patch._AiohttpClientSessionRequest),
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class VcrPatchLeak:
|
||||
patch_points: tuple[str, ...]
|
||||
cassette_paths: tuple[str, ...]
|
||||
|
||||
|
||||
def _cassette_paths_wrapped_into(fn: object) -> tuple[str, ...]:
|
||||
if not inspect.isfunction(fn):
|
||||
return ()
|
||||
cassette: Final = inspect.getclosurevars(fn).nonlocals.get("cassette")
|
||||
own: Final = (str(cassette._path),) if isinstance(cassette, vcr.cassette.Cassette) else ()
|
||||
return own + _cassette_paths_wrapped_into(getattr(fn, "__wrapped__", None))
|
||||
|
||||
|
||||
def detect_vcr_patch_leak() -> VcrPatchLeak | None:
|
||||
leaked: Final = tuple(point for point in _VCR_PATCH_POINTS if point.is_patched())
|
||||
if not leaked:
|
||||
return None
|
||||
return VcrPatchLeak(
|
||||
patch_points=tuple(point.name for point in leaked),
|
||||
cassette_paths=tuple(
|
||||
dict.fromkeys(path for point in leaked for path in _cassette_paths_wrapped_into(point.current()))
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def restore_vcr_patch_points() -> None:
|
||||
for point in _VCR_PATCH_POINTS:
|
||||
point.restore()
|
||||
|
||||
|
||||
def guard_vcr_patch_points(item: pytest.Item, teardown_failed: bool) -> None:
|
||||
leak: Final = detect_vcr_patch_leak()
|
||||
if leak is None:
|
||||
return
|
||||
restore_vcr_patch_points()
|
||||
if teardown_failed:
|
||||
return
|
||||
pytest.fail(
|
||||
f"{item.nodeid} finished with a vcrpy cassette still patched into "
|
||||
f"{', '.join(leak.patch_points)} (cassettes: {', '.join(leak.cassette_paths) or 'unknown'}); "
|
||||
"the originals were restored so later tests are unaffected",
|
||||
pytrace=False,
|
||||
)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def rewound_new_episodes_cassette(cassette_dir: Path) -> Iterator[vcr.cassette.Cassette]:
|
||||
cassette_path: Final = cassette_dir / "rewound_owner.yaml"
|
||||
cassette_path.write_text("interactions: []\nversion: 1\n")
|
||||
recorder: Final = vcr.VCR(cassette_library_dir=str(cassette_dir))
|
||||
with recorder.use_cassette(cassette_path.name, record_mode="new_episodes") as cassette:
|
||||
yield cassette
|
||||
|
|
|
|||
25
tests/capturing_transport.py
Normal file
25
tests/capturing_transport.py
Normal file
|
|
@ -0,0 +1,25 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
|
||||
_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object])
|
||||
|
||||
|
||||
class CapturingTransport(httpx.AsyncBaseTransport, httpx.BaseTransport):
|
||||
def __init__(self, response: BaseModel) -> None:
|
||||
self._response: Final = response
|
||||
self.request_bodies: tuple[Mapping[str, object], ...] = ()
|
||||
|
||||
def handle_request(self, request: httpx.Request) -> httpx.Response:
|
||||
return self._respond(request.read())
|
||||
|
||||
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
|
||||
return self._respond(await request.aread())
|
||||
|
||||
def _respond(self, body: bytes) -> httpx.Response:
|
||||
self.request_bodies = (*self.request_bodies, _JSON_OBJECT.validate_json(body))
|
||||
return httpx.Response(200, json=self._response.model_dump(mode="json"))
|
||||
|
|
@ -7,6 +7,8 @@
|
|||
|
||||
import asyncio
|
||||
import importlib
|
||||
from collections.abc import Generator
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -20,6 +22,7 @@ from tests._vcr_conftest_common import ( # noqa: E402,F401
|
|||
emit_cassette_cache_session_banner,
|
||||
emit_vcr_classification_summary,
|
||||
emit_vcr_diagnostic_log,
|
||||
guard_vcr_patch_points,
|
||||
install_live_call_probe,
|
||||
record_vcr_outcome,
|
||||
register_persister_if_enabled,
|
||||
|
|
@ -37,17 +40,12 @@ def fake_openai_endpoint():
|
|||
|
||||
# Per-item respx detection (``apply_vcr_auto_marker_to_items``) handles
|
||||
# the vast majority of respx-vs-vcrpy conflicts automatically. The entries
|
||||
# below are the persister's and the WebSocket VCR's own unit-test files, which
|
||||
# exercise ``save_cassette`` / ``load_cassette`` against fakeredis and must not
|
||||
# themselves run under a live cassette context.
|
||||
# below are the persister's, the WebSocket VCR's, and the cassette patch-leak
|
||||
# guard's own unit-test files, which exercise ``save_cassette`` /
|
||||
# ``load_cassette`` against fakeredis or enter cassettes themselves and must
|
||||
# not run under a live cassette context.
|
||||
_VCR_AUTO_MARKER_SKIP_FILES = frozenset(
|
||||
{"test_vcr_redis_persister.py", "test_ws_vcr.py"}
|
||||
)
|
||||
|
||||
_VCR_INCOMPATIBLE_NODEID_SUFFIXES: tuple[str, ...] = (
|
||||
"test_nvidia_nim.py::test_embedding_nvidia_nim",
|
||||
"test_litellm_proxy_provider.py::test_litellm_gateway_from_sdk_embedding[False]",
|
||||
"test_litellm_proxy_provider.py::test_litellm_gateway_from_sdk_embedding[True]",
|
||||
{"test_vcr_redis_persister.py", "test_ws_vcr.py", "test_vcr_leak_guard.py"}
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -77,6 +75,17 @@ def _vcr_outcome_gate(request, vcr):
|
|||
record_vcr_outcome(request, vcr)
|
||||
|
||||
|
||||
@pytest.hookimpl(wrapper=True, trylast=True)
|
||||
def pytest_runtest_teardown(item: pytest.Item) -> Generator[None, object, object]:
|
||||
try:
|
||||
result: Final = yield
|
||||
except BaseException:
|
||||
guard_vcr_patch_points(item, teardown_failed=True)
|
||||
raise
|
||||
guard_vcr_patch_points(item, teardown_failed=False)
|
||||
return result
|
||||
|
||||
|
||||
def pytest_configure(config):
|
||||
_verbose_state.remember_pluginmanager(config)
|
||||
reset_vcr_diag_dir()
|
||||
|
|
@ -172,7 +181,6 @@ def pytest_collection_modifyitems(config, items):
|
|||
apply_vcr_auto_marker_to_items(
|
||||
items,
|
||||
skip_files=_VCR_AUTO_MARKER_SKIP_FILES,
|
||||
skip_nodeid_suffixes=_VCR_INCOMPATIBLE_NODEID_SUFFIXES,
|
||||
)
|
||||
|
||||
custom_logger_tests = [
|
||||
|
|
|
|||
|
|
@ -2,6 +2,8 @@ import json
|
|||
import re
|
||||
from datetime import datetime
|
||||
from io import BytesIO
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
|
||||
|
|
@ -12,7 +14,12 @@ import pytest
|
|||
from unittest.mock import MagicMock, patch
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler, AsyncHTTPHandler
|
||||
import pytest_asyncio
|
||||
from openai import AsyncOpenAI
|
||||
from openai import AsyncOpenAI, OpenAI
|
||||
from openai.types import CreateEmbeddingResponse, Embedding
|
||||
from openai.types.create_embedding_response import Usage
|
||||
|
||||
from tests.capturing_transport import CapturingTransport
|
||||
from tests._vcr_conftest_common import rewound_new_episodes_cassette
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -87,62 +94,61 @@ async def test_litellm_gateway_from_sdk_structured_output():
|
|||
assert "json_schema" in json_schema
|
||||
|
||||
|
||||
@pytest.mark.parametrize("is_async", [False, True])
|
||||
_GATEWAY_EMBEDDING_RESPONSE: Final = CreateEmbeddingResponse(
|
||||
object="list",
|
||||
data=(Embedding(object="embedding", index=0, embedding=(0.1, 0.2, 0.3)),),
|
||||
model="my-vllm-model",
|
||||
usage=Usage(prompt_tokens=2, total_tokens=2),
|
||||
)
|
||||
|
||||
|
||||
async def _gateway_embedding_via_injected_client(
|
||||
is_async: bool,
|
||||
) -> tuple[CapturingTransport, litellm.EmbeddingResponse]:
|
||||
transport: Final = CapturingTransport(_GATEWAY_EMBEDDING_RESPONSE)
|
||||
response: Final = (
|
||||
await litellm.aembedding(
|
||||
model="litellm_proxy/my-vllm-model",
|
||||
input="Hello world",
|
||||
client=AsyncOpenAI(api_key="fake-key", http_client=httpx.AsyncClient(transport=transport)),
|
||||
api_base="my-custom-api-base",
|
||||
)
|
||||
if is_async
|
||||
else litellm.embedding(
|
||||
model="litellm_proxy/my-vllm-model",
|
||||
input="Hello world",
|
||||
client=OpenAI(api_key="fake-key", http_client=httpx.Client(transport=transport)),
|
||||
api_base="my-custom-api-base",
|
||||
)
|
||||
)
|
||||
return transport, response
|
||||
|
||||
|
||||
@pytest.mark.parametrize("is_async", (False, True))
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_gateway_from_sdk_embedding(is_async):
|
||||
async def test_litellm_gateway_from_sdk_embedding(is_async: bool):
|
||||
litellm.set_verbose = True
|
||||
litellm._turn_on_debug()
|
||||
|
||||
captured_bodies = []
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
captured_bodies.append(json.loads(request.content))
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}],
|
||||
"model": "my-vllm-model",
|
||||
"usage": {"prompt_tokens": 2, "total_tokens": 2},
|
||||
},
|
||||
)
|
||||
|
||||
if is_async:
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
openai_client = AsyncOpenAI(
|
||||
api_key="fake-key",
|
||||
http_client=httpx.AsyncClient(transport=httpx.MockTransport(handler)),
|
||||
)
|
||||
response = await litellm.aembedding(
|
||||
model="litellm_proxy/my-vllm-model",
|
||||
input="Hello world",
|
||||
client=openai_client,
|
||||
api_base="my-custom-api-base",
|
||||
)
|
||||
else:
|
||||
from openai import OpenAI
|
||||
|
||||
openai_client = OpenAI(
|
||||
api_key="fake-key",
|
||||
http_client=httpx.Client(transport=httpx.MockTransport(handler)),
|
||||
)
|
||||
response = litellm.embedding(
|
||||
model="litellm_proxy/my-vllm-model",
|
||||
input="Hello world",
|
||||
client=openai_client,
|
||||
api_base="my-custom-api-base",
|
||||
)
|
||||
|
||||
request_body = captured_bodies[0]
|
||||
print("Request body - {}".format(request_body))
|
||||
transport, response = await _gateway_embedding_via_injected_client(is_async)
|
||||
|
||||
request_body: Final = transport.request_bodies[0]
|
||||
assert "Hello world" == request_body["input"]
|
||||
assert "my-vllm-model" == request_body["model"]
|
||||
assert "encoding_format" not in request_body
|
||||
assert response.data[0]["embedding"] == [0.1, 0.2, 0.3]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_gateway_from_sdk_embedding_under_foreign_cassette(tmp_path: Path):
|
||||
with rewound_new_episodes_cassette(tmp_path):
|
||||
sync_transport, _ = await _gateway_embedding_via_injected_client(is_async=False)
|
||||
async_transport, _ = await _gateway_embedding_via_injected_client(is_async=True)
|
||||
|
||||
assert tuple(body["input"] for body in sync_transport.request_bodies) == ("Hello world",)
|
||||
assert tuple(body["input"] for body in async_transport.request_bodies) == ("Hello world",)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("is_async", [False, True])
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_gateway_from_sdk_image_generation(is_async):
|
||||
|
|
|
|||
|
|
@ -1,17 +1,21 @@
|
|||
import json
|
||||
from datetime import datetime
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from openai.types import CreateEmbeddingResponse, Embedding
|
||||
from openai.types.create_embedding_response import Usage as EmbeddingUsage
|
||||
from unittest.mock import patch, MagicMock
|
||||
|
||||
import litellm
|
||||
from litellm import Choices, Message, ModelResponse, EmbeddingResponse, Usage
|
||||
from litellm import completion
|
||||
from base_rerank_unit_tests import BaseLLMRerankTest
|
||||
from tests.capturing_transport import CapturingTransport
|
||||
|
||||
|
||||
def test_completion_nvidia_nim():
|
||||
|
|
@ -63,33 +67,23 @@ def test_embedding_nvidia_nim():
|
|||
litellm.set_verbose = True
|
||||
from openai import OpenAI
|
||||
|
||||
captured_bodies = []
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
captured_bodies.append(json.loads(request.content))
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}],
|
||||
"model": "nvidia/nv-embedqa-e5-v5",
|
||||
"usage": {"prompt_tokens": 6, "total_tokens": 6},
|
||||
},
|
||||
transport: Final = CapturingTransport(
|
||||
CreateEmbeddingResponse(
|
||||
object="list",
|
||||
data=(Embedding(object="embedding", index=0, embedding=(0.1, 0.2, 0.3)),),
|
||||
model="nvidia/nv-embedqa-e5-v5",
|
||||
usage=EmbeddingUsage(prompt_tokens=6, total_tokens=6),
|
||||
)
|
||||
|
||||
client = OpenAI(
|
||||
api_key="fake-api-key",
|
||||
http_client=httpx.Client(transport=httpx.MockTransport(handler)),
|
||||
)
|
||||
response = litellm.embedding(
|
||||
client: Final = OpenAI(api_key="fake-api-key", http_client=httpx.Client(transport=transport))
|
||||
response: Final = litellm.embedding(
|
||||
model="nvidia_nim/nvidia/nv-embedqa-e5-v5",
|
||||
input="What is the meaning of life?",
|
||||
input_type="passage",
|
||||
dimensions=1024,
|
||||
client=client,
|
||||
)
|
||||
request_body = captured_bodies[0]
|
||||
print("request_body: ", request_body)
|
||||
request_body: Final = transport.request_bodies[0]
|
||||
assert request_body["input"] == "What is the meaning of life?"
|
||||
assert request_body["model"] == "nvidia/nv-embedqa-e5-v5"
|
||||
assert request_body["input_type"] == "passage"
|
||||
|
|
|
|||
72
tests/llm_translation/test_vcr_leak_guard.py
Normal file
72
tests/llm_translation/test_vcr_leak_guard.py
Normal file
|
|
@ -0,0 +1,72 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import httpx2
|
||||
import pytest
|
||||
|
||||
from tests._vcr_conftest_common import (
|
||||
detect_vcr_patch_leak,
|
||||
guard_vcr_patch_points,
|
||||
restore_vcr_patch_points,
|
||||
rewound_new_episodes_cassette,
|
||||
)
|
||||
|
||||
_ORIGINAL_MOCK_HANDLE_ASYNC_REQUEST: Final = httpx.MockTransport.handle_async_request
|
||||
_ORIGINAL_HTTPX2_MOCK_HANDLE_ASYNC_REQUEST: Final = httpx2.MockTransport.handle_async_request
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def leaked_cassette_dir(tmp_path: Path):
|
||||
context: Final = rewound_new_episodes_cassette(tmp_path)
|
||||
context.__enter__()
|
||||
yield tmp_path
|
||||
context.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_no_leak_when_no_cassette_is_active():
|
||||
assert detect_vcr_patch_leak() is None
|
||||
|
||||
|
||||
def test_leaked_cassette_is_detected_named_and_restorable(leaked_cassette_dir: Path):
|
||||
leak: Final = detect_vcr_patch_leak()
|
||||
|
||||
assert leak is not None
|
||||
assert {"httpx.MockTransport.handle_async_request", "aiohttp.client.ClientSession._request"} <= set(
|
||||
leak.patch_points
|
||||
)
|
||||
assert leak.cassette_paths == (str(leaked_cassette_dir / "rewound_owner.yaml"),)
|
||||
|
||||
restore_vcr_patch_points()
|
||||
|
||||
assert detect_vcr_patch_leak() is None
|
||||
assert httpx.MockTransport.handle_async_request is _ORIGINAL_MOCK_HANDLE_ASYNC_REQUEST
|
||||
|
||||
|
||||
def test_leak_is_detected_on_every_transport_family_vcrpy_patches(leaked_cassette_dir: Path):
|
||||
leak: Final = detect_vcr_patch_leak()
|
||||
|
||||
assert leak is not None
|
||||
assert "httpx2.MockTransport.handle_async_request" in leak.patch_points
|
||||
assert httpx2.MockTransport.handle_async_request is not _ORIGINAL_HTTPX2_MOCK_HANDLE_ASYNC_REQUEST
|
||||
|
||||
restore_vcr_patch_points()
|
||||
|
||||
assert httpx2.MockTransport.handle_async_request is _ORIGINAL_HTTPX2_MOCK_HANDLE_ASYNC_REQUEST
|
||||
|
||||
|
||||
def test_guard_fails_the_leaking_test_and_restores_the_originals(request, leaked_cassette_dir: Path):
|
||||
with pytest.raises(pytest.fail.Exception, match=re.escape(request.node.nodeid)) as failure:
|
||||
guard_vcr_patch_points(request.node, teardown_failed=False)
|
||||
|
||||
assert str(leaked_cassette_dir / "rewound_owner.yaml") in str(failure.value)
|
||||
assert detect_vcr_patch_leak() is None
|
||||
|
||||
|
||||
def test_guard_restores_silently_when_the_teardown_already_failed(request, leaked_cassette_dir: Path):
|
||||
guard_vcr_patch_points(request.node, teardown_failed=True)
|
||||
|
||||
assert detect_vcr_patch_leak() is None
|
||||
|
|
@ -93,7 +93,6 @@ _VCR_INCOMPATIBLE_FILES = frozenset(
|
|||
# carry no real provider cost.
|
||||
_VCR_INCOMPATIBLE_NODEID_SUFFIXES: tuple[str, ...] = (
|
||||
"test_router.py::test_router_text_completion_client",
|
||||
"test_embedding.py::test_encoding_format_omitted_by_default_for_openai_sdk",
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -15,6 +15,9 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||
|
||||
import litellm
|
||||
from litellm import completion, completion_cost, embedding
|
||||
from openai.types import CreateEmbeddingResponse
|
||||
from openai.types.create_embedding_response import Usage as EmbeddingUsage
|
||||
from tests.capturing_transport import CapturingTransport
|
||||
from tests.fake_openai_endpoint import FAKE_OPENAI_API_BASE
|
||||
|
||||
litellm.set_verbose = False
|
||||
|
|
@ -1268,23 +1271,15 @@ def test_encoding_format_omitted_by_default_for_openai_sdk(monkeypatch):
|
|||
Optional global override: `LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT`.
|
||||
"""
|
||||
monkeypatch.delenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", raising=False)
|
||||
captured_bodies = []
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
captured_bodies.append(json.loads(request.content))
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}],
|
||||
"model": "text-embedding-ada-002",
|
||||
"usage": {"prompt_tokens": 1, "total_tokens": 1},
|
||||
},
|
||||
transport = CapturingTransport(
|
||||
CreateEmbeddingResponse(
|
||||
object="list",
|
||||
data=(Embedding(object="embedding", index=0, embedding=(0.1, 0.2, 0.3)),),
|
||||
model="text-embedding-ada-002",
|
||||
usage=EmbeddingUsage(prompt_tokens=1, total_tokens=1),
|
||||
)
|
||||
|
||||
client = openai.OpenAI(
|
||||
api_key="sk-test", http_client=httpx.Client(transport=httpx.MockTransport(handler))
|
||||
)
|
||||
client = openai.OpenAI(api_key="sk-test", http_client=httpx.Client(transport=transport))
|
||||
|
||||
response = embedding(
|
||||
model="text-embedding-ada-002",
|
||||
|
|
@ -1294,7 +1289,7 @@ def test_encoding_format_omitted_by_default_for_openai_sdk(monkeypatch):
|
|||
)
|
||||
|
||||
assert response.data[0]["embedding"] == [0.1, 0.2, 0.3]
|
||||
assert "encoding_format" not in captured_bodies[0], (
|
||||
assert "encoding_format" not in transport.request_bodies[0], (
|
||||
"encoding_format should be omitted from the upstream request when not provided by user"
|
||||
)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue