Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 9 additions & 5 deletions flashinfer/artifacts.py
Original file line number Diff line number Diff line change
Expand Up @@ -207,10 +207,13 @@ def get_checksums(subdirs):
for line in f:
sha256, filename = line.strip().split()

# Distinguish between all meta info header files
if ".h" in filename:
filename = safe_urljoin(subdir, filename)
checksums[filename] = sha256
# Key every entry by its full path. Bare filenames are not
# unique across subdirs: per-arch cubin pins publish the same
# kernel names built from different sources, so a flat dict
# lets whichever subdir is processed last silently overwrite
# the earlier one's hashes, failing verification for every
# shared name.
checksums[safe_urljoin(subdir, filename)] = sha256
return checksums


Expand Down Expand Up @@ -268,7 +271,8 @@ def get_subdir_file_list() -> Generator[tuple[str, str], None, None]:
checksum_path = safe_urljoin(cubin_dir, "checksums.txt")
yield (checksum_path, CheckSumHash.map_checksums[checksum_path])
for name in get_available_cubin_files(safe_urljoin(base, cubin_dir)):
yield (safe_urljoin(cubin_dir, name), checksums[name])
full_path = safe_urljoin(cubin_dir, name)
yield (full_path, checksums[full_path])
for name in get_available_header_files(safe_urljoin(base, cubin_dir)):
full_path = safe_urljoin(cubin_dir, name)
yield (full_path, checksums[full_path])
Expand Down
64 changes: 64 additions & 0 deletions tests/test_artifacts.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
from flashinfer.artifacts import (
ArtifactPath,
get_available_cubin_files,
get_checksums,
get_subdir_file_list,
)

Expand Down Expand Up @@ -413,3 +414,66 @@ def test_get_subdir_file_list(monkeypatch, tmp_path):
assert len(meta_info_headers) == 3, (
f"Meta info headers count mismatch. Expected 3, got {len(meta_info_headers)}. Headers found: {meta_info_headers}"
)


@responses.activate
def test_get_checksums_does_not_collide_across_subdirs(monkeypatch, tmp_path):
"""Two pins publishing the same kernel filename must keep separate hashes.

Per-arch cubin pins ship identically-named sm100f/sm103a kernels built from
different sources. Keying the checksum map by bare filename let whichever
subdir was processed last overwrite the earlier one, so every shared name
was then verified against the wrong pin's hash.
"""
from flashinfer import artifacts

temp_cubin_dir = tmp_path / "cubins"
temp_cubin_dir.mkdir(exist_ok=True)
monkeypatch.setattr(
artifacts, "FLASHINFER_CUBINS_REPOSITORY", test_cubin_repository
)
monkeypatch.setattr(artifacts, "FLASHINFER_CUBIN_DIR", temp_cubin_dir)

shared_kernel = "Bmm_Bfloat16_E2m1E2m1_Fp32_shared_name_sm100f.cubin"
subdir_a = "pin-a/batched_gemm-aaaaaaa-bbbbbbb/"
subdir_b = "pin-b/batched_gemm-ccccccc-ddddddd/"
hash_a = "1111111111111111111111111111111111111111111111111111111111111111"
hash_b = "2222222222222222222222222222222222222222222222222222222222222222"
header_a = "3333333333333333333333333333333333333333333333333333333333333333"
header_b = "4444444444444444444444444444444444444444444444444444444444444444"

for subdir, kernel_hash, header_hash in (
(subdir_a, hash_a, header_a),
(subdir_b, hash_b, header_b),
):
responses.add(
responses.GET,
safe_urljoin(test_cubin_repository, safe_urljoin(subdir, "checksums.txt")),
body=(
f"{header_hash} include/flashinferMetaInfo.h\n"
f"{kernel_hash} {shared_kernel}\n"
),
status=200,
)

checksums = get_checksums([subdir_a, subdir_b])

key_a = safe_urljoin(subdir_a, shared_kernel)
key_b = safe_urljoin(subdir_b, shared_kernel)

# The bare filename must not be a key at all -- that is what collided.
assert shared_kernel not in checksums

assert checksums[key_a] == hash_a, (
f"{subdir_a} kernel resolved to {checksums[key_a]}, expected {hash_a}"
)
assert checksums[key_b] == hash_b, (
f"{subdir_b} kernel resolved to {checksums[key_b]}, expected {hash_b}"
)
assert checksums[key_a] != checksums[key_b], (
"both pins resolved to the same checksum -- the per-pin hashes collided"
)

# Headers were already namespaced before this fix; keep them covered.
assert checksums[safe_urljoin(subdir_a, "include/flashinferMetaInfo.h")] == header_a
assert checksums[safe_urljoin(subdir_b, "include/flashinferMetaInfo.h")] == header_b
Loading