diff --git a/CHANGELOG.md b/CHANGELOG.md index b1f6d6654..2457d06a8 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Bug Fixes +- keep strongly identified SentencePiece tokenizer `.model` artifacts, including custom-unknown models with disabled special tokens, out of XGBoost routing while preserving fail-closed `.model` ambiguity for malformed tails. - tolerate picklescan reports that omit private metadata - stream bounded Llamafile runtime coverage across preview gaps and report incomplete runtime reads or bounds - treat canonical storage persistent IDs as informational only after legacy PyTorch framing, storage tuples, and storage payloads validate completely diff --git a/modelaudit/core.py b/modelaudit/core.py index 304f29a1d..2d0bc13ab 100644 --- a/modelaudit/core.py +++ b/modelaudit/core.py @@ -82,9 +82,11 @@ def shared_source_sensitive_caches() -> Iterator[None]: ONNX_ROUTING_INCONCLUSIVE_FORMAT, PICKLE_ROUTING_INCONCLUSIVE_FORMAT, PROTOBUF_MODEL_CANDIDATE_FORMAT, + SENTENCEPIECE_MODEL_PROTO_INCONCLUSIVE_FORMAT, TENSORFLOW_PROTOBUF_ROUTING_INCONCLUSIVE_FORMAT, XGBOOST_UBJSON_ROUTING_INCONCLUSIVE_FORMAT, XML_MODEL_INCONCLUSIVE_FORMAT, + _is_malformed_sentencepiece_model_proto_candidate_file, detect_file_format, detect_file_format_from_magic, detect_flax_msgpack_overlap_routes, @@ -96,6 +98,7 @@ def shared_source_sensitive_caches() -> Iterator[None]: is_executorch_archive, is_keras_zip_archive, is_pytorch_zip_archive, + is_sentencepiece_model_proto_file, is_skops_archive, is_torchserve_mar_archive, should_defer_safetensors_header_limit_hash, @@ -277,6 +280,7 @@ def is_covered(target: Path) -> bool: _FORMAT_DETECTION_READ_FAILED_REASON = "format_detection_read_failed" _XML_MODEL_ROUTING_INCOMPLETE_REASON = "xml_model_routing_incomplete" _PROTOBUF_MODEL_ROUTING_INCOMPLETE_REASON = "protobuf_model_routing_incomplete" +_SENTENCEPIECE_MODEL_PROTO_ROUTING_INCOMPLETE_REASON = "sentencepiece_model_proto_routing_incomplete" _LLAMAFILE_ROUTING_INCOMPLETE_REASON = "llamafile_routing_incomplete" _MXNET_SYMBOL_ROUTING_INCOMPLETE_REASON = "mxnet_symbol_routing_incomplete" _PICKLE_ROUTING_INCOMPLETE_REASON = "pickle_routing_incomplete" @@ -1504,6 +1508,26 @@ def _make_incomplete_protobuf_model_result(path: str) -> ScanResult: return result +def _make_incomplete_sentencepiece_model_proto_result(path: str) -> ScanResult: + """Fail closed when a SentencePiece-like protobuf fails ownership validation.""" + result = ScanResult(scanner_name="unknown") + result.add_check( + name="SentencePiece ModelProto Routing", + passed=False, + message=( + "SentencePiece ModelProto routing was inconclusive because the payload " + "looked like a tokenizer protobuf but failed ownership validation" + ), + severity=IssueSeverity.INFO, + location=path, + details={"format": SENTENCEPIECE_MODEL_PROTO_INCONCLUSIVE_FORMAT, "path": path}, + ) + _mark_inconclusive_scan_outcome(result, _SENTENCEPIECE_MODEL_PROTO_ROUTING_INCOMPLETE_REASON) + _mark_operational_scan_error(result, _SENTENCEPIECE_MODEL_PROTO_ROUTING_INCOMPLETE_REASON) + result.finish(success=False) + return result + + def _make_incomplete_llamafile_routing_result(path: str, config: dict[str, Any]) -> ScanResult: """Fail closed when an executable Llamafile marker probe cannot complete.""" result = ScanResult(scanner_name="unknown") @@ -3654,6 +3678,7 @@ def _scan_file_internal(path: str, config: dict[str, Any] | None = None) -> Scan format_probe_error is None and header_format in {"unknown", "pytorch_binary"} and pytorch_binary_supplemental_scanner_id is None + and not (ext == ".model" and _is_malformed_sentencepiece_model_proto_candidate_file(path)) and scanner_selection.allows("zip") and ZipScanner.can_handle(path) ): @@ -3727,6 +3752,13 @@ def _scan_file_internal(path: str, config: dict[str, Any] | None = None) -> Scan if nested_xgboost_route == "xgboost": config[XGBOOST_CONTENT_ROUTED_UBJSON_CONFIG_KEY] = True is_xgboost_pickle_spoof = ext in _XGBOOST_BINARY_EXTENSIONS and header_format == "pickle" + sentencepiece_model_proto_owned = ( + format_probe_error is None + and ext == ".model" + and header_format == "unknown" + and magic_format == "unknown" + and is_sentencepiece_model_proto_file(path) + ) # Record telemetry for file type detection detected_format = header_format if header_format != "unknown" else ext_format record_file_type_detected(path, detected_format) @@ -3772,6 +3804,14 @@ def _scan_file_internal(path: str, config: dict[str, Any] | None = None) -> Scan if sr.bytes_scanned == 0 and file_size > 0: sr.bytes_scanned = file_size return sr + if ( + header_format == SENTENCEPIECE_MODEL_PROTO_INCONCLUSIVE_FORMAT + or magic_format == SENTENCEPIECE_MODEL_PROTO_INCONCLUSIVE_FORMAT + ): + sr = _make_incomplete_sentencepiece_model_proto_result(path) + if sr.bytes_scanned == 0 and file_size > 0: + sr.bytes_scanned = file_size + return sr if header_format == PICKLE_ROUTING_INCONCLUSIVE_FORMAT or magic_format == PICKLE_ROUTING_INCONCLUSIVE_FORMAT: sr = _make_incomplete_pickle_routing_result(path) merge_safetensors_overlap_analysis( @@ -3990,7 +4030,7 @@ def _scan_file_internal(path: str, config: dict[str, Any] | None = None) -> Scan and scanner_selection.allows(fallback_scanner_id) ): scanner_class = _registry.load_scanner_by_id(fallback_scanner_id) - elif scanner_class is None: + elif scanner_class is None and not sentencepiece_model_proto_owned: scanner_class = _registry.get_scanner_for_path( path, scanner_selection=scanner_selection if scanner_selection.active else None, @@ -4066,7 +4106,11 @@ def _scan_file_internal(path: str, config: dict[str, Any] | None = None) -> Scan kind=SCANNER_SELECTION_PREFERRED_KIND, ) else: - if unavailable_preferred_scanner_id is None and scanner_selection.active: + if ( + unavailable_preferred_scanner_id is None + and scanner_selection.active + and not sentencepiece_model_proto_owned + ): candidate_scanner_id = skipped_preferred_scanner_id if candidate_scanner_id is None: candidate_scanner_class = _registry.get_scanner_for_path(path) @@ -4126,6 +4170,8 @@ def _scan_file_internal(path: str, config: dict[str, Any] | None = None) -> Scan sr = _make_incomplete_xgboost_ubjson_routing_result(path) elif magic_format == ONNX_ROUTING_INCONCLUSIVE_FORMAT: sr = _make_incomplete_onnx_routing_result(path) + elif magic_format == SENTENCEPIECE_MODEL_PROTO_INCONCLUSIVE_FORMAT: + sr = _make_incomplete_sentencepiece_model_proto_result(path) elif magic_format == PICKLE_ROUTING_INCONCLUSIVE_FORMAT: sr = _make_incomplete_pickle_routing_result(path) elif magic_format == TENSORFLOW_PROTOBUF_ROUTING_INCONCLUSIVE_FORMAT: diff --git a/modelaudit/scanners/__init__.py b/modelaudit/scanners/__init__.py index 57fba2129..5aa6460b6 100644 --- a/modelaudit/scanners/__init__.py +++ b/modelaudit/scanners/__init__.py @@ -367,6 +367,11 @@ def get_scanner_for_path( # If stricter extension-specific scanners all decline, fall back to the # generic ZIP scanner so helper-level routing does not drop coverage. if is_zip_file and (scanner_selection is None or scanner_selection.allows("zip")): + if file_ext == ".model": + from modelaudit.utils.file.detection import _is_malformed_sentencepiece_model_proto_candidate_file + + if _is_malformed_sentencepiece_model_proto_candidate_file(path): + return None scanner_class = self._load_scanner("zip") if scanner_class and scanner_class.can_handle(path): return scanner_class diff --git a/modelaudit/scanners/archive_dispatch.py b/modelaudit/scanners/archive_dispatch.py index cee4653a2..4475847ca 100644 --- a/modelaudit/scanners/archive_dispatch.py +++ b/modelaudit/scanners/archive_dispatch.py @@ -38,9 +38,11 @@ ONNX_ROUTING_INCONCLUSIVE_FORMAT, PICKLE_ROUTING_INCONCLUSIVE_FORMAT, PROTOBUF_MODEL_CANDIDATE_FORMAT, + SENTENCEPIECE_MODEL_PROTO_INCONCLUSIVE_FORMAT, TENSORFLOW_PROTOBUF_ROUTING_INCONCLUSIVE_FORMAT, XGBOOST_UBJSON_ROUTING_INCONCLUSIVE_FORMAT, XML_MODEL_INCONCLUSIVE_FORMAT, + _is_malformed_sentencepiece_model_proto_candidate_file, detect_file_format, detect_file_format_from_magic, detect_flax_msgpack_overlap_routes, @@ -87,6 +89,7 @@ def _build_header_format_to_scanner_id() -> dict[str, str]: _RECOGNIZED_FORMAT_SCANNER_UNAVAILABLE_REASON = "recognized_format_scanner_unavailable" _XML_MODEL_ROUTING_INCOMPLETE_REASON = "xml_model_routing_incomplete" _PROTOBUF_MODEL_ROUTING_INCOMPLETE_REASON = "protobuf_model_routing_incomplete" +_SENTENCEPIECE_MODEL_PROTO_ROUTING_INCOMPLETE_REASON = "sentencepiece_model_proto_routing_incomplete" _LLAMAFILE_ROUTING_INCOMPLETE_REASON = "llamafile_routing_incomplete" _MXNET_SYMBOL_ROUTING_INCOMPLETE_REASON = "mxnet_symbol_routing_incomplete" _PICKLE_ROUTING_INCOMPLETE_REASON = "pickle_routing_incomplete" @@ -426,6 +429,26 @@ def _make_incomplete_protobuf_model_result(path: str) -> ScanResult: return result +def _make_incomplete_sentencepiece_model_proto_result(path: str) -> ScanResult: + """Fail closed when a nested SentencePiece-like protobuf fails ownership validation.""" + result = ScanResult(scanner_name="unknown") + result.add_check( + name="SentencePiece ModelProto Routing", + passed=False, + message=( + "SentencePiece ModelProto routing was inconclusive because the payload " + "looked like a tokenizer protobuf but failed ownership validation" + ), + severity=IssueSeverity.INFO, + location=path, + details={"format": SENTENCEPIECE_MODEL_PROTO_INCONCLUSIVE_FORMAT, "path": path}, + ) + mark_inconclusive_scan_result(result, _SENTENCEPIECE_MODEL_PROTO_ROUTING_INCOMPLETE_REASON) + mark_operational_scan_error(result, _SENTENCEPIECE_MODEL_PROTO_ROUTING_INCOMPLETE_REASON) + result.finish(success=False) + return result + + def _deduplicate_exact_merged_findings(result: ScanResult) -> None: """Remove identical output emitted by composed subtype and ZIP analyses.""" @@ -1199,8 +1222,12 @@ def with_safetensors_overlap(result: ScanResult) -> ScanResult: except (TypeError, ValueError): max_zip_entries = ZipScanner.DEFAULT_MAX_ENTRIES max_zip_directory_size = ZipScanner.central_directory_size_limit(raw_config) + malformed_sentencepiece_model_candidate = Path( + path + ).suffix.lower() == ".model" and _is_malformed_sentencepiece_model_proto_candidate_file(path) if ( - (not is_safetensors_hdf5_overlap or hdf5_signature_offset in (None, 0)) + not malformed_sentencepiece_model_candidate + and (not is_safetensors_hdf5_overlap or hdf5_signature_offset in (None, 0)) and allows_zip_structure_analysis(scanner_selection, path) and ZipScanner.requires_preflight_result( path, @@ -1395,6 +1422,11 @@ def with_safetensors_overlap(result: ScanResult) -> ScanResult: if trusted_content_format == XML_MODEL_INCONCLUSIVE_FORMAT: return with_safetensors_overlap(_make_incomplete_xml_model_result(path)) + if ( + routed_content_format == SENTENCEPIECE_MODEL_PROTO_INCONCLUSIVE_FORMAT + or trusted_content_format == SENTENCEPIECE_MODEL_PROTO_INCONCLUSIVE_FORMAT + ): + return with_safetensors_overlap(_make_incomplete_sentencepiece_model_proto_result(path)) if routed_content_format == PROTOBUF_MODEL_CANDIDATE_FORMAT: return with_safetensors_overlap(_make_incomplete_protobuf_model_result(path)) if trusted_content_format == PROTOBUF_MODEL_CANDIDATE_FORMAT and routed_content_format != "unknown": diff --git a/modelaudit/scanners/compressed_scanner.py b/modelaudit/scanners/compressed_scanner.py index 85466853e..88afa30b2 100644 --- a/modelaudit/scanners/compressed_scanner.py +++ b/modelaudit/scanners/compressed_scanner.py @@ -15,6 +15,10 @@ from .. import core from ..scanner_results import INCONCLUSIVE_SCAN_OUTCOME, mark_inconclusive_scan_result from ..utils.file._compression import is_zlib_header +from ..utils.file.detection import ( + _is_malformed_sentencepiece_model_proto_candidate_file, + is_sentencepiece_model_proto_file, +) from ._archive_config import get_archive_depth from ._archive_locations import rewrite_extracted_member_location from ._archive_outcomes import member_scan_incomplete @@ -181,7 +185,44 @@ def _derive_inner_suffix(path: str) -> str: wrapper_path.name[: -len(wrapper_path.suffix)] if wrapper_path.suffix else wrapper_path.name ) inferred_suffix = Path(stem_without_wrapper).suffix - return inferred_suffix or ".bin" + if inferred_suffix: + return inferred_suffix + if stem_without_wrapper.lower() in {"tokenizer", "spiece", "sentencepiece"}: + return "" + return ".bin" + + @staticmethod + def _uses_tokenizer_extensionless_inner(path: str) -> bool: + wrapper_path = Path(path) + if not wrapper_path.suffix or not CompressedScanner._has_declared_wrapper_extension(path): + return False + stem_without_wrapper = wrapper_path.name[: -len(wrapper_path.suffix)] + return not Path(stem_without_wrapper).suffix and stem_without_wrapper.lower() in { + "tokenizer", + "spiece", + "sentencepiece", + } + + @staticmethod + def _is_sentencepiece_candidate(path: str) -> bool: + return is_sentencepiece_model_proto_file(path) or _is_malformed_sentencepiece_model_proto_candidate_file(path) + + @staticmethod + def _replace_temp_suffix(path: str, suffix: str) -> str: + fd, replacement_path = tempfile.mkstemp(suffix=suffix) + os.close(fd) + os.unlink(replacement_path) + os.replace(path, replacement_path) + return replacement_path + + @classmethod + def _route_tokenizer_extensionless_or_bin(cls, wrapper_path: str, temp_paths: list[str]) -> list[str]: + if not cls._uses_tokenizer_extensionless_inner(wrapper_path): + return temp_paths + return [ + temp_path if cls._is_sentencepiece_candidate(temp_path) else cls._replace_temp_suffix(temp_path, ".bin") + for temp_path in temp_paths + ] @staticmethod def _derive_inner_display_name(path: str) -> str: @@ -916,6 +957,9 @@ def scan(self, path: str) -> ScanResult: decompressed_bytes = 0 try: temp_path, member_temp_paths, decompressed_bytes = self._decompress_to_tempfiles(path, expected_codec) + routed_temp_paths = self._route_tokenizer_extensionless_or_bin(path, [temp_path, *member_temp_paths]) + temp_path = routed_temp_paths[0] + member_temp_paths = routed_temp_paths[1:] temp_paths = [temp_path, *member_temp_paths] result.metadata["decompressed_bytes"] = decompressed_bytes result.metadata["compressed_member_count"] = len(member_temp_paths) diff --git a/modelaudit/scanners/xgboost_scanner.py b/modelaudit/scanners/xgboost_scanner.py index 07693d64c..0042ab57a 100644 --- a/modelaudit/scanners/xgboost_scanner.py +++ b/modelaudit/scanners/xgboost_scanner.py @@ -30,7 +30,7 @@ from typing import Any, ClassVar, cast from ..scanner_selection import add_scanner_selection_skip_check, policy_from_config -from ..utils.file.detection import has_jax_json_checkpoint_structure +from ..utils.file.detection import has_jax_json_checkpoint_structure, is_sentencepiece_model_proto_file from .base import INCONCLUSIVE_SCAN_OUTCOME, BaseScanner, IssueSeverity, ScanResult logger = logging.getLogger(__name__) @@ -513,9 +513,10 @@ def can_handle(cls, path: str) -> bool: if file_ext == ".json": return cls._is_xgboost_json(path) or cls._is_probable_xgboost_json_candidate(path) - # For .model files, accept (generic extension) + # For .model files, keep ambiguous inputs on the XGBoost route but do + # not claim strongly identified SentencePiece tokenizer ModelProto files. if file_ext == ".model": - return True + return not is_sentencepiece_model_proto_file(path) # Check for XGBoost files without extension if file_ext == "": diff --git a/modelaudit/utils/file/detection.py b/modelaudit/utils/file/detection.py index 6c4d6fe2d..c40177f55 100644 --- a/modelaudit/utils/file/detection.py +++ b/modelaudit/utils/file/detection.py @@ -13,6 +13,7 @@ import zlib from collections.abc import Callable, Iterator from dataclasses import dataclass, field +from functools import lru_cache from io import BytesIO, StringIO from pathlib import Path, PurePosixPath from typing import Any, BinaryIO, Literal, cast @@ -65,6 +66,7 @@ "inconclusive", ] _TensorFlowOuterHint = Literal["unknown", "tf_metagraph", "tf_savedmodel"] +_SentencePieceModelProtoRoute = Literal["unknown", "strong", "malformed_candidate"] _GzipTarTrailingStatus = Literal["invalid", "nonzero"] _TORCH7_SIGNATURE_READ_BYTES = 4096 _TORCH7_ASCII_HEADER_MAX_LINE_BYTES = 4096 @@ -90,6 +92,65 @@ _PROTO_GROUP_MAX_ROUTING_FIELDS = 512 _PROTO_GROUP_MAX_ROUTING_DEPTH = 8 _COREML_PROTO_SIGNATURE_READ_BYTES = 1024 * 1024 +_SENTENCEPIECE_MODEL_PROTO_READ_BYTES = 10 * 1024 * 1024 +_SENTENCEPIECE_MODEL_PROTO_CACHE_FINGERPRINT_BYTES = 4096 +_SENTENCEPIECE_MODEL_MAX_FIELDS = 512 * 1024 +_SENTENCEPIECE_MIN_STRONG_PIECES = 8 +_SENTENCEPIECE_MAX_PIECE_FIELDS = 16 +_SENTENCEPIECE_MAX_PIECE_MESSAGE_BYTES = 4096 +_SENTENCEPIECE_MAX_PIECE_TEXT_BYTES = 512 +_SENTENCEPIECE_MAX_TRAINER_SPEC_FIELDS = 512 +_SENTENCEPIECE_MAX_TRAINER_SPEC_MESSAGE_BYTES = 64 * 1024 +_SENTENCEPIECE_MAX_TRAINER_SPEC_TEXT_BYTES = 4096 +_SENTENCEPIECE_UNKNOWN_PIECE_TYPE = 2 +_SENTENCEPIECE_BYTE_PIECE_TYPE = 6 +_SENTENCEPIECE_BYTE_FALLBACK_PIECE_COUNT = 256 +_SENTENCEPIECE_IDENTITY_TOKENS = frozenset({"", "", "", "", "", ""}) +_SENTENCEPIECE_BYTE_FALLBACK_RE = re.compile(r"^<0x[0-9A-Fa-f]{2}>$") +_SENTENCEPIECE_TRAINER_SPEC_VARINT_FIELDS = frozenset( + { + 3, + 4, + 6, + 11, + 12, + 13, + 14, + 16, + 17, + 18, + 19, + 20, + 21, + 22, + 23, + 24, + 25, + 26, + 32, + 33, + 34, + 35, + 40, + 41, + 42, + 43, + 49, + 50, + 52, + } +) +_SENTENCEPIECE_TRAINER_SPEC_STRING_FIELDS = frozenset({1, 2, 5, 7, 30, 31, 36, 44, 45, 46, 47, 48, 53, 54}) +_SENTENCEPIECE_TRAINER_SPEC_FIXED32_FIELDS = frozenset({10, 15, 51}) +_SENTENCEPIECE_TRAINER_SPEC_FIXED64_FIELDS: frozenset[int] = frozenset() +_SENTENCEPIECE_NORMALIZER_SPEC_WIRE_TYPES = { + 1: 2, + 2: 2, + 3: 0, + 4: 0, + 5: 0, + 6: 2, +} _COREML_PROTO_PREFIX_WIRE_TYPES = frozenset({0, 1, 2, 3, 5}) _COREML_GROUP_BUDGET_EXHAUSTED: Literal["budget_exhausted"] = "budget_exhausted" _COREML_GROUP_INCOMPLETE: Literal["incomplete"] = "incomplete" @@ -231,6 +292,7 @@ } XML_MODEL_INCONCLUSIVE_FORMAT = "xml_model_inconclusive" PROTOBUF_MODEL_CANDIDATE_FORMAT = "protobuf_model_candidate" +SENTENCEPIECE_MODEL_PROTO_INCONCLUSIVE_FORMAT = "sentencepiece_model_proto_inconclusive" JAX_JSON_CHECKPOINT_STRUCTURE_READ_BYTES = 1024 * 1024 JAX_JSON_CHECKPOINT_ROUTING_READ_BYTES = 2 * JAX_JSON_CHECKPOINT_STRUCTURE_READ_BYTES _JAX_JSON_CHECKPOINT_IDENTITY_KEYS = frozenset( @@ -1397,6 +1459,810 @@ def _skip_proto_value(data: bytes, offset: int, wire_type: int, end: int | None return None +@dataclass +class _SentencePieceTrainerSpecSignals: + model_type: int | None = None + vocab_size: int | None = None + unk_id: int = 0 + unk_id_explicit: bool = False + unk_piece: str | None = None + unk_piece_explicit: bool = False + byte_fallback: bool = False + byte_fallback_explicit: bool = False + + @property + def has_core_metadata(self) -> bool: + return self.model_type is not None and self.vocab_size is not None + + def merge_from(self, other: "_SentencePieceTrainerSpecSignals") -> None: + if other.model_type is not None: + self.model_type = other.model_type + if other.vocab_size is not None: + self.vocab_size = other.vocab_size + if other.unk_id_explicit: + self.unk_id = other.unk_id + self.unk_id_explicit = True + if other.unk_piece_explicit: + self.unk_piece = other.unk_piece + self.unk_piece_explicit = True + if other.byte_fallback_explicit: + self.byte_fallback = other.byte_fallback + self.byte_fallback_explicit = True + + +@dataclass +class _SentencePiecePieceProtoSignals: + piece_text: str | None = None + piece_type: int | None = None + decoded_text_bytes: int = 0 + + +def _decode_proto_int32_varint(value: int) -> int: + """Decode proto2 int32 values that may be sign-extended into a uint64 varint.""" + if value >= 1 << 63: + value -= 1 << 64 + return value + + +def _decode_bounded_proto_string(data: bytes, start: int, end: int, *, max_bytes: int) -> str | None: + if end < start or end - start > max_bytes: + return None + try: + value = data[start:end].decode("utf-8") + except UnicodeDecodeError: + return None + if not value or "\x00" in value: + return None + return value + + +def _parse_sentencepiece_piece_proto(data: bytes, start: int, end: int) -> tuple[str, int | None] | None: + """Return the token text and optional type for one SentencePiece piece.""" + if end - start > _SENTENCEPIECE_MAX_PIECE_MESSAGE_BYTES: + return None + + offset = start + fields_seen = 0 + piece_text: str | None = None + piece_type: int | None = None + has_score = False + while offset < end and fields_seen < _SENTENCEPIECE_MAX_PIECE_FIELDS: + tag_result = _read_proto_varint(data, offset, end) + if tag_result is None: + return None + tag, value_offset = tag_result + field_number = tag >> 3 + wire_type = tag & 0x07 + if field_number == 0: + return None + + if field_number == 1 and wire_type == 2: + bounds = _read_length_delimited_proto_value(data, value_offset, end) + if bounds is None: + return None + length, value_start, _value_end, actual_value_end = bounds + if length == 0 or length > _SENTENCEPIECE_MAX_PIECE_TEXT_BYTES or actual_value_end > end: + return None + piece_text = _decode_bounded_proto_string( + data, + value_start, + actual_value_end, + max_bytes=_SENTENCEPIECE_MAX_PIECE_TEXT_BYTES, + ) + if piece_text is None: + return None + offset = actual_value_end + elif field_number == 2 and wire_type == 5: + fixed32_end = value_offset + 4 + if fixed32_end > end: + return None + has_score = True + offset = fixed32_end + elif field_number == 3 and wire_type == 0: + type_result = _read_proto_varint(data, value_offset, end) + if type_result is None: + return None + piece_type, offset = type_result + if not 1 <= piece_type <= 6: + return None + else: + skipped_offset = _skip_proto_value(data, value_offset, wire_type, end) + if skipped_offset is None: + return None + offset = skipped_offset + fields_seen += 1 + + if offset != end or fields_seen >= _SENTENCEPIECE_MAX_PIECE_FIELDS: + return None + if piece_text is None or not has_score: + return None + return piece_text, piece_type + + +def _parse_sentencepiece_trainer_spec_proto( + data: bytes, + start: int, + end: int, +) -> _SentencePieceTrainerSpecSignals | None: + """Parse enough TrainerSpec structure to identify custom unknown-piece models.""" + if end - start > _SENTENCEPIECE_MAX_TRAINER_SPEC_MESSAGE_BYTES: + return None + + offset = start + fields_seen = 0 + signals = _SentencePieceTrainerSpecSignals() + while offset < end and fields_seen < _SENTENCEPIECE_MAX_TRAINER_SPEC_FIELDS: + tag_result = _read_proto_varint(data, offset, end) + if tag_result is None: + return None + tag, value_offset = tag_result + field_number = tag >> 3 + wire_type = tag & 0x07 + if field_number == 0: + return None + + if field_number in _SENTENCEPIECE_TRAINER_SPEC_VARINT_FIELDS: + if wire_type != 0: + return None + value_result = _read_proto_varint(data, value_offset, end) + if value_result is None: + return None + value, offset = value_result + if field_number == 3 and 1 <= value <= 4: + signals.model_type = value + elif field_number == 4 and value > 0: + signals.vocab_size = value + elif field_number == 35: + if value not in {0, 1}: + return None + signals.byte_fallback = bool(value) + signals.byte_fallback_explicit = True + elif field_number == 40: + signals.unk_id = _decode_proto_int32_varint(value) + signals.unk_id_explicit = True + elif field_number in _SENTENCEPIECE_TRAINER_SPEC_STRING_FIELDS: + if wire_type != 2: + return None + bounds = _read_length_delimited_proto_value(data, value_offset, end) + if bounds is None: + return None + _length, value_start, _sampled_value_end, actual_value_end = bounds + if actual_value_end > end: + return None + if field_number == 45: + signals.unk_piece = _decode_bounded_proto_string( + data, + value_start, + actual_value_end, + max_bytes=_SENTENCEPIECE_MAX_TRAINER_SPEC_TEXT_BYTES, + ) + if signals.unk_piece is None: + return None + signals.unk_piece_explicit = True + offset = actual_value_end + elif field_number in _SENTENCEPIECE_TRAINER_SPEC_FIXED32_FIELDS: + if wire_type != 5: + return None + offset = value_offset + 4 + if offset > end: + return None + elif field_number in _SENTENCEPIECE_TRAINER_SPEC_FIXED64_FIELDS: + if wire_type != 1: + return None + offset = value_offset + 8 + if offset > end: + return None + else: + next_offset = _skip_proto_value(data, value_offset, wire_type, end) + if next_offset is None: + return None + offset = next_offset + + fields_seen += 1 + + if offset != end or fields_seen >= _SENTENCEPIECE_MAX_TRAINER_SPEC_FIELDS: + return None + return signals + + +def _is_well_formed_sentencepiece_submessage( + data: bytes, + start: int, + end: int, + *, + expected_wire_types: dict[int, int] | None = None, + max_fields: int = _SENTENCEPIECE_MAX_TRAINER_SPEC_FIELDS, +) -> bool: + offset = start + fields_seen = 0 + while offset < end and fields_seen < max_fields: + tag_result = _read_proto_varint(data, offset, end) + if tag_result is None: + return False + tag, value_offset = tag_result + field_number = tag >> 3 + wire_type = tag & 0x07 + if field_number == 0: + return False + if expected_wire_types is not None and expected_wire_types.get(field_number, wire_type) != wire_type: + return False + next_offset = _skip_proto_value(data, value_offset, wire_type, end) + if next_offset is None: + return False + offset = next_offset + fields_seen += 1 + + return offset == end and fields_seen < max_fields + + +def _is_well_formed_sentencepiece_submessage_stream( + stream: BinaryIO, + end_offset: int, + *, + expected_wire_types: dict[int, int] | None = None, + max_fields: int = _SENTENCEPIECE_MAX_TRAINER_SPEC_FIELDS, +) -> bool: + fields_seen = 0 + while stream.tell() < end_offset and fields_seen < max_fields: + tag = _read_proto_varint_stream(stream, end_offset) + if tag is None: + return False + field_number = tag >> 3 + wire_type = tag & 0x07 + if field_number == 0: + return False + if expected_wire_types is not None and expected_wire_types.get(field_number, wire_type) != wire_type: + return False + skip_status = _skip_proto_stream_value( + stream, + wire_type, + end_offset, + field_number=field_number, + ) + if skip_status is not True: + return False + fields_seen += 1 + + return stream.tell() == end_offset and fields_seen < max_fields + + +def _is_sentencepiece_special_identity_piece(piece: str) -> bool: + return piece in _SENTENCEPIECE_IDENTITY_TOKENS + + +def _is_sentencepiece_byte_fallback_piece(piece: str) -> bool: + return _SENTENCEPIECE_BYTE_FALLBACK_RE.fullmatch(piece) is not None + + +def _has_strong_sentencepiece_model_proto_evidence( + *, + piece_count: int, + typed_piece_count: int, + special_identity_piece_count: int, + unknown_piece_count: int, + unknown_piece_index: int | None, + unknown_piece_text: str | None, + byte_piece_count: int, + byte_piece_texts: set[str], + malformed_byte_piece: bool, + trainer_spec: _SentencePieceTrainerSpecSignals | None, +) -> bool: + if unknown_piece_count != 1 or unknown_piece_index is None or unknown_piece_text is None: + return False + if malformed_byte_piece: + return False + if byte_piece_count: + if trainer_spec is None or not trainer_spec.byte_fallback: + return False + if ( + byte_piece_count != _SENTENCEPIECE_BYTE_FALLBACK_PIECE_COUNT + or len(byte_piece_texts) != _SENTENCEPIECE_BYTE_FALLBACK_PIECE_COUNT + ): + return False + elif trainer_spec is not None and trainer_spec.byte_fallback: + return False + + if ( + trainer_spec is not None + and trainer_spec.has_core_metadata + and trainer_spec.vocab_size == piece_count + and 0 <= trainer_spec.unk_id < piece_count + and trainer_spec.unk_id == unknown_piece_index + and (not trainer_spec.unk_piece_explicit or trainer_spec.unk_piece == unknown_piece_text) + ): + return True + + if piece_count < _SENTENCEPIECE_MIN_STRONG_PIECES: + return False + return typed_piece_count >= 3 and special_identity_piece_count >= 3 + + +def _has_sufficient_sentencepiece_piece_scan_evidence( + *, + piece_count: int, + unknown_piece_count: int, + unknown_piece_index: int | None, + unknown_piece_text: str | None, + byte_piece_count: int, + byte_piece_texts: set[str], + malformed_byte_piece: bool, +) -> bool: + if piece_count < _SENTENCEPIECE_MIN_STRONG_PIECES: + return False + if unknown_piece_count != 1 or unknown_piece_index is None or unknown_piece_text is None: + return False + if malformed_byte_piece: + return False + return not byte_piece_count or ( + byte_piece_count == _SENTENCEPIECE_BYTE_FALLBACK_PIECE_COUNT + and len(byte_piece_texts) == _SENTENCEPIECE_BYTE_FALLBACK_PIECE_COUNT + ) + + +def _has_strong_sentencepiece_model_proto_prefix(data: bytes, *, sample_is_prefix: bool = False) -> bool: + """Recognize a SentencePiece ModelProto from repeated scored pieces.""" + offset = 0 + fields_seen = 0 + piece_count = 0 + typed_piece_count = 0 + special_identity_piece_count = 0 + unknown_piece_count = 0 + unknown_piece_index: int | None = None + unknown_piece_text: str | None = None + byte_piece_count = 0 + byte_piece_texts: set[str] = set() + malformed_byte_piece = False + trainer_spec: _SentencePieceTrainerSpecSignals | None = None + strong_match = False + + def accept_incomplete_prefix() -> bool: + return False + + while offset < len(data) and fields_seen < _SENTENCEPIECE_MODEL_MAX_FIELDS: + tag_result = _read_proto_varint(data, offset) + if tag_result is None: + return accept_incomplete_prefix() + tag, value_offset = tag_result + field_number = tag >> 3 + wire_type = tag & 0x07 + if field_number == 0: + return False + if wire_type not in {0, 1, 2, 5}: + return False + if field_number in {1, 2, 3, 4, 5} and wire_type != 2: + return False + + if field_number == 1: + bounds = _read_length_delimited_proto_value(data, value_offset) + if bounds is None: + return accept_incomplete_prefix() + length, value_start, _sampled_value_end, actual_value_end = bounds + if length == 0 or actual_value_end > len(data): + return accept_incomplete_prefix() + parsed_piece = _parse_sentencepiece_piece_proto(data, value_start, actual_value_end) + if parsed_piece is None: + return False + piece, piece_type = parsed_piece + piece_index = piece_count + piece_count += 1 + if piece_type is not None: + typed_piece_count += 1 + if _is_sentencepiece_special_identity_piece(piece): + special_identity_piece_count += 1 + if piece_type == _SENTENCEPIECE_UNKNOWN_PIECE_TYPE: + unknown_piece_count += 1 + unknown_piece_index = piece_index + unknown_piece_text = piece + elif piece_type == _SENTENCEPIECE_BYTE_PIECE_TYPE: + byte_piece_count += 1 + if _is_sentencepiece_byte_fallback_piece(piece): + byte_piece_texts.add(piece) + else: + malformed_byte_piece = True + offset = actual_value_end + elif field_number == 2: + bounds = _read_length_delimited_proto_value(data, value_offset) + if bounds is None: + return accept_incomplete_prefix() + _length, value_start, _sampled_value_end, actual_value_end = bounds + if actual_value_end > len(data): + return accept_incomplete_prefix() + parsed_trainer_spec = _parse_sentencepiece_trainer_spec_proto(data, value_start, actual_value_end) + if parsed_trainer_spec is None: + return False + if trainer_spec is None: + trainer_spec = parsed_trainer_spec + else: + trainer_spec.merge_from(parsed_trainer_spec) + offset = actual_value_end + elif field_number == 3: + bounds = _read_length_delimited_proto_value(data, value_offset) + if bounds is None: + return accept_incomplete_prefix() + _length, value_start, _sampled_value_end, actual_value_end = bounds + if actual_value_end > len(data): + return accept_incomplete_prefix() + if not _is_well_formed_sentencepiece_submessage( + data, + value_start, + actual_value_end, + expected_wire_types=_SENTENCEPIECE_NORMALIZER_SPEC_WIRE_TYPES, + ): + return False + offset = actual_value_end + elif field_number in {4, 5}: + bounds = _read_length_delimited_proto_value(data, value_offset) + if bounds is None: + return accept_incomplete_prefix() + _length, value_start, _sampled_value_end, actual_value_end = bounds + if actual_value_end > len(data): + return accept_incomplete_prefix() + if not _is_well_formed_sentencepiece_submessage(data, value_start, actual_value_end): + return False + offset = actual_value_end + else: + return False + + fields_seen += 1 + strong_match = _has_strong_sentencepiece_model_proto_evidence( + piece_count=piece_count, + typed_piece_count=typed_piece_count, + special_identity_piece_count=special_identity_piece_count, + unknown_piece_count=unknown_piece_count, + unknown_piece_index=unknown_piece_index, + unknown_piece_text=unknown_piece_text, + byte_piece_count=byte_piece_count, + byte_piece_texts=byte_piece_texts, + malformed_byte_piece=malformed_byte_piece, + trainer_spec=trainer_spec, + ) + + return ( + strong_match and not sample_is_prefix and offset == len(data) and fields_seen < _SENTENCEPIECE_MODEL_MAX_FIELDS + ) + + +def _read_bounded_sentencepiece_submessage(stream: BinaryIO, value_end: int, *, max_bytes: int) -> bytes | None: + length = value_end - stream.tell() + if length < 0 or length > max_bytes: + return None + payload = stream.read(length) + return payload if len(payload) == length else None + + +def _parse_sentencepiece_piece_proto_stream( + stream: BinaryIO, + value_end: int, + *, + decode_text: bool, + max_decoded_text_bytes: int, +) -> _SentencePiecePieceProtoSignals | None: + """Validate one piece submessage while avoiding unnecessary text reads.""" + if value_end - stream.tell() > _SENTENCEPIECE_MAX_PIECE_MESSAGE_BYTES: + return None + + fields_seen = 0 + text_bounds: tuple[int, int] | None = None + piece_type: int | None = None + has_score = False + while stream.tell() < value_end and fields_seen < _SENTENCEPIECE_MAX_PIECE_FIELDS: + tag = _read_proto_varint_stream(stream, value_end) + if tag is None: + return None + field_number = tag >> 3 + wire_type = tag & 0x07 + if field_number == 0: + return None + + if field_number == 1 and wire_type == 2: + bounds = _read_proto_length_delimited_bounds_stream(stream, value_end) + if bounds is None: + return None + length, value_start, actual_value_end = bounds + if length == 0 or length > _SENTENCEPIECE_MAX_PIECE_TEXT_BYTES: + return None + text_bounds = (value_start, actual_value_end) + stream.seek(actual_value_end) + elif field_number == 2 and wire_type == 5: + fixed32_end = stream.tell() + 4 + if fixed32_end > value_end: + return None + has_score = True + stream.seek(fixed32_end) + elif field_number == 3 and wire_type == 0: + parsed_type = _read_proto_varint_stream(stream, value_end) + if parsed_type is None or not 1 <= parsed_type <= 6: + return None + piece_type = parsed_type + else: + skip_status = _skip_proto_stream_value( + stream, + wire_type, + value_end, + field_number=field_number, + ) + if skip_status is not True: + return None + fields_seen += 1 + + if stream.tell() != value_end or fields_seen >= _SENTENCEPIECE_MAX_PIECE_FIELDS: + return None + if text_bounds is None or not has_score: + return None + + should_decode_text = decode_text or piece_type in { + _SENTENCEPIECE_UNKNOWN_PIECE_TYPE, + _SENTENCEPIECE_BYTE_PIECE_TYPE, + 3, + 4, + 5, + } + if not should_decode_text: + return _SentencePiecePieceProtoSignals(piece_type=piece_type) + + text_start, text_end = text_bounds + text_length = text_end - text_start + if text_length > max_decoded_text_bytes: + return None + stream.seek(text_start) + payload = stream.read(text_length) + if len(payload) != text_length: + return None + stream.seek(value_end) + piece_text = _decode_bounded_proto_string( + payload, + 0, + len(payload), + max_bytes=_SENTENCEPIECE_MAX_PIECE_TEXT_BYTES, + ) + if piece_text is None: + return None + return _SentencePiecePieceProtoSignals( + piece_text=piece_text, + piece_type=piece_type, + decoded_text_bytes=text_length, + ) + + +def _parse_sentencepiece_trainer_spec_proto_stream( + stream: BinaryIO, + value_end: int, +) -> _SentencePieceTrainerSpecSignals | None: + payload = _read_bounded_sentencepiece_submessage( + stream, + value_end, + max_bytes=_SENTENCEPIECE_MAX_TRAINER_SPEC_MESSAGE_BYTES, + ) + if payload is None: + return None + return _parse_sentencepiece_trainer_spec_proto(payload, 0, len(payload)) + + +def _classify_sentencepiece_model_proto_stream( + stream: BinaryIO, + file_size: int, + *, + max_decoded_piece_text_bytes: int = _SENTENCEPIECE_MODEL_PROTO_READ_BYTES, +) -> _SentencePieceModelProtoRoute: + offset = 0 + fields_seen = 0 + piece_count = 0 + typed_piece_count = 0 + special_identity_piece_count = 0 + unknown_piece_count = 0 + unknown_piece_index: int | None = None + unknown_piece_text: str | None = None + byte_piece_count = 0 + byte_piece_texts: set[str] = set() + malformed_byte_piece = False + trainer_spec: _SentencePieceTrainerSpecSignals | None = None + strong_match = False + decoded_piece_text_bytes = 0 + + def reject_candidate() -> _SentencePieceModelProtoRoute: + return "malformed_candidate" if piece_count else "unknown" + + while offset < file_size and fields_seen < _SENTENCEPIECE_MODEL_MAX_FIELDS: + tag = _read_proto_varint_stream(stream, file_size) + if tag is None: + return reject_candidate() + field_number = tag >> 3 + wire_type = tag & 0x07 + if field_number == 0: + return reject_candidate() + + if field_number == 1: + if wire_type != 2: + return reject_candidate() + bounds = _read_proto_length_delimited_bounds_stream(stream, file_size) + if bounds is None: + return reject_candidate() + _length, _value_start, actual_value_end = bounds + decode_piece_text = not _has_sufficient_sentencepiece_piece_scan_evidence( + piece_count=piece_count, + unknown_piece_count=unknown_piece_count, + unknown_piece_index=unknown_piece_index, + unknown_piece_text=unknown_piece_text, + byte_piece_count=byte_piece_count, + byte_piece_texts=byte_piece_texts, + malformed_byte_piece=malformed_byte_piece, + ) + parsed_piece = _parse_sentencepiece_piece_proto_stream( + stream, + actual_value_end, + decode_text=decode_piece_text, + max_decoded_text_bytes=max_decoded_piece_text_bytes - decoded_piece_text_bytes, + ) + if parsed_piece is None: + return reject_candidate() + decoded_piece_text_bytes += parsed_piece.decoded_text_bytes + piece = parsed_piece.piece_text + piece_type = parsed_piece.piece_type + piece_index = piece_count + piece_count += 1 + if piece_type is not None: + typed_piece_count += 1 + if piece is not None and _is_sentencepiece_special_identity_piece(piece): + special_identity_piece_count += 1 + if piece_type is not None: + if piece_type == _SENTENCEPIECE_UNKNOWN_PIECE_TYPE: + if piece is None: + return reject_candidate() + unknown_piece_count += 1 + unknown_piece_index = piece_index + unknown_piece_text = piece + elif piece_type == _SENTENCEPIECE_BYTE_PIECE_TYPE: + if piece is None: + return reject_candidate() + byte_piece_count += 1 + if _is_sentencepiece_byte_fallback_piece(piece): + byte_piece_texts.add(piece) + else: + malformed_byte_piece = True + stream.seek(actual_value_end) + offset = actual_value_end + elif field_number == 2: + if wire_type != 2: + return reject_candidate() + bounds = _read_proto_length_delimited_bounds_stream(stream, file_size) + if bounds is None: + return reject_candidate() + _length, _value_start, actual_value_end = bounds + parsed_trainer_spec = _parse_sentencepiece_trainer_spec_proto_stream(stream, actual_value_end) + if parsed_trainer_spec is None: + return reject_candidate() + if trainer_spec is None: + trainer_spec = parsed_trainer_spec + else: + trainer_spec.merge_from(parsed_trainer_spec) + stream.seek(actual_value_end) + offset = actual_value_end + elif field_number == 3: + if wire_type != 2: + return reject_candidate() + bounds = _read_proto_length_delimited_bounds_stream(stream, file_size) + if bounds is None: + return reject_candidate() + _length, _value_start, actual_value_end = bounds + if not _is_well_formed_sentencepiece_submessage_stream( + stream, + actual_value_end, + expected_wire_types=_SENTENCEPIECE_NORMALIZER_SPEC_WIRE_TYPES, + ): + return reject_candidate() + stream.seek(actual_value_end) + offset = actual_value_end + elif field_number in {4, 5}: + if wire_type != 2: + return reject_candidate() + bounds = _read_proto_length_delimited_bounds_stream(stream, file_size) + if bounds is None: + return reject_candidate() + _length, _value_start, actual_value_end = bounds + if not _is_well_formed_sentencepiece_submessage_stream(stream, actual_value_end): + return reject_candidate() + stream.seek(actual_value_end) + offset = actual_value_end + else: + return reject_candidate() + + fields_seen += 1 + strong_match = _has_strong_sentencepiece_model_proto_evidence( + piece_count=piece_count, + typed_piece_count=typed_piece_count, + special_identity_piece_count=special_identity_piece_count, + unknown_piece_count=unknown_piece_count, + unknown_piece_index=unknown_piece_index, + unknown_piece_text=unknown_piece_text, + byte_piece_count=byte_piece_count, + byte_piece_texts=byte_piece_texts, + malformed_byte_piece=malformed_byte_piece, + trainer_spec=trainer_spec, + ) + + if strong_match and stream.tell() == file_size and fields_seen < _SENTENCEPIECE_MODEL_MAX_FIELDS: + return "strong" + return "malformed_candidate" if piece_count else "unknown" + + +@lru_cache(maxsize=1024) +def _classify_sentencepiece_model_proto_file_cached( + path: str, + size: int, + mtime_ns: int, + ctime_ns: int, + fingerprint_head: bytes, + fingerprint_tail: bytes, +) -> _SentencePieceModelProtoRoute: + del mtime_ns, ctime_ns, fingerprint_head, fingerprint_tail + file_path = Path(path) + try: + with file_path.open("rb") as handle: + if size <= _SENTENCEPIECE_MODEL_PROTO_READ_BYTES: + payload = handle.read(size) + if len(payload) != size: + return "unknown" + return _classify_sentencepiece_model_proto_stream(BytesIO(payload), size) + return _classify_sentencepiece_model_proto_stream( + handle, + size, + max_decoded_piece_text_bytes=_SENTENCEPIECE_MODEL_PROTO_READ_BYTES, + ) + except OSError: + return "unknown" + + +def _sentencepiece_model_proto_cache_fingerprint(file_path: Path, size: int) -> tuple[bytes, bytes]: + try: + with file_path.open("rb") as handle: + head = handle.read(min(size, _SENTENCEPIECE_MODEL_PROTO_CACHE_FINGERPRINT_BYTES)) + if size <= _SENTENCEPIECE_MODEL_PROTO_CACHE_FINGERPRINT_BYTES: + return head, b"" + handle.seek(max(size - _SENTENCEPIECE_MODEL_PROTO_CACHE_FINGERPRINT_BYTES, 0)) + tail = handle.read(_SENTENCEPIECE_MODEL_PROTO_CACHE_FINGERPRINT_BYTES) + return head, tail + except OSError: + return b"", b"" + + +def _classify_sentencepiece_model_proto_file(path: str | Path) -> _SentencePieceModelProtoRoute: + file_path = Path(path) + try: + if not file_path.is_file(): + return "unknown" + stat = file_path.stat() + if stat.st_size < 32: + return "unknown" + fingerprint_head, fingerprint_tail = _sentencepiece_model_proto_cache_fingerprint(file_path, stat.st_size) + return _classify_sentencepiece_model_proto_file_cached( + str(file_path.resolve()), + stat.st_size, + stat.st_mtime_ns, + stat.st_ctime_ns, + fingerprint_head, + fingerprint_tail, + ) + except OSError: + return "unknown" + + +def _is_malformed_sentencepiece_model_proto_candidate_file(path: str | Path) -> bool: + return _classify_sentencepiece_model_proto_file(path) == "malformed_candidate" + + +def _should_fail_closed_malformed_sentencepiece_model_proto_file(path: str | Path) -> bool: + return Path(path).suffix.lower() in {"", ".proto"} and _is_malformed_sentencepiece_model_proto_candidate_file(path) + + +def _should_treat_sentencepiece_model_proto_file_as_unknown(path: str | Path) -> bool: + return Path(path).suffix.lower() in {".model", ".proto"} and is_sentencepiece_model_proto_file(path) + + +def is_sentencepiece_model_proto_file(path: str | Path) -> bool: + """Return True for strongly identified SentencePiece tokenizer ModelProto files.""" + return _classify_sentencepiece_model_proto_file(path) == "strong" + + def _skip_coreml_proto_group( data: bytes, offset: int, @@ -6138,6 +7004,12 @@ def detect_format_from_magic_bytes( if xgboost_route is not None: return xgboost_route + if file_path is not None and _should_fail_closed_malformed_sentencepiece_model_proto_file(file_path): + return SENTENCEPIECE_MODEL_PROTO_INCONCLUSIVE_FORMAT + + if file_path is not None and _should_treat_sentencepiece_model_proto_file_as_unknown(file_path): + return "unknown" + renamed_tensorflow_format = "unknown" if file_path is not None: renamed_tensorflow_format = _detect_renamed_tensorflow_protobuf( @@ -6294,6 +7166,9 @@ def detect_file_format_from_magic(path: str) -> str: if xgboost_route is not None: return xgboost_route + if _should_fail_closed_malformed_sentencepiece_model_proto_file(file_path): + return SENTENCEPIECE_MODEL_PROTO_INCONCLUSIVE_FORMAT + if ( _allows_renamed_binary_content_route(file_path) and _detect_executorch_content_route(file_path, magic8) == "executorch" @@ -6318,6 +7193,9 @@ def detect_file_format_from_magic(path: str) -> str: if _could_be_content_routed_flax_msgpack(file_path): return "flax_msgpack" + if _should_treat_sentencepiece_model_proto_file_as_unknown(file_path): + return "unknown" + renamed_tensorflow_format = _detect_renamed_tensorflow_protobuf(file_path, size) if renamed_tensorflow_format != "unknown": return renamed_tensorflow_format @@ -6476,6 +7354,12 @@ def detect_file_format_for_skip_filter(path: str) -> str: if xgboost_route is not None: return xgboost_route + if _should_fail_closed_malformed_sentencepiece_model_proto_file(file_path): + return SENTENCEPIECE_MODEL_PROTO_INCONCLUSIVE_FORMAT + + if _should_treat_sentencepiece_model_proto_file_as_unknown(file_path): + return "unknown" + if ( _allows_renamed_binary_content_route(file_path) and _detect_executorch_content_route(file_path, magic8) == "executorch" @@ -6686,6 +7570,12 @@ def detect_file_format(path: str) -> str: if xgboost_route is not None: return xgboost_route + if _should_fail_closed_malformed_sentencepiece_model_proto_file(file_path): + return SENTENCEPIECE_MODEL_PROTO_INCONCLUSIVE_FORMAT + + if _should_treat_sentencepiece_model_proto_file_as_unknown(file_path): + return "unknown" + renamed_tensorflow_format = _detect_renamed_tensorflow_protobuf( file_path, size, @@ -6740,6 +7630,8 @@ def detect_file_format(path: str) -> str: ) if xgboost_route is not None: return xgboost_route + if _should_fail_closed_malformed_sentencepiece_model_proto_file(file_path): + return SENTENCEPIECE_MODEL_PROTO_INCONCLUSIVE_FORMAT # For .bin files, do more sophisticated detection if ext == ".bin": magic64 = read_magic_bytes(path, 64) diff --git a/tests/assets/samples/sentencepiece/custom_unknown_disabled_specials.model b/tests/assets/samples/sentencepiece/custom_unknown_disabled_specials.model new file mode 100644 index 000000000..bfd54cf36 Binary files /dev/null and b/tests/assets/samples/sentencepiece/custom_unknown_disabled_specials.model differ diff --git a/tests/assets/samples/sentencepiece/custom_unknown_disabled_specials_byte_fallback.model b/tests/assets/samples/sentencepiece/custom_unknown_disabled_specials_byte_fallback.model new file mode 100644 index 000000000..cde4c60cd Binary files /dev/null and b/tests/assets/samples/sentencepiece/custom_unknown_disabled_specials_byte_fallback.model differ diff --git a/tests/scanners/test_compressed_scanner.py b/tests/scanners/test_compressed_scanner.py index 811e7f492..aae44d384 100644 --- a/tests/scanners/test_compressed_scanner.py +++ b/tests/scanners/test_compressed_scanner.py @@ -129,6 +129,53 @@ def test_compressed_scanner_can_handle_requires_matching_signature(tmp_path: Pat assert CompressedScanner.can_handle(str(invalid_gzip_path)) is False +@pytest.mark.parametrize( + ("filename", "expected_suffix"), + [ + ("model.gz", ".bin"), + ("weights.xz", ".bin"), + ("tokenizer.gz", ""), + ("spiece.gz", ""), + ("tokenizer.model.gz", ".model"), + ], +) +def test_declared_compressed_inner_suffix_preserves_routing_intent( + filename: str, + expected_suffix: str, +) -> None: + assert CompressedScanner._derive_inner_suffix(filename) == expected_suffix + + +def test_bare_compressed_raw_binary_routes_inner_as_pytorch_binary(tmp_path: Path) -> None: + wrapper = tmp_path / "model.gz" + wrapper.write_bytes(gzip.compress(b"raw binary weights" + b"\0" * 128)) + + result = scan_model_directory_or_file(str(wrapper), cache_enabled=False) + metadata = result.file_metadata[str(wrapper)] + metadata_extra = metadata.model_extra or {} + + assert result.success is True + assert metadata_extra["scanner_dependency_ids"] == ["compressed", "pytorch_binary"] + assert metadata_extra["decompressed_bytes"] == 146 + + +def test_tokenizer_gzip_raw_binary_routes_inner_as_pytorch_binary_security_scan(tmp_path: Path) -> None: + wrapper = tmp_path / "tokenizer.gz" + wrapper.write_bytes(gzip.compress(b"\0" * 50 + b"CONFIDENTIAL_DATA" + b"\0" * 50)) + + result = scan_model_directory_or_file( + str(wrapper), + blacklist_patterns=["CONFIDENTIAL"], + cache_enabled=False, + ) + metadata = result.file_metadata[str(wrapper)] + metadata_extra = metadata.model_extra or {} + + assert determine_exit_code(result) == 1 + assert metadata_extra["scanner_dependency_ids"] == ["compressed", "pytorch_binary"] + assert any(issue.rule_code == "S1001" and "CONFIDENTIAL" in issue.message for issue in result.issues) + + def test_compressed_scanner_can_handle_header_routed_misnamed_wrapper(tmp_path: Path) -> None: disguised_gzip_path = tmp_path / "model.jpg" disguised_gzip_path.write_bytes(gzip.compress(pickle.dumps({"weights": [1, 2, 3]}))) diff --git a/tests/scanners/test_xgboost_scanner.py b/tests/scanners/test_xgboost_scanner.py index 2c02d3dfe..213f466b2 100644 --- a/tests/scanners/test_xgboost_scanner.py +++ b/tests/scanners/test_xgboost_scanner.py @@ -10,8 +10,10 @@ """ import copy +import gzip import io import json +import os import pickle import struct import subprocess as real_subprocess @@ -24,14 +26,24 @@ from unittest.mock import ANY, Mock, patch import pytest +from click.testing import CliRunner from modelaudit.cache import get_cache_manager, reset_cache_manager -from modelaudit.core import determine_exit_code, scan_file, scan_model_directory_or_file +from modelaudit.cli import cli +from modelaudit.core import determine_exit_code, scan_file, scan_model_directory_or_file, scan_model_streaming from modelaudit.models import ModelAuditResultModel from modelaudit.scanners.base import INCONCLUSIVE_SCAN_OUTCOME, CheckStatus, IssueSeverity, ScanResult from modelaudit.scanners.tar_scanner import TarScanner from modelaudit.scanners.xgboost_scanner import XGBOOST_JSON_ROUTING_CHUNK_BYTES, XGBoostScanner from modelaudit.scanners.zip_scanner import ZipScanner +from modelaudit.utils.file import detection as file_detection +from modelaudit.utils.file.detection import ( + detect_file_format, + detect_file_format_for_skip_filter, + detect_file_format_from_magic, +) +from modelaudit.utils.helpers.file_iterator import iterate_files_streaming +from tests.cli_output import parse_click_json_output class FakeBooster: @@ -53,6 +65,239 @@ def _headerless_legacy_binary_header() -> bytes: return struct.pack(" bytes: + encoded = bytearray() + while value >= 0x80: + encoded.append((value & 0x7F) | 0x80) + value >>= 7 + encoded.append(value) + return bytes(encoded) + + +def _proto_field(field_number: int, wire_type: int, payload: bytes) -> bytes: + return _proto_varint((field_number << 3) | wire_type) + payload + + +def _unknown_sentencepiece_field(payload: bytes) -> bytes: + return _proto_field(99, 2, _proto_varint(len(payload)) + payload) + + +def _small_zip_payload() -> bytes: + archive = io.BytesIO() + with zipfile.ZipFile(archive, "w") as zip_archive: + zip_archive.writestr("payload.py", "print('nested')") + return archive.getvalue() + + +def _sentencepiece_piece(piece: str, piece_type: int) -> bytes: + piece_payload = ( + _proto_field(1, 2, _proto_varint(len(piece.encode("utf-8"))) + piece.encode("utf-8")) + + _proto_field(2, 5, struct.pack(" bytes: + trainer_spec = _proto_field(3, 0, _proto_varint(model_type)) + _proto_field(4, 0, _proto_varint(vocab_size)) + if byte_fallback: + trainer_spec += _proto_field(35, 0, _proto_varint(1)) + if unk_id is not None: + trainer_spec += _proto_field(40, 0, _proto_varint(unk_id)) + if unk_piece is not None: + encoded_unk_piece = unk_piece.encode("utf-8") + trainer_spec += _proto_field(45, 2, _proto_varint(len(encoded_unk_piece)) + encoded_unk_piece) + return _proto_field(2, 2, _proto_varint(len(trainer_spec + extra_fields)) + trainer_spec + extra_fields) + + +def _sentencepiece_trainer_spec_field(payload: bytes) -> bytes: + return _proto_field(2, 2, _proto_varint(len(payload)) + payload) + + +def _sentencepiece_model_proto(*, include_trainer_spec: bool = False) -> bytes: + pieces = [ + ("", 2), + ("", 3), + ("", 3), + ("", 3), + ("the", 1), + ("of", 1), + ("and", 1), + ("to", 1), + ] + model = b"".join(_sentencepiece_piece(piece, piece_type) for piece, piece_type in pieces) + if include_trainer_spec: + model += _sentencepiece_trainer_spec(vocab_size=len(pieces)) + return model + + +def _sentencepiece_model_proto_with_default_unknown_metadata() -> bytes: + pieces = [ + ("", 2), + ("hello", 1), + ("world", 1), + ("token", 1), + ] + return b"".join( + _sentencepiece_piece(piece, piece_type) for piece, piece_type in pieces + ) + _sentencepiece_trainer_spec( + vocab_size=len(pieces), + unk_id=None, + unk_piece=None, + ) + + +def _sentencepiece_model_proto_with_split_trainer_spec_metadata() -> bytes: + pieces = [ + ("", 2), + ("hello", 1), + ("world", 1), + ("token", 1), + ] + trainer_spec_head = _proto_field(3, 0, _proto_varint(1)) + _proto_field(4, 0, _proto_varint(len(pieces))) + custom_unknown = b"" + trainer_spec_tail = _proto_field(45, 2, _proto_varint(len(custom_unknown)) + custom_unknown) + return ( + b"".join(_sentencepiece_piece(piece, piece_type) for piece, piece_type in pieces) + + _sentencepiece_trainer_spec_field(trainer_spec_head) + + _sentencepiece_trainer_spec_field(trainer_spec_tail) + ) + + +def _sentencepiece_model_proto_with_duplicate_unknown_pieces() -> bytes: + pieces = [ + ("", 2), + ("", 2), + ("", 3), + ("", 3), + ("", 3), + ("the", 1), + ("of", 1), + ("and", 1), + ] + return b"".join(_sentencepiece_piece(piece, piece_type) for piece, piece_type in pieces) + + +def _sentencepiece_model_proto_with_byte_pieces_without_fallback() -> bytes: + pieces = [ + ("", 2), + ("", 3), + ("", 3), + ("<0x00>", 6), + ("<0x01>", 6), + ("<0x02>", 6), + ("<0x03>", 6), + ("<0x04>", 6), + ] + return b"".join(_sentencepiece_piece(piece, piece_type) for piece, piece_type in pieces) + + +def _sentencepiece_model_proto_with_trainer_varint_field_52() -> bytes: + return _sentencepiece_model_proto() + _sentencepiece_trainer_spec( + vocab_size=8, + extra_fields=_proto_field(52, 0, _proto_varint(7)), + ) + + +def _sentencepiece_model_proto_with_trainer_string_field_54() -> bytes: + seed_sentencepieces_file = b"seed_sentencepieces.tsv" + return _sentencepiece_model_proto() + _sentencepiece_trainer_spec( + vocab_size=8, + extra_fields=_proto_field(54, 2, _proto_varint(len(seed_sentencepieces_file)) + seed_sentencepieces_file), + ) + + +def _large_real_sentencepiece_model_proto_shape() -> bytes: + """Mirror large uMT5-style spiece.model prefixes without committing a large fixture.""" + pieces = [ + ("", 3), + ("", 3), + ("", 3), + ("", 2), + ("[eod]", 4), + ("[web]", 4), + ("[wiki]", 4), + ("[translate]", 4), + ] + pieces.extend((f"<0x{byte:02X}>", 6) for byte in range(256)) + pieces.extend((f"t{index:05d}-{'x' * 470}", 1) for index in range(24000)) + return b"".join( + _sentencepiece_piece(piece, piece_type) for piece, piece_type in pieces + ) + _sentencepiece_trainer_spec( + vocab_size=len(pieces), + unk_id=3, + byte_fallback=True, + ) + + +_SENTENCEPIECE_FIXTURE_DIR = Path(__file__).resolve().parents[1] / "assets" / "samples" / "sentencepiece" +_SENTENCEPIECE_OFFICIAL_FIXTURES = ( + "custom_unknown_disabled_specials.model", + "custom_unknown_disabled_specials_byte_fallback.model", +) +_SENTENCEPIECE_SEMANTIC_MALFORMED_FACTORIES = ( + pytest.param(_sentencepiece_model_proto_with_duplicate_unknown_pieces, id="duplicate-unknown-pieces"), + pytest.param(_sentencepiece_model_proto_with_byte_pieces_without_fallback, id="byte-pieces-without-fallback"), +) +_SENTENCEPIECE_MALFORMED_TAILS = ( + pytest.param(b"\x12\x80", id="truncated-length-varint"), + pytest.param(b"\x80", id="truncated-tag-varint"), + pytest.param(_proto_field(2, 2, _proto_varint(4) + b"x"), id="truncated-length-payload"), + pytest.param(_proto_field(2, 0, _proto_varint(0)), id="wrong-wire-trainer-spec"), + pytest.param(_proto_field(3, 0, _proto_varint(0)), id="wrong-wire-normalizer-spec"), + pytest.param(_proto_field(4, 0, _proto_varint(0)), id="wrong-wire-self-test-data"), + pytest.param(_proto_field(5, 0, _proto_varint(0)), id="wrong-wire-denormalizer-spec"), + pytest.param(_proto_varint((2 << 3) | 6), id="invalid-wire-type-6"), + pytest.param(_proto_varint((2 << 3) | 7), id="invalid-wire-type-7"), + pytest.param(b"\x00not a protobuf tail", id="arbitrary-tail"), + pytest.param( + _unknown_sentencepiece_field(pickle.dumps({"payload": "pickle"}, protocol=4)), id="unknown-field-pickle" + ), + pytest.param(_unknown_sentencepiece_field(_small_zip_payload()), id="unknown-field-zip"), +) + + +def _write_official_sentencepiece_fixture(tmp_path: Path, fixture_name: str) -> Path: + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes((_SENTENCEPIECE_FIXTURE_DIR / fixture_name).read_bytes()) + return tokenizer_model + + +def _write_sentencepiece_archive(tmp_path: Path, archive_kind: str, member_name: str, payload: bytes) -> Path: + if archive_kind == "zip": + archive_file = tmp_path / "tokenizer-bundle.zip" + with zipfile.ZipFile(archive_file, "w") as archive: + archive.writestr(member_name, payload) + return archive_file + if archive_kind == "tar": + archive_file = tmp_path / "tokenizer-bundle.tar" + with tarfile.open(archive_file, "w") as archive: + info = tarfile.TarInfo(member_name) + info.size = len(payload) + archive.addfile(info, io.BytesIO(payload)) + return archive_file + if archive_kind == "tar.gz": + archive_file = tmp_path / "tokenizer-bundle.tar.gz" + with tarfile.open(archive_file, "w:gz") as archive: + info = tarfile.TarInfo(member_name) + info.size = len(payload) + archive.addfile(info, io.BytesIO(payload)) + return archive_file + if archive_kind == "gz": + archive_file = tmp_path / f"{Path(member_name).name}.gz" + archive_file.write_bytes(gzip.compress(payload)) + return archive_file + raise AssertionError(f"Unhandled archive kind: {archive_kind}") + + @pytest.fixture def xgboost_scanner() -> XGBoostScanner: """Create an XGBoost scanner instance.""" @@ -165,6 +410,18 @@ def _assert_inconclusive_metadata(result: ModelAuditResultModel, path: Path, rea assert reason in metadata.get("scan_outcome_reasons", []) +def _assert_no_xgboost_s1004(result: ModelAuditResultModel) -> None: + assert "xgboost" not in result.scanner_names + assert determine_exit_code(result) == 0 + assert not any(issue.rule_code == "S1004" for issue in result.issues) + + +def _assert_xgboost_s1004(result: ModelAuditResultModel) -> None: + assert "xgboost" in result.scanner_names + assert determine_exit_code(result) == 2 + assert any(issue.rule_code == "S1004" for issue in result.issues) + + def _ubjson_key(key: bytes) -> bytes: return b"U" + bytes([len(key)]) + key @@ -286,9 +543,9 @@ def _xgboost_ubjson_deep_before_counted_null_array_probe() -> bytes: class TestXGBoostScannerBasic: """Test basic XGBoost scanner functionality.""" - def test_can_handle_supported_extensions(self, temp_dir): + def test_can_handle_supported_extensions(self, temp_dir: Path) -> None: """Test that scanner handles supported XGBoost file extensions.""" - # .bst, .model, .ubj are accepted based on extension + # .bst, .ubj, and ambiguous .model files are accepted based on extension. for ext in [".bst", ".model", ".ubj"]: test_file = temp_dir / f"test{ext}" test_file.write_text("dummy content") @@ -299,6 +556,84 @@ def test_can_handle_supported_extensions(self, temp_dir): json_file.write_text(json.dumps({"version": [1, 5, 2], "learner": {"gradient_booster": {}}})) assert XGBoostScanner.can_handle(str(json_file)) + def test_can_handle_rejects_strong_sentencepiece_tokenizer_model(self, tmp_path: Path) -> None: + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes(_sentencepiece_model_proto()) + + assert not XGBoostScanner.can_handle(str(tokenizer_model)) + + @pytest.mark.parametrize("fixture_name", _SENTENCEPIECE_OFFICIAL_FIXTURES) + def test_can_handle_rejects_dependency_sentencepiece_with_disabled_specials( + self, tmp_path: Path, fixture_name: str + ) -> None: + tokenizer_model = _write_official_sentencepiece_fixture(tmp_path, fixture_name) + + assert not XGBoostScanner.can_handle(str(tokenizer_model)) + + def test_can_handle_rejects_sentencepiece_with_proto2_default_unknown_metadata(self, tmp_path: Path) -> None: + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes(_sentencepiece_model_proto_with_default_unknown_metadata()) + + assert not XGBoostScanner.can_handle(str(tokenizer_model)) + + def test_can_handle_rejects_strong_sentencepiece_with_well_formed_tail(self, tmp_path: Path) -> None: + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes(_sentencepiece_model_proto() + _proto_field(2, 2, _proto_varint(0))) + + assert not XGBoostScanner.can_handle(str(tokenizer_model)) + + def test_can_handle_rejects_sentencepiece_with_trainer_varint_field_52(self, tmp_path: Path) -> None: + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes(_sentencepiece_model_proto_with_trainer_varint_field_52()) + + assert not XGBoostScanner.can_handle(str(tokenizer_model)) + + def test_can_handle_rejects_sentencepiece_with_trainer_string_field_54(self, tmp_path: Path) -> None: + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes(_sentencepiece_model_proto_with_trainer_string_field_54()) + + assert not XGBoostScanner.can_handle(str(tokenizer_model)) + + def test_can_handle_keeps_malformed_sentencepiece_like_model_on_xgboost_route(self, tmp_path: Path) -> None: + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes(b"\x0a\x0e\x0a\x05\x15\x00" + (b"\0" * 64)) + + assert XGBoostScanner.can_handle(str(tokenizer_model)) + + @pytest.mark.parametrize("payload_factory", _SENTENCEPIECE_SEMANTIC_MALFORMED_FACTORIES) + def test_can_handle_keeps_semantically_malformed_sentencepiece_on_xgboost_route( + self, tmp_path: Path, payload_factory: Callable[[], bytes] + ) -> None: + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes(payload_factory()) + + assert XGBoostScanner.can_handle(str(tokenizer_model)) + + def test_can_handle_keeps_capped_sentencepiece_prefix_on_xgboost_route( + self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch + ) -> None: + prefix = _sentencepiece_model_proto() + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes(prefix + b"\x12\x80") + monkeypatch.setattr(file_detection, "_SENTENCEPIECE_MODEL_PROTO_READ_BYTES", len(prefix)) + + assert XGBoostScanner.can_handle(str(tokenizer_model)) + + @pytest.mark.parametrize("tail", _SENTENCEPIECE_MALFORMED_TAILS) + def test_can_handle_keeps_strong_sentencepiece_with_malformed_tail_on_xgboost_route( + self, tmp_path: Path, tail: bytes + ) -> None: + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes(_sentencepiece_model_proto() + tail) + + assert XGBoostScanner.can_handle(str(tokenizer_model)) + + def test_can_handle_keeps_xgboost_model_extension_controls(self, tmp_path: Path) -> None: + binary_model = tmp_path / "xgboost.model" + binary_model.write_bytes(b"binf" + (b"\0" * 60)) + + assert XGBoostScanner.can_handle(str(binary_model)) + def test_cannot_handle_unsupported_extensions(self, temp_dir): """Test that scanner rejects unsupported file extensions.""" unsupported_extensions = [".txt", ".pkl", ".h5", ".onnx"] @@ -1666,6 +2001,787 @@ def test_binary_read_failure_core_is_operational_not_security_finding(self, tmp_ _assert_inconclusive_metadata(result, binary_file, "xgboost_read_failed") assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) + def test_sentencepiece_tokenizer_model_core_is_not_xgboost_false_positive(self, tmp_path: Path) -> None: + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes(_sentencepiece_model_proto()) + + direct = scan_file(str(tokenizer_model), config={"cache_enabled": False}) + aggregate = scan_model_directory_or_file(str(tmp_path), cache_enabled=False) + + assert direct.success is True + assert direct.scanner_name == "unknown" + assert "xgboost" not in aggregate.scanner_names + assert determine_exit_code(aggregate) == 0 + assert not any(issue.rule_code == "S1004" for issue in aggregate.issues) + + @pytest.mark.parametrize("archive_kind", ["zip", "tar"]) + def test_sentencepiece_tokenizer_model_nested_archive_is_not_xgboost_false_positive( + self, tmp_path: Path, archive_kind: str + ) -> None: + member_name = "models/tokenizer.model" + archive_file = tmp_path / f"tokenizer-bundle.{archive_kind}" + payload = _sentencepiece_model_proto() + if archive_kind == "zip": + with zipfile.ZipFile(archive_file, "w") as archive: + archive.writestr(member_name, payload) + direct_archive = ZipScanner({"cache_enabled": False}).scan(str(archive_file)) + else: + with tarfile.open(archive_file, "w") as archive: + info = tarfile.TarInfo(member_name) + info.size = len(payload) + archive.addfile(info, io.BytesIO(payload)) + direct_archive = TarScanner({"cache_enabled": False}).scan(str(archive_file)) + + aggregate = scan_model_directory_or_file(str(archive_file), cache_enabled=False) + + assert direct_archive.success is True + assert not any(issue.rule_code == "S1004" for issue in direct_archive.issues) + _assert_no_xgboost_s1004(aggregate) + + def test_sentencepiece_tokenizer_model_cli_xgboost_selection_is_not_false_positive(self, tmp_path: Path) -> None: + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes(_sentencepiece_model_proto()) + + cli_result = CliRunner().invoke( + cli, + ["scan", "--no-cache", "--format", "json", "--scanners", "xgboost", str(tmp_path)], + env={"PROMPTFOO_DISABLE_TELEMETRY": "1"}, + ) + cli_payload = parse_click_json_output(cli_result.output) + + assert cli_result.exit_code == 0 + assert "xgboost" not in cli_payload.get("scanner_names", []) + assert not any(issue.get("rule_code") == "S1004" for issue in cli_payload.get("issues", [])) + + def test_extensionless_sentencepiece_tokenizer_core_is_not_xgboost_false_positive(self, tmp_path: Path) -> None: + tokenizer_model = tmp_path / "tokenizer" + tokenizer_model.write_bytes(_sentencepiece_model_proto()) + + direct = scan_file(str(tokenizer_model), config={"cache_enabled": False}) + aggregate = scan_model_directory_or_file(str(tmp_path), cache_enabled=False) + + assert detect_file_format(str(tokenizer_model)) == "unknown" + assert detect_file_format_from_magic(str(tokenizer_model)) == "unknown" + assert detect_file_format_for_skip_filter(str(tokenizer_model)) == "unknown" + assert direct.success is True + assert direct.scanner_name == "unknown" + _assert_no_xgboost_s1004(aggregate) + + def test_sentencepiece_tokenizer_proto_core_is_clean_unknown(self, tmp_path: Path) -> None: + tokenizer_proto = tmp_path / "tokenizer.proto" + payload = _sentencepiece_model_proto() + tokenizer_proto.write_bytes(payload) + + direct = scan_file(str(tokenizer_proto), config={"cache_enabled": False}) + aggregate = scan_model_directory_or_file(str(tmp_path), cache_enabled=False) + + assert ( + file_detection.detect_format_from_magic_bytes( + payload[:4], + payload[:8], + payload[:16], + len(payload), + tokenizer_proto, + ) + == "unknown" + ) + assert detect_file_format(str(tokenizer_proto)) == "unknown" + assert detect_file_format_from_magic(str(tokenizer_proto)) == "unknown" + assert detect_file_format_for_skip_filter(str(tokenizer_proto)) == "unknown" + assert direct.success is True + assert direct.scanner_name == "unknown" + _assert_no_xgboost_s1004(aggregate) + + @pytest.mark.parametrize("payload_factory", _SENTENCEPIECE_SEMANTIC_MALFORMED_FACTORIES) + def test_extensionless_malformed_sentencepiece_model_fails_closed_without_xgboost( + self, tmp_path: Path, payload_factory: Callable[[], bytes] + ) -> None: + tokenizer_model = tmp_path / "tokenizer" + tokenizer_model.write_bytes(payload_factory()) + expected_format = file_detection.SENTENCEPIECE_MODEL_PROTO_INCONCLUSIVE_FORMAT + expected_reason = "sentencepiece_model_proto_routing_incomplete" + + direct = scan_file(str(tokenizer_model), config={"cache_enabled": False}) + aggregate = scan_model_directory_or_file(str(tmp_path), cache_enabled=False) + streaming = scan_model_streaming( + file_generator=iterate_files_streaming(tmp_path), + scan_root=str(tmp_path), + delete_after_scan=False, + cache_enabled=False, + skip_file_types=True, + ) + cli_result = CliRunner().invoke( + cli, + ["scan", "--stream", "--no-cache", "--format", "json", str(tmp_path)], + env={"PROMPTFOO_DISABLE_TELEMETRY": "1"}, + ) + cli_payload = parse_click_json_output(cli_result.output) + + assert detect_file_format(str(tokenizer_model)) == expected_format + assert detect_file_format_from_magic(str(tokenizer_model)) == expected_format + assert detect_file_format_for_skip_filter(str(tokenizer_model)) == expected_format + assert direct.success is False + assert direct.scanner_name == "unknown" + assert expected_reason in direct.metadata["scan_outcome_reasons"] + assert not any(issue.rule_code == "S1004" for issue in direct.issues) + assert "xgboost" not in aggregate.scanner_names + assert "xgboost" not in streaming.scanner_names + _assert_inconclusive_metadata(aggregate, tokenizer_model, expected_reason) + _assert_inconclusive_metadata(streaming, tokenizer_model, expected_reason) + assert determine_exit_code(aggregate) == 2 + assert determine_exit_code(streaming) == 2 + assert cli_result.exit_code == 2 + assert "xgboost" not in cli_payload.get("scanner_names", []) + assert any( + expected_reason in issue.get("details", {}).get("scan_outcome_reason", "") + or "SentencePiece ModelProto routing was inconclusive" in issue.get("message", "") + for issue in cli_payload.get("issues", []) + ) + + @pytest.mark.parametrize( + "payload", + [ + pytest.param(_sentencepiece_model_proto() + b"\x12\x80", id="truncated-tail"), + pytest.param( + _sentencepiece_model_proto() + _proto_field(2, 0, _proto_varint(0)), + id="wrong-wire-trainer-spec", + ), + pytest.param(_sentencepiece_model_proto_with_duplicate_unknown_pieces(), id="duplicate-unknown-pieces"), + ], + ) + def test_malformed_sentencepiece_proto_fails_closed_without_xgboost(self, tmp_path: Path, payload: bytes) -> None: + tokenizer_proto = tmp_path / "tokenizer.proto" + tokenizer_proto.write_bytes(payload) + expected_format = file_detection.SENTENCEPIECE_MODEL_PROTO_INCONCLUSIVE_FORMAT + expected_reason = "sentencepiece_model_proto_routing_incomplete" + + direct = scan_file(str(tokenizer_proto), config={"cache_enabled": False}) + aggregate = scan_model_directory_or_file(str(tmp_path), cache_enabled=False) + cli_result = CliRunner().invoke( + cli, + ["scan", "--no-cache", "--format", "json", str(tmp_path)], + env={"PROMPTFOO_DISABLE_TELEMETRY": "1"}, + ) + cli_payload = parse_click_json_output(cli_result.output) + + assert ( + file_detection.detect_format_from_magic_bytes( + payload[:4], + payload[:8], + payload[:16], + len(payload), + tokenizer_proto, + ) + == expected_format + ) + assert detect_file_format(str(tokenizer_proto)) == expected_format + assert detect_file_format_from_magic(str(tokenizer_proto)) == expected_format + assert detect_file_format_for_skip_filter(str(tokenizer_proto)) == expected_format + assert direct.success is False + assert direct.scanner_name == "unknown" + assert expected_reason in direct.metadata["scan_outcome_reasons"] + assert not any(issue.rule_code == "S1004" for issue in direct.issues) + assert "xgboost" not in aggregate.scanner_names + assert not any(issue.rule_code == "S1004" for issue in aggregate.issues) + _assert_inconclusive_metadata(aggregate, tokenizer_proto, expected_reason) + assert determine_exit_code(aggregate) == 2 + assert cli_result.exit_code == 2 + assert "xgboost" not in cli_payload.get("scanner_names", []) + assert any( + expected_reason in issue.get("details", {}).get("scan_outcome_reason", "") + or "SentencePiece ModelProto routing was inconclusive" in issue.get("message", "") + for issue in cli_payload.get("issues", []) + ) + + def test_model_card_text_mentions_do_not_trigger_xgboost_sentencepiece_routing(self, tmp_path: Path) -> None: + model_card = tmp_path / "README.md" + model_card.write_text( + "This repository ships a SentencePiece tokenizer.model file. " + "The words learner and version are documentation, not an XGBoost model.", + encoding="utf-8", + ) + + result = scan_model_directory_or_file(str(tmp_path), cache_enabled=False) + + assert determine_exit_code(result) == 0 + assert "xgboost" not in result.scanner_names + assert not any(issue.rule_code == "S1004" for issue in result.issues) + + def test_sentencepiece_payload_with_xgboost_only_suffix_still_fails_closed(self, tmp_path: Path) -> None: + deceptive_model = tmp_path / "tokenizer.bst" + deceptive_model.write_bytes(_sentencepiece_model_proto()) + + result = scan_file(str(deceptive_model), config={"cache_enabled": False}) + + assert result.success is False + assert result.scanner_name == "xgboost" + assert "xgboost_binary_structure_unrecognized" in result.metadata["scan_outcome_reasons"] + assert any(issue.rule_code == "S1004" for issue in result.issues) + + def test_sentencepiece_ownership_rechecks_same_path_replacement(self, tmp_path: Path) -> None: + tokenizer_model = tmp_path / "tokenizer.model" + valid_tokenizer = _sentencepiece_model_proto() + tokenizer_model.write_bytes(valid_tokenizer) + original_stat = tokenizer_model.stat() + + valid_result = scan_file(str(tokenizer_model), config={"cache_enabled": False}) + + assert valid_result.success is True + assert valid_result.scanner_name == "unknown" + + replacement = b"custom xgboost binary gbtree reg:squarederror" + tokenizer_model.write_bytes(replacement.ljust(len(valid_tokenizer), b"\0")) + os.utime(tokenizer_model, ns=(original_stat.st_atime_ns, original_stat.st_mtime_ns)) + + replaced_result = scan_file(str(tokenizer_model), config={"cache_enabled": False}) + + assert replaced_result.success is False + assert replaced_result.scanner_name == "xgboost" + assert "xgboost_binary_structure_unrecognized" in replaced_result.metadata["scan_outcome_reasons"] + assert any(issue.rule_code == "S1004" for issue in replaced_result.issues) + + def test_sentencepiece_model_with_proto2_default_unknown_metadata_is_not_xgboost_false_positive( + self, tmp_path: Path + ) -> None: + # Synthetic regression for HuggingFaceH4/zephyr-7b-beta@892b3d7... tokenizer.model. + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes(_sentencepiece_model_proto_with_default_unknown_metadata()) + + direct = scan_file(str(tokenizer_model), config={"cache_enabled": False}) + aggregate = scan_model_directory_or_file(str(tmp_path), cache_enabled=False) + streaming = scan_model_streaming( + file_generator=iterate_files_streaming(tmp_path), + scan_root=str(tmp_path), + delete_after_scan=False, + cache_enabled=False, + skip_file_types=True, + ) + cli_result = CliRunner().invoke( + cli, + ["scan", "--stream", "--no-cache", "--format", "json", str(tmp_path)], + env={"PROMPTFOO_DISABLE_TELEMETRY": "1"}, + ) + cli_payload = parse_click_json_output(cli_result.output) + + assert direct.success is True + assert direct.scanner_name == "unknown" + _assert_no_xgboost_s1004(aggregate) + _assert_no_xgboost_s1004(streaming) + assert cli_result.exit_code == 0 + assert "xgboost" not in cli_payload.get("scanner_names", []) + assert not any(issue.get("rule_code") == "S1004" for issue in cli_payload.get("issues", [])) + + def test_sentencepiece_model_with_repeated_trainer_spec_merge_is_not_xgboost_false_positive( + self, tmp_path: Path + ) -> None: + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes(_sentencepiece_model_proto_with_split_trainer_spec_metadata()) + + direct = scan_file(str(tokenizer_model), config={"cache_enabled": False}) + aggregate = scan_model_directory_or_file(str(tmp_path), cache_enabled=False) + streaming = scan_model_streaming( + file_generator=iterate_files_streaming(tmp_path), + scan_root=str(tmp_path), + delete_after_scan=False, + cache_enabled=False, + skip_file_types=True, + ) + cli_result = CliRunner().invoke( + cli, + ["scan", "--stream", "--no-cache", "--format", "json", str(tmp_path)], + env={"PROMPTFOO_DISABLE_TELEMETRY": "1"}, + ) + cli_payload = parse_click_json_output(cli_result.output) + + assert direct.success is True + assert direct.scanner_name == "unknown" + _assert_no_xgboost_s1004(aggregate) + _assert_no_xgboost_s1004(streaming) + assert cli_result.exit_code == 0 + assert "xgboost" not in cli_payload.get("scanner_names", []) + assert not any(issue.get("rule_code") == "S1004" for issue in cli_payload.get("issues", [])) + + @pytest.mark.parametrize("archive_kind", ["zip", "tar", "gz"]) + def test_archived_sentencepiece_repeated_trainer_spec_merge_is_not_xgboost_false_positive( + self, tmp_path: Path, archive_kind: str + ) -> None: + member_name = "models/tokenizer.model" + archive_file = _write_sentencepiece_archive( + tmp_path, + archive_kind, + member_name, + _sentencepiece_model_proto_with_split_trainer_spec_metadata(), + ) + + result = scan_model_directory_or_file(str(archive_file), cache_enabled=False) + cli_result = CliRunner().invoke( + cli, + ["scan", "--no-cache", "--format", "json", str(archive_file)], + env={"PROMPTFOO_DISABLE_TELEMETRY": "1"}, + ) + cli_payload = parse_click_json_output(cli_result.output) + + assert determine_exit_code(result) == 0 + assert "xgboost" not in result.scanner_names + assert not any(issue.rule_code == "S1004" for issue in result.issues) + assert cli_result.exit_code == 0 + assert "xgboost" not in cli_payload.get("scanner_names", []) + assert not any(issue.get("rule_code") == "S1004" for issue in cli_payload.get("issues", [])) + + def test_sentencepiece_model_with_trainer_varint_field_52_is_not_xgboost_false_positive( + self, tmp_path: Path + ) -> None: + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes(_sentencepiece_model_proto_with_trainer_varint_field_52()) + + direct = scan_file(str(tokenizer_model), config={"cache_enabled": False}) + aggregate = scan_model_directory_or_file(str(tmp_path), cache_enabled=False) + streaming = scan_model_streaming( + file_generator=iterate_files_streaming(tmp_path), + scan_root=str(tmp_path), + delete_after_scan=False, + cache_enabled=False, + skip_file_types=True, + ) + cli_result = CliRunner().invoke( + cli, + ["scan", "--stream", "--no-cache", "--format", "json", str(tmp_path)], + env={"PROMPTFOO_DISABLE_TELEMETRY": "1"}, + ) + cli_payload = parse_click_json_output(cli_result.output) + + assert direct.success is True + assert direct.scanner_name == "unknown" + _assert_no_xgboost_s1004(aggregate) + _assert_no_xgboost_s1004(streaming) + assert cli_result.exit_code == 0 + assert "xgboost" not in cli_payload.get("scanner_names", []) + assert not any(issue.get("rule_code") == "S1004" for issue in cli_payload.get("issues", [])) + + def test_sentencepiece_model_with_trainer_string_field_54_is_not_xgboost_false_positive( + self, tmp_path: Path + ) -> None: + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes(_sentencepiece_model_proto_with_trainer_string_field_54()) + + direct = scan_file(str(tokenizer_model), config={"cache_enabled": False}) + aggregate = scan_model_directory_or_file(str(tmp_path), cache_enabled=False) + streaming = scan_model_streaming( + file_generator=iterate_files_streaming(tmp_path), + scan_root=str(tmp_path), + delete_after_scan=False, + cache_enabled=False, + skip_file_types=True, + ) + cli_result = CliRunner().invoke( + cli, + ["scan", "--stream", "--no-cache", "--format", "json", str(tmp_path)], + env={"PROMPTFOO_DISABLE_TELEMETRY": "1"}, + ) + cli_payload = parse_click_json_output(cli_result.output) + + assert direct.success is True + assert direct.scanner_name == "unknown" + _assert_no_xgboost_s1004(aggregate) + _assert_no_xgboost_s1004(streaming) + assert cli_result.exit_code == 0 + assert "xgboost" not in cli_payload.get("scanner_names", []) + assert not any(issue.get("rule_code") == "S1004" for issue in cli_payload.get("issues", [])) + + def test_large_real_sentencepiece_model_shape_is_not_tensorflow_or_xgboost_false_positive( + self, tmp_path: Path + ) -> None: + # Mirrors baidu/NAVA@16c20287... Wan2.2-TI2V-5B/google/umt5-xxl/spiece.model: + # a large SentencePiece ModelProto whose full payload exceeds the ownership read cap. + tokenizer_model = tmp_path / "spiece.model" + tokenizer_model.write_bytes(_large_real_sentencepiece_model_proto_shape()) + + direct = scan_file(str(tokenizer_model), config={"cache_enabled": False}) + aggregate = scan_model_directory_or_file(str(tmp_path), cache_enabled=False) + streaming = scan_model_streaming( + file_generator=iterate_files_streaming(tmp_path), + scan_root=str(tmp_path), + delete_after_scan=False, + cache_enabled=False, + skip_file_types=True, + ) + cli_result = CliRunner().invoke( + cli, + ["scan", "--stream", "--no-cache", "--format", "json", str(tmp_path)], + env={"PROMPTFOO_DISABLE_TELEMETRY": "1"}, + ) + cli_payload = parse_click_json_output(cli_result.output) + + assert tokenizer_model.stat().st_size > file_detection._SENTENCEPIECE_MODEL_PROTO_READ_BYTES + assert detect_file_format(str(tokenizer_model)) == "unknown" + assert detect_file_format_from_magic(str(tokenizer_model)) == "unknown" + assert detect_file_format_for_skip_filter(str(tokenizer_model)) == "unknown" + assert direct.success is True + assert direct.scanner_name == "unknown" + _assert_no_xgboost_s1004(aggregate) + _assert_no_xgboost_s1004(streaming) + assert cli_result.exit_code == 0 + assert "xgboost" not in cli_payload.get("scanner_names", []) + assert not any(issue.get("rule_code") == "S1004" for issue in cli_payload.get("issues", [])) + + def test_large_sentencepiece_ownership_probe_skips_unsampled_piece_text(self) -> None: + class _CountingBytesIO(io.BytesIO): + def __init__(self, payload: bytes) -> None: + super().__init__(payload) + self.bytes_read = 0 + + def read(self, size: int | None = -1) -> bytes: + payload = super().read(size) + self.bytes_read += len(payload) + return payload + + payload = _large_real_sentencepiece_model_proto_shape() + stream = _CountingBytesIO(payload) + + route = file_detection._classify_sentencepiece_model_proto_stream(stream, len(payload)) + + assert len(payload) > file_detection._SENTENCEPIECE_MODEL_PROTO_READ_BYTES + assert route == "strong" + assert stream.bytes_read < file_detection._SENTENCEPIECE_MODEL_PROTO_READ_BYTES + + def test_sentencepiece_ownership_probe_reuses_direct_scan_route( + self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch + ) -> None: + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes(_large_real_sentencepiece_model_proto_shape()) + original_classifier = file_detection._classify_sentencepiece_model_proto_stream + calls = 0 + + def counting_classifier( + stream: Any, + file_size: int, + *, + max_decoded_piece_text_bytes: int = file_detection._SENTENCEPIECE_MODEL_PROTO_READ_BYTES, + ) -> Any: + nonlocal calls + calls += 1 + return original_classifier( + stream, + file_size, + max_decoded_piece_text_bytes=max_decoded_piece_text_bytes, + ) + + file_detection._classify_sentencepiece_model_proto_file_cached.cache_clear() + monkeypatch.setattr(file_detection, "_classify_sentencepiece_model_proto_stream", counting_classifier) + + direct = scan_file(str(tokenizer_model), config={"cache_enabled": False}) + + assert direct.success is True + assert direct.scanner_name == "unknown" + assert calls == 1 + + @pytest.mark.parametrize("fixture_name", _SENTENCEPIECE_OFFICIAL_FIXTURES) + def test_dependency_sentencepiece_model_with_disabled_specials_is_not_xgboost_false_positive( + self, tmp_path: Path, fixture_name: str + ) -> None: + tokenizer_model = _write_official_sentencepiece_fixture(tmp_path, fixture_name) + + direct = scan_file(str(tokenizer_model), config={"cache_enabled": False}) + aggregate = scan_model_directory_or_file(str(tmp_path), cache_enabled=False) + streaming = scan_model_streaming( + file_generator=iterate_files_streaming(tmp_path), + scan_root=str(tmp_path), + delete_after_scan=False, + cache_enabled=False, + skip_file_types=True, + ) + cli_result = CliRunner().invoke( + cli, + ["scan", "--stream", "--no-cache", "--format", "json", str(tmp_path)], + env={"PROMPTFOO_DISABLE_TELEMETRY": "1"}, + ) + cli_payload = parse_click_json_output(cli_result.output) + + assert direct.success is True + assert direct.scanner_name == "unknown" + _assert_no_xgboost_s1004(aggregate) + _assert_no_xgboost_s1004(streaming) + assert cli_result.exit_code == 0 + assert "xgboost" not in cli_payload.get("scanner_names", []) + assert not any(issue.get("rule_code") == "S1004" for issue in cli_payload.get("issues", [])) + + @pytest.mark.parametrize("archive_kind", ["zip", "tar", "tar.gz", "gz"]) + def test_archived_sentencepiece_model_is_not_xgboost_false_positive( + self, tmp_path: Path, archive_kind: str + ) -> None: + archive_file = _write_sentencepiece_archive( + tmp_path, + archive_kind, + "models/tokenizer.model", + _sentencepiece_model_proto(), + ) + + result = scan_model_directory_or_file(str(archive_file), cache_enabled=False) + cli_result = CliRunner().invoke( + cli, + ["scan", "--no-cache", "--format", "json", str(archive_file)], + env={"PROMPTFOO_DISABLE_TELEMETRY": "1"}, + ) + cli_payload = parse_click_json_output(cli_result.output) + + assert determine_exit_code(result) == 0 + assert "xgboost" not in result.scanner_names + assert not any(issue.rule_code == "S1004" for issue in result.issues) + assert cli_result.exit_code == 0 + assert "xgboost" not in cli_payload.get("scanner_names", []) + assert not any(issue.get("rule_code") == "S1004" for issue in cli_payload.get("issues", [])) + + @pytest.mark.parametrize("archive_kind", ["zip", "tar", "tar.gz", "gz"]) + @pytest.mark.parametrize("payload_factory", _SENTENCEPIECE_SEMANTIC_MALFORMED_FACTORIES) + def test_extensionless_archived_malformed_sentencepiece_model_fails_closed_without_xgboost( + self, tmp_path: Path, archive_kind: str, payload_factory: Callable[[], bytes] + ) -> None: + member_name = "models/tokenizer" + archive_file = _write_sentencepiece_archive(tmp_path, archive_kind, member_name, payload_factory()) + expected_location = f"{archive_file} -> tokenizer" if archive_kind == "gz" else f"{archive_file}:{member_name}" + expected_reason = "sentencepiece_model_proto_routing_incomplete" + + result = scan_model_directory_or_file(str(archive_file), cache_enabled=False) + cli_result = CliRunner().invoke( + cli, + ["scan", "--no-cache", "--format", "json", str(archive_file)], + env={"PROMPTFOO_DISABLE_TELEMETRY": "1"}, + ) + cli_payload = parse_click_json_output(cli_result.output) + metadata = result.file_metadata[str(archive_file)] + + assert determine_exit_code(result) == 2 + assert "xgboost" not in result.scanner_names + assert expected_reason in metadata.get("scan_outcome_reasons", []) + assert any( + issue.location == expected_location + and "SentencePiece ModelProto routing was inconclusive" in str(issue.message) + for issue in result.issues + ) + assert cli_result.exit_code == 2 + assert "xgboost" not in cli_payload.get("scanner_names", []) + assert any( + issue.get("location") == expected_location + and "SentencePiece ModelProto routing was inconclusive" in issue.get("message", "") + for issue in cli_payload.get("issues", []) + ) + + def test_malformed_sentencepiece_like_model_still_fails_closed_in_xgboost(self, tmp_path: Path) -> None: + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes(b"\x0a\x0e\x0a\x05\x15\x00" + (b"\0" * 64)) + + result = scan_file(str(tokenizer_model), config={"cache_enabled": False}) + + assert result.success is False + assert result.scanner_name == "xgboost" + assert "xgboost_binary_structure_unrecognized" in result.metadata["scan_outcome_reasons"] + assert any(issue.rule_code == "S1004" for issue in result.issues) + + @pytest.mark.parametrize("payload_factory", _SENTENCEPIECE_SEMANTIC_MALFORMED_FACTORIES) + def test_semantically_malformed_sentencepiece_model_still_fails_closed_in_xgboost( + self, tmp_path: Path, payload_factory: Callable[[], bytes] + ) -> None: + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes(payload_factory()) + + direct = scan_file(str(tokenizer_model), config={"cache_enabled": False}) + aggregate = scan_model_directory_or_file(str(tmp_path), cache_enabled=False) + streaming = scan_model_streaming( + file_generator=iterate_files_streaming(tmp_path), + scan_root=str(tmp_path), + delete_after_scan=False, + cache_enabled=False, + skip_file_types=True, + ) + cli_result = CliRunner().invoke( + cli, + ["scan", "--stream", "--no-cache", "--format", "json", str(tmp_path)], + env={"PROMPTFOO_DISABLE_TELEMETRY": "1"}, + ) + cli_payload = parse_click_json_output(cli_result.output) + + assert direct.success is False + assert direct.scanner_name == "xgboost" + assert "xgboost_binary_structure_unrecognized" in direct.metadata["scan_outcome_reasons"] + assert any(issue.rule_code == "S1004" for issue in direct.issues) + _assert_xgboost_s1004(aggregate) + _assert_xgboost_s1004(streaming) + assert cli_result.exit_code == 2 + assert "xgboost" in cli_payload.get("scanner_names", []) + assert any(issue.get("rule_code") == "S1004" for issue in cli_payload.get("issues", [])) + + @pytest.mark.parametrize("archive_kind", ["zip", "tar", "tar.gz", "gz"]) + @pytest.mark.parametrize("payload_factory", _SENTENCEPIECE_SEMANTIC_MALFORMED_FACTORIES) + def test_archived_semantically_malformed_sentencepiece_model_still_fails_closed_in_xgboost( + self, tmp_path: Path, archive_kind: str, payload_factory: Callable[[], bytes] + ) -> None: + member_name = "models/tokenizer.model" + archive_file = _write_sentencepiece_archive(tmp_path, archive_kind, member_name, payload_factory()) + + result = scan_model_directory_or_file(str(archive_file), cache_enabled=False) + cli_result = CliRunner().invoke( + cli, + ["scan", "--no-cache", "--format", "json", str(archive_file)], + env={"PROMPTFOO_DISABLE_TELEMETRY": "1"}, + ) + cli_payload = parse_click_json_output(cli_result.output) + expected_location = ( + f"{archive_file} -> tokenizer.model" if archive_kind == "gz" else f"{archive_file}:{member_name}" + ) + + assert determine_exit_code(result) == 2 + assert any(issue.rule_code == "S1004" and issue.location == expected_location for issue in result.issues) + assert cli_result.exit_code == 2 + assert any( + issue.get("rule_code") == "S1004" and issue.get("location") == expected_location + for issue in cli_payload.get("issues", []) + ) + + @pytest.mark.parametrize("archive_kind", ["zip", "tar", "tar.gz", "gz"]) + @pytest.mark.parametrize( + "tail", + [ + pytest.param(_unknown_sentencepiece_field(pickle.dumps({"payload": "pickle"}, protocol=4)), id="pickle"), + pytest.param(_unknown_sentencepiece_field(_small_zip_payload()), id="zip"), + ], + ) + def test_archived_sentencepiece_unknown_field_payload_still_fails_closed_in_xgboost( + self, tmp_path: Path, archive_kind: str, tail: bytes + ) -> None: + member_name = "models/tokenizer.model" + archive_file = _write_sentencepiece_archive( + tmp_path, + archive_kind, + member_name, + _sentencepiece_model_proto() + tail, + ) + expected_location = ( + f"{archive_file} -> tokenizer.model" if archive_kind == "gz" else f"{archive_file}:{member_name}" + ) + + result = scan_model_directory_or_file(str(archive_file), cache_enabled=False) + cli_result = CliRunner().invoke( + cli, + ["scan", "--no-cache", "--format", "json", str(archive_file)], + env={"PROMPTFOO_DISABLE_TELEMETRY": "1"}, + ) + cli_payload = parse_click_json_output(cli_result.output) + + assert determine_exit_code(result) == 2 + assert any(issue.rule_code == "S1004" and issue.location == expected_location for issue in result.issues) + assert cli_result.exit_code == 2 + assert any( + issue.get("rule_code") == "S1004" and issue.get("location") == expected_location + for issue in cli_payload.get("issues", []) + ) + + @pytest.mark.parametrize("tail", _SENTENCEPIECE_MALFORMED_TAILS) + def test_sentencepiece_model_with_malformed_tail_still_fails_closed_in_xgboost( + self, tmp_path: Path, tail: bytes + ) -> None: + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes(_sentencepiece_model_proto() + tail) + + direct = scan_file(str(tokenizer_model), config={"cache_enabled": False}) + aggregate = scan_model_directory_or_file(str(tmp_path), cache_enabled=False) + streaming = scan_model_streaming( + file_generator=iterate_files_streaming(tmp_path), + scan_root=str(tmp_path), + delete_after_scan=False, + cache_enabled=False, + skip_file_types=True, + ) + cli_result = CliRunner().invoke( + cli, + ["scan", "--stream", "--no-cache", "--format", "json", str(tmp_path)], + env={"PROMPTFOO_DISABLE_TELEMETRY": "1"}, + ) + cli_payload = parse_click_json_output(cli_result.output) + + assert direct.success is False + assert direct.scanner_name == "xgboost" + assert "xgboost_binary_structure_unrecognized" in direct.metadata["scan_outcome_reasons"] + assert any(issue.rule_code == "S1004" for issue in direct.issues) + _assert_xgboost_s1004(aggregate) + _assert_xgboost_s1004(streaming) + assert cli_result.exit_code == 2 + assert "xgboost" in cli_payload.get("scanner_names", []) + assert any(issue.get("rule_code") == "S1004" for issue in cli_payload.get("issues", [])) + + def test_capped_sentencepiece_prefix_with_unread_tail_still_fails_closed_in_xgboost( + self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch + ) -> None: + prefix = _sentencepiece_model_proto() + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes(prefix + b"\x12\x80") + monkeypatch.setattr(file_detection, "_SENTENCEPIECE_MODEL_PROTO_READ_BYTES", len(prefix)) + + direct = scan_file(str(tokenizer_model), config={"cache_enabled": False}) + aggregate = scan_model_directory_or_file(str(tmp_path), cache_enabled=False) + streaming = scan_model_streaming( + file_generator=iterate_files_streaming(tmp_path), + scan_root=str(tmp_path), + delete_after_scan=False, + cache_enabled=False, + skip_file_types=True, + ) + cli_result = CliRunner().invoke( + cli, + ["scan", "--stream", "--no-cache", "--format", "json", str(tmp_path)], + env={"PROMPTFOO_DISABLE_TELEMETRY": "1"}, + ) + cli_payload = parse_click_json_output(cli_result.output) + + assert direct.success is False + assert direct.scanner_name == "xgboost" + assert "xgboost_binary_structure_unrecognized" in direct.metadata["scan_outcome_reasons"] + assert any(issue.rule_code == "S1004" for issue in direct.issues) + _assert_xgboost_s1004(aggregate) + _assert_xgboost_s1004(streaming) + assert cli_result.exit_code == 2 + assert "xgboost" in cli_payload.get("scanner_names", []) + assert any(issue.get("rule_code") == "S1004" for issue in cli_payload.get("issues", [])) + + def test_large_sentencepiece_wrong_wire_type_metadata_still_fails_closed_in_xgboost(self, tmp_path: Path) -> None: + tokenizer_model = tmp_path / "tokenizer.model" + tokenizer_model.write_bytes( + _large_real_sentencepiece_model_proto_shape() + _proto_field(2, 0, _proto_varint(0)) + ) + + direct = scan_file(str(tokenizer_model), config={"cache_enabled": False}) + aggregate = scan_model_directory_or_file(str(tmp_path), cache_enabled=False) + + assert tokenizer_model.stat().st_size > file_detection._SENTENCEPIECE_MODEL_PROTO_READ_BYTES + assert direct.success is False + assert direct.scanner_name == "xgboost" + assert "xgboost_binary_structure_unrecognized" in direct.metadata["scan_outcome_reasons"] + assert any(issue.rule_code == "S1004" for issue in direct.issues) + _assert_xgboost_s1004(aggregate) + + def test_xgboost_binary_model_extension_still_scans_cleanly(self, tmp_path: Path) -> None: + binary_model = tmp_path / "native.model" + binary_model.write_bytes(b"binf" + (b"\0" * 60)) + + result = scan_file(str(binary_model), config={"cache_enabled": False}) + + assert result.success is True + assert result.scanner_name == "xgboost" + assert not result.issues + + def test_xgboost_shaped_model_extension_still_fails_closed( + self, tmp_path: Path, valid_xgboost_json: dict[str, Any] + ) -> None: + valid_xgboost_json["learner"]["malicious_code"] = "os.system('touch pwned')" + model_file = tmp_path / "tokenizer.model" + model_file.write_text(json.dumps(valid_xgboost_json), encoding="utf-8") + + result = scan_file(str(model_file), config={"cache_enabled": False}) + + assert result.success is False + assert result.scanner_name == "xgboost" + assert "xgboost_binary_structure_unrecognized" in result.metadata["scan_outcome_reasons"] + assert any(issue.rule_code == "S1004" for issue in result.issues) + def test_extensionless_ubjson_nested_zip_detects_malicious_content( self, tmp_path: Path, valid_xgboost_json: dict[str, Any] ) -> None: diff --git a/tests/test_committed_fixture_hygiene.py b/tests/test_committed_fixture_hygiene.py index 11b53c059..7d5038af4 100644 --- a/tests/test_committed_fixture_hygiene.py +++ b/tests/test_committed_fixture_hygiene.py @@ -17,6 +17,7 @@ SAFETENSORS_DIR = ASSETS_DIR / "samples" / "safetensors" KERAS_DIR = ASSETS_DIR / "samples" / "keras" JINJA2_DIR = ASSETS_DIR / "samples" / "jinja2" +SENTENCEPIECE_DIR = ASSETS_DIR / "samples" / "sentencepiece" EXPECTED_SAFETENSORS_FIXTURES = { "tests/assets/samples/safetensors/malicious_import.safetensors", @@ -68,6 +69,11 @@ "tests/assets/samples/jinja2/yaml/model_config.yaml", } +EXPECTED_SENTENCEPIECE_FIXTURES = { + "tests/assets/samples/sentencepiece/custom_unknown_disabled_specials.model", + "tests/assets/samples/sentencepiece/custom_unknown_disabled_specials_byte_fallback.model", +} + LARGE_ASSET_BYTES = 100 * 1024 LARGE_ASSET_ALLOWLIST = { "tests/assets/samples/pickles/safe_large_model.pkl": "safe pickle regression corpus large-file fixture", @@ -136,6 +142,12 @@ def test_jinja2_corpus_matches_routed_inventory() -> None: assert jinja2_files == EXPECTED_JINJA2_FIXTURES +def test_sentencepiece_corpus_matches_routed_inventory() -> None: + sentencepiece_files = {_repo_relative(path) for path in _tracked_under(SENTENCEPIECE_DIR)} + + assert sentencepiece_files == EXPECTED_SENTENCEPIECE_FIXTURES + + def test_pickle_bypass_poc_generators_are_not_committed() -> None: bypass_poc_files = sorted(_repo_relative(path) for path in _tracked_asset_paths() if _is_bypass_poc_generator(path))