From 9b1b223186c3c6d9711b919f1ba936e2139733c5 Mon Sep 17 00:00:00 2001 From: Kbz-8 Date: Mon, 20 Jul 2026 22:07:41 +0200 Subject: [PATCH] adding backend agnostic IR --- .gitea/workflows/Build.yml | 23 +- .gitea/workflows/Test.yml | 76 +- build.zig | 44 +- src/ape/ApeInstance.zig | 2 +- src/ape/lib.zig | 4 +- src/compiler/README.md | 612 ++++++++++++ src/compiler/ir/Builder.zig | 227 +++++ src/compiler/ir/Rewriter.zig | 323 +++++++ src/compiler/ir/cfg.zig | 198 ++++ src/compiler/ir/constant.zig | 18 + src/compiler/ir/id.zig | 84 ++ src/compiler/ir/instruction.zig | 228 +++++ src/compiler/ir/ir.zig | 58 ++ src/compiler/ir/module.zig | 235 +++++ src/compiler/ir/parser/Lexer.zig | 196 ++++ src/compiler/ir/parser/ast.zig | 93 ++ src/compiler/ir/parser/lower.zig | 380 ++++++++ src/compiler/ir/parser/parser.zig | 767 +++++++++++++++ src/compiler/ir/pass_manager.zig | 83 ++ src/compiler/ir/printer.zig | 336 +++++++ src/compiler/ir/type.zig | 106 ++ src/compiler/ir/validator/dominance.zig | 93 ++ src/compiler/ir/validator/validator.zig | 470 +++++++++ src/compiler/ir/value.zig | 31 + src/compiler/ir/visitor.zig | 98 ++ src/compiler/root.zig | 826 ++++++++++++++++ src/compiler/spirv/Parser.zig | 155 +++ src/compiler/spirv/root.zig | 16 + src/compiler/spirv/spirv.zig | 169 ++++ src/compiler/spirv/translator.zig | 1074 +++++++++++++++++++++ src/intel/FlintCommandBuffer.zig | 41 +- src/intel/FlintDescriptorPool.zig | 29 +- src/intel/FlintDescriptorSet.zig | 137 ++- src/intel/FlintImage.zig | 4 +- src/intel/FlintInstance.zig | 4 +- src/intel/FlintPhysicalDevice.zig | 12 +- src/intel/FlintPipeline.zig | 14 +- src/intel/FlintShaderModule.zig | 6 + src/intel/blitter.zig | 2 +- src/intel/i915/i915.zig | 1 - src/intel/kmd.zig | 6 +- src/intel/lib.zig | 10 +- src/phi/PhiCommandBuffer.zig | 2 +- src/phi/PhiInstance.zig | 2 +- src/phi/PhiPhysicalDevice.zig | 10 +- src/phi/PhiTransport.zig | 4 +- src/phi/lib.zig | 6 +- src/phi/pci_ids.zig | 1 - src/phi/scif.zig | 4 +- src/software/SoftBuffer.zig | 10 +- src/software/SoftCommandBuffer.zig | 16 +- src/software/SoftImage.zig | 2 +- src/software/SoftInstance.zig | 2 +- src/software/SoftPhysicalDevice.zig | 54 +- src/software/device/Device.zig | 22 +- src/software/device/clip.zig | 8 +- src/software/device/fragment.zig | 18 +- src/software/device/vertex_dispatcher.zig | 6 +- src/software/lib.zig | 40 +- src/vulkan/CommandBuffer.zig | 6 +- src/vulkan/CommandPool.zig | 6 +- src/vulkan/Instance.zig | 18 +- src/vulkan/PhysicalDevice.zig | 2 +- src/vulkan/PipelineCache.zig | 2 +- src/vulkan/PipelineLayout.zig | 10 +- src/vulkan/lib.zig | 12 +- src/vulkan/logger.zig | 6 +- src/vulkan/wsi/WaylandSurfaceKHR.zig | 2 +- src/vulkan/wsi/clients/wayland.zig | 38 +- test/test_runner.zig | 8 +- 70 files changed, 7358 insertions(+), 250 deletions(-) create mode 100644 src/compiler/README.md create mode 100644 src/compiler/ir/Builder.zig create mode 100644 src/compiler/ir/Rewriter.zig create mode 100644 src/compiler/ir/cfg.zig create mode 100644 src/compiler/ir/constant.zig create mode 100644 src/compiler/ir/id.zig create mode 100644 src/compiler/ir/instruction.zig create mode 100644 src/compiler/ir/ir.zig create mode 100644 src/compiler/ir/module.zig create mode 100644 src/compiler/ir/parser/Lexer.zig create mode 100644 src/compiler/ir/parser/ast.zig create mode 100644 src/compiler/ir/parser/lower.zig create mode 100644 src/compiler/ir/parser/parser.zig create mode 100644 src/compiler/ir/pass_manager.zig create mode 100644 src/compiler/ir/printer.zig create mode 100644 src/compiler/ir/type.zig create mode 100644 src/compiler/ir/validator/dominance.zig create mode 100644 src/compiler/ir/validator/validator.zig create mode 100644 src/compiler/ir/value.zig create mode 100644 src/compiler/ir/visitor.zig create mode 100644 src/compiler/root.zig create mode 100644 src/compiler/spirv/Parser.zig create mode 100644 src/compiler/spirv/root.zig create mode 100644 src/compiler/spirv/spirv.zig create mode 100644 src/compiler/spirv/translator.zig diff --git a/.gitea/workflows/Build.yml b/.gitea/workflows/Build.yml index 76ca3b0..e30f360 100644 --- a/.gitea/workflows/Build.yml +++ b/.gitea/workflows/Build.yml @@ -2,9 +2,9 @@ name: Build on: push: - branches: [ "master" ] + branches: ["master"] pull_request: - branches: [ "master" ] + branches: ["master"] jobs: build: @@ -52,6 +52,9 @@ jobs: echo "$PATH" ls -la /opt/mpss/3.8.6/sysroots/k1om-mpss-linux/bin | head -50 + - name: ZLint pass + uses: DonIsaac/zlint-action@v0.0.1 + - name: Building Ape run: zig build ape --release=safe @@ -79,3 +82,19 @@ jobs: local-path: "./zig-out/docs" remote-path: "/www" sync: full + + - name: Generating IR docs + run: zig build docs-ir + + - name: Deploying IR docs + if: ${{ github.event_name == 'push' && github.ref == 'refs/heads/master' }} + uses: milanmk/actions-file-deployer@master + with: + remote-protocol: sftp + remote-host: ${{ secrets.SFTP_HOST_DOCS_IR }} + remote-user: ${{ secrets.SFTP_USER_DOCS_IR }} + remote-password: ${{ secrets.SFTP_PASSWORD_DOCS_IR }} + remote-port: 6969 + local-path: "./zig-out/docs-ir" + remote-path: "/www" + sync: full diff --git a/.gitea/workflows/Test.yml b/.gitea/workflows/Test.yml index 5b20b27..e617410 100644 --- a/.gitea/workflows/Test.yml +++ b/.gitea/workflows/Test.yml @@ -2,9 +2,9 @@ name: Test on: push: - branches: [ "master" ] + branches: ["master"] pull_request: - branches: [ "master" ] + branches: ["master"] jobs: build_and_test: @@ -12,51 +12,57 @@ jobs: if: "!contains(github.event.head_commit.message, 'ci skip') && !contains(github.event.head_commit.message, 'test skip')" steps: - - uses: actions/checkout@v4 - - uses: https://codeberg.org/mlugg/setup-zig@v2 + - uses: actions/checkout@v4 + - uses: https://codeberg.org/mlugg/setup-zig@v2 - - name: Install MPSS - run: | - set -euxo pipefail + - name: Install MPSS + run: | + set -euxo pipefail - apt-get update - apt-get install -y unzip curl + apt-get update + apt-get install -y unzip curl - MPSS_URL="https://drive.kbz8.me/dav/public-files/KSKeiOKxnFPyfws/mpss.zip" - MPSS_ZIP="/tmp/mpss.zip" - MPSS_BIN="/opt/mpss/3.8.6/sysroots/k1om-mpss-linux/bin" + MPSS_URL="https://drive.kbz8.me/dav/public-files/KSKeiOKxnFPyfws/mpss.zip" + MPSS_ZIP="/tmp/mpss.zip" + MPSS_BIN="/opt/mpss/3.8.6/sysroots/k1om-mpss-linux/bin" - curl --fail --location --show-error --progress-bar "$MPSS_URL" -o "$MPSS_ZIP" + curl --fail --location --show-error --progress-bar "$MPSS_URL" -o "$MPSS_ZIP" - size="$(stat -c%s "$MPSS_ZIP")" - test "$size" -gt 300000000 + size="$(stat -c%s "$MPSS_ZIP")" + test "$size" -gt 300000000 - unzip -t "$MPSS_ZIP" >/dev/null + unzip -t "$MPSS_ZIP" >/dev/null - rm -rf /opt/mpss - unzip -q "$MPSS_ZIP" -d /opt - chmod -R a+rX /opt/mpss + rm -rf /opt/mpss + unzip -q "$MPSS_ZIP" -d /opt + chmod -R a+rX /opt/mpss - test -d "$MPSS_BIN" + test -d "$MPSS_BIN" - # Append as fallback for following CI steps. - echo "PATH=$PATH:$MPSS_BIN" >> "$GITHUB_ENV" + # Append as fallback for following CI steps. + echo "PATH=$PATH:$MPSS_BIN" >> "$GITHUB_ENV" - echo "MPSS installed." + echo "MPSS installed." - - name: Check MPSS PATH - run: | - echo "$PATH" - ls -la /opt/mpss/3.8.6/sysroots/k1om-mpss-linux/bin | head -50 + - name: Check MPSS PATH + run: | + echo "$PATH" + ls -la /opt/mpss/3.8.6/sysroots/k1om-mpss-linux/bin | head -50 - - name: Ape Tests - run: zig build test-ape --release=safe + - name: Installing dependencies + run: apt update && apt install spirv-tools - - name: Soft Tests - run: zig build test-soft --release=safe + - name: Ape Tests + run: zig build test-ape --release=safe - - name: Flint Tests - run: zig build test-flint --release=safe + - name: Soft Tests + run: zig build test-soft --release=safe - - name: Phi Tests - run: zig build test-phi --release=safe + - name: Flint Tests + run: zig build test-flint --release=safe + + - name: Phi Tests + run: zig build test-phi --release=safe + + - name: IR Tests + run: zig build test-ir --release=safe diff --git a/build.zig b/build.zig index 7d2f78a..eaaafa9 100644 --- a/build.zig +++ b/build.zig @@ -71,6 +71,37 @@ pub fn build(b: *std.Build) !void { const target = b.standardTargetOptions(.{}); const optimize = b.standardOptimizeOption(.{}); + const ir_mod = b.createModule(.{ + .root_source_file = b.path("src/compiler/root.zig"), + .target = target, + .optimize = optimize, + }); + + const ir_tests = b.addTest(.{ + .root_module = ir_mod, + .test_runner = .{ + .path = b.path("test/test_runner.zig"), + .mode = .simple, + }, + }); + const run_ir_tests = b.addRunArtifact(ir_tests); + const ir_test_step = b.step("test-ir", "Run shared shader ir tests"); + ir_test_step.dependOn(&run_ir_tests.step); + + const ir_autodoc_test = b.addObject(.{ + .name = "lib", + .root_module = ir_mod, + }); + + const ir_install_docs = b.addInstallDirectory(.{ + .source_dir = ir_autodoc_test.getEmittedDocs(), + .install_dir = .prefix, + .install_subdir = "docs-ir", + }); + + const ir_docs_step = b.step("docs-ir", "Build and install the documentation or shader IR"); + ir_docs_step.dependOn(&ir_install_docs.step); + const base_mod = b.createModule(.{ .root_source_file = b.path("src/vulkan/lib.zig"), .target = target, @@ -357,7 +388,7 @@ fn addMultithreadedCTS(b: *std.Build, target: std.Build.ResolvedTarget, impl: *c run.addArg("run"); run.addArg("--timeout"); - run.addArg("300"); + run.addArg("60"); run.addArg("--deqp"); run.addArg(cts_exe_path); run.addArg("--caselist"); @@ -470,17 +501,22 @@ fn optionsSoft(b: *std.Build, options: *Step.Options) !void { // Flint specialized functions fn customFlint( - _: *std.Build, + b: *std.Build, _: *Step.Compile, lib_mod: *std.Build.Module, _: *std.Build.Module, _: *std.Build.Module, base_c_mod: *std.Build.Module, - _: std.Build.ResolvedTarget, - _: std.builtin.OptimizeMode, + target: std.Build.ResolvedTarget, + optimize: std.builtin.OptimizeMode, _: bool, ) !void { lib_mod.addImport("intel_c", base_c_mod); + lib_mod.addImport("shader_compiler", b.createModule(.{ + .root_source_file = b.path("src/compiler/root.zig"), + .target = target, + .optimize = optimize, + })); } fn optionsFlint(b: *std.Build, options: *Step.Options) !void { diff --git a/src/ape/ApeInstance.zig b/src/ape/ApeInstance.zig index 96e8de6..6724671 100644 --- a/src/ape/ApeInstance.zig +++ b/src/ape/ApeInstance.zig @@ -14,7 +14,7 @@ pub const Interface = base.Instance; interface: Interface, backend_instances: std.ArrayList(*Interface), -pub const EXTENSIONS = soft.Instance.EXTENSIONS; +pub const extensions = soft.Instance.extensions; pub fn create(allocator: std.mem.Allocator, infos: *const vk.InstanceCreateInfo) VkError!*Interface { const self = allocator.create(Self) catch return VkError.OutOfHostMemory; diff --git a/src/ape/lib.zig b/src/ape/lib.zig index 06555ba..facf48f 100644 --- a/src/ape/lib.zig +++ b/src/ape/lib.zig @@ -10,9 +10,9 @@ pub const ApeInstance = @import("ApeInstance.zig"); pub const Instance = ApeInstance; -pub const DRIVER_NAME = "Ape"; +pub const driver_name = "Ape"; -pub const VULKAN_VERSION = vk.makeApiVersion( +pub const vulkan_version = vk.makeApiVersion( 0, config.ape_vulkan_version.major, config.ape_vulkan_version.minor, diff --git a/src/compiler/README.md b/src/compiler/README.md new file mode 100644 index 0000000..7382e1c --- /dev/null +++ b/src/compiler/README.md @@ -0,0 +1,612 @@ +# Backend-Agnostic Shader IR + +> Note: this IR is still foundational and incomplete. Its representation may +> change as the compiler gains features. Backends should not treat it as a +> stable ABI yet. + +This directory contains the backend-agnostic shader intermediate representation. +It sits between SPIR-V and the backends that consume it. +Format-specific details are removed while types, values, control flow, +interfaces, and semantic operations remain. + +## Overview + +- One `Module` describes one selected shader entry point. +- The current stages are `vertex`, `fragment`, and `compute`. +- Values are typed and use static single assignment (SSA). +- Blocks act as control-flow graph labels, but also own instructions, + parameters, a terminator, and structured-control metadata. +- Block parameters serve the same purpose as `OpPhi` in SPIR-V. +- Instructions speak normalized meanings, not source-format opcodes. +- Types and constants are interned. Equal things share one identity. +- IDs are stable. Erasure leaves a tombstone; no dead ID is reused for another object. +- The printer is for debugging and tests. Its output can be parsed back into a + validated module, but it is not yet a stable interchange format. + +## Printer syntax + +The printer uses these prefixes: + +| Prefix | Meaning | Example | +| ------- | ---------------------------------------------------------------- | --------------------- | +| `%id` | An SSA value, whether constant, parameter, or instruction result | `%3`, `%merged_value` | +| `@name` | A function or interface declaration | `@main`, `@out_color` | +| `.name` | A basic block | `.entry`, `.merge` | +| `#N` | A constant-store identity used within composite constants | `#2` | + +Names are annotations rather than identity. When a name is absent, invalid for +the textual grammar, or duplicated, the printer uses a numeric `%N` value +reference. Other unnamed objects use forms such as `@fn0`, `@interface1`, and +`.b2`. The parser accepts both numeric and identifier-shaped value references. + +An instruction that produces a value prints its result type explicitly: + +```text +%result: = +``` + +The type annotation makes operations such as `bitcast` and heterogeneous +`composite_construct` unambiguous when the text is parsed. Constants, function +parameters, and block parameters already carry their types in their own forms. + +The outer structure has this shape: + +```text +shader @ +{ + + + + fn @() -> + { + .(): + + + + } +} +``` + +Execution modes, resources, source locations, and structured-control metadata +exist in memory, but the printer does not display them yet. + +## Parsing + +`ir.parser` accepts the complete syntax emitted by the printer: + +```zig +var from_string = try ir.parser.parseString(allocator, source); +defer from_string.deinit(); + +var from_file = try ir.parser.parseFile(allocator, io, "shader.ir"); +defer from_file.deinit(); +``` + +Use `parseFileInDir` when the path is relative to an existing `std.Io.Dir`. +Each parser entry point owns the returned module with the supplied allocator and +runs the IR validator before returning it. Parse, reference-resolution, or +validation failures are returned as errors. Because the printer omits the +metadata listed above, a print/parse round trip preserves the displayed IR but +cannot recover those hidden fields. + +## Types + +Types are interned in the module and printed inline; their `TypeId` is hidden. + +| Kind | Printed form | Meaning | +| ---------------- | ------------------------------- | ------------------------------------------------------------------------- | +| Void | `void` | No value. Used mainly for functions that return nothing. | +| Boolean | `bool` | A truth value. | +| Signed integer | `i32` | A signed integer of the written bit width. | +| Unsigned integer | `u32` | An unsigned integer of the written bit width. | +| Floating point | `f32` | A floating value of the written bit width. | +| Vector | `vec4[f32]` | A fixed number of equal scalar elements. Its length must be at least two. | +| Array | `array[u32, 8]` | A fixed number of equal elements. Its length must not be zero. | +| Structure | `struct[f32, vec4[f32]]` | An ordered sequence of potentially different member types. | +| Pointer | `ptr[workgroup, u32]` | A pointer to a type within an address space. | +| Resource handle | `resourceHandle[sampled_image]` | An opaque handle for later resource operations. | + +The current address spaces are `function`, `private`, `workgroup`, +`input`, `output`, `uniform`, `storage`, `push_constant`, and `physical`. + +The current resource kinds are `uniform_buffer`, `storage_buffer`, +`sampled_image`, `storage_image`, and `sampler`. A resource handle may also +carry an optional data type in memory, although the printer omits that type. + +## Constants + +Constants live at module scope and also have ordinary numeric or named `%id` value identities. +The parser accepts direct decimal integer and floating-point values, including +signed values and floating-point exponents, as source-level convenience syntax: + +```text +%0: constant u32 = 42 +%1: constant i32 = -7 +%2: constant f32 = 1.5e2 +``` + +Direct integers must fit their declared width and signedness. Direct floats are +rounded to their declared `f16`, `f32`, or `f64` representation. The canonical +printer always emits integer and float bit patterns so reparsing cannot silently +change the stored value. + +```text +%id: constant = +``` + +| Form | Meaning | Printed example | +| ------------ | ----------------------------------- | --------------------------------------- | +| Boolean | `true` or `false` | `%0: constant bool = true` | +| Integer bits | The fixed-width integer bit pattern | `%1: constant u32 = bits(0x2a)` | +| Float bits | The IEEE-like bit pattern as stored | `%2: constant f32 = bits(0x3f800000)` | +| Null | The null value of its type | `%3: constant ptr[private, u32] = null` | +| Undef | An unconstrained value | `%4: constant u32 = undef` | +| Composite | A sequence of other constants | `%5: constant vec2[u32] = [#1, #1]` | + +## Functions, blocks, and SSA + +A function owns typed parameters, an ordered list of blocks, one entry block, +and one return type. The first block created by the builder becomes the entry +block. The validator requires the entry block to have no predecessor. + +Every block must end in exactly one terminator. A value defined by an instruction +must dominate every use, and within one block it must be written before it is +used. Constants and standalone undef values are module-wide; function and block +parameters cannot be used by another function. + +### Block parameters and Phi lowering + +A merge block does not contain a `phi` instruction. Instead, it declares a +parameter, and every incoming edge passes one argument of the same type: + +```text +.left(): + branch .merge(%3) + +.right(): + branch .merge(%4) + +.merge(%5: u32): + %6: u32 = integer_multiply %5, %2 + return +``` + +`%5` therefore receives the value supplied by the selected edge. The number and +types of edge arguments must exactly match the target block's parameters. + +### Structured control + +A block may carry one of these unprinted metadata values: + +- `none`: no structured-control promise. +- `selection`: names one merge block. +- `loop`: names both merge and continue blocks. + +The SPIR-V translator preserves `OpSelectionMerge` and `OpLoopMerge` in this +metadata. They are not terminators and do not create graph edges themselves. + +## Common instruction rules + +An instruction belongs to one block, has zero or one result, and may carry a +source location. Except for `store_interface` and `call`, current operations are +treated as side-effect free by the rewriter. A block's terminator is stored +separately from its ordinary instructions. + +Most arithmetic operations are intended for scalars or vectors of their named +category and act component by component where vectors are allowed. The current +foundational validator often checks only that operand and result types match. +The stricter integer, float, boolean, bit-width, and vector-shape requirements +below describe semantic intent and still need more complete validation. + +## Unary opcodes + +Form: + +```text +%result: = %operand +``` + +| Opcode | Arity | Description | Usage | Small printed example | +| ------------- | ----: | ---------------------------- | ----------------------------------------------------------------- | --------------------------- | +| `negate` | 1 | Changes the arithmetic sign. | Signed integer or floating operand; the result has the same type. | `%2: i32 = negate %1` | +| `logical_not` | 1 | Inverts a boolean value. | Boolean operand and boolean result. | `%2: bool = logical_not %1` | +| `bitwise_not` | 1 | Inverts every bit. | Integer operand; the result has the same type. | `%2: u32 = bitwise_not %1` | + +`negate` is one normalized opcode: the operand type distinguishes integer +negation from floating negation. + +## Binary opcodes + +Form: + +```text +%result: = %lhs, %rhs +``` + +The two operands and result currently must have the same IR type. + +### Integer arithmetic + +| Opcode | Description | Usage | Small printed example | +| ------------------ | ------------------------------------------------------- | --------------------------------------------------------------------------- | ----------------------------------- | +| `integer_add` | Adds fixed-width integers. | Integer operands of one type. | `%3: u32 = integer_add %1, %2` | +| `integer_subtract` | Subtracts the right operand from the left. | Integer operands of one type. | `%3: u32 = integer_subtract %1, %2` | +| `integer_multiply` | Multiplies fixed-width integers. | Integer operands of one type. | `%3: u32 = integer_multiply %1, %2` | +| `unsigned_divide` | Divides unsigned integers. | Unsigned integer operands. | `%3: u32 = unsigned_divide %1, %2` | +| `signed_divide` | Divides signed integers. | Signed integer operands. | `%3: i32 = signed_divide %1, %2` | +| `unsigned_modulo` | Produces the unsigned remainder. | Unsigned integer operands. | `%3: u32 = unsigned_modulo %1, %2` | +| `signed_modulo` | Produces signed modulo, whose sign follows the divisor. | Signed integer operands; this corresponds to SPIR-V `OpSMod`, not `OpSRem`. | `%3: i32 = signed_modulo %1, %2` | + +Integer addition, subtraction, and multiplication are signedness-neutral at the +opcode level; the type retains signedness. Exceptional division, overflow, +and poison rules are not yet separately recorded by the IR. + +### Floating arithmetic + +| Opcode | Description | Usage | Small printed example | +| ---------------- | ------------------------------------------------------------ | ------------------------------------------------------------------- | --------------------------------- | +| `float_add` | Adds floating-point values. | Floating operands of one type. | `%3: f32 = float_add %1, %2` | +| `float_subtract` | Subtracts the right operand from the left. | Floating operands of one type. | `%3: f32 = float_subtract %1, %2` | +| `float_multiply` | Multiplies floating-point values. | Floating operands of one type. | `%3: f32 = float_multiply %1, %2` | +| `float_divide` | Divides the left operand by the right. | Floating operands of one type. | `%3: f32 = float_divide %1, %2` | +| `float_modulo` | Produces floating modulo, similar to `x - y * floor(x / y)`. | Floating operands of one type; this corresponds to SPIR-V `OpFMod`. | `%3: f32 = float_modulo %1, %2` | + +Instructions do not yet carry fast-math flags, rounding modes, contraction +permission, or NaN guarantees. + +### Shifts and bitwise arithmetic + +| Opcode | Description | Usage | Small printed example | +| ------------------------ | -------------------------------------------------- | ------------------------------------------------------------- | ----------------------------------------- | +| `shift_left` | Shifts bits left and fills the low bits with zero. | Integer operands; the right operand supplies the shift count. | `%3: u32 = shift_left %1, %2` | +| `logical_shift_right` | Shifts right and fills high bits with zero. | Integer operands interpreted without sign extension. | `%3: u32 = logical_shift_right %1, %2` | +| `arithmetic_shift_right` | Shifts right while repeating the sign bit. | Signed integer value and integer shift count. | `%3: i32 = arithmetic_shift_right %1, %2` | +| `bitwise_and` | Keeps bits set in both operands. | Integer operands of one type. | `%3: u32 = bitwise_and %1, %2` | +| `bitwise_or` | Keeps bits set in either operand. | Integer operands of one type. | `%3: u32 = bitwise_or %1, %2` | +| `bitwise_xor` | Keeps bits set in exactly one operand. | Integer operands of one type. | `%3: u32 = bitwise_xor %1, %2` | + +The validator currently requires the shift count to have the same complete IR +type as the shifted value. More flexible shift typing is not implemented yet. + +### Boolean conjunction + +| Opcode | Description | Usage | Small printed example | +| ------------- | ----------------------------------------- | ---------------------------- | ------------------------------- | +| `logical_and` | Is true only when both operands are true. | Boolean operands and result. | `%3: bool = logical_and %1, %2` | +| `logical_or` | Is true when either operand is true. | Boolean operands and result. | `%3: bool = logical_or %1, %2` | + +## Comparison opcodes + +All comparisons are printed with the `cmp_` prefix as one opcode token: + +```text +%result: bool = cmp_ %lhs, %rhs +``` + +The operands must share one type, and the current validator requires the result +to be the scalar `bool` type. + +| Opcode | Description | Usage | Small printed example | +| ------------------------------- | ------------------------------------------------------- | ---------------------------------------- | ------------------------------------------------- | +| `cmp_equal` | Tests whether two booleans or integers are equal. | Equal-typed boolean or integer operands. | `%3: bool = cmp_equal %1, %2` | +| `cmp_not_equal` | Tests whether two booleans or integers differ. | Equal-typed boolean or integer operands. | `%3: bool = cmp_not_equal %1, %2` | +| `cmp_unsigned_less` | Compares integer bit patterns as unsigned. | Unsigned integer operands. | `%3: bool = cmp_unsigned_less %1, %2` | +| `cmp_signed_less` | Compares integers as signed. | Signed integer operands. | `%3: bool = cmp_signed_less %1, %2` | +| `cmp_ordered_float_equal` | Is true when neither operand is NaN and they are equal. | Floating operands. | `%3: bool = cmp_ordered_float_equal %1, %2` | +| `cmp_unordered_float_equal` | Is true when either operand is NaN, or they are equal. | Floating operands. | `%3: bool = cmp_unordered_float_equal %1, %2` | +| `cmp_ordered_float_not_equal` | Is true when neither operand is NaN and they differ. | Floating operands. | `%3: bool = cmp_ordered_float_not_equal %1, %2` | +| `cmp_unordered_float_not_equal` | Is true when either operand is NaN, or they differ. | Floating operands. | `%3: bool = cmp_unordered_float_not_equal %1, %2` | +| `cmp_ordered_float_less` | Is true when neither operand is NaN and left is less. | Floating operands. | `%3: bool = cmp_ordered_float_less %1, %2` | +| `cmp_unordered_float_less` | Is true when either operand is NaN, or left is less. | Floating operands. | `%3: bool = cmp_unordered_float_less %1, %2` | + +There are no greater-than opcodes in the current instruction set. Swap the +operands and use the appropriate less-than form. Less-or-equal forms are also +not defined yet. + +## Other opcodes + +### `select` + +Selects one of two equal-typed values using a boolean condition. It does not +change control flow. + +```text +%4: u32 = select %1, %2, %3 +``` + +Here `%1` is `bool`; `%2`, `%3`, and `%4` share one type. + +### `bitcast` + +Reinterprets an operand's bits as the result type without performing a numeric +conversion. + +```text +%2: f32 = bitcast %1 +``` + +The intended source and destination have equal total bit width. The current +validator only requires that both values exist; it does not yet prove equal +width. + +### `composite_construct` + +Constructs a vector or structure from its immediate elements. + +```text +%5: vec4[f32] = composite_construct %1, %2, %3, %4 +``` + +For a vector, every element must have the vector's element type and their count +must equal its length. For a structure, each element must match the member at +the same position. The validator does not support array construction yet. + +### `composite_extract` + +Traverses one or more literal indices through a vector, array, or structure and +returns the selected nested member. + +```text +%4: f32 = composite_extract %3[1][0] +``` + +At least one index is required. Every index must lie within its composite, and +the result type must equal the selected member type. + +### `load_interface` + +Reads one declared shader input. It cannot read an output declaration. + +```text +%1: vec4[f32] = load_interface @in_color +``` + +The result type must equal the interface variable's type. An optional dynamic +`element_index` exists in memory for future arrayed interfaces, but the current +printer does not show it and the validator does not use it to change the +result type. + +### `store_interface` + +Writes one declared shader output. It produces no SSA result and cannot write an +input declaration. + +```text +store_interface @out_color, %1 +``` + +The stored value must equal the interface variable's type. As with +`load_interface`, an optional unprinted `element_index` is reserved for later +arrayed-interface work. This operation has side effects. + +### `call` + +Invokes another IR function. Arguments must match the callee's parameters in +number, order, and type. + +```text +%4: vec4[f32] = call @shade(%1, %2) +call @observe(%4) +``` + +A non-void callee requires a result of its return type; a void callee forbids +one. Calls are conservatively treated as side-effecting. The operation exists in +the common IR, although the current SPIR-V translator rejects +`OpFunctionCall`. + +## Terminators + +Terminators yield no ordinary instruction result. They alone determine outgoing +control-flow edges. + +| Terminator | Description | Usage | Small printed example | +| -------------------- | ----------------------------------------------- | ------------------------------------------------------------------------------- | ------------------------------------------ | +| `branch` | Unconditionally transfers control to one block. | Pass exactly one argument for every target parameter. | `branch .merge(%3)` | +| `conditional_branch` | Selects one of two edges using a boolean. | The condition is `bool`; each edge independently matches its target parameters. | `conditional_branch %1, .yes(%2), .no(%3)` | +| `return` (void) | Ends a void function. | The enclosing return type is `void`. | `return` | +| `return` (value) | Ends a function and returns a value. | The value type equals the function return type. | `return %3` | +| `discard` | Discards the fragment invocation. | Fragment stage only. | `discard` | +| `unreachable` | States that execution cannot reach this point. | Any function; no successors. | `unreachable` | + +The Zig union names the return forms `return_void` and `return_value`; the +printer renders both as the overloaded `return` spelling shown above. + +## Interfaces and builtins + +An interface variable has a type, a direction (`input` or `output`), and one +semantic. Its semantic attributes are enclosed in the direction's brackets: + +- A location: `location(N), component(C), index(I)`. +- A builtin: `builtin(name)`. + +The currently supported builtins are `position`, `vertex_index`, `instance_index`, +`frag_coord`, `frag_depth`, and `global_invocation_id`. + +Printed declarations resemble these: + +```text +@in_color: vec4[f32] = input[location(0), component(0), index(0)] +@position: vec4[f32] = output[builtin(position)] +``` + +## Complete examples + +These examples use the printer's exact grammar and indentation. Their numeric +value IDs are illustrative but follow the same module-wide numbering used by +the printer. + +### A compute shader that adds two constants + +```text +shader compute @main +{ + %0: constant u32 = bits(0x1) + %1: constant u32 = bits(0x2) + + fn @main() -> void + { + .entry(): + %2: u32 = integer_add %0, %1 + return + + } +} +``` + +### A vertex interface passed through + +```text +shader vertex @main +{ + @in_color: vec4[f32] = input[location(0), component(0), index(0)] + @out_color: vec4[f32] = output[location(0), component(0), index(0)] + + fn @main() -> void + { + .entry(): + %0: vec4[f32] = load_interface @in_color + store_interface @out_color, %0 + return + + } +} +``` + +### A selection whose Phi becomes a block parameter + +```text +shader compute @main +{ + %0: constant bool = true + %1: constant u32 = bits(0x1) + %2: constant u32 = bits(0x2) + + fn @main() -> void + { + .entry(): + conditional_branch %0, .left(), .right() + + .left(): + %3: u32 = integer_add %1, %2 + branch .merge(%3) + + .right(): + %4: u32 = integer_subtract %2, %1 + branch .merge(%4) + + .merge(%5: u32): + %6: u32 = integer_multiply %5, %2 + return + + } +} +``` + +The selection's merge metadata is not visible in this output, although it +remains attached to `.entry` in memory. + +## Validator guarantees + +When `validator.validate` succeeds, it has proved the following: + +- The module has a live entry point. +- All referenced types, constants, values, functions, blocks, and instructions + are live. +- Parent links and SSA definition links agree in both directions. +- Every function has an entry block, and no edge targets that entry block. +- Every block has a terminator. +- Edges remain within their function and exactly match target block parameters. +- Returns agree with function return types; `discard` appears only in a fragment shader. +- Interface loads read inputs; interface stores write outputs. +- Operation-specific result presence and the foundational type equalities hold. +- Structured merge and continue targets belong to the same function. + +The validator does not yet prove every semantic category listed in the opcode +reference. In particular, several arithmetic opcodes can currently be built +with an inappropriate but equal operand type; bitcast widths are not compared; +shift-count rules are rudimentary; and floating-point execution modes are not +attached to operations. Backends should explicitly require the properties and +validation needed by their lowering. + +## Properties and passes + +The module carries independent property bits: + +- `valid_cfg` +- `valid_ssa` +- `structured_control_flow` +- `no_function_calls` +- `no_local_memory` +- `no_matrix_types` +- `no_large_composites` +- `explicit_resource_offsets` + +A pass declares properties that it requires, produces, and invalidates. The pass +manager rejects a pass whose requirements are missing, applies its property +changes, and runs the validator after every pass by default. + +## Builder, rewriter, and visitor + +`Builder.zig` is the normal entry point for construction. It interns types and +constants, stores copied slices in the module arena, adds functions and blocks, +appends instructions, and assigns block terminators. + +`parser/root.zig` owns the public parsing entry points and recursive-descent grammar. +Its implementation details are split by responsibility: `parser/Lexer.zig` +tokenizes input, `parser/ast.zig` holds the temporary syntax model, and +`parser/lower.zig` resolves that model into the common IR. + +`Rewriter.zig` provides the first safe mutation operations: + +- Count and replace SSA uses without changing definitions. +- Erase a dead, side-effect-free instruction. +- Redirect edges with a complete new argument list. +- Add a block parameter while adding every incoming edge argument. +- Remove a block parameter while removing its incoming arguments. + +`visitor.zig` walks module declarations, functions, blocks, instructions, +terminators, and SSA uses in hierarchical order. `cfg.zig` computes +predecessors, reachability, and dominance with a deliberately simple quadratic +matrix. `validator/dominance.zig` contains the SSA dominance checks built on +that analysis. This is suitable for the foundation but is not intended as the +final large-shader implementation. + +## SPIR-V frontend + +The compiler currently provides a word parser and an initial translator in +`spirv/`. The parser validates the header, word counts, truncation, and literal +strings. The translator selects one entry point and lowers a defined subset: + +- Vertex, fragment, and compute stages. +- Basic scalar, vector, array, structure, pointer, and function types. +- Ordinary and composite constants; unapplied specialization constants are + refused. +- Functions, blocks, branches, structured merge marks, and returns. +- `OpPhi` into block parameters and edge arguments. +- The arithmetic, comparison, select, bitcast, and composite operations named + in the reference above where mappings currently exist. +- Decorated stage inputs and outputs, with interface load and store. +- `OpName` debug names for functions, blocks, parameters, constants, and + instruction results when they are valid textual IR identifiers. + +Symbolic identifiers in SPIR-V assembly are assembler syntax and are not stored +in the binary by `spirv-as`; add `OpName` instructions when those names must +survive translation. Unsupported source instructions return an error; they are +not preserved as opaque SPIR-V. This prevents silent mistranslation. + +## Running tests + +From the repository root, run: + +```sh +zig build test-ir +``` + +The SPIR-V translation tests keep their assembly as multiline strings beside +their assertions and pipe it through `spirv-as`. Therefore SPIRV-Tools must be +on `PATH` when those tests run. Only malformed-binary parser tests use raw words, +because an assembler cannot produce intentionally malformed instructions. + +## Documentation + +A complete codebase documentation can be found [here](https://vulkan-driver.kbz8.me/docs/ir/). diff --git a/src/compiler/ir/Builder.zig b/src/compiler/ir/Builder.zig new file mode 100644 index 0000000..8183b64 --- /dev/null +++ b/src/compiler/ir/Builder.zig @@ -0,0 +1,227 @@ +const std = @import("std"); +const ids = @import("id.zig"); +const type_ir = @import("type.zig"); +const constant_ir = @import("constant.zig"); +const inst_ir = @import("instruction.zig"); +const module_ir = @import("module.zig"); + +const Self = @This(); + +module: *module_ir.Module, + +pub fn init(module: *module_ir.Module) Self { + return .{ .module = module }; +} + +fn copyName(self: *Self, name: ?[]const u8) !?[]const u8 { + return if (name) |text| try self.module.allocator().dupe(u8, text) else null; +} + +pub fn internType(self: *Self, ty: type_ir.Type) !ids.TypeId { + return self.module.internType(ty); +} + +pub fn internConstant(self: *Self, ty: ids.TypeId, candidate: constant_ir.ConstantValue) !ids.ValueId { + for (self.module.constants.entries.items, 0..) |entry, index| { + const existing = entry orelse continue; + + if (existing.type == ty and constantEql(existing.value, candidate)) { + const constant_id = ids.ConstantId.fromIndex(index); + + for (self.module.values.entries.items, 0..) |value_entry, value_index| { + const value = value_entry orelse continue; + + if (value.definition == .constant and value.definition.constant == constant_id) + return ids.ValueId.fromIndex(value_index); + } + } + } + + var owned = candidate; + if (candidate == .composite) + owned.composite = try self.module.allocator().dupe(ids.ConstantId, candidate.composite); + + const constant_id = try self.module.constants.add(self.module.allocator(), .{ .type = ty, .value = owned }); + return self.module.values.add(self.module.allocator(), .{ + .type = ty, + .definition = .{ .constant = constant_id }, + }); +} + +pub fn setValueName(self: *Self, value_id: ids.ValueId, name: ?[]const u8) !void { + const text = name orelse return; + const value = self.module.values.getMut(value_id) orelse return error.InvalidValue; + + if (value.name == null) + value.name = try self.copyName(text); +} + +pub fn addFunction(self: *Self, return_type: ids.TypeId, name: ?[]const u8) !ids.FunctionId { + return self.module.functions.add(self.module.allocator(), .{ + .return_type = return_type, + .name = try self.copyName(name), + }); +} + +pub fn setEntryPoint(self: *Self, function: ids.FunctionId) void { + self.module.entry_point = function; +} + +pub fn addFunctionParameter(self: *Self, function_id: ids.FunctionId, ty: ids.TypeId, name: ?[]const u8) !ids.ValueId { + const function = self.module.functions.getMut(function_id) orelse return error.InvalidFunction; + const index: u32 = @intCast(function.parameters.items.len); + + const value_id = try self.module.values.add(self.module.allocator(), .{ + .type = ty, + .definition = .{ .function_parameter = .{ .function = function_id, .index = index } }, + .name = try self.copyName(name), + }); + + try function.parameter_types.append(self.module.allocator(), ty); + try function.parameters.append(self.module.allocator(), value_id); + + return value_id; +} + +pub fn addBlock(self: *Self, function_id: ids.FunctionId, name: ?[]const u8) !ids.BlockId { + const function = self.module.functions.getMut(function_id) orelse return error.InvalidFunction; + + const block_id = try self.module.blocks.add(self.module.allocator(), .{ + .parent_function = function_id, + .name = try self.copyName(name), + }); + + try function.blocks.append(self.module.allocator(), block_id); + if (function.entry_block == null) + function.entry_block = block_id; + + return block_id; +} + +pub fn addBlockParameter(self: *Self, block_id: ids.BlockId, ty: ids.TypeId, name: ?[]const u8) !ids.ValueId { + const block = self.module.blocks.getMut(block_id) orelse return error.InvalidBlock; + const index: u32 = @intCast(block.parameters.items.len); + + const value_id = try self.module.values.add(self.module.allocator(), .{ + .type = ty, + .definition = .{ .block_parameter = .{ .block = block_id, .index = index } }, + .name = try self.copyName(name), + }); + + try block.parameters.append(self.module.allocator(), value_id); + + return value_id; +} + +pub fn appendInstruction( + self: *Self, + block_id: ids.BlockId, + result_type: ?ids.TypeId, + operation: inst_ir.Operation, + name: ?[]const u8, +) !?ids.ValueId { + const block = self.module.blocks.getMut(block_id) orelse return error.InvalidBlock; + const owned_operation = try self.copyOperation(operation); + + const instruction_id = try self.module.instructions.add(self.module.allocator(), .{ + .parent_block = block_id, + .result = null, + .operation = owned_operation, + }); + errdefer _ = self.module.instructions.remove(instruction_id); + + const result = if (result_type) |ty| + try self.module.values.add(self.module.allocator(), .{ + .type = ty, + .definition = .{ .instruction = instruction_id }, + .name = try self.copyName(name), + }) + else + null; + + self.module.instructions.getMut(instruction_id).?.result = result; + try block.instructions.append(self.module.allocator(), instruction_id); + + return result; +} + +pub fn setTerminator(self: *Self, block_id: ids.BlockId, terminator: module_ir.Terminator) !void { + const block = self.module.blocks.getMut(block_id) orelse return error.InvalidBlock; + + if (block.terminator != null) + return error.TerminatorAlreadySet; + + block.terminator = try self.copyTerminator(terminator); +} + +pub fn addInterfaceVariable( + self: *Self, + ty: ids.TypeId, + direction: module_ir.InterfaceDirection, + semantic: module_ir.InterfaceSemantic, + name: ?[]const u8, +) !ids.InterfaceVariableId { + return self.module.interface_variables.add(self.module.allocator(), .{ + .type = ty, + .direction = direction, + .semantic = semantic, + .name = try self.copyName(name), + }); +} + +pub fn edge(self: *Self, target: ids.BlockId, arguments: []const ids.ValueId) !module_ir.Edge { + return .{ + .target = target, + .arguments = try self.module.allocator().dupe(ids.ValueId, arguments), + }; +} + +fn copyOperation(self: *Self, operation: inst_ir.Operation) !inst_ir.Operation { + return switch (operation) { + .composite_construct => |op| .{ + .composite_construct = .{ + .elements = try self.module.allocator().dupe(ids.ValueId, op.elements), + }, + }, + .composite_extract => |op| .{ + .composite_extract = .{ + .composite = op.composite, + .indices = try self.module.allocator().dupe(u32, op.indices), + }, + }, + .call => |op| .{ + .call = .{ + .function = op.function, + .arguments = try self.module.allocator().dupe(ids.ValueId, op.arguments), + }, + }, + else => operation, + }; +} + +fn copyTerminator(self: *Self, terminator: module_ir.Terminator) !module_ir.Terminator { + return switch (terminator) { + .branch => |edge_value| .{ + .branch = try self.edge(edge_value.target, edge_value.arguments), + }, + .conditional_branch => |branch| .{ + .conditional_branch = .{ + .condition = branch.condition, + .true_edge = try self.edge(branch.true_edge.target, branch.true_edge.arguments), + .false_edge = try self.edge(branch.false_edge.target, branch.false_edge.arguments), + }, + }, + else => terminator, + }; +} + +fn constantEql(a: constant_ir.ConstantValue, b: constant_ir.ConstantValue) bool { + return switch (a) { + .boolean => |value| b == .boolean and value == b.boolean, + .integer_bits => |value| b == .integer_bits and value == b.integer_bits, + .float_bits => |value| b == .float_bits and value == b.float_bits, + .null => b == .null, + .undef => b == .undef, + .composite => |value| b == .composite and std.mem.eql(ids.ConstantId, value, b.composite), + }; +} diff --git a/src/compiler/ir/Rewriter.zig b/src/compiler/ir/Rewriter.zig new file mode 100644 index 0000000..deffd00 --- /dev/null +++ b/src/compiler/ir/Rewriter.zig @@ -0,0 +1,323 @@ +const std = @import("std"); +const ids = @import("id.zig"); +const module_ir = @import("module.zig"); +const Builder = @import("Builder.zig"); + +const Self = @This(); + +pub const Error = std.mem.Allocator.Error || error{ + InvalidValue, + InvalidInstruction, + InvalidBlock, + InvalidFunction, + TypeMismatch, + ResultStillUsed, + SideEffectingInstruction, + InstructionNotOwnedByBlock, + InvalidParameterIndex, + MissingIncomingValue, + UnexpectedIncomingValue, +}; + +pub const IncomingValue = struct { + predecessor: ids.BlockId, + value: ids.ValueId, +}; + +const UseCountContext = struct { + needle: ids.ValueId, + count: usize = 0, +}; + +module: *module_ir.Module, + +pub fn init(module: *module_ir.Module) Self { + return .{ .module = module }; +} + +pub fn countUses(self: *const Self, value: ids.ValueId) usize { + var context: UseCountContext = .{ .needle = value }; + for (self.module.instructions.entries.items) |entry| { + const instruction = entry orelse continue; + instruction.operation.visitValueUses(&context, countUse); + } + + for (self.module.blocks.entries.items) |entry| { + const block = entry orelse continue; + if (block.terminator) |terminator| + module_ir.visitTerminatorValueUses(terminator, &context, countUse); + } + + return context.count; +} + +pub fn replaceAllUses(self: *Self, old: ids.ValueId, replacement: ids.ValueId) Error!usize { + const old_value = self.module.values.get(old) orelse return error.InvalidValue; + const replacement_value = self.module.values.get(replacement) orelse return error.InvalidValue; + + if (old_value.type != replacement_value.type) + return error.TypeMismatch; + + if (old == replacement) + return 0; + + var count: usize = 0; + for (self.module.instructions.entries.items) |*entry| { + const instruction = if (entry.*) |*value| value else continue; + count += try instruction.operation.replaceValueUses(self.module.allocator(), old, replacement); + } + + for (self.module.blocks.entries.items) |*entry| { + const block = if (entry.*) |*value| value else continue; + if (block.terminator) |*terminator| + count += try module_ir.replaceTerminatorValueUses(self.module.allocator(), terminator, old, replacement); + } + + return count; +} + +pub fn eraseInstruction(self: *Self, instruction_id: ids.InstructionId) Error!void { + const instruction = self.module.instructions.get(instruction_id) orelse return error.InvalidInstruction; + + if (instruction.operation.hasSideEffects()) + return error.SideEffectingInstruction; + + if (instruction.result) |result| { + if (self.countUses(result) != 0) + return error.ResultStillUsed; + } + + const block = self.module.blocks.getMut(instruction.parent_block) orelse return error.InvalidBlock; + var owned_index: ?usize = null; + + for (block.instructions.items, 0..) |candidate, index| { + if (candidate == instruction_id) { + owned_index = index; + break; + } + } + + _ = block.instructions.orderedRemove(owned_index orelse return error.InstructionNotOwnedByBlock); + if (instruction.result) |result| + _ = self.module.values.remove(result); + _ = self.module.instructions.remove(instruction_id); +} + +pub fn redirectEdges( + self: *Self, + source: ids.BlockId, + old_target: ids.BlockId, + new_target: ids.BlockId, + new_arguments: []const ids.ValueId, +) Error!usize { + const source_block = self.module.blocks.get(source) orelse return error.InvalidBlock; + const target_block = self.module.blocks.get(new_target) orelse return error.InvalidBlock; + + if (source_block.parent_function != target_block.parent_function) + return error.InvalidFunction; + + try self.validateArguments(target_block, new_arguments); + + const mutable_source = self.module.blocks.getMut(source).?; + const terminator = if (mutable_source.terminator) |*value| value else return error.InvalidBlock; + + var count: usize = 0; + + switch (terminator.*) { + .branch => |*edge| { + if (try self.redirectOne(edge, old_target, new_target, new_arguments)) + count += 1; + }, + .conditional_branch => |*branch| { + if (try self.redirectOne(&branch.true_edge, old_target, new_target, new_arguments)) + count += 1; + + if (try self.redirectOne(&branch.false_edge, old_target, new_target, new_arguments)) + count += 1; + }, + else => {}, + } + + return count; +} + +pub fn addBlockParameter( + self: *Self, + block_id: ids.BlockId, + ty: ids.TypeId, + name: ?[]const u8, + incoming: []const IncomingValue, +) Error!ids.ValueId { + const block = self.module.blocks.get(block_id) orelse return error.InvalidBlock; + const function = self.module.functions.get(block.parent_function) orelse return error.InvalidFunction; + + for (incoming) |item| { + const value = self.module.values.get(item.value) orelse return error.InvalidValue; + + if (value.type != ty) + return error.TypeMismatch; + + if (!functionHasEdgeTo(self.module, function, item.predecessor, block_id)) + return error.UnexpectedIncomingValue; + } + for (function.blocks.items) |predecessor| { + const edge_count = countEdgesTo(self.module.blocks.get(predecessor).?, block_id); + if (edge_count != 0 and findIncoming(incoming, predecessor) == null) + return error.MissingIncomingValue; + } + + var builder = Builder.init(self.module); + const parameter = try builder.addBlockParameter(block_id, ty, name); + + for (function.blocks.items) |predecessor| { + const incoming_value = findIncoming(incoming, predecessor) orelse continue; + try self.appendArgumentToEdges(predecessor, block_id, incoming_value); + } + + return parameter; +} + +pub fn removeBlockParameter( + self: *Self, + block_id: ids.BlockId, + parameter_index: usize, + replacement: ids.ValueId, +) Error!void { + const block = self.module.blocks.get(block_id) orelse return error.InvalidBlock; + if (parameter_index >= block.parameters.items.len) return error.InvalidParameterIndex; + const parameter = block.parameters.items[parameter_index]; + if (parameter == replacement) return error.InvalidValue; + _ = try self.replaceAllUses(parameter, replacement); + + const function = self.module.functions.get(block.parent_function) orelse return error.InvalidFunction; + for (function.blocks.items) |predecessor| { + try self.removeArgumentFromEdges(predecessor, block_id, parameter_index); + } + + const mutable_block = self.module.blocks.getMut(block_id).?; + _ = mutable_block.parameters.orderedRemove(parameter_index); + + for (mutable_block.parameters.items[parameter_index..], parameter_index..) |value_id, index| { + const value = self.module.values.getMut(value_id) orelse return error.InvalidValue; + value.definition.block_parameter.index = @intCast(index); + } + + _ = self.module.values.remove(parameter); +} + +fn validateArguments(self: *const Self, target: *const module_ir.Block, arguments: []const ids.ValueId) Error!void { + if (arguments.len != target.parameters.items.len) return error.TypeMismatch; + for (arguments, target.parameters.items) |argument, parameter| { + const argument_value = self.module.values.get(argument) orelse return error.InvalidValue; + const parameter_value = self.module.values.get(parameter) orelse return error.InvalidValue; + if (argument_value.type != parameter_value.type) return error.TypeMismatch; + } +} + +fn redirectOne( + self: *Self, + edge: *module_ir.Edge, + old_target: ids.BlockId, + new_target: ids.BlockId, + arguments: []const ids.ValueId, +) !bool { + if (edge.target != old_target) + return false; + + edge.target = new_target; + edge.arguments = try self.module.allocator().dupe(ids.ValueId, arguments); + return true; +} + +fn appendArgumentToEdges(self: *Self, predecessor: ids.BlockId, target: ids.BlockId, value: ids.ValueId) !void { + const block = self.module.blocks.getMut(predecessor) orelse return error.InvalidBlock; + const terminator = if (block.terminator) |*item| item else return error.InvalidBlock; + + switch (terminator.*) { + .branch => |*edge| { + if (edge.target == target) + try self.appendEdgeArgument(edge, value); + }, + .conditional_branch => |*branch| { + if (branch.true_edge.target == target) + try self.appendEdgeArgument(&branch.true_edge, value); + + if (branch.false_edge.target == target) + try self.appendEdgeArgument(&branch.false_edge, value); + }, + else => {}, + } +} + +fn appendEdgeArgument(self: *Self, edge: *module_ir.Edge, value: ids.ValueId) !void { + const arguments = try self.module.allocator().alloc(ids.ValueId, edge.arguments.len + 1); + @memcpy(arguments[0..edge.arguments.len], edge.arguments); + arguments[edge.arguments.len] = value; + edge.arguments = arguments; +} + +fn removeArgumentFromEdges(self: *Self, predecessor: ids.BlockId, target: ids.BlockId, index: usize) !void { + const block = self.module.blocks.getMut(predecessor) orelse return error.InvalidBlock; + const terminator = if (block.terminator) |*item| item else return error.InvalidBlock; + + switch (terminator.*) { + .branch => |*edge| { + if (edge.target == target) + try self.removeEdgeArgument(edge, index); + }, + .conditional_branch => |*branch| { + if (branch.true_edge.target == target) + try self.removeEdgeArgument(&branch.true_edge, index); + + if (branch.false_edge.target == target) + try self.removeEdgeArgument(&branch.false_edge, index); + }, + else => {}, + } +} + +fn removeEdgeArgument(self: *Self, edge: *module_ir.Edge, index: usize) !void { + if (index >= edge.arguments.len) + return error.InvalidParameterIndex; + + const arguments = try self.module.allocator().alloc(ids.ValueId, edge.arguments.len - 1); + @memcpy(arguments[0..index], edge.arguments[0..index]); + @memcpy(arguments[index..], edge.arguments[index + 1 ..]); + edge.arguments = arguments; +} + +fn countUse(context: *UseCountContext, value: ids.ValueId) void { + if (value == context.needle) + context.count += 1; +} + +fn findIncoming(incoming: []const IncomingValue, predecessor: ids.BlockId) ?ids.ValueId { + for (incoming) |item| { + if (item.predecessor == predecessor) + return item.value; + } + return null; +} + +fn functionHasEdgeTo( + module: *const module_ir.Module, + function: *const module_ir.Function, + predecessor: ids.BlockId, + target: ids.BlockId, +) bool { + for (function.blocks.items) |block_id| { + if (block_id != predecessor) + continue; + return countEdgesTo(module.blocks.get(block_id) orelse return false, target) != 0; + } + return false; +} + +fn countEdgesTo(block: *const module_ir.Block, target: ids.BlockId) usize { + const terminator = block.terminator orelse return 0; + return switch (terminator) { + .branch => |edge| @intFromBool(edge.target == target), + .conditional_branch => |branch| @intFromBool(branch.true_edge.target == target) + @intFromBool(branch.false_edge.target == target), + else => 0, + }; +} diff --git a/src/compiler/ir/cfg.zig b/src/compiler/ir/cfg.zig new file mode 100644 index 0000000..cd1aadf --- /dev/null +++ b/src/compiler/ir/cfg.zig @@ -0,0 +1,198 @@ +const std = @import("std"); +const ids = @import("id.zig"); +const module_ir = @import("module.zig"); + +pub const Self = @This(); + +pub const Error = std.mem.Allocator.Error || error{ + InvalidFunction, + MissingEntryBlock, + InvalidBlock, + MissingTerminator, + CrossFunctionEdge, +}; + +allocator: std.mem.Allocator, +blocks: []ids.BlockId, +predecessors_by_block: []std.ArrayList(ids.BlockId), +reachable: []bool, +dominators: []bool, + +pub fn init(allocator: std.mem.Allocator, module: *const module_ir.Module, function_id: ids.FunctionId) Error!Self { + const function = module.functions.get(function_id) orelse return error.InvalidFunction; + const entry = function.entry_block orelse return error.MissingEntryBlock; + const blocks = try allocator.dupe(ids.BlockId, function.blocks.items); + errdefer allocator.free(blocks); + + const predecessor_lists = try allocator.alloc(std.ArrayList(ids.BlockId), blocks.len); + errdefer allocator.free(predecessor_lists); + + for (predecessor_lists) |*list| + list.* = .empty; + + errdefer for (predecessor_lists) |*list| list.deinit(allocator); + + const reachable = try allocator.alloc(bool, blocks.len); + errdefer allocator.free(reachable); + @memset(reachable, false); + + const dominators = try allocator.alloc(bool, blocks.len * blocks.len); + errdefer allocator.free(dominators); + @memset(dominators, false); + + var self: Self = .{ + .allocator = allocator, + .blocks = blocks, + .predecessors_by_block = predecessor_lists, + .reachable = reachable, + .dominators = dominators, + }; + + try self.buildPredecessors(module); + try self.buildReachability(module, entry); + self.buildDominators(entry); + + return self; +} + +pub fn deinit(self: *Self) void { + for (self.predecessors_by_block) |*list| + list.deinit(self.allocator); + + self.allocator.free(self.predecessors_by_block); + self.allocator.free(self.blocks); + self.allocator.free(self.reachable); + self.allocator.free(self.dominators); + self.* = undefined; +} + +pub fn predecessors(self: *const Self, block: ids.BlockId) ?[]const ids.BlockId { + const index = self.indexOf(block) orelse return null; + return self.predecessors_by_block[index].items; +} + +pub fn isReachable(self: *const Self, block: ids.BlockId) bool { + const index = self.indexOf(block) orelse return false; + return self.reachable[index]; +} + +pub fn dominates(self: *const Self, dominator: ids.BlockId, block: ids.BlockId) bool { + const dominator_index = self.indexOf(dominator) orelse return false; + const block_index = self.indexOf(block) orelse return false; + return self.dominators[block_index * self.blocks.len + dominator_index]; +} + +fn buildPredecessors(self: *Self, module: *const module_ir.Module) Error!void { + for (self.blocks) |source| { + const block = module.blocks.get(source) orelse return error.InvalidBlock; + const terminator = block.terminator orelse return error.MissingTerminator; + switch (terminator) { + .branch => |edge| try self.addPredecessor(edge.target, source), + .conditional_branch => |branch| { + try self.addPredecessor(branch.true_edge.target, source); + try self.addPredecessor(branch.false_edge.target, source); + }, + else => {}, + } + } +} + +fn buildReachability(self: *Self, module: *const module_ir.Module, entry: ids.BlockId) Error!void { + var queue: std.ArrayList(ids.BlockId) = .empty; + defer queue.deinit(self.allocator); + try queue.append(self.allocator, entry); + self.reachable[self.indexOf(entry) orelse return error.InvalidBlock] = true; + + var cursor: usize = 0; + while (cursor < queue.items.len) : (cursor += 1) { + const block = module.blocks.get(queue.items[cursor]) orelse return error.InvalidBlock; + const terminator = block.terminator orelse return error.MissingTerminator; + switch (terminator) { + .branch => |edge| try self.markReachable(&queue, edge.target), + .conditional_branch => |branch| { + try self.markReachable(&queue, branch.true_edge.target); + try self.markReachable(&queue, branch.false_edge.target); + }, + else => {}, + } + } +} + +fn buildDominators(self: *Self, entry: ids.BlockId) void { + const entry_index = self.indexOf(entry).?; + const count = self.blocks.len; + + for (0..count) |block_index| { + if (!self.reachable[block_index]) { + self.setDominates(block_index, block_index, true); + } else if (block_index == entry_index) { + self.setDominates(block_index, entry_index, true); + } else { + for (0..count) |candidate| { + if (self.reachable[candidate]) + self.setDominates(block_index, candidate, true); + } + } + } + + var changed = true; + while (changed) { + changed = false; + for (0..count) |block_index| { + if (!self.reachable[block_index] or block_index == entry_index) + continue; + + for (0..count) |candidate| { + var new_value = candidate == block_index; + if (!new_value) { + var saw_reachable_predecessor = false; + new_value = true; + for (self.predecessors_by_block[block_index].items) |predecessor| { + const predecessor_index = self.indexOf(predecessor).?; + + if (!self.reachable[predecessor_index]) + continue; + + saw_reachable_predecessor = true; + new_value = new_value and self.getDominates(predecessor_index, candidate); + } + new_value = new_value and saw_reachable_predecessor; + } + if (self.getDominates(block_index, candidate) != new_value) { + self.setDominates(block_index, candidate, new_value); + changed = true; + } + } + } + } +} + +fn addPredecessor(self: *Self, target: ids.BlockId, source: ids.BlockId) Error!void { + const target_index = self.indexOf(target) orelse return error.CrossFunctionEdge; + try self.predecessors_by_block[target_index].append(self.allocator, source); +} + +fn markReachable(self: *Self, queue: *std.ArrayList(ids.BlockId), target: ids.BlockId) Error!void { + const target_index = self.indexOf(target) orelse return error.CrossFunctionEdge; + + if (self.reachable[target_index]) + return; + + self.reachable[target_index] = true; + try queue.append(self.allocator, target); +} + +fn indexOf(self: *const Self, block: ids.BlockId) ?usize { + for (self.blocks, 0..) |candidate, index| { + if (candidate == block) return index; + } + return null; +} + +fn getDominates(self: *const Self, block_index: usize, candidate_index: usize) bool { + return self.dominators[block_index * self.blocks.len + candidate_index]; +} + +fn setDominates(self: *Self, block_index: usize, candidate_index: usize, value: bool) void { + self.dominators[block_index * self.blocks.len + candidate_index] = value; +} diff --git a/src/compiler/ir/constant.zig b/src/compiler/ir/constant.zig new file mode 100644 index 0000000..76b30f0 --- /dev/null +++ b/src/compiler/ir/constant.zig @@ -0,0 +1,18 @@ +const ids = @import("id.zig"); + +pub const ConstantId = ids.ConstantId; +pub const TypeId = ids.TypeId; + +pub const ConstantValue = union(enum) { + boolean: bool, + integer_bits: u64, + float_bits: u64, + null, + undef, + composite: []const ConstantId, +}; + +pub const Constant = struct { + type: TypeId, + value: ConstantValue, +}; diff --git a/src/compiler/ir/id.zig b/src/compiler/ir/id.zig new file mode 100644 index 0000000..b10ac96 --- /dev/null +++ b/src/compiler/ir/id.zig @@ -0,0 +1,84 @@ +const std = @import("std"); + +pub const TypeTag = opaque {}; +pub const ConstantTag = opaque {}; +pub const ValueTag = opaque {}; +pub const InstructionTag = opaque {}; +pub const BlockTag = opaque {}; +pub const FunctionTag = opaque {}; +pub const InterfaceVariableTag = opaque {}; +pub const ResourceTag = opaque {}; + +pub const TypeId = Id(TypeTag); +pub const ConstantId = Id(ConstantTag); +pub const ValueId = Id(ValueTag); +pub const InstructionId = Id(InstructionTag); +pub const BlockId = Id(BlockTag); +pub const FunctionId = Id(FunctionTag); +pub const InterfaceVariableId = Id(InterfaceVariableTag); +pub const ResourceId = Id(ResourceTag); + +pub fn Id(comptime Tag: type) type { + return enum(u32) { + _, + + pub const tag_type = Tag; + + pub fn fromIndex(item_index: usize) @This() { + std.debug.assert(item_index <= std.math.maxInt(u32)); + return @enumFromInt(item_index); + } + + pub fn index(self: @This()) usize { + return @intFromEnum(self); + } + }; +} + +pub fn Store(comptime IdType: type, comptime T: type) type { + return struct { + const Self = @This(); + + entries: std.ArrayList(?T) = .empty, + + pub fn add(self: *Self, allocator: std.mem.Allocator, value: T) !IdType { + const id = IdType.fromIndex(self.entries.items.len); + try self.entries.append(allocator, value); + return id; + } + + pub fn get(self: *const Self, id: IdType) ?*const T { + if (id.index() >= self.entries.items.len) + return null; + + const entry = &self.entries.items[id.index()]; + return if (entry.*) |*value| value else null; + } + + pub fn getMut(self: *Self, id: IdType) ?*T { + if (id.index() >= self.entries.items.len) + return null; + + const entry = &self.entries.items[id.index()]; + return if (entry.*) |*value| value else null; + } + + /// Removing an object leaves a tombstone as IDs are deliberately not recycled + /// so they are never silently redirected to a different object. + pub fn remove(self: *Self, id: IdType) bool { + if (id.index() >= self.entries.items.len) + return false; + + const entry = &self.entries.items[id.index()]; + if (entry.* == null) + return false; + + entry.* = null; + return true; + } + + pub fn isLive(self: *const Self, id: IdType) bool { + return self.get(id) != null; + } + }; +} diff --git a/src/compiler/ir/instruction.zig b/src/compiler/ir/instruction.zig new file mode 100644 index 0000000..53cdb6e --- /dev/null +++ b/src/compiler/ir/instruction.zig @@ -0,0 +1,228 @@ +const std = @import("std"); +const ids = @import("id.zig"); + +pub const TypeId = ids.TypeId; +pub const ValueId = ids.ValueId; +pub const BlockId = ids.BlockId; +pub const FunctionId = ids.FunctionId; +pub const InterfaceVariableId = ids.InterfaceVariableId; + +pub const SourceLocation = struct { + file: ?[]const u8 = null, + line: u32, + column: u32, +}; + +pub const UnaryOpcode = enum { + negate, + logical_not, + bitwise_not, +}; + +pub const BinaryOpcode = enum { + integer_add, + integer_subtract, + integer_multiply, + unsigned_divide, + signed_divide, + unsigned_modulo, + signed_modulo, + float_add, + float_subtract, + float_multiply, + float_divide, + float_modulo, + shift_left, + logical_shift_right, + arithmetic_shift_right, + bitwise_and, + bitwise_or, + bitwise_xor, + logical_and, + logical_or, +}; + +pub const CompareOpcode = enum { + equal, + not_equal, + unsigned_less, + signed_less, + ordered_float_equal, + unordered_float_equal, + ordered_float_not_equal, + unordered_float_not_equal, + ordered_float_less, + unordered_float_less, +}; + +pub const Unary = struct { + opcode: UnaryOpcode, + operand: ValueId, +}; + +pub const Binary = struct { + opcode: BinaryOpcode, + lhs: ValueId, + rhs: ValueId, +}; + +pub const Compare = struct { + opcode: CompareOpcode, + lhs: ValueId, + rhs: ValueId, +}; + +pub const Select = struct { + condition: ValueId, + true_value: ValueId, + false_value: ValueId, +}; + +pub const CompositeConstruct = struct { + elements: []const ValueId, +}; + +pub const CompositeExtract = struct { + composite: ValueId, + indices: []const u32, +}; + +pub const LoadInterface = struct { + variable: InterfaceVariableId, + element_index: ?ValueId = null, +}; + +pub const StoreInterface = struct { + variable: InterfaceVariableId, + value: ValueId, + element_index: ?ValueId = null, +}; + +pub const Call = struct { + function: FunctionId, + arguments: []const ValueId, +}; + +pub const Operation = union(enum) { + unary: Unary, + binary: Binary, + compare: Compare, + select: Select, + bitcast: ValueId, + composite_construct: CompositeConstruct, + composite_extract: CompositeExtract, + load_interface: LoadInterface, + store_interface: StoreInterface, + call: Call, + + pub fn visitValueUses(self: Operation, context: anytype, comptime visitor: anytype) void { + switch (self) { + .unary => |op| visitor(context, op.operand), + .binary => |op| { + visitor(context, op.lhs); + visitor(context, op.rhs); + }, + .compare => |op| { + visitor(context, op.lhs); + visitor(context, op.rhs); + }, + .select => |op| { + visitor(context, op.condition); + visitor(context, op.true_value); + visitor(context, op.false_value); + }, + .bitcast => |operand| visitor(context, operand), + .composite_construct => |op| for (op.elements) |element| visitor(context, element), + .composite_extract => |op| visitor(context, op.composite), + .load_interface => |op| if (op.element_index) |index| visitor(context, index), + .store_interface => |op| { + visitor(context, op.value); + if (op.element_index) |index| + visitor(context, index); + }, + .call => |op| { + for (op.arguments) |argument| + visitor(context, argument); + }, + } + } + + pub fn replaceValueUses(self: *Operation, allocator: std.mem.Allocator, old: ValueId, replacement: ValueId) !usize { + var count: usize = 0; + switch (self.*) { + .unary => |*op| replaceOne(&op.operand, old, replacement, &count), + .binary => |*op| { + replaceOne(&op.lhs, old, replacement, &count); + replaceOne(&op.rhs, old, replacement, &count); + }, + .compare => |*op| { + replaceOne(&op.lhs, old, replacement, &count); + replaceOne(&op.rhs, old, replacement, &count); + }, + .select => |*op| { + replaceOne(&op.condition, old, replacement, &count); + replaceOne(&op.true_value, old, replacement, &count); + replaceOne(&op.false_value, old, replacement, &count); + }, + .bitcast => |*operand| replaceOne(operand, old, replacement, &count), + .composite_construct => |*op| op.elements = try replaceSlice(allocator, op.elements, old, replacement, &count), + .composite_extract => |*op| replaceOne(&op.composite, old, replacement, &count), + .load_interface => |*op| { + if (op.element_index) |*index| + replaceOne(index, old, replacement, &count); + }, + .store_interface => |*op| { + replaceOne(&op.value, old, replacement, &count); + if (op.element_index) |*index| + replaceOne(index, old, replacement, &count); + }, + .call => |*op| op.arguments = try replaceSlice(allocator, op.arguments, old, replacement, &count), + } + return count; + } + + pub fn hasSideEffects(self: Operation) bool { + return switch (self) { + .store_interface, .call => true, + else => false, + }; + } +}; + +pub const Instruction = struct { + parent_block: BlockId, + result: ?ValueId, + operation: Operation, + source: ?SourceLocation = null, +}; + +fn replaceOne(operand: *ValueId, old: ValueId, replacement: ValueId, count: *usize) void { + if (operand.* != old) return; + operand.* = replacement; + count.* += 1; +} + +fn replaceSlice( + allocator: std.mem.Allocator, + operands: []const ValueId, + old: ValueId, + replacement: ValueId, + count: *usize, +) ![]const ValueId { + var occurrences: usize = 0; + for (operands) |operand| if (operand == old) { + occurrences += 1; + }; + + if (occurrences == 0) + return operands; + + const copy = try allocator.dupe(ValueId, operands); + for (copy) |*operand| { + if (operand.* == old) + operand.* = replacement; + } + + count.* += occurrences; + return copy; +} diff --git a/src/compiler/ir/ir.zig b/src/compiler/ir/ir.zig new file mode 100644 index 0000000..b6f1755 --- /dev/null +++ b/src/compiler/ir/ir.zig @@ -0,0 +1,58 @@ +//! ## Intermediate Representation +//! +//! The IR is the compiler's target-independent shader representation. It keeps +//! shader structure explicit while hiding SPIR-V's binary encoding and large +//! instruction surface. +//! +//! `module.Module` owns types, constants, values, instructions, blocks, functions, +//! interfaces, and resources in typed ID stores. Values are in SSA form: each value +//! is defined by a constant, parameter, instruction, or `undef`. +//! +//! Control flow is represented with basic blocks and terminators. Phi-like values +//! are modeled as block parameters, with branch edge arguments supplying incoming +//! values. Blocks can also record structured selection or loop metadata. +//! +//! Use `Builder` to construct modules, `validator.validate` to check invariants, +//! `cfg` for control-flow queries, `Rewriter` for common edits, and +//! `parser`/`printer` for the textual IR format used by tests and debugging. +//! +//! Here's a simple text representation of a shader module that `parser.Parser` and `printer` can handle/produce: +//! ``` +//! shader vertex @main +//! { +//! @color: vec4[f32] = input[location(0), component(0), index(0)] +//! @out_color: vec4[f32] = output[location(0), component(0), index(0)] +//! %0: constant bool = true +//! %1: constant f32 = bits(0x3f800000) +//! +//! fn @main() -> void +//! { +//! .entry(): +//! %3: vec4[f32] = load_interface @color +//! conditional_branch %0, .pass(), .merge(%3) +//! +//! .pass(): +//! %4: vec4[f32] = composite_construct %1, %1, %1, %1 +//! branch .merge(%4) +//! +//! .merge(%2: vec4[f32]): +//! store_interface @out_color, %2 +//! return +//! } +//! } +//! ``` + +pub const Builder = @import("Builder.zig"); +pub const Rewriter = @import("Rewriter.zig"); +pub const cfg = @import("cfg.zig"); +pub const constant = @import("constant.zig"); +pub const id = @import("id.zig"); +pub const instruction = @import("instruction.zig"); +pub const module = @import("module.zig"); +pub const parser = @import("parser/parser.zig"); +pub const pass_manager = @import("pass_manager.zig"); +pub const printer = @import("printer.zig"); +pub const types = @import("type.zig"); +pub const validator = @import("validator/validator.zig"); +pub const value = @import("value.zig"); +pub const visitor = @import("visitor.zig"); diff --git a/src/compiler/ir/module.zig b/src/compiler/ir/module.zig new file mode 100644 index 0000000..c77cca3 --- /dev/null +++ b/src/compiler/ir/module.zig @@ -0,0 +1,235 @@ +const std = @import("std"); +const ids = @import("id.zig"); +const types = @import("type.zig"); +const constants = @import("constant.zig"); +const values = @import("value.zig"); +const instructions = @import("instruction.zig"); + +pub const Stage = enum { + vertex, + fragment, + compute, +}; + +pub const ExecutionModes = struct { + workgroup_size: ?[3]u32 = null, + early_fragment_tests: bool = false, +}; + +pub const Properties = packed struct { + valid_cfg: bool = false, + valid_ssa: bool = false, + structured_control_flow: bool = false, + no_function_calls: bool = false, + no_local_memory: bool = false, + no_matrix_types: bool = false, + no_large_composites: bool = false, + explicit_resource_offsets: bool = false, + _padding: u24 = 0, +}; + +pub const ConstantStore = ids.Store(ids.ConstantId, constants.Constant); +pub const ValueStore = ids.Store(ids.ValueId, values.Value); +pub const InstructionStore = ids.Store(ids.InstructionId, instructions.Instruction); +pub const BlockStore = ids.Store(ids.BlockId, Block); +pub const FunctionStore = ids.Store(ids.FunctionId, Function); +pub const InterfaceVariableStore = ids.Store(ids.InterfaceVariableId, InterfaceVariable); +pub const ResourceStore = ids.Store(ids.ResourceId, Resource); +pub const TypeStore = ids.Store(ids.TypeId, types.Type); + +pub const Edge = struct { + target: ids.BlockId, + arguments: []const ids.ValueId, +}; + +pub const Terminator = union(enum) { + branch: Edge, + conditional_branch: struct { + condition: ids.ValueId, + true_edge: Edge, + false_edge: Edge, + }, + return_void, + return_value: ids.ValueId, + discard, + @"unreachable", +}; + +pub const StructuredControl = union(enum) { + none, + selection: struct { merge_block: ids.BlockId }, + loop: struct { + merge_block: ids.BlockId, + continue_block: ids.BlockId, + }, +}; + +pub const Block = struct { + parent_function: ids.FunctionId, + parameters: std.ArrayList(ids.ValueId) = .empty, + instructions: std.ArrayList(ids.InstructionId) = .empty, + terminator: ?Terminator = null, + structured_control: StructuredControl = .none, + name: ?[]const u8 = null, +}; + +pub const Function = struct { + return_type: ids.TypeId, + parameter_types: std.ArrayList(ids.TypeId) = .empty, + parameters: std.ArrayList(ids.ValueId) = .empty, + blocks: std.ArrayList(ids.BlockId) = .empty, + entry_block: ?ids.BlockId = null, + name: ?[]const u8 = null, +}; + +pub const InterfaceDirection = enum { + input, + output, +}; + +pub const Builtin = enum { + position, + vertex_index, + instance_index, + frag_coord, + frag_depth, + global_invocation_id, +}; + +pub const InterfaceSemantic = union(enum) { + location: struct { location: u32, component: u8 = 0, index: u8 = 0 }, + builtin: Builtin, +}; + +pub const InterfaceVariable = struct { + type: ids.TypeId, + direction: InterfaceDirection, + semantic: InterfaceSemantic, + name: ?[]const u8 = null, +}; + +pub const Resource = struct { + kind: types.ResourceKind, + set: u32, + binding: u32, + type: ids.TypeId, + name: ?[]const u8 = null, +}; + +pub const Module = struct { + arena: std.heap.ArenaAllocator, + stage: Stage, + entry_point: ?ids.FunctionId = null, + execution_modes: ExecutionModes = .{}, + types: TypeStore = .{}, + constants: ConstantStore = .{}, + values: ValueStore = .{}, + instructions: InstructionStore = .{}, + blocks: BlockStore = .{}, + functions: FunctionStore = .{}, + interface_variables: InterfaceVariableStore = .{}, + resources: ResourceStore = .{}, + properties: Properties = .{}, + + pub fn init(backing_allocator: std.mem.Allocator, stage: Stage) Module { + return .{ + .arena = std.heap.ArenaAllocator.init(backing_allocator), + .stage = stage, + }; + } + + pub fn deinit(self: *Module) void { + self.arena.deinit(); + self.* = undefined; + } + + pub fn allocator(self: *Module) std.mem.Allocator { + return self.arena.allocator(); + } + + pub fn backingAllocator(self: *const Module) std.mem.Allocator { + return self.arena.child_allocator; + } + + pub fn internType(self: *Module, candidate: types.Type) !ids.TypeId { + for (self.types.entries.items, 0..) |entry, index| { + if (entry) |existing| { + if (existing.eql(candidate)) + return ids.TypeId.fromIndex(index); + } + } + + var owned = candidate; + if (candidate == .structure) { + owned.structure.members = try self.allocator().dupe(ids.TypeId, candidate.structure.members); + } + + return self.types.add(self.allocator(), owned); + } + + pub fn typeOf(self: *const Module, value_id: ids.ValueId) ?ids.TypeId { + const value = self.values.get(value_id) orelse return null; + return value.type; + } +}; + +pub fn visitTerminatorValueUses(terminator: Terminator, context: anytype, comptime visitor: anytype) void { + switch (terminator) { + .branch => |edge| { + for (edge.arguments) |argument| + visitor(context, argument); + }, + .conditional_branch => |branch| { + visitor(context, branch.condition); + + for (branch.true_edge.arguments) |argument| + visitor(context, argument); + + for (branch.false_edge.arguments) |argument| + visitor(context, argument); + }, + .return_value => |value| visitor(context, value), + else => {}, + } +} + +pub fn replaceTerminatorValueUses(allocator: std.mem.Allocator, terminator: *Terminator, old: ids.ValueId, replacement: ids.ValueId) !usize { + var count: usize = 0; + switch (terminator.*) { + .branch => |*edge| try replaceEdgeUses(allocator, edge, old, replacement, &count), + .conditional_branch => |*branch| { + replaceOne(&branch.condition, old, replacement, &count); + try replaceEdgeUses(allocator, &branch.true_edge, old, replacement, &count); + try replaceEdgeUses(allocator, &branch.false_edge, old, replacement, &count); + }, + .return_value => |*value| replaceOne(value, old, replacement, &count), + else => {}, + } + return count; +} + +fn replaceEdgeUses(allocator: std.mem.Allocator, edge: *Edge, old: ids.ValueId, replacement: ids.ValueId, count: *usize) !void { + var occurrences: usize = 0; + for (edge.arguments) |argument| if (argument == old) { + occurrences += 1; + }; + + if (occurrences == 0) + return; + + const copy = try allocator.dupe(ids.ValueId, edge.arguments); + for (copy) |*argument| { + if (argument.* == old) + argument.* = replacement; + } + + edge.arguments = copy; + count.* += occurrences; +} + +fn replaceOne(operand: *ids.ValueId, old: ids.ValueId, replacement: ids.ValueId, count: *usize) void { + if (operand.* != old) + return; + operand.* = replacement; + count.* += 1; +} diff --git a/src/compiler/ir/parser/Lexer.zig b/src/compiler/ir/parser/Lexer.zig new file mode 100644 index 0000000..448a385 --- /dev/null +++ b/src/compiler/ir/parser/Lexer.zig @@ -0,0 +1,196 @@ +const std = @import("std"); + +const Self = @This(); + +source: []const u8, +cursor: usize = 0, +lookahead: ?Token = null, + +pub const TokenTag = enum { + eof, + invalid, + identifier, + number, + value_ref, + constant_ref, + at_name, + dot_name, + left_brace, + right_brace, + left_paren, + right_paren, + left_square, + right_square, + colon, + comma, + equal, + arrow, +}; + +pub const Token = struct { + tag: TokenTag, + text: []const u8, +}; + +pub fn init(source: []const u8) Self { + return .{ + .source = source, + }; +} + +pub fn peek(self: *Self) Token { + if (self.lookahead == null) + self.lookahead = self.lex(); + return self.lookahead.?; +} + +pub fn take(self: *Self) Token { + const token = self.peek(); + self.lookahead = null; + return token; +} + +fn lex(self: *Self) Token { + while (self.cursor < self.source.len and std.ascii.isWhitespace(self.source[self.cursor])) + self.cursor += 1; + + if (self.cursor == self.source.len) { + return .{ + .tag = .eof, + .text = self.source[self.cursor..self.cursor], + }; + } + + const start = self.cursor; + const byte = self.source[self.cursor]; + self.cursor += 1; + + switch (byte) { + '{' => return self.simpleToken(.left_brace, start), + '}' => return self.simpleToken(.right_brace, start), + '(' => return self.simpleToken(.left_paren, start), + ')' => return self.simpleToken(.right_paren, start), + '[' => return self.simpleToken(.left_square, start), + ']' => return self.simpleToken(.right_square, start), + ':' => return self.simpleToken(.colon, start), + ',' => return self.simpleToken(.comma, start), + '=' => return self.simpleToken(.equal, start), + '-', '+' => { + if (byte == '-' and self.cursor < self.source.len and self.source[self.cursor] == '>') { + self.cursor += 1; + return .{ + .tag = .arrow, + .text = self.source[start..self.cursor], + }; + } + + if (self.cursor < self.source.len and std.ascii.isDigit(self.source[self.cursor])) + return self.numberToken(start); + + return self.simpleToken(.invalid, start); + }, + '%', '@', '.' => { + const tag: TokenTag = switch (byte) { + '%' => .value_ref, + '@' => .at_name, + '.' => .dot_name, + else => unreachable, + }; + const content_start = self.cursor; + + if (byte == '%' and self.cursor < self.source.len and std.ascii.isDigit(self.source[self.cursor])) { + while (self.cursor < self.source.len and std.ascii.isDigit(self.source[self.cursor])) self.cursor += 1; + } else { + while (self.cursor < self.source.len and isNameByte(self.source[self.cursor])) self.cursor += 1; + } + + if (self.cursor == content_start) + return self.simpleToken(.invalid, start); + + return .{ + .tag = tag, + .text = self.source[content_start..self.cursor], + }; + }, + '#' => { + const number_start = self.cursor; + while (self.cursor < self.source.len and std.ascii.isDigit(self.source[self.cursor])) + self.cursor += 1; + + if (self.cursor == number_start) + return self.simpleToken(.invalid, start); + + return .{ + .tag = .constant_ref, + .text = self.source[number_start..self.cursor], + }; + }, + else => {}, + } + + if (std.ascii.isDigit(byte)) + return self.numberToken(start); + + if (isNameStart(byte)) { + while (self.cursor < self.source.len and isNameByte(self.source[self.cursor])) + self.cursor += 1; + + return .{ + .tag = .identifier, + .text = self.source[start..self.cursor], + }; + } + + return self.simpleToken(.invalid, start); +} + +fn numberToken(self: *Self, start: usize) Token { + var number_start = start; + if (self.source[number_start] == '-' or self.source[number_start] == '+') + number_start += 1; + + self.cursor = number_start; + + if (self.source[number_start] == '0' and number_start + 1 < self.source.len and self.source[number_start + 1] == 'x') { + self.cursor = number_start + 2; + while (self.cursor < self.source.len and std.ascii.isHex(self.source[self.cursor])) + self.cursor += 1; + } else { + while (self.cursor < self.source.len and std.ascii.isDigit(self.source[self.cursor])) + self.cursor += 1; + + if (self.cursor < self.source.len and self.source[self.cursor] == '.') { + self.cursor += 1; + while (self.cursor < self.source.len and std.ascii.isDigit(self.source[self.cursor])) + self.cursor += 1; + } + + if (self.cursor < self.source.len and (self.source[self.cursor] == 'e' or self.source[self.cursor] == 'E')) { + self.cursor += 1; + if (self.cursor < self.source.len and (self.source[self.cursor] == '-' or self.source[self.cursor] == '+')) + self.cursor += 1; + while (self.cursor < self.source.len and std.ascii.isDigit(self.source[self.cursor])) + self.cursor += 1; + } + } + + return .{ + .tag = .number, + .text = self.source[start..self.cursor], + }; +} + +fn simpleToken(self: *Self, tag: TokenTag, start: usize) Token { + return .{ + .tag = tag, + .text = self.source[start..self.cursor], + }; +} + +fn isNameStart(byte: u8) bool { + return std.ascii.isAlphabetic(byte) or byte == '_'; +} + +fn isNameByte(byte: u8) bool { + return std.ascii.isAlphanumeric(byte) or byte == '_'; +} diff --git a/src/compiler/ir/parser/ast.zig b/src/compiler/ir/parser/ast.zig new file mode 100644 index 0000000..d896ebc --- /dev/null +++ b/src/compiler/ir/parser/ast.zig @@ -0,0 +1,93 @@ +const std = @import("std"); +const ids = @import("../id.zig"); +const inst_ir = @import("../instruction.zig"); +const module_ir = @import("../module.zig"); + +pub const ValueRef = []const u8; + +pub const ParsedModule = struct { + entry_point_name: ?[]const u8, + interfaces: std.ArrayList(ParsedInterface) = .empty, + constants: std.ArrayList(ParsedConstant) = .empty, + functions: std.ArrayList(ParsedFunction) = .empty, +}; + +pub const ParsedInterface = struct { + direction: module_ir.InterfaceDirection, + name: []const u8, + ty: ids.TypeId, + semantic: module_ir.InterfaceSemantic, +}; + +pub const ParsedConstantValue = union(enum) { + boolean: bool, + integer_bits: u64, + float_bits: u64, + null_value, + undef, + composite: []const u32, +}; + +pub const ParsedConstant = struct { + printed_value: ValueRef, + ty: ids.TypeId, + value: ParsedConstantValue, +}; + +pub const ParsedParameter = struct { + printed_value: ValueRef, + ty: ids.TypeId, +}; + +pub const ParsedInstruction = struct { + printed_result: ?ValueRef, + result_type: ?ids.TypeId, + operation: ParsedOperation, +}; + +pub const ParsedBlock = struct { + name: []const u8, + parameters: std.ArrayList(ParsedParameter) = .empty, + instructions: std.ArrayList(ParsedInstruction) = .empty, + terminator: ?ParsedTerminator = null, + actual: ?ids.BlockId = null, +}; + +pub const ParsedFunction = struct { + name: []const u8, + return_type: ids.TypeId, + parameters: std.ArrayList(ParsedParameter) = .empty, + blocks: std.ArrayList(ParsedBlock) = .empty, + actual: ?ids.FunctionId = null, +}; + +pub const ParsedEdge = struct { + block_name: []const u8, + arguments: []const ValueRef, +}; + +pub const ParsedTerminator = union(enum) { + branch: ParsedEdge, + conditional_branch: struct { + condition: ValueRef, + true_edge: ParsedEdge, + false_edge: ParsedEdge, + }, + return_void, + return_value: ValueRef, + discard, + unreachable_value, +}; + +pub const ParsedOperation = union(enum) { + unary: struct { opcode: inst_ir.UnaryOpcode, operand: ValueRef }, + binary: struct { opcode: inst_ir.BinaryOpcode, lhs: ValueRef, rhs: ValueRef }, + compare: struct { opcode: inst_ir.CompareOpcode, lhs: ValueRef, rhs: ValueRef }, + select: struct { condition: ValueRef, true_value: ValueRef, false_value: ValueRef }, + bitcast: ValueRef, + composite_construct: []const ValueRef, + composite_extract: struct { composite: ValueRef, indices: []const u32 }, + load_interface: []const u8, + store_interface: struct { interface_name: []const u8, value: ValueRef }, + call: struct { function_name: []const u8, arguments: []const ValueRef }, +}; diff --git a/src/compiler/ir/parser/lower.zig b/src/compiler/ir/parser/lower.zig new file mode 100644 index 0000000..2ee4f34 --- /dev/null +++ b/src/compiler/ir/parser/lower.zig @@ -0,0 +1,380 @@ +const std = @import("std"); +const Builder = @import("../Builder.zig"); +const constant_ir = @import("../constant.zig"); +const ids = @import("../id.zig"); +const inst_ir = @import("../instruction.zig"); +const module_ir = @import("../module.zig"); +const ast = @import("ast.zig"); + +const ValueRef = ast.ValueRef; +const ParsedModule = ast.ParsedModule; +const ParsedOperation = ast.ParsedOperation; +const ParsedTerminator = ast.ParsedTerminator; +const ParsedEdge = ast.ParsedEdge; + +const LoweredOperation = struct { + operation: inst_ir.Operation, + inferred_type: ?ids.TypeId, +}; + +pub fn lower(allocator: std.mem.Allocator, module: *module_ir.Module, parsed: *ParsedModule) !void { + var builder = Builder.init(module); + var values: std.StringHashMapUnmanaged(ids.ValueId) = .empty; + var constants: std.AutoHashMapUnmanaged(u32, ids.ConstantId) = .empty; + var interfaces: std.StringHashMapUnmanaged(ids.InterfaceVariableId) = .empty; + var functions: std.StringHashMapUnmanaged(ids.FunctionId) = .empty; + + for (parsed.interfaces.items) |interface| { + if (interfaces.contains(interface.name)) + return error.DuplicateName; + + const id = try builder.addInterfaceVariable(interface.ty, interface.direction, interface.semantic, interface.name); + try interfaces.put(allocator, interface.name, id); + } + + for (parsed.constants.items, 0..) |constant, constant_index| { + const value: constant_ir.ConstantValue = switch (constant.value) { + .boolean => |item| .{ .boolean = item }, + .integer_bits => |item| .{ .integer_bits = item }, + .float_bits => |item| .{ .float_bits = item }, + .null_value => .null, + .undef => .undef, + .composite => |printed_elements| blk: { + var elements: std.ArrayList(ids.ConstantId) = .empty; + for (printed_elements) |printed_element| { + try elements.append(allocator, constants.get(printed_element) orelse return error.UnknownConstant); + } + break :blk .{ .composite = elements.items }; + }, + }; + + const value_id = try builder.internConstant(constant.ty, value); + try builder.setValueName(value_id, valueName(constant.printed_value)); + try putValue(allocator, &values, constant.printed_value, value_id); + const definition = module.values.get(value_id).?.definition; + + if (definition != .constant) + return error.InvalidResult; + + try constants.put(allocator, @intCast(constant_index), definition.constant); + } + + for (parsed.functions.items) |*function| { + if (functions.contains(function.name)) + return error.DuplicateName; + + const function_id = try builder.addFunction(function.return_type, function.name); + function.actual = function_id; + try functions.put(allocator, function.name, function_id); + + for (function.parameters.items) |parameter| { + const value_id = try builder.addFunctionParameter(function_id, parameter.ty, valueName(parameter.printed_value)); + try putValue(allocator, &values, parameter.printed_value, value_id); + } + } + + if (parsed.entry_point_name) |entry_name| + builder.setEntryPoint(functions.get(entry_name) orelse return error.UnknownFunction); + + for (parsed.functions.items) |*function| { + var block_names: std.StringHashMapUnmanaged(ids.BlockId) = .empty; + for (function.blocks.items) |*block| { + if (block_names.contains(block.name)) + return error.DuplicateName; + + const block_id = try builder.addBlock(function.actual.?, block.name); + block.actual = block_id; + try block_names.put(allocator, block.name, block_id); + + for (block.parameters.items) |parameter| { + const value_id = try builder.addBlockParameter(block_id, parameter.ty, valueName(parameter.printed_value)); + try putValue(allocator, &values, parameter.printed_value, value_id); + } + } + + for (function.blocks.items) |block| { + for (block.instructions.items) |instruction| { + const lowered = try lowerOperation(allocator, module, &values, &interfaces, &functions, instruction.operation); + const result_type = instruction.result_type orelse lowered.inferred_type; + + if (instruction.printed_result != null and result_type == null) + return error.MissingResultType; + + if (instruction.printed_result == null and result_type != null) + return error.InvalidResult; + + const result = try builder.appendInstruction( + block.actual.?, + result_type, + lowered.operation, + if (instruction.printed_result) |printed_result| valueName(printed_result) else null, + ); + if (instruction.printed_result) |printed_result| + try putValue(allocator, &values, printed_result, result orelse return error.InvalidResult) + else if (result != null) + return error.InvalidResult; + } + const terminator = try lowerTerminator(allocator, &builder, &values, &block_names, block.terminator orelse return error.MissingTerminator); + try builder.setTerminator(block.actual.?, terminator); + } + } +} + +fn lowerOperation( + allocator: std.mem.Allocator, + module: *module_ir.Module, + values: *const std.StringHashMapUnmanaged(ids.ValueId), + interfaces: *const std.StringHashMapUnmanaged(ids.InterfaceVariableId), + functions: *const std.StringHashMapUnmanaged(ids.FunctionId), + parsed: ParsedOperation, +) !LoweredOperation { + return switch (parsed) { + .unary => |op| blk: { + const operand = resolveValue(values, op.operand) orelse return error.UnknownValue; + + break :blk .{ + .operation = .{ + .unary = .{ + .opcode = op.opcode, + .operand = operand, + }, + }, + .inferred_type = module.typeOf(operand), + }; + }, + .binary => |op| blk: { + const lhs = resolveValue(values, op.lhs) orelse return error.UnknownValue; + const rhs = resolveValue(values, op.rhs) orelse return error.UnknownValue; + + break :blk .{ + .operation = .{ + .binary = .{ + .opcode = op.opcode, + .lhs = lhs, + .rhs = rhs, + }, + }, + .inferred_type = module.typeOf(lhs), + }; + }, + .compare => |op| blk: { + const lhs = resolveValue(values, op.lhs) orelse return error.UnknownValue; + const rhs = resolveValue(values, op.rhs) orelse return error.UnknownValue; + + break :blk .{ + .operation = .{ + .compare = .{ + .opcode = op.opcode, + .lhs = lhs, + .rhs = rhs, + }, + }, + .inferred_type = try module.internType(.boolean), + }; + }, + .select => |op| blk: { + const condition = resolveValue(values, op.condition) orelse return error.UnknownValue; + const true_value = resolveValue(values, op.true_value) orelse return error.UnknownValue; + const false_value = resolveValue(values, op.false_value) orelse return error.UnknownValue; + + break :blk .{ + .operation = .{ + .select = .{ + .condition = condition, + .true_value = true_value, + .false_value = false_value, + }, + }, + .inferred_type = module.typeOf(true_value), + }; + }, + .bitcast => |printed_operand| blk: { + const operand = resolveValue(values, printed_operand) orelse return error.UnknownValue; + break :blk .{ .operation = .{ .bitcast = operand }, .inferred_type = module.typeOf(operand) }; + }, + .composite_construct => |printed_elements| blk: { + var elements: std.ArrayList(ids.ValueId) = .empty; + var element_types: std.ArrayList(ids.TypeId) = .empty; + for (printed_elements) |printed_element| { + const element = resolveValue(values, printed_element) orelse return error.UnknownValue; + try elements.append(allocator, element); + try element_types.append(allocator, module.typeOf(element) orelse return error.UnknownValue); + } + + break :blk .{ + .operation = .{ + .composite_construct = .{ + .elements = elements.items, + }, + }, + .inferred_type = try inferCompositeType(module, element_types.items), + }; + }, + .composite_extract => |op| blk: { + const composite = resolveValue(values, op.composite) orelse return error.UnknownValue; + + break :blk .{ + .operation = .{ + .composite_extract = .{ + .composite = composite, + .indices = op.indices, + }, + }, + .inferred_type = try extractedType(module, module.typeOf(composite) orelse return error.UnknownValue, op.indices), + }; + }, + .load_interface => |name| blk: { + const interface_id = interfaces.get(name) orelse return error.UnknownInterface; + + break :blk .{ + .operation = .{ + .load_interface = .{ + .variable = interface_id, + }, + }, + .inferred_type = module.interface_variables.get(interface_id).?.type, + }; + }, + .store_interface => |op| blk: { + const interface_id = interfaces.get(op.interface_name) orelse return error.UnknownInterface; + const value = resolveValue(values, op.value) orelse return error.UnknownValue; + + break :blk .{ + .operation = .{ + .store_interface = .{ + .variable = interface_id, + .value = value, + }, + }, + .inferred_type = null, + }; + }, + .call => |op| blk: { + const function_id = functions.get(op.function_name) orelse return error.UnknownFunction; + var arguments: std.ArrayList(ids.ValueId) = .empty; + + for (op.arguments) |printed_argument| + try arguments.append(allocator, resolveValue(values, printed_argument) orelse return error.UnknownValue); + + const return_type = module.functions.get(function_id).?.return_type; + const return_ir_type = module.types.get(return_type) orelse return error.InvalidType; + + break :blk .{ + .operation = .{ + .call = .{ + .function = function_id, + .arguments = arguments.items, + }, + }, + .inferred_type = if (return_ir_type.* == .void) null else return_type, + }; + }, + }; +} + +fn lowerTerminator( + allocator: std.mem.Allocator, + builder: *Builder, + values: *const std.StringHashMapUnmanaged(ids.ValueId), + blocks: *const std.StringHashMapUnmanaged(ids.BlockId), + parsed: ParsedTerminator, +) !module_ir.Terminator { + return switch (parsed) { + .branch => |edge| .{ + .branch = try lowerEdge(allocator, builder, values, blocks, edge), + }, + .conditional_branch => |branch| .{ + .conditional_branch = .{ + .condition = resolveValue(values, branch.condition) orelse return error.UnknownValue, + .true_edge = try lowerEdge(allocator, builder, values, blocks, branch.true_edge), + .false_edge = try lowerEdge(allocator, builder, values, blocks, branch.false_edge), + }, + }, + .return_void => .return_void, + .return_value => |printed_value| .{ + .return_value = resolveValue(values, printed_value) orelse return error.UnknownValue, + }, + .discard => .discard, + .unreachable_value => .@"unreachable", + }; +} + +fn lowerEdge( + allocator: std.mem.Allocator, + builder: *Builder, + values: *const std.StringHashMapUnmanaged(ids.ValueId), + blocks: *const std.StringHashMapUnmanaged(ids.BlockId), + parsed: ParsedEdge, +) !module_ir.Edge { + var arguments: std.ArrayList(ids.ValueId) = .empty; + for (parsed.arguments) |printed_argument| + try arguments.append(allocator, resolveValue(values, printed_argument) orelse return error.UnknownValue); + + return builder.edge(blocks.get(parsed.block_name) orelse return error.UnknownBlock, arguments.items); +} + +fn inferCompositeType(module: *module_ir.Module, element_types: []const ids.TypeId) !ids.TypeId { + if (element_types.len >= 2 and element_types.len <= std.math.maxInt(u8)) { + const first = element_types[0]; + for (element_types[1..]) |element_type| { + if (element_type != first) + return module.internType(.{ + .structure = .{ + .members = element_types, + }, + }); + } + return module.internType(.{ + .vector = .{ + .element_type = first, + .length = @intCast(element_types.len), + }, + }); + } + return module.internType(.{ + .structure = .{ + .members = element_types, + }, + }); +} + +fn extractedType(module: *const module_ir.Module, root_type: ids.TypeId, indices: []const u32) !ids.TypeId { + var current = root_type; + for (indices) |index| { + const ty = module.types.get(current) orelse return error.InvalidType; + current = switch (ty.*) { + .vector => |vector| if (index < vector.length) + vector.element_type + else + return error.InvalidCompositeIndex, + .array => |array| if (index < array.length) + array.element_type + else + return error.InvalidCompositeIndex, + .structure => |structure| if (index < structure.members.len) + structure.members[index] + else + return error.InvalidCompositeIndex, + else => return error.InvalidCompositeIndex, + }; + } + return current; +} + +fn putValue(allocator: std.mem.Allocator, values: *std.StringHashMapUnmanaged(ids.ValueId), printed: ValueRef, actual: ids.ValueId) !void { + if (values.contains(printed)) + return error.DuplicateValue; + try values.put(allocator, printed, actual); +} + +fn resolveValue(values: *const std.StringHashMapUnmanaged(ids.ValueId), printed: ValueRef) ?ids.ValueId { + return values.get(printed); +} + +fn valueName(reference: ValueRef) ?[]const u8 { + for (reference) |byte| { + if (!std.ascii.isDigit(byte)) + return reference; + } + return null; +} diff --git a/src/compiler/ir/parser/parser.zig b/src/compiler/ir/parser/parser.zig new file mode 100644 index 0000000..81b344b --- /dev/null +++ b/src/compiler/ir/parser/parser.zig @@ -0,0 +1,767 @@ +const std = @import("std"); +const ids = @import("../id.zig"); +const inst_ir = @import("../instruction.zig"); +const module_ir = @import("../module.zig"); +const type_ir = @import("../type.zig"); +const validator = @import("../validator/validator.zig"); +const Lexer = @import("Lexer.zig"); + +const ast = @import("ast.zig"); +const lowerer = @import("lower.zig"); + +pub const Error = error{ + UnexpectedToken, + InvalidNumber, + InvalidStage, + InvalidType, + InvalidOpcode, + InvalidSemantic, + DuplicateValue, + DuplicateName, + UnknownValue, + UnknownConstant, + UnknownInterface, + UnknownFunction, + UnknownBlock, + MissingTerminator, + MissingResultType, + InvalidResult, + InvalidCompositeIndex, +}; + +pub const max_file_size = 64 * 1024 * 1024; + +const ValueRef = ast.ValueRef; +const ParsedModule = ast.ParsedModule; +const ParsedInterface = ast.ParsedInterface; +const ParsedConstantValue = ast.ParsedConstantValue; +const ParsedConstant = ast.ParsedConstant; +const ParsedParameter = ast.ParsedParameter; +const ParsedInstruction = ast.ParsedInstruction; +const ParsedBlock = ast.ParsedBlock; +const ParsedFunction = ast.ParsedFunction; +const ParsedEdge = ast.ParsedEdge; +const ParsedTerminator = ast.ParsedTerminator; +const ParsedOperation = ast.ParsedOperation; +const Token = Lexer.Token; +const TokenTag = Lexer.TokenTag; + +const Parser = struct { + lexer: Lexer, + allocator: std.mem.Allocator, + module: ?*module_ir.Module = null, + + fn parseInterface(self: *Parser) !ParsedInterface { + const name = (try self.expect(.at_name)).text; + + try self.expectDiscard(.colon); + const ty = try self.parseType(); + try self.expectDiscard(.equal); + + const direction_token = try self.expect(.identifier); + const direction = std.meta.stringToEnum(module_ir.InterfaceDirection, direction_token.text) orelse return error.InvalidSemantic; + try self.expectDiscard(.left_square); + + const semantic_name = (try self.expect(.identifier)).text; + const semantic: module_ir.InterfaceSemantic = if (std.mem.eql(u8, semantic_name, "location")) blk: { + try self.expectDiscard(.left_paren); + const location = try self.parseUnsigned(u32, .number); + try self.expectDiscard(.right_paren); + + try self.expectDiscard(.comma); + try self.expectIdentifier("component"); + try self.expectDiscard(.left_paren); + const component = try self.parseUnsigned(u8, .number); + try self.expectDiscard(.right_paren); + + try self.expectDiscard(.comma); + try self.expectIdentifier("index"); + try self.expectDiscard(.left_paren); + const index = try self.parseUnsigned(u8, .number); + try self.expectDiscard(.right_paren); + + break :blk .{ + .location = .{ + .location = location, + .component = component, + .index = index, + }, + }; + } else if (std.mem.eql(u8, semantic_name, "builtin")) blk: { + try self.expectDiscard(.left_paren); + const builtin_name = (try self.expect(.identifier)).text; + try self.expectDiscard(.right_paren); + + const builtin = std.meta.stringToEnum(module_ir.Builtin, builtin_name) orelse return error.InvalidSemantic; + break :blk .{ .builtin = builtin }; + } else return error.InvalidSemantic; + + try self.expectDiscard(.right_square); + + return .{ + .direction = direction, + .name = name, + .ty = ty, + .semantic = semantic, + }; + } + + fn parseConstant(self: *Parser) !ParsedConstant { + const printed_value = try self.parseValueRef(); + + try self.expectDiscard(.colon); + try self.expectIdentifier("constant"); + + const ty = try self.parseType(); + try self.expectDiscard(.equal); + + const token = try self.peek(); + + const value: ParsedConstantValue = switch (token.tag) { + .identifier => blk: { + const word = (try self.take()).text; + if (std.mem.eql(u8, word, "true")) + break :blk .{ .boolean = true }; + + if (std.mem.eql(u8, word, "false")) + break :blk .{ .boolean = false }; + + if (std.mem.eql(u8, word, "null")) + break :blk .null_value; + + if (std.mem.eql(u8, word, "undef")) + break :blk .undef; + + if (std.mem.eql(u8, word, "bits")) { + try self.expectDiscard(.left_paren); + const bits = try self.parseUnsigned(u64, .number); + try self.expectDiscard(.right_paren); + + const ir_type = self.module.?.types.get(ty) orelse return error.InvalidType; + + break :blk switch (ir_type.*) { + .integer => .{ + .integer_bits = bits, + }, + .floating => .{ + .float_bits = bits, + }, + else => return error.InvalidType, + }; + } + return error.UnexpectedToken; + }, + .left_square => .{ + .composite = try self.parseConstantList(), + }, + .number => try self.parseDirectConstant(ty), + else => return error.UnexpectedToken, + }; + + return .{ + .printed_value = printed_value, + .ty = ty, + .value = value, + }; + } + + fn parseDirectConstant(self: *Parser, ty: ids.TypeId) !ParsedConstantValue { + const text = (try self.expect(.number)).text; + const ir_type = self.module.?.types.get(ty) orelse return error.InvalidType; + + return switch (ir_type.*) { + .integer => |integer| .{ .integer_bits = try parseIntegerLiteral(integer, text) }, + .floating => |float| .{ .float_bits = try parseFloatLiteral(float.bits, text) }, + else => error.InvalidType, + }; + } + + fn parseFunction(self: *Parser) !ParsedFunction { + try self.expectIdentifier("fn"); + const name = (try self.expect(.at_name)).text; + try self.expectDiscard(.left_paren); + + var parameters: std.ArrayList(ParsedParameter) = .empty; + if ((try self.peek()).tag != .right_paren) { + while (true) { + const printed_value = try self.parseValueRef(); + try self.expectDiscard(.colon); + + const ty = try self.parseType(); + try parameters.append(self.allocator, .{ .printed_value = printed_value, .ty = ty }); + + if (!try self.consume(.comma)) + break; + } + } + + try self.expectDiscard(.right_paren); + try self.expectDiscard(.arrow); + + const return_type = try self.parseType(); + + var function: ParsedFunction = .{ + .name = name, + .return_type = return_type, + .parameters = parameters, + }; + + try self.expectDiscard(.left_brace); + + while ((try self.peek()).tag != .right_brace) { + if ((try self.peek()).tag != .dot_name) + return error.UnexpectedToken; + try function.blocks.append(self.allocator, try self.parseBlock()); + } + + try self.expectDiscard(.right_brace); + + return function; + } + + fn parseBlock(self: *Parser) !ParsedBlock { + var block: ParsedBlock = .{ + .name = (try self.expect(.dot_name)).text, + }; + + try self.expectDiscard(.left_paren); + + if ((try self.peek()).tag != .right_paren) { + while (true) { + const printed_value = try self.parseValueRef(); + + try self.expectDiscard(.colon); + const ty = try self.parseType(); + try block.parameters.append(self.allocator, .{ .printed_value = printed_value, .ty = ty }); + + if (!try self.consume(.comma)) + break; + } + } + + try self.expectDiscard(.right_paren); + try self.expectDiscard(.colon); + + while (block.terminator == null) { + const token = try self.peek(); + + if (token.tag == .right_brace or token.tag == .dot_name) + return error.MissingTerminator; + + if (token.tag == .value_ref) { + try block.instructions.append(self.allocator, try self.parseInstruction(true)); + continue; + } + + if (token.tag != .identifier) + return error.UnexpectedToken; + + if (isTerminatorName(token.text)) { + block.terminator = try self.parseTerminator(); + } else { + try block.instructions.append(self.allocator, try self.parseInstruction(false)); + } + } + + return block; + } + + fn parseInstruction(self: *Parser, has_result: bool) !ParsedInstruction { + const printed_result = if (has_result) try self.parseValueRef() else null; + const result_type = if (has_result and try self.consume(.colon)) try self.parseType() else null; + + if (has_result) + try self.expectDiscard(.equal); + + return .{ + .printed_result = printed_result, + .result_type = result_type, + .operation = try self.parseOperation(), + }; + } + + fn parseOperation(self: *Parser) !ParsedOperation { + const name = (try self.expect(.identifier)).text; + + if (std.mem.startsWith(u8, name, "cmp_")) { + const opcode_name = name["cmp_".len..]; + const opcode = std.meta.stringToEnum(inst_ir.CompareOpcode, opcode_name) orelse return error.InvalidOpcode; + const lhs = try self.parseValueRef(); + + try self.expectDiscard(.comma); + + return .{ + .compare = .{ + .opcode = opcode, + .lhs = lhs, + .rhs = try self.parseValueRef(), + }, + }; + } + + if (std.meta.stringToEnum(inst_ir.UnaryOpcode, name)) |opcode| { + return .{ .unary = .{ .opcode = opcode, .operand = try self.parseValueRef() } }; + } + + if (std.meta.stringToEnum(inst_ir.BinaryOpcode, name)) |opcode| { + const lhs = try self.parseValueRef(); + try self.expectDiscard(.comma); + return .{ .binary = .{ .opcode = opcode, .lhs = lhs, .rhs = try self.parseValueRef() } }; + } + + if (std.mem.eql(u8, name, "select")) { + const condition = try self.parseValueRef(); + try self.expectDiscard(.comma); + const true_value = try self.parseValueRef(); + try self.expectDiscard(.comma); + return .{ .select = .{ + .condition = condition, + .true_value = true_value, + .false_value = try self.parseValueRef(), + } }; + } + + if (std.mem.eql(u8, name, "bitcast")) + return .{ .bitcast = try self.parseValueRef() }; + + if (std.mem.eql(u8, name, "composite_construct")) + return .{ .composite_construct = try self.parseTrailingValueList() }; + + if (std.mem.eql(u8, name, "composite_extract")) { + const composite = try self.parseValueRef(); + var indices: std.ArrayList(u32) = .empty; + + while (try self.consume(.left_square)) { + try indices.append(self.allocator, try self.parseUnsigned(u32, .number)); + try self.expectDiscard(.right_square); + } + + if (indices.items.len == 0) + return error.InvalidCompositeIndex; + return .{ + .composite_extract = .{ + .composite = composite, + .indices = indices.items, + }, + }; + } + + if (std.mem.eql(u8, name, "load_interface")) + return .{ + .load_interface = (try self.expect(.at_name)).text, + }; + + if (std.mem.eql(u8, name, "store_interface")) { + const interface_name = (try self.expect(.at_name)).text; + try self.expectDiscard(.comma); + return .{ + .store_interface = .{ + .interface_name = interface_name, + .value = try self.parseValueRef(), + }, + }; + } + + if (std.mem.eql(u8, name, "call")) { + const function_name = (try self.expect(.at_name)).text; + try self.expectDiscard(.left_paren); + + const arguments = try self.parseDelimitedValueList(.right_paren); + try self.expectDiscard(.right_paren); + + return .{ + .call = .{ + .function_name = function_name, + .arguments = arguments, + }, + }; + } + + return error.InvalidOpcode; + } + + fn parseTerminator(self: *Parser) !ParsedTerminator { + const name = (try self.expect(.identifier)).text; + + if (std.mem.eql(u8, name, "branch")) + return .{ .branch = try self.parseEdge() }; + + if (std.mem.eql(u8, name, "conditional_branch")) { + const condition = try self.parseValueRef(); + try self.expectDiscard(.comma); + + const true_edge = try self.parseEdge(); + try self.expectDiscard(.comma); + + return .{ + .conditional_branch = .{ + .condition = condition, + .true_edge = true_edge, + .false_edge = try self.parseEdge(), + }, + }; + } + + if (std.mem.eql(u8, name, "return")) { + if ((try self.peek()).tag == .value_ref) + return .{ .return_value = try self.parseValueRef() }; + return .return_void; + } + + if (std.mem.eql(u8, name, "discard")) + return .discard; + + if (std.mem.eql(u8, name, "unreachable")) + return .unreachable_value; + + return error.UnexpectedToken; + } + + fn parseEdge(self: *Parser) !ParsedEdge { + const block_name = (try self.expect(.dot_name)).text; + try self.expectDiscard(.left_paren); + + const arguments = try self.parseDelimitedValueList(.right_paren); + try self.expectDiscard(.right_paren); + + return .{ + .block_name = block_name, + .arguments = arguments, + }; + } + + fn parseType(self: *Parser) !ids.TypeId { + const token = try self.expect(.identifier); + const module = self.module.?; + + if (std.mem.eql(u8, token.text, "void")) + return module.internType(.void); + + if (std.mem.eql(u8, token.text, "bool")) + return module.internType(.boolean); + + if (std.mem.startsWith(u8, token.text, "vec")) { + const length = parseTextUnsigned(u8, token.text[3..]) catch return error.InvalidType; + try self.expectDiscard(.left_square); + + const element_type = try self.parseType(); + try self.expectDiscard(.right_square); + + return module.internType(.{ + .vector = .{ + .element_type = element_type, + .length = length, + }, + }); + } + + if (std.mem.eql(u8, token.text, "array")) { + try self.expectDiscard(.left_square); + + const element_type = try self.parseType(); + try self.expectDiscard(.comma); + + const length = try self.parseUnsigned(u32, .number); + try self.expectDiscard(.right_square); + + return module.internType(.{ + .array = .{ + .element_type = element_type, + .length = length, + }, + }); + } + + if (std.mem.eql(u8, token.text, "struct")) { + try self.expectDiscard(.left_square); + var members: std.ArrayList(ids.TypeId) = .empty; + + if ((try self.peek()).tag != .right_square) { + while (true) { + try members.append(self.allocator, try self.parseType()); + if (!try self.consume(.comma)) + break; + } + } + + try self.expectDiscard(.right_square); + + return module.internType(.{ + .structure = .{ + .members = members.items, + }, + }); + } + + if (std.mem.eql(u8, token.text, "ptr")) { + try self.expectDiscard(.left_square); + + const address_name = (try self.expect(.identifier)).text; + const address_space = std.meta.stringToEnum(type_ir.AddressSpace, address_name) orelse return error.InvalidType; + try self.expectDiscard(.comma); + + const pointee_type = try self.parseType(); + try self.expectDiscard(.right_square); + + return module.internType(.{ + .pointer = .{ + .address_space = address_space, + .pointee_type = pointee_type, + }, + }); + } + + if (std.mem.eql(u8, token.text, "resourceHandle")) { + try self.expectDiscard(.left_square); + + const kind_name = (try self.expect(.identifier)).text; + const kind = std.meta.stringToEnum(type_ir.ResourceKind, kind_name) orelse return error.InvalidType; + + try self.expectDiscard(.right_square); + + return module.internType(.{ + .resource_handle = .{ + .kind = kind, + }, + }); + } + + if (token.text.len > 1 and (token.text[0] == 'i' or token.text[0] == 'u')) { + const bits = parseTextUnsigned(u16, token.text[1..]) catch return error.InvalidType; + return module.internType(.{ .integer = .{ + .bits = bits, + .signedness = if (token.text[0] == 'i') .signed else .unsigned, + } }); + } + + if (token.text.len > 1 and token.text[0] == 'f') { + const bits = parseTextUnsigned(u16, token.text[1..]) catch return error.InvalidType; + return module.internType(.{ + .floating = .{ + .bits = bits, + }, + }); + } + return error.InvalidType; + } + + fn parseConstantList(self: *Parser) ![]const u32 { + try self.expectDiscard(.left_square); + var values: std.ArrayList(u32) = .empty; + + if ((try self.peek()).tag != .right_square) { + while (true) { + try values.append(self.allocator, try self.parseUnsigned(u32, .constant_ref)); + if (!try self.consume(.comma)) + break; + } + } + + try self.expectDiscard(.right_square); + return values.items; + } + + fn parseTrailingValueList(self: *Parser) ![]const ValueRef { + var values: std.ArrayList(ValueRef) = .empty; + if ((try self.peek()).tag != .value_ref) + return values.items; + + while (true) { + try values.append(self.allocator, try self.parseValueRef()); + if (!try self.consume(.comma)) + break; + } + + return values.items; + } + + fn parseDelimitedValueList(self: *Parser, closing: TokenTag) ![]const ValueRef { + var values: std.ArrayList(ValueRef) = .empty; + if ((try self.peek()).tag == closing) + return values.items; + + while (true) { + try values.append(self.allocator, try self.parseValueRef()); + if (!try self.consume(.comma)) + break; + } + + return values.items; + } + + fn parseValueRef(self: *Parser) !ValueRef { + return (try self.expect(.value_ref)).text; + } + + fn parseUnsigned(self: *Parser, comptime T: type, tag: TokenTag) !T { + const token = try self.expect(tag); + return parseTextUnsigned(T, token.text) catch error.InvalidNumber; + } + + fn expectIdentifier(self: *Parser, expected: []const u8) !void { + const token = try self.expect(.identifier); + if (!std.mem.eql(u8, token.text, expected)) + return error.UnexpectedToken; + } + + fn expectDiscard(self: *Parser, tag: TokenTag) !void { + _ = try self.expect(tag); + } + + fn expect(self: *Parser, tag: TokenTag) !Token { + const token = try self.take(); + if (token.tag != tag) + return error.UnexpectedToken; + return token; + } + + fn consume(self: *Parser, tag: TokenTag) !bool { + if ((try self.peek()).tag != tag) + return false; + _ = try self.take(); + return true; + } + + fn peek(self: *Parser) !Token { + return self.lexer.peek(); + } + + fn take(self: *Parser) !Token { + const token = self.lexer.take(); + if (token.tag == .invalid) + return error.UnexpectedToken; + return token; + } +}; + +fn isTerminatorName(name: []const u8) bool { + return std.mem.eql(u8, name, "branch") or + std.mem.eql(u8, name, "conditional_branch") or + std.mem.eql(u8, name, "return") or + std.mem.eql(u8, name, "discard") or + std.mem.eql(u8, name, "unreachable"); +} + +fn parseIntegerLiteral(integer: type_ir.IntegerType, text: []const u8) !u64 { + if (integer.bits == 0 or integer.bits > 64) + return error.InvalidType; + + if (integer.signedness == .unsigned) { + const value = std.fmt.parseInt(u64, text, 10) catch return error.InvalidNumber; + if (integer.bits < 64) { + const shift: u6 = @intCast(integer.bits); + const maximum = (@as(u64, 1) << shift) - 1; + if (value > maximum) + return error.InvalidNumber; + } + return value; + } + + const value = std.fmt.parseInt(i64, text, 10) catch return error.InvalidNumber; + if (integer.bits < 64) { + const sign_shift: u6 = @intCast(integer.bits - 1); + const magnitude = @as(i64, 1) << sign_shift; + if (value < -magnitude or value > magnitude - 1) + return error.InvalidNumber; + + const width: u6 = @intCast(integer.bits); + const mask = (@as(u64, 1) << width) - 1; + return @as(u64, @bitCast(value)) & mask; + } + + return @bitCast(value); +} + +fn parseFloatLiteral(bits: u16, text: []const u8) !u64 { + return switch (bits) { + 16 => blk: { + const value = std.fmt.parseFloat(f16, text) catch return error.InvalidNumber; + break :blk @as(u16, @bitCast(value)); + }, + 32 => blk: { + const value = std.fmt.parseFloat(f32, text) catch return error.InvalidNumber; + break :blk @as(u32, @bitCast(value)); + }, + 64 => blk: { + const value = std.fmt.parseFloat(f64, text) catch return error.InvalidNumber; + break :blk @as(u64, @bitCast(value)); + }, + else => error.InvalidType, + }; +} + +fn parseTextUnsigned(comptime T: type, text: []const u8) !T { + const base: u8 = if (std.mem.startsWith(u8, text, "0x")) 16 else 10; + const digits = if (base == 16) text[2..] else text; + + if (digits.len == 0) + return error.InvalidNumber; + + return std.fmt.parseInt(T, digits, base); +} + +pub fn parseString(backing_allocator: std.mem.Allocator, source: []const u8) !module_ir.Module { + var temporary = std.heap.ArenaAllocator.init(backing_allocator); + defer temporary.deinit(); + const temporary_allocator = temporary.allocator(); + + var parser: Parser = .{ + .lexer = .init(source), + .allocator = temporary_allocator, + }; + + try parser.expectIdentifier("shader"); + const stage_token = try parser.expect(.identifier); + const stage = std.meta.stringToEnum(module_ir.Stage, stage_token.text) orelse return error.InvalidStage; + + var module = module_ir.Module.init(backing_allocator, stage); + errdefer module.deinit(); + parser.module = &module; + + const entry_point_name = if ((try parser.peek()).tag == .at_name) + (try parser.take()).text + else + null; + + try parser.expectDiscard(.left_brace); + + var parsed: ParsedModule = .{ .entry_point_name = entry_point_name }; + while ((try parser.peek()).tag != .right_brace) { + const token = try parser.peek(); + switch (token.tag) { + .value_ref => try parsed.constants.append(temporary_allocator, try parser.parseConstant()), + .at_name => try parsed.interfaces.append(temporary_allocator, try parser.parseInterface()), + .identifier => { + if (std.mem.eql(u8, token.text, "fn")) { + try parsed.functions.append(temporary_allocator, try parser.parseFunction()); + } else { + return error.UnexpectedToken; + } + }, + else => return error.UnexpectedToken, + } + } + try parser.expectDiscard(.right_brace); + try parser.expectDiscard(.eof); + + try lowerer.lower(temporary_allocator, &module, &parsed); + try validator.validate(&module); + return module; +} + +pub fn parseFile(backing_allocator: std.mem.Allocator, io: std.Io, path: []const u8) !module_ir.Module { + return parseFileInDir(backing_allocator, io, std.Io.Dir.cwd(), path); +} + +pub fn parseFileInDir(backing_allocator: std.mem.Allocator, io: std.Io, directory: std.Io.Dir, path: []const u8) !module_ir.Module { + const file = try directory.openFile(io, path, .{}); + defer file.close(io); + + var buffer: [4096]u8 = @splat(0); + var reader = file.reader(io, &buffer); + + const source = try reader.interface.allocRemaining(backing_allocator, .limited(max_file_size)); + defer backing_allocator.free(source); + + return parseString(backing_allocator, source); +} diff --git a/src/compiler/ir/pass_manager.zig b/src/compiler/ir/pass_manager.zig new file mode 100644 index 0000000..2a0eda1 --- /dev/null +++ b/src/compiler/ir/pass_manager.zig @@ -0,0 +1,83 @@ +const std = @import("std"); +const module_ir = @import("module.zig"); +const validator = @import("validator/validator.zig"); + +pub const Context = struct { + allocator: std.mem.Allocator, + validate_after_each_pass: bool = true, +}; + +pub const Pass = struct { + name: []const u8, + required: module_ir.Properties = .{}, + produced: module_ir.Properties = .{}, + invalidated: module_ir.Properties = .{}, + run: *const fn (module: *module_ir.Module, context: *Context) anyerror!bool, +}; + +pub const Manager = struct { + allocator: std.mem.Allocator, + passes: std.ArrayList(Pass) = .empty, + + pub fn init(allocator: std.mem.Allocator) Manager { + return .{ .allocator = allocator }; + } + + pub fn deinit(self: *Manager) void { + self.passes.deinit(self.allocator); + self.* = undefined; + } + + pub fn add(self: *Manager, pass: Pass) !void { + try self.passes.append(self.allocator, pass); + } + + pub fn run(self: *Manager, module: *module_ir.Module, context: *Context) !bool { + var changed = false; + for (self.passes.items) |pass| { + if (!satisfies(module.properties, pass.required)) + return error.RequiredPropertyMissing; + + changed = (try pass.run(module, context)) or changed; + applyInvalidated(&module.properties, pass.invalidated); + applyProduced(&module.properties, pass.produced); + + if (context.validate_after_each_pass) + try validator.validate(module); + } + return changed; + } +}; + +fn satisfies(actual: module_ir.Properties, required: module_ir.Properties) bool { + inline for (property_names) |name| { + if (@field(required, name) and !@field(actual, name)) + return false; + } + return true; +} + +fn applyProduced(properties: *module_ir.Properties, produced: module_ir.Properties) void { + inline for (property_names) |name| { + if (@field(produced, name)) + @field(properties, name) = true; + } +} + +fn applyInvalidated(properties: *module_ir.Properties, invalidated: module_ir.Properties) void { + inline for (property_names) |name| { + if (@field(invalidated, name)) + @field(properties, name) = false; + } +} + +const property_names = .{ + "valid_cfg", + "valid_ssa", + "structured_control_flow", + "no_function_calls", + "no_local_memory", + "no_matrix_types", + "no_large_composites", + "explicit_resource_offsets", +}; diff --git a/src/compiler/ir/printer.zig b/src/compiler/ir/printer.zig new file mode 100644 index 0000000..2350ca7 --- /dev/null +++ b/src/compiler/ir/printer.zig @@ -0,0 +1,336 @@ +const std = @import("std"); +const ids = @import("id.zig"); +const inst_ir = @import("instruction.zig"); +const module_ir = @import("module.zig"); + +const indent = " "; + +pub fn write(module: *const module_ir.Module, writer: *std.Io.Writer) std.Io.Writer.Error!void { + try writer.print("shader {t}", .{module.stage}); + + if (module.entry_point) |entry| { + try writer.writeByte(' '); + try writeFunctionRef(module, writer, entry); + } + + try writer.writeAll("\n{\n"); + + for (module.interface_variables.entries.items, 0..) |entry, index| { + const variable = entry orelse continue; + + try writer.writeAll(indent); + try writeNamedRef(writer, variable.name, "interface", index); + try writer.writeAll(": "); + try writeType(module, writer, variable.type); + try writer.print(" = {t}[", .{variable.direction}); + + switch (variable.semantic) { + .location => |location| try writer.print("location({d}), component({d}), index({d})", .{ location.location, location.component, location.index }), + .builtin => |builtin| try writer.print("builtin({t})", .{builtin}), + } + + try writer.writeAll("]\n"); + } + + for (module.constants.entries.items, 0..) |entry, constant_index| { + const constant = entry orelse continue; + const value_id = constantValueId(module, ids.ConstantId.fromIndex(constant_index)) orelse continue; + + try writer.writeAll(indent); + try writeValueRef(module, writer, value_id); + try writer.writeAll(": constant "); + try writeType(module, writer, constant.type); + try writer.writeAll(" = "); + + switch (constant.value) { + .boolean => |value| try writer.print("{}", .{value}), + .integer_bits => |bits| try writer.print("bits(0x{x})", .{bits}), + .float_bits => |bits| try writer.print("bits(0x{x})", .{bits}), + .null => try writer.writeAll("null"), + .undef => try writer.writeAll("undef"), + .composite => |elements| { + try writer.writeByte('['); + for (elements, 0..) |element, index| { + if (index != 0) try writer.writeAll(", "); + try writer.print("#{d}", .{element.index()}); + } + try writer.writeByte(']'); + }, + } + + try writer.writeByte('\n'); + } + + for (module.functions.entries.items, 0..) |entry, function_index| { + const function = entry orelse continue; + + try writer.writeAll("\n" ++ indent ++ "fn "); + try writeNamedRef(writer, function.name, "fn", function_index); + try writer.writeByte('('); + for (function.parameters.items, 0..) |parameter, index| { + if (index != 0) + try writer.writeAll(", "); + try writeValueRef(module, writer, parameter); + try writer.writeAll(": "); + try writeType(module, writer, function.parameter_types.items[index]); + } + try writer.writeAll(") -> "); + try writeType(module, writer, function.return_type); + try writer.writeAll("\n" ++ indent ++ "{\n"); + + for (function.blocks.items) |block_id| { + const block = module.blocks.get(block_id) orelse continue; + + try writer.writeAll(indent ** 2); + try writeBlockRef(module, writer, block_id); + try writer.writeByte('('); + for (block.parameters.items, 0..) |parameter, index| { + if (index != 0) try writer.writeAll(", "); + try writeValueRef(module, writer, parameter); + try writer.writeAll(": "); + try writeType(module, writer, module.typeOf(parameter).?); + } + try writer.writeAll("):\n"); + + for (block.instructions.items) |instruction_id| { + const instruction = module.instructions.get(instruction_id) orelse continue; + try writer.writeAll(indent ** 3); + if (instruction.result) |result| { + try writeValueRef(module, writer, result); + try writer.writeAll(": "); + try writeType(module, writer, module.typeOf(result).?); + try writer.writeAll(" = "); + } + try writeOperation(module, writer, instruction.operation); + try writer.writeByte('\n'); + } + + if (block.terminator) |terminator| { + try writer.writeAll(indent ** 3); + try writeTerminator(module, writer, terminator); + try writer.writeAll("\n\n"); + } else { + try writer.writeAll(indent ** 3 ++ "\n\n"); + } + } + try writer.writeAll(indent ++ "}\n"); + } + try writer.writeAll("}\n"); +} + +pub fn allocPrint(allocator: std.mem.Allocator, module: *const module_ir.Module) ![]u8 { + var output: std.Io.Writer.Allocating = .init(allocator); + defer output.deinit(); + try write(module, &output.writer); + return output.toOwnedSlice(); +} + +fn writeType(module: *const module_ir.Module, writer: *std.Io.Writer, type_id: ids.TypeId) !void { + const ty = module.types.get(type_id) orelse { + try writer.print("", .{type_id.index()}); + return; + }; + + switch (ty.*) { + .void => try writer.writeAll("void"), + .boolean => try writer.writeAll("bool"), + .integer => |integer| try writer.print("{s}{d}", .{ if (integer.signedness == .signed) "i" else "u", integer.bits }), + .floating => |float| try writer.print("f{d}", .{float.bits}), + .vector => |vector| { + try writer.print("vec{d}[", .{vector.length}); + try writeType(module, writer, vector.element_type); + try writer.writeByte(']'); + }, + .array => |array| { + try writer.writeAll("array["); + try writeType(module, writer, array.element_type); + try writer.print(", {d}]", .{array.length}); + }, + .structure => |structure| { + try writer.writeAll("struct["); + for (structure.members, 0..) |member, index| { + if (index != 0) try writer.writeAll(", "); + try writeType(module, writer, member); + } + try writer.writeByte(']'); + }, + .pointer => |pointer| { + try writer.print("ptr[{t}, ", .{pointer.address_space}); + try writeType(module, writer, pointer.pointee_type); + try writer.writeByte(']'); + }, + .resource_handle => |handle| try writer.print("resourceHandle[{t}]", .{handle.kind}), + } +} + +fn writeOperation(module: *const module_ir.Module, writer: *std.Io.Writer, operation: inst_ir.Operation) !void { + switch (operation) { + .unary => |op| { + try writer.print("{t} ", .{op.opcode}); + try writeValueRef(module, writer, op.operand); + }, + .binary => |op| { + try writer.print("{t} ", .{op.opcode}); + try writeValueRef(module, writer, op.lhs); + try writer.writeAll(", "); + try writeValueRef(module, writer, op.rhs); + }, + .compare => |op| { + try writer.print("cmp_{t} ", .{op.opcode}); + try writeValueRef(module, writer, op.lhs); + try writer.writeAll(", "); + try writeValueRef(module, writer, op.rhs); + }, + .select => |op| { + try writer.writeAll("select "); + try writeValueRef(module, writer, op.condition); + try writer.writeAll(", "); + try writeValueRef(module, writer, op.true_value); + try writer.writeAll(", "); + try writeValueRef(module, writer, op.false_value); + }, + .bitcast => |value| { + try writer.writeAll("bitcast "); + try writeValueRef(module, writer, value); + }, + .composite_construct => |op| { + try writer.writeAll("composite_construct "); + try writeValueList(module, writer, op.elements); + }, + .composite_extract => |op| { + try writer.writeAll("composite_extract "); + try writeValueRef(module, writer, op.composite); + for (op.indices) |index| try writer.print("[{d}]", .{index}); + }, + .load_interface => |op| { + try writer.writeAll("load_interface "); + const variable = module.interface_variables.get(op.variable); + try writeNamedRef(writer, if (variable) |v| v.name else null, "interface", op.variable.index()); + }, + .store_interface => |op| { + try writer.writeAll("store_interface "); + const variable = module.interface_variables.get(op.variable); + try writeNamedRef(writer, if (variable) |v| v.name else null, "interface", op.variable.index()); + try writer.writeAll(", "); + try writeValueRef(module, writer, op.value); + }, + .call => |op| { + try writer.writeAll("call "); + try writeFunctionRef(module, writer, op.function); + try writer.writeByte('('); + try writeValueList(module, writer, op.arguments); + try writer.writeByte(')'); + }, + } +} + +fn writeTerminator(module: *const module_ir.Module, writer: *std.Io.Writer, terminator: module_ir.Terminator) !void { + switch (terminator) { + .branch => |edge| { + try writer.writeAll("branch "); + try writeEdge(module, writer, edge); + }, + .conditional_branch => |branch| { + try writer.writeAll("conditional_branch "); + try writeValueRef(module, writer, branch.condition); + try writer.writeAll(", "); + try writeEdge(module, writer, branch.true_edge); + try writer.writeAll(", "); + try writeEdge(module, writer, branch.false_edge); + }, + .return_void => try writer.writeAll("return"), + .return_value => |value| { + try writer.writeAll("return "); + try writeValueRef(module, writer, value); + }, + .discard => try writer.writeAll("discard"), + .@"unreachable" => try writer.writeAll("unreachable"), + } +} + +fn writeEdge(module: *const module_ir.Module, writer: *std.Io.Writer, edge: module_ir.Edge) !void { + try writeBlockRef(module, writer, edge.target); + try writer.writeByte('('); + try writeValueList(module, writer, edge.arguments); + try writer.writeByte(')'); +} + +fn writeValueList(module: *const module_ir.Module, writer: *std.Io.Writer, value_ids: []const ids.ValueId) !void { + for (value_ids, 0..) |value, index| { + if (index != 0) + try writer.writeAll(", "); + try writeValueRef(module, writer, value); + } +} + +fn writeValueRef(module: *const module_ir.Module, writer: *std.Io.Writer, value_id: ids.ValueId) !void { + try writer.writeByte('%'); + + const value = module.values.get(value_id); + if (value) |data| { + if (data.name) |name| { + if (isValidName(name) and isUniqueValueName(module, value_id, name)) { + try writer.writeAll(name); + return; + } + } + } + + try writer.print("{d}", .{value_id.index()}); +} + +fn isValidName(name: []const u8) bool { + if (name.len == 0 or (!std.ascii.isAlphabetic(name[0]) and name[0] != '_')) + return false; + + for (name[1..]) |byte| { + if (!std.ascii.isAlphanumeric(byte) and byte != '_') + return false; + } + return true; +} + +fn isUniqueValueName(module: *const module_ir.Module, value_id: ids.ValueId, name: []const u8) bool { + for (module.values.entries.items, 0..) |entry, index| { + if (index == value_id.index()) + continue; + + const other = entry orelse continue; + if (other.name) |other_name| { + if (std.mem.eql(u8, name, other_name)) + return false; + } + } + return true; +} + +fn writeBlockRef(module: *const module_ir.Module, writer: *std.Io.Writer, block: ids.BlockId) !void { + const value = module.blocks.get(block); + try writeNamedRef(writer, if (value) |b| b.name else null, "b", block.index()); +} + +fn writeFunctionRef(module: *const module_ir.Module, writer: *std.Io.Writer, function: ids.FunctionId) !void { + const value = module.functions.get(function); + try writeNamedRef(writer, if (value) |f| f.name else null, "fn", function.index()); +} + +fn writeNamedRef(writer: *std.Io.Writer, name: ?[]const u8, fallback: []const u8, index: usize) !void { + try writer.writeByte(if (std.mem.eql(u8, fallback, "b")) '.' else '@'); + if (name) |text| { + if (isValidName(text)) { + try writer.writeAll(text); + return; + } + } + try writer.print("{s}{d}", .{ fallback, index }); +} + +fn constantValueId(module: *const module_ir.Module, constant_id: ids.ConstantId) ?ids.ValueId { + for (module.values.entries.items, 0..) |entry, index| { + const value = entry orelse continue; + if (value.definition == .constant and value.definition.constant == constant_id) + return ids.ValueId.fromIndex(index); + } + + return null; +} diff --git a/src/compiler/ir/type.zig b/src/compiler/ir/type.zig new file mode 100644 index 0000000..439894f --- /dev/null +++ b/src/compiler/ir/type.zig @@ -0,0 +1,106 @@ +const std = @import("std"); +const ids = @import("id.zig"); + +pub const TypeId = ids.TypeId; + +pub const Signedness = enum { signed, unsigned }; + +pub const IntegerType = struct { + bits: u16, + signedness: Signedness, +}; + +pub const FloatType = struct { + bits: u16, +}; + +pub const VectorType = struct { + element_type: TypeId, + length: u8, +}; + +pub const ArrayType = struct { + element_type: TypeId, + length: u32, +}; + +pub const StructureType = struct { + members: []const TypeId, +}; + +pub const AddressSpace = enum { + function, + private, + workgroup, + input, + output, + uniform, + storage, + push_constant, + physical, +}; + +pub const PointerType = struct { + address_space: AddressSpace, + pointee_type: TypeId, +}; + +pub const ResourceKind = enum { + uniform_buffer, + storage_buffer, + sampled_image, + storage_image, + sampler, +}; + +pub const ResourceHandleType = struct { + kind: ResourceKind, + data_type: ?TypeId = null, +}; + +pub const Type = union(enum) { + void, + boolean, + integer: IntegerType, + floating: FloatType, + vector: VectorType, + array: ArrayType, + structure: StructureType, + pointer: PointerType, + resource_handle: ResourceHandleType, + + pub fn eql(a: Type, b: Type) bool { + return switch (a) { + .void => b == .void, + .boolean => b == .boolean, + .integer => |value| switch (b) { + .integer => |other| std.meta.eql(value, other), + else => false, + }, + .floating => |value| switch (b) { + .floating => |other| std.meta.eql(value, other), + else => false, + }, + .vector => |value| switch (b) { + .vector => |other| std.meta.eql(value, other), + else => false, + }, + .array => |value| switch (b) { + .array => |other| std.meta.eql(value, other), + else => false, + }, + .structure => |value| switch (b) { + .structure => |other| std.mem.eql(TypeId, value.members, other.members), + else => false, + }, + .pointer => |value| switch (b) { + .pointer => |other| std.meta.eql(value, other), + else => false, + }, + .resource_handle => |value| switch (b) { + .resource_handle => |other| std.meta.eql(value, other), + else => false, + }, + }; + } +}; diff --git a/src/compiler/ir/validator/dominance.zig b/src/compiler/ir/validator/dominance.zig new file mode 100644 index 0000000..4eba3e5 --- /dev/null +++ b/src/compiler/ir/validator/dominance.zig @@ -0,0 +1,93 @@ +const std = @import("std"); +const cfg = @import("../cfg.zig"); +const ids = @import("../id.zig"); +const module_ir = @import("../module.zig"); + +pub const Error = std.mem.Allocator.Error || error{ + InvalidBlock, + DefinitionDoesNotDominateUse, +}; + +const DominanceUseContext = struct { + module: *const module_ir.Module, + analysis: *const cfg, + function_id: ids.FunctionId, + use_block: ids.BlockId, + use_index: usize, + valid: bool = true, +}; + +pub fn validate(module: *const module_ir.Module, function_id: ids.FunctionId) Error!void { + var analysis = cfg.init(module.backingAllocator(), module, function_id) catch |err| switch (err) { + error.OutOfMemory => return error.OutOfMemory, + else => return error.InvalidBlock, + }; + defer analysis.deinit(); + + const function = module.functions.get(function_id).?; + for (function.blocks.items) |block_id| { + const block = module.blocks.get(block_id).?; + + for (block.instructions.items, 0..) |instruction_id, instruction_index| { + const instruction = module.instructions.get(instruction_id).?; + var context: DominanceUseContext = .{ + .module = module, + .analysis = &analysis, + .function_id = function_id, + .use_block = block_id, + .use_index = instruction_index, + }; + + instruction.operation.visitValueUses(&context, checkDominanceUse); + + if (!context.valid) + return error.DefinitionDoesNotDominateUse; + } + + var context: DominanceUseContext = .{ + .module = module, + .analysis = &analysis, + .function_id = function_id, + .use_block = block_id, + .use_index = block.instructions.items.len, + }; + + module_ir.visitTerminatorValueUses(block.terminator.?, &context, checkDominanceUse); + + if (!context.valid) + return error.DefinitionDoesNotDominateUse; + } +} + +fn checkDominanceUse(context: *DominanceUseContext, value_id: ids.ValueId) void { + if (!context.valid) + return; + + const value = context.module.values.get(value_id) orelse { + context.valid = false; + return; + }; + + context.valid = switch (value.definition) { + .constant, .undef => true, + .function_parameter => |definition| definition.function == context.function_id, + .block_parameter => |definition| context.analysis.dominates(definition.block, context.use_block), + .instruction => |instruction_id| blk: { + const definition = context.module.instructions.get(instruction_id) orelse break :blk false; + + if (!context.analysis.dominates(definition.parent_block, context.use_block)) + break :blk false; + + if (definition.parent_block != context.use_block) + break :blk true; + + const block = context.module.blocks.get(context.use_block) orelse break :blk false; + for (block.instructions.items, 0..) |candidate, definition_index| { + if (candidate == instruction_id) + break :blk definition_index < context.use_index; + } + + break :blk false; + }, + }; +} diff --git a/src/compiler/ir/validator/validator.zig b/src/compiler/ir/validator/validator.zig new file mode 100644 index 0000000..4c4d0a0 --- /dev/null +++ b/src/compiler/ir/validator/validator.zig @@ -0,0 +1,470 @@ +const std = @import("std"); +const ids = @import("../id.zig"); +const type_ir = @import("../type.zig"); +const inst_ir = @import("../instruction.zig"); +const module_ir = @import("../module.zig"); +const dominance = @import("dominance.zig"); + +pub const ValidationError = error{ + MissingEntryPoint, + InvalidEntryPoint, + InvalidType, + InvalidConstant, + InvalidValue, + InvalidFunction, + InvalidBlock, + InvalidInstruction, + MissingFunctionEntryBlock, + MissingTerminator, + EntryBlockHasPredecessor, + WrongParent, + WrongDefinition, + WrongParameterIndex, + WrongResultPresence, + WrongOperandType, + WrongResultType, + WrongBranchArgumentCount, + WrongBranchArgumentType, + CrossFunctionReference, + WrongReturnType, + WrongInterfaceDirection, + InvalidStructuredControl, + DefinitionDoesNotDominateUse, +}; + +pub const Error = ValidationError || std.mem.Allocator.Error; + +/// Early validator for the foundational IR. It covers object ownership, CFG +/// edges, single definitions, function boundaries, and the currently modeled +/// operation types, and SSA dominance. +pub fn validate(module: *const module_ir.Module) Error!void { + const entry_point = module.entry_point orelse return error.MissingEntryPoint; + if (!module.functions.isLive(entry_point)) + return error.InvalidEntryPoint; + + for (module.types.entries.items) |entry| { + const ty = entry orelse continue; + try validateType(module, ty); + } + + for (module.constants.entries.items) |entry| { + const constant = entry orelse continue; + + if (!module.types.isLive(constant.type)) + return error.InvalidType; + + if (constant.value == .composite) { + for (constant.value.composite) |element| { + if (!module.constants.isLive(element)) + return error.InvalidConstant; + } + } + } + + for (module.values.entries.items, 0..) |entry, value_index| { + const value = entry orelse continue; + + if (!module.types.isLive(value.type)) + return error.InvalidType; + + const value_id = ids.ValueId.fromIndex(value_index); + + switch (value.definition) { + .constant => |id| { + const constant = module.constants.get(id) orelse return error.InvalidConstant; + if (constant.type != value.type) + return error.WrongResultType; + }, + .function_parameter => |definition| { + const function = module.functions.get(definition.function) orelse return error.InvalidFunction; + if (definition.index >= function.parameters.items.len or function.parameters.items[definition.index] != value_id) + return error.WrongParameterIndex; + }, + .block_parameter => |definition| { + const block = module.blocks.get(definition.block) orelse return error.InvalidBlock; + if (definition.index >= block.parameters.items.len or block.parameters.items[definition.index] != value_id) + return error.WrongParameterIndex; + }, + .instruction => |instruction_id| { + const instruction = module.instructions.get(instruction_id) orelse return error.InvalidInstruction; + if (instruction.result != value_id) + return error.WrongDefinition; + }, + .undef => {}, + } + } + + for (module.interface_variables.entries.items) |entry| { + const variable = entry orelse continue; + if (!module.types.isLive(variable.type)) + return error.InvalidType; + } + + for (module.resources.entries.items) |entry| { + const resource = entry orelse continue; + if (!module.types.isLive(resource.type)) + return error.InvalidType; + } + + for (module.functions.entries.items, 0..) |entry, function_index| { + const function = entry orelse continue; + const function_id = ids.FunctionId.fromIndex(function_index); + + if (!module.types.isLive(function.return_type)) + return error.InvalidType; + + if (function.parameter_types.items.len != function.parameters.items.len) + return error.WrongParameterIndex; + + for (function.parameter_types.items, function.parameters.items, 0..) |parameter_type, parameter_id, index| { + const parameter = module.values.get(parameter_id) orelse return error.InvalidValue; + + if (parameter.type != parameter_type) + return error.WrongResultType; + + if (parameter.definition != .function_parameter or + parameter.definition.function_parameter.function != function_id or + parameter.definition.function_parameter.index != index) + return error.WrongDefinition; + } + + const entry_block = function.entry_block orelse return error.MissingFunctionEntryBlock; + const entry_block_value = module.blocks.get(entry_block) orelse return error.InvalidBlock; + if (entry_block_value.parent_function != function_id) return error.WrongParent; + + for (function.blocks.items) |block_id| { + const block = module.blocks.get(block_id) orelse return error.InvalidBlock; + + if (block.parent_function != function_id) + return error.WrongParent; + + try validateBlock(module, function_id, block_id, block); + } + + for (function.blocks.items) |block_id| { + const block = module.blocks.get(block_id).?; + + if (block.terminator) |terminator| { + if (targetsBlock(terminator, entry_block)) + return error.EntryBlockHasPredecessor; + } + } + + try dominance.validate(module, function_id); + } +} + +fn validateType(module: *const module_ir.Module, ty: type_ir.Type) ValidationError!void { + switch (ty) { + .vector => |vector| { + if (!module.types.isLive(vector.element_type) or vector.length < 2) + return error.InvalidType; + }, + .array => |array| { + if (!module.types.isLive(array.element_type) or array.length == 0) + return error.InvalidType; + }, + .structure => |structure| for (structure.members) |member| { + if (!module.types.isLive(member)) + return error.InvalidType; + }, + .pointer => |pointer| { + if (!module.types.isLive(pointer.pointee_type)) + return error.InvalidType; + }, + .resource_handle => |handle| if (handle.data_type) |data_type| { + if (!module.types.isLive(data_type)) + return error.InvalidType; + }, + else => {}, + } +} + +fn validateBlock( + module: *const module_ir.Module, + function_id: ids.FunctionId, + block_id: ids.BlockId, + block: *const module_ir.Block, +) ValidationError!void { + for (block.parameters.items, 0..) |parameter_id, index| { + const parameter = module.values.get(parameter_id) orelse return error.InvalidValue; + if (parameter.definition != .block_parameter or + parameter.definition.block_parameter.block != block_id or + parameter.definition.block_parameter.index != index) + return error.WrongDefinition; + } + + switch (block.structured_control) { + .none => {}, + .selection => |selection| try validateTarget(module, function_id, selection.merge_block), + .loop => |loop| { + try validateTarget(module, function_id, loop.merge_block); + try validateTarget(module, function_id, loop.continue_block); + }, + } + + for (block.instructions.items) |instruction_id| { + const instruction = module.instructions.get(instruction_id) orelse return error.InvalidInstruction; + + if (instruction.parent_block != block_id) + return error.WrongParent; + + if (instruction.result) |result_id| { + const result = module.values.get(result_id) orelse return error.InvalidValue; + + if (result.definition != .instruction or result.definition.instruction != instruction_id) + return error.WrongDefinition; + } + + try validateOperation(module, function_id, instruction); + } + + const terminator = block.terminator orelse return error.MissingTerminator; + try validateTerminator(module, function_id, terminator); +} + +fn validateOperation(module: *const module_ir.Module, function_id: ids.FunctionId, instruction: *const inst_ir.Instruction) ValidationError!void { + const result_type = if (instruction.result) |result| module.typeOf(result) orelse return error.InvalidValue else null; + + switch (instruction.operation) { + .unary => |op| { + const operand_type = try operandType(module, function_id, op.operand); + + if (result_type == null) + return error.WrongResultPresence; + + if (result_type.? != operand_type) + return error.WrongResultType; + }, + .binary => |op| { + const lhs_type = try operandType(module, function_id, op.lhs); + const rhs_type = try operandType(module, function_id, op.rhs); + + if (lhs_type != rhs_type) + return error.WrongOperandType; + + if (result_type == null or result_type.? != lhs_type) + return error.WrongResultType; + }, + .compare => |op| { + const lhs_type = try operandType(module, function_id, op.lhs); + if (try operandType(module, function_id, op.rhs) != lhs_type) + return error.WrongOperandType; + + const result = result_type orelse return error.WrongResultPresence; + if (!isBoolean(module, result)) + return error.WrongResultType; + }, + .select => |op| { + if (!isBoolean(module, try operandType(module, function_id, op.condition))) + return error.WrongOperandType; + + const true_type = try operandType(module, function_id, op.true_value); + if (try operandType(module, function_id, op.false_value) != true_type) + return error.WrongOperandType; + + if (result_type == null or result_type.? != true_type) + return error.WrongResultType; + }, + .bitcast => |operand| { + _ = try operandType(module, function_id, operand); + if (result_type == null) + return error.WrongResultPresence; + }, + .composite_construct => |op| { + const result = result_type orelse return error.WrongResultPresence; + const ty = module.types.get(result) orelse return error.InvalidType; + + switch (ty.*) { + .vector => |vector| { + if (op.elements.len != vector.length) + return error.WrongOperandType; + + for (op.elements) |element| { + if (try operandType(module, function_id, element) != vector.element_type) + return error.WrongOperandType; + } + }, + .structure => |structure| { + if (op.elements.len != structure.members.len) + return error.WrongOperandType; + + for (op.elements, structure.members) |element, member_type| { + if (try operandType(module, function_id, element) != member_type) + return error.WrongOperandType; + } + }, + else => return error.WrongResultType, + } + }, + .composite_extract => |op| { + const composite_type = try operandType(module, function_id, op.composite); + const extracted_type = try indexedType(module, composite_type, op.indices); + + if (result_type == null or result_type.? != extracted_type) + return error.WrongResultType; + }, + .load_interface => |op| { + const variable = module.interface_variables.get(op.variable) orelse return error.InvalidValue; + + if (variable.direction != .input) + return error.WrongInterfaceDirection; + + if (op.element_index) |index| + _ = try operandType(module, function_id, index); + + if (result_type == null or result_type.? != variable.type) + return error.WrongResultType; + }, + .store_interface => |op| { + if (result_type != null) + return error.WrongResultPresence; + + const variable = module.interface_variables.get(op.variable) orelse return error.InvalidValue; + + if (variable.direction != .output) + return error.WrongInterfaceDirection; + + if (try operandType(module, function_id, op.value) != variable.type) + return error.WrongOperandType; + + if (op.element_index) |index| + _ = try operandType(module, function_id, index); + }, + .call => |op| { + const callee = module.functions.get(op.function) orelse return error.InvalidFunction; + + if (op.arguments.len != callee.parameter_types.items.len) + return error.WrongOperandType; + + for (op.arguments, callee.parameter_types.items) |argument, parameter_type| { + if (try operandType(module, function_id, argument) != parameter_type) + return error.WrongOperandType; + } + + const return_type = module.types.get(callee.return_type) orelse return error.InvalidType; + + if (return_type.* == .void) { + if (result_type != null) + return error.WrongResultPresence; + } else if (result_type == null or result_type.? != callee.return_type) + return error.WrongResultType; + }, + } +} + +fn validateTerminator(module: *const module_ir.Module, function_id: ids.FunctionId, terminator: module_ir.Terminator) ValidationError!void { + const function = module.functions.get(function_id) orelse return error.InvalidFunction; + + switch (terminator) { + .branch => |edge| try validateEdge(module, function_id, edge), + .conditional_branch => |branch| { + if (!isBoolean(module, try operandType(module, function_id, branch.condition))) + return error.WrongOperandType; + + try validateEdge(module, function_id, branch.true_edge); + try validateEdge(module, function_id, branch.false_edge); + }, + .return_void => { + if (module.types.get(function.return_type).?.* != .void) + return error.WrongReturnType; + }, + .return_value => |value| { + if (try operandType(module, function_id, value) != function.return_type) + return error.WrongReturnType; + }, + .discard => { + if (module.stage != .fragment) + return error.WrongReturnType; + }, + .@"unreachable" => {}, + } +} + +fn validateEdge(module: *const module_ir.Module, function_id: ids.FunctionId, edge: module_ir.Edge) ValidationError!void { + const target = module.blocks.get(edge.target) orelse return error.InvalidBlock; + + if (target.parent_function != function_id) + return error.CrossFunctionReference; + + if (edge.arguments.len != target.parameters.items.len) + return error.WrongBranchArgumentCount; + + for (edge.arguments, target.parameters.items) |argument, parameter| { + if (try operandType(module, function_id, argument) != module.typeOf(parameter).?) + return error.WrongBranchArgumentType; + } +} + +fn validateTarget(module: *const module_ir.Module, function_id: ids.FunctionId, target_id: ids.BlockId) ValidationError!void { + const target = module.blocks.get(target_id) orelse return error.InvalidStructuredControl; + if (target.parent_function != function_id) + return error.InvalidStructuredControl; +} + +fn operandType(module: *const module_ir.Module, function_id: ids.FunctionId, value_id: ids.ValueId) ValidationError!ids.TypeId { + const value = module.values.get(value_id) orelse return error.InvalidValue; + const owner = valueFunction(module, value_id) catch return error.InvalidValue; + + if (owner) |actual| { + if (actual != function_id) + return error.CrossFunctionReference; + } + return value.type; +} + +fn valueFunction(module: *const module_ir.Module, value_id: ids.ValueId) ValidationError!?ids.FunctionId { + const value = module.values.get(value_id) orelse return error.InvalidValue; + + return switch (value.definition) { + .constant, .undef => null, + .function_parameter => |definition| definition.function, + .block_parameter => |definition| (module.blocks.get(definition.block) orelse return error.InvalidBlock).parent_function, + .instruction => |instruction_id| blk: { + const instruction = module.instructions.get(instruction_id) orelse return error.InvalidInstruction; + const block = module.blocks.get(instruction.parent_block) orelse return error.InvalidBlock; + break :blk block.parent_function; + }, + }; +} + +fn indexedType(module: *const module_ir.Module, root: ids.TypeId, indices: []const u32) ValidationError!ids.TypeId { + if (indices.len == 0) + return error.WrongOperandType; + + var current = root; + + for (indices) |index| { + const ty = module.types.get(current) orelse return error.InvalidType; + current = switch (ty.*) { + .vector => |vector| if (index < vector.length) + vector.element_type + else + return error.WrongOperandType, + .array => |array| if (index < array.length) + array.element_type + else + return error.WrongOperandType, + .structure => |structure| if (index < structure.members.len) + structure.members[index] + else + return error.WrongOperandType, + + else => return error.WrongOperandType, + }; + } + return current; +} + +fn isBoolean(module: *const module_ir.Module, type_id: ids.TypeId) bool { + const ty = module.types.get(type_id) orelse return false; + return ty.* == .boolean; +} + +fn targetsBlock(terminator: module_ir.Terminator, target: ids.BlockId) bool { + return switch (terminator) { + .branch => |edge| edge.target == target, + .conditional_branch => |branch| branch.true_edge.target == target or branch.false_edge.target == target, + else => false, + }; +} diff --git a/src/compiler/ir/value.zig b/src/compiler/ir/value.zig new file mode 100644 index 0000000..07ee5d5 --- /dev/null +++ b/src/compiler/ir/value.zig @@ -0,0 +1,31 @@ +const ids = @import("id.zig"); + +pub const TypeId = ids.TypeId; +pub const ConstantId = ids.ConstantId; +pub const InstructionId = ids.InstructionId; +pub const FunctionId = ids.FunctionId; +pub const BlockId = ids.BlockId; + +pub const FunctionParameterDefinition = struct { + function: FunctionId, + index: u32, +}; + +pub const BlockParameterDefinition = struct { + block: BlockId, + index: u32, +}; + +pub const Definition = union(enum) { + constant: ConstantId, + function_parameter: FunctionParameterDefinition, + block_parameter: BlockParameterDefinition, + instruction: InstructionId, + undef, +}; + +pub const Value = struct { + type: TypeId, + definition: Definition, + name: ?[]const u8 = null, +}; diff --git a/src/compiler/ir/visitor.zig b/src/compiler/ir/visitor.zig new file mode 100644 index 0000000..3242293 --- /dev/null +++ b/src/compiler/ir/visitor.zig @@ -0,0 +1,98 @@ +const ids = @import("id.zig"); +const instruction_ir = @import("instruction.zig"); +const module_ir = @import("module.zig"); + +pub const Visitor = struct { + context: ?*anyopaque = null, + visitInterfaceVariable: ?*const fn (?*anyopaque, ids.InterfaceVariableId, *const module_ir.InterfaceVariable) anyerror!void = null, + visitResource: ?*const fn (?*anyopaque, ids.ResourceId, *const module_ir.Resource) anyerror!void = null, + visitFunction: ?*const fn (?*anyopaque, ids.FunctionId, *const module_ir.Function) anyerror!void = null, + visitBlock: ?*const fn (?*anyopaque, ids.BlockId, *const module_ir.Block) anyerror!void = null, + visitInstruction: ?*const fn (?*anyopaque, ids.InstructionId, *const instruction_ir.Instruction) anyerror!void = null, + visitTerminator: ?*const fn (?*anyopaque, ids.BlockId, module_ir.Terminator) anyerror!void = null, + visitValueUse: ?*const fn (?*anyopaque, ids.BlockId, ?ids.InstructionId, ids.ValueId) anyerror!void = null, +}; + +const UseContext = struct { + visitor_context: ?*anyopaque, + callback: *const fn (?*anyopaque, ids.BlockId, ?ids.InstructionId, ids.ValueId) anyerror!void, + block: ids.BlockId, + instruction: ?ids.InstructionId, + failure: ?anyerror = null, +}; + +pub fn walk(module: *const module_ir.Module, visitor: Visitor) !void { + for (module.interface_variables.entries.items, 0..) |entry, index| { + const variable = entry orelse continue; + if (visitor.visitInterfaceVariable) |callback| + try callback(visitor.context, ids.InterfaceVariableId.fromIndex(index), &variable); + } + + for (module.resources.entries.items, 0..) |entry, index| { + const resource = entry orelse continue; + if (visitor.visitResource) |callback| + try callback(visitor.context, ids.ResourceId.fromIndex(index), &resource); + } + + for (module.functions.entries.items, 0..) |entry, function_index| { + const function = entry orelse continue; + const function_id = ids.FunctionId.fromIndex(function_index); + + if (visitor.visitFunction) |callback| + try callback(visitor.context, function_id, &function); + + for (function.blocks.items) |block_id| { + const block = module.blocks.get(block_id) orelse continue; + + if (visitor.visitBlock) |callback| + try callback(visitor.context, block_id, block); + + for (block.instructions.items) |instruction_id| { + const instruction = module.instructions.get(instruction_id) orelse continue; + + if (visitor.visitInstruction) |callback| + try callback(visitor.context, instruction_id, instruction); + + if (visitor.visitValueUse) |callback| { + var context: UseContext = .{ + .visitor_context = visitor.context, + .callback = callback, + .block = block_id, + .instruction = instruction_id, + }; + instruction.operation.visitValueUses(&context, visitUse); + + if (context.failure) |failure| + return failure; + } + } + + if (block.terminator) |terminator| { + if (visitor.visitTerminator) |callback| + try callback(visitor.context, block_id, terminator); + + if (visitor.visitValueUse) |callback| { + var context: UseContext = .{ + .visitor_context = visitor.context, + .callback = callback, + .block = block_id, + .instruction = null, + }; + module_ir.visitTerminatorValueUses(terminator, &context, visitUse); + + if (context.failure) |failure| + return failure; + } + } + } + } +} + +fn visitUse(context: *UseContext, value: ids.ValueId) void { + if (context.failure != null) + return; + + context.callback(context.visitor_context, context.block, context.instruction, value) catch |err| { + context.failure = err; + }; +} diff --git a/src/compiler/root.zig b/src/compiler/root.zig new file mode 100644 index 0000000..c8b1c6e --- /dev/null +++ b/src/compiler/root.zig @@ -0,0 +1,826 @@ +//! ## Shader compiler infrastructure. +//! +//! This module exposes the project-specific intermediate representation in +//! `ir` and the SPIR-V frontend in `spirv`. +//! +//! Together they form +//! the first stage of the compiler pipeline: SPIR-V binary modules are decoded, +//! translated into a smaller and easier-to-transform IR, validated, and then made +//! available to later optimization or code-generation passes. + +const std = @import("std"); + +pub const ir = @import("ir/ir.zig"); +pub const spirv = @import("spirv/root.zig"); + +const VisitorStatistics = struct { + functions: usize = 0, + blocks: usize = 0, +}; + +test "IR builder generation" { + // shader vertex @main + // { + // @color: vec4[f32] = input[location(0), component(0), index(0)] + // @out_color: vec4[f32] = output[location(0), component(0), index(0)] + // %0: constant bool = true + // %1: constant f32 = bits(0x3f800000) + // + // fn @main() -> void + // { + // .entry(): + // %3: vec4[f32] = load_interface @color + // conditional_branch %0, .pass(), .merge(%3) + // + // .pass(): + // %4: vec4[f32] = composite_construct %1, %1, %1, %1 + // branch .merge(%4) + // + // .merge(%2: vec4[f32]): + // store_interface @out_color, %2 + // return + // } + // } + + var module = ir.module.Module.init(std.testing.allocator, .vertex); + defer module.deinit(); + var builder = ir.Builder.init(&module); + + const void_type = try builder.internType(.void); + const bool_type = try builder.internType(.boolean); + const f32_type = try builder.internType(.{ .floating = .{ .bits = 32 } }); + const duplicate_f32 = try builder.internType(.{ .floating = .{ .bits = 32 } }); + try std.testing.expectEqual(f32_type, duplicate_f32); + const vec4_type = try builder.internType(.{ .vector = .{ .element_type = f32_type, .length = 4 } }); + + const true_value = try builder.internConstant(bool_type, .{ .boolean = true }); + const one = try builder.internConstant(f32_type, .{ .float_bits = @as(u32, @bitCast(@as(f32, 1.0))) }); + + const input = try builder.addInterfaceVariable(vec4_type, .input, .{ .location = .{ .location = 0 } }, "color"); + const output = try builder.addInterfaceVariable(vec4_type, .output, .{ .location = .{ .location = 0 } }, "out_color"); + const main = try builder.addFunction(void_type, "main"); + builder.setEntryPoint(main); + const entry = try builder.addBlock(main, "entry"); + const pass = try builder.addBlock(main, "pass"); + const merge = try builder.addBlock(main, "merge"); + const merged = try builder.addBlockParameter(merge, vec4_type, "merged"); + + const loaded = (try builder.appendInstruction(entry, vec4_type, .{ + .load_interface = .{ .variable = input }, + }, "loaded")).?; + try builder.setTerminator(entry, .{ .conditional_branch = .{ + .condition = true_value, + .true_edge = try builder.edge(pass, &.{}), + .false_edge = try builder.edge(merge, &.{loaded}), + } }); + + const splat = (try builder.appendInstruction(pass, vec4_type, .{ + .composite_construct = .{ .elements = &.{ one, one, one, one } }, + }, "white")).?; + try builder.setTerminator(pass, .{ .branch = try builder.edge(merge, &.{splat}) }); + _ = try builder.appendInstruction(merge, null, .{ + .store_interface = .{ .variable = output, .value = merged }, + }, null); + try builder.setTerminator(merge, .return_void); + + try ir.validator.validate(&module); + + var control_flow = try ir.cfg.init(std.testing.allocator, &module, main); + defer control_flow.deinit(); + try std.testing.expectEqual(@as(usize, 2), control_flow.predecessors(merge).?.len); + try std.testing.expect(control_flow.dominates(entry, merge)); + try std.testing.expect(!control_flow.dominates(pass, merge)); + + const text = try ir.printer.allocPrint(std.testing.allocator, &module); + + defer std.testing.allocator.free(text); + try std.testing.expect(std.mem.indexOf(u8, text, "shader vertex @main") != null); + try std.testing.expect(std.mem.indexOf(u8, text, "@color: vec4[f32] = input[location(0), component(0), index(0)]") != null); + try std.testing.expect(std.mem.indexOf(u8, text, "@out_color: vec4[f32] = output[location(0), component(0), index(0)]") != null); + try std.testing.expect(std.mem.indexOf(u8, text, "conditional_branch %0, .pass(), .merge(%loaded)") != null); + try std.testing.expect(std.mem.indexOf(u8, text, ".merge(%merged: vec4[f32])") != null); + try std.testing.expect(std.mem.indexOf(u8, text, "store_interface @out_color, %merged") != null); + + var parsed = try ir.parser.parseString(std.testing.allocator, text); + defer parsed.deinit(); + const round_trip = try ir.printer.allocPrint(std.testing.allocator, &parsed); + defer std.testing.allocator.free(round_trip); + try std.testing.expectEqualStrings(text, round_trip); + + const io = std.Options.debug_io; + const path = ".zig-cache/ir-parser-round-trip.ir"; + const file = try std.Io.Dir.cwd().createFile(io, path, .{ .truncate = true }); + { + defer file.close(io); + var file_buffer: [4096]u8 = @splat(0); + var file_writer = file.writer(io, &file_buffer); + try file_writer.interface.writeAll(text); + try file_writer.interface.flush(); + } + defer std.Io.Dir.cwd().deleteFile(io, path) catch @panic("Caught an error while handling an error"); + + var parsed_file = try ir.parser.parseFile(std.testing.allocator, io, path); + defer parsed_file.deinit(); + const file_round_trip = try ir.printer.allocPrint(std.testing.allocator, &parsed_file); + defer std.testing.allocator.free(file_round_trip); + try std.testing.expectEqualStrings(text, file_round_trip); +} + +test "IR parse interface" { + const source = + \\ shader vertex @main + \\ { + \\ @in_color: vec4[f32] = input[location(0), component(0), index(0)] + \\ @out_color: vec4[f32] = output[location(0), component(0), index(0)] + \\ @position: vec4[f32] = output[builtin(position)] + \\ + \\ fn @main() -> void + \\ { + \\ .entry(): + \\ return + \\ } + \\ } + ; + + var module = try ir.parser.parseString(std.testing.allocator, source); + defer module.deinit(); + const printed = try ir.printer.allocPrint(std.testing.allocator, &module); + defer std.testing.allocator.free(printed); + + try std.testing.expect(std.mem.indexOf(u8, printed, "@in_color: vec4[f32] = input[location(0), component(0), index(0)]") != null); + try std.testing.expect(std.mem.indexOf(u8, printed, "@out_color: vec4[f32] = output[location(0), component(0), index(0)]") != null); + try std.testing.expect(std.mem.indexOf(u8, printed, "@position: vec4[f32] = output[builtin(position)]") != null); +} + +test "IR parse types, operations, calls, terminators" { + const source = + \\ shader fragment @main + \\ { + \\ %0: constant bool = true + \\ %1: constant u32 = bits(0x1) + \\ %2: constant u32 = bits(0x2) + \\ %3: constant f32 = bits(0x3f800000) + \\ %4: constant array[u32, 2] = [#1, #2] + \\ %5: constant struct[u32, u32] = [#1, #2] + \\ %6: constant ptr[private, u32] = null + \\ %7: constant resourceHandle[sampler] = null + \\ + \\ fn @main() -> void + \\ { + \\ .entry(): + \\ %9: u32 = bitwise_not %1 + \\ %10: u32 = integer_add %9, %2 + \\ %11: bool = cmp_equal %1, %2 + \\ %12: u32 = select %11, %1, %2 + \\ %13: u32 = bitcast %12 + \\ %14: vec2[u32] = composite_construct %1, %2 + \\ %15: u32 = composite_extract %14[0] + \\ %16: f32 = negate %3 + \\ %17: f32 = float_add %3, %16 + \\ %18: u32 = call @helper(%15) + \\ return + \\ } + \\ + \\ fn @helper(%8: u32) -> u32 + \\ { + \\ .entry(): + \\ return %8 + \\ } + \\ + \\ fn @discarder() -> void + \\ { + \\ .entry(): + \\ discard + \\ } + \\ + \\ fn @dead() -> void + \\ { + \\ .entry(): + \\ unreachable + \\ } + \\ } + ; + + var module = try ir.parser.parseString(std.testing.allocator, source); + defer module.deinit(); + const printed = try ir.printer.allocPrint(std.testing.allocator, &module); + defer std.testing.allocator.free(printed); + + var reparsed = try ir.parser.parseString(std.testing.allocator, printed); + defer reparsed.deinit(); + const printed_again = try ir.printer.allocPrint(std.testing.allocator, &reparsed); + defer std.testing.allocator.free(printed_again); + try std.testing.expectEqualStrings(printed, printed_again); + try std.testing.expect(std.mem.indexOf(u8, printed, "cmp_equal %1, %2") != null); + try std.testing.expect(std.mem.indexOf(u8, printed, "cmp.") == null); +} + +test "IR parse named value IDs" { + const source = + \\ shader compute @main + \\ { + \\ %one_value: constant u32 = bits(0x1) + \\ + \\ fn @main() -> void + \\ { + \\ .entry(): + \\ %sum_value: u32 = integer_add %one_value, %one_value + \\ branch .merge(%sum_value) + \\ + \\ .merge(%merged_value: u32): + \\ %product_value: u32 = integer_multiply %merged_value, %one_value + \\ return + \\ } + \\ } + ; + + var module = try ir.parser.parseString(std.testing.allocator, source); + defer module.deinit(); + const printed = try ir.printer.allocPrint(std.testing.allocator, &module); + defer std.testing.allocator.free(printed); + + try std.testing.expect(std.mem.indexOf(u8, printed, "%one_value: constant u32 = bits(0x1)") != null); + try std.testing.expect(std.mem.indexOf(u8, printed, "%sum_value: u32 = integer_add %one_value, %one_value") != null); + try std.testing.expect(std.mem.indexOf(u8, printed, "branch .merge(%sum_value)") != null); + try std.testing.expect(std.mem.indexOf(u8, printed, ".merge(%merged_value: u32)") != null); + try std.testing.expect(std.mem.indexOf(u8, printed, "%product_value: u32 = integer_multiply %merged_value, %one_value") != null); + + var reparsed = try ir.parser.parseString(std.testing.allocator, printed); + defer reparsed.deinit(); + const printed_again = try ir.printer.allocPrint(std.testing.allocator, &reparsed); + defer std.testing.allocator.free(printed_again); + try std.testing.expectEqualStrings(printed, printed_again); +} + +test "IR parse numeric constants" { + const source = + \\ shader compute @main + \\ { + \\ %0: constant u8 = 255 + \\ %1: constant i8 = -1 + \\ %2: constant f16 = 1.5 + \\ %3: constant f32 = -0.0 + \\ %4: constant f64 = 2.5e0 + \\ + \\ fn @main() -> void + \\ { + \\ .entry(): + \\ return + \\ } + \\ } + ; + + var module = try ir.parser.parseString(std.testing.allocator, source); + defer module.deinit(); + const printed = try ir.printer.allocPrint(std.testing.allocator, &module); + defer std.testing.allocator.free(printed); + + try std.testing.expect(std.mem.indexOf(u8, printed, "%0: constant u8 = bits(0xff)") != null); + try std.testing.expect(std.mem.indexOf(u8, printed, "%1: constant i8 = bits(0xff)") != null); + try std.testing.expect(std.mem.indexOf(u8, printed, "%2: constant f16 = bits(0x3e00)") != null); + try std.testing.expect(std.mem.indexOf(u8, printed, "%3: constant f32 = bits(0x80000000)") != null); + try std.testing.expect(std.mem.indexOf(u8, printed, "%4: constant f64 = bits(0x4004000000000000)") != null); + + const out_of_range = + \\ shader compute @main + \\ { + \\ %0: constant u8 = 256 + \\ fn @main() -> void + \\ { + \\ .entry(): + \\ return + \\ } + \\ } + ; + try std.testing.expectError(error.InvalidNumber, ir.parser.parseString(std.testing.allocator, out_of_range)); +} + +test "IR parser error: unknown value" { + const source = + \\ shader compute @main + \\ { + \\ fn @main() -> void + \\ { + \\ .entry(): + \\ return %99 + \\ } + \\ } + ; + try std.testing.expectError(error.UnknownValue, ir.parser.parseString(std.testing.allocator, source)); +} + +test "Validator error: wrong block argument count" { + // shader compute @main + // { + // fn @main() -> void + // { + // .entry(): + // branch .merge() + // + // .merge(%0: u32): + // return + // } + // } + + var module = ir.module.Module.init(std.testing.allocator, .compute); + defer module.deinit(); + var builder = ir.Builder.init(&module); + + const void_type = try builder.internType(.void); + const u32_type = try builder.internType(.{ .integer = .{ .bits = 32, .signedness = .unsigned } }); + const main = try builder.addFunction(void_type, "main"); + builder.setEntryPoint(main); + const entry = try builder.addBlock(main, "entry"); + const merge = try builder.addBlock(main, "merge"); + _ = try builder.addBlockParameter(merge, u32_type, null); + try builder.setTerminator(entry, .{ .branch = try builder.edge(merge, &.{}) }); + try builder.setTerminator(merge, .return_void); + + try std.testing.expectError(error.WrongBranchArgumentCount, ir.validator.validate(&module)); +} + +test "Central store IDs disposal" { + var module = ir.module.Module.init(std.testing.allocator, .fragment); + defer module.deinit(); + const first = try module.internType(.boolean); + try std.testing.expect(module.types.remove(first)); + const second = try module.internType(.boolean); + try std.testing.expect(first.index() != second.index()); + try std.testing.expect(module.types.get(first) == null); +} + +test "Validator error: SSA definition does not dominate its use" { + // shader compute @main + // { + // %0: constant bool = true + // %1: constant u32 = bits(0x1) + // + // fn @main() -> void + // { + // .entry(): + // conditional_branch %0, .left(), .right() + // + // .left(): + // %2: u32 = integer_add %1, %1 + // branch .merge() + // + // .right(): + // branch .merge() + // + // .merge(): + // %3: u32 = integer_multiply %2, %1 + // return + // } + // } + + var module = ir.module.Module.init(std.testing.allocator, .compute); + defer module.deinit(); + var builder = ir.Builder.init(&module); + + const void_type = try builder.internType(.void); + const bool_type = try builder.internType(.boolean); + const u32_type = try builder.internType(.{ .integer = .{ .bits = 32, .signedness = .unsigned } }); + const condition = try builder.internConstant(bool_type, .{ .boolean = true }); + const one = try builder.internConstant(u32_type, .{ .integer_bits = 1 }); + const main = try builder.addFunction(void_type, "main"); + builder.setEntryPoint(main); + const entry = try builder.addBlock(main, "entry"); + const left = try builder.addBlock(main, "left"); + const right = try builder.addBlock(main, "right"); + const merge = try builder.addBlock(main, "merge"); + + try builder.setTerminator( + entry, + .{ + .conditional_branch = .{ + .condition = condition, + .true_edge = try builder.edge(left, &.{}), + .false_edge = try builder.edge(right, &.{}), + }, + }, + ); + + const left_value = (try builder.appendInstruction(left, u32_type, .{ + .binary = .{ + .opcode = .integer_add, + .lhs = one, + .rhs = one, + }, + }, null)).?; + + try builder.setTerminator(left, .{ .branch = try builder.edge(merge, &.{}) }); + try builder.setTerminator(right, .{ .branch = try builder.edge(merge, &.{}) }); + + _ = try builder.appendInstruction(merge, u32_type, .{ + .binary = .{ + .opcode = .integer_multiply, + .lhs = left_value, + .rhs = one, + }, + }, null); + + try builder.setTerminator(merge, .return_void); + + try std.testing.expectError(error.DefinitionDoesNotDominateUse, ir.validator.validate(&module)); +} + +test "Rewriter replace all ID uses, safely erase dead instruction" { + // shader compute @main + // { + // %0: constant u32 = bits(0x1) + // %1: constant u32 = bits(0x2) + // + // fn @main() -> void + // { + // .entry(): + // %2: u32 = integer_add %0, %1 + // %3: u32 = integer_multiply %2, %1 + // return + // } + // } + + var module = ir.module.Module.init(std.testing.allocator, .compute); + defer module.deinit(); + + var builder = ir.Builder.init(&module); + + const void_type = try builder.internType(.void); + const u32_type = try builder.internType(.{ .integer = .{ .bits = 32, .signedness = .unsigned } }); + + const one = try builder.internConstant(u32_type, .{ .integer_bits = 1 }); + const two = try builder.internConstant(u32_type, .{ .integer_bits = 2 }); + + const main = try builder.addFunction(void_type, "main"); + builder.setEntryPoint(main); + + const entry = try builder.addBlock(main, "entry"); + const sum = (try builder.appendInstruction(entry, u32_type, .{ + .binary = .{ + .opcode = .integer_add, + .lhs = one, + .rhs = two, + }, + }, null)).?; + _ = try builder.appendInstruction(entry, u32_type, .{ + .binary = .{ + .opcode = .integer_multiply, + .lhs = sum, + .rhs = two, + }, + }, null); + try builder.setTerminator(entry, .return_void); + + try ir.validator.validate(&module); + + const sum_instruction = module.values.get(sum).?.definition.instruction; + var rewriter = ir.Rewriter.init(&module); + + try std.testing.expectEqual(@as(usize, 1), try rewriter.replaceAllUses(sum, one)); + try rewriter.eraseInstruction(sum_instruction); + + try std.testing.expect(module.values.get(sum) == null); + try std.testing.expect(module.instructions.get(sum_instruction) == null); + + try ir.validator.validate(&module); +} + +test "Rewriter add block parameter and sync branch calls" { + // shader compute @main + // { + // %0: constant u32 = bits(0x1) + // + // fn @main() -> void + // { + // .entry(): + // branch .merge() + // + // .merge(): + // return + // + // .alternate(): + // return + // } + // } + + var module = ir.module.Module.init(std.testing.allocator, .compute); + defer module.deinit(); + + var builder = ir.Builder.init(&module); + + const void_type = try builder.internType(.void); + const u32_type = try builder.internType(.{ .integer = .{ .bits = 32, .signedness = .unsigned } }); + + const one = try builder.internConstant(u32_type, .{ .integer_bits = 1 }); + + const main = try builder.addFunction(void_type, "main"); + builder.setEntryPoint(main); + + const entry = try builder.addBlock(main, "entry"); + const merge = try builder.addBlock(main, "merge"); + const alternate = try builder.addBlock(main, "alternate"); + + try builder.setTerminator(entry, .{ .branch = try builder.edge(merge, &.{}) }); + try builder.setTerminator(merge, .return_void); + try builder.setTerminator(alternate, .return_void); + + var rewriter = ir.Rewriter.init(&module); + + const parameter = try rewriter.addBlockParameter(merge, u32_type, "incoming", &.{ + .{ + .predecessor = entry, + .value = one, + }, + }); + const merge_edge = module.blocks.get(entry).?.terminator.?.branch; + try std.testing.expectEqualSlices(ir.id.ValueId, &.{one}, merge_edge.arguments); + + _ = try builder.appendInstruction(merge, u32_type, .{ + .binary = .{ + .opcode = .integer_add, + .lhs = parameter, + .rhs = one, + }, + }, null); + try ir.validator.validate(&module); + + try rewriter.removeBlockParameter(merge, 0, one); + + try std.testing.expectEqual(@as(usize, 0), module.blocks.get(merge).?.parameters.items.len); + try std.testing.expectEqual(@as(usize, 0), module.blocks.get(entry).?.terminator.?.branch.arguments.len); + + try ir.validator.validate(&module); + + try std.testing.expectEqual(@as(usize, 1), try rewriter.redirectEdges(entry, merge, alternate, &.{})); + try std.testing.expectEqual(alternate, module.blocks.get(entry).?.terminator.?.branch.target); + + try ir.validator.validate(&module); +} + +fn establishNoCalls(_: *ir.module.Module, _: *ir.pass_manager.Context) !bool { + return false; +} + +fn countVisitedFunction(context: ?*anyopaque, _: ir.id.FunctionId, _: *const ir.module.Function) !void { + const statistics: *VisitorStatistics = @ptrCast(@alignCast(context.?)); + statistics.functions += 1; +} + +fn countVisitedBlock(context: ?*anyopaque, _: ir.id.BlockId, _: *const ir.module.Block) !void { + const statistics: *VisitorStatistics = @ptrCast(@alignCast(context.?)); + statistics.blocks += 1; +} + +test "Pass manager track independent IR properties" { + // shader compute @main + // { + // fn @main() -> void + // { + // .entry(): + // return + // } + // } + + var module = ir.module.Module.init(std.testing.allocator, .compute); + defer module.deinit(); + + var builder = ir.Builder.init(&module); + + const void_type = try builder.internType(.void); + + const main = try builder.addFunction(void_type, "main"); + builder.setEntryPoint(main); + + const entry = try builder.addBlock(main, "entry"); + try builder.setTerminator(entry, .return_void); + + module.properties.valid_cfg = true; + + var manager = ir.pass_manager.Manager.init(std.testing.allocator); + defer manager.deinit(); + + try manager.add(.{ + .name = "establish-no-calls", + .required = .{ .valid_cfg = true }, + .produced = .{ .no_function_calls = true }, + .run = establishNoCalls, + }); + + var context: ir.pass_manager.Context = .{ .allocator = std.testing.allocator }; + + try std.testing.expect(!try manager.run(&module, &context)); + try std.testing.expect(module.properties.no_function_calls); + + var statistics: VisitorStatistics = .{}; + + try ir.visitor.walk(&module, .{ + .context = &statistics, + .visitFunction = countVisitedFunction, + .visitBlock = countVisitedBlock, + }); + + try std.testing.expectEqual(@as(usize, 1), statistics.functions); + try std.testing.expectEqual(@as(usize, 1), statistics.blocks); +} + +test "SPIR-V parser error: zero-word instruction" { + const words = [_]u32{ + spirv.spec.magic_number, + 0x0001_0000, + 0, + 2, + 0, + instructionWord(.nop, 0), + }; + try std.testing.expectError(error.ZeroWordInstruction, spirv.Parser.init(&words)); + + const truncated = [_]u32{ + spirv.spec.magic_number, + 0x0001_0000, + 0, + 2, + 0, + instructionWord(.i_add, 5), + 1, + }; + try std.testing.expectError(error.TruncatedInstruction, spirv.Parser.init(&truncated)); +} + +test "SPIR-V structured branches and OpPhi to block parameters" { + const assembly = + \\ OpCapability Shader + \\ OpMemoryModel Logical GLSL450 + \\ OpEntryPoint GLCompute %main "main" + \\ OpExecutionMode %main LocalSize 1 1 1 + \\ OpName %main "main" + \\ OpName %entry "entry" + \\ OpName %true "true" + \\ OpName %one "one" + \\ OpName %then "then" + \\ OpName %then_value "then_value" + \\ OpName %else "else" + \\ OpName %else_value "else_value" + \\ OpName %merge "merge" + \\ OpName %merged "merged" + \\ OpName %product "product" + \\ + \\ %void = OpTypeVoid + \\ %bool = OpTypeBool + \\ %uint = OpTypeInt 32 0 + \\ %fn_void = OpTypeFunction %void + \\ %true = OpConstantTrue %bool + \\ %one = OpConstant %uint 1 + \\ + \\ %main = OpFunction %void None %fn_void + \\ %entry = OpLabel + \\ OpSelectionMerge %merge None + \\ OpBranchConditional %true %then %else + \\ %then = OpLabel + \\ %then_value = OpIAdd %uint %one %one + \\ OpBranch %merge + \\ %else = OpLabel + \\ %else_value = OpISub %uint %one %one + \\ OpBranch %merge + \\ %merge = OpLabel + \\ %merged = OpPhi %uint %then_value %then %else_value %else + \\ %product = OpIMul %uint %merged %one + \\ OpReturn + \\ OpFunctionEnd + ; + const words = try assembleSpirv(std.testing.allocator, assembly); + defer std.testing.allocator.free(words); + + var module = try spirv.translator.translate(std.testing.allocator, words, .{ .entry_point = "main" }); + defer module.deinit(); + + try std.testing.expectEqual(ir.module.Stage.compute, module.stage); + try std.testing.expectEqual([3]u32{ 1, 1, 1 }, module.execution_modes.workgroup_size.?); + try std.testing.expect(module.properties.valid_cfg); + try std.testing.expect(module.properties.valid_ssa); + + const function = module.functions.get(module.entry_point.?).?; + try std.testing.expectEqual(@as(usize, 4), function.blocks.items.len); + const entry = module.blocks.get(function.blocks.items[0]).?; + try std.testing.expect(entry.structured_control == .selection); + const merge = module.blocks.get(function.blocks.items[3]).?; + try std.testing.expectEqual(@as(usize, 1), merge.parameters.items.len); + try std.testing.expectEqual(@as(usize, 1), merge.instructions.items.len); + const multiply = module.instructions.get(merge.instructions.items[0]).?; + try std.testing.expectEqual(ir.instruction.BinaryOpcode.integer_multiply, multiply.operation.binary.opcode); + + const text = try ir.printer.allocPrint(std.testing.allocator, &module); + defer std.testing.allocator.free(text); + + try std.testing.expect(std.mem.indexOf(u8, text, "%one: constant u32 = bits(0x1)") != null); + try std.testing.expect(std.mem.indexOf(u8, text, "%true: constant bool = true") != null); + try std.testing.expect(std.mem.indexOf(u8, text, "conditional_branch %true, .then(), .else()") != null); + try std.testing.expect(std.mem.indexOf(u8, text, "%then_value: u32 = integer_add %one, %one") != null); + try std.testing.expect(std.mem.indexOf(u8, text, "branch .merge(%then_value)") != null); + try std.testing.expect(std.mem.indexOf(u8, text, "%else_value: u32 = integer_subtract %one, %one") != null); + try std.testing.expect(std.mem.indexOf(u8, text, ".merge(%merged: u32)") != null); + try std.testing.expect(std.mem.indexOf(u8, text, "%product: u32 = integer_multiply %merged, %one") != null); + try std.testing.expect(std.mem.indexOf(u8, text, "integerMultiply") == null); + + var parsed = try ir.parser.parseString(std.testing.allocator, text); + defer parsed.deinit(); + const round_trip = try ir.printer.allocPrint(std.testing.allocator, &parsed); + defer std.testing.allocator.free(round_trip); + try std.testing.expectEqualStrings(text, round_trip); +} + +test "SPIR-V decorated vertex interfaces and load-store operations" { + const assembly = + \\ OpCapability Shader + \\ OpMemoryModel Logical GLSL450 + \\ OpEntryPoint Vertex %main "main" %in_color %out_color + \\ OpName %in_color "in_color" + \\ OpName %out_color "out_color" + \\ OpDecorate %in_color Location 0 + \\ OpDecorate %out_color Location 0 + \\ + \\ %void = OpTypeVoid + \\ %float = OpTypeFloat 32 + \\ %vec4 = OpTypeVector %float 4 + \\ %input_vec4 = OpTypePointer Input %vec4 + \\ %output_vec4 = OpTypePointer Output %vec4 + \\ %fn_void = OpTypeFunction %void + \\ %in_color = OpVariable %input_vec4 Input + \\ %out_color = OpVariable %output_vec4 Output + \\ + \\ %main = OpFunction %void None %fn_void + \\ %entry = OpLabel + \\ %color = OpLoad %vec4 %in_color + \\ OpStore %out_color %color + \\ OpReturn + \\ OpFunctionEnd + ; + const words = try assembleSpirv(std.testing.allocator, assembly); + defer std.testing.allocator.free(words); + + var module = try spirv.translator.translate(std.testing.allocator, words, .{ .entry_point = "main" }); + defer module.deinit(); + try std.testing.expectEqual(ir.module.Stage.vertex, module.stage); + try std.testing.expectEqual(@as(usize, 2), module.interface_variables.entries.items.len); + + const text = try ir.printer.allocPrint(std.testing.allocator, &module); + defer std.testing.allocator.free(text); + try std.testing.expect(std.mem.indexOf(u8, text, "load_interface @in_color") != null); + try std.testing.expect(std.mem.indexOf(u8, text, "store_interface @out_color") != null); +} + +fn instructionWord(opcode: spirv.spec.Opcode, word_count: u16) u32 { + return (@as(u32, word_count) << 16) | @intFromEnum(opcode); +} + +fn assembleSpirv(allocator: std.mem.Allocator, assembly: []const u8) ![]u32 { + var io_backend: std.Io.Threaded = .init(allocator, .{}); + defer io_backend.deinit(); + const io = io_backend.io(); + + var child = try std.process.spawn(io, .{ + .argv = &.{ "spirv-as", "--target-env", "spv1.0", "-o", "-", "-" }, + .stdin = .pipe, + .stdout = .pipe, + .stderr = .pipe, + }); + defer child.kill(io); + + { + const stdin = child.stdin.?; + var stdin_writer = stdin.writer(io, &.{}); + try stdin_writer.interface.writeAll(assembly); + try stdin_writer.interface.flush(); + stdin.close(io); + child.stdin = null; + } + + var stdout_buffer: [4096]u8 = undefined; + var stdout_reader = child.stdout.?.reader(io, &stdout_buffer); + const binary = try stdout_reader.interface.allocRemaining(allocator, .limited(1024 * 1024)); + defer allocator.free(binary); + + var stderr_buffer: [4096]u8 = undefined; + var stderr_reader = child.stderr.?.reader(io, &stderr_buffer); + const stderr = try stderr_reader.interface.allocRemaining(allocator, .limited(64 * 1024)); + defer allocator.free(stderr); + + const term = try child.wait(io); + switch (term) { + .exited => |code| if (code != 0) { + std.log.err("spirv-as failed:\n{s}", .{stderr}); + return error.SpirvAssemblyFailed; + }, + else => { + std.log.err("spirv-as terminated unexpectedly:\n{s}", .{stderr}); + return error.SpirvAssemblyFailed; + }, + } + + if (binary.len % @sizeOf(u32) != 0) return error.InvalidSpirvBinaryLength; + const words = try allocator.alloc(u32, binary.len / @sizeOf(u32)); + errdefer allocator.free(words); + for (words, 0..) |*word, index| { + const offset = index * @sizeOf(u32); + word.* = std.mem.readInt(u32, binary[offset..][0..4], .little); + } + return words; +} diff --git a/src/compiler/spirv/Parser.zig b/src/compiler/spirv/Parser.zig new file mode 100644 index 0000000..48be5df --- /dev/null +++ b/src/compiler/spirv/Parser.zig @@ -0,0 +1,155 @@ +const spirv = @import("spirv.zig"); + +const Self = @This(); + +pub const Error = error{ + HeaderTooShort, + InvalidMagic, + ByteSwappedModule, + InvalidVersion, + InvalidIdBound, + InvalidSchema, + ZeroWordInstruction, + TruncatedInstruction, + UnterminatedString, +}; + +pub const Header = struct { + version: u32, + generator: u32, + bound: u32, + schema: u32, + + pub inline fn major(self: Header) u8 { + return @truncate(self.version >> 16); + } + + pub inline fn minor(self: Header) u8 { + return @truncate(self.version >> 8); + } +}; + +pub const Instruction = struct { + opcode: spirv.Opcode, + operands: []const u32, + word_offset: usize, + + pub fn operand(self: Instruction, index: usize) ?u32 { + return if (index < self.operands.len) self.operands[index] else null; + } +}; + +pub const Iterator = struct { + words: []const u32, + cursor: usize = spirv.header_word_count, + + pub fn next(self: *Iterator) Error!?Instruction { + if (self.cursor == self.words.len) + return null; + + const first_word = self.words[self.cursor]; + const word_count: usize = first_word >> 16; + + if (word_count == 0) + return error.ZeroWordInstruction; + if (word_count > self.words.len - self.cursor) + return error.TruncatedInstruction; + + const instruction: Instruction = .{ + .opcode = @enumFromInt(@as(u16, @truncate(first_word))), + .operands = self.words[self.cursor + 1 .. self.cursor + word_count], + .word_offset = self.cursor, + }; + self.cursor += word_count; + return instruction; + } +}; + +words: []const u32, +header: Header, + +pub fn init(words: []const u32) Error!Self { + if (words.len < spirv.header_word_count) + return error.HeaderTooShort; + if (words[0] == spirv.byte_swapped_magic_number) + return error.ByteSwappedModule; + if (words[0] != spirv.magic_number) + return error.InvalidMagic; + + const header: Header = .{ + .version = words[1], + .generator = words[2], + .bound = words[3], + .schema = words[4], + }; + + if (header.major() != 1 or header.minor() > 6 or (header.version & 0xff00_00ff) != 0) + return error.InvalidVersion; + if (header.bound == 0) + return error.InvalidIdBound; + if (header.schema != 0) + return error.InvalidSchema; + + var self: Self = .{ .words = words, .header = header }; + var instruction_iterator = self.iterator(); + while (try instruction_iterator.next()) |_| {} + return self; +} + +pub fn iterator(self: Self) Iterator { + return .{ .words = self.words }; +} + +pub fn literalStringWordCount(words: []const u32) Error!usize { + for (words, 0..) |word, word_index| { + inline for (0..4) |byte_index| { + if (@as(u8, @truncate(word >> (byte_index * 8))) == 0) + return word_index + 1; + } + } + return error.UnterminatedString; +} + +pub fn literalStringEquals(words: []const u32, expected: []const u8) Error!bool { + var byte_cursor: usize = 0; + for (words) |word| { + inline for (0..4) |byte_index| { + const byte: u8 = @truncate(word >> (byte_index * 8)); + + if (byte == 0) + return byte_cursor == expected.len; + if (byte_cursor >= expected.len or byte != expected[byte_cursor]) + return false; + + byte_cursor += 1; + } + } + return error.UnterminatedString; +} + +pub fn copyLiteralString(allocator: anytype, words: []const u32) ![]u8 { + const word_count = try literalStringWordCount(words); + var byte_count: usize = 0; + outer: for (words[0..word_count]) |word| { + inline for (0..4) |byte_index| { + if (@as(u8, @truncate(word >> (byte_index * 8))) == 0) + break :outer; + + byte_count += 1; + } + } + + const result = try allocator.alloc(u8, byte_count); + var cursor: usize = 0; + outer: for (words[0..word_count]) |word| { + inline for (0..4) |byte_index| { + const byte: u8 = @truncate(word >> (byte_index * 8)); + if (byte == 0) + break :outer; + + result[cursor] = byte; + cursor += 1; + } + } + return result; +} diff --git a/src/compiler/spirv/root.zig b/src/compiler/spirv/root.zig new file mode 100644 index 0000000..c4b3186 --- /dev/null +++ b/src/compiler/spirv/root.zig @@ -0,0 +1,16 @@ +//! ## SPIR-V frontend +//! +//! This namespace contains the SPIR-V parser and translator used to import shader +//! modules into the compiler IR. +//! +//! `Parser` validates the SPIR-V header and iterates over binary instructions. +//! `spec` exposes a minimalistic SPIR-V header translation. +//! +//! The main entry point is `translator.translate`, which finds the requested entry +//! point, maps its execution model to an IR shader stage, lowers supported types, +//! constants, interfaces, instructions, and structured control flow, then validates +//! the generated IR module. + +pub const Parser = @import("Parser.zig"); +pub const translator = @import("translator.zig"); +pub const spec = @import("spirv.zig"); diff --git a/src/compiler/spirv/spirv.zig b/src/compiler/spirv/spirv.zig new file mode 100644 index 0000000..455fb29 --- /dev/null +++ b/src/compiler/spirv/spirv.zig @@ -0,0 +1,169 @@ +pub const magic_number: u32 = 0x07230203; +pub const byte_swapped_magic_number: u32 = 0x03022307; +pub const header_word_count: usize = 5; + +pub const Opcode = enum(u32) { + nop = 0, + undef = 1, + name = 5, + member_name = 6, + string = 7, + line = 8, + extension = 10, + ext_inst_import = 11, + ext_inst = 12, + memory_model = 14, + entry_point = 15, + execution_mode = 16, + capability = 17, + + type_void = 19, + type_bool = 20, + type_int = 21, + type_float = 22, + type_vector = 23, + type_matrix = 24, + type_image = 25, + type_sampler = 26, + type_sampled_image = 27, + type_array = 28, + type_runtime_array = 29, + type_struct = 30, + type_opaque = 31, + type_pointer = 32, + type_function = 33, + + constant_true = 41, + constant_false = 42, + constant = 43, + constant_composite = 44, + constant_null = 46, + spec_constant_true = 48, + spec_constant_false = 49, + spec_constant = 50, + spec_constant_composite = 51, + spec_constant_op = 52, + + function = 54, + function_parameter = 55, + function_end = 56, + function_call = 57, + variable = 59, + load = 61, + store = 62, + access_chain = 65, + + decorate = 71, + member_decorate = 72, + vector_shuffle = 79, + composite_construct = 80, + composite_extract = 81, + composite_insert = 82, + copy_object = 83, + + convert_f_to_u = 109, + convert_f_to_s = 110, + convert_s_to_f = 111, + convert_u_to_f = 112, + u_convert = 113, + s_convert = 114, + f_convert = 115, + bitcast = 124, + + s_negate = 126, + f_negate = 127, + i_add = 128, + f_add = 129, + i_sub = 130, + f_sub = 131, + i_mul = 132, + f_mul = 133, + u_div = 134, + s_div = 135, + f_div = 136, + u_mod = 137, + s_rem = 138, + s_mod = 139, + f_rem = 140, + f_mod = 141, + shift_right_logical = 194, + shift_right_arithmetic = 195, + shift_left_logical = 196, + bitwise_or = 197, + bitwise_xor = 198, + bitwise_and = 199, + logical_equal = 164, + logical_not_equal = 165, + logical_or = 166, + logical_and = 167, + logical_not = 168, + select = 169, + i_equal = 170, + i_not_equal = 171, + u_less_than = 176, + s_less_than = 177, + f_ord_equal = 180, + f_unord_equal = 181, + f_ord_not_equal = 182, + f_unord_not_equal = 183, + f_ord_less_than = 184, + f_unord_less_than = 185, + + phi = 245, + loop_merge = 246, + selection_merge = 247, + label = 248, + branch = 249, + branch_conditional = 250, + @"switch" = 251, + kill = 252, + return_ = 253, + return_value = 254, + @"unreachable" = 255, + no_line = 317, + + _, +}; + +pub const ExecutionModel = enum(u32) { + vertex = 0, + tessellation_control = 1, + tessellation_evaluation = 2, + geometry = 3, + fragment = 4, + gl_compute = 5, + kernel = 6, + _, +}; + +pub const StorageClass = enum(u32) { + uniform_constant = 0, + input = 1, + uniform = 2, + output = 3, + workgroup = 4, + cross_workgroup = 5, + private = 6, + function = 7, + generic = 8, + push_constant = 9, + atomic_counter = 10, + image = 11, + storage_buffer = 12, + physical_storage_buffer = 5349, + _, +}; + +pub const ExecutionMode = enum(u32) { + early_fragment_tests = 9, + local_size = 17, + _, +}; + +pub const Decoration = enum(u32) { + built_in = 11, + location = 30, + component = 31, + index = 32, + _, +}; diff --git a/src/compiler/spirv/translator.zig b/src/compiler/spirv/translator.zig new file mode 100644 index 0000000..78b1856 --- /dev/null +++ b/src/compiler/spirv/translator.zig @@ -0,0 +1,1074 @@ +const std = @import("std"); +const Parser = @import("Parser.zig"); +const spirv = @import("spirv.zig"); +const ir = @import("../ir/ir.zig"); + +pub const Options = struct { + entry_point: []const u8, +}; + +pub const TranslationError = error{ + EntryPointNotFound, + AmbiguousEntryPoint, + UnsupportedExecutionModel, + InvalidInstruction, + InvalidId, + DuplicateId, + MissingDefinition, + MissingFunction, + InvalidFunctionType, + InvalidFunctionParameter, + InvalidBlock, + InvalidPhi, + MissingPhiIncomingValue, + UnsupportedType, + UnsupportedConstant, + SpecializationConstantsNotApplied, + UnsupportedOpcode, +}; + +const EntryPoint = struct { + model: spirv.ExecutionModel, + function_id: u32, + interface_ids: []const u32, +}; + +const Decorations = struct { + location: ?u32 = null, + component: u8 = 0, + index: u8 = 0, + builtin: ?u32 = null, +}; + +const PhiInfo = struct { + target_label: u32, + incoming_words: []const u32, +}; + +const Context = struct { + scratch: std.mem.Allocator, + parser: Parser, + module: *ir.module.Module, + builder: ir.Builder, + bound: usize, + + type_defs: []?Parser.Instruction, + value_defs: []?Parser.Instruction, + variable_defs: []?Parser.Instruction, + names: []?[]const u8, + decorations: []Decorations, + + types: []?ir.id.TypeId, + values: []?ir.id.ValueId, + blocks: []?ir.id.BlockId, + interfaces: []?ir.id.InterfaceVariableId, + phi_infos: std.ArrayList(PhiInfo) = .empty, + + fn idIndex(self: *const Context, id: u32) TranslationError!usize { + if (id == 0 or id >= self.bound) + return error.InvalidId; + return id; + } + + fn recordDefinition(self: *Context, definitions: []?Parser.Instruction, id: u32, instruction: Parser.Instruction) TranslationError!void { + const index = try self.idIndex(id); + if (definitions[index] != null) + return error.DuplicateId; + definitions[index] = instruction; + } + + fn nameOf(self: *const Context, id: u32) ?[]const u8 { + const index = self.idIndex(id) catch return null; + return self.names[index]; + } + + fn translateType(self: *Context, spv_id: u32) anyerror!ir.id.TypeId { + const index = try self.idIndex(spv_id); + if (self.types[index]) |translated| + return translated; + + const instruction = self.type_defs[index] orelse return error.MissingDefinition; + const operands = instruction.operands; + + const translated = switch (instruction.opcode) { + .type_void => blk: { + try expectOperandCount(operands, 1); + break :blk try self.builder.internType(.void); + }, + .type_bool => blk: { + try expectOperandCount(operands, 1); + break :blk try self.builder.internType(.boolean); + }, + .type_int => blk: { + try expectOperandCount(operands, 3); + + if (operands[1] == 0 or operands[1] > 64 or operands[2] > 1) + return error.UnsupportedType; + + break :blk try self.builder.internType(.{ + .integer = .{ + .bits = @intCast(operands[1]), + .signedness = if (operands[2] == 0) .unsigned else .signed, + }, + }); + }, + .type_float => blk: { + try expectOperandCount(operands, 2); + + if (operands[1] != 16 and operands[1] != 32 and operands[1] != 64) + return error.UnsupportedType; + + break :blk try self.builder.internType(.{ + .floating = .{ + .bits = @intCast(operands[1]), + }, + }); + }, + .type_vector => blk: { + try expectOperandCount(operands, 3); + + if (operands[2] < 2 or operands[2] > std.math.maxInt(u8)) + return error.UnsupportedType; + + break :blk try self.builder.internType(.{ + .vector = .{ + .element_type = try self.translateType(operands[1]), + .length = @intCast(operands[2]), + }, + }); + }, + .type_array => blk: { + try expectOperandCount(operands, 3); + + const length_value = try self.translateValue(operands[2]); + const length_constant_value = self.module.values.get(length_value) orelse return error.InvalidId; + + if (length_constant_value.definition != .constant) + return error.UnsupportedType; + + const constant = self.module.constants.get(length_constant_value.definition.constant) orelse return error.InvalidId; + if (constant.value != .integer_bits or constant.value.integer_bits == 0 or constant.value.integer_bits > std.math.maxInt(u32)) + return error.UnsupportedType; + + break :blk try self.builder.internType(.{ + .array = .{ + .element_type = try self.translateType(operands[1]), + .length = @intCast(constant.value.integer_bits), + }, + }); + }, + .type_struct => blk: { + if (operands.len == 0) + return error.InvalidInstruction; + + const members = try self.scratch.alloc(ir.id.TypeId, operands.len - 1); + + for (operands[1..], members) |member_id, *member| + member.* = try self.translateType(member_id); + + break :blk try self.builder.internType(.{ + .structure = .{ + .members = members, + }, + }); + }, + .type_pointer => blk: { + try expectOperandCount(operands, 3); + break :blk try self.builder.internType(.{ + .pointer = .{ + .address_space = try translateStorageClass(@enumFromInt(operands[1])), + .pointee_type = try self.translateType(operands[2]), + }, + }); + }, + + else => return error.UnsupportedType, + }; + self.types[index] = translated; + return translated; + } + + fn translateValue(self: *Context, spv_id: u32) anyerror!ir.id.ValueId { + const index = try self.idIndex(spv_id); + + if (self.values[index]) |translated| + return translated; + + const instruction = self.value_defs[index] orelse return error.MissingDefinition; + const operands = instruction.operands; + + const translated = switch (instruction.opcode) { + .undef => blk: { + try expectOperandCount(operands, 2); + break :blk try self.module.values.add(self.module.allocator(), .{ + .type = try self.translateType(operands[0]), + .definition = .undef, + .name = if (self.nameOf(spv_id)) |name| try self.module.allocator().dupe(u8, name) else null, + }); + }, + .constant_true => blk: { + try expectOperandCount(operands, 2); + break :blk try self.builder.internConstant(try self.translateType(operands[0]), .{ .boolean = true }); + }, + .constant_false => blk: { + try expectOperandCount(operands, 2); + break :blk try self.builder.internConstant(try self.translateType(operands[0]), .{ .boolean = false }); + }, + .constant => blk: { + if (operands.len < 3 or operands.len > 4) + return error.InvalidInstruction; + + const ty = try self.translateType(operands[0]); + const type_data = self.module.types.get(ty) orelse return error.InvalidId; + const bits = try literalBits(operands[2..]); + + break :blk switch (type_data.*) { + .integer => try self.builder.internConstant(ty, .{ .integer_bits = bits }), + .floating => try self.builder.internConstant(ty, .{ .float_bits = bits }), + else => return error.UnsupportedConstant, + }; + }, + .constant_null => blk: { + try expectOperandCount(operands, 2); + break :blk try self.builder.internConstant(try self.translateType(operands[0]), .null); + }, + .constant_composite => blk: { + if (operands.len < 2) + return error.InvalidInstruction; + + const elements = try self.scratch.alloc(ir.id.ConstantId, operands.len - 2); + + for (operands[2..], elements) |element_id, *element| { + const element_value = self.module.values.get(try self.translateValue(element_id)) orelse return error.InvalidId; + + if (element_value.definition != .constant) + return error.UnsupportedConstant; + + element.* = element_value.definition.constant; + } + + break :blk try self.builder.internConstant( + try self.translateType(operands[0]), + .{ .composite = elements }, + ); + }, + .spec_constant_true, + .spec_constant_false, + .spec_constant, + .spec_constant_composite, + .spec_constant_op, + => return error.SpecializationConstantsNotApplied, + + else => return error.MissingDefinition, + }; + try self.builder.setValueName(translated, self.nameOf(spv_id)); + self.values[index] = translated; + return translated; + } + + fn setValue(self: *Context, spv_id: u32, translated_value: ir.id.ValueId) TranslationError!void { + const index = try self.idIndex(spv_id); + + if (self.values[index] != null) + return error.DuplicateId; + + self.values[index] = translated_value; + } + + fn resolveValue(self: *Context, spv_id: u32) anyerror!ir.id.ValueId { + return self.translateValue(spv_id); + } + + fn block(self: *const Context, spv_id: u32) TranslationError!ir.id.BlockId { + const index = try self.idIndex(spv_id); + return self.blocks[index] orelse error.InvalidBlock; + } + + fn interfaceVariable(self: *const Context, spv_id: u32) TranslationError!ir.id.InterfaceVariableId { + const index = try self.idIndex(spv_id); + return self.interfaces[index] orelse error.UnsupportedOpcode; + } +}; + +pub fn translate(allocator: std.mem.Allocator, words: []const u32, options: Options) !ir.module.Module { + const parser = try Parser.init(words); + const entry_point = try findEntryPoint(parser, options.entry_point); + const stage = try translateStage(entry_point.model); + + var module = ir.module.Module.init(allocator, stage); + errdefer module.deinit(); + + var scratch_arena = std.heap.ArenaAllocator.init(allocator); + defer scratch_arena.deinit(); + const scratch = scratch_arena.allocator(); + const bound: usize = parser.header.bound; + + var context: Context = .{ + .scratch = scratch, + .parser = parser, + .module = &module, + .builder = ir.Builder.init(&module), + .bound = bound, + .type_defs = try allocOptional(Parser.Instruction, scratch, bound), + .value_defs = try allocOptional(Parser.Instruction, scratch, bound), + .variable_defs = try allocOptional(Parser.Instruction, scratch, bound), + .names = try allocOptional([]const u8, scratch, bound), + .decorations = try scratch.alloc(Decorations, bound), + .types = try allocOptional(ir.id.TypeId, scratch, bound), + .values = try allocOptional(ir.id.ValueId, scratch, bound), + .blocks = try allocOptional(ir.id.BlockId, scratch, bound), + .interfaces = try allocOptional(ir.id.InterfaceVariableId, scratch, bound), + }; + @memset(context.decorations, .{}); + defer context.phi_infos.deinit(scratch); + + try collectDeclarations(&context); + try translateInterfaces(&context, entry_point.interface_ids); + try applyExecutionModes(&context, entry_point.function_id); + try translateFunction(&context, entry_point.function_id, options.entry_point); + try ir.validator.validate(&module); + module.properties.valid_cfg = true; + module.properties.valid_ssa = true; + module.properties.structured_control_flow = true; + module.properties.no_function_calls = true; + return module; +} + +fn collectDeclarations(context: *Context) !void { + var iterator = context.parser.iterator(); + while (try iterator.next()) |instruction| { + const operands = instruction.operands; + if (isTypeOpcode(instruction.opcode)) { + if (operands.len == 0) + return error.InvalidInstruction; + + try context.recordDefinition(context.type_defs, operands[0], instruction); + continue; + } + if (isConstantOpcode(instruction.opcode) or instruction.opcode == .undef) { + if (operands.len < 2) + return error.InvalidInstruction; + + try context.recordDefinition(context.value_defs, operands[1], instruction); + continue; + } + if (instruction.opcode == .name) { + if (operands.len < 2) + return error.InvalidInstruction; + + const index = try context.idIndex(operands[0]); + context.names[index] = try Parser.copyLiteralString(context.scratch, operands[1..]); + continue; + } + if (instruction.opcode == .variable) { + if (operands.len < 3) + return error.InvalidInstruction; + + try context.recordDefinition(context.variable_defs, operands[1], instruction); + continue; + } + if (instruction.opcode == .decorate) { + try collectDecoration(context, operands); + } + } +} + +fn collectDecoration(context: *Context, operands: []const u32) !void { + if (operands.len < 2) return error.InvalidInstruction; + const index = try context.idIndex(operands[0]); + const decoration: spirv.Decoration = @enumFromInt(operands[1]); + switch (decoration) { + .built_in => { + try expectOperandCount(operands, 3); + context.decorations[index].builtin = operands[2]; + }, + .location => { + try expectOperandCount(operands, 3); + context.decorations[index].location = operands[2]; + }, + .component => { + try expectOperandCount(operands, 3); + if (operands[2] > std.math.maxInt(u8)) return error.InvalidInstruction; + context.decorations[index].component = @intCast(operands[2]); + }, + .index => { + try expectOperandCount(operands, 3); + if (operands[2] > std.math.maxInt(u8)) return error.InvalidInstruction; + context.decorations[index].index = @intCast(operands[2]); + }, + else => {}, + } +} + +fn translateInterfaces(context: *Context, interface_ids: []const u32) !void { + for (interface_ids) |spv_id| { + const index = try context.idIndex(spv_id); + const variable = context.variable_defs[index] orelse return error.MissingDefinition; + + if (variable.operands.len < 3 or variable.operands.len > 4) + return error.InvalidInstruction; + + const storage_class: spirv.StorageClass = @enumFromInt(variable.operands[2]); + const direction: ir.module.InterfaceDirection = switch (storage_class) { + .input => .input, + .output => .output, + else => continue, + }; + + const pointer_index = try context.idIndex(variable.operands[0]); + const pointer = context.type_defs[pointer_index] orelse return error.MissingDefinition; + + if (pointer.opcode != .type_pointer) + return error.InvalidInstruction; + + try expectOperandCount(pointer.operands, 3); + + if (pointer.operands[1] != variable.operands[2]) + return error.InvalidInstruction; + + const decoration = context.decorations[index]; + + if (decoration.location != null and decoration.builtin != null) + return error.InvalidInstruction; + + const semantic: ir.module.InterfaceSemantic = if (decoration.location) |location| + .{ + .location = .{ + .location = location, + .component = decoration.component, + .index = decoration.index, + }, + } + else if (decoration.builtin) |builtin| + .{ + .builtin = try translateBuiltin(builtin), + } + else + return error.InvalidInstruction; + + context.interfaces[index] = try context.builder.addInterfaceVariable( + try context.translateType(pointer.operands[2]), + direction, + semantic, + context.nameOf(spv_id), + ); + } +} + +fn findEntryPoint(parser: Parser, requested_name: []const u8) !EntryPoint { + var found: ?EntryPoint = null; + var iterator = parser.iterator(); + while (try iterator.next()) |instruction| { + if (instruction.opcode != .entry_point) + continue; + + if (instruction.operands.len < 3) + return error.InvalidInstruction; + + const string_words = try Parser.literalStringWordCount(instruction.operands[2..]); + + if (2 + string_words > instruction.operands.len) + return error.InvalidInstruction; + + if (!try Parser.literalStringEquals(instruction.operands[2 .. 2 + string_words], requested_name)) + continue; + + if (found != null) + return error.AmbiguousEntryPoint; + + found = .{ + .model = @enumFromInt(instruction.operands[0]), + .function_id = instruction.operands[1], + .interface_ids = instruction.operands[2 + string_words ..], + }; + } + return found orelse error.EntryPointNotFound; +} + +fn applyExecutionModes(context: *Context, entry_function: u32) !void { + var iterator = context.parser.iterator(); + while (try iterator.next()) |instruction| { + if (instruction.opcode != .execution_mode) + continue; + + if (instruction.operands.len < 2) + return error.InvalidInstruction; + + if (instruction.operands[0] != entry_function) + continue; + + const mode: spirv.ExecutionMode = @enumFromInt(instruction.operands[1]); + + switch (mode) { + .early_fragment_tests => context.module.execution_modes.early_fragment_tests = true, + .local_size => { + try expectOperandCount(instruction.operands, 5); + context.module.execution_modes.workgroup_size = instruction.operands[2..5].*; + }, + else => {}, + } + } +} + +fn translateFunction(context: *Context, spv_function: u32, entry_name: []const u8) !void { + const function_instruction = try findFunction(context.parser, spv_function); + try expectOperandCount(function_instruction.operands, 4); + const function_type = try functionTypeDefinition(context, function_instruction.operands[3]); + + if (function_type.operands.len < 2 or function_type.operands[1] != function_instruction.operands[0]) + return error.InvalidFunctionType; + + const function = try context.builder.addFunction( + try context.translateType(function_instruction.operands[0]), + context.nameOf(spv_function) orelse entry_name, + ); + context.builder.setEntryPoint(function); + + try predeclareFunction(context, spv_function, function, function_type.operands[2..]); + try translateFunctionInstructions(context, spv_function); + try translateFunctionControlFlow(context, spv_function); +} + +fn predeclareFunction(context: *Context, spv_function: u32, function: ir.id.FunctionId, parameter_types: []const u32) !void { + var active = false; + var parameter_index: usize = 0; + var current_label: ?u32 = null; + var iterator = context.parser.iterator(); + while (try iterator.next()) |instruction| { + if (instruction.opcode == .function) { + active = instruction.operands.len >= 2 and instruction.operands[1] == spv_function; + continue; + } + + if (!active) + continue; + + switch (instruction.opcode) { + .function_parameter => { + try expectOperandCount(instruction.operands, 2); + + if (parameter_index >= parameter_types.len or parameter_types[parameter_index] != instruction.operands[0]) + return error.InvalidFunctionParameter; + + const value = try context.builder.addFunctionParameter( + function, + try context.translateType(instruction.operands[0]), + context.nameOf(instruction.operands[1]), + ); + + try context.setValue(instruction.operands[1], value); + parameter_index += 1; + }, + .label => { + try expectOperandCount(instruction.operands, 1); + const label_id = instruction.operands[0]; + const index = try context.idIndex(label_id); + + if (context.blocks[index] != null) + return error.DuplicateId; + + context.blocks[index] = try context.builder.addBlock(function, context.nameOf(label_id)); + current_label = label_id; + }, + .phi => { + if (instruction.operands.len < 4 or (instruction.operands.len - 2) % 2 != 0) + return error.InvalidPhi; + + const label = current_label orelse return error.InvalidPhi; + const value = try context.builder.addBlockParameter( + try context.block(label), + try context.translateType(instruction.operands[0]), + context.nameOf(instruction.operands[1]), + ); + + try context.setValue(instruction.operands[1], value); + try context.phi_infos.append(context.scratch, .{ + .target_label = label, + .incoming_words = instruction.operands[2..], + }); + }, + .function_end => break, + + else => {}, + } + } + if (parameter_index != parameter_types.len) + return error.InvalidFunctionParameter; +} + +fn translateFunctionInstructions(context: *Context, spv_function: u32) !void { + var active = false; + var current_block: ?ir.id.BlockId = null; + var iterator = context.parser.iterator(); + + while (try iterator.next()) |instruction| { + if (instruction.opcode == .function) { + active = instruction.operands.len >= 2 and instruction.operands[1] == spv_function; + continue; + } + + if (!active) + continue; + + switch (instruction.opcode) { + .label => current_block = try context.block(instruction.operands[0]), + + .function_parameter, + .phi, + .selection_merge, + .loop_merge, + .branch, + .branch_conditional, + .return_, + .return_value, + .kill, + .@"unreachable", + => {}, + + .function_end => break, + + .nop, + .line, + .no_line, + => {}, + + else => try translateInstruction(context, current_block orelse return error.InvalidBlock, instruction), + } + } +} + +fn translateInstruction(context: *Context, block: ir.id.BlockId, instruction: Parser.Instruction) !void { + const operands = instruction.operands; + switch (instruction.opcode) { + .undef => { + try expectOperandCount(operands, 2); + _ = try context.translateValue(operands[1]); + }, + .copy_object => { + try expectOperandCount(operands, 3); + const source = try context.resolveValue(operands[2]); + + if (context.module.typeOf(source) != try context.translateType(operands[0])) + return error.InvalidInstruction; + + try context.setValue(operands[1], source); + }, + .load => { + if (operands.len < 3) + return error.InvalidInstruction; + + const result = (try context.builder.appendInstruction(block, try context.translateType(operands[0]), .{ + .load_interface = .{ .variable = try context.interfaceVariable(operands[2]) }, + }, context.nameOf(operands[1]))).?; + + try context.setValue(operands[1], result); + }, + .store => { + if (operands.len < 2) + return error.InvalidInstruction; + + _ = try context.builder.appendInstruction(block, null, .{ + .store_interface = .{ + .variable = try context.interfaceVariable(operands[0]), + .value = try context.resolveValue(operands[1]), + }, + }, null); + }, + .s_negate, .f_negate, .logical_not => { + try expectOperandCount(operands, 3); + + const result = (try context.builder.appendInstruction(block, try context.translateType(operands[0]), .{ + .unary = .{ + .opcode = if (instruction.opcode == .logical_not) .logical_not else .negate, + .operand = try context.resolveValue(operands[2]), + }, + }, context.nameOf(operands[1]))).?; + + try context.setValue(operands[1], result); + }, + .i_add, + .i_sub, + .i_mul, + .u_div, + .s_div, + .u_mod, + .s_mod, + .f_add, + .f_sub, + .f_mul, + .f_div, + .f_mod, + .shift_left_logical, + .shift_right_logical, + .shift_right_arithmetic, + .bitwise_and, + .bitwise_or, + .bitwise_xor, + .logical_and, + .logical_or, + => { + try expectOperandCount(operands, 4); + + const result = (try context.builder.appendInstruction(block, try context.translateType(operands[0]), .{ + .binary = .{ + .opcode = translateBinaryOpcode(instruction.opcode), + .lhs = try context.resolveValue(operands[2]), + .rhs = try context.resolveValue(operands[3]), + }, + }, context.nameOf(operands[1]))).?; + + try context.setValue(operands[1], result); + }, + .logical_equal, + .logical_not_equal, + .i_equal, + .i_not_equal, + .u_less_than, + .s_less_than, + .f_ord_equal, + .f_unord_equal, + .f_ord_not_equal, + .f_unord_not_equal, + .f_ord_less_than, + .f_unord_less_than, + => { + try expectOperandCount(operands, 4); + + const result = (try context.builder.appendInstruction(block, try context.translateType(operands[0]), .{ + .compare = .{ + .opcode = translateCompareOpcode(instruction.opcode), + .lhs = try context.resolveValue(operands[2]), + .rhs = try context.resolveValue(operands[3]), + }, + }, context.nameOf(operands[1]))).?; + + try context.setValue(operands[1], result); + }, + .select => { + try expectOperandCount(operands, 5); + + const result = (try context.builder.appendInstruction(block, try context.translateType(operands[0]), .{ + .select = .{ + .condition = try context.resolveValue(operands[2]), + .true_value = try context.resolveValue(operands[3]), + .false_value = try context.resolveValue(operands[4]), + }, + }, context.nameOf(operands[1]))).?; + + try context.setValue(operands[1], result); + }, + .bitcast => { + try expectOperandCount(operands, 3); + + const result = (try context.builder.appendInstruction(block, try context.translateType(operands[0]), .{ + .bitcast = try context.resolveValue(operands[2]), + }, context.nameOf(operands[1]))).?; + + try context.setValue(operands[1], result); + }, + .composite_construct => { + if (operands.len < 2) + return error.InvalidInstruction; + + const elements = try context.scratch.alloc(ir.id.ValueId, operands.len - 2); + + for (operands[2..], elements) |element_id, *element| + element.* = try context.resolveValue(element_id); + + const result = (try context.builder.appendInstruction(block, try context.translateType(operands[0]), .{ + .composite_construct = .{ + .elements = elements, + }, + }, context.nameOf(operands[1]))).?; + + try context.setValue(operands[1], result); + }, + .composite_extract => { + if (operands.len < 4) + return error.InvalidInstruction; + + const result = (try context.builder.appendInstruction(block, try context.translateType(operands[0]), .{ + .composite_extract = .{ + .composite = try context.resolveValue(operands[2]), + .indices = operands[3..], + }, + }, context.nameOf(operands[1]))).?; + + try context.setValue(operands[1], result); + }, + .function_call => return error.UnsupportedOpcode, + + else => return error.UnsupportedOpcode, + } +} + +fn translateFunctionControlFlow(context: *Context, spv_function: u32) !void { + var active = false; + var current_label: ?u32 = null; + var iterator = context.parser.iterator(); + while (try iterator.next()) |instruction| { + if (instruction.opcode == .function) { + active = instruction.operands.len >= 2 and instruction.operands[1] == spv_function; + continue; + } + + if (!active) + continue; + + const operands = instruction.operands; + switch (instruction.opcode) { + .label => { + try expectOperandCount(operands, 1); + current_label = operands[0]; + }, + .selection_merge => { + try expectOperandCount(operands, 2); + const block = context.module.blocks.getMut(try context.block(current_label orelse return error.InvalidBlock)).?; + block.structured_control = .{ + .selection = .{ + .merge_block = try context.block(operands[0]), + }, + }; + }, + .loop_merge => { + try expectOperandCount(operands, 3); + const block = context.module.blocks.getMut(try context.block(current_label orelse return error.InvalidBlock)).?; + block.structured_control = .{ + .loop = .{ + .merge_block = try context.block(operands[0]), + .continue_block = try context.block(operands[1]), + }, + }; + }, + .branch => { + try expectOperandCount(operands, 1); + const predecessor = current_label orelse return error.InvalidBlock; + try context.builder.setTerminator(try context.block(predecessor), .{ + .branch = try makeEdge(context, predecessor, operands[0]), + }); + }, + .branch_conditional => { + if (operands.len < 3 or operands.len > 5) return error.InvalidInstruction; + const predecessor = current_label orelse return error.InvalidBlock; + try context.builder.setTerminator(try context.block(predecessor), .{ .conditional_branch = .{ + .condition = try context.resolveValue(operands[0]), + .true_edge = try makeEdge(context, predecessor, operands[1]), + .false_edge = try makeEdge(context, predecessor, operands[2]), + } }); + }, + .return_ => { + try expectOperandCount(operands, 0); + try context.builder.setTerminator(try context.block(current_label orelse return error.InvalidBlock), .return_void); + }, + .return_value => { + try expectOperandCount(operands, 1); + try context.builder.setTerminator( + try context.block(current_label orelse return error.InvalidBlock), + .{ .return_value = try context.resolveValue(operands[0]) }, + ); + }, + .kill => { + try expectOperandCount(operands, 0); + try context.builder.setTerminator(try context.block(current_label orelse return error.InvalidBlock), .discard); + }, + .@"unreachable" => { + try expectOperandCount(operands, 0); + try context.builder.setTerminator(try context.block(current_label orelse return error.InvalidBlock), .@"unreachable"); + }, + .@"switch" => return error.UnsupportedOpcode, + .function_end => break, + + else => {}, + } + } +} + +fn makeEdge(context: *Context, predecessor_label: u32, target_label: u32) !ir.module.Edge { + var arguments: std.ArrayList(ir.id.ValueId) = .empty; + defer arguments.deinit(context.scratch); + + for (context.phi_infos.items) |phi| { + if (phi.target_label != target_label) + continue; + + var incoming: ?u32 = null; + var index: usize = 0; + + while (index < phi.incoming_words.len) : (index += 2) { + if (phi.incoming_words[index + 1] == predecessor_label) { + if (incoming != null) + return error.InvalidPhi; + + incoming = phi.incoming_words[index]; + } + } + try arguments.append(context.scratch, try context.resolveValue(incoming orelse return error.MissingPhiIncomingValue)); + } + + return context.builder.edge(try context.block(target_label), arguments.items); +} + +fn findFunction(parser: Parser, function_id: u32) !Parser.Instruction { + var iterator = parser.iterator(); + + while (try iterator.next()) |instruction| { + if (instruction.opcode == .function and instruction.operands.len >= 2 and instruction.operands[1] == function_id) + return instruction; + } + + return error.MissingFunction; +} + +fn functionTypeDefinition(context: *Context, type_id: u32) !Parser.Instruction { + const index = try context.idIndex(type_id); + const instruction = context.type_defs[index] orelse return error.InvalidFunctionType; + + if (instruction.opcode != .type_function) + return error.InvalidFunctionType; + + return instruction; +} + +fn translateStage(model: spirv.ExecutionModel) TranslationError!ir.module.Stage { + return switch (model) { + .vertex => .vertex, + .fragment => .fragment, + .gl_compute => .compute, + + else => error.UnsupportedExecutionModel, + }; +} + +fn translateStorageClass(storage_class: spirv.StorageClass) TranslationError!ir.types.AddressSpace { + return switch (storage_class) { + .function => .function, + .private => .private, + .workgroup => .workgroup, + .input => .input, + .output => .output, + .uniform, .uniform_constant => .uniform, + .storage_buffer => .storage, + .push_constant => .push_constant, + .physical_storage_buffer => .physical, + + else => error.UnsupportedType, + }; +} + +fn translateBuiltin(builtin: u32) TranslationError!ir.module.Builtin { + return switch (builtin) { + 0 => .position, + 15 => .frag_coord, + 22 => .frag_depth, + 28 => .global_invocation_id, + 42 => .vertex_index, + 43 => .instance_index, + + else => error.UnsupportedOpcode, + }; +} + +fn translateBinaryOpcode(opcode: spirv.Opcode) ir.instruction.BinaryOpcode { + return switch (opcode) { + .i_add => .integer_add, + .i_sub => .integer_subtract, + .i_mul => .integer_multiply, + .u_div => .unsigned_divide, + .s_div => .signed_divide, + .u_mod => .unsigned_modulo, + .s_mod => .signed_modulo, + .f_add => .float_add, + .f_sub => .float_subtract, + .f_mul => .float_multiply, + .f_div => .float_divide, + .f_mod => .float_modulo, + .shift_left_logical => .shift_left, + .shift_right_logical => .logical_shift_right, + .shift_right_arithmetic => .arithmetic_shift_right, + .bitwise_and => .bitwise_and, + .bitwise_or => .bitwise_or, + .bitwise_xor => .bitwise_xor, + .logical_and => .logical_and, + .logical_or => .logical_or, + + else => unreachable, + }; +} + +fn translateCompareOpcode(opcode: spirv.Opcode) ir.instruction.CompareOpcode { + return switch (opcode) { + .logical_equal, .i_equal => .equal, + .logical_not_equal, .i_not_equal => .not_equal, + .u_less_than => .unsigned_less, + .s_less_than => .signed_less, + .f_ord_equal => .ordered_float_equal, + .f_unord_equal => .unordered_float_equal, + .f_ord_not_equal => .ordered_float_not_equal, + .f_unord_not_equal => .unordered_float_not_equal, + .f_ord_less_than => .ordered_float_less, + .f_unord_less_than => .unordered_float_less, + + else => unreachable, + }; +} + +fn literalBits(words: []const u32) TranslationError!u64 { + return switch (words.len) { + 1 => words[0], + 2 => @as(u64, words[0]) | (@as(u64, words[1]) << 32), + + else => error.UnsupportedConstant, + }; +} + +fn isTypeOpcode(opcode: spirv.Opcode) bool { + return switch (opcode) { + .type_void, + .type_bool, + .type_int, + .type_float, + .type_vector, + .type_matrix, + .type_image, + .type_sampler, + .type_sampled_image, + .type_array, + .type_runtime_array, + .type_struct, + .type_opaque, + .type_pointer, + .type_function, + => true, + + else => false, + }; +} + +fn isConstantOpcode(opcode: spirv.Opcode) bool { + return switch (opcode) { + .constant_true, + .constant_false, + .constant, + .constant_composite, + .constant_null, + .spec_constant_true, + .spec_constant_false, + .spec_constant, + .spec_constant_composite, + .spec_constant_op, + => true, + + else => false, + }; +} + +fn expectOperandCount(operands: []const u32, expected: usize) TranslationError!void { + if (operands.len != expected) + return error.InvalidInstruction; +} + +fn allocOptional(comptime T: type, allocator: std.mem.Allocator, count: usize) ![]?T { + const values = try allocator.alloc(?T, count); + @memset(values, null); + return values; +} diff --git a/src/intel/FlintCommandBuffer.zig b/src/intel/FlintCommandBuffer.zig index 0c99444..8bc86a4 100644 --- a/src/intel/FlintCommandBuffer.zig +++ b/src/intel/FlintCommandBuffer.zig @@ -5,6 +5,8 @@ const kmd = @import("kmd.zig"); const VkError = base.VkError; const FlintDevice = @import("FlintDevice.zig"); +const FlintDescriptorSet = @import("FlintDescriptorSet.zig"); +const FlintPipeline = @import("FlintPipeline.zig"); const MemoryRange = @import("MemoryRange.zig"); @@ -17,6 +19,8 @@ pub const Interface = base.CommandBuffer; interface: Interface, batch: std.ArrayList(u32), relocations: std.ArrayList(kmd.Relocation), +bound_compute_pipeline: ?*FlintPipeline, +bound_compute_descriptor_sets: [base.vulkan_max_descriptor_sets]?*FlintDescriptorSet, pub fn create(device: *base.Device, allocator: std.mem.Allocator, info: *const vk.CommandBufferAllocateInfo) VkError!*Self { const self = allocator.create(Self) catch return VkError.OutOfHostMemory; @@ -80,6 +84,8 @@ pub fn create(device: *base.Device, allocator: std.mem.Allocator, info: *const v .interface = interface, .batch = .empty, .relocations = .empty, + .bound_compute_pipeline = null, + .bound_compute_descriptor_sets = @splat(null), }; return self; } @@ -96,8 +102,7 @@ pub fn submitGpuBatch(self: *Self, syncs: []const kmd.SyncDependency) VkError!vo try self.interface.submit(); defer self.interface.finish() catch @panic("Caught an error while handling an error"); - if (self.batch.items.len == 0) return; - + // Empty command buffers still need a no-op submission to carry queue synchronization. const device: *FlintDevice = @alignCast(@fieldParentPtr("interface", self.interface.owner)); const allocator = self.interface.host_allocator.allocator(); try device.kmd.submitBatch(self.interface.owner.io(), allocator, self.batch.items, self.relocations.items, syncs); @@ -122,6 +127,8 @@ pub fn reset(interface: *Interface, flags: vk.CommandBufferResetFlags) VkError!v self.batch.clearRetainingCapacity(); self.relocations.clearRetainingCapacity(); } + self.bound_compute_pipeline = null; + self.bound_compute_descriptor_sets = @splat(null); } pub fn emit(self: *Self, dword: u32) VkError!void { @@ -165,18 +172,27 @@ pub fn beginRenderPass(interface: *Interface, render_pass: *base.RenderPass, fra _ = clear_values; } -pub fn bindDescriptorSets(interface: *Interface, bind_point: vk.PipelineBindPoint, first_set: u32, sets: [base.VULKAN_MAX_DESCRIPTOR_SETS]?*base.DescriptorSet, dynamic_offsets: []const u32) VkError!void { - _ = interface; - _ = bind_point; - _ = first_set; - _ = sets; +pub fn bindDescriptorSets(interface: *Interface, bind_point: vk.PipelineBindPoint, first_set: u32, sets: [base.vulkan_max_descriptor_sets]?*base.DescriptorSet, dynamic_offsets: []const u32) VkError!void { + const self: *Self = @alignCast(@fieldParentPtr("interface", interface)); + if (bind_point != .compute) return; + if (first_set >= base.vulkan_max_descriptor_sets) return VkError.ValidationFailed; + + for (sets, 0..) |set, index| { + const base_set = set orelse break; + const destination = first_set + index; + if (destination >= base.vulkan_max_descriptor_sets) return VkError.ValidationFailed; + self.bound_compute_descriptor_sets[destination] = @alignCast(@fieldParentPtr("interface", base_set)); + } _ = dynamic_offsets; } pub fn bindPipeline(interface: *Interface, bind_point: vk.PipelineBindPoint, pipeline: *base.Pipeline) VkError!void { - _ = interface; - _ = bind_point; - _ = pipeline; + const self: *Self = @alignCast(@fieldParentPtr("interface", interface)); + if (bind_point != .compute) + return; + + const flint_pipeline: *FlintPipeline = @alignCast(@fieldParentPtr("interface", pipeline)); + self.bound_compute_pipeline = flint_pipeline; } pub fn bindIndexBuffer(interface: *Interface, buffer: *base.Buffer, offset: usize, index_type: vk.IndexType) VkError!void { @@ -269,10 +285,7 @@ pub fn copyQueryPoolResults(interface: *Interface, pool: *base.QueryPool, first: } pub fn dispatch(interface: *Interface, group_count_x: u32, group_count_y: u32, group_count_z: u32) VkError!void { - _ = interface; - _ = group_count_x; - _ = group_count_y; - _ = group_count_z; + try dispatchBase(interface, 0, 0, 0, group_count_x, group_count_y, group_count_z); } pub fn dispatchBase(interface: *Interface, base_group_x: u32, base_group_y: u32, base_group_z: u32, group_count_x: u32, group_count_y: u32, group_count_z: u32) VkError!void { diff --git a/src/intel/FlintDescriptorPool.zig b/src/intel/FlintDescriptorPool.zig index 6640041..cc01128 100644 --- a/src/intel/FlintDescriptorPool.zig +++ b/src/intel/FlintDescriptorPool.zig @@ -3,11 +3,15 @@ const vk = @import("vulkan"); const base = @import("base"); const VkError = base.VkError; +const VulkanAllocator = base.VulkanAllocator; + +const FlintDescriptorSet = @import("FlintDescriptorSet.zig"); const Self = @This(); pub const Interface = base.DescriptorPool; interface: Interface, +sets: std.ArrayList(*FlintDescriptorSet), pub fn create(device: *base.Device, allocator: std.mem.Allocator, info: *const vk.DescriptorPoolCreateInfo) VkError!*Self { const self = allocator.create(Self) catch return VkError.OutOfHostMemory; @@ -24,29 +28,42 @@ pub fn create(device: *base.Device, allocator: std.mem.Allocator, info: *const v self.* = .{ .interface = interface, + .sets = std.ArrayList(*FlintDescriptorSet).initCapacity(allocator, info.max_sets) catch return VkError.OutOfHostMemory, }; return self; } pub fn allocateDescriptorSet(interface: *Interface, layout: *base.DescriptorSetLayout) VkError!*base.DescriptorSet { const self: *Self = @alignCast(@fieldParentPtr("interface", interface)); - _ = self; - _ = layout; - return VkError.Unknown; + if (self.sets.items.len == self.sets.capacity) return VkError.OutOfPoolMemory; + + const allocator = VulkanAllocator.init(null, .object).allocator(); + const set = try FlintDescriptorSet.create(interface.owner, allocator, layout); + self.sets.appendAssumeCapacity(set); + return &set.interface; } pub fn destroy(interface: *Interface, allocator: std.mem.Allocator) void { const self: *Self = @alignCast(@fieldParentPtr("interface", interface)); + const set_allocator = VulkanAllocator.init(null, .object).allocator(); + for (self.sets.items) |set| set.interface.destroy(set_allocator); + self.sets.deinit(allocator); allocator.destroy(self); } pub fn freeDescriptorSet(interface: *Interface, set: *base.DescriptorSet) VkError!void { const self: *Self = @alignCast(@fieldParentPtr("interface", interface)); - _ = self; - _ = set; + const flint_set: *FlintDescriptorSet = @alignCast(@fieldParentPtr("interface", set)); + const index = std.mem.indexOfScalar(*FlintDescriptorSet, self.sets.items, flint_set) orelse return VkError.ValidationFailed; + _ = self.sets.orderedRemove(index); + + const allocator = VulkanAllocator.init(null, .object).allocator(); + set.destroy(allocator); } pub fn reset(interface: *Interface, _: vk.DescriptorPoolResetFlags) VkError!void { const self: *Self = @alignCast(@fieldParentPtr("interface", interface)); - _ = self; + const allocator = VulkanAllocator.init(null, .object).allocator(); + for (self.sets.items) |set| set.interface.destroy(allocator); + self.sets.clearRetainingCapacity(); } diff --git a/src/intel/FlintDescriptorSet.zig b/src/intel/FlintDescriptorSet.zig index 7472fb9..88de776 100644 --- a/src/intel/FlintDescriptorSet.zig +++ b/src/intel/FlintDescriptorSet.zig @@ -3,11 +3,25 @@ const vk = @import("vulkan"); const base = @import("base"); const VkError = base.VkError; +const NonDispatchable = base.NonDispatchable; const Self = @This(); pub const Interface = base.DescriptorSet; +pub const DescriptorBuffer = struct { + buffer: ?*base.Buffer, + offset: vk.DeviceSize, + size: vk.DeviceSize, +}; + +const Descriptor = union(enum) { + buffer: []DescriptorBuffer, + unsupported, +}; + interface: Interface, +heap: []u8, +descriptors: []Descriptor, pub fn create(device: *base.Device, allocator: std.mem.Allocator, layout: *base.DescriptorSetLayout) VkError!*Self { const self = allocator.create(Self) catch return VkError.OutOfHostMemory; @@ -21,26 +35,139 @@ pub fn create(device: *base.Device, allocator: std.mem.Allocator, layout: *base. .write = write, }; + var heap_size = layout.bindings.len * @sizeOf(Descriptor); + for (layout.bindings) |binding| { + heap_size += switch (binding.descriptor_type) { + .uniform_buffer, + .uniform_buffer_dynamic, + .storage_buffer, + .storage_buffer_dynamic, + => binding.array_size * @sizeOf(DescriptorBuffer), + + else => 0, + }; + } + + const heap = allocator.alloc(u8, heap_size) catch return VkError.OutOfHostMemory; + errdefer allocator.free(heap); + + var fixed = std.heap.FixedBufferAllocator.init(heap); + const descriptors = fixed.allocator().alloc(Descriptor, layout.bindings.len) catch return VkError.OutOfHostMemory; + for (descriptors, layout.bindings) |*descriptor, binding| { + descriptor.* = switch (binding.descriptor_type) { + .uniform_buffer, .uniform_buffer_dynamic, .storage_buffer, .storage_buffer_dynamic => blk: { + const buffers = fixed.allocator().alloc(DescriptorBuffer, binding.array_size) catch return VkError.OutOfHostMemory; + + for (buffers) |*buffer| { + buffer.* = .{ + .buffer = null, + .offset = 0, + .size = 0, + }; + } + + break :blk .{ .buffer = buffers }; + }, + else => .unsupported, + }; + } + self.* = .{ .interface = interface, + .heap = heap, + .descriptors = descriptors, }; return self; } pub fn destroy(interface: *Interface, allocator: std.mem.Allocator) void { const self: *Self = @alignCast(@fieldParentPtr("interface", interface)); + allocator.free(self.heap); allocator.destroy(self); } pub fn copy(interface: *Interface, src_interface: *const Interface, data: vk.CopyDescriptorSet) VkError!void { const self: *Self = @alignCast(@fieldParentPtr("interface", interface)); - _ = self; - _ = src_interface; - _ = data; + const src: *const Self = @alignCast(@fieldParentPtr("interface", src_interface)); + + if (data.dst_binding >= self.descriptors.len or data.src_binding >= src.descriptors.len) + return VkError.ValidationFailed; + + const dst = switch (self.descriptors[data.dst_binding]) { + .buffer => |buffers| buffers, + .unsupported => return VkError.FeatureNotPresent, + }; + + const source = switch (src.descriptors[data.src_binding]) { + .buffer => |buffers| buffers, + .unsupported => return VkError.FeatureNotPresent, + }; + + const dst_start: usize = @intCast(data.dst_array_element); + const src_start: usize = @intCast(data.src_array_element); + const count: usize = @intCast(data.descriptor_count); + + if (dst_start > dst.len or count > dst.len - dst_start) + return VkError.ValidationFailed; + if (src_start > source.len or count > source.len - src_start) + return VkError.ValidationFailed; + + @memcpy(dst[dst_start .. dst_start + count], source[src_start .. src_start + count]); } pub fn write(interface: *Interface, write_data: vk.WriteDescriptorSet) VkError!void { const self: *Self = @alignCast(@fieldParentPtr("interface", interface)); - _ = self; - _ = write_data; + + switch (write_data.descriptor_type) { + .uniform_buffer, + .uniform_buffer_dynamic, + .storage_buffer, + .storage_buffer_dynamic, + => { + if (write_data.dst_binding >= self.descriptors.len) + return VkError.ValidationFailed; + + const descriptors = switch (self.descriptors[write_data.dst_binding]) { + .buffer => |buffers| buffers, + .unsupported => return VkError.FeatureNotPresent, + }; + + const start: usize = @intCast(write_data.dst_array_element); + const count: usize = @intCast(write_data.descriptor_count); + + if (start > descriptors.len or count > descriptors.len - start) + return VkError.ValidationFailed; + + for (write_data.p_buffer_info, 0..write_data.descriptor_count) |buffer_info, index| { + const descriptor = &descriptors[start + index]; + descriptor.* = .{ .buffer = null, .offset = buffer_info.offset, .size = buffer_info.range }; + + if (buffer_info.buffer == .null_handle) + continue; + + const buffer = try NonDispatchable(base.Buffer).fromHandleObject(buffer_info.buffer); + + if (descriptor.offset > buffer.size) + return VkError.ValidationFailed; + if (descriptor.size == vk.WHOLE_SIZE) + descriptor.size = buffer.size - descriptor.offset; + if (descriptor.size > buffer.size - descriptor.offset) + return VkError.ValidationFailed; + + descriptor.buffer = buffer; + } + }, + + else => return VkError.FeatureNotPresent, + } +} + +pub fn getBuffer(self: *const Self, binding: u32, array_element: u32) VkError!DescriptorBuffer { + if (binding >= self.descriptors.len) return VkError.ValidationFailed; + const buffers = switch (self.descriptors[binding]) { + .buffer => |items| items, + .unsupported => return VkError.FeatureNotPresent, + }; + if (array_element >= buffers.len) return VkError.ValidationFailed; + return buffers[array_element]; } diff --git a/src/intel/FlintImage.zig b/src/intel/FlintImage.zig index 5c002e3..cca13a5 100644 --- a/src/intel/FlintImage.zig +++ b/src/intel/FlintImage.zig @@ -46,8 +46,8 @@ pub fn destroy(interface: *Interface, allocator: std.mem.Allocator) void { } pub fn getMemoryRequirements(_: *Interface, requirements: *vk.MemoryRequirements) VkError!void { - requirements.alignment = lib.IMAGE_MEMORY_ALIGNMENT; - requirements.size = std.mem.alignForward(vk.DeviceSize, requirements.size, lib.IMAGE_MEMORY_ALIGNMENT); + requirements.alignment = lib.image_memory_alignment; + requirements.size = std.mem.alignForward(vk.DeviceSize, requirements.size, lib.image_memory_alignment); } pub fn copyToMemory(interface: *const Interface, dst: []u8, subresource: vk.ImageSubresourceLayers) VkError!void { diff --git a/src/intel/FlintInstance.zig b/src/intel/FlintInstance.zig index 93fa669..f161f16 100644 --- a/src/intel/FlintInstance.zig +++ b/src/intel/FlintInstance.zig @@ -26,7 +26,7 @@ fn castExtension(comptime ext: vk.ApiInfo) vk.ExtensionProperties { return props; } -pub const EXTENSIONS = [_]vk.ExtensionProperties{ +pub const extensions = [_]vk.ExtensionProperties{ castExtension(vk.extensions.khr_device_group_creation), castExtension(vk.extensions.khr_get_physical_device_properties_2), castExtension(vk.extensions.khr_surface), @@ -73,7 +73,7 @@ fn requestPhysicalDevices(interface: *Interface, allocator: std.mem.Allocator, d if (drm_device.node_type != .render or std.meta.activeTag(drm_device.device_info) != .pci or - drm_device.device_info.pci.vendor_id != lib.INTEL_PCI_VENDOR_ID) + drm_device.device_info.pci.vendor_id != lib.intel_pci_vendor_id) continue; const version = device.getVersion(io_var, allocator) catch continue; diff --git a/src/intel/FlintPhysicalDevice.zig b/src/intel/FlintPhysicalDevice.zig index db4865f..3866663 100644 --- a/src/intel/FlintPhysicalDevice.zig +++ b/src/intel/FlintPhysicalDevice.zig @@ -23,7 +23,7 @@ fn castExtension(comptime ext: vk.ApiInfo) vk.ExtensionProperties { return props; } -pub const EXTENSIONS = [_]vk.ExtensionProperties{ +pub const extensions = [_]vk.ExtensionProperties{ castExtension(vk.extensions.khr_swapchain), }; @@ -53,9 +53,9 @@ pub fn create(allocator: std.mem.Allocator, instance: *base.Instance, drm_device .getSurfaceSupportKHR = getSurfaceSupportKHR, }; - interface.props.api_version = @bitCast(lib.VULKAN_VERSION); - interface.props.vendor_id = lib.INTEL_PCI_VENDOR_ID; - interface.props.driver_version = @bitCast(base.DRIVER_VERSION); + interface.props.api_version = @bitCast(lib.vulkan_version); + interface.props.vendor_id = lib.intel_pci_vendor_id; + interface.props.driver_version = @bitCast(base.driver_version); interface.props.device_id = drm_device.device_info.pci.device_id; interface.props.device_type = .integrated_gpu; @@ -256,10 +256,10 @@ pub fn enumerateExtensionProperties(_: *const Interface, layer_name: ?[]const u8 return VkError.LayerNotPresent; } - const available = EXTENSIONS.len; + const available = extensions.len; if (p_properties) |properties| { const write_count = @min(count.*, available); - for (EXTENSIONS[0..write_count], properties[0..write_count]) |ext, *prop| { + for (extensions[0..write_count], properties[0..write_count]) |ext, *prop| { prop.* = ext; } count.* = @intCast(write_count); diff --git a/src/intel/FlintPipeline.zig b/src/intel/FlintPipeline.zig index a61048b..80a819c 100644 --- a/src/intel/FlintPipeline.zig +++ b/src/intel/FlintPipeline.zig @@ -8,6 +8,7 @@ const Self = @This(); pub const Interface = base.Pipeline; interface: Interface, +host_allocator: base.VulkanAllocator, pub fn createCompute(device: *base.Device, allocator: std.mem.Allocator, cache: ?*base.PipelineCache, info: *const vk.ComputePipelineCreateInfo) VkError!*Self { const self = allocator.create(Self) catch return VkError.OutOfHostMemory; @@ -15,8 +16,12 @@ pub fn createCompute(device: *base.Device, allocator: std.mem.Allocator, cache: var interface = try Interface.initCompute(device, allocator, cache, info); interface.vtable = &.{ .destroy = destroy }; + self.* = .{ + .interface = interface, + .host_allocator = base.VulkanAllocator.from(allocator).clone(), + }; + errdefer self.interface.layout.unref(allocator); - self.* = .{ .interface = interface }; return self; } @@ -27,7 +32,12 @@ pub fn createGraphics(device: *base.Device, allocator: std.mem.Allocator, cache: var interface = try Interface.initGraphics(device, allocator, cache, info); interface.vtable = &.{ .destroy = destroy }; - self.* = .{ .interface = interface }; + self.* = .{ + .interface = interface, + .host_allocator = base.VulkanAllocator.from(allocator).clone(), + }; + errdefer self.interface.layout.unref(allocator); + return self; } diff --git a/src/intel/FlintShaderModule.zig b/src/intel/FlintShaderModule.zig index c91d96a..970e099 100644 --- a/src/intel/FlintShaderModule.zig +++ b/src/intel/FlintShaderModule.zig @@ -8,6 +8,7 @@ const Self = @This(); pub const Interface = base.ShaderModule; interface: Interface, +code: []u32, ref_count: std.atomic.Value(usize), pub fn create(device: *base.Device, allocator: std.mem.Allocator, info: *const vk.ShaderModuleCreateInfo) VkError!*Self { @@ -16,9 +17,13 @@ pub fn create(device: *base.Device, allocator: std.mem.Allocator, info: *const v var interface = try Interface.init(device, allocator, info); interface.vtable = &.{ .destroy = destroy }; + if (info.code_size % @sizeOf(u32) != 0) return VkError.ValidationFailed; + const code = allocator.dupe(u32, info.p_code[0 .. info.code_size / @sizeOf(u32)]) catch return VkError.OutOfHostMemory; + errdefer allocator.free(code); self.* = .{ .interface = interface, + .code = code, .ref_count = std.atomic.Value(usize).init(1), }; return self; @@ -30,6 +35,7 @@ pub fn destroy(interface: *Interface, allocator: std.mem.Allocator) void { } pub fn drop(self: *Self, allocator: std.mem.Allocator) void { + allocator.free(self.code); allocator.destroy(self); } diff --git a/src/intel/blitter.zig b/src/intel/blitter.zig index 8618850..d54fd44 100644 --- a/src/intel/blitter.zig +++ b/src/intel/blitter.zig @@ -32,7 +32,7 @@ pub fn blitImageRegion(cmd: *FlintCommandBuffer, src: *base.Image, dst: *base.Im if (base.format.isCompressed(src_format) or base.format.isCompressed(dst_format)) return VkError.FormatNotSupported; - // XY_SRC_COPY_BLT does not perform component conversion. Different Vulkan + // xy_src_copy_blt does not perform component conversion. Different Vulkan // names are safe only when they describe the same bytes and values. if (!bitwiseCompatibleFormats(src_format, dst_format) or src_texel_size != dst_texel_size) return VkError.FormatNotSupported; diff --git a/src/intel/i915/i915.zig b/src/intel/i915/i915.zig index b07bb47..7ba8bc8 100644 --- a/src/intel/i915/i915.zig +++ b/src/intel/i915/i915.zig @@ -1,4 +1,3 @@ - pub const command_base = 0x40; pub const gem_create = 0x1b; pub const gem_mmap_gtt = 0x24; diff --git a/src/intel/kmd.zig b/src/intel/kmd.zig index 299b775..423de09 100644 --- a/src/intel/kmd.zig +++ b/src/intel/kmd.zig @@ -9,7 +9,7 @@ const i915_kmd = @import("i915/kmd.zig"); const xe = @import("xe/kmd.zig"); const VkError = base.VkError; -const IOCTL = std.os.linux.IOCTL; +const ioctl = std.os.linux.IOCTL; pub const xy_src_copy_blt: u32 = (2 << 29) | (0x53 << 22) | 8; pub const xy_blt_write_alpha: u32 = 1 << 21; @@ -160,9 +160,9 @@ pub const Memory = union(KmdType) { }; pub inline fn drmIoctlIow(nr: u8, comptime T: type) u32 { - return IOCTL.IOW('d', nr, T); + return ioctl.IOW('d', nr, T); } pub inline fn drmIoctlIowr(nr: u8, comptime T: type) u32 { - return IOCTL.IOWR('d', nr, T); + return ioctl.IOWR('d', nr, T); } diff --git a/src/intel/lib.zig b/src/intel/lib.zig index c5532d1..5146a75 100644 --- a/src/intel/lib.zig +++ b/src/intel/lib.zig @@ -35,13 +35,13 @@ pub const FlintShaderModule = @import("FlintShaderModule.zig"); pub const Instance = FlintInstance; -pub const DRIVER_NAME = "Flint"; +pub const driver_name = "Flint"; -pub const PHYSICAL_DEVICE_DEFAULT_NAME = "Unkown Intel device"; +pub const physical_device_default_name = "Unkown Intel device"; -pub const INTEL_PCI_VENDOR_ID = 0x8086; +pub const intel_pci_vendor_id = 0x8086; -pub const VULKAN_VERSION = vk.makeApiVersion( +pub const vulkan_version = vk.makeApiVersion( 0, config.flint_vulkan_version.major, config.flint_vulkan_version.minor, @@ -49,7 +49,7 @@ pub const VULKAN_VERSION = vk.makeApiVersion( ); /// GEM buffer objects are page based -pub const IMAGE_MEMORY_ALIGNMENT = std.heap.page_size_max; +pub const image_memory_alignment = std.heap.page_size_max; pub const KmdType = enum { invalid, diff --git a/src/phi/PhiCommandBuffer.zig b/src/phi/PhiCommandBuffer.zig index d358639..41e4d8a 100644 --- a/src/phi/PhiCommandBuffer.zig +++ b/src/phi/PhiCommandBuffer.zig @@ -154,7 +154,7 @@ pub fn beginRenderPass(interface: *Interface, render_pass: *base.RenderPass, fra _ = clear_values; } -pub fn bindDescriptorSets(interface: *Interface, bind_point: vk.PipelineBindPoint, first_set: u32, sets: [base.VULKAN_MAX_DESCRIPTOR_SETS]?*base.DescriptorSet, dynamic_offsets: []const u32) VkError!void { +pub fn bindDescriptorSets(interface: *Interface, bind_point: vk.PipelineBindPoint, first_set: u32, sets: [base.vulkan_max_descriptor_sets]?*base.DescriptorSet, dynamic_offsets: []const u32) VkError!void { const self: *Self = @alignCast(@fieldParentPtr("interface", interface)); self.cmd_count += 1; _ = bind_point; diff --git a/src/phi/PhiInstance.zig b/src/phi/PhiInstance.zig index 637ceeb..ea3c6eb 100644 --- a/src/phi/PhiInstance.zig +++ b/src/phi/PhiInstance.zig @@ -27,7 +27,7 @@ fn castExtension(comptime ext: vk.ApiInfo) vk.ExtensionProperties { return props; } -pub const EXTENSIONS = [_]vk.ExtensionProperties{ +pub const extensions = [_]vk.ExtensionProperties{ castExtension(vk.extensions.khr_device_group_creation), castExtension(vk.extensions.khr_get_physical_device_properties_2), castExtension(vk.extensions.khr_surface), diff --git a/src/phi/PhiPhysicalDevice.zig b/src/phi/PhiPhysicalDevice.zig index b325022..d99fc0c 100644 --- a/src/phi/PhiPhysicalDevice.zig +++ b/src/phi/PhiPhysicalDevice.zig @@ -24,7 +24,7 @@ fn castExtension(comptime ext: vk.ApiInfo) vk.ExtensionProperties { return props; } -pub const EXTENSIONS = [_]vk.ExtensionProperties{ +pub const extensions = [_]vk.ExtensionProperties{ castExtension(vk.extensions.khr_device_group), castExtension(vk.extensions.khr_swapchain), }; @@ -55,8 +55,8 @@ pub fn create(allocator: std.mem.Allocator, instance: *base.Instance, mic_device .getSurfaceSupportKHR = getSurfaceSupportKHR, }; - interface.props.api_version = @bitCast(lib.VULKAN_VERSION); - interface.props.driver_version = @bitCast(base.DRIVER_VERSION); + interface.props.api_version = @bitCast(lib.vulkan_version); + interface.props.driver_version = @bitCast(base.driver_version); interface.props.device_type = .other; @memset(interface.props.device_name[0..], 0); @@ -292,10 +292,10 @@ pub fn enumerateExtensionProperties(_: *const Interface, layer_name: ?[]const u8 return VkError.LayerNotPresent; } - const available = EXTENSIONS.len; + const available = extensions.len; if (p_properties) |properties| { const write_count = @min(count.*, available); - for (EXTENSIONS[0..write_count], properties[0..write_count]) |ext, *prop| { + for (extensions[0..write_count], properties[0..write_count]) |ext, *prop| { prop.* = ext; } count.* = @intCast(write_count); diff --git a/src/phi/PhiTransport.zig b/src/phi/PhiTransport.zig index 7ff4a82..ebc4e83 100644 --- a/src/phi/PhiTransport.zig +++ b/src/phi/PhiTransport.zig @@ -96,7 +96,7 @@ pub fn statusToErr(status: c_int) VkError { fn writeAll(self: *Self, bytes: []const u8) VkError!void { var offset: usize = 0; while (offset < bytes.len) { - const written = scif.send(self.epd, bytes[offset..].ptr, bytes.len - offset, scif.SEND_BLOCK); + const written = scif.send(self.epd, bytes[offset..].ptr, bytes.len - offset, scif.send_block); if (written <= 0) { return VkError.InitializationFailed; } @@ -107,7 +107,7 @@ fn writeAll(self: *Self, bytes: []const u8) VkError!void { fn readAll(self: *Self, bytes: []u8) VkError!void { var offset: usize = 0; while (offset < bytes.len) { - const read = scif.recv(self.epd, bytes[offset..].ptr, bytes.len - offset, scif.RECV_BLOCK); + const read = scif.recv(self.epd, bytes[offset..].ptr, bytes.len - offset, scif.recv_block); if (read <= 0) { return VkError.InitializationFailed; } diff --git a/src/phi/lib.zig b/src/phi/lib.zig index da38a15..42d492d 100644 --- a/src/phi/lib.zig +++ b/src/phi/lib.zig @@ -38,11 +38,11 @@ pub const PhiShaderModule = @import("PhiShaderModule.zig"); pub const Instance = PhiInstance; -pub const DRIVER_NAME = "Phi"; +pub const driver_name = "Phi"; -pub const PHYSICAL_DEVICE_DEFAULT_NAME = "Intel(R) Xeon Phi(TM) Coprocessor"; +pub const physical_device_default_name = "Intel(R) Xeon Phi(TM) Coprocessor"; -pub const VULKAN_VERSION = vk.makeApiVersion( +pub const vulkan_version = vk.makeApiVersion( 0, config.phi_vulkan_version.major, config.phi_vulkan_version.minor, diff --git a/src/phi/pci_ids.zig b/src/phi/pci_ids.zig index 36e2308..c1dbeec 100644 --- a/src/phi/pci_ids.zig +++ b/src/phi/pci_ids.zig @@ -1,4 +1,3 @@ - const PciInfo = struct { id: u16, name: []const u8, diff --git a/src/phi/scif.zig b/src/phi/scif.zig index b8e3c5e..8f4a943 100644 --- a/src/phi/scif.zig +++ b/src/phi/scif.zig @@ -10,8 +10,8 @@ pub const PortId = extern struct { port: u16, }; -pub const SEND_BLOCK = 1; -pub const RECV_BLOCK = 1; +pub const send_block = 1; +pub const recv_block = 1; // SAFETY: load assigns every function pointer before the public wrappers can be used. var scif_open: *const fn () callconv(.c) epd_t = undefined; diff --git a/src/software/SoftBuffer.zig b/src/software/SoftBuffer.zig index 28eb2d7..34b1cf4 100644 --- a/src/software/SoftBuffer.zig +++ b/src/software/SoftBuffer.zig @@ -17,7 +17,7 @@ pub fn create(device: *base.Device, allocator: std.mem.Allocator, info: *const v var interface = try Interface.init(device, allocator, info); interface.allowed_memory_types = std.bit_set.IntegerBitSet(32).initEmpty(); - interface.allowed_memory_types.set(lib.MEMORY_TYPE_GENERIC_BIT); + interface.allowed_memory_types.set(lib.memory_type_generic_bit); interface.vtable = &.{ .destroy = destroy, @@ -36,15 +36,15 @@ pub fn destroy(interface: *Interface, allocator: std.mem.Allocator) void { } pub fn getMemoryRequirements(interface: *Interface, requirements: *vk.MemoryRequirements) void { - requirements.alignment = lib.MEMORY_REQUIREMENTS_BUFFER_ALIGNMENT; + requirements.alignment = lib.memory_requirements_buffer_alignment; if (interface.usage.uniform_texel_buffer_bit or interface.usage.storage_texel_buffer_bit) { - requirements.alignment = @max(requirements.alignment, lib.MIN_TEXEL_BUFFER_ALIGNMENT); + requirements.alignment = @max(requirements.alignment, lib.min_texel_buffer_alignment); } if (interface.usage.storage_buffer_bit) { - requirements.alignment = @max(requirements.alignment, lib.MIN_STORAGE_BUFFER_ALIGNMENT); + requirements.alignment = @max(requirements.alignment, lib.min_storage_buffer_alignment); } if (interface.usage.uniform_buffer_bit) { - requirements.alignment = @max(requirements.alignment, lib.MIN_UNIFORM_BUFFER_ALIGNMENT); + requirements.alignment = @max(requirements.alignment, lib.min_uniform_buffer_alignment); } } diff --git a/src/software/SoftCommandBuffer.zig b/src/software/SoftCommandBuffer.zig index f4153ee..3ba0bd4 100644 --- a/src/software/SoftCommandBuffer.zig +++ b/src/software/SoftCommandBuffer.zig @@ -407,7 +407,7 @@ pub fn beginRenderPass(interface: *Interface, render_pass: *base.RenderPass, fra self.commands.append(allocator, .{ .ptr = cmd, .vtable = &.{ .execute = CommandImpl.execute } }) catch return VkError.OutOfHostMemory; } -pub fn bindDescriptorSets(interface: *Interface, bind_point: vk.PipelineBindPoint, first_set: u32, sets: [base.VULKAN_MAX_DESCRIPTOR_SETS]?*base.DescriptorSet, dynamic_offsets: []const u32) VkError!void { +pub fn bindDescriptorSets(interface: *Interface, bind_point: vk.PipelineBindPoint, first_set: u32, sets: [base.vulkan_max_descriptor_sets]?*base.DescriptorSet, dynamic_offsets: []const u32) VkError!void { const self: *Self = @alignCast(@fieldParentPtr("interface", interface)); const allocator = self.command_allocator.allocator(); @@ -416,7 +416,7 @@ pub fn bindDescriptorSets(interface: *Interface, bind_point: vk.PipelineBindPoin bind_point: vk.PipelineBindPoint, first_set: u32, - sets: [base.VULKAN_MAX_DESCRIPTOR_SETS]?*base.DescriptorSet, + sets: [base.vulkan_max_descriptor_sets]?*base.DescriptorSet, dynamic_offsets: []const u32, pub fn execute(context: *anyopaque, device: *ExecutionDevice) VkError!void { @@ -430,7 +430,7 @@ pub fn bindDescriptorSets(interface: *Interface, bind_point: vk.PipelineBindPoin state.sets[i] = soft_set; const dynamic_count = soft_set.interface.layout.dynamic_descriptor_count; - if (dynamic_count > ExecutionDevice.MAX_DYNAMIC_DESCRIPTORS_PER_SET or + if (dynamic_count > ExecutionDevice.max_dynamic_descriptors_per_set or dynamic_offset_index + dynamic_count > impl.dynamic_offsets.len) { return VkError.ValidationFailed; @@ -496,7 +496,7 @@ pub fn bindIndexBuffer(interface: *Interface, buffer: *base.Buffer, offset: usiz pub fn execute(context: *anyopaque, device: *ExecutionDevice) VkError!void { const impl: *Impl = @ptrCast(@alignCast(context)); - device.pipeline_states[ExecutionDevice.GRAPHICS_PIPELINE_STATE].data.graphics.index_buffer = .{ + device.pipeline_states[ExecutionDevice.graphics_pipeline_state].data.graphics.index_buffer = .{ .buffer = impl.buffer, .offset = impl.offset, .index_type = impl.index_type, @@ -527,7 +527,7 @@ pub fn bindVertexBuffer(interface: *Interface, index: usize, buffer: *base.Buffe pub fn execute(context: *anyopaque, device: *ExecutionDevice) VkError!void { const impl: *Impl = @ptrCast(@alignCast(context)); - device.pipeline_states[ExecutionDevice.GRAPHICS_PIPELINE_STATE].data.graphics.vertex_buffers[impl.index] = .{ + device.pipeline_states[ExecutionDevice.graphics_pipeline_state].data.graphics.vertex_buffers[impl.index] = .{ .buffer = impl.buffer, .offset = impl.offset, .size = 0, @@ -1191,7 +1191,7 @@ pub fn pushConstants(interface: *Interface, stages: vk.ShaderStageFlags, offset: pub fn execute(context: *anyopaque, device: *ExecutionDevice) VkError!void { const impl: *Impl = @ptrCast(@alignCast(context)); - const size = @min(lib.PUSH_CONSTANT_SIZE - impl.offset, impl.blob.len); + const size = @min(lib.push_constant_size - impl.offset, impl.blob.len); if (impl.stages.vertex_bit or impl.stages.tessellation_control_bit or @@ -1199,12 +1199,12 @@ pub fn pushConstants(interface: *Interface, stages: vk.ShaderStageFlags, offset: impl.stages.geometry_bit or impl.stages.fragment_bit) { - const state = &device.pipeline_states[ExecutionDevice.GRAPHICS_PIPELINE_STATE]; + const state = &device.pipeline_states[ExecutionDevice.graphics_pipeline_state]; @memcpy(state.push_constant_blob[impl.offset .. impl.offset + size], impl.blob[0..size]); } if (impl.stages.compute_bit) { - const state = &device.pipeline_states[ExecutionDevice.COMPUTE_PIPELINE_STATE]; + const state = &device.pipeline_states[ExecutionDevice.compute_pipeline_state]; @memcpy(state.push_constant_blob[impl.offset .. impl.offset + size], impl.blob[0..size]); } } diff --git a/src/software/SoftImage.zig b/src/software/SoftImage.zig index 4aa47a5..881926b 100644 --- a/src/software/SoftImage.zig +++ b/src/software/SoftImage.zig @@ -97,7 +97,7 @@ pub fn destroy(interface: *Interface, allocator: std.mem.Allocator) void { } pub fn getMemoryRequirements(_: *Interface, requirements: *vk.MemoryRequirements) VkError!void { - requirements.alignment = lib.MEMORY_REQUIREMENTS_IMAGE_ALIGNMENT; + requirements.alignment = lib.memory_requirements_image_alignment; } pub fn getClearFormat(self: *Self) VkError!vk.Format { diff --git a/src/software/SoftInstance.zig b/src/software/SoftInstance.zig index 6a20c7f..3002029 100644 --- a/src/software/SoftInstance.zig +++ b/src/software/SoftInstance.zig @@ -24,7 +24,7 @@ fn castExtension(comptime ext: vk.ApiInfo) vk.ExtensionProperties { return props; } -pub const EXTENSIONS = [_]vk.ExtensionProperties{ +pub const extensions = [_]vk.ExtensionProperties{ castExtension(vk.extensions.khr_device_group_creation), castExtension(vk.extensions.khr_get_physical_device_properties_2), castExtension(vk.extensions.khr_surface), diff --git a/src/software/SoftPhysicalDevice.zig b/src/software/SoftPhysicalDevice.zig index 0eb9e33..050517d 100644 --- a/src/software/SoftPhysicalDevice.zig +++ b/src/software/SoftPhysicalDevice.zig @@ -22,7 +22,7 @@ fn castExtension(comptime ext: vk.ApiInfo) vk.ExtensionProperties { return props; } -pub const EXTENSIONS = [_]vk.ExtensionProperties{ +pub const extensions = [_]vk.ExtensionProperties{ castExtension(vk.extensions.khr_device_group), castExtension(vk.extensions.khr_swapchain), }; @@ -56,11 +56,11 @@ pub fn create(allocator: std.mem.Allocator, instance: *base.Instance) VkError!*S .getSurfaceSupportKHR = getSurfaceSupportKHR, }; - interface.props.api_version = @bitCast(lib.VULKAN_VERSION); - interface.props.driver_version = @bitCast(base.DRIVER_VERSION); - interface.props.device_id = lib.DEVICE_ID; + interface.props.api_version = @bitCast(lib.vulkan_version); + interface.props.driver_version = @bitCast(base.driver_version); + interface.props.device_id = lib.device_id; interface.props.device_type = .cpu; - interface.props.pipeline_cache_uuid = lib.PIPELINE_CACHE_UUID; + interface.props.pipeline_cache_uuid = lib.pipeline_cache_uuid; interface.props.limits = .{ .max_image_dimension_1d = 4096, .max_image_dimension_2d = 4096, @@ -70,12 +70,12 @@ pub fn create(allocator: std.mem.Allocator, instance: *base.Instance) VkError!*S .max_texel_buffer_elements = 65536, .max_uniform_buffer_range = 16384, .max_storage_buffer_range = 134217728, - .max_push_constants_size = lib.PUSH_CONSTANT_SIZE, + .max_push_constants_size = lib.push_constant_size, .max_memory_allocation_count = std.math.maxInt(u32), .max_sampler_allocation_count = 4096, .buffer_image_granularity = 131072, .sparse_address_space_size = 0, - .max_bound_descriptor_sets = base.VULKAN_MAX_DESCRIPTOR_SETS, + .max_bound_descriptor_sets = base.vulkan_max_descriptor_sets, .max_per_stage_descriptor_samplers = 16, .max_per_stage_descriptor_uniform_buffers = 12, .max_per_stage_descriptor_storage_buffers = 4, @@ -91,8 +91,8 @@ pub fn create(allocator: std.mem.Allocator, instance: *base.Instance) VkError!*S .max_descriptor_set_sampled_images = 96, .max_descriptor_set_storage_images = 24, .max_descriptor_set_input_attachments = 4, - .max_vertex_input_attributes = lib.MAX_VERTEX_INPUT_ATTRIBUTES, - .max_vertex_input_bindings = lib.MAX_VERTEX_INPUT_BINDINGS, + .max_vertex_input_attributes = lib.max_vertex_input_attributes, + .max_vertex_input_bindings = lib.max_vertex_input_bindings, .max_vertex_input_attribute_offset = 2047, .max_vertex_input_binding_stride = 2048, .max_vertex_output_components = 64, @@ -182,7 +182,7 @@ pub fn create(allocator: std.mem.Allocator, instance: *base.Instance) VkError!*S }; interface.mem_props.memory_heap_count = 1; interface.mem_props.memory_heaps[0] = .{ - .size = std.process.totalSystemMemory() catch lib.PHYSICAL_DEVICE_FALLBACK_HEAP_SIZE, + .size = std.process.totalSystemMemory() catch lib.physical_device_fallback_heap_size, .flags = .{ .device_local_bit = true }, }; @@ -240,12 +240,12 @@ pub fn create(allocator: std.mem.Allocator, instance: *base.Instance) VkError!*S } } - break :blk command_allocator.dupe(u8, lib.PHYSICAL_DEVICE_DEFAULT_NAME) catch return VkError.OutOfHostMemory; + break :blk command_allocator.dupe(u8, lib.physical_device_default_name) catch return VkError.OutOfHostMemory; }; defer command_allocator.free(name); var writer = std.Io.Writer.fixed(device_name[0 .. vk.MAX_PHYSICAL_DEVICE_NAME_SIZE - 1]); - writer.print("{s} [" ++ lib.DRIVER_NAME ++ " ApeDriver]", .{name}) catch return VkError.InitializationFailed; + writer.print("{s} [" ++ lib.driver_name ++ " ApeDriver]", .{name}) catch return VkError.InitializationFailed; } @memcpy(&interface.props.device_name, &device_name); @@ -277,10 +277,10 @@ pub fn enumerateExtensionProperties(_: *const Interface, layer_name: ?[]const u8 return VkError.LayerNotPresent; } - const available = EXTENSIONS.len; + const available = extensions.len; if (p_properties) |properties| { const write_count = @min(count.*, available); - for (EXTENSIONS[0..write_count], properties[0..write_count]) |ext, *prop| { + for (extensions[0..write_count], properties[0..write_count]) |ext, *prop| { prop.* = ext; } count.* = @intCast(write_count); @@ -752,26 +752,26 @@ pub fn getImageFormatProperties( var properties: vk.ImageFormatProperties = .{ .max_extent = .{ .width = 0, .height = 0, .depth = 1 }, .max_mip_levels = 1, - .max_array_layers = lib.MAX_IMAGE_ARRAY_LAYERS, + .max_array_layers = lib.max_image_array_layers, .sample_counts = .{ .@"1_bit" = true }, .max_resource_size = std.math.maxInt(u32), }; switch (image_type) { .@"1d" => { - properties.max_mip_levels = lib.MAX_IMAGE_LEVELS_1D; - properties.max_extent.width = 1 << (lib.MAX_IMAGE_LEVELS_1D - 1); + properties.max_mip_levels = lib.max_image_levels_1d; + properties.max_extent.width = 1 << (lib.max_image_levels_1d - 1); properties.max_extent.height = 1; }, .@"2d" => { if (flags.cube_compatible_bit) { - properties.max_mip_levels = lib.MAX_IMAGE_LEVELS_CUBE; - properties.max_extent.width = 1 << (lib.MAX_IMAGE_LEVELS_CUBE - 1); - properties.max_extent.height = 1 << (lib.MAX_IMAGE_LEVELS_CUBE - 1); + properties.max_mip_levels = lib.max_image_levels_cube; + properties.max_extent.width = 1 << (lib.max_image_levels_cube - 1); + properties.max_extent.height = 1 << (lib.max_image_levels_cube - 1); } else { - properties.max_mip_levels = lib.MAX_IMAGE_LEVELS_2D; - properties.max_extent.width = 1 << (lib.MAX_IMAGE_LEVELS_2D - 1); - properties.max_extent.height = 1 << (lib.MAX_IMAGE_LEVELS_2D - 1); + properties.max_mip_levels = lib.max_image_levels_2d; + properties.max_extent.width = 1 << (lib.max_image_levels_2d - 1); + properties.max_extent.height = 1 << (lib.max_image_levels_2d - 1); const format_properties = try interface.getFormatProperties(format); const format_features = if (tiling == .linear) format_properties.linear_tiling_features else format_properties.optimal_tiling_features; @@ -781,10 +781,10 @@ pub fn getImageFormatProperties( } }, .@"3d" => { - properties.max_mip_levels = lib.MAX_IMAGE_LEVELS_3D; - properties.max_extent.width = 1 << (lib.MAX_IMAGE_LEVELS_3D - 1); - properties.max_extent.height = 1 << (lib.MAX_IMAGE_LEVELS_3D - 1); - properties.max_extent.depth = 1 << (lib.MAX_IMAGE_LEVELS_3D - 1); + properties.max_mip_levels = lib.max_image_levels_3d; + properties.max_extent.width = 1 << (lib.max_image_levels_3d - 1); + properties.max_extent.height = 1 << (lib.max_image_levels_3d - 1); + properties.max_extent.depth = 1 << (lib.max_image_levels_3d - 1); properties.max_array_layers = 1; }, else => return VkError.FormatNotSupported, diff --git a/src/software/device/Device.zig b/src/software/device/Device.zig index c807ad2..759116f 100644 --- a/src/software/device/Device.zig +++ b/src/software/device/Device.zig @@ -13,9 +13,9 @@ const Renderer = @import("Renderer.zig"); const Self = @This(); -pub const GRAPHICS_PIPELINE_STATE = 0; -pub const COMPUTE_PIPELINE_STATE = 1; -pub const MAX_DYNAMIC_DESCRIPTORS_PER_SET = 64; +pub const graphics_pipeline_state = 0; +pub const compute_pipeline_state = 1; +pub const max_dynamic_descriptors_per_set = 64; pub const ActiveOcclusionQuery = struct { pool: *base.QueryPool, @@ -24,14 +24,14 @@ pub const ActiveOcclusionQuery = struct { pub const PipelineState = struct { pipeline: ?*SoftPipeline, - sets: [base.VULKAN_MAX_DESCRIPTOR_SETS]?*SoftDescriptorSet, - dynamic_offsets: [base.VULKAN_MAX_DESCRIPTOR_SETS][MAX_DYNAMIC_DESCRIPTORS_PER_SET]u32, - push_constant_blob: [lib.PUSH_CONSTANT_SIZE]u8, + sets: [base.vulkan_max_descriptor_sets]?*SoftDescriptorSet, + dynamic_offsets: [base.vulkan_max_descriptor_sets][max_dynamic_descriptors_per_set]u32, + push_constant_blob: [lib.push_constant_size]u8, data: union { compute: struct {}, graphics: struct { index_buffer: Renderer.IndexBuffer, - vertex_buffers: [lib.MAX_VERTEX_INPUT_BINDINGS]Renderer.VertexBuffer, + vertex_buffers: [lib.max_vertex_input_bindings]Renderer.VertexBuffer, }, }, }; @@ -48,11 +48,11 @@ pub fn setup(self: *Self, device: *SoftDevice) void { for (self.pipeline_states[0..], 0..) |*state, i| { state.* = .{ .pipeline = null, - .sets = [_]?*SoftDescriptorSet{null} ** base.VULKAN_MAX_DESCRIPTOR_SETS, - .dynamic_offsets = [_][MAX_DYNAMIC_DESCRIPTORS_PER_SET]u32{[_]u32{0} ** MAX_DYNAMIC_DESCRIPTORS_PER_SET} ** base.VULKAN_MAX_DESCRIPTOR_SETS, + .sets = [_]?*SoftDescriptorSet{null} ** base.vulkan_max_descriptor_sets, + .dynamic_offsets = [_][max_dynamic_descriptors_per_set]u32{[_]u32{0} ** max_dynamic_descriptors_per_set} ** base.vulkan_max_descriptor_sets, .push_constant_blob = @splat(0), .data = switch (i) { - GRAPHICS_PIPELINE_STATE => .{ + graphics_pipeline_state => .{ .graphics = .{ // SAFETY: indexed draws bind the index buffer before the renderer reads it. .index_buffer = undefined, @@ -60,7 +60,7 @@ pub fn setup(self: *Self, device: *SoftDevice) void { .vertex_buffers = undefined, }, }, - COMPUTE_PIPELINE_STATE => .{ .compute = .{} }, + compute_pipeline_state => .{ .compute = .{} }, else => unreachable, }, }; diff --git a/src/software/device/clip.zig b/src/software/device/clip.zig index c3341ac..867567f 100644 --- a/src/software/device/clip.zig +++ b/src/software/device/clip.zig @@ -10,7 +10,7 @@ const Renderer = @import("Renderer.zig"); const Vertex = Renderer.Vertex; const VkError = base.VkError; -const INTERFACE_BLOB_PADDING = @sizeOf(F32x4); +const interface_blob_padding = @sizeOf(F32x4); const ClipPlane = enum { left, @@ -21,7 +21,7 @@ const ClipPlane = enum { far, }; -const MAX_CLIPPED_POLYGON_VERTICES = 16; +const max_clipped_polygon_vertices = 16; pub const ClippedLine = struct { v0: Vertex, @@ -29,7 +29,7 @@ pub const ClippedLine = struct { }; const ClippedPolygon = struct { - vertices: [MAX_CLIPPED_POLYGON_VERTICES]Vertex = std.mem.zeroes([MAX_CLIPPED_POLYGON_VERTICES]Vertex), + vertices: [max_clipped_polygon_vertices]Vertex = std.mem.zeroes([max_clipped_polygon_vertices]Vertex), len: usize = 0, fn append(self: *@This(), vertex: Vertex) VkError!void { @@ -146,7 +146,7 @@ fn isVertexInsidePlane(vertex: *const Vertex, plane: ClipPlane) bool { fn interpolateBlob(allocator: std.mem.Allocator, a: []const u8, b: []const u8, size: usize, t: f32) VkError![]u8 { const len = @min(size, a.len, b.len); - const result = allocator.alloc(u8, len + INTERFACE_BLOB_PADDING) catch return VkError.OutOfDeviceMemory; + const result = allocator.alloc(u8, len + interface_blob_padding) catch return VkError.OutOfDeviceMemory; @memset(result, 0); var byte_index: usize = 0; diff --git a/src/software/device/fragment.zig b/src/software/device/fragment.zig index ab77904..cf75ef8 100644 --- a/src/software/device/fragment.zig +++ b/src/software/device/fragment.zig @@ -11,8 +11,8 @@ const Renderer = @import("Renderer.zig"); const SoftPipeline = @import("../SoftPipeline.zig"); const SpvRuntimeError = spv.Runtime.RuntimeError; -const INTERFACE_BLOB_PADDING = @sizeOf(zm.F32x4); -const PROCESSED_INPUTS_STACK_CAPACITY = 4096; +const interface_blob_padding = @sizeOf(zm.F32x4); +const processed_inputs_stack_capacity = 4096; pub const InvocationResult = struct { outputs: [spv.SPIRV_MAX_OUTPUT_LOCATIONS][@sizeOf(zm.F32x4)]u8, @@ -131,7 +131,7 @@ pub fn shaderInvocation( const entry = try rt.getEntryPointByName(shader.entry); - var processed_inputs_stack: [PROCESSED_INPUTS_STACK_CAPACITY]bool = undefined; + var processed_inputs_stack: [processed_inputs_stack_capacity]bool = undefined; const processed_inputs = if (rt.results.len <= processed_inputs_stack.len) processed_inputs_stack[0..rt.results.len] else @@ -169,7 +169,7 @@ pub fn shaderInvocation( else input.size; if (input.size == 0) { - const zeroes = allocator.alloc(u8, memory_size + INTERFACE_BLOB_PADDING) catch return SpvRuntimeError.OutOfMemory; + const zeroes = allocator.alloc(u8, memory_size + interface_blob_padding) catch return SpvRuntimeError.OutOfMemory; @memset(zeroes, 0); fragment_inputs[location][component] = .{ .blob = zeroes, @@ -296,7 +296,7 @@ fn writeAggregateInputLocations( }, } - const bytes = allocator.alloc(u8, memory_size + INTERFACE_BLOB_PADDING) catch return SpvRuntimeError.OutOfMemory; + const bytes = allocator.alloc(u8, memory_size + interface_blob_padding) catch return SpvRuntimeError.OutOfMemory; defer allocator.free(bytes); @memset(bytes, 0); @@ -357,7 +357,7 @@ fn readFragmentOutput( return SpvRuntimeError.OutOfBounds; const memory_size = try element.getPlainMemorySize(); - const output = allocator.alloc(u8, memory_size + INTERFACE_BLOB_PADDING) catch return SpvRuntimeError.OutOfMemory; + const output = allocator.alloc(u8, memory_size + interface_blob_padding) catch return SpvRuntimeError.OutOfMemory; defer allocator.free(output); @memset(output, 0); @@ -370,14 +370,14 @@ fn readFragmentOutput( } const memory_size = try rt.getResultMemorySize(result_word); - if (memory_size <= INTERFACE_BLOB_PADDING) { - var output = std.mem.zeroes([INTERFACE_BLOB_PADDING]u8); + if (memory_size <= interface_blob_padding) { + var output = std.mem.zeroes([interface_blob_padding]u8); try rt.readOutput(output[0..memory_size], result_word); try copyFragmentOutputBytes(outputs, output[0..memory_size], location, component); return; } - const output = allocator.alloc(u8, memory_size + INTERFACE_BLOB_PADDING) catch return SpvRuntimeError.OutOfMemory; + const output = allocator.alloc(u8, memory_size + interface_blob_padding) catch return SpvRuntimeError.OutOfMemory; defer allocator.free(output); @memset(output, 0); diff --git a/src/software/device/vertex_dispatcher.zig b/src/software/device/vertex_dispatcher.zig index 3957f8b..c4fc5cd 100644 --- a/src/software/device/vertex_dispatcher.zig +++ b/src/software/device/vertex_dispatcher.zig @@ -12,7 +12,7 @@ const SoftPipeline = @import("../SoftPipeline.zig"); const blitter = @import("blitter.zig"); const VkError = base.VkError; -const INTERFACE_BLOB_PADDING = @sizeOf(F32x4); +const interface_blob_padding = @sizeOf(F32x4); pub const RunData = struct { allocator: std.mem.Allocator, @@ -214,7 +214,7 @@ fn readVertexOutput(data: RunData, output: *Renderer.Vertex, rt: *spv.Runtime, l output.outputs[location][component] = .{ .interpolation_type = interpolation_type, .centroid = centroid, - .blob = data.allocator.alloc(u8, memory_size + INTERFACE_BLOB_PADDING) catch return VkError.OutOfDeviceMemory, + .blob = data.allocator.alloc(u8, memory_size + interface_blob_padding) catch return VkError.OutOfDeviceMemory, .size = memory_size, }; @memset(output.outputs[location][component].?.blob, 0); @@ -264,7 +264,7 @@ fn writeVertexOutputValue( output.outputs[location][component] = .{ .interpolation_type = interpolation_type, .centroid = centroid, - .blob = data.allocator.alloc(u8, memory_size + INTERFACE_BLOB_PADDING) catch return VkError.OutOfDeviceMemory, + .blob = data.allocator.alloc(u8, memory_size + interface_blob_padding) catch return VkError.OutOfDeviceMemory, .size = memory_size, }; @memset(output.outputs[location][component].?.blob, 0); diff --git a/src/software/lib.zig b/src/software/lib.zig index 243c0d2..0d18e73 100644 --- a/src/software/lib.zig +++ b/src/software/lib.zig @@ -36,45 +36,45 @@ pub const SoftShaderModule = @import("SoftShaderModule.zig"); pub const Instance = SoftInstance; -pub const DRIVER_NAME = "Soft"; +pub const driver_name = "Soft"; -pub const VULKAN_VERSION = vk.makeApiVersion( +pub const vulkan_version = vk.makeApiVersion( 0, config.soft_vulkan_version.major, config.soft_vulkan_version.minor, config.soft_vulkan_version.patch, ); -pub const DEVICE_ID = 0x600DCAFE; -pub const PIPELINE_CACHE_UUID: [vk.UUID_SIZE]u8 = "ApeSoftCacheUUID".*; +pub const device_id = 0x600DCAFE; +pub const pipeline_cache_uuid: [vk.UUID_SIZE]u8 = "ApeSoftCacheUUID".*; /// Generic system memory. -pub const MEMORY_TYPE_GENERIC_BIT = 0; +pub const memory_type_generic_bit = 0; /// 16 bytes for 128-bit vector types. -pub const MEMORY_REQUIREMENTS_BUFFER_ALIGNMENT = 16; +pub const memory_requirements_buffer_alignment = 16; -pub const MEMORY_REQUIREMENTS_IMAGE_ALIGNMENT = 256; +pub const memory_requirements_image_alignment = 256; /// Vulkan 1.2 requires buffer offset alignment to be at most 256. -pub const MIN_TEXEL_BUFFER_ALIGNMENT = 256; +pub const min_texel_buffer_alignment = 256; /// Vulkan 1.2 requires buffer offset alignment to be at most 256. -pub const MIN_UNIFORM_BUFFER_ALIGNMENT = 256; +pub const min_uniform_buffer_alignment = 256; /// Vulkan 1.2 requires buffer offset alignment to be at most 256. -pub const MIN_STORAGE_BUFFER_ALIGNMENT = 256; +pub const min_storage_buffer_alignment = 256; -pub const MAX_VERTEX_INPUT_BINDINGS = 16; -pub const MAX_VERTEX_INPUT_ATTRIBUTES = 16; +pub const max_vertex_input_bindings = 16; +pub const max_vertex_input_attributes = 16; -pub const PUSH_CONSTANT_SIZE = 128; +pub const push_constant_size = 128; -pub const MAX_IMAGE_LEVELS_1D = 15; -pub const MAX_IMAGE_LEVELS_2D = 15; -pub const MAX_IMAGE_LEVELS_3D = 12; -pub const MAX_IMAGE_LEVELS_CUBE = 15; -pub const MAX_IMAGE_ARRAY_LAYERS = 2048; +pub const max_image_levels_1d = 15; +pub const max_image_levels_2d = 15; +pub const max_image_levels_3d = 12; +pub const max_image_levels_cube = 15; +pub const max_image_array_layers = 2048; -pub const PHYSICAL_DEVICE_DEFAULT_NAME = "Ape software device"; -pub const PHYSICAL_DEVICE_FALLBACK_HEAP_SIZE = 0x10000000; // 256MB +pub const physical_device_default_name = "Ape software device"; +pub const physical_device_fallback_heap_size = 0x10000000; // 256MB pub const std_options = base.std_options; diff --git a/src/vulkan/CommandBuffer.zig b/src/vulkan/CommandBuffer.zig index 3e4800a..d36f6cb 100644 --- a/src/vulkan/CommandBuffer.zig +++ b/src/vulkan/CommandBuffer.zig @@ -44,7 +44,7 @@ pub const DispatchTable = struct { begin: *const fn (*Self, *const vk.CommandBufferBeginInfo) VkError!void, beginQuery: *const fn (*Self, *QueryPool, u32, vk.QueryControlFlags) VkError!void, beginRenderPass: *const fn (*Self, *RenderPass, *Framebuffer, vk.Rect2D, ?[]const vk.ClearValue) VkError!void, - bindDescriptorSets: *const fn (*Self, vk.PipelineBindPoint, u32, [lib.VULKAN_MAX_DESCRIPTOR_SETS]?*DescriptorSet, []const u32) VkError!void, + bindDescriptorSets: *const fn (*Self, vk.PipelineBindPoint, u32, [lib.vulkan_max_descriptor_sets]?*DescriptorSet, []const u32) VkError!void, bindPipeline: *const fn (*Self, vk.PipelineBindPoint, *Pipeline) VkError!void, bindIndexBuffer: *const fn (*Self, *Buffer, usize, vk.IndexType) VkError!void, bindVertexBuffer: *const fn (*Self, usize, *Buffer, usize) VkError!void, @@ -192,10 +192,10 @@ pub inline fn beginQuery(self: *Self, pool: *QueryPool, query: u32, flags: vk.Qu } pub fn bindDescriptorSets(self: *Self, bind_point: vk.PipelineBindPoint, first_set: u32, sets: []const vk.DescriptorSet, dynamic_offsets: []const u32) VkError!void { - if (sets.len > lib.VULKAN_MAX_DESCRIPTOR_SETS or first_set > lib.VULKAN_MAX_DESCRIPTOR_SETS or first_set + sets.len > lib.VULKAN_MAX_DESCRIPTOR_SETS) + if (sets.len > lib.vulkan_max_descriptor_sets or first_set > lib.vulkan_max_descriptor_sets or first_set + sets.len > lib.vulkan_max_descriptor_sets) return VkError.ValidationFailed; - var inner_sets: [lib.VULKAN_MAX_DESCRIPTOR_SETS]?*DescriptorSet = @splat(null); + var inner_sets: [lib.vulkan_max_descriptor_sets]?*DescriptorSet = @splat(null); for (sets, inner_sets[0..sets.len]) |set, *inner_set| { inner_set.* = try NonDispatchable(DescriptorSet).fromHandleObject(set); } diff --git a/src/vulkan/CommandPool.zig b/src/vulkan/CommandPool.zig index 7720f7d..1a02527 100644 --- a/src/vulkan/CommandPool.zig +++ b/src/vulkan/CommandPool.zig @@ -13,7 +13,7 @@ pub const ObjectType: vk.ObjectType = .command_pool; /// Base capacity of the command buffer pool. /// Every increase of the capacity will be by this amount. -pub const BUFFER_POOL_BASE_CAPACITY = 64; +pub const buffer_pool_base_capacity = 64; owner: *Device, flags: vk.CommandPoolCreateFlags, @@ -42,7 +42,7 @@ pub fn init(device: *Device, allocator: std.mem.Allocator, info: *const vk.Comma .flags = info.flags, .queue_family_index = info.queue_family_index, .host_allocator = VulkanAllocator.from(allocator).clone(), - .buffers = std.ArrayList(*Dispatchable(CommandBuffer)).initCapacity(allocator, BUFFER_POOL_BASE_CAPACITY) catch return VkError.OutOfHostMemory, + .buffers = std.ArrayList(*Dispatchable(CommandBuffer)).initCapacity(allocator, buffer_pool_base_capacity) catch return VkError.OutOfHostMemory, .first_free_buffer_index = 0, // SAFETY: the backend assigns the vtable before returning the command pool. .vtable = undefined, @@ -54,7 +54,7 @@ pub fn allocateCommandBuffers(self: *Self, info: *const vk.CommandBufferAllocate if (self.buffers.items.len < self.first_free_buffer_index + info.command_buffer_count) { while (self.buffers.capacity < self.buffers.items.len + info.command_buffer_count) { - self.buffers.ensureUnusedCapacity(allocator, BUFFER_POOL_BASE_CAPACITY) catch return VkError.OutOfHostMemory; + self.buffers.ensureUnusedCapacity(allocator, buffer_pool_base_capacity) catch return VkError.OutOfHostMemory; } const original_len = self.buffers.items.len; errdefer { diff --git a/src/vulkan/Instance.zig b/src/vulkan/Instance.zig index b18e68d..208765b 100644 --- a/src/vulkan/Instance.zig +++ b/src/vulkan/Instance.zig @@ -13,8 +13,8 @@ const root = @import("root"); comptime { if (!builtin.is_test) { - if (!@hasDecl(root, "VULKAN_VERSION")) { - @compileError("Missing VULKAN_VERSION in module root"); + if (!@hasDecl(root, "vulkan_version")) { + @compileError("Missing vulkan_version in module root"); } } } @@ -23,7 +23,7 @@ const Self = @This(); pub const ObjectType: vk.ObjectType = .instance; /// Dummy -pub const EXTENSIONS = [_]vk.ExtensionProperties{}; +pub const extensions = [_]vk.ExtensionProperties{}; physical_devices: std.ArrayList(*Dispatchable(PhysicalDevice)), @@ -58,7 +58,7 @@ pub fn validateCreateInfo(info: *const vk.InstanceCreateInfo) VkError!void { const supported: vk.Version = if (comptime builtin.is_test) vk.API_VERSION_1_0 else - @bitCast(root.VULKAN_VERSION); + @bitCast(root.vulkan_version); if (requested.variant != 0 or requested.major > supported.major or (requested.major == supported.major and requested.minor > supported.minor)) { return VkError.IncompatibleDriver; } @@ -78,7 +78,7 @@ pub fn validateCreateInfo(info: *const vk.InstanceCreateInfo) VkError!void { const supported_extensions = if (comptime !@hasDecl(root, "Instance")) &[_]vk.ExtensionProperties{} else - root.Instance.EXTENSIONS[0..]; + root.Instance.extensions[0..]; for (0..info.enabled_extension_count) |i| { const name = utils.boundedName(names[i], vk.MAX_EXTENSION_NAME_SIZE) orelse return VkError.ExtensionNotPresent; @@ -124,11 +124,11 @@ pub fn enumerateExtensionProperties(layer_name: ?[]const u8, count: *u32, p_prop return VkError.LayerNotPresent; } - if (comptime !builtin.is_test and @hasDecl(root.Instance, "EXTENSIONS")) { - const available = root.Instance.EXTENSIONS.len; + if (comptime !builtin.is_test and @hasDecl(root.Instance, "extensions")) { + const available = root.Instance.extensions.len; if (p_properties) |properties| { const write_count = @min(count.*, available); - for (root.Instance.EXTENSIONS[0..write_count], properties[0..write_count]) |ext, *prop| { + for (root.Instance.extensions[0..write_count], properties[0..write_count]) |ext, *prop| { prop.* = ext; } count.* = @intCast(write_count); @@ -145,7 +145,7 @@ pub fn enumerateVersion(version: *u32) VkError!void { if (comptime builtin.is_test) { version.* = @bitCast(vk.makeApiVersion(0, 1, 0, 0)); } else { - version.* = @bitCast(root.VULKAN_VERSION); + version.* = @bitCast(root.vulkan_version); } } diff --git a/src/vulkan/PhysicalDevice.zig b/src/vulkan/PhysicalDevice.zig index ac8446b..1a5f20a 100644 --- a/src/vulkan/PhysicalDevice.zig +++ b/src/vulkan/PhysicalDevice.zig @@ -40,7 +40,7 @@ pub fn init(allocator: std.mem.Allocator, instance: *Instance) VkError!Self { .props = .{ .api_version = 0, .driver_version = 0, - .vendor_id = root.VULKAN_VENDOR_ID, + .vendor_id = root.vulkan_vendor_id, .device_id = 0, .device_type = .other, .device_name = @as([vk.MAX_PHYSICAL_DEVICE_NAME_SIZE]u8, @splat(0)), diff --git a/src/vulkan/PipelineCache.zig b/src/vulkan/PipelineCache.zig index 366c76a..15ad7b7 100644 --- a/src/vulkan/PipelineCache.zig +++ b/src/vulkan/PipelineCache.zig @@ -94,7 +94,7 @@ fn makeHeader(device: *const Device) Header { return .{ .header_size = @sizeOf(Header), .header_version = .one, - .vendor_id = @intCast(root.VULKAN_VENDOR_ID), + .vendor_id = @intCast(root.vulkan_vendor_id), .device_id = device.physical_device.props.device_id, .pipeline_cache_uuid = device.physical_device.props.pipeline_cache_uuid, }; diff --git a/src/vulkan/PipelineLayout.zig b/src/vulkan/PipelineLayout.zig index 4063fec..4a141e2 100644 --- a/src/vulkan/PipelineLayout.zig +++ b/src/vulkan/PipelineLayout.zig @@ -16,12 +16,12 @@ owner: *Device, set_count: usize, -set_layouts: [lib.VULKAN_MAX_DESCRIPTOR_SETS]?*DescriptorSetLayout, +set_layouts: [lib.vulkan_max_descriptor_sets]?*DescriptorSetLayout, -dynamic_descriptor_offsets: [lib.VULKAN_MAX_DESCRIPTOR_SETS]usize, +dynamic_descriptor_offsets: [lib.vulkan_max_descriptor_sets]usize, push_ranges_count: usize, -push_ranges: [lib.VULKAN_MAX_PUSH_CONSTANT_RANGES]vk.PushConstantRange, +push_ranges: [lib.vulkan_max_push_constant_ranges]vk.PushConstantRange, ref_count: std.atomic.Value(usize), @@ -36,8 +36,8 @@ pub fn init(device: *Device, allocator: std.mem.Allocator, info: *const vk.Pipel var self: Self = .{ .owner = device, .set_count = info.set_layout_count, - .set_layouts = [_]?*DescriptorSetLayout{null} ** lib.VULKAN_MAX_DESCRIPTOR_SETS, - .dynamic_descriptor_offsets = [_]usize{0} ** lib.VULKAN_MAX_DESCRIPTOR_SETS, + .set_layouts = [_]?*DescriptorSetLayout{null} ** lib.vulkan_max_descriptor_sets, + .dynamic_descriptor_offsets = [_]usize{0} ** lib.vulkan_max_descriptor_sets, .push_ranges_count = info.push_constant_range_count, .push_ranges = @splat(std.mem.zeroes(vk.PushConstantRange)), .ref_count = std.atomic.Value(usize).init(1), diff --git a/src/vulkan/lib.zig b/src/vulkan/lib.zig index 390157f..52732ef 100644 --- a/src/vulkan/lib.zig +++ b/src/vulkan/lib.zig @@ -52,11 +52,11 @@ pub const ShaderModule = @import("ShaderModule.zig"); pub const SurfaceKHR = @import("wsi/SurfaceKHR.zig"); pub const SwapchainKHR = @import("wsi/SwapchainKHR.zig"); -pub const VULKAN_VENDOR_ID: i32 = @intFromEnum(vk.VendorId.ape); +pub const vulkan_vendor_id: i32 = @intFromEnum(vk.VendorId.ape); /// Default driver name -pub const DRIVER_NAME = "Unnamed Ape Driver"; -pub const DRIVER_VERSION = vk.makeApiVersion( +pub const driver_name = "Unnamed Ape Driver"; +pub const driver_version = vk.makeApiVersion( 0, config.driver_version.major, config.driver_version.minor, @@ -64,10 +64,10 @@ pub const DRIVER_VERSION = vk.makeApiVersion( ); /// Dummy fallback -pub const VULKAN_VERSION = vk.makeApiVersion(0, 1, 0, 0); +pub const vulkan_version = vk.makeApiVersion(0, 1, 0, 0); /// Maximum number of descriptor sets per pipeline -pub const VULKAN_MAX_DESCRIPTOR_SETS = 8; +pub const vulkan_max_descriptor_sets = 8; /// The number of push constant ranges is effectively bounded /// by the number of possible shader stages. Not the number of stages that can @@ -80,7 +80,7 @@ pub const VULKAN_MAX_DESCRIPTOR_SETS = 8; /// - VK_SHADER_STAGE_GEOMETRY_BIT /// - VK_SHADER_STAGE_FRAGMENT_BIT /// - VK_SHADER_STAGE_COMPUTE_BIT -pub const VULKAN_MAX_PUSH_CONSTANT_RANGES = 6; +pub const vulkan_max_push_constant_ranges = 6; pub const std_options: std.Options = .{ .log_level = .debug, diff --git a/src/vulkan/logger.zig b/src/vulkan/logger.zig index 6defb19..29c82ae 100644 --- a/src/vulkan/logger.zig +++ b/src/vulkan/logger.zig @@ -5,8 +5,8 @@ const lib = @import("lib.zig"); comptime { if (!builtin.is_test) { - if (!@hasDecl(root, "DRIVER_NAME")) { - @compileError("Missing DRIVER_NAME in module root"); + if (!@hasDecl(root, "driver_name")) { + @compileError("Missing driver_name in module root"); } } } @@ -90,7 +90,7 @@ pub fn log(comptime level: std.log.Level, comptime scope: @EnumLiteral(), compti writer.writeAll("[ApeDriver") catch continue; if (!builtin.is_test) { term.setColor(.cyan) catch @panic("Caught an error while handling an error"); - writer.print(" {s} ", .{root.DRIVER_NAME}) catch continue; + writer.print(" {s} ", .{root.driver_name}) catch continue; } term.setColor(.yellow) catch @panic("Caught an error while handling an error"); writer.print("{d}:{d}:{d}.{d:0>3}.{d:0>3}", .{ now_hour, now_min, now_sec, now_ms, now_us }) catch continue; diff --git a/src/vulkan/wsi/WaylandSurfaceKHR.zig b/src/vulkan/wsi/WaylandSurfaceKHR.zig index e5627d1..ce0ed76 100644 --- a/src/vulkan/wsi/WaylandSurfaceKHR.zig +++ b/src/vulkan/wsi/WaylandSurfaceKHR.zig @@ -111,7 +111,7 @@ pub fn attachImage(interface: *Interface, allocator: std.mem.Allocator, image: * const pool = wayland.wl_shm_create_pool(shm, fd, @intCast(size)) orelse return VkError.Unknown; defer wayland.wl_shm_pool_destroy(pool); - const buffer = wayland.wl_shm_pool_create_buffer(pool, 0, @intCast(width), @intCast(height), @intCast(stride), wayland.WL_SHM_FORMAT_ARGB8888) orelse return VkError.Unknown; + const buffer = wayland.wl_shm_pool_create_buffer(pool, 0, @intCast(width), @intCast(height), @intCast(stride), wayland.wl_shm_format_argb8888) orelse return VkError.Unknown; errdefer wayland.wl_buffer_destroy(buffer); wl_image.* = .{ diff --git a/src/vulkan/wsi/clients/wayland.zig b/src/vulkan/wsi/clients/wayland.zig index ae78351..c300eb8 100644 --- a/src/vulkan/wsi/clients/wayland.zig +++ b/src/vulkan/wsi/clients/wayland.zig @@ -11,16 +11,16 @@ pub const wl_registry_listener = extern struct { global_remove: ?*const fn (data: ?*anyopaque, wl_registry: ?*wl_registry, name: u32) callconv(.c) void = null, }; -pub const WL_BUFFER_DESTROY: c_int = 0; -pub const WL_DISPLAY_GET_REGISTRY: c_int = 1; -pub const WL_REGISTRY_BIND: c_int = 0; -pub const WL_SHM_CREATE_POOL: c_int = 0; -pub const WL_SHM_FORMAT_ARGB8888: c_int = 0; -pub const WL_SHM_POOL_CREATE_BUFFER: c_int = 0; -pub const WL_SHM_POOL_DESTROY: c_int = 1; -pub const WL_SURFACE_ATTACH: c_int = 1; -pub const WL_SURFACE_COMMIT: c_int = 6; -pub const WL_SURFACE_DAMAGE: c_int = 2; +pub const wl_buffer_destroy_opcode: c_int = 0; +pub const wl_display_get_registry_opcode: c_int = 1; +pub const wl_registry_bind_opcode: c_int = 0; +pub const wl_shm_create_pool_opcode: c_int = 0; +pub const wl_shm_format_argb8888: c_int = 0; +pub const wl_shm_pool_create_buffer_opcode: c_int = 0; +pub const wl_shm_pool_destroy_opcode: c_int = 1; +pub const wl_surface_attach_opcode: c_int = 1; +pub const wl_surface_commit_opcode: c_int = 6; +pub const wl_surface_damage_opcode: c_int = 2; pub const wl_buffer = opaque {}; pub const wl_callback = opaque {}; @@ -104,7 +104,7 @@ pub fn unload() void { pub fn wl_registry_bind(registry: *wl_registry, name: u32, interface: *const wl_interface, version: u32) ?*wl_proxy { return wl_proxy_marshal_flags( @ptrCast(@alignCast(registry)), - WL_REGISTRY_BIND, + wl_registry_bind_opcode, interface, version, 0, @@ -118,7 +118,7 @@ pub fn wl_registry_bind(registry: *wl_registry, name: u32, interface: *const wl_ pub fn wl_display_get_registry(display: *wl_display) ?*wl_registry { return @ptrCast(@alignCast(wl_proxy_marshal_flags( @ptrCast(@alignCast(display)), - WL_DISPLAY_GET_REGISTRY, + wl_display_get_registry_opcode, wl_registry_interface, wl_proxy_get_version(@ptrCast(@alignCast(display))), 0, @@ -133,7 +133,7 @@ pub fn wl_registry_add_listener(registry: *wl_registry, listener: *const wl_regi pub fn wl_shm_create_pool(shm: *wl_shm, fd: i32, size: i32) callconv(.c) ?*wl_shm_pool { return @ptrCast(@alignCast(wl_proxy_marshal_flags( @ptrCast(@alignCast(shm)), - WL_SHM_CREATE_POOL, + wl_shm_create_pool_opcode, wl_shm_pool_interface, wl_proxy_get_version(@ptrCast(@alignCast(shm))), 0, @@ -146,7 +146,7 @@ pub fn wl_shm_create_pool(shm: *wl_shm, fd: i32, size: i32) callconv(.c) ?*wl_sh pub fn wl_shm_pool_destroy(shm_pool: *wl_shm_pool) void { _ = wl_proxy_marshal_flags( @ptrCast(@alignCast(shm_pool)), - WL_SHM_POOL_DESTROY, + wl_shm_pool_destroy_opcode, null, wl_proxy_get_version(@ptrCast(@alignCast(shm_pool))), @bitCast(@as(c_int, @as(c_int, 1) << @intCast(@as(c_int, 0)))), @@ -156,7 +156,7 @@ pub fn wl_shm_pool_destroy(shm_pool: *wl_shm_pool) void { pub fn wl_shm_pool_create_buffer(shm_pool: *wl_shm_pool, offset: i32, width: i32, height: i32, stride: i32, format: u32) ?*wl_buffer { return @ptrCast(@alignCast(wl_proxy_marshal_flags( @ptrCast(@alignCast(shm_pool)), - WL_SHM_POOL_CREATE_BUFFER, + wl_shm_pool_create_buffer_opcode, wl_buffer_interface, wl_proxy_get_version(@ptrCast(@alignCast(shm_pool))), 0, @@ -172,7 +172,7 @@ pub fn wl_shm_pool_create_buffer(shm_pool: *wl_shm_pool, offset: i32, width: i32 pub fn wl_buffer_destroy(buffer: *wl_buffer) void { _ = wl_proxy_marshal_flags( @ptrCast(@alignCast(buffer)), - WL_BUFFER_DESTROY, + wl_buffer_destroy_opcode, null, wl_proxy_get_version(@ptrCast(@alignCast(buffer))), @bitCast(@as(c_int, @as(c_int, 1) << @intCast(@as(c_int, 0)))), @@ -182,7 +182,7 @@ pub fn wl_buffer_destroy(buffer: *wl_buffer) void { pub fn wl_surface_attach(surface: *wl_surface, buffer: *wl_buffer, x: i32, y: i32) void { _ = wl_proxy_marshal_flags( @ptrCast(@alignCast(surface)), - WL_SURFACE_ATTACH, + wl_surface_attach_opcode, null, wl_proxy_get_version(@ptrCast(@alignCast(surface))), 0, @@ -195,7 +195,7 @@ pub fn wl_surface_attach(surface: *wl_surface, buffer: *wl_buffer, x: i32, y: i3 pub fn wl_surface_damage(surface: *wl_surface, x: i32, y: i32, width: i32, height: i32) void { _ = wl_proxy_marshal_flags( @ptrCast(@alignCast(surface)), - WL_SURFACE_DAMAGE, + wl_surface_damage_opcode, null, wl_proxy_get_version(@ptrCast(@alignCast(surface))), 0, @@ -209,7 +209,7 @@ pub fn wl_surface_damage(surface: *wl_surface, x: i32, y: i32, width: i32, heigh pub fn wl_surface_commit(surface: *wl_surface) void { _ = wl_proxy_marshal_flags( @ptrCast(@alignCast(surface)), - WL_SURFACE_COMMIT, + wl_surface_commit_opcode, null, wl_proxy_get_version(@ptrCast(@alignCast(surface))), 0, diff --git a/test/test_runner.zig b/test/test_runner.zig index 3925e34..9f53859 100644 --- a/test/test_runner.zig +++ b/test/test_runner.zig @@ -5,7 +5,7 @@ const builtin = @import("builtin"); const Allocator = std.mem.Allocator; -const BORDER = "=" ** 80; +const border = "=" ** 80; // use in custom panic handler var current_test: ?[]const u8 = null; @@ -64,7 +64,7 @@ pub fn main() !void { if (std.testing.allocator_instance.deinit() == .leak) { leak += 1; - Printer.status(.fail, "\n{s}\n\"{s}\" - Memory Leak\n{s}\n", .{ BORDER, friendly_name, BORDER }); + Printer.status(.fail, "\n{s}\n\"{s}\" - Memory Leak\n{s}\n", .{ border, friendly_name, border }); } if (result) |_| { @@ -77,7 +77,7 @@ pub fn main() !void { else => { status = .fail; fail += 1; - Printer.status(.fail, "\n{s}\n\"{s}\" - {s}\n{s}\n", .{ BORDER, friendly_name, @errorName(err), BORDER }); + Printer.status(.fail, "\n{s}\n\"{s}\" - {s}\n{s}\n", .{ border, friendly_name, @errorName(err), border }); if (@errorReturnTrace()) |trace| { std.debug.dumpErrorReturnTrace(trace); } @@ -214,7 +214,7 @@ const SlowTracker = struct { pub const panic = std.debug.FullPanic(struct { pub fn panicFn(msg: []const u8, first_trace_addr: ?usize) noreturn { if (current_test) |ct| { - Printer.fmt("\x1b[31m{s}\npanic running \"{s}\"\n{s}\x1b[0m\n", .{ BORDER, ct, BORDER }); + Printer.fmt("\x1b[31m{s}\npanic running \"{s}\"\n{s}\x1b[0m\n", .{ border, ct, border }); } std.debug.defaultPanic(msg, first_trace_addr); }