diff --git a/.github/ISSUE_TEMPLATE/bug_report.md b/.github/ISSUE_TEMPLATE/bug_report.md new file mode 100644 index 0000000..95e7494 --- /dev/null +++ b/.github/ISSUE_TEMPLATE/bug_report.md @@ -0,0 +1,32 @@ +--- +name: Bug report +about: Create a report to help us improve +title: '' +labels: bug +assignees: pomponchik + +--- + +## Short description + +Replace this text with a short description of the error and the behavior that you expected to see instead. + + +## Describe the bug in detail + +Please add a test that reproduces the bug (i.e., currently fails): + +```python +def test_your_bug(): + ... +``` + +When writing the test, please ensure compatibility with the [`pytest`](https://docs.pytest.org/) framework. + +If for some reason you cannot describe the error in the test format, describe the steps to reproduce it here. + + +## Environment + - OS: ... + - Python version (the output of the `python --version` command): ... + - Version of this package: ... diff --git a/.github/ISSUE_TEMPLATE/documentation.md b/.github/ISSUE_TEMPLATE/documentation.md new file mode 100644 index 0000000..5f5fdc0 --- /dev/null +++ b/.github/ISSUE_TEMPLATE/documentation.md @@ -0,0 +1,26 @@ +--- +name: Documentation fix +about: Add something to the documentation, delete it, or change it +title: '' +labels: documentation +assignees: pomponchik +--- + +## It's cool that you're here! + +Documentation is an important part of the project; we strive to make it high-quality and keep it up to date. Please adjust this template by outlining your proposal. + + +## Type of action + +What do you want to do: remove something, add something, or change something? + + +## Where? + +Specify which part of the documentation you want to change. For example, the name of an existing documentation section or a line number in `README.md`. + + +## The essence + +Please describe the essence of the proposed change. diff --git a/.github/ISSUE_TEMPLATE/feature_request.md b/.github/ISSUE_TEMPLATE/feature_request.md new file mode 100644 index 0000000..117d79f --- /dev/null +++ b/.github/ISSUE_TEMPLATE/feature_request.md @@ -0,0 +1,17 @@ +--- +name: Feature request +about: Suggest an idea for this project +title: '' +labels: enhancement +assignees: pomponchik + +--- + +## Short description + +What do you propose and why do you consider it important? + + +## Some details + +If you can, provide code examples that will show how your proposal will work. Also, if you can, indicate which alternative approaches you have considered. And finally, describe how you propose to verify that your idea is implemented correctly, if at all possible. diff --git a/.github/ISSUE_TEMPLATE/question.md b/.github/ISSUE_TEMPLATE/question.md new file mode 100644 index 0000000..6f86494 --- /dev/null +++ b/.github/ISSUE_TEMPLATE/question.md @@ -0,0 +1,12 @@ +--- +name: Question or consultation +about: Ask anything about this project +title: '' +labels: question +assignees: pomponchik + +--- + +## Your question + +Here you can freely describe your question about the project. Please read the documentation provided before doing this, and ask the question only if it is not answered there. In addition, please keep in mind that this is a free non-commercial project and user support is optional for its author. Response times are not guaranteed. diff --git a/.github/workflows/lint.yml b/.github/workflows/lint.yml new file mode 100644 index 0000000..7c46fca --- /dev/null +++ b/.github/workflows/lint.yml @@ -0,0 +1,63 @@ +name: Lint + +on: + push + +jobs: + build: + + runs-on: ubuntu-latest + strategy: + matrix: + python-version: ['3.8', '3.9', '3.10', '3.11', '3.12', '3.13', '3.14', '3.14t', '3.15'] + + steps: + - uses: actions/checkout@v4 + + - name: Set up Python ${{ matrix.python-version }} + uses: actions/setup-python@v5 + with: + python-version: ${{ matrix.python-version }} + # Prerelease CPython ABIs can change; keep 3.15 aligned with current wheels. + allow-prereleases: true + check-latest: ${{ matrix.python-version == '3.15' }} + + - name: Set up uv + uses: astral-sh/setup-uv@v7 + with: + enable-cache: true + + - name: Install dependencies + shell: bash + run: uv pip install --system -r requirements_dev.txt + + - name: Install the library + shell: bash + run: uv pip install --system . + + - name: Run ruff + shell: bash + run: ruff check wasmgpu + + - name: Run ruff for tests + shell: bash + run: ruff check tests + + - name: Run mypy + shell: bash + run: >- + mypy + --show-error-codes + --strict + --disallow-any-decorated + --disallow-any-explicit + --disallow-any-expr + --disallow-any-generics + --disallow-any-unimported + --disallow-subclassing-any + --warn-return-any + wasmgpu + + - name: Run mypy for tests + shell: bash + run: mypy --exclude '^tests/typing/' tests diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml new file mode 100644 index 0000000..641ef68 --- /dev/null +++ b/.github/workflows/release.yml @@ -0,0 +1,39 @@ +name: Release + +on: + push: + branches: + - main + +jobs: + pypi-publish: + name: upload release to PyPI + runs-on: ubuntu-latest + # Specifying a GitHub environment is optional, but strongly encouraged + environment: release + permissions: + # IMPORTANT: this permission is mandatory for trusted publishing + id-token: write + steps: + - uses: actions/checkout@v4 + + - name: Set up Python 3.13 + uses: actions/setup-python@v5 + with: + python-version: '3.13' + + - name: Set up uv + uses: astral-sh/setup-uv@v7 + with: + enable-cache: true + + - name: Install dependencies + shell: bash + run: uv pip install --system -r requirements_dev.txt + + - name: Build the project + shell: bash + run: python -m build . + + - name: Publish package distributions to PyPI + uses: pypa/gh-action-pypi-publish@release/v1 diff --git a/.github/workflows/tests_and_coverage.yml b/.github/workflows/tests_and_coverage.yml new file mode 100644 index 0000000..55cf4c2 --- /dev/null +++ b/.github/workflows/tests_and_coverage.yml @@ -0,0 +1,68 @@ +name: Tests + +on: + push + +jobs: + build: + + runs-on: ${{ matrix.os }} + strategy: + matrix: + os: [macos-latest, ubuntu-latest, windows-latest] + python-version: ['3.8', '3.9', '3.10', '3.11', '3.12', '3.13', '3.14', '3.14t', '3.15'] + + steps: + - uses: actions/checkout@v4 + - name: Set up Python ${{ matrix.python-version }} + uses: actions/setup-python@v5 + with: + python-version: ${{ matrix.python-version }} + # Prerelease CPython ABIs can change; keep 3.15 aligned with current wheels. + allow-prereleases: true + check-latest: ${{ matrix.python-version == '3.15' }} + + - name: Set up uv + uses: astral-sh/setup-uv@v7 + with: + enable-cache: true + + - name: Install dependencies + shell: bash + run: | + uv pip install --system -r requirements_dev.txt + uv pip install --system . + + - name: Test Python validation and configuration (no GPU required) + shell: bash + run: | + coverage erase + coverage run -m pytest -m 'not gpu' --cache-clear + coverage combine + coverage report -m + coverage xml + # This report covers only tests that require no GPU. GPU execution and + # shader conformance are verified locally; these jobs do not claim that + # skipped GPU tests passed. Never enable a software execution fallback. + + - name: Build source distribution and wheel + shell: bash + run: python -m build + + - name: Upload Python-only coverage to Coveralls + if: runner.os == 'Linux' && matrix.python-version == '3.11' + env: + COVERALLS_REPO_TOKEN: ${{ secrets.COVERALLS_REPO_TOKEN }} + uses: coverallsapp/github-action@v2 + with: + format: cobertura + file: coverage.xml + flag-name: python-only + continue-on-error: true + + - name: Upload Python-only coverage report + if: runner.os == 'Linux' && matrix.python-version == '3.11' + uses: actions/upload-artifact@v4 + with: + name: python-only-coverage + path: coverage.xml diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..b427988 --- /dev/null +++ b/.gitignore @@ -0,0 +1,21 @@ +.DS_Store +__pycache__ +venv +.pytest_cache +build +dist +*.egg-info +test.py +.coverage +.coverage.* +.idea +.ruff_cache +.mutmut-cache +.mypy_cache +html +CLAUDE.md +.claude +mutants +planning_features.md +coverage.xml +/tests/benchmark-results.json diff --git a/MANIFEST.in b/MANIFEST.in new file mode 100644 index 0000000..890aa1c --- /dev/null +++ b/MANIFEST.in @@ -0,0 +1,3 @@ +include requirements_dev.txt +recursive-include tests *.py *.c *.rs *.wasm *.zip *.md *.json +exclude tests/benchmark-results.json diff --git a/README.md b/README.md index 0f918f4..01e6809 100644 --- a/README.md +++ b/README.md @@ -1 +1,219 @@ -# wasmgpu \ No newline at end of file +# wasmgpu + +Run independent WebAssembly instances on a hardware GPU from Python 3.8+. + +```python +import wasmgpu + +module = wasmgpu.Module("worker.wasm") +with module.spawn(100_000) as instances: + results = instances.call("process", inputs) +``` + +The runtime is implemented in this repository: + +```text +Python → wgpu-py → WGSL bytecode interpreter → wgpu-native → Metal / Vulkan / DX12 +``` + +Python parses and validates the module, uploads initial state, dispatches work, +and reads results. **All guest instructions and WASI services run on the GPU.** +There is no CPU interpreter, CPU fallback, host WASI service loop, or dependency +on Wasmtime in the installed package. Software GPU adapters are rejected. + +This is an experimental engine with the supported profile below, **not a +complete implementation of every WebAssembly proposal or a security sandbox**. +Apple M4 / Metal has been tested. Vulkan and DX12 use the same shader but have +not yet been verified on physical hardware in this project. + +## Install + +Use a virtual environment for installation, development and tests: + +```sh +python3 -m venv venv +venv/bin/python -m pip install -e . +venv/bin/python -m pip install -r requirements_dev.txt +``` + +On Windows, use `venv\Scripts\python.exe` instead. Dependency markers select +wgpu-py 0.18.0 for Python 3.8, 0.24.0 for 3.9, 0.31.0 for 3.10, and 0.32.0 for +3.11+. Actual Metal execution was verified with Python 3.8, 3.9, 3.10 and 3.11. The other +Python/backend combinations still require hardware testing. + +## Calls and state + +`Module` accepts a path or binary WASM bytes. It does not compile Rust, C or WAT. +`module.exports` lists export names and kinds. Instantiate once, then reuse the +instances: linear memory, globals, tables, files and descriptors persist between +calls. Instances have independent state, including across internal GPU batches. + +```python +with module.spawn(3) as instances: + a = instances.call("one_argument", [10, 20, 30]) + b = instances.call("two_arguments", [(1, 2), (3, 4), (5, 6)]) + c = instances.call("no_arguments") +``` + +There must be exactly one input row per instance. Input validation completes +before any instance executes. Results are a list in instance order: scalars for +one result, tuples for multiple results, and `None` for no results. Integers +return as signed i32/i64; integer inputs accept signed or unsigned bit patterns. +Floating inputs/outputs are Python floats; Python may quiet signalling f32 NaNs +at this API boundary. Only null references can be supplied from Python. + +```python +instances.write_memory(1024, b"payload", instance=0) +content = instances.read_memory(1024, 7, instance=0) +``` + +These are explicit data transfers. Reading or writing memory does not execute +guest code on the CPU. A core WASM start function executes during `spawn`. +For WASI reactor modules, call an exported `_initialize` once if the compiler +requires it; for command modules, explicitly call `_start`. + +`Trap` provides `traps` (instance index → reason) and `results` (including +successful peers). Completed effects before a trap persist; calls are not +transactions. A later call can reuse the instances. `proc_exit` raises a trap +and records the per-instance code in `instances.exit_codes`. Explicitly close +instances, preferably with a context manager, to release GPU buffers. + +## Embedded WASI Preview 1 + +Files are **byte contents embedded into each instance's GPU filesystem**. +There are no host directory mounts, host file operations or network access. + +```python +module = wasmgpu.Module("worker.wasm", files={"data/input.txt": b"1.25\n2.5\n"}) +wasi = wasmgpu.Wasi( + args=["worker", "data/input.txt"], + env={"MODE": "batch"}, + stdin=b"input stream\n", + storage_size=256 * 1024, + max_files=64, + max_fds=64, + seed=123, +) +with module.spawn(8, wasi=wasi, memory_pages=32, stack_size=4096) as instances: + instances.call("process") + output_file = instances.read_file("result.txt", instance=0) + stdout = instances.stdout # list of captured bytes, one per instance + stderr = instances.stderr +``` + +`Wasi(files=...)` overrides same-named `Module(files=...)` entries. Configuration +and contents are copied at instantiation. Root is preopened at fd 3 as `.`; +fd 0/1/2 are emulated stdin/stdout/stderr. Paths use `/`, are limited to 255 UTF-8 +bytes after resolution, and cannot escape their directory capability. +`max_files` includes directories, symlinks and four reserved entries; +`storage_size` includes stdin, stdout, stderr and all file contents. + +The shader implements file creation, reads/writes and positioned I/O, seeks, +truncation/allocation, descriptor rights, metadata/timestamps, directory +enumeration, rename, hardlinks, symlinks and unlink. Deleted storage is reclaimed; +open descriptors and hardlinks keep their inode alive. `read_file` is an +inspection API for a regular file's stored path; guest code resolves symlinks. + +Other operating-system services have explicit virtual semantics: + +- Arguments/environment come from the embedded configuration. +- All clocks use a per-instance virtual counter. It starts at `clock_epoch_ns` + (default 0) and advances by `clock_resolution_ns` (default 1) per interpreter + instruction. It never reads the host clock. +- `poll_oneoff` reports ready virtual descriptors or advances virtual time to + the earliest clock deadline, without sleeping on the host. +- `random_get` uses ChaCha20 on the GPU. The seed is a nonzero u32 or a 32-byte + key; an instance index supplies its nonce. Streams are deterministic and + independent of batch size. The default seed is public and provides **no + unpredictable system entropy**. +- `sched_yield` is a no-op in the isolated instance model. Signals terminate the + affected invocation; no host process is signalled. +- No virtual sockets are provisioned. Socket imports return `BADF` for invalid + descriptors and `NOTSOCK` for existing non-socket descriptors. They never open + host sockets. `sync`/`datasync` operate on the in-memory filesystem only. + +All 46 Preview 1 imports have signature validation and GPU dispatch. This is an +emulated environment, not a promise of an ordinary operating system or complete +WASI conformance. Preview 2/3 and arbitrary host imports are unsupported. + +## Supported WASM profile and limits + +Supported: all scalar MVP numeric instructions; i32/i64; software IEEE-754 f32 +and f64 including subnormals, signed zero and rounding; direct/indirect recursive +calls; blocks/loops/branches; multi-value; mutable globals; memory32; multiple +funcref/externref tables; reference instructions; sign extension; saturating +conversions; bulk memory/table operations and passive/declarative segments. + +Currently unsupported: SIMD, threads/shared memory, exceptions, tail calls, GC, +typed function references, memory64, multiple linear memories, imported +memories/tables/globals, and module linking. Unsupported features fail explicitly. +They never trigger execution through a CPU engine. Tables and linear memory have +fixed GPU growth budgets; `grow` returns -1 when the reserved capacity is reached. +The current memory32 addressing implementation caps memory at 65,535 pages. + +`spawn` exposes resource controls: + +| Option | Default | Meaning | +|---|---:|---| +| `memory_pages` | up to 16, at least declared minimum | Reserved 64 KiB pages per instance, capped by the module maximum | +| `table_elements` | up to 256, at least each declared minimum | Growth capacity per table, capped by each declared maximum | +| `stack_size` | 256 | 64-bit value slots per instance, including locals | +| `call_depth` | 64 | Nested call frames per instance | +| `fuel` | 10,000,000 | Invocation budget; bulk work also consumes fuel | +| `quantum` | 4096 | Interpreter instructions per dispatch before resumption | +| `batch_size` | device-derived | Instances per dispatch/buffer group | +| `max_resident_bytes` | 512 MiB | Total resident buffer allocation budget | + +All persistent instance state stays on the GPU. Internal batching respects device +buffer limits; it does not page state to a CPU runtime. Thus 100,000 tiny workers +are practical, but 100,000 workers with 1 MiB private memory require about 100 GiB +before stacks/files. Excessive allocations fail before allocation. `resident_bytes` +and `adapter_info` expose the allocation estimate and selected hardware. + +A dispatch quantum is not a real-time deadline. Large individual bulk operations +and WASI operations can take longer than a scalar instruction. Tune resource +budgets for trusted workloads; this runtime is not suitable for hostile modules. + +## Verification and benchmarks + +```sh +venv/bin/python -m pytest tests -q # hardware GPU required; absence fails +venv/bin/python -m pytest tests -m 'not gpu' # parser/configuration/failure tests only +WASMGPU_COVERAGE_BRANCH=true venv/bin/python -m coverage run -m pytest tests -q +venv/bin/python -m coverage combine +venv/bin/python -m coverage report -m +venv/bin/python -m tests.benchmark # run alone, without concurrent GPU jobs +venv/bin/python -m build +``` + +All 358 tests passed on Apple M4 / Metal with Python 3.11, including 40 official +spec suites plus API, numeric, WASI and compiled-code cases. This run covered +100% of Python statements and branches; that percentage does not measure WGSL. +Python 3.8/3.9/3.10 each passed 315 cases with their pinned GPU backend; Python 3.11 +also runs every official suite. Python 3.15rc2 passed the checks that require no GPU. +Tests compare scalar operations with Wasmtime, exercise 100,003 concurrent +instances, and run actual compiled C and Rust fixtures with allocation, internal +calls, f64, libc/Rust formatting and embedded files. Forty unmodified official +WebAssembly 2.0 core suites contribute over 20,000 module/action/assertion commands, +including malformed modules and precise NaN bit checks. See +[fixture provenance](tests/fixtures/README.md). They are a selected subset, not the +complete spec suite; Python coverage does not measure WGSL instruction coverage. + +[Benchmark script](tests/benchmark.py) records medians after warmup, hardware and +versions in the local `tests/benchmark-results.json` file. This generated report +is ignored by Git and excluded from distributions; use `--output` to choose a +different path. GPU timing includes Python +packing, transfers, every dispatch/synchronization and result decoding. Two CPU +baselines use Wasmtime: one Python call per worker, and a single WASM loop that +writes every output followed by reading those outputs into a Python list. Module +compilation/instantiation is excluded from call timings; GPU spawn is separate. + +Compare both CPU baselines when interpreting results: reducing Python call +overhead alone does not establish an acceleration of WASM computation over CPU +JIT execution. Performance depends on the workload, instance count and hardware. + +CI runs only Python tests that require no GPU, static checks, and package builds, +as configured for ordinary hosted runners. The Linux / Python 3.11 job uploads its +coverage report to Coveralls with the `python-only` flag and saves it as a GitHub +artifact. This report does not claim GPU tests have run. Full conformance and +performance tests must be run locally with a hardware GPU. diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..aa0c63e --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,73 @@ +[build-system] +requires = ["setuptools==68.0.0"] +build-backend = "setuptools.build_meta" + +[project] +name = "wasmgpu" +version = "0.0.1" +authors = [ + { name="Evgeniy Blinov", email="zheni-b@yandex.ru" }, +] +description = 'Batched WebAssembly execution in a GPU compute-shader interpreter' +readme = "README.md" +requires-python = ">=3.8" +dependencies = [ + "wgpu==0.18.0; python_version < '3.9'", + "wgpu==0.24.0; python_version >= '3.9' and python_version < '3.10'", + "wgpu==0.31.0; python_version >= '3.10' and python_version < '3.11'", + "wgpu==0.32.0; python_version >= '3.11'", +] +classifiers = [ + "Operating System :: OS Independent", + 'Operating System :: MacOS :: MacOS X', + 'Operating System :: Microsoft :: Windows', + 'Operating System :: POSIX', + 'Operating System :: POSIX :: Linux', + 'Programming Language :: Python', + 'Programming Language :: Python :: 3.8', + 'Programming Language :: Python :: 3.9', + 'Programming Language :: Python :: 3.10', + 'Programming Language :: Python :: 3.11', + 'Programming Language :: Python :: 3.12', + 'Programming Language :: Python :: 3.13', + 'Programming Language :: Python :: 3.14', + 'Programming Language :: Python :: 3.15', + 'Programming Language :: Python :: Free Threading', + 'License :: OSI Approved :: MIT License', + 'Intended Audience :: Developers', + 'Topic :: Software Development :: Libraries', + 'Typing :: Typed', +] +keywords = [ + 'WASM', +] + +[tool.setuptools.package-data] +"wasmgpu" = ["py.typed", "*.wgsl"] + +[tool.setuptools.packages.find] +include = ["wasmgpu*"] + +[tool.mutmut] +paths_to_mutate=["wasmgpu"] + +[tool.coverage.run] +branch = "${WASMGPU_COVERAGE_BRANCH-false}" +omit = ["*tests*"] +parallel = true +plugins = ["coverage_pyver_pragma"] +source = ["wasmgpu"] + +[tool.pytest.ini_options] +norecursedirs = ["build", "mutants"] +markers = ["gpu: requires a real hardware GPU; never uses a software adapter"] + +[tool.ruff] +lint.ignore = ['E501', 'E712', 'PTH123', 'PTH118', 'PLR2004', 'PTH107', 'SIM105', 'SIM102', 'RET503', 'PLR0912', 'C901', 'E731', 'F821'] +lint.select = ["ERA001", "YTT", "ASYNC", "BLE", "B", "A", "COM", "INP", "PIE", "T20", "PT", "RSE", "RET", "SIM", "SLOT", "TID252", "ARG", "PTH", "I", "C90", "N", "E", "W", "D201", "D202", "D419", "F", "PL", "PLE", "PLR", "PLW", "RUF", "TRY201", "TRY400", "TRY401"] +lint.isort.combine-as-imports = true +format.quote-style = "single" + +[project.urls] +'Source' = 'https://github.com/mutating/wasmgpu' +'Tracker' = 'https://github.com/mutating/wasmgpu/issues' diff --git a/requirements_dev.txt b/requirements_dev.txt new file mode 100644 index 0000000..2eb4788 --- /dev/null +++ b/requirements_dev.txt @@ -0,0 +1,19 @@ +pytest==8.3.5 +pytest-xdist==3.6.1; python_version < '3.9' +pytest-xdist==3.8.0; python_version >= '3.9' +coverage==7.6.1; python_version == '3.8' +coverage==7.6.10; python_version >= '3.9' +coverage-pyver-pragma==0.4.0 +build==1.2.2.post1 +mypy==1.14.1 +pytest-mypy-testing==0.1.3 +ruff==0.14.6 +mutmut==3.2.3 +cosmic-ray==8.3.15; python_version < '3.9' +cosmic-ray==8.4.6; python_version >= '3.9' +full_match==0.0.3 +locklib==0.0.22 + +# Wasmtime is exclusively an independent test/benchmark reference. +wasmtime==24.0.0; python_version < '3.9' +wasmtime==49.0.0; python_version >= '3.9' diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/benchmark.py b/tests/benchmark.py new file mode 100644 index 0000000..916f2d3 --- /dev/null +++ b/tests/benchmark.py @@ -0,0 +1,118 @@ +"""Reproducible, end-to-end GPU / Wasmtime comparison; run as a module.""" +from __future__ import annotations + +import argparse +import hashlib +import json +import platform +import statistics +import struct +import time +from datetime import datetime, timezone +from importlib.metadata import version +from pathlib import Path + +import wasmtime + +import wasmgpu + +WAT = '''(module + (memory (export "memory") 0 64) + (func $process (export "process") (param $x i32) (param $rounds i32) (result i32) + block $done loop $again + local.get $rounds i32.eqz br_if $done + local.get $x i32.const 1664525 i32.mul i32.const 1013904223 i32.add + local.get $x i32.const 13 i32.rotl i32.xor local.set $x + local.get $rounds i32.const 1 i32.sub local.set $rounds br $again + end end local.get $x) + (func (export "batch") (param $count i32) (param $rounds i32) (result i32) + (local $i i32) (local $sum i32) (local $value i32) + block $done loop $again + local.get $i local.get $count i32.ge_u br_if $done + local.get $i local.get $rounds call $process local.set $value + local.get $i i32.const 4 i32.mul local.get $value i32.store + local.get $sum local.get $value i32.xor local.set $sum + local.get $i i32.const 1 i32.add local.set $i br $again + end end local.get $sum))''' + + +def median_seconds(function, repetitions): + times = [] + for _ in range(repetitions): + start = time.perf_counter() + function() + times.append(time.perf_counter() - start) + return statistics.median(times) + + +def run(count, rounds, repetitions): + wasm = bytes(wasmtime.wat2wasm(WAT)) + engine = wasmtime.Engine() + store = wasmtime.Store(engine) + reference = wasmtime.Instance(store, wasmtime.Module(engine, wasm), []) + function = reference.exports(store)['process'] + batched = reference.exports(store)['batch'] + memory = reference.exports(store)['memory'] + memory.grow(store, (count * 4 + 65535) // 65536) + + def cpu_batch(): + batched(store, count, rounds) + return list(struct.unpack('<' + 'i' * count, memory.read(store, 0, count * 4))) + rows = [(i, rounds) for i in range(count)] + expected = [function(store, *row) for row in rows] + start = time.perf_counter() + # The workload has no state: reusing a Wasmtime instance avoids charging it + # artificial setup costs. Both CPU paths execute exactly the same WASM code. + with wasmgpu.Module(wasm).spawn(count, stack_size=32, call_depth=4, quantum=65536, memory_pages=0) as instances: + spawn_seconds = time.perf_counter() - start + actual = instances.call('process', rows) + if actual != expected: + raise AssertionError('GPU results differ from Wasmtime') + checksum = 0 + for value in actual: + checksum ^= value + if batched(store, count, rounds) != checksum: + raise AssertionError('batched CPU results differ') + # Warm up all paths before measuring. GPU timing includes Python input + # packing, upload, every dispatch/synchronization, readback and decoding. + [function(store, *row) for row in rows] + if cpu_batch() != expected: + raise AssertionError('CPU batched output differs') + gpu_seconds = median_seconds(lambda: instances.call('process', rows), repetitions) + cpu_seconds = median_seconds(lambda: [function(store, *row) for row in rows], repetitions) + cpu_batch_seconds = median_seconds(cpu_batch, repetitions) + return { + 'instances': count, 'rounds': rounds, 'repetitions': repetitions, + 'gpu_call_seconds': gpu_seconds, 'wasmtime_python_loop_seconds': cpu_seconds, + 'wasmtime_wasm_loop_seconds': cpu_batch_seconds, + 'speedup_vs_python_loop': cpu_seconds / gpu_seconds, + 'speedup_vs_wasm_loop': cpu_batch_seconds / gpu_seconds, + 'spawn_seconds': spawn_seconds, 'resident_bytes': instances.resident_bytes, + 'checksum': checksum, 'adapter': instances.adapter_info, + 'wasm_sha256': hashlib.sha256(wasm).hexdigest(), + 'gpu_config': {'stack_size': 32, 'call_depth': 4, 'quantum': 65536, 'memory_pages': 0}, + } + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument('--counts', type=int, nargs='+', default=[1000, 10000, 100000]) + parser.add_argument('--rounds', type=int, nargs='+', default=[1, 100]) + parser.add_argument('--repetitions', type=int, default=5) + parser.add_argument('--output', type=Path, default=Path('tests/benchmark-results.json')) + args = parser.parse_args() + if min(*args.counts, *args.rounds, args.repetitions) < 1: + parser.error('counts, rounds and repetitions must be positive') + report = { + 'platform': platform.platform(), 'python': platform.python_version(), + 'wgpu': version('wgpu'), 'wasmtime': version('wasmtime'), + 'workload': 'i32 LCG + rotate/xor; stateless; medians after warmup', + 'timing': 'GPU call includes input packing/upload/dispatch/readback/decoding; CPU wasm loop includes reading every output into a Python list; spawn reported separately', + 'results': [run(count, rounds, args.repetitions) for rounds in args.rounds for count in args.counts], + } + report['completed_at_utc'] = datetime.now(timezone.utc).isoformat() + args.output.write_text(json.dumps(report, indent=2) + '\n') + + +if __name__ == '__main__': + main() diff --git a/tests/build_fixtures.py b/tests/build_fixtures.py new file mode 100644 index 0000000..c67495f --- /dev/null +++ b/tests/build_fixtures.py @@ -0,0 +1,45 @@ +"""Rebuild offline conformance fixtures with an explicitly supplied WABT/spec checkout.""" +from __future__ import annotations + +import argparse +import json +import subprocess +import tempfile +import zipfile +from pathlib import Path + +SPEC_COMMIT = '05ca4182176763112561ae20153975c12bd689e4' # WebAssembly/spec v2.0.0 +SUITES = ['i32', 'i64', 'f32', 'f64', 'f32_cmp', 'f64_cmp', 'conversions', 'float_exprs', 'float_literals', 'float_memory', 'int_exprs', 'int_literals', 'block', 'loop', 'if', 'br', 'br_if', 'br_table', 'call', 'call_indirect', 'return', 'select', 'local_get', 'local_set', 'local_tee', 'memory', 'memory_copy', 'memory_fill', 'memory_init', 'memory_grow', 'memory_size', 'load', 'store', 'align', 'address', 'const', 'end', 'func', 'nop', 'unreachable', 'unwind'] + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument('--spec', type=Path, required=True) + parser.add_argument('--wast2json', type=Path, required=True) + parser.add_argument('--output', type=Path, default=Path('tests/fixtures/core-spec.zip')) + args = parser.parse_args() + revision = subprocess.check_output(['git', '-C', str(args.spec), 'rev-parse', 'HEAD'], text=True).strip() + if revision != SPEC_COMMIT: + parser.error(f'spec checkout must be at {SPEC_COMMIT}, found {revision}') + with tempfile.TemporaryDirectory() as directory: + output = Path(directory) + for name in SUITES: + source = args.spec / 'test' / 'core' / (name + '.wast') + if source.exists(): + subprocess.run([str(args.wast2json), str(source), '-o', str(output / (name + '.json'))], check=True) + with zipfile.ZipFile(args.output, 'w', zipfile.ZIP_DEFLATED, compresslevel=9) as archive: + for path in sorted(output.iterdir()): + contents = path.read_bytes() + if path.suffix == '.json': + data = json.loads(contents) + data['source_filename'] = 'test/core/' + path.stem + '.wast' + contents = (json.dumps(data, separators=(',', ':')) + '\n').encode() + info = zipfile.ZipInfo(path.name, (2020, 1, 1, 0, 0, 0)) + info.compress_type = zipfile.ZIP_DEFLATED + archive.writestr(info, contents) + info = zipfile.ZipInfo('LICENSE', (2020, 1, 1, 0, 0, 0)) + archive.writestr(info, (args.spec / 'LICENSE').read_bytes()) + + +if __name__ == '__main__': + main() diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..0bbf627 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,40 @@ +from __future__ import annotations + +import struct + +import pytest +import wasmtime + +import wasmgpu + + +@pytest.fixture(scope='session') +def engine(): + return wasmtime.Engine() + + +def binary(wat): + return bytes(wasmtime.wat2wasm(wat)) + + +def oracle(engine, wat, name, rows): + module = wasmtime.Module(engine, binary(wat)) + store = wasmtime.Store(engine) + instance = wasmtime.Instance(store, module, []) + function = instance.exports(store)[name] + results = [] + for row in rows: + try: + results.append(function(store, *row)) + except wasmtime.Trap as error: + results.append(error) + return results + + +def gpu(wat, rows, name='run', **options): + with wasmgpu.Module(binary(wat)).spawn(len(rows), **options) as instances: + return instances.call(name, rows) + + +def float_bits(value, ty): + return struct.pack(' +#include +#include +#include + +static uint64_t factorial(uint32_t n) { + if (n < 2) return 1; + return n * factorial(n - 1); +} + +static double transform(double x) { return x * 1.25 + 0.5; } +static double (*volatile callback)(double) = transform; + +__attribute__((export_name("process"))) +double process(int n) { + double *values = malloc((size_t)n * sizeof(double)); + if (values == NULL) return -1.0; + double result = 0; + for (int i = 0; i < n; ++i) { + values[i] = callback((double)i); + result += values[i]; + } + free(values); + return result + (double)factorial((uint32_t)n % 12); +} + +__attribute__((export_name("file_process"))) +int file_process(void) { + FILE *input = fopen("numbers.txt", "r"); + if (input == NULL) return -1; + char line[80]; + double sum = 0; + while (fgets(line, sizeof(line), input)) sum += strtod(line, NULL); + fclose(input); + FILE *output = fopen("result.txt", "w"); + if (output == NULL) return -2; + fprintf(output, "sum=%.3f\n", sum); + fclose(output); + printf("processed\n"); + return (int)(sum * 1000); +} diff --git a/tests/fixtures/worker.rs b/tests/fixtures/worker.rs new file mode 100644 index 0000000..e99bbea --- /dev/null +++ b/tests/fixtures/worker.rs @@ -0,0 +1,13 @@ +#[no_mangle] +pub extern "C" fn process(n: u32) -> f64 { + let numbers: Vec = (0..n).map(|i| f64::from(i) * 1.25 + 0.5).collect(); + numbers.iter().sum() +} + +#[no_mangle] +pub extern "C" fn file_process() -> i32 { + let content = std::fs::read_to_string("numbers.txt").unwrap(); + let sum: f64 = content.lines().map(|line| line.parse::().unwrap()).sum(); + std::fs::write("rust-result.txt", format!("sum={sum:.3}\n")).unwrap(); + (sum * 1000.0) as i32 +} diff --git a/tests/fixtures/worker.wasm b/tests/fixtures/worker.wasm new file mode 100755 index 0000000..272565d Binary files /dev/null and b/tests/fixtures/worker.wasm differ diff --git a/tests/test_api.py b/tests/test_api.py new file mode 100644 index 0000000..bf6c0de --- /dev/null +++ b/tests/test_api.py @@ -0,0 +1,128 @@ +from __future__ import annotations + +import pytest + +import wasmgpu +from wasmgpu import runtime + +from .conftest import binary + +pytestmark = pytest.mark.gpu + + +@pytest.fixture +def instances(): + wat = '''(module (memory (export "memory") 1 2) + (func (export "none")) + (func (export "int") (param i32) (result i32) local.get 0) + (func (export "float") (param f32) (result f32) local.get 0) + (func (export "reference") (param externref) (result externref) local.get 0))''' + with wasmgpu.Module(binary(wat)).spawn(2) as guest: + yield guest + + +@pytest.mark.parametrize(('name', 'rows', 'error'), [ + ('int', None, TypeError), ('int', [1], ValueError), ('int', [1, True], TypeError), + ('int', [1, 2**32], OverflowError), ('int', [1, -2**31 - 1], OverflowError), + ('int', [(1, 2), (3, 4)], TypeError), ('int', ['x', 'y'], TypeError), + ('none', [1, 2], TypeError), ('float', [True, 0], TypeError), + ('float', [1e100, 0], OverflowError), ('reference', [1, None], TypeError), + ('missing', None, KeyError), ('memory', None, KeyError), +]) +def test_call_validation(instances, name, rows, error): + with pytest.raises(error): + instances.call(name, rows) + assert instances.call('int', [12, 34]) == [12, 34] + + +def test_reference_and_float_values(instances): + assert instances.call('reference', [None, None]) == [None, None] + assert instances.call('float', [1, -2.5]) == [1.0, -2.5] + assert instances.call('int', [2**32 - 1, -2**31]) == [-1, -2**31] + assert instances.call('none') == [None, None] + + +@pytest.mark.parametrize(('offset', 'size', 'index', 'error'), [ + (-1, 1, 0, ValueError), (0, -1, 0, ValueError), (False, 1, 0, TypeError), + (0, 1, -1, ValueError), (0, 1, 2, IndexError), (65536, 1, 0, IndexError), + (65537, 0, 0, IndexError), (1, 65536, 0, IndexError), +]) +def test_memory_read_validation(instances, offset, size, index, error): + with pytest.raises(error): + instances.read_memory(offset, size, instance=index) + + +def test_zero_length_read_write_and_absent_files(instances): + assert instances.read_memory(65536, 0) == b'' + instances.write_memory(65536, b'') + assert instances.stdout == [b'', b''] + assert instances.stderr == [b'', b''] + with pytest.raises(FileNotFoundError): + instances.read_file('missing') + with pytest.raises(IndexError): + instances.write_memory(65536, b'x') + + +@pytest.mark.parametrize(('wat', 'options', 'error'), [ + ('(module (memory 1 2))', {'memory_pages': 0}, ValueError), + ('(module (memory 1 2))', {'memory_pages': 3}, ValueError), + ('(module (table 10 funcref))', {'table_elements': 5}, ValueError), + ('(module (func (local i32 i32)))', {'stack_size': 1}, wasmgpu.ResourceLimitError), + ('(module)', {'batch_size': 2**32 - 1}, wasmgpu.ResourceLimitError), + ('(module)', {'batch_size': 0}, ValueError), + ('(module (memory 0 10000))', {'memory_pages': 10000, 'max_resident_bytes': 2**32 - 1}, wasmgpu.ResourceLimitError), + ('(module (memory 0) (data (i32.const 0) "x"))', {}, wasmgpu.Trap), + ('(module (table 0 funcref) (func $x) (elem (i32.const 0) $x))', {}, wasmgpu.Trap), + ('(module (func $x unreachable) (start $x))', {}, wasmgpu.Trap), +]) +def test_instance_resource_and_initialization_failures(wat, options, error): + with pytest.raises(error): + wasmgpu.Module(binary(wat)).spawn(1, **options) + + +def test_gpu_allocation_failure_releases_partial_buffers(monkeypatch): + context = runtime._context() + original = context.buffer + created = [] + destroyed = [] + + def allocate(*args, **kwargs): + if len(created) == 3: + raise RuntimeError('simulated GPU allocation failure') + buffer = original(*args, **kwargs) + created.append(buffer) + destroy = buffer.destroy + + def record_destroy(): + destroyed.append(buffer) + destroy() + monkeypatch.setattr(buffer, 'destroy', record_destroy) + return buffer + + monkeypatch.setattr(context, 'buffer', allocate) + with pytest.raises(RuntimeError, match='simulated GPU allocation failure'): + wasmgpu.Module(binary('(module)')).spawn(1) + assert set(created) == set(destroyed) + assert len(destroyed) == len(created) + + +def test_environment_and_batch_buffers_count_toward_budget(monkeypatch): + module = wasmgpu.Module(binary('(module)')) + with module.spawn(1) as base: + budget = base.resident_bytes + with pytest.raises(wasmgpu.ResourceLimitError, match='configuration buffers'): + module.spawn(1, max_resident_bytes=budget - 1) + with pytest.raises(wasmgpu.ResourceLimitError): + module.spawn(1, max_resident_bytes=budget, wasi=wasmgpu.Wasi(args=['x' * 4096])) + context = runtime._context() + monkeypatch.setitem(context.limits, 'max-buffer-size', 128) + with pytest.raises(wasmgpu.ResourceLimitError, match='program and guest environment'): + module.spawn(1, wasi=wasmgpu.Wasi(args=['x' * 4096])) + + +def test_program_buffer_allocation_failure_closes_cleanly(monkeypatch): + def fail(*_args, **_kwargs): + raise RuntimeError('device allocation failed') + monkeypatch.setattr(runtime._context(), 'buffer', fail) + with pytest.raises(RuntimeError, match='device allocation failed'): + wasmgpu.Module(binary('(module)')).spawn(1) diff --git a/tests/test_compiled.py b/tests/test_compiled.py new file mode 100644 index 0000000..2c0ade4 --- /dev/null +++ b/tests/test_compiled.py @@ -0,0 +1,73 @@ +from __future__ import annotations + +from pathlib import Path + +import pytest +import wasmtime + +import wasmgpu + +pytestmark = pytest.mark.gpu + +FIXTURE = Path(__file__).parent / 'fixtures' / 'worker.wasm' +RUST_FIXTURE = FIXTURE.with_name('worker-rust.wasm') + + +@pytest.mark.parametrize('fixture', [FIXTURE, RUST_FIXTURE], ids=['C', 'Rust']) +def test_malloc_and_double_arithmetic(engine, fixture): + data = fixture.read_bytes() + store = wasmtime.Store(engine) + linker = wasmtime.Linker(engine) + linker.define_wasi() + store.set_wasi(wasmtime.WasiConfig()) + reference = linker.instantiate(store, wasmtime.Module(engine, data)) + if '_initialize' in reference.exports(store): + reference.exports(store)['_initialize'](store) + values = [0, 1, 2, 5, 10, 33, 100] + expected = [reference.exports(store)['process'](store, n) for n in values] + with wasmgpu.Module(data).spawn(len(values), stack_size=2048, call_depth=128, memory_pages=32) as instances: + if '_initialize' in instances.module.exports: + instances.call('_initialize') + assert instances.call('process', values) == expected + + +def test_c_stdio_strtod_printf_and_embedded_files(engine, tmp_path): + data = FIXTURE.read_bytes() + content = b'1.25\n2.5\n-0.125\n' + (tmp_path / 'numbers.txt').write_bytes(content) + store = wasmtime.Store(engine) + wasi = wasmtime.WasiConfig() + wasi.preopen_dir(str(tmp_path), '.') + wasi.stdout_file = str(tmp_path / 'stdout') + store.set_wasi(wasi) + linker = wasmtime.Linker(engine) + linker.define_wasi() + reference = linker.instantiate(store, wasmtime.Module(engine, data)) + reference.exports(store)['_initialize'](store) + expected = reference.exports(store)['file_process'](store) + expected_file = (tmp_path / 'result.txt').read_bytes() + with wasmgpu.Module(data, files={'numbers.txt': content}).spawn(2, stack_size=2048, call_depth=128, memory_pages=32) as instances: + instances.call('_initialize') + assert instances.call('file_process') == [expected] * 2 + assert instances.read_file('result.txt') == expected_file + assert instances.read_file('result.txt', instance=1) == expected_file + assert instances.stdout == [b'processed\n'] * 2 + + +def test_rust_std_files_parsing_and_formatting(engine, tmp_path): + data = RUST_FIXTURE.read_bytes() + content = b'1.25\n2.5\n-0.125\n' + (tmp_path / 'numbers.txt').write_bytes(content) + store = wasmtime.Store(engine) + wasi = wasmtime.WasiConfig() + wasi.preopen_dir(str(tmp_path), '.') + store.set_wasi(wasi) + linker = wasmtime.Linker(engine) + linker.define_wasi() + reference = linker.instantiate(store, wasmtime.Module(engine, data)) + expected = reference.exports(store)['file_process'](store) + expected_file = (tmp_path / 'rust-result.txt').read_bytes() + with wasmgpu.Module(data, files={'numbers.txt': content}).spawn(2, stack_size=4096, call_depth=128, memory_pages=32) as instances: + assert instances.call('file_process') == [expected] * 2 + assert instances.read_file('rust-result.txt') == expected_file + assert instances.read_file('rust-result.txt', instance=1) == expected_file diff --git a/tests/test_execution.py b/tests/test_execution.py new file mode 100644 index 0000000..ea7f160 --- /dev/null +++ b/tests/test_execution.py @@ -0,0 +1,231 @@ +from __future__ import annotations + +import ast +import struct +from pathlib import Path + +import pytest + +import wasmgpu + +from .conftest import binary, gpu, oracle + +pytestmark = pytest.mark.gpu + + +def test_recursive_calls(engine): + wat = '''(module (func $factorial (export "run") (param i32) (result i64) + local.get 0 i32.eqz if (result i64) i64.const 1 + else local.get 0 i64.extend_i32_u local.get 0 i32.const 1 i32.sub call $factorial i64.mul end))''' + rows = [(n,) for n in range(21)] + assert gpu(wat, rows) == oracle(engine, wat, 'run', rows) + + +def test_loop_branch_and_locals(engine): + wat = '''(module (func (export "run") (param $n i32) (result i32) (local $acc i32) + block $done loop $again + local.get $n i32.eqz br_if $done + local.get $acc local.get $n i32.add local.set $acc + local.get $n i32.const 1 i32.sub local.set $n br $again + end end local.get $acc))''' + rows = [(n,) for n in range(50)] + assert gpu(wat, rows, quantum=7) == oracle(engine, wat, 'run', rows) + + +def test_br_table(engine): + wat = '''(module (func (export "run") (param i32) (result i32) + block $out (result i32) + block $two block $one block $zero + local.get 0 br_table $zero $one $two + end i32.const 10 br $out + end i32.const 20 br $out + end i32.const 30 + end))''' + rows = [(0,), (1,), (2,), (3,), (-1,)] + assert gpu(wat, rows) == oracle(engine, wat, 'run', rows) + + +def test_multi_value_blocks_calls_and_returns(): + wat = '''(module + (type $pair (func (param i32 i64) (result i64 i32))) + (func $swap (type $pair) local.get 1 local.get 0) + (func (export "run") (param i32 i64) (result i64 i32) + local.get 0 local.get 1 block (type $pair) call $swap end))''' + assert gpu(wat, [(7, 99), (-2, -100)]) == [(99, 7), (-100, -2)] + + +def test_branch_carries_results_and_discards_operands(): + wat = '''(module (func (export "run") (result i32) + block (result i32) i32.const 999 i32.const 42 br 0 end))''' + assert gpu(wat, [(), ()]) == [42, 42] + + +def test_indirect_calls_and_equivalent_type_indices(): + wat = '''(module + (type $a (func (param i32) (result i32))) + (type $b (func (param i32) (result i32))) + (table 3 funcref) + (func $double (type $a) local.get 0 i32.const 2 i32.mul) + (func $wrong (result i64) i64.const 1) + (elem (i32.const 0) $double $wrong) + (func (export "run") (param i32 i32) (result i32) + local.get 0 local.get 1 call_indirect (type $b)))''' + with wasmgpu.Module(binary(wat)).spawn(4) as instances: + with pytest.raises(wasmgpu.Trap) as info: + instances.call('run', [(21, 0), (1, 1), (1, 2), (1, 3)]) + assert info.value.results == [42, None, None, None] + assert info.value.traps == {1: 'indirect call type mismatch', 2: 'uninitialized element', 3: 'out of bounds table access'} + + +def test_persistent_and_isolated_memory_globals(): + wat = '''(module (memory (export "memory") 1 2) (global $g (mut i32) (i32.const 1)) + (data (i32.const 3) "abcd") + (func (export "run") (param i32) (result i32) + global.get $g local.get 0 i32.add global.set $g + i32.const 3 global.get $g i32.store + i32.const 3 i32.load))''' + module = wasmgpu.Module(binary(wat)) + with module.spawn(3, batch_size=2, memory_pages=2) as first, module.spawn(1) as second: + assert first.read_memory(3, 4) == b'abcd' + assert first.call('run', [2, 10, 20]) == [3, 11, 21] + assert first.call('run', [2, 10, 20]) == [5, 21, 41] + assert first.read_memory(3, 4, instance=2) == struct.pack(' first[0] >= 1000 + assert instances.read_memory(32, 16) != first_random + with wasmgpu.Module(binary(wat)).spawn(2, wasi=wasmgpu.Wasi(seed=81, clock_epoch_ns=1000)) as instances: + assert instances.call('run') == first + assert instances.read_memory(32, 16) == first_random + + +def test_proc_exit_is_per_instance(): + wat = '''(module (import "wasi_snapshot_preview1" "proc_exit" (func $exit (param i32))) + (func (export "run") (param i32) local.get 0 call $exit))''' + with wasmgpu.Module(binary(wat)).spawn(2) as instances: + with pytest.raises(wasmgpu.Trap, match='proc_exit'): + instances.call('run', [0, 7]) + assert instances.exit_codes == [0, 7] + + +def test_wasi_invalid_pointers_return_errno(): + wat = '''(module (import "wasi_snapshot_preview1" "random_get" (func $random (param i32 i32) (result i32))) + (memory 1) (func (export "run") (param i32) (result i32) local.get 0 i32.const 8 call $random))''' + with wasmgpu.Module(binary(wat)).spawn(3) as instances: + assert instances.call('run', [65536, -1, 65529]) == [21, 21, 21] + + +def test_initial_files_are_copied_and_not_host_paths(): + content = bytearray(b'hello') + module = wasmgpu.Module(binary(READ_WRITE), files={'input.txt': content}) + content[:] = b'wrong' + with module.spawn(1) as instances: + assert instances.read_file('./input.txt') == b'hello' + with pytest.raises(TypeError, match='bytes'): + wasmgpu.Wasi(files={'input.txt': '/tmp/host-file'}) + with pytest.raises(ValueError, match='escapes'): + wasmgpu.Wasi(files={'../secret': b'value'}) diff --git a/tests/test_wasi_services.py b/tests/test_wasi_services.py new file mode 100644 index 0000000..8e328ae --- /dev/null +++ b/tests/test_wasi_services.py @@ -0,0 +1,276 @@ +from __future__ import annotations + +import struct + +import pytest + +import wasmgpu +from wasmgpu.wasi import _SIGNATURES + +from .conftest import binary + +pytestmark = pytest.mark.gpu + + +def wasi_module(): + imports = [] + exports = [] + for index, (name, (args, results)) in enumerate(_SIGNATURES.items()): + params = ' '.join('i32' if ty == 'i' else 'i64' for ty in args) + signature = (f'(param {params})' if params else '') + ('(result i32)' if results else '') + imports.append(f'(import "wasi_snapshot_preview1" "{name}" (func {signature}))') + exports.append(f'(export "{name}" (func {index}))') + return wasmgpu.Module(binary('(module ' + ''.join(imports + exports) + '(memory 1))')) + + +@pytest.fixture +def guest(): + with wasi_module().spawn(1, wasi=wasmgpu.Wasi(files={'dir/input': b'abcdef', 'outside': b'xyz'})) as instances: + yield instances + + +def invoke(guest, name, *args): + return guest.call(name, [args])[0] + + +def text(guest, value, offset=0): + blob = value.encode() + guest.write_memory(offset, blob) + return offset, len(blob) + + +def u32(guest, offset): + return int.from_bytes(guest.read_memory(offset, 4), 'little') + + +def open_file(guest, name='dir/input', *, rights=(1 << 30) - 1, flags=0, oflags=0, fd=3, follow=1): # noqa: PLR0913 - Mirrors the WASI path_open arguments. + pointer, size = text(guest, name) + error = invoke(guest, 'path_open', fd, follow, pointer, size, oflags, rights, rights, flags, 512) + assert error == 0 + return u32(guest, 512) + + +def test_descriptors_positioned_io_and_flags(guest): + fd = open_file(guest) + guest.write_memory(128, struct.pack('= 10**12 + subscription[8] = 1 + struct.pack_into(' None: + self.data = bytes(data) + self.pos = 0 + + def take(self, count: int) -> bytes: + if count < 0 or count > len(self.data) - self.pos: + raise ValidationError('unexpected end of WebAssembly') + result = self.data[self.pos:self.pos + count] + self.pos += count + return result + + def byte(self) -> int: + return self.take(1)[0] + + def leb(self, bits: int = 32, signed: bool = False) -> int: + value = 0 + for shift in range(0, bits, 7): + byte = self.byte() + value |= (byte & 127) << shift + if byte < 128: + if signed and byte & 64: + value -= 1 << (shift + 7) + low = -(1 << (bits - 1)) if signed else 0 + high = (1 << (bits - (1 if signed else 0))) - 1 + if not low <= value <= high: + raise ValidationError('integer representation out of range') + return value + raise ValidationError('integer representation too long') + + def vec(self, read: Callable[[], T]) -> list[T]: + count = self.leb() + if count > len(self.data) - self.pos: + raise ValidationError('vector length exceeds section size') + return [read() for _ in range(count)] + + def name(self) -> str: + try: + return self.take(self.leb()).decode('utf-8') + except UnicodeDecodeError as error: + raise ValidationError('invalid UTF-8 name') from error + + def finish(self) -> None: + if self.pos != len(self.data): + raise ValidationError('trailing bytes in section or function') + + +def value_type(reader: Reader) -> int: + result = reader.byte() + if result not in VALUE_TYPES: + raise UnsupportedFeatureError(f'unsupported value type 0x{result:02x}') + return result + + +def limits(reader: Reader, memory: bool = False) -> Limits: + flags = reader.leb() + if flags not in (0, 1): + raise UnsupportedFeatureError('shared memory, memory64 and custom page sizes are not supported') + minimum = reader.leb() + maximum = reader.leb() if flags else None + if maximum is not None and minimum > maximum: + raise ValidationError('minimum exceeds maximum') + if memory and (minimum > 65536 or (maximum is not None and maximum > 65536)): + raise ValidationError('memory32 exceeds 65536 pages') + return minimum, maximum + + +@dataclass +class Function: + type_index: int + locals: list[int] = field(default_factory=list) + instructions: list[list[int]] = field(default_factory=list) + imported: tuple[str, str] | None = None + offset: int = 0 + max_stack: int = 0 + + +@dataclass +class Control: + kind: int + height: int + params: list[int] + results: list[int] + start: int + patches: list[int] = field(default_factory=list) + unreachable: bool = False + else_patch: int | None = None + has_else: bool = False + + +class BinaryModule: + def __init__(self, data: bytes) -> None: + self.types: list[Signature] = [] + self.functions: list[Function] = [] + self.declared_functions: set[int] = set() + self.exports: dict[str, tuple[int, int]] = {} + self.globals: list[tuple[int, int, int]] = [] + self.memory: Limits | None = None + self.tables: list[tuple[int, Limits]] = [] + self.data: list[tuple[int | None, bytes]] = [] + self.elements: list[Element] = [] + self.start: int | None = None + self.data_count: int | None = None + reader = Reader(data) + if reader.take(8) != b'\0asm\x01\0\0\0': + raise ValidationError('expected a WebAssembly 1 binary module') + order = {1: 1, 2: 2, 3: 3, 4: 4, 5: 5, 6: 6, 7: 7, 8: 8, 9: 9, 12: 10, 10: 11, 11: 12} + last = 0 + bodies: list[Reader] = [] + while reader.pos < len(reader.data): + section_id = reader.byte() + section = Reader(reader.take(reader.leb())) + if section_id == 0: + section.name() + continue + if section_id not in order: + raise UnsupportedFeatureError(f'unsupported section {section_id}') + if order[section_id] <= last: + raise ValidationError('duplicate or out-of-order section') + last = order[section_id] + parsed_bodies = self._read_section(section_id, section) + if section_id == 10: + bodies = parsed_bodies + section.finish() + defined = [fn for fn in self.functions if fn.imported is None] + if len(bodies) != len(defined): + raise ValidationError('function and code section lengths differ') + if self.data_count is not None and self.data_count != len(self.data): + raise ValidationError('data count does not match data section') + for fn, body in zip(defined, bodies): + params, _ = self.types[fn.type_index] + fn.locals = list(params) + for _ in range(body.leb()): + count, ty = body.leb(), value_type(body) + if count > 65536 - len(fn.locals): + raise UnsupportedFeatureError('more than 65536 locals per function') + fn.locals.extend([ty] * count) + Compiler(self, fn, body).compile() + body.finish() + + def _read_section(self, section_id: int, section: Reader) -> list[Reader]: + if section_id == 1: + self.types = section.vec(lambda: self._type(section)) + elif section_id == 2: + section.vec(lambda: self._import(section)) + elif section_id == 3: + self.functions.extend(section.vec(lambda: Function(self._type_index(section)))) + elif section_id == 4: + self.tables = section.vec(lambda: (value_type(section), limits(section))) + if any(ty not in (FUNCREF, EXTERNREF) for ty, _ in self.tables): + raise ValidationError('table element type must be a reference') + elif section_id == 5: + memories = section.vec(lambda: limits(section, memory=True)) + if len(memories) > 1: + raise UnsupportedFeatureError('multiple memories are not supported') + self.memory = memories[0] if memories else None + elif section_id == 6: + section.vec(lambda: self._global(section)) + elif section_id == 7: + section.vec(lambda: self._export(section)) + elif section_id == 8: + self.start = self._index(section, self.functions, 'start function') + if self.signature(self.start) != ([], []): + raise ValidationError('start function must have no arguments or results') + elif section_id == 9: + self.elements = section.vec(lambda: self._element(section)) + elif section_id == 10: + return section.vec(lambda: Reader(section.take(section.leb()))) + elif section_id == 11: + self.data = section.vec(lambda: self._data(section)) + else: # The caller validated the section id; only data-count (12) remains. + self.data_count = section.leb() + return [] + + @staticmethod + def _index(reader: Reader, values: Sequence[T], name: str) -> int: + index = reader.leb() + if index >= len(values): + raise ValidationError(f'{name} index out of range: {index}') + return index + + def _type_index(self, reader: Reader) -> int: + return self._index(reader, self.types, 'type') + + def signature(self, index: int) -> Signature: + return self.types[self.functions[index].type_index] + + @staticmethod + def _type(reader: Reader) -> Signature: + if reader.byte() != 0x60: + raise UnsupportedFeatureError('GC and recursive types are not supported') + return reader.vec(lambda: value_type(reader)), reader.vec(lambda: value_type(reader)) + + def _import(self, reader: Reader) -> None: + module, name, kind = reader.name(), reader.name(), reader.byte() + if kind != 0: + raise UnsupportedFeatureError('only function imports are supported') + self.functions.append(Function(self._type_index(reader), imported=(module, name))) + + def _constant(self, reader: Reader, expected: int) -> int: + op = reader.byte() + if op in (0x41, 0x42): + ty = I32 if op == 0x41 else I64 + val = reader.leb(32 if ty == I32 else 64, signed=True) + value = val & ((1 << (32 if ty == I32 else 64)) - 1) + elif op in (0x43, 0x44): + ty = F32 if op == 0x43 else F64 + value = int.from_bytes(reader.take(4 if ty == F32 else 8), 'little') + elif op == 0xD0: + ty, value = value_type(reader), 0 + if ty not in (FUNCREF, EXTERNREF): + raise ValidationError('ref.null requires a reference type') + elif op == 0xD2: + ty, value = FUNCREF, self._index(reader, self.functions, 'function') + 1 + self.declared_functions.add(value - 1) + else: + raise UnsupportedFeatureError(f'unsupported constant expression opcode 0x{op:02x}') + if ty != expected or reader.byte() != 0x0B: + raise ValidationError('invalid constant expression') + return value + + def _global(self, reader: Reader) -> None: + ty, mutable = value_type(reader), reader.byte() + if mutable > 1: + raise ValidationError('invalid global mutability') + self.globals.append((ty, mutable, self._constant(reader, ty))) + + def _export(self, reader: Reader) -> None: + name, kind, index = reader.name(), reader.byte(), reader.leb() + counts = {0: len(self.functions), 1: len(self.tables), 2: int(self.memory is not None), 3: len(self.globals)} + if kind not in counts or index >= counts[kind] or name in self.exports: + raise ValidationError('invalid or duplicate export') + self.exports[name] = (kind, index) + if kind == 0: + self.declared_functions.add(index) + + def _data(self, reader: Reader) -> tuple[int | None, bytes]: + flags = reader.leb() + if flags not in (0, 1, 2): + raise ValidationError('invalid data segment flags') + if flags == 2 and reader.leb() != 0: + raise ValidationError('memory index out of range') + offset = None if flags == 1 else self._constant(reader, I32) + if offset is not None and self.memory is None: + raise ValidationError('data segment requires a memory') + return offset, reader.take(reader.leb()) + + def _element(self, reader: Reader) -> Element: + flags = reader.leb() + if flags > 7: + raise ValidationError('invalid element segment flags') + table_index = reader.leb() if flags in (2, 6) else 0 + offset = self._constant(reader, I32) if flags % 2 == 0 else None + if offset is not None and table_index >= len(self.tables): + raise ValidationError('active element segment requires a table') + if flags & 4: + ty = value_type(reader) if flags != 4 else FUNCREF + values = reader.vec(lambda: self._constant(reader, ty)) + else: + ty = FUNCREF + if flags != 0 and reader.byte() != 0: + raise ValidationError('invalid element kind') + values = reader.vec(lambda: self._index(reader, self.functions, 'function') + 1) + if ty not in (FUNCREF, EXTERNREF) or (offset is not None and ty != self.tables[table_index][0]): + raise ValidationError('invalid element type') + if ty == FUNCREF: + self.declared_functions.update(value - 1 for value in values if value) + return offset, values, flags in (3, 7), ty, table_index + + +def numeric_signature(op: int) -> Signature | None: # noqa: PLR0911 - Opcode signature dispatch. + if op == 0x45: + return [I32], [I32] + if 0x46 <= op <= 0x4F: + return [I32, I32], [I32] + if op == 0x50: + return [I64], [I32] + if 0x51 <= op <= 0x5A: + return [I64, I64], [I32] + if 0x5B <= op <= 0x60: + return [F32, F32], [I32] + if 0x61 <= op <= 0x66: + return [F64, F64], [I32] + for start, unary_end, end, ty in ((0x67, 0x69, 0x78, I32), (0x79, 0x7B, 0x8A, I64), (0x8B, 0x91, 0x98, F32), (0x99, 0x9F, 0xA6, F64)): + if start <= op <= end: + return [ty] * (1 if op <= unary_end else 2), [ty] + conversions = { + 0xA7: (I64, I32), 0xA8: (F32, I32), 0xA9: (F32, I32), 0xAA: (F64, I32), 0xAB: (F64, I32), + 0xAC: (I32, I64), 0xAD: (I32, I64), 0xAE: (F32, I64), 0xAF: (F32, I64), 0xB0: (F64, I64), 0xB1: (F64, I64), + 0xB2: (I32, F32), 0xB3: (I32, F32), 0xB4: (I64, F32), 0xB5: (I64, F32), 0xB6: (F64, F32), + 0xB7: (I32, F64), 0xB8: (I32, F64), 0xB9: (I64, F64), 0xBA: (I64, F64), 0xBB: (F32, F64), + 0xBC: (F32, I32), 0xBD: (F64, I64), 0xBE: (I32, F32), 0xBF: (I64, F64), + 0xC0: (I32, I32), 0xC1: (I32, I32), 0xC2: (I64, I64), 0xC3: (I64, I64), 0xC4: (I64, I64), + } + if op in conversions: + arg, result = conversions[op] + return [arg], [result] + if 0xFC00 <= op <= 0xFC07: + offset = op - 0xFC00 + return [F32 if offset % 4 < 2 else F64], [I32 if offset < 4 else I64] + return None + + +class Compiler: + def __init__(self, module: BinaryModule, function: Function, reader: Reader) -> None: + self.module, self.fn, self.reader = module, function, reader + self.stack: list[int | None] = [] + self.controls: list[Control] = [] + + def emit(self, op: int, a: int = 0, b: int = 0, c: int = 0) -> int: + self.fn.instructions.append([op, a, b, c]) + return len(self.fn.instructions) - 1 + + def pop(self, expected: int | None = None) -> int | None: + control = self.controls[-1] + if len(self.stack) == control.height and control.unreachable: + return expected + if len(self.stack) <= control.height: + raise ValidationError('operand stack underflow') + actual = self.stack.pop() + if expected is not None and actual is not None and actual != expected: + raise ValidationError('operand type mismatch') + return actual + + def effect(self, params: Sequence[int | None], results: Sequence[int | None]) -> None: + for ty in reversed(params): + self.pop(ty) + self.stack.extend(results) + self.fn.max_stack = max(self.fn.max_stack, len(self.stack)) + + def unreachable(self) -> None: + frame = self.controls[-1] + del self.stack[frame.height:] + frame.unreachable = True + + def target(self, depth: int) -> Control: + if depth >= len(self.controls): + raise ValidationError('branch depth out of range') + return self.controls[-1 - depth] + + def branch(self, target: Control, conditional: bool = False) -> list[int]: + types = target.params if target.kind == 0x03 else target.results + self.effect(types, types) + index = self.emit(BRANCH_IF if conditional else BRANCH, target.start, target.height, len(types)) + if target.kind != 0x03: + target.patches.append(index) + return types + + def require_memory(self) -> None: + if self.module.memory is None: + raise ValidationError('instruction requires memory') + + def table_index(self) -> int: + return self.module._index(self.reader, self.module.tables, 'table') + + def memory_index(self) -> None: + self.require_memory() + if self.reader.leb() != 0: + raise ValidationError('memory index out of range') + + def compile(self) -> None: # noqa: PLR0915 - One validation rule per opcode family. + _, results = self.module.types[self.fn.type_index] + self.controls.append(Control(0xFF, 0, [], results, 0)) + params: list[int] + returns: list[int] + while self.controls: + op = self.reader.byte() + if op in (0x02, 0x03, 0x04): + if op == 0x04: + self.pop(I32) + bt = self.reader.leb(33, signed=True) + if bt == -64: + params, returns = [], [] + elif bt < 0 and (bt & 127) in VALUE_TYPES: + params, returns = [], [bt & 127] + elif 0 <= bt < len(self.module.types): + params, returns = self.module.types[bt] + else: + raise ValidationError('invalid block type') + self.effect(params, []) + control = Control(op, len(self.stack), params, returns, len(self.fn.instructions)) + self.stack.extend(params) + if op == 0x04: + control.else_patch = self.emit(IF_ZERO) + self.controls.append(control) + elif op in (0x05, 0x0B): + control = self.controls[-1] + self.effect(control.results, []) + if len(self.stack) != control.height: + raise ValidationError('incorrect block result arity') + if op == 0x05: + if control.kind != 0x04 or control.has_else: + raise ValidationError('unexpected else') + control.has_else = True + end_jump = self.emit(JUMP) + assert control.else_patch is not None + self.fn.instructions[control.else_patch][1] = len(self.fn.instructions) + control.patches.append(end_jump) + control.else_patch = None + control.unreachable = False + self.stack.extend(control.params) + else: + if control.kind == 0x04 and not control.has_else and control.params != control.results: + raise ValidationError('if without else must have matching parameter/result types') + end = len(self.fn.instructions) + for index in control.patches: + self.fn.instructions[index][1] = end + if control.else_patch is not None: + self.fn.instructions[control.else_patch][1] = end + self.controls.pop() + self.stack.extend(control.results) + if not self.controls: + self.emit(RETURN) + elif op in (0x0C, 0x0D): + if op == 0x0D: + self.pop(I32) + self.branch(self.target(self.reader.leb()), op == 0x0D) + if op == 0x0C: + self.unreachable() + elif op == 0x0E: + targets = self.reader.vec(lambda: self.target(self.reader.leb())) + targets.append(self.target(self.reader.leb())) + self.pop(I32) + index = self.emit(BRANCH_TABLE, len(self.fn.instructions) + 1, len(targets)) + branch_signature = None + for target in targets: + types = self.branch(target) + if branch_signature is not None and branch_signature != types: + raise ValidationError('br_table target types differ') + branch_signature = types + self.unreachable() + elif op == RETURN: + self.effect(self.controls[0].results, []) + self.emit(RETURN) + self.unreachable() + elif op in (0x10, 0x11): + table = 0 + if op == 0x10: + index = self.module._index(self.reader, self.module.functions, 'function') + params, returns = self.module.signature(index) + else: + index = self.module._type_index(self.reader) + table = self.table_index() + if self.module.tables[table][0] != FUNCREF: + raise ValidationError('call_indirect requires funcref table') + self.pop(I32) + params, returns = self.module.types[index] + self.effect(params, returns) + self.emit(op, index, table) + elif op == 0x00: + self.emit(op) + self.unreachable() + elif op == 0x01: + pass + elif op == 0x1A: + self.pop() + self.emit(op) + elif op in (0x1B, 0x1C): + select_types = self.reader.vec(lambda: value_type(self.reader)) if op == 0x1C else None + if select_types is not None and len(select_types) != 1: + raise ValidationError('typed select requires one type') + self.pop(I32) + rhs = self.pop(select_types[0] if select_types else None) + lhs = self.pop(rhs) + if select_types is None and lhs in (FUNCREF, EXTERNREF): + raise ValidationError('reference select requires explicit type') + self.stack.append(lhs if lhs is not None else rhs) + self.emit(0x1B) + elif 0x20 <= op <= 0x24: + global_op = op >= 0x23 + index = self.module._index(self.reader, self.module.globals if global_op else self.fn.locals, 'variable') + ty = self.module.globals[index][0] if global_op else self.fn.locals[index] + if op == 0x24 and not self.module.globals[index][1]: + raise ValidationError('cannot set immutable global') + self.effect([] if op in (0x20, 0x23) else [ty], [ty] if op in (0x20, 0x22, 0x23) else []) + self.emit(op, index) + elif op in (0x25, 0x26): + table = self.table_index() + ty = self.module.tables[table][0] + self.effect([I32] if op == 0x25 else [I32, ty], [ty] if op == 0x25 else []) + self.emit(op, table) + elif 0x28 <= op <= 0x3E: + self.require_memory() + align, offset = self.reader.leb(), self.reader.leb() + widths = [4, 8, 4, 8, 1, 1, 2, 2, 1, 1, 2, 2, 4, 4, 4, 8, 4, 8, 1, 2, 1, 2, 4] + width = widths[op - 0x28] + if (1 << min(align, 32)) > width: + raise ValidationError('memory alignment exceeds natural alignment') + if op <= 0x35: + ty = I32 if op in (0x28, 0x2C, 0x2D, 0x2E, 0x2F) else F32 if op == 0x2A else F64 if op == 0x2B else I64 + self.effect([I32], [ty]) + else: + ty = I32 if op in (0x36, 0x3A, 0x3B) else F32 if op == 0x38 else F64 if op == 0x39 else I64 + self.effect([I32, ty], []) + self.emit(op, offset, width) + elif op in (0x3F, 0x40): + self.memory_index() + self.effect([I32] if op == 0x40 else [], [I32]) + self.emit(op) + elif 0x41 <= op <= 0x44: + ty = (I32, I64, F32, F64)[op - 0x41] + if op <= 0x42: + val = self.reader.leb(32 if op == 0x41 else 64, signed=True) & ((1 << 64) - 1) + else: + val = int.from_bytes(self.reader.take(4 if op == 0x43 else 8), 'little') + self.effect([], [ty]) + self.emit(op, val & 0xFFFFFFFF, (val >> 32) if ty in (I64, F64) else 0) + elif op == 0xD0: + ty = value_type(self.reader) + if ty not in (FUNCREF, EXTERNREF): + raise ValidationError('invalid reference type') + self.effect([], [ty]) + self.emit(0x41) + elif op == 0xD1: + reference_type = self.pop() + if reference_type not in (None, FUNCREF, EXTERNREF): + raise ValidationError('ref.is_null requires reference') + self.effect([], [I32]) + self.emit(0x45) + elif op == 0xD2: + index = self.module._index(self.reader, self.module.functions, 'function') + if index not in self.module.declared_functions: + raise ValidationError('ref.func requires a declared function reference') + self.effect([], [FUNCREF]) + self.emit(0x41, index + 1) + elif op == 0xFC: + sub = self.reader.leb() + if sub <= 7: + conversion = numeric_signature(0xFC00 + sub) + assert conversion is not None + self.effect(*conversion) + self.emit(0xFC00 + sub, 1) + else: + self.bulk(sub) + else: + signature = numeric_signature(op) + if signature is None: + raise UnsupportedFeatureError(f'unsupported opcode 0x{op:02x}') + self.effect(*signature) + self.emit(op, len(signature[0])) + + def bulk(self, sub: int) -> None: + a, b = 0, 0 + if sub in (8, 9): + if self.module.data_count is None: + raise ValidationError('memory.init/data.drop requires data count section') + a = self.module._index(self.reader, self.module.data, 'data segment') + if sub == 8: + self.memory_index() + self.effect([I32, I32, I32], []) + elif sub in (10, 11): + self.memory_index() + if sub == 10: + self.memory_index() + self.effect([I32, I32, I32], []) + elif sub in (12, 13): + a = self.module._index(self.reader, self.module.elements, 'element segment') + if sub == 12: + b = self.table_index() + if self.module.elements[a][3] != self.module.tables[b][0]: + raise ValidationError('table.init element type mismatch') + self.effect([I32, I32, I32], []) + elif sub in (14, 15, 16, 17): + a = self.table_index() + if sub == 14: + b = self.table_index() + if self.module.tables[a][0] != self.module.tables[b][0]: + raise ValidationError('table.copy element type mismatch') + ty = self.module.tables[a][0] + args = {14: [I32, I32, I32], 15: [ty, I32], 16: [], 17: [I32, ty, I32]}[sub] + self.effect(args, [I32] if sub in (15, 16) else []) + else: + raise UnsupportedFeatureError(f'unsupported 0xfc opcode {sub}') + self.emit(0xFC00 + sub, a, b) diff --git a/wasmgpu/errors.py b/wasmgpu/errors.py new file mode 100644 index 0000000..27c4b19 --- /dev/null +++ b/wasmgpu/errors.py @@ -0,0 +1,35 @@ +"""Public errors. GPU traps never cause execution to fall back to the CPU.""" + + +from __future__ import annotations + +from .types import Result + + +class WasmGPUError(Exception): + """Base class for runtime errors.""" + + +class ValidationError(WasmGPUError, ValueError): + """Malformed or ill-typed WebAssembly.""" + + +class UnsupportedFeatureError(WasmGPUError, NotImplementedError): + """A valid feature outside this runtime's supported profile.""" + + +class GPUUnavailableError(WasmGPUError): + """No hardware GPU is available.""" + + +class ResourceLimitError(WasmGPUError): + """An explicit runtime or device resource limit was reached.""" + + +class Trap(WasmGPUError): # noqa: N818 - WebAssembly calls these traps. + """One or more instances trapped; successful peers are in ``results``.""" + + def __init__(self, traps: dict[int, str], results: list[Result]) -> None: + self.traps = traps + self.results = results + super().__init__('; '.join(f'instance {index}: {reason}' for index, reason in traps.items())) diff --git a/wasmgpu/filesystem.wgsl b/wasmgpu/filesystem.wgsl new file mode 100644 index 0000000..6571b4f --- /dev/null +++ b/wasmgpu/filesystem.wgsl @@ -0,0 +1,687 @@ +// Per-instance filesystem and WASI Preview 1 services, executed on the GPU. +// Names, descriptors, inode metadata and file contents are all in heap storage. +var path_raw: array; +var path_tail: array; +var path_text: array; +var path_length: u32; +fn arg(base: u32, index: u32) -> u32 { return get_value(base + index).x; } +fn fs_entry(index: u32) -> u32 { return config.fs_offset + 48u + index * 12u; } +fn fs_names() -> u32 { return config.fs_offset + 48u + config.fs_files * 12u; } +fn fs_descriptor(fd: u32) -> u32 { return fs_names() + config.fs_files * 64u + fd * 8u; } +fn fs_data() -> u32 { return fs_descriptor(config.fs_fds); } +fn heap_byte(base: u32, index: u32) -> u32 { return (read_heap(base + index / 4u) >> ((index & 3u) * 8u)) & 255u; } +fn set_heap_byte(base: u32, index: u32, value: u32) { + let address = base + index / 4u; let shift = (index & 3u) * 8u; + write_heap(address, (read_heap(address) & ~(255u << shift)) | ((value & 255u) << shift)); +} +fn load32(address: u32) -> u32 { + return read_byte(address) | (read_byte(address + 1u) << 8u) | (read_byte(address + 2u) << 16u) | (read_byte(address + 3u) << 24u); +} +fn store32(address: u32, value: u32) { for (var i = 0u; i < 4u; i += 1u) { write_byte(address + i, value >> (8u * i)); } } +fn store64(address: u32, value: vec2u) { store32(address, value.x); store32(address + 4u, value.y); } +fn fs_file(fd: u32) -> u32 { + if fd >= config.fs_fds { return 0xffffffffu; } + return read_heap(fs_descriptor(fd)) - 1u; +} +fn fs_inode(index: u32) -> u32 { return fs_entry(read_heap(fs_entry(index) + 5u)); } +fn fs_right(fd: u32, right: u32) -> bool { return (read_heap(fs_descriptor(fd) + 4u) & right) == right; } +fn fs_charge(amount: u32) -> bool { + if amount > vm.fuel { fail(8u); return false; } + vm.fuel -= amount; return true; +} +fn fs_open_entry(index: u32) -> bool { + for (var fd = 0u; fd < config.fs_fds; fd += 1u) { if fs_file(fd) == index { return true; } } + return false; +} +fn fs_live_inode(index: u32) -> bool { + for (var i = 0u; i < config.fs_files; i += 1u) { + let entry = fs_entry(i); let kind = read_heap(entry); + if kind == 0u || read_heap(entry + 5u) != index { continue; } + if kind < 4u || fs_open_entry(i) { return true; } + } + return false; +} +fn fs_collect() { + for (var i = 4u; i < config.fs_files; i += 1u) { + if read_heap(fs_entry(i)) != 4u || fs_open_entry(i) { continue; } + if read_heap(fs_entry(i) + 5u) == i && fs_live_inode(i) { continue; } + write_heap(fs_entry(i), 0u); write_heap(fs_entry(i) + 4u, 0u); + } +} +fn fs_compact() -> bool { + fs_collect(); + var scan = 0u; var packed = 0u; + for (var step = 0u; step < config.fs_files; step += 1u) { + var chosen = 0xffffffffu; var position = 0xffffffffu; + for (var i = 0u; i < config.fs_files; i += 1u) { + let entry = fs_entry(i); let start = read_heap(entry + 3u); + if read_heap(entry) == 0u || read_heap(entry + 5u) != i || read_heap(entry + 4u) == 0u { continue; } + if start >= scan && start < position { position = start; chosen = i; } + } + if chosen == 0xffffffffu { break; } + let entry = fs_entry(chosen); let length = read_heap(entry + 2u); + if !fs_charge(length) { return false; } + for (var i = 0u; i < length; i += 1u) { set_heap_byte(fs_data(), packed + i, heap_byte(fs_data(), position + i)); } + write_heap(entry + 3u, packed); write_heap(entry + 4u, length); + scan = position + 1u; packed += length; + } + write_heap(config.fs_offset, packed); return true; +} +fn fs_resize(index: u32, length: u32) -> u32 { + let entry = fs_inode(index); let old_size = read_heap(entry + 2u); + var start = read_heap(entry + 3u); var capacity = read_heap(entry + 4u); + if length > config.fs_bytes { return 51u; } + if length > capacity { + var cursor = read_heap(config.fs_offset); + if length - capacity > config.fs_bytes - cursor { + if !fs_compact() { return 29u; } + start = read_heap(entry + 3u); capacity = read_heap(entry + 4u); cursor = read_heap(config.fs_offset); + if length - capacity > config.fs_bytes - cursor { return 51u; } + } + let wanted = min(capacity + config.fs_bytes - cursor, max(length, min(config.fs_bytes, max(64u, capacity * 2u)))); + if capacity == 0u { start = cursor; } + let tail = start + capacity; let extra = wanted - capacity; + if !fs_charge(cursor - tail) { return 29u; } + for (var i = cursor; i > tail; i -= 1u) { set_heap_byte(fs_data(), i - 1u + extra, heap_byte(fs_data(), i - 1u)); } + for (var i = 0u; i < config.fs_files; i += 1u) { + let other = fs_entry(i); + if other == entry || read_heap(other) == 0u || read_heap(other + 5u) != i || read_heap(other + 4u) == 0u { continue; } + if read_heap(other + 3u) >= tail { write_heap(other + 3u, read_heap(other + 3u) + extra); } + } + write_heap(entry + 3u, start); write_heap(entry + 4u, wanted); + write_heap(config.fs_offset, cursor + extra); + } + if length > old_size { + if !fs_charge(length - old_size) { return 29u; } + for (var i = old_size; i < length; i += 1u) { set_heap_byte(fs_data(), start + i, 0u); } + } + write_heap(entry + 2u, length); + let now = vec2u(read_heap(config.fs_offset + 1u), read_heap(config.fs_offset + 2u)); + write_heap(entry + 8u, now.x); write_heap(entry + 9u, now.y); + write_heap(entry + 10u, now.x); write_heap(entry + 11u, now.y); + return 0u; +} +fn fs_resolve(fd: u32, address: u32, size: u32, follow_final: bool) -> u32 { + let directory = fs_file(fd); + if directory == 0xffffffffu { return 8u; } + if read_heap(fs_entry(directory)) != 2u { return 54u; } + if !bounds(address, size) { return 21u; } + if size == 0u { return 44u; } + if size > 511u { return 37u; } + if read_byte(address) == 47u { return 76u; } + let floor = read_heap(fs_entry(directory) + 1u); + path_length = floor; + for (var j = 0u; j < floor; j += 1u) { path_text[j] = heap_byte(fs_names() + directory * 64u, j); } + var raw_length = size; + for (var j = 0u; j < size; j += 1u) { + let byte = read_byte(address + j); + if byte == 0u { return 28u; } + path_raw[j] = byte; + } + var i = 0u; var links = 0u; + while i < raw_length { + if path_raw[i] == 47u { i += 1u; continue; } + let begin = i; + while i < raw_length && path_raw[i] != 47u { i += 1u; } + let length = i - begin; + if length == 1u && path_raw[begin] == 46u { continue; } + if length == 2u && path_raw[begin] == 46u && path_raw[begin + 1u] == 46u { + if path_length <= floor { return 76u; } + while path_length > floor && path_text[path_length - 1u] != 47u { path_length -= 1u; } + if path_length > floor { path_length -= 1u; } + continue; + } + let parent_length = path_length; + if path_length + length + u32(path_length != 0u) > 255u { return 37u; } + if path_length != 0u { path_text[path_length] = 47u; path_length += 1u; } + for (var j = 0u; j < length; j += 1u) { path_text[path_length] = path_raw[begin + j]; path_length += 1u; } + let index = fs_lookup(); + let intermediate = i < raw_length; + if index == 0xffffffffu { if intermediate { return 44u; } continue; } + let kind = read_heap(fs_entry(index)); + if kind == 3u && (intermediate || follow_final) { + links += 1u; if links > 40u { return 32u; } + let entry = fs_inode(index); let link_size = read_heap(entry + 2u); let start = read_heap(entry + 3u); + if link_size == 0u { return 44u; } + if heap_byte(fs_data(), start) == 47u { return 76u; } + let tail_size = raw_length - i; + if link_size + tail_size > 511u { return 37u; } + for (var j = 0u; j < tail_size; j += 1u) { path_tail[j] = path_raw[i + j]; } + for (var j = 0u; j < link_size; j += 1u) { path_raw[j] = heap_byte(fs_data(), start + j); } + for (var j = 0u; j < tail_size; j += 1u) { path_raw[link_size + j] = path_tail[j]; } + raw_length = link_size + tail_size; i = 0u; path_length = parent_length; + } else if intermediate && kind != 2u { return 54u; } + } + return 0u; +} +fn fs_path(fd: u32, address: u32, size: u32) -> u32 { return fs_resolve(fd, address, size, false); } +fn fs_lookup() -> u32 { + if path_length == 0u { return 3u; } + for (var index = 4u; index < config.fs_files; index += 1u) { + let entry = fs_entry(index); let kind = read_heap(entry); + if kind == 0u || kind >= 4u || read_heap(entry + 1u) != path_length { continue; } + var equal = true; + for (var i = 0u; i < path_length; i += 1u) { if heap_byte(fs_names() + index * 64u, i) != path_text[i] { equal = false; break; } } + if equal { return index; } + } + return 0xffffffffu; +} +fn fs_parent_exists() -> bool { + let original = path_length; + while path_length > 0u && path_text[path_length - 1u] != 47u { path_length -= 1u; } + if path_length > 0u { path_length -= 1u; } + let parent = fs_lookup(); path_length = original; + return parent != 0xffffffffu && read_heap(fs_entry(parent)) == 2u; +} +fn fs_save_name(index: u32) { + write_heap(fs_entry(index) + 1u, path_length); + for (var i = 0u; i < path_length; i += 1u) { set_heap_byte(fs_names() + index * 64u, i, path_text[i]); } +} +fn fs_create(kind: u32) -> u32 { + fs_collect(); + for (var index = 4u; index < config.fs_files; index += 1u) { + if read_heap(fs_entry(index)) == 0u { + let entry = fs_entry(index); + for (var j = 0u; j < 12u; j += 1u) { write_heap(entry + j, 0u); } + write_heap(entry, kind); write_heap(entry + 5u, index); fs_save_name(index); + return index; + } + } + return 0xffffffffu; +} +fn fs_child(index: u32, parent: u32) -> bool { + let kind = read_heap(fs_entry(index)); let length = read_heap(fs_entry(parent) + 1u); + if kind == 0u || kind >= 4u || read_heap(fs_entry(index) + 1u) <= length { return false; } + if heap_byte(fs_names() + index * 64u, length) != 47u { return false; } + for (var i = 0u; i < length; i += 1u) { + if heap_byte(fs_names() + index * 64u, i) != heap_byte(fs_names() + parent * 64u, i) { return false; } + } + return true; +} +fn fs_rename(index: u32, existing: u32) -> u32 { + if existing == index { return 0u; } + if index == 3u || existing == 3u { return 10u; } + let directory = read_heap(fs_entry(index)) == 2u; + let old_length = read_heap(fs_entry(index) + 1u); + if directory && path_length > old_length && path_text[old_length] == 47u { + var descendant = true; + for (var i = 0u; i < old_length; i += 1u) { descendant = descendant && path_text[i] == heap_byte(fs_names() + index * 64u, i); } + if descendant { return 28u; } + } + if existing != 0xffffffffu { + let existing_directory = read_heap(fs_entry(existing)) == 2u; + if directory && !existing_directory { return 54u; } + if !directory && existing_directory { return 31u; } + if fs_inode(existing) == fs_inode(index) { return 0u; } + if existing_directory { + for (var i = 4u; i < config.fs_files; i += 1u) { if fs_child(i, existing) { return 55u; } } + } + } + if directory { + for (var i = 4u; i < config.fs_files; i += 1u) { + if fs_child(i, index) && read_heap(fs_entry(i) + 1u) - old_length + path_length > 255u { return 37u; } + } + for (var i = 4u; i < config.fs_files; i += 1u) { + if !fs_child(i, index) { continue; } + let suffix = read_heap(fs_entry(i) + 1u) - old_length; + for (var j = 0u; j < suffix; j += 1u) { path_tail[j] = heap_byte(fs_names() + i * 64u, old_length + j); } + for (var j = 0u; j < path_length; j += 1u) { set_heap_byte(fs_names() + i * 64u, j, path_text[j]); } + for (var j = 0u; j < suffix; j += 1u) { set_heap_byte(fs_names() + i * 64u, path_length + j, path_tail[j]); } + write_heap(fs_entry(i) + 1u, path_length + suffix); + } + } + if existing != 0xffffffffu { write_heap(fs_entry(existing), 4u); } + fs_save_name(index); return 0u; +} +fn fs_stat(index: u32, address: u32) { + let entry = fs_inode(index); let kind = read_heap(fs_entry(index)); + for (var i = 0u; i < 64u; i += 1u) { write_byte(address + i, 0u); } + store64(address, vec2u(1u, 0u)); store64(address + 8u, vec2u(read_heap(fs_entry(index) + 5u) + 1u, 0u)); + write_byte(address + 16u, select(select(select(4u, 3u, kind == 2u), 7u, kind == 3u), 2u, index < 3u)); + var links = 0u; + for (var i = 3u; i < config.fs_files; i += 1u) { if read_heap(fs_entry(i)) > 0u && read_heap(fs_entry(i)) < 4u && fs_inode(i) == entry { links += 1u; } } + store64(address + 24u, vec2u(select(links, 1u, index < 3u), 0u)); store64(address + 32u, vec2u(read_heap(entry + 2u), 0u)); + for (var i = 0u; i < 6u; i += 1u) { store32(address + 40u + i * 4u, read_heap(entry + 6u + i)); } +} +fn fs_io(syscall: u32, base: u32) -> u32 { + let fd = arg(base, 0u); let index = fs_file(fd); + if index == 0xffffffffu { return 8u; } + if read_heap(fs_entry(index)) == 2u { return 31u; } + let writing = syscall == WASI_FD_WRITE || syscall == WASI_FD_PWRITE; + if !fs_right(fd, select(2u, 64u, writing)) { return 76u; } + let positioned = syscall == WASI_FD_PREAD || syscall == WASI_FD_PWRITE; + let vectors = arg(base, 1u); let count = arg(base, 2u); let result_ptr = arg(base, select(3u, 4u, positioned)); + if count > 0x1fffffffu || !bounds(vectors, count * 8u) || !bounds(result_ptr, 4u) { return 21u; } + var total = 0u; + for (var i = 0u; i < count; i += 1u) { + let pointer = load32(vectors + i * 8u); let length = load32(vectors + i * 8u + 4u); + if !bounds(pointer, length) { return 21u; } + if total + length < total { return 28u; } + total += length; + } + if !fs_charge(total) { return 29u; } + let entry = fs_inode(index); let descriptor = fs_descriptor(fd); + var position = vec2u(read_heap(descriptor + 1u), read_heap(descriptor + 2u)); + if positioned { position = get_value(base + 3u); } + if writing && (read_heap(descriptor + 3u) & 1u) != 0u { position = vec2u(read_heap(entry + 2u), 0u); } + if position.y != 0u { if writing { return 27u; } store32(result_ptr, 0u); return 0u; } + if writing && total != 0u { + if position.x > config.fs_bytes || total > config.fs_bytes - position.x { return 51u; } + let error = fs_resize(index, max(read_heap(entry + 2u), position.x + total)); + if error != 0u { return error; } + } + var transferred = 0u; + let start = read_heap(entry + 3u); let file_length = read_heap(entry + 2u); + for (var i = 0u; i < count; i += 1u) { + let pointer = load32(vectors + i * 8u); var size = load32(vectors + i * 8u + 4u); + if !bounds(pointer, size) { return 21u; } + if !writing { size = min(size, file_length - min(position.x, file_length)); } + for (var j = 0u; j < size; j += 1u) { + if writing { set_heap_byte(fs_data(), start + position.x + j, read_byte(pointer + j)); } + else { write_byte(pointer + j, heap_byte(fs_data(), start + position.x + j)); } + } + position.x += size; transferred += size; + } + if !positioned { write_heap(descriptor + 1u, position.x); write_heap(descriptor + 2u, position.y); } + store32(result_ptr, transferred); return 0u; +} +fn fs_poll(base: u32) -> u32 { + let input_ptr = arg(base, 0u); let output_ptr = arg(base, 1u); let count = arg(base, 2u); let result_ptr = arg(base, 3u); + if count == 0u || count > 0x05555555u { return 28u; } + if !bounds(input_ptr, count * 48u) || !bounds(output_ptr, count * 32u) || !bounds(result_ptr, 4u) { return 21u; } + if !fs_charge(count) { return 29u; } + let now = vec2u(read_heap(config.fs_offset + 1u), read_heap(config.fs_offset + 2u)); + var earliest = vec2u(0xffffffffu); var immediate = false; + for (var i = 0u; i < count; i += 1u) { + let subscription = input_ptr + i * 48u; let kind = read_byte(subscription + 8u); + if kind > 2u { return 28u; } + if kind != 0u { immediate = true; continue; } + let clock = load32(subscription + 16u); let flags = read_byte(subscription + 40u) | (read_byte(subscription + 41u) << 8u); + if clock > 3u || flags > 1u { immediate = true; continue; } + var deadline = vec2u(load32(subscription + 24u), load32(subscription + 28u)); + if flags == 0u { deadline = add64(now, deadline); } + if lt64(deadline, earliest) { earliest = deadline; } + } + let awake = select(earliest, now, immediate || lt64(earliest, now)); + write_heap(config.fs_offset + 1u, awake.x); write_heap(config.fs_offset + 2u, awake.y); + var events = 0u; + for (var i = 0u; i < count; i += 1u) { + let subscription = input_ptr + i * 48u; let kind = read_byte(subscription + 8u); + var error = 0u; var available = 0u; + if kind == 0u { + let flags = read_byte(subscription + 40u) | (read_byte(subscription + 41u) << 8u); + if load32(subscription + 16u) > 3u || flags > 1u { error = 28u; } + else { + var deadline = vec2u(load32(subscription + 24u), load32(subscription + 28u)); + if flags == 0u { deadline = add64(now, deadline); } + if lt64(awake, deadline) { continue; } + } + } else { + let fd = load32(subscription + 16u); let index = fs_file(fd); + if index == 0xffffffffu { error = 8u; } + else if !fs_right(fd, 1u << 27u) { error = 76u; } + else { + let size = read_heap(fs_inode(index) + 2u); let position = read_heap(fs_descriptor(fd) + 1u); + available = select(size - min(size, position), config.fs_bytes - min(config.fs_bytes, position), kind == 2u); + } + } + let event = output_ptr + events * 32u; + for (var j = 0u; j < 32u; j += 1u) { write_byte(event + j, 0u); } + store64(event, vec2u(load32(subscription), load32(subscription + 4u))); + write_byte(event + 8u, error); write_byte(event + 10u, kind); store64(event + 16u, vec2u(available, 0u)); + events += 1u; + } + store32(result_ptr, events); return 0u; +} +fn fs_set_times(index: u32, atime: vec2u, mtime: vec2u, flags: u32) -> u32 { + if flags > 15u || (flags & 3u) == 3u || (flags & 12u) == 12u { return 28u; } + let entry = fs_inode(index); + let now = vec2u(read_heap(config.fs_offset + 1u), read_heap(config.fs_offset + 2u)); + if (flags & 3u) != 0u { + let value = select(atime, now, (flags & 2u) != 0u); + write_heap(entry + 6u, value.x); write_heap(entry + 7u, value.y); + } + if (flags & 12u) != 0u { + let value = select(mtime, now, (flags & 8u) != 0u); + write_heap(entry + 8u, value.x); write_heap(entry + 9u, value.y); + } + write_heap(entry + 10u, now.x); write_heap(entry + 11u, now.y); return 0u; +} +fn fs_readdir(base: u32) -> u32 { + let fd = arg(base, 0u); let directory = fs_file(fd); + if directory == 0xffffffffu { return 8u; } + if read_heap(fs_entry(directory)) != 2u { return 54u; } + if !fs_right(fd, 1u << 14u) { return 76u; } + let pointer = arg(base, 1u); let size = arg(base, 2u); let cookie = get_value(base + 3u); let result_ptr = arg(base, 4u); + if !bounds(pointer, size) || !bounds(result_ptr, 4u) { return 21u; } + if !fs_charge(size) { return 29u; } + let prefix = read_heap(fs_entry(directory) + 1u); var written = 0u; + if cookie.y == 0u { + for (var index = max(4u, cookie.x); index < config.fs_files; index += 1u) { + let entry = fs_entry(index); let kind = read_heap(entry); let length = read_heap(entry + 1u); + if kind == 0u || kind >= 4u || length <= prefix { continue; } + let name_start = fs_names() + index * 64u; + var child = true; + for (var j = 0u; j < prefix; j += 1u) { child = child && heap_byte(name_start, j) == heap_byte(fs_names() + directory * 64u, j); } + if prefix != 0u { child = child && heap_byte(name_start, prefix) == 47u; } + let basename = prefix + u32(prefix != 0u); + for (var j = basename; j < length; j += 1u) { if heap_byte(name_start, j) == 47u { child = false; } } + if !child { continue; } + let name_length = length - basename; + for (var j = 0u; j < 24u + name_length && written < size; j += 1u) { + var value = 0u; + if j < 4u { value = (index + 1u) >> (j * 8u); } + else if j >= 8u && j < 12u { value = (read_heap(entry + 5u) + 1u) >> ((j - 8u) * 8u); } + else if j >= 16u && j < 20u { value = name_length >> ((j - 16u) * 8u); } + else if j == 20u { value = select(select(4u, 3u, kind == 2u), 7u, kind == 3u); } + else if j >= 24u { value = heap_byte(name_start, basename + j - 24u); } + write_byte(pointer + written, value); written += 1u; + } + if written == size { break; } + } + } + store32(result_ptr, written); return 0u; +} +// ChaCha20 (RFC 8439), deterministic keyed streams with an instance nonce. +// No entropy or clock is requested from the operating system. +fn rng_rotate(x: u32, n: u32) -> u32 { return (x << n) | (x >> (32u - n)); } +fn rng_quarter(v: vec4u) -> vec4u { + var a = v.x; var b = v.y; var c = v.z; var d = v.w; + a += b; d = rng_rotate(d ^ a, 16u); c += d; b = rng_rotate(b ^ c, 12u); + a += b; d = rng_rotate(d ^ a, 8u); c += d; b = rng_rotate(b ^ c, 7u); + return vec4u(a, b, c, d); +} +fn rng_block() { + var original: array; + original[0] = 0x61707865u; original[1] = 0x3320646eu; original[2] = 0x79622d32u; original[3] = 0x6b206574u; + for (var i = 0u; i < 8u; i += 1u) { original[4u + i] = read_heap(config.fs_offset + 8u + i); } + original[12] = read_heap(config.fs_offset + 16u); + for (var i = 0u; i < 3u; i += 1u) { original[13u + i] = read_heap(config.fs_offset + 18u + i); } + var x = original; + for (var round = 0u; round < 10u; round += 1u) { + for (var i = 0u; i < 4u; i += 1u) { + let q = rng_quarter(vec4u(x[i], x[i + 4u], x[i + 8u], x[i + 12u])); + x[i] = q.x; x[i + 4u] = q.y; x[i + 8u] = q.z; x[i + 12u] = q.w; + } + for (var i = 0u; i < 4u; i += 1u) { + let b = 4u + (i + 1u) % 4u; let c = 8u + (i + 2u) % 4u; let d = 12u + (i + 3u) % 4u; + let q = rng_quarter(vec4u(x[i], x[b], x[c], x[d])); x[i] = q.x; x[b] = q.y; x[c] = q.z; x[d] = q.w; + } + } + for (var i = 0u; i < 16u; i += 1u) { write_heap(config.fs_offset + 21u + i, x[i] + original[i]); } + let counter = read_heap(config.fs_offset + 16u) + 1u; + write_heap(config.fs_offset + 16u, counter); + if counter == 0u { write_heap(config.fs_offset + 19u, read_heap(config.fs_offset + 19u) + 1u); } + write_heap(config.fs_offset + 37u, 0u); +} +fn wasi_dispatch(syscall: u32, base: u32) -> u32 { + if syscall == WASI_PROC_EXIT { vm.exit_code = arg(base, 0u); fail(12u); return 0u; } + if syscall == WASI_PROC_RAISE { + if arg(base, 0u) > 30u { return 28u; } + if arg(base, 0u) != 0u { vm.exit_code = arg(base, 0u); fail(13u); } + return 0u; + } + if syscall == WASI_SCHED_YIELD { return 0u; } + if syscall == WASI_POLL_ONEOFF { return fs_poll(base); } + if syscall == WASI_FD_READDIR { return fs_readdir(base); } + if syscall == WASI_CLOCK_TIME_GET || syscall == WASI_CLOCK_RES_GET { + if arg(base, 0u) > 3u { return 28u; } + let pointer = arg(base, select(2u, 1u, syscall == WASI_CLOCK_RES_GET)); + if !bounds(pointer, 8u) { return 21u; } + var clock = vec2u(read_heap(config.fs_offset + 1u), read_heap(config.fs_offset + 2u)); + if syscall == WASI_CLOCK_RES_GET { clock = vec2u(read_heap(config.fs_offset + 4u), 0u); } + store64(pointer, clock); return 0u; + } + if syscall == WASI_RANDOM_GET { + let pointer = arg(base, 0u); let size = arg(base, 1u); + if !bounds(pointer, size) { return 21u; } + if !fs_charge(size) { return 29u; } + var cursor = read_heap(config.fs_offset + 37u); + for (var i = 0u; i < size; i += 1u) { + if cursor == 64u { rng_block(); cursor = 0u; } + write_byte(pointer + i, heap_byte(config.fs_offset + 21u, cursor)); cursor += 1u; + } + write_heap(config.fs_offset + 37u, cursor); return 0u; + } + if syscall == WASI_FD_WRITE || syscall == WASI_FD_READ || syscall == WASI_FD_PWRITE || syscall == WASI_FD_PREAD { return fs_io(syscall, base); } + if syscall == WASI_ARGS_GET || syscall == WASI_ARGS_SIZES_GET || syscall == WASI_ENVIRON_GET || syscall == WASI_ENVIRON_SIZES_GET { + let environment = syscall == WASI_ENVIRON_GET || syscall == WASI_ENVIRON_SIZES_GET; + let sizes = syscall == WASI_ARGS_SIZES_GET || syscall == WASI_ENVIRON_SIZES_GET; + let slot = select(8u, 10u, environment); let info = program[slot]; let count = program[slot + 1u]; + let a = arg(base, 0u); let b = arg(base, 1u); + var total = 0u; + for (var i = 0u; i < count; i += 1u) { total += program[info + i * 2u + 1u]; } + if sizes { + if !bounds(a, 4u) || !bounds(b, 4u) { return 21u; } + store32(a, count); store32(b, total); + } else { + if !bounds(a, count * 4u) || !bounds(b, total) { return 21u; } + if !fs_charge(total) { return 29u; } + var cursor = b; + for (var i = 0u; i < count; i += 1u) { + store32(a + i * 4u, cursor); + let start = program[info + i * 2u]; let length = program[info + i * 2u + 1u]; + for (var j = 0u; j < length; j += 1u) { write_byte(cursor + j, (program[start + j / 4u] >> ((j & 3u) * 8u)) & 255u); } + cursor += length; + } + } + return 0u; + } + if syscall == WASI_PATH_OPEN { + let fd = arg(base, 0u); let result_ptr = arg(base, 8u); + if !bounds(result_ptr, 4u) { return 21u; } + let error = fs_resolve(fd, arg(base, 2u), arg(base, 3u), arg(base, 1u) == 1u); if error != 0u { return error; } + if !fs_right(fd, 1u << 13u) { return 76u; } + let oflags = arg(base, 4u); let flags = arg(base, 7u); + if oflags > 15u || flags > 31u || arg(base, 1u) > 1u { return 28u; } + let rights = get_value(base + 5u); let inheriting = get_value(base + 6u); + let allowed = vec2u(read_heap(fs_descriptor(fd) + 6u), read_heap(fs_descriptor(fd) + 7u)); + if any((rights & ~allowed) != vec2u(0u)) || any((inheriting & ~allowed) != vec2u(0u)) { return 76u; } + var new_fd = 0xffffffffu; + for (var i = 0u; i < config.fs_fds; i += 1u) { if read_heap(fs_descriptor(i)) == 0u { new_fd = i; break; } } + if new_fd == 0xffffffffu { return 33u; } + var index = fs_lookup(); + if index == 0xffffffffu { + if (oflags & 1u) == 0u { return 44u; } + if (oflags & 2u) != 0u { return 54u; } + if !fs_right(fd, 1u << 10u) { return 76u; } + if !fs_parent_exists() { return 44u; } + index = fs_create(1u); if index == 0xffffffffu { return 51u; } + } else if (oflags & 5u) == 5u { return 20u; } + let kind = read_heap(fs_entry(index)); + if (oflags & 2u) != 0u && kind != 2u { return 54u; } + if kind == 3u { return 32u; } + if (oflags & 8u) != 0u { + if (rights.x & 64u) == 0u { return 76u; } + if kind == 2u { return 31u; } + let resize_error = fs_resize(index, 0u); if resize_error != 0u { return resize_error; } + } + let descriptor = fs_descriptor(new_fd); + write_heap(descriptor, index + 1u); write_heap(descriptor + 1u, 0u); write_heap(descriptor + 2u, 0u); + write_heap(descriptor + 3u, flags); write_heap(descriptor + 4u, rights.x); write_heap(descriptor + 5u, rights.y); + write_heap(descriptor + 6u, inheriting.x); write_heap(descriptor + 7u, inheriting.y); + store32(result_ptr, new_fd); return 0u; + } + if syscall == WASI_PATH_CREATE_DIRECTORY || syscall == WASI_PATH_UNLINK_FILE || syscall == WASI_PATH_REMOVE_DIRECTORY { + let fd = arg(base, 0u); let error = fs_path(fd, arg(base, 1u), arg(base, 2u)); if error != 0u { return error; } + let right = select(select(1u << 25u, 1u << 26u, syscall == WASI_PATH_UNLINK_FILE), 1u << 9u, syscall == WASI_PATH_CREATE_DIRECTORY); + if !fs_right(fd, right) { return 76u; } + let index = fs_lookup(); + if syscall == WASI_PATH_CREATE_DIRECTORY { + if index != 0xffffffffu { return 20u; } + if !fs_parent_exists() { return 44u; } + return select(0u, 51u, fs_create(2u) == 0xffffffffu); + } + if index == 0xffffffffu { return 44u; } + let kind = read_heap(fs_entry(index)); + if syscall == WASI_PATH_UNLINK_FILE && kind == 2u { return 31u; } + if syscall == WASI_PATH_REMOVE_DIRECTORY { + if kind != 2u { return 54u; } + if index == 3u { return 10u; } + for (var i = 4u; i < config.fs_files; i += 1u) { + if read_heap(fs_entry(i)) == 0u || read_heap(fs_entry(i)) >= 4u || read_heap(fs_entry(i) + 1u) <= path_length { continue; } + var child = heap_byte(fs_names() + i * 64u, path_length) == 47u; + for (var j = 0u; j < path_length; j += 1u) { child = child && heap_byte(fs_names() + i * 64u, j) == path_text[j]; } + if child { return 55u; } + } + } + write_heap(fs_entry(index), 4u); return 0u; + } + if syscall == WASI_PATH_FILESTAT_GET { + if !bounds(arg(base, 4u), 64u) { return 21u; } + if arg(base, 1u) > 1u { return 28u; } + let error = fs_resolve(arg(base, 0u), arg(base, 2u), arg(base, 3u), arg(base, 1u) == 1u); if error != 0u { return error; } + if !fs_right(arg(base, 0u), 1u << 18u) { return 76u; } + let index = fs_lookup(); if index == 0xffffffffu { return 44u; } + fs_stat(index, arg(base, 4u)); return 0u; + } + if syscall == WASI_PATH_FILESTAT_SET_TIMES { + if arg(base, 1u) > 1u { return 28u; } + let fd = arg(base, 0u); let error = fs_resolve(fd, arg(base, 2u), arg(base, 3u), arg(base, 1u) == 1u); if error != 0u { return error; } + if !fs_right(fd, 1u << 20u) { return 76u; } + let index = fs_lookup(); if index == 0xffffffffu { return 44u; } + return fs_set_times(index, get_value(base + 4u), get_value(base + 5u), arg(base, 6u)); + } + if syscall == WASI_PATH_SYMLINK { + let old_ptr = arg(base, 0u); let old_size = arg(base, 1u); let fd = arg(base, 2u); + if !bounds(old_ptr, old_size) { return 21u; } + if old_size == 0u { return 44u; } + for (var i = 0u; i < old_size; i += 1u) { if read_byte(old_ptr + i) == 0u { return 28u; } } + let error = fs_path(fd, arg(base, 3u), arg(base, 4u)); if error != 0u { return error; } + if !fs_right(fd, 1u << 24u) { return 76u; } + if fs_lookup() != 0xffffffffu { return 20u; } + if !fs_parent_exists() { return 44u; } + let index = fs_create(3u); if index == 0xffffffffu { return 51u; } + let resized = fs_resize(index, old_size); + if resized != 0u { write_heap(fs_entry(index), 0u); return resized; } + let start = read_heap(fs_inode(index) + 3u); + for (var i = 0u; i < old_size; i += 1u) { set_heap_byte(fs_data(), start + i, read_byte(old_ptr + i)); } + return 0u; + } + if syscall == WASI_PATH_READLINK { + let fd = arg(base, 0u); let error = fs_path(fd, arg(base, 1u), arg(base, 2u)); if error != 0u { return error; } + if !fs_right(fd, 1u << 15u) { return 76u; } + let index = fs_lookup(); if index == 0xffffffffu { return 44u; } + if read_heap(fs_entry(index)) != 3u { return 28u; } + let pointer = arg(base, 3u); let size = arg(base, 4u); let result_ptr = arg(base, 5u); + if !bounds(pointer, size) || !bounds(result_ptr, 4u) { return 21u; } + let entry = fs_inode(index); let amount = min(size, read_heap(entry + 2u)); let start = read_heap(entry + 3u); + if !fs_charge(amount) { return 29u; } + for (var i = 0u; i < amount; i += 1u) { write_byte(pointer + i, heap_byte(fs_data(), start + i)); } + store32(result_ptr, amount); return 0u; + } + if syscall == WASI_PATH_RENAME || syscall == WASI_PATH_LINK { + let linking = syscall == WASI_PATH_LINK; let old_fd = arg(base, 0u); + if linking && arg(base, 1u) > 1u { return 28u; } + var error = fs_resolve(old_fd, arg(base, select(1u, 2u, linking)), arg(base, select(2u, 3u, linking)), linking && arg(base, 1u) == 1u); + if error != 0u { return error; } + if !fs_right(old_fd, select(1u << 16u, 1u << 11u, linking)) { return 76u; } + let index = fs_lookup(); if index == 0xffffffffu { return 44u; } + if linking && read_heap(fs_entry(index)) == 2u { return 63u; } + let new_fd = arg(base, select(3u, 4u, linking)); + error = fs_path(new_fd, arg(base, select(4u, 5u, linking)), arg(base, select(5u, 6u, linking))); + if error != 0u { return error; } + if !fs_right(new_fd, select(1u << 17u, 1u << 12u, linking)) { return 76u; } + if !fs_parent_exists() { return 44u; } + let existing = fs_lookup(); + if linking { + if existing != 0xffffffffu { return 20u; } + let linked = fs_create(read_heap(fs_entry(index))); if linked == 0xffffffffu { return 51u; } + write_heap(fs_entry(linked) + 5u, read_heap(fs_entry(index) + 5u)); + } else { return fs_rename(index, existing); } + return 0u; + } + // Descriptor services. + let fd = arg(base, 0u); let index = fs_file(fd); + if index == 0xffffffffu { return 8u; } + if syscall == WASI_SOCK_ACCEPT || syscall == WASI_SOCK_RECV || syscall == WASI_SOCK_SEND || syscall == WASI_SOCK_SHUTDOWN { return 57u; } + let descriptor = fs_descriptor(fd); let entry = fs_inode(index); + if syscall == WASI_FD_CLOSE { write_heap(descriptor, 0u); return 0u; } + if syscall == WASI_FD_RENUMBER { + let destination = arg(base, 1u); if destination >= config.fs_fds { return 8u; } + if destination != fd { + for (var i = 0u; i < 8u; i += 1u) { write_heap(fs_descriptor(destination) + i, read_heap(descriptor + i)); } + write_heap(descriptor, 0u); + } + return 0u; + } + if syscall == WASI_FD_PRESTAT_GET || syscall == WASI_FD_PRESTAT_DIR_NAME { + if fd != 3u || index != 3u { return 8u; } + let pointer = arg(base, 1u); + if syscall == WASI_FD_PRESTAT_GET { + if !bounds(pointer, 8u) { return 21u; } + store32(pointer, 0u); store32(pointer + 4u, 1u); + } else { + if arg(base, 2u) < 1u { return 37u; } + if !bounds(pointer, 1u) { return 21u; } + write_byte(pointer, 46u); + } + return 0u; + } + if syscall == WASI_FD_FDSTAT_GET { + let pointer = arg(base, 1u); if !bounds(pointer, 24u) { return 21u; } + for (var i = 0u; i < 24u; i += 1u) { write_byte(pointer + i, 0u); } + write_byte(pointer, select(select(4u, 3u, read_heap(fs_entry(index)) == 2u), 2u, index < 3u)); + write_byte(pointer + 2u, read_heap(descriptor + 3u)); + store64(pointer + 8u, vec2u(read_heap(descriptor + 4u), read_heap(descriptor + 5u))); + store64(pointer + 16u, vec2u(read_heap(descriptor + 6u), read_heap(descriptor + 7u))); + return 0u; + } + if syscall == WASI_FD_FDSTAT_SET_FLAGS { + if !fs_right(fd, 1u << 3u) { return 76u; } + let flags = arg(base, 1u); if flags > 31u { return 28u; } + write_heap(descriptor + 3u, flags); return 0u; + } + if syscall == WASI_FD_FDSTAT_SET_RIGHTS { + let rights = get_value(base + 1u); let inheriting = get_value(base + 2u); + if (rights.x & ~read_heap(descriptor + 4u)) != 0u || (rights.y & ~read_heap(descriptor + 5u)) != 0u || (inheriting.x & ~read_heap(descriptor + 6u)) != 0u || (inheriting.y & ~read_heap(descriptor + 7u)) != 0u { return 76u; } + write_heap(descriptor + 4u, rights.x); write_heap(descriptor + 5u, rights.y); + write_heap(descriptor + 6u, inheriting.x); write_heap(descriptor + 7u, inheriting.y); return 0u; + } + if syscall == WASI_FD_SEEK || syscall == WASI_FD_TELL { + if !fs_right(fd, select(1u << 5u, 1u << 2u, syscall == WASI_FD_SEEK)) { return 76u; } + let result_ptr = arg(base, select(1u, 3u, syscall == WASI_FD_SEEK)); + if !bounds(result_ptr, 8u) { return 21u; } + var position = vec2u(read_heap(descriptor + 1u), read_heap(descriptor + 2u)); + if syscall == WASI_FD_SEEK { + let whence = arg(base, 2u); let offset = get_value(base + 1u); + if whence > 2u { return 28u; } + if whence == 0u { position = vec2u(0u); } + if whence == 2u { position = vec2u(read_heap(entry + 2u), 0u); } + if (offset.y >> 31u) != 0u && lt64(position, neg64(offset)) { return 28u; } + let next = add64(position, offset); + if (offset.y >> 31u) == 0u && lt64(next, position) { return 61u; } + position = next; write_heap(descriptor + 1u, position.x); write_heap(descriptor + 2u, position.y); + } + store64(result_ptr, position); return 0u; + } + if syscall == WASI_FD_FILESTAT_GET { + if !fs_right(fd, 1u << 21u) { return 76u; } + if !bounds(arg(base, 1u), 64u) { return 21u; } + fs_stat(index, arg(base, 1u)); return 0u; + } + if syscall == WASI_FD_FILESTAT_SET_TIMES { + if !fs_right(fd, 1u << 23u) { return 76u; } + return fs_set_times(index, get_value(base + 1u), get_value(base + 2u), arg(base, 3u)); + } + if syscall == WASI_FD_FILESTAT_SET_SIZE || syscall == WASI_FD_ALLOCATE { + if !fs_right(fd, select(1u << 22u, 1u << 8u, syscall == WASI_FD_ALLOCATE)) { return 76u; } + var length = get_value(base + 1u); + if syscall == WASI_FD_ALLOCATE { + let end = add64(length, get_value(base + 2u)); + if lt64(end, length) { return 61u; } + length = end; length.x = max(length.x, read_heap(entry + 2u)); + } + if length.y != 0u { return 27u; } + if read_heap(fs_entry(index)) == 2u { return 31u; } + return fs_resize(index, length.x); + } + if syscall == WASI_FD_SYNC || syscall == WASI_FD_DATASYNC { + return select(76u, 0u, fs_right(fd, select(1u, 1u << 4u, syscall == WASI_FD_SYNC))); + } + if syscall == WASI_FD_ADVISE { + if arg(base, 3u) > 5u { return 28u; } + return select(76u, 0u, fs_right(fd, 1u << 7u)); + } + return 58u; +} diff --git a/wasmgpu/numeric.wgsl b/wasmgpu/numeric.wgsl new file mode 100644 index 0000000..e83a68d --- /dev/null +++ b/wasmgpu/numeric.wgsl @@ -0,0 +1,261 @@ +// Integer and IEEE-754 operations. vec2 words are little endian. +// Floating point is implemented with integers, including subnormals and f64. +// All arithmetic in this file executes on the GPU. +fn add64(a: vec2u, b: vec2u) -> vec2u { + let lo = a.x + b.x; + return vec2u(lo, a.y + b.y + u32(lo < a.x)); +} +fn neg64(a: vec2u) -> vec2u { return add64(~a, vec2u(1u, 0u)); } +fn sub64(a: vec2u, b: vec2u) -> vec2u { return add64(a, neg64(b)); } +fn zero64(a: vec2u) -> bool { return (a.x | a.y) == 0u; } +fn lt64(a: vec2u, b: vec2u) -> bool { return a.y < b.y || (a.y == b.y && a.x < b.x); } +fn eq64(a: vec2u, b: vec2u) -> bool { return all(a == b); } +fn shl64(a: vec2u, n: u32) -> vec2u { + if n == 0u { return a; } + if n >= 64u { return vec2u(0u); } + if n >= 32u { return vec2u(0u, a.x << (n - 32u)); } + return vec2u(a.x << n, (a.y << n) | (a.x >> (32u - n))); +} +fn shr64(a: vec2u, n: u32) -> vec2u { + if n == 0u { return a; } + if n >= 64u { return vec2u(0u); } + if n >= 32u { return vec2u(a.y >> (n - 32u), 0u); } + return vec2u((a.x >> n) | (a.y << (32u - n)), a.y >> n); +} +fn jam64(a: vec2u, n: u32) -> vec2u { + let shifted = shr64(a, n); + return shifted | vec2u(u32(!eq64(shl64(shifted, n), a)), 0u); +} +fn bit64(a: vec2u, n: i32) -> u32 { + if n < 0 || n >= 64 { return 0u; } + return shr64(a, u32(n)).x & 1u; +} +fn clz64(a: vec2u) -> u32 { + if a.y != 0u { return countLeadingZeros(a.y); } + return 32u + countLeadingZeros(a.x); +} +fn mul32(a: u32, b: u32) -> vec2u { + let a0 = a & 65535u; let a1 = a >> 16u; + let b0 = b & 65535u; let b1 = b >> 16u; + let p0 = a0 * b0; + let p1 = a1 * b0 + (p0 >> 16u); + let p2 = a0 * b1 + (p1 & 65535u); + return vec2u((p0 & 65535u) | (p2 << 16u), a1 * b1 + (p1 >> 16u) + (p2 >> 16u)); +} +fn mul64(a: vec2u, b: vec2u) -> vec2u { + let low = mul32(a.x, b.x); + return vec2u(low.x, low.y + a.x * b.y + a.y * b.x); +} +struct Division { quotient: vec2u, remainder: vec2u } +fn div64(a: vec2u, b: vec2u) -> Division { + var q = vec2u(0u); var rem = vec2u(0u); + for (var i = 63; i >= 0; i -= 1) { + let carry = rem.y >> 31u; + rem = shl64(rem, 1u) | vec2u(bit64(a, i), 0u); + if carry != 0u || !lt64(rem, b) { + rem = sub64(rem, b); + q |= shl64(vec2u(1u, 0u), u32(i)); + } + } + return Division(q, rem); +} +fn sar64(a: vec2u, n: u32) -> vec2u { + let shift = n & 63u; + if shift == 0u { return a; } + let sign = a.y >> 31u; + return shr64(a, shift) | select(vec2u(0u), shl64(vec2u(0xffffffffu), 64u - shift), sign != 0u); +} +fn add128(a: vec4u, b: vec4u) -> vec4u { + let lo = add64(a.xy, b.xy); + let hi = add64(add64(a.zw, b.zw), vec2u(u32(lt64(lo, a.xy)), 0u)); + return vec4u(lo, hi); +} +fn mul128(a: vec2u, b: vec2u) -> vec4u { + let p0 = mul32(a.x, b.x); let p1 = mul32(a.x, b.y); + let p2 = mul32(a.y, b.x); let p3 = mul32(a.y, b.y); + return add128(add128(vec4u(p0, p3), vec4u(0u, p1, 0u)), vec4u(0u, p2, 0u)); +} +struct SoftFloat { sign: u32, exponent: i32, sig: vec2u, kind: u32 } +fn unpack_float(bits: vec2u, single: bool) -> SoftFloat { + var sign = bits.y >> 31u; + var exponent = i32((bits.y >> 20u) & 2047u); + var sig = vec2u(bits.x, bits.y & 0xfffffu); + var max_exp = 2047; var bias = 1023; + if single { + sign = bits.x >> 31u; exponent = i32((bits.x >> 23u) & 255u); + sig = shl64(vec2u(bits.x & 0x7fffffu, 0u), 29u); + max_exp = 255; bias = 127; + } + if exponent == max_exp { return SoftFloat(sign, 0, sig, select(2u, 1u, zero64(sig))); } + if exponent == 0 { + exponent = 1 - bias; + if !zero64(sig) { + let shift = clz64(sig) - 11u; + sig = shl64(sig, shift); exponent -= i32(shift); + } + } else { + exponent -= bias; sig.y |= 0x100000u; + } + return SoftFloat(sign, exponent, sig, 0u); +} +fn float_special(sign: u32, nan: bool, single: bool) -> vec2u { + if single { return vec2u((sign << 31u) | 0x7f800000u | select(0u, 0x400000u, nan), 0u); } + return vec2u(0u, (sign << 31u) | 0x7ff00000u | select(0u, 0x80000u, nan)); +} +fn float_zero(sign: u32, single: bool) -> vec2u { + if single { return vec2u(sign << 31u, 0u); } + return vec2u(0u, sign << 31u); +} +fn pack_float(sign: u32, original_exp: i32, original_sig: vec2u, single: bool) -> vec2u { + if zero64(original_sig) { return float_zero(sign, single); } + var sig = original_sig; var exponent = original_exp; + while sig.y >= 0x1000000u { sig = jam64(sig, 1u); exponent += 1; } + while sig.y < 0x800000u { sig = shl64(sig, 1u); exponent -= 1; } + let min_exp = select(-1022, -126, single); let max_exp = select(1023, 127, single); + if exponent < min_exp { sig = jam64(sig, u32(min_exp - exponent)); exponent = min_exp; } + if single { sig = jam64(sig, 29u); } + let round = sig.x & 7u; + sig = shr64(sig, 3u); + if round > 4u || (round == 4u && (sig.x & 1u) != 0u) { sig = add64(sig, vec2u(1u, 0u)); } + let mantissa_bits = select(52u, 23u, single); + if bit64(sig, i32(mantissa_bits + 1u)) != 0u { sig = shr64(sig, 1u); exponent += 1; } + if exponent > max_exp { return float_special(sign, false, single); } + let encoded_exp = select(0u, u32(exponent + max_exp), bit64(sig, i32(mantissa_bits)) != 0u); + if single { return vec2u((sign << 31u) | (encoded_exp << 23u) | (sig.x & 0x7fffffu), 0u); } + return vec2u(sig.x, (sign << 31u) | (encoded_exp << 20u) | (sig.y & 0xfffffu)); +} +fn soft_add(x: vec2u, y: vec2u, subtract: bool, single: bool) -> vec2u { + var a = unpack_float(x, single); var b = unpack_float(y, single); + if subtract { b.sign ^= 1u; } + if a.kind == 2u || b.kind == 2u { return float_special(0u, true, single); } + if a.kind == 1u { + return float_special(a.sign, b.kind == 1u && b.sign != a.sign, single); + } + if b.kind == 1u { return float_special(b.sign, false, single); } + if zero64(a.sig) && zero64(b.sig) { return float_zero(a.sign & b.sign, single); } + if zero64(a.sig) { return pack_float(b.sign, b.exponent, shl64(b.sig, 3u), single); } + if zero64(b.sig) { return pack_float(a.sign, a.exponent, shl64(a.sig, 3u), single); } + if a.exponent < b.exponent || (a.exponent == b.exponent && lt64(a.sig, b.sig)) { + let swap = a; a = b; b = swap; + } + let aa = shl64(a.sig, 3u); let bb = jam64(shl64(b.sig, 3u), u32(a.exponent - b.exponent)); + if a.sign == b.sign { return pack_float(a.sign, a.exponent, add64(aa, bb), single); } + let difference = sub64(aa, bb); + return pack_float(select(a.sign, 0u, zero64(difference)), a.exponent, difference, single); +} +fn soft_mul(x: vec2u, y: vec2u, single: bool) -> vec2u { + let a = unpack_float(x, single); let b = unpack_float(y, single); let sign = a.sign ^ b.sign; + if a.kind == 2u || b.kind == 2u { return float_special(0u, true, single); } + if a.kind == 1u || b.kind == 1u { + return float_special(sign, (a.kind == 0u && zero64(a.sig)) || (b.kind == 0u && zero64(b.sig)), single); + } + let product = mul128(a.sig, b.sig); + // Product / 2^49 gives a significand with three rounding bits. + let sig = vec2u((product.y >> 17u) | (product.z << 15u), (product.z >> 17u) | (product.w << 15u)); + let sticky = u32(product.x != 0u || (product.y & 0x1ffffu) != 0u); + return pack_float(sign, a.exponent + b.exponent, sig | vec2u(sticky, 0u), single); +} +fn soft_div(x: vec2u, y: vec2u, single: bool) -> vec2u { + let a = unpack_float(x, single); let b = unpack_float(y, single); let sign = a.sign ^ b.sign; + if a.kind == 2u || b.kind == 2u || (a.kind == 1u && b.kind == 1u) || (a.kind == 0u && b.kind == 0u && zero64(a.sig) && zero64(b.sig)) { + return float_special(0u, true, single); + } + if a.kind == 1u || (b.kind == 0u && zero64(b.sig)) { return float_special(sign, false, single); } + if b.kind == 1u || zero64(a.sig) { return float_zero(sign, single); } + var rem = a.sig; var quotient = vec2u(0u); + for (var i = 0u; i < 57u; i += 1u) { + quotient = shl64(quotient, 1u); + if !lt64(rem, b.sig) { rem = sub64(rem, b.sig); quotient.x |= 1u; } + rem = shl64(rem, 1u); + } + quotient.x |= u32(!zero64(rem)); + return pack_float(sign, a.exponent - b.exponent - 1, quotient, single); +} +fn soft_sqrt(x: vec2u, single: bool) -> vec2u { + let a = unpack_float(x, single); + if a.kind == 2u { return float_special(0u, true, single); } + if a.kind == 0u && zero64(a.sig) { return x; } + if a.sign != 0u { return float_special(0u, true, single); } + if a.kind == 1u { return x; } + let shift = 58 + (a.exponent & 1); + var root = vec2u(0u); var rem = vec2u(0u); + for (var i = 55; i >= 0; i -= 1) { + let pair = (bit64(a.sig, i * 2 + 1 - shift) << 1u) | bit64(a.sig, i * 2 - shift); + rem = shl64(rem, 2u) | vec2u(pair, 0u); + let test = shl64(root, 2u) | vec2u(1u, 0u); + root = shl64(root, 1u); + if !lt64(rem, test) { rem = sub64(rem, test); root.x |= 1u; } + } + root.x |= u32(!zero64(rem)); + return pack_float(0u, a.exponent >> 1, root, single); +} +// Comparison: 0 equal, 1 less, 2 greater, 3 unordered. +fn float_compare(x: vec2u, y: vec2u, single: bool) -> u32 { + let a = unpack_float(x, single); let b = unpack_float(y, single); + if a.kind == 2u || b.kind == 2u { return 3u; } + if a.kind == 0u && b.kind == 0u && zero64(a.sig) && zero64(b.sig) { return 0u; } + if eq64(x, y) { return 0u; } + if a.sign != b.sign { return select(2u, 1u, a.sign != 0u); } + return select(2u, 1u, lt64(x, y) != (a.sign != 0u)); +} +fn soft_round(x: vec2u, mode: u32, single: bool) -> vec2u { + let a = unpack_float(x, single); + if a.kind == 2u { return float_special(0u, true, single); } + if a.kind == 1u || zero64(a.sig) || a.exponent >= 52 { return x; } + let shift = u32(max(0, 52 - a.exponent)); + var integer = shr64(a.sig, shift); + let remainder = sub64(a.sig, shl64(integer, shift)); + if !zero64(remainder) { + if (mode == 0u && a.sign == 0u) || (mode == 1u && a.sign != 0u) { + integer = add64(integer, vec2u(1u, 0u)); + } else if mode == 3u && shift <= 53u { + let half = shl64(vec2u(1u, 0u), shift - 1u); + if lt64(half, remainder) || (eq64(half, remainder) && (integer.x & 1u) != 0u) { + integer = add64(integer, vec2u(1u, 0u)); + } + } + } + return pack_float(a.sign, 52, shl64(integer, 3u), single); +} +fn integer_to_float(x: vec2u, is_signed: bool, wide: bool, single: bool) -> vec2u { + var integer = x; var sign = 0u; + if wide { + if is_signed && (x.y >> 31u) != 0u { sign = 1u; integer = neg64(x); } + } else { + integer.y = 0u; + if is_signed && (x.x >> 31u) != 0u { sign = 1u; integer.x = 0u - x.x; } + } + if zero64(integer) { return float_zero(0u, single); } + let top = 63u - clz64(integer); + var extended = vec2u(0u); + if top > 55u { extended = jam64(integer, top - 55u); } + else { extended = shl64(integer, 55u - top); } + return pack_float(sign, i32(top), extended, single); +} +fn convert_float(x: vec2u, from_single: bool) -> vec2u { + let a = unpack_float(x, from_single); + if a.kind != 0u { return float_special(a.sign, a.kind == 2u, !from_single); } + return pack_float(a.sign, a.exponent, shl64(a.sig, 3u), !from_single); +} +struct IntConversion { value: vec2u, trap: u32 } +fn float_to_integer(x: vec2u, single: bool, wide: bool, is_signed: bool, saturate: bool) -> IntConversion { + let a = unpack_float(x, single); + let size = select(32u, 64u, wide); + let max_unsigned = select(vec2u(0xffffffffu, 0u), vec2u(0xffffffffu), wide); + let minimum = shl64(vec2u(1u, 0u), size - 1u); + let maximum = select(max_unsigned, sub64(minimum, vec2u(1u, 0u)), is_signed); + if a.kind == 2u { return IntConversion(vec2u(0u), select(6u, 0u, saturate)); } + var magnitude = vec2u(0u); + if a.exponent >= 52 { magnitude = shl64(a.sig, u32(a.exponent - 52)); } + else { magnitude = shr64(a.sig, u32(52 - a.exponent)); } + var invalid = a.kind == 1u || a.exponent >= i32(size); + if is_signed { invalid = invalid || lt64(select(maximum, minimum, a.sign != 0u), magnitude); } + else { invalid = invalid || (a.sign != 0u && !zero64(magnitude)); } + if invalid { + if !saturate { return IntConversion(vec2u(0u), 5u); } + return IntConversion(select(maximum, select(vec2u(0u), minimum, is_signed), a.sign != 0u), 0u); + } + var result = select(magnitude, neg64(magnitude), a.sign != 0u); + if !wide { result.y = 0u; } + return IntConversion(result, 0u); +} diff --git a/wasmgpu/operations.wgsl b/wasmgpu/operations.wgsl new file mode 100644 index 0000000..b9bfa4e --- /dev/null +++ b/wasmgpu/operations.wgsl @@ -0,0 +1,150 @@ +struct NumericResult { value: vec2u, trap: u32 } +fn numeric(op: u32, a: vec2u, b: vec2u) -> NumericResult { + var result = vec2u(0u); + var trap = 0u; + let ai = bitcast(a.x); let bi = bitcast(b.x); + let shift = b.x & 31u; + switch op { + case 0x45u: { result.x = u32(a.x == 0u); } + case 0x46u: { result.x = u32(a.x == b.x); } + case 0x47u: { result.x = u32(a.x != b.x); } + case 0x48u: { result.x = u32(ai < bi); } + case 0x49u: { result.x = u32(a.x < b.x); } + case 0x4au: { result.x = u32(ai > bi); } + case 0x4bu: { result.x = u32(a.x > b.x); } + case 0x4cu: { result.x = u32(ai <= bi); } + case 0x4du: { result.x = u32(a.x <= b.x); } + case 0x4eu: { result.x = u32(ai >= bi); } + case 0x4fu: { result.x = u32(a.x >= b.x); } + case 0x50u: { result.x = u32(zero64(a)); } + case 0x51u: { result.x = u32(eq64(a, b)); } + case 0x52u: { result.x = u32(!eq64(a, b)); } + case 0x53u, 0x55u, 0x57u, 0x59u: { + let less = lt64(a ^ vec2u(0u, 0x80000000u), b ^ vec2u(0u, 0x80000000u)); + let equal = eq64(a, b); + switch op { + case 0x53u: { result.x = u32(less); } + case 0x55u: { result.x = u32(!less && !equal); } + case 0x57u: { result.x = u32(less || equal); } + default: { result.x = u32(!less); } + } + } + case 0x54u: { result.x = u32(lt64(a, b)); } + case 0x56u: { result.x = u32(lt64(b, a)); } + case 0x58u: { result.x = u32(!lt64(b, a)); } + case 0x5au: { result.x = u32(!lt64(a, b)); } + case 0x67u: { result.x = countLeadingZeros(a.x); } + case 0x68u: { result.x = countTrailingZeros(a.x); } + case 0x69u: { result.x = countOneBits(a.x); } + case 0x6au: { result.x = a.x + b.x; } + case 0x6bu: { result.x = a.x - b.x; } + case 0x6cu: { result.x = a.x * b.x; } + case 0x6du: { + if b.x == 0u { trap = 4u; } + else if a.x == 0x80000000u && b.x == 0xffffffffu { trap = 5u; } + else { result.x = bitcast(ai / bi); } + } + case 0x6eu: { if b.x == 0u { trap = 4u; } else { result.x = a.x / b.x; } } + case 0x6fu: { if b.x == 0u { trap = 4u; } else { result.x = bitcast(ai % bi); } } + case 0x70u: { if b.x == 0u { trap = 4u; } else { result.x = a.x % b.x; } } + case 0x71u: { result.x = a.x & b.x; } + case 0x72u: { result.x = a.x | b.x; } + case 0x73u: { result.x = a.x ^ b.x; } + case 0x74u: { result.x = a.x << shift; } + case 0x75u: { result.x = bitcast(ai >> shift); } + case 0x76u: { result.x = a.x >> shift; } + case 0x77u: { result.x = (a.x << shift) | (a.x >> ((32u - shift) & 31u)); } + case 0x78u: { result.x = (a.x >> shift) | (a.x << ((32u - shift) & 31u)); } + case 0x79u: { result.x = clz64(a); } + case 0x7au: { result.x = select(countTrailingZeros(a.x), 32u + countTrailingZeros(a.y), a.x == 0u); } + case 0x7bu: { result.x = countOneBits(a.x) + countOneBits(a.y); } + case 0x7cu: { result = add64(a, b); } + case 0x7du: { result = sub64(a, b); } + case 0x7eu: { result = mul64(a, b); } + case 0x7fu, 0x80u, 0x81u, 0x82u: { + let is_signed = op == 0x7fu || op == 0x81u; + let remainder = op >= 0x81u; + if zero64(b) { trap = 4u; } + else if is_signed && !remainder && eq64(a, vec2u(0u, 0x80000000u)) && all(b == vec2u(0xffffffffu)) { trap = 5u; } + else { + let aneg = is_signed && (a.y >> 31u) != 0u; let bneg = is_signed && (b.y >> 31u) != 0u; + let division = div64(select(a, neg64(a), aneg), select(b, neg64(b), bneg)); + result = select(division.quotient, division.remainder, remainder); + if select(aneg != bneg, aneg, remainder) { result = neg64(result); } + } + } + case 0x83u: { result = a & b; } + case 0x84u: { result = a | b; } + case 0x85u: { result = a ^ b; } + case 0x86u: { result = shl64(a, b.x & 63u); } + case 0x87u: { result = sar64(a, b.x); } + case 0x88u: { result = shr64(a, b.x & 63u); } + case 0x89u: { result = shl64(a, b.x & 63u) | shr64(a, (64u - b.x) & 63u); } + case 0x8au: { result = shr64(a, b.x & 63u) | shl64(a, (64u - b.x) & 63u); } + case 0xa7u: { result.x = a.x; } + case 0xacu: { result = vec2u(a.x, select(0u, 0xffffffffu, ai < 0)); } + case 0xadu: { result.x = a.x; } + case 0xb6u: { result = convert_float(a, false); } + case 0xbbu: { result = convert_float(a, true); } + case 0xbcu, 0xbeu: { result.x = a.x; } + case 0xbdu, 0xbfu: { result = a; } + case 0xc0u: { result.x = bitcast(bitcast(a.x << 24u) >> 24); } + case 0xc1u: { result.x = bitcast(bitcast(a.x << 16u) >> 16); } + case 0xc2u, 0xc3u, 0xc4u: { + let bits = select(select(32u, 16u, op == 0xc3u), 8u, op == 0xc2u); + result = sar64(shl64(a, 64u - bits), 64u - bits); + } + default: { + if op >= 0x5bu && op <= 0x66u { + let single = op <= 0x60u; + let compare = float_compare(a, b, single); + switch (op - select(0x61u, 0x5bu, single)) { + case 0u: { result.x = u32(compare == 0u); } + case 1u: { result.x = u32(compare != 0u); } + case 2u: { result.x = u32(compare == 1u); } + case 3u: { result.x = u32(compare == 2u); } + case 4u: { result.x = u32(compare <= 1u); } + default: { result.x = u32(compare == 0u || compare == 2u); } + } + } else if op >= 0x8bu && op <= 0xa6u { + let single = op <= 0x98u; + let kind = op - select(0x99u, 0x8bu, single); + let mask = select(vec2u(0u, 0x80000000u), vec2u(0x80000000u, 0u), single); + switch kind { + case 0u: { result = a & ~mask; } + case 1u: { result = a ^ mask; } + case 2u, 3u, 4u, 5u: { result = soft_round(a, kind - 2u, single); } + case 6u: { result = soft_sqrt(a, single); } + case 7u: { result = soft_add(a, b, false, single); } + case 8u: { result = soft_add(a, b, true, single); } + case 9u: { result = soft_mul(a, b, single); } + case 10u: { result = soft_div(a, b, single); } + case 11u, 12u: { + let compare = float_compare(a, b, single); + if compare == 3u { result = float_special(0u, true, single); } + else if compare == 0u { result = select(a & b, a | b, kind == 11u); } + else { result = select(b, a, compare == select(2u, 1u, kind == 11u)); } + } + default: { result = (a & ~mask) | (b & mask); } + } + } else if (op >= 0xb2u && op <= 0xb5u) || (op >= 0xb7u && op <= 0xbau) { + let single = op <= 0xb5u; + let kind = op - select(0xb7u, 0xb2u, single); + result = integer_to_float(a, (kind & 1u) == 0u, kind >= 2u, single); + } else { + var single = true; var wide = false; var is_signed = true; var saturate = false; + if op >= 0xfc00u && op <= 0xfc07u { + let kind = op - 0xfc00u; + single = (kind & 3u) < 2u; wide = kind >= 4u; is_signed = (kind & 1u) == 0u; saturate = true; + } else if op >= 0xa8u && op <= 0xabu { + single = op <= 0xa9u; is_signed = (op & 1u) == 0u; + } else if op >= 0xaeu && op <= 0xb1u { + single = op <= 0xafu; is_signed = (op & 1u) == 0u; wide = true; + } else { return NumericResult(vec2u(0u), 11u); } + let conversion = float_to_integer(a, single, wide, is_signed, saturate); + result = conversion.value; trap = conversion.trap; + } + } + } + return NumericResult(result, trap); +} diff --git a/wasmgpu/py.typed b/wasmgpu/py.typed new file mode 100644 index 0000000..e69de29 diff --git a/wasmgpu/runtime.py b/wasmgpu/runtime.py new file mode 100644 index 0000000..b62a980 --- /dev/null +++ b/wasmgpu/runtime.py @@ -0,0 +1,600 @@ +"""GPU execution and Python API. No CPU WebAssembly executor is used here.""" +from __future__ import annotations + +import copy +import importlib +import os +import struct +import threading +from array import array +from pathlib import Path +from typing import Iterable, Mapping, Sequence, Tuple, cast + +from .binary import F32, F64, I32, I64, BinaryModule +from .errors import GPUUnavailableError, ResourceLimitError, Trap +from .types import ( + AdapterInfo, + Backend, + Blob, + Buffer, + RequestAdapter, + RequestDevice, + Result, + Row, + Scalar, +) +from .wasi import WASI_IDS, Wasi, normalize_path + +_TRAPS = { + 1: 'unreachable', 2: 'out of bounds memory access', 3: 'out of bounds table access', + 4: 'integer divide by zero', 5: 'integer overflow', 6: 'invalid conversion to integer', + 7: 'stack exhausted', 8: 'fuel exhausted', 9: 'uninitialized element', + 10: 'indirect call type mismatch', 11: 'invalid runtime state', 12: 'WASI proc_exit', 13: 'WASI proc_raise', +} +_CONTEXT: _Context | None = None +_CONTEXT_LOCK = threading.Lock() + + +def _words(blob: Blob) -> array[int]: + words = array('I') + words.frombytes(blob) + return words + + +def _positive(value: int, name: str, zero: bool = False) -> int: + if isinstance(value, bool) or not isinstance(value, int): + raise TypeError(f'{name} must be an integer') + if value < (0 if zero else 1) or value > 0xFFFFFFFF: + raise ValueError(f'{name} out of range') + return value + + +def _encode(value: Scalar, ty: int) -> int: + if ty in (I32, I64): + bits = 32 if ty == I32 else 64 + if isinstance(value, bool) or not isinstance(value, int): + raise TypeError('integer WASM arguments require Python int') + if not -(1 << (bits - 1)) <= value < (1 << bits): + raise OverflowError(f'integer does not fit i{bits}') + return value & ((1 << bits) - 1) + if ty in (F32, F64): + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise TypeError('floating WASM arguments require int or float') + try: + return int.from_bytes(struct.pack(' Scalar: + if ty in (I32, I64): + bits = 32 if ty == I32 else 64 + value &= (1 << bits) - 1 + return value - (1 << bits) if value >> (bits - 1) else value + if ty in (F32, F64): + size = 4 if ty == F32 else 8 + return cast(Tuple[float, ...], struct.unpack(' None: + wgpu = cast(Backend, importlib.import_module('wgpu')) + self.wgpu = wgpu + try: + request_adapter = cast(RequestAdapter, cast(object, getattr(wgpu.gpu, 'request_adapter_sync', None)) or cast(object, wgpu.gpu.request_adapter)) + self.adapter = request_adapter(power_preference='high-performance') + if self.adapter is None or str(self.adapter.info.get('adapter_type', '')).lower() == 'cpu': + raise GPUUnavailableError('a hardware GPU is required; CPU adapters are rejected') + request_device = cast(RequestDevice, cast(object, getattr(self.adapter, 'request_device_sync', None)) or cast(object, self.adapter.request_device)) + requested = {key: min(256 * 1024 * 1024, value) for key, value in self.adapter.limits.items() + if key.replace('_', '-') in ('max-storage-buffer-binding-size', 'max-buffer-size')} + self.device = request_device(required_limits=requested) + self.limits = {key.replace('_', '-'): value for key, value in self.device.limits.items()} + except GPUUnavailableError: + raise + except Exception as error: + raise GPUUnavailableError(f'could not initialize a hardware GPU: {error}') from error + constants = '\n'.join(f'const WASI_{name.upper()}: u32 = {index}u;' for name, index in WASI_IDS.items()) + source = constants + '\n' + '\n'.join((Path(__file__).parent / name).read_text() for name in ('numeric.wgsl', 'operations.wgsl', 'vm.wgsl', 'filesystem.wgsl')) + shader = self.device.create_shader_module(label='wasmgpu interpreter', code=source) + self.pipeline = self.device.create_compute_pipeline(layout='auto', compute={'module': shader, 'entry_point': 'run'}) + + def buffer(self, size: int = 0, data: Blob | array[int] | None = None, uniform: bool = False) -> Buffer: + flags = self.wgpu.BufferUsage + usage = flags.COPY_DST | flags.COPY_SRC | (flags.UNIFORM if uniform else flags.STORAGE) + if data is not None: + return self.device.create_buffer_with_data(data=data, usage=usage) + return self.device.create_buffer(size=max(16, size), usage=usage) + + +def _context() -> _Context: + global _CONTEXT # noqa: PLW0603 - Lazy process-wide device, protected by a lock. + with _CONTEXT_LOCK: + if _CONTEXT is None: + _CONTEXT = _Context() + return _CONTEXT + + +def _program(module: BinaryModule) -> array[int]: + canonical = [module.types.index(signature) for signature in module.types] + words = [0] * 16 + code: list[int] = [] + for fn in module.functions: + params, results = module.types[fn.type_index] + instructions = fn.instructions + if fn.imported is not None: + fn.locals = list(params) + instructions = [[0x20, index, 0, 0] for index in range(len(params))] + [[0x10, module.functions.index(fn), 0, 0], [0x0F, 0, 0, 0]] + fn.offset = len(code) // 4 + words.extend([fn.offset, len(params), len(results), len(fn.locals), canonical[fn.type_index], WASI_IDS[fn.imported[1]] if fn.imported else 0, 0, 0]) + for op, operand, b, c in instructions: + a = operand + if 0x100 <= op <= 0x104: + a += fn.offset + elif op == 0x11: + a = canonical[a] + code.extend([op, a & 0xFFFFFFFF, b & 0xFFFFFFFF, c & 0xFFFFFFFF]) + words[0] = len(words) + words.extend(code) + words[1] = len(words) + words.extend([0] * (len(module.data) * 2)) + words[2] = len(words) + words.extend([0] * (len(module.elements) * 2)) + words[3] = len(module.functions) + for i, (_, blob) in enumerate(module.data): + words[words[1] + i * 2:words[1] + i * 2 + 2] = [len(words), len(blob)] + words.extend(_words(blob + b'\0' * (-len(blob) % 4))) + for i, (_, elements, _, _, _) in enumerate(module.elements): + words[words[2] + i * 2:words[2] + i * 2 + 2] = [len(words), len(elements)] + words.extend(elements) + return array('I', words) + + +class Module: + """Validate a binary WASM file. Function bodies execute only on the GPU. + + ``source`` is a path or binary bytes. WAT conversion is intentionally not a + runtime dependency; tests use Wasmtime's assembler. + """ + + def __init__(self, source: str | os.PathLike[str] | Blob, *, files: Mapping[str, Blob] | None = None) -> None: + if isinstance(source, (str, os.PathLike)): + data = Path(source).read_bytes() + elif isinstance(source, (bytes, bytearray, memoryview)): + data = bytes(source) + else: + raise TypeError('Module requires a filesystem path or WASM bytes') + self._binary = BinaryModule(data) + Wasi.validate_imports(self._binary) + self._program = _program(self._binary) + self._files = Wasi(files=files).files + + @property + def exports(self) -> dict[str, str]: + """Export names mapped to kind names.""" + kinds = ('function', 'table', 'memory', 'global') + return {name: kinds[kind] for name, (kind, _) in self._binary.exports.items()} + + def spawn(self, count: int, *, memory_pages: int | None = None, table_elements: int | None = None, # noqa: PLR0913 - Explicit independent resource budgets. + stack_size: int = 256, call_depth: int = 64, fuel: int = 10_000_000, quantum: int = 4096, + batch_size: int | None = None, max_resident_bytes: int = 512 * 1024 * 1024, + wasi: Wasi | None = None) -> Instances: + """Create isolated, persistent instances on a hardware GPU. + + memory_pages/table_elements bound growth, and do not alter initial sizes. + The resident allocation budget is checked before any GPU allocation. + """ + return Instances(self, count, memory_pages=memory_pages, table_elements=table_elements, + stack_size=stack_size, call_depth=call_depth, fuel=fuel, quantum=quantum, + batch_size=batch_size, max_resident_bytes=max_resident_bytes, wasi=wasi) + + +class _Batch: + def __init__(self, owner: Instances, count: int, first: int) -> None: + self.owner, self.count, self.first = owner, count, first + self.context = owner._context + device = self.context.device + binary = owner.module._binary + words = array('I', [0]) * (owner._heap_words * count) + + def repeat(offset: int, value: int) -> None: + words[offset * count:(offset + 1) * count] = array('I', [value]) * count + + for index, (_, _, value) in enumerate(binary.globals): + repeat(index * 2, value & 0xFFFFFFFF) + repeat(index * 2 + 1, value >> 32) + initial_memory = bytearray((binary.memory[0] if binary.memory else 0) * 65536) + for index, (offset, blob) in enumerate(binary.data): + if offset is not None: + if offset > len(initial_memory) or len(blob) > len(initial_memory) - offset: + raise Trap({first: 'out of bounds memory access during instantiation'}, []) + initial_memory[offset:offset + len(blob)] = blob + repeat(owner._data_flags + index, 1) + for index, value in enumerate(_words(initial_memory)): + if value: + repeat(owner._memory_offset + index, value) + for index, (_, (initial_size, _)) in enumerate(binary.tables): + repeat(owner._table_lengths + index, initial_size) + for index, (offset, elements, declarative, _, table) in enumerate(binary.elements): + if offset is not None: + initial_size = binary.tables[table][1][0] + if offset > initial_size or len(elements) > initial_size - offset: + raise Trap({first: 'out of bounds table access during instantiation'}, []) + for element_index, value in enumerate(elements): + repeat(owner._table_offsets[table] + offset + element_index, value) + if offset is not None or declarative: + repeat(owner._element_flags + index, 1) + for offset, value in enumerate(owner._filesystem): + if value: + repeat(owner._fs_offset + offset, value) + if owner._filesystem: + for lane in range(count): + words[(owner._fs_offset + 18) * count + lane] = first + lane + self.buffers: list[Buffer] = [] + try: + assert owner._program_buffer is not None + self.buffers = [owner._program_buffer] + config = array('I', [count, owner.stack_size, owner.call_depth, owner.memory_pages, + owner._memory_offset, owner._table_offset, owner.table_elements, owner._data_flags, + owner._element_flags, owner.quantum, owner._output_slots, 0, + owner._fs_offset, owner._wasi.max_files, owner._wasi.max_fds, owner._wasi.storage_size]) + self.buffers.append(self.context.buffer(data=config, uniform=True)) + self.buffers.append(self.context.buffer(size=count * owner.stack_size * 8)) + self.buffers.append(self.context.buffer(data=words)) + self.buffers.append(self.context.buffer(size=count * owner.call_depth * 16)) + state = array('I', [0]) * (count * 12) + for lane in range(count): + state[lane * 12 + 5] = binary.memory[0] if binary.memory else 0 + state[lane * 12 + 6] = 0 + self.buffers.append(self.context.buffer(data=state)) + self.buffers.append(self.context.buffer(size=count * owner._output_slots * 8)) + self.state = state + self.group = device.create_bind_group(layout=self.context.pipeline.get_bind_group_layout(0), entries=[ + {'binding': index, 'resource': {'buffer': buffer, 'offset': 0, 'size': buffer.size}} + for index, buffer in enumerate(self.buffers) + ]) + except Exception: + self.close() + raise + + def close(self) -> None: + for buffer in self.buffers[1:]: + buffer.destroy() + self.buffers = [] + + def execute(self, function_index: int, inputs: Sequence[tuple[int, ...]], fuel: int, raw: bool = False) -> tuple[list[Result], dict[int, str]]: + owner = self.owner + device = self.context.device + fn = owner.module._binary.functions[function_index] + _, returns = owner.module._binary.signature(function_index) + arguments = array('Q', [0]) * (self.count * max(1, len(fn.locals))) + for lane, row in enumerate(inputs): + for index, value in enumerate(row): + arguments[index * self.count + lane] = value + base = lane * 12 + pages, table_len = self.state[base + 5], self.state[base + 6] + self.state[base:base + 12] = array('I', [fn.offset, len(fn.locals), 0, function_index, 0, pages, table_len, 0, 0, fuel, 0, 0]) + device.queue.write_buffer(self.buffers[2], 0, arguments) + device.queue.write_buffer(self.buffers[5], 0, self.state) + while True: + encoder = device.create_command_encoder(label='wasmgpu dispatch') + compute = encoder.begin_compute_pass() + compute.set_pipeline(self.context.pipeline) + compute.set_bind_group(0, self.group) + compute.dispatch_workgroups((self.count + 63) // 64) + compute.end() + device.queue.submit([encoder.finish()]) + self.state = array('I') + self.state.frombytes(device.queue.read_buffer(self.buffers[5])) + statuses = self.state[7::12] + if all(status in (1, 2) for status in statuses): + break + output = array('Q') + output.frombytes(device.queue.read_buffer(self.buffers[6])) + results: list[Result] = [] + traps: dict[int, str] = {} + for lane in range(self.count): + if self.state[lane * 12 + 7] == 2: + code = self.state[lane * 12 + 8] + if code == 12: + owner.exit_codes[self.first + lane] = self.state[lane * 12 + 10] + traps[self.first + lane] = _TRAPS[code] + results.append(None) + else: + output_row = tuple(output[index * self.count + lane] if raw else _decode(output[index * self.count + lane], ty) for index, ty in enumerate(returns)) + results.append(output_row[0] if len(output_row) == 1 else output_row if output_row else None) + return results, traps + + + +class Instances: + """A collection of isolated GPU WASM instances. Calls preserve memory/state.""" + + def __init__(self, module: Module, count: int, *, memory_pages: int | None, table_elements: int | None, # noqa: PLR0913, PLR0915 - Checked instance allocation. + stack_size: int, call_depth: int, fuel: int, quantum: int, batch_size: int | None, + max_resident_bytes: int, wasi: Wasi | None) -> None: + self.module = module + self.count = _positive(count, 'count', zero=True) + self.stack_size = _positive(stack_size, 'stack_size') + self.call_depth = _positive(call_depth, 'call_depth') + self.fuel = _positive(fuel, 'fuel') + self.quantum = _positive(quantum, 'quantum') + self._closed = False + self._lock = threading.RLock() + self._batches: list[_Batch] = [] + self._program_buffer: Buffer | None = None + binary = module._binary + initial_pages, declared_pages = binary.memory or (0, 0) + self.memory_pages = min(declared_pages if declared_pages is not None else 65535, max(initial_pages, 16)) if memory_pages is None else _positive(memory_pages, 'memory_pages', zero=True) + if self.memory_pages < initial_pages or self.memory_pages > 65535 or (declared_pages is not None and self.memory_pages > declared_pages): + raise ValueError('memory_pages must be within the declared memory limits and at most 65535') + if table_elements is not None: + _positive(table_elements, 'table_elements', zero=True) + self._table_capacities = [min(maximum if maximum is not None else 0xFFFFFFFF, + max(minimum, 256) if table_elements is None else table_elements) + for _, (minimum, maximum) in binary.tables] + if any(capacity < limits[0] for capacity, (_, limits) in zip(self._table_capacities, binary.tables)): + raise ValueError('table_elements is below a declared table minimum') + self.table_elements = sum(self._table_capacities) + self._table_lengths = len(binary.globals) * 2 + self._table_offset = self._table_lengths + len(binary.tables) + self._table_offsets = [] + next_offset = self._table_offset + for capacity in self._table_capacities: + self._table_offsets.append(next_offset) + next_offset += capacity + self._data_flags = self._table_offset + self.table_elements + self._element_flags = self._data_flags + len(binary.data) + self._memory_offset = self._element_flags + len(binary.elements) + self._fs_offset = self._memory_offset + self.memory_pages * 16384 + self._wasi = copy.deepcopy(wasi) if wasi is not None else Wasi() + if not isinstance(self._wasi, Wasi): + raise TypeError('wasi must be a Wasi configuration') + uses_wasi = any(fn.imported is not None for fn in binary.functions) or bool(module._files) or bool(self._wasi.files) + filesystem_words = 48 + self._wasi.max_files * 76 + self._wasi.max_fds * 8 + (self._wasi.storage_size + 3) // 4 if uses_wasi else 0 + self._heap_words = max(4, self._fs_offset + filesystem_words) + self._output_slots = max([8] + [max(len(args), len(results)) for args, results in binary.types]) + for fn in binary.functions: + if len(fn.locals) + fn.max_stack > self.stack_size: + raise ResourceLimitError('stack_size is smaller than a function requires') + per_instance = self.stack_size * 8 + self.call_depth * 16 + self._heap_words * 4 + 48 + self._output_slots * 8 + if isinstance(max_resident_bytes, bool) or not isinstance(max_resident_bytes, int): + raise TypeError('max_resident_bytes must be an integer') + if max_resident_bytes <= 0: + raise ValueError('max_resident_bytes must be positive') + budget = max_resident_bytes + environment_strings = [*self._wasi.args, *(f'{key}={value}' for key, value in self._wasi.env.items())] + program_bytes = 4 * (len(module._program) + len(binary.tables) * 4 + sum(2 + (len(value.encode()) + 4) // 4 for value in environment_strings)) + self.resident_bytes = per_instance * self.count + program_bytes + if self.resident_bytes > budget: + raise ResourceLimitError(f'{self.count} instances require approximately {self.resident_bytes} bytes; budget is {budget}') + if uses_wasi: + entries = self._wasi._initial(module._files) + if sum((len(content) + 3) & ~3 for _, _, content in entries) > self._wasi.storage_size: + raise ResourceLimitError('embedded files exceed WASI storage_size') + self.exit_codes: list[int | None] = [None] * count + self._context = _context() + limits = self._context.limits + if program_bytes > min(limits['max-storage-buffer-binding-size'], limits['max-buffer-size']): + raise ResourceLimitError('program and guest environment exceed the device buffer limit') + largest = max(self.stack_size * 8, self.call_depth * 16, self._heap_words * 4, 48, self._output_slots * 8) + max_batch = min(limits['max-storage-buffer-binding-size'] // largest, + limits['max-compute-workgroups-per-dimension'] * 64, + (128 * 1024 * 1024) // per_instance) + if max_batch < 1: + raise ResourceLimitError('a single instance exceeds device buffer limits') + self.batch_size = max_batch if batch_size is None else _positive(batch_size, 'batch_size') + if self.batch_size > max_batch: + raise ResourceLimitError(f'batch_size exceeds device/runtime limit {max_batch}') + self.resident_bytes += 64 * ((count + self.batch_size - 1) // self.batch_size) + if self.resident_bytes > budget: + raise ResourceLimitError('batch configuration buffers exceed the resident allocation budget') + self._filesystem = self._build_filesystem() if uses_wasi else array('I') + try: + program = array('I', module._program) + program[12] = int(bool(self._filesystem)) + program[4] = len(program) + for index, (offset, capacity) in enumerate(zip(self._table_offsets, self._table_capacities)): + program.extend([self._table_lengths + index, offset, capacity, 0]) + for slot, strings in ((8, self._wasi.args), (10, [f'{key}={value}' for key, value in self._wasi.env.items()])): + program[slot], program[slot + 1] = len(program), len(strings) + info_offset = len(program) + program.extend([0] * (len(strings) * 2)) + for index, string in enumerate(strings): + blob = string.encode() + b'\0' + program[info_offset + index * 2:info_offset + index * 2 + 2] = array('I', [len(program), len(blob)]) + program.extend(_words(blob + b'\0' * (-len(blob) % 4))) + self._program_buffer = self._context.buffer(data=program) + for first in range(0, count, self.batch_size): + self._batches.append(_Batch(self, min(self.batch_size, count - first), first)) + if binary.start is not None: + self._call(binary.start, [()] * count, self.fuel) + except Exception: + self.close() + raise + + @property + def adapter_info(self) -> AdapterInfo: + assert self._context.adapter is not None + return dict(self._context.adapter.info) + + @property + def stdout(self) -> list[bytes]: + return [self._read_file_index(1, instance) for instance in range(self.count)] + + @property + def stderr(self) -> list[bytes]: + return [self._read_file_index(2, instance) for instance in range(self.count)] + + def _build_filesystem(self) -> array[int]: + config = self._wasi + entries = config._initial(self.module._files) + names_offset = 48 + config.max_files * 12 + fds_offset = names_offset + config.max_files * 64 + data_offset = fds_offset + config.max_fds * 8 + words = array('I', [0]) * (data_offset + (config.storage_size + 3) // 4) + cursor = 0 + for index, (kind, name, content) in enumerate(entries): + capacity = (len(content) + 3) & ~3 + encoded = name.encode() + words[48 + index * 12:48 + index * 12 + 6] = array('I', [kind, len(encoded), len(content), cursor, capacity, index]) + for start, blob in ((names_offset + index * 64, encoded), (data_offset + cursor // 4, content)): + packed = blob + b'\0' * (-len(blob) % 4) + values = array('I') + values.frombytes(packed) + words[start:start + len(values)] = values + cursor += capacity + words[0] = cursor + words[1] = config.clock_epoch_ns & 0xFFFFFFFF + words[2] = config.clock_epoch_ns >> 32 + words[8:16] = _words(config.random_key) + words[37] = 64 + words[4] = config.clock_resolution_ns + for fd in range(4): + rights = (1 << 1) | (1 << 21) | (1 << 27) if fd == 0 else (1 << 6) | (1 << 21) | (1 << 27) if fd in (1, 2) else (1 << 30) - 1 + words[fds_offset + fd * 8:fds_offset + (fd + 1) * 8] = array('I', [fd + 1, 0, 0, 0, rights, 0, (1 << 30) - 1, 0]) + return words + + def _read_file_index(self, index: int, instance: int) -> bytes: + with self._lock: + self._check_open() + if not self._filesystem: + return b'' + batch, lane = self._locate(instance) + raw = self._context.device.queue.read_buffer(batch.buffers[3]) + words = _words(raw)[lane::batch.count] + entry = self._fs_offset + 48 + index * 12 + inode = words[entry + 5] + entry = self._fs_offset + 48 + inode * 12 + length, start = words[entry + 2], words[entry + 3] + data_offset = self._fs_offset + 48 + self._wasi.max_files * 76 + self._wasi.max_fds * 8 + content = array('I', words[data_offset + start // 4:data_offset + (start + length + 3) // 4]).tobytes() + return content[start % 4:start % 4 + length] + + def read_file(self, path: str, *, instance: int = 0) -> bytes: + """Read a file from one instance's GPU filesystem.""" + with self._lock: + self._check_open() + encoded_path = normalize_path(path).encode() + batch, lane = self._locate(instance) + if self._filesystem: + raw = self._context.device.queue.read_buffer(batch.buffers[3]) + words = _words(raw)[lane::batch.count] + names = self._fs_offset + 48 + self._wasi.max_files * 12 + for index in range(4, self._wasi.max_files): + entry = self._fs_offset + 48 + index * 12 + name = array('I', words[names + index * 64:names + (index + 1) * 64]).tobytes()[:words[entry + 1]] + if words[entry] == 1 and name == encoded_path: + return self._read_file_index(index, instance) + raise FileNotFoundError(path) + + def __len__(self) -> int: + return self.count + + def _check_open(self) -> None: + if self._closed: + raise RuntimeError('instances are closed') + + def call(self, name: str, inputs: Iterable[Row] | None = None, *, fuel: int | None = None) -> list[Result]: + """Invoke one export per instance, in input order. + + A scalar per instance is accepted for a single argument. For multiple + arguments use tuples; no-argument functions accept omitted inputs. + The whole batch is validated before executing any instance. + """ + with self._lock: + self._check_open() + binary = self.module._binary + if name not in binary.exports or binary.exports[name][0] != 0: + raise KeyError(f'no exported function {name!r}') + index = binary.exports[name][1] + params, _ = binary.signature(index) + if inputs is None: + if params: + raise TypeError('inputs are required for a function with parameters') + rows: list[Row] = [()] * self.count + else: + rows = list(inputs) + if len(rows) != self.count: + raise ValueError(f'expected {self.count} input rows, got {len(rows)}') + encoded: list[tuple[int, ...]] = [] + for row in rows: + normalized = (row,) if len(params) == 1 and not isinstance(row, (tuple, list)) else row + if not isinstance(normalized, (tuple, list)) or len(normalized) != len(params): + raise TypeError(f'each input row must contain {len(params)} arguments') + encoded.append(tuple(_encode(value, ty) for value, ty in zip(cast(Sequence[Scalar], normalized), params))) + return self._call(index, encoded, self.fuel if fuel is None else _positive(fuel, 'fuel')) + + def _call(self, index: int, inputs: Sequence[tuple[int, ...]], fuel: int, raw: bool = False) -> list[Result]: + results: list[Result] = [] + traps: dict[int, str] = {} + for batch in self._batches: + output, errors = batch.execute(index, inputs[batch.first:batch.first + batch.count], fuel, raw=raw) + results.extend(output) + traps.update(errors) + if traps: + raise Trap(traps, results) + return results + + def _locate(self, instance: int) -> tuple[_Batch, int]: + _positive(instance, 'instance', zero=True) + if instance >= self.count: + raise IndexError('instance index out of range') + batch = self._batches[instance // self.batch_size] + return batch, instance - batch.first + + def read_memory(self, offset: int, size: int, *, instance: int = 0) -> bytes: + """Read bytes from an instance's current linear memory.""" + with self._lock: + self._check_open() + _positive(offset, 'offset', zero=True) + _positive(size, 'size', zero=True) + batch, lane = self._locate(instance) + length = batch.state[lane * 12 + 5] * 65536 + if offset > length or size > length - offset: + raise IndexError('memory range out of bounds') + if size == 0: + return b'' + start, end = offset // 4, (offset + size + 3) // 4 + data = self._context.device.queue.read_buffer(batch.buffers[3], (self._memory_offset + start) * batch.count * 4, (end - start) * batch.count * 4) + words = array('I') + words.frombytes(data) + return words[lane::batch.count].tobytes()[offset % 4:offset % 4 + size] + + def write_memory(self, offset: int, data: Blob, *, instance: int = 0) -> None: + """Write bytes without executing WASM on the host.""" + with self._lock: + self._check_open() + data = bytes(data) + self.read_memory(offset, len(data), instance=instance) + if not data: + return + batch, lane = self._locate(instance) + start, end = offset // 4, (offset + len(data) + 3) // 4 + address = (self._memory_offset + start) * batch.count * 4 + words = array('I') + words.frombytes(self._context.device.queue.read_buffer(batch.buffers[3], address, (end - start) * batch.count * 4)) + existing = bytearray(words[lane::batch.count].tobytes()) + existing[offset % 4:offset % 4 + len(data)] = data + words[lane::batch.count] = _words(existing) + self._context.device.queue.write_buffer(batch.buffers[3], address, words) + + def close(self) -> None: + with self._lock: + if not self._closed: + for batch in self._batches: + batch.close() + if self._program_buffer is not None: + self._program_buffer.destroy() + self._closed = True + + def __enter__(self) -> Instances: + self._check_open() + return self + + def __exit__(self, *_: object) -> None: + self.close() diff --git a/wasmgpu/types.py b/wasmgpu/types.py new file mode 100644 index 0000000..0633a77 --- /dev/null +++ b/wasmgpu/types.py @@ -0,0 +1,85 @@ +"""Value types and the narrow wgpu interface used by the runtime. + +wgpu releases supported on Python 3.8+ do not publish typing information. These +protocols describe the common API, including the pre-0.19 synchronous spelling. +""" +from __future__ import annotations + +from array import array +from typing import Mapping, Protocol, Sequence, Tuple, Union + +Blob = Union[bytes, bytearray, memoryview] +Scalar = Union[int, float, None] +Result = Union[Scalar, Tuple[Scalar, ...]] +Row = Union[Scalar, Sequence[Scalar]] +AdapterInfo = Mapping[str, Union[str, int]] + + +class Buffer(Protocol): + size: int + def destroy(self) -> None: ... + + +class Queue(Protocol): + def write_buffer(self, buffer: Buffer, buffer_offset: int, data: Blob | array[int]) -> None: ... + def read_buffer(self, buffer: Buffer, buffer_offset: int = 0, size: int | None = None) -> memoryview: ... + def submit(self, command_buffers: Sequence[object]) -> None: ... + + +class Pipeline(Protocol): + def get_bind_group_layout(self, index: int) -> object: ... + + +class ComputePass(Protocol): + def set_pipeline(self, pipeline: Pipeline) -> None: ... + def set_bind_group(self, index: int, bind_group: object) -> None: ... + def dispatch_workgroups(self, workgroup_count_x: int) -> None: ... + def end(self) -> None: ... + + +class Encoder(Protocol): + def begin_compute_pass(self) -> ComputePass: ... + def finish(self) -> object: ... + + +class Device(Protocol): + queue: Queue + limits: Mapping[str, int] + + def create_buffer(self, *, size: int, usage: int) -> Buffer: ... + def create_buffer_with_data(self, *, data: Blob | array[int], usage: int) -> Buffer: ... + def create_shader_module(self, *, label: str, code: str) -> object: ... + def create_compute_pipeline(self, *, layout: str, compute: Mapping[str, object]) -> Pipeline: ... + def create_bind_group(self, *, layout: object, entries: Sequence[Mapping[str, object]]) -> object: ... + def create_command_encoder(self, *, label: str) -> Encoder: ... + + +class Adapter(Protocol): + info: AdapterInfo + limits: Mapping[str, int] + + def request_device(self, *, required_limits: Mapping[str, int]) -> Device: ... + + +class RequestDevice(Protocol): + def __call__(self, *, required_limits: Mapping[str, int]) -> Device: ... + + +class RequestAdapter(Protocol): + def __call__(self, *, power_preference: str) -> Adapter | None: ... + + +class BufferUsage(Protocol): + COPY_DST: int + COPY_SRC: int + UNIFORM: int + STORAGE: int + + +class GPU(Protocol): + def request_adapter(self, *, power_preference: str) -> Adapter | None: ... + + +class Backend(Protocol): + gpu: GPU + BufferUsage: BufferUsage diff --git a/wasmgpu/vm.wgsl b/wasmgpu/vm.wgsl new file mode 100644 index 0000000..7ae7ba3 --- /dev/null +++ b/wasmgpu/vm.wgsl @@ -0,0 +1,222 @@ +struct Config { + count: u32, stack_cap: u32, frame_cap: u32, memory_cap: u32, + memory_offset: u32, table_offset: u32, table_cap: u32, data_flags: u32, + element_flags: u32, quantum: u32, output_slots: u32, reserved: u32, + fs_offset: u32, fs_files: u32, fs_fds: u32, fs_bytes: u32, +} +struct State { + pc: u32, sp: u32, base: u32, function_id: u32, + depth: u32, pages: u32, table_len: u32, status: u32, + trap: u32, fuel: u32, exit_code: u32, reserved: u32, +} +@group(0) @binding(0) var program: array; +@group(0) @binding(1) var config: Config; +@group(0) @binding(2) var values: array; +@group(0) @binding(3) var heap: array; +@group(0) @binding(4) var frames: array; +@group(0) @binding(5) var states: array; +@group(0) @binding(6) var output: array; +var lane: u32; +var vm: State; +fn fail(code: u32) { vm.trap = code; vm.status = 2u; } +fn get_value(index: u32) -> vec2u { return values[index * config.count + lane]; } +fn set_value(index: u32, value: vec2u) { values[index * config.count + lane] = value; } +fn push(value: vec2u) { + if vm.sp >= config.stack_cap { fail(7u); return; } + set_value(vm.sp, value); vm.sp += 1u; +} +fn pop() -> vec2u { + if vm.sp == 0u { fail(11u); return vec2u(0u); } + vm.sp -= 1u; return get_value(vm.sp); +} +fn read_heap(index: u32) -> u32 { return heap[index * config.count + lane]; } +fn write_heap(index: u32, value: u32) { heap[index * config.count + lane] = value; } +fn bounds(address: u32, size: u32) -> bool { + // memory_cap is limited to 65535 pages by the host (u32 byte addressing). + let length = vm.pages * 65536u; + return address <= length && size <= length - address; +} +fn read_byte(address: u32) -> u32 { + return (read_heap(config.memory_offset + address / 4u) >> ((address & 3u) * 8u)) & 255u; +} +fn write_byte(address: u32, value: u32) { + let index = config.memory_offset + address / 4u; let shift = (address & 3u) * 8u; + write_heap(index, (read_heap(index) & ~(255u << shift)) | ((value & 255u) << shift)); +} +fn branch(destination_pc: u32, height: u32, arity: u32) { + let metadata = 16u + vm.function_id * 8u; + let destination = vm.base + program[metadata + 3u] + height; + for (var i = 0u; i < arity; i += 1u) { set_value(destination + i, get_value(vm.sp - arity + i)); } + vm.sp = destination + arity; vm.pc = destination_pc; +} +fn invoke(function_id: u32) { + if function_id >= program[3] { fail(3u); return; } + let metadata = 16u + function_id * 8u; + let params = program[metadata + 1u]; let locals = program[metadata + 3u]; + let base = vm.sp - params; + if program[metadata + 5u] != 0u { + let syscall = program[metadata + 5u]; + let answer = wasi_dispatch(syscall, base); + vm.sp = base; + if program[metadata + 2u] != 0u { push(vec2u(answer, 0u)); } + return; + } + if vm.depth >= config.frame_cap || base + locals > config.stack_cap { fail(7u); return; } + frames[vm.depth * config.count + lane] = vec4u(vm.pc, vm.base, vm.function_id, base); + vm.depth += 1u; vm.base = base; vm.function_id = function_id; + for (var i = params; i < locals; i += 1u) { set_value(base + i, vec2u(0u)); } + vm.sp = base + locals; vm.pc = program[metadata]; +} +fn return_function() { + let result_count = program[16u + vm.function_id * 8u + 2u]; + if vm.depth == 0u { + for (var i = 0u; i < result_count; i += 1u) { output[i * config.count + lane] = get_value(vm.sp - result_count + i); } + vm.status = 1u; return; + } + vm.depth -= 1u; + let frame = frames[vm.depth * config.count + lane]; + for (var i = 0u; i < result_count; i += 1u) { set_value(frame.w + i, get_value(vm.sp - result_count + i)); } + vm.sp = frame.w + result_count; vm.pc = frame.x; vm.base = frame.y; vm.function_id = frame.z; +} +fn memory_operation(op: u32, offset: u32, width: u32) { + let store = op >= 0x36u; + var value = vec2u(0u); if store { value = pop(); } + let pointer = pop().x; + let address = pointer + offset; + if address < pointer || !bounds(address, width) { fail(2u); return; } + if store { + for (var i = 0u; i < width; i += 1u) { write_byte(address + i, shr64(value, i * 8u).x); } + } else { + for (var i = 0u; i < width; i += 1u) { value |= shl64(vec2u(read_byte(address + i), 0u), i * 8u); } + let is_signed = op == 0x2cu || op == 0x2eu || op == 0x30u || op == 0x32u || op == 0x34u; + if is_signed { value = sar64(shl64(value, 64u - width * 8u), 64u - width * 8u); } + if op == 0x28u || op == 0x2au || (op >= 0x2cu && op <= 0x2fu) { value.y = 0u; } + push(value); + } +} +fn table_size(table: u32) -> u32 { return read_heap(program[program[4] + table * 4u]); } +fn table_offset(table: u32) -> u32 { return program[program[4] + table * 4u + 1u]; } +fn bulk_operation(op: u32, index: u32, other: u32) { + if op == 9u { write_heap(config.data_flags + index, 1u); return; } + if op == 13u { write_heap(config.element_flags + index, 1u); return; } + if op == 16u { push(vec2u(table_size(index), 0u)); return; } + if op == 15u { + let delta = pop().x; let value = pop().x; let old = table_size(index); + if delta > program[program[4] + index * 4u + 2u] - old { push(vec2u(0xffffffffu, 0u)); return; } + if !fs_charge(delta) { return; } + write_heap(program[program[4] + index * 4u], old + delta); + for (var i = old; i < old + delta; i += 1u) { write_heap(table_offset(index) + i, value); } + push(vec2u(old, 0u)); return; + } + let size = pop().x; let source = pop().x; let destination = pop().x; + let table = op == 12u || op == 14u || op == 17u; + if !fs_charge(size) { return; } + var length = vm.pages * 65536u; var offset = 0u; + if table { let id = select(index, other, op == 12u); length = table_size(id); offset = table_offset(id); } + if destination > length || size > length - destination { fail(select(2u, 3u, table)); return; } + if op == 8u || op == 12u { + let info = program[select(1u, 2u, table)] + index * 2u; + let data_offset = program[info]; + let dropped = read_heap(select(config.data_flags, config.element_flags, table) + index) != 0u; + let data_length = select(program[info + 1u], 0u, dropped); + if source > data_length || size > data_length - source { fail(select(2u, 3u, table)); return; } + for (var i = 0u; i < size; i += 1u) { + if table { write_heap(offset + destination + i, program[data_offset + source + i]); } + else { let address = source + i; write_byte(destination + i, (program[data_offset + address / 4u] >> ((address & 3u) * 8u)) & 255u); } + } + } else if op == 11u || op == 17u { + for (var i = 0u; i < size; i += 1u) { + if table { write_heap(offset + destination + i, source); } + else { write_byte(destination + i, source); } + } + } else { + var source_offset = 0u; + if table { length = table_size(other); source_offset = table_offset(other); } + if source > length || size > length - source { fail(select(2u, 3u, table)); return; } + for (var step = 0u; step < size; step += 1u) { + let i = select(step, size - 1u - step, destination > source && (!table || index == other)); + if table { write_heap(offset + destination + i, read_heap(source_offset + source + i)); } + else { write_byte(destination + i, read_byte(source + i)); } + } + } +} +@compute @workgroup_size(64) +fn run(@builtin(global_invocation_id) id: vec3u) { + lane = id.x; + if lane >= config.count { return; } + vm = states[lane]; + if vm.status != 0u { return; } + for (var step = 0u; step < config.quantum; step += 1u) { + if vm.fuel == 0u { fail(8u); break; } + vm.fuel -= 1u; + if program[12] != 0u { + let clock = add64(vec2u(read_heap(config.fs_offset + 1u), read_heap(config.fs_offset + 2u)), vec2u(read_heap(config.fs_offset + 4u), 0u)); + write_heap(config.fs_offset + 1u, clock.x); write_heap(config.fs_offset + 2u, clock.y); + } + let location = program[0] + vm.pc * 4u; + let op = program[location]; let a = program[location + 1u]; + let b = program[location + 2u]; let c = program[location + 3u]; + vm.pc += 1u; + switch op { + case 0u: { fail(1u); } + case 0x100u: { vm.pc = a; } + case 0x101u: { if pop().x == 0u { vm.pc = a; } } + case 0x102u: { branch(a, b, c); } + case 0x103u: { if pop().x != 0u { branch(a, b, c); } } + case 0x104u: { + let destination_pc = program[0] + (a + min(pop().x, b - 1u)) * 4u; + branch(program[destination_pc + 1u], program[destination_pc + 2u], program[destination_pc + 3u]); + } + case 0x0fu: { return_function(); } + case 0x10u: { invoke(a); } + case 0x11u: { + let index = pop().x; + if index >= table_size(b) { fail(3u); } + else { + let reference = read_heap(table_offset(b) + index); + if reference == 0u { fail(9u); } + else if program[16u + (reference - 1u) * 8u + 4u] != a { fail(10u); } + else { invoke(reference - 1u); } + } + } + case 0x1au: { let ignored = pop(); } + case 0x1bu: { let condition = pop().x; let rhs = pop(); let lhs = pop(); push(select(rhs, lhs, condition != 0u)); } + case 0x20u: { push(get_value(vm.base + a)); } + case 0x21u: { let value = pop(); set_value(vm.base + a, value); } + case 0x22u: { set_value(vm.base + a, get_value(vm.sp - 1u)); } + case 0x23u: { push(vec2u(read_heap(a * 2u), read_heap(a * 2u + 1u))); } + case 0x24u: { let value = pop(); write_heap(a * 2u, value.x); write_heap(a * 2u + 1u, value.y); } + case 0x25u: { + let index = pop().x; + if index >= table_size(a) { fail(3u); } else { push(vec2u(read_heap(table_offset(a) + index), 0u)); } + } + case 0x26u: { + let value = pop().x; let index = pop().x; + if index >= table_size(a) { fail(3u); } else { write_heap(table_offset(a) + index, value); } + } + case 0x3fu: { push(vec2u(vm.pages, 0u)); } + case 0x40u: { + let delta = pop().x; let old = vm.pages; + if delta > config.memory_cap - old { push(vec2u(0xffffffffu, 0u)); } + else { + vm.pages += delta; + // The complete reserved arena was zero-initialized; pages never shrink. + push(vec2u(old, 0u)); + } + } + case 0x41u, 0x42u, 0x43u, 0x44u: { push(vec2u(a, b)); } + default: { + if op >= 0x28u && op <= 0x3eu { memory_operation(op, a, b); } + else if op >= 0xfc08u && op <= 0xfc11u { bulk_operation(op - 0xfc00u, a, b); } + else { + var rhs = vec2u(0u); + if a == 2u { rhs = pop(); } + let lhs = pop(); let answer = numeric(op, lhs, rhs); + if answer.trap != 0u { fail(answer.trap); } else { push(answer.value); } + } + } + } + if vm.status != 0u { break; } + } + states[lane] = vm; +} diff --git a/wasmgpu/wasi.py b/wasmgpu/wasi.py new file mode 100644 index 0000000..6c3301c --- /dev/null +++ b/wasmgpu/wasi.py @@ -0,0 +1,131 @@ +"""Embedded WASI configuration. Files and their descriptors live on the GPU. + +Clocks, random data, streams and filesystem operations are emulated in the +compute shader. No guest operation calls an operating-system service. +""" +from __future__ import annotations + +from pathlib import PurePosixPath +from typing import Iterable, Mapping + +from .binary import I32, I64, BinaryModule +from .errors import UnsupportedFeatureError, ValidationError +from .types import Blob + +_SIGNATURES = { + 'args_get': ('ii', 'i'), 'args_sizes_get': ('ii', 'i'), + 'environ_get': ('ii', 'i'), 'environ_sizes_get': ('ii', 'i'), + 'clock_res_get': ('ii', 'i'), 'clock_time_get': ('iIi', 'i'), + 'fd_advise': ('iIIi', 'i'), 'fd_allocate': ('iII', 'i'), 'fd_close': ('i', 'i'), + 'fd_datasync': ('i', 'i'), 'fd_fdstat_get': ('ii', 'i'), 'fd_fdstat_set_flags': ('ii', 'i'), + 'fd_fdstat_set_rights': ('iII', 'i'), 'fd_filestat_get': ('ii', 'i'), + 'fd_filestat_set_size': ('iI', 'i'), 'fd_filestat_set_times': ('iIIi', 'i'), + 'fd_pread': ('iiiIi', 'i'), 'fd_prestat_get': ('ii', 'i'), 'fd_prestat_dir_name': ('iii', 'i'), + 'fd_pwrite': ('iiiIi', 'i'), 'fd_read': ('iiii', 'i'), 'fd_readdir': ('iiiIi', 'i'), + 'fd_renumber': ('ii', 'i'), 'fd_seek': ('iIii', 'i'), 'fd_sync': ('i', 'i'), + 'fd_tell': ('ii', 'i'), 'fd_write': ('iiii', 'i'), + 'path_create_directory': ('iii', 'i'), 'path_filestat_get': ('iiiii', 'i'), + 'path_filestat_set_times': ('iiiiIIi', 'i'), 'path_link': ('iiiiiii', 'i'), + 'path_open': ('iiiiiIIii', 'i'), 'path_readlink': ('iiiiii', 'i'), + 'path_remove_directory': ('iii', 'i'), 'path_rename': ('iiiiii', 'i'), + 'path_symlink': ('iiiii', 'i'), 'path_unlink_file': ('iii', 'i'), + 'poll_oneoff': ('iiii', 'i'), 'proc_exit': ('i', ''), 'proc_raise': ('i', 'i'), + 'random_get': ('ii', 'i'), 'sched_yield': ('', 'i'), + 'sock_accept': ('iii', 'i'), 'sock_recv': ('iiiiii', 'i'), + 'sock_send': ('iiiii', 'i'), 'sock_shutdown': ('ii', 'i'), +} + +WASI_IDS = {name: index + 1 for index, name in enumerate(_SIGNATURES)} + + +def normalize_path(path: str) -> str: + if not isinstance(path, str) or '\0' in path: + raise ValueError('embedded file names must be strings without NUL') + parts: list[str] = [] + for part in path.split('/'): + if part in ('', '.'): + continue + if part == '..': + if not parts: + raise ValueError('path escapes the instance filesystem') + parts.pop() + else: + parts.append(part) + result = '/'.join(parts) + if len(result.encode()) > 255: + raise ValueError('embedded paths are limited to 255 UTF-8 bytes') + return result + + +class Wasi: + """WASI Preview 1 with a private filesystem resident in GPU memory. + + files maps guest paths to bytes; it never refers to host paths. Each instance + receives its own copy. storage_size is the total file-data arena capacity. + The root directory is preopened as descriptor 3, named '.'. + """ + + def __init__(self, *, files: Mapping[str, Blob] | None = None, args: Iterable[str] = (), # noqa: PLR0913 - Guest environment and explicit quotas. + env: Mapping[str, str] | None = None, stdin: Blob = b'', storage_size: int = 65536, + max_files: int = 32, max_fds: int = 64, seed: int | bytes = 1, clock_epoch_ns: int = 0, + clock_resolution_ns: int = 1) -> None: + self.files: dict[str, bytes] = {} + for name, content in (files or {}).items(): + path = normalize_path(name) + if not path or path in self.files: + raise ValueError('empty or duplicate embedded file name') + if not isinstance(content, (bytes, bytearray, memoryview)): + raise TypeError('embedded file contents must be bytes, not host paths') + self.files[path] = bytes(content) + self.args = tuple(str(value) for value in args) + self.env = {str(key): str(value) for key, value in (env or {}).items()} + self.stdin = bytes(stdin) + if any('\0' in arg for arg in self.args) or any('\0' in key + value or '=' in key for key, value in self.env.items()): + raise ValueError('WASI arguments/environment contain invalid characters') + for name, value in [('storage_size', storage_size), ('max_files', max_files), ('max_fds', max_fds)]: + if isinstance(value, bool) or not isinstance(value, int) or not 4 <= value <= 0x7fffffff: + raise ValueError(f'{name} must be an integer between 4 and 2147483647') + self.storage_size, self.max_files, self.max_fds = storage_size, max_files, max_fds + if isinstance(seed, bytes): + if len(seed) != 32: + raise ValueError('a byte seed must contain exactly 32 bytes') + self.random_key = seed + else: + if isinstance(seed, bool) or not isinstance(seed, int) or not 0 < seed <= 0xFFFFFFFF: + raise ValueError('seed must be a nonzero u32 or 32 bytes') + self.random_key = seed.to_bytes(32, 'little') + if isinstance(clock_epoch_ns, bool) or not isinstance(clock_epoch_ns, int) or not 0 <= clock_epoch_ns < 1 << 64: + raise ValueError('clock_epoch_ns must fit u64') + if isinstance(clock_resolution_ns, bool) or not isinstance(clock_resolution_ns, int) or not 0 < clock_resolution_ns < 1 << 32: + raise ValueError('clock_resolution_ns must be a positive u32') + self.seed, self.clock_epoch_ns, self.clock_resolution_ns = seed, clock_epoch_ns, clock_resolution_ns + + @staticmethod + def validate_imports(module: BinaryModule) -> None: + for fn in module.functions: + if fn.imported is None: + continue + namespace, name = fn.imported + if namespace != 'wasi_snapshot_preview1' or name not in _SIGNATURES: + raise UnsupportedFeatureError(f'unsupported import {namespace}.{name}') + args, results = _SIGNATURES[name] + expected = ([I32 if char == 'i' else I64 for char in args], [I32 for _ in results]) + if module.types[fn.type_index] != expected: + raise ValidationError(f'incorrect WASI signature for {name}') + + def _initial(self, extra_files: Mapping[str, bytes]) -> list[tuple[int, str, bytes]]: + files = dict(extra_files) + files.update(self.files) + directories = {''} + for path in files: + for parent in PurePosixPath(path).parents: + if str(parent) != '.': + directories.add(str(parent)) + if directories.intersection(files): + raise ValueError('an embedded path is both a file and a directory') + entries = [(1, '', self.stdin), (1, '', b''), (1, '', b''), (2, '', b'')] + entries.extend((2, name, b'') for name in sorted(directories - {''})) + entries.extend((1, name, content) for name, content in sorted(files.items())) + if len(entries) > self.max_files: + raise ValueError('max_files is too small for the embedded files and directories') + return entries