From 38463561e6f3987a162e523fa92e8af51ad135e9 Mon Sep 17 00:00:00 2001 From: ZhouGuangyuan Date: Tue, 15 Sep 2026 00:44:04 +0800 Subject: [PATCH] feat: propagate result relations before ABI conversion --- doc/Function-attributes.md | 19 +- internal/build/build.go | 6 + internal/cabi/result_attributes_test.go | 131 +++++++++ internal/funcattrs/llvm.go | 8 +- internal/funcattrs/source.go | 46 ++- internal/funcattrs/source_test.go | 25 +- internal/funcattrs/values.go | 277 ++++++++++++++++++ internal/funcattrs/values_test.go | 202 +++++++++++++ runtime/internal/runtime/z_error.go | 2 +- runtime/internal/runtime/z_string.go | 2 + .../internal/test/result_attributes_test.go | 106 +++++++ .../internal/test/value_attributes_test.go | 2 +- ssa/decl.go | 2 +- ssa/package.go | 10 +- ssa/value_attributes.go | 41 ++- 15 files changed, 861 insertions(+), 18 deletions(-) create mode 100644 internal/cabi/result_attributes_test.go create mode 100644 internal/funcattrs/values.go create mode 100644 internal/funcattrs/values_test.go create mode 100644 runtime/internal/test/result_attributes_test.go diff --git a/doc/Function-attributes.md b/doc/Function-attributes.md index b5e5dd01ee..7518d80723 100644 --- a/doc/Function-attributes.md +++ b/doc/Function-attributes.md @@ -21,7 +21,7 @@ func Fatal(message string) { The properties apply to definitions and imported declarations, including methods, generic instances and declarations connected by linkname. They are preserved by ABI conversion and cached builds. The runtime uses the same source directives as ordinary packages. -## Parameter and single-result attributes +## Parameter and result attributes Use `//llgo:param(name|index)`, `//llgo:receiver`, or `//llgo:result(name|index)` to select a whole Go value. Indices start at zero; the receiver is separate. A function with exactly one result also permits `//llgo:result`. Both comment spacings accepted for function attributes are accepted here. @@ -43,5 +43,20 @@ func Checked(p *int) *int { Source selectors are resolved before exported signatures can lose parameter names. Imports, linkname declarations and generic instances preserve the guarantees; concrete generic types are checked at instantiation. ABI conversion preserves attributes on the corresponding scalar values when other arguments are packed or passed indirectly. -Result attributes currently require exactly one result. Multiple-result guarantees, `sameas`, `access` and `noalias` are subsequent steps of #2590. +### Result relations and multiple results + +`sameas(name)` states that an integer or pointer result equals the entry value of the ordinary parameter named `name`. Integer types must match; pointer conversions must preserve the pointer. Like other result guarantees, it applies after normal return, including changes made by deferred functions. + +```go +//llgo:result(out) nonnull sameas(p) +//llgo:result(count) range(0,64) +func Make(p *int, n uint32) (out *int, count uint32) { + if p == nil { panic("nil pointer") } + return p, n & 63 +} +``` + +With multiple Go results, select each whole result by name or zero-based index. The selector remains stable if the ABI packs or splits the results or returns them through caller-provided storage. The call and its possible panic remain; only the guaranteed results can be simplified. `sameas` uses the evaluated input value, even if memory supplying that input changes during the call. + +`access` and `noalias` are the final step of #2590. diff --git a/internal/build/build.go b/internal/build/build.go index 08884bfbe4..d3dccf4c94 100644 --- a/internal/build/build.go +++ b/internal/build/build.go @@ -2191,6 +2191,9 @@ func buildMainLink(ctx *context, pkg *packages.Package, preparation *mainLinkPre ctx.stripDarwinLTOLocals = false entryPkg := genMainModule(ctx, llssa.PkgRuntime, pkg, &preparation.gen) cExports := preparation.gen.cExports + if err := ctx.prog.MaterializeValueAttributes(entryPkg.LPkg.Module()); err != nil { + return nil, err + } if len(cExports) != 0 { llabi.LowerLargeAggregates(ctx.prog.TargetData(), entryPkg.LPkg.Module()) ctx.cTransformer.TransformModule(entryPkg.LPkg.Path(), entryPkg.LPkg.Module()) @@ -2853,6 +2856,9 @@ func compilePackageModule(ctx *context, aPkg *aPackage, externs []string, verbos ret := aPkg.LPkg ctx.cTransformer.SetSkipFuncs(cabiSkipFuncsForPlan9Asm(ctx, pkgPath, ret.Module())) + if err := ctx.prog.MaterializeValueAttributes(ret.Module()); err != nil { + return err + } llabi.LowerLargeAggregates(ctx.prog.TargetData(), ret.Module()) ctx.cTransformer.TransformModule(ret.Path(), ret.Module()) ctx.cTransformer.SetSkipFuncs(nil) diff --git a/internal/cabi/result_attributes_test.go b/internal/cabi/result_attributes_test.go new file mode 100644 index 0000000000..a8ef6fa1f6 --- /dev/null +++ b/internal/cabi/result_attributes_test.go @@ -0,0 +1,131 @@ +package cabi + +import ( + "go/ast" + "go/parser" + "go/token" + "go/types" + "strings" + "testing" + + "github.com/xgo-dev/llgo/internal/funcattrs" + llssa "github.com/xgo-dev/llgo/ssa" + "github.com/xgo-dev/llvm" +) + +func contractTestDeclaration(t *testing.T, prog llssa.Program, pkg llssa.Package, source, name string) llssa.Function { + t.Helper() + fset := token.NewFileSet() + file, err := parser.ParseFile(fset, "contracts.go", "package contracts\n"+source, parser.ParseComments) + if err != nil { + t.Fatal(err) + } + info := &types.Info{Defs: make(map[*ast.Ident]types.Object)} + if _, err := (&types.Config{}).Check("contracts", fset, []*ast.File{file}, info); err != nil { + t.Fatal(err) + } + decl := file.Decls[len(file.Decls)-1].(*ast.FuncDecl) + attrs, err := funcattrs.Parse(fset, decl) + if err != nil { + t.Fatal(err) + } + if err := prog.SetValueAttributes(name, attrs); err != nil { + t.Fatal(err) + } + return pkg.NewFunc(name, info.Defs[decl.Name].Type().(*types.Signature), llssa.InGo) +} + +func TestValueContractsSurviveByvalSretAndPackedABI(t *testing.T) { + llvm.InitializeAllTargets() + llvm.InitializeAllTargetInfos() + llvm.InitializeAllTargetMCs() + llvm.InitializeAllAsmPrinters() + for _, target := range []llssa.Target{ + {GOOS: "linux", GOARCH: "amd64"}, + {GOOS: "darwin", GOARCH: "arm64"}, + {GOOS: "windows", GOARCH: "386"}, + {GOOS: "wasip1", GOARCH: "wasm32"}, + } { + t.Run(target.GOARCH, func(t *testing.T) { + prog := llssa.NewProgram(&target) + defer prog.Dispose() + if target.GOARCH == "386" { + prog.TypeSizes(types.SizesFor("gc", "386")) + } + pkg := prog.NewPackage("contracts", "contracts") + callee := contractTestDeclaration(t, prog, pkg, ` +type Pair struct { P *int; N int64; Extra [24]byte } +//llgo:param(p) nonnull +//llgo:result(q) nonnull sameas(p) +//llgo:result(n) range(0,7) +func F(input Pair, p *int) (q *int, n int64, output Pair) +`, "contracts.F") + calleeType := callee.Type.RawType().(*types.Signature) + callerSig := types.NewSignatureType(nil, nil, nil, calleeType.Params(), + types.NewTuple(types.NewVar(token.NoPos, nil, "", types.Typ[types.Bool])), false) + caller := pkg.NewFunc("contracts.Caller", callerSig, llssa.InGo) + b := caller.MakeBody(1) + r := b.Call(callee.Expr, caller.Param(0), caller.Param(1)) + pointer := b.Extract(r, 0) + badPointer := b.BinOp(token.EQL, pointer, prog.Nil(pointer.Type)) + n := b.Extract(r, 1) + negative := b.BinOp(token.LSS, n, prog.IntVal(0, n.Type)) + large := b.BinOp(token.GEQ, n, prog.IntVal(7, n.Type)) + b.Return(b.BinOp(token.OR, badPointer, b.BinOp(token.OR, negative, large))) + b.EndBuild() + + packed := contractTestDeclaration(t, prog, pkg, ` +type Packed struct { N int8; Other [7]byte } +//llgo:param(n) range(-3,4) +func PackedInput(p Packed, n int8) bool +`, "contracts.PackedInput") + pbody := packed.MakeBody(1) + value := packed.Param(1) + pbody.Return(pbody.BinOp(token.GTR, value, prog.IntVal(3, value.Type))) + pbody.EndBuild() + + mod := pkg.Module() + logicalType := mod.NamedFunction("contracts.F").GlobalValueType() + if err := prog.MaterializeValueAttributes(mod); err != nil { + t.Fatal(err) + } + NewTransformer(prog, mod.Target(), "", false).TransformModule("contracts", mod) + physical := mod.NamedFunction("contracts.F") + + if physical.GlobalValueType() == logicalType { + t.Fatal("test did not exercise aggregate ABI rewriting") + } + if physical.GetEnumAttributeAtIndex(1, llvm.AttributeKindID("sret")).IsNil() { + t.Fatalf("large source results were not transported via sret:\n%s", physical.String()) + } + if target.GOARCH == "amd64" && physical.GetEnumAttributeAtIndex(2, llvm.AttributeKindID("byval")).IsNil() { + t.Fatalf("amd64 input did not exercise native byval:\n%s", physical.String()) + } + if err := llvm.VerifyModule(mod, llvm.ReturnStatusAction); err != nil { + t.Fatalf("invalid transformed contracts: %v\n%s", err, mod.String()) + } + options := llvm.NewPassBuilderOptions() + defer options.Dispose() + if err := mod.RunPasses("default", prog.TargetMachine(), options); err != nil { + t.Fatal(err) + } + for _, name := range []string{"contracts.Caller", "contracts.PackedInput"} { + body := mod.NamedFunction(name).String() + if !strings.Contains(body, "ret i1 false") { + t.Fatalf("%s did not consume reconstructed value facts:\n%s", name, body) + } + } + if !strings.Contains(mod.NamedFunction("contracts.Caller").String(), "@contracts.F(") { + t.Fatal("result facts removed the unknown external invocation") + } + if err := llvm.VerifyModule(mod, llvm.ReturnStatusAction); err != nil { + t.Fatal(err) + } + assembly, err := prog.TargetMachine().EmitToMemoryBuffer(mod, llvm.AssemblyFile) + if err != nil { + t.Fatalf("cannot emit transformed contract assembly: %v", err) + } + assembly.Dispose() + }) + } +} diff --git a/internal/funcattrs/llvm.go b/internal/funcattrs/llvm.go index 5e70f74dc0..099cb25dbf 100644 --- a/internal/funcattrs/llvm.go +++ b/internal/funcattrs/llvm.go @@ -13,6 +13,9 @@ func Apply(ctx llvm.Context, fn llvm.Value, sig *types.Signature, attrs []Attrib return err } for _, attr := range attrs { + if attr.Name == "sameas" || (attr.Target.Scope == Result && sig.Results().Len() != 1) { + continue + } index := 1 + environment if attr.Target.Scope == Result { index = 0 @@ -72,7 +75,10 @@ func Apply(ctx llvm.Context, fn llvm.Value, sig *types.Signature, attrs []Attrib // CopyValueAttributes is used only when ABI conversion preserves the whole value. func CopyValueAttributes(from, to llvm.Value, old, new int) { - for _, name := range []string{"nonnull", "range"} { + for _, name := range []string{"nonnull", "range", "returned"} { + if name == "returned" && from.GlobalValueType().ReturnType() != to.GlobalValueType().ReturnType() { + continue + } if attr := from.GetEnumAttributeAtIndex(old, llvm.AttributeKindID(name)); !attr.IsNil() { to.AddAttributeAtIndex(new, attr) } diff --git a/internal/funcattrs/source.go b/internal/funcattrs/source.go index b98581b9c3..57026be7e0 100644 --- a/internal/funcattrs/source.go +++ b/internal/funcattrs/source.go @@ -60,6 +60,7 @@ type Attribute struct { Args string // Canonical spelling; backends consume the typed operands below. Position token.Position Range *RangeBounds + From *Target // Ordinary parameter entry value used by sameas. } func (a Attribute) Error(format string, args ...any) error { @@ -76,7 +77,7 @@ func IsSourceDirective(d directive.Directive) bool { name = name[:i] } switch name { - case "param", "result", "receiver", "nonnull", "range", "nonnegative": + case "param", "result", "receiver", "nonnull", "range", "nonnegative", "sameas": return true } return false @@ -101,9 +102,7 @@ func Parse(fset *token.FileSet, decl *ast.FuncDecl) ([]Attribute, error) { if err != nil { return nil, base.Error("%v", err) } - if base.Target.Scope == Result && fieldCount(decl.Type.Results) != 1 { - return nil, base.Error("result attributes currently require exactly one source result") - } + words = words[1:] if len(words) == 0 { return nil, base.Error("expected an attribute after selector") @@ -118,6 +117,16 @@ func Parse(fset *token.FileSet, decl *ast.FuncDecl) ([]Attribute, error) { return nil, a.Error("%v", err) } + if a.Name == "sameas" { + if !token.IsIdentifier(a.Args) || a.Args == "_" { + return nil, a.Error("sameas expects a parameter name") + } + index, err := selectIndex(decl.Type.Params, a.Args) + if err != nil { + return nil, a.Error("sameas: %v", err) + } + a.From = &Target{Scope: Parameter, Index: index} + } a, err = normalize(a) if err != nil { return nil, err @@ -287,6 +296,8 @@ func normalize(a Attribute) (Attribute, error) { valid = input || result case "range": valid, takesArgs = input || result, true + case "sameas": + valid, takesArgs = result, true default: return a, a.Error("unsupported attribute %q", a.Name) } @@ -306,6 +317,12 @@ func normalize(a Attribute) (Attribute, error) { a.Args = lo.String() + "," + hi.String() } + case "sameas": + if a.From == nil || a.From.Scope != Parameter { + err = fmt.Errorf("sameas requires a parameter name") + } else { + a.Args = a.From.String() + } } if err != nil { return a, a.Error("%v", err) @@ -325,6 +342,10 @@ func Merge(attrs ...[]Attribute) ([]Attribute, error) { return nil, err } + if a.From != nil { + from := *a.From + a.From = &from + } duplicate := false for _, prev := range out { if !prev.Target.Equal(a.Target) { @@ -437,6 +458,23 @@ func Validate(attrs []Attribute, sig *types.Signature, intBits int, deferTypePar return a.Error("%s requires a pointer, got %s", a.Name, t) } + case "sameas": + if !pointer(t) && !integer(t) { + return a.Error("sameas requires an integer or pointer result, got %s", t) + } + if a.From == nil { + return a.Error("sameas requires an input selector") + } + from, err := ResolveTarget(sig, *a.From) + if err != nil { + return a.Error("sameas input: %v", err) + } + if unresolved(from) && deferTypeParams { + continue + } + if !(pointer(t) && pointer(from)) && !types.Identical(types.Unalias(t), types.Unalias(from)) { + return a.Error("sameas requires compatible pointer types or identical integer source types, got %s and %s", from, t) + } case "range", "nonnegative": if _, _, _, err := IntegerRange(a, t, intBits); err != nil { return err diff --git a/internal/funcattrs/source_test.go b/internal/funcattrs/source_test.go index 6f282c0997..f883f437fe 100644 --- a/internal/funcattrs/source_test.go +++ b/internal/funcattrs/source_test.go @@ -112,7 +112,6 @@ func TestSourceValueDiagnostics(t *testing.T) { {"receiver nonnull", "func F() {}", "requires a method"}, {"param(0) nonnull", "func F(p string) {}", "requires a pointer"}, {"result", "func F() {}", "exactly one source result"}, - {"result(0) nonnull", "func F() (*int,*int) { return nil,nil }", "exactly one source result"}, {"result", "func F() *int { return nil }", "expected an attribute"}, {"param(0) nonnull)", "func F(p *int) {}", "unbalanced"}, {"param(0) nonnull(", "func F(p *int) {}", "unbalanced"}, @@ -158,3 +157,27 @@ func TestGenericValueAttributeValidation(t *testing.T) { } } } + +func TestSourceResultRelations(t *testing.T) { + for _, tc := range []struct{ selector, signature, want string }{ + {"result sameas(p)", "func F(p *int) *int { return p }", ""}, + {"result(out) sameas(p) nonnull", "func F(p *int) (out *int,n int) { return p,0 }", ""}, + {"result sameas(p)", "func F(p int32) uint32 { return 0 }", "identical integer"}, + {"result sameas(p)", "func F(p bool) bool { return p }", "integer or pointer"}, + {"result sameas(0)", "func F(p int) int { return p }", "parameter name"}, + {"result sameas(missing)", "func F(p int) int { return p }", "unknown source value"}, + {"param(p) sameas(p)", "func F(p int) int { return p }", "not supported on parameter"}, + } { + attrs, sig, err := parseTest(t, "//llgo:"+tc.selector+"\n"+tc.signature) + if err == nil { + err = Validate(attrs, sig, 64, false) + } + if tc.want == "" { + if err != nil { + t.Fatal(err) + } + } else if err == nil || !strings.Contains(err.Error(), tc.want) { + t.Fatalf("%s: %v", tc.selector, err) + } + } +} diff --git a/internal/funcattrs/values.go b/internal/funcattrs/values.go new file mode 100644 index 0000000000..a5d22ec9f4 --- /dev/null +++ b/internal/funcattrs/values.go @@ -0,0 +1,277 @@ +package funcattrs + +import ( + "fmt" + "go/types" + + "github.com/xgo-dev/llvm" +) + +type resultFact struct { + Name string + Path []int // The whole Go result within the logical return tuple. + From int // LLVM input index for sameas, or -1 for an independent value fact. + Lower, Upper uint64 +} + +// ValuePlan belongs to one LLVM function before ABI conversion. +type ValuePlan []resultFact + +// ResultPathResolver accounts for padding wrappers in a logical Go result tuple. +type ResultPathResolver func(index int) []int + +func PrepareResultAttributes(ctx llvm.Context, fn llvm.Value, sig *types.Signature, attrs []Attribute, environment, intBits int, resolver ResultPathResolver) (ValuePlan, error) { + var plan ValuePlan + for _, source := range attrs { + if source.Target.Scope != Result { + continue + } + fact := resultFact{Name: source.Name, From: -1} + if sig.Results().Len() > 1 { + fact.Path = []int{source.Target.Index} + if resolver != nil { + fact.Path = resolver(source.Target.Index) + } + } + typ := fn.GlobalValueType().ReturnType() + for _, index := range fact.Path { + if typ.TypeKind() != llvm.StructTypeKind { + return nil, source.Error("selected result has no LLVM tuple") + } + fields := typ.StructElementTypes() + if index < 0 || index >= len(fields) { + return nil, source.Error("selected result has no LLVM value") + } + typ = fields[index] + } + selected, err := ResolveTarget(sig, source.Target) + if err != nil { + return nil, source.Error("%v", err) + } + switch source.Name { + case "nonnull": + if typ.TypeKind() != llvm.PointerTypeKind { + return nil, source.Error("selected result is not an LLVM pointer") + } + case "range", "nonnegative": + bits, bounds, full, err := IntegerRange(source, selected, intBits) + if err != nil { + return nil, err + } + if typ.TypeKind() != llvm.IntegerTypeKind || typ.IntTypeWidth() != bits { + return nil, source.Error("selected result does not have its source integer width") + } + if full { + continue + } + fact.Lower, fact.Upper = bounds[0], bounds[1] + case "sameas": + if source.From == nil { + return nil, source.Error("sameas requires a parameter name") + } + fact.From = source.From.Index + environment + if sig.Recv() != nil { + fact.From++ + } + if fact.From < 0 || fact.From >= fn.ParamsCount() || fn.Param(fact.From).Type() != typ { + return nil, source.Error("sameas values do not share an LLVM representation") + } + // Current collectors do not move objects or native stacks. + if len(fact.Path) == 0 { + fn.AddAttributeAtIndex(fact.From+1, ctx.CreateEnumAttribute(llvm.AttributeKindID("returned"), 0)) + } + } + plan = append(plan, fact) + } + return plan, nil +} + +// MaterializeValueContracts applies result guarantees on normal continuations +// before ABI conversion reconstructs the values. No plan is serialized to bitcode. +func MaterializeValueContracts(m llvm.Module, plans map[llvm.Value]ValuePlan) error { + if len(plans) == 0 { + return nil + } + ctx := m.Context() + b := ctx.NewBuilder() + defer b.Dispose() + var calls []llvm.Value + for fn := m.FirstFunction(); !fn.IsNil(); fn = llvm.NextFunction(fn) { + for bb := fn.FirstBasicBlock(); !bb.IsNil(); bb = llvm.NextBasicBlock(bb) { + for instr := bb.FirstInstruction(); !instr.IsNil(); instr = llvm.NextInstruction(instr) { + if !instr.IsACallInst().IsNil() || !instr.IsAInvokeInst().IsNil() { + if _, ok := plans[instr.CalledValue()]; ok { + calls = append(calls, instr) + } + } + } + } + } + for _, call := range calls { + if call.CalledFunctionType() != call.CalledValue().GlobalValueType() { + return fmt.Errorf("value contract call to %s has an incompatible logical prototype", call.CalledValue().Name()) + } + } + for _, call := range calls { + materializeCallResults(b, call, plans[call.CalledValue()]) + } + + return nil +} + +func extractValue(b llvm.Builder, value llvm.Value, path []int) llvm.Value { + for _, index := range path { + value = b.CreateExtractValue(value, index, "contract.value") + } + return value +} + +func materializeCallResults(b llvm.Builder, call llvm.Value, plan ValuePlan) { + results := plan + continuation := normalContinuation(b, call) + // Capture existing users before constructing replacement aggregates: a + // blanket RAUW would replace their old-call operand with themselves. + var users []llvm.Value + seen := make(map[llvm.Value]bool) + for use := call.FirstUse(); !use.IsNil(); use = use.NextUse() { + if user := use.User(); !seen[user] { + seen[user] = true + users = append(users, user) + } + } + b.SetInsertPointBefore(continuation) + result := call + for _, contract := range results { + if contract.From < 0 { + continue + } + // Arguments are already evaluated SSA snapshots. In particular, never + // reload a mutable parameter home after the call. Pointer forwarding + // retains the original access provenance, unlike address equality alone. + input := call.Operand(contract.From) + if len(contract.Path) == 0 { + result = input + continue + } + // Forward existing leaf projections, not the aggregate's unrelated + // fields or padding. Rebuilding a giant aggregate with insertvalue would + // defeat the ABI's indirect copies solely to carry a relation. Whole + // value transfers can retain their original result: its field already + // has the promised identity. The equality fact also relates later + // projections without inventing pointer access provenance. + leaves := matchingProjections(call, contract.Path) + actual := extractValue(b, call, contract.Path) + equal := b.CreateICmp(llvm.IntEQ, actual, input, "contract.same") + b.CreateIntrinsic(call.Type().Context().VoidType(), llvm.LookupIntrinsicID("llvm.assume"), []llvm.Value{equal}, "") + for _, leaf := range leaves { + leaf.ReplaceAllUsesWith(input) + } + } + for _, contract := range results { + if contract.From < 0 { + value := extractValue(b, result, contract.Path) + emitValueFact(b, value, contract) + } + } + if result != call { + for _, user := range users { + for i := 0; i < user.OperandsCount(); i++ { + if user.Operand(i) == call { + user.SetOperand(i, result) + } + } + } + } + // Postconditions require a normal continuation; retaining a mandatory + // tail-call form would leave no legal place for reconstruction and facts. + if !call.IsACallInst().IsNil() { + call.SetTailCall(false) + } +} + +func matchingProjections(value llvm.Value, path []int) []llvm.Value { + var matches []llvm.Value + for use := value.FirstUse(); !use.IsNil(); use = use.NextUse() { + extract := use.User().IsAExtractValueInst() + if extract.IsNil() { + continue + } + indices := extract.Indices() + if len(indices) > len(path) { + continue + } + prefix := true + for i, index := range indices { + if int(index) != path[i] { + prefix = false + break + } + } + if !prefix { + continue + } + if len(indices) == len(path) { + matches = append(matches, extract) + } else { + matches = append(matches, matchingProjections(extract, path[len(indices):])...) + } + } + return matches +} + +// normalContinuation makes invoke facts local to its normal edge. Splitting +// unconditionally also handles shared normal destinations and their PHIs. +func normalContinuation(b llvm.Builder, call llvm.Value) llvm.Value { + if call.IsAInvokeInst().IsNil() { + return llvm.NextInstruction(call) + } + parent := call.InstructionParent() + normal := call.Successor(0) + ctx := call.Type().Context() + edge := ctx.AddBasicBlock(parent.Parent(), "contract.normal") + b.SetInsertPointAtEnd(edge) + branch := b.CreateBr(normal) + for i := 0; i < call.OperandsCount(); i++ { + if call.Operand(i) == normal.AsValue() { + call.SetOperand(i, edge.AsValue()) + break + } + } + for phi := normal.FirstInstruction(); !phi.IsNil() && !phi.IsAPHINode().IsNil(); { + next := llvm.NextInstruction(phi) + values := make([]llvm.Value, phi.IncomingCount()) + blocks := make([]llvm.BasicBlock, len(values)) + for i := range values { + values[i], blocks[i] = phi.IncomingValue(i), phi.IncomingBlock(i) + if blocks[i] == parent { + blocks[i] = edge + } + } + b.SetInsertPointBefore(phi) + replacement := b.CreatePHI(phi.Type(), "contract.merge") + replacement.AddIncoming(values, blocks) + phi.ReplaceAllUsesWith(replacement) + phi.EraseFromParentAsInstruction() + phi = next + } + return branch +} + +func emitValueFact(b llvm.Builder, value llvm.Value, contract resultFact) { + ctx := value.Type().Context() + var predicate llvm.Value + switch contract.Name { + case "nonnull": + predicate = b.CreateICmp(llvm.IntNE, value, llvm.ConstNull(value.Type()), "contract.nonnull") + case "range", "nonnegative": + // Subtraction in the source bit width makes signed intervals that cross + // zero work without constraining any ABI carrier or its padding bits. + lower := llvm.ConstInt(value.Type(), contract.Lower, false) + width := llvm.ConstInt(value.Type(), contract.Upper-contract.Lower, false) + distance := b.CreateSub(value, lower, "contract.range.offset") + predicate = b.CreateICmp(llvm.IntULT, distance, width, "contract.range") + default: + return + } + b.CreateIntrinsic(ctx.VoidType(), llvm.LookupIntrinsicID("llvm.assume"), []llvm.Value{predicate}, "") +} diff --git a/internal/funcattrs/values_test.go b/internal/funcattrs/values_test.go new file mode 100644 index 0000000000..2b7aad434e --- /dev/null +++ b/internal/funcattrs/values_test.go @@ -0,0 +1,202 @@ +package funcattrs + +import ( + "os" + "path/filepath" + "strings" + "testing" + + "github.com/xgo-dev/llvm" +) + +func valueTestModule(t *testing.T, ir string) llvm.Module { + t.Helper() + ctx := llvm.NewContext() + t.Cleanup(ctx.Dispose) + path := filepath.Join(t.TempDir(), "contracts.ll") + if err := os.WriteFile(path, []byte(ir), 0o644); err != nil { + t.Fatal(err) + } + buf, err := llvm.NewMemoryBufferFromFile(path) + if err != nil { + t.Fatal(err) + } + mod, err := ctx.ParseIR(buf) + if err != nil { + t.Fatal(err) + } + t.Cleanup(mod.Dispose) + return mod +} + +func attachValueTestContracts(t *testing.T, mod llvm.Module, name, source string) map[llvm.Value]ValuePlan { + t.Helper() + attrs, sig, err := parseTest(t, source) + if err != nil { + t.Fatal(err) + } + if err := Apply(mod.Context(), mod.NamedFunction(name), sig, attrs, 0, 64); err != nil { + t.Fatal(err) + } + plan, err := PrepareResultAttributes(mod.Context(), mod.NamedFunction(name), sig, attrs, 0, 64, nil) + if err != nil { + t.Fatal(err) + } + if len(plan) == 0 { + return nil + } + return map[llvm.Value]ValuePlan{mod.NamedFunction(name): plan} +} + +func optimizeValueTest(t *testing.T, mod llvm.Module) { + t.Helper() + if err := llvm.VerifyModule(mod, llvm.ReturnStatusAction); err != nil { + t.Fatalf("invalid contract IR: %v\n%s", err, mod.String()) + } + options := llvm.NewPassBuilderOptions() + defer options.Dispose() + if err := mod.RunPasses("default", llvm.TargetMachine{}, options); err != nil { + t.Fatal(err) + } + if err := llvm.VerifyModule(mod, llvm.ReturnStatusAction); err != nil { + t.Fatalf("invalid optimized contract IR: %v\n%s", err, mod.String()) + } +} + +func TestValueContractsMaterializeMultipleResults(t *testing.T) { + mod := valueTestModule(t, ` +declare { ptr, i8 } @F(ptr) + +define i1 @caller(ptr %p) { +entry: + %r = call { ptr, i8 } @F(ptr %p) + %result.p = extractvalue { ptr, i8 } %r, 0 + %result.n = extractvalue { ptr, i8 } %r, 1 + %nil = icmp eq ptr %result.p, null + %small = icmp slt i8 %result.n, -3 + %large = icmp sge i8 %result.n, 4 + %range.bad = or i1 %small, %large + %bad = or i1 %nil, %range.bad + ret i1 %bad +} + +define i1 @relation(ptr %p) { +entry: + %r = call { ptr, i8 } @F(ptr %p) + %result.p = extractvalue { ptr, i8 } %r, 0 + %same = icmp eq ptr %result.p, %p + ret i1 %same +} +`) + plans := attachValueTestContracts(t, mod, "F", ` +//llgo:result(out) nonnull sameas(p) +//llgo:result(n) range(-3,4) +func F(p *int) (out *int, n int8) +`) + if err := MaterializeValueContracts(mod, plans); err != nil { + t.Fatal(err) + } + optimizeValueTest(t, mod) + for name, want := range map[string]string{"caller": "ret i1 false", "relation": "ret i1 true"} { + body := mod.NamedFunction(name).String() + if !strings.Contains(body, want) || !strings.Contains(body, "@F(") { + t.Fatalf("%s failed to use the result facts while preserving the call:\n%s", name, body) + } + } +} + +func TestValueContractsEntryIntersection(t *testing.T) { + mod := valueTestModule(t, ` +define i1 @Input(ptr %pointer, i8 %n, i16 %element) { +entry: + %nil = icmp eq ptr %pointer, null + %negative = icmp slt i8 %n, 0 + %large = icmp sge i8 %n, 5 + %element.bad = icmp sgt i16 %element, 9 + %a = or i1 %nil, %negative + %c = or i1 %a, %large + %bad = or i1 %c, %element.bad + ret i1 %bad +} +`) + plans := attachValueTestContracts(t, mod, "Input", ` +//llgo:param(p) nonnull +//llgo:param(n) range(-4,5) nonnegative +//llgo:param(element) range(0,10) +func Input(p *int, n int8, element int16) bool +`) + if err := MaterializeValueContracts(mod, plans); err != nil { + t.Fatal(err) + } + optimizeValueTest(t, mod) + if body := mod.NamedFunction("Input").String(); !strings.Contains(body, "ret i1 false") { + t.Fatalf("entry facts or their intersection did not reach LLVM:\n%s", body) + } +} + +func TestResultContractDoesNotRemovePanickingCall(t *testing.T) { + mod := valueTestModule(t, ` +declare ptr @Checked(ptr) +define i1 @nil_input() { +entry: + %r = call ptr @Checked(ptr null) + %nonnull = icmp ne ptr %r, null + ret i1 %nonnull +} +`) + plans := attachValueTestContracts(t, mod, "Checked", ` +//llgo:result(0) nonnull sameas(p) +func Checked(p *int) *int +`) + if err := MaterializeValueContracts(mod, plans); err != nil { + t.Fatal(err) + } + optimizeValueTest(t, mod) + if body := mod.NamedFunction("nil_input").String(); !strings.Contains(body, "@Checked(ptr null)") { + t.Fatalf("normal-result contract deleted the required call on nil input:\n%s", body) + } +} + +func TestInvokeResultFactsOnlyOnNormalEdge(t *testing.T) { + mod := valueTestModule(t, ` +declare i32 @__gxx_personality_v0(...) +declare ptr @Checked(ptr) +define ptr @caller(i1 %should.call, ptr %p) personality ptr @__gxx_personality_v0 { +entry: + br i1 %should.call, label %try, label %join +try: + %r = invoke ptr @Checked(ptr %p) to label %join unwind label %unwind +join: + %result = phi ptr [ %r, %try ], [ null, %entry ] + ret ptr %result +unwind: + %exception = landingpad { ptr, i32 } cleanup + resume { ptr, i32 } %exception +} +`) + plans := attachValueTestContracts(t, mod, "Checked", ` +//llgo:result(0) nonnull sameas(p) +func Checked(p *int) *int +`) + if err := MaterializeValueContracts(mod, plans); err != nil { + t.Fatal(err) + } + if err := llvm.VerifyModule(mod, llvm.ReturnStatusAction); err != nil { + t.Fatalf("invoke result materialization invalid: %v\n%s", err, mod.String()) + } + fn := mod.NamedFunction("caller") + assumptions := 0 + for bb := fn.FirstBasicBlock(); !bb.IsNil(); bb = llvm.NextBasicBlock(bb) { + for instr := bb.FirstInstruction(); !instr.IsNil(); instr = llvm.NextInstruction(instr) { + if !instr.IsACallInst().IsNil() && instr.CalledValue().Name() == "llvm.assume" { + assumptions++ + if !strings.HasPrefix(bb.AsValue().Name(), "contract.normal") { + t.Fatalf("postcondition escaped normal invoke edge:\n%s", fn.String()) + } + } + } + } + if assumptions != 1 || !strings.Contains(fn.String(), "[ null, %entry ]") { + t.Fatalf("normal edge or unrelated PHI input was lost:\n%s", fn.String()) + } +} diff --git a/runtime/internal/runtime/z_error.go b/runtime/internal/runtime/z_error.go index f69239b63d..e8c88ac4ce 100644 --- a/runtime/internal/runtime/z_error.go +++ b/runtime/internal/runtime/z_error.go @@ -94,7 +94,7 @@ func AssertNilDeref(b bool) { } } -//llgo:result nonnull +//llgo:result nonnull sameas(ptr) func AssertNilDerefPtr(ptr unsafe.Pointer) unsafe.Pointer { AssertNilDeref(ptr == nil) return ptr diff --git a/runtime/internal/runtime/z_string.go b/runtime/internal/runtime/z_string.go index 839d19fb86..1fef395566 100644 --- a/runtime/internal/runtime/z_string.go +++ b/runtime/internal/runtime/z_string.go @@ -47,6 +47,8 @@ func StringCat(a, b String) String { // ----------------------------------------------------------------------------- // CStrCopy copies a Go string to a C string buffer and returns it. +// +//llgo:result sameas(dest) func CStrCopy(dest unsafe.Pointer, s String) *int8 { n := s.len c.Memcpy(dest, s.data, uintptr(n)) diff --git a/runtime/internal/test/result_attributes_test.go b/runtime/internal/test/result_attributes_test.go new file mode 100644 index 0000000000..e1087c5f05 --- /dev/null +++ b/runtime/internal/test/result_attributes_test.go @@ -0,0 +1,106 @@ +package test + +import "testing" + +type attributeContainer struct { + P *int + N uint32 + Payload [64]byte +} + +//go:noinline +//llgo:param(p) nonnull +//llgo:result(pointer) nonnull sameas(p) +//llgo:result(count) range(0,64) +func attributeAggregate(input attributeContainer, p *int, n uint32) (pointer *int, count uint32, result attributeContainer) { + input.P = p + input.N = n & 63 + return p, input.N, input +} + +//go:noinline +//llgo:result(0) sameas(p) +func attributeEntrySnapshot(p *int, source *attributeContainer, replacement *int) *int { + source.P = replacement + return p +} + +type attributePacked struct { + Signed int8 + Count uint8 + Bytes [6]byte +} + +//go:noinline +//llgo:param(n) range(-3,5) +//llgo:result(signed) sameas(n) +//llgo:result(count) range(0,64) +func attributePackedRoundTrip(input attributePacked, n int8) (signed int8, count uint8, result attributePacked) { + input.Signed = n + input.Count &= 63 + return n, input.Count, input +} + +type attributeLarge struct { + P *int + Payload [10000]uint64 + Count uint32 +} + +//go:noinline +//llgo:param(p) nonnull +func attributeLargeResult(p *int) (out attributeLarge) { + out.P = p + out.Payload[0], out.Payload[9999] = 19, 101 + out.Count = 7 + return +} + +func TestSourceContractsAfterABI(t *testing.T) { + value := 17 + in := attributeContainer{P: &value, N: 999, Payload: [64]byte{0: 1, 63: 2}} + for n := uint32(0); n < 130; n++ { + p, count, out := attributeAggregate(in, &value, n) + if p != &value || count != n%64 || out.P != &value || out.N != n%64 || out.Payload != in.Payload { + t.Fatal("indirect arguments or multiple results changed across ABI conversion") + } + if in.N != 999 { + t.Fatal("callee changed the caller's by-value input") + } + } + replacement := 29 + old := attributeEntrySnapshot(in.P, &in, &replacement) + if old != &value || in.P != &replacement { + t.Fatal("sameas reloaded changed input memory") + } + for n := int8(-3); n < 5; n++ { + in := attributePacked{Signed: n, Count: 193, Bytes: [6]byte{0: 11, 5: 253}} + signed, count, out := attributePackedRoundTrip(in, n) + if signed != n || count != 1 || out.Signed != n || out.Count != 1 || out.Bytes != in.Bytes { + t.Fatalf("packed argument or multiple results changed unrelated data: %#v", out) + } + } + large := attributeLargeResult(&value) + if large.P != &value || large.Count != 7 || large.Payload[0] != 19 || large.Payload[9999] != 101 { + t.Fatal("large indirect result lost values") + } +} + +//go:noinline +//llgo:result(out) sameas(p) +func deferredResultRelation(p *int) (out *int, n int) { + defer func() { + if recover() != nil { + out = p + n = 17 + } + }() + panic("deferred result") +} + +func TestResultRelationAfterDeferredRecovery(t *testing.T) { + x := 42 + if p, n := deferredResultRelation(&x); p != &x || n != 17 { + t.Fatalf("deferred result: %p %d", p, n) + } +} diff --git a/runtime/internal/test/value_attributes_test.go b/runtime/internal/test/value_attributes_test.go index 08b51c4b8c..e6bf92845d 100644 --- a/runtime/internal/test/value_attributes_test.go +++ b/runtime/internal/test/value_attributes_test.go @@ -3,7 +3,7 @@ package test import "testing" //go:noinline -//llgo:result nonnull +//llgo:result nonnull sameas(p) func checkedValueAttribute[T any](p *T) *T { if p == nil { panic("nil value attribute") diff --git a/ssa/decl.go b/ssa/decl.go index 635427579c..f49703265d 100644 --- a/ssa/decl.go +++ b/ssa/decl.go @@ -384,7 +384,7 @@ func (p Package) newFunc( } fn := llvm.AddFunction(p.mod, llvmName, t.ll) p.Prog.applyFunctionAttributes(fn, name) - p.Prog.applyValueAttributes(fn, name, sig, env != nil) + p.Prog.applyValueAttributes(fn, name, sig, env != nil, bg) switch name { case "github.com/xgo-dev/llgo/runtime/internal/runtime.AllocU", "github.com/xgo-dev/llgo/runtime/internal/runtime.AllocZ", diff --git a/ssa/package.go b/ssa/package.go index b293d177ed..bcd8a151da 100644 --- a/ssa/package.go +++ b/ssa/package.go @@ -26,6 +26,7 @@ import ( "unsafe" "github.com/xgo-dev/llgo/internal/env" + "github.com/xgo-dev/llgo/internal/funcattrs" "github.com/xgo-dev/llgo/internal/meta" "github.com/xgo-dev/llgo/internal/optlevel" "github.com/xgo-dev/llgo/ssa/abi" @@ -117,10 +118,11 @@ func Initialize(flags InitFlags) { // ----------------------------------------------------------------------------- type aProgram struct { - ctx llvm.Context - typs typeutil.Map // rawType -> Type - sizes types.Sizes // provided by Go compiler - gocvt goTypes + valuePlans map[llvm.Value]funcattrs.ValuePlan + ctx llvm.Context + typs typeutil.Map // rawType -> Type + sizes types.Sizes // provided by Go compiler + gocvt goTypes patchType func(types.Type) types.Type diff --git a/ssa/value_attributes.go b/ssa/value_attributes.go index 387d425399..9f8e9f2a32 100644 --- a/ssa/value_attributes.go +++ b/ssa/value_attributes.go @@ -1,10 +1,11 @@ package ssa import ( - "github.com/xgo-dev/llgo/internal/funcattrs" - "github.com/xgo-dev/llvm" "go/types" "sort" + + "github.com/xgo-dev/llgo/internal/funcattrs" + "github.com/xgo-dev/llvm" ) func (p Program) SetValueAttributes(name string, attrs []funcattrs.Attribute) error { @@ -42,7 +43,7 @@ func (p Program) valueAttributes(name string) ([]funcattrs.Attribute, error) { return funcattrs.Merge(sets...) } -func (p Program) applyValueAttributes(fn llvm.Value, name string, sig *types.Signature, environment bool) { +func (p Program) applyValueAttributes(fn llvm.Value, name string, sig *types.Signature, environment bool, bg Background) { attrs, err := p.valueAttributes(name) if err != nil { panic(err) @@ -57,4 +58,38 @@ func (p Program) applyValueAttributes(fn llvm.Value, name string, sig *types.Sig if err := funcattrs.Apply(p.ctx, fn, sig, attrs, offset, p.Int().ll.IntTypeWidth()); err != nil { panic(err) } + plan, err := funcattrs.PrepareResultAttributes(p.ctx, fn, sig, attrs, offset, p.Int().ll.IntTypeWidth(), func(index int) []int { + path := []int{index} + converted := p.FuncDecl(sig, bg).raw.Type.(*types.Signature) + if layout, ok := p.structLayout(p.retType(converted)); ok && layout.wrapped[index] { + path = append(path, 0) + } + return path + }) + if err != nil { + panic(err) + } + if len(plan) != 0 { + if p.valuePlans == nil { + p.valuePlans = make(map[llvm.Value]funcattrs.ValuePlan) + } + p.valuePlans[fn] = plan + } +} + +// MaterializeValueAttributes consumes only plans belonging to this module. +func (p Program) MaterializeValueAttributes(m llvm.Module) error { + plans := make(map[llvm.Value]funcattrs.ValuePlan) + for fn := m.FirstFunction(); !fn.IsNil(); fn = llvm.NextFunction(fn) { + if plan, ok := p.valuePlans[fn]; ok { + plans[fn] = plan + } + } + if err := funcattrs.MaterializeValueContracts(m, plans); err != nil { + return err + } + for fn := range plans { + delete(p.valuePlans, fn) + } + return nil }