mirror of
https://github.com/netbirdio/gvisor.git
synced 2026-05-22 17:12:49 -07:00
Implement automated marshalling for slices of Marshallable types.
PiperOrigin-RevId: 304119255
This commit is contained in:
committed by
gVisor bot
parent
d25036ad17
commit
840980aeba
@@ -304,7 +304,7 @@ func (t *Task) rseqAddrInterrupt() {
|
||||
}
|
||||
|
||||
var cs linux.RSeqCriticalSection
|
||||
if err := cs.CopyIn(t, critAddr); err != nil {
|
||||
if _, err := cs.CopyIn(t, critAddr); err != nil {
|
||||
t.Debugf("Failed to copy critical section from %#x for rseq: %v", critAddr, err)
|
||||
t.forceSignal(linux.SIGSEGV, false /* unconditional */)
|
||||
t.SendSignal(SignalInfoPriv(linux.SIGSEGV))
|
||||
|
||||
@@ -115,7 +115,8 @@ func stat(t *kernel.Task, d *fs.Dirent, dirPath bool, statAddr usermem.Addr) err
|
||||
return err
|
||||
}
|
||||
s := statFromAttrs(t, d.Inode.StableAttr, uattr)
|
||||
return s.CopyOut(t, statAddr)
|
||||
_, err = s.CopyOut(t, statAddr)
|
||||
return err
|
||||
}
|
||||
|
||||
// fstat implements fstat for the given *fs.File.
|
||||
@@ -125,7 +126,8 @@ func fstat(t *kernel.Task, f *fs.File, statAddr usermem.Addr) error {
|
||||
return err
|
||||
}
|
||||
s := statFromAttrs(t, f.Dirent.Inode.StableAttr, uattr)
|
||||
return s.CopyOut(t, statAddr)
|
||||
_, err = s.CopyOut(t, statAddr)
|
||||
return err
|
||||
}
|
||||
|
||||
// Statx implements linux syscall statx(2).
|
||||
|
||||
@@ -101,14 +101,14 @@ func EpollCtl(t *kernel.Task, args arch.SyscallArguments) (uintptr, *kernel.Sysc
|
||||
var event linux.EpollEvent
|
||||
switch op {
|
||||
case linux.EPOLL_CTL_ADD:
|
||||
if err := event.CopyIn(t, eventAddr); err != nil {
|
||||
if _, err := event.CopyIn(t, eventAddr); err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
return 0, nil, ep.AddInterest(file, fd, event)
|
||||
case linux.EPOLL_CTL_DEL:
|
||||
return 0, nil, ep.DeleteInterest(file, fd)
|
||||
case linux.EPOLL_CTL_MOD:
|
||||
if err := event.CopyIn(t, eventAddr); err != nil {
|
||||
if _, err := event.CopyIn(t, eventAddr); err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
return 0, nil, ep.ModifyInterest(file, fd, event)
|
||||
|
||||
@@ -374,7 +374,8 @@ func copyOutTimespecRemaining(t *kernel.Task, startNs ktime.Time, timeout time.D
|
||||
}
|
||||
remaining := timeoutRemaining(t, startNs, timeout)
|
||||
tsRemaining := linux.NsecToTimespec(remaining.Nanoseconds())
|
||||
return tsRemaining.CopyOut(t, timespecAddr)
|
||||
_, err := tsRemaining.CopyOut(t, timespecAddr)
|
||||
return err
|
||||
}
|
||||
|
||||
// copyOutTimevalRemaining copies the time remaining in timeout to timevalAddr.
|
||||
@@ -386,7 +387,8 @@ func copyOutTimevalRemaining(t *kernel.Task, startNs ktime.Time, timeout time.Du
|
||||
}
|
||||
remaining := timeoutRemaining(t, startNs, timeout)
|
||||
tvRemaining := linux.NsecToTimeval(remaining.Nanoseconds())
|
||||
return tvRemaining.CopyOut(t, timevalAddr)
|
||||
_, err := tvRemaining.CopyOut(t, timevalAddr)
|
||||
return err
|
||||
}
|
||||
|
||||
// pollRestartBlock encapsulates the state required to restart poll(2) via
|
||||
@@ -477,7 +479,7 @@ func Select(t *kernel.Task, args arch.SyscallArguments) (uintptr, *kernel.Syscal
|
||||
timeout := time.Duration(-1)
|
||||
if timevalAddr != 0 {
|
||||
var timeval linux.Timeval
|
||||
if err := timeval.CopyIn(t, timevalAddr); err != nil {
|
||||
if _, err := timeval.CopyIn(t, timevalAddr); err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
if timeval.Sec < 0 || timeval.Usec < 0 {
|
||||
@@ -519,7 +521,7 @@ func Pselect(t *kernel.Task, args arch.SyscallArguments) (uintptr, *kernel.Sysca
|
||||
panic(fmt.Sprintf("unsupported sizeof(void*): %d", t.Arch().Width()))
|
||||
}
|
||||
var maskStruct sigSetWithSize
|
||||
if err := maskStruct.CopyIn(t, maskWithSizeAddr); err != nil {
|
||||
if _, err := maskStruct.CopyIn(t, maskWithSizeAddr); err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
if err := setTempSignalSet(t, usermem.Addr(maskStruct.sigsetAddr), uint(maskStruct.sizeofSigset)); err != nil {
|
||||
@@ -554,7 +556,7 @@ func copyTimespecInToDuration(t *kernel.Task, timespecAddr usermem.Addr) (time.D
|
||||
timeout := time.Duration(-1)
|
||||
if timespecAddr != 0 {
|
||||
var timespec linux.Timespec
|
||||
if err := timespec.CopyIn(t, timespecAddr); err != nil {
|
||||
if _, err := timespec.CopyIn(t, timespecAddr); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if !timespec.Valid() {
|
||||
@@ -573,7 +575,7 @@ func setTempSignalSet(t *kernel.Task, maskAddr usermem.Addr, maskSize uint) erro
|
||||
return syserror.EINVAL
|
||||
}
|
||||
var mask linux.SignalSet
|
||||
if err := mask.CopyIn(t, maskAddr); err != nil {
|
||||
if _, err := mask.CopyIn(t, maskAddr); err != nil {
|
||||
return err
|
||||
}
|
||||
mask &^= kernel.UnblockableSignals
|
||||
|
||||
@@ -226,7 +226,7 @@ func Utime(t *kernel.Task, args arch.SyscallArguments) (uintptr, *kernel.Syscall
|
||||
opts.Stat.Mtime.Nsec = linux.UTIME_NOW
|
||||
} else {
|
||||
var times linux.Utime
|
||||
if err := times.CopyIn(t, timesAddr); err != nil {
|
||||
if _, err := times.CopyIn(t, timesAddr); err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
opts.Stat.Atime.Sec = times.Actime
|
||||
|
||||
@@ -91,7 +91,8 @@ func fstatat(t *kernel.Task, dirfd int32, pathAddr, statAddr usermem.Addr, flags
|
||||
}
|
||||
var stat linux.Stat
|
||||
convertStatxToUserStat(t, &statx, &stat)
|
||||
return stat.CopyOut(t, statAddr)
|
||||
_, err = stat.CopyOut(t, statAddr)
|
||||
return err
|
||||
}
|
||||
start = dirfile.VirtualDentry()
|
||||
start.IncRef()
|
||||
@@ -111,7 +112,8 @@ func fstatat(t *kernel.Task, dirfd int32, pathAddr, statAddr usermem.Addr, flags
|
||||
}
|
||||
var stat linux.Stat
|
||||
convertStatxToUserStat(t, &statx, &stat)
|
||||
return stat.CopyOut(t, statAddr)
|
||||
_, err = stat.CopyOut(t, statAddr)
|
||||
return err
|
||||
}
|
||||
|
||||
func timespecFromStatxTimestamp(sxts linux.StatxTimestamp) linux.Timespec {
|
||||
@@ -140,7 +142,8 @@ func Fstat(t *kernel.Task, args arch.SyscallArguments) (uintptr, *kernel.Syscall
|
||||
}
|
||||
var stat linux.Stat
|
||||
convertStatxToUserStat(t, &statx, &stat)
|
||||
return 0, nil, stat.CopyOut(t, statAddr)
|
||||
_, err = stat.CopyOut(t, statAddr)
|
||||
return 0, nil, err
|
||||
}
|
||||
|
||||
// Statx implements Linux syscall statx(2).
|
||||
@@ -199,7 +202,8 @@ func Statx(t *kernel.Task, args arch.SyscallArguments) (uintptr, *kernel.Syscall
|
||||
return 0, nil, err
|
||||
}
|
||||
userifyStatx(t, &statx)
|
||||
return 0, nil, statx.CopyOut(t, statxAddr)
|
||||
_, err = statx.CopyOut(t, statxAddr)
|
||||
return 0, nil, err
|
||||
}
|
||||
start = dirfile.VirtualDentry()
|
||||
start.IncRef()
|
||||
@@ -218,7 +222,8 @@ func Statx(t *kernel.Task, args arch.SyscallArguments) (uintptr, *kernel.Syscall
|
||||
return 0, nil, err
|
||||
}
|
||||
userifyStatx(t, &statx)
|
||||
return 0, nil, statx.CopyOut(t, statxAddr)
|
||||
_, err = statx.CopyOut(t, statxAddr)
|
||||
return 0, nil, err
|
||||
}
|
||||
|
||||
func userifyStatx(t *kernel.Task, statx *linux.Statx) {
|
||||
@@ -359,8 +364,8 @@ func Statfs(t *kernel.Task, args arch.SyscallArguments) (uintptr, *kernel.Syscal
|
||||
if err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
|
||||
return 0, nil, statfs.CopyOut(t, bufAddr)
|
||||
_, err = statfs.CopyOut(t, bufAddr)
|
||||
return 0, nil, err
|
||||
}
|
||||
|
||||
// Fstatfs implements Linux syscall fstatfs(2).
|
||||
@@ -378,6 +383,6 @@ func Fstatfs(t *kernel.Task, args arch.SyscallArguments) (uintptr, *kernel.Sysca
|
||||
if err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
|
||||
return 0, nil, statfs.CopyOut(t, bufAddr)
|
||||
_, err = statfs.CopyOut(t, bufAddr)
|
||||
return 0, nil, err
|
||||
}
|
||||
|
||||
@@ -161,6 +161,10 @@ func AlignmentCheck(t *testing.T, typ reflect.Type) (ok bool, delta uint64) {
|
||||
if typ.NumField() > 0 && nextXOff != int(typ.Size()) {
|
||||
implicitPad := int(typ.Size()) - nextXOff
|
||||
f := typ.Field(typ.NumField() - 1) // Final field
|
||||
if tag, ok := f.Tag.Lookup("marshal"); ok && tag == "unaligned" {
|
||||
// Final field explicitly marked unaligned.
|
||||
break
|
||||
}
|
||||
t.Fatalf("Suspect offset for field %s.%s at the end of %s, detected an implicit %d byte padding from offset %d to %d at the end of the struct; either add %d bytes of explict padding at end of the struct or tag the final field %s as `marshal:\"unaligned\"`.",
|
||||
typ.Name(), f.Name, typ.Name(), implicitPad, nextXOff, typ.Size(), implicitPad, f.Name)
|
||||
}
|
||||
|
||||
@@ -53,9 +53,10 @@ go_marshal = rule(
|
||||
|
||||
# marshal_deps are the dependencies requied by generated code.
|
||||
marshal_deps = [
|
||||
"//tools/go_marshal/marshal",
|
||||
"//pkg/gohacks",
|
||||
"//pkg/safecopy",
|
||||
"//pkg/usermem",
|
||||
"//tools/go_marshal/marshal",
|
||||
]
|
||||
|
||||
# marshal_test_deps are required by test targets.
|
||||
|
||||
@@ -28,12 +28,6 @@ import (
|
||||
"gvisor.dev/gvisor/tools/tags"
|
||||
)
|
||||
|
||||
const (
|
||||
marshalImport = "gvisor.dev/gvisor/tools/go_marshal/marshal"
|
||||
safecopyImport = "gvisor.dev/gvisor/pkg/safecopy"
|
||||
usermemImport = "gvisor.dev/gvisor/pkg/usermem"
|
||||
)
|
||||
|
||||
// List of identifiers we use in generated code that may conflict with a
|
||||
// similarly-named source identifier. Abort gracefully when we see these to
|
||||
// avoid potentially confusing compilation failures in generated code.
|
||||
@@ -44,8 +38,8 @@ const (
|
||||
// All recievers are single letters, so we don't allow import aliases to be a
|
||||
// single letter.
|
||||
var badIdents = []string{
|
||||
"addr", "blk", "buf", "dst", "dsts", "err", "hdr", "idx", "inner", "len",
|
||||
"ptr", "src", "srcs", "task", "val",
|
||||
"addr", "blk", "buf", "dst", "dsts", "count", "err", "hdr", "idx", "inner",
|
||||
"length", "limit", "ptr", "size", "src", "srcs", "task", "val",
|
||||
// All single-letter identifiers.
|
||||
}
|
||||
|
||||
@@ -110,9 +104,10 @@ func NewGenerator(srcs []string, out, outTest, pkg string, imports []string) (*G
|
||||
g.imports.add("reflect")
|
||||
g.imports.add("runtime")
|
||||
g.imports.add("unsafe")
|
||||
g.imports.add(marshalImport)
|
||||
g.imports.add(safecopyImport)
|
||||
g.imports.add(usermemImport)
|
||||
g.imports.add("gvisor.dev/gvisor/pkg/gohacks")
|
||||
g.imports.add("gvisor.dev/gvisor/pkg/safecopy")
|
||||
g.imports.add("gvisor.dev/gvisor/pkg/usermem")
|
||||
g.imports.add("gvisor.dev/gvisor/tools/go_marshal/marshal")
|
||||
|
||||
return &g, nil
|
||||
}
|
||||
@@ -194,10 +189,73 @@ func (g *Generator) parse() ([]*ast.File, []*token.FileSet, error) {
|
||||
return files, fsets, nil
|
||||
}
|
||||
|
||||
// sliceAPI carries information about the '+marshal slice' directive.
|
||||
type sliceAPI struct {
|
||||
// Comment node in the AST containing the +marshal tag.
|
||||
comment *ast.Comment
|
||||
// Identifier fragment to use when naming generated functions for the slice
|
||||
// API.
|
||||
ident string
|
||||
// Whether the generated functions should reference the newtype name, or the
|
||||
// inner type name. Only meaningful on newtype declarations on primitives.
|
||||
inner bool
|
||||
}
|
||||
|
||||
// marshallableType carries information about a type marked with the '+marshal'
|
||||
// directive.
|
||||
type marshallableType struct {
|
||||
spec *ast.TypeSpec
|
||||
slice *sliceAPI
|
||||
}
|
||||
|
||||
func newMarshallableType(fset *token.FileSet, tagLine *ast.Comment, spec *ast.TypeSpec) marshallableType {
|
||||
mt := marshallableType{
|
||||
spec: spec,
|
||||
slice: nil,
|
||||
}
|
||||
|
||||
var unhandledTags []string
|
||||
|
||||
for _, tag := range strings.Fields(strings.TrimPrefix(tagLine.Text, "// +marshal")) {
|
||||
if strings.HasPrefix(tag, "slice:") {
|
||||
tokens := strings.Split(tag, ":")
|
||||
if len(tokens) < 2 || len(tokens) > 3 {
|
||||
abortAt(fset.Position(tagLine.Slash), fmt.Sprintf("+marshal directive has invalid 'slice' clause. Expecting format 'slice:<IDENTIFIER>[:inner]', got '%v'", tag))
|
||||
}
|
||||
if len(tokens[1]) == 0 {
|
||||
abortAt(fset.Position(tagLine.Slash), "+marshal slice directive has empty identifier argument. Expecting '+marshal slice:identifier'")
|
||||
}
|
||||
|
||||
sa := &sliceAPI{
|
||||
comment: tagLine,
|
||||
ident: tokens[1],
|
||||
}
|
||||
mt.slice = sa
|
||||
|
||||
if len(tokens) == 3 {
|
||||
if tokens[2] != "inner" {
|
||||
abortAt(fset.Position(tagLine.Slash), "+marshal slice directive has an invalid argument. Expecting '+marshal slice:<IDENTIFIER>[:inner]'")
|
||||
}
|
||||
sa.inner = true
|
||||
}
|
||||
|
||||
continue
|
||||
}
|
||||
|
||||
unhandledTags = append(unhandledTags, tag)
|
||||
}
|
||||
|
||||
if len(unhandledTags) > 0 {
|
||||
abortAt(fset.Position(tagLine.Slash), fmt.Sprintf("+marshal directive contained the following unknown clauses: %v", strings.Join(unhandledTags, " ")))
|
||||
}
|
||||
|
||||
return mt
|
||||
}
|
||||
|
||||
// collectMarshallableTypes walks the parsed AST and collects a list of type
|
||||
// declarations for which we need to generate the Marshallable interface.
|
||||
func (g *Generator) collectMarshallableTypes(a *ast.File, f *token.FileSet) []*ast.TypeSpec {
|
||||
var types []*ast.TypeSpec
|
||||
func (g *Generator) collectMarshallableTypes(a *ast.File, f *token.FileSet) []marshallableType {
|
||||
var types []marshallableType
|
||||
for _, decl := range a.Decls {
|
||||
gdecl, ok := decl.(*ast.GenDecl)
|
||||
// Type declaration?
|
||||
@@ -212,9 +270,11 @@ func (g *Generator) collectMarshallableTypes(a *ast.File, f *token.FileSet) []*a
|
||||
}
|
||||
// Does the comment contain a "+marshal" line?
|
||||
marked := false
|
||||
var tagLine *ast.Comment
|
||||
for _, c := range gdecl.Doc.List {
|
||||
if c.Text == "// +marshal" {
|
||||
if strings.HasPrefix(c.Text, "// +marshal") {
|
||||
marked = true
|
||||
tagLine = c
|
||||
break
|
||||
}
|
||||
}
|
||||
@@ -229,20 +289,17 @@ func (g *Generator) collectMarshallableTypes(a *ast.File, f *token.FileSet) []*a
|
||||
switch t.Type.(type) {
|
||||
case *ast.StructType:
|
||||
debugfAt(f.Position(t.Pos()), "Collected marshallable struct %s.\n", t.Name.Name)
|
||||
types = append(types, t)
|
||||
continue
|
||||
case *ast.Ident: // Newtype on primitive.
|
||||
debugfAt(f.Position(t.Pos()), "Collected marshallable newtype on primitive %s.\n", t.Name.Name)
|
||||
types = append(types, t)
|
||||
continue
|
||||
case *ast.ArrayType: // Newtype on array.
|
||||
debugfAt(f.Position(t.Pos()), "Collected marshallable newtype on array %s.\n", t.Name.Name)
|
||||
types = append(types, t)
|
||||
continue
|
||||
default:
|
||||
// A user specifically requested marshalling on this type, but we
|
||||
// don't support it.
|
||||
abortAt(f.Position(t.Pos()), fmt.Sprintf("Marshalling codegen was requested on type '%s', but go-marshal doesn't support this kind of declaration.\n", t.Name))
|
||||
}
|
||||
// A user specifically requested marshalling on this type, but we
|
||||
// don't support it.
|
||||
abortAt(f.Position(t.Pos()), fmt.Sprintf("Marshalling codegen was requested on type '%s', but go-marshal doesn't support this kind of declaration.\n", t.Name))
|
||||
types = append(types, newMarshallableType(f, tagLine, t))
|
||||
|
||||
}
|
||||
}
|
||||
return types
|
||||
@@ -281,19 +338,28 @@ func (g *Generator) collectImports(a *ast.File, f *token.FileSet) map[string]imp
|
||||
|
||||
}
|
||||
|
||||
func (g *Generator) generateOne(t *ast.TypeSpec, fset *token.FileSet) *interfaceGenerator {
|
||||
i := newInterfaceGenerator(t, fset)
|
||||
switch ty := t.Type.(type) {
|
||||
func (g *Generator) generateOne(t marshallableType, fset *token.FileSet) *interfaceGenerator {
|
||||
i := newInterfaceGenerator(t.spec, fset)
|
||||
switch ty := t.spec.Type.(type) {
|
||||
case *ast.StructType:
|
||||
i.validateStruct(t, ty)
|
||||
i.validateStruct(t.spec, ty)
|
||||
i.emitMarshallableForStruct(ty)
|
||||
if t.slice != nil {
|
||||
i.emitMarshallableSliceForStruct(ty, t.slice)
|
||||
}
|
||||
case *ast.Ident:
|
||||
i.validatePrimitiveNewtype(ty)
|
||||
i.emitMarshallableForPrimitiveNewtype(ty)
|
||||
if t.slice != nil {
|
||||
i.emitMarshallableSliceForPrimitiveNewtype(ty, t.slice)
|
||||
}
|
||||
case *ast.ArrayType:
|
||||
i.validateArrayNewtype(t.Name, ty)
|
||||
i.validateArrayNewtype(t.spec.Name, ty)
|
||||
// After validate, we can safely call arrayLen.
|
||||
i.emitMarshallableForArrayNewtype(t.Name, ty.Elt.(*ast.Ident), arrayLen(ty))
|
||||
i.emitMarshallableForArrayNewtype(t.spec.Name, ty.Elt.(*ast.Ident), arrayLen(ty))
|
||||
if t.slice != nil {
|
||||
abortAt(fset.Position(t.slice.comment.Slash), fmt.Sprintf("Array type marked as '+marshal slice:...', but this is not supported. Perhaps fold one of the dimensions?"))
|
||||
}
|
||||
default:
|
||||
// This should've been filtered out by collectMarshallabeTypes.
|
||||
panic(fmt.Sprintf("Unexpected type %+v", ty))
|
||||
@@ -303,9 +369,9 @@ func (g *Generator) generateOne(t *ast.TypeSpec, fset *token.FileSet) *interface
|
||||
|
||||
// generateOneTestSuite generates a test suite for the automatically generated
|
||||
// implementations type t.
|
||||
func (g *Generator) generateOneTestSuite(t *ast.TypeSpec) *testGenerator {
|
||||
i := newTestGenerator(t)
|
||||
i.emitTests()
|
||||
func (g *Generator) generateOneTestSuite(t marshallableType) *testGenerator {
|
||||
i := newTestGenerator(t.spec)
|
||||
i.emitTests(t.slice)
|
||||
return i
|
||||
}
|
||||
|
||||
|
||||
@@ -163,3 +163,65 @@ func (g *interfaceGenerator) unmarshalScalar(accessor, typ, bufVar string) {
|
||||
g.recordPotentiallyNonPackedField(accessor)
|
||||
}
|
||||
}
|
||||
|
||||
// emitCastToByteSlice unsafely casts an arbitrary type's underlying memory to a
|
||||
// byte slice, bypassing escape analysis. The caller is responsible for ensuring
|
||||
// srcPtr lives until they're done with dstVar, the runtime does not consider
|
||||
// dstVar dependent on srcPtr due to the escape analysis bypass.
|
||||
//
|
||||
// srcPtr must be a pointer.
|
||||
//
|
||||
// This function uses internally uses the identifier "hdr", and cannot be used
|
||||
// in a context where it is already bound.
|
||||
func (g *interfaceGenerator) emitCastToByteSlice(srcPtr, dstVar, lenExpr string) {
|
||||
g.recordUsedImport("gohacks")
|
||||
g.emit("// Construct a slice backed by dst's underlying memory.\n")
|
||||
g.emit("var %s []byte\n", dstVar)
|
||||
g.emit("hdr := (*reflect.SliceHeader)(unsafe.Pointer(&%s))\n", dstVar)
|
||||
g.emit("hdr.Data = uintptr(gohacks.Noescape(unsafe.Pointer(%s)))\n", srcPtr)
|
||||
g.emit("hdr.Len = %s\n", lenExpr)
|
||||
g.emit("hdr.Cap = %s\n\n", lenExpr)
|
||||
}
|
||||
|
||||
// emitCastToByteSlice unsafely casts a slice with elements of an abitrary type
|
||||
// to a byte slice. As part of the cast, the byte slice is made to look
|
||||
// independent of the src slice by bypassing escape analysis. This means the
|
||||
// byte slice can be used without causing the source to escape. The caller is
|
||||
// responsible for ensuring srcPtr lives until they're done with dstVar, as the
|
||||
// runtime no longer considers dstVar dependent on srcPtr and is free to GC it.
|
||||
//
|
||||
// srcPtr must be a pointer.
|
||||
//
|
||||
// This function uses internally uses the identifiers "ptr", "val" and "hdr",
|
||||
// and cannot be used in a context where these identifiers are already bound.
|
||||
func (g *interfaceGenerator) emitCastSliceToByteSlice(srcPtr, dstVar, lenExpr string) {
|
||||
g.emitNoEscapeSliceDataPointer(srcPtr, "val")
|
||||
|
||||
g.emit("// Construct a slice backed by dst's underlying memory.\n")
|
||||
g.emit("var %s []byte\n", dstVar)
|
||||
g.emit("hdr := (*reflect.SliceHeader)(unsafe.Pointer(&%s))\n", dstVar)
|
||||
g.emit("hdr.Data = uintptr(val)\n")
|
||||
g.emit("hdr.Len = %s\n", lenExpr)
|
||||
g.emit("hdr.Cap = %s\n\n", lenExpr)
|
||||
}
|
||||
|
||||
// emitNoEscapeSliceDataPointer unsafely casts a slice's data pointer to an
|
||||
// unsafe.Pointer, bypassing escape analysis. The caller is responsible for
|
||||
// ensuring srcPtr lives until they're done with dstVar, as the runtime no
|
||||
// longer considers dstVar dependent on srcPtr and is free to GC it.
|
||||
//
|
||||
// srcPtr must be a pointer.
|
||||
//
|
||||
// This function uses internally uses the identifier "ptr" cannot be used in a
|
||||
// context where this identifier is already bound.
|
||||
func (g *interfaceGenerator) emitNoEscapeSliceDataPointer(srcPtr, dstVar string) {
|
||||
g.recordUsedImport("gohacks")
|
||||
g.emit("ptr := unsafe.Pointer(%s)\n", srcPtr)
|
||||
g.emit("%s := gohacks.Noescape(unsafe.Pointer((*reflect.SliceHeader)(ptr).Data))\n\n", dstVar)
|
||||
}
|
||||
|
||||
func (g *interfaceGenerator) emitKeepAlive(ptrVar string) {
|
||||
g.emit("// Since we bypassed the compiler's escape analysis, indicate that %s\n", ptrVar)
|
||||
g.emit("// must live until the use above.\n")
|
||||
g.emit("runtime.KeepAlive(%s)\n", ptrVar)
|
||||
}
|
||||
|
||||
@@ -104,79 +104,43 @@ func (g *interfaceGenerator) emitMarshallableForArrayNewtype(n, elt *ast.Ident,
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
|
||||
g.emit("// CopyOut implements marshal.Marshallable.CopyOut.\n")
|
||||
g.emit("func (%s *%s) CopyOut(task marshal.Task, addr usermem.Addr) error {\n", g.r, g.typeName())
|
||||
g.emit("// CopyOutN implements marshal.Marshallable.CopyOutN.\n")
|
||||
g.emit("func (%s *%s) CopyOutN(task marshal.Task, addr usermem.Addr, limit int) (int, error) {\n", g.r, g.typeName())
|
||||
g.inIndent(func() {
|
||||
// Fast serialization.
|
||||
g.emit("// Bypass escape analysis on %s. The no-op arithmetic operation on the\n", g.r)
|
||||
g.emit("// pointer makes the compiler think val doesn't depend on %s.\n", g.r)
|
||||
g.emit("// See src/runtime/stubs.go:noescape() in the golang toolchain.\n")
|
||||
g.emit("ptr := unsafe.Pointer(%s)\n", g.r)
|
||||
g.emit("val := uintptr(ptr)\n")
|
||||
g.emit("val = val^0\n\n")
|
||||
g.emitCastToByteSlice(g.r, "buf", fmt.Sprintf("%s.SizeBytes()", g.r))
|
||||
|
||||
g.emit("// Construct a slice backed by %s's underlying memory.\n", g.r)
|
||||
g.emit("var buf []byte\n")
|
||||
g.emit("hdr := (*reflect.SliceHeader)(unsafe.Pointer(&buf))\n")
|
||||
g.emit("hdr.Data = val\n")
|
||||
g.emit("hdr.Len = %s.SizeBytes()\n", g.r)
|
||||
g.emit("hdr.Cap = %s.SizeBytes()\n\n", g.r)
|
||||
g.emit("length, err := task.CopyOutBytes(addr, buf[:limit])\n")
|
||||
g.emitKeepAlive(g.r)
|
||||
g.emit("return length, err\n")
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
|
||||
g.emit("_, err := task.CopyOutBytes(addr, buf)\n")
|
||||
g.emit("// Since we bypassed the compiler's escape analysis, indicate that %s\n", g.r)
|
||||
g.emit("// must live until after the CopyOutBytes.\n")
|
||||
g.emit("runtime.KeepAlive(%s)\n", g.r)
|
||||
g.emit("return err\n")
|
||||
g.emit("// CopyOut implements marshal.Marshallable.CopyOut.\n")
|
||||
g.emit("func (%s *%s) CopyOut(task marshal.Task, addr usermem.Addr) (int, error) {\n", g.r, g.typeName())
|
||||
g.inIndent(func() {
|
||||
g.emit("return %s.CopyOutN(task, addr, %s.SizeBytes())\n", g.r, g.r)
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
|
||||
g.emit("// CopyIn implements marshal.Marshallable.CopyIn.\n")
|
||||
g.emit("func (%s *%s) CopyIn(task marshal.Task, addr usermem.Addr) error {\n", g.r, g.typeName())
|
||||
g.emit("func (%s *%s) CopyIn(task marshal.Task, addr usermem.Addr) (int, error) {\n", g.r, g.typeName())
|
||||
g.inIndent(func() {
|
||||
g.emit("// Bypass escape analysis on %s. The no-op arithmetic operation on the\n", g.r)
|
||||
g.emit("// pointer makes the compiler think val doesn't depend on %s.\n", g.r)
|
||||
g.emit("// See src/runtime/stubs.go:noescape() in the golang toolchain.\n")
|
||||
g.emit("ptr := unsafe.Pointer(%s)\n", g.r)
|
||||
g.emit("val := uintptr(ptr)\n")
|
||||
g.emit("val = val^0\n\n")
|
||||
g.emitCastToByteSlice(g.r, "buf", fmt.Sprintf("%s.SizeBytes()", g.r))
|
||||
|
||||
g.emit("// Construct a slice backed by %s's underlying memory.\n", g.r)
|
||||
g.emit("var buf []byte\n")
|
||||
g.emit("hdr := (*reflect.SliceHeader)(unsafe.Pointer(&buf))\n")
|
||||
g.emit("hdr.Data = val\n")
|
||||
g.emit("hdr.Len = %s.SizeBytes()\n", g.r)
|
||||
g.emit("hdr.Cap = %s.SizeBytes()\n\n", g.r)
|
||||
|
||||
g.emit("_, err := task.CopyInBytes(addr, buf)\n")
|
||||
g.emit("// Since we bypassed the compiler's escape analysis, indicate that %s\n", g.r)
|
||||
g.emit("// must live until after the CopyInBytes.\n")
|
||||
g.emit("runtime.KeepAlive(%s)\n", g.r)
|
||||
g.emit("return err\n")
|
||||
g.emit("length, err := task.CopyInBytes(addr, buf)\n")
|
||||
g.emitKeepAlive(g.r)
|
||||
g.emit("return length, err\n")
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
|
||||
g.emit("// WriteTo implements io.WriterTo.WriteTo.\n")
|
||||
g.emit("func (%s *%s) WriteTo(w io.Writer) (int64, error) {\n", g.r, g.typeName())
|
||||
g.inIndent(func() {
|
||||
g.emit("// Bypass escape analysis on %s. The no-op arithmetic operation on the\n", g.r)
|
||||
g.emit("// pointer makes the compiler think val doesn't depend on %s.\n", g.r)
|
||||
g.emit("// See src/runtime/stubs.go:noescape() in the golang toolchain.\n")
|
||||
g.emit("ptr := unsafe.Pointer(%s)\n", g.r)
|
||||
g.emit("val := uintptr(ptr)\n")
|
||||
g.emit("val = val^0\n\n")
|
||||
g.emitCastToByteSlice(g.r, "buf", fmt.Sprintf("%s.SizeBytes()", g.r))
|
||||
|
||||
g.emit("// Construct a slice backed by %s's underlying memory.\n", g.r)
|
||||
g.emit("var buf []byte\n")
|
||||
g.emit("hdr := (*reflect.SliceHeader)(unsafe.Pointer(&buf))\n")
|
||||
g.emit("hdr.Data = val\n")
|
||||
g.emit("hdr.Len = %s.SizeBytes()\n", g.r)
|
||||
g.emit("hdr.Cap = %s.SizeBytes()\n\n", g.r)
|
||||
|
||||
g.emit("len, err := w.Write(buf)\n")
|
||||
g.emit("// Since we bypassed the compiler's escape analysis, indicate that %s\n", g.r)
|
||||
g.emit("// must live until after the Write.\n")
|
||||
g.emit("runtime.KeepAlive(%s)\n", g.r)
|
||||
g.emit("return int64(len), err\n")
|
||||
g.emit("length, err := w.Write(buf)\n")
|
||||
g.emitKeepAlive(g.r)
|
||||
g.emit("return int64(length), err\n")
|
||||
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
|
||||
@@ -150,80 +150,133 @@ func (g *interfaceGenerator) emitMarshallableForPrimitiveNewtype(nt *ast.Ident)
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
|
||||
g.emit("// CopyOut implements marshal.Marshallable.CopyOut.\n")
|
||||
g.emit("func (%s *%s) CopyOut(task marshal.Task, addr usermem.Addr) error {\n", g.r, g.typeName())
|
||||
g.emit("// CopyOutN implements marshal.Marshallable.CopyOutN.\n")
|
||||
g.emit("func (%s *%s) CopyOutN(task marshal.Task, addr usermem.Addr, limit int) (int, error) {\n", g.r, g.typeName())
|
||||
g.inIndent(func() {
|
||||
// Fast serialization.
|
||||
g.emit("// Bypass escape analysis on %s. The no-op arithmetic operation on the\n", g.r)
|
||||
g.emit("// pointer makes the compiler think val doesn't depend on %s.\n", g.r)
|
||||
g.emit("// See src/runtime/stubs.go:noescape() in the golang toolchain.\n")
|
||||
g.emit("ptr := unsafe.Pointer(%s)\n", g.r)
|
||||
g.emit("val := uintptr(ptr)\n")
|
||||
g.emit("val = val^0\n\n")
|
||||
g.emitCastToByteSlice(g.r, "buf", fmt.Sprintf("%s.SizeBytes()", g.r))
|
||||
|
||||
g.emit("// Construct a slice backed by %s's underlying memory.\n", g.r)
|
||||
g.emit("var buf []byte\n")
|
||||
g.emit("hdr := (*reflect.SliceHeader)(unsafe.Pointer(&buf))\n")
|
||||
g.emit("hdr.Data = val\n")
|
||||
g.emit("hdr.Len = %s.SizeBytes()\n", g.r)
|
||||
g.emit("hdr.Cap = %s.SizeBytes()\n\n", g.r)
|
||||
g.emit("length, err := task.CopyOutBytes(addr, buf[:limit])\n")
|
||||
g.emitKeepAlive(g.r)
|
||||
g.emit("return length, err\n")
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
|
||||
g.emit("_, err := task.CopyOutBytes(addr, buf)\n")
|
||||
g.emit("// Since we bypassed the compiler's escape analysis, indicate that %s\n", g.r)
|
||||
g.emit("// must live until after the CopyOutBytes.\n")
|
||||
g.emit("runtime.KeepAlive(%s)\n", g.r)
|
||||
g.emit("return err\n")
|
||||
g.emit("// CopyOut implements marshal.Marshallable.CopyOut.\n")
|
||||
g.emit("func (%s *%s) CopyOut(task marshal.Task, addr usermem.Addr) (int, error) {\n", g.r, g.typeName())
|
||||
g.inIndent(func() {
|
||||
g.emit("return %s.CopyOutN(task, addr, %s.SizeBytes())\n", g.r, g.r)
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
|
||||
g.emit("// CopyIn implements marshal.Marshallable.CopyIn.\n")
|
||||
g.emit("func (%s *%s) CopyIn(task marshal.Task, addr usermem.Addr) error {\n", g.r, g.typeName())
|
||||
g.emit("func (%s *%s) CopyIn(task marshal.Task, addr usermem.Addr) (int, error) {\n", g.r, g.typeName())
|
||||
g.inIndent(func() {
|
||||
g.emit("// Bypass escape analysis on %s. The no-op arithmetic operation on the\n", g.r)
|
||||
g.emit("// pointer makes the compiler think val doesn't depend on %s.\n", g.r)
|
||||
g.emit("// See src/runtime/stubs.go:noescape() in the golang toolchain.\n")
|
||||
g.emit("ptr := unsafe.Pointer(%s)\n", g.r)
|
||||
g.emit("val := uintptr(ptr)\n")
|
||||
g.emit("val = val^0\n\n")
|
||||
g.emitCastToByteSlice(g.r, "buf", fmt.Sprintf("%s.SizeBytes()", g.r))
|
||||
|
||||
g.emit("// Construct a slice backed by %s's underlying memory.\n", g.r)
|
||||
g.emit("var buf []byte\n")
|
||||
g.emit("hdr := (*reflect.SliceHeader)(unsafe.Pointer(&buf))\n")
|
||||
g.emit("hdr.Data = val\n")
|
||||
g.emit("hdr.Len = %s.SizeBytes()\n", g.r)
|
||||
g.emit("hdr.Cap = %s.SizeBytes()\n\n", g.r)
|
||||
|
||||
g.emit("_, err := task.CopyInBytes(addr, buf)\n")
|
||||
g.emit("// Since we bypassed the compiler's escape analysis, indicate that %s\n", g.r)
|
||||
g.emit("// must live until after the CopyInBytes.\n")
|
||||
g.emit("runtime.KeepAlive(%s)\n", g.r)
|
||||
g.emit("return err\n")
|
||||
g.emit("length, err := task.CopyInBytes(addr, buf)\n")
|
||||
g.emitKeepAlive(g.r)
|
||||
g.emit("return length, err\n")
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
|
||||
g.emit("// WriteTo implements io.WriterTo.WriteTo.\n")
|
||||
g.emit("func (%s *%s) WriteTo(w io.Writer) (int64, error) {\n", g.r, g.typeName())
|
||||
g.inIndent(func() {
|
||||
g.emit("// Bypass escape analysis on %s. The no-op arithmetic operation on the\n", g.r)
|
||||
g.emit("// pointer makes the compiler think val doesn't depend on %s.\n", g.r)
|
||||
g.emit("// See src/runtime/stubs.go:noescape() in the golang toolchain.\n")
|
||||
g.emit("ptr := unsafe.Pointer(%s)\n", g.r)
|
||||
g.emit("val := uintptr(ptr)\n")
|
||||
g.emit("val = val^0\n\n")
|
||||
g.emitCastToByteSlice(g.r, "buf", fmt.Sprintf("%s.SizeBytes()", g.r))
|
||||
|
||||
g.emit("// Construct a slice backed by %s's underlying memory.\n", g.r)
|
||||
g.emit("var buf []byte\n")
|
||||
g.emit("hdr := (*reflect.SliceHeader)(unsafe.Pointer(&buf))\n")
|
||||
g.emit("hdr.Data = val\n")
|
||||
g.emit("hdr.Len = %s.SizeBytes()\n", g.r)
|
||||
g.emit("hdr.Cap = %s.SizeBytes()\n\n", g.r)
|
||||
|
||||
g.emit("len, err := w.Write(buf)\n")
|
||||
g.emit("// Since we bypassed the compiler's escape analysis, indicate that %s\n", g.r)
|
||||
g.emit("// must live until after the Write.\n")
|
||||
g.emit("runtime.KeepAlive(%s)\n", g.r)
|
||||
g.emit("return int64(len), err\n")
|
||||
g.emit("length, err := w.Write(buf)\n")
|
||||
g.emitKeepAlive(g.r)
|
||||
g.emit("return int64(length), err\n")
|
||||
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
}
|
||||
|
||||
func (g *interfaceGenerator) emitMarshallableSliceForPrimitiveNewtype(nt *ast.Ident, slice *sliceAPI) {
|
||||
g.recordUsedImport("marshal")
|
||||
g.recordUsedImport("usermem")
|
||||
g.recordUsedImport("reflect")
|
||||
g.recordUsedImport("runtime")
|
||||
g.recordUsedImport("unsafe")
|
||||
|
||||
eltType := g.typeName()
|
||||
if slice.inner {
|
||||
eltType = nt.Name
|
||||
}
|
||||
|
||||
g.emit("// Copy%sIn copies in a slice of %s objects from the task's memory.\n", slice.ident, eltType)
|
||||
g.emit("func Copy%sIn(task marshal.Task, addr usermem.Addr, dst []%s) (int, error) {\n", slice.ident, eltType)
|
||||
g.inIndent(func() {
|
||||
g.emit("count := len(dst)\n")
|
||||
g.emit("if count == 0 {\n")
|
||||
g.inIndent(func() {
|
||||
g.emit("return 0, nil\n")
|
||||
})
|
||||
g.emit("}\n")
|
||||
g.emit("size := (*%s)(nil).SizeBytes()\n\n", g.typeName())
|
||||
|
||||
g.emitCastSliceToByteSlice("&dst", "buf", "size * count")
|
||||
|
||||
g.emit("length, err := task.CopyInBytes(addr, buf)\n")
|
||||
g.emitKeepAlive("dst")
|
||||
g.emit("return length, err\n")
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
|
||||
g.emit("// Copy%sOut copies a slice of %s objects to the task's memory.\n", slice.ident, eltType)
|
||||
g.emit("func Copy%sOut(task marshal.Task, addr usermem.Addr, src []%s) (int, error) {\n", slice.ident, eltType)
|
||||
g.inIndent(func() {
|
||||
g.emit("count := len(src)\n")
|
||||
g.emit("if count == 0 {\n")
|
||||
g.inIndent(func() {
|
||||
g.emit("return 0, nil\n")
|
||||
})
|
||||
g.emit("}\n")
|
||||
g.emit("size := (*%s)(nil).SizeBytes()\n\n", g.typeName())
|
||||
|
||||
g.emitCastSliceToByteSlice("&src", "buf", "size * count")
|
||||
|
||||
g.emit("length, err := task.CopyOutBytes(addr, buf)\n")
|
||||
g.emitKeepAlive("src")
|
||||
g.emit("return length, err\n")
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
|
||||
g.emit("// MarshalUnsafe%s is like %s.MarshalUnsafe, but for a []%s.\n", slice.ident, g.typeName(), g.typeName())
|
||||
g.emit("func MarshalUnsafe%s(src []%s, dst []byte) (int, error) {\n", slice.ident, g.typeName())
|
||||
g.inIndent(func() {
|
||||
g.emit("count := len(src)\n")
|
||||
g.emit("if count == 0 {\n")
|
||||
g.inIndent(func() {
|
||||
g.emit("return 0, nil\n")
|
||||
})
|
||||
g.emit("}\n")
|
||||
g.emit("size := (*%s)(nil).SizeBytes()\n\n", g.typeName())
|
||||
|
||||
g.emitNoEscapeSliceDataPointer("&src", "val")
|
||||
|
||||
g.emit("length, err := safecopy.CopyIn(dst[:(size*count)], val)\n")
|
||||
g.emitKeepAlive("src")
|
||||
g.emit("return length, err\n")
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
|
||||
g.emit("// UnmarshalUnsafe%s is like %s.UnmarshalUnsafe, but for a []%s.\n", slice.ident, g.typeName(), g.typeName())
|
||||
g.emit("func UnmarshalUnsafe%s(dst []%s, src []byte) (int, error) {\n", slice.ident, g.typeName())
|
||||
g.inIndent(func() {
|
||||
g.emit("count := len(dst)\n")
|
||||
g.emit("if count == 0 {\n")
|
||||
g.inIndent(func() {
|
||||
g.emit("return 0, nil\n")
|
||||
})
|
||||
g.emit("}\n")
|
||||
g.emit("size := (*%s)(nil).SizeBytes()\n\n", g.typeName())
|
||||
|
||||
g.emitNoEscapeSliceDataPointer("&dst", "val")
|
||||
|
||||
g.emit("length, err := safecopy.CopyOut(val, src[:(size*count)])\n")
|
||||
g.emitKeepAlive("dst")
|
||||
g.emit("return length, err\n")
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
}
|
||||
|
||||
@@ -72,20 +72,24 @@ func (g *interfaceGenerator) validateStruct(ts *ast.TypeSpec, st *ast.StructType
|
||||
})
|
||||
}
|
||||
|
||||
func (g *interfaceGenerator) emitMarshallableForStruct(st *ast.StructType) {
|
||||
// Is g.t a packed struct without consideing field types?
|
||||
thisPacked := true
|
||||
func (g *interfaceGenerator) isStructPacked(st *ast.StructType) bool {
|
||||
packed := true
|
||||
forEachStructField(st, func(f *ast.Field) {
|
||||
if f.Tag != nil {
|
||||
if f.Tag.Value == "`marshal:\"unaligned\"`" {
|
||||
if thisPacked {
|
||||
if packed {
|
||||
debugfAt(g.f.Position(g.t.Pos()),
|
||||
fmt.Sprintf("Marking type '%s' as not packed due to tag `marshal:\"unaligned\"`.\n", g.t.Name))
|
||||
thisPacked = false
|
||||
packed = false
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
return packed
|
||||
}
|
||||
|
||||
func (g *interfaceGenerator) emitMarshallableForStruct(st *ast.StructType) {
|
||||
thisPacked := g.isStructPacked(st)
|
||||
|
||||
g.emit("// SizeBytes implements marshal.Marshallable.SizeBytes.\n")
|
||||
g.emit("func (%s *%s) SizeBytes() int {\n", g.r, g.typeName())
|
||||
@@ -302,17 +306,16 @@ func (g *interfaceGenerator) emitMarshallableForStruct(st *ast.StructType) {
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
|
||||
g.emit("// CopyOut implements marshal.Marshallable.CopyOut.\n")
|
||||
g.emit("// CopyOutN implements marshal.Marshallable.CopyOutN.\n")
|
||||
g.recordUsedImport("marshal")
|
||||
g.recordUsedImport("usermem")
|
||||
g.emit("func (%s *%s) CopyOut(task marshal.Task, addr usermem.Addr) error {\n", g.r, g.typeName())
|
||||
g.emit("func (%s *%s) CopyOutN(task marshal.Task, addr usermem.Addr, limit int) (int, error) {\n", g.r, g.typeName())
|
||||
g.inIndent(func() {
|
||||
fallback := func() {
|
||||
g.emit("// Type %s doesn't have a packed layout in memory, fall back to MarshalBytes.\n", g.typeName())
|
||||
g.emit("buf := task.CopyScratchBuffer(%s.SizeBytes())\n", g.r)
|
||||
g.emit("%s.MarshalBytes(buf)\n", g.r)
|
||||
g.emit("_, err := task.CopyOutBytes(addr, buf)\n")
|
||||
g.emit("return err\n")
|
||||
g.emit("return task.CopyOutBytes(addr, buf[:limit])\n")
|
||||
}
|
||||
if thisPacked {
|
||||
g.recordUsedImport("reflect")
|
||||
@@ -324,48 +327,39 @@ func (g *interfaceGenerator) emitMarshallableForStruct(st *ast.StructType) {
|
||||
g.emit("}\n\n")
|
||||
}
|
||||
// Fast serialization.
|
||||
g.emit("// Bypass escape analysis on %s. The no-op arithmetic operation on the\n", g.r)
|
||||
g.emit("// pointer makes the compiler think val doesn't depend on %s.\n", g.r)
|
||||
g.emit("// See src/runtime/stubs.go:noescape() in the golang toolchain.\n")
|
||||
g.emit("ptr := unsafe.Pointer(%s)\n", g.r)
|
||||
g.emit("val := uintptr(ptr)\n")
|
||||
g.emit("val = val^0\n\n")
|
||||
g.emitCastToByteSlice(g.r, "buf", fmt.Sprintf("%s.SizeBytes()", g.r))
|
||||
|
||||
g.emit("// Construct a slice backed by %s's underlying memory.\n", g.r)
|
||||
g.emit("var buf []byte\n")
|
||||
g.emit("hdr := (*reflect.SliceHeader)(unsafe.Pointer(&buf))\n")
|
||||
g.emit("hdr.Data = val\n")
|
||||
g.emit("hdr.Len = %s.SizeBytes()\n", g.r)
|
||||
g.emit("hdr.Cap = %s.SizeBytes()\n\n", g.r)
|
||||
|
||||
g.emit("_, err := task.CopyOutBytes(addr, buf)\n")
|
||||
g.emit("// Since we bypassed the compiler's escape analysis, indicate that %s\n", g.r)
|
||||
g.emit("// must live until after the CopyOutBytes.\n")
|
||||
g.emit("runtime.KeepAlive(%s)\n", g.r)
|
||||
g.emit("return err\n")
|
||||
g.emit("length, err := task.CopyOutBytes(addr, buf[:limit])\n")
|
||||
g.emitKeepAlive(g.r)
|
||||
g.emit("return length, err\n")
|
||||
} else {
|
||||
fallback()
|
||||
}
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
|
||||
g.emit("// CopyOut implements marshal.Marshallable.CopyOut.\n")
|
||||
g.recordUsedImport("marshal")
|
||||
g.recordUsedImport("usermem")
|
||||
g.emit("func (%s *%s) CopyOut(task marshal.Task, addr usermem.Addr) (int, error) {\n", g.r, g.typeName())
|
||||
g.inIndent(func() {
|
||||
g.emit("return %s.CopyOutN(task, addr, %s.SizeBytes())\n", g.r, g.r)
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
|
||||
g.emit("// CopyIn implements marshal.Marshallable.CopyIn.\n")
|
||||
g.recordUsedImport("marshal")
|
||||
g.recordUsedImport("usermem")
|
||||
g.emit("func (%s *%s) CopyIn(task marshal.Task, addr usermem.Addr) error {\n", g.r, g.typeName())
|
||||
g.emit("func (%s *%s) CopyIn(task marshal.Task, addr usermem.Addr) (int, error) {\n", g.r, g.typeName())
|
||||
g.inIndent(func() {
|
||||
fallback := func() {
|
||||
g.emit("// Type %s doesn't have a packed layout in memory, fall back to UnmarshalBytes.\n", g.typeName())
|
||||
g.emit("buf := task.CopyScratchBuffer(%s.SizeBytes())\n", g.r)
|
||||
g.emit("_, err := task.CopyInBytes(addr, buf)\n")
|
||||
g.emit("if err != nil {\n")
|
||||
g.inIndent(func() {
|
||||
g.emit("return err\n")
|
||||
})
|
||||
g.emit("}\n")
|
||||
|
||||
g.emit("length, err := task.CopyInBytes(addr, buf)\n")
|
||||
g.emit("// Unmarshal unconditionally. If we had a short copy-in, this results in a\n")
|
||||
g.emit("// partially unmarshalled struct.\n")
|
||||
g.emit("%s.UnmarshalBytes(buf)\n", g.r)
|
||||
g.emit("return nil\n")
|
||||
g.emit("return length, err\n")
|
||||
}
|
||||
if thisPacked {
|
||||
g.recordUsedImport("reflect")
|
||||
@@ -377,25 +371,11 @@ func (g *interfaceGenerator) emitMarshallableForStruct(st *ast.StructType) {
|
||||
g.emit("}\n\n")
|
||||
}
|
||||
// Fast deserialization.
|
||||
g.emit("// Bypass escape analysis on %s. The no-op arithmetic operation on the\n", g.r)
|
||||
g.emit("// pointer makes the compiler think val doesn't depend on %s.\n", g.r)
|
||||
g.emit("// See src/runtime/stubs.go:noescape() in the golang toolchain.\n")
|
||||
g.emit("ptr := unsafe.Pointer(%s)\n", g.r)
|
||||
g.emit("val := uintptr(ptr)\n")
|
||||
g.emit("val = val^0\n\n")
|
||||
g.emitCastToByteSlice(g.r, "buf", fmt.Sprintf("%s.SizeBytes()", g.r))
|
||||
|
||||
g.emit("// Construct a slice backed by %s's underlying memory.\n", g.r)
|
||||
g.emit("var buf []byte\n")
|
||||
g.emit("hdr := (*reflect.SliceHeader)(unsafe.Pointer(&buf))\n")
|
||||
g.emit("hdr.Data = val\n")
|
||||
g.emit("hdr.Len = %s.SizeBytes()\n", g.r)
|
||||
g.emit("hdr.Cap = %s.SizeBytes()\n\n", g.r)
|
||||
|
||||
g.emit("_, err := task.CopyInBytes(addr, buf)\n")
|
||||
g.emit("// Since we bypassed the compiler's escape analysis, indicate that %s\n", g.r)
|
||||
g.emit("// must live until after the CopyInBytes.\n")
|
||||
g.emit("runtime.KeepAlive(%s)\n", g.r)
|
||||
g.emit("return err\n")
|
||||
g.emit("length, err := task.CopyInBytes(addr, buf)\n")
|
||||
g.emitKeepAlive(g.r)
|
||||
g.emit("return length, err\n")
|
||||
} else {
|
||||
fallback()
|
||||
}
|
||||
@@ -410,8 +390,8 @@ func (g *interfaceGenerator) emitMarshallableForStruct(st *ast.StructType) {
|
||||
g.emit("// Type %s doesn't have a packed layout in memory, fall back to MarshalBytes.\n", g.typeName())
|
||||
g.emit("buf := make([]byte, %s.SizeBytes())\n", g.r)
|
||||
g.emit("%s.MarshalBytes(buf)\n", g.r)
|
||||
g.emit("n, err := w.Write(buf)\n")
|
||||
g.emit("return int64(n), err\n")
|
||||
g.emit("length, err := w.Write(buf)\n")
|
||||
g.emit("return int64(length), err\n")
|
||||
}
|
||||
if thisPacked {
|
||||
g.recordUsedImport("reflect")
|
||||
@@ -423,25 +403,199 @@ func (g *interfaceGenerator) emitMarshallableForStruct(st *ast.StructType) {
|
||||
g.emit("}\n\n")
|
||||
}
|
||||
// Fast serialization.
|
||||
g.emit("// Bypass escape analysis on %s. The no-op arithmetic operation on the\n", g.r)
|
||||
g.emit("// pointer makes the compiler think val doesn't depend on %s.\n", g.r)
|
||||
g.emit("// See src/runtime/stubs.go:noescape() in the golang toolchain.\n")
|
||||
g.emit("ptr := unsafe.Pointer(%s)\n", g.r)
|
||||
g.emit("val := uintptr(ptr)\n")
|
||||
g.emit("val = val^0\n\n")
|
||||
g.emitCastToByteSlice(g.r, "buf", fmt.Sprintf("%s.SizeBytes()", g.r))
|
||||
|
||||
g.emit("// Construct a slice backed by %s's underlying memory.\n", g.r)
|
||||
g.emit("var buf []byte\n")
|
||||
g.emit("hdr := (*reflect.SliceHeader)(unsafe.Pointer(&buf))\n")
|
||||
g.emit("hdr.Data = val\n")
|
||||
g.emit("hdr.Len = %s.SizeBytes()\n", g.r)
|
||||
g.emit("hdr.Cap = %s.SizeBytes()\n\n", g.r)
|
||||
|
||||
g.emit("len, err := w.Write(buf)\n")
|
||||
g.emit("// Since we bypassed the compiler's escape analysis, indicate that %s\n", g.r)
|
||||
g.emit("// must live until after the Write.\n")
|
||||
g.emit("runtime.KeepAlive(%s)\n", g.r)
|
||||
g.emit("return int64(len), err\n")
|
||||
g.emit("length, err := w.Write(buf)\n")
|
||||
g.emitKeepAlive(g.r)
|
||||
g.emit("return int64(length), err\n")
|
||||
} else {
|
||||
fallback()
|
||||
}
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
}
|
||||
|
||||
func (g *interfaceGenerator) emitMarshallableSliceForStruct(st *ast.StructType, slice *sliceAPI) {
|
||||
thisPacked := g.isStructPacked(st)
|
||||
|
||||
if slice.inner {
|
||||
abortAt(g.f.Position(slice.comment.Slash), fmt.Sprintf("The ':inner' argument to '+marshal slice:%s:inner' is only applicable to newtypes on primitives. Remove it from this struct declaration.", slice.ident))
|
||||
}
|
||||
|
||||
g.recordUsedImport("marshal")
|
||||
g.recordUsedImport("usermem")
|
||||
|
||||
g.emit("// Copy%sIn copies in a slice of %s objects from the task's memory.\n", slice.ident, g.typeName())
|
||||
g.emit("func Copy%sIn(task marshal.Task, addr usermem.Addr, dst []%s) (int, error) {\n", slice.ident, g.typeName())
|
||||
g.inIndent(func() {
|
||||
g.emit("count := len(dst)\n")
|
||||
g.emit("if count == 0 {\n")
|
||||
g.inIndent(func() {
|
||||
g.emit("return 0, nil\n")
|
||||
})
|
||||
g.emit("}\n")
|
||||
g.emit("size := (*%s)(nil).SizeBytes()\n\n", g.typeName())
|
||||
|
||||
fallback := func() {
|
||||
g.emit("// Type %s doesn't have a packed layout in memory, fall back to UnmarshalBytes.\n", g.typeName())
|
||||
g.emit("buf := task.CopyScratchBuffer(size * count)\n")
|
||||
g.emit("length, err := task.CopyInBytes(addr, buf)\n\n")
|
||||
|
||||
g.emit("// Unmarshal as much as possible, even on error. First handle full objects.\n")
|
||||
g.emit("limit := length/size\n")
|
||||
g.emit("for idx := 0; idx < limit; idx++ {\n")
|
||||
g.inIndent(func() {
|
||||
g.emit("dst[idx].UnmarshalBytes(buf[size*idx:size*(idx+1)])\n")
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
|
||||
g.emit("// Handle any final partial object.\n")
|
||||
g.emit("if length < size*count && length%size != 0 {\n")
|
||||
g.inIndent(func() {
|
||||
g.emit("idx := limit\n")
|
||||
g.emit("dst[idx].UnmarshalBytes(buf[size*idx:size*(idx+1)])\n")
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
|
||||
g.emit("return length, err\n")
|
||||
}
|
||||
if thisPacked {
|
||||
g.recordUsedImport("reflect")
|
||||
g.recordUsedImport("runtime")
|
||||
g.recordUsedImport("unsafe")
|
||||
if _, ok := g.areFieldsPackedExpression(); ok {
|
||||
g.emit("if !dst[0].Packed() {\n")
|
||||
g.inIndent(fallback)
|
||||
g.emit("}\n\n")
|
||||
}
|
||||
// Fast deserialization.
|
||||
g.emitCastSliceToByteSlice("&dst", "buf", "size * count")
|
||||
|
||||
g.emit("length, err := task.CopyInBytes(addr, buf)\n")
|
||||
g.emitKeepAlive("dst")
|
||||
g.emit("return length, err\n")
|
||||
} else {
|
||||
fallback()
|
||||
}
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
|
||||
g.emit("// Copy%sOut copies a slice of %s objects to the task's memory.\n", slice.ident, g.typeName())
|
||||
g.emit("func Copy%sOut(task marshal.Task, addr usermem.Addr, src []%s) (int, error) {\n", slice.ident, g.typeName())
|
||||
g.inIndent(func() {
|
||||
g.emit("count := len(src)\n")
|
||||
g.emit("if count == 0 {\n")
|
||||
g.inIndent(func() {
|
||||
g.emit("return 0, nil\n")
|
||||
})
|
||||
g.emit("}\n")
|
||||
g.emit("size := (*%s)(nil).SizeBytes()\n\n", g.typeName())
|
||||
|
||||
fallback := func() {
|
||||
g.emit("// Type %s doesn't have a packed layout in memory, fall back to MarshalBytes.\n", g.typeName())
|
||||
g.emit("buf := task.CopyScratchBuffer(size * count)\n")
|
||||
g.emit("for idx := 0; idx < count; idx++ {\n")
|
||||
g.inIndent(func() {
|
||||
g.emit("src[idx].MarshalBytes(buf[size*idx:size*(idx+1)])\n")
|
||||
})
|
||||
g.emit("}\n")
|
||||
g.emit("return task.CopyOutBytes(addr, buf)\n")
|
||||
}
|
||||
if thisPacked {
|
||||
g.recordUsedImport("reflect")
|
||||
g.recordUsedImport("runtime")
|
||||
g.recordUsedImport("unsafe")
|
||||
if _, ok := g.areFieldsPackedExpression(); ok {
|
||||
g.emit("if !src[0].Packed() {\n")
|
||||
g.inIndent(fallback)
|
||||
g.emit("}\n\n")
|
||||
}
|
||||
// Fast serialization.
|
||||
g.emitCastSliceToByteSlice("&src", "buf", "size * count")
|
||||
|
||||
g.emit("length, err := task.CopyOutBytes(addr, buf)\n")
|
||||
g.emitKeepAlive("src")
|
||||
g.emit("return length, err\n")
|
||||
} else {
|
||||
fallback()
|
||||
}
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
|
||||
g.emit("// MarshalUnsafe%s is like %s.MarshalUnsafe, but for a []%s.\n", slice.ident, g.typeName(), g.typeName())
|
||||
g.emit("func MarshalUnsafe%s(src []%s, dst []byte) (int, error) {\n", slice.ident, g.typeName())
|
||||
g.inIndent(func() {
|
||||
g.emit("count := len(src)\n")
|
||||
g.emit("if count == 0 {\n")
|
||||
g.inIndent(func() {
|
||||
g.emit("return 0, nil\n")
|
||||
})
|
||||
g.emit("}\n")
|
||||
g.emit("size := (*%s)(nil).SizeBytes()\n\n", g.typeName())
|
||||
|
||||
fallback := func() {
|
||||
g.emit("// Type %s doesn't have a packed layout in memory, fall back to MarshalBytes.\n", g.typeName())
|
||||
g.emit("for idx := 0; idx < count; idx++ {\n")
|
||||
g.inIndent(func() {
|
||||
g.emit("src[idx].MarshalBytes(dst[size*idx:(size)*(idx+1)])\n")
|
||||
})
|
||||
g.emit("}\n")
|
||||
g.emit("return size * count, nil\n")
|
||||
}
|
||||
if thisPacked {
|
||||
g.recordUsedImport("reflect")
|
||||
g.recordUsedImport("runtime")
|
||||
g.recordUsedImport("unsafe")
|
||||
if _, ok := g.areFieldsPackedExpression(); ok {
|
||||
g.emit("if !src[0].Packed() {\n")
|
||||
g.inIndent(fallback)
|
||||
g.emit("}\n\n")
|
||||
}
|
||||
g.emitNoEscapeSliceDataPointer("&src", "val")
|
||||
|
||||
g.emit("length, err := safecopy.CopyIn(dst[:(size*count)], val)\n")
|
||||
g.emitKeepAlive("src")
|
||||
g.emit("return length, err\n")
|
||||
} else {
|
||||
fallback()
|
||||
}
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
|
||||
g.emit("// UnmarshalUnsafe%s is like %s.UnmarshalUnsafe, but for a []%s.\n", slice.ident, g.typeName(), g.typeName())
|
||||
g.emit("func UnmarshalUnsafe%s(dst []%s, src []byte) (int, error) {\n", slice.ident, g.typeName())
|
||||
g.inIndent(func() {
|
||||
g.emit("count := len(dst)\n")
|
||||
g.emit("if count == 0 {\n")
|
||||
g.inIndent(func() {
|
||||
g.emit("return 0, nil\n")
|
||||
})
|
||||
g.emit("}\n")
|
||||
g.emit("size := (*%s)(nil).SizeBytes()\n\n", g.typeName())
|
||||
|
||||
fallback := func() {
|
||||
g.emit("// Type %s doesn't have a packed layout in memory, fall back to UnmarshalBytes.\n", g.typeName())
|
||||
g.emit("for idx := 0; idx < count; idx++ {\n")
|
||||
g.inIndent(func() {
|
||||
g.emit("dst[idx].UnmarshalBytes(src[size*idx:size*(idx+1)])\n")
|
||||
})
|
||||
g.emit("}\n")
|
||||
g.emit("return size * count, nil\n")
|
||||
}
|
||||
if thisPacked {
|
||||
g.recordUsedImport("reflect")
|
||||
g.recordUsedImport("runtime")
|
||||
g.recordUsedImport("unsafe")
|
||||
if _, ok := g.areFieldsPackedExpression(); ok {
|
||||
g.emit("if !dst[0].Packed() {\n")
|
||||
g.inIndent(fallback)
|
||||
g.emit("}\n\n")
|
||||
}
|
||||
g.emitNoEscapeSliceDataPointer("&dst", "val")
|
||||
|
||||
g.emit("length, err := safecopy.CopyOut(val, src[:(size*count)])\n")
|
||||
g.emitKeepAlive("dst")
|
||||
g.emit("return length, err\n")
|
||||
} else {
|
||||
fallback()
|
||||
}
|
||||
|
||||
@@ -30,6 +30,11 @@ var standardImports = []string{
|
||||
"gvisor.dev/gvisor/tools/go_marshal/analysis",
|
||||
}
|
||||
|
||||
var sliceAPIImports = []string{
|
||||
"encoding/binary",
|
||||
"gvisor.dev/gvisor/pkg/usermem",
|
||||
}
|
||||
|
||||
type testGenerator struct {
|
||||
sourceBuffer
|
||||
|
||||
@@ -58,6 +63,11 @@ func newTestGenerator(t *ast.TypeSpec) *testGenerator {
|
||||
for _, i := range standardImports {
|
||||
g.imports.add(i).markUsed()
|
||||
}
|
||||
// These imports are used if a type requests the slice API. Don't
|
||||
// mark them as used by default.
|
||||
for _, i := range sliceAPIImports {
|
||||
g.imports.add(i)
|
||||
}
|
||||
|
||||
return g
|
||||
}
|
||||
@@ -132,6 +142,42 @@ func (g *testGenerator) emitTestMarshalUnmarshalPreservesData() {
|
||||
})
|
||||
}
|
||||
|
||||
func (g *testGenerator) emitTestMarshalUnmarshalSlicePreservesData(slice *sliceAPI) {
|
||||
for _, name := range []string{"binary", "usermem"} {
|
||||
if !g.imports.markUsed(name) {
|
||||
panic(fmt.Sprintf("Generated test for '%s' referenced a non-existent import with local name '%s'", g.typeName(), name))
|
||||
}
|
||||
}
|
||||
|
||||
g.inTestFunction("TestSafeMarshalUnmarshalSlicePreservesData", func() {
|
||||
g.emit("var x, y, yUnsafe [8]%s\n", g.typeName())
|
||||
g.emit("analysis.RandomizeValue(&x)\n\n")
|
||||
g.emit("size := (*%s)(nil).SizeBytes() * len(x)\n", g.typeName())
|
||||
g.emit("buf := bytes.NewBuffer(make([]byte, size))\n")
|
||||
g.emit("buf.Reset()\n")
|
||||
g.emit("if err := binary.Write(buf, usermem.ByteOrder, x[:]); err != nil {\n")
|
||||
g.inIndent(func() {
|
||||
g.emit("t.Fatal(fmt.Sprintf(\"binary.Write failed: %v\", err))\n")
|
||||
})
|
||||
g.emit("}\n")
|
||||
g.emit("bufUnsafe := make([]byte, size)\n")
|
||||
g.emit("MarshalUnsafe%s(x[:], bufUnsafe)\n\n", slice.ident)
|
||||
|
||||
g.emit("UnmarshalUnsafe%s(y[:], buf.Bytes())\n", slice.ident)
|
||||
g.emit("if !reflect.DeepEqual(x, y) {\n")
|
||||
g.inIndent(func() {
|
||||
g.emit("t.Fatal(fmt.Sprintf(\"Data corrupted across binary.Write/UnmarshalUnsafeSlice cycle:\\nBefore: %+v\\nAfter: %+v\\n\", x, y))\n")
|
||||
})
|
||||
g.emit("}\n")
|
||||
g.emit("UnmarshalUnsafe%s(yUnsafe[:], bufUnsafe)\n", slice.ident)
|
||||
g.emit("if !reflect.DeepEqual(x, yUnsafe) {\n")
|
||||
g.inIndent(func() {
|
||||
g.emit("t.Fatal(fmt.Sprintf(\"Data corrupted across MarshalUnsafeSlice/UnmarshalUnsafeSlice cycle:\\nBefore: %+v\\nAfter: %+v\\n\", x, yUnsafe))\n")
|
||||
})
|
||||
g.emit("}\n\n")
|
||||
})
|
||||
}
|
||||
|
||||
func (g *testGenerator) emitTestWriteToUnmarshalPreservesData() {
|
||||
g.inTestFunction("TestWriteToUnmarshalPreservesData", func() {
|
||||
g.emit("var x, y, yUnsafe %s\n", g.typeName())
|
||||
@@ -170,12 +216,16 @@ func (g *testGenerator) emitTestSizeBytesOnTypedNilPtr() {
|
||||
})
|
||||
}
|
||||
|
||||
func (g *testGenerator) emitTests() {
|
||||
func (g *testGenerator) emitTests(slice *sliceAPI) {
|
||||
g.emitTestNonZeroSize()
|
||||
g.emitTestSuspectAlignment()
|
||||
g.emitTestMarshalUnmarshalPreservesData()
|
||||
g.emitTestWriteToUnmarshalPreservesData()
|
||||
g.emitTestSizeBytesOnTypedNilPtr()
|
||||
|
||||
if slice != nil {
|
||||
g.emitTestMarshalUnmarshalSlicePreservesData(slice)
|
||||
}
|
||||
}
|
||||
|
||||
func (g *testGenerator) write(out io.Writer) error {
|
||||
|
||||
@@ -344,22 +344,25 @@ func newImportTable() *importTable {
|
||||
// result in a panic.
|
||||
func (i *importTable) merge(other *importTable) {
|
||||
for name, im := range other.is {
|
||||
if dup, ok := i.is[name]; ok && !dup.equivalent(im) {
|
||||
panic(fmt.Sprintf("Found colliding import statements: ours: %+v, other's: %+v", dup, im))
|
||||
}
|
||||
dup, ok := i.is[name]
|
||||
if ok {
|
||||
// When merging two imports, if either are marked used, the merged entry
|
||||
// should also be marked used.
|
||||
im.used = im.used || dup.used
|
||||
|
||||
if !dup.equivalent(im) {
|
||||
panic(fmt.Sprintf("Found colliding import statements: ours: %+v, other's: %+v", dup, im))
|
||||
}
|
||||
}
|
||||
i.is[name] = im
|
||||
}
|
||||
}
|
||||
|
||||
func (i *importTable) addStmt(s *importStmt) *importStmt {
|
||||
if old, ok := i.is[s.name]; ok && !old.equivalent(s) {
|
||||
// A collision should always be between an import inserted by the
|
||||
// go-marshal tool and an import from the original source file (assuming
|
||||
// the original source file was valid). We could theoretically handle
|
||||
// the collision by assigning a local name to our import. However, this
|
||||
// would need to be plumbed throughout the generator. Given that
|
||||
// collisions should be rare, simply panic on collision.
|
||||
// We could theoretically handle the collision by assigning a local name
|
||||
// to one of the imports. However, this is a non-trivial transformation.
|
||||
// Given that collisions should be rare, simply panic on collision.
|
||||
panic(fmt.Sprintf("Import collision: old: %s as %v; new: %v as %v", old.path, old.name, s.path, s.name))
|
||||
}
|
||||
i.is[s.name] = s
|
||||
|
||||
@@ -42,7 +42,11 @@ type Task interface {
|
||||
CopyInBytes(addr usermem.Addr, b []byte) (int, error)
|
||||
}
|
||||
|
||||
// Marshallable represents a type that can be marshalled to and from memory.
|
||||
// Marshallable represents operations on a type that can be marshalled to and
|
||||
// from memory.
|
||||
//
|
||||
// go-marshal automatically generates implementations for this interface for
|
||||
// types marked as '+marshal'.
|
||||
type Marshallable interface {
|
||||
io.WriterTo
|
||||
|
||||
@@ -54,12 +58,18 @@ type Marshallable interface {
|
||||
// likely make use of the type of these fields).
|
||||
SizeBytes() int
|
||||
|
||||
// MarshalBytes serializes a copy of a type to dst. dst must be at least
|
||||
// SizeBytes() long.
|
||||
// MarshalBytes serializes a copy of a type to dst. dst may be smaller than
|
||||
// SizeBytes(), which results in a part of the struct being marshalled. Note
|
||||
// that this may have unexpected results for non-packed types, as implicit
|
||||
// padding needs to be taken into account when reasoning about how much of
|
||||
// the type is serialized.
|
||||
MarshalBytes(dst []byte)
|
||||
|
||||
// UnmarshalBytes deserializes a type from src. src must be at least
|
||||
// SizeBytes() long.
|
||||
// UnmarshalBytes deserializes a type from src. src may be smaller than
|
||||
// SizeBytes(), which results in a partially deserialized struct. Note that
|
||||
// this may have unexpected results for non-packed types, as implicit
|
||||
// padding needs to be taken into account when reasoning about how much of
|
||||
// the type is deserialized.
|
||||
UnmarshalBytes(src []byte)
|
||||
|
||||
// Packed returns true if the marshalled size of the type is the same as the
|
||||
@@ -67,13 +77,20 @@ type Marshallable interface {
|
||||
// starting at unaligned addresses (should always be true by default for ABI
|
||||
// structs, verified by automatically generated tests when using
|
||||
// go_marshal), and has no fields marked `marshal:"unaligned"`.
|
||||
//
|
||||
// Packed must return the same result for all possible values of the type
|
||||
// implementing it. Violating this constraint implies the type doesn't have
|
||||
// a static memory layout, and will lead to memory corruption.
|
||||
// Go-marshal-generated code reuses the result of Packed for multiple values
|
||||
// of the same type.
|
||||
Packed() bool
|
||||
|
||||
// MarshalUnsafe serializes a type by bulk copying its in-memory
|
||||
// representation to the dst buffer. This is only safe to do when the type
|
||||
// has no implicit padding, see Marshallable.Packed. When Packed would
|
||||
// return false, MarshalUnsafe should fall back to the safer but slower
|
||||
// MarshalBytes.
|
||||
// MarshalBytes. dst may be smaller than SizeBytes(), see comment for
|
||||
// MarshalBytes for implications.
|
||||
MarshalUnsafe(dst []byte)
|
||||
|
||||
// UnmarshalUnsafe deserializes a type by directly copying to the underlying
|
||||
@@ -82,7 +99,8 @@ type Marshallable interface {
|
||||
// This allows much faster unmarshalling of types which have no implicit
|
||||
// padding, see Marshallable.Packed. When Packed would return false,
|
||||
// UnmarshalUnsafe should fall back to the safer but slower unmarshal
|
||||
// mechanism implemented in UnmarshalBytes.
|
||||
// mechanism implemented in UnmarshalBytes. src may be smaller than
|
||||
// SizeBytes(), see comment for UnmarshalBytes for implications.
|
||||
UnmarshalUnsafe(src []byte)
|
||||
|
||||
// CopyIn deserializes a Marshallable type from a task's memory. This may
|
||||
@@ -91,12 +109,79 @@ type Marshallable interface {
|
||||
// marshalled does not escape. The implementation should avoid creating
|
||||
// extra copies in memory by directly deserializing to the object's
|
||||
// underlying memory.
|
||||
CopyIn(task Task, addr usermem.Addr) error
|
||||
//
|
||||
// If the copy-in from the task memory is only partially successful, CopyIn
|
||||
// should still attempt to deserialize as much data as possible. See comment
|
||||
// for UnmarshalBytes.
|
||||
CopyIn(task Task, addr usermem.Addr) (int, error)
|
||||
|
||||
// CopyOut serializes a Marshallable type to a task's memory. This may only
|
||||
// be called from a task goroutine. This is more efficient than calling
|
||||
// MarshalUnsafe on Marshallable.Packed types, as the type being serialized
|
||||
// does not escape. The implementation should avoid creating extra copies in
|
||||
// memory by directly serializing from the object's underlying memory.
|
||||
CopyOut(task Task, addr usermem.Addr) error
|
||||
//
|
||||
// The copy-out to the task memory may be partially successful, in which
|
||||
// case CopyOut returns how much data was serialized. See comment for
|
||||
// MarshalBytes for implications.
|
||||
CopyOut(task Task, addr usermem.Addr) (int, error)
|
||||
|
||||
// CopyOutN is like CopyOut, but explicitly requests a partial
|
||||
// copy-out. Note that this may yield unexpected results for non-packed
|
||||
// types and the caller may only want to allow this for packed types. See
|
||||
// comment on MarshalBytes.
|
||||
//
|
||||
// The limit must be less than or equal to SizeBytes().
|
||||
CopyOutN(task Task, addr usermem.Addr, limit int) (int, error)
|
||||
}
|
||||
|
||||
// go-marshal generates additional functions for a type based on additional
|
||||
// clauses to the +marshal directive. They are documented below.
|
||||
//
|
||||
// Slice API
|
||||
// =========
|
||||
//
|
||||
// Adding a "slice" clause to the +marshal directive for structs or newtypes on
|
||||
// primitives like this:
|
||||
//
|
||||
// // +marshal slice:FooSlice
|
||||
// type Foo struct { ... }
|
||||
//
|
||||
// Generates four additional functions for marshalling slices of Foos like this:
|
||||
//
|
||||
// // MarshalUnsafeFooSlice is like Foo.MarshalUnsafe, buf for a []Foo. It's
|
||||
// // more efficient that repeatedly calling calling Foo.MarshalUnsafe over a
|
||||
// // []Foo in a loop.
|
||||
// func MarshalUnsafeFooSlice(src []Foo, dst []byte) (int, error) { ... }
|
||||
//
|
||||
// // UnmarshalUnsafeFooSlice is like Foo.UnmarshalUnsafe, buf for a []Foo. It's
|
||||
// // more efficient that repeatedly calling calling Foo.UnmarshalUnsafe over a
|
||||
// // []Foo in a loop.
|
||||
// func UnmarshalUnsafeFooSlice(dst []Foo, src []byte) (int, error) { ... }
|
||||
//
|
||||
// // CopyFooSliceIn copies in a slice of Foo objects from the task's memory.
|
||||
// func CopyFooSliceIn(task marshal.Task, addr usermem.Addr, dst []Foo) (int, error) { ... }
|
||||
//
|
||||
// // CopyFooSliceIn copies out a slice of Foo objects to the task's memory.
|
||||
// func CopyFooSliceOut(task marshal.Task, addr usermem.Addr, src []Foo) (int, error) { ... }
|
||||
//
|
||||
// The name of the functions are of the format "Copy%sIn" and "Copy%sOut", where
|
||||
// %s is the first argument to the slice clause. This directive is not supported
|
||||
// for newtypes on arrays.
|
||||
//
|
||||
// The slice clause also takes an optional second argument, which must be the
|
||||
// value "inner":
|
||||
//
|
||||
// // +marshal slice:Int32Slice:inner
|
||||
// type Int32 int32
|
||||
//
|
||||
// This is only valid on newtypes on primitives, and causes the generated
|
||||
// functions to accept slices of the inner type instead:
|
||||
//
|
||||
// func CopyInt32SliceIn(task marshal.Task, addr usermem.Addr, dst []int32) (int, error) { ... }
|
||||
//
|
||||
// Without "inner", they would instead be:
|
||||
//
|
||||
// func CopyInt32SliceIn(task marshal.Task, addr usermem.Addr, dst []Int32) (int, error) { ... }
|
||||
//
|
||||
// This may help avoid a cast depending on how the generated functions are used.
|
||||
|
||||
@@ -0,0 +1,18 @@
|
||||
load("//tools:defs.bzl", "go_library")
|
||||
|
||||
licenses(["notice"])
|
||||
|
||||
go_library(
|
||||
name = "primitive",
|
||||
srcs = [
|
||||
"primitive.go",
|
||||
],
|
||||
marshal = True,
|
||||
visibility = [
|
||||
"//:sandbox",
|
||||
],
|
||||
deps = [
|
||||
"//pkg/usermem",
|
||||
"//tools/go_marshal/marshal",
|
||||
],
|
||||
)
|
||||
@@ -0,0 +1,175 @@
|
||||
// Copyright 2020 The gVisor Authors.
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
// Package primitive defines marshal.Marshallable implementations for primitive
|
||||
// types.
|
||||
package primitive
|
||||
|
||||
import (
|
||||
"gvisor.dev/gvisor/pkg/usermem"
|
||||
"gvisor.dev/gvisor/tools/go_marshal/marshal"
|
||||
)
|
||||
|
||||
// Int16 is a marshal.Marshallable implementation for int16.
|
||||
//
|
||||
// +marshal slice:Int16Slice:inner
|
||||
type Int16 int16
|
||||
|
||||
// Uint16 is a marshal.Marshallable implementation for uint16.
|
||||
//
|
||||
// +marshal slice:Uint16Slice:inner
|
||||
type Uint16 uint16
|
||||
|
||||
// Int32 is a marshal.Marshallable implementation for int32.
|
||||
//
|
||||
// +marshal slice:Int32Slice:inner
|
||||
type Int32 int32
|
||||
|
||||
// Uint32 is a marshal.Marshallable implementation for uint32.
|
||||
//
|
||||
// +marshal slice:Uint32Slice:inner
|
||||
type Uint32 uint32
|
||||
|
||||
// Int64 is a marshal.Marshallable implementation for int64.
|
||||
//
|
||||
// +marshal slice:Int64Slice:inner
|
||||
type Int64 int64
|
||||
|
||||
// Uint64 is a marshal.Marshallable implementation for uint64.
|
||||
//
|
||||
// +marshal slice:Uint64Slice:inner
|
||||
type Uint64 uint64
|
||||
|
||||
// Below, we define some convenience functions for marshalling primitive types
|
||||
// using the newtypes above, without requiring superfluous casts.
|
||||
|
||||
// 16-bit integers
|
||||
|
||||
// CopyInt16In is a convenient wrapper for copying in an int16 from the task's
|
||||
// memory.
|
||||
func CopyInt16In(task marshal.Task, addr usermem.Addr, dst *int16) (int, error) {
|
||||
var buf Int16
|
||||
n, err := buf.CopyIn(task, addr)
|
||||
if err != nil {
|
||||
return n, err
|
||||
}
|
||||
*dst = int16(buf)
|
||||
return n, nil
|
||||
}
|
||||
|
||||
// CopyInt16Out is a convenient wrapper for copying out an int16 to the task's
|
||||
// memory.
|
||||
func CopyInt16Out(task marshal.Task, addr usermem.Addr, src int16) (int, error) {
|
||||
srcP := Int16(src)
|
||||
return srcP.CopyOut(task, addr)
|
||||
}
|
||||
|
||||
// CopyUint16In is a convenient wrapper for copying in a uint16 from the task's
|
||||
// memory.
|
||||
func CopyUint16In(task marshal.Task, addr usermem.Addr, dst *uint16) (int, error) {
|
||||
var buf Uint16
|
||||
n, err := buf.CopyIn(task, addr)
|
||||
if err != nil {
|
||||
return n, err
|
||||
}
|
||||
*dst = uint16(buf)
|
||||
return n, nil
|
||||
}
|
||||
|
||||
// CopyUint16Out is a convenient wrapper for copying out a uint16 to the task's
|
||||
// memory.
|
||||
func CopyUint16Out(task marshal.Task, addr usermem.Addr, src uint16) (int, error) {
|
||||
srcP := Uint16(src)
|
||||
return srcP.CopyOut(task, addr)
|
||||
}
|
||||
|
||||
// 32-bit integers
|
||||
|
||||
// CopyInt32In is a convenient wrapper for copying in an int32 from the task's
|
||||
// memory.
|
||||
func CopyInt32In(task marshal.Task, addr usermem.Addr, dst *int32) (int, error) {
|
||||
var buf Int32
|
||||
n, err := buf.CopyIn(task, addr)
|
||||
if err != nil {
|
||||
return n, err
|
||||
}
|
||||
*dst = int32(buf)
|
||||
return n, nil
|
||||
}
|
||||
|
||||
// CopyInt32Out is a convenient wrapper for copying out an int32 to the task's
|
||||
// memory.
|
||||
func CopyInt32Out(task marshal.Task, addr usermem.Addr, src int32) (int, error) {
|
||||
srcP := Int32(src)
|
||||
return srcP.CopyOut(task, addr)
|
||||
}
|
||||
|
||||
// CopyUint32In is a convenient wrapper for copying in a uint32 from the task's
|
||||
// memory.
|
||||
func CopyUint32In(task marshal.Task, addr usermem.Addr, dst *uint32) (int, error) {
|
||||
var buf Uint32
|
||||
n, err := buf.CopyIn(task, addr)
|
||||
if err != nil {
|
||||
return n, err
|
||||
}
|
||||
*dst = uint32(buf)
|
||||
return n, nil
|
||||
}
|
||||
|
||||
// CopyUint32Out is a convenient wrapper for copying out a uint32 to the task's
|
||||
// memory.
|
||||
func CopyUint32Out(task marshal.Task, addr usermem.Addr, src uint32) (int, error) {
|
||||
srcP := Uint32(src)
|
||||
return srcP.CopyOut(task, addr)
|
||||
}
|
||||
|
||||
// 64-bit integers
|
||||
|
||||
// CopyInt64In is a convenient wrapper for copying in an int64 from the task's
|
||||
// memory.
|
||||
func CopyInt64In(task marshal.Task, addr usermem.Addr, dst *int64) (int, error) {
|
||||
var buf Int64
|
||||
n, err := buf.CopyIn(task, addr)
|
||||
if err != nil {
|
||||
return n, err
|
||||
}
|
||||
*dst = int64(buf)
|
||||
return n, nil
|
||||
}
|
||||
|
||||
// CopyInt64Out is a convenient wrapper for copying out an int64 to the task's
|
||||
// memory.
|
||||
func CopyInt64Out(task marshal.Task, addr usermem.Addr, src int64) (int, error) {
|
||||
srcP := Int64(src)
|
||||
return srcP.CopyOut(task, addr)
|
||||
}
|
||||
|
||||
// CopyUint64In is a convenient wrapper for copying in a uint64 from the task's
|
||||
// memory.
|
||||
func CopyUint64In(task marshal.Task, addr usermem.Addr, dst *uint64) (int, error) {
|
||||
var buf Uint64
|
||||
n, err := buf.CopyIn(task, addr)
|
||||
if err != nil {
|
||||
return n, err
|
||||
}
|
||||
*dst = uint64(buf)
|
||||
return n, nil
|
||||
}
|
||||
|
||||
// CopyUint64Out is a convenient wrapper for copying out a uint64 to the task's
|
||||
// memory.
|
||||
func CopyUint64Out(task marshal.Task, addr usermem.Addr, src uint64) (int, error) {
|
||||
srcP := Uint64(src)
|
||||
return srcP.CopyOut(task, addr)
|
||||
}
|
||||
@@ -39,3 +39,17 @@ go_binary(
|
||||
"//tools/go_marshal/marshal",
|
||||
],
|
||||
)
|
||||
|
||||
go_test(
|
||||
name = "marshal_test",
|
||||
size = "small",
|
||||
srcs = ["marshal_test.go"],
|
||||
deps = [
|
||||
":test",
|
||||
"//pkg/syserror",
|
||||
"//pkg/usermem",
|
||||
"//tools/go_marshal/analysis",
|
||||
"//tools/go_marshal/marshal",
|
||||
"@com_github_google_go-cmp//cmp:go_default_library",
|
||||
],
|
||||
)
|
||||
|
||||
@@ -176,3 +176,45 @@ func BenchmarkGoMarshalUnsafe(b *testing.B) {
|
||||
panic(fmt.Sprintf("Data corruption across marshal/unmarshal cycle:\nBefore: %+v\nAfter: %+v\n", s1, s2))
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkBinarySlice(b *testing.B) {
|
||||
var s1, s2 [64]test.Stat
|
||||
analysis.RandomizeValue(&s1)
|
||||
|
||||
size := binary.Size(s1)
|
||||
|
||||
b.ResetTimer()
|
||||
|
||||
for n := 0; n < b.N; n++ {
|
||||
buf := make([]byte, 0, size)
|
||||
buf = binary.Marshal(buf, usermem.ByteOrder, &s1)
|
||||
binary.Unmarshal(buf, usermem.ByteOrder, &s2)
|
||||
}
|
||||
|
||||
b.StopTimer()
|
||||
|
||||
// Sanity check, make sure the values were preserved.
|
||||
if !reflect.DeepEqual(s1, s2) {
|
||||
panic(fmt.Sprintf("Data corruption across marshal/unmarshal cycle:\nBefore: %+v\nAfter: %+v\n", s1, s2))
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkGoMarshalUnsafeSlice(b *testing.B) {
|
||||
var s1, s2 [64]test.Stat
|
||||
analysis.RandomizeValue(&s1)
|
||||
|
||||
b.ResetTimer()
|
||||
|
||||
for n := 0; n < b.N; n++ {
|
||||
buf := make([]byte, (*test.Stat)(nil).SizeBytes()*len(s1))
|
||||
test.MarshalUnsafeStatSlice(s1[:], buf)
|
||||
test.UnmarshalUnsafeStatSlice(s2[:], buf)
|
||||
}
|
||||
|
||||
b.StopTimer()
|
||||
|
||||
// Sanity check, make sure the values were preserved.
|
||||
if !reflect.DeepEqual(s1, s2) {
|
||||
panic(fmt.Sprintf("Data corruption across marshal/unmarshal cycle:\nBefore: %+v\nAfter: %+v\n", s1, s2))
|
||||
}
|
||||
}
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user