Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions .github/workflows/webgpu.yml
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,8 @@ jobs:
cache-dependency-path: backend/accelerated/webgpu/web/package-lock.json
- run: npm ci
- run: npx playwright install chromium
- name: wasm runtime unit tests
run: GOOS=js GOARCH=wasm go test -exec="$(go env GOROOT)/lib/wasm/go_js_wasm_exec" ../internal/wasmruntime
- name: generated shaders and accelerators are up to date
run: |
(cd ../internal/generator && go run .)
Expand Down
14 changes: 11 additions & 3 deletions backend/accelerated/webgpu/internal/wasmruntime/runtime.go
Original file line number Diff line number Diff line change
Expand Up @@ -327,13 +327,21 @@ func bytesArg(args []js.Value, index int) ([]byte, error) {
if len(args) <= index {
return nil, fmt.Errorf("missing bytes argument")
}
n := args[index].Get("byteLength")
v := args[index]
if v.Type() != js.TypeObject || !v.InstanceOf(js.Global().Get("Uint8Array")) {
return nil, fmt.Errorf("expected Uint8Array")
}
n := v.Get("byteLength")
if n.Type() != js.TypeNumber {
return nil, fmt.Errorf("expected Uint8Array")
}
out := make([]byte, n.Int())
size := n.Int()
if size < 0 {
return nil, fmt.Errorf("expected Uint8Array")
}
out := make([]byte, size)
if len(out) > 0 {
js.CopyBytesToGo(out, args[index])
js.CopyBytesToGo(out, v)
}
return out, nil
}
Expand Down
70 changes: 70 additions & 0 deletions backend/accelerated/webgpu/internal/wasmruntime/runtime_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,70 @@
//go:build js && wasm

package wasmruntime

import (
"bytes"
"syscall/js"
"testing"
)

func TestBytesArgRejectsNonUint8Array(t *testing.T) {
objectWithByteLength := js.Global().Get("Object").New()
objectWithByteLength.Set("byteLength", 3)

tests := []struct {
name string
value js.Value
}{
{name: "number", value: js.ValueOf(123)},
{name: "string", value: js.ValueOf("bytes")},
{name: "boolean", value: js.ValueOf(true)},
{name: "null", value: js.Null()},
{name: "undefined", value: js.Undefined()},
{name: "object with byteLength", value: objectWithByteLength},
{name: "ArrayBuffer", value: js.Global().Get("ArrayBuffer").New(3)},
{name: "Uint8ClampedArray", value: js.Global().Get("Uint8ClampedArray").New(3)},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := bytesArg([]js.Value{tt.value}, 0)
if err == nil || err.Error() != "expected Uint8Array" {
t.Fatalf("bytesArg() = (%v, %v), want (nil, expected Uint8Array error)", got, err)
}
if got != nil {
t.Fatalf("bytesArg() returned %v, want nil", got)
}
})
}
}

func TestBytesArgCopiesUint8Array(t *testing.T) {
want := []byte{0, 1, 127, 255}
value := js.Global().Get("Uint8Array").New(len(want))
js.CopyBytesToJS(value, want)

got, err := bytesArg([]js.Value{value}, 0)
if err != nil {
t.Fatalf("bytesArg() error = %v", err)
}
if !bytes.Equal(got, want) {
t.Fatalf("bytesArg() = %v, want %v", got, want)
}
}

func TestBytesArgAcceptsEmptyUint8Array(t *testing.T) {
got, err := bytesArg([]js.Value{js.Global().Get("Uint8Array").New(0)}, 0)
if err != nil {
t.Fatalf("bytesArg() error = %v", err)
}
if len(got) != 0 {
t.Fatalf("bytesArg() = %v, want empty bytes", got)
}
}

func TestBytesArgRequiresArgument(t *testing.T) {
got, err := bytesArg(nil, 0)
if err == nil || err.Error() != "missing bytes argument" {
t.Fatalf("bytesArg() = (%v, %v), want (nil, missing bytes argument error)", got, err)
}
}