Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
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
26 changes: 7 additions & 19 deletions wgpu/adapter.go
Original file line number Diff line number Diff line change
Expand Up @@ -79,21 +79,10 @@ var (
adapterCallbackOnce sync.Once
)

// adapterCallbackHandler is the Go function called by C code via ffi.NewCallback.
// Windows x64 ABI: args in RCX, RDX, R8, R9, then stack.
// Signature: void(status uint32, adapter uintptr, message *StringView, userdata1 uintptr, userdata2 uintptr)
// Note: On Windows x64 ABI, structs > 8 bytes are passed by pointer.
// goffi v0.2.1+ requires all args to be uintptr and exactly one uintptr return.
func adapterCallbackHandler(status uintptr, adapter uintptr, message uintptr, userdata1, userdata2 uintptr) uintptr {
// Extract message string (message is pointer to StringView on Windows)
var msg string
if message != 0 {
sv := (*StringView)(ptrFromUintptr(message))
if sv.Data != 0 && sv.Length > 0 && sv.Length < 1<<20 {
msg = unsafe.String((*byte)(ptrFromUintptr(sv.Data)), int(sv.Length))
}
}

// handleAdapterCallback completes a request after the platform callback entry
// normalizes the ABI-specific WGPUStringView representation.
// userdata2 is reserved by WebGPU and discarded by the platform entry.
func handleAdapterCallback(status uintptr, adapter uintptr, message StringView, userdata1 uintptr) uintptr {
// Find and complete the request
adapterRequestsMu.Lock()
req, ok := adapterRequests[userdata1]
Expand All @@ -108,16 +97,15 @@ func adapterCallbackHandler(status uintptr, adapter uintptr, message uintptr, us
trackResource(adapter, "Adapter")
req.adapter = &Adapter{handle: adapter}
}
req.message = msg
req.message = stringViewToString(message)
close(req.done)
}
return 0
}

// initAdapterCallback creates the C callback function pointer using goffi.
// goffi v0.2.1+ properly handles Windows x64 calling convention.
// initAdapterCallback creates the platform-correct C callback function pointer.
func initAdapterCallback() {
adapterCallbackPtr = ffi.NewCallback(adapterCallbackHandler)
adapterCallbackPtr = ffi.NewCallback(adapterCallbackEntry)
}

// RequestAdapter requests a GPU adapter from the instance.
Expand Down
21 changes: 6 additions & 15 deletions wgpu/buffer.go
Original file line number Diff line number Diff line change
Expand Up @@ -67,18 +67,9 @@ var (
mapCallbackOnce sync.Once
)

// mapCallbackHandler is the Go function called by C code via ffi.NewCallback.
// Signature: void(status uint32, message *StringView, userdata1 uintptr, userdata2 uintptr)
func mapCallbackHandler(status uintptr, message uintptr, userdata1, userdata2 uintptr) uintptr {
// Extract message string
var msg string
if message != 0 {
sv := (*StringView)(ptrFromUintptr(message))
if sv.Data != 0 && sv.Length > 0 && sv.Length < 1<<20 {
msg = unsafe.String((*byte)(ptrFromUintptr(sv.Data)), int(sv.Length))
}
}

// handleMapCallback completes a request after the platform callback entry
// normalizes the ABI-specific WGPUStringView representation.
func handleMapCallback(status uintptr, message StringView, userdata1 uintptr) uintptr {
// Find and complete the request
mapRequestsMu.Lock()
req, ok := mapRequests[userdata1]
Expand All @@ -89,15 +80,15 @@ func mapCallbackHandler(status uintptr, message uintptr, userdata1, userdata2 ui

if ok && req != nil {
req.status = MapAsyncStatus(status)
req.message = msg
req.message = stringViewToString(message)
close(req.done)
}
return 0
}

// initMapCallback creates the C callback function pointer using goffi.
// initMapCallback creates the platform-correct C callback function pointer.
func initMapCallback() {
mapCallbackPtr = ffi.NewCallback(mapCallbackHandler)
mapCallbackPtr = ffi.NewCallback(mapCallbackEntry)
}

// BufferDescriptor describes a GPU buffer to create.
Expand Down
28 changes: 28 additions & 0 deletions wgpu/callback_flat.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,28 @@
//go:build ((linux || darwin || freebsd) && (amd64 || arm64)) || (windows && arm64)

package wgpu

// Callback entry implementations support amd64 and arm64 on Linux, macOS,
// FreeBSD, and Windows. Windows amd64 uses callback_windows_amd64.go;
// wgpu-native does not support other architectures.
//
// Unix amd64/arm64 and Windows ARM64 ABIs pass the two-word WGPUStringView
// callback argument by value in integer registers. goffi callbacks expose
// those words as separate uintptr arguments, so each entry reconstructs the
// view before invoking shared logic.

func adapterCallbackEntry(status, adapter, messageData, messageLength, userdata1, _ uintptr) uintptr {
return handleAdapterCallback(status, adapter, StringView{Data: messageData, Length: messageLength}, userdata1)
}

func deviceCallbackEntry(status, device, messageData, messageLength, userdata1, _ uintptr) uintptr {
return handleDeviceCallback(status, device, StringView{Data: messageData, Length: messageLength}, userdata1)
}

func mapCallbackEntry(status, messageData, messageLength, userdata1, _ uintptr) uintptr {
return handleMapCallback(status, StringView{Data: messageData, Length: messageLength}, userdata1)
}

func errorScopeCallbackEntry(status, errType, messageData, messageLength, userdata1, _ uintptr) uintptr {
return handleErrorScopeCallback(status, errType, StringView{Data: messageData, Length: messageLength}, userdata1)
}
122 changes: 122 additions & 0 deletions wgpu/callback_flat_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,122 @@
//go:build ((linux || darwin || freebsd) && (amd64 || arm64)) || (windows && arm64)

package wgpu

import (
"testing"
"unsafe"
)

func TestABICallbackEntriesPreserveStringViewAndUserdata(t *testing.T) {
message := []byte("callback message")
messageData := uintptr(unsafe.Pointer(&message[0]))
messageLength := uintptr(len(message))

t.Run("adapter", func(t *testing.T) {
const requestID = uintptr(101)
req := &adapterRequest{done: make(chan struct{})}
adapterRequestsMu.Lock()
adapterRequests[requestID] = req
adapterRequestsMu.Unlock()
t.Cleanup(func() {
adapterRequestsMu.Lock()
delete(adapterRequests, requestID)
adapterRequestsMu.Unlock()
})

adapterCallbackEntry(7, 0, messageData, messageLength, requestID, 0)

assertCallbackCompleted(t, req.done, req.message)
if req.status != RequestAdapterStatus(7) {
t.Fatalf("status = %d, want 7", req.status)
}
})

t.Run("device", func(t *testing.T) {
const requestID = uintptr(102)
req := &deviceRequest{done: make(chan struct{})}
deviceRequestsMu.Lock()
deviceRequests[requestID] = req
deviceRequestsMu.Unlock()
t.Cleanup(func() {
deviceRequestsMu.Lock()
delete(deviceRequests, requestID)
deviceRequestsMu.Unlock()
})

deviceCallbackEntry(8, 0, messageData, messageLength, requestID, 0)

assertCallbackCompleted(t, req.done, req.message)
if req.status != RequestDeviceStatus(8) {
t.Fatalf("status = %d, want 8", req.status)
}
})

t.Run("buffer map", func(t *testing.T) {
const requestID = uintptr(103)
req := &mapRequest{done: make(chan struct{})}
mapRequestsMu.Lock()
mapRequests[requestID] = req
mapRequestsMu.Unlock()
t.Cleanup(func() {
mapRequestsMu.Lock()
delete(mapRequests, requestID)
mapRequestsMu.Unlock()
})

mapCallbackEntry(9, messageData, messageLength, requestID, 0)

assertCallbackCompleted(t, req.done, req.message)
if req.status != MapAsyncStatus(9) {
t.Fatalf("status = %d, want 9", req.status)
}
})

t.Run("error scope", func(t *testing.T) {
const requestID = uintptr(104)
result := &errorScopeResult{done: make(chan struct{})}
errorScopeResultsMu.Lock()
errorScopeResults[requestID] = result
errorScopeResultsMu.Unlock()
t.Cleanup(func() {
errorScopeResultsMu.Lock()
delete(errorScopeResults, requestID)
errorScopeResultsMu.Unlock()
})

errorScopeCallbackEntry(10, 11, messageData, messageLength, requestID, 0)

assertCallbackCompleted(t, result.done, result.message)
if result.status != PopErrorScopeStatus(10) {
t.Fatalf("status = %d, want 10", result.status)
}
if result.errType != ErrorType(11) {
t.Fatalf("error type = %d, want 11", result.errType)
}
})
}

func TestABICallbackEntriesHandleMessageEdges(t *testing.T) {
t.Run("null and empty", func(t *testing.T) {
const requestID = uintptr(105)
req := registerTestAdapterRequest(t, requestID)

adapterCallbackEntry(0, 0, 0, 0, requestID, 0)

assertCallbackMessage(t, req.done, req.message, "")
})

t.Run("zero length with non-null data", func(t *testing.T) {
const requestID = uintptr(106)
req := registerTestAdapterRequest(t, requestID)
message := []byte("ignored")

adapterCallbackEntry(0, 0, uintptr(unsafe.Pointer(&message[0])), 0, requestID, 0)

assertCallbackMessage(t, req.done, req.message, "")
})

t.Run("unknown userdata", func(t *testing.T) {
adapterCallbackEntry(0, 0, 0, 0, ^uintptr(0), 0)
})
}
60 changes: 60 additions & 0 deletions wgpu/callback_test_helpers_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,60 @@
package wgpu

import "testing"

func TestABICallbackInitializers(t *testing.T) {
tests := []struct {
name string
init func()
target *uintptr
}{
{name: "adapter", init: initAdapterCallback, target: &adapterCallbackPtr},
{name: "device", init: initDeviceCallback, target: &deviceCallbackPtr},
{name: "buffer map", init: initMapCallback, target: &mapCallbackPtr},
{name: "error scope", init: initErrorScopeCallback, target: &errorScopeCallbackPtr},
}

for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
original := *test.target
t.Cleanup(func() {
*test.target = original
})

test.init()
if *test.target == 0 {
t.Fatal("callback pointer is zero")
}
})
}
}

func registerTestAdapterRequest(t *testing.T, requestID uintptr) *adapterRequest {
t.Helper()
req := &adapterRequest{done: make(chan struct{})}
adapterRequestsMu.Lock()
adapterRequests[requestID] = req
adapterRequestsMu.Unlock()
t.Cleanup(func() {
adapterRequestsMu.Lock()
delete(adapterRequests, requestID)
adapterRequestsMu.Unlock()
})
return req
}

func assertCallbackCompleted(t *testing.T, done <-chan struct{}, message string) {
assertCallbackMessage(t, done, message, "callback message")
}

func assertCallbackMessage(t *testing.T, done <-chan struct{}, message, want string) {
t.Helper()
select {
case <-done:
default:
t.Fatal("callback did not complete the registered request")
}
if message != want {
t.Fatalf("message = %q, want %q", message, want)
}
}
30 changes: 30 additions & 0 deletions wgpu/callback_windows_amd64.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,30 @@
//go:build windows && amd64

package wgpu

// Windows x64 passes a WGPUStringView callback argument indirectly because
// the aggregate is larger than one register. Normalize that pointer into the
// same value form used by the shared callback logic.

func adapterCallbackEntry(status, adapter, message, userdata1, _ uintptr) uintptr {
return handleAdapterCallback(status, adapter, callbackStringView(message), userdata1)
}

func deviceCallbackEntry(status, device, message, userdata1, _ uintptr) uintptr {
return handleDeviceCallback(status, device, callbackStringView(message), userdata1)
}

func mapCallbackEntry(status, message, userdata1, _ uintptr) uintptr {
return handleMapCallback(status, callbackStringView(message), userdata1)
}

func errorScopeCallbackEntry(status, errType, message, userdata1, _ uintptr) uintptr {
return handleErrorScopeCallback(status, errType, callbackStringView(message), userdata1)
}

func callbackStringView(message uintptr) StringView {
if message == 0 {
return StringView{}
}
return *(*StringView)(ptrFromUintptr(message))
}
40 changes: 40 additions & 0 deletions wgpu/callback_windows_amd64_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,40 @@
//go:build windows && amd64

package wgpu

import (
"testing"
"unsafe"
)

func TestABICallbackStringViewWindowsAMD64(t *testing.T) {
if got := callbackStringView(0); got != (StringView{}) {
t.Fatalf("callbackStringView(0) = %#v, want empty", got)
}

message := []byte("callback message")
want := StringView{
Data: uintptr(unsafe.Pointer(&message[0])),
Length: uintptr(len(message)),
}
if got := callbackStringView(uintptr(unsafe.Pointer(&want))); got != want {
t.Fatalf("callbackStringView(valid) = %#v, want %#v", got, want)
}
}

func TestABIAdapterCallbackEntryWindowsAMD64(t *testing.T) {
const requestID = uintptr(201)
req := registerTestAdapterRequest(t, requestID)
message := []byte("callback message")
view := StringView{
Data: uintptr(unsafe.Pointer(&message[0])),
Length: uintptr(len(message)),
}

adapterCallbackEntry(7, 0, uintptr(unsafe.Pointer(&view)), requestID, 0)

assertCallbackCompleted(t, req.done, req.message)
if req.status != RequestAdapterStatus(7) {
t.Fatalf("status = %d, want 7", req.status)
}
}
Loading
Loading