diff --git a/.github/workflows/build-aot.yml b/.github/workflows/build-aot.yml new file mode 100644 index 0000000..b58f2ec --- /dev/null +++ b/.github/workflows/build-aot.yml @@ -0,0 +1,145 @@ +name: Build AOT + +# Traces a snapshot, translates its hot blocks, and emits the engine modules +# every other workflow consumes. See docs/design/release-strategy.md. +# +# The output is wasm32-wasip2 and therefore host-independent: one build serves +# all five platform targets. Running this inside build-binaries' matrix would +# pay the whole pipeline five times. + +on: + workflow_call: + inputs: + snapshot: + description: 'Snapshot id to trace against' + required: false + default: 'alpine-3.23.0-256mb' + type: string + workflow_dispatch: + inputs: + snapshot: + description: 'Snapshot id to trace against' + required: false + default: 'alpine-3.23.0-256mb' + type: string + +jobs: + build-aot: + name: Trace and translate + runs-on: ubuntu-latest + + steps: + - uses: actions/checkout@v4 + + - name: Resolve snapshot in registry + id: registry + env: + SNAPSHOT: ${{ inputs.snapshot }} + run: | + python3 - <<'PY' >> "$GITHUB_OUTPUT" + import json, os, urllib.request + + snap_id = os.environ["SNAPSHOT"] + registry = json.load(urllib.request.urlopen( + "https://registry.vpod.sh/v1/snapshots.json", timeout=60)) + entry = next((s for s in registry["snapshots"] if s["id"] == snap_id), None) + if entry is None: + raise SystemExit(f"snapshot '{snap_id}' not in registry") + print(f"sha256={entry['sha256']}") + print(f"url={entry['url']}") + PY + + - name: Cache AOT modules + id: cache + uses: actions/cache@v4 + with: + path: dist/aot + key: >- + aot-${{ inputs.snapshot }}-${{ steps.registry.outputs.sha256 }}-${{ hashFiles( + 'crates/riscv-core/**', + 'crates/machine/**', + 'crates/vpod-translate/**', + 'crates/wasi-component/**', + 'scripts/aot-snapshot.sh', + 'scripts/build-wasm.sh') }} + + - name: Install Rust + if: steps.cache.outputs.cache-hit != 'true' + uses: dtolnay/rust-toolchain@stable + with: + targets: wasm32-wasip2 + + - name: Cache cargo + if: steps.cache.outputs.cache-hit != 'true' + uses: actions/cache@v4 + with: + path: | + ~/.cargo/registry + ~/.cargo/git + target + key: aot-cargo-${{ hashFiles('**/Cargo.lock') }} + + - name: Fetch snapshot + if: steps.cache.outputs.cache-hit != 'true' + env: + SNAPSHOT: ${{ inputs.snapshot }} + SNAPSHOT_URL: ${{ steps.registry.outputs.url }} + SNAPSHOT_SHA256: ${{ steps.registry.outputs.sha256 }} + run: | + python3 -m pip install --quiet lz4 + python3 - <<'PY' + import hashlib, os, shutil, urllib.request + from pathlib import Path + import lz4.frame + + snap_id = os.environ["SNAPSHOT"] + url = os.environ["SNAPSHOT_URL"] + sha256 = os.environ["SNAPSHOT_SHA256"] + + dist = Path("dist"); dist.mkdir(exist_ok=True) + packed = dist / f"{snap_id}.snap.lz4" + with urllib.request.urlopen(url, timeout=300) as r, open(packed, "wb") as f: + shutil.copyfileobj(r, f) + + digest = hashlib.sha256(packed.read_bytes()).hexdigest() + if digest != sha256: + raise SystemExit(f"checksum mismatch: {digest} != {sha256}") + + out = dist / f"{snap_id}.snap" + with lz4.frame.open(str(packed), "rb") as src, open(out, "wb") as dst: + shutil.copyfileobj(src, dst) + packed.unlink() + print(f"{out} ({out.stat().st_size/1e6:.1f} MB)") + PY + + - name: Trace and translate + if: steps.cache.outputs.cache-hit != 'true' + run: ./scripts/aot-snapshot.sh "dist/${{ inputs.snapshot }}.snap" --force + + - name: Build engine modules + if: steps.cache.outputs.cache-hit != 'true' + run: ./scripts/build-wasm.sh + + - name: Collect + if: steps.cache.outputs.cache-hit != 'true' + run: | + mkdir -p dist/aot + cp sdks/python/vpod/vpod_wasi_lib_aot.wasm dist/aot/ + cp sdks/python/vpod/vpod_wasi_lib.wasm dist/aot/ + cp target/wasm32-wasip2/release/vpod-wasi-cli.wasm dist/aot/ + printf '%s %s\n' "${{ inputs.snapshot }}" "${{ steps.registry.outputs.sha256 }}" > dist/aot/TRACED_SNAPSHOT + ls -lh dist/aot + + - name: Verify + run: | + for f in vpod_wasi_lib_aot.wasm vpod_wasi_lib.wasm vpod-wasi-cli.wasm; do + test -s "dist/aot/$f" || { echo "missing dist/aot/$f" >&2; exit 1; } + done + + - name: Upload + uses: actions/upload-artifact@v4 + with: + name: aot-wasm + path: dist/aot + retention-days: 1 + overwrite: true diff --git a/.github/workflows/build-binaries.yml b/.github/workflows/build-binaries.yml index 4d2eaca..19cf23c 100644 --- a/.github/workflows/build-binaries.yml +++ b/.github/workflows/build-binaries.yml @@ -1,21 +1,41 @@ name: Build Binaries on: - release: - types: [published] + workflow_call: + inputs: + tag: + description: 'Release tag to upload binaries to (e.g. v0.1.3)' + required: true + type: string + snapshot: + description: 'Snapshot the AOT module is traced against' + required: false + default: 'alpine-3.23.0-256mb' + type: string workflow_dispatch: inputs: tag: description: 'Release tag to upload binaries to (e.g. v0.1.3)' + required: true + type: string + snapshot: + description: 'Snapshot the AOT module is traced against' required: false + default: 'alpine-3.23.0-256mb' type: string permissions: contents: write jobs: + build-aot: + uses: ./.github/workflows/build-aot.yml + with: + snapshot: ${{ inputs.snapshot }} + build: name: ${{ matrix.target }} + needs: build-aot runs-on: ${{ matrix.os }} strategy: fail-fast: false @@ -44,14 +64,14 @@ jobs: - name: Inject version shell: bash run: | - V="${{ github.event_name == 'release' && github.ref_name || inputs.tag }}" + V="${{ inputs.tag }}" V="${V#v}" perl -i -pe "s/^version = .*/version = \"$V\"/" crates/vpod/Cargo.toml - name: Install Rust toolchain uses: dtolnay/rust-toolchain@stable with: - targets: wasm32-wasip2,${{ matrix.target }} + targets: ${{ matrix.target }} - name: Install cargo-zigbuild (Linux aarch64) if: matrix.zigbuild @@ -60,8 +80,15 @@ jobs: cargo install cargo-zigbuild rustup target add aarch64-unknown-linux-gnu - - name: Build WASM component - run: cargo build --release --target wasm32-wasip2 -p wasi-component + - name: Fetch AOT engine module + uses: actions/download-artifact@v4 + with: + name: aot-wasm + path: dist/aot + + - name: Place module for build.rs + shell: bash + run: cp dist/aot/vpod-wasi-cli.wasm crates/vpod/vpod-wasi-cli.wasm - name: Build binary (native) if: "!matrix.zigbuild" @@ -92,7 +119,7 @@ jobs: - name: Upload to release uses: softprops/action-gh-release@v2 with: - tag_name: ${{ github.event_name == 'release' && github.ref_name || inputs.tag }} + tag_name: ${{ inputs.tag }} files: ${{ env.ASSET }} - name: Upload artifact diff --git a/.github/workflows/build-snapshots.yml b/.github/workflows/build-snapshots.yml new file mode 100644 index 0000000..d971f37 --- /dev/null +++ b/.github/workflows/build-snapshots.yml @@ -0,0 +1,180 @@ +name: Build Snapshots + +# Builds the registry snapshots and (optionally) publishes them to R2. +# See docs/design/release-strategy.md — this is "step 0": any snapshot change +# must be in the registry BEFORE build-aot traces against it, or the release +# ships an engine whose page hashes never match (silent interpreter fallback). +# +# Deliberately NOT run on every release: snapshot builds are not reproducible +# (apk pulls whatever Alpine serves that day), so every rebuild changes the +# bytes — and republishing changed bytes under the same ids silently degrades +# AOT for every previously shipped engine. Publish only when guest content +# actually changed, and always cut a release right after. +# +# The registry stores lz4-compressed files under the .snap name; sha256 and +# size in snapshots.json are of the *compressed* bytes (what the SDK and +# build-aot.yml verify before decompressing). + +on: + workflow_dispatch: + inputs: + version: + description: 'Registry manifest version (e.g. 0.5.0)' + required: true + type: string + publish: + description: 'Upload to R2 (off = dry run, artifacts only)' + required: false + default: false + type: boolean + workflow_call: + inputs: + version: + required: true + type: string + publish: + required: false + default: false + type: boolean + secrets: + R2_ACCOUNT_ID: + required: false + R2_ACCESS_KEY_ID: + required: false + R2_SECRET_ACCESS_KEY: + required: false + R2_BUCKET: + required: false + +jobs: + build-snapshots: + name: Build and publish snapshots + runs-on: ubuntu-latest + timeout-minutes: 300 + + steps: + - uses: actions/checkout@v4 + + - name: Install host tools + run: | + sudo apt-get update -q + sudo apt-get install -y -q libarchive-tools cpio lz4 + + - name: Install Zig + uses: mlugg/setup-zig@v1 + + - name: Install Rust + uses: dtolnay/rust-toolchain@stable + + - name: Cache cargo + uses: actions/cache@v4 + with: + path: | + ~/.cargo/registry + ~/.cargo/git + target + key: snapshots-cargo-${{ hashFiles('**/Cargo.lock') }} + + - name: Build default snapshot (alpine-3.23.0-256mb) + run: ./scripts/build-default-snapshot.sh + + - name: Build 512MB base snapshot (vsnap-base-512mb) + run: ./scripts/build-default-snapshot.sh --ram 512 + + - name: Build data snapshot (vsnap-data-512mb) + run: ./scripts/build-data-snapshot.sh + + - name: Compress for the registry + run: | + mkdir -p dist/registry + lz4 -9 -f dist/alpine-3.23.0-256mb.snap dist/registry/alpine-3.23.0-256mb.snap + cp dist/registry/alpine-3.23.0-256mb.snap dist/registry/vsnap-base-256mb.snap + lz4 -9 -f dist/alpine-3.23.0-512mb.snap dist/registry/vsnap-base-512mb.snap + lz4 -9 -f dist/vsnap-data-512mb.snap dist/registry/vsnap-data-512mb.snap + ls -lh dist/registry + + - name: Generate snapshots.json + env: + VERSION: ${{ inputs.version }} + run: | + python3 - <<'PY' + import hashlib, json, os, urllib.request + from pathlib import Path + + registry_dir = Path("dist/registry") + manifest = json.load(urllib.request.urlopen( + "https://registry.vpod.sh/v1/snapshots.json", timeout=60)) + + manifest["version"] = os.environ["VERSION"] + + # A snapshot id we built but the live manifest doesn't know yet is a + # new variant: take its metadata from the repo template, pointing the + # url at the production path. + template = json.loads( + Path("docs/v0.4.1-registry/test/snapshots.json").read_text()) + known_ids = {entry["id"] for entry in manifest["snapshots"]} + for entry in template["snapshots"]: + if entry["id"] in known_ids: + continue + if not (registry_dir / f"{entry['id']}.snap").exists(): + continue + entry["url"] = f"https://registry.vpod.sh/v1/{entry['id']}.snap" + manifest["snapshots"].append(entry) + print(f" {entry['id']}: new snapshot, added from template") + + for entry in manifest["snapshots"]: + built = registry_dir / f"{entry['id']}.snap" + if not built.exists(): + print(f" {entry['id']}: not rebuilt, keeping registry entry") + continue + data = built.read_bytes() + entry["sha256"] = hashlib.sha256(data).hexdigest() + entry["size"] = len(data) + print(f" {entry['id']}: sha256={entry['sha256'][:12]}… size={entry['size']}") + + (registry_dir / "snapshots.json").write_text( + json.dumps(manifest, indent=2) + "\n") + PY + + - name: Upload artifact + uses: actions/upload-artifact@v4 + with: + name: registry-snapshots + path: dist/registry/snapshots.json + retention-days: 30 + + - name: Publish to R2 + if: ${{ inputs.publish }} + env: + AWS_ACCESS_KEY_ID: ${{ secrets.R2_ACCESS_KEY_ID }} + AWS_SECRET_ACCESS_KEY: ${{ secrets.R2_SECRET_ACCESS_KEY }} + AWS_DEFAULT_REGION: auto + ENDPOINT: https://${{ secrets.R2_ACCOUNT_ID }}.r2.cloudflarestorage.com + BUCKET: ${{ secrets.R2_BUCKET }} + run: | + for f in dist/registry/*.snap; do + aws s3 cp "$f" "s3://$BUCKET/v1/$(basename "$f")" --endpoint-url "$ENDPOINT" + done + aws s3 cp dist/registry/snapshots.json "s3://$BUCKET/v1/snapshots.json" \ + --endpoint-url "$ENDPOINT" --content-type application/json + + - name: Verify the live registry + if: ${{ inputs.publish }} + run: | + python3 - <<'PY' + import hashlib, json, urllib.request + from pathlib import Path + + local = json.loads(Path("dist/registry/snapshots.json").read_text()) + live = json.load(urllib.request.urlopen( + "https://registry.vpod.sh/v1/snapshots.json", timeout=60)) + if live != local: + raise SystemExit("live snapshots.json does not match what we uploaded") + + for entry in local["snapshots"]: + with urllib.request.urlopen(entry["url"], timeout=300) as r: + digest = hashlib.sha256(r.read()).hexdigest() + if digest != entry["sha256"]: + raise SystemExit(f"{entry['id']}: served bytes do not match manifest") + print(f"{entry['id']}: ok") + PY diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 13e8f08..d8972b6 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -26,6 +26,9 @@ jobs: - name: Check formatting run: cargo fmt --all -- --check + - name: Generate AOT stub + run: ./scripts/aot-stub.sh + - name: Build WASM component run: cargo build --release --target wasm32-wasip2 -p wasi-component @@ -45,6 +48,9 @@ jobs: with: targets: wasm32-wasip2 + - name: Generate AOT stub + run: ./scripts/aot-stub.sh + - name: Build WASM component run: cargo build --release --target wasm32-wasip2 -p wasi-component @@ -65,6 +71,11 @@ jobs: - name: Install bsdtar run: sudo apt-get update && sudo apt-get install -y libarchive-tools + - name: Install Zig + uses: mlugg/setup-zig@v2 + with: + version: 0.14.0 + - name: Build Default Snapshot run: ./scripts/build-default-snapshot.sh @@ -81,6 +92,9 @@ jobs: with: targets: wasm32-wasip2 + - name: Generate AOT stub + run: ./scripts/aot-stub.sh + - name: Build WASM library run: cargo build --release --target wasm32-wasip2 -p wasi-component --lib diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 5e544eb..458310d 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -7,58 +7,72 @@ on: description: 'Version to publish (e.g. 0.2.0)' required: true type: string - publish_crate: - description: 'Publish vpod to crates.io' + snapshot: + description: 'Snapshot to trace the AOT module against' required: false - default: true - type: boolean + default: 'alpine-3.23.0-256mb' + type: string publish_python: description: 'Publish Python SDK to PyPI' required: false default: true type: boolean + build_snapshots: + description: 'Rebuild + publish registry snapshots first (only when guest content changed — republishing changes their sha256 and degrades AOT for older releases)' + required: false + default: false + type: boolean + +permissions: + contents: write jobs: - publish-crate: - name: Publish – crates.io (vpod) - if: ${{ inputs.publish_crate }} - runs-on: ubuntu-latest + build-snapshots: + if: ${{ inputs.build_snapshots }} + uses: ./.github/workflows/build-snapshots.yml + with: + version: ${{ inputs.version }} + publish: true + secrets: inherit + + build-aot: + needs: [build-snapshots] + if: ${{ !cancelled() && !failure() }} + uses: ./.github/workflows/build-aot.yml + with: + snapshot: ${{ inputs.snapshot }} + + # Must exist before build-binaries can attach assets to it. + create-release: + name: Create draft release + runs-on: ubuntu-latest steps: - - name: Checkout repo - uses: actions/checkout@v4 - - - name: Inject version - run: | - V="${{ inputs.version }}" - sed -i "s/^version = \".*\"/version = \"$V\"/" crates/vpod/Cargo.toml - - - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable - with: - targets: wasm32-wasip2 + - uses: actions/checkout@v4 - - name: Build WASM component (required by vpod build.rs) - run: | - cargo build --release --target wasm32-wasip2 -p wasi-component - cp target/wasm32-wasip2/release/vpod-wasi-cli.wasm crates/vpod/ - - - name: Cache Cargo registry - uses: actions/cache@v4 - with: - path: | - ~/.cargo/registry - ~/.cargo/git - key: ${{ runner.os }}-cargo-${{ hashFiles('**/Cargo.lock') }} - - - name: Publish to crates.io + - name: Create draft env: - CARGO_REGISTRY_TOKEN: ${{ secrets.CARGO_REGISTRY_TOKEN }} - run: cargo publish -p vpod --allow-dirty + GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} + TAG: v${{ inputs.version }} + run: | + if gh release view "$TAG" >/dev/null 2>&1; then + echo "release $TAG already exists — reusing it" + exit 0 + fi + gh release create "$TAG" --draft --generate-notes --title "$TAG" + + build-binaries: + needs: [build-aot, create-release] + uses: ./.github/workflows/build-binaries.yml + with: + tag: v${{ inputs.version }} + snapshot: ${{ inputs.snapshot }} publish-python: name: Publish – PyPI (Python SDK) - if: ${{ inputs.publish_python }} + # An explicit `if` replaces the implicit success() check, so guard it back. + if: ${{ !cancelled() && !failure() && inputs.publish_python }} + needs: build-aot runs-on: ubuntu-latest steps: @@ -71,15 +85,17 @@ jobs: sed -i "s/^version = \".*\"/version = \"$V\"/" sdks/python/pyproject.toml sed -i "s/^__version__ = \".*\"/__version__ = \"$V\"/" sdks/python/vpod/__init__.py - - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable + - name: Fetch AOT engine modules + uses: actions/download-artifact@v4 with: - targets: wasm32-wasip2 + name: aot-wasm + path: dist/aot - - name: Build WASM component for Python SDK + - name: Place modules in the SDK run: | - cargo build --release --target wasm32-wasip2 -p wasi-component --lib - cp target/wasm32-wasip2/release/vpod_wasi_lib.wasm sdks/python/vpod/ + cp dist/aot/vpod_wasi_lib_aot.wasm sdks/python/vpod/ + cp dist/aot/vpod_wasi_lib.wasm sdks/python/vpod/ + ls -lh sdks/python/vpod/*.wasm - name: Set up Python uses: actions/setup-python@v5 @@ -93,9 +109,42 @@ jobs: working-directory: sdks/python run: python -m build + - name: Verify the wheel carries both modules + working-directory: sdks/python + run: | + WHEEL=$(ls dist/*.whl) + python3 -c " + import sys, zipfile + names = zipfile.ZipFile('$WHEEL').namelist() + for want in ('vpod/vpod_wasi_lib_aot.wasm', 'vpod/vpod_wasi_lib.wasm'): + if want not in names: + sys.exit(f'{want} missing from the wheel') + print(want, 'ok') + " + - name: Publish to PyPI working-directory: sdks/python env: TWINE_USERNAME: __token__ TWINE_PASSWORD: ${{ secrets.PYPI_API_TOKEN }} run: twine upload dist/* + + attach-provenance: + name: Attach AOT provenance + needs: [build-aot, create-release] + runs-on: ubuntu-latest + steps: + - name: Fetch AOT artifact + uses: actions/download-artifact@v4 + with: + name: aot-wasm + path: dist/aot + + - name: Upload provenance to draft + env: + GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} + GH_REPO: ${{ github.repository }} + TAG: v${{ inputs.version }} + run: | + cp dist/aot/TRACED_SNAPSHOT aot-provenance.txt + gh release upload "$TAG" aot-provenance.txt --clobber diff --git a/.gitignore b/.gitignore index 5baf059..8750298 100644 --- a/.gitignore +++ b/.gitignore @@ -2,6 +2,7 @@ # will have compiled files and executables debug target +target-shadow # These are backup files generated by rustfmt **/*.rs.bk @@ -29,3 +30,4 @@ __pycache__/ *.wasm +crates/riscv-core/src/aot/generated.rs diff --git a/Cargo.lock b/Cargo.lock index 4f68d36..bb67b22 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -11,6 +11,41 @@ dependencies = [ "gimli", ] +[[package]] +name = "aead" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d122413f284cf2d62fb1b7db97e02edb8cda96d769b16e443a4f6195e35662b0" +dependencies = [ + "crypto-common", + "generic-array", +] + +[[package]] +name = "aes" +version = "0.8.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b169f7a6d4742236a0a00c541b845991d0ac43e546831af1249753ab4c3aa3a0" +dependencies = [ + "cfg-if", + "cipher", + "cpufeatures", +] + +[[package]] +name = "aes-gcm" +version = "0.10.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "831010a0f742e1209b3bcea8fab6a8e149051ba6099432c8cb2cc117dec3ead1" +dependencies = [ + "aead", + "aes", + "cipher", + "ctr", + "ghash", + "subtle", +] + [[package]] name = "aho-corasick" version = "1.1.4" @@ -129,6 +164,41 @@ version = "1.1.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1505bd5d3d116872e7271a6d4e16d81d0c8570876c8de68093a09ac269d8aac0" +[[package]] +name = "autocfg" +version = "1.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f2032f911046de80f0a198e0901378627c33f59ea0ac00e363d481118bd70a53" + +[[package]] +name = "aws-lc-rs" +version = "1.17.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4342d8937fc7e5dd9b1c60292261c0670c882a2cd1719cfc11b1af41731e32ad" +dependencies = [ + "aws-lc-sys", + "zeroize", +] + +[[package]] +name = "aws-lc-sys" +version = "0.42.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6d9ceb1da931507a12f4fccea479dccd00da1943e1b4ae72d8e502d707361444" +dependencies = [ + "cc", + "cmake", + "dunce", + "fs_extra", + "pkg-config", +] + +[[package]] +name = "base16ct" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4c7f02d4ea65f2c1853089ffd8d2787bdbc63de2f0d29dedbcf8ccdfa0ccd4cf" + [[package]] name = "base64" version = "0.21.7" @@ -141,6 +211,12 @@ version = "0.22.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "72b3254f16251a8381aa12e40e3c4d2f0199f8c6508fbecb9d91f575e0fbb8c6" +[[package]] +name = "base64ct" +version = "1.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2af50177e190e07a26ab74f8b1efbfe2ef87da2116221318cb1c2e82baf7de06" + [[package]] name = "bitflags" version = "2.11.1" @@ -279,6 +355,41 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724" +[[package]] +name = "chacha20" +version = "0.9.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c3613f74bd2eac03dad61bd53dbe620703d4371614fe0bc3b9f04dd36fe4e818" +dependencies = [ + "cfg-if", + "cipher", + "cpufeatures", +] + +[[package]] +name = "chacha20poly1305" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "10cd79432192d1c0f4e1a0fef9527696cc039165d729fb41b3f4f4f354c2dc35" +dependencies = [ + "aead", + "chacha20", + "cipher", + "poly1305", + "zeroize", +] + +[[package]] +name = "cipher" +version = "0.4.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "773f3b9af64447d2ce9850330c473515014aa235e6a783b02db81ff39e4a3dad" +dependencies = [ + "crypto-common", + "inout", + "zeroize", +] + [[package]] name = "clap" version = "4.6.1" @@ -319,6 +430,15 @@ version = "1.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c8d4a3bb8b1e0c1050499d1815f5ab16d04f0959b233085fb31653fbfc9d98f9" +[[package]] +name = "cmake" +version = "0.1.58" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c0f78a02292a74a88ac736019ab962ece0bc380e3f977bf72e376c5d78ff0678" +dependencies = [ + "cc", +] + [[package]] name = "cobs" version = "0.3.0" @@ -347,6 +467,12 @@ dependencies = [ "windows-sys 0.59.0", ] +[[package]] +name = "const-oid" +version = "0.9.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c2459377285ad874054d797f3ccebf984978aa39129f6eafde5cdc8315b612f8" + [[package]] name = "core-foundation-sys" version = "0.8.7" @@ -531,6 +657,18 @@ version = "0.8.21" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d0a5c400df2834b80a4c3327b3aad3a4c4cd4de0629063962b03235697506a28" +[[package]] +name = "crypto-bigint" +version = "0.5.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0dc92fb57ca44df6db8059111ab3af99a63d5d0f8375d9972e319a379c6bab76" +dependencies = [ + "generic-array", + "rand_core 0.6.4", + "subtle", + "zeroize", +] + [[package]] name = "crypto-common" version = "0.1.7" @@ -541,6 +679,42 @@ dependencies = [ "typenum", ] +[[package]] +name = "ctr" +version = "0.9.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0369ee1ad671834580515889b80f2ea915f23b8be8d0daa4bbaf2ac5c7590835" +dependencies = [ + "cipher", +] + +[[package]] +name = "curve25519-dalek" +version = "4.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "97fb8b7c4503de7d6ae7b42ab72a5a59857b4c937ec27a3d4539dba95b5ab2be" +dependencies = [ + "cfg-if", + "cpufeatures", + "curve25519-dalek-derive", + "digest", + "fiat-crypto", + "rustc_version", + "subtle", + "zeroize", +] + +[[package]] +name = "curve25519-dalek-derive" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f46882e17999c6cc590af592290432be3bce0428cb0d5f8b6715e4dc7b383eb3" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "debugid" version = "0.8.0" @@ -550,6 +724,30 @@ dependencies = [ "uuid", ] +[[package]] +name = "der" +version = "0.7.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e7c1832837b905bbfb5101e07cc24c8deddf52f93225eee6ead5f4d63d53ddcb" +dependencies = [ + "const-oid", + "der_derive", + "flagset", + "pem-rfc7468", + "zeroize", +] + +[[package]] +name = "der_derive" +version = "0.7.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8034092389675178f570469e6c3b0465d3d30b4505c294a6550db47f3c17ad18" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "digest" version = "0.10.7" @@ -557,7 +755,9 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" dependencies = [ "block-buffer", + "const-oid", "crypto-common", + "subtle", ] [[package]] @@ -654,12 +854,77 @@ dependencies = [ "syn", ] +[[package]] +name = "dunce" +version = "1.0.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92773504d58c093f6de2459af4af33faa518c13451eb8f2b5698ed3d36e7c813" + +[[package]] +name = "ecdsa" +version = "0.16.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ee27f32b5c5292967d2d4a9d7f1e0b0aed2c15daded5a60300e4abb9d8020bca" +dependencies = [ + "der", + "digest", + "elliptic-curve", + "rfc6979", + "signature", + "spki", +] + +[[package]] +name = "ed25519" +version = "2.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "115531babc129696a58c64a4fef0a8bf9e9698629fb97e9e40767d235cfbcd53" +dependencies = [ + "pkcs8", + "signature", +] + +[[package]] +name = "ed25519-dalek" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "70e796c081cee67dc755e1a36a0a172b897fab85fc3f6bc48307991f64e4eca9" +dependencies = [ + "curve25519-dalek", + "ed25519", + "serde", + "sha2", + "subtle", + "zeroize", +] + [[package]] name = "either" version = "1.16.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "91622ff5e7162018101f2fea40d6ebf4a78bbe5a49736a2020649edf9693679e" +[[package]] +name = "elliptic-curve" +version = "0.13.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b5e6043086bf7973472e0c7dff2142ea0b680d30e18d9cc40f267efbf222bd47" +dependencies = [ + "base16ct", + "crypto-bigint", + "digest", + "ff", + "generic-array", + "group", + "hkdf", + "pem-rfc7468", + "pkcs8", + "rand_core 0.6.4", + "sec1", + "subtle", + "zeroize", +] + [[package]] name = "embedded-io" version = "0.4.0" @@ -743,12 +1008,34 @@ dependencies = [ "windows-sys 0.59.0", ] +[[package]] +name = "ff" +version = "0.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c0b50bfb653653f9ca9095b427bed08ab8d75a137839d9ad64eb11810d5b6393" +dependencies = [ + "rand_core 0.6.4", + "subtle", +] + +[[package]] +name = "fiat-crypto" +version = "0.2.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "28dea519a9695b9977216879a3ebfddf92f1c08c05d984f8996aecd6ecdc811d" + [[package]] name = "find-msvc-tools" version = "0.1.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582" +[[package]] +name = "flagset" +version = "0.4.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b7ac824320a75a52197e8f2d787f6a38b6718bb6897a35142d749af3c0e8f4fe" + [[package]] name = "foldhash" version = "0.1.5" @@ -781,6 +1068,12 @@ dependencies = [ "windows-sys 0.59.0", ] +[[package]] +name = "fs_extra" +version = "1.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c" + [[package]] name = "futures" version = "0.3.32" @@ -886,6 +1179,7 @@ checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a" dependencies = [ "typenum", "version_check", + "zeroize", ] [[package]] @@ -915,6 +1209,16 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "ghash" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0d8a4362ccb29cb0b265253fb0a2728f592895ee6854fd9bc13f2ffda266ff1" +dependencies = [ + "opaque-debug", + "polyval", +] + [[package]] name = "gimli" version = "0.31.1" @@ -926,6 +1230,17 @@ dependencies = [ "stable_deref_trait", ] +[[package]] +name = "group" +version = "0.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0f9ef7462f7c099f518d754361858f86d8a07af53ba9af0fe635bbccb151a63" +dependencies = [ + "ff", + "rand_core 0.6.4", + "subtle", +] + [[package]] name = "hashbrown" version = "0.15.5" @@ -957,6 +1272,24 @@ version = "0.4.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70" +[[package]] +name = "hkdf" +version = "0.12.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7b5f8eb2ad728638ea2c7d47a21db23b7b58a72ed6a38256b8a1849f15fbbdf7" +dependencies = [ + "hmac", +] + +[[package]] +name = "hmac" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6c49c37c09c17a53d937dfbb742eb3a961d65a994e6bcdcf37e7399d0cc8ab5e" +dependencies = [ + "digest", +] + [[package]] name = "http" version = "1.4.1" @@ -1029,7 +1362,7 @@ dependencies = [ "tokio", "tokio-rustls", "tower-service", - "webpki-roots", + "webpki-roots 1.0.7", ] [[package]] @@ -1213,6 +1546,15 @@ dependencies = [ "web-time", ] +[[package]] +name = "inout" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "879f10e63c20629ecabbb64a8010319738c66a5cd0c29b02d63d272b03751d01" +dependencies = [ + "generic-array", +] + [[package]] name = "io-extras" version = "0.18.4" @@ -1322,6 +1664,15 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "lazy_static" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bbd2bcb4c963f2ddae06a2efc7e9f3591312473c50c6685e1f298068316e66fe" +dependencies = [ + "spin", +] + [[package]] name = "leb128" version = "0.2.6" @@ -1407,11 +1758,21 @@ dependencies = [ name = "machine" version = "0.0.0" dependencies = [ + "const-oid", + "der", "dirs 6.0.0", + "libc", "log", + "p256", + "rand_core 0.6.4", "riscv-core", + "rustls", + "rustls-pki-types", + "rustls-rustcrypto", "serde", "serde_json", + "webpki-roots 0.26.11", + "x509-cert", ] [[package]] @@ -1468,6 +1829,51 @@ dependencies = [ "riscv-core", ] +[[package]] +name = "num-bigint-dig" +version = "0.8.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e661dda6640fad38e827a6d4a310ff4763082116fe217f279885c97f511bb0b7" +dependencies = [ + "lazy_static", + "libm", + "num-integer", + "num-iter", + "num-traits", + "rand 0.8.6", + "smallvec", + "zeroize", +] + +[[package]] +name = "num-integer" +version = "0.1.46" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7969661fd2958a5cb096e56c8e1ad0444ac2bbcd0061bd28660485a44879858f" +dependencies = [ + "num-traits", +] + +[[package]] +name = "num-iter" +version = "0.1.46" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c92800bd69a1eac91786bcfe9da64a897eb72911b8dc3095decbd07429e8048b" +dependencies = [ + "num-integer", + "num-traits", +] + +[[package]] +name = "num-traits" +version = "0.2.19" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "071dfc062690e90b734c0b2273ce72ad0ffa95f0c74596bc250dcfd960262841" +dependencies = [ + "autocfg", + "libm", +] + [[package]] name = "number_prefix" version = "0.4.0" @@ -1507,18 +1913,57 @@ version = "1.70.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe" +[[package]] +name = "opaque-debug" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381" + [[package]] name = "option-ext" version = "0.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "04744f49eae99ab78e0d5c0b603ab218f515ea8cfe5a456d7629ad883a3b6e7d" +[[package]] +name = "p256" +version = "0.13.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c9863ad85fa8f4460f9c48cb909d38a0d689dba1f6f6988a5e3e0d31071bcd4b" +dependencies = [ + "ecdsa", + "elliptic-curve", + "primeorder", + "sha2", +] + +[[package]] +name = "p384" +version = "0.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fe42f1670a52a47d448f14b6a5c61dd78fce51856e68edaa38f7ae3a46b8d6b6" +dependencies = [ + "ecdsa", + "elliptic-curve", + "primeorder", + "sha2", +] + [[package]] name = "paste" version = "1.0.15" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "57c0d7b74b563b49d38dae00a0c37d4d6de9b432382b2892f0574ddcae73fd0a" +[[package]] +name = "pem-rfc7468" +version = "0.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "88b39c9bfcfc231068454382784bb460aae594343fb030d46e9f50a645418412" +dependencies = [ + "base64ct", +] + [[package]] name = "percent-encoding" version = "2.3.2" @@ -1531,12 +1976,67 @@ version = "0.2.17" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd" +[[package]] +name = "pkcs1" +version = "0.7.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c8ffb9f10fa047879315e6625af03c164b16962a5368d724ed16323b68ace47f" +dependencies = [ + "der", + "pkcs8", + "spki", +] + +[[package]] +name = "pkcs5" +version = "0.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e847e2c91a18bfa887dd028ec33f2fe6f25db77db3619024764914affe8b69a6" +dependencies = [ + "der", + "spki", +] + +[[package]] +name = "pkcs8" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f950b2377845cebe5cf8b5165cb3cc1a5e0fa5cfa3e1f7f55707d8fd82e0a7b7" +dependencies = [ + "der", + "pkcs5", + "spki", +] + [[package]] name = "pkg-config" version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "19f132c84eca552bf34cab8ec81f1c1dcc229b811638f9d283dceabe58c5569e" +[[package]] +name = "poly1305" +version = "0.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8159bd90725d2df49889a078b54f4f79e87f1f8a8444194cdca81d38f5393abf" +dependencies = [ + "cpufeatures", + "opaque-debug", + "universal-hash", +] + +[[package]] +name = "polyval" +version = "0.6.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9d1fe60d06143b2430aa532c94cfe9e29783047f06c0d7fd359a9a51b729fa25" +dependencies = [ + "cfg-if", + "cpufeatures", + "opaque-debug", + "universal-hash", +] + [[package]] name = "portable-atomic" version = "1.13.1" @@ -1592,6 +2092,15 @@ dependencies = [ "syn", ] +[[package]] +name = "primeorder" +version = "0.13.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "353e1ca18966c16d9deb1c69278edbc5f194139612772bd9537af60ac231e1e6" +dependencies = [ + "elliptic-curve", +] + [[package]] name = "proc-macro2" version = "1.0.106" @@ -1875,7 +2384,17 @@ dependencies = [ "wasm-bindgen-futures", "wasm-streams", "web-sys", - "webpki-roots", + "webpki-roots 1.0.7", +] + +[[package]] +name = "rfc6979" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8dd2a808d456c4a54e300a23e9f5a67e122c3024119acbfd73e3bf664491cb2" +dependencies = [ + "hmac", + "subtle", ] [[package]] @@ -1898,6 +2417,28 @@ version = "0.0.0" dependencies = [ "env_logger", "log", + "rustc-hash", +] + +[[package]] +name = "rsa" +version = "0.9.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b8573f03f5883dcaebdfcf4725caa1ecb9c15b2ef50c43a07b816e06799bb12d" +dependencies = [ + "const-oid", + "digest", + "num-bigint-dig", + "num-integer", + "num-traits", + "pkcs1", + "pkcs8", + "rand_core 0.6.4", + "sha2", + "signature", + "spki", + "subtle", + "zeroize", ] [[package]] @@ -1912,6 +2453,15 @@ version = "2.1.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "94300abf3f1ae2e2b8ffb7b58043de3d399c73fa6f4b73826402a5c457614dbe" +[[package]] +name = "rustc_version" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cfcb3a22ef46e85b45de6ee7e79d063319ebb6594faafcf1c225ea92ab6e9b92" +dependencies = [ + "semver", +] + [[package]] name = "rustix" version = "0.38.44" @@ -1954,10 +2504,11 @@ version = "0.23.40" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ef86cd5876211988985292b91c96a8f2d298df24e75989a43a3c73f2d4d8168b" dependencies = [ + "aws-lc-rs", "once_cell", "ring", "rustls-pki-types", - "rustls-webpki", + "rustls-webpki 0.103.13", "subtle", "zeroize", ] @@ -1972,12 +2523,54 @@ dependencies = [ "zeroize", ] +[[package]] +name = "rustls-rustcrypto" +version = "0.0.2-alpha" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f12052947763ab8515f753315357599e9b0b4dab3b8ba15f30f725fe6d025557" +dependencies = [ + "aead", + "aes-gcm", + "chacha20poly1305", + "crypto-common", + "der", + "digest", + "ecdsa", + "ed25519-dalek", + "hmac", + "p256", + "p384", + "paste", + "pkcs8", + "rand_core 0.6.4", + "rsa", + "rustls", + "rustls-pki-types", + "rustls-webpki 0.102.8", + "sec1", + "sha2", + "signature", + "x25519-dalek", +] + +[[package]] +name = "rustls-webpki" +version = "0.102.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "64ca1bc8749bd4cf37b5ce386cc146580777b4e8572c7b97baf22c83f444bee9" +dependencies = [ + "ring", + "rustls-pki-types", + "untrusted", +] + [[package]] name = "rustls-webpki" version = "0.103.13" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "61c429a8649f110dddef65e2a5ad240f747e85f7758a6bccc7e5777bd33f756e" dependencies = [ + "aws-lc-rs", "ring", "rustls-pki-types", "untrusted", @@ -1995,6 +2588,20 @@ version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f" +[[package]] +name = "sec1" +version = "0.7.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d3e97a565f76233a6003f9f5c54be1d9c5bdfa3eccfb189469f11ec4901c47dc" +dependencies = [ + "base16ct", + "der", + "generic-array", + "pkcs8", + "subtle", + "zeroize", +] + [[package]] name = "semver" version = "1.0.28" @@ -2069,6 +2676,17 @@ dependencies = [ "serde", ] +[[package]] +name = "sha1" +version = "0.10.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a978451301f4db1d02937a4ab3ccce137717b81826e79b7d49ffe3244a13c3b8" +dependencies = [ + "cfg-if", + "cpufeatures", + "digest", +] + [[package]] name = "sha2" version = "0.10.9" @@ -2095,6 +2713,16 @@ version = "2.0.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f8fadd59c855ef2080decdef8ff161eb6661b86933c9d82e5ba29dc602a55aba" +[[package]] +name = "signature" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77549399552de45a898a580c1b41d445bf730df867cc44e6c0233bbc4b8329de" +dependencies = [ + "digest", + "rand_core 0.6.4", +] + [[package]] name = "slab" version = "0.4.12" @@ -2120,6 +2748,22 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "spin" +version = "0.9.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3763264f6b73151db08c50ff20d7d8a0b8796e021cdea7ceedad07b80155fa0e" + +[[package]] +name = "spki" +version = "0.7.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d91ed6c858b01f942cd56b37a94b3e0a1798290327d1236e4d9cf4eaca44d29d" +dependencies = [ + "base64ct", + "der", +] + [[package]] name = "sptr" version = "0.3.2" @@ -2271,6 +2915,27 @@ version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" +[[package]] +name = "tls_codec" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0de2e01245e2bb89d6f05801c564fa27624dbd7b1846859876c7dad82e90bf6b" +dependencies = [ + "tls_codec_derive", + "zeroize", +] + +[[package]] +name = "tls_codec_derive" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2d2e76690929402faae40aebdda620a2c0e25dd6d3b9afe48867dfd95991f4bd" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "tokio" version = "1.52.3" @@ -2472,6 +3137,16 @@ version = "0.2.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ebc1c04c71510c7f702b52b7c350734c9ff1295c464a03335b00bb84fc54f853" +[[package]] +name = "universal-hash" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc1de2c688dc15305988b563c3854064043356019f97a4b46276fe734c4f07ea" +dependencies = [ + "crypto-common", + "subtle", +] + [[package]] name = "untrusted" version = "0.9.0" @@ -2538,6 +3213,14 @@ dependencies = [ "windows-sys 0.59.0", ] +[[package]] +name = "vpod-translate" +version = "0.1.0" +dependencies = [ + "lz4_flex", + "riscv-core", +] + [[package]] name = "want" version = "0.3.1" @@ -3091,6 +3774,15 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "webpki-roots" +version = "0.26.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "521bc38abb08001b01866da9f51eb7c5d647a19260e00054a8c7fd5f9e57f7a9" +dependencies = [ + "webpki-roots 1.0.7", +] + [[package]] name = "webpki-roots" version = "1.0.7" @@ -3553,6 +4245,31 @@ version = "0.6.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1ffae5123b2d3fc086436f8834ae3ab053a283cfac8fe0a0b8eaae044768a4c4" +[[package]] +name = "x25519-dalek" +version = "2.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c7e468321c81fb07fa7f4c636c3972b9100f0346e5b6a9f2bd0603a52f7ed277" +dependencies = [ + "curve25519-dalek", + "rand_core 0.6.4", + "zeroize", +] + +[[package]] +name = "x509-cert" +version = "0.2.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1301e935010a701ae5f8655edc0ad17c44bad3ac5ce8c39185f75453b720ae94" +dependencies = [ + "const-oid", + "der", + "sha1", + "signature", + "spki", + "tls_codec", +] + [[package]] name = "yoke" version = "0.8.2" @@ -3622,6 +4339,20 @@ name = "zeroize" version = "1.8.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b97154e67e32c85465826e8bcc1c59429aaaf107c1e4a9e53c8d8ccd5eff88d0" +dependencies = [ + "zeroize_derive", +] + +[[package]] +name = "zeroize_derive" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3c50655cbb0fe3fc43170059e702f1ce5e19b84cec58dc87b037a09935c2f328" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] [[package]] name = "zerotrie" diff --git a/Cargo.toml b/Cargo.toml index a852c7d..30fc44d 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,10 @@ [workspace] -members = ["crates/riscv-core", "crates/machine", "crates/native-cli", "crates/wasi-component", "crates/vpod"] +members = ["crates/riscv-core", "crates/machine", "crates/native-cli", "crates/wasi-component", "crates/vpod", "crates/vpod-translate"] resolver = "3" [workspace.package] edition = "2024" + +[profile.release] +lto = "fat" +codegen-units = 1 diff --git a/README.md b/README.md index a6c6fa9..56f8036 100644 --- a/README.md +++ b/README.md @@ -25,9 +25,13 @@ A `vpod` is a lightweight, portable sandbox that gives an untrusted process an i ## How it works -A `vpod` runs a RISC‑V virtual machine compiled to WebAssembly, implementing the RV64GC specification. When you start a `vpod`, it boots from a snapshot, a saved VM state ready in under a second. +A `vpod` runs a complete RISC‑V system (RV64GC, single vCPU) compiled to WebAssembly. Inside it boots a real Linux kernel with a real userspace, so shells, tools and daemons all behave like they would on actual hardware. -The WASM component communicates with the host through WASI 0.2, providing controlled access to filesystem, networking, and standard I/O while keeping all execution state (CPU registers, memory, filesystem) isolated inside the sandbox. +**Snapshots.** Instead of booting Linux from scratch, a `vpod` restores a snapshot: a saved machine state (CPU registers, RAM, filesystem) captured right after boot. Restoring one takes well under a second. Suspend works the same way in reverse, only dirty memory pages are written back to disk, so you can pause a sandbox and resume it later, even from another process. + +**Ahead-of-time translation.** Pure instruction-by-instruction emulation is slow, and WebAssembly rules out a runtime JIT. So at snapshot build time, the hottest guest code paths are translated from RISC‑V into native code that gets compiled into the WASM module itself. At runtime the emulator dispatches into these translated blocks when the guest code matches, and falls back to the interpreter when it doesn't. This is worth roughly 5x on CPU-bound work, with zero effect on isolation: translated code goes through the same MMU and memory checks as interpreted code. + +**The WASI boundary.** The WASM component talks to the host exclusively through WASI 0.2. The guest never sees host file descriptors, sockets, or memory: filesystem access goes through explicitly mounted directories, and networking goes through a user-mode network stack inside the component that only ever asks the host for plain outbound sockets. Everything else (guest kernel, processes, memory) lives inside the WASM linear memory and dies with it. ### RV64GC Specification @@ -45,38 +49,6 @@ Reduces code size by 30%, improving instruction fetch speed and memory efficienc ## Getting started -### CLI - -```bash -curl -fsSL https://install.vpod.sh | sh -``` - ->
-> Or install via PowerShell (windows) -> -> ```bash -> irm https://install.vpod.sh | iex -> ``` -> ->
- ->
-> Or install via cargo -> -> ```bash -> cargo install vpod -> ``` -> ->
- -```bash -# Pull a snapshot -vpod pull alpine:latest - -# Start an interactive shell -vpod -``` - ### Python SDK ```bash @@ -109,47 +81,112 @@ with Sandbox.create() as sandbox: > [!IMPORTANT] > The first call to `Sandbox.create()` downloads the default snapshot (`alpine`) and caches it locally if not already present. -For more details, see the [full documentation](https://docs.vpod.sh/quickstart). +> [!TIP] +> The default snapshot ships with `apk` for system packages and `uv` for Python packages, so you can install what you need at runtime. + + +### CLI + +```bash +curl -fsSL https://install.vpod.sh | sh +``` + +>
+> Or install via PowerShell (windows) +> +> ```bash +> irm https://install.vpod.sh | iex +> ``` +> +>
+ +```bash +# Pull a snapshot +vpod pull alpine:latest + +# Start an interactive shell +vpod +``` + +## Documentation +Visit [Vpod documentation](https://docs.vpod.sh/quickstart). ## Limitations -- **Emulation overhead**: No hardware acceleration in the WASM component. CPU-intensive workloads may run slower than native. -- **No GPU access**: CUDA, Metal, and hardware ML accelerators are not yet available. Support may be added in the future with wasi-nn. -- **Env vars don't cross between shell and Python**: `sandbox.commands.run("export FOO=bar")` is not visible in `sandbox.code.run(...)`. Use the filesystem to share data between the two. +- **Emulation overhead**: There is no hardware virtualization inside WebAssembly, so all guest code is emulated. The overhead depends entirely on the workload: I/O-bound and network-bound work runs close to native speed, while heavy CPU-bound work runs noticeably slower even with AOT translation. If your workload is mostly "run a tool, read a file, call an API", you won't notice. +- **No GPU access**: CUDA, Metal, and hardware ML accelerators are not available. Support may be added in the future with wasi-nn. ## Contributing -**Prerequisites** -- Rust (latest stable) -- Python 3.10+ +Contributions are welcome, from bug reports to new device support. Open an [issue](https://github.com/capsulerun/vpod/issues/new) to discuss anything substantial before building it. + +### Repository layout + +| Path | What it is | +|:---|:---| +| `crates/riscv-core` | RV64GC decoder and executor, MMU, and the AOT block runtime | +| `crates/machine` | The machine model: RAM (copy-on-write), UART, PLIC/CLINT, virtio devices, snapshot save/restore | +| `crates/wasi-component` | The WASM component (WASI 0.2) that wraps the machine for sandboxed use | +| `crates/vpod` | The host CLI (`vpod`), which runs the WASM component | +| `crates/native-cli` | A native (non-WASM) build of the emulator, used for development and debugging | +| `crates/vpod-translate` | The AOT translator: turns traced hot RISC‑V code into Rust at snapshot build time | +| `sdks/python` | The Python SDK (`pip install vpod`) | +| `scripts/` | Build scripts for the WASM component, snapshots, and AOT translation | + +### Prerequisites + +- **Rust** (latest stable) with the `wasm32-wasip2` target: `rustup target add wasm32-wasip2` +- **Python 3.10+** for the SDK +- **Zig** (0.14) and **bsdtar**, only needed if you build snapshots yourself + +### Development setup -**Development setup** ```bash -# Build WASM component +# One-time: generate the AOT stub (a fresh clone has no translated blocks) +./scripts/aot-stub.sh + +# Build the WASM component (library + CLI) ./scripts/build-wasm.sh -# Install CLI +# Install the host CLI cargo install --path crates/vpod -# Install Python SDK in dev mode -pip install -e sdks/python[dev] - -# Run tests -cargo test # Rust tests -pytest sdks/python/tests/ -v -m integration # Integration tests (requires WASM build) +# Install the Python SDK in dev mode +pip install -e "sdks/python[dev]" ``` -**Building snapshots** +### Running tests -The project uses pre-built Alpine snapshots from `registry.vpod.sh`. To build a custom snapshot: +CI runs these on every PR, so run them before pushing: ```bash -./scripts/build-default-snapshot.sh +cargo fmt --all -- --check # formatting +cargo clippy --all-targets --all-features -- -D warnings +cargo test --all # Rust tests + +# Python SDK integration tests (needs the WASM library in place) +cp target/wasm32-wasip2/release/vpod_wasi_lib.wasm sdks/python/vpod/ +pytest sdks/python/tests/ -v -m integration ``` -This creates `dist/alpine-3.23.0-256mb.snap`. +### Building snapshots + +The project uses pre-built Alpine snapshots from `registry.vpod.sh`, so you normally don't need this. To build one locally: + +```bash +./scripts/build-default-snapshot.sh # dist/alpine-3.23.0-256mb.snap +./scripts/build-data-snapshot.sh # 512 MB variant with numpy/pandas/scipy +``` > [!IMPORTANT] -> To use it locally, uncomment lines in `resolve_snapshot()` in `crates/vpod/src/main.rs`. +> To use a locally built snapshot, uncomment the lines in `resolve_snapshot()` in `crates/vpod/src/main.rs`. + +Snapshot builds can also run the AOT pass (`scripts/aot-snapshot.sh `), which traces a representative workload, translates the hot blocks, and rebuilds the emulator with them baked in. It takes a while; the stub from `aot-stub.sh` is fine for everyday development, everything works the same, just slower. + +### Pull requests + +- Keep PRs focused: one change per PR. +- `fmt`, `clippy` and the test suite must pass (CI enforces all three). +- If you touch the emulator's execution or memory paths, say how you validated correctness (the test suite at minimum; for subtle changes a boot plus a real workload in the guest is a good sanity check). ## License diff --git a/crates/machine/Cargo.toml b/crates/machine/Cargo.toml index 606b58b..a7a0efe 100644 --- a/crates/machine/Cargo.toml +++ b/crates/machine/Cargo.toml @@ -9,3 +9,18 @@ log = "0.4" serde = { version = "1", features = ["derive"] } serde_json = "1" dirs = "6" +rustls = { version = "0.23", default-features = false, features = ["std", "tls12"] } +rustls-rustcrypto = "0.0.2-alpha" +rustls-pki-types = "1" +webpki-roots = "0.26" +x509-cert = { version = "0.2", features = ["builder"] } +p256 = { version = "0.13", features = ["ecdsa", "pem", "pkcs8"] } +der = { version = "0.7", features = ["pem"] } +const-oid = { version = "0.9", features = ["db"] } +rand_core = { version = "0.6", features = ["getrandom"] } + +[dev-dependencies] +rustls = { version = "0.23", default-features = false, features = ["std", "tls12", "ring", "aws-lc-rs"] } + +[target.'cfg(unix)'.dev-dependencies] +libc = "0.2" diff --git a/crates/machine/assets/tls/vpod-ca-cert.pem b/crates/machine/assets/tls/vpod-ca-cert.pem new file mode 100644 index 0000000..ec19dac --- /dev/null +++ b/crates/machine/assets/tls/vpod-ca-cert.pem @@ -0,0 +1,10 @@ +-----BEGIN CERTIFICATE----- +MIIBYDCCAQegAwIBAgIBATAKBggqhkjOPQQDAjAYMRYwFAYDVQQDDA12cG9kIGxv +Y2FsIENBMB4XDTI2MDcxNTA3MjY1N1oXDTM2MDcxMjA3MjY1N1owGDEWMBQGA1UE +AwwNdnBvZCBsb2NhbCBDQTBZMBMGByqGSM49AgEGCCqGSM49AwEHA0IABOYt0tJ8 +FrlckHPCsLPRqE5J0vYaKHdlV0AniqCz5OwLGgMaLZKGGTVJB42w5L1WqKKqp/tm +Gwu0M49x3QxuFtajQjBAMB0GA1UdDgQWBBS2hEVRKfu/KzxnjhsIYN7H/is+5DAP +BgNVHRMBAf8EBTADAQH/MA4GA1UdDwEB/wQEAwIBBjAKBggqhkjOPQQDAgNHADBE +AiB3UGbX5SbemKzH2S/Sukz9hbPibxZpp3oMdvMgKYJc1QIgc5e3X7oymhHe5c51 +23V/+lL6Nzu4UcuirnZyZLCZHAY= +-----END CERTIFICATE----- diff --git a/crates/machine/assets/tls/vpod-ca-key.pem b/crates/machine/assets/tls/vpod-ca-key.pem new file mode 100644 index 0000000..379a1a9 --- /dev/null +++ b/crates/machine/assets/tls/vpod-ca-key.pem @@ -0,0 +1,5 @@ +-----BEGIN PRIVATE KEY----- +MIGHAgEAMBMGByqGSM49AgEGCCqGSM49AwEHBG0wawIBAQQgZgXEZ+vmg2LpSUb2 +HFGLOhd+IAp5AVfMgo7gNe/1JYahRANCAATmLdLSfBa5XJBzwrCz0ahOSdL2Gih3 +ZVdAJ4qgs+TsCxoDGi2Shhk1SQeNsOS9Vqiiqqf7ZhsLtDOPcd0MbhbW +-----END PRIVATE KEY----- diff --git a/crates/machine/src/clint.rs b/crates/machine/src/clint.rs index 1db1b46..4bf8cd9 100644 --- a/crates/machine/src/clint.rs +++ b/crates/machine/src/clint.rs @@ -39,6 +39,25 @@ impl Clint { self.mtime += nanos / NANOS_PER_TICK; } + pub fn nanos_until_timer(&self) -> Option { + const NANOS_PER_TICK: u64 = 1_000_000_000 / TIMER_FREQUENCY; + + if self.mtimecmp == u64::MAX || self.mtimecmp <= self.mtime { + None + } else { + Some((self.mtimecmp - self.mtime) * NANOS_PER_TICK) + } + } + + pub fn fast_forward_to_timer(&mut self) -> bool { + if self.mtimecmp != u64::MAX && self.mtimecmp > self.mtime { + self.mtime = self.mtimecmp; + true + } else { + false + } + } + pub fn mtime(&self) -> u64 { self.mtime } diff --git a/crates/machine/src/cow_ram.rs b/crates/machine/src/cow_ram.rs index 2183be1..85311db 100644 --- a/crates/machine/src/cow_ram.rs +++ b/crates/machine/src/cow_ram.rs @@ -1,13 +1,21 @@ // Copy-on-write guest RAM. use std::sync::Arc; +use std::sync::atomic::{AtomicU64, Ordering}; pub const PAGE_SIZE: usize = 4096; +static EPOCH_SOURCE: AtomicU64 = AtomicU64::new(1); + +fn next_epoch() -> u64 { + EPOCH_SOURCE.fetch_add(1, Ordering::Relaxed) +} + pub struct CowRam { base: Arc>, pages: Vec>>, len: usize, mask: u64, + epoch: u64, } impl CowRam { @@ -27,6 +35,7 @@ impl CowRam { pages: vec![None; num_pages], len, mask: ram_size - 1, + epoch: next_epoch(), } } @@ -43,6 +52,7 @@ impl CowRam { pages: vec![None; num_pages], len: logical_len, mask: ram_size - 1, + epoch: next_epoch(), } } @@ -52,9 +62,25 @@ impl CowRam { pages: vec![None; self.pages.len()], len: self.len, mask: self.mask, + epoch: next_epoch(), } } + #[inline(always)] + pub fn epoch(&self) -> u64 { + self.epoch + } + + #[inline(always)] + pub fn page_ptr(&self, page: usize) -> *const u8 { + self.page_ref(page).as_ptr() + } + + #[inline(always)] + pub fn page_mut_ptr(&mut self, page: usize) -> *mut u8 { + self.page_mut(page).as_mut_ptr() + } + #[inline(always)] pub fn len(&self) -> usize { self.len @@ -80,14 +106,16 @@ impl CowRam { #[inline(always)] fn page_mut(&mut self, page: usize) -> &mut [u8] { - let base = &self.base; - - self.pages[page].get_or_insert_with(|| { + if self.pages[page].is_none() { + let start = page * PAGE_SIZE; let mut owned = vec![0u8; PAGE_SIZE].into_boxed_slice(); - owned.copy_from_slice(&base[page * PAGE_SIZE..(page + 1) * PAGE_SIZE]); + owned.copy_from_slice(&self.base[start..start + PAGE_SIZE]); + self.pages[page] = Some(owned); + + self.epoch = next_epoch(); + } - owned - }) + self.pages[page].as_mut().unwrap() } #[inline(always)] @@ -245,5 +273,6 @@ impl CowRam { self.base = Arc::new(padded); self.pages = vec![None; num_pages]; + self.epoch = next_epoch(); } } diff --git a/crates/machine/src/machine_bus.rs b/crates/machine/src/machine_bus.rs index a6e3159..56b2e6d 100644 --- a/crates/machine/src/machine_bus.rs +++ b/crates/machine/src/machine_bus.rs @@ -232,6 +232,7 @@ impl MachineBus { } impl SystemBus for MachineBus { + #[inline] fn read_byte(&mut self, address: u64) -> u8 { if address >= RAM_BASE && address < RAM_BASE + self.ram_mask + 1 { return self.ram_read_u8(address); @@ -284,10 +285,12 @@ impl SystemBus for MachineBus { 0 } + #[inline] fn read_halfword(&mut self, address: u64) -> u16 { u16::from_le_bytes([self.read_byte(address), self.read_byte(address + 1)]) } + #[inline(always)] fn read_word(&mut self, address: u64) -> u32 { if address >= RAM_BASE && address + 3 < RAM_BASE + self.ram.len() as u64 { let index = (address - RAM_BASE) as usize; @@ -331,6 +334,7 @@ impl SystemBus for MachineBus { ]) } + #[inline(always)] fn read_doubleword(&mut self, address: u64) -> u64 { if address >= RAM_BASE && address + 7 < RAM_BASE + self.ram.len() as u64 { let index = (address - RAM_BASE) as usize; @@ -341,6 +345,7 @@ impl SystemBus for MachineBus { (self.read_word(address) as u64) | ((self.read_word(address + 4) as u64) << 32) } + #[inline] fn write_byte(&mut self, address: u64, value: u8) { if address >= RAM_BASE && address < RAM_BASE + self.ram_mask + 1 { self.ram_write_u8(address, value); @@ -377,12 +382,14 @@ impl SystemBus for MachineBus { } } + #[inline] fn write_halfword(&mut self, address: u64, value: u16) { let [low_byte, high_byte] = value.to_le_bytes(); self.write_byte(address, low_byte); self.write_byte(address + 1, high_byte); } + #[inline(always)] fn write_word(&mut self, address: u64, value: u32) { let [byte_0, byte_1, byte_2, byte_3] = value.to_le_bytes(); @@ -456,6 +463,7 @@ impl SystemBus for MachineBus { self.write_byte(address + 3, byte_3); } + #[inline(always)] fn write_doubleword(&mut self, address: u64, value: u64) { if address >= RAM_BASE && address + 7 < RAM_BASE + self.ram.len() as u64 { let index = (address - RAM_BASE) as usize; @@ -466,6 +474,44 @@ impl SystemBus for MachineBus { self.write_word(address, value as u32); self.write_word(address + 4, (value >> 32) as u32); } + + fn ram_load_page(&mut self, address: u64) -> Option<*const u8> { + let page_base = address & !0xfff; + if page_base < RAM_BASE { + return None; + } + + let page = ((page_base - RAM_BASE) as usize) >> 12; + if page >= self.ram.num_pages() { + return None; + } + + Some(self.ram.page_ptr(page)) + } + + fn ram_store_page(&mut self, address: u64) -> Option<*mut u8> { + let page_base = address & !0xfff; + if page_base < RAM_BASE { + return None; + } + + let page = ((page_base - RAM_BASE) as usize) >> 12; + if page >= self.ram.num_pages() { + return None; + } + + Some(self.ram.page_mut_ptr(page)) + } + + #[inline(always)] + fn ram_epoch(&self) -> u64 { + self.ram.epoch() + } + + #[inline(always)] + fn timer_interrupt_pending(&self) -> Option { + Some(self.clint.get_interrupt_status().0) + } } fn kernel_entry_and_offset(kernel: &[u8]) -> (u64, u64) { diff --git a/crates/machine/src/snapshot.rs b/crates/machine/src/snapshot.rs index 6f6f926..73cf4de 100644 --- a/crates/machine/src/snapshot.rs +++ b/crates/machine/src/snapshot.rs @@ -348,9 +348,9 @@ fn save_hart(hart: &Hart, writer: &mut impl Write) -> io::Result<()> { } } - writer.write_all(&hart.fetch_vpage.to_le_bytes())?; - writer.write_all(&hart.fetch_ppage.to_le_bytes())?; - writer.write_all(&hart.fetch_satp.to_le_bytes())?; + writer.write_all(&u64::MAX.to_le_bytes())?; + writer.write_all(&0u64.to_le_bytes())?; + writer.write_all(&u64::MAX.to_le_bytes())?; save_csr(hart, writer) } @@ -387,14 +387,25 @@ fn restore_hart(hart: &mut Hart, reader: &mut impl Read) -> io::Result<()> { Some(u64::from_le_bytes(buffer_u64)) }; - reader.read_exact(&mut buffer_u64)?; - hart.fetch_vpage = u64::from_le_bytes(buffer_u64); + let mut cached_vpage = [0u8; 8]; + reader.read_exact(&mut cached_vpage)?; - reader.read_exact(&mut buffer_u64)?; - hart.fetch_ppage = u64::from_le_bytes(buffer_u64); + let mut cached_ppage = [0u8; 8]; + reader.read_exact(&mut cached_ppage)?; - reader.read_exact(&mut buffer_u64)?; - hart.fetch_satp = u64::from_le_bytes(buffer_u64); + let mut cached_satp = [0u8; 8]; + reader.read_exact(&mut cached_satp)?; + + hart.fetch_tlb.flush(); + + let vpage = u64::from_le_bytes(cached_vpage); + if vpage != u64::MAX { + hart.fetch_tlb.insert( + vpage, + u64::from_le_bytes(cached_ppage), + u64::from_le_bytes(cached_satp), + ); + } hart.invalidate_icache(); diff --git a/crates/machine/src/virtio/https_gateway.rs b/crates/machine/src/virtio/https_gateway.rs new file mode 100644 index 0000000..1878efa --- /dev/null +++ b/crates/machine/src/virtio/https_gateway.rs @@ -0,0 +1,633 @@ +//Gateway for guest connections to :443 + +use std::collections::VecDeque; +use std::io::{Read, Write}; +use std::net::{Ipv4Addr, SocketAddrV4, TcpStream}; + +use rustls::pki_types::ServerName; +use rustls::{ClientConfig, ClientConnection}; +use std::sync::Arc; + +use super::tls_proxy::{Timing, TlsContext, TlsProxy}; + +pub const PREAMBLE_PREFIX: &[u8] = b"VPOD-CONNECT "; +const PREAMBLE_MAX: usize = 280; + +enum GatewayState { + Sniffing(Vec), + Tls(Box), + Plain(Box), + Failed, +} + +pub struct HttpsGateway { + state: GatewayState, + ctx: TlsContext, + upstream_config: Arc, + dst_ip: [u8; 4], + timing: Option, +} + +impl HttpsGateway { + pub fn new(ctx: &TlsContext, dst_ip: [u8; 4]) -> Self { + Self { + state: GatewayState::Sniffing(Vec::new()), + ctx: ctx.clone(), + upstream_config: ctx.upstream_config(), + dst_ip, + timing: Timing::new(), + } + } + + #[cfg(test)] + fn new_test(ctx: &TlsContext, dst_ip: [u8; 4], upstream_config: Arc) -> Self { + Self { + state: GatewayState::Sniffing(Vec::new()), + ctx: ctx.clone(), + upstream_config, + dst_ip, + timing: None, + } + } + + pub fn failed(&self) -> bool { + match &self.state { + GatewayState::Failed => true, + GatewayState::Tls(p) => p.failed(), + GatewayState::Plain(b) => b.failed, + GatewayState::Sniffing(_) => false, + } + } + + pub fn eof(&self) -> bool { + match &self.state { + GatewayState::Plain(b) => b.upstream_closed && b.to_guest.is_empty(), + _ => false, + } + } + + pub fn has_pending(&self) -> bool { + match &self.state { + GatewayState::Tls(p) => p.has_pending(), + GatewayState::Plain(b) => !b.to_guest.is_empty(), + _ => false, + } + } + + pub fn push_from_guest(&mut self, bytes: &[u8]) { + match &mut self.state { + GatewayState::Sniffing(buffered) => { + buffered.extend_from_slice(bytes); + self.decide_dialect(); + } + GatewayState::Tls(p) => p.push_from_guest(bytes), + GatewayState::Plain(b) => b.push_from_guest(bytes), + GatewayState::Failed => {} + } + } + + pub fn pull_to_guest(&mut self, buf: &mut [u8]) -> Option { + match &mut self.state { + GatewayState::Tls(p) => p.pull_to_guest(buf), + GatewayState::Plain(b) => b.pull_to_guest(buf), + _ => None, + } + } + + pub fn shutdown_write(&mut self) {} + + fn decide_dialect(&mut self) { + let GatewayState::Sniffing(buffered) = &self.state else { + return; + }; + if buffered.is_empty() { + return; + } + + let could_be_preamble = PREAMBLE_PREFIX + .starts_with(&buffered[..buffered.len().min(PREAMBLE_PREFIX.len())]) + || buffered.starts_with(PREAMBLE_PREFIX); + + if !could_be_preamble { + let GatewayState::Sniffing(buffered) = + std::mem::replace(&mut self.state, GatewayState::Failed) + else { + unreachable!(); + }; + + match TlsProxy::with_timing(&self.ctx, self.dst_ip, self.timing.take()) { + Ok(mut proxy) => { + proxy.push_from_guest(&buffered); + self.state = GatewayState::Tls(Box::new(proxy)); + } + Err(e) => { + log::warn!("https_gateway: TLS proxy init failed: {e}"); + } + } + + return; + } + + let Some(newline) = buffered.iter().position(|&b| b == b'\n') else { + if buffered.len() > PREAMBLE_MAX { + self.state = GatewayState::Failed; + } + + return; + }; + + let GatewayState::Sniffing(buffered) = + std::mem::replace(&mut self.state, GatewayState::Failed) + else { + unreachable!(); + }; + + let (line, remainder) = buffered.split_at(newline + 1); + + let Some((host, port)) = parse_preamble(line) else { + log::warn!("https_gateway: malformed preamble line"); + return; + }; + + if let Some(t) = &mut self.timing { + t.mark(&format!( + "preamble received (plaintext bridge to {host}:{port})" + )); + } + + match PlainBridge::connect( + self.upstream_config.clone(), + self.dst_ip, + port, + &host, + self.timing.take(), + ) { + Ok(mut bridge) => { + if !remainder.is_empty() { + bridge.push_from_guest(remainder); + } + self.state = GatewayState::Plain(Box::new(bridge)); + } + Err(e) => { + log::warn!("https_gateway: plaintext bridge to {host}:{port} failed: {e}"); + } + } + } +} + +fn parse_preamble(line: &[u8]) -> Option<(String, u16)> { + let line = std::str::from_utf8(line).ok()?; + let rest = line.strip_prefix(std::str::from_utf8(PREAMBLE_PREFIX).unwrap())?; + let mut parts = rest.trim_end().split(' '); + let host = parts.next()?.to_string(); + let port: u16 = parts.next()?.parse().ok()?; + if host.is_empty() || parts.next().is_some() { + return None; + } + Some((host, port)) +} + +struct PlainBridge { + client: ClientConnection, + upstream: TcpStream, + to_guest: VecDeque, + failed: bool, + upstream_closed: bool, + timing: Option, + upstream_hs_done: bool, + marked_upstream_hs: bool, + first_reply: bool, +} + +impl PlainBridge { + fn connect( + client_config: Arc, + dst_ip: [u8; 4], + port: u16, + host: &str, + timing: Option, + ) -> Result { + let addr = SocketAddrV4::new(Ipv4Addr::from(dst_ip), port); + + // See if better alternative for cfg + #[cfg(target_family = "wasm")] + let stream = TcpStream::connect(addr); + + #[cfg(not(target_family = "wasm"))] + let stream = TcpStream::connect_timeout(&addr.into(), std::time::Duration::from_secs(10)); + + let stream = stream.map_err(|e| e.to_string())?; + stream.set_nonblocking(true).ok(); + stream.set_nodelay(true).ok(); + + let server_name = ServerName::try_from(host.to_string()).map_err(|e| e.to_string())?; + let client = + ClientConnection::new(client_config, server_name).map_err(|e| e.to_string())?; + + let bridge = Self { + client, + upstream: stream, + to_guest: VecDeque::new(), + failed: false, + upstream_closed: false, + timing, + upstream_hs_done: false, + marked_upstream_hs: false, + first_reply: false, + }; + if let Some(t) = &bridge.timing { + t.mark("upstream TCP connected"); + } + Ok(bridge) + } + + fn push_from_guest(&mut self, bytes: &[u8]) { + if self.client.writer().write_all(bytes).is_err() { + self.failed = true; + return; + } + + self.pump(); + } + + fn pull_to_guest(&mut self, buf: &mut [u8]) -> Option { + self.pump(); + if self.to_guest.is_empty() { + return None; + } + + let n = self.to_guest.len().min(buf.len()); + for slot in buf.iter_mut().take(n) { + *slot = self.to_guest.pop_front().unwrap(); + } + + Some(n) + } + + fn pump(&mut self) { + if self.failed { + return; + } + + if !self.client.is_handshaking() { + self.upstream_hs_done = true; + } + + while !self.upstream_closed && self.client.wants_write() { + match self.client.write_tls(&mut self.upstream) { + Ok(0) => break, + Ok(_) => {} + Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => break, + Err(e) => { + if self.upstream_hs_done { + log::debug!("plain_bridge: upstream write ended: {e}"); + self.upstream_closed = true; + break; + } + + log::warn!("plain_bridge: upstream write error: {e}"); + self.failed = true; + return; + } + } + } + + let mut buf = [0u8; 16384]; + while !self.upstream_closed { + loop { + match self.client.reader().read(&mut buf) { + Ok(0) => break, + Ok(n) => self.to_guest.extend(&buf[..n]), + Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => break, + Err(_) => break, + } + } + + match self.client.read_tls(&mut self.upstream) { + Ok(0) => self.upstream_closed = true, + Ok(_) => { + if let Err(e) = self.client.process_new_packets() { + if self.upstream_hs_done { + log::debug!("plain_bridge: upstream closed uncleanly: {e}"); + self.upstream_closed = true; + } else { + log::warn!("plain_bridge: upstream tls error: {e}"); + self.failed = true; + return; + } + } + } + Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => break, + Err(e) => { + if self.upstream_hs_done { + log::debug!("plain_bridge: upstream read ended: {e}"); + self.upstream_closed = true; + } else { + log::warn!("plain_bridge: upstream read error: {e}"); + self.failed = true; + return; + } + } + } + } + + loop { + match self.client.reader().read(&mut buf) { + Ok(0) => break, + Ok(n) => self.to_guest.extend(&buf[..n]), + Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => break, + Err(_) => break, + } + } + + if let Some(t) = &mut self.timing + && self.upstream_hs_done + && !self.marked_upstream_hs + { + self.marked_upstream_hs = true; + t.mark("upstream handshake done"); + } + + if let Some(t) = &mut self.timing + && self.upstream_hs_done + && !self.first_reply + && !self.to_guest.is_empty() + { + self.first_reply = true; + t.mark("first reply byte to guest (plaintext)"); + } + } +} + +#[cfg(test)] +mod tests { + use super::super::tls_proxy::{ + ca_cert_pem, client_config_trusting, spawn_test_upstream, spawn_test_upstream_rst, + spawn_test_upstream_streaming, + }; + use super::*; + use std::time::Duration; + + const UPSTREAM_REPLY: &[u8] = b"HTTP/1.0 200 OK\r\nContent-Length: 5\r\n\r\nhello"; + + fn drain(gateway: &mut HttpsGateway, into: &mut Vec) { + let mut buf = [0u8; 16384]; + while let Some(n) = gateway.pull_to_guest(&mut buf) { + into.extend_from_slice(&buf[..n]); + } + } + + #[test] + fn preamble_bridges_plaintext_to_real_tls_upstream() { + let (port, up_ca, up) = spawn_test_upstream(UPSTREAM_REPLY); + let ctx = TlsContext::new().unwrap(); + let mut gateway = + HttpsGateway::new_test(&ctx, [127, 0, 0, 1], client_config_trusting(&up_ca)); + + let wire = format!("VPOD-CONNECT localhost {port}\nGET / HTTP/1.0\r\n\r\n"); + for byte in wire.as_bytes() { + gateway.push_from_guest(std::slice::from_ref(byte)); + assert!(!gateway.failed(), "gateway failed mid-preamble"); + } + + let mut got = Vec::new(); + for _ in 0..2000 { + drain(&mut gateway, &mut got); + if gateway.eof() { + break; + } + std::thread::sleep(Duration::from_millis(1)); + } + let _ = up.join(); + + assert!( + got.windows(5).any(|w| w == b"hello"), + "expected upstream reply, got {got:?}" + ); + assert!(gateway.eof(), "upstream close must surface as EOF"); + } + + #[test] + fn abrupt_upstream_rst_still_delivers_full_reply() { + const BODY_LEN: usize = 8000; + static BIG_REPLY: std::sync::OnceLock> = std::sync::OnceLock::new(); + let reply: &'static [u8] = BIG_REPLY.get_or_init(|| { + let mut v = + format!("HTTP/1.0 200 OK\r\nContent-Length: {BODY_LEN}\r\n\r\n").into_bytes(); + + v.extend(std::iter::repeat_n(b'x', BODY_LEN)); + v + }); + + let (port, up_ca, up) = spawn_test_upstream_rst(reply); + let ctx = TlsContext::new().unwrap(); + let mut gateway = + HttpsGateway::new_test(&ctx, [127, 0, 0, 1], client_config_trusting(&up_ca)); + + let wire = format!("VPOD-CONNECT localhost {port}\nGET / HTTP/1.0\r\n\r\n"); + gateway.push_from_guest(wire.as_bytes()); + + let mut got = Vec::new(); + for _ in 0..5000 { + drain(&mut gateway, &mut got); + if gateway.eof() { + break; + } + std::thread::sleep(Duration::from_millis(1)); + } + drain(&mut gateway, &mut got); + let _ = up.join(); + + assert!(!gateway.failed(), "RST after data must not fail the bridge"); + assert_eq!( + got.len(), + reply.len(), + "reply truncated: got {} of {} bytes", + got.len(), + reply.len() + ); + } + + #[test] + fn large_response_delivered_in_full_without_truncation() { + const BODY: usize = 1_000_000; + let (port, up_ca, up) = spawn_test_upstream_streaming(BODY); + let ctx = TlsContext::new().unwrap(); + let mut gateway = + HttpsGateway::new_test(&ctx, [127, 0, 0, 1], client_config_trusting(&up_ca)); + + gateway.push_from_guest( + format!("VPOD-CONNECT localhost {port}\nGET / HTTP/1.0\r\n\r\n").as_bytes(), + ); + + let mut got = Vec::new(); + for _ in 0..20000 { + drain(&mut gateway, &mut got); + if gateway.eof() { + break; + } + std::thread::sleep(Duration::from_millis(1)); + } + drain(&mut gateway, &mut got); + let _ = up.join(); + + assert!(!gateway.failed(), "bridge failed on large response"); + let body = got + .windows(4) + .position(|w| w == b"\r\n\r\n") + .map(|i| got.len() - (i + 4)); + assert_eq!(body, Some(BODY), "large body truncated: {body:?} of {BODY}"); + } + + #[test] + fn client_hello_promotes_to_terminating_proxy() { + use rustls::ClientConnection; + use rustls::pki_types::ServerName; + + let ctx = TlsContext::new().unwrap(); + let mut gateway = + HttpsGateway::new_test(&ctx, [127, 0, 0, 1], client_config_trusting(ca_cert_pem())); + + let guest_cfg = client_config_trusting(ca_cert_pem()); + let mut guest = + ClientConnection::new(guest_cfg, ServerName::try_from("localhost").unwrap()).unwrap(); + + let mut handshaken = false; + for _ in 0..200 { + let mut out = Vec::new(); + while guest.wants_write() { + guest.write_tls(&mut out).unwrap(); + } + + for chunk in out.chunks(7) { + gateway.push_from_guest(chunk); + } + + let mut buf = [0u8; 16384]; + while let Some(n) = gateway.pull_to_guest(&mut buf) { + let mut slice = &buf[..n]; + while !slice.is_empty() { + guest.read_tls(&mut slice).unwrap(); + } + guest.process_new_packets().unwrap(); + } + + if !guest.is_handshaking() { + handshaken = true; + break; + } + } + assert!( + handshaken, + "guest handshake did not complete through gateway" + ); + } + + #[test] + fn garbage_first_bytes_fail_like_a_bad_client_hello() { + let ctx = TlsContext::new().unwrap(); + let mut gateway = + HttpsGateway::new_test(&ctx, [127, 0, 0, 1], client_config_trusting(ca_cert_pem())); + gateway.push_from_guest(b"GET / HTTP/1.0\r\n\r\n"); + assert!(gateway.failed(), "plaintext HTTP on :443 must be rejected"); + } + + #[test] + fn malformed_preamble_fails_connection() { + let ctx = TlsContext::new().unwrap(); + let mut gateway = + HttpsGateway::new_test(&ctx, [127, 0, 0, 1], client_config_trusting(ca_cert_pem())); + gateway.push_from_guest(b"VPOD-CONNECT missing-port\n"); + assert!(gateway.failed()); + } + + #[test] + fn parse_preamble_accepts_host_and_port() { + assert_eq!( + parse_preamble(b"VPOD-CONNECT example.com 443\n"), + Some(("example.com".to_string(), 443)) + ); + assert_eq!( + parse_preamble(b"VPOD-CONNECT example.com 443\r\n"), + Some(("example.com".to_string(), 443)) + ); + } + + #[test] + fn parse_preamble_rejects_malformed_lines() { + assert_eq!(parse_preamble(b"VPOD-CONNECT example.com\n"), None); + assert_eq!(parse_preamble(b"VPOD-CONNECT 443\n"), None); + assert_eq!(parse_preamble(b"VPOD-CONNECT a b c\n"), None); + assert_eq!(parse_preamble(b"GET / HTTP/1.0\n"), None); + } +} + +#[cfg(test)] +mod native_probe { + use super::*; + use std::io::{Read, Write}; + use std::net::TcpStream; + use std::time::{Duration, Instant}; + + #[test] + #[ignore] + fn plain_bridge_loop_against_real_host() { + let ctx = TlsContext::new().unwrap(); + let stream = TcpStream::connect_timeout( + &"172.66.147.243:443".parse().unwrap(), + Duration::from_secs(5), + ) + .unwrap(); + stream.set_nonblocking(true).unwrap(); + stream.set_nodelay(true).ok(); + let mut upstream = stream; + let mut client = ClientConnection::new( + ctx.upstream_config(), + ServerName::try_from("example.com".to_string()).unwrap(), + ) + .unwrap(); + client + .writer() + .write_all(b"GET / HTTP/1.0\r\nHost: example.com\r\n\r\n") + .unwrap(); + + let start = Instant::now(); + let mut got = Vec::new(); + while start.elapsed() < Duration::from_secs(8) { + while client.wants_write() { + match client.write_tls(&mut upstream) { + Ok(0) => break, + Ok(n) => println!("wrote {n}"), + Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => break, + Err(e) => panic!("write err {e}"), + } + } + match client.read_tls(&mut upstream) { + Ok(0) => break, + Ok(n) => { + println!("read {n}"); + client.process_new_packets().unwrap(); + } + Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => {} + Err(e) => panic!("read err {e}"), + } + let mut buf = [0u8; 4096]; + if let Ok(n) = client.reader().read(&mut buf) { + got.extend_from_slice(&buf[..n]); + if n > 0 { + println!("plaintext {n}"); + } + } + if got.len() > 100 { + break; + } + std::thread::sleep(Duration::from_millis(2)); + } + println!( + "got {} plaintext bytes: {:?}", + got.len(), + String::from_utf8_lossy(&got[..got.len().min(60)]) + ); + assert!(!got.is_empty()); + } +} diff --git a/crates/machine/src/virtio/mod.rs b/crates/machine/src/virtio/mod.rs index d2d0a69..e3bc51e 100644 --- a/crates/machine/src/virtio/mod.rs +++ b/crates/machine/src/virtio/mod.rs @@ -1,8 +1,10 @@ pub mod blk; pub mod console; pub mod fs; +pub mod https_gateway; pub mod net; pub mod slirp; +pub mod tls_proxy; use crate::RAM_BASE; use crate::cow_ram::CowRam; diff --git a/crates/machine/src/virtio/slirp.rs b/crates/machine/src/virtio/slirp.rs index 305db6d..1ff9259 100644 --- a/crates/machine/src/virtio/slirp.rs +++ b/crates/machine/src/virtio/slirp.rs @@ -1,10 +1,68 @@ +use std::collections::BTreeMap; use std::collections::HashMap; use std::collections::VecDeque; use std::io::{Read, Write}; use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4, TcpStream, ToSocketAddrs, UdpSocket}; use std::time::{Duration, Instant}; +use super::https_gateway::HttpsGateway; use super::net::NetworkBackend; +use super::tls_proxy::TlsContext; + +const HTTPS_PORT: u16 = 443; + +enum Transport { + Raw(TcpStream), + Https(Box), +} + +impl Transport { + fn write(&mut self, buf: &[u8]) -> std::io::Result { + match self { + Transport::Raw(s) => s.write(buf), + Transport::Https(g) => { + g.push_from_guest(buf); + if g.failed() { + Err(std::io::Error::from(std::io::ErrorKind::BrokenPipe)) + } else { + Ok(buf.len()) + } + } + } + } + + fn read(&mut self, buf: &mut [u8]) -> std::io::Result { + match self { + Transport::Raw(s) => s.read(buf), + Transport::Https(g) => { + if g.failed() { + return Err(std::io::Error::from(std::io::ErrorKind::BrokenPipe)); + } + match g.pull_to_guest(buf) { + Some(n) => Ok(n), + None if g.eof() => Ok(0), + None => Err(std::io::Error::from(std::io::ErrorKind::WouldBlock)), + } + } + } + } + + fn has_pending(&self) -> bool { + match self { + Transport::Raw(_) => false, + Transport::Https(g) => g.has_pending(), + } + } + + fn shutdown_write(&mut self) { + match self { + Transport::Raw(s) => { + s.shutdown(std::net::Shutdown::Write).ok(); + } + Transport::Https(g) => g.shutdown_write(), + } + } +} const GW_IP: [u8; 4] = [10, 0, 2, 2]; const GUEST_IP: [u8; 4] = [10, 0, 2, 15]; @@ -69,7 +127,7 @@ enum TcpState { struct TcpConn { state: TcpState, - stream: TcpStream, + transport: Transport, guest_mac: [u8; 6], src_ip: [u8; 4], dst_ip: [u8; 4], @@ -85,6 +143,8 @@ struct TcpConn { rcv_nxt: u32, rcv_wnd: u32, wnd_shift: u8, + + ooo_buf: BTreeMap>, } #[derive(Debug, Clone, PartialEq, Eq, Hash)] @@ -125,10 +185,19 @@ pub struct SlirpBackend { udp_conns: HashMap, dns_pending: Vec, dhcp_xid: u32, + tls: Option, } impl SlirpBackend { pub fn new(guest_mac: [u8; 6]) -> Self { + let tls = match TlsContext::new() { + Ok(ctx) => Some(ctx), + Err(e) => { + log::warn!("tls_proxy: disabled, terminator init failed: {e}"); + None + } + }; + Self { guest_mac, rx_pending: VecDeque::new(), @@ -136,6 +205,7 @@ impl SlirpBackend { udp_conns: HashMap::new(), dns_pending: Vec::new(), dhcp_xid: 0, + tls, } } @@ -151,7 +221,7 @@ impl SlirpBackend { while !conn.write_buf.is_empty() { let (a, b) = conn.write_buf.as_slices(); let slice = if !a.is_empty() { a } else { b }; - match conn.stream.write(slice) { + match conn.transport.write(slice) { Ok(n) => { conn.write_buf.drain(..n); } @@ -170,7 +240,7 @@ impl SlirpBackend { if !remove { let mut buf = [0u8; 16384]; loop { - match conn.stream.read(&mut buf) { + match conn.transport.read(&mut buf) { Ok(0) => { if conn.snd_buf.is_empty() { frames.push(make_tcp_frame( @@ -261,6 +331,35 @@ impl SlirpBackend { } } + fn drain_ooo_buf(conn: &mut TcpConn) { + loop { + let mut progressed = false; + let seqs: Vec = conn.ooo_buf.keys().cloned().collect(); + for seq in seqs { + let data_len = conn.ooo_buf[&seq].len() as u32; + let end_seq = seq.wrapping_add(data_len); + let gap = seq.wrapping_sub(conn.rcv_nxt) as i32; + let new_bytes = end_seq.wrapping_sub(conn.rcv_nxt) as i32; + + if gap > 0 { + continue; + } + + let data = conn.ooo_buf.remove(&seq).unwrap(); + if new_bytes > 0 { + let skip = data.len() - new_bytes as usize; + conn.write_buf.extend(&data[skip..]); + conn.rcv_nxt = end_seq; + progressed = true; + } + } + + if !progressed { + break; + } + } + } + fn drain_snd_buf(conn: &mut TcpConn, frames: &mut Vec>) { loop { let in_flight = conn.snd_nxt.wrapping_sub(conn.snd_una); @@ -598,33 +697,40 @@ impl SlirpBackend { let wnd_shift = parse_wnd_scale(payload, tcp_hlen); - let addr = SocketAddrV4::new(Ipv4Addr::from(dst_ip), dst_port); - - #[cfg(target_family = "wasm")] - let stream_result = TcpStream::connect(addr); - - #[cfg(not(target_family = "wasm"))] - let stream_result = TcpStream::connect_timeout(&addr.into(), Duration::from_secs(5)); - - let stream = match stream_result { - Ok(s) => s, - Err(_) => { - self.rx_pending.push_back(make_tcp_frame( - &src_mac, - &dst_ip, - &src_ip, - dst_port, - src_port, - 0, - seq_guest.wrapping_add(1), - RST | ACK, - &[], - )); - return; - } + let tls_ctx = self.tls.as_ref().filter(|_| dst_port == HTTPS_PORT); + let transport = if let Some(ctx) = tls_ctx { + Transport::Https(Box::new(HttpsGateway::new(ctx, dst_ip))) + } else { + let addr = SocketAddrV4::new(Ipv4Addr::from(dst_ip), dst_port); + + #[cfg(target_family = "wasm")] + let stream_result = TcpStream::connect(addr); + + #[cfg(not(target_family = "wasm"))] + let stream_result = + TcpStream::connect_timeout(&addr.into(), Duration::from_secs(5)); + + let stream = match stream_result { + Ok(s) => s, + Err(_) => { + self.rx_pending.push_back(make_tcp_frame( + &src_mac, + &dst_ip, + &src_ip, + dst_port, + src_port, + 0, + seq_guest.wrapping_add(1), + RST | ACK, + &[], + )); + return; + } + }; + stream.set_nonblocking(true).ok(); + stream.set_nodelay(true).ok(); + Transport::Raw(stream) }; - stream.set_nonblocking(true).ok(); - stream.set_nodelay(true).ok(); let isn_host: u32 = generate_isn(&src_ip, src_port, &dst_ip, dst_port); @@ -644,7 +750,7 @@ impl SlirpBackend { key, TcpConn { state: TcpState::Established, - stream, + transport, guest_mac: src_mac, src_ip, dst_ip, @@ -657,6 +763,7 @@ impl SlirpBackend { rcv_nxt: seq_guest.wrapping_add(1), rcv_wnd: (window as u32) << wnd_shift, wnd_shift, + ooo_buf: BTreeMap::new(), }, ); @@ -678,11 +785,22 @@ impl SlirpBackend { if !data.is_empty() { let end_seq = seq_guest.wrapping_add(data.len() as u32); let new_bytes = end_seq.wrapping_sub(conn.rcv_nxt) as i32; + let gap = seq_guest.wrapping_sub(conn.rcv_nxt) as i32; - if new_bytes > 0 { + if new_bytes > 0 && gap <= 0 { let skip = data.len() - new_bytes as usize; conn.write_buf.extend(&data[skip..]); conn.rcv_nxt = end_seq; + Self::drain_ooo_buf(conn); + } else if gap > 0 { + const MAX_OOO_BYTES: usize = 256 * 1024; + let buffered: usize = conn.ooo_buf.values().map(Vec::len).sum(); + + if buffered + data.len() <= MAX_OOO_BYTES { + conn.ooo_buf + .entry(seq_guest) + .or_insert_with(|| data.to_vec()); + } } self.rx_pending.push_back(make_tcp_frame( @@ -711,7 +829,14 @@ impl SlirpBackend { ACK, &[], )); - conn.state = TcpState::Closed; + match conn.state { + TcpState::Established => { + conn.transport.shutdown_write(); + } + TcpState::FinWait | TcpState::Closed => { + conn.state = TcpState::Closed; + } + } } } else { self.rx_pending.push_back(make_tcp_frame( @@ -757,7 +882,7 @@ impl NetworkBackend for SlirpBackend { for conn in self.tcp_conns.values() { match conn.state { TcpState::Established | TcpState::FinWait => { - if !conn.snd_buf.is_empty() { + if !conn.snd_buf.is_empty() || conn.transport.has_pending() { return true; } } @@ -1175,4 +1300,59 @@ mod tests { ); assert_eq!(u16::from_be_bytes([reply[6], reply[7]]), 0); // no answers } + + #[test] + fn out_of_order_guest_segments_are_reassembled_before_upstream_write() { + let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap(); + let port = listener.local_addr().unwrap().port(); + + let guest_mac = [0x02, 0, 0, 0, 0, 0x15]; + let mut slirp = SlirpBackend::new(guest_mac); + let dst_ip = [127, 0, 0, 1]; + let sport = 45001u16; + let guest_isn = 1000u32; + + slirp.send(&make_tcp_frame( + &guest_mac, + &GUEST_IP, + &dst_ip, + sport, + port, + guest_isn, + 0, + SYN, + &[], + )); + let (mut upstream, _) = listener.accept().unwrap(); + upstream.set_nonblocking(true).unwrap(); + + let syn_ack = slirp.recv().expect("SYN-ACK"); + let host_isn = u32::from_be_bytes(syn_ack[38..42].try_into().unwrap()); + let ack = host_isn.wrapping_add(1); + + let part1 = b"hello "; + let part2 = b"world"; + let seq1 = guest_isn.wrapping_add(1); + let seq2 = seq1.wrapping_add(part1.len() as u32); + + slirp.send(&make_tcp_frame( + &guest_mac, &GUEST_IP, &dst_ip, sport, port, seq2, ack, ACK, part2, + )); + slirp.send(&make_tcp_frame( + &guest_mac, &GUEST_IP, &dst_ip, sport, port, seq1, ack, ACK, part1, + )); + + let mut got = Vec::new(); + let deadline = Instant::now() + Duration::from_secs(2); + while got.len() < part1.len() + part2.len() && Instant::now() < deadline { + while slirp.recv().is_some() {} + let mut buf = [0u8; 64]; + match upstream.read(&mut buf) { + Ok(0) => break, + Ok(n) => got.extend_from_slice(&buf[..n]), + Err(_) => std::thread::sleep(Duration::from_millis(5)), + } + } + assert_eq!(got, b"hello world"); + } } diff --git a/crates/machine/src/virtio/tls_proxy/mod.rs b/crates/machine/src/virtio/tls_proxy/mod.rs new file mode 100644 index 0000000..75d548c --- /dev/null +++ b/crates/machine/src/virtio/tls_proxy/mod.rs @@ -0,0 +1,601 @@ +// Transparent TLS-terminating proxy and prepare transfer to the host + +use std::collections::HashMap; +use std::collections::VecDeque; +use std::io::{Read, Write}; +use std::net::{Ipv4Addr, SocketAddrV4, TcpStream}; +use std::str::FromStr; +use std::sync::{Arc, Mutex}; +use std::time::{Duration, Instant, SystemTime}; + +use const_oid::db::rfc5280::ID_KP_SERVER_AUTH; +use der::asn1::Ia5String; +use der::{DecodePem, Encode}; +use p256::ecdsa::{DerSignature, SigningKey}; +use p256::pkcs8::{DecodePrivateKey, EncodePrivateKey}; +use rand_core::{OsRng, RngCore}; +use rustls::crypto::CryptoProvider; +use rustls::pki_types::{CertificateDer, PrivateKeyDer, ServerName}; +use rustls::server::{ClientHello, ResolvesServerCert}; +use rustls::sign::CertifiedKey; +use rustls::{ClientConfig, ClientConnection, RootCertStore, ServerConfig, ServerConnection}; +use x509_cert::Certificate; +use x509_cert::builder::{Builder, CertificateBuilder, Profile}; +use x509_cert::ext::pkix::name::GeneralName; +use x509_cert::ext::pkix::{ExtendedKeyUsage, SubjectAltName}; +use x509_cert::name::Name; +use x509_cert::serial_number::SerialNumber; +use x509_cert::spki::SubjectPublicKeyInfoOwned; +use x509_cert::time::{Time, Validity}; + +#[cfg(test)] +mod testutil; + +#[cfg(test)] +mod tests; + +#[cfg(test)] +pub(crate) use testutil::{ + client_config_trusting, spawn_test_upstream, spawn_test_upstream_rst, + spawn_test_upstream_streaming, +}; + +const CA_KEY_PEM: &str = include_str!("../../../assets/tls/vpod-ca-key.pem"); +pub const CA_CERT_PEM: &str = include_str!("../../../assets/tls/vpod-ca-cert.pem"); + +#[derive(Clone)] +pub struct TlsContext { + server_config: Arc, + client_config: Arc, +} + +impl TlsContext { + pub fn new() -> Result { + let mut provider = rustls_rustcrypto::provider(); + provider + .cipher_suites + .sort_by_key(|cs| u8::from(!format!("{:?}", cs.suite()).contains("CHACHA20"))); + let provider = Arc::new(provider); + + let resolver = Arc::new(SniResolver::from_pems( + provider.clone(), + CA_KEY_PEM, + CA_CERT_PEM, + )?); + + let mut server_config = ServerConfig::builder_with_provider(provider.clone()) + .with_safe_default_protocol_versions() + .map_err(|e| e.to_string())? + .with_no_client_auth() + .with_cert_resolver(resolver); + + server_config.ignore_client_order = true; + + let mut roots = RootCertStore::empty(); + roots.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned()); + let client_config = ClientConfig::builder_with_provider(provider) + .with_safe_default_protocol_versions() + .map_err(|e| e.to_string())? + .with_root_certificates(roots) + .with_no_client_auth(); + + Ok(Self { + server_config: Arc::new(server_config), + client_config: Arc::new(client_config), + }) + } + + pub(crate) fn upstream_config(&self) -> Arc { + self.client_config.clone() + } +} + +struct SniResolver { + provider: Arc, + ca_key: SigningKey, + ca_issuer: Name, + cache: Mutex>>, +} + +impl std::fmt::Debug for SniResolver { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("SniResolver").finish() + } +} + +impl SniResolver { + fn from_pems( + provider: Arc, + ca_key_pem: &str, + ca_cert_pem: &str, + ) -> Result { + let ca_key = SigningKey::from_pkcs8_pem(ca_key_pem).map_err(|e| e.to_string())?; + let ca_cert = Certificate::from_pem(ca_cert_pem).map_err(|e| e.to_string())?; + Ok(Self { + provider, + ca_key, + ca_issuer: ca_cert.tbs_certificate.subject, + cache: Mutex::new(HashMap::new()), + }) + } + + fn certified_key_for(&self, sni: &str) -> Result, String> { + if let Some(hit) = self.cache.lock().unwrap().get(sni) { + return Ok(hit.clone()); + } + + let (leaf_der, key_der) = self.mint_leaf(sni)?; + let signing_key = self + .provider + .key_provider + .load_private_key(PrivateKeyDer::try_from(key_der).map_err(|e| e.to_string())?) + .map_err(|e| e.to_string())?; + let certified = Arc::new(CertifiedKey::new( + vec![CertificateDer::from(leaf_der)], + signing_key, + )); + + self.cache + .lock() + .unwrap() + .insert(sni.to_string(), certified.clone()); + Ok(certified) + } + + fn mint_leaf(&self, sni: &str) -> Result<(Vec, Vec), String> { + let leaf_key = SigningKey::random(&mut OsRng); + let spki = SubjectPublicKeyInfoOwned::from_key(*leaf_key.verifying_key()) + .map_err(|e| e.to_string())?; + let subject = Name::from_str(&format!("CN={sni}")).map_err(|e| e.to_string())?; + + let now = SystemTime::now(); + let not_before = + Time::try_from(now - Duration::from_secs(24 * 3600)).map_err(|e| e.to_string())?; + let not_after = Time::try_from(now + Duration::from_secs(825 * 24 * 3600)) + .map_err(|e| e.to_string())?; + let validity = Validity { + not_before, + not_after, + }; + + let serial = SerialNumber::from(rand_serial()); + + let mut builder = CertificateBuilder::new( + Profile::Leaf { + issuer: self.ca_issuer.clone(), + enable_key_agreement: false, + enable_key_encipherment: false, + }, + serial, + validity, + subject, + spki, + &self.ca_key, + ) + .map_err(|e| e.to_string())?; + + builder + .add_extension(&SubjectAltName(vec![GeneralName::DnsName( + Ia5String::new(sni).map_err(|e| e.to_string())?, + )])) + .map_err(|e| e.to_string())?; + builder + .add_extension(&ExtendedKeyUsage(vec![ID_KP_SERVER_AUTH])) + .map_err(|e| e.to_string())?; + + let leaf: Certificate = builder.build::().map_err(|e| e.to_string())?; + let leaf_der = leaf.to_der().map_err(|e| e.to_string())?; + let key_der = leaf_key + .to_pkcs8_der() + .map_err(|e| e.to_string())? + .as_bytes() + .to_vec(); + Ok((leaf_der, key_der)) + } +} + +impl ResolvesServerCert for SniResolver { + fn resolve(&self, client_hello: ClientHello) -> Option> { + let sni = client_hello.server_name()?; + match self.certified_key_for(sni) { + Ok(ck) => Some(ck), + Err(e) => { + log::warn!("tls_proxy: leaf minting failed for {sni}: {e}"); + None + } + } + } +} + +fn rand_serial() -> u64 { + let mut b = [0u8; 8]; + OsRng.fill_bytes(&mut b); + (u64::from_be_bytes(b) >> 1) | 1 +} + +pub struct TlsProxy { + server: ServerConnection, + client: Option, + upstream: Option, + dst_ip: [u8; 4], + upstream_port: u16, + client_config: Arc, + to_guest: VecDeque, + failed: bool, + upstream_closed: bool, + close_notified: bool, + timing: Option, +} + +pub(crate) struct Timing { + start: Instant, + first_guest_bytes: bool, + serverhello_sent: bool, + guest_hs_done: bool, + upstream_connected: bool, + upstream_hs_done: bool, + first_reply: bool, +} + +impl Timing { + pub(crate) fn new() -> Option { + match std::env::var("VPOD_TLS_TIMING") { + Ok(v) if v != "0" && !v.is_empty() => Some(Self { + start: Instant::now(), + first_guest_bytes: false, + serverhello_sent: false, + guest_hs_done: false, + upstream_connected: false, + upstream_hs_done: false, + first_reply: false, + }), + _ => None, + } + } + + pub(crate) fn mark(&self, label: &str) { + eprintln!( + "[tls-timing] +{:>6.1}ms {label}", + self.start.elapsed().as_secs_f64() * 1000.0 + ); + } +} + +impl Drop for TlsProxy { + fn drop(&mut self) { + if let Some(t) = &self.timing { + t.mark("connection closed (proxy dropped)"); + } + } +} + +impl TlsProxy { + pub fn new(ctx: &TlsContext, dst_ip: [u8; 4]) -> Result { + Self::with_timing(ctx, dst_ip, Timing::new()) + } + + pub(crate) fn with_timing( + ctx: &TlsContext, + dst_ip: [u8; 4], + timing: Option, + ) -> Result { + let mut server = + ServerConnection::new(ctx.server_config.clone()).map_err(|e| e.to_string())?; + server.set_buffer_limit(None); + + Ok(Self { + server, + client: None, + upstream: None, + dst_ip, + upstream_port: 443, + client_config: ctx.client_config.clone(), + to_guest: VecDeque::new(), + failed: false, + upstream_closed: false, + close_notified: false, + timing, + }) + } + + pub fn failed(&self) -> bool { + self.failed + } + + pub fn push_from_guest(&mut self, mut bytes: &[u8]) { + if let Some(t) = &mut self.timing + && !t.first_guest_bytes + && !bytes.is_empty() + { + t.first_guest_bytes = true; + t.mark("first guest bytes (ClientHello)"); + } + while !bytes.is_empty() { + match self.server.read_tls(&mut bytes) { + Ok(0) => break, + Ok(_) => {} + Err(_) => { + self.failed = true; + return; + } + } + } + self.pump(); + } + + pub fn pull_to_guest(&mut self, buf: &mut [u8]) -> Option { + self.pump(); + if self.to_guest.is_empty() { + return None; + } + let n = self.to_guest.len().min(buf.len()); + for slot in buf.iter_mut().take(n) { + *slot = self.to_guest.pop_front().unwrap(); + } + Some(n) + } + + pub fn has_pending(&self) -> bool { + !self.to_guest.is_empty() || self.server.wants_write() + } + + fn pump(&mut self) { + if self.failed { + return; + } + + if let Err(e) = self.server.process_new_packets() { + log::warn!("tls_proxy: guest-side handshake failed: {e}"); + self.failed = true; + return; + } + + if let Some(t) = &mut self.timing + && !t.guest_hs_done + && !self.server.is_handshaking() + { + t.guest_hs_done = true; + t.mark("guest handshake done"); + } + + if self.client.is_none() + && !self.server.is_handshaking() + && let Some(sni) = self.server.server_name().map(|s| s.to_string()) + { + self.connect_upstream(&sni); + + if let (Some(t), true) = (&mut self.timing, self.client.is_some()) + && !t.upstream_connected + { + t.upstream_connected = true; + t.mark("upstream TCP connected"); + } + } + + if let Some(client) = self.client.as_mut() { + let mut buf = [0u8; 16384]; + + loop { + match self.server.reader().read(&mut buf) { + Ok(0) => break, + Ok(n) => { + if client.writer().write_all(&buf[..n]).is_err() { + self.failed = true; + return; + } + } + Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => break, + Err(_) => break, + } + } + } + + self.service_upstream(); + + if let Some(client) = self.client.as_mut() { + let mut buf = [0u8; 16384]; + loop { + match client.reader().read(&mut buf) { + Ok(0) => break, + Ok(n) => { + if self.server.writer().write_all(&buf[..n]).is_err() { + self.failed = true; + return; + } + } + Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => break, + Err(_) => break, + } + } + } + + if self.upstream_closed && !self.close_notified { + self.server.send_close_notify(); + self.close_notified = true; + } + + let mut wrote_to_guest = false; + while self.server.wants_write() { + let mut out = Vec::new(); + match self.server.write_tls(&mut out) { + Ok(0) => break, + Ok(_) => { + wrote_to_guest = true; + self.to_guest.extend(out); + } + Err(_) => { + self.failed = true; + break; + } + } + } + + if let Some(t) = &mut self.timing + && !t.serverhello_sent + && wrote_to_guest + && self.server.is_handshaking() + { + t.serverhello_sent = true; + t.mark("ServerHello flight sent to guest"); + } + + if let Some(t) = &mut self.timing + && t.upstream_hs_done + && !t.first_reply + && !self.to_guest.is_empty() + { + t.first_reply = true; + t.mark("first reply byte to guest"); + } + } + + fn connect_upstream(&mut self, sni: &str) { + let addr = SocketAddrV4::new(Ipv4Addr::from(self.dst_ip), self.upstream_port); + + #[cfg(target_family = "wasm")] + let stream = TcpStream::connect(addr); + + #[cfg(not(target_family = "wasm"))] + let stream = TcpStream::connect_timeout(&addr.into(), Duration::from_secs(10)); + + let stream = match stream { + Ok(s) => s, + Err(_) => { + self.failed = true; + return; + } + }; + stream.set_nonblocking(true).ok(); + stream.set_nodelay(true).ok(); + + let server_name = match ServerName::try_from(sni.to_string()) { + Ok(n) => n, + Err(_) => { + self.failed = true; + return; + } + }; + + match ClientConnection::new(self.client_config.clone(), server_name) { + Ok(mut c) => { + c.set_buffer_limit(None); + self.client = Some(c); + self.upstream = Some(stream); + } + Err(_) => self.failed = true, + } + } + + fn service_upstream(&mut self) { + if self.client.is_none() || self.upstream.is_none() { + return; + } + + let handshook = !self.client.as_ref().unwrap().is_handshaking(); + while !self.upstream_closed && self.client.as_ref().unwrap().wants_write() { + let sock = self.upstream.as_mut().unwrap(); + match self.client.as_mut().unwrap().write_tls(sock) { + Ok(0) => break, + Ok(_) => {} + Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => break, + Err(_) => { + if handshook { + self.upstream_closed = true; + break; + } + + self.failed = true; + return; + } + } + } + + while !self.upstream_closed { + let sock = self.upstream.as_mut().unwrap(); + match self.client.as_mut().unwrap().read_tls(sock) { + Ok(0) => { + self.upstream_closed = true; + } + Ok(_) => { + if self.client.as_mut().unwrap().process_new_packets().is_err() { + if handshook { + self.upstream_closed = true; + } else { + self.failed = true; + return; + } + } + } + Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => break, + Err(_) => { + if handshook { + self.upstream_closed = true; + } else { + self.failed = true; + return; + } + } + } + + let mut buf = [0u8; 16384]; + loop { + match self.client.as_mut().unwrap().reader().read(&mut buf) { + Ok(0) => break, + Ok(n) => { + if self.server.writer().write_all(&buf[..n]).is_err() { + self.failed = true; + return; + } + } + Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => break, + Err(_) => break, + } + } + } + + let upstream_handshaking = self + .client + .as_ref() + .map(|c| c.is_handshaking()) + .unwrap_or(true); + if let Some(t) = &mut self.timing + && t.upstream_connected + && !t.upstream_hs_done + && !upstream_handshaking + { + t.upstream_hs_done = true; + t.mark("upstream handshake done"); + } + } +} + +pub fn ca_cert_pem() -> &'static str { + CA_CERT_PEM +} + +#[cfg(test)] +impl TlsProxy { + fn new_test( + server_config: Arc, + client_config: Arc, + dst_ip: [u8; 4], + upstream_port: u16, + ) -> Self { + let mut server = ServerConnection::new(server_config).unwrap(); + server.set_buffer_limit(None); + + Self { + server, + client: None, + upstream: None, + dst_ip, + upstream_port, + client_config, + to_guest: VecDeque::new(), + failed: false, + upstream_closed: false, + close_notified: false, + timing: None, + } + } +} diff --git a/crates/machine/src/virtio/tls_proxy/tests.rs b/crates/machine/src/virtio/tls_proxy/tests.rs new file mode 100644 index 0000000..dbc52c1 --- /dev/null +++ b/crates/machine/src/virtio/tls_proxy/tests.rs @@ -0,0 +1,816 @@ +use super::testutil::{ + client_config_trusting, generate_ca_pems, mint_leaf_with_ca, spawn_test_upstream, + spawn_test_upstream_streaming, +}; +use super::*; +use std::net::TcpListener; +use std::thread; + +const UPSTREAM_REPLY: &[u8] = b"HTTP/1.0 200 OK\r\nContent-Length: 5\r\n\r\nhello"; + +fn provider() -> Arc { + Arc::new(rustls_rustcrypto::provider()) +} + +#[test] +fn resolver_mints_loadable_leaf_for_sni() { + let ctx = TlsContext::new().expect("terminator init"); + let _ = ServerConnection::new(ctx.server_config.clone()).unwrap(); +} + +#[test] +fn end_to_end_bridge_delivers_upstream_reply() { + // its own CA + a localhost leaf the proxy + let (up_ca_key, up_ca_cert) = generate_ca_pems(); + let (leaf_der, leaf_key) = mint_leaf_with_ca(&up_ca_key, &up_ca_cert, "localhost"); + + let up_server_cfg = ServerConfig::builder_with_provider(provider()) + .with_safe_default_protocol_versions() + .unwrap() + .with_no_client_auth() + .with_single_cert( + vec![CertificateDer::from(leaf_der)], + PrivateKeyDer::try_from(leaf_key).unwrap(), + ) + .unwrap(); + let up_server_cfg = Arc::new(up_server_cfg); + + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let port = listener.local_addr().unwrap().port(); + + let up = thread::spawn(move || { + listener.set_nonblocking(false).ok(); + let (mut sock, _) = listener.accept().unwrap(); + sock.set_read_timeout(Some(Duration::from_secs(10))).ok(); + let mut conn = ServerConnection::new(up_server_cfg).unwrap(); + + let mut req = Vec::new(); + loop { + if conn.wants_write() { + conn.write_tls(&mut sock).unwrap(); + continue; + } + + if conn.is_handshaking() { + conn.read_tls(&mut sock).unwrap(); + conn.process_new_packets().unwrap(); + continue; + } + + let mut buf = [0u8; 1024]; + match conn.reader().read(&mut buf) { + Ok(n) if n > 0 => { + req.extend_from_slice(&buf[..n]); + if req.windows(4).any(|w| w == b"\r\n\r\n") { + break; + } + continue; + } + _ => {} + } + conn.read_tls(&mut sock).unwrap(); + conn.process_new_packets().unwrap(); + } + + conn.writer().write_all(UPSTREAM_REPLY).unwrap(); + while conn.wants_write() { + conn.write_tls(&mut sock).unwrap(); + } + + conn.send_close_notify(); + while conn.wants_write() { + conn.write_tls(&mut sock).unwrap(); + } + }); + + let ctx = TlsContext::new().unwrap(); + let mut up_roots = RootCertStore::empty(); + let up_ca_der = Certificate::from_pem(&up_ca_cert) + .unwrap() + .to_der() + .unwrap(); + + up_roots.add(CertificateDer::from(up_ca_der)).unwrap(); + let proxy_client_cfg = Arc::new( + ClientConfig::builder_with_provider(provider()) + .with_safe_default_protocol_versions() + .unwrap() + .with_root_certificates(up_roots) + .with_no_client_auth(), + ); + + let mut proxy = TlsProxy::new_test( + ctx.server_config.clone(), + proxy_client_cfg, + [127, 0, 0, 1], + port, + ); + + let mut guest_roots = RootCertStore::empty(); + let vpod_ca_der = Certificate::from_pem(CA_CERT_PEM) + .unwrap() + .to_der() + .unwrap(); + guest_roots.add(CertificateDer::from(vpod_ca_der)).unwrap(); + let guest_cfg = Arc::new( + ClientConfig::builder_with_provider(provider()) + .with_safe_default_protocol_versions() + .unwrap() + .with_root_certificates(guest_roots) + .with_no_client_auth(), + ); + + let mut guest = + ClientConnection::new(guest_cfg, ServerName::try_from("localhost").unwrap()).unwrap(); + + let mut request_sent = false; + let mut got = Vec::new(); + for _ in 0..2000 { + let mut out = Vec::new(); + while guest.wants_write() { + guest.write_tls(&mut out).unwrap(); + } + if !out.is_empty() { + proxy.push_from_guest(&out); + } + + let mut buf = [0u8; 16384]; + while let Some(n) = proxy.pull_to_guest(&mut buf) { + let mut slice = &buf[..n]; + while !slice.is_empty() { + guest.read_tls(&mut slice).unwrap(); + } + guest.process_new_packets().unwrap(); + } + + if !request_sent && !guest.is_handshaking() { + guest.writer().write_all(b"GET / HTTP/1.0\r\n\r\n").unwrap(); + request_sent = true; + } + + let mut buf = [0u8; 1024]; + if let Ok(n) = guest.reader().read(&mut buf) + && n > 0 + { + got.extend_from_slice(&buf[..n]); + } + + if got.windows(5).any(|w| w == b"hello") { + break; + } + + if proxy.failed() { + panic!("proxy failed before delivering reply"); + } + thread::sleep(Duration::from_millis(1)); + } + + let _ = up.join(); + assert!( + got.windows(5).any(|w| w == b"hello"), + "did not receive upstream reply, got {got:?}" + ); +} + +#[test] +fn ring_backed_client_interops_with_rustcrypto_server() { + let (port, up_ca, up) = spawn_test_upstream(UPSTREAM_REPLY); + + let ctx = TlsContext::new().unwrap(); + let mut proxy = TlsProxy::new_test( + ctx.server_config.clone(), + client_config_trusting(&up_ca), + [127, 0, 0, 1], + port, + ); + + let mut guest_roots = RootCertStore::empty(); + let vpod_ca_der = Certificate::from_pem(CA_CERT_PEM) + .unwrap() + .to_der() + .unwrap(); + guest_roots.add(CertificateDer::from(vpod_ca_der)).unwrap(); + let ring_cfg = Arc::new( + ClientConfig::builder_with_provider(Arc::new(rustls::crypto::ring::default_provider())) + .with_safe_default_protocol_versions() + .unwrap() + .with_root_certificates(guest_roots) + .with_no_client_auth(), + ); + + let mut guest = + ClientConnection::new(ring_cfg, ServerName::try_from("localhost").unwrap()).unwrap(); + + let mut request_sent = false; + let mut got = Vec::new(); + for _ in 0..2000 { + let mut out = Vec::new(); + while guest.wants_write() { + guest.write_tls(&mut out).unwrap(); + } + if !out.is_empty() { + proxy.push_from_guest(&out); + } + + let mut buf = [0u8; 16384]; + while let Some(n) = proxy.pull_to_guest(&mut buf) { + let mut slice = &buf[..n]; + while !slice.is_empty() { + guest.read_tls(&mut slice).unwrap(); + guest.process_new_packets().unwrap(); + } + } + + if !request_sent && !guest.is_handshaking() { + guest.writer().write_all(b"GET / HTTP/1.0\r\n\r\n").unwrap(); + request_sent = true; + } + + let mut buf = [0u8; 1024]; + if let Ok(n) = guest.reader().read(&mut buf) + && n > 0 + { + got.extend_from_slice(&buf[..n]); + } + + if got.windows(5).any(|w| w == b"hello") { + break; + } + + if proxy.failed() { + panic!("proxy failed against a ring-backed client"); + } + thread::sleep(Duration::from_millis(1)); + } + + let _ = up.join(); + assert!( + got.windows(5).any(|w| w == b"hello"), + "did not receive upstream reply via ring client, got {got:?}" + ); +} + +#[test] +fn hello_retry_request_path_completes() { + // Modern rustls clients (uv via aws-lc-rs) lead with an X25519MLKEM768 + // post-quantum keyshare. Our rustcrypto server doesn't support it, so + // it must send a HelloRetryRequest and complete on a classic group. + // OpenSSL clients lead with X25519 and never take this path — an HRR + // bug in the alpha rustcrypto provider only bites rustls clients. + let (port, up_ca, up) = spawn_test_upstream(UPSTREAM_REPLY); + + let ctx = TlsContext::new().unwrap(); + let mut proxy = TlsProxy::new_test( + ctx.server_config.clone(), + client_config_trusting(&up_ca), + [127, 0, 0, 1], + port, + ); + + let mut guest_roots = RootCertStore::empty(); + let vpod_ca_der = Certificate::from_pem(CA_CERT_PEM) + .unwrap() + .to_der() + .unwrap(); + guest_roots.add(CertificateDer::from(vpod_ca_der)).unwrap(); + + let mut awslc_provider = rustls::crypto::aws_lc_rs::default_provider(); + awslc_provider.kx_groups = vec![ + rustls::crypto::aws_lc_rs::kx_group::X25519MLKEM768, + rustls::crypto::aws_lc_rs::kx_group::X25519, + rustls::crypto::aws_lc_rs::kx_group::SECP256R1, + ]; + let hrr_cfg = Arc::new( + ClientConfig::builder_with_provider(Arc::new(awslc_provider)) + .with_safe_default_protocol_versions() + .unwrap() + .with_root_certificates(guest_roots) + .with_no_client_auth(), + ); + + let mut guest = + ClientConnection::new(hrr_cfg, ServerName::try_from("localhost").unwrap()).unwrap(); + + let mut request_sent = false; + let mut got = Vec::new(); + for _ in 0..2000 { + let mut out = Vec::new(); + while guest.wants_write() { + guest.write_tls(&mut out).unwrap(); + } + if !out.is_empty() { + proxy.push_from_guest(&out); + } + + let mut buf = [0u8; 16384]; + while let Some(n) = proxy.pull_to_guest(&mut buf) { + let mut slice = &buf[..n]; + while !slice.is_empty() { + guest.read_tls(&mut slice).unwrap(); + guest.process_new_packets().unwrap(); + } + } + + if !request_sent && !guest.is_handshaking() { + guest.writer().write_all(b"GET / HTTP/1.0\r\n\r\n").unwrap(); + request_sent = true; + } + + let mut buf = [0u8; 1024]; + if let Ok(n) = guest.reader().read(&mut buf) + && n > 0 + { + got.extend_from_slice(&buf[..n]); + } + + if got.windows(5).any(|w| w == b"hello") { + break; + } + + if proxy.failed() { + panic!("proxy failed on the HelloRetryRequest path"); + } + thread::sleep(Duration::from_millis(1)); + } + + let _ = up.join(); + assert!( + got.windows(5).any(|w| w == b"hello"), + "HRR handshake did not complete, got {got:?}" + ); +} + +#[test] +fn real_uv_binary_handshakes_with_proxy() { + use std::io::ErrorKind; + use std::process::Command; + use std::sync::atomic::{AtomicBool, Ordering}; + + if std::env::var("VPOD_UV_REPRO").is_err() { + eprintln!("skipping real-uv repro (set VPOD_UV_REPRO=1 to run)"); + return; + } + + let (up_ca_key, up_ca_cert) = generate_ca_pems(); + let (leaf_der, leaf_key) = mint_leaf_with_ca(&up_ca_key, &up_ca_cert, "localhost"); + let up_cfg = Arc::new( + ServerConfig::builder_with_provider(provider()) + .with_safe_default_protocol_versions() + .unwrap() + .with_no_client_auth() + .with_single_cert( + vec![CertificateDer::from(leaf_der)], + PrivateKeyDer::try_from(leaf_key).unwrap(), + ) + .unwrap(), + ); + + let up_listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let up_port = up_listener.local_addr().unwrap().port(); + thread::spawn(move || { + for sock in up_listener.incoming().flatten() { + let cfg = up_cfg.clone(); + thread::spawn(move || { + let mut sock = sock; + sock.set_read_timeout(Some(Duration::from_secs(10))).ok(); + let mut conn = ServerConnection::new(cfg).unwrap(); + let mut req = Vec::new(); + loop { + if conn.wants_write() { + if conn.write_tls(&mut sock).is_err() { + return; + } + continue; + } + let mut buf = [0u8; 4096]; + match conn.reader().read(&mut buf) { + Ok(n) if n > 0 => { + req.extend_from_slice(&buf[..n]); + if req.windows(4).any(|w| w == b"\r\n\r\n") { + break; + } + continue; + } + _ => {} + } + if conn.read_tls(&mut sock).is_err() || conn.process_new_packets().is_err() { + return; + } + } + conn.writer() + .write_all( + b"HTTP/1.1 404 Not Found\r\nContent-Length: 0\r\nConnection: close\r\n\r\n", + ) + .unwrap(); + conn.send_close_notify(); + while conn.wants_write() { + if conn.write_tls(&mut sock).is_err() { + return; + } + } + }); + } + }); + + let ctx = TlsContext::new().unwrap(); + let proxy_client_cfg = client_config_trusting(&up_ca_cert); + let front = TcpListener::bind("127.0.0.1:0").unwrap(); + let front_port = front.local_addr().unwrap().port(); + let any_proxy_failed = Arc::new(AtomicBool::new(false)); + + { + let ctx = ctx.clone(); + let failed_flag = any_proxy_failed.clone(); + thread::spawn(move || { + for sock in front.incoming().flatten() { + let mut proxy = TlsProxy::new_test( + ctx.server_config.clone(), + proxy_client_cfg.clone(), + [127, 0, 0, 1], + up_port, + ); + let failed_flag = failed_flag.clone(); + thread::spawn(move || { + let mut sock = sock; + sock.set_nonblocking(true).ok(); + let mut buf = [0u8; 16384]; + loop { + match sock.read(&mut buf) { + Ok(0) => return, + Ok(n) => proxy.push_from_guest(&buf[..n]), + Err(ref e) if e.kind() == ErrorKind::WouldBlock => {} + Err(_) => return, + } + while let Some(n) = proxy.pull_to_guest(&mut buf) { + if sock.write_all(&buf[..n]).is_err() { + return; + } + } + if proxy.failed() { + failed_flag.store(true, Ordering::SeqCst); + return; + } + thread::sleep(Duration::from_millis(1)); + } + }); + } + }); + } + + // Point real uv at the proxy, trusting the vpod CA. + let ca_path = std::env::temp_dir().join("vpod-uv-repro-ca.pem"); + std::fs::write(&ca_path, CA_CERT_PEM).unwrap(); + let out = Command::new("uv") + .args([ + "pip", + "compile", + "-", + "--no-cache", + "--index-url", + &format!("https://localhost:{front_port}/simple/"), + ]) + .env("SSL_CERT_FILE", &ca_path) + .env("UV_HTTP_TIMEOUT", "20") + .stdin(std::process::Stdio::piped()) + .stdout(std::process::Stdio::piped()) + .stderr(std::process::Stdio::piped()) + .spawn() + .and_then(|mut child| { + child.stdin.take().unwrap().write_all(b"six\n").unwrap(); + child.wait_with_output() + }) + .expect("failed to run uv"); + + let stderr = String::from_utf8_lossy(&out.stderr); + eprintln!("uv exit: {:?}\nuv stderr:\n{stderr}", out.status); + + assert!( + !any_proxy_failed.load(Ordering::SeqCst), + "TlsProxy failed against real uv (see [tls-proxy] lines above)" + ); + let lower = stderr.to_lowercase(); + assert!( + !lower.contains("handshake") && !lower.contains("decrypt") && !lower.contains("tls"), + "uv reported a TLS-level error:\n{stderr}" + ); +} + +#[test] +fn server_prefers_chacha20_when_client_offers_it() { + // The guest offers AES-256 first, then ChaCha20 (see captured + // ClientHello). Our server must ignore that order and pick ChaCha20 to + // dodge the guest's broken AES-GCM. Verify the negotiated suite. + let ctx = TlsContext::new().unwrap(); + + let mut guest_roots = RootCertStore::empty(); + let vpod_ca_der = Certificate::from_pem(CA_CERT_PEM) + .unwrap() + .to_der() + .unwrap(); + guest_roots.add(CertificateDer::from(vpod_ca_der)).unwrap(); + let mut prov = rustls::crypto::ring::default_provider(); + prov.kx_groups = vec![rustls::crypto::ring::kx_group::X25519]; + // AES-256 first, exactly like the guest — server must still choose ChaCha. + prov.cipher_suites = vec![ + rustls::crypto::ring::cipher_suite::TLS13_AES_256_GCM_SHA384, + rustls::crypto::ring::cipher_suite::TLS13_AES_128_GCM_SHA256, + rustls::crypto::ring::cipher_suite::TLS13_CHACHA20_POLY1305_SHA256, + ]; + let cfg = Arc::new( + ClientConfig::builder_with_provider(Arc::new(prov)) + .with_safe_default_protocol_versions() + .unwrap() + .with_root_certificates(guest_roots) + .with_no_client_auth(), + ); + let mut guest = ClientConnection::new(cfg, ServerName::try_from("localhost").unwrap()).unwrap(); + let mut server = ServerConnection::new(ctx.server_config.clone()).unwrap(); + + // Drive the handshake to completion in-memory. + for _ in 0..20 { + let mut c2s = Vec::new(); + while guest.wants_write() { + guest.write_tls(&mut c2s).unwrap(); + } + let mut cur = c2s.as_slice(); + while !cur.is_empty() { + server.read_tls(&mut cur).unwrap(); + } + server.process_new_packets().unwrap(); + + let mut s2c = Vec::new(); + while server.wants_write() { + server.write_tls(&mut s2c).unwrap(); + } + let mut cur = s2c.as_slice(); + while !cur.is_empty() { + guest.read_tls(&mut cur).unwrap(); + } + guest.process_new_packets().unwrap(); + + if !guest.is_handshaking() && !server.is_handshaking() { + break; + } + } + + assert_eq!( + server.negotiated_cipher_suite().map(|s| s.suite()), + Some(rustls::CipherSuite::TLS13_CHACHA20_POLY1305_SHA256), + "server should prefer ChaCha20 even though the client offered AES-256 first" + ); +} + +#[test] +fn mirrors_guest_handshake_x25519_aes256_no_hrr() { + // Reproduce the guest's EXACT handshake shape from the captured bytes: + // ring client (like uv on riscv), x25519-only keyshare so there's no + // HelloRetryRequest (guest sent one ClientHello), AES-256 offered first + // so the server honors client order and picks AES_256_GCM_SHA384, plus + // a large app-data request. If this decrypts, our server handles the + // guest's handshake correctly and the fault is in the guest's TLS + // client, not ours. + let (port, up_ca, up) = spawn_test_upstream(UPSTREAM_REPLY); + + let ctx = TlsContext::new().unwrap(); + let mut proxy = TlsProxy::new_test( + ctx.server_config.clone(), + client_config_trusting(&up_ca), + [127, 0, 0, 1], + port, + ); + + let mut guest_roots = RootCertStore::empty(); + let vpod_ca_der = Certificate::from_pem(CA_CERT_PEM) + .unwrap() + .to_der() + .unwrap(); + guest_roots.add(CertificateDer::from(vpod_ca_der)).unwrap(); + + let mut prov = rustls::crypto::ring::default_provider(); + prov.kx_groups = vec![rustls::crypto::ring::kx_group::X25519]; + prov.cipher_suites = vec![ + rustls::crypto::ring::cipher_suite::TLS13_AES_256_GCM_SHA384, + rustls::crypto::ring::cipher_suite::TLS13_AES_128_GCM_SHA256, + rustls::crypto::ring::cipher_suite::TLS13_CHACHA20_POLY1305_SHA256, + ]; + let cfg = Arc::new( + ClientConfig::builder_with_provider(Arc::new(prov)) + .with_safe_default_protocol_versions() + .unwrap() + .with_root_certificates(guest_roots) + .with_no_client_auth(), + ); + let mut guest = ClientConnection::new(cfg, ServerName::try_from("localhost").unwrap()).unwrap(); + + let mut request = b"GET /simple/six/ HTTP/1.1\r\nHost: files.pythonhosted.org\r\n".to_vec(); + for i in 0..20 { + request + .extend_from_slice(format!("X-Padding-Header-{i}: {}\r\n", "a".repeat(30)).as_bytes()); + } + request.extend_from_slice(b"\r\n"); + + let mut request_sent = false; + let mut got = Vec::new(); + for _ in 0..2000 { + let mut out = Vec::new(); + while guest.wants_write() { + guest.write_tls(&mut out).unwrap(); + } + if !out.is_empty() { + proxy.push_from_guest(&out); + } + + let mut buf = [0u8; 16384]; + while let Some(n) = proxy.pull_to_guest(&mut buf) { + let mut slice = &buf[..n]; + while !slice.is_empty() { + guest.read_tls(&mut slice).unwrap(); + guest.process_new_packets().unwrap(); + } + } + + if !request_sent && !guest.is_handshaking() { + guest.writer().write_all(&request).unwrap(); + request_sent = true; + } + + let mut buf = [0u8; 1024]; + if let Ok(n) = guest.reader().read(&mut buf) + && n > 0 + { + got.extend_from_slice(&buf[..n]); + } + + if got.windows(5).any(|w| w == b"hello") { + break; + } + + assert!( + !proxy.failed(), + "mirrored guest handshake: app-data failed to decrypt" + ); + thread::sleep(Duration::from_millis(1)); + } + + let _ = up.join(); + assert!( + got.windows(5).any(|w| w == b"hello"), + "mirrored guest request did not round-trip, got {got:?}" + ); +} + +#[test] +fn second_connection_with_session_resumption_succeeds() { + let ctx = TlsContext::new().unwrap(); + + let mut guest_roots = RootCertStore::empty(); + let vpod_ca_der = Certificate::from_pem(CA_CERT_PEM) + .unwrap() + .to_der() + .unwrap(); + guest_roots.add(CertificateDer::from(vpod_ca_der)).unwrap(); + + let ring_cfg = Arc::new( + ClientConfig::builder_with_provider(Arc::new(rustls::crypto::ring::default_provider())) + .with_safe_default_protocol_versions() + .unwrap() + .with_root_certificates(guest_roots) + .with_no_client_auth(), + ); + + for attempt in 0..2 { + let (port, up_ca, up) = spawn_test_upstream(UPSTREAM_REPLY); + let mut proxy = TlsProxy::new_test( + ctx.server_config.clone(), + client_config_trusting(&up_ca), + [127, 0, 0, 1], + port, + ); + let mut guest = + ClientConnection::new(ring_cfg.clone(), ServerName::try_from("localhost").unwrap()) + .unwrap(); + + let mut request_sent = false; + let mut got = Vec::new(); + for _ in 0..2000 { + let mut out = Vec::new(); + while guest.wants_write() { + guest.write_tls(&mut out).unwrap(); + } + if !out.is_empty() { + proxy.push_from_guest(&out); + } + + let mut buf = [0u8; 16384]; + while let Some(n) = proxy.pull_to_guest(&mut buf) { + let mut slice = &buf[..n]; + while !slice.is_empty() { + guest.read_tls(&mut slice).unwrap(); + guest.process_new_packets().unwrap(); + } + } + + if !request_sent && !guest.is_handshaking() { + guest.writer().write_all(b"GET / HTTP/1.0\r\n\r\n").unwrap(); + request_sent = true; + } + + let mut buf = [0u8; 1024]; + if let Ok(n) = guest.reader().read(&mut buf) + && n > 0 + { + got.extend_from_slice(&buf[..n]); + } + + if got.windows(5).any(|w| w == b"hello") { + break; + } + + if proxy.failed() { + panic!("proxy failed on connection #{attempt} (resumption)"); + } + thread::sleep(Duration::from_millis(1)); + } + + let _ = up.join(); + assert!( + got.windows(5).any(|w| w == b"hello"), + "connection #{attempt} did not get reply, got {got:?}" + ); + } +} + +#[test] +fn large_response_delivered_in_full_without_truncation() { + const BODY: usize = 1_000_000; + let (port, up_ca, up) = spawn_test_upstream_streaming(BODY); + + let ctx = TlsContext::new().unwrap(); + let mut proxy = TlsProxy::new_test( + ctx.server_config.clone(), + client_config_trusting(&up_ca), + [127, 0, 0, 1], + port, + ); + + let mut guest = ClientConnection::new( + client_config_trusting(CA_CERT_PEM), + ServerName::try_from("localhost").unwrap(), + ) + .unwrap(); + + guest.set_buffer_limit(None); + + let mut request_sent = false; + let mut got = Vec::new(); + for _ in 0..40000 { + let mut out = Vec::new(); + while guest.wants_write() { + guest.write_tls(&mut out).unwrap(); + } + if !out.is_empty() { + proxy.push_from_guest(&out); + } + + let mut buf = [0u8; 16384]; + while let Some(n) = proxy.pull_to_guest(&mut buf) { + let mut slice = &buf[..n]; + while !slice.is_empty() { + guest.read_tls(&mut slice).unwrap(); + // Process + drain after every read_tls so neither the record + // deframer nor the plaintext buffer overflows on a big body. + guest.process_new_packets().unwrap(); + let mut pt = [0u8; 16384]; + loop { + match guest.reader().read(&mut pt) { + Ok(m) if m > 0 => got.extend_from_slice(&pt[..m]), + _ => break, + } + } + } + } + + if !request_sent && !guest.is_handshaking() { + guest.writer().write_all(b"GET / HTTP/1.0\r\n\r\n").unwrap(); + request_sent = true; + } + + let body = got + .windows(4) + .position(|w| w == b"\r\n\r\n") + .map(|i| got.len() - (i + 4)); + if body == Some(BODY) { + break; + } + + assert!(!proxy.failed(), "proxy failed on large response"); + thread::sleep(Duration::from_millis(1)); + } + + let _ = up.join(); + let body = got + .windows(4) + .position(|w| w == b"\r\n\r\n") + .map(|i| got.len() - (i + 4)); + assert_eq!(body, Some(BODY), "large body truncated: {body:?} of {BODY}"); +} diff --git a/crates/machine/src/virtio/tls_proxy/testutil.rs b/crates/machine/src/virtio/tls_proxy/testutil.rs new file mode 100644 index 0000000..2e3afaf --- /dev/null +++ b/crates/machine/src/virtio/tls_proxy/testutil.rs @@ -0,0 +1,288 @@ +use super::*; + +pub(crate) fn mint_leaf_with_ca( + ca_key_pem: &str, + ca_cert_pem: &str, + sni: &str, +) -> (Vec, Vec) { + let provider = Arc::new(rustls_rustcrypto::provider()); + let resolver = SniResolver::from_pems(provider, ca_key_pem, ca_cert_pem).unwrap(); + resolver.mint_leaf(sni).unwrap() +} + +pub(crate) fn generate_ca_pems() -> (String, String) { + use der::pem::LineEnding; + use x509_cert::der::EncodePem; + + let key = SigningKey::random(&mut OsRng); + let spki = SubjectPublicKeyInfoOwned::from_key(*key.verifying_key()).unwrap(); + let name = Name::from_str("CN=vpod local CA").unwrap(); + let validity = Validity::from_now(Duration::from_secs(10 * 365 * 24 * 3600)).unwrap(); + let builder = CertificateBuilder::new( + Profile::Root, + SerialNumber::from(1u32), + validity, + name, + spki, + &key, + ) + .unwrap(); + let cert: Certificate = builder.build::().unwrap(); + ( + key.to_pkcs8_pem(Default::default()).unwrap().to_string(), + cert.to_pem(LineEnding::LF).unwrap(), + ) +} + +pub(crate) fn spawn_test_upstream( + reply: &'static [u8], +) -> (u16, String, std::thread::JoinHandle<()>) { + use std::net::TcpListener; + + let (up_ca_key, up_ca_cert) = generate_ca_pems(); + let (leaf_der, leaf_key) = mint_leaf_with_ca(&up_ca_key, &up_ca_cert, "localhost"); + + let up_server_cfg = + ServerConfig::builder_with_provider(Arc::new(rustls_rustcrypto::provider())) + .with_safe_default_protocol_versions() + .unwrap() + .with_no_client_auth() + .with_single_cert( + vec![CertificateDer::from(leaf_der)], + PrivateKeyDer::try_from(leaf_key).unwrap(), + ) + .unwrap(); + let up_server_cfg = Arc::new(up_server_cfg); + + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let port = listener.local_addr().unwrap().port(); + + let handle = std::thread::spawn(move || { + listener.set_nonblocking(false).ok(); + let (mut sock, _) = listener.accept().unwrap(); + + sock.set_read_timeout(Some(Duration::from_secs(10))).ok(); + let mut conn = ServerConnection::new(up_server_cfg).unwrap(); + let mut req = Vec::new(); + + loop { + if conn.wants_write() { + conn.write_tls(&mut sock).unwrap(); + continue; + } + + if conn.is_handshaking() { + conn.read_tls(&mut sock).unwrap(); + conn.process_new_packets().unwrap(); + continue; + } + + let mut buf = [0u8; 1024]; + match conn.reader().read(&mut buf) { + Ok(n) if n > 0 => { + req.extend_from_slice(&buf[..n]); + if req.windows(4).any(|w| w == b"\r\n\r\n") { + break; + } + continue; + } + _ => {} + } + conn.read_tls(&mut sock).unwrap(); + conn.process_new_packets().unwrap(); + } + conn.writer().write_all(reply).unwrap(); + while conn.wants_write() { + conn.write_tls(&mut sock).unwrap(); + } + conn.send_close_notify(); + while conn.wants_write() { + conn.write_tls(&mut sock).unwrap(); + } + }); + + (port, up_ca_cert, handle) +} + +pub(crate) fn spawn_test_upstream_streaming( + body_len: usize, +) -> (u16, String, std::thread::JoinHandle<()>) { + use std::net::TcpListener; + + let (up_ca_key, up_ca_cert) = generate_ca_pems(); + let (leaf_der, leaf_key) = mint_leaf_with_ca(&up_ca_key, &up_ca_cert, "localhost"); + let up_server_cfg = + ServerConfig::builder_with_provider(Arc::new(rustls_rustcrypto::provider())) + .with_safe_default_protocol_versions() + .unwrap() + .with_no_client_auth() + .with_single_cert( + vec![CertificateDer::from(leaf_der)], + PrivateKeyDer::try_from(leaf_key).unwrap(), + ) + .unwrap(); + let up_server_cfg = Arc::new(up_server_cfg); + + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let port = listener.local_addr().unwrap().port(); + + let handle = std::thread::spawn(move || { + listener.set_nonblocking(false).ok(); + + let (mut sock, _) = listener.accept().unwrap(); + sock.set_read_timeout(Some(Duration::from_secs(10))).ok(); + + let mut conn = ServerConnection::new(up_server_cfg).unwrap(); + let mut req = Vec::new(); + + loop { + if conn.wants_write() { + conn.write_tls(&mut sock).unwrap(); + continue; + } + + if conn.is_handshaking() { + conn.read_tls(&mut sock).unwrap(); + conn.process_new_packets().unwrap(); + continue; + } + + let mut buf = [0u8; 1024]; + match conn.reader().read(&mut buf) { + Ok(n) if n > 0 => { + req.extend_from_slice(&buf[..n]); + if req.windows(4).any(|w| w == b"\r\n\r\n") { + break; + } + continue; + } + _ => {} + } + + conn.read_tls(&mut sock).unwrap(); + conn.process_new_packets().unwrap(); + } + + let header = format!("HTTP/1.0 200 OK\r\nContent-Length: {body_len}\r\n\r\n"); + conn.writer().write_all(header.as_bytes()).unwrap(); + let mut sent = 0usize; + let chunk = vec![b'x'; 16384]; + while sent < body_len { + let n = chunk.len().min(body_len - sent); + conn.writer().write_all(&chunk[..n]).unwrap(); + sent += n; + + while conn.wants_write() { + conn.write_tls(&mut sock).unwrap(); + } + } + + conn.send_close_notify(); + while conn.wants_write() { + conn.write_tls(&mut sock).unwrap(); + } + }); + + (port, up_ca_cert, handle) +} + +pub(crate) fn spawn_test_upstream_rst( + reply: &'static [u8], +) -> (u16, String, std::thread::JoinHandle<()>) { + use std::net::TcpListener; + + let (up_ca_key, up_ca_cert) = generate_ca_pems(); + let (leaf_der, leaf_key) = mint_leaf_with_ca(&up_ca_key, &up_ca_cert, "localhost"); + let up_server_cfg = + ServerConfig::builder_with_provider(Arc::new(rustls_rustcrypto::provider())) + .with_safe_default_protocol_versions() + .unwrap() + .with_no_client_auth() + .with_single_cert( + vec![CertificateDer::from(leaf_der)], + PrivateKeyDer::try_from(leaf_key).unwrap(), + ) + .unwrap(); + let up_server_cfg = Arc::new(up_server_cfg); + + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let port = listener.local_addr().unwrap().port(); + + let handle = std::thread::spawn(move || { + listener.set_nonblocking(false).ok(); + + let (mut sock, _) = listener.accept().unwrap(); + sock.set_read_timeout(Some(Duration::from_secs(10))).ok(); + + let mut conn = ServerConnection::new(up_server_cfg).unwrap(); + let mut req = Vec::new(); + loop { + if conn.wants_write() { + conn.write_tls(&mut sock).unwrap(); + continue; + } + + if conn.is_handshaking() { + conn.read_tls(&mut sock).unwrap(); + conn.process_new_packets().unwrap(); + continue; + } + + let mut buf = [0u8; 1024]; + match conn.reader().read(&mut buf) { + Ok(n) if n > 0 => { + req.extend_from_slice(&buf[..n]); + if req.windows(4).any(|w| w == b"\r\n\r\n") { + break; + } + continue; + } + _ => {} + } + + conn.read_tls(&mut sock).unwrap(); + conn.process_new_packets().unwrap(); + } + + conn.writer().write_all(reply).unwrap(); + while conn.wants_write() { + conn.write_tls(&mut sock).unwrap(); + } + + std::thread::sleep(Duration::from_millis(500)); + + { + use std::os::fd::AsRawFd; + let linger = libc::linger { + l_onoff: 1, + l_linger: 0, + }; + unsafe { + libc::setsockopt( + sock.as_raw_fd(), + libc::SOL_SOCKET, + libc::SO_LINGER, + &linger as *const libc::linger as *const libc::c_void, + std::mem::size_of::() as libc::socklen_t, + ); + } + } + drop(sock); + }); + + (port, up_ca_cert, handle) +} + +pub(crate) fn client_config_trusting(ca_pem: &str) -> Arc { + let mut roots = RootCertStore::empty(); + let ca_der = Certificate::from_pem(ca_pem).unwrap().to_der().unwrap(); + + roots.add(CertificateDer::from(ca_der)).unwrap(); + Arc::new( + ClientConfig::builder_with_provider(Arc::new(rustls_rustcrypto::provider())) + .with_safe_default_protocol_versions() + .unwrap() + .with_root_certificates(roots) + .with_no_client_auth(), + ) +} diff --git a/crates/native-cli/Cargo.toml b/crates/native-cli/Cargo.toml index 9a34bf4..538bbae 100644 --- a/crates/native-cli/Cargo.toml +++ b/crates/native-cli/Cargo.toml @@ -3,6 +3,11 @@ name = "native-cli" version = "0.0.0" edition.workspace = true +[features] +perf = ["riscv-core/perf-counters"] +aot = ["riscv-core/aot"] +aot-trace = ["riscv-core/aot-trace"] + [dependencies] machine = { path = "../machine" } riscv-core = { path = "../riscv-core" } diff --git a/crates/native-cli/src/main.rs b/crates/native-cli/src/main.rs index 056b3b3..ec420dd 100644 --- a/crates/native-cli/src/main.rs +++ b/crates/native-cli/src/main.rs @@ -17,7 +17,7 @@ fn usage() -> ! { eprintln!( "usage: vpod [--bios ] [--initrd ] [--disk ] [--net] [--agent] \ [--setup ] [--ram ] [--bootargs ] \ - [--snapshot-save ] [--snapshot-load ]" + [--snapshot-save ] [--snapshot-load ] [--no-aot]" ); std::process::exit(1); } @@ -65,6 +65,7 @@ fn main() { let mut ram_mb: u64 = 256; let mut bootargs = "root=/dev/ram0 rw console=ttyS0 earlycon".to_string(); let mut trace_insns: u64 = 0; + let mut no_aot = false; let first = args.next().unwrap_or_else(|| usage()); let mut kernel_path: Option = None; @@ -104,6 +105,7 @@ fn main() { .and_then(|s| s.parse().ok()) .unwrap_or_else(|| usage()); } + "--no-aot" => no_aot = true, "--trace" => { trace_insns = args.next().and_then(|s| s.parse().ok()).unwrap_or(64); } @@ -210,6 +212,20 @@ fn main() { 0 }); + #[cfg(feature = "aot")] + if no_aot { + eprintln!("[vpod] aot dispatch disabled (--no-aot)"); + } else { + hart.blocks.aot_init(riscv_core::aot::AOT_PAGE_HASHES); + eprintln!( + "[vpod] aot dispatch enabled: {} translated pages (content-hash keyed)", + riscv_core::aot::AOT_PAGE_HASHES.len() + ); + } + + #[cfg(not(feature = "aot"))] + let _ = no_aot; + if !setup_cmds.is_empty() { run_setup::run( &mut bus, @@ -228,4 +244,38 @@ fn main() { restored_flags, ); } + + #[cfg(feature = "aot")] + { + use std::sync::atomic::Ordering; + let calls = riscv_core::aot::DISPATCH_CALLS.load(Ordering::Relaxed); + let retired = riscv_core::aot::DISPATCH_RETIRED.load(Ordering::Relaxed); + eprintln!( + "\r[vpod] aot: {retired} insns retired in {calls} dispatches ({:.1} insns/dispatch) | total instret {}", + retired as f64 / calls.max(1) as f64, + hart.csr.instret + ); + if let Some(report) = riscv_core::perf::report() { + eprintln!("{report}"); + } + } + + #[cfg(feature = "aot-trace")] + if let Ok(path) = std::env::var("VPOD_AOT_TRACE") { + let mut counts: Vec<(u64, u64)> = hart.blocks.trace_counts().collect(); + counts.sort_by_key(|b| std::cmp::Reverse(b.1)); + + let mut out = String::new(); + for (pa, n) in counts { + out.push_str(&format!("{pa:x} {n}\n")); + } + + match std::fs::write(&path, &out) { + Ok(()) => eprintln!( + "\r[vpod] aot trace: {} block pcs -> {path}", + out.lines().count() + ), + Err(e) => eprintln!("\r[vpod] aot trace write failed: {e}"), + } + } } diff --git a/crates/native-cli/src/run_interactive.rs b/crates/native-cli/src/run_interactive.rs index 6505321..3ab4c8a 100644 --- a/crates/native-cli/src/run_interactive.rs +++ b/crates/native-cli/src/run_interactive.rs @@ -22,6 +22,10 @@ pub fn run( eprintln!("[vpod] Press Ctrl-C to exit the emulator."); } + if cfg!(feature = "perf") { + eprintln!("[vpod] perf counters on: Ctrl-P prints and resets them."); + } + if trace_insns > 0 { run_trace(bus, hart, trace_insns); return; @@ -53,7 +57,24 @@ pub fn run( terminal::poll_stdin(bus, snap_save, hart, snap_flags); bus.flush_console_to_stdout(); - match hart.run(bus, interval) { + + if hart.is_waiting { + hart.is_waiting = false; + + if !bus.net_rx_pending() { + const MAX_IDLE_SLEEP_NS: u64 = 1_000_000; + let sleep_ns = bus + .clint + .nanos_until_timer() + .unwrap_or(MAX_IDLE_SLEEP_NS) + .min(MAX_IDLE_SLEEP_NS); + + std::thread::sleep(std::time::Duration::from_nanos(sleep_ns)); + bus.clint.advance_by_nanos(sleep_ns); + } + } + + match hart.run_until_wait(bus, interval) { StepResult::Ok => {} StepResult::Trap(cause) => { eprintln!( diff --git a/crates/native-cli/src/run_setup.rs b/crates/native-cli/src/run_setup.rs index fd16bc1..3a4ed07 100644 --- a/crates/native-cli/src/run_setup.rs +++ b/crates/native-cli/src/run_setup.rs @@ -18,6 +18,7 @@ pub fn run( bus.uart.capture_tx.set(true); bus.uart_data.capture_tx.set(true); + push_line(bus, b""); wait_for_prompt(bus, hart, true); for (i, cmd) in cmds.iter().enumerate() { @@ -95,11 +96,18 @@ fn wait_for_prompt(bus: &mut MachineBus, hart: &mut Hart, verbose: bool) -> Vec< for _ in 0..8_000_000u32 { if hart.is_waiting { hart.is_waiting = false; + + if !bus.net_rx_pending() && !bus.net_has_active_connections() { + bus.clint.fast_forward_to_timer(); + } else if !bus.net_rx_pending() { + std::thread::sleep(std::time::Duration::from_micros(200)); + bus.clint.advance_by_nanos(200_000); + } } let interval = if bus.net_rx_pending() { 4096 } else { STEP }; bus.clint.advance_by_instructions(interval); bus.poll(hart); - match hart.run(bus, interval) { + match hart.run_until_wait(bus, interval) { StepResult::Ok => {} StepResult::Trap(cause) => { eprintln!("[vpod-setup] trap {:?} at pc={:#x}", cause, hart.regs.pc); @@ -146,7 +154,10 @@ fn python_init(bus: &mut MachineBus, hart: &mut Hart) -> bool { wait_for_prompt(bus, hart, false); drain_all(bus); - push_line(bus, b"python3 /usr/lib/vpod/pyrunner.py &"); + push_line( + bus, + b"PYR=/usr/bin/python3.real; [ -x $PYR ] || PYR=python3; $PYR /usr/lib/vpod/pyrunner.py &", + ); wait_for_prompt(bus, hart, false); drain_all(bus); @@ -158,12 +169,19 @@ fn python_init(bus: &mut MachineBus, hart: &mut Hart) -> bool { for _ in 0..2_000_000u32 { if hart.is_waiting { hart.is_waiting = false; + + if !bus.net_rx_pending() && !bus.net_has_active_connections() { + bus.clint.fast_forward_to_timer(); + } else if !bus.net_rx_pending() { + std::thread::sleep(std::time::Duration::from_micros(200)); + bus.clint.advance_by_nanos(200_000); + } } bus.clint.advance_by_instructions(8192); bus.poll(hart); - match hart.run(bus, 8192) { + match hart.run_until_wait(bus, 8192) { riscv_core::StepResult::Ok => {} _ => break, } diff --git a/crates/native-cli/src/terminal.rs b/crates/native-cli/src/terminal.rs index f010e63..6b7f080 100644 --- a/crates/native-cli/src/terminal.rs +++ b/crates/native-cli/src/terminal.rs @@ -58,6 +58,11 @@ pub fn poll_stdin(bus: &mut MachineBus, snap_path: Option<&PathBuf>, hart: &Hart match b { 0x1d | 0x03 => { eprintln!("\r\n[vpod] exiting."); + + if let Some(report) = riscv_core::perf::report() { + eprint!("{}", report.replace('\n', "\r\n")); + } + std::process::exit(0); } 0x13 => { @@ -65,6 +70,10 @@ pub fn poll_stdin(bus: &mut MachineBus, snap_path: Option<&PathBuf>, hart: &Hart super::save_snapshot(bus, hart, path, snap_flags); } } + 0x10 => match riscv_core::perf::report() { + Some(report) => eprint!("\r\n{}", report.replace('\n', "\r\n")), + None => eprintln!("\r\n[vpod] perf counters empty or disabled"), + }, _ => bus.uart.push_rx(b), } } diff --git a/crates/riscv-core/Cargo.toml b/crates/riscv-core/Cargo.toml index ee27d23..86aedf7 100644 --- a/crates/riscv-core/Cargo.toml +++ b/crates/riscv-core/Cargo.toml @@ -3,8 +3,14 @@ name = "riscv-core" version = "0.0.0" edition = "2024" +[features] +perf-counters = [] +aot = [] +aot-trace = [] + [dependencies] log = "0.4" +rustc-hash = "2.1" [dev-dependencies] env_logger = "0.11" diff --git a/crates/riscv-core/src/aot.rs b/crates/riscv-core/src/aot.rs new file mode 100644 index 0000000..bd1d7a5 --- /dev/null +++ b/crates/riscv-core/src/aot.rs @@ -0,0 +1,20 @@ +// AOT-translated guest blocks + +#[allow(unused_variables, unreachable_code, clippy::all)] +pub mod generated { + include!("aot/generated.rs"); +} + +pub use generated::AOT_PAGE_HASHES; +pub use generated::dispatch; + +use std::sync::atomic::{AtomicU64, Ordering}; + +pub static DISPATCH_CALLS: AtomicU64 = AtomicU64::new(0); +pub static DISPATCH_RETIRED: AtomicU64 = AtomicU64::new(0); + +#[inline(always)] +pub fn note_dispatch(retired: u64) { + DISPATCH_CALLS.fetch_add(1, Ordering::Relaxed); + DISPATCH_RETIRED.fetch_add(retired, Ordering::Relaxed); +} diff --git a/crates/riscv-core/src/block.rs b/crates/riscv-core/src/block.rs index 68dd58f..2c989fe 100644 --- a/crates/riscv-core/src/block.rs +++ b/crates/riscv-core/src/block.rs @@ -5,7 +5,9 @@ use crate::decode::{CompressedInstruction, Instruction, sign_extend}; use crate::execute::{ExecContext, take_exception}; use crate::extensions as ext; use crate::mmu::{Mmu, MmuFault}; +use crate::perf; use crate::system_bus::SystemBus; +use crate::trap::StepResult; pub const MAX_BLOCK_OPS: usize = 64; @@ -63,6 +65,14 @@ pub enum AluKind { ZextH, Rev8, OrcB, + Sh1add, + Sh2add, + Sh3add, + AddUw, + Sh1addUw, + Sh2addUw, + Sh3addUw, + SlliUw, } #[derive(Clone, Copy, Debug)] @@ -143,6 +153,10 @@ pub enum Op { rs1: u8, imm: i64, }, + + Fallback { + raw: u32, + }, } #[derive(Clone, Copy, Debug)] @@ -172,8 +186,28 @@ impl Slot { pub struct BlockCache { slots: Vec, page_bitmap: Vec, + code_generation: u64, + + aot_hashes: rustc_hash::FxHashMap, + aot_page_table: Vec, + aot_evict_generation: u64, + aot_hash_probes: u64, + aot_hash_matches: u64, + + #[cfg(feature = "aot-trace")] + trace: std::collections::HashMap, +} + +pub enum AotResolve { + Hit(u64), + Miss, + Unknown, } +const AOT_UNKNOWN: u64 = u64::MAX; +const AOT_MISS: u64 = u64::MAX - 1; +const AOT_TABLE_MAX_PAGES: u64 = 1 << 20; + impl Default for BlockCache { fn default() -> Self { Self::new() @@ -191,9 +225,79 @@ impl BlockCache { Self { slots, page_bitmap: Vec::new(), + code_generation: 1, + aot_hashes: rustc_hash::FxHashMap::default(), + aot_page_table: Vec::new(), + aot_evict_generation: 0, + aot_hash_probes: 0, + aot_hash_matches: 0, + #[cfg(feature = "aot-trace")] + trace: std::collections::HashMap::new(), } } + pub fn aot_init(&mut self, hashes: &[(u64, u64)]) { + self.aot_hashes = hashes.iter().copied().collect(); + self.aot_page_table.clear(); + self.aot_evict_generation = self.aot_evict_generation.wrapping_add(1); + self.aot_hash_probes = 0; + self.aot_hash_matches = 0; + } + + pub fn aot_match_stats(&self) -> (u64, u64) { + (self.aot_hash_probes, self.aot_hash_matches) + } + + pub fn aot_enabled(&self) -> bool { + !self.aot_hashes.is_empty() + } + + #[inline(always)] + pub fn aot_resolve(&mut self, page: u64) -> AotResolve { + match self.aot_page_table.get(page as usize) { + Some(&AOT_UNKNOWN) => AotResolve::Unknown, + Some(&AOT_MISS) => AotResolve::Miss, + Some(&orig) => AotResolve::Hit(orig), + None => { + if self.aot_hashes.is_empty() || page >= AOT_TABLE_MAX_PAGES { + AotResolve::Miss + } else { + AotResolve::Unknown + } + } + } + } + + pub fn aot_record_hash(&mut self, page: u64, hash: u64) -> Option { + let orig = self.aot_hashes.get(&hash).copied(); + + self.aot_hash_probes += 1; + if orig.is_some() { + self.aot_hash_matches += 1; + } + + if page < AOT_TABLE_MAX_PAGES { + if self.aot_page_table.len() <= page as usize { + self.aot_page_table.resize(page as usize + 1, AOT_UNKNOWN); + } + self.aot_page_table[page as usize] = orig.unwrap_or(AOT_MISS); + } + self.mark_page(page); + + orig + } + + #[cfg(feature = "aot-trace")] + #[inline(always)] + pub fn trace_exec(&mut self, physical_address: u64) { + *self.trace.entry(physical_address).or_insert(0) += 1; + } + + #[cfg(feature = "aot-trace")] + pub fn trace_counts(&self) -> impl Iterator + '_ { + self.trace.iter().map(|(&pa, &n)| (pa, n)) + } + #[inline(always)] fn slot_index(physical_address: u64) -> usize { ((physical_address >> 1) as usize) & (CACHE_SLOTS - 1) @@ -230,6 +334,19 @@ impl BlockCache { } self.page_bitmap[word] |= 1 << (page % 64); + + self.code_generation = self.code_generation.wrapping_add(1); + } + + #[inline(always)] + pub fn code_generation(&self) -> u64 { + self.code_generation + } + + #[cfg(debug_assertions)] + pub fn page_has_code(&self, page: u64) -> bool { + let word = (page / 64) as usize; + word < self.page_bitmap.len() && self.page_bitmap[word] & (1 << (page % 64)) != 0 } #[inline(always)] @@ -242,7 +359,19 @@ impl BlockCache { } } + #[inline(always)] + pub fn aot_evict_generation(&self) -> u64 { + self.aot_evict_generation + } + fn evict_page(&mut self, page: u64) { + perf::note_store_page_eviction(); + + self.aot_evict_generation = self.aot_evict_generation.wrapping_add(1); + if let Some(entry) = self.aot_page_table.get_mut(page as usize) { + *entry = AOT_UNKNOWN; + } + for slot in &mut self.slots { if slot.tag != u64::MAX && slot.tag >> 12 == page { *slot = Slot::EMPTY; @@ -258,6 +387,9 @@ impl BlockCache { } self.page_bitmap.clear(); + + self.aot_page_table.clear(); + self.aot_evict_generation = self.aot_evict_generation.wrapping_add(1); } } @@ -280,14 +412,20 @@ pub fn decode_block(bus: &mut B, physical_address: u64) -> Option< } let high_halfword = bus.read_halfword(physical_address + offset_in_block + 2) as u32; + let raw = low_halfword | (high_halfword << 16); - match decode_full(low_halfword | (high_halfword << 16)) { + match decode_full(raw) { Some(op) => (op, 4u8), - None => break, + None => { + if fallback_stops_block(raw) { + break; + } + (Op::Fallback { raw }, 4u8) + } } }; - let terminator = matches!(op, Op::Branch { .. } | Op::Jal { .. } | Op::Jalr { .. }); + let terminator = matches!(op, Op::Jal { .. } | Op::Jalr { .. }); ops.push(DecodedInsn { op, pc_off: offset_in_block as u16, @@ -310,6 +448,11 @@ pub fn decode_block(bus: &mut B, physical_address: u64) -> Option< }) } +#[inline(always)] +fn fallback_stops_block(raw: u32) -> bool { + (raw & 0x7f == 0x73 && (raw >> 12) & 0x7 == 0) || raw & 0x707f == 0x100f +} + fn decode_full(raw: u32) -> Option { let inst = Instruction(raw); let rd = inst.rd() as u8; @@ -442,6 +585,7 @@ fn decode_full(raw: u32) -> Option { let kind_imm = match funct3 { 0x0 => Some((AluKind::Addw, imm)), + 0x1 if funct7 == 0x04 || funct7 == 0x05 => Some((AluKind::SlliUw, imm & 0x3f)), 0x1 => match funct7 { 0x00 => Some((AluKind::Sllw, shamt)), 0x30 => match shamt { @@ -491,7 +635,10 @@ fn decode_full(raw: u32) -> Option { (0x5, 0x05) => AluKind::Max, (0x6, 0x05) => AluKind::Minu, (0x7, 0x05) => AluKind::Maxu, - (0x4, 0x04) => AluKind::ZextH, + // Zba: shifted add + (0x2, 0x10) => AluKind::Sh1add, + (0x4, 0x10) => AluKind::Sh2add, + (0x6, 0x10) => AluKind::Sh3add, _ => return None, }; Some(Op::AluReg { kind, rd, rs1, rs2 }) @@ -510,6 +657,13 @@ fn decode_full(raw: u32) -> Option { (0x7, 0x01) => AluKind::Remuw, (0x1, 0x30) => AluKind::Rolw, (0x5, 0x30) => AluKind::Rorw, + // Zbb: zext.h is `packw rd, rs1, x0`, so it only decodes with rs2 == 0. + (0x4, 0x04) if rs2 == 0 => AluKind::ZextH, + // Zba: zero-extended shifted add + (0x0, 0x04) => AluKind::AddUw, + (0x2, 0x10) => AluKind::Sh1addUw, + (0x4, 0x10) => AluKind::Sh2addUw, + (0x6, 0x10) => AluKind::Sh3addUw, _ => return None, }; @@ -848,7 +1002,7 @@ fn decode_compressed(raw: u16) -> Option { } #[inline(always)] -fn alu(kind: AluKind, lhs: u64, rhs: u64) -> u64 { +pub(crate) fn alu(kind: AluKind, lhs: u64, rhs: u64) -> u64 { match kind { AluKind::Add => lhs.wrapping_add(rhs), AluKind::Sub => lhs.wrapping_sub(rhs), @@ -911,6 +1065,15 @@ fn alu(kind: AluKind, lhs: u64, rhs: u64) -> u64 { AluKind::SextH => lhs as i16 as i64 as u64, AluKind::ZextH => lhs as u16 as u64, AluKind::Rev8 => lhs.swap_bytes(), + // Zba + AluKind::Sh1add => (lhs << 1).wrapping_add(rhs), + AluKind::Sh2add => (lhs << 2).wrapping_add(rhs), + AluKind::Sh3add => (lhs << 3).wrapping_add(rhs), + AluKind::AddUw => (lhs as u32 as u64).wrapping_add(rhs), + AluKind::Sh1addUw => ((lhs as u32 as u64) << 1).wrapping_add(rhs), + AluKind::Sh2addUw => ((lhs as u32 as u64) << 2).wrapping_add(rhs), + AluKind::Sh3addUw => ((lhs as u32 as u64) << 3).wrapping_add(rhs), + AluKind::SlliUw => (lhs as u32 as u64) << (rhs & 0x3f), AluKind::OrcB => { let mut result = 0u64; for i in 0..8 { @@ -927,9 +1090,10 @@ pub fn exec_block( ctx: &mut ExecContext, block: &Block, entry_pc: u64, - satp: u64, -) -> u64 { + mut satp: u64, +) -> (u64, StepResult) { let mut retired_instructions: u64 = 0; + let mut pending: u64 = 0; for decoded_instruction in block.ops.iter() { let pc = entry_pc.wrapping_add(decoded_instruction.pc_off as u64); @@ -954,12 +1118,11 @@ pub fn exec_block( Op::Load { kind, rd, rs1, imm } => { let virtual_address = ctx.regs.read(rs1 as usize).wrapping_add(imm as u64); - ctx.regs.pc = pc; - match do_load(ctx, satp, kind, virtual_address) { + match do_load(ctx, satp, kind, virtual_address, pc) { Ok(v) => ctx.regs.write(rd as usize, v), Err(()) => { - ctx.csr.instret = ctx.csr.instret.wrapping_add(retired_instructions); - return retired_instructions + 1; + ctx.csr.instret = ctx.csr.instret.wrapping_add(pending); + return (retired_instructions + 1, StepResult::Ok); } } } @@ -971,10 +1134,9 @@ pub fn exec_block( } => { let virtual_address = ctx.regs.read(rs1 as usize).wrapping_add(imm as u64); let val = ctx.regs.read(rs2 as usize); - ctx.regs.pc = pc; - if do_store(ctx, satp, kind, virtual_address, val).is_err() { - ctx.csr.instret = ctx.csr.instret.wrapping_add(retired_instructions); - return retired_instructions + 1; + if do_store(ctx, satp, kind, virtual_address, val, pc).is_err() { + ctx.csr.instret = ctx.csr.instret.wrapping_add(pending); + return (retired_instructions + 1, StepResult::Ok); } } Op::Branch { @@ -994,16 +1156,11 @@ pub fn exec_block( BranchKind::Bgeu => a >= b, }; - ctx.regs.pc = if taken { - pc.wrapping_add(offset as u64) - } else { - pc.wrapping_add(decoded_instruction.ilen as u64) - }; - - retired_instructions += 1; - ctx.csr.instret = ctx.csr.instret.wrapping_add(retired_instructions); - - return retired_instructions; + if taken { + ctx.regs.pc = pc.wrapping_add(offset as u64); + ctx.csr.instret = ctx.csr.instret.wrapping_add(pending + 1); + return (retired_instructions + 1, StepResult::Ok); + } } Op::Jal { rd, offset } => { ctx.regs.write( @@ -1012,10 +1169,9 @@ pub fn exec_block( ); ctx.regs.pc = pc.wrapping_add(offset as u64); - retired_instructions += 1; - ctx.csr.instret = ctx.csr.instret.wrapping_add(retired_instructions); + ctx.csr.instret = ctx.csr.instret.wrapping_add(pending + 1); - return retired_instructions; + return (retired_instructions + 1, StepResult::Ok); } Op::Jalr { rd, rs1, imm } => { let target = ctx.regs.read(rs1 as usize).wrapping_add(imm as u64) & !1; @@ -1025,26 +1181,49 @@ pub fn exec_block( ); ctx.regs.pc = target; + ctx.csr.instret = ctx.csr.instret.wrapping_add(pending + 1); + + return (retired_instructions + 1, StepResult::Ok); + } + Op::Fallback { raw } => { + perf::note_fallback_op(); + ctx.regs.pc = pc; + ctx.csr.instret = ctx.csr.instret.wrapping_add(pending); + pending = 0; + + let result = crate::execute::exec_raw(ctx, raw, pc); retired_instructions += 1; - ctx.csr.instret = ctx.csr.instret.wrapping_add(retired_instructions); - return retired_instructions; + let next_pc = pc.wrapping_add(decoded_instruction.ilen as u64); + if !matches!(result, StepResult::Ok) || ctx.regs.pc != next_pc || *ctx.is_waiting { + return (retired_instructions, result); + } + + let new_satp = effective_satp(*ctx.priv_mode, ctx.csr.satp); + if new_satp != satp { + return (retired_instructions, result); + } + satp = new_satp; + continue; } } retired_instructions += 1; + pending += 1; } ctx.regs.pc = entry_pc.wrapping_add(block.byte_len as u64); - ctx.csr.instret = ctx.csr.instret.wrapping_add(retired_instructions); - retired_instructions + ctx.csr.instret = ctx.csr.instret.wrapping_add(pending); + (retired_instructions, StepResult::Ok) } -fn do_load( +#[inline(always)] +pub(crate) fn do_load( ctx: &mut ExecContext, satp: u64, kind: LoadKind, virtual_address: u64, + pc: u64, ) -> Result { let size: u64 = match kind { LoadKind::Lb | LoadKind::Lbu => 1, @@ -1053,15 +1232,65 @@ fn do_load( LoadKind::Ld => 8, }; + let offset_in_page = virtual_address & 0xFFF; + if offset_in_page + size <= 0x1000 + && let Some(host_page) = + ctx.mmu + .load_fast_lookup(virtual_address, satp, ctx.bus.ram_epoch()) + { + perf::note_load_fast_hit(); + + let raw = unsafe { + let p = host_page.add(offset_in_page as usize); + match size { + 1 => *p as u64, + 2 => u16::from_le((p as *const u16).read_unaligned()) as u64, + 4 => u32::from_le((p as *const u32).read_unaligned()) as u64, + _ => u64::from_le((p as *const u64).read_unaligned()), + } + }; + + #[cfg(debug_assertions)] + { + let shadow = raw_load(ctx.mmu, ctx.bus, satp, virtual_address, size) + .expect("load fast-path hit but slow path faulted"); + assert_eq!( + raw, shadow, + "load fast path diverged at va={virtual_address:#x} size={size}" + ); + } + + return Ok(extend_load(kind, raw)); + } + + do_load_slow(ctx, satp, kind, virtual_address, size, pc) +} + +#[cold] +#[inline(never)] +fn do_load_slow( + ctx: &mut ExecContext, + satp: u64, + kind: LoadKind, + virtual_address: u64, + size: u64, + pc: u64, +) -> Result { let raw = match raw_load(ctx.mmu, ctx.bus, satp, virtual_address, size) { Ok(v) => v, Err(f) => { + ctx.regs.pc = pc; take_exception(ctx, f.mcause(), f.tval()); return Err(()); } }; - Ok(match kind { + Ok(extend_load(kind, raw)) +} + +#[inline(always)] +fn extend_load(kind: LoadKind, raw: u64) -> u64 { + match kind { LoadKind::Lb => raw as u8 as i8 as i64 as u64, LoadKind::Lbu => raw as u8 as u64, LoadKind::Lh => raw as u16 as i16 as i64 as u64, @@ -1069,15 +1298,17 @@ fn do_load( LoadKind::Lw => raw as u32 as i32 as i64 as u64, LoadKind::Lwu => raw as u32 as u64, LoadKind::Ld => raw, - }) + } } -fn do_store( +#[inline(always)] +pub(crate) fn do_store( ctx: &mut ExecContext, satp: u64, kind: StoreKind, virtual_address: u64, val: u64, + pc: u64, ) -> Result<(), ()> { let size: u64 = match kind { StoreKind::Sb => 1, @@ -1086,6 +1317,63 @@ fn do_store( StoreKind::Sd => 8, }; + let offset_in_page = virtual_address & 0xFFF; + if offset_in_page + size <= 0x1000 + && let Some(host_page) = ctx.mmu.store_fast_lookup( + virtual_address, + satp, + ctx.bus.ram_epoch(), + ctx.blocks.code_generation(), + ) + { + perf::note_store_fast_hit(); + + #[cfg(debug_assertions)] + { + let physical_address = ctx + .mmu + .translate_store(virtual_address, satp, ctx.bus) + .expect("store fast-path hit but slow path faulted"); + assert!( + !ctx.blocks.page_has_code(physical_address >> 12), + "store fast path would skip a required SMC eviction at va={virtual_address:#x}" + ); + let shadow = ctx + .bus + .ram_store_page(physical_address) + .expect("store fast-path hit on a non-RAM page"); + assert_eq!( + shadow as usize, host_page as usize, + "store fast path diverged at va={virtual_address:#x} size={size}" + ); + } + + unsafe { + let p = host_page.add(offset_in_page as usize); + match size { + 1 => *p = val as u8, + 2 => (p as *mut u16).write_unaligned((val as u16).to_le()), + 4 => (p as *mut u32).write_unaligned((val as u32).to_le()), + _ => (p as *mut u64).write_unaligned(val.to_le()), + } + } + + return Ok(()); + } + + do_store_slow(ctx, satp, virtual_address, val, size, pc) +} + +#[cold] +#[inline(never)] +fn do_store_slow( + ctx: &mut ExecContext, + satp: u64, + virtual_address: u64, + val: u64, + size: u64, + pc: u64, +) -> Result<(), ()> { match raw_store( ctx.mmu, ctx.bus, @@ -1097,12 +1385,142 @@ fn do_store( ) { Ok(()) => Ok(()), Err(f) => { + ctx.regs.pc = pc; take_exception(ctx, f.mcause(), f.tval()); Err(()) } } } +#[inline(always)] +pub(crate) fn _load_span_page( + ctx: &mut ExecContext, + satp: u64, + span_lo: u64, + span_hi: u64, +) -> Option<*const u8> { + if span_lo >> 12 != span_hi >> 12 { + return None; + } + + ctx.mmu.load_fast_lookup(span_lo, satp, ctx.bus.ram_epoch()) +} + +#[inline(always)] +pub(crate) fn _store_span_page( + ctx: &mut ExecContext, + satp: u64, + span_lo: u64, + span_hi: u64, +) -> Option<*mut u8> { + if span_lo >> 12 != span_hi >> 12 { + return None; + } + + ctx.mmu.store_fast_lookup( + span_lo, + satp, + ctx.bus.ram_epoch(), + ctx.blocks.code_generation(), + ) +} + +#[inline(always)] +pub(crate) fn _span_load( + ctx: &mut ExecContext, + host_page: *const u8, + satp: u64, + kind: LoadKind, + virtual_address: u64, +) -> u64 { + let size: u64 = match kind { + LoadKind::Lb | LoadKind::Lbu => 1, + LoadKind::Lh | LoadKind::Lhu => 2, + LoadKind::Lw | LoadKind::Lwu => 4, + LoadKind::Ld => 8, + }; + let offset_in_page = virtual_address & 0xFFF; + debug_assert!(offset_in_page + size <= 0x1000); + perf::note_load_fast_hit(); + + let raw = unsafe { + let p = host_page.add(offset_in_page as usize); + match size { + 1 => *p as u64, + 2 => u16::from_le((p as *const u16).read_unaligned()) as u64, + 4 => u32::from_le((p as *const u32).read_unaligned()) as u64, + _ => u64::from_le((p as *const u64).read_unaligned()), + } + }; + + #[cfg(debug_assertions)] + { + let shadow = raw_load(ctx.mmu, ctx.bus, satp, virtual_address, size) + .expect("span load covered by page check but slow path faulted"); + assert_eq!( + raw, shadow, + "span load diverged at va={virtual_address:#x} size={size}" + ); + } + #[cfg(not(debug_assertions))] + let _ = (ctx, satp); + + extend_load(kind, raw) +} + +#[inline(always)] +pub(crate) fn _span_store( + ctx: &mut ExecContext, + host_page: *mut u8, + satp: u64, + kind: StoreKind, + virtual_address: u64, + val: u64, +) { + let size: u64 = match kind { + StoreKind::Sb => 1, + StoreKind::Sh => 2, + StoreKind::Sw => 4, + StoreKind::Sd => 8, + }; + let offset_in_page = virtual_address & 0xFFF; + debug_assert!(offset_in_page + size <= 0x1000); + perf::note_store_fast_hit(); + + #[cfg(debug_assertions)] + { + let physical_address = ctx + .mmu + .translate_store(virtual_address, satp, ctx.bus) + .expect("span store covered by page check but slow path faulted"); + assert!( + !ctx.blocks.page_has_code(physical_address >> 12), + "span store would skip a required SMC eviction at va={virtual_address:#x}" + ); + let shadow = ctx + .bus + .ram_store_page(physical_address) + .expect("span store hit on a non-RAM page"); + assert_eq!( + shadow as usize, host_page as usize, + "span store diverged at va={virtual_address:#x} size={size}" + ); + } + #[cfg(not(debug_assertions))] + let _ = (ctx, satp); + + unsafe { + let p = host_page.add(offset_in_page as usize); + match size { + 1 => *p = val as u8, + 2 => (p as *mut u16).write_unaligned((val as u16).to_le()), + 4 => (p as *mut u32).write_unaligned((val as u32).to_le()), + _ => (p as *mut u64).write_unaligned(val.to_le()), + } + } +} + +#[inline] pub fn raw_load( mmu: &mut Mmu, bus: &mut B, @@ -1110,7 +1528,11 @@ pub fn raw_load( virtual_address: u64, size: u64, ) -> Result { - if (virtual_address & 0xFFF) + size > 0x1000 { + perf::note_load(); + let offset_in_page = virtual_address & 0xFFF; + + if offset_in_page + size > 0x1000 { + perf::note_cross_page(); let mut buf = [0u8; 8]; for i in 0..size { @@ -1119,19 +1541,20 @@ pub fn raw_load( buf[i as usize] = bus.read_byte(byte_physical_address); } - Ok(u64::from_le_bytes(buf)) - } else { - let physical_address = mmu.translate_load(virtual_address, satp, bus)?; - - Ok(match size { - 1 => bus.read_byte(physical_address) as u64, - 2 => bus.read_halfword(physical_address) as u64, - 4 => bus.read_word(physical_address) as u64, - _ => bus.read_doubleword(physical_address), - }) + return Ok(u64::from_le_bytes(buf)); } + + let physical_address = mmu.translate_load(virtual_address, satp, bus)?; + + Ok(match size { + 1 => bus.read_byte(physical_address) as u64, + 2 => bus.read_halfword(physical_address) as u64, + 4 => bus.read_word(physical_address) as u64, + _ => bus.read_doubleword(physical_address), + }) } +#[inline] pub fn raw_store( mmu: &mut Mmu, bus: &mut B, @@ -1141,7 +1564,9 @@ pub fn raw_store( val: u64, size: u64, ) -> Result<(), MmuFault> { + perf::note_store(); if (virtual_address & 0xFFF) + size > 0x1000 { + perf::note_cross_page(); let bytes = val.to_le_bytes(); for i in 0..size { @@ -1160,6 +1585,17 @@ pub fn raw_store( 4 => bus.write_word(physical_address, val as u32), _ => bus.write_doubleword(physical_address, val), } + + // Fill after the write: the write already materialized/dirty-tracked + // the page, so ram_epoch is stable and notify_store just ran (any + // code on this page was evicted; code_generation guards refills). + mmu.store_fast_fill( + virtual_address >> 12, + satp, + physical_address, + blocks.code_generation(), + bus, + ); } Ok(()) } @@ -1182,6 +1618,55 @@ mod tests { mem } + const ZBA_PROGRAM: &[(u64, u32)] = &[ + (0x00, 0x00500093), // addi x1, x0, 5 + (0x04, 0x06400113), // addi x2, x0, 100 + (0x08, 0x2020a233), // sh1add x4, x1, x2 + (0x0c, 0x2020c2b3), // sh2add x5, x1, x2 + (0x10, 0x2020e1b3), // sh3add x3, x1, x2 + (0x14, 0xfff00313), // addi x6, x0, -1 + (0x18, 0x082303bb), // add.uw x7, x6, x2 + (0x1c, 0x0843141b), // slli.uw x8, x6, 4 + (0x20, 0x202364bb), // sh3add.uw x9, x6, x2 + (0x24, 0x0803453b), // zext.h x10, x6 + (0x28, 0x00000073), // ecall + ]; + + const ZBA_EXPECTED: &[(usize, u64)] = &[ + (4, 110), // (5 << 1) + 100 + (5, 120), // (5 << 2) + 100 + (3, 140), // (5 << 3) + 100 + (7, 0x1_0000_0063), // zext32(-1) + 100 + (8, 0xf_ffff_fff0), // zext32(-1) << 4 + (9, 0x8_0000_005c), // (zext32(-1) << 3) + 100 + (10, 0xffff), // zext.h(-1) + ]; + + #[test] + fn zba_block_path() { + let mut mem = mem_with(ZBA_PROGRAM); + let mut hart = Hart::new(0); + + assert_eq!(hart.run(&mut mem, 32), StepResult::Ok); + assert_eq!(hart.regs.pc, 0x28); + for &(reg, want) in ZBA_EXPECTED { + assert_eq!(hart.regs.read(reg), want, "x{reg} on the block path"); + } + } + + #[test] + fn zba_step_path_matches_block_path() { + let mut mem = mem_with(ZBA_PROGRAM); + let mut hart = Hart::new(0); + + while hart.regs.pc != 0x28 { + assert_eq!(hart.step(&mut mem), StepResult::Ok); + } + for &(reg, want) in ZBA_EXPECTED { + assert_eq!(hart.regs.read(reg), want, "x{reg} on the single-step path"); + } + } + #[test] fn block_loop_sum() { let mut mem = mem_with(&[ diff --git a/crates/riscv-core/src/csr.rs b/crates/riscv-core/src/csr.rs index 196cd42..23e73c2 100644 --- a/crates/riscv-core/src/csr.rs +++ b/crates/riscv-core/src/csr.rs @@ -404,12 +404,18 @@ impl Csr { self.mstatus = (self.mstatus & preserved) | (value & MSTATUS_WRITE_MASK); } + #[inline(always)] pub fn pending_interrupt(&self, priv_mode: PrivMode) -> Option { let pending = self.mip & self.mie; if pending == 0 { return None; } + self.pending_interrupt_slow(priv_mode, pending) + } + #[cold] + #[inline(never)] + fn pending_interrupt_slow(&self, priv_mode: PrivMode, pending: u64) -> Option { let mie_bit = (self.mstatus & MSTATUS_MIE) != 0; let sie_bit = (self.mstatus & MSTATUS_SIE) != 0; diff --git a/crates/riscv-core/src/execute.rs b/crates/riscv-core/src/execute.rs index d223784..3b61754 100644 --- a/crates/riscv-core/src/execute.rs +++ b/crates/riscv-core/src/execute.rs @@ -2,19 +2,78 @@ use crate::block::{self, BlockCache}; use crate::csr::{ - Csr, MSTATUS_FS, MSTATUS_MIE, MSTATUS_MPIE, MSTATUS_MPP, MSTATUS_SIE, MSTATUS_SPIE, + Csr, MIP_MTIP, MSTATUS_FS, MSTATUS_MIE, MSTATUS_MPIE, MSTATUS_MPP, MSTATUS_SIE, MSTATUS_SPIE, MSTATUS_SPP, PrivMode, }; use crate::decode::{Instruction, sign_extend}; use crate::extensions as ext; use crate::gpr::Gpr; use crate::mmu::{Mmu, MmuFault}; +use crate::perf; use crate::system_bus::SystemBus; use crate::trap::{StepResult, TrapCause}; +#[cfg(feature = "aot")] +use crate::aot; + pub const ICACHE_SIZE: usize = 4096; const ICACHE_TAG_SHIFT: u32 = 1 + ICACHE_SIZE.trailing_zeros(); +pub const FETCH_TLB_SIZE: usize = 256; + +#[derive(Clone, Copy)] +pub struct FetchTlbEntry { + pub vpage: u64, + pub ppage: u64, + pub satp: u64, +} + +pub struct FetchTlb { + pub entries: Box<[FetchTlbEntry; FETCH_TLB_SIZE]>, +} + +impl FetchTlb { + pub fn new() -> Self { + Self { + entries: Box::new( + [FetchTlbEntry { + vpage: u64::MAX, + ppage: 0, + satp: u64::MAX, + }; FETCH_TLB_SIZE], + ), + } + } + + #[inline(always)] + pub fn lookup(&self, vpage: u64, satp: u64) -> Option { + let entry = &self.entries[(vpage as usize) & (FETCH_TLB_SIZE - 1)]; + if entry.vpage == vpage && entry.satp == satp { + Some(entry.ppage) + } else { + None + } + } + + #[inline(always)] + pub fn insert(&mut self, vpage: u64, ppage: u64, satp: u64) { + self.entries[(vpage as usize) & (FETCH_TLB_SIZE - 1)] = + FetchTlbEntry { vpage, ppage, satp }; + } + + pub fn flush(&mut self) { + for entry in self.entries.iter_mut() { + entry.vpage = u64::MAX; + } + } +} + +impl Default for FetchTlb { + fn default() -> Self { + Self::new() + } +} + const OP_LUI: u32 = 0x37; const OP_AUIPC: u32 = 0x17; const OP_JAL: u32 = 0x6f; @@ -64,9 +123,7 @@ pub struct ExecContext<'a, B: SystemBus> { pub bus: &'a mut B, pub priv_mode: &'a mut PrivMode, pub lr_addr: &'a mut Option, - pub fetch_vpage: &'a mut u64, - pub fetch_ppage: &'a mut u64, - pub fetch_satp: &'a mut u64, + pub fetch_tlb: &'a mut FetchTlb, pub icache_tags: &'a mut Box<[u64; ICACHE_SIZE]>, pub icache_data: &'a mut Box<[u32; ICACHE_SIZE]>, @@ -75,22 +132,39 @@ pub struct ExecContext<'a, B: SystemBus> { } fn invalidate_fetch_cache(ctx: &mut ExecContext) { - *ctx.fetch_vpage = u64::MAX; + ctx.fetch_tlb.flush(); ctx.icache_tags.fill(u64::MAX); } pub fn run(ctx: &mut ExecContext, max_steps: u64) -> StepResult { + run_impl::(ctx, max_steps) +} + +pub fn run_until_wait(ctx: &mut ExecContext, max_steps: u64) -> StepResult { + run_impl::(ctx, max_steps) +} + +fn run_impl( + ctx: &mut ExecContext, + max_steps: u64, +) -> StepResult { let mut remaining = max_steps as i64; while remaining > 0 { if let Some(irq) = ctx.csr.pending_interrupt(*ctx.priv_mode) { - *ctx.fetch_vpage = u64::MAX; + if irq == 7 && ctx.bus.timer_interrupt_pending() == Some(false) { + ctx.csr.mip &= !MIP_MTIP; + continue; + } + + ctx.fetch_tlb.flush(); *ctx.is_waiting = false; match take_interrupt(ctx, irq) { StepResult::Ok => {} other => return other, } + remaining -= 1; continue; } @@ -99,14 +173,22 @@ pub fn run(ctx: &mut ExecContext, max_steps: u64) -> StepResult let effective_satp = block::effective_satp(*ctx.priv_mode, ctx.csr.satp); let virtual_page = pc >> 12; - let fetch_pa = if virtual_page == *ctx.fetch_vpage && effective_satp == *ctx.fetch_satp { - (*ctx.fetch_ppage << 12) | (pc & 0xfff) + let fetch_pa = if let Some(ppage) = ctx.fetch_tlb.lookup(virtual_page, effective_satp) { + perf::note_fetch_page_hit(); + #[cfg(debug_assertions)] + debug_assert_eq!( + ctx.mmu + .translate_fetch(pc, effective_satp, ctx.bus) + .map(|pa| pa >> 12), + Ok(ppage), + "fetch TLB hit disagrees with slow-path translation" + ); + (ppage << 12) | (pc & 0xfff) } else { + perf::note_fetch_translate(); match ctx.mmu.translate_fetch(pc, effective_satp, ctx.bus) { Ok(pa) => { - *ctx.fetch_vpage = virtual_page; - *ctx.fetch_ppage = pa >> 12; - *ctx.fetch_satp = effective_satp; + ctx.fetch_tlb.insert(virtual_page, pa >> 12, effective_satp); pa } Err(fault) => { @@ -120,33 +202,135 @@ pub fn run(ctx: &mut ExecContext, max_steps: u64) -> StepResult } }; + #[cfg(feature = "aot")] + if let Some(key_pa) = aot_page_key(ctx, fetch_pa) { + let priv_at_entry = *ctx.priv_mode; + if let Some(retired) = aot::dispatch( + ctx, + key_pa, + pc, + effective_satp, + remaining as u64, + fetch_pa >> 12, + ) { + aot::note_dispatch(retired); + perf::note_retired(priv_at_entry, retired); + remaining -= retired as i64; + continue; + } + } + + #[cfg(feature = "aot-trace")] + ctx.blocks.trace_exec(fetch_pa); + if let Some(cached) = ctx.blocks.lookup(fetch_pa) { - remaining -= block::exec_block(ctx, &cached, pc, effective_satp) as i64; + perf::note_block_hit(); + let priv_at_entry = *ctx.priv_mode; + let (retired, result) = block::exec_block(ctx, &cached, pc, effective_satp); + perf::note_retired(priv_at_entry, retired); + remaining -= retired as i64; + + match result { + StepResult::Ok => {} + other => return other, + } + + if STOP_ON_WFI && *ctx.is_waiting { + return StepResult::Ok; + } + continue; } if let Some(decoded) = block::decode_block(ctx.bus, fetch_pa) { + perf::note_block_decode(); + let cached = ctx.blocks.insert(fetch_pa, decoded); - remaining -= block::exec_block(ctx, &cached, pc, effective_satp) as i64; + let priv_at_entry = *ctx.priv_mode; + let (retired, result) = block::exec_block(ctx, &cached, pc, effective_satp); + + perf::note_retired(priv_at_entry, retired); + + remaining -= retired as i64; + match result { + StepResult::Ok => {} + other => return other, + } + if STOP_ON_WFI && *ctx.is_waiting { + return StepResult::Ok; + } continue; } + perf::note_single_step(); + #[cfg(feature = "perf-counters")] + { + let low = ctx.bus.read_halfword(fetch_pa) as u32; + let raw = if low & 0x3 == 0x3 { + low | ((ctx.bus.read_halfword(fetch_pa + 2) as u32) << 16) + } else { + low + }; + perf::note_single_step_op(raw); + } + perf::note_retired(*ctx.priv_mode, 1); + match step(ctx) { StepResult::Ok => {} other => return other, } + if STOP_ON_WFI && *ctx.is_waiting { + return StepResult::Ok; + } + remaining -= 1; } StepResult::Ok } +#[cfg(feature = "aot")] +#[inline(always)] +pub(crate) fn aot_page_key(ctx: &mut ExecContext, fetch_pa: u64) -> Option { + use crate::block::AotResolve; + + let page = fetch_pa >> 12; + match ctx.blocks.aot_resolve(page) { + AotResolve::Hit(orig) => Some((orig << 12) | (fetch_pa & 0xfff)), + AotResolve::Miss => None, + AotResolve::Unknown => aot_page_key_hash(ctx, fetch_pa, page), + } +} + +#[cfg(feature = "aot")] +#[cold] +#[inline(never)] +fn aot_page_key_hash( + ctx: &mut ExecContext, + fetch_pa: u64, + page: u64, +) -> Option { + let base = page << 12; + let mut hash = 0xcbf2_9ce4_8422_2325u64; + for i in 0..512u64 { + hash ^= ctx.bus.read_doubleword(base + i * 8); + hash = hash.wrapping_mul(0x0000_0100_0000_01b3); + } + ctx.blocks + .aot_record_hash(page, hash) + .map(|orig| (orig << 12) | (fetch_pa & 0xfff)) +} + pub fn step(ctx: &mut ExecContext) -> StepResult { if let Some(irq) = ctx.csr.pending_interrupt(*ctx.priv_mode) { - *ctx.fetch_vpage = u64::MAX; - *ctx.is_waiting = false; - return take_interrupt(ctx, irq); + if irq == 7 && ctx.bus.timer_interrupt_pending() == Some(false) { + ctx.csr.mip &= !MIP_MTIP; + } else { + ctx.fetch_tlb.flush(); + *ctx.is_waiting = false; + return take_interrupt(ctx, irq); + } } let pc = ctx.regs.pc; @@ -158,16 +342,16 @@ pub fn step(ctx: &mut ExecContext) -> StepResult { let virtual_page = pc >> 12; let fetch_physical_address = - if virtual_page == *ctx.fetch_vpage && effective_satp == *ctx.fetch_satp { - (*ctx.fetch_ppage << 12) | (pc & 0xfff) + if let Some(ppage) = ctx.fetch_tlb.lookup(virtual_page, effective_satp) { + (ppage << 12) | (pc & 0xfff) } else { let physical_address = match ctx.mmu.translate_fetch(pc, effective_satp, ctx.bus) { Ok(pa) => pa, Err(fault) => return trap_from_mmu(ctx, fault), }; - *ctx.fetch_vpage = virtual_page; - *ctx.fetch_ppage = physical_address >> 12; - *ctx.fetch_satp = effective_satp; + + ctx.fetch_tlb + .insert(virtual_page, physical_address >> 12, effective_satp); physical_address }; @@ -183,9 +367,8 @@ pub fn step(ctx: &mut ExecContext) -> StepResult { Err(f) => return trap_from_mmu(ctx, f), }; - *ctx.fetch_vpage = next_pc >> 12; - *ctx.fetch_ppage = next_pa >> 12; - *ctx.fetch_satp = effective_satp; + ctx.fetch_tlb + .insert(next_pc >> 12, next_pa >> 12, effective_satp); let hi = ctx.bus.read_halfword(next_pa) as u32; lo | (hi << 16) @@ -203,6 +386,14 @@ pub fn step(ctx: &mut ExecContext) -> StepResult { } }; + exec_raw(ctx, instruction_encoding, pc) +} + +pub(crate) fn exec_raw( + ctx: &mut ExecContext, + instruction_encoding: u32, + pc: u64, +) -> StepResult { let result = if instruction_encoding & 0x3 != 0x3 { exec_compressed(ctx, instruction_encoding as u16) } else { @@ -486,29 +677,34 @@ fn exec_full( let shamt = (imm & 0x1f) as u32; let funct7 = inst.funct7(); - let val: i32 = match inst.funct3() { - 0x0 => rs1.wrapping_add(imm as u32) as i32, - 0x1 => match funct7 { - 0x00 => (rs1 << shamt) as i32, // slliw - 0x30 => match shamt { - // Zbb: clzw/ctzw/cpopw - 0 => rs1.leading_zeros() as i32, // clzw - 1 => rs1.trailing_zeros() as i32, // ctzw - 2 => rs1.count_ones() as i32, // cpopw + let val: u64 = if inst.funct3() == 0x1 && (funct7 == 0x04 || funct7 == 0x05) { + (rs1 as u64) << ((imm & 0x3f) as u32) + } else { + let word: i32 = match inst.funct3() { + 0x0 => rs1.wrapping_add(imm as u32) as i32, + 0x1 => match funct7 { + 0x00 => (rs1 << shamt) as i32, // slliw + 0x30 => match shamt { + // Zbb: clzw/ctzw/cpopw + 0 => rs1.leading_zeros() as i32, // clzw + 1 => rs1.trailing_zeros() as i32, // ctzw + 2 => rs1.count_ones() as i32, // cpopw + _ => return StepResult::Trap(TrapCause::IllegalInstruction(raw)), + }, + _ => return StepResult::Trap(TrapCause::IllegalInstruction(raw)), + }, + 0x5 => match funct7 { + 0x00 => (rs1 >> shamt) as i32, // srliw + 0x20 => (rs1 as i32) >> shamt, // sraiw + 0x30 => rs1.rotate_right(shamt) as i32, // roriw (Zbb) _ => return StepResult::Trap(TrapCause::IllegalInstruction(raw)), }, _ => return StepResult::Trap(TrapCause::IllegalInstruction(raw)), - }, - 0x5 => match funct7 { - 0x00 => (rs1 >> shamt) as i32, // srliw - 0x20 => (rs1 as i32) >> shamt, // sraiw - 0x30 => rs1.rotate_right(shamt) as i32, // roriw (Zbb) - _ => return StepResult::Trap(TrapCause::IllegalInstruction(raw)), - }, - _ => return StepResult::Trap(TrapCause::IllegalInstruction(raw)), + }; + word as i64 as u64 }; - ctx.regs.write(inst.rd(), val as i64 as u64); + ctx.regs.write(inst.rd(), val); ctx.regs.pc = pc.wrapping_add(4); } @@ -561,8 +757,10 @@ fn exec_full( } // max (0x6, 0x05) => rs1.min(rs2), // minu (0x7, 0x05) => rs1.max(rs2), // maxu - // Zbb: zext.h (pack rs2=x0, funct7=0x04) - (0x4, 0x04) => rs1 as u16 as u64, // zext.h + // Zba: shifted add + (0x2, 0x10) => (rs1 << 1).wrapping_add(rs2), // sh1add + (0x4, 0x10) => (rs1 << 2).wrapping_add(rs2), // sh2add + (0x6, 0x10) => (rs1 << 3).wrapping_add(rs2), // sh3add _ => return StepResult::Trap(TrapCause::IllegalInstruction(raw)), }; @@ -590,6 +788,13 @@ fn exec_full( (0x5, 0x30) => { ((rs1 as u32).rotate_right((rs2 & 0x1f) as u32)) as i32 as i64 as u64 } + // Zbb: zext.h is `packw rd, rs1, x0`, so it only decodes with rs2 == 0. + (0x4, 0x04) if inst.rs2() == 0 => rs1 as u16 as u64, + // Zba: zero-extended shifted add + (0x0, 0x04) => (rs1 as u32 as u64).wrapping_add(rs2), // add.uw + (0x2, 0x10) => ((rs1 as u32 as u64) << 1).wrapping_add(rs2), // sh1add.uw + (0x4, 0x10) => ((rs1 as u32 as u64) << 2).wrapping_add(rs2), // sh2add.uw + (0x6, 0x10) => ((rs1 as u32 as u64) << 3).wrapping_add(rs2), // sh3add.uw _ => return StepResult::Trap(TrapCause::IllegalInstruction(raw)), }; @@ -752,9 +957,6 @@ fn exec_system(ctx: &mut ExecContext, inst: Instruction, raw: u ctx.csr.mstatus |= MSTATUS_SPIE; *ctx.priv_mode = PrivMode::from_bits(spp); ctx.regs.pc = ctx.csr.sepc; - ctx.mmu.flush(); - - invalidate_fetch_cache(ctx); return StepResult::Ok; } 0x302 => { @@ -772,13 +974,12 @@ fn exec_system(ctx: &mut ExecContext, inst: Instruction, raw: u ctx.csr.mstatus |= MSTATUS_MPIE; *ctx.priv_mode = PrivMode::from_bits(mpp); ctx.regs.pc = ctx.csr.mepc; - ctx.mmu.flush(); - - invalidate_fetch_cache(ctx); return StepResult::Ok; } 0x105 => { // WFI + perf::note_wfi(); + *ctx.is_waiting = true; ctx.regs.pc = pc.wrapping_add(4); @@ -1346,8 +1547,6 @@ fn take_interrupt(ctx: &mut ExecContext, irq_bit: u64) -> StepR ctx.regs.pc = ctx.csr.mtvec & !3; } - ctx.mmu.flush(); - invalidate_fetch_cache(ctx); StepResult::Ok } @@ -1383,9 +1582,6 @@ pub fn take_exception(ctx: &mut ExecContext, cause: u64, tval: *ctx.priv_mode = PrivMode::M; ctx.regs.pc = ctx.csr.mtvec & !3; } - - ctx.mmu.flush(); - invalidate_fetch_cache(ctx); } // Zbb: orc.b diff --git a/crates/riscv-core/src/hart.rs b/crates/riscv-core/src/hart.rs index 6ee6524..416467a 100644 --- a/crates/riscv-core/src/hart.rs +++ b/crates/riscv-core/src/hart.rs @@ -1,6 +1,6 @@ use crate::block::BlockCache; use crate::csr::{Csr, PrivMode}; -use crate::execute::{self, ExecContext}; +use crate::execute::{self, ExecContext, FetchTlb}; use crate::execute::ICACHE_SIZE; @@ -15,9 +15,7 @@ pub struct Hart { pub mmu: Mmu, pub priv_mode: PrivMode, pub lr_addr: Option, - pub fetch_vpage: u64, - pub fetch_ppage: u64, - pub fetch_satp: u64, + pub fetch_tlb: FetchTlb, pub icache_tags: Box<[u64; ICACHE_SIZE]>, pub icache_data: Box<[u32; ICACHE_SIZE]>, @@ -33,9 +31,7 @@ impl Hart { mmu: Mmu::new(), priv_mode: PrivMode::M, lr_addr: None, - fetch_vpage: u64::MAX, - fetch_ppage: 0, - fetch_satp: u64::MAX, + fetch_tlb: FetchTlb::new(), icache_tags: Box::new([u64::MAX; ICACHE_SIZE]), icache_data: Box::new([0u32; ICACHE_SIZE]), @@ -57,9 +53,7 @@ impl Hart { bus, priv_mode: &mut self.priv_mode, lr_addr: &mut self.lr_addr, - fetch_vpage: &mut self.fetch_vpage, - fetch_ppage: &mut self.fetch_ppage, - fetch_satp: &mut self.fetch_satp, + fetch_tlb: &mut self.fetch_tlb, icache_tags: &mut self.icache_tags, icache_data: &mut self.icache_data, @@ -79,9 +73,7 @@ impl Hart { bus, priv_mode: &mut self.priv_mode, lr_addr: &mut self.lr_addr, - fetch_vpage: &mut self.fetch_vpage, - fetch_ppage: &mut self.fetch_ppage, - fetch_satp: &mut self.fetch_satp, + fetch_tlb: &mut self.fetch_tlb, icache_tags: &mut self.icache_tags, icache_data: &mut self.icache_data, is_waiting: &mut self.is_waiting, @@ -90,4 +82,22 @@ impl Hart { execute::run(&mut ctx, max_steps) } + + pub fn run_until_wait(&mut self, bus: &mut impl SystemBus, max_steps: u64) -> StepResult { + let mut ctx = ExecContext { + regs: &mut self.regs, + csr: &mut self.csr, + mmu: &mut self.mmu, + bus, + priv_mode: &mut self.priv_mode, + lr_addr: &mut self.lr_addr, + fetch_tlb: &mut self.fetch_tlb, + icache_tags: &mut self.icache_tags, + icache_data: &mut self.icache_data, + is_waiting: &mut self.is_waiting, + blocks: &mut self.blocks, + }; + + execute::run_until_wait(&mut ctx, max_steps) + } } diff --git a/crates/riscv-core/src/lib.rs b/crates/riscv-core/src/lib.rs index 06ada71..5f98865 100644 --- a/crates/riscv-core/src/lib.rs +++ b/crates/riscv-core/src/lib.rs @@ -1,3 +1,6 @@ +#[cfg(feature = "aot")] +pub mod aot; + pub mod block; pub mod csr; pub mod decode; @@ -6,6 +9,7 @@ pub mod extensions; pub mod gpr; pub mod hart; pub mod mmu; +pub mod perf; pub mod system_bus; pub mod trap; diff --git a/crates/riscv-core/src/mmu.rs b/crates/riscv-core/src/mmu.rs index 1f64079..d36a86c 100644 --- a/crates/riscv-core/src/mmu.rs +++ b/crates/riscv-core/src/mmu.rs @@ -1,3 +1,4 @@ +use crate::perf; use crate::system_bus::SystemBus; const PTE_V: u64 = 1 << 0; @@ -26,9 +27,64 @@ impl TlbEntry { }; } +#[derive(Clone, Copy)] +struct LoadFastEntry { + virt_page_num: u64, + satp: u64, + ram_epoch: u64, + epoch: u32, + host_page: *const u8, +} + +impl LoadFastEntry { + const EMPTY: Self = Self { + virt_page_num: u64::MAX, + satp: 0, + ram_epoch: 0, + epoch: 0, + host_page: std::ptr::null(), + }; +} + +unsafe impl Send for LoadFastEntry {} +unsafe impl Sync for LoadFastEntry {} + +#[derive(Clone, Copy)] +struct StoreFastEntry { + virt_page_num: u64, + satp: u64, + ram_epoch: u64, + code_generation: u64, + epoch: u32, + host_page: *mut u8, +} + +impl StoreFastEntry { + const EMPTY: Self = Self { + virt_page_num: u64::MAX, + satp: 0, + ram_epoch: 0, + code_generation: 0, + epoch: 0, + host_page: std::ptr::null_mut(), + }; +} + +unsafe impl Send for StoreFastEntry {} +unsafe impl Sync for StoreFastEntry {} + pub struct Mmu { tlb: [TlbEntry; TLB_SIZE], + load_fast: [LoadFastEntry; TLB_SIZE], + store_fast: [StoreFastEntry; TLB_SIZE], epoch: u32, + + load_vpage: u64, + load_ppage: u64, + load_satp: u64, + store_vpage: u64, + store_ppage: u64, + store_satp: u64, } #[derive(Debug, Clone, PartialEq, Eq)] @@ -74,7 +130,15 @@ impl Mmu { pub fn new() -> Self { Self { tlb: [TlbEntry::EMPTY; TLB_SIZE], + load_fast: [LoadFastEntry::EMPTY; TLB_SIZE], + store_fast: [StoreFastEntry::EMPTY; TLB_SIZE], epoch: 1, + load_vpage: u64::MAX, + load_ppage: 0, + load_satp: u64::MAX, + store_vpage: u64::MAX, + store_ppage: 0, + store_satp: u64::MAX, } } @@ -83,6 +147,8 @@ impl Mmu { if self.epoch == 0 { self.epoch = 1; } + self.load_vpage = u64::MAX; + self.store_vpage = u64::MAX; } #[inline(always)] @@ -97,6 +163,90 @@ impl Mmu { } } + #[inline(always)] + pub fn load_fast_lookup( + &self, + virtual_address: u64, + satp: u64, + ram_epoch: u64, + ) -> Option<*const u8> { + let virt_page_num = virtual_address >> 12; + let entry = &self.load_fast[(virt_page_num & TLB_MASK) as usize]; + + if entry.virt_page_num == virt_page_num + && entry.satp == satp + && entry.epoch == self.epoch + && entry.ram_epoch == ram_epoch + { + Some(entry.host_page) + } else { + None + } + } + + #[inline(always)] + pub fn store_fast_lookup( + &self, + virtual_address: u64, + satp: u64, + ram_epoch: u64, + code_generation: u64, + ) -> Option<*mut u8> { + let virt_page_num = virtual_address >> 12; + let entry = &self.store_fast[(virt_page_num & TLB_MASK) as usize]; + + if entry.virt_page_num == virt_page_num + && entry.satp == satp + && entry.epoch == self.epoch + && entry.ram_epoch == ram_epoch + && entry.code_generation == code_generation + { + Some(entry.host_page) + } else { + None + } + } + + #[inline] + pub fn store_fast_fill( + &mut self, + virt_page_num: u64, + satp: u64, + physical_address: u64, + code_generation: u64, + bus: &mut impl SystemBus, + ) { + if let Some(host_page) = bus.ram_store_page(physical_address) { + self.store_fast[(virt_page_num & TLB_MASK) as usize] = StoreFastEntry { + virt_page_num, + satp, + ram_epoch: bus.ram_epoch(), + code_generation, + epoch: self.epoch, + host_page, + }; + } + } + + #[inline] + fn load_fast_fill( + &mut self, + virt_page_num: u64, + satp: u64, + physical_address: u64, + bus: &mut impl SystemBus, + ) { + if let Some(host_page) = bus.ram_load_page(physical_address) { + self.load_fast[(virt_page_num & TLB_MASK) as usize] = LoadFastEntry { + virt_page_num, + satp, + ram_epoch: bus.ram_epoch(), + epoch: self.epoch, + host_page, + }; + } + } + #[inline(always)] fn insert(&mut self, virt_page_num: u64, phys_page_num: u64, flags: u64) { let slot = (virt_page_num & TLB_MASK) as usize; @@ -115,6 +265,7 @@ impl Mmu { bus: &mut impl SystemBus, ) -> Result { if satp >> 60 == 0 { + perf::note_bare_translate(); return Ok(virtual_address); } @@ -122,9 +273,11 @@ impl Mmu { if let Some((ppn, flags)) = self.lookup(vpn) && flags & PTE_X != 0 { + perf::note_tlb_hit(); return Ok((ppn << 12) | (virtual_address & 0xfff)); } + perf::note_tlb_walk(); self.walk(virtual_address, satp, false, true, bus) .map_err(|_| MmuFault::InstructionPageFault(virtual_address)) } @@ -135,19 +288,42 @@ impl Mmu { satp: u64, bus: &mut impl SystemBus, ) -> Result { + let vpn = virtual_address >> 12; + if satp >> 60 == 0 { + perf::note_bare_translate(); + self.load_fast_fill(vpn, satp, virtual_address, bus); return Ok(virtual_address); } - let vpn = virtual_address >> 12; + if vpn == self.load_vpage && satp == self.load_satp { + perf::note_tlb_hit(); + let pa = (self.load_ppage << 12) | (virtual_address & 0xfff); + self.load_fast_fill(vpn, satp, pa, bus); + return Ok(pa); + } + if let Some((ppn, flags)) = self.lookup(vpn) && flags & PTE_R != 0 { - return Ok((ppn << 12) | (virtual_address & 0xfff)); + perf::note_tlb_hit(); + self.load_vpage = vpn; + self.load_ppage = ppn; + self.load_satp = satp; + let pa = (ppn << 12) | (virtual_address & 0xfff); + self.load_fast_fill(vpn, satp, pa, bus); + return Ok(pa); } - self.walk(virtual_address, satp, false, false, bus) - .map_err(|_| MmuFault::LoadPageFault(virtual_address)) + perf::note_tlb_walk(); + let pa = self + .walk(virtual_address, satp, false, false, bus) + .map_err(|_| MmuFault::LoadPageFault(virtual_address))?; + self.load_vpage = vpn; + self.load_ppage = pa >> 12; + self.load_satp = satp; + self.load_fast_fill(vpn, satp, pa, bus); + Ok(pa) } pub fn translate_store( @@ -157,18 +333,34 @@ impl Mmu { bus: &mut impl SystemBus, ) -> Result { if satp >> 60 == 0 { + perf::note_bare_translate(); return Ok(virtual_address); } let vpn = virtual_address >> 12; + if vpn == self.store_vpage && satp == self.store_satp { + perf::note_tlb_hit(); + return Ok((self.store_ppage << 12) | (virtual_address & 0xfff)); + } + if let Some((ppn, flags)) = self.lookup(vpn) && flags & (PTE_W | PTE_D) == (PTE_W | PTE_D) { + perf::note_tlb_hit(); + self.store_vpage = vpn; + self.store_ppage = ppn; + self.store_satp = satp; return Ok((ppn << 12) | (virtual_address & 0xfff)); } - self.walk(virtual_address, satp, true, false, bus) - .map_err(|_| MmuFault::StorePageFault(virtual_address)) + perf::note_tlb_walk(); + let pa = self + .walk(virtual_address, satp, true, false, bus) + .map_err(|_| MmuFault::StorePageFault(virtual_address))?; + self.store_vpage = vpn; + self.store_ppage = pa >> 12; + self.store_satp = satp; + Ok(pa) } fn walk( diff --git a/crates/riscv-core/src/perf.rs b/crates/riscv-core/src/perf.rs new file mode 100644 index 0000000..e28cee6 --- /dev/null +++ b/crates/riscv-core/src/perf.rs @@ -0,0 +1,210 @@ +// Feature-gated profiling counters + +use std::sync::atomic::AtomicU64; + +#[cfg(feature = "perf-counters")] +use std::sync::atomic::Ordering::Relaxed; + +use crate::csr::PrivMode; + +macro_rules! counters { + ($($name:ident),* $(,)?) => { + $(pub static $name: AtomicU64 = AtomicU64::new(0);)* + + #[cfg(feature = "perf-counters")] + const ALL: &[(&str, &AtomicU64)] = &[$((stringify!($name), &$name)),*]; + }; +} + +counters!( + INSNS_MMODE, + INSNS_KERNEL, + INSNS_USER, + WFI_EXECUTED, + BLOCK_HITS, + BLOCK_DECODES, + SINGLE_STEPS, + FETCH_PAGE_HITS, + FETCH_TRANSLATES, + LOADS, + LOAD_FAST_HITS, + STORES, + STORE_FAST_HITS, + CROSS_PAGE_ACCESSES, + TLB_HITS, + TLB_WALKS, + BARE_TRANSLATES, + STORE_PAGE_EVICTIONS, + FALLBACK_OPS, + SS_SYSTEM, + SS_AMO, + SS_FP, + SS_MISC_MEM, + SS_OTHER, +); + +#[inline(always)] +pub fn note_single_step_op(raw: u32) { + #[cfg(feature = "perf-counters")] + { + let counter = if raw & 0x3 != 0x3 { + &SS_OTHER // compressed insn the C-decoder refused + } else { + match raw & 0x7f { + 0x73 => &SS_SYSTEM, // csr / ecall / ebreak / sret / wfi / sfence + 0x2f => &SS_AMO, // atomics / lr / sc + 0x0f => &SS_MISC_MEM, // fence / fence.i + 0x07 | 0x27 | 0x43 | 0x47 | 0x4b | 0x4f | 0x53 => &SS_FP, + _ => &SS_OTHER, + } + }; + counter.fetch_add(1, Relaxed); + } + #[cfg(not(feature = "perf-counters"))] + let _ = raw; +} + +#[inline(always)] +pub fn note_retired(priv_mode: PrivMode, count: u64) { + #[cfg(feature = "perf-counters")] + { + let counter = match priv_mode { + PrivMode::M => &INSNS_MMODE, + PrivMode::S => &INSNS_KERNEL, + PrivMode::U => &INSNS_USER, + }; + counter.fetch_add(count, Relaxed); + } + #[cfg(not(feature = "perf-counters"))] + let _ = (priv_mode, count); +} + +macro_rules! note_fns { + ($($fn_name:ident => $counter:ident),* $(,)?) => { + $( + #[inline(always)] + pub fn $fn_name() { + #[cfg(feature = "perf-counters")] + $counter.fetch_add(1, Relaxed); + } + )* + }; +} + +note_fns!( + note_wfi => WFI_EXECUTED, + note_block_hit => BLOCK_HITS, + note_block_decode => BLOCK_DECODES, + note_single_step => SINGLE_STEPS, + note_fetch_page_hit => FETCH_PAGE_HITS, + note_fetch_translate => FETCH_TRANSLATES, + note_load => LOADS, + note_load_fast_hit => LOAD_FAST_HITS, + note_store => STORES, + note_store_fast_hit => STORE_FAST_HITS, + note_cross_page => CROSS_PAGE_ACCESSES, + note_tlb_hit => TLB_HITS, + note_tlb_walk => TLB_WALKS, + note_bare_translate => BARE_TRANSLATES, + note_store_page_eviction => STORE_PAGE_EVICTIONS, + note_fallback_op => FALLBACK_OPS, +); + +pub fn report() -> Option { + #[cfg(feature = "perf-counters")] + { + let get = |c: &AtomicU64| c.swap(0, Relaxed); + let values: Vec<(&str, u64)> = ALL.iter().map(|(n, c)| (*n, get(c))).collect(); + let total_insns: u64 = values[..3].iter().map(|(_, v)| v).sum(); + if total_insns == 0 { + return None; + } + + let v = |name: &str| values.iter().find(|(n, _)| *n == name).unwrap().1; + let pct = |part: u64, whole: u64| { + if whole == 0 { + 0.0 + } else { + 100.0 * part as f64 / whole as f64 + } + }; + + let total_loads = v("LOADS") + v("LOAD_FAST_HITS"); + let total_stores = v("STORES") + v("STORE_FAST_HITS"); + let mem_accesses = total_loads + total_stores; + let translations = v("TLB_HITS") + v("TLB_WALKS") + v("BARE_TRANSLATES"); + let block_entries = v("BLOCK_HITS") + v("BLOCK_DECODES"); + + let mut out = String::from("── vpod perf counters ──────────────────────\n"); + out.push_str(&format!( + "retired: {} total | M(sbi) {:.1}% | S(kernel) {:.1}% | U(user) {:.1}%\n", + total_insns, + pct(v("INSNS_MMODE"), total_insns), + pct(v("INSNS_KERNEL"), total_insns), + pct(v("INSNS_USER"), total_insns), + )); + out.push_str(&format!( + "wfi executed: {} (kernel-idle marker)\n", + v("WFI_EXECUTED") + )); + out.push_str(&format!( + "blocks: {} hits | {} decodes ({:.2}% miss) | {} single-step fallbacks | {:.1} insns/block\n", + v("BLOCK_HITS"), + v("BLOCK_DECODES"), + pct(v("BLOCK_DECODES"), block_entries), + v("SINGLE_STEPS"), + if block_entries == 0 { 0.0 } else { total_insns as f64 / block_entries as f64 }, + )); + out.push_str(&format!( + "in-block fallback ops: {}\nsingle-step by op: {} system | {} amo | {} fp | {} fence | {} other\n", + v("FALLBACK_OPS"), + v("SS_SYSTEM"), + v("SS_AMO"), + v("SS_FP"), + v("SS_MISC_MEM"), + v("SS_OTHER"), + )); + #[cfg(feature = "aot")] + { + let calls = crate::aot::DISPATCH_CALLS.swap(0, Relaxed); + let retired = crate::aot::DISPATCH_RETIRED.swap(0, Relaxed); + out.push_str(&format!( + "aot: {} dispatches | {} insns retired ({:.1}% of all retired) | {:.1} insns/dispatch\n", + calls, + retired, + pct(retired, total_insns), + if calls == 0 { 0.0 } else { retired as f64 / calls as f64 }, + )); + } + out.push_str(&format!( + "fetch: {} page-cache hits | {} translates ({:.2}% miss)\n", + v("FETCH_PAGE_HITS"), + v("FETCH_TRANSLATES"), + pct( + v("FETCH_TRANSLATES"), + v("FETCH_PAGE_HITS") + v("FETCH_TRANSLATES") + ), + )); + out.push_str(&format!( + "memory: {} loads ({:.1}% fast-path) | {} stores ({:.1}% fast-path) | {:.1}% of insns | {} cross-page\n", + total_loads, + pct(v("LOAD_FAST_HITS"), total_loads), + total_stores, + pct(v("STORE_FAST_HITS"), total_stores), + pct(mem_accesses, total_insns), + v("CROSS_PAGE_ACCESSES"), + )); + out.push_str(&format!( + "softmmu: {} tlb hits | {} walks ({:.4}% miss) | {} bare | {} store-page evictions\n", + v("TLB_HITS"), + v("TLB_WALKS"), + pct(v("TLB_WALKS"), translations), + v("BARE_TRANSLATES"), + v("STORE_PAGE_EVICTIONS"), + )); + Some(out) + } + + #[cfg(not(feature = "perf-counters"))] + None +} diff --git a/crates/riscv-core/src/system_bus.rs b/crates/riscv-core/src/system_bus.rs index 8f38688..0d84bcd 100644 --- a/crates/riscv-core/src/system_bus.rs +++ b/crates/riscv-core/src/system_bus.rs @@ -1,4 +1,5 @@ // External communication linking the hart to RAM and peripherals (disk, network). +use std::sync::atomic::{AtomicU64, Ordering}; pub trait SystemBus { fn read_byte(&mut self, address: u64) -> u8; @@ -10,11 +11,32 @@ pub trait SystemBus { fn write_halfword(&mut self, address: u64, val: u16); fn write_word(&mut self, address: u64, val: u32); fn write_doubleword(&mut self, address: u64, val: u64); + + fn ram_load_page(&mut self, address: u64) -> Option<*const u8> { + let _ = address; + None + } + + fn ram_store_page(&mut self, address: u64) -> Option<*mut u8> { + let _ = address; + None + } + + fn ram_epoch(&self) -> u64 { + 0 + } + + fn timer_interrupt_pending(&self) -> Option { + None + } } +static FLAT_EPOCH_SOURCE: AtomicU64 = AtomicU64::new(1); + pub struct FlatMemory { data: Vec, mask: u64, + epoch: u64, } impl FlatMemory { @@ -26,6 +48,7 @@ impl FlatMemory { Self { data: vec![0u8; size_bytes], mask: (size_bytes - 1) as u64, + epoch: FLAT_EPOCH_SOURCE.fetch_add(1, Ordering::Relaxed), } } @@ -86,4 +109,26 @@ impl SystemBus for FlatMemory { let i = self.idx(address); self.data[i..i + 8].copy_from_slice(&val.to_le_bytes()); } + + fn ram_load_page(&mut self, address: u64) -> Option<*const u8> { + if self.data.len() < 0x1000 { + return None; + } + + let page_index = self.idx(address) & !0xfff; + Some(self.data[page_index..].as_ptr()) + } + + fn ram_store_page(&mut self, address: u64) -> Option<*mut u8> { + if self.data.len() < 0x1000 { + return None; + } + + let page_index = self.idx(address) & !0xfff; + Some(self.data[page_index..].as_mut_ptr()) + } + + fn ram_epoch(&self) -> u64 { + self.epoch + } } diff --git a/crates/riscv-core/tests/differential.rs b/crates/riscv-core/tests/differential.rs new file mode 100644 index 0000000..f4df9de --- /dev/null +++ b/crates/riscv-core/tests/differential.rs @@ -0,0 +1,141 @@ +// Differential harness: interpreter vs AOT lockstep. + +#![cfg(feature = "aot")] + +use riscv_core::Hart; +use riscv_core::system_bus::FlatMemory; + +fn assert_state_eq(a: &Hart, b: &Hart, entry: u64, step: u32) { + let ctx = |what: &str| { + format!( + "{what} diverged: program 0x{entry:x}, block step {step}, interp pc=0x{:x} aot pc=0x{:x}", + a.regs.pc, b.regs.pc + ) + }; + + assert_eq!(a.regs.pc, b.regs.pc, "{}", ctx("pc")); + for i in 0..32 { + assert_eq!(a.regs.read(i), b.regs.read(i), "{}", ctx(&format!("x{i}"))); + } + assert_eq!(a.priv_mode, b.priv_mode, "{}", ctx("priv_mode")); + assert_eq!(a.csr.instret, b.csr.instret, "{}", ctx("instret")); + assert_eq!(a.csr.mcause, b.csr.mcause, "{}", ctx("mcause")); + assert_eq!(a.csr.mepc, b.csr.mepc, "{}", ctx("mepc")); + assert_eq!(a.csr.mstatus, b.csr.mstatus, "{}", ctx("mstatus")); + assert_eq!(a.csr.mtval, b.csr.mtval, "{}", ctx("mtval")); + assert_eq!(a.is_waiting, b.is_waiting, "{}", ctx("is_waiting")); +} + +#[test] +fn aot_differential_lockstep() { + let Ok(dir) = std::env::var("VPOD_DIFF_DIR") else { + eprintln!("VPOD_DIFF_DIR not set — run via scripts/aot-diff.sh; skipping"); + return; + }; + + let ram = std::fs::read(format!("{dir}/ram.bin")).expect("ram.bin"); + let entries: Vec = std::fs::read_to_string(format!("{dir}/entries.txt")) + .expect("entries.txt") + .lines() + .filter_map(|l| u64::from_str_radix(l.trim(), 16).ok()) + .collect(); + let trap_pa = ram.len() as u64 - 4096; + + assert!(!entries.is_empty(), "no program entries"); + assert!( + !riscv_core::aot::AOT_PAGE_HASHES.is_empty(), + "generated.rs is the stub — regenerate via scripts/aot-diff.sh" + ); + + let mut terminated = 0usize; + for &entry in &entries { + let mut mem_interp = FlatMemory::new(ram.len()); + mem_interp.load_at(0, &ram); + let mut mem_aot = FlatMemory::new(ram.len()); + mem_aot.load_at(0, &ram); + + let mut interp = Hart::new(entry); + let mut aot = Hart::new(entry); + interp.csr.mtvec = trap_pa; + aot.csr.mtvec = trap_pa; + aot.blocks.aot_init(riscv_core::aot::AOT_PAGE_HASHES); + + for step in 0..20_000u32 { + interp.run(&mut mem_interp, 1); + aot.run(&mut mem_aot, 1); + assert_state_eq(&interp, &aot, entry, step); + + if interp.is_waiting { + break; + } + } + + if interp.is_waiting { + terminated += 1; + } + + for &chunk in &[7u64, 97, 1024] { + let mut mem_interp = FlatMemory::new(ram.len()); + mem_interp.load_at(0, &ram); + let mut mem_aot = FlatMemory::new(ram.len()); + mem_aot.load_at(0, &ram); + + let mut interp = Hart::new(entry); + let mut aot = Hart::new(entry); + interp.csr.mtvec = trap_pa; + aot.csr.mtvec = trap_pa; + aot.blocks.aot_init(riscv_core::aot::AOT_PAGE_HASHES); + + for step in 0..(40_000 / chunk as u32).max(64) { + interp.run(&mut mem_interp, chunk); + aot.run(&mut mem_aot, chunk); + assert_state_eq(&interp, &aot, entry, step); + + if interp.is_waiting { + break; + } + } + } + + { + let reloc_len = ram.len() * 2; + let reloc_entry = ram.len() as u64 + (entry & !0xfff); + + let mut mem_interp = FlatMemory::new(reloc_len); + mem_interp.load_at(0, &ram); + let code_page: Vec = + ram[(entry & !0xfff) as usize..((entry & !0xfff) + 4096) as usize].to_vec(); + mem_interp.load_at(reloc_entry as usize, &code_page); + let mut mem_aot = FlatMemory::new(reloc_len); + mem_aot.load_at(0, &ram); + mem_aot.load_at(reloc_entry as usize, &code_page); + + let mut interp = Hart::new(reloc_entry); + let mut aot = Hart::new(reloc_entry); + interp.csr.mtvec = trap_pa; + aot.csr.mtvec = trap_pa; + aot.blocks.aot_init(riscv_core::aot::AOT_PAGE_HASHES); + + for step in 0..2_000u32 { + interp.run(&mut mem_interp, 512); + aot.run(&mut mem_aot, 512); + assert_state_eq(&interp, &aot, entry, step); + + if interp.is_waiting { + break; + } + } + } + } + + let dispatched = riscv_core::aot::DISPATCH_RETIRED.load(std::sync::atomic::Ordering::Relaxed); + assert!( + dispatched > 0, + "AOT dispatch never fired — the comparison was interpreter vs interpreter" + ); + + eprintln!( + "[diff] {} programs ({terminated} reached the trap vector), {dispatched} insns retired via aot, lockstep state identical throughout", + entries.len() + ); +} diff --git a/crates/vpod-translate/Cargo.toml b/crates/vpod-translate/Cargo.toml new file mode 100644 index 0000000..ac7cce5 --- /dev/null +++ b/crates/vpod-translate/Cargo.toml @@ -0,0 +1,8 @@ +[package] +name = "vpod-translate" +version = "0.1.0" +edition.workspace = true + +[dependencies] +riscv-core = { path = "../riscv-core" } +lz4_flex = "0.13.1" diff --git a/crates/vpod-translate/src/main.rs b/crates/vpod-translate/src/main.rs new file mode 100644 index 0000000..8561cb8 --- /dev/null +++ b/crates/vpod-translate/src/main.rs @@ -0,0 +1,1197 @@ +// snapshot RAM + physical-address trace -> generated Rust + +use std::collections::BTreeSet; +use std::fmt::Write as _; +use std::io::Read; + +use riscv_core::block::{self, AluKind, Block, Op}; +use riscv_core::system_bus::SystemBus; + +const RAM_BASE: u64 = 0x8000_0000; + +struct SnapshotRam { + ram: Vec, + base: u64, +} + +impl SnapshotRam { + fn idx(&self, address: u64) -> Option { + let off = address.checked_sub(self.base)? as usize; + (off < self.ram.len()).then_some(off) + } + + fn contains(&self, address: u64, len: u64) -> bool { + address >= self.base && address + len <= self.base + self.ram.len() as u64 + } +} + +impl SystemBus for SnapshotRam { + fn read_byte(&mut self, address: u64) -> u8 { + self.idx(address).map_or(0, |i| self.ram[i]) + } + + fn read_halfword(&mut self, address: u64) -> u16 { + self.idx(address).map_or(0, |i| { + u16::from_le_bytes(self.ram[i..i + 2].try_into().unwrap()) + }) + } + + fn read_word(&mut self, address: u64) -> u32 { + self.idx(address).map_or(0, |i| { + u32::from_le_bytes(self.ram[i..i + 4].try_into().unwrap()) + }) + } + + fn read_doubleword(&mut self, address: u64) -> u64 { + self.idx(address).map_or(0, |i| { + u64::from_le_bytes(self.ram[i..i + 8].try_into().unwrap()) + }) + } + + fn write_byte(&mut self, _: u64, _: u8) {} + fn write_halfword(&mut self, _: u64, _: u16) {} + fn write_word(&mut self, _: u64, _: u32) {} + fn write_doubleword(&mut self, _: u64, _: u64) {} +} + +fn load_snapshot_ram(path: &str) -> Vec { + let raw = std::fs::read(path).unwrap_or_else(|e| { + eprintln!("cannot read snapshot {path}: {e}"); + std::process::exit(1); + }); + + let mut plain; + let bytes: &[u8] = if raw.starts_with(b"VPOD") { + &raw + } else { + plain = Vec::new(); + lz4_flex::frame::FrameDecoder::new(&raw[..]) + .read_to_end(&mut plain) + .unwrap_or_else(|e| { + eprintln!("snapshot is neither raw VPOD nor lz4: {e}"); + std::process::exit(1); + }); + &plain + }; + + assert_eq!(&bytes[0..4], b"VPOD", "bad snapshot magic"); + assert_eq!(bytes[4], 1, "unsupported snapshot version"); + + let ram_size = u64::from_le_bytes(bytes[6..14].try_into().unwrap()) as usize; + bytes[14..14 + ram_size].to_vec() +} + +fn alu_call(kind: AluKind, lhs: &str, rhs: &str) -> String { + format!("crate::block::alu(crate::block::AluKind::{kind:?}, {lhs}, {rhs})") +} + +struct RegAlloc { + declared: BTreeSet, + written: BTreeSet, +} + +impl RegAlloc { + fn new() -> Self { + Self { + declared: BTreeSet::new(), + written: BTreeSet::new(), + } + } + + fn read(&mut self, out: &mut String, r: u8) -> String { + if r == 0 { + return "0u64".to_string(); + } + if self.declared.insert(r) { + writeln!(out, " let mut x{r}: u64 = ctx.regs.read({r});").unwrap(); + } + format!("x{r}") + } + + fn write_target(&mut self, out: &mut String, r: u8) -> Option { + if r == 0 { + return None; + } + if self.declared.insert(r) { + writeln!(out, " let mut x{r}: u64;").unwrap(); + } + self.written.insert(r); + Some(format!("x{r}")) + } + + fn writeback(&self) -> String { + self.written + .iter() + .map(|&r| format!("ctx.regs.write({r}, x{r}); ")) + .collect() + } + + fn reload(&mut self) -> String { + self.written.clear(); + self.declared + .iter() + .map(|&r| format!("x{r} = ctx.regs.read({r}); ")) + .collect() + } +} + +fn load_size(kind: block::LoadKind) -> i64 { + use block::LoadKind::*; + match kind { + Lb | Lbu => 1, + Lh | Lhu => 2, + Lw | Lwu => 4, + Ld => 8, + } +} + +fn store_size(kind: block::StoreKind) -> i64 { + use block::StoreKind::*; + match kind { + Sb => 1, + Sh => 2, + Sw => 4, + Sd => 8, + } +} + +fn access_info(op: &Op) -> Option<(u8, i64, i64, bool)> { + match *op { + Op::Load { kind, rs1, imm, .. } => Some((rs1, imm, imm + load_size(kind) - 1, false)), + Op::Store { kind, rs1, imm, .. } => Some((rs1, imm, imm + store_size(kind) - 1, true)), + _ => None, + } +} + +fn detect_group(ops: &[block::DecodedInsn], start: usize) -> Option<(usize, i64, i64)> { + if std::env::var_os("VPOD_TRANSLATE_NO_BATCH").is_some() { + return None; + } + let (base, mut lo, mut hi, _) = access_info(&ops[start].op)?; + let mut len = 1usize; + + loop { + if let Op::Load { rd, .. } = ops[start + len - 1].op + && rd == base + && rd != 0 + { + break; + } + let Some(next) = ops.get(start + len) else { + break; + }; + let Some((next_base, next_lo, next_hi, _)) = access_info(&next.op) else { + break; + }; + + if next_base != base { + break; + } + + let new_lo = lo.min(next_lo); + let new_hi = hi.max(next_hi); + if new_hi - new_lo + 1 > 4096 { + break; + } + + lo = new_lo; + hi = new_hi; + len += 1; + } + + (len >= 2).then_some((len, lo, hi)) +} + +struct CachedGroup { + base: u8, + lo: i64, + hi: i64, + store_check: bool, + members: Vec, +} + +fn op_clobbers(op: &Op) -> Option { + match *op { + Op::Lui { rd, .. } + | Op::Auipc { rd, .. } + | Op::AluImm { rd, .. } + | Op::AluReg { rd, .. } + | Op::Load { rd, .. } => (rd != 0).then_some(rd), + _ => None, + } +} + +fn plan_cached_groups(ops: &[block::DecodedInsn]) -> Vec { + // only to test + if std::env::var_os("VPOD_TRANSLATE_NO_REGION").is_some() { + return Vec::new(); + } + + let mut in_consecutive_group = vec![false; ops.len()]; + let mut i = 0; + while i < ops.len() { + if let Some((len, _, _)) = detect_group(ops, i) { + in_consecutive_group[i..i + len].fill(true); + i += len; + } else { + i += 1; + } + } + + struct OpenGroup { + group: CachedGroup, + interior_store_seen: bool, + } + let mut open: std::collections::BTreeMap = std::collections::BTreeMap::new(); + let mut done: Vec = Vec::new(); + let close = |open: &mut std::collections::BTreeMap, + done: &mut Vec, + base: u8| { + if let Some(o) = open.remove(&base) + && o.group.members.len() >= 2 + { + done.push(o.group); + } + }; + + for (i, insn) in ops.iter().enumerate() { + match insn.op { + Op::Fallback { .. } | Op::Jal { .. } | Op::Jalr { .. } => { + let bases: Vec = open.keys().copied().collect(); + for b in bases { + close(&mut open, &mut done, b); + } + continue; + } + Op::Branch { .. } => continue, + _ => {} + } + + if let Some((base, lo, hi, is_store)) = access_info(&insn.op) { + if is_store { + for o in open.values_mut() { + o.interior_store_seen = true; + } + } + + if !in_consecutive_group[i] { + let fits = open + .get(&base) + .map(|o| o.group.hi.max(hi) - o.group.lo.min(lo) < 4096); + + if fits == Some(false) { + close(&mut open, &mut done, base); + } + + match open.get_mut(&base) { + Some(o) => { + o.group.lo = o.group.lo.min(lo); + o.group.hi = o.group.hi.max(hi); + o.group.members.push(i); + if is_store || o.interior_store_seen { + o.group.store_check = true; + } + } + None => { + open.insert( + base, + OpenGroup { + group: CachedGroup { + base, + lo, + hi, + store_check: is_store, + members: vec![i], + }, + interior_store_seen: false, + }, + ); + } + } + } + } + + if let Some(rd) = op_clobbers(&insn.op) { + close(&mut open, &mut done, rd); + } + } + + for (_, o) in open { + if o.group.members.len() >= 2 { + done.push(o.group); + } + } + + done +} + +#[allow(clippy::too_many_arguments)] +fn emit_cached_member( + out: &mut String, + regs: &mut RegAlloc, + group: &CachedGroup, + group_id: usize, + is_first: bool, + insn: &block::DecodedInsn, + pc: &str, + retired: u64, + flushed: u64, +) { + let base_str = regs.read(out, group.base); + + if is_first { + let helper = if group.store_check { + "_store_span_page" + } else { + "_load_span_page" + }; + writeln!( + out, + " let page_g{group_id} = crate::block::{helper}(ctx, satp, {base_str}.wrapping_add({}i64 as u64), {base_str}.wrapping_add({}i64 as u64));", + group.lo, group.hi + ) + .unwrap(); + } + + match insn.op { + Op::Load { kind, rd, imm, .. } => { + let fault_writeback = regs.writeback(); + let bind = match regs.write_target(out, rd) { + None => "Ok(_) => {}".to_string(), + Some(dst) => format!("Ok(v) => {dst} = v,"), + }; + writeln!( + out, + " let va = {base_str}.wrapping_add({imm}i64 as u64);\n match if let Some(p) = page_g{group_id} {{ Ok(crate::block::_span_load(ctx, p as *const u8, satp, crate::block::LoadKind::{kind:?}, va)) }} else {{ crate::block::do_load(ctx, satp, crate::block::LoadKind::{kind:?}, va, {pc}) }} {{\n {bind}\n Err(()) => {{ {fault_writeback}ctx.csr.instret = ctx.csr.instret.wrapping_add({}); return ({}, u64::MAX); }}\n }}", + retired - flushed, + retired + 1 + ) + .unwrap(); + } + Op::Store { kind, rs2, imm, .. } => { + let val = regs.read(out, rs2); + let fault_writeback = regs.writeback(); + writeln!( + out, + " let va = {base_str}.wrapping_add({imm}i64 as u64);\n let store_ok = if let Some(p) = page_g{group_id} {{ crate::block::_span_store(ctx, p, satp, crate::block::StoreKind::{kind:?}, va, {val}); true }} else {{ crate::block::do_store(ctx, satp, crate::block::StoreKind::{kind:?}, va, {val}, {pc}).is_ok() }};\n if !store_ok {{\n {fault_writeback}ctx.csr.instret = ctx.csr.instret.wrapping_add({}); return ({}, u64::MAX);\n }}", + retired - flushed, + retired + 1 + ) + .unwrap(); + } + _ => unreachable!("cached group members are loads/stores only"), + } +} + +#[allow(clippy::too_many_arguments)] +fn emit_group( + out: &mut String, + regs: &mut RegAlloc, + ops: &[block::DecodedInsn], + start: usize, + len: usize, + span_lo: i64, + span_hi: i64, + retired: u64, + flushed: u64, +) { + let members = &ops[start..start + len]; + let (base, ..) = access_info(&members[0].op).unwrap(); + let base_str = regs.read(out, base); + let has_store = members.iter().any(|m| matches!(m.op, Op::Store { .. })); + + for m in members { + match m.op { + Op::Load { rd, .. } => { + if rd != 0 { + regs.read(out, rd); + let _ = regs.write_target(out, rd); + } + } + Op::Store { rs2, .. } => { + regs.read(out, rs2); + } + _ => unreachable!("group members are loads/stores only"), + } + } + + writeln!( + out, + " let span_lo = {base_str}.wrapping_add({span_lo}i64 as u64);\n let span_hi = {base_str}.wrapping_add({span_hi}i64 as u64);" + ) + .unwrap(); + + let helper = if has_store { + "_store_span_page" + } else { + "_load_span_page" + }; + writeln!( + out, + " if let Some(span_page) = crate::block::{helper}(ctx, satp, span_lo, span_hi) {{" + ) + .unwrap(); + + for m in members { + match m.op { + Op::Load { kind, rd, imm, .. } => { + let dst = if rd == 0 { + "let _".to_string() + } else { + format!("x{rd}") + }; + let page = if has_store { + "span_page as *const u8" + } else { + "span_page" + }; + writeln!( + out, + " {dst} = crate::block::_span_load(ctx, {page}, satp, crate::block::LoadKind::{kind:?}, {base_str}.wrapping_add({imm}i64 as u64));" + ) + .unwrap(); + } + Op::Store { kind, rs2, imm, .. } => { + let val = if rs2 == 0 { + "0u64".to_string() + } else { + format!("x{rs2}") + }; + writeln!( + out, + " crate::block::_span_store(ctx, span_page, satp, crate::block::StoreKind::{kind:?}, {base_str}.wrapping_add({imm}i64 as u64), {val});" + ) + .unwrap(); + } + _ => unreachable!(), + } + } + + writeln!(out, " }} else {{").unwrap(); + + for (j, m) in members.iter().enumerate() { + let member_retired = retired + j as u64; + let pc = format!("entry_pc.wrapping_add({})", m.pc_off); + match m.op { + Op::Load { kind, rd, imm, .. } => { + let fault_writeback = regs.writeback(); + let bind = if rd == 0 { + "Ok(_) => {}".to_string() + } else { + format!("Ok(v) => x{rd} = v,") + }; + writeln!( + out, + " let va = {base_str}.wrapping_add({imm}i64 as u64);\n match crate::block::do_load(ctx, satp, crate::block::LoadKind::{kind:?}, va, {pc}) {{\n {bind}\n Err(()) => {{ {fault_writeback}ctx.csr.instret = ctx.csr.instret.wrapping_add({}); return ({}, u64::MAX); }}\n }}", + member_retired - flushed, + member_retired + 1 + ) + .unwrap(); + } + Op::Store { kind, rs2, imm, .. } => { + let val = if rs2 == 0 { + "0u64".to_string() + } else { + format!("x{rs2}") + }; + let fault_writeback = regs.writeback(); + writeln!( + out, + " let va = {base_str}.wrapping_add({imm}i64 as u64);\n if crate::block::do_store(ctx, satp, crate::block::StoreKind::{kind:?}, va, {val}, {pc}).is_err() {{\n {fault_writeback}ctx.csr.instret = ctx.csr.instret.wrapping_add({}); return ({}, u64::MAX);\n }}", + member_retired - flushed, + member_retired + 1 + ) + .unwrap(); + } + _ => unreachable!(), + } + } + + writeln!(out, " }}").unwrap(); +} + +fn emit_block(out: &mut String, pa: u64, entry_seen: &BTreeSet, blk: &Block) { + let _ = entry_seen; + writeln!( + out, + "#[allow(unused_variables, unused_mut, unused_assignments)]\nfn block_{pa:x}(ctx: &mut ExecContext, entry_pc: u64, satp: u64) -> (u64, u64) {{" + ) + .unwrap(); + + let mut regs = RegAlloc::new(); + + let page_off = pa & 0xfff; + let next_expr = |insn_off: u64, delta: i64| -> String { + let target = page_off as i64 + insn_off as i64 + delta; + if (0..4096).contains(&target) { + format!( + "entry_pc.wrapping_add({}u64)", + (insn_off as i64 + delta) as u64 + ) + } else { + "u64::MAX".to_string() + } + }; + + let mut retired: u64 = 0; + let mut flushed: u64 = 0; + + let cached_groups = plan_cached_groups(&blk.ops); + let mut cgroup_of: std::collections::BTreeMap = std::collections::BTreeMap::new(); + for (group_id, group) in cached_groups.iter().enumerate() { + for &member in &group.members { + cgroup_of.insert(member, group_id); + } + } + + let mut i = 0usize; + while i < blk.ops.len() { + let insn = &blk.ops[i]; + let off = insn.pc_off as u64; + let ilen = insn.ilen as u64; + let pc = format!("entry_pc.wrapping_add({off})"); + + if let Some((group_len, span_lo, span_hi)) = detect_group(&blk.ops, i) { + emit_group( + out, &mut regs, &blk.ops, i, group_len, span_lo, span_hi, retired, flushed, + ); + retired += group_len as u64; + i += group_len; + continue; + } + + if let Some(&group_id) = cgroup_of.get(&i) { + let group = &cached_groups[group_id]; + emit_cached_member( + out, + &mut regs, + group, + group_id, + group.members[0] == i, + insn, + &pc, + retired, + flushed, + ); + retired += 1; + i += 1; + continue; + } + + match insn.op { + Op::Lui { rd, imm } => { + if let Some(dst) = regs.write_target(out, rd) { + writeln!(out, " {dst} = {imm}i64 as u64;").unwrap(); + } + } + Op::Auipc { rd, imm } => { + if let Some(dst) = regs.write_target(out, rd) { + writeln!(out, " {dst} = {pc}.wrapping_add({imm}i64 as u64);").unwrap(); + } + } + Op::AluImm { kind, rd, rs1, imm } => { + if rd != 0 { + let a = regs.read(out, rs1); + let expr = alu_call(kind, &a, &format!("{imm}i64 as u64")); + let dst = regs.write_target(out, rd).unwrap(); + writeln!(out, " {dst} = {expr};").unwrap(); + } + } + Op::AluReg { kind, rd, rs1, rs2 } => { + if rd != 0 { + let a = regs.read(out, rs1); + let b = regs.read(out, rs2); + let expr = alu_call(kind, &a, &b); + let dst = regs.write_target(out, rd).unwrap(); + writeln!(out, " {dst} = {expr};").unwrap(); + } + } + Op::Load { kind, rd, rs1, imm } => { + let base = regs.read(out, rs1); + let fault_writeback = regs.writeback(); + let bind = match regs.write_target(out, rd) { + None => "Ok(_) => {}".to_string(), + Some(dst) => format!("Ok(v) => {dst} = v,"), + }; + writeln!( + out, + " let va = {base}.wrapping_add({imm}i64 as u64);\n match crate::block::do_load(ctx, satp, crate::block::LoadKind::{kind:?}, va, {pc}) {{\n {bind}\n Err(()) => {{ {fault_writeback}ctx.csr.instret = ctx.csr.instret.wrapping_add({}); return ({}, u64::MAX); }}\n }}", + retired - flushed, + retired + 1 + ) + .unwrap(); + } + Op::Store { + kind, + rs1, + rs2, + imm, + } => { + let base = regs.read(out, rs1); + let val = regs.read(out, rs2); + let fault_writeback = regs.writeback(); + writeln!( + out, + " let va = {base}.wrapping_add({imm}i64 as u64);\n if crate::block::do_store(ctx, satp, crate::block::StoreKind::{kind:?}, va, {val}, {pc}).is_err() {{\n {fault_writeback}ctx.csr.instret = ctx.csr.instret.wrapping_add({}); return ({}, u64::MAX);\n }}", + retired - flushed, + retired + 1 + ) + .unwrap(); + } + Op::Branch { + kind, + rs1, + rs2, + offset, + } => { + let cond = { + let a = regs.read(out, rs1); + let b = regs.read(out, rs2); + match kind { + block::BranchKind::Beq => format!("{a} == {b}"), + block::BranchKind::Bne => format!("{a} != {b}"), + block::BranchKind::Blt => format!("({a} as i64) < ({b} as i64)"), + block::BranchKind::Bge => format!("({a} as i64) >= ({b} as i64)"), + block::BranchKind::Bltu => format!("{a} < {b}"), + block::BranchKind::Bgeu => format!("{a} >= {b}"), + } + }; + let writeback = regs.writeback(); + let taken_next = next_expr(off, offset); + writeln!( + out, + " let taken = {cond};\n if taken {{\n {writeback}ctx.regs.pc = {pc}.wrapping_add({offset}i64 as u64);\n ctx.csr.instret = ctx.csr.instret.wrapping_add({});\n return ({}, {taken_next});\n }}", + retired + 1 - flushed, + retired + 1 + ) + .unwrap(); + retired += 1; + } + Op::Jal { rd, offset } => { + let next = next_expr(off, offset); + if let Some(dst) = regs.write_target(out, rd) { + writeln!(out, " {dst} = {pc}.wrapping_add({ilen});").unwrap(); + } + writeln!( + out, + " {}ctx.regs.pc = {pc}.wrapping_add({offset}i64 as u64);\n ctx.csr.instret = ctx.csr.instret.wrapping_add({});\n return ({}, {next});", + regs.writeback(), + retired + 1 - flushed, + retired + 1 + ) + .unwrap(); + + retired += 1; + } + Op::Jalr { rd, rs1, imm } => { + let base = regs.read(out, rs1); + writeln!( + out, + " let target = {base}.wrapping_add({imm}i64 as u64) & !1;" + ) + .unwrap(); + if let Some(dst) = regs.write_target(out, rd) { + writeln!(out, " {dst} = {pc}.wrapping_add({ilen});").unwrap(); + } + + writeln!( + out, + " {}ctx.regs.pc = target;\n ctx.csr.instret = ctx.csr.instret.wrapping_add({});\n return ({}, if target >> 12 == entry_pc >> 12 {{ target }} else {{ u64::MAX }});", + regs.writeback(), + retired + 1 - flushed, + retired + 1 + ) + .unwrap(); + + retired += 1; + } + Op::Fallback { raw } => { + writeln!( + out, + " {}ctx.regs.pc = {pc};\n ctx.csr.instret = ctx.csr.instret.wrapping_add({});\n let r = crate::execute::exec_raw(ctx, {raw}u32, {pc});\n debug_assert!(matches!(r, crate::trap::StepResult::Ok));\n if ctx.regs.pc != {pc}.wrapping_add({ilen}) {{\n return ({}, u64::MAX);\n }}\n {}let satp = crate::block::effective_satp(*ctx.priv_mode, ctx.csr.satp);", + regs.writeback(), + retired - flushed, + retired + 1, + regs.reload() + ) + .unwrap(); + + flushed = retired + 1; + retired += 1; + } + } + + if !matches!( + insn.op, + Op::Branch { .. } | Op::Jal { .. } | Op::Jalr { .. } | Op::Fallback { .. } + ) { + retired += 1; + } + i += 1; + } + + let fall_next = next_expr(blk.byte_len as u64, 0); + writeln!( + out, + " {}ctx.regs.pc = entry_pc.wrapping_add({});\n ctx.csr.instret = ctx.csr.instret.wrapping_add({});\n ({retired}, {fall_next})\n}}\n", + regs.writeback(), + blk.byte_len, + retired - flushed + ) + .unwrap(); +} + +fn translate_set(bus: &mut SnapshotRam, pas: &BTreeSet, hot: &BTreeSet, out_path: &str) { + let mut out = String::new(); + out.push_str( + "// Generated by vpod-translate. Do not edit.\nuse crate::execute::ExecContext;\nuse crate::system_bus::SystemBus;\n\n", + ); + + let mut entries: Vec = Vec::new(); + let mut pages: BTreeSet = BTreeSet::new(); + let mut skipped = 0usize; + + for &pa in pas { + if !bus.contains(pa, 4) { + skipped += 1; + continue; + } + match block::decode_block(bus, pa) { + Some(blk) => { + emit_block(&mut out, pa, pas, &blk); + entries.push(pa); + pages.insert(pa >> 12); + } + None => skipped += 1, + } + } + + let mut by_page: std::collections::BTreeMap, Vec)> = + std::collections::BTreeMap::new(); + + for &pa in &entries { + let (hot_pas, cold_pas) = by_page.entry(pa >> 12).or_default(); + if hot.contains(&pa) { + hot_pas.push(pa); + } else { + cold_pas.push(pa); + } + } + + let mut dispatch = String::new(); + for (page, (_, cold_pas)) in &by_page { + if cold_pas.is_empty() { + continue; + } + writeln!( + dispatch, + "#[inline(never)]\nfn page_{page:x}(ctx: &mut ExecContext, pa: u64, pc: u64, satp: u64) -> Option<(u64, u64)> {{\n Some(match pa {{" + ) + .unwrap(); + for pa in cold_pas { + writeln!(dispatch, " 0x{pa:x} => block_{pa:x}(ctx, pc, satp),").unwrap(); + } + dispatch.push_str(" _ => return None,\n })\n}\n\n"); + } + + dispatch.push_str(concat!( + "#[inline(never)]\npub fn dispatch(ctx: &mut ExecContext, pa_in: u64, entry_pc: u64, satp: u64, fuel: u64, rt_page: u64) -> Option {\n", + " let mut pa = pa_in;\n let mut pc = entry_pc;\n let mut total = 0u64;\n", + " let mut satp = satp;\n let mut rt_page = rt_page;\n", + " let mut chain_generation = ctx.blocks.aot_evict_generation();\n loop {\n", + " let step = match pa >> 12 {\n", + )); + + for (page, (hot_pas, cold_pas)) in &by_page { + let cold_arm = if cold_pas.is_empty() { + "None".to_string() + } else { + format!("page_{page:x}(ctx, pa, pc, satp)") + }; + if hot_pas.is_empty() { + writeln!(dispatch, " 0x{page:x} => {cold_arm},").unwrap(); + continue; + } + writeln!(dispatch, " 0x{page:x} => match pa {{").unwrap(); + for pa in hot_pas { + writeln!( + dispatch, + " 0x{pa:x} => Some(block_{pa:x}(ctx, pc, satp))," + ) + .unwrap(); + } + writeln!( + dispatch, + " _ => {cold_arm},\n }}," + ) + .unwrap(); + } + + dispatch.push_str(concat!( + " _ => None,\n };\n", + " let (retired, next) = match step {\n", + " Some(v) => v,\n", + " None => break,\n", + " };\n", + " total += retired;\n", + " if total >= fuel {\n return Some(total);\n }\n", + " if next == u64::MAX {\n", + " if ctx.csr.pending_interrupt(*ctx.priv_mode).is_some() {\n", + " return Some(total);\n", + " }\n", + " satp = crate::block::effective_satp(*ctx.priv_mode, ctx.csr.satp);\n", + " pc = ctx.regs.pc;\n", + " let vpage = pc >> 12;\n", + " let fetch_pa = if let Some(ppage) = ctx.fetch_tlb.lookup(vpage, satp) {\n", + " crate::perf::note_fetch_page_hit();\n", + " debug_assert_eq!(\n", + " ctx.mmu.translate_fetch(pc, satp, ctx.bus).map(|pa| pa >> 12),\n", + " Ok(ppage),\n", + " \"fetch TLB hit disagrees with slow-path translation\"\n", + " );\n", + " (ppage << 12) | (pc & 0xfff)\n", + " } else {\n", + " crate::perf::note_fetch_translate();\n", + " match ctx.mmu.translate_fetch(pc, satp, ctx.bus) {\n", + " Ok(fetch_pa) => {\n", + " ctx.fetch_tlb.insert(vpage, fetch_pa >> 12, satp);\n", + " fetch_pa\n", + " }\n", + " Err(_) => return Some(total),\n", + " }\n", + " };\n", + " match crate::execute::aot_page_key(ctx, fetch_pa) {\n", + " Some(key_pa) => {\n", + " pa = key_pa;\n", + " rt_page = fetch_pa >> 12;\n", + " chain_generation = ctx.blocks.aot_evict_generation();\n", + " continue;\n", + " }\n", + " None => return Some(total),\n", + " }\n", + " }\n", + " pa = (pa & !0xfffu64) | (next & 0xfff);\n", + " if ctx.blocks.aot_evict_generation() != chain_generation {\n", + " match crate::execute::aot_page_key(ctx, rt_page << 12) {\n", + " Some(k) if k >> 12 == pa >> 12 => {}\n", + " _ => return Some(total),\n", + " }\n", + " }\n", + " pc = next;\n }\n", + " if total == 0 { None } else { Some(total) }\n}\n\n", + )); + + let mut pages_str = String::from("pub const AOT_PAGE_HASHES: &[(u64, u64)] = &[\n"); + for &page in &pages { + let mut hash = 0xcbf2_9ce4_8422_2325u64; + for i in 0..512u64 { + hash ^= bus.read_doubleword((page << 12) + i * 8); + hash = hash.wrapping_mul(0x0000_0100_0000_01b3); + } + writeln!(pages_str, " (0x{hash:x}, 0x{page:x}),").unwrap(); + } + + pages_str.push_str("];\n"); + + out.push_str(&dispatch); + out.push_str(&pages_str); + + std::fs::write(out_path, out).unwrap_or_else(|e| { + eprintln!("cannot write {out_path}: {e}"); + std::process::exit(1); + }); + + eprintln!( + "[vpod-translate] {} blocks on {} pages ({} pcs skipped) -> {}", + entries.len(), + pages.len(), + skipped, + out_path + ); +} + +struct Rng(u64); + +impl Rng { + fn next(&mut self) -> u64 { + let mut x = self.0; + x ^= x << 13; + x ^= x >> 7; + x ^= x << 17; + self.0 = x; + x + } + + fn below(&mut self, n: u64) -> u64 { + self.next() % n + } +} + +fn r_type(op: u32, f3: u32, f7: u32, rd: u32, rs1: u32, rs2: u32) -> u32 { + op | (rd << 7) | (f3 << 12) | (rs1 << 15) | (rs2 << 20) | (f7 << 25) +} + +fn i_type(op: u32, f3: u32, rd: u32, rs1: u32, imm: u32) -> u32 { + op | (rd << 7) | (f3 << 12) | (rs1 << 15) | ((imm & 0xfff) << 20) +} + +fn s_type(op: u32, f3: u32, rs1: u32, rs2: u32, imm: u32) -> u32 { + op | ((imm & 0x1f) << 7) | (f3 << 12) | (rs1 << 15) | (rs2 << 20) | (((imm >> 5) & 0x7f) << 25) +} + +fn b_type(f3: u32, rs1: u32, rs2: u32, imm: u32) -> u32 { + 0x63 | (((imm >> 11) & 1) << 7) + | (((imm >> 1) & 0xf) << 8) + | (f3 << 12) + | (rs1 << 15) + | (rs2 << 20) + | (((imm >> 5) & 0x3f) << 25) + | (((imm >> 12) & 1) << 31) +} + +fn rand_rd(rng: &mut Rng) -> u32 { + let r = 1 + rng.below(15) as u32; + if r == 10 { 11 } else { r } +} + +const PROGRAM_INSNS: usize = 48; +const PAGE: usize = 4096; + +fn gen_program(rng: &mut Rng, data_page: u32, code_page: u32) -> Vec { + let mut insns: Vec = Vec::new(); + + insns.push(0x37 | (10 << 7) | (data_page << 12)); + insns.push(0x37 | (11 << 7) | (code_page << 12)); + + while insns.len() < PROGRAM_INSNS { + let remaining = PROGRAM_INSNS - insns.len(); + let insn = match rng.below(12) { + // alu imm + 0..=2 => { + let f3 = [0u32, 2, 3, 4, 6, 7][rng.below(6) as usize]; + i_type(0x13, f3, rand_rd(rng), rand_rd(rng), rng.next() as u32) + } + // shifts + 3 => { + let (f3, top) = [(1u32, 0u32), (5, 0), (5, 0x400)][rng.below(3) as usize]; + let shamt = (rng.below(64) as u32) | top; + i_type(0x13, f3, rand_rd(rng), rand_rd(rng), shamt) + } + // alu reg + 4..=5 => { + let (f3, f7) = [ + (0u32, 0u32), + (0, 0x20), + (1, 0), + (2, 0), + (3, 0), + (4, 0), + (5, 0), + (5, 0x20), + (6, 0), + (7, 0), + (0, 1), + (4, 1), + (5, 1), + (6, 1), + (7, 1), + // Zba: sh1add / sh2add / sh3add + (2, 0x10), + (4, 0x10), + (6, 0x10), + ][rng.below(18) as usize]; + r_type(0x33, f3, f7, rand_rd(rng), rand_rd(rng), rand_rd(rng)) + } + // lui / auipc + 6 => { + let op = if rng.below(2) == 0 { 0x37 } else { 0x17 }; + op | (rand_rd(rng) << 7) | ((rng.next() as u32) & 0xfffff000) + } + // load from data page + 7 => { + let f3 = [0u32, 1, 2, 3, 4, 5, 6][rng.below(7) as usize]; + i_type(0x03, f3, rand_rd(rng), 10, (rng.next() as u32) & 0x7f8) + } + // store to data page + 8 => { + let f3 = [0u32, 1, 2, 3][rng.below(4) as usize]; + s_type(0x23, f3, 10, rand_rd(rng), (rng.next() as u32) & 0x7f8) + } + // forward branch + 9 => { + let f3 = [0u32, 1, 4, 5, 6, 7][rng.below(6) as usize]; + let max_skip = remaining.saturating_sub(1).clamp(1, 4); + + let offset = (4 + 4 * rng.below(max_skip as u64) as u32) & 0x1ffe; + b_type(f3, rand_rd(rng), rand_rd(rng), offset) + } + // Zba + 10 => match rng.below(3) { + 0 => { + let f3 = [0u32, 2, 4, 6][rng.below(4) as usize]; + let f7 = if f3 == 0 { 0x04 } else { 0x10 }; + r_type(0x3b, f3, f7, rand_rd(rng), rand_rd(rng), rand_rd(rng)) + } + 1 => { + // slli.uw carries funct6 = 0x02 above its 6-bit shamt. + let shamt = rng.below(64) as u32; + i_type(0x1b, 1, rand_rd(rng), rand_rd(rng), (0x02 << 6) | shamt) + } + _ => { + let rs2 = if rng.below(4) == 0 { rand_rd(rng) } else { 0 }; + r_type(0x3b, 4, 0x04, rand_rd(rng), rand_rd(rng), rs2) + } + }, + // self-modifying store + _ => { + let target_insn = + (insns.len() as u64 + 1 + rng.below((PROGRAM_INSNS - insns.len()) as u64)) + .min(PROGRAM_INSNS as u64 - 1); + s_type(0x23, 2, 11, rand_rd(rng), (target_insn as u32) * 4) + } + }; + insns.push(insn); + } + + insns.push(0x0000_0073); + insns +} + +fn run_gen(args: &[String]) { + if args.len() != 6 { + eprintln!( + "usage: vpod-translate gen " + ); + std::process::exit(1); + } + let dir = &args[2]; + let num_programs: usize = args[3].parse().expect("bad num-programs"); + let seed: u64 = args[4].parse().expect("bad seed"); + let mut rng = Rng(seed.max(1)); + + let ram_len = ((2 * num_programs + 1) * PAGE).next_power_of_two(); + let mut ram = vec![0u8; ram_len]; + let trap_pa = (ram_len - PAGE) as u64; + ram[trap_pa as usize..trap_pa as usize + 4].copy_from_slice(&0x1050_0073u32.to_le_bytes()); + + let mut entry_pcs: Vec = Vec::new(); + let mut pas: BTreeSet = BTreeSet::new(); + + for i in 0..num_programs { + let code_base = 2 * i * PAGE; + let data_page = (2 * i + 1) as u32; + let insns = gen_program(&mut rng, data_page, (2 * i) as u32); + + for (j, insn) in insns.iter().enumerate() { + ram[code_base + 4 * j..code_base + 4 * j + 4].copy_from_slice(&insn.to_le_bytes()); + pas.insert((code_base + 4 * j) as u64); + } + entry_pcs.push(code_base as u64); + } + + std::fs::create_dir_all(dir).expect("cannot create out dir"); + std::fs::write(format!("{dir}/ram.bin"), &ram).expect("cannot write ram.bin"); + let entries_txt: String = entry_pcs.iter().map(|pc| format!("{pc:x}\n")).collect(); + std::fs::write(format!("{dir}/entries.txt"), entries_txt).expect("cannot write entries.txt"); + + let hot: BTreeSet = pas.iter().copied().step_by(2).collect(); + + let mut bus = SnapshotRam { ram, base: 0 }; + translate_set(&mut bus, &pas, &hot, &args[5]); + + eprintln!( + "[vpod-translate] gen: {num_programs} programs (seed {seed}), ram {} KiB, trap vector at 0x{trap_pa:x} -> {dir}", + ram_len / 1024 + ); +} + +fn main() { + let args: Vec = std::env::args().collect(); + + if args.len() >= 2 && args[1] == "gen" { + run_gen(&args); + return; + } + + if args.len() < 4 { + eprintln!( + "usage: vpod-translate [--max-blocks N] [--hot-blocks N] [--coverage PCT]" + ); + std::process::exit(1); + } + + let mut max_blocks: usize = 16384; + let mut hot_blocks: usize = 4096; + let mut coverage: f64 = 100.0; + let mut i = 4; + while i < args.len() { + match args[i].as_str() { + "--max-blocks" => { + max_blocks = args[i + 1].parse().expect("bad --max-blocks"); + i += 2; + } + "--hot-blocks" => { + hot_blocks = args[i + 1].parse().expect("bad --hot-blocks"); + i += 2; + } + "--coverage" => { + coverage = args[i + 1].parse().expect("bad --coverage"); + i += 2; + } + other => { + eprintln!("unknown argument: {other}"); + std::process::exit(1); + } + } + } + + let ram = load_snapshot_ram(&args[1]); + let trace = std::fs::read_to_string(&args[2]).unwrap_or_else(|e| { + eprintln!("cannot read trace {}: {e}", args[2]); + std::process::exit(1); + }); + + let mut bus = SnapshotRam { + ram, + base: RAM_BASE, + }; + + let mut counted: Vec<(u64, u64)> = trace + .lines() + .filter_map(|l| { + let mut parts = l.split_whitespace(); + let pa = u64::from_str_radix(parts.next()?.trim_start_matches("0x"), 16).ok()?; + let n = parts.next().and_then(|c| c.parse().ok()).unwrap_or(1u64); + Some((pa, n)) + }) + .collect(); + + counted.sort_by_key(|b| std::cmp::Reverse(b.1)); + + let total: u64 = counted.iter().map(|&(_, n)| n).sum(); + let target = (total as f64 * coverage / 100.0) as u64; + let mut cumulative = 0u64; + let mut pas: BTreeSet = BTreeSet::new(); + let mut hot: BTreeSet = BTreeSet::new(); + + for &(pa, n) in counted.iter().take(max_blocks) { + if cumulative >= target { + break; + } + cumulative += n; + pas.insert(pa); + if hot.len() < hot_blocks { + hot.insert(pa); + } + } + + eprintln!( + "[vpod-translate] selected {} of {} traced blocks ({:.2}% of {} block executions)", + pas.len(), + counted.len(), + cumulative as f64 / total.max(1) as f64 * 100.0, + total + ); + + translate_set(&mut bus, &pas, &hot, &args[3]); +} diff --git a/crates/vpod/src/start.rs b/crates/vpod/src/start.rs index be28ab3..4270360 100644 --- a/crates/vpod/src/start.rs +++ b/crates/vpod/src/start.rs @@ -306,6 +306,7 @@ pub fn run(cfg: RunConfig) -> Result<()> { let mut builder = WasiCtxBuilder::new(); builder.inherit_stdin().inherit_stdout().inherit_stderr(); builder.args(&wasm_args); + builder.preopened_dir(&snap_dir, "snap", DirPerms::READ, FilePerms::READ)?; for (i, mount) in cfg.mounts.iter().enumerate() { diff --git a/crates/wasi-component/Cargo.toml b/crates/wasi-component/Cargo.toml index bd46bf9..1927784 100644 --- a/crates/wasi-component/Cargo.toml +++ b/crates/wasi-component/Cargo.toml @@ -14,7 +14,7 @@ path = "src/main.rs" [dependencies] machine = { path = "../machine" } -riscv-core = { path = "../riscv-core" } +riscv-core = { path = "../riscv-core", features = ["aot"] } log = "0.4" wasi = "0.14.7" lz4_flex = "0.13.1" diff --git a/crates/wasi-component/src/api/session.rs b/crates/wasi-component/src/api/session.rs index 2ab9825..0f7160d 100644 --- a/crates/wasi-component/src/api/session.rs +++ b/crates/wasi-component/src/api/session.rs @@ -12,6 +12,27 @@ use riscv_core::Hart; const PYRUNNER_SENTINEL: &str = "---VPOD_DONE---"; +const AOT_MISMATCH_PROBE_THRESHOLD: u64 = 64; + +fn warn_if_aot_mismatch(hart: &Hart) { + static WARNED: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false); + + if !hart.blocks.aot_enabled() { + return; + } + let (probes, matches) = hart.blocks.aot_match_stats(); + if probes >= AOT_MISMATCH_PROBE_THRESHOLD + && matches == 0 + && !WARNED.swap(true, std::sync::atomic::Ordering::Relaxed) + { + eprintln!( + "[vpod] warning: the bundled AOT module does not match this snapshot; \ + running at interpreter speed. Upgrade the vpod package or re-pull \ + the snapshot so the two agree." + ); + } +} + pub struct Session { pub bus: MachineBus, pub hart: Hart, @@ -150,6 +171,8 @@ impl SessionManager { bus.uart.drain_tx(); } + warn_if_aot_mismatch(&hart); + let id = self.next_id.get(); self.next_id.set(id + 1); self.sessions.borrow_mut().insert( @@ -223,7 +246,7 @@ impl SessionManager { .trim_end() .to_string(); - let ctrl_bytes = session.bus.uart_ctrl.drain_tx(); + let ctrl_bytes = repl::drain_ctrl_with_grace(&mut session.bus, &mut session.hart); let exit_code = match ctrl_bytes.first() { Some(byte) => *byte as u32, None => { @@ -265,7 +288,7 @@ impl SessionManager { let mut timed_out = false; let exit_code = if session.is_shell { - let ctrl_bytes = session.bus.uart_ctrl.drain_tx(); + let ctrl_bytes = repl::drain_ctrl_with_grace(&mut session.bus, &mut session.hart); match ctrl_bytes.first() { Some(byte) => *byte as u32, None => { @@ -381,6 +404,8 @@ impl SessionManager { (true, false, false, b"# ".to_vec()) }; + warn_if_aot_mismatch(&hart); + let id = self.next_id.get(); self.next_id.set(id + 1); self.sessions.borrow_mut().insert( @@ -401,7 +426,7 @@ impl SessionManager { } fn restart_pyrunner(session: &mut Session) { - let restart = b"pkill -9 -f pyrunner.py; python3 /usr/lib/vpod/pyrunner.py &\n"; + let restart = b"pkill -9 -f pyrunner.py; PYR=/usr/bin/python3.real; [ -x $PYR ] || PYR=python3; $PYR /usr/lib/vpod/pyrunner.py &\n"; for byte in restart { session.bus.uart.push_rx(*byte); } diff --git a/crates/wasi-component/src/repl.rs b/crates/wasi-component/src/repl.rs index 2e1810c..48c154d 100644 --- a/crates/wasi-component/src/repl.rs +++ b/crates/wasi-component/src/repl.rs @@ -5,7 +5,9 @@ use wasi::clocks::wall_clock; use wasi::io::poll; const STEP: u64 = 8192; +const RUN_STEP: u64 = 524_288; +const MAX_TIMER_WARP_NS: u64 = 100_000_000; const NET_YIELD_NS: u64 = 5_000_000; // 5 ms // TO TEST : the time UART must be quiet after last output before declare the command @@ -69,6 +71,28 @@ pub fn settle(bus: &mut MachineBus, hart: &mut Hart, wall_ns: u64) { } } +pub fn drain_ctrl_with_grace(bus: &mut MachineBus, hart: &mut Hart) -> Vec { + for _ in 0..2000u32 { + let bytes = bus.uart_ctrl.drain_tx(); + if !bytes.is_empty() { + return bytes; + } + + if hart.is_waiting { + hart.is_waiting = false; + } + + bus.clint.advance_by_instructions(STEP); + bus.poll(hart); + + if let StepResult::Trap(_) | StepResult::Halt = hart.run(bus, STEP) { + break; + } + } + + bus.uart_ctrl.drain_tx() +} + pub fn wait_for_prompt(bus: &mut MachineBus, hart: &mut Hart, prompt: &[u8]) { let mut buffer = Vec::new(); @@ -121,29 +145,35 @@ pub fn capture_output( hart.is_waiting = false; if !bus.has_pending_io() { - let before = monotonic_clock::now(); - let timeout = monotonic_clock::subscribe_duration(NET_YIELD_NS); - poll::poll(&[&timeout]); - - let idle_ns = monotonic_clock::now().saturating_sub(before); - bus.clint.advance_by_nanos(idle_ns); - - if !data_channel - && sentinel.is_none() - && got_output - && !bus.net_rx_pending() - && !bus.net_has_active_connections() - && monotonic_clock::now().saturating_sub(last_output_ns) >= QUIET_PERIOD_NS - { - break; + if matches!(bus.clint.nanos_until_timer(), Some(ns) if ns <= MAX_TIMER_WARP_NS) { + bus.clint.fast_forward_to_timer(); + bus.poll(hart); + } else { + let before = monotonic_clock::now(); + let timeout = monotonic_clock::subscribe_duration(NET_YIELD_NS); + poll::poll(&[&timeout]); + + let idle_ns = monotonic_clock::now().saturating_sub(before); + bus.clint.advance_by_nanos(idle_ns); + + if !data_channel + && sentinel.is_none() + && got_output + && !bus.net_rx_pending() + && !bus.net_has_active_connections() + && monotonic_clock::now().saturating_sub(last_output_ns) >= QUIET_PERIOD_NS + { + break; + } } } } - bus.clint.advance_by_instructions(STEP); + let step = if bus.net_rx_pending() { STEP } else { RUN_STEP }; + bus.clint.advance_by_instructions(step); bus.poll(hart); - match hart.run(bus, STEP) { + match hart.run_until_wait(bus, step) { StepResult::Ok => {} StepResult::Trap(_) | StepResult::Halt => { if stop_on_ctrl && !bus.uart_ctrl.tx_is_empty() { diff --git a/crates/wasi-component/src/run_interactive.rs b/crates/wasi-component/src/run_interactive.rs index 350f740..cc1f429 100644 --- a/crates/wasi-component/src/run_interactive.rs +++ b/crates/wasi-component/src/run_interactive.rs @@ -10,7 +10,7 @@ const OUTPUT_HOLD_CYCLES: u32 = 3; pub fn run(bus: &mut MachineBus, hart: &mut Hart) { let stdin = wasi::cli::stdin::get_stdin(); - const POLL_INTERVAL_ACTIVE: u64 = 32768; + const POLL_INTERVAL_ACTIVE: u64 = 131_072; const POLL_INTERVAL_IDLE: u64 = 8192; const POLL_INTERVAL_NET: u64 = 4096; const IDLE_TIMEOUT_NS: u64 = 50_000_000; @@ -21,6 +21,8 @@ pub fn run(bus: &mut MachineBus, hart: &mut Hart) { let mut hold_cycles = 0u32; let mut active_ticks = 0u32; + let stdin_pollable = stdin.subscribe(); + loop { let interval = if bus.net_rx_pending() { POLL_INTERVAL_NET @@ -33,7 +35,7 @@ pub fn run(bus: &mut MachineBus, hart: &mut Hart) { bus.clint.advance_by_instructions(interval); bus.poll(hart); - if poll_stdin(bus, &stdin) { + if poll_stdin(bus, &stdin, &stdin_pollable) { bus.poll(hart); idle_ticks = 0; active_ticks = 512; @@ -123,8 +125,7 @@ fn flush_pending(pending: &mut Vec) { pending.clear(); } -fn poll_stdin(bus: &mut MachineBus, stdin: &InputStream) -> bool { - let pollable = stdin.subscribe(); +fn poll_stdin(bus: &mut MachineBus, stdin: &InputStream, pollable: &poll::Pollable) -> bool { if !pollable.ready() { return false; } diff --git a/crates/wasi-component/src/vm.rs b/crates/wasi-component/src/vm.rs index 01f4a8c..d24b066 100644 --- a/crates/wasi-component/src/vm.rs +++ b/crates/wasi-component/src/vm.rs @@ -4,7 +4,7 @@ use machine::machine_bus::MachineBus; use machine::snapshot; use machine::virtio::fs::Mount; use riscv_core::Hart; -use std::io::{BufReader, Read}; +use std::io::BufReader; use std::path::{Path, PathBuf}; #[derive(Clone)] @@ -77,29 +77,25 @@ enum Compression { Raw, } -fn detect_compression(path: &Path) -> Result { - let mut file = - std::fs::File::open(path).map_err(|e| format!("failed to open {:?}: {e}", path))?; - let mut magic = [0u8; 4]; - let n = file - .read(&mut magic) - .map_err(|e| format!("failed to read file magic: {e}"))?; - if n >= 4 && magic == [0x04, 0x22, 0x4D, 0x18] { - Ok(Compression::Lz4) +fn detect_compression_bytes(data: &[u8]) -> Compression { + if data.len() >= 4 && data[..4] == [0x04, 0x22, 0x4D, 0x18] { + Compression::Lz4 } else { - Ok(Compression::Raw) + Compression::Raw } } pub fn _read_base_and_tail(path: &Path) -> Result<(CowRam, Vec, u8), String> { - let file = std::fs::File::open(path) - .map_err(|e| format!("failed to open snapshot {:?}: {e}", path))?; + let data = + std::fs::read(path).map_err(|e| format!("failed to read snapshot {:?}: {e}", path))?; + let compression = detect_compression_bytes(&data); + let cursor = std::io::Cursor::new(data); - match detect_compression(path)? { + match compression { Compression::Lz4 => { - snapshot::load_base_and_tail(&mut BufReader::new(FrameDecoder::new(file))) + snapshot::load_base_and_tail(&mut BufReader::new(FrameDecoder::new(cursor))) } - Compression::Raw => snapshot::load_base_and_tail(&mut BufReader::new(file)), + Compression::Raw => snapshot::load_base_and_tail(&mut BufReader::new(cursor)), } .map_err(|e| format!("failed to load snapshot base: {e}")) } @@ -113,7 +109,9 @@ pub fn _bus_from_base( let mut bus = MachineBus::new(ram_size, base.clone_shared()); bus.attach_net(); bus.attach_fs(vec![]); - let hart = Hart::new(0x1000); + + let mut hart = Hart::new(0x1000); + hart.blocks.aot_init(riscv_core::aot::AOT_PAGE_HASHES); for (i, m) in mounts.iter().enumerate() { if let Some(fs) = bus.fs_devices.get_mut(i) { @@ -140,6 +138,7 @@ pub fn _load(config: _VmConfig) -> Result<(MachineBus, Hart, u8), String> { bus.attach_net(); bus.attach_fs(vec![]); let mut hart = Hart::new(0x1000); + hart.blocks.aot_init(riscv_core::aot::AOT_PAGE_HASHES); if let Some(disk_path) = config.disk { let file = std::fs::OpenOptions::new() @@ -152,17 +151,19 @@ pub fn _load(config: _VmConfig) -> Result<(MachineBus, Hart, u8), String> { .map_err(|e| format!("failed to attach disk: {e}"))?; } - let snapshot_file = std::fs::File::open(config.snapshot) - .map_err(|e| format!("failed to open snapshot {:?}: {e}", config.snapshot))?; + let snapshot_data = std::fs::read(config.snapshot) + .map_err(|e| format!("failed to read snapshot {:?}: {e}", config.snapshot))?; + let compression = detect_compression_bytes(&snapshot_data); + let snapshot_cursor = std::io::Cursor::new(snapshot_data); - let flags = match detect_compression(config.snapshot)? { + let flags = match compression { Compression::Lz4 => snapshot::restore( &mut bus, &mut hart, - &mut BufReader::new(FrameDecoder::new(snapshot_file)), + &mut BufReader::new(FrameDecoder::new(snapshot_cursor)), ), Compression::Raw => { - snapshot::restore(&mut bus, &mut hart, &mut BufReader::new(snapshot_file)) + snapshot::restore(&mut bus, &mut hart, &mut BufReader::new(snapshot_cursor)) } } .map_err(|e| format!("failed to restore snapshot: {e}"))?; diff --git a/guest/tls/vpod_ssl_client.c b/guest/tls/vpod_ssl_client.c new file mode 100644 index 0000000..3da83eb --- /dev/null +++ b/guest/tls/vpod_ssl_client.c @@ -0,0 +1,138 @@ +/* + * vpod ssl_client, replacement for busybox's ssl_client that does no TLS. + */ + +#include +#include +#include +#include +#include +#include +#include +#include + +#define REAL_SSL_CLIENT "/usr/bin/ssl_client.real" + +static void run_real_ssl_client(char **argv) { + execv(REAL_SSL_CLIENT, argv); + fprintf(stderr, "vpod ssl_client: exec %s: %s\n", REAL_SSL_CLIENT, strerror(errno)); + + exit(1); +} + +static int write_all(int fd, const char *buf, size_t len) { + while (len > 0) { + ssize_t n = write(fd, buf, len); + if (n < 0) { + if (errno == EINTR) + continue; + return -1; + } + + buf += n; + len -= (size_t)n; + } + + return 0; +} + +static int splice_ready(int src, int dst, int *open_flag) { + char buf[16384]; + ssize_t n = read(src, buf, sizeof(buf)); + + if (n < 0) + return (errno == EINTR || errno == EAGAIN) ? 0 : -1; + + if (n == 0) { + *open_flag = 0; + return 0; + } + + return write_all(dst, buf, (size_t)n); +} + +int main(int argc, char **argv) { + int net_fd = -1; + const char *sni = NULL; + + for (int i = 1; i < argc; i++) { + if ((strcmp(argv[i], "-s") == 0 || strcmp(argv[i], "-h") == 0) && i + 1 < argc) { + net_fd = atoi(argv[++i]); + } else if ((strncmp(argv[i], "-s", 2) == 0 || strncmp(argv[i], "-h", 2) == 0) && argv[i][2] != '\0') { + net_fd = atoi(argv[i] + 2); + } else if (strcmp(argv[i], "-n") == 0 && i + 1 < argc) { + sni = argv[++i]; + } else if (strncmp(argv[i], "-n", 2) == 0 && argv[i][2] != '\0') { + sni = argv[i] + 2; + } + /* other busybox options */ + } + + if (net_fd < 0 || sni == NULL) { + run_real_ssl_client(argv); + } + + struct sockaddr_in peer; + socklen_t peer_len = sizeof(peer); + unsigned port = 443; + + if (getpeername(net_fd, (struct sockaddr *)&peer, &peer_len) == 0 && + peer.sin_family == AF_INET) { + port = ntohs(peer.sin_port); + } + + if (port != 443) { + run_real_ssl_client(argv); + } + + char preamble[300]; + int len = snprintf(preamble, sizeof(preamble), "VPOD-CONNECT %s %u\n", sni, port); + if (len <= 0 || (size_t)len >= sizeof(preamble) || + write_all(net_fd, preamble, (size_t)len) != 0) { + fprintf(stderr, "vpod ssl_client: preamble write failed\n"); + return 1; + } + + int stdin_open = 1, net_open = 1; + while (net_open) { + struct pollfd fds[2]; + nfds_t nfds = 0; + int stdin_slot = -1, net_slot = -1; + + if (stdin_open) { + stdin_slot = (int)nfds; + fds[nfds].fd = 0; + fds[nfds].events = POLLIN; + nfds++; + } + + net_slot = (int)nfds; + fds[nfds].fd = net_fd; + fds[nfds].events = POLLIN; + nfds++; + + if (poll(fds, nfds, -1) < 0) { + if (errno == EINTR) + continue; + + return 1; + } + + if (stdin_slot >= 0 && (fds[stdin_slot].revents & (POLLIN | POLLHUP))) { + if (splice_ready(0, net_fd, &stdin_open) != 0) + return 1; + + if (!stdin_open) + shutdown(net_fd, SHUT_WR); + } + if (fds[net_slot].revents & (POLLIN | POLLHUP)) { + if (splice_ready(net_fd, 1, &net_open) != 0) + return 1; + } + + if (fds[net_slot].revents & (POLLERR | POLLNVAL)) + return 1; + } + + return 0; +} diff --git a/guest/warmpy/pydaemon.py b/guest/warmpy/pydaemon.py new file mode 100644 index 0000000..95248c3 --- /dev/null +++ b/guest/warmpy/pydaemon.py @@ -0,0 +1,385 @@ +"""vpod warm-python prefork server. (for python3 shell cmd) + +Protocol: see guest/warmpy/vpod_python_shim.c. +""" + +import array +import importlib +import io +import os +import runpy +import signal +import socket +import struct +import sys +import traceback +import warnings + +SOCK_PATH = os.environ.get("VPOD_PYD_SOCK", "/run/vpod-pyd.sock") + +CHILD_EXECUTABLE = os.environ.get("VPOD_PYD_EXECUTABLE", "/usr/bin/python3.real") +MAGIC = b"VPY1" + +for _mod in ( + "abc", "base64", "codecs", "collections", "encodings.utf_8", + "functools", "json", "re", "shutil", "subprocess", "types", +): + try: + __import__(_mod) + except ImportError: + pass + + +WARM_IMPORTS_FILE = os.environ.get( + "VPOD_PYD_WARM_IMPORTS", "/etc/vpod/pydaemon-warm-imports" +) + + +def warm_heavy_tools(): + """Pre-import expensive module trees so forked children inherit them via + copy-on-write.""" + modules = [] + try: + with open(WARM_IMPORTS_FILE) as f: + for line in f: + name = line.split("#", 1)[0].strip() + if name: + modules.append(name) + except OSError: + pass + for name in modules: + try: + __import__(name) + except Exception: + pass + + +warm_heavy_tools() + + +def _watch_dirs(): + """The directories a runtime `pip install` writes into, plus the warm-list + file itself (so appending a module to it takes effect live).""" + dirs = {WARM_IMPORTS_FILE} + try: + import site + + for getter in ("getsitepackages", "getusersitepackages"): + fn = getattr(site, getter, None) + if not fn: + continue + try: + result = fn() + except Exception: + continue + dirs.update([result] if isinstance(result, str) else result) + + except Exception: + pass + for entry in sys.path: + if entry and "site-packages" in entry: + dirs.add(entry) + return sorted(dirs) + + +def _dir_signature(dirs): + signature = [] + for path in dirs: + try: + signature.append((path, os.stat(path).st_mtime_ns)) + except OSError: + signature.append((path, -1)) + return tuple(signature) + + +_WATCH_DIRS = _watch_dirs() +_last_signature = _dir_signature(_WATCH_DIRS) + + +def refresh_import_caches(): + """Called in the daemon (once per accept).""" + global _last_signature + signature = _dir_signature(_WATCH_DIRS) + if signature != _last_signature: + importlib.invalidate_caches() + _last_signature = signature + warm_heavy_tools() + + +def recv_request(conn): + """returns ([stdin_fd, stdout_fd, stderr_fd], payload).""" + fds = array.array("i") + header = b"" + + while len(header) < 8: + data, ancdata, _flags, _addr = conn.recvmsg( + 8 - len(header), socket.CMSG_SPACE(3 * 4) + ) + if not data: + raise ConnectionError("client closed during header") + header += data + for level, ctype, cdata in ancdata: + if level == socket.SOL_SOCKET and ctype == socket.SCM_RIGHTS: + fds.frombytes(cdata[: len(cdata) - len(cdata) % 4]) + + if header[:4] != MAGIC: + raise ValueError(f"bad magic: {header[:4]!r}") + + if len(fds) != 3: + raise ValueError(f"expected 3 fds, got {len(fds)}") + + (payload_len,) = struct.unpack("= len(args): + return None + value = args[i] + if flag == "c": + req.mode = "c" + req.target = value + req.args = ["-c"] + args[i + 1 :] + return req + if flag == "m": + req.mode = "m" + req.target = value + req.args = [value] + args[i + 1 :] + return req + if flag == "W": + req.warn_options.append(value) + else: + pass + break + elif flag == "u": + req.unbuffered = True + elif flag == "E": + req.ignore_env = True + elif flag in ("b", "B", "q", "s"): + pass + else: + return None + + j += 1 + i += 1 + req.mode = "stdin" + return req + + +def wire_stdio(fds, unbuffered): + for i, fd in enumerate(fds): + if fd != i: + os.dup2(fd, i) + for fd in set(fds): + if fd > 2: + os.close(fd) + + stdin_raw = io.FileIO(0, "rb", closefd=False) + sys.stdin = sys.__stdin__ = io.TextIOWrapper( + io.BufferedReader(stdin_raw), encoding="utf-8", errors="strict" + ) + + def make_writer(fd): + raw = io.FileIO(fd, "wb", closefd=False) + if unbuffered: + return io.TextIOWrapper(raw, encoding="utf-8", write_through=True) + return io.TextIOWrapper( + io.BufferedWriter(raw), encoding="utf-8", + line_buffering=os.isatty(fd), + ) + + sys.stdout = sys.__stdout__ = make_writer(1) + sys.stderr = sys.__stderr__ = make_writer(2) + + +def run_child(fds, req, argv0, cwd, env): + exit_code = 0 + try: + signal.signal(signal.SIGINT, signal.default_int_handler) + for sig in (signal.SIGTERM, signal.SIGHUP, signal.SIGQUIT, + signal.SIGPIPE, signal.SIGUSR1, signal.SIGUSR2): + signal.signal(sig, signal.SIG_DFL) + + os.chdir(cwd) + os.environ.clear() + os.environ.update(env) + + unbuffered = req.unbuffered or ( + not req.ignore_env and bool(env.get("PYTHONUNBUFFERED")) + ) + wire_stdio(fds, unbuffered) + + sys.executable = CHILD_EXECUTABLE + sys.argv = req.args or [argv0] + + if not req.ignore_env: + for path in reversed(env.get("PYTHONPATH", "").split(":")): + if path and path not in sys.path: + sys.path.insert(0, path) + for warn_option in req.warn_options: + warnings._processoptions([warn_option]) + + if req.mode == "c": + sys.path.insert(0, "") + exec(compile(req.target, "", "exec"), + {"__name__": "__main__", "__doc__": None, "__package__": None, + "__spec__": None, "__builtins__": __builtins__}) + elif req.mode == "m": + sys.path.insert(0, os.getcwd()) + runpy.run_module(req.target, run_name="__main__", alter_sys=True) + elif req.mode == "script": + runpy.run_path(req.target, run_name="__main__") + else: + source = sys.stdin.read() + sys.path.insert(0, "") + exec(compile(source, "", "exec"), + {"__name__": "__main__", "__doc__": None, "__package__": None, + "__spec__": None, "__builtins__": __builtins__}) + + except SystemExit as exc: + if exc.code is None: + exit_code = 0 + elif isinstance(exc.code, int): + exit_code = exc.code + else: + print(exc.code, file=sys.stderr) + exit_code = 1 + except BaseException: + traceback.print_exc() + exit_code = 1 + finally: + try: + sys.stdout.flush() + except Exception: + pass + try: + sys.stderr.flush() + except Exception: + pass + os._exit(exit_code & 0xFF) + + +def handle_connection(conn): + fds, payload = recv_request(conn) + argv, cwd, env = parse_payload(payload) + req = parse_argv(argv) + + if req is None or (req.mode == "stdin" and os.isatty(fds[0])): + conn.sendall(b"F") + return + + child = os.fork() + if child == 0: + conn.close() + run_child(fds, req, argv[0] if argv else "python3", cwd, env) + + for fd in set(fds): + os.close(fd) + conn.sendall(b"P" + struct.pack(" +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +extern char **environ; + +static char **g_argv; + +static const char *real_python(void) { + const char *p = getenv("VPOD_PYTHON_REAL"); + return (p && *p) ? p : "/usr/bin/python3.real"; +} + +static const char *sock_path(void) { + const char *p = getenv("VPOD_PYD_SOCK"); + return (p && *p) ? p : "/run/vpod-pyd.sock"; +} + +static void fallback(void) { + signal(SIGPIPE, SIG_DFL); + execv(real_python(), g_argv); + perror("vpod-python-shim: exec python3.real"); + _exit(127); +} + +static void venv_passthrough(void) { + const char *execfn = (const char *)getauxval(AT_EXECFN); + if (!execfn || !*execfn) return; + + char cfg[4096]; + if (strlen(execfn) + sizeof("/../pyvenv.cfg") > sizeof(cfg)) return; + strcpy(cfg, execfn); + + char *slash = strrchr(cfg, '/'); + if (!slash) return; + strcpy(slash, "/../pyvenv.cfg"); + + if (access(cfg, F_OK) == 0) { + g_argv[0] = (char *)execfn; + fallback(); + } +} + +static volatile pid_t child_pid = 0; + +static void forward_signal(int sig) { + pid_t pid = child_pid; + if (pid > 0) { + kill(pid, sig); + } +} + +static void install_forwarders(void) { + static const int sigs[] = {SIGINT, SIGTERM, SIGHUP, SIGQUIT, SIGUSR1, SIGUSR2}; + struct sigaction sa; + memset(&sa, 0, sizeof(sa)); + + sa.sa_handler = forward_signal; + sa.sa_flags = SA_RESTART; + sigemptyset(&sa.sa_mask); + + for (size_t i = 0; i < sizeof(sigs) / sizeof(sigs[0]); i++) { + sigaction(sigs[i], &sa, NULL); + } +} + +static int full_read(int fd, void *buf, size_t n) { + char *p = buf; + while (n > 0) { + ssize_t r = read(fd, p, n); + if (r < 0) { + if (errno == EINTR) continue; + return -1; + } + if (r == 0) return -1; + p += r; + n -= (size_t)r; + } + return 0; +} + +static int full_write(int fd, const void *buf, size_t n) { + const char *p = buf; + while (n > 0) { + ssize_t r = write(fd, p, n); + if (r < 0) { + if (errno == EINTR) continue; + return -1; + } + p += r; + n -= (size_t)r; + } + return 0; +} + +static char *build_payload(int argc, char **argv, uint32_t *out_len) { + char argc_str[16]; + int argc_len = snprintf(argc_str, sizeof(argc_str), "%d", argc); + + char cwd[4096]; + if (!getcwd(cwd, sizeof(cwd))) { + strcpy(cwd, "/"); + } + + size_t len = (size_t)argc_len + 1; + for (int i = 0; i < argc; i++) len += strlen(argv[i]) + 1; + len += strlen(cwd) + 1; + for (char **e = environ; *e; e++) len += strlen(*e) + 1; + + char *buf = malloc(len); + if (!buf) return NULL; + + char *p = buf; + memcpy(p, argc_str, (size_t)argc_len + 1); + p += argc_len + 1; + for (int i = 0; i < argc; i++) { + size_t l = strlen(argv[i]) + 1; + memcpy(p, argv[i], l); + p += l; + } + + size_t l = strlen(cwd) + 1; + memcpy(p, cwd, l); + p += l; + for (char **e = environ; *e; e++) { + l = strlen(*e) + 1; + memcpy(p, *e, l); + p += l; + } + + *out_len = (uint32_t)len; + return buf; +} + +static int send_request(int fd, int argc, char **argv) { + uint32_t payload_len; + char *payload = build_payload(argc, argv, &payload_len); + if (!payload) return -1; + + unsigned char header[8] = {'V', 'P', 'Y', '1'}; + header[4] = (unsigned char)(payload_len & 0xff); + header[5] = (unsigned char)((payload_len >> 8) & 0xff); + header[6] = (unsigned char)((payload_len >> 16) & 0xff); + header[7] = (unsigned char)((payload_len >> 24) & 0xff); + + int fds[3] = {0, 1, 2}; + union { + char buf[CMSG_SPACE(sizeof(fds))]; + struct cmsghdr align; + } cmsg_storage; + memset(&cmsg_storage, 0, sizeof(cmsg_storage)); + + struct iovec iov = {.iov_base = header, .iov_len = sizeof(header)}; + struct msghdr msg; + memset(&msg, 0, sizeof(msg)); + msg.msg_iov = &iov; + msg.msg_iovlen = 1; + msg.msg_control = cmsg_storage.buf; + msg.msg_controllen = sizeof(cmsg_storage.buf); + + struct cmsghdr *cmsg = CMSG_FIRSTHDR(&msg); + cmsg->cmsg_level = SOL_SOCKET; + cmsg->cmsg_type = SCM_RIGHTS; + cmsg->cmsg_len = CMSG_LEN(sizeof(fds)); + memcpy(CMSG_DATA(cmsg), fds, sizeof(fds)); + + ssize_t sent; + do { + sent = sendmsg(fd, &msg, 0); + } while (sent < 0 && errno == EINTR); + if (sent != (ssize_t)sizeof(header)) { + free(payload); + return -1; + } + + int rc = full_write(fd, payload, payload_len); + free(payload); + return rc; +} + +int main(int argc, char **argv) { + g_argv = argv; + + venv_passthrough(); + + int fd = socket(AF_UNIX, SOCK_STREAM, 0); + if (fd < 0) fallback(); + + struct sockaddr_un addr; + memset(&addr, 0, sizeof(addr)); + addr.sun_family = AF_UNIX; + const char *path = sock_path(); + if (strlen(path) >= sizeof(addr.sun_path)) fallback(); + strcpy(addr.sun_path, path); + + if (connect(fd, (struct sockaddr *)&addr, sizeof(addr)) < 0) fallback(); + + signal(SIGPIPE, SIG_IGN); + if (send_request(fd, argc, argv) < 0) fallback(); + + int started = 0; + for (;;) { + unsigned char tag; + if (full_read(fd, &tag, 1) < 0) { + if (!started) fallback(); + + fprintf(stderr, "vpod-python-shim: server connection lost\n"); + return 1; + } + if (tag == 'F') { + fallback(); + } else if (tag == 'P') { + unsigned char b[4]; + if (full_read(fd, b, 4) < 0) return 1; + child_pid = (pid_t)(b[0] | (b[1] << 8) | ((uint32_t)b[2] << 16) | + ((uint32_t)b[3] << 24)); + started = 1; + install_forwarders(); + } else if (tag == 'X') { + unsigned char b[4]; + if (full_read(fd, b, 4) < 0) return 1; + uint32_t status = (uint32_t)b[0] | ((uint32_t)b[1] << 8) | + ((uint32_t)b[2] << 16) | ((uint32_t)b[3] << 24); + + if (WIFEXITED(status)) return WEXITSTATUS(status); + if (WIFSIGNALED(status)) { + int sig = WTERMSIG(status); + signal(sig, SIG_DFL); + raise(sig); + return 128 + sig; + } + return 1; + } else { + if (!started) fallback(); + return 1; + } + } +} diff --git a/scripts/aot-snapshot.sh b/scripts/aot-snapshot.sh new file mode 100755 index 0000000..a24d45f --- /dev/null +++ b/scripts/aot-snapshot.sh @@ -0,0 +1,119 @@ +#!/bin/sh +set -e + +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" +ROOT="$(cd "$SCRIPT_DIR/.." && pwd)" + +SNAP="" +WORKLOAD="default" +MAX_BLOCKS="" +COVERAGE="" +FORCE=0 + +while [ $# -gt 0 ]; do + case "$1" in + --workload) WORKLOAD="$2"; shift 2 ;; + --max-blocks) MAX_BLOCKS="$2"; shift 2 ;; + --coverage) COVERAGE="$2"; shift 2 ;; + --force) FORCE=1; shift ;; + -*) echo "unknown arg: $1" >&2; exit 1 ;; + *) SNAP="$1"; shift ;; + esac +done + +if [ -z "$SNAP" ]; then + echo "usage: $0 [--workload default|data] [--max-blocks N] [--coverage PCT] [--force]" >&2 + exit 1 +fi +if [ ! -f "$SNAP" ]; then + echo "error: snapshot not found: $SNAP" >&2 + exit 1 +fi +case "$WORKLOAD" in + default|data) ;; + *) echo "error: unknown workload '$WORKLOAD' (expected: default, data)" >&2; exit 1 ;; +esac + +GENERATED="$ROOT/crates/riscv-core/src/aot/generated.rs" +VPOD="$ROOT/target/release/vpod-native" +AOT_TRACE="$ROOT/dist/.aot-trace.txt" + +if [ -f "$GENERATED" ] && [ "$FORCE" = "0" ]; then + LINES=$(wc -l < "$GENERATED" | tr -d ' ') + if [ "$LINES" -gt 100 ]; then + echo "note: overwriting existing translation ($LINES lines) in 3s — Ctrl-C to abort," >&2 + echo " or back it up first; it is gitignored and unrecoverable." >&2 + sleep 3 + fi +fi + +echo "=== AOT translation pass ===" +echo "Snapshot : $SNAP" +echo "Workload : $WORKLOAD" +echo "" + +echo "── AOT: tracing representative workload on the snapshot..." +(cd "$ROOT" && cargo build --release -p native-cli --features aot-trace) + + +set -- \ + --setup "python3 -c 'print(sum(i*i for i in range(200000)))'" \ + --setup "python3 -c 'exec(\"s=0\nfor i in range(200000): s=(s+i*i)^(i&0xff)\nprint(s)\")'" + +case "$WORKLOAD" in + default) + set -- "$@" \ + --setup "python3 -c 'import json,os; print(json.dumps({\"cwd\": os.getcwd()}))'" + ;; + data) + set -- "$@" \ + --setup "python3 -c 'import numpy as np; a = np.arange(100000); print(int((a * a).sum()))'" \ + --setup "python3 -c 'import pandas as pd; df = pd.DataFrame({\"x\": range(20000)}); print(int(df.x.sum()))'" + ;; +esac + +set -- "$@" \ + --setup "i=0; while [ \$i -lt 100 ]; do echo x > /tmp/aot-\$i; i=\$((\$i+1)); done; cat /tmp/aot-* | wc -l; rm -f /tmp/aot-*" \ + --setup "uv venv /tmp/aot-venv && rm -rf /tmp/aot-venv" + + +set -- "$@" \ + --setup "apk update && apk add jq && echo '{\"a\":[1,2,3]}' | jq -c '.a | add' && echo VPOD_AOT_APK_OK" + +TRACE_LOG="$ROOT/dist/.aot-trace-run.log" +VPOD_AOT_TRACE="$AOT_TRACE" "$VPOD" --snapshot-load "$SNAP" --net "$@" 2>&1 | tee "$TRACE_LOG" + +if [ ! -s "$AOT_TRACE" ]; then + echo "error: aot trace is empty — the workload did not run" >&2 + exit 1 +fi + +if ! grep -q VPOD_AOT_APK_OK "$TRACE_LOG"; then + echo "" >&2 + echo "error: the apk trace step did not complete — apk's code would be left" >&2 + echo " untranslated (it is ~1.6B guest insns, 98% emulated CPU)." >&2 + echo " Check network/DNS from the guest; see $TRACE_LOG" >&2 + exit 1 +fi +rm -f "$TRACE_LOG" + +echo "── AOT: translating hot blocks..." +(cd "$ROOT" && cargo build --release -p vpod-translate) + +TRANSLATE_ARGS="" +[ -n "$MAX_BLOCKS" ] && TRANSLATE_ARGS="$TRANSLATE_ARGS --max-blocks $MAX_BLOCKS" +[ -n "$COVERAGE" ] && TRANSLATE_ARGS="$TRANSLATE_ARGS --coverage $COVERAGE" + +# shellcheck disable=SC2086 +"$ROOT/target/release/vpod-translate" $TRANSLATE_ARGS "$SNAP" "$AOT_TRACE" "$GENERATED" + +echo "── AOT: rebuilding vpod with translated blocks..." +(cd "$ROOT" && cargo build --release -p native-cli --features aot) +rm -f "$AOT_TRACE" + +echo "" +echo "=== Done ===" +echo "" +echo "Translation: $GENERATED" +echo "Native vpod rebuilt with --features aot." +echo "For the wasm component, run: ./scripts/build-wasm.sh" diff --git a/scripts/aot-stub.sh b/scripts/aot-stub.sh new file mode 100755 index 0000000..b5200bd --- /dev/null +++ b/scripts/aot-stub.sh @@ -0,0 +1,50 @@ +#!/bin/sh +# +# Writes a no-op crates/riscv-core/src/aot/generated.rs. +# +# riscv-core's `aot` feature is hardwired on by wasi-component, so anything +# building the component needs generated.rs to exist. A real one is produced by +# vpod-translate from a snapshot trace and is gitignored, so CI (and a fresh +# clone) has none. This stub satisfies the API: dispatch always declines and +# every block falls through to the interpreter. +# +# Refuses to clobber a real translation unless --force, because generated.rs +# cannot be restored from git and needs the full trace pipeline to rebuild. + +set -e + +ROOT="$(cd "$(dirname "$0")/.." && pwd)" +OUT="$ROOT/crates/riscv-core/src/aot/generated.rs" + +if [ -f "$OUT" ] && [ "$1" != "--force" ]; then + LINES=$(wc -l < "$OUT" | tr -d ' ') + if [ "$LINES" -gt 100 ]; then + echo "refusing to overwrite $OUT ($LINES lines — looks like a real translation)" >&2 + echo "it is gitignored and unrecoverable; re-run with --force if you mean it." >&2 + exit 1 + fi +fi + +mkdir -p "$(dirname "$OUT")" +cat > "$OUT" <<'EOF' +// AOT stub (scripts/aot-stub.sh). No translated blocks: dispatch always +// declines, so execution falls through to the interpreter. Replaced by +// vpod-translate output in a real snapshot build. +use crate::execute::ExecContext; +use crate::system_bus::SystemBus; + +pub fn dispatch( + _ctx: &mut ExecContext, + _pa_in: u64, + _entry_pc: u64, + _satp: u64, + _fuel: u64, + _rt_page: u64, +) -> Option { + None +} + +pub const AOT_PAGE_HASHES: &[(u64, u64)] = &[]; +EOF + +echo "wrote AOT stub: $OUT" diff --git a/scripts/build-data-snapshot.sh b/scripts/build-data-snapshot.sh new file mode 100755 index 0000000..71a6419 --- /dev/null +++ b/scripts/build-data-snapshot.sh @@ -0,0 +1,388 @@ +#!/bin/sh + +set -e + +cleanup() { + rm -rf "$ROOT/dist/agent-minirootfs" "$ROOT/dist/agent-mini.cpio.gz" "$ROOT/dist/agent-overlay.cpio.gz" +} +trap cleanup EXIT + +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" +ROOT="$(cd "$SCRIPT_DIR/.." && pwd)" + +ALPINE_VERSION="3.23.0" +OUT="$ROOT/dist/rootfs.cpio.gz" +RAM_MB=512 + +while [ $# -gt 0 ]; do + case "$1" in + --version) ALPINE_VERSION="$2"; shift 2 ;; + --out) OUT="$2"; shift 2 ;; + --ram) RAM_MB="$2"; shift 2 ;; + *) echo "unknown arg: $1"; exit 1 ;; + esac +done + +ALPINE_MINOR="${ALPINE_VERSION%.*}" +ALPINE_DIR="$ROOT/dist/alpine-standard-${ALPINE_VERSION}-riscv64" +ISO_URL="https://dl-cdn.alpinelinux.org/alpine/v${ALPINE_MINOR}/releases/riscv64/alpine-standard-${ALPINE_VERSION}-riscv64.iso" +MINIROOTFS_URL="https://dl-cdn.alpinelinux.org/alpine/v${ALPINE_MINOR}/releases/riscv64/alpine-minirootfs-${ALPINE_VERSION}-riscv64.tar.gz" +ISO="$ROOT/dist/alpine-standard-${ALPINE_VERSION}-riscv64.iso" +MINIROOTFS="$ROOT/dist/alpine-minirootfs-${ALPINE_VERSION}-riscv64.tar.gz" +KERNEL="$ALPINE_DIR/kernel" +INITRAMFS_LTS="$ALPINE_DIR/initramfs-lts" +OPENSBI_VERSION="1.6" +OPENSBI_FW="$ROOT/dist/fw_jump.bin" +OPENSBI_URL="https://github.com/riscv-software-src/opensbi/releases/download/v${OPENSBI_VERSION}/opensbi-${OPENSBI_VERSION}-rv-bin.tar.xz" +OPENSBI_TAR="$ROOT/dist/opensbi-${OPENSBI_VERSION}-rv-bin.tar.xz" +OVERLAY="$ROOT/dist/agent-overlay" +VPOD="$ROOT/target/release/vpod-native" + +echo "=== Capsulev snapshot builder ===" +echo "Alpine : ${ALPINE_VERSION}" +echo "RAM : ${RAM_MB} MB" +echo "Out : ${OUT}" +echo "" + +echo "── Checking host tools..." +MISSING="" +for cmd in curl bsdtar cpio gzip cargo zig; do + command -v "$cmd" >/dev/null || MISSING="$MISSING $cmd" +done +if [ -n "$MISSING" ]; then + echo "ERROR: missing tools:$MISSING" + echo " macOS : brew install libarchive zig" + echo " Debian : apt install libarchive-tools bsdtar cpio zig" + echo " Fedora : dnf install bsdtar libarchive zig" + echo " Windows: use WSL2 and follow the Linux instructions" + exit 1 +fi +echo " OK" + +mkdir -p "$ROOT/dist" "$ALPINE_DIR" + + +echo "── Building vpod..." +(cd "$ROOT" && cargo build --release -p native-cli --bin vpod-native) + + +if [ ! -f "$OPENSBI_FW" ]; then + echo "── Downloading OpenSBI ${OPENSBI_VERSION} pre-built firmware..." + curl -L --progress-bar -o "$OPENSBI_TAR" "$OPENSBI_URL" + bsdtar -xf "$OPENSBI_TAR" -C "$ROOT/dist" \ + "opensbi-${OPENSBI_VERSION}-rv-bin/share/opensbi/lp64/generic/firmware/fw_jump.bin" + mv "$ROOT/dist/opensbi-${OPENSBI_VERSION}-rv-bin/share/opensbi/lp64/generic/firmware/fw_jump.bin" \ + "$OPENSBI_FW" + rm -rf "$OPENSBI_TAR" "$ROOT/dist/opensbi-${OPENSBI_VERSION}-rv-bin" + echo " OpenSBI firmware: $(du -sh "$OPENSBI_FW" | cut -f1)" +else + echo "── OpenSBI firmware already present, skipping." +fi + +if [ ! -f "$ISO" ]; then + echo "── Downloading Alpine standard ISO ${ALPINE_VERSION}..." + curl -L --progress-bar -o "$ISO" "$ISO_URL" +else + echo "── Alpine ISO already present, skipping download." +fi + +if [ ! -f "$KERNEL" ] || [ ! -f "$INITRAMFS_LTS" ]; then + echo "── Extracting kernel and initramfs-lts from ISO..." + bsdtar -xf "$ISO" -C "$ALPINE_DIR" \ + --include "boot/vmlinuz-lts" \ + --include "boot/initramfs-lts" \ + --strip-components=1 + + MAGIC=$(dd if="$ALPINE_DIR/vmlinuz-lts" bs=2 count=1 2>/dev/null | od -A n -t x1 | tr -d ' \n') + if [ "$MAGIC" = "1f8b" ]; then + gzip -dc "$ALPINE_DIR/vmlinuz-lts" > "$KERNEL" + else + cp "$ALPINE_DIR/vmlinuz-lts" "$KERNEL" + fi + echo " kernel : $(du -sh "$KERNEL" | cut -f1)" + echo " initramfs : $(du -sh "$INITRAMFS_LTS" | cut -f1)" +else + echo "── Kernel and initramfs already extracted, skipping." +fi + +if [ ! -f "$MINIROOTFS" ]; then + echo "── Downloading Alpine minirootfs ${ALPINE_VERSION}..." + curl -L --progress-bar -o "$MINIROOTFS" "$MINIROOTFS_URL" +else + echo "── Minirootfs already present, skipping download." +fi + + +echo "── Building overlay..." +rm -rf "$OVERLAY" +mkdir -p "$OVERLAY/sbin" "$OVERLAY/etc/apk" "$OVERLAY/usr/lib/vpod" + +printf 'https://dl-cdn.alpinelinux.org/alpine/v%s/main\nhttps://dl-cdn.alpinelinux.org/alpine/v%s/community\n' \ + "$ALPINE_MINOR" "$ALPINE_MINOR" > "$OVERLAY/etc/apk/repositories" +echo 'nameserver 8.8.8.8' > "$OVERLAY/etc/resolv.conf" + +mkdir -p "$OVERLAY/usr/local/share/ca-certificates" +cp "$ROOT/crates/machine/assets/tls/vpod-ca-cert.pem" \ + "$OVERLAY/usr/local/share/ca-certificates/vpod-ca.crt" + +mkdir -p "$OVERLAY/etc/ssl/vpod" +cp "$ROOT/crates/machine/assets/tls/vpod-ca-cert.pem" \ + "$OVERLAY/etc/ssl/vpod/ca-only.pem" + + +echo "── Cross-compiling vpod ssl_client (riscv64-musl, static)..." +zig cc -target riscv64-linux-musl -Os -static -s \ + -o "$OVERLAY/usr/lib/vpod/vpod-ssl-client" \ + "$ROOT/guest/tls/vpod_ssl_client.c" +chmod +x "$OVERLAY/usr/lib/vpod/vpod-ssl-client" + +mkdir -p "$OVERLAY/etc/vpod" +cat > "$OVERLAY/etc/vpod/pydaemon-warm-imports" << 'WARM_EOF' +#Warm set for the python3 shim/daemon path (commands.run("python3 ...") +#numpy +#pandas +WARM_EOF +cat > "$OVERLAY/etc/vpod/pyrunner-warm-imports" << 'WARM_EOF' +# Warm set for pyrunner, the persistent code.run() interpreter — the primary +numpy +pandas +scipy +scipy.stats +scipy.optimize +scipy.interpolate +scipy.fft +WARM_EOF + +mkdir -p "$OVERLAY/etc/uv" +cat > "$OVERLAY/etc/uv/uv.toml" << 'UV_EOF' +python-preference = "only-system" +UV_EOF + +echo "── Cross-compiling vpod python shim (riscv64-musl, dynamic)..." +zig cc -target riscv64-linux-musl -Os -dynamic -s \ + -o "$OVERLAY/usr/lib/vpod/vpod-python-shim" \ + "$ROOT/guest/warmpy/vpod_python_shim.c" +chmod +x "$OVERLAY/usr/lib/vpod/vpod-python-shim" +cp "$ROOT/guest/warmpy/pydaemon.py" "$OVERLAY/usr/lib/vpod/pydaemon.py" + +cat > "$OVERLAY/sbin/init" << 'INIT_EOF' +#!/bin/sh + +export PATH=/usr/bin:/usr/sbin:/bin:/sbin + +mount -t proc proc /proc +mount -t sysfs sysfs /sys +mount -t devtmpfs devtmpfs /dev +mount -t tmpfs tmpfs /tmp + +hostname vpod +ip link set lo up 2>/dev/null || true + +modprobe virtio_mmio 2>/dev/null || true +modprobe virtio_net 2>/dev/null || true +modprobe virtio_blk 2>/dev/null || true +modprobe virtiofs 2>/dev/null || true + +ip link set eth0 up 2>/dev/null || true +ip addr add 10.0.2.15/24 dev eth0 2>/dev/null || true +ip route add default via 10.0.2.2 2>/dev/null || true +echo "nameserver 10.0.2.2" > /etc/resolv.conf + +if [ -c /dev/hvc0 ]; then + ( + echo "VPOD_READY" >/dev/hvc0 + while IFS= read -r cmd /dev/hvc0 2>&1 + printf 'VPOD_EXIT:%d\n' "$?" >/dev/hvc0 + done + ) & +fi + +export TERM=dumb +export HOME=/root + +export SSL_CERT_FILE=/etc/ssl/vpod/ca-only.pem +export REQUESTS_CA_BUNDLE=/etc/ssl/vpod/ca-only.pem +export PIP_CERT=/etc/ssl/vpod/ca-only.pem +export NODE_EXTRA_CA_CERTS=/etc/ssl/vpod/ca-only.pem + +export ENV='' +unset HISTFILE +set +o history 2>/dev/null || true +exec setsid sh -c 'HISTFILE=/dev/null HISTSIZE=0 HOME=/root SSL_CERT_FILE=/etc/ssl/vpod/ca-only.pem exec sh /dev/ttyS0 2>&1' +INIT_EOF +chmod +x "$OVERLAY/sbin/init" +ln -sf /sbin/init "$OVERLAY/init" + +cat > "$OVERLAY/usr/lib/vpod/pyrunner.py" << 'PYRUNNER_EOF' +import sys, io, traceback, base64 + +try: + import ssl, urllib.request +except ImportError: + pass + +# pyrunner's own warm list (pydaemon has a separate one): pyrunner is a +# persistent process, so one import at startup (snapshot build time) makes +# it warm for every code.run(). +try: + with open("/etc/vpod/pyrunner-warm-imports") as _f: + for _line in _f: + _name = _line.split("#", 1)[0].strip() + if _name: + try: + __import__(_name) + except Exception: + pass +except OSError: + pass + +_globals = {} +_sentinel = "---VPOD_DONE---" +_real_stdout = sys.stdout +_real_stderr = sys.stderr +_data_out = open("/dev/ttyS3", "w") +_data_in = open("/dev/ttyS3", "r", buffering=1) +_exit_code_out = open("/dev/ttyS2", "wb", buffering=0) + +while True: + _line = _data_in.readline() + if not _line: + break + _line = _line.rstrip("\n") + if not _line: + continue + try: + _code = base64.b64decode(_line).decode() + except Exception: + _code = _line + + _buf = io.StringIO() + sys.stdout = _buf + sys.stderr = _buf + _exit_code = 0 + try: + exec(compile(_code, "", "exec"), _globals) + except SystemExit as _e: + if isinstance(_e.code, int): + _exit_code = _e.code & 0xFF + elif _e.code is not None: + _buf.write(str(_e.code) + "\n") + _exit_code = 1 + except Exception: + _buf.write(traceback.format_exc()) + _exit_code = 1 + finally: + sys.stdout = _real_stdout + sys.stderr = _real_stderr + + _val = _buf.getvalue() + if _val: + _data_out.write(_val) + _exit_code_out.write(bytes([_exit_code])) + _data_out.write(_sentinel + "\n") + _data_out.flush() +PYRUNNER_EOF + +echo "── Repacking minirootfs as cpio..." +MINI_WORK="$ROOT/dist/agent-minirootfs" +rm -rf "$MINI_WORK" +mkdir -p "$MINI_WORK" +bsdtar -xf "$MINIROOTFS" -C "$MINI_WORK" --no-same-owner + +echo "── Extracting kernel modules into minirootfs..." +mkdir -p "$MINI_WORK/lib" +gunzip -c "$INITRAMFS_LTS" | (cd "$MINI_WORK" && cpio -idmu --quiet 'usr/lib/modules/*') 2>/dev/null || true +if [ -d "$MINI_WORK/usr/lib/modules" ] && [ ! -e "$MINI_WORK/lib/modules" ]; then + ln -sf /usr/lib/modules "$MINI_WORK/lib/modules" +fi + +echo "── Packing rootfs.cpio.gz (minirootfs + modules + overlay)..." +PART_MINI="$ROOT/dist/agent-mini.cpio.gz" +PART_OVL="$ROOT/dist/agent-overlay.cpio.gz" +(cd "$MINI_WORK" && find . | sort | cpio -H newc -o --quiet) | gzip -9 > "$PART_MINI" +(cd "$OVERLAY" && find . | sort | cpio -H newc -o --quiet) | gzip -9 > "$PART_OVL" +cat "$PART_MINI" "$PART_OVL" > "$OUT" +rm -f "$PART_MINI" "$PART_OVL" +echo " Done: $OUT ($(du -sh "$OUT" | cut -f1))" + +SNAP="$ROOT/dist/vsnap-data-${RAM_MB}mb.snap" +BOOTARGS="root=/dev/ram0 rw console=ttyS0 earlycon init=/sbin/init" + +echo "── Booting guest to pre-install ca-certificates + python3 + data stack..." + +CA_MARKER="$(sed -n '2p' "$ROOT/crates/machine/assets/tls/vpod-ca-cert.pem")" +BUILD_LOG="$ROOT/dist/.snapshot-build.log" +NOW="$(date -u '+%Y-%m-%d %H:%M:%S')" + + +SETUP_CMD="" +SETUP_CMD="${SETUP_CMD}date -s '$NOW'; " +SETUP_CMD="${SETUP_CMD}sed -i 's|https://|http://|g' /etc/apk/repositories; " +SETUP_CMD="${SETUP_CMD}apk update --allow-untrusted; " + +SETUP_CMD="${SETUP_CMD}apk add --allow-untrusted ca-certificates python3 py3-pip uv py3-numpy py3-pandas py3-scipy; " +SETUP_CMD="${SETUP_CMD}rm -f /usr/lib/python3.*/EXTERNALLY-MANAGED; mkdir -p /root/.cache; " + +SETUP_CMD="${SETUP_CMD}update-ca-certificates; " +SETUP_CMD="${SETUP_CMD}grep -qF '$CA_MARKER' /etc/ssl/certs/ca-certificates.crt || cat /usr/local/share/ca-certificates/vpod-ca.crt >> /etc/ssl/certs/ca-certificates.crt; " +SETUP_CMD="${SETUP_CMD}if grep -qF '$CA_MARKER' /etc/ssl/certs/ca-certificates.crt; then echo VPOD_CA_INSTALLED; else echo VPOD_CA_FAILED; fi; " +SETUP_CMD="${SETUP_CMD}sed -i 's|http://|https://|g' /etc/apk/repositories; " +SETUP_CMD="${SETUP_CMD}cp /usr/bin/ssl_client /usr/bin/ssl_client.real && cp /usr/lib/vpod/vpod-ssl-client /usr/bin/ssl_client && chmod +x /usr/bin/ssl_client && echo VPOD_SSL_CLIENT_SWAPPED; " +SETUP_CMD="${SETUP_CMD}PYBIN=\$(readlink -f /usr/bin/python3) && cp \$PYBIN /usr/bin/python3.real && cp /usr/lib/vpod/vpod-python-shim \$PYBIN && chmod +x \$PYBIN /usr/bin/python3.real && echo VPOD_PY_SHIM_INSTALLED; " +SETUP_CMD="${SETUP_CMD}/usr/bin/python3.real /usr/lib/vpod/pydaemon.py /dev/null 2>&1 & " +SETUP_CMD="${SETUP_CMD}n=0; while [ ! -S /run/vpod-pyd.sock ] && [ \$n -lt 300 ]; do sleep 0.1; n=\$((n+1)); done; " +SETUP_CMD="${SETUP_CMD}[ -S /run/vpod-pyd.sock ] && /usr/bin/python3 -c 'print(\"VPOD_PYD_READY\")'; " + +SETUP_CMD="${SETUP_CMD}sync" + +"$VPOD" \ + "$KERNEL" \ + --bios "$OPENSBI_FW" \ + --initrd "$OUT" \ + --ram "$RAM_MB" \ + --bootargs "$BOOTARGS" \ + --net \ + --setup "$SETUP_CMD" \ + --snapshot-save "$SNAP" \ + --snapshot-python 2>&1 | tee "$BUILD_LOG" + +if ! grep -q VPOD_CA_INSTALLED "$BUILD_LOG"; then + echo "" >&2 + echo "error: the vpod proxy CA is not in the guest trust store." >&2 + echo " HTTPS interception (:443) would fail at runtime. Aborting the build." >&2 + echo " (see $BUILD_LOG for the guest setup output)" >&2 + exit 1 +fi +if ! grep -q VPOD_SSL_CLIENT_SWAPPED "$BUILD_LOG"; then + echo "" >&2 + echo "error: the vpod ssl_client was not installed over busybox's." >&2 + echo " wget https would pay the full guest-TLS cost. Aborting the build." >&2 + echo " (see $BUILD_LOG for the guest setup output)" >&2 + exit 1 +fi +if ! grep -q VPOD_PY_SHIM_INSTALLED "$BUILD_LOG"; then + echo "" >&2 + echo "error: the vpod python shim was not installed over the real python3." >&2 + echo " every python3 invocation would pay the ~0.65s cold start. Aborting the build." >&2 + echo " (see $BUILD_LOG for the guest setup output)" >&2 + exit 1 +fi +if ! grep -q VPOD_PYD_READY "$BUILD_LOG"; then + echo "" >&2 + echo "error: the warm-python daemon did not come up (no socket, or the" >&2 + echo " shim→daemon round-trip failed). Aborting the build." >&2 + echo " (see $BUILD_LOG for the guest setup output)" >&2 + exit 1 +fi +rm -f "$BUILD_LOG" + + +echo "" +echo "=== Done ===" +echo "" +echo "Snapshot: $SNAP" +echo "" +echo "Next (optional), to build and bake AOT-translated blocks:" +echo " ./scripts/aot-snapshot.sh \"$SNAP\" --workload data" diff --git a/scripts/build-default-snapshot.sh b/scripts/build-default-snapshot.sh index 0cc2219..25453ba 100755 --- a/scripts/build-default-snapshot.sh +++ b/scripts/build-default-snapshot.sh @@ -46,14 +46,14 @@ echo "" echo "── Checking host tools..." MISSING="" -for cmd in curl bsdtar cpio gzip cargo; do +for cmd in curl bsdtar cpio gzip cargo zig; do command -v "$cmd" >/dev/null || MISSING="$MISSING $cmd" done if [ -n "$MISSING" ]; then echo "ERROR: missing tools:$MISSING" - echo " macOS : brew install libarchive" - echo " Debian : apt install libarchive-tools bsdtar cpio" - echo " Fedora : dnf install bsdtar libarchive" + echo " macOS : brew install libarchive zig" + echo " Debian : apt install libarchive-tools bsdtar cpio zig" + echo " Fedora : dnf install bsdtar libarchive zig" echo " Windows: use WSL2 and follow the Linux instructions" exit 1 fi @@ -63,7 +63,7 @@ mkdir -p "$ROOT/dist" "$ALPINE_DIR" echo "── Building vpod..." -(cd "$ROOT" && cargo build --release --bin vpod-native) +(cd "$ROOT" && cargo build --release -p native-cli --bin vpod-native) if [ ! -f "$OPENSBI_FW" ]; then @@ -121,6 +121,41 @@ printf 'https://dl-cdn.alpinelinux.org/alpine/v%s/main\nhttps://dl-cdn.alpinelin "$ALPINE_MINOR" "$ALPINE_MINOR" > "$OVERLAY/etc/apk/repositories" echo 'nameserver 8.8.8.8' > "$OVERLAY/etc/resolv.conf" +mkdir -p "$OVERLAY/usr/local/share/ca-certificates" +cp "$ROOT/crates/machine/assets/tls/vpod-ca-cert.pem" \ + "$OVERLAY/usr/local/share/ca-certificates/vpod-ca.crt" + +mkdir -p "$OVERLAY/etc/ssl/vpod" +cp "$ROOT/crates/machine/assets/tls/vpod-ca-cert.pem" \ + "$OVERLAY/etc/ssl/vpod/ca-only.pem" + + +echo "── Cross-compiling vpod ssl_client (riscv64-musl, static)..." +zig cc -target riscv64-linux-musl -Os -static -s \ + -o "$OVERLAY/usr/lib/vpod/vpod-ssl-client" \ + "$ROOT/guest/tls/vpod_ssl_client.c" +chmod +x "$OVERLAY/usr/lib/vpod/vpod-ssl-client" + +mkdir -p "$OVERLAY/etc/vpod" +cat > "$OVERLAY/etc/vpod/pydaemon-warm-imports" << 'WARM_EOF' +# Warm set for the python3 shim/daemon path (commands.run("python3 ...") +WARM_EOF +cat > "$OVERLAY/etc/vpod/pyrunner-warm-imports" << 'WARM_EOF' +# Warm set for pyrunner, the persistent code.run() interpreter. +WARM_EOF + +mkdir -p "$OVERLAY/etc/uv" +cat > "$OVERLAY/etc/uv/uv.toml" << 'UV_EOF' +python-preference = "only-system" +UV_EOF + +echo "── Cross-compiling vpod python shim (riscv64-musl, dynamic)..." +zig cc -target riscv64-linux-musl -Os -dynamic -s \ + -o "$OVERLAY/usr/lib/vpod/vpod-python-shim" \ + "$ROOT/guest/warmpy/vpod_python_shim.c" +chmod +x "$OVERLAY/usr/lib/vpod/vpod-python-shim" +cp "$ROOT/guest/warmpy/pydaemon.py" "$OVERLAY/usr/lib/vpod/pydaemon.py" + cat > "$OVERLAY/sbin/init" << 'INIT_EOF' #!/bin/sh @@ -156,11 +191,17 @@ if [ -c /dev/hvc0 ]; then fi export TERM=dumb -export SSL_CERT_FILE=/etc/ssl/certs/ca-certificates.crt +export HOME=/root + +export SSL_CERT_FILE=/etc/ssl/vpod/ca-only.pem +export REQUESTS_CA_BUNDLE=/etc/ssl/vpod/ca-only.pem +export PIP_CERT=/etc/ssl/vpod/ca-only.pem +export NODE_EXTRA_CA_CERTS=/etc/ssl/vpod/ca-only.pem + export ENV='' unset HISTFILE set +o history 2>/dev/null || true -exec setsid sh -c 'HISTFILE=/dev/null HISTSIZE=0 SSL_CERT_FILE=/etc/ssl/certs/ca-certificates.crt exec sh /dev/ttyS0 2>&1' +exec setsid sh -c 'HISTFILE=/dev/null HISTSIZE=0 HOME=/root SSL_CERT_FILE=/etc/ssl/vpod/ca-only.pem exec sh /dev/ttyS0 2>&1' INIT_EOF chmod +x "$OVERLAY/sbin/init" ln -sf /sbin/init "$OVERLAY/init" @@ -168,6 +209,23 @@ ln -sf /sbin/init "$OVERLAY/init" cat > "$OVERLAY/usr/lib/vpod/pyrunner.py" << 'PYRUNNER_EOF' import sys, io, traceback, base64 +try: + import ssl, urllib.request +except ImportError: + pass + +try: + with open("/etc/vpod/pyrunner-warm-imports") as _f: + for _line in _f: + _name = _line.split("#", 1)[0].strip() + if _name: + try: + __import__(_name) + except Exception: + pass +except OSError: + pass + _globals = {} _sentinel = "---VPOD_DONE---" _real_stdout = sys.stdout @@ -237,11 +295,36 @@ cat "$PART_MINI" "$PART_OVL" > "$OUT" rm -f "$PART_MINI" "$PART_OVL" echo " Done: $OUT ($(du -sh "$OUT" | cut -f1))" -# SNAP="$ROOT/dist/vsnap-base-${RAM_MB}mb.snap" -SNAP="$ROOT/dist/alpine-3.23.0-256mb.snap" +SNAP="$ROOT/dist/alpine-3.23.0-${RAM_MB}mb.snap" BOOTARGS="root=/dev/ram0 rw console=ttyS0 earlycon init=/sbin/init" echo "── Booting guest to pre-install ca-certificates + python3..." + +CA_MARKER="$(sed -n '2p' "$ROOT/crates/machine/assets/tls/vpod-ca-cert.pem")" +BUILD_LOG="$ROOT/dist/.snapshot-build.log" +NOW="$(date -u '+%Y-%m-%d %H:%M:%S')" + + +SETUP_CMD="" +SETUP_CMD="${SETUP_CMD}date -s '$NOW'; " +SETUP_CMD="${SETUP_CMD}sed -i 's|https://|http://|g' /etc/apk/repositories; " +SETUP_CMD="${SETUP_CMD}apk update --allow-untrusted; " + +SETUP_CMD="${SETUP_CMD}apk add --allow-untrusted ca-certificates python3 uv; " +SETUP_CMD="${SETUP_CMD}rm -f /usr/lib/python3.*/EXTERNALLY-MANAGED; mkdir -p /root/.cache; " + +SETUP_CMD="${SETUP_CMD}update-ca-certificates; " +SETUP_CMD="${SETUP_CMD}grep -qF '$CA_MARKER' /etc/ssl/certs/ca-certificates.crt || cat /usr/local/share/ca-certificates/vpod-ca.crt >> /etc/ssl/certs/ca-certificates.crt; " +SETUP_CMD="${SETUP_CMD}if grep -qF '$CA_MARKER' /etc/ssl/certs/ca-certificates.crt; then echo VPOD_CA_INSTALLED; else echo VPOD_CA_FAILED; fi; " +SETUP_CMD="${SETUP_CMD}sed -i 's|http://|https://|g' /etc/apk/repositories; " +SETUP_CMD="${SETUP_CMD}cp /usr/bin/ssl_client /usr/bin/ssl_client.real && cp /usr/lib/vpod/vpod-ssl-client /usr/bin/ssl_client && chmod +x /usr/bin/ssl_client && echo VPOD_SSL_CLIENT_SWAPPED; " +SETUP_CMD="${SETUP_CMD}PYBIN=\$(readlink -f /usr/bin/python3) && cp \$PYBIN /usr/bin/python3.real && cp /usr/lib/vpod/vpod-python-shim \$PYBIN && chmod +x \$PYBIN /usr/bin/python3.real && echo VPOD_PY_SHIM_INSTALLED; " +SETUP_CMD="${SETUP_CMD}/usr/bin/python3.real /usr/lib/vpod/pydaemon.py /dev/null 2>&1 & " +SETUP_CMD="${SETUP_CMD}n=0; while [ ! -S /run/vpod-pyd.sock ] && [ \$n -lt 300 ]; do sleep 0.1; n=\$((n+1)); done; " +SETUP_CMD="${SETUP_CMD}[ -S /run/vpod-pyd.sock ] && /usr/bin/python3 -c 'print(\"VPOD_PYD_READY\")'; " + +SETUP_CMD="${SETUP_CMD}sync" + "$VPOD" \ "$KERNEL" \ --bios "$OPENSBI_FW" \ @@ -249,11 +332,46 @@ echo "── Booting guest to pre-install ca-certificates + python3..." --ram "$RAM_MB" \ --bootargs "$BOOTARGS" \ --net \ - --setup "date -s '$(date -u '+%Y-%m-%d %H:%M:%S')'; sed -i 's|https://|http://|g' /etc/apk/repositories; apk update --allow-untrusted; apk add --allow-untrusted ca-certificates python3; sed -i 's|http://|https://|g' /etc/apk/repositories; sync" \ + --setup "$SETUP_CMD" \ --snapshot-save "$SNAP" \ - --snapshot-python + --snapshot-python 2>&1 | tee "$BUILD_LOG" + +if ! grep -q VPOD_CA_INSTALLED "$BUILD_LOG"; then + echo "" >&2 + echo "error: the vpod proxy CA is not in the guest trust store." >&2 + echo " HTTPS interception (:443) would fail at runtime. Aborting the build." >&2 + echo " (see $BUILD_LOG for the guest setup output)" >&2 + exit 1 +fi +if ! grep -q VPOD_SSL_CLIENT_SWAPPED "$BUILD_LOG"; then + echo "" >&2 + echo "error: the vpod ssl_client was not installed over busybox's." >&2 + echo " wget https would pay the full guest-TLS cost. Aborting the build." >&2 + echo " (see $BUILD_LOG for the guest setup output)" >&2 + exit 1 +fi +if ! grep -q VPOD_PY_SHIM_INSTALLED "$BUILD_LOG"; then + echo "" >&2 + echo "error: the vpod python shim was not installed over the real python3." >&2 + echo " every python3 invocation would pay the ~0.65s cold start. Aborting the build." >&2 + echo " (see $BUILD_LOG for the guest setup output)" >&2 + exit 1 +fi +if ! grep -q VPOD_PYD_READY "$BUILD_LOG"; then + echo "" >&2 + echo "error: the warm-python daemon did not come up (no socket, or the" >&2 + echo " shim→daemon round-trip failed). Aborting the build." >&2 + echo " (see $BUILD_LOG for the guest setup output)" >&2 + exit 1 +fi +rm -f "$BUILD_LOG" + + echo "" echo "=== Done ===" echo "" echo "Snapshot: $SNAP" +echo "" +echo "Next (optional), to build and bake AOT-translated blocks:" +echo " ./scripts/aot-snapshot.sh \"$SNAP\"" diff --git a/scripts/build-registry-bundle.sh b/scripts/build-registry-bundle.sh new file mode 100755 index 0000000..5ad6086 --- /dev/null +++ b/scripts/build-registry-bundle.sh @@ -0,0 +1,47 @@ +#!/usr/bin/env bash + +set -euo pipefail + +cd "$(dirname "$0")/.." + +TEMPLATE="${TEMPLATE:-docs/v0.4.1-registry/test/snapshots.json}" +OUT="${OUT:-dist/registry-bundle}" + +if [ "${SKIP_BUILD:-0}" != "1" ]; then + ./scripts/build-default-snapshot.sh + ./scripts/build-default-snapshot.sh --ram 512 + ./scripts/build-data-snapshot.sh +fi + +for snap in dist/alpine-3.23.0-256mb.snap dist/alpine-3.23.0-512mb.snap dist/vsnap-data-512mb.snap; do + [ -f "$snap" ] || { echo "error: $snap missing (build it or drop SKIP_BUILD)"; exit 1; } +done + +mkdir -p "$OUT" +lz4 -9 -f dist/alpine-3.23.0-256mb.snap "$OUT/alpine-3.23.0-256mb.snap" +cp "$OUT/alpine-3.23.0-256mb.snap" "$OUT/vsnap-base-256mb.snap" +lz4 -9 -f dist/alpine-3.23.0-512mb.snap "$OUT/vsnap-base-512mb.snap" +lz4 -9 -f dist/vsnap-data-512mb.snap "$OUT/vsnap-data-512mb.snap" + +TEMPLATE="$TEMPLATE" OUT="$OUT" python3 - <<'PY' +import hashlib, json, os +from pathlib import Path + +out = Path(os.environ["OUT"]) +manifest = json.loads(Path(os.environ["TEMPLATE"]).read_text()) + +if os.environ.get("VERSION"): + manifest["version"] = os.environ["VERSION"] + +for entry in manifest["snapshots"]: + data = (out / f"{entry['id']}.snap").read_bytes() + entry["sha256"] = hashlib.sha256(data).hexdigest() + entry["size"] = len(data) + print(f" {entry['id']}: sha256={entry['sha256'][:12]}… size={entry['size']:,}") + +(out / "snapshots.json").write_text(json.dumps(manifest, indent=2) + "\n") +PY + +echo +echo "bundle ready in $OUT:" +ls -lh "$OUT" diff --git a/scripts/build-wasm.sh b/scripts/build-wasm.sh index 6a03ff0..85abb73 100755 --- a/scripts/build-wasm.sh +++ b/scripts/build-wasm.sh @@ -1,6 +1,15 @@ #!/usr/bin/env bash + set -euo pipefail +ROOT="$(cd "$(dirname "$0")/.." && pwd)" +GENERATED="$ROOT/crates/riscv-core/src/aot/generated.rs" +SDK_DIR="$ROOT/sdks/python/vpod" +OUT_DIR="$ROOT/target/wasm32-wasip2/release" + +BASE_TARGET_DIR="$ROOT/target/base-wasm" +BACKUP="$ROOT/dist/.generated.rs.aot-backup" + RUSTUP_HOME="${RUSTUP_HOME:-$HOME/.rustup}" STABLE_TC=$(rustup toolchain list | grep '^stable-' | head -1 | awk '{print $1}') @@ -13,16 +22,69 @@ RUSTUP_CARGO="$RUSTUP_HOME/toolchains/$STABLE_TC/bin/cargo" RUSTUP_RUSTC="$RUSTUP_HOME/toolchains/$STABLE_TC/bin/rustc" echo "[build-wasm] toolchain: $STABLE_TC" -echo "[build-wasm] building wasi-component (cli)..." -RUSTC="$RUSTUP_RUSTC" "$RUSTUP_CARGO" build -p wasi-component --bin vpod-wasi-cli --release --target wasm32-wasip2 -echo "[build-wasm] building wasi-component (library)..." -RUSTC="$RUSTUP_RUSTC" "$RUSTUP_CARGO" build -p wasi-component --lib --release --target wasm32-wasip2 +if [[ -f "$BACKUP" ]]; then + echo "error: a previous run left $BACKUP behind, which means it was killed" >&2 + echo " mid-build with generated.rs stubbed. That backup is the only copy" >&2 + echo " of the translation — restore it before continuing:" >&2 + echo " cp -p '$BACKUP' '$GENERATED' && rm '$BACKUP'" >&2 + exit 1 +fi + +if [[ ! -f "$GENERATED" ]]; then + echo "[build-wasm] no generated.rs — writing stub" + "$ROOT/scripts/aot-stub.sh" +fi + +HAVE_AOT=0 +if [[ $(wc -l < "$GENERATED" | tr -d ' ') -gt 100 ]]; then + HAVE_AOT=1 +fi + +build_lib() { + RUSTC="$RUSTUP_RUSTC" "$RUSTUP_CARGO" build -p wasi-component --lib --release --target wasm32-wasip2 +} + +if [[ "$HAVE_AOT" == "1" ]]; then + echo "[build-wasm] building wasi-component (cli)..." + RUSTC="$RUSTUP_RUSTC" "$RUSTUP_CARGO" build -p wasi-component --bin vpod-wasi-cli --release --target wasm32-wasip2 + + echo "[build-wasm] building wasi-component (library, AOT)..." + build_lib + cp "$OUT_DIR/vpod_wasi_lib.wasm" "$SDK_DIR/vpod_wasi_lib_aot.wasm" + + echo "[build-wasm] building wasi-component (library, base)..." + mkdir -p "$(dirname "$BACKUP")" + + cp -p "$GENERATED" "$BACKUP" + trap 'cp -p "$BACKUP" "$GENERATED" && rm -f "$BACKUP"' EXIT INT TERM + + "$ROOT/scripts/aot-stub.sh" --force + CARGO_TARGET_DIR="$BASE_TARGET_DIR" build_lib + cp "$BASE_TARGET_DIR/wasm32-wasip2/release/vpod_wasi_lib.wasm" \ + "$SDK_DIR/vpod_wasi_lib.wasm" + + cp -p "$BACKUP" "$GENERATED" + rm -f "$BACKUP" + trap - EXIT INT TERM +else + echo "[build-wasm] generated.rs is a stub — building base tier only" + echo "[build-wasm] building wasi-component (cli)..." + RUSTC="$RUSTUP_RUSTC" "$RUSTUP_CARGO" build -p wasi-component --bin vpod-wasi-cli --release --target wasm32-wasip2 + + echo "[build-wasm] building wasi-component (library, base)..." + build_lib + cp "$OUT_DIR/vpod_wasi_lib.wasm" "$SDK_DIR/vpod_wasi_lib.wasm" + rm -f "$SDK_DIR/vpod_wasi_lib_aot.wasm" +fi echo "[build-wasm] building vpod host..." cargo build -p vpod --release echo "[build-wasm] done" -echo " cli: target/wasm32-wasip2/release/vpod-wasi-cli.wasm" -echo " library: target/wasm32-wasip2/release/vpod_wasi_lib.wasm" -echo " host: target/release/vpod" +if [[ "$HAVE_AOT" == "1" ]]; then + echo " library (AOT): $SDK_DIR/vpod_wasi_lib_aot.wasm" +fi +echo " library (base): $SDK_DIR/vpod_wasi_lib.wasm" +echo " cli: $OUT_DIR/vpod-wasi-cli.wasm" +echo " host: $ROOT/target/release/vpod" diff --git a/sdks/python/README.md b/sdks/python/README.md index 586db79..f3165ae 100644 --- a/sdks/python/README.md +++ b/sdks/python/README.md @@ -6,6 +6,8 @@ CI + +[Documentation](https://docs.vpod.sh/quickstart) • [Issues](https://github.com/capsulerun/vpod/issues/new)
@@ -26,7 +28,7 @@ pip install vpod ### Persistent session (Recommended) -All calls share the same running VM. Using a context manager (`with`) automatically cleans up resources when done: +All calls share the same running sandbox. Using a context manager (`with`) automatically cleans up resources when done: ```python from vpod import Sandbox @@ -76,12 +78,12 @@ Pause a running sandbox and resume it later — no daemon, no background process from vpod import Sandbox with Sandbox.create() as sbx: - sbx.commands.run("pip install numpy") + sbx.commands.run("uv pip install --system requests") instance_id = sbx.suspend() # Later (even from a new process): sbx = Sandbox.resume(instance_id) -sbx.code.run("import numpy; print(numpy.__version__)") +sbx.code.run("import requests; print(requests.__version__)") ``` | Method | Description | @@ -106,7 +108,7 @@ sbx.close() # Clean up the sandbox process ## Snapshots -The first call to `Sandbox.create()` downloads the VM snapshot and caches it locally. Subsequent calls use the cache instantly. +The first call to `Sandbox.create()` downloads the snapshot and caches it locally. Subsequent calls use the cache instantly. To pre-download (e.g. in a Dockerfile or CI setup): @@ -125,20 +127,9 @@ snapshots.pull("alpine:latest") |:---|:---|:---|:---| | `alpine` | 3.23.0 | Minimal Alpine Linux snapshot. | 256 MB | | `vsnap-base` | 0.1.0 | Alpine-based general-purpose snapshot with Python. | 256 MB | +| `vsnap-base-512mb` | 0.1.0 | Same as `vsnap-base` with more memory headroom, for web servers and larger installs. | 512 MB | | `vsnap-data` | 0.1.0 | Alpine-based snapshot with `numpy`, `pandas`, and `scipy`. | 512 MB | -## How it works - -A `vpod` runs a RISC‑V virtual machine compiled to WebAssembly. The core implements the **RV64GC** specification: - -- **G (General-purpose)**: I/M/A/F/D extensions for integer, multiply/divide, atomics, and floating-point. -- **C (Compressed)**: 30% smaller code size, improving memory efficiency. - -The WASM component communicates with the host through WASI 0.2, providing controlled access to networking and I/O while keeping all execution state isolated inside the sandbox. - -## Limitations - -- **Emulation overhead**: No hardware acceleration in WASM. CPU-intensive workloads run slower than native. -- **No GPU access**: CUDA, Metal, and hardware ML accelerators are not yet available. +## Documentation -For full documentation and to report issues, visit the [main GitHub repository](https://github.com/capsulerun/vpod). +Visit the [Vpod documentation](https://docs.vpod.sh/quickstart) for the full guide and API reference. To report issues or contribute, head to the [main GitHub repository](https://github.com/capsulerun/vpod). diff --git a/sdks/python/vpod/_component.py b/sdks/python/vpod/_component.py index 1774eec..5935f6f 100644 --- a/sdks/python/vpod/_component.py +++ b/sdks/python/vpod/_component.py @@ -1,8 +1,12 @@ import hashlib +import os +import subprocess +import sys +import threading from pathlib import Path from platformdirs import user_data_dir -from wasmtime import Engine, Store, WasiConfig +from wasmtime import Config, Engine, Store, WasiConfig from wasmtime._bindings import ( wasi_config_allow_ip_name_lookup, wasi_config_inherit_network, @@ -11,8 +15,15 @@ ) from wasmtime.component import Component, Linker -_BUNDLED_WASM = Path(__file__).parent / "vpod_wasi_lib.wasm" -_REPO_WASM = Path(__file__).parents[4] / "target" / "wasm32-wasip2" / "release" / "vpod_wasi_lib.wasm" + +_PKG_DIR = Path(__file__).parent +_TARGET_DIR = Path(__file__).parents[4] / "target" / "wasm32-wasip2" / "release" + +_AOT_CANDIDATES = ( + _PKG_DIR / "vpod_wasi_lib_aot.wasm", + _TARGET_DIR / "vpod_wasi_lib_aot.wasm", +) +_BASE_CANDIDATES = (_PKG_DIR / "vpod_wasi_lib.wasm", _TARGET_DIR / "vpod_wasi_lib.wasm") try: from importlib.metadata import version @@ -24,44 +35,71 @@ _component = None _linker = None _instance_cache = {} +_precompile_started = set() +_thread_compile_started = set() +_cwasm_path_cache = {} +_active_tier = None +_load_lock = threading.RLock() -def locate_wasm() -> Path: - for candidate in (_BUNDLED_WASM, _REPO_WASM): +def _first_existing(candidates) -> Path | None: + for candidate in candidates: if candidate.exists(): return candidate + return None + + +def locate_wasm() -> Path: + """The best engine module available: AOT if we have it, else interpreter-only.""" + found = _first_existing(_AOT_CANDIDATES) or _first_existing(_BASE_CANDIDATES) + if found is not None: + return found raise FileNotFoundError( f"WASM library not found. Either:\n" - f" - Run: cargo build -p wasi-component --lib --release --target wasm32-wasip2\n" - f" - Or copy vpod_wasi_lib.wasm to {_BUNDLED_WASM}" + f" - Run: ./scripts/build-wasm.sh\n" + f" - Or copy vpod_wasi_lib.wasm to {_PKG_DIR}" ) +def _tier_of(wasm_path: Path) -> str: + return "aot" if wasm_path.name.endswith("_aot.wasm") else "base" + + def _cwasm_cache_path(wasm_path: Path) -> Path: - wasm_bytes = wasm_path.read_bytes() - digest = hashlib.sha256(wasm_bytes).hexdigest()[:16] + """Cache file for a module, named so the directory explains itself.""" + cached = _cwasm_path_cache.get(wasm_path) + if cached is not None: + return cached + + digest = hashlib.sha256(wasm_path.read_bytes()).hexdigest()[:16] base = Path(user_data_dir()) or Path.home() / ".local" / "share" cache_dir = base / "vpod" cache_dir.mkdir(parents=True, exist_ok=True) - return cache_dir / f"component-{_VERSION}-{digest}.cwasm" + path = cache_dir / f"component-{_VERSION}-{_tier_of(wasm_path)}-{digest}.cwasm" + _cwasm_path_cache[wasm_path] = path + return path -def _load_or_compile_component(engine: Engine, wasm_path: Path) -> Component: - cache_path = _cwasm_cache_path(wasm_path) +def _write_cwasm_atomically(cache_path: Path, serialized: bytes) -> None: + """Publish the .cwasm by rename, never by writing into its final name.""" + tmp_path = cache_path.with_suffix(f".{os.getpid()}.tmp") + + try: + tmp_path.write_bytes(serialized) + os.replace(tmp_path, cache_path) + except Exception: + tmp_path.unlink(missing_ok=True) + raise - if cache_path.exists(): - try: - return Component.deserialize_file(engine, str(cache_path)) - except Exception: - cache_path.unlink(missing_ok=True) +def _compile_and_cache(engine: Engine, wasm_path: Path) -> Component: component = Component.from_file(engine, str(wasm_path)) try: - serialized = component.serialize() - cache_path.write_bytes(serialized) + cache_path = _cwasm_cache_path(wasm_path) + _write_cwasm_atomically(cache_path, component.serialize()) _prune_stale_cwasm(cache_path) except Exception: pass @@ -69,6 +107,95 @@ def _load_or_compile_component(engine: Engine, wasm_path: Path) -> Component: return component +def _load_cached(engine: Engine, wasm_path: Path) -> Component | None: + cache_path = _cwasm_cache_path(wasm_path) + if not cache_path.exists(): + return None + + try: + return Component.deserialize_file(engine, str(cache_path)) + except Exception: + cache_path.unlink(missing_ok=True) + return None + + +def _precompile_in_background(wasm_path: Path, parallel: bool = False, fallback: bool = False) -> None: + """Warm the AOT module's cache out of process.""" + key = str(wasm_path) + if key in _precompile_started: + return + _precompile_started.add(key) + + cmd = [sys.executable, "-m", "vpod._precompile", str(wasm_path)] + if parallel: + cmd.append("--parallel") + if fallback: + cmd.append("--fallback") + + try: + subprocess.Popen( + cmd, + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + stdin=subprocess.DEVNULL, + start_new_session=True, + cwd=str(Path(__file__).parents[1]), + ) + except Exception: + pass + + +def _promote_to_aot(engine: Engine, component: Component) -> None: + """Swap the process onto an already-compiled AOT component.""" + global _engine, _component, _linker, _active_tier + + with _load_lock: + if _active_tier == "aot": + return + + linker = Linker(engine) + linker.add_wasip2() + wasmtime_component_linker_add_wasi_http(linker.ptr()) + + _engine, _component, _linker = engine, component, linker + _active_tier = "aot" + _instance_cache.clear() + + +def _compile_aot_in_thread(wasm_path: Path) -> None: + """Compile the AOT module in-process and hand it over without a disk trip.""" + key = str(wasm_path) + if key in _thread_compile_started: + return + _thread_compile_started.add(key) + + def work() -> None: + try: + config = Config() + config.parallel_compilation = os.environ.get("VPOD_AOT_EAGER") != "0" + engine = Engine(config) + component = Component.from_file(engine, str(wasm_path)) + _promote_to_aot(engine, component) + + cache_path = _cwasm_cache_path(wasm_path) + if not cache_path.exists(): + _write_cwasm_atomically(cache_path, component.serialize()) + _prune_stale_cwasm(cache_path) + retire_base_cwasm() + except Exception: + pass + + threading.Thread(target=work, daemon=True, name="vpod-aot-compile").start() + + +def prewarm() -> None: + aot_path = _first_existing(_AOT_CANDIDATES) + if aot_path is None or _cwasm_cache_path(aot_path).exists(): + return + + _precompile_in_background(aot_path, parallel=True) + + def _prune_stale_cwasm(active_cache_path: Path) -> None: caches = sorted( active_cache_path.parent.glob("component-*.cwasm"), @@ -80,22 +207,94 @@ def _prune_stale_cwasm(active_cache_path: Path) -> None: stale.unlink(missing_ok=True) +def retire_base_cwasm() -> None: + """Drop the base tier's cache once the AOT one exists.""" + base_path = _first_existing(_BASE_CANDIDATES) + if base_path is None or _first_existing(_AOT_CANDIDATES) is None: + return + + try: + _cwasm_cache_path(base_path).unlink(missing_ok=True) + except OSError: + pass + + +def _select_component(engine: Engine, wasm_path: Path) -> Component: + """Pick a tier and load it, preferring speed-now over speed-later.""" + global _active_tier + + cached = _load_cached(engine, wasm_path) + if cached is not None: + _active_tier = _tier_of(wasm_path) + return cached + + base_path = _first_existing(_BASE_CANDIDATES) + is_aot = base_path is not None and base_path != wasm_path + blocking = os.environ.get("VPOD_AOT_BLOCKING") == "1" + + if is_aot and not blocking: + _precompile_in_background(wasm_path, fallback=True) + + base = _load_cached(engine, base_path) + if base is None: + base = _compile_and_cache(engine, base_path) + _active_tier = "base" + + _compile_aot_in_thread(wasm_path) + return base + + _active_tier = _tier_of(wasm_path) + return _compile_and_cache(engine, wasm_path) + + +def _maybe_upgrade_tier() -> None: + """Move new sandboxes to the AOT tier once its .cwasm lands.""" + global _engine, _component, _linker, _active_tier + + with _load_lock: + if _component is None or _active_tier != "base": + return + + aot_path = _first_existing(_AOT_CANDIDATES) + if aot_path is None or not _cwasm_cache_path(aot_path).exists(): + return + + engine = Engine() + component = _load_cached(engine, aot_path) + if component is None: + return + + linker = Linker(engine) + linker.add_wasip2() + wasmtime_component_linker_add_wasi_http(linker.ptr()) + + _engine, _component, _linker = engine, component, linker + _active_tier = "aot" + _instance_cache.clear() + + +def active_tier() -> str | None: + """Which tier this process is actually running on: 'aot', 'base', or None.""" + return _active_tier + + def _get_or_load_component(wasm_path: Path): global _engine, _component, _linker - if _engine is None: - _engine = Engine() + with _load_lock: + if _engine is None: + _engine = Engine() - if _component is None: - _component = _load_or_compile_component(_engine, wasm_path) + if _component is None: + _component = _select_component(_engine, wasm_path) - if _linker is None: - _linker = Linker(_engine) - _linker.add_wasip2() + if _linker is None: + _linker = Linker(_engine) + _linker.add_wasip2() - wasmtime_component_linker_add_wasi_http(_linker.ptr()) + wasmtime_component_linker_add_wasi_http(_linker.ptr()) - return _engine, _component, _linker + return _engine, _component, _linker def _resolve_exports(store, instance): @@ -103,6 +302,8 @@ def _resolve_exports(store, instance): if iface_index is None: raise RuntimeError("WASM component does not export 'vpod:sandbox/executor@0.1.0'") + lock = threading.RLock() + def get_export(name: str): idx = instance.get_export_index(store, name, iface_index) if idx is None: @@ -110,7 +311,12 @@ def get_export(name: str): func = instance.get_func(store, idx) if func is None: raise RuntimeError(f"WASM export '{name}' is not a function") - return lambda *args: func(store, *args) + + def call(*args): + with lock: + return func(store, *args) + + return call return {name: get_export(name) for name in ("session-start", "session-exec", "session-close", "session-suspend", "session-resume")} @@ -128,6 +334,8 @@ def load_component(wasm_path: Path, snapshot_path: Path = None, mount_dirs: list key = _instance_key(snap_dir, mount_dirs) + _maybe_upgrade_tier() + if key in _instance_cache: return _instance_cache[key] diff --git a/sdks/python/vpod/_precompile.py b/sdks/python/vpod/_precompile.py new file mode 100644 index 0000000..196533a --- /dev/null +++ b/sdks/python/vpod/_precompile.py @@ -0,0 +1,42 @@ +import sys +import time +from pathlib import Path + +from wasmtime import Config, Engine + +from ._component import _cwasm_cache_path, _compile_and_cache, retire_base_cwasm + + +def main() -> int: + args = sys.argv[1:] + parallel = "--parallel" in args + fallback = "--fallback" in args + args = [a for a in args if a not in ("--parallel", "--fallback")] + + if len(args) != 1: + print("usage: python -m vpod._precompile [--parallel] [--fallback]", file=sys.stderr) + return 2 + + wasm_path = Path(args[0]) + if not wasm_path.exists(): + return 1 + + if fallback: + deadline = time.monotonic() + 30 + while time.monotonic() < deadline: + if _cwasm_cache_path(wasm_path).exists(): + retire_base_cwasm() + return 0 + time.sleep(1) + + config = Config() + config.parallel_compilation = parallel + + _compile_and_cache(Engine(config), wasm_path) + + retire_base_cwasm() + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/sdks/python/vpod/code.py b/sdks/python/vpod/code.py index 7e3a733..13c8353 100644 --- a/sdks/python/vpod/code.py +++ b/sdks/python/vpod/code.py @@ -5,8 +5,8 @@ class Code: """Code execution interface for a sandbox — persistent Python REPL.""" - def __init__(self, exports, snapshot_path: str, get_session_id): - self._exports = exports + def __init__(self, get_exports, snapshot_path: str, get_session_id): + self._get_exports = get_exports self._snapshot_path = snapshot_path self._get_session_id = get_session_id @@ -19,7 +19,7 @@ def run(self, code: str, timeout: int = 120) -> CodeExecution: "Use 'with Sandbox.create() as sandbox:'" ) - result = unwrap_result(self._exports["session-exec"](session_id, "\x00" + code, timeout)) + result = unwrap_result(self._get_exports()["session-exec"](session_id, "\x00" + code, timeout)) output = result.stdout if hasattr(result, 'stdout') else str(result) stderr = result.stderr if hasattr(result, 'stderr') else "" diff --git a/sdks/python/vpod/commands.py b/sdks/python/vpod/commands.py index b47314b..2a77b97 100644 --- a/sdks/python/vpod/commands.py +++ b/sdks/python/vpod/commands.py @@ -5,14 +5,14 @@ class Commands: """Shell command execution interface for a sandbox.""" - def __init__(self, exports, snapshot_path: str, get_session_id): - self._exports = exports + def __init__(self, get_exports, snapshot_path: str, get_session_id): + self._get_exports = get_exports self._snapshot_path = snapshot_path self._get_session_id = get_session_id def run(self, command: str, timeout: int = 120) -> CommandResult: session_id = self._get_session_id() - exec = self._exports["session-exec"] + exec = self._get_exports()["session-exec"] result = unwrap_result(exec(session_id, command, timeout)) return CommandResult( diff --git a/sdks/python/vpod/sandbox.py b/sdks/python/vpod/sandbox.py index 4564a87..82dc462 100644 --- a/sdks/python/vpod/sandbox.py +++ b/sdks/python/vpod/sandbox.py @@ -7,7 +7,7 @@ from . import snapshots from .snapshots import cache_dir -from ._component import load_component, locate_wasm +from ._component import _maybe_upgrade_tier, active_tier, load_component, locate_wasm from ._result import unwrap_result as _unwrap_result from .code import Code from .commands import Commands @@ -45,21 +45,24 @@ def __init__(self, snapshot: str = "alpine:latest", mounts: dict[str, str] | Non wasm_path = locate_wasm() self._snapshot_path = "snap/" + snapshot_path.name + self._snapshot_file = snapshot_path self._mounts = _parse_mounts(mounts) if mounts else [] mount_dirs = [m["host_path"] for m in self._mounts] self._store, self._exports = load_component(wasm_path, snapshot_path, mount_dirs or None) self._shell_session_id: Optional[int] = None self._in_context = False + self._tier = active_tier() + self._migrating = False self.commands = Commands( - self._exports, + lambda: self._exports, self._snapshot_path, self._get_shell_session_id, ) self.code = Code( - self._exports, + lambda: self._exports, self._snapshot_path, self._get_code_session_id, ) @@ -68,22 +71,77 @@ def __init__(self, snapshot: str = "alpine:latest", mounts: dict[str, str] | Non def create(cls, snapshot: str = "vsnap-base:latest", mounts: dict[str, str] | None = None) -> "Sandbox": return cls(snapshot, mounts=mounts) + def _mount_entries(self) -> list: + mount_entries = [] + for i, m in enumerate(self._mounts): + entry = object.__new__(type("MountEntry", (), {})) + object.__setattr__(entry, "host-alias", f"mount{i}") + object.__setattr__(entry, "guest-path", m["guest_path"]) + object.__setattr__(entry, "writable", m["writable"]) + mount_entries.append(entry) + return mount_entries + def _get_shell_session_id(self) -> int: + self._maybe_upgrade_engine() if self._shell_session_id is None: - mount_entries = [] - for i, m in enumerate(self._mounts): - entry = object.__new__(type("MountEntry", (), {})) - object.__setattr__(entry, "host-alias", f"mount{i}") - object.__setattr__(entry, "guest-path", m["guest_path"]) - object.__setattr__(entry, "writable", m["writable"]) - mount_entries.append(entry) - result = self._exports["session-start"]( - self._snapshot_path, _DEFAULT_SHELL, _DEFAULT_PROMPT, mount_entries + self._snapshot_path, _DEFAULT_SHELL, _DEFAULT_PROMPT, self._mount_entries() ) self._shell_session_id = int(_unwrap_result(result)) return self._shell_session_id + def _maybe_upgrade_engine(self) -> None: + """Hop this sandbox onto the AOT tier once its cache lands.""" + if self._tier != "base" or self._migrating: + return + + _maybe_upgrade_tier() + if active_tier() != "aot": + return + + self._migrating = True + try: + mount_dirs = [m["host_path"] for m in self._mounts] or None + + if self._shell_session_id is None: + try: + self._store, self._exports = load_component( + locate_wasm(), self._snapshot_file, mount_dirs + ) + except Exception: + return + else: + old_store, old_exports = self._store, self._exports + try: + instance_id = self.suspend() + except Exception: + return + delta_rel = f"instances/{instance_id}/delta.bin" + + def _resume(exports) -> None: + result = exports["session-resume"]( + self._snapshot_path, delta_rel, _DEFAULT_SHELL, + _DEFAULT_PROMPT, self._mount_entries(), + ) + self._shell_session_id = int(_unwrap_result(result)) + + try: + self._store, self._exports = load_component( + locate_wasm(), self._snapshot_file, mount_dirs + ) + _resume(self._exports) + except Exception: + self._store, self._exports = old_store, old_exports + _resume(old_exports) + Sandbox.destroy(instance_id) + return + + Sandbox.destroy(instance_id) + + self._tier = "aot" + finally: + self._migrating = False + def _get_code_session_id(self) -> Optional[int]: if not self._in_context: return None @@ -141,6 +199,7 @@ def resume(cls, instance_id: str, mounts: dict[str, str] | None = None) -> "Sand snapshot_file = meta["snapshot"].removeprefix("snap/") override = os.environ.get("VPOD_SNAPSHOT") + if override and Path(override).exists(): snapshot_path = Path(override) else: @@ -182,13 +241,16 @@ def resume(cls, instance_id: str, mounts: dict[str, str] | None = None) -> "Sand instance = cls.__new__(cls) instance._snapshot_path = snap_rel + instance._snapshot_file = snapshot_path + instance._tier = active_tier() + instance._migrating = False instance._mounts = saved_mounts instance._store = store instance._exports = exports instance._shell_session_id = session_id instance._in_context = True - instance.commands = Commands(exports, snap_rel, instance._get_shell_session_id) - instance.code = Code(exports, snap_rel, instance._get_code_session_id) + instance.commands = Commands(lambda: instance._exports, snap_rel, instance._get_shell_session_id) + instance.code = Code(lambda: instance._exports, snap_rel, instance._get_code_session_id) Sandbox.destroy(instance_id) return instance @@ -206,8 +268,10 @@ def destroy(instance_id: str) -> None: def list_instances() -> list[dict]: """List all suspended instances.""" manifest_path = INSTANCES_DIR / "manifest.json" + if not manifest_path.exists(): return [] + return json.loads(manifest_path.read_text()) @staticmethod @@ -222,8 +286,10 @@ def _update_manifest(instance_id: str, state: str) -> None: @staticmethod def _remove_from_manifest(instance_id: str) -> None: manifest_path = INSTANCES_DIR / "manifest.json" + if not manifest_path.exists(): return + entries = json.loads(manifest_path.read_text()) entries = [e for e in entries if e["id"] != instance_id] manifest_path.write_text(json.dumps(entries, indent=2)) diff --git a/sdks/python/vpod/snapshots.py b/sdks/python/vpod/snapshots.py index fad981b..2fb5c31 100644 --- a/sdks/python/vpod/snapshots.py +++ b/sdks/python/vpod/snapshots.py @@ -53,6 +53,10 @@ def pull(name: str = "vsnap-base:latest") -> Path: meta.unlink(missing_ok=True) dest.parent.mkdir(parents=True, exist_ok=True) + + from ._component import prewarm + prewarm() + _download_and_decompress(snapshot["url"], dest, snapshot["sha256"]) meta.write_text(snapshot["sha256"]) _prune_stale_snapshots(registry) @@ -95,6 +99,7 @@ def _snapshots_referenced_by_instances() -> set[str]: _REGISTRY_TTL = 86400 _REGISTRY_CACHE = cache_dir() / "snapshots.json" +_REGISTRY_VERSION_MARKER = cache_dir() / "snapshots.json.sdkver" def catalog() -> list[dict]: @@ -102,8 +107,15 @@ def catalog() -> list[dict]: return fetch_registry() +def _registry_cache_version_matches() -> bool: + try: + return _REGISTRY_VERSION_MARKER.read_text().strip() == _version() + except OSError: + return False + + def fetch_registry() -> list[dict]: - if _REGISTRY_CACHE.exists(): + if _REGISTRY_CACHE.exists() and _registry_cache_version_matches(): age = time.time() - _REGISTRY_CACHE.stat().st_mtime if age < _REGISTRY_TTL: return json.loads(_REGISTRY_CACHE.read_text())["snapshots"] @@ -120,6 +132,7 @@ def fetch_registry() -> list[dict]: _REGISTRY_CACHE.parent.mkdir(parents=True, exist_ok=True) _REGISTRY_CACHE.write_bytes(data) + _REGISTRY_VERSION_MARKER.write_text(_version()) return json.loads(data)["snapshots"] except Exception as e: if _REGISTRY_CACHE.exists():