diff --git a/func.go b/func.go index 632b3470..8a3e9871 100644 --- a/func.go +++ b/func.go @@ -48,8 +48,8 @@ func RegisterLibFunc(fptr any, handle uintptr, name string) { // // These conversions describe how a Go type in the fptr will be used to call // the C function. It is important to note that there is no way to verify that fptr -// matches the C function. This also holds true for struct types where the padding -// needs to be ensured to match that of C; RegisterFunc does not verify this. +// matches the C function. This also holds true for struct types, whose memory layout +// must match the C one; RegisterFunc does not verify this. // // # Type Conversions (Go <=> C) // @@ -101,9 +101,12 @@ func RegisterLibFunc(fptr any, handle uintptr, name string) { // // # Structs // -// Purego can handle the most common structs that have fields of builtin types like int8, uint16, float32, etc. However, -// it does not support aligning fields properly. It is therefore the responsibility of the caller to ensure -// that all padding is added to the Go struct to match the C one. See `BoolStructFn` in struct_test.go for an example. +// Purego can handle the most common structs that have fields of builtin types like int8, uint16, float32, etc. +// Each field is placed at the offset it has in the Go struct's memory image, so the padding the Go compiler +// inserts is preserved and explicit padding fields are not needed. The Go struct must still be declared with the +// same fields, in the same order, as the C one, and should include a field named _ of type +// [structs.HostLayout] to guarantee that layout. +// Purego does not verify that the two match. // // On Apple ARM64 platforms (macOS and iOS), purego handles proper alignment of struct arguments // when passing them on the stack, following the C ABI's byte-level packing rules. diff --git a/struct_amd64.go b/struct_amd64.go index 85731fa3..4f2a0ba8 100644 --- a/struct_amd64.go +++ b/struct_amd64.go @@ -4,7 +4,6 @@ package purego import ( - "math" "reflect" "runtime" "unsafe" @@ -137,23 +136,10 @@ func addStruct(v reflect.Value, numInts, numFloats, numStack *int, addInt, addFl return keepAlive } - // if greater than 64 bytes place on stack - if v.Type().Size() > 8*8 { - placeStack(v, addStack) - return keepAlive - } - var ( - savedNumFloats = *numFloats - savedNumInts = *numInts - savedNumStack = *numStack - ) - placeOnStack := postMerger(v.Type()) || !tryPlaceRegister(v, addFloat, addInt) - if placeOnStack { - // reset any values placed in registers - *numFloats = savedNumFloats - *numInts = savedNumInts - *numStack = savedNumStack + if postMerger(v.Type()) { placeStack(v, addStack) + } else { + tryPlaceRegister(v, addFloat, addInt) } return keepAlive } @@ -191,122 +177,17 @@ func postMerger(t reflect.Type) (passInMemory bool) { return true // Go does not have an SSE/SSEUP type so this is always true } -func tryPlaceRegister(v reflect.Value, addFloat func(uintptr), addInt func(uintptr)) (ok bool) { - ok = true - var val uint64 - var shift byte // # of bits to shift - var flushed bool - class := _NO_CLASS - flushIfNeeded := func() { - if flushed { - return - } - flushed = true - if class == _SSE { - addFloat(uintptr(val)) +// tryPlaceRegister passes a struct of at most two eightbytes in its ABI classes. +func tryPlaceRegister(v reflect.Value, addFloat func(uintptr), addInt func(uintptr)) { + var buf [2]uintptr + reflect.NewAt(v.Type(), unsafe.Pointer(&buf[0])).Elem().Set(v) + for i := uintptr(0); i*8 < v.Type().Size(); i++ { + if classifyEightbyte(v.Type(), i*8, i*8+8) == _SSE { + addFloat(buf[i]) } else { - addInt(uintptr(val)) + addInt(buf[i]) } - val = 0 - shift = 0 - class = _NO_CLASS } - var place func(v reflect.Value) - place = func(v reflect.Value) { - var numFields int - if v.Kind() == reflect.Struct { - numFields = v.Type().NumField() - } else { - numFields = v.Type().Len() - } - - for i := range numFields { - if v.Kind() == reflect.Struct && !isABIField(v.Type().Field(i)) { - continue - } - flushed = false - var f reflect.Value - if v.Kind() == reflect.Struct { - f = v.Field(i) - } else { - f = v.Index(i) - } - switch f.Kind() { - case reflect.Struct: - place(f) - case reflect.Bool: - if f.Bool() { - val |= 1 << shift - } - shift += 8 - class |= _INTEGER - case reflect.Pointer, reflect.UnsafePointer: - val = uint64(f.Pointer()) - shift = 64 - class = _INTEGER - case reflect.Int8: - val |= uint64(f.Int()&0xFF) << shift - shift += 8 - class |= _INTEGER - case reflect.Int16: - val |= uint64(f.Int()&0xFFFF) << shift - shift += 16 - class |= _INTEGER - case reflect.Int32: - val |= uint64(f.Int()&0xFFFF_FFFF) << shift - shift += 32 - class |= _INTEGER - case reflect.Int64, reflect.Int: - val = uint64(f.Int()) - shift = 64 - class = _INTEGER - case reflect.Uint8: - val |= f.Uint() << shift - shift += 8 - class |= _INTEGER - case reflect.Uint16: - val |= f.Uint() << shift - shift += 16 - class |= _INTEGER - case reflect.Uint32: - val |= f.Uint() << shift - shift += 32 - class |= _INTEGER - case reflect.Uint64, reflect.Uint, reflect.Uintptr: - val = f.Uint() - shift = 64 - class = _INTEGER - case reflect.Float32: - val |= uint64(math.Float32bits(float32(f.Float()))) << shift - shift += 32 - class |= _SSE - case reflect.Float64: - if v.Type().Size() > 16 { - ok = false - return - } - val = uint64(math.Float64bits(f.Float())) - shift = 64 - class = _SSE - case reflect.Array: - place(f) - default: - panic("purego: unsupported kind " + f.Kind().String()) - } - - if shift == 64 { - flushIfNeeded() - } else if shift > 64 { - // Should never happen, but may if we forget to reset shift after flush (or forget to flush), - // better fall apart here, than corrupt arguments. - panic("purego: tryPlaceRegisters shift > 64") - } - } - } - - place(v) - flushIfNeeded() - return ok } func placeStack(v reflect.Value, addStack func(uintptr)) { diff --git a/struct_arm64.go b/struct_arm64.go index b776bcd7..2f15c7f1 100644 --- a/struct_arm64.go +++ b/struct_arm64.go @@ -107,119 +107,39 @@ func placeRegisters(v reflect.Value, addFloat func(uintptr), addInt func(uintptr } func placeRegistersArm64(v reflect.Value, addFloat func(uintptr), addInt func(uintptr)) { - var val uint64 - var shift byte - var flushed bool - class := _NO_CLASS - var place func(v reflect.Value) - place = func(v reflect.Value) { - var numFields int - if v.Kind() == reflect.Struct { - numFields = v.Type().NumField() - } else { - numFields = v.Type().Len() - } - for k := range numFields { - if v.Kind() == reflect.Struct && !isABIField(v.Type().Field(k)) { - continue - } - flushed = false - var f reflect.Value - if v.Kind() == reflect.Struct { - f = v.Field(k) - } else { - f = v.Index(k) - } - align := byte(f.Type().Align()*8 - 1) - shift = (shift + align) &^ align - if shift >= 64 { - shift = 0 - flushed = true - if class == _FLOAT { - addFloat(uintptr(val)) - } else { - addInt(uintptr(val)) - } - val = 0 - class = _NO_CLASS - } - switch f.Type().Kind() { + if isHFA(v.Type()) { + var place func(reflect.Value) + place = func(v reflect.Value) { + switch v.Kind() { case reflect.Struct: - place(f) - case reflect.Bool: - if f.Bool() { - val |= 1 << shift + for i := range v.NumField() { + if isABIField(v.Type().Field(i)) { + place(v.Field(i)) + } } - shift += 8 - class |= _INT - case reflect.Uint8: - val |= f.Uint() << shift - shift += 8 - class |= _INT - case reflect.Uint16: - val |= f.Uint() << shift - shift += 16 - class |= _INT - case reflect.Uint32: - val |= f.Uint() << shift - shift += 32 - class |= _INT - case reflect.Uint64, reflect.Uint, reflect.Uintptr: - addInt(uintptr(f.Uint())) - shift = 0 - flushed = true - class = _NO_CLASS - case reflect.Int8: - val |= uint64(f.Int()&0xFF) << shift - shift += 8 - class |= _INT - case reflect.Int16: - val |= uint64(f.Int()&0xFFFF) << shift - shift += 16 - class |= _INT - case reflect.Int32: - val |= uint64(f.Int()&0xFFFF_FFFF) << shift - shift += 32 - class |= _INT - case reflect.Int64, reflect.Int: - addInt(uintptr(f.Int())) - shift = 0 - flushed = true - class = _NO_CLASS - case reflect.Float32: - if class == _FLOAT { - addFloat(uintptr(val)) - val = 0 - shift = 0 + case reflect.Array: + for i := range v.Len() { + place(v.Index(i)) } - val |= uint64(math.Float32bits(float32(f.Float()))) << shift - shift += 32 - class |= _FLOAT + case reflect.Float32: + addFloat(uintptr(math.Float32bits(float32(v.Float())))) case reflect.Float64: - addFloat(uintptr(math.Float64bits(float64(f.Float())))) - shift = 0 - flushed = true - class = _NO_CLASS - case reflect.Pointer, reflect.UnsafePointer: - addInt(f.Pointer()) - shift = 0 - flushed = true - class = _NO_CLASS - case reflect.Array: - place(f) + addFloat(uintptr(math.Float64bits(v.Float()))) default: - panic("purego: unsupported kind " + f.Kind().String()) + panic("purego: unsupported HFA kind " + v.Kind().String()) } } + place(v) + return } - place(v) - if !flushed { - if class == _FLOAT { - addFloat(uintptr(val)) - } else { - addInt(uintptr(val)) - } + + // Non-HFA composites use integer registers, including their padding. + if !v.CanAddr() { + addressable := reflect.New(v.Type()).Elem() + addressable.Set(v) + v = addressable } + copyStruct8ByteChunks(v.Addr().UnsafePointer(), v.Type().Size(), addInt) } func placeStack(v reflect.Value, keepAlive []any, addInt func(uintptr)) []any { @@ -311,11 +231,8 @@ func isHVA(t reflect.Type) bool { } // copyStruct8ByteChunks copies struct memory in 8-byte chunks to the provided callback. -// This is used for Darwin ARM64's byte-level packing of non-HFA/HVA structs. +// This preserves padding for non-HFA composites. func copyStruct8ByteChunks(ptr unsafe.Pointer, size uintptr, addChunk func(uintptr)) { - if !isDarwin { - panic("purego: should only be called on darwin") - } for offset := uintptr(0); offset < size; offset += 8 { var chunk uintptr remaining := size - offset @@ -332,37 +249,9 @@ func copyStruct8ByteChunks(ptr unsafe.Pointer, size uintptr, addChunk func(uintp } } -// placeRegisters implements Darwin ARM64 calling convention for struct arguments. -// -// For HFA/HVA structs, each element must go in a separate register (or stack slot for elements -// that don't fit in registers). We use placeRegistersArm64 for this. -// -// For non-HFA/HVA structs, Darwin uses byte-level packing. We copy the struct memory in -// 8-byte chunks, which works correctly for both register and stack placement. +// placeRegistersDarwin uses the same register layout as other ARM64 platforms. func placeRegistersDarwin(v reflect.Value, addFloat func(uintptr), addInt func(uintptr)) { - if !isDarwin { - panic("purego: placeRegistersDarwin should only be called on darwin") - } - // Check if this is an HFA/HVA - hfa := isHFA(v.Type()) - hva := isHVA(v.Type()) - - // For HFA/HVA structs, use the standard ARM64 logic which places each element separately - if hfa || hva { - placeRegistersArm64(v, addFloat, addInt) - return - } - - // For non-HFA/HVA structs, use byte-level copying - // If the value is not addressable, create an addressable copy - if !v.CanAddr() { - addressable := reflect.New(v.Type()).Elem() - addressable.Set(v) - v = addressable - } - ptr := unsafe.Pointer(v.Addr().Pointer()) - size := v.Type().Size() - copyStruct8ByteChunks(ptr, size, addInt) + placeRegistersArm64(v, addFloat, addInt) } // shouldBundleStackArgs determines if we need to start C-style packing for diff --git a/struct_test.go b/struct_test.go index f5695932..6cf661bf 100644 --- a/struct_test.go +++ b/struct_test.go @@ -613,7 +613,7 @@ func TestRegisterFunc_structArgs(t *testing.T) { type BoolFloat struct { _ structs.HostLayout b bool - _ [3]byte // purego won't do padding for you so make sure it aligns properly with C struct + _ [3]byte // redundant with the padding Go inserts, but mirrors the C layout f float32 } var BoolFloatFn func(BoolFloat) float32 @@ -944,6 +944,197 @@ func TestRegisterFunc_structArgs(t *testing.T) { } runtime.KeepAlive(ptr) } + t.Run("CharLong", func(t *testing.T) { + // The wide field must not overwrite the pending small field. + type CharLong struct { + _ structs.HostLayout + A int8 + B int64 + } + var fn func(CharLong) CharLong + register(&fn, lib, "IdentityCharLong", func(s CharLong) CharLong { + return s + }) + expected := CharLong{A: 0x7f, B: -0x0102030405060708} + if ret := fn(expected); ret != expected { + t.Fatalf("IdentityCharLong returned %+v wanted %+v", ret, expected) + } + }) + t.Run("CharDouble", func(t *testing.T) { + // Preserve the small field before the aligned wide field. + type CharDouble struct { + _ structs.HostLayout + A int8 + B float64 + } + var fn func(CharDouble) CharDouble + register(&fn, lib, "IdentityCharDouble", func(s CharDouble) CharDouble { + return s + }) + expected := CharDouble{A: -7, B: -123.5} + if ret := fn(expected); ret != expected { + t.Fatalf("IdentityCharDouble returned %+v wanted %+v", ret, expected) + } + }) + t.Run("CharPointer", func(t *testing.T) { + // Preserve the small field before the aligned wide field. + type CharPointer struct { + _ structs.HostLayout + A int8 + B unsafe.Pointer + } + var fn func(CharPointer) CharPointer + register(&fn, lib, "IdentityCharPointer", func(s CharPointer) CharPointer { + return s + }) + value := int64(0x12345678) + expected := CharPointer{A: -7, B: unsafe.Pointer(&value)} + if ret := fn(expected); ret != expected { + t.Fatalf("IdentityCharPointer returned %+v wanted %+v", ret, expected) + } + runtime.KeepAlive(&value) + }) + t.Run("CharLongBetweenPrims", func(t *testing.T) { + // The struct must consume exactly two register slots so the + // trailing scalar is not shifted. + type CharLong struct { + _ structs.HostLayout + A int8 + B int64 + } + var fn func(int64, CharLong, int64) CharLong + register(&fn, lib, "IdentityCharLongBetweenPrims", func(x int64, s CharLong, y int64) CharLong { + return s + }) + expected := CharLong{A: -1, B: 0x1122334455667788} + if ret := fn(1, expected, 2); ret != expected { + t.Fatalf("IdentityCharLongBetweenPrims returned %+v wanted %+v", ret, expected) + } + }) + t.Run("CharInt", func(t *testing.T) { + // The int32 must be placed at its padded offset, not back-to-back + // with the int8. + type CharInt struct { + _ structs.HostLayout + A int8 + B int32 + } + var fn func(CharInt) CharInt + register(&fn, lib, "IdentityCharInt", func(s CharInt) CharInt { + return s + }) + expected := CharInt{A: 0x01, B: 0x02030405} + if ret := fn(expected); ret != expected { + t.Fatalf("IdentityCharInt returned %+v wanted %+v", ret, expected) + } + }) + t.Run("NestedSmallTail", func(t *testing.T) { + // A sibling in the next eightbyte after a padded nested struct. + type NestedSmallTail struct { + _ structs.HostLayout + I struct { + _ structs.HostLayout + A int8 + B int32 + } + C int8 + } + var fn func(NestedSmallTail) NestedSmallTail + register(&fn, lib, "IdentityNestedSmallTail", func(s NestedSmallTail) NestedSmallTail { + return s + }) + expected := NestedSmallTail{ + I: struct { + _ structs.HostLayout + A int8 + B int32 + }{A: 0x11, B: 0x22334455}, + C: 0x66, + } + if ret := fn(expected); ret != expected { + t.Fatalf("IdentityNestedSmallTail returned %+v wanted %+v", ret, expected) + } + }) + t.Run("NestedIntsPlusOne", func(t *testing.T) { + // The sibling must be merged into the eightbyte the nested + // struct left pending. + type inner struct { + _ structs.HostLayout + X int32 + Y int32 + Z int32 + } + type NestedIntsPlusOne struct { + _ structs.HostLayout + A inner + B int32 + } + var sum func(NestedIntsPlusOne) int64 + register(&sum, lib, "SumNestedIntsPlusOne", func(s NestedIntsPlusOne) int64 { + return int64(s.A.X) + int64(s.A.Y) + int64(s.A.Z) + int64(s.B) + }) + if ret := sum(NestedIntsPlusOne{A: inner{X: 1, Y: 2, Z: 3}, B: 4}); ret != 10 { + t.Fatalf("SumNestedIntsPlusOne returned %d wanted 10", ret) + } + }) + t.Run("NestedPadTail", func(t *testing.T) { + // The sibling after the nested struct's trailing padding must + // still be flushed. + type inner struct { + _ structs.HostLayout + X int32 + Y int8 + } + type NestedPadTail struct { + _ structs.HostLayout + A inner + B int8 + } + var sum func(NestedPadTail) int64 + register(&sum, lib, "SumNestedPadTail", func(s NestedPadTail) int64 { + return int64(s.A.X) + int64(s.A.Y) + int64(s.B) + }) + if ret := sum(NestedPadTail{A: inner{X: 1, Y: 2}, B: 3}); ret != 6 { + t.Fatalf("SumNestedPadTail returned %d wanted 6", ret) + } + }) + t.Run("ArrayIntsPlusOne", func(t *testing.T) { + // The array counterpart of NestedIntsPlusOne. + type ArrayIntsPlusOne struct { + _ structs.HostLayout + A [3]int32 + B int32 + } + var sum func(ArrayIntsPlusOne) int64 + register(&sum, lib, "SumArrayIntsPlusOne", func(s ArrayIntsPlusOne) int64 { + return int64(s.A[0]) + int64(s.A[1]) + int64(s.A[2]) + int64(s.B) + }) + if ret := sum(ArrayIntsPlusOne{A: [3]int32{1, 2, 3}, B: 4}); ret != 10 { + t.Fatalf("SumArrayIntsPlusOne returned %d wanted 10", ret) + } + }) + t.Run("BoolFloatNoPadding", func(t *testing.T) { + // Fields are placed at their in-memory offsets, so no + // explicit padding is needed after the bool. + type BoolFloat struct { + _ structs.HostLayout + b bool + f float32 + } + var fn func(BoolFloat) float32 + register(&fn, lib, "BoolFloat", func(s BoolFloat) float32 { + if s.b { + return s.f + } + return -s.f + }) + if ret := fn(BoolFloat{b: true, f: 10}); ret != expectedFloat { + t.Fatalf("BoolFloat returned %f wanted %f", ret, expectedFloat) + } + if ret := fn(BoolFloat{b: false, f: 10}); ret != -expectedFloat { + t.Fatalf("BoolFloat returned %f wanted %f", ret, -expectedFloat) + } + }) }) } } diff --git a/testdata/structtest/struct_test.c b/testdata/structtest/struct_test.c index b513b816..257c4c67 100644 --- a/testdata/structtest/struct_test.c +++ b/testdata/structtest/struct_test.c @@ -453,3 +453,91 @@ struct Mixed5Args { struct Mixed5Args IdentityMixed5Args(struct Mixed5Args s) { return s; } + +struct CharLong { + int8_t a; + int64_t b; +}; + +struct CharLong IdentityCharLong(struct CharLong s) { + return s; +} + +struct CharLong IdentityCharLongBetweenPrims(int64_t x, struct CharLong s, int64_t y) { + (void) x; + (void) y; + return s; +} + +struct CharInt { + int8_t a; + int32_t b; +}; + +struct CharInt IdentityCharInt(struct CharInt s) { + return s; +} + +struct NestedSmallTail { + struct { + int8_t a; + int32_t b; + } i; + int8_t c; +}; + +struct NestedSmallTail IdentityNestedSmallTail(struct NestedSmallTail s) { + return s; +} + +struct NestedIntsPlusOne { + struct { + int32_t x; + int32_t y; + int32_t z; + } a; + int32_t b; +}; + +int64_t SumNestedIntsPlusOne(struct NestedIntsPlusOne s) { + return (int64_t) s.a.x + s.a.y + s.a.z + s.b; +} + +struct ArrayIntsPlusOne { + int32_t a[3]; + int32_t b; +}; + +int64_t SumArrayIntsPlusOne(struct ArrayIntsPlusOne s) { + return (int64_t) s.a[0] + s.a[1] + s.a[2] + s.b; +} + +struct NestedPadTail { + struct { + int32_t x; + int8_t y; + } a; + int8_t b; +}; + +int64_t SumNestedPadTail(struct NestedPadTail s) { + return (int64_t) s.a.x + s.a.y + s.b; +} + +struct CharDouble { + int8_t a; + double b; +}; + +struct CharDouble IdentityCharDouble(struct CharDouble s) { + return s; +} + +struct CharPointer { + int8_t a; + void * b; +}; + +struct CharPointer IdentityCharPointer(struct CharPointer s) { + return s; +}