mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
fix(ci): enforce release wheel metadata contract
This commit is contained in:
parent
bb824dc921
commit
1dea3f3666
2 changed files with 266 additions and 14 deletions
110
.github/scripts/verify_linux_native_wheel.py
vendored
110
.github/scripts/verify_linux_native_wheel.py
vendored
|
|
@ -6,21 +6,44 @@ import re
|
|||
import subprocess
|
||||
import sys
|
||||
import zipfile
|
||||
from email import policy
|
||||
from email.parser import BytesParser
|
||||
from itertools import product
|
||||
from pathlib import Path, PurePosixPath
|
||||
from typing import Final
|
||||
from types import ModuleType
|
||||
from typing import Final, cast
|
||||
|
||||
EXPECTED_PYTHON_TAG: Final = "cp310"
|
||||
EXPECTED_ABI_TAG: Final = "abi3"
|
||||
EXPECTED_PLATFORM_TAG: Final = "linux_x86_64"
|
||||
|
||||
|
||||
def _loads_native_module(native_path: Path) -> bool:
|
||||
def _dist_info_directory(member: zipfile.ZipInfo) -> str | None:
|
||||
parts: Final = PurePosixPath(member.filename).parts
|
||||
if not parts or not parts[0].endswith(".dist-info"):
|
||||
return None
|
||||
return parts[0]
|
||||
|
||||
|
||||
def _wheel_metadata_tags(archive: zipfile.ZipFile, members: tuple[zipfile.ZipInfo, ...]) -> tuple[str, ...]:
|
||||
if len(members) != 1:
|
||||
return ()
|
||||
metadata: Final = BytesParser(policy=policy.default).parsebytes(archive.read(members[0]))
|
||||
tags: Final = cast(list[str], metadata.get_all("Tag", []))
|
||||
return tuple(tag.strip() for tag in tags)
|
||||
|
||||
|
||||
def _load_native_module(native_path: Path) -> ModuleType | None:
|
||||
module_spec: Final = importlib.util.spec_from_file_location("litellm.rust_bridge._native", native_path)
|
||||
if module_spec is None or module_spec.loader is None:
|
||||
return False
|
||||
return None
|
||||
try:
|
||||
native_module: Final = importlib.util.module_from_spec(module_spec)
|
||||
module_spec.loader.exec_module(native_module)
|
||||
except Exception as error:
|
||||
sys.stderr.write(f"native module load failed: {error}\n")
|
||||
return False
|
||||
return True
|
||||
return None
|
||||
return native_module
|
||||
|
||||
|
||||
def main() -> int:
|
||||
|
|
@ -29,8 +52,39 @@ def main() -> int:
|
|||
return 2
|
||||
|
||||
wheel: Final = Path(sys.argv[1])
|
||||
wheel_tags: Final = wheel.stem.rsplit("-", maxsplit=3)
|
||||
if len(wheel_tags) != 4:
|
||||
sys.stderr.write(f"cannot parse wheel tags from {wheel.name}\n")
|
||||
return 1
|
||||
|
||||
wheel_identity: Final = wheel_tags[0].split("-")
|
||||
if len(wheel_identity) != 2 or wheel_identity[0] != "litellm" or not wheel_identity[1]:
|
||||
sys.stderr.write(f"unexpected wheel identity: {wheel_tags[0]}\n")
|
||||
return 1
|
||||
|
||||
expected_dist_info_directory: Final = f"{wheel_tags[0]}.dist-info"
|
||||
expected_dist_info_directories: Final = frozenset((expected_dist_info_directory,))
|
||||
python_tag: Final = wheel_tags[1]
|
||||
abi_tag: Final = wheel_tags[2]
|
||||
platform_tag: Final = wheel_tags[3]
|
||||
expanded_filename_tags: Final = frozenset(
|
||||
"-".join(tag) for tag in product(python_tag.split("."), abi_tag.split("."), platform_tag.split("."))
|
||||
)
|
||||
|
||||
with zipfile.ZipFile(wheel) as archive:
|
||||
wheel_members: Final = archive.infolist()
|
||||
dist_info_directories: Final = frozenset(
|
||||
directory for member in wheel_members if (directory := _dist_info_directory(member)) is not None
|
||||
)
|
||||
required_dist_info_files: Final = ("METADATA", "RECORD", "WHEEL")
|
||||
dist_info_file_counts: Final = {
|
||||
filename: sum(member.filename == f"{expected_dist_info_directory}/{filename}" for member in wheel_members)
|
||||
for filename in required_dist_info_files
|
||||
}
|
||||
wheel_metadata_members: Final = tuple(
|
||||
member for member in wheel_members if member.filename == f"{expected_dist_info_directory}/WHEEL"
|
||||
)
|
||||
wheel_metadata_tags: Final = _wheel_metadata_tags(archive, wheel_metadata_members)
|
||||
native_members: Final = tuple(
|
||||
member
|
||||
for member in wheel_members
|
||||
|
|
@ -52,14 +106,10 @@ def main() -> int:
|
|||
native_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
native_path.write_bytes(archive.read(native_member))
|
||||
|
||||
wheel_tags: Final = wheel.stem.rsplit("-", maxsplit=3)
|
||||
if len(wheel_tags) != 4:
|
||||
sys.stderr.write(f"cannot parse wheel tags from {wheel.name}\n")
|
||||
return 1
|
||||
|
||||
python_tag: Final = wheel_tags[1]
|
||||
abi_tag: Final = wheel_tags[2]
|
||||
platform_tag: Final = wheel_tags[3]
|
||||
wheel_metadata_tags_match: Final = (
|
||||
len(wheel_metadata_tags) == len(expanded_filename_tags)
|
||||
and frozenset(wheel_metadata_tags) == expanded_filename_tags
|
||||
)
|
||||
commit_sha: Final = os.environ.get("RELEASE_WHEEL_COMMIT_SHA", os.environ.get("GITHUB_SHA", "unknown"))
|
||||
rustc_version: Final = subprocess.run(
|
||||
("rustc", "--version"),
|
||||
|
|
@ -120,14 +170,26 @@ def main() -> int:
|
|||
text=True,
|
||||
).stdout
|
||||
extension_entry_point_present: Final = "PyInit__native" in dynamic_symbols
|
||||
native_module_loads: Final = _loads_native_module(native_path)
|
||||
native_module: Final = _load_native_module(native_path)
|
||||
native_module_loads: Final = native_module is not None
|
||||
panic_test_hook_absent: Final = native_module is not None and not hasattr(native_module, "_panic_for_test")
|
||||
native_size_limit: Final = 20_000_000
|
||||
native_size_within_limit: Final = native_member.file_size <= native_size_limit
|
||||
validations: Final = (
|
||||
(f"Python tag is {EXPECTED_PYTHON_TAG}", python_tag == EXPECTED_PYTHON_TAG),
|
||||
(f"ABI tag is {EXPECTED_ABI_TAG}", abi_tag == EXPECTED_ABI_TAG),
|
||||
(f"Platform tag is {EXPECTED_PLATFORM_TAG}", platform_tag == EXPECTED_PLATFORM_TAG),
|
||||
("Wheel dist-info directory matches the filename", dist_info_directories == expected_dist_info_directories),
|
||||
(
|
||||
"Required dist-info files are present exactly once",
|
||||
all(count == 1 for count in dist_info_file_counts.values()),
|
||||
),
|
||||
("Wheel metadata tags match the filename", wheel_metadata_tags_match),
|
||||
("Debug sections are absent", debug_sections_absent),
|
||||
("Static symbol table is absent", static_symbol_table_absent),
|
||||
("Python extension entry point is present", extension_entry_point_present),
|
||||
("Native module loads", native_module_loads),
|
||||
("Production module omits the panic test hook", panic_test_hook_absent),
|
||||
("Native extension does not exceed 20 MB", native_size_within_limit),
|
||||
("Wheel contents are valid", not unexpected_members),
|
||||
)
|
||||
|
|
@ -146,6 +208,26 @@ def main() -> int:
|
|||
sys.stderr.write(f"{native_member.filename} contains a static symbol table\n")
|
||||
if not extension_entry_point_present:
|
||||
sys.stderr.write("native extension does not export PyInit__native\n")
|
||||
if python_tag != EXPECTED_PYTHON_TAG:
|
||||
sys.stderr.write(f"unexpected Python tag: expected {EXPECTED_PYTHON_TAG}, found {python_tag}\n")
|
||||
if abi_tag != EXPECTED_ABI_TAG:
|
||||
sys.stderr.write(f"unexpected ABI tag: expected {EXPECTED_ABI_TAG}, found {abi_tag}\n")
|
||||
if platform_tag != EXPECTED_PLATFORM_TAG:
|
||||
sys.stderr.write(f"unexpected platform tag: expected {EXPECTED_PLATFORM_TAG}, found {platform_tag}\n")
|
||||
if dist_info_directories != expected_dist_info_directories:
|
||||
sys.stderr.write(
|
||||
f"unexpected dist-info directories: expected {[expected_dist_info_directory]}, "
|
||||
f"found {sorted(dist_info_directories)}\n"
|
||||
)
|
||||
if any(count != 1 for count in dist_info_file_counts.values()):
|
||||
sys.stderr.write(f"required dist-info file counts are invalid: {dist_info_file_counts}\n")
|
||||
elif not wheel_metadata_tags_match:
|
||||
sys.stderr.write(
|
||||
f"WHEEL tags do not match filename: expected {sorted(expanded_filename_tags)}, "
|
||||
f"found {sorted(wheel_metadata_tags)}\n"
|
||||
)
|
||||
if native_module is not None and not panic_test_hook_absent:
|
||||
sys.stderr.write("production native module exposes _panic_for_test\n")
|
||||
if not native_size_within_limit:
|
||||
sys.stderr.write(f"native extension exceeds 20 MB: {native_member.file_size / 1_000_000:.2f} MB\n")
|
||||
if unexpected_members:
|
||||
|
|
|
|||
170
tests/test_litellm/test_verify_linux_native_wheel.py
Normal file
170
tests/test_litellm/test_verify_linux_native_wheel.py
Normal file
|
|
@ -0,0 +1,170 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
import subprocess
|
||||
import sys
|
||||
import zipfile
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
from types import ModuleType
|
||||
from typing import Final, Protocol, cast
|
||||
|
||||
import pytest
|
||||
|
||||
_REPO_ROOT: Final = Path(__file__).resolve().parents[2]
|
||||
_MODULE_PATH: Final = _REPO_ROOT / ".github" / "scripts" / "verify_linux_native_wheel.py"
|
||||
|
||||
|
||||
class _VerifierModule(Protocol):
|
||||
subprocess: ModuleType
|
||||
_load_native_module: Callable[[Path], ModuleType | None]
|
||||
main: Callable[[], int]
|
||||
|
||||
|
||||
_SPEC: Final = importlib.util.spec_from_file_location("verify_linux_native_wheel", _MODULE_PATH)
|
||||
assert _SPEC is not None and _SPEC.loader is not None
|
||||
_LOADED_VERIFIER: Final = importlib.util.module_from_spec(_SPEC)
|
||||
sys.modules[_SPEC.name] = _LOADED_VERIFIER
|
||||
_SPEC.loader.exec_module(_LOADED_VERIFIER)
|
||||
verifier: Final = cast(_VerifierModule, _LOADED_VERIFIER)
|
||||
|
||||
_EXPECTED_TAG: Final = "cp310-abi3-linux_x86_64"
|
||||
_NATIVE_MEMBER: Final = "litellm/rust_bridge/_native.abi3.so"
|
||||
_DIST_INFO: Final = "litellm-1.100.0.dist-info"
|
||||
|
||||
|
||||
def _write_wheel(
|
||||
tmp_path: Path,
|
||||
*,
|
||||
filename_tag: str,
|
||||
metadata_tags: tuple[str, ...] | None = (_EXPECTED_TAG,),
|
||||
dist_info: str = _DIST_INFO,
|
||||
duplicate_wheel: bool = False,
|
||||
) -> Path:
|
||||
wheel: Final = tmp_path / f"litellm-1.100.0-{filename_tag}.whl"
|
||||
with zipfile.ZipFile(wheel, "w", compression=zipfile.ZIP_DEFLATED) as archive:
|
||||
archive.writestr(_NATIVE_MEMBER, b"synthetic native extension")
|
||||
archive.writestr(
|
||||
f"{dist_info}/METADATA",
|
||||
"Metadata-Version: 2.1\nName: litellm\nVersion: 1.100.0\n",
|
||||
)
|
||||
archive.writestr(
|
||||
f"{dist_info}/RECORD",
|
||||
f"{_NATIVE_MEMBER},,\n{dist_info}/WHEEL,,\n",
|
||||
)
|
||||
if metadata_tags is not None:
|
||||
wheel_metadata: Final = (
|
||||
"Wheel-Version: 1.0\nGenerator: regression-test\nRoot-Is-Purelib: false\n"
|
||||
+ "".join(f"Tag: {tag}\n" for tag in metadata_tags)
|
||||
)
|
||||
archive.writestr(f"{dist_info}/WHEEL", wheel_metadata)
|
||||
if duplicate_wheel:
|
||||
archive.writestr(f"{dist_info}/WHEEL", wheel_metadata)
|
||||
return wheel
|
||||
|
||||
|
||||
def _fake_subprocess_run(command: tuple[str, ...], **_: object) -> subprocess.CompletedProcess[str]:
|
||||
if command == ("rustc", "--version"):
|
||||
stdout = "rustc 1.98.0 (regression-test)\n"
|
||||
elif "--sections" in command:
|
||||
stdout = "[ 1] .text PROGBITS\n"
|
||||
elif "--dyn-syms" in command:
|
||||
stdout = "PyInit__native\n"
|
||||
else:
|
||||
raise AssertionError(f"unexpected subprocess command: {command}")
|
||||
return subprocess.CompletedProcess(command, 0, stdout=stdout, stderr="")
|
||||
|
||||
|
||||
def _run_verifier(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
wheel: Path,
|
||||
*,
|
||||
exposes_panic: bool = False,
|
||||
) -> int:
|
||||
native_module: Final = ModuleType("litellm.rust_bridge._native")
|
||||
if exposes_panic:
|
||||
setattr(native_module, "_panic_for_test", lambda: None)
|
||||
|
||||
def _fake_load_native_module(_: Path) -> ModuleType:
|
||||
return native_module
|
||||
|
||||
monkeypatch.setattr(verifier, "_load_native_module", _fake_load_native_module)
|
||||
monkeypatch.setattr(verifier.subprocess, "run", _fake_subprocess_run)
|
||||
monkeypatch.setattr(sys, "argv", [str(_MODULE_PATH), str(wheel)])
|
||||
monkeypatch.setenv("GITHUB_STEP_SUMMARY", str(wheel.parent / "summary.md"))
|
||||
return verifier.main()
|
||||
|
||||
|
||||
def test_accepts_expected_release_wheel_tags(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
wheel: Final = _write_wheel(tmp_path, filename_tag=_EXPECTED_TAG)
|
||||
|
||||
assert _run_verifier(monkeypatch, wheel) == 0
|
||||
|
||||
|
||||
def test_rejects_cp312_version_specific_wheel(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
tag: Final = "cp312-cp312-linux_x86_64"
|
||||
wheel: Final = _write_wheel(tmp_path, filename_tag=tag, metadata_tags=(tag,))
|
||||
|
||||
assert _run_verifier(monkeypatch, wheel) == 1
|
||||
|
||||
|
||||
def test_rejects_non_linux_platform_tag(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
tag: Final = "cp310-abi3-win_amd64"
|
||||
wheel: Final = _write_wheel(tmp_path, filename_tag=tag, metadata_tags=(tag,))
|
||||
|
||||
assert _run_verifier(monkeypatch, wheel) == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"metadata_tags",
|
||||
[None, ("cp312-cp312-linux_x86_64",)],
|
||||
ids=["missing", "mismatched"],
|
||||
)
|
||||
def test_rejects_missing_or_mismatched_wheel_metadata_tag(
|
||||
tmp_path: Path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
metadata_tags: tuple[str, ...] | None,
|
||||
) -> None:
|
||||
wheel: Final = _write_wheel(tmp_path, filename_tag=_EXPECTED_TAG, metadata_tags=metadata_tags)
|
||||
|
||||
assert _run_verifier(monkeypatch, wheel) == 1
|
||||
|
||||
|
||||
def test_rejects_wheel_metadata_from_wrong_dist_info_directory(
|
||||
tmp_path: Path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
wheel: Final = _write_wheel(
|
||||
tmp_path,
|
||||
filename_tag=_EXPECTED_TAG,
|
||||
dist_info="decoy-1.0.0.dist-info",
|
||||
)
|
||||
|
||||
assert _run_verifier(monkeypatch, wheel) == 1
|
||||
|
||||
|
||||
def test_rejects_duplicate_wheel_metadata_tags(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
wheel: Final = _write_wheel(
|
||||
tmp_path,
|
||||
filename_tag=_EXPECTED_TAG,
|
||||
metadata_tags=(_EXPECTED_TAG, _EXPECTED_TAG),
|
||||
)
|
||||
|
||||
assert _run_verifier(monkeypatch, wheel) == 1
|
||||
|
||||
|
||||
def test_rejects_duplicate_wheel_metadata_file(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
with pytest.warns(UserWarning, match="Duplicate name"):
|
||||
wheel: Final = _write_wheel(
|
||||
tmp_path,
|
||||
filename_tag=_EXPECTED_TAG,
|
||||
duplicate_wheel=True,
|
||||
)
|
||||
|
||||
assert _run_verifier(monkeypatch, wheel) == 1
|
||||
|
||||
|
||||
def test_rejects_production_module_exposing_panic_hook(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
wheel: Final = _write_wheel(tmp_path, filename_tag=_EXPECTED_TAG)
|
||||
|
||||
assert _run_verifier(monkeypatch, wheel, exposes_panic=True) == 1
|
||||
Loading…
Add table
Reference in a new issue