diff --git a/litellm/rust_bridge/verify_linux_native_wheel.py b/litellm/rust_bridge/verify_linux_native_wheel.py index 783de87a8b2..899e2a211c0 100644 --- a/litellm/rust_bridge/verify_linux_native_wheel.py +++ b/litellm/rust_bridge/verify_linux_native_wheel.py @@ -7,12 +7,10 @@ import subprocess import sys import zipfile from collections.abc import Callable, Mapping, Sequence -from email import policy -from email.parser import BytesParser from itertools import product from pathlib import Path, PurePosixPath -from types import ModuleType -from typing import Final, Protocol, cast +from types import MappingProxyType, ModuleType +from typing import Final, Protocol EXPECTED_PYTHON_TAG: Final = "cp310" EXPECTED_ABI_TAG: Final = "abi3" @@ -50,9 +48,8 @@ def _dist_info_directory(member: zipfile.ZipInfo) -> str | None: 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) + lines: Final = archive.read(members[0]).splitlines() + return tuple(line.removeprefix(b"Tag:").strip().decode("ascii") for line in lines if line.startswith(b"Tag:")) def _load_native_module(native_path: Path) -> ModuleType | None: @@ -62,7 +59,7 @@ def _load_native_module(native_path: Path) -> ModuleType | None: try: native_module: Final = importlib.util.module_from_spec(module_spec) module_spec.loader.exec_module(native_module) - except Exception as error: + except Exception as error: # noqa: BLE001 # native module initialization can raise arbitrary exceptions sys.stderr.write(f"native module load failed: {error}\n") return None return native_module @@ -106,10 +103,14 @@ def main( 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 - } + dist_info_file_counts: Final = MappingProxyType( + { + 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" ) @@ -233,36 +234,46 @@ def main( if summary_path is not None: Path(summary_path).write_text(verified_report) - if debug_sections: - sys.stderr.write(f"{native_member.filename} contains debug sections: {', '.join(debug_sections)}\n") - if not static_symbol_table_absent: - 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" + invalid_dist_info_files: Final = any(count != 1 for count in dist_info_file_counts.values()) + validation_errors: Final = tuple( + message + for failed, message in ( + (bool(debug_sections), f"{native_member.filename} contains debug sections: {', '.join(debug_sections)}"), + (not static_symbol_table_absent, f"{native_member.filename} contains a static symbol table"), + (not extension_entry_point_present, "native extension does not export PyInit__native"), + ( + python_tag != EXPECTED_PYTHON_TAG, + f"unexpected Python tag: expected {EXPECTED_PYTHON_TAG}, found {python_tag}", + ), + (abi_tag != EXPECTED_ABI_TAG, f"unexpected ABI tag: expected {EXPECTED_ABI_TAG}, found {abi_tag}"), + ( + platform_tag != EXPECTED_PLATFORM_TAG, + f"unexpected platform tag: expected {EXPECTED_PLATFORM_TAG}, found {platform_tag}", + ), + ( + dist_info_directories != expected_dist_info_directories, + f"unexpected dist-info directories: expected {expected_dist_info_directory}, " + f"found {', '.join(sorted(dist_info_directories))}", + ), + (invalid_dist_info_files, f"required dist-info file counts are invalid: {dist_info_file_counts}"), + ( + not invalid_dist_info_files and not wheel_metadata_tags_match, + f"WHEEL tags do not match filename: expected {', '.join(sorted(expanded_filename_tags))}, " + f"found {', '.join(sorted(wheel_metadata_tags))}", + ), + ( + native_module is not None and not panic_test_hook_absent, + "production native module exposes _panic_for_test", + ), + ( + not native_size_within_limit, + f"native extension exceeds 20 MB: {native_member.file_size / 1_000_000:.2f} MB", + ), + (bool(unexpected_members), f"wheel contains unexpected build artifacts: {', '.join(unexpected_members)}"), ) - 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: - sys.stderr.write(f"wheel contains unexpected build artifacts: {', '.join(unexpected_members)}\n") + if failed + ) + sys.stderr.write("".join(f"{message}\n" for message in validation_errors)) return 0 if all(passed for _, passed in validations) else 1 diff --git a/tests/test_litellm/rust_bridge/test_verify_linux_native_wheel.py b/tests/test_litellm/rust_bridge/test_verify_linux_native_wheel.py index a4291ce0a65..8d0082dddc1 100644 --- a/tests/test_litellm/rust_bridge/test_verify_linux_native_wheel.py +++ b/tests/test_litellm/rust_bridge/test_verify_linux_native_wheel.py @@ -3,7 +3,7 @@ from __future__ import annotations import subprocess import zipfile from pathlib import Path -from types import ModuleType +from types import MappingProxyType, ModuleType from typing import Final import pytest @@ -47,16 +47,26 @@ def _write_wheel( return wheel -def _fake_subprocess_run(command: tuple[str, ...], **_: object) -> subprocess.CompletedProcess[str]: +def _fake_subprocess_run( + command: tuple[str, ...], + *, + check: bool, + capture_output: bool, + text: bool, +) -> subprocess.CompletedProcess[str]: + assert check and capture_output and text 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="") + return subprocess.CompletedProcess(command, 0, stdout="rustc 1.98.0 (regression-test)\n", stderr="") + if "--sections" in command: + return subprocess.CompletedProcess(command, 0, stdout="[ 1] .text PROGBITS\n", stderr="") + if "--dyn-syms" in command: + return subprocess.CompletedProcess(command, 0, stdout="PyInit__native\n", stderr="") + raise AssertionError(f"unexpected subprocess command: {command}") + + +class _NativeModuleWithPanicHook(ModuleType): + def _panic_for_test(self) -> None: + return None def _run_verifier( @@ -64,14 +74,16 @@ def _run_verifier( *, exposes_panic: bool = False, ) -> int: - native_module: Final = ModuleType("litellm.rust_bridge._native") - if exposes_panic: - setattr(native_module, "_panic_for_test", lambda: None) + native_module: Final = ( + _NativeModuleWithPanicHook("litellm.rust_bridge._native") + if exposes_panic + else ModuleType("litellm.rust_bridge._native") + ) def _fake_load_native_module(_: Path) -> ModuleType: return native_module - environment: Final = {"GITHUB_STEP_SUMMARY": str(wheel.parent / "summary.md")} + environment: Final = MappingProxyType({"GITHUB_STEP_SUMMARY": str(wheel.parent / "summary.md")}) return verifier.main( (str(_MODULE_PATH), str(wheel)), environment, @@ -102,8 +114,8 @@ def test_rejects_non_linux_platform_tag(tmp_path: Path) -> None: @pytest.mark.parametrize( "metadata_tags", - [None, ("cp312-cp312-linux_x86_64",)], - ids=["missing", "mismatched"], + (None, ("cp312-cp312-linux_x86_64",)), + ids=("missing", "mismatched"), ) def test_rejects_missing_or_mismatched_wheel_metadata_tag( tmp_path: Path,