Source file src/runtime/_mkmalloc/mkmalloc.go

     1  // Copyright 2025 The Go Authors. All rights reserved.
     2  // Use of this source code is governed by a BSD-style
     3  // license that can be found in the LICENSE file.
     4  
     5  package main
     6  
     7  import (
     8  	"bytes"
     9  	"flag"
    10  	"fmt"
    11  	"go/ast"
    12  	"go/format"
    13  	"go/parser"
    14  	"go/token"
    15  	"log"
    16  	"os"
    17  	"strings"
    18  
    19  	"golang.org/x/tools/go/ast/astutil"
    20  
    21  	internalastutil "runtime/_mkmalloc/astutil"
    22  )
    23  
    24  var stdout = flag.Bool("stdout", false, "write sizeclasses source to stdout instead of sizeclasses.go")
    25  
    26  func makeSizeToSizeClass(classes []class) []uint8 {
    27  	sc := uint8(0)
    28  	ret := make([]uint8, benchmarkMax+1)
    29  	for i := range ret {
    30  		if i > classes[sc].size {
    31  			sc++
    32  		}
    33  		ret[i] = sc
    34  	}
    35  	return ret
    36  }
    37  
    38  func main() {
    39  	log.SetFlags(0)
    40  	log.SetPrefix("mkmalloc: ")
    41  
    42  	classes := makeClasses()
    43  	sizeToSizeClass := makeSizeToSizeClass(classes)
    44  
    45  	if *stdout {
    46  		if _, err := os.Stdout.Write(mustFormat(generateSizeClasses(classes))); err != nil {
    47  			log.Fatal(err)
    48  		}
    49  		return
    50  	}
    51  
    52  	sizeclasesesfile := "../../internal/runtime/gc/sizeclasses.go"
    53  	if err := os.WriteFile(sizeclasesesfile, mustFormat(generateSizeClasses(classes)), 0666); err != nil {
    54  		log.Fatal(err)
    55  	}
    56  
    57  	outfile := "../malloc_generated.go"
    58  	if err := os.WriteFile(outfile, mustFormat(inline(specializedMallocConfig(classes, sizeToSizeClass))), 0666); err != nil {
    59  		log.Fatal(err)
    60  	}
    61  
    62  	tablefile := "../malloc_tables_generated.go"
    63  	if err := os.WriteFile(tablefile, mustFormat(generateTable(sizeToSizeClass)), 0666); err != nil {
    64  		log.Fatal(err)
    65  	}
    66  
    67  	benchmarkFile := "../malloc_bench_generated_test.go"
    68  	if err := os.WriteFile(benchmarkFile, mustFormat(append(inline(benchmarkConfig(classes, sizeToSizeClass)), []byte(generateTopBenchmark(classes, sizeToSizeClass))...)), 0666); err != nil {
    69  		log.Fatal(err)
    70  	}
    71  
    72  }
    73  
    74  // withLineNumbers returns b with line numbers added to help debugging.
    75  func withLineNumbers(b []byte) []byte {
    76  	var buf bytes.Buffer
    77  	i := 1
    78  	for line := range bytes.Lines(b) {
    79  		fmt.Fprintf(&buf, "%d: %s", i, line)
    80  		i++
    81  	}
    82  	return buf.Bytes()
    83  }
    84  
    85  // mustFormat formats the input source, or exits if there's an error.
    86  func mustFormat(b []byte) []byte {
    87  	formatted, err := format.Source(b)
    88  	if err != nil {
    89  		log.Fatalf("error formatting source: %v\nsource:\n%s\n", err, withLineNumbers(b))
    90  	}
    91  	return formatted
    92  }
    93  
    94  // generatorConfig is the configuration for the generator. It uses the given file to find
    95  // its templates, and generates each of the functions specified by specs.
    96  type generatorConfig struct {
    97  	file  string
    98  	specs []spec
    99  }
   100  
   101  // spec is the specification for a function for the inliner to produce. The function gets
   102  // the given name, and is produced by starting with the function with the name given by
   103  // templateFunc and applying each of the ops.
   104  type spec struct {
   105  	name         string
   106  	templateFunc string
   107  	ops          []op
   108  }
   109  
   110  // replacementKind specifies the operation to ben done by a op.
   111  type replacementKind int
   112  
   113  const (
   114  	inlineFunc = replacementKind(iota)
   115  	subBasicLit
   116  	foldCondition
   117  	subIdent
   118  	deleteConst
   119  )
   120  
   121  // op is a single inlining operation for the inliner. Any calls to the function
   122  // from are replaced with the inlined body of to. For non-functions, uses of from are
   123  // replaced with the basic literal expression given by to.
   124  type op struct {
   125  	kind replacementKind
   126  	from string
   127  	to   string
   128  }
   129  
   130  func smallScanNoHeaderSCFuncName(sc, scMax uint8) string {
   131  	if sc == 0 || sc > scMax {
   132  		return "mallocPanic"
   133  	}
   134  	return fmt.Sprintf("mallocgcSmallScanNoHeaderSC%d", sc)
   135  }
   136  
   137  const tinyFuncName = "mallocgcTinySC2"
   138  
   139  func smallNoScanSCFuncName(sc, scMax uint8) string {
   140  	if sc < 2 || sc > scMax {
   141  		return "mallocPanic"
   142  	}
   143  	return fmt.Sprintf("mallocgcSmallNoScanSC%d", sc)
   144  }
   145  
   146  // specializedMallocConfig produces an inlining config to stamp out the definitions of the size-specialized
   147  // malloc functions to be written by mkmalloc.
   148  func specializedMallocConfig(classes []class, sizeToSizeClass []uint8) generatorConfig {
   149  	config := generatorConfig{file: "../malloc_stubs.go"}
   150  
   151  	// Only generate specialized functions for sizes up to specializedMallocMax.
   152  	// We've noticed limited benefit (or sometimes worse performance) for specialized
   153  	// functions for larger sizes, and having too many functions causes icache issues.
   154  	scMax := sizeToSizeClass[specializedMallocMax]
   155  
   156  	str := fmt.Sprint
   157  
   158  	// allocations with pointer bits
   159  	{
   160  		const noscan = 0
   161  		for sc := uint8(0); sc <= scMax; sc++ {
   162  			if sc == 0 {
   163  				continue
   164  			}
   165  			name := smallScanNoHeaderSCFuncName(sc, scMax)
   166  			elemsize := classes[sc].size
   167  			config.specs = append(config.specs, spec{
   168  				templateFunc: "mallocStub",
   169  				name:         name,
   170  				ops: []op{
   171  					{inlineFunc, "inlinedMalloc", "smallStub"},
   172  					{inlineFunc, "postMallocgc", "postMallocgc"},
   173  					{foldCondition, "isNoScan_", str(false)},
   174  					{inlineFunc, "heapSetTypeNoHeaderStub", "heapSetTypeNoHeaderStub"},
   175  					{inlineFunc, "nextFreeFastStub", "nextFreeFastStub"},
   176  					{inlineFunc, "writeHeapBitsSmallStub", "writeHeapBitsSmallStub"},
   177  					{foldCondition, "isSlowPath_", str(false)},
   178  					{subBasicLit, "elemsize_", str(elemsize)},
   179  					{subBasicLit, "sizeclass_", str(sc)},
   180  					{subBasicLit, "noscanint_", str(noscan)},
   181  					{foldCondition, "isTiny_", str(false)},
   182  					{subIdent, "mallocgcSlowPathStub", "mallocgcSmallScanSlowPath"},
   183  				},
   184  			})
   185  		}
   186  	}
   187  
   188  	// allocations without pointer bits
   189  	{
   190  		const noscan = 1
   191  
   192  		// tiny
   193  		tinySizeClass := sizeToSizeClass[tinySize]
   194  		{
   195  			name := tinyFuncName
   196  			elemsize := classes[tinySizeClass].size
   197  			config.specs = append(config.specs, spec{
   198  				templateFunc: "mallocStub",
   199  				name:         name,
   200  				ops: []op{
   201  					{inlineFunc, "inlinedMalloc", "tinyStub"},
   202  					{inlineFunc, "nextFreeFastTiny", "nextFreeFastTiny"},
   203  					{inlineFunc, "postMallocgc", "postMallocgc"},
   204  					{inlineFunc, "nextFreeFastStub", "nextFreeFastStub"},
   205  					{foldCondition, "isSlowPath_", str(false)},
   206  					{subBasicLit, "elemsize_", str(elemsize)},
   207  					{subBasicLit, "sizeclass_", str(tinySizeClass)},
   208  					{subBasicLit, "noscanint_", str(noscan)},
   209  					{foldCondition, "isTiny_", str(true)},
   210  				},
   211  			})
   212  		}
   213  
   214  		// non-tiny
   215  		for sc := uint8(tinySizeClass); sc <= scMax; sc++ {
   216  			name := smallNoScanSCFuncName(sc, scMax)
   217  			elemsize := classes[sc].size
   218  			config.specs = append(config.specs, spec{
   219  				templateFunc: "mallocStub",
   220  				name:         name,
   221  				ops: []op{
   222  					{inlineFunc, "inlinedMalloc", "smallStub"},
   223  					{inlineFunc, "postMallocgc", "postMallocgc"},
   224  					{foldCondition, "isNoScan_", str(true)},
   225  					{inlineFunc, "nextFreeFastStub", "nextFreeFastStub"},
   226  					{foldCondition, "isSlowPath_", str(false)},
   227  					{subBasicLit, "elemsize_", str(elemsize)},
   228  					{subBasicLit, "sizeclass_", str(sc)},
   229  					{subBasicLit, "noscanint_", str(noscan)},
   230  					{foldCondition, "isTiny_", str(false)},
   231  					{subIdent, "mallocgcSlowPathStub", "mallocgcSmallNoScanSlowPath"},
   232  				},
   233  			})
   234  		}
   235  	}
   236  
   237  	// Non-size-specialized fallbacks in case we can't do the fast path.
   238  	config.specs = append(config.specs, spec{
   239  		templateFunc: "mallocStub",
   240  		name:         "mallocgcTinySlowPath",
   241  		ops: []op{
   242  			{inlineFunc, "inlinedMalloc", "tinyStub"},
   243  			{inlineFunc, "postMallocgc", "postMallocgc"},
   244  			{inlineFunc, "nextFreeFastTiny", "nextFreeFastTiny"},
   245  			{inlineFunc, "deductAssistCredit", "deductAssistCredit"},
   246  			{foldCondition, "isSlowPath_", str(true)},
   247  			{foldCondition, "isTiny_", str(true)},
   248  			{subBasicLit, "elemsize_", str(classes[sizeToSizeClass[tinySize]].size)},
   249  		},
   250  	})
   251  	config.specs = append(config.specs, spec{
   252  		templateFunc: "mallocgcSlowPathStub",
   253  		name:         "mallocgcSmallScanSlowPath",
   254  		ops: []op{
   255  			{inlineFunc, "mallocStub", "mallocStub"},
   256  			{inlineFunc, "inlinedMalloc", "smallStub"},
   257  			{inlineFunc, "heapSetTypeNoHeaderStub", "heapSetTypeNoHeaderStub"},
   258  			{inlineFunc, "writeHeapBitsSmallStub", "writeHeapBitsSmallStub"},
   259  			{inlineFunc, "postMallocgc", "postMallocgc"},
   260  			{inlineFunc, "nextFreeFastStub", "nextFreeFastStub"},
   261  			{inlineFunc, "deductAssistCredit", "deductAssistCredit"},
   262  			{foldCondition, "isSlowPath_", str(true)},
   263  			{foldCondition, "isTiny_", str(false)},
   264  			{foldCondition, "isNoScan_", str(false)},
   265  
   266  			// Remove constants used by size-specialized variants.
   267  			{deleteConst, "elemsize", ""},
   268  			{deleteConst, "sizeclass", ""},
   269  			{deleteConst, "spc", ""},
   270  		},
   271  	})
   272  	config.specs = append(config.specs, spec{
   273  		templateFunc: "mallocgcSlowPathStub",
   274  		name:         "mallocgcSmallNoScanSlowPath",
   275  		ops: []op{
   276  			{inlineFunc, "mallocStub", "mallocStub"},
   277  			{inlineFunc, "inlinedMalloc", "smallStub"},
   278  			{inlineFunc, "postMallocgc", "postMallocgc"},
   279  			{inlineFunc, "nextFreeFastStub", "nextFreeFastStub"},
   280  			{inlineFunc, "deductAssistCredit", "deductAssistCredit"},
   281  			{foldCondition, "isSlowPath_", str(true)},
   282  			{foldCondition, "isTiny_", str(false)},
   283  			{foldCondition, "isNoScan_", str(true)},
   284  
   285  			// Remove constants used by size-specialized variants.
   286  			{deleteConst, "elemsize", ""},
   287  			{deleteConst, "sizeclass", ""},
   288  			{deleteConst, "spc", ""},
   289  		},
   290  	})
   291  
   292  	return config
   293  }
   294  
   295  // inline applies the inlining operations given by the config.
   296  func inline(config generatorConfig) []byte {
   297  	var out bytes.Buffer
   298  
   299  	// Read the template file in.
   300  	fset := token.NewFileSet()
   301  	f, err := parser.ParseFile(fset, config.file, nil, parser.SkipObjectResolution)
   302  	if err != nil {
   303  		log.Fatalf("parsing %s: %v", config.file, err)
   304  	}
   305  
   306  	// Collect the function and import declarations. The function
   307  	// declarations in the template file provide both the templates
   308  	// that will be stamped out, and the functions that will be inlined
   309  	// into them. The imports from the template file will be copied
   310  	// straight to the output.
   311  	funcDecls := map[string]*ast.FuncDecl{}
   312  	importDecls := []*ast.GenDecl{}
   313  	for _, decl := range f.Decls {
   314  		switch decl := decl.(type) {
   315  		case *ast.FuncDecl:
   316  			funcDecls[decl.Name.Name] = decl
   317  		case *ast.GenDecl:
   318  			if decl.Tok.String() == "import" {
   319  				importDecls = append(importDecls, decl)
   320  				continue
   321  			}
   322  		}
   323  	}
   324  
   325  	// Write out the package and import declarations.
   326  	out.WriteString("// Code generated by mkmalloc.go; DO NOT EDIT.\n")
   327  	out.WriteString("// See overview in malloc_stubs.go.\n\n")
   328  	out.WriteString("package " + f.Name.Name + "\n\n")
   329  	for _, importDecl := range importDecls {
   330  		out.Write(mustFormatNode(fset, importDecl))
   331  		out.WriteString("\n\n")
   332  	}
   333  
   334  	// Produce each of the inlined functions specified by specs.
   335  	for _, spec := range config.specs {
   336  		// Start with a renamed copy of the template function.
   337  		containingFuncCopy := internalastutil.CloneNode(funcDecls[spec.templateFunc])
   338  		if containingFuncCopy == nil {
   339  			log.Fatal("did not find", spec.templateFunc)
   340  		}
   341  		containingFuncCopy.Name.Name = spec.name
   342  
   343  		// Apply each of the ops given by the specs
   344  		stamped := ast.Node(containingFuncCopy)
   345  		for _, repl := range spec.ops {
   346  			switch repl.kind {
   347  			case inlineFunc:
   348  				if toDecl, ok := funcDecls[repl.to]; ok {
   349  					stamped = inlineFunction(stamped, repl.from, toDecl)
   350  				}
   351  			case subBasicLit:
   352  				stamped = substituteWithBasicLit(stamped, repl.from, repl.to)
   353  			case foldCondition:
   354  				stamped = foldIfCondition(stamped, repl.from, repl.to)
   355  			case subIdent:
   356  				stamped = substituteIdent(stamped, repl.from, repl.to)
   357  			case deleteConst:
   358  				stamped = deleteConstDecl(stamped, repl.from)
   359  			default:
   360  				log.Fatalf("unknown op kind %v", repl.kind)
   361  			}
   362  		}
   363  
   364  		stamped = cleanLabels(stamped)
   365  
   366  		out.Write(mustFormatNode(fset, stamped))
   367  		out.WriteString("\n\n")
   368  	}
   369  
   370  	return out.Bytes()
   371  }
   372  
   373  // substituteWithBasicLit recursively renames identifiers in the provided AST
   374  // according to 'from' and 'to'.
   375  func substituteWithBasicLit(node ast.Node, from, to string) ast.Node {
   376  	// The op is a substitution of an identifier with a basic literal.
   377  	toExpr, err := parser.ParseExpr(to)
   378  	if err != nil {
   379  		log.Fatalf("parsing expr %q: %v", to, err)
   380  	}
   381  	toLit, ok := toExpr.(*ast.BasicLit)
   382  	if !ok {
   383  		log.Fatalf("op 'to' expr %q is not a basic literal", to)
   384  	}
   385  	return astutil.Apply(node, func(cursor *astutil.Cursor) bool {
   386  		if ident, ok := cursor.Node().(*ast.Ident); ok && ident.Name == from {
   387  			replacement := *toLit
   388  			replacement.ValuePos = ident.NamePos
   389  			cursor.Replace(new(replacement))
   390  		}
   391  		return true
   392  	}, nil)
   393  }
   394  
   395  // substituteIdent replaces the ident named 'from' to 'to'.
   396  func substituteIdent(node ast.Node, from, to string) ast.Node {
   397  	return astutil.Apply(node, func(cursor *astutil.Cursor) bool {
   398  		if ident, ok := cursor.Node().(*ast.Ident); ok && ident.Name == from {
   399  			cursor.Replace(&ast.Ident{Name: to, NamePos: ident.NamePos})
   400  		}
   401  		return true
   402  	}, nil)
   403  }
   404  
   405  // foldIfCondition replaces 'from' with 'to', which must be "true" or "false".
   406  // It then applies simplifications to any boolean expressions that have literal
   407  // true or false values, from the bottom up. Any if statements that have a condition
   408  // that is a literal true or false after the simplification will be replaced with
   409  // their bodies (in the true case) or deleted (in the false case).
   410  func foldIfCondition(node ast.Node, from, to string) ast.Node {
   411  	boolLit := func(n ast.Expr) (v, ok bool) {
   412  		if ident, ok := ast.Unparen(n).(*ast.Ident); ok {
   413  			switch ident.Name {
   414  			case "true":
   415  				return true, true
   416  			case "false":
   417  				return false, true
   418  			}
   419  			return false, false
   420  		}
   421  		return false, false
   422  	}
   423  	handleIfs := func(cursor *astutil.Cursor) bool {
   424  		switch n := cursor.Node().(type) {
   425  		case *ast.Ident:
   426  			// First, do the replacement.
   427  			if n.Name == from {
   428  				cursor.Replace(&ast.Ident{Name: to, NamePos: n.NamePos})
   429  			}
   430  		case *ast.UnaryExpr:
   431  			if n.Op == token.NOT {
   432  				if b, ok := boolLit(n.X); ok {
   433  					name := "true"
   434  					if b {
   435  						name = "false"
   436  					}
   437  					cursor.Replace(&ast.Ident{Name: name, NamePos: n.Pos()})
   438  				}
   439  			}
   440  		case *ast.BinaryExpr:
   441  			xBool, xOk := boolLit(n.X)
   442  			yBool, yOk := boolLit(n.Y)
   443  			if n.Op == token.LAND {
   444  				switch {
   445  				case xOk && !xBool || yOk && !yBool:
   446  					cursor.Replace(&ast.Ident{Name: "false", NamePos: n.Pos()})
   447  				case xOk && xBool:
   448  					cursor.Replace(n.Y)
   449  				case yOk && yBool:
   450  					cursor.Replace(n.X)
   451  				}
   452  			} else if n.Op == token.LOR {
   453  				switch {
   454  				case xOk && xBool || yOk && yBool:
   455  					cursor.Replace(&ast.Ident{Name: "true", NamePos: n.Pos()})
   456  				case xOk && !xBool:
   457  					cursor.Replace(n.Y)
   458  				case yOk && !yBool:
   459  					cursor.Replace(n.X)
   460  				}
   461  			}
   462  		case *ast.IfStmt:
   463  			if v, ok := boolLit(n.Cond); ok {
   464  				if cursor.Index() < 0 {
   465  					replacement := ast.Node(&ast.EmptyStmt{})
   466  					if v {
   467  						replacement = n.Body
   468  					}
   469  					cursor.Replace(replacement)
   470  					break
   471  				}
   472  				if v {
   473  					for _, stmt := range n.Body.List {
   474  						cursor.InsertBefore(stmt)
   475  					}
   476  				} else if n.Else != nil {
   477  					if block, ok := n.Else.(*ast.BlockStmt); ok {
   478  						for i := len(block.List) - 1; i >= 0; i-- {
   479  							cursor.InsertAfter(block.List[i])
   480  						}
   481  					}
   482  				}
   483  				cursor.Delete()
   484  			}
   485  		case *ast.LabeledStmt:
   486  			// This case isn't necessary but it moves the code
   487  			// out of the block so that it looks cleaner.
   488  			if inner, ok := n.Stmt.(*ast.BlockStmt); ok {
   489  				if len(inner.List) == 0 {
   490  					cursor.Delete()
   491  					break
   492  				}
   493  				list := inner.List
   494  				n.Stmt = list[0]
   495  				for i := len(list) - 1; i > 0; i-- {
   496  					cursor.InsertAfter(list[i])
   497  				}
   498  			}
   499  		}
   500  		return true
   501  	}
   502  	return astutil.Apply(node, nil, handleIfs)
   503  }
   504  
   505  func cleanLabels(node ast.Node) ast.Node {
   506  	found := map[string]bool{}
   507  	ast.Inspect(node, func(node ast.Node) bool {
   508  		if branch, ok := node.(*ast.BranchStmt); ok {
   509  			if branch.Label != nil {
   510  				found[branch.Label.Name] = true
   511  			}
   512  		}
   513  		return true
   514  	})
   515  	return astutil.Apply(node, nil, func(cursor *astutil.Cursor) bool {
   516  		if lstmt, ok := cursor.Node().(*ast.LabeledStmt); ok {
   517  			if !found[lstmt.Label.Name] {
   518  				if _, ok := lstmt.Stmt.(*ast.EmptyStmt); ok {
   519  					cursor.Delete()
   520  				} else {
   521  					cursor.Replace(lstmt.Stmt)
   522  				}
   523  			}
   524  		}
   525  		return true
   526  	})
   527  }
   528  
   529  // reports whether this is a non-grouped constant decl named 'name'.
   530  func isNamedConstDecl(node ast.Node, name string) bool {
   531  	declStmt, ok := node.(*ast.DeclStmt)
   532  	if !ok {
   533  		return false
   534  	}
   535  
   536  	genDecl, ok := declStmt.Decl.(*ast.GenDecl)
   537  	if !ok || genDecl.Tok != token.CONST {
   538  		return false
   539  	}
   540  
   541  	if len(genDecl.Specs) != 1 {
   542  		return false
   543  	}
   544  	vs, ok := genDecl.Specs[0].(*ast.ValueSpec)
   545  	if !ok || len(vs.Names) != 1 || len(vs.Values) != 1 {
   546  		return false
   547  	}
   548  
   549  	return vs.Names[0].Name == name
   550  }
   551  
   552  // deleteConstDecl removes const declarations whose name matches the given name.
   553  // It only applies to declaration statements with a single declaration.
   554  func deleteConstDecl(node ast.Node, name string) ast.Node {
   555  	return astutil.Apply(node, func(cursor *astutil.Cursor) bool {
   556  		if isNamedConstDecl(cursor.Node(), name) {
   557  			cursor.Delete()
   558  		}
   559  		return true
   560  	}, nil)
   561  }
   562  
   563  // inlineFunction recursively replaces calls to the function 'from' with the body of the function
   564  // 'toDecl'. All calls to 'from' must either have no return values and appear in standalone expression statements
   565  // or otherwise must appear in assignment statements.
   566  // The replacement is very simple: it doesn't substitute the arguments for the parameters, so the
   567  // arguments to the function call must be the same identifier as the parameters to the function
   568  // declared by 'toDecl'. If there are any calls to from where that's not the case there will be a fatal error.
   569  func inlineFunction(node ast.Node, from string, toDecl *ast.FuncDecl) ast.Node {
   570  	return astutil.Apply(node, func(cursor *astutil.Cursor) bool {
   571  		switch node := cursor.Node().(type) {
   572  		case *ast.AssignStmt:
   573  			// TODO(matloob) CHECK function args have same name
   574  			// as parameters (or parameter is "_").
   575  			if len(node.Rhs) == 1 && isCallTo(node.Rhs[0], from) {
   576  				args := node.Rhs[0].(*ast.CallExpr).Args
   577  				if !argsMatchParameters(args, toDecl.Type.Params) {
   578  					log.Fatalf("applying op: arguments to %v don't match parameter names of %v: %v", from, toDecl.Name, debugPrint(args...))
   579  				}
   580  				replaceAssignment(cursor, node, toDecl)
   581  			}
   582  			return false
   583  		case *ast.ExprStmt:
   584  			if callExpr, ok := node.X.(*ast.CallExpr); ok && isCallTo(callExpr, from) {
   585  				if !argsMatchParameters(callExpr.Args, toDecl.Type.Params) {
   586  					log.Fatalf("applying op: arguments to %v don't match parameter names of %v: %v", from, toDecl.Name, debugPrint(callExpr.Args...))
   587  				}
   588  				if toDecl.Type.Results != nil {
   589  					log.Fatalf("applying op: call to %v, which does not appear in an assignment, is replaced with %v which has return values: %v", from, toDecl.Name, debugPrint(callExpr.Args...))
   590  				}
   591  				replaceCallExprStmt(cursor, toDecl)
   592  			}
   593  			return false
   594  		case *ast.ReturnStmt:
   595  			if len(node.Results) == 1 && isCallTo(node.Results[0], from) {
   596  				args := node.Results[0].(*ast.CallExpr).Args
   597  				if !argsMatchParameters(args, toDecl.Type.Params) {
   598  					log.Fatalf("applying op: arguments to %v don't match parameter names of %v: %v", from, toDecl.Name, debugPrint(args...))
   599  				}
   600  				replaceTailCall(cursor, toDecl)
   601  			}
   602  			return false
   603  		case *ast.CallExpr:
   604  			if isCallTo(node, from) {
   605  				switch cursor.Parent().(type) {
   606  				case *ast.AssignStmt, *ast.ExprStmt:
   607  				default:
   608  					log.Fatalf("applying op: all calls to function %q being replaced must appear in an assignment or expression statement, appears in %T", from, cursor.Parent())
   609  				}
   610  			}
   611  		}
   612  		return true
   613  	}, nil)
   614  }
   615  
   616  // argsMatchParameters reports whether the arguments given by args are all identifiers
   617  // whose names are the same as the corresponding parameters in params.
   618  func argsMatchParameters(args []ast.Expr, params *ast.FieldList) bool {
   619  	var paramIdents []*ast.Ident
   620  	for _, f := range params.List {
   621  		paramIdents = append(paramIdents, f.Names...)
   622  	}
   623  
   624  	if len(args) != len(paramIdents) {
   625  		return false
   626  	}
   627  
   628  	for i := range args {
   629  		if !isIdentWithName(args[i], paramIdents[i].Name) {
   630  			return false
   631  		}
   632  	}
   633  
   634  	return true
   635  }
   636  
   637  // isIdentWithName reports whether the expression is an identifier with the given name.
   638  func isIdentWithName(expr ast.Node, name string) bool {
   639  	ident, ok := expr.(*ast.Ident)
   640  	if !ok {
   641  		return false
   642  	}
   643  	return ident.Name == name
   644  }
   645  
   646  // isCallTo reports whether the expression is a call expression to the function with the given name.
   647  func isCallTo(expr ast.Expr, name string) bool {
   648  	callexpr, ok := expr.(*ast.CallExpr)
   649  	if !ok {
   650  		return false
   651  	}
   652  	return isIdentWithName(callexpr.Fun, name)
   653  }
   654  
   655  // replaceCallExprStmt replaces a standalone expression statement calling a function with no
   656  // return values with the body of the function.
   657  func replaceCallExprStmt(cursor *astutil.Cursor, funcdecl *ast.FuncDecl) {
   658  	body := internalastutil.CloneNode(funcdecl.Body)
   659  	for _, stmt := range body.List {
   660  		cursor.InsertBefore(stmt)
   661  	}
   662  	cursor.Delete()
   663  }
   664  
   665  func replaceTailCall(cursor *astutil.Cursor, funcdecl *ast.FuncDecl) {
   666  	if !hasTerminatingReturn(funcdecl.Body) {
   667  		log.Fatal("function being inlined must have a return at the end")
   668  	}
   669  
   670  	body := internalastutil.CloneNode(funcdecl.Body)
   671  	if len(body.List) < 1 {
   672  		log.Fatal("replacing with empty bodied function")
   673  	}
   674  
   675  	// The op happens in two steps: first we insert the body of the function being inlined (except for
   676  	// the final return) before the assignment, and then we change the assignment statement to replace the function call
   677  	// with the expressions being returned.
   678  
   679  	// Insert the body up to the final return.
   680  	for _, stmt := range body.List {
   681  		cursor.InsertBefore(stmt)
   682  	}
   683  	cursor.Delete()
   684  }
   685  
   686  // replaceAssignment replaces an assignment statement where the right hand side is a function call
   687  // whose arguments have the same names as the parameters to funcdecl with the body of funcdecl.
   688  // It sets the left hand side of the assignment to the return values of the function.
   689  func replaceAssignment(cursor *astutil.Cursor, assign *ast.AssignStmt, funcdecl *ast.FuncDecl) {
   690  	if !hasTerminatingReturn(funcdecl.Body) {
   691  		log.Fatal("function being inlined must have a return at the end")
   692  	}
   693  
   694  	body := internalastutil.CloneNode(funcdecl.Body)
   695  	if hasTerminatingAndNonterminatingReturn(funcdecl.Body) {
   696  		// The function has multiple return points. Add the code that we'd continue with in the caller
   697  		// after each of the return points. The calling function must have a terminating return
   698  		// so we don't continue execution in the replaced function after we finish executing the
   699  		// continue block that we add.
   700  		body = addContinues(cursor, assign, body, everythingFollowingInParent(cursor)).(*ast.BlockStmt)
   701  	}
   702  
   703  	if len(body.List) < 1 {
   704  		log.Fatal("replacing with empty bodied function")
   705  	}
   706  
   707  	// The op happens in two steps: first we insert the body of the function being inlined (except for
   708  	// the final return) before the assignment, and then we change the assignment statement to replace the function call
   709  	// with the expressions being returned.
   710  
   711  	// Determine the expressions being returned.
   712  	beforeReturn, ret := body.List[:len(body.List)-1], body.List[len(body.List)-1]
   713  	returnStmt, ok := ret.(*ast.ReturnStmt)
   714  	if !ok {
   715  		log.Fatal("last stmt in function we're replacing with should be a return")
   716  	}
   717  	results := returnStmt.Results
   718  
   719  	// Insert the body up to the final return.
   720  	for _, stmt := range beforeReturn {
   721  		cursor.InsertBefore(stmt)
   722  	}
   723  
   724  	// Rewrite the assignment statement.
   725  	replaceWithAssignment(cursor, assign.Lhs, results, assign.Tok)
   726  }
   727  
   728  // hasTerminatingReturn reparts whether the block ends in a return statement.
   729  func hasTerminatingReturn(block *ast.BlockStmt) bool {
   730  	_, ok := block.List[len(block.List)-1].(*ast.ReturnStmt)
   731  	return ok
   732  }
   733  
   734  // hasTerminatingAndNonterminatingReturn reports whether the block ends in a return
   735  // statement, and also has a return elsewhere in it.
   736  func hasTerminatingAndNonterminatingReturn(block *ast.BlockStmt) bool {
   737  	if !hasTerminatingReturn(block) {
   738  		return false
   739  	}
   740  	var ret bool
   741  	for i := range block.List[:len(block.List)-1] {
   742  		ast.Inspect(block.List[i], func(node ast.Node) bool {
   743  			_, ok := node.(*ast.ReturnStmt)
   744  			if ok {
   745  				ret = true
   746  				return false
   747  			}
   748  			return true
   749  		})
   750  	}
   751  	return ret
   752  }
   753  
   754  // everythingFollowingInParent returns a block with everything in the parent block node of the cursor after
   755  // the cursor itself. The cursor must point to an element in a block node's list.
   756  func everythingFollowingInParent(cursor *astutil.Cursor) *ast.BlockStmt {
   757  	parent := cursor.Parent()
   758  	block, ok := parent.(*ast.BlockStmt)
   759  	if !ok {
   760  		log.Fatal("internal error: in everythingFollowingInParent, cursor doesn't point to element in block list")
   761  	}
   762  
   763  	blockcopy := internalastutil.CloneNode(block)      // get a clean copy
   764  	blockcopy.List = blockcopy.List[cursor.Index()+1:] // and remove everything before and including stmt
   765  
   766  	if _, ok := blockcopy.List[len(blockcopy.List)-1].(*ast.ReturnStmt); !ok {
   767  		log.Printf("%s", mustFormatNode(token.NewFileSet(), blockcopy))
   768  		log.Fatal("internal error: parent doesn't end in a return")
   769  	}
   770  	return blockcopy
   771  }
   772  
   773  // in the case that there's a return in the body being inlined (toBlock), addContinues
   774  // replaces those returns that are not at the end of the function with the code in the
   775  // caller after the function call that execution would continue with after the return.
   776  // The block being added must end in a return.
   777  func addContinues(cursor *astutil.Cursor, assignNode *ast.AssignStmt, toBlock *ast.BlockStmt, continueBlock *ast.BlockStmt) ast.Node {
   778  	if !hasTerminatingReturn(continueBlock) {
   779  		log.Fatal("the block being continued to in addContinues must end in a return")
   780  	}
   781  	applyFunc := func(cursor *astutil.Cursor) bool {
   782  		ret, ok := cursor.Node().(*ast.ReturnStmt)
   783  		if !ok {
   784  			return true
   785  		}
   786  
   787  		if cursor.Parent() == toBlock && cursor.Index() == len(toBlock.List)-1 {
   788  			return false
   789  		}
   790  
   791  		// This is the opposite of replacing a function call with the body. First
   792  		// we replace the return statement with the assignment from the caller, and
   793  		// then add the code we continue with.
   794  		replaceWithAssignment(cursor, assignNode.Lhs, ret.Results, assignNode.Tok)
   795  		cursor.InsertAfter(internalastutil.CloneNode(continueBlock))
   796  
   797  		return false
   798  	}
   799  	return astutil.Apply(toBlock, applyFunc, nil)
   800  }
   801  
   802  // debugPrint prints out the expressions given by nodes for debugging.
   803  func debugPrint(nodes ...ast.Expr) string {
   804  	var b strings.Builder
   805  	for i, node := range nodes {
   806  		b.Write(mustFormatNode(token.NewFileSet(), node))
   807  		if i != len(nodes)-1 {
   808  			b.WriteString(", ")
   809  		}
   810  	}
   811  	return b.String()
   812  }
   813  
   814  // mustFormatNode produces the formatted Go code for the given node.
   815  func mustFormatNode(fset *token.FileSet, node any) []byte {
   816  	var buf bytes.Buffer
   817  	format.Node(&buf, fset, node)
   818  	return buf.Bytes()
   819  }
   820  
   821  // mustMatchExprs makes sure that the expression lists have the same length,
   822  // and returns the lists of the expressions on the lhs and rhs where the
   823  // identifiers are not the same. These are used to produce assignment statements
   824  // where the expressions on the right are assigned to the identifiers on the left.
   825  func mustMatchExprs(lhs []ast.Expr, rhs []ast.Expr) ([]ast.Expr, []ast.Expr) {
   826  	if len(lhs) != len(rhs) {
   827  		log.Fatal("exprs don't match", debugPrint(lhs...), debugPrint(rhs...))
   828  	}
   829  
   830  	var newLhs, newRhs []ast.Expr
   831  	for i := range lhs {
   832  		lhsIdent, ok1 := lhs[i].(*ast.Ident)
   833  		rhsIdent, ok2 := rhs[i].(*ast.Ident)
   834  		if ok1 && ok2 && lhsIdent.Name == rhsIdent.Name {
   835  			continue
   836  		}
   837  		newLhs = append(newLhs, lhs[i])
   838  		newRhs = append(newRhs, rhs[i])
   839  	}
   840  
   841  	return newLhs, newRhs
   842  }
   843  
   844  // replaceWithAssignment replaces the node pointed to by the cursor with an assignment of the
   845  // left hand side to the righthand side, removing any redundant assignments of a variable to itself,
   846  // and replacing an assignment to a single basic literal with a constant declaration.
   847  func replaceWithAssignment(cursor *astutil.Cursor, lhs, rhs []ast.Expr, tok token.Token) {
   848  	newLhs, newRhs := mustMatchExprs(lhs, rhs)
   849  	if len(newLhs) == 0 {
   850  		cursor.Delete()
   851  		return
   852  	}
   853  	if len(newRhs) == 1 {
   854  		if lit, ok := newRhs[0].(*ast.BasicLit); ok {
   855  			constDecl := &ast.DeclStmt{
   856  				Decl: &ast.GenDecl{
   857  					Tok: token.CONST,
   858  					Specs: []ast.Spec{
   859  						&ast.ValueSpec{
   860  							Names:  []*ast.Ident{newLhs[0].(*ast.Ident)},
   861  							Values: []ast.Expr{lit},
   862  						},
   863  					},
   864  				},
   865  			}
   866  			cursor.Replace(constDecl)
   867  			return
   868  		}
   869  	}
   870  	newAssignment := &ast.AssignStmt{
   871  		Lhs: newLhs,
   872  		Rhs: newRhs,
   873  		Tok: tok,
   874  	}
   875  	cursor.Replace(newAssignment)
   876  }
   877  
   878  // generateTable generates the file with the jump tables for the specialized malloc functions.
   879  func generateTable(sizeToSizeClass []uint8) []byte {
   880  	scMax := sizeToSizeClass[specializedMallocMax]
   881  
   882  	var b bytes.Buffer
   883  	fmt.Fprintf(&b, `// Code generated by mkmalloc.go; DO NOT EDIT.
   884  //go:build !plan9
   885  
   886  package runtime
   887  
   888  import "unsafe"
   889  
   890  var mallocScanTable = [%d]func(size uintptr, typ *_type, needzero bool) unsafe.Pointer{`, specializedMallocMax+1)
   891  
   892  	for i := range uintptr(specializedMallocMax + 1) {
   893  		fmt.Fprintf(&b, "%s,\n", smallScanNoHeaderSCFuncName(sizeToSizeClass[i], scMax))
   894  	}
   895  
   896  	fmt.Fprintf(&b, `
   897  }
   898  
   899  var mallocNoScanTable = [%d]func(size uintptr, typ *_type, needzero bool) unsafe.Pointer{`, specializedMallocMax+1)
   900  	for i := range uintptr(specializedMallocMax + 1) {
   901  		if i < 16 {
   902  			fmt.Fprintf(&b, "%s,\n", "mallocPanic")
   903  		} else {
   904  			fmt.Fprintf(&b, "%s,\n", smallNoScanSCFuncName(sizeToSizeClass[i], scMax))
   905  		}
   906  	}
   907  
   908  	fmt.Fprintln(&b, `
   909  }`)
   910  
   911  	return b.Bytes()
   912  }
   913  
   914  // Generate benchmarks for all potentially small sizes
   915  // (sizes for which smallScanNoHeader would be called)
   916  // gc.MinSizeForMallocHeader is defined as goarch.PtrSize * goarch.PtrBits.
   917  
   918  const benchmarkMax = maxPtrSize * maxPtrBits
   919  
   920  // benchmarkConfig produces an inlining config to stamp out microbenchmarks.
   921  func benchmarkConfig(classes []class, sizeToSizeClass []uint8) generatorConfig {
   922  	config := generatorConfig{file: "../malloc_stubs_test.go"}
   923  
   924  	scMax := sizeToSizeClass[benchmarkMax]
   925  
   926  	str := fmt.Sprint
   927  
   928  	for sc := uint8(1); sc <= scMax; sc++ {
   929  		elemsize := classes[sc].size
   930  		config.specs = append(config.specs, spec{
   931  			templateFunc: "benchmarkStub",
   932  			name:         fmt.Sprintf("benchmarkMallocgcNoscan%d", elemsize),
   933  			ops: []op{
   934  				{subBasicLit, "size_", str(elemsize)},
   935  				{foldCondition, "noscan_", str(true)},
   936  			},
   937  		})
   938  		config.specs = append(config.specs, spec{
   939  			templateFunc: "benchmarkStub",
   940  			name:         fmt.Sprintf("benchmarkMallocgcScan%d", elemsize),
   941  			ops: []op{
   942  				{subBasicLit, "size_", str(elemsize)},
   943  				{foldCondition, "noscan_", str(false)},
   944  			},
   945  		})
   946  		config.specs = append(config.specs, spec{
   947  			templateFunc: "benchmarkScanSliceStub",
   948  			name:         fmt.Sprintf("benchmarkMallocgcScanSlice%d", elemsize),
   949  			ops:          []op{{subBasicLit, "size_", str(elemsize)}},
   950  		})
   951  	}
   952  
   953  	for size := 1; size < tinySize; size++ {
   954  		config.specs = append(config.specs, spec{
   955  			templateFunc: "benchmarkStubTiny",
   956  			name:         fmt.Sprintf("benchmarkMallocgcTiny%d", size),
   957  			ops:          []op{{subBasicLit, "size_", str(size)}, {foldCondition, "noscan_", str(true)}},
   958  		})
   959  	}
   960  
   961  	return config
   962  }
   963  
   964  func generateTopBenchmark(classes []class, sizeToSizeClass []uint8) string {
   965  	scMax := sizeToSizeClass[benchmarkMax]
   966  	bench := `func BenchmarkMallocgc(b *testing.B) {
   967  		b.Run("scan=noscan", func(b *testing.B) {
   968  `
   969  	for size := 1; size < tinySize; size++ {
   970  		bench += fmt.Sprintf(`b.Run("size=%d", benchmarkMallocgcTiny%d)`, size, size) + "\n"
   971  	}
   972  	for sc := uint8(2); sc <= scMax; sc++ {
   973  		elemsize := classes[sc].size
   974  		bench += fmt.Sprintf(`b.Run("size=%d", benchmarkMallocgcNoscan%d)`, elemsize, elemsize) + "\n"
   975  	}
   976  	bench += `})
   977  		b.Run("scan=scan", func(b *testing.B) {
   978  `
   979  	for sc := uint8(1); sc <= scMax; sc++ {
   980  		elemsize := classes[sc].size
   981  		bench += fmt.Sprintf(`b.Run("size=%d", benchmarkMallocgcScan%d)`, elemsize, elemsize) + "\n"
   982  
   983  	}
   984  	bench += `})
   985  		b.Run("scan=scanslice", func(b *testing.B) {
   986  `
   987  	for sc := uint8(1); sc <= scMax; sc++ {
   988  		elemsize := classes[sc].size
   989  		bench += fmt.Sprintf(`b.Run("size=%d", benchmarkMallocgcScanSlice%d)`, elemsize, elemsize) + "\n"
   990  	}
   991  	bench += `})
   992  }`
   993  
   994  	return bench
   995  }
   996  

View as plain text