diff --git a/VERSION b/VERSION index 7c0802125aecdb..8a3457effbde79 100644 --- a/VERSION +++ b/VERSION @@ -1,2 +1,2 @@ -go1.27.1 -time 2026-08-28T16:20:06Z +go1.27.2 +time 2026-10-02T20:28:03Z diff --git a/doc/godebug.md b/doc/godebug.md index de47d3a91cbfe4..49b513c4cb559e 100644 --- a/doc/godebug.md +++ b/doc/godebug.md @@ -154,6 +154,22 @@ for example, see the [runtime documentation](/pkg/runtime#hdr-Environment_Variables) and the [go command documentation](/cmd/go#hdr-Build_and_test_caching). +### Go 1.28 + +Go 1.28 changed +[`net/http.FileServer`](/pkg/net/http#FileServer), +[`net/http.FileServerFS`](/pkg/net/http#FileServerFS), +[`net/http.ServeContent`](/pkg/net/http#ServeContent), +[`net/http.ServeFile`](/pkg/net/http#ServeFile), and +[`net/http.ServeFileFS`](/pkg/net/http#ServeFileFS) to +limit the maximum number of ranges in a Range header. +When the number of ranges in a header exceeds the new +`httpservecontentmaxranges` setting, the header is ignored. +The default value is `httpservecontentmaxranges=200`. +Setting `httpservecontentmaxranges=0` disables the limit. +To avoid denial of service attacks, this setting and default +was backported to Go 1.27.2 and Go 1.26.9. + ### Go 1.27 Go 1.27 removed the `gotypesalias` setting, as noted in the [Go 1.22](#go-122) section. @@ -199,7 +215,7 @@ We plan to remove this setting in Go 1.31. Go 1.26 added a new `httpcookiemaxnum` setting that controls the maximum number of cookies that net/http will accept when parsing HTTP headers. If the number of -cookie in a header exceeds the number set in `httpcookiemaxnum`, cookie parsing +cookies in a header exceeds the number set in `httpcookiemaxnum`, cookie parsing will fail early. The default value is `httpcookiemaxnum=3000`. Setting `httpcookiemaxnum=0` will allow the cookie parsing to accept an indefinite number of cookies. To avoid denial of service attacks, this setting and default diff --git a/lib/fips140/fips140.sum b/lib/fips140/fips140.sum index 050957af603e13..5dff95b1b62e9e 100644 --- a/lib/fips140/fips140.sum +++ b/lib/fips140/fips140.sum @@ -1,4 +1,7 @@ -# SHA256 checksums of snapshot zip files in this directory. +# Checksums of snapshot zip files in this directory. +# Each line is "NAME SHA256HEX H1HASH", where SHA256HEX is the SHA256 +# checksum of the zip file and H1HASH is the module zip hash of its +# contents, as recorded in go.sum files and the module cache. # These checksums are included in the FIPS security policy # (validation instructions sent to the lab) and MUST NOT CHANGE. # That is, the zip files themselves must not change. @@ -9,5 +12,5 @@ # # go test cmd/go/internal/fips140 -update # -v1.0.0-c2097c7c.zip daf3614e0406f67ae6323c902db3f953a1effb199142362a039e7526dfb9368b -v1.26.0.zip 9b28f847fdf1db4a36cb2b2f8ec09443c039383f085630a03ecfaddf6db7ea23 +v1.0.0-c2097c7c.zip daf3614e0406f67ae6323c902db3f953a1effb199142362a039e7526dfb9368b h1:YapsF69QWGwygaBqdI0rp7IfE5a6tAwepNAHONVcWqQ= +v1.26.0.zip 9b28f847fdf1db4a36cb2b2f8ec09443c039383f085630a03ecfaddf6db7ea23 h1:dtoPX1ALGGp4rMLzyh6oqIkYRXnXxRkpPWu56l5DFpM= diff --git a/src/cmd/cgo/internal/testplugin/plugin_test.go b/src/cmd/cgo/internal/testplugin/plugin_test.go index 3216073edbcb2d..4f0458b7dee472 100644 --- a/src/cmd/cgo/internal/testplugin/plugin_test.go +++ b/src/cmd/cgo/internal/testplugin/plugin_test.go @@ -430,3 +430,14 @@ func TestIssue75102(t *testing.T) { goCmd(t, "build", "-o", "issue75102.exe", "./issue75102/main.go") run(t, "./issue75102.exe") } + +func TestIssue81303(t *testing.T) { + // Issue 81303: the itab copies of a plugin must not hide the itabs + // of the host for the same interface/type pairs. + globalSkip(t) + goCmd(t, "build", "-buildmode=plugin", "-o", "issue81303p1.so", "./issue81303/plugin1.go") + goCmd(t, "build", "-buildmode=plugin", "-o", "issue81303p2.so", "./issue81303/plugin2.go") + goCmd(t, "build", "-buildmode=plugin", "-o", "issue81303p3.so", "./issue81303/plugin3.go") + goCmd(t, "build", "-o", "issue81303.exe", "./issue81303/main.go") + run(t, "./issue81303.exe") +} diff --git a/src/cmd/cgo/internal/testplugin/testdata/issue81303/main.go b/src/cmd/cgo/internal/testplugin/testdata/issue81303/main.go new file mode 100644 index 00000000000000..7b2ea1595f18ac --- /dev/null +++ b/src/cmd/cgo/internal/testplugin/testdata/issue81303/main.go @@ -0,0 +1,39 @@ +// Copyright 2026 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +// Issue 81303: a plugin has its own copies of the itabs of the host. +// The runtime added these copies to the itab table as second entries +// for the same interface/type pairs. After the table grew, a lookup +// could return the copy from the plugin. A type switch compares the +// itab with the itab of the host, so it took the default case. +// +// Each plugin imports package p, which uses go/ast, so each plugin +// adds a copy of every go/ast itab of the host. The host walks a +// syntax tree before and after each plugin loads. The three plugins +// are the same because plugin.Open loads a plugin path only once. + +package main + +import ( + "log" + "plugin" + + "testplugin/issue81303/p" +) + +func main() { + p.Walk() + for _, name := range []string{"issue81303p1.so", "issue81303p2.so", "issue81303p3.so"} { + pl, err := plugin.Open(name) + if err != nil { + log.Fatal(err) + } + f, err := pl.Lookup("F") + if err != nil { + log.Fatal(err) + } + f.(func())() + p.Walk() + } +} diff --git a/src/cmd/cgo/internal/testplugin/testdata/issue81303/p/p.go b/src/cmd/cgo/internal/testplugin/testdata/issue81303/p/p.go new file mode 100644 index 00000000000000..960d35e4376d5d --- /dev/null +++ b/src/cmd/cgo/internal/testplugin/testdata/issue81303/p/p.go @@ -0,0 +1,48 @@ +// Copyright 2026 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +package p + +import ( + "go/ast" + "go/parser" + "go/token" + "log" +) + +const src = `package p + +import "fmt" + +type T struct{ x int } + +func (t *T) M(n int) (r int) { + defer func() { r++ }() + go fmt.Println(n) + if n > 0 { + r = n + } else { + r = -n + } + for i := 0; i < n; i++ { + r += t.x + } + switch v := any(n).(type) { + case int: + r = v + } + return +} +` + +// Walk parses src and walks the syntax tree. ast.Walk converts each +// node to ast.Node, which looks up the itab in the itab table, and +// then compares that itab with the itab of the host in a type switch. +func Walk() { + f, err := parser.ParseFile(token.NewFileSet(), "src.go", src, parser.SkipObjectResolution) + if err != nil { + log.Fatal(err) + } + ast.Inspect(f, func(ast.Node) bool { return true }) +} diff --git a/src/cmd/cgo/internal/testplugin/testdata/issue81303/plugin1.go b/src/cmd/cgo/internal/testplugin/testdata/issue81303/plugin1.go new file mode 100644 index 00000000000000..c97a98583684ea --- /dev/null +++ b/src/cmd/cgo/internal/testplugin/testdata/issue81303/plugin1.go @@ -0,0 +1,11 @@ +// Copyright 2026 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +package main + +import "testplugin/issue81303/p" + +func main() {} + +func F() { p.Walk() } diff --git a/src/cmd/cgo/internal/testplugin/testdata/issue81303/plugin2.go b/src/cmd/cgo/internal/testplugin/testdata/issue81303/plugin2.go new file mode 100644 index 00000000000000..c97a98583684ea --- /dev/null +++ b/src/cmd/cgo/internal/testplugin/testdata/issue81303/plugin2.go @@ -0,0 +1,11 @@ +// Copyright 2026 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +package main + +import "testplugin/issue81303/p" + +func main() {} + +func F() { p.Walk() } diff --git a/src/cmd/cgo/internal/testplugin/testdata/issue81303/plugin3.go b/src/cmd/cgo/internal/testplugin/testdata/issue81303/plugin3.go new file mode 100644 index 00000000000000..c97a98583684ea --- /dev/null +++ b/src/cmd/cgo/internal/testplugin/testdata/issue81303/plugin3.go @@ -0,0 +1,11 @@ +// Copyright 2026 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +package main + +import "testplugin/issue81303/p" + +func main() {} + +func F() { p.Walk() } diff --git a/src/cmd/compile/internal/importer/ureader.go b/src/cmd/compile/internal/importer/ureader.go index 70ee01be39675b..fbe275be04c62b 100644 --- a/src/cmd/compile/internal/importer/ureader.go +++ b/src/cmd/compile/internal/importer/ureader.go @@ -6,6 +6,9 @@ package importer import ( + "cmp" + "slices" + "cmd/compile/internal/base" "cmd/compile/internal/syntax" "cmd/compile/internal/types2" @@ -459,12 +462,26 @@ func (pr *pkgReader) objIdx(idx pkgbits.Index) (*types2.Package, string) { // about it, so maybe we can avoid worrying about that here. underlying := r.typ().Underlying() - methods := make([]*types2.Func, r.Len()) - for i := range methods { - methods[i] = r.method(true) + type indexedMethod struct { + index int // (or -1 in V4) + fn *types2.Func } + var methods []indexedMethod if r.Version().Has(pkgbits.GenericMethods) { + // V4 (go1.27.0) emitted all non-generic methods + // before all generic ones, discarding source + // order: a bug (go.dev/issue/81188). + // V5 (go1.27.x) fixes it by emitting an explicit + // index along with each method. + + // ordinary methods + for range r.Len() { + idx, m := r.method(true) + methods = append(methods, indexedMethod{idx, m}) + } + + // generic methods for range r.Len() { // Careful: objIdx is used to read in package-scoped declarations, which // methods are not. Instead, decode it here. This makes it easier to @@ -479,17 +496,36 @@ func (pr *pkgReader) objIdx(idx pkgbits.Index) (*types2.Package, string) { pkg, name := t.selector() rtparams := t.typeParamNames(true, true) recv := t.param() + methodIdx := -1 + if r.Version().Has(pkgbits.PreserveMethodOrder) { + methodIdx = t.Len() + } tparams := t.typeParamNames(true, false) sig := t.signature(recv, rtparams, tparams) r.delayed = append(r.delayed, t.delayed...) // propagate before retiring - pr.retireReader(t) - methods = append(methods, types2.NewFunc(pos, pkg, name, sig)) + methods = append(methods, indexedMethod{methodIdx, types2.NewFunc(pos, pkg, name, sig)}) + } + + if r.Version().Has(pkgbits.PreserveMethodOrder) { + slices.SortFunc(methods, func(a, b indexedMethod) int { + return cmp.Compare(a.index, b.index) + }) + } + } else { + for range r.Len() { + _, m := r.method(true) + methods = append(methods, indexedMethod{-1, m}) } } - return tparams, underlying, methods, r.delayed + funcs := make([]*types2.Func, len(methods)) + for i, m := range methods { + funcs[i] = m.fn + } + + return tparams, underlying, funcs, r.delayed }) case pkgbits.ObjVar: @@ -606,8 +642,12 @@ func (r *reader) typeParamNames(isLazy bool, isGenMeth bool) []*types2.TypeParam return tparams } -func (r *reader) method(isLazy bool) *types2.Func { +func (r *reader) method(isLazy bool) (int, *types2.Func) { r.Sync(pkgbits.SyncMethod) + idx := -1 + if r.Version().Has(pkgbits.PreserveMethodOrder) { + idx = r.Len() + } pos := r.pos() pkg, name := r.selector() @@ -615,7 +655,7 @@ func (r *reader) method(isLazy bool) *types2.Func { sig := r.signature(r.param(), rtparams, nil) _ = r.pos() // TODO(mdempsky): Remove; this is a hacker for linker.go. - return types2.NewFunc(pos, pkg, name, sig) + return idx, types2.NewFunc(pos, pkg, name, sig) } func (r *reader) qualifiedIdent() (*types2.Package, string) { return r.ident(pkgbits.SyncSym) } diff --git a/src/cmd/compile/internal/noder/reader.go b/src/cmd/compile/internal/noder/reader.go index 4bafec1f721980..405f9c44ad372d 100644 --- a/src/cmd/compile/internal/noder/reader.go +++ b/src/cmd/compile/internal/noder/reader.go @@ -799,6 +799,9 @@ func (pr *pkgReader) objIdxMayFail(idx index, implicits, explicits []*types.Type sel = r.selector() r.recvTypeParamNames() recv = r.param() + if r.Version().Has(pkgbits.PreserveMethodOrder) { + _ = r.Len() // method index not needed in compiler + } } else { if sym.Name == "init" { sym = Renameinit() @@ -1130,6 +1133,9 @@ func (r *reader) typeParamNames() { func (r *reader) method(rext *reader) *types.Field { r.Sync(pkgbits.SyncMethod) + if r.Version().Has(pkgbits.PreserveMethodOrder) { + _ = r.Len() // method index not needed in compiler + } npos := r.pos() sym := r.selector() r.typeParamNames() diff --git a/src/cmd/compile/internal/noder/unified.go b/src/cmd/compile/internal/noder/unified.go index 9ec2e40e4883f5..b8b0a0a144a6ec 100644 --- a/src/cmd/compile/internal/noder/unified.go +++ b/src/cmd/compile/internal/noder/unified.go @@ -25,8 +25,8 @@ import ( ) // uirVersion is the unified IR version to use for encoding/decoding. -// Use V4 for generic methods. -const uirVersion = pkgbits.V4 +// Use V5 for generic method index ordering. +const uirVersion = pkgbits.V5 // localPkgReader holds the package reader used for reading the local // package. It exists so the unified IR linker can refer back to it diff --git a/src/cmd/compile/internal/noder/writer.go b/src/cmd/compile/internal/noder/writer.go index bcc68bbc81b910..8ca91973e3d632 100644 --- a/src/cmd/compile/internal/noder/writer.go +++ b/src/cmd/compile/internal/noder/writer.go @@ -11,6 +11,7 @@ import ( "go/version" "internal/buildcfg" "internal/pkgbits" + "log" "os" "slices" "strings" @@ -82,8 +83,9 @@ type pkgWriter struct { // Maps from types2.Objects back to their syntax.Decl. - funDecls map[*types2.Func]*syntax.FuncDecl - typDecls map[*types2.TypeName]typeDeclGen + funDecls map[*types2.Func]*syntax.FuncDecl + typDecls map[*types2.TypeName]typeDeclGen + methodIdx map[*types2.Func]int // method declaration order, for x/tools decoder (#81188) // linknames maps package-scope objects to their linker symbol name, // if specified by a //go:linkname or //go:linknamestd directive. @@ -114,8 +116,9 @@ func newPkgWriter(m posMap, pkg *types2.Package, info *types2.Info, otherInfo ma posBasesIdx: make(map[*syntax.PosBase]index), - funDecls: make(map[*types2.Func]*syntax.FuncDecl), - typDecls: make(map[*types2.TypeName]typeDeclGen), + funDecls: make(map[*types2.Func]*syntax.FuncDecl), + typDecls: make(map[*types2.TypeName]typeDeclGen), + methodIdx: make(map[*types2.Func]int), linknames: make(map[types2.Object]struct { remote string @@ -863,7 +866,7 @@ func (w *writer) doObj(wext *writer, obj types2.Object) pkgbits.CodeObj { // Unified IR panics are the worst; this is a huge help in debugging them. defer func() { if p := recover(); p != nil { - fmt.Printf("Intercepted unified IR writer panic for function %s, repanicking", obj.FullName()) + log.Printf("Intercepted unified IR writer panic for function %s, repanicking", obj.FullName()) panic(p) } }() @@ -879,6 +882,9 @@ func (w *writer) doObj(wext *writer, obj types2.Object) pkgbits.CodeObj { w.selector(obj) w.typeParamNames(sig.RecvTypeParams()) w.param(sig.Recv()) + if w.Version().Has(pkgbits.PreserveMethodOrder) { + w.Len(w.p.methodIdx[obj.Origin()]) + } } else { if w.Version().Has(pkgbits.GenericMethods) { w.Bool(false) // function @@ -920,6 +926,7 @@ func (w *writer) doObj(wext *writer, obj types2.Object) pkgbits.CodeObj { var methods, gmethods []*types2.Func for i := range named.NumMethods() { m := named.Method(i) + w.p.methodIdx[m] = i if isGenericMethod(m.Type()) { gmethods = append(gmethods, m) } else { @@ -1054,6 +1061,9 @@ func (w *writer) method(wext *writer, meth *types2.Func) { sig := meth.Type().(*types2.Signature) w.Sync(pkgbits.SyncMethod) + if w.Version().Has(pkgbits.PreserveMethodOrder) { + w.Len(w.p.methodIdx[meth.Origin()]) + } w.pos(meth) w.selector(meth) w.typeParamNames(sig.RecvTypeParams()) diff --git a/src/cmd/compile/internal/riscv64/ssa.go b/src/cmd/compile/internal/riscv64/ssa.go index fff3c13e2a4044..f01f1bbbcf3cbf 100644 --- a/src/cmd/compile/internal/riscv64/ssa.go +++ b/src/cmd/compile/internal/riscv64/ssa.go @@ -772,10 +772,8 @@ func ssaGenValue(s *ssagen.State, v *ssa.Value) { case ssa.OpRISCV64LoweredZero: ptr := v.Args[0].Reg() - sc := v.AuxValAndOff() - n := sc.Val64() - - mov, sz := largestMove(sc.Off64()) + n, align := v.AuxSizeAndAlign() + mov, sz := largestMove(align) // mov ZERO, (offset)(Rarg0) var off int64 @@ -797,9 +795,8 @@ func ssaGenValue(s *ssagen.State, v *ssa.Value) { case ssa.OpRISCV64LoweredZeroLoop: ptr := v.Args[0].Reg() - sc := v.AuxValAndOff() - n := sc.Val64() - mov, sz := largestMove(sc.Off64()) + n, align := v.AuxSizeAndAlign() + mov, sz := largestMove(align) chunk := 8 * sz if n <= 3*chunk { @@ -808,9 +805,21 @@ func ssaGenValue(s *ssagen.State, v *ssa.Value) { tmp := v.RegTmp() + if n >= 1<<31 { + p := s.Prog(riscv.AMOV) + p.From.Type = obj.TYPE_CONST + p.From.Offset = n - n%chunk + p.To.Type = obj.TYPE_REG + p.To.Reg = tmp + } p := s.Prog(riscv.AADD) - p.From.Type = obj.TYPE_CONST - p.From.Offset = n - n%chunk + if n >= 1<<31 { + p.From.Type = obj.TYPE_REG + p.From.Reg = tmp + } else { + p.From.Type = obj.TYPE_CONST + p.From.Offset = n - n%chunk + } p.Reg = ptr p.To.Type = obj.TYPE_REG p.To.Reg = tmp @@ -859,9 +868,8 @@ func ssaGenValue(s *ssagen.State, v *ssa.Value) { break } - sa := v.AuxValAndOff() - n := sa.Val64() - mov, sz := largestMove(sa.Off64()) + n, align := v.AuxSizeAndAlign() + mov, sz := largestMove(align) var off int64 tmp := int16(riscv.REG_X5) @@ -888,9 +896,8 @@ func ssaGenValue(s *ssagen.State, v *ssa.Value) { break } - sc := v.AuxValAndOff() - n := sc.Val64() - mov, sz := largestMove(sc.Off64()) + n, align := v.AuxSizeAndAlign() + mov, sz := largestMove(align) chunk := 8 * sz if n <= 3*chunk { @@ -898,9 +905,21 @@ func ssaGenValue(s *ssagen.State, v *ssa.Value) { } tmp := int16(riscv.REG_X5) + if n >= 1<<31 { + p := s.Prog(riscv.AMOV) + p.From.Type = obj.TYPE_CONST + p.From.Offset = n - n%chunk + p.To.Type = obj.TYPE_REG + p.To.Reg = riscv.REG_X6 + } p := s.Prog(riscv.AADD) - p.From.Type = obj.TYPE_CONST - p.From.Offset = n - n%chunk + if n >= 1<<31 { + p.From.Type = obj.TYPE_REG + p.From.Reg = riscv.REG_X6 + } else { + p.From.Type = obj.TYPE_CONST + p.From.Offset = n - n%chunk + } p.Reg = src p.To.Type = obj.TYPE_REG p.To.Reg = riscv.REG_X6 diff --git a/src/cmd/compile/internal/slice/slice.go b/src/cmd/compile/internal/slice/slice.go index 11f71e2ebca57c..8669b0ed0cc7a4 100644 --- a/src/cmd/compile/internal/slice/slice.go +++ b/src/cmd/compile/internal/slice/slice.go @@ -286,6 +286,18 @@ func analyze(fn *ir.Func) { } } + // do walks n and everything below it, recording what happens to the + // slice variables we are tracking. It counts every mention of such a + // variable in allUses, and the subset of those mentions that this pass + // understands to preserve exclusivity in okUses. A variable whose two + // counts end up equal is only ever used in ways we understand; any + // other variable is dropped at the end of the analysis. Uses that + // definitely destroy exclusivity, like &s[i], stop the tracking right + // away by clearing s.Opt, which makes tracking report the variable as + // no longer being considered. + // + // It is always used as an ir.DoChildren visitor and always returns + // false, so that the whole function body is walked. var do func(ir.Node) bool do = func(n ir.Node) bool { if n == nil { @@ -324,13 +336,34 @@ func analyze(fn *ir.Func) { } case ir.OADDR: n := n.(*ir.AddrExpr) - if n.X.Op() == ir.OINDEX { - n := n.X.(*ir.IndexExpr) - if i := tracking(n.X); i != nil { - // &s[i] is definitely a nonexclusive transition. - // (We need this case because s[i] is ok, but &s[i] is not.) - i.s.Opt = nil + // Walk down to the object whose interior we're taking the + // address of. Field selectors and array indexes don't leave + // that object, so &s[i], &s[i].f, and &s[i].f[j] all end up + // pointing into s's backing store. + // (We need this because s[i] is ok, but &s[i] is not.) + x := n.X + for x != nil { + switch x.Op() { + case ir.ODOT: + // &x.f points into x. + x = x.(*ir.SelectorExpr).X + continue + case ir.OINDEX: + idx := x.(*ir.IndexExpr) + if idx.X.Type().IsArray() { + // &a[i] points into a. + // Note: for a pointer to an array, or for a + // slice, the address points into a different + // object instead, so we stop here. + x = idx.X + continue + } + if i := tracking(idx.X); i != nil { + // &s[i] is definitely a nonexclusive transition. + i.s.Opt = nil + } } + break } case ir.ORETURN: n := n.(*ir.ReturnStmt) diff --git a/src/cmd/compile/internal/ssa/_gen/RISCV64.rules b/src/cmd/compile/internal/ssa/_gen/RISCV64.rules index 853cd2fbeca39a..5270b26862b32f 100644 --- a/src/cmd/compile/internal/ssa/_gen/RISCV64.rules +++ b/src/cmd/compile/internal/ssa/_gen/RISCV64.rules @@ -376,11 +376,11 @@ // Unroll zeroing in medium size (at most 192 bytes i.e. 3 cachelines) (Zero [s] {t} ptr mem) && s <= 24*moveSize(t.Alignment(), config) => - (LoweredZero [makeValAndOff(int32(s),int32(t.Alignment()))] ptr mem) + (LoweredZero [s] {t.Alignment()} ptr mem) // Generic zeroing uses a loop (Zero [s] {t} ptr mem) && s > 24*moveSize(t.Alignment(), config) => - (LoweredZeroLoop [makeValAndOff(int32(s),int32(t.Alignment()))] ptr mem) + (LoweredZeroLoop [s] {t.Alignment()} ptr mem) // Checks (IsNonNil ...) => (SNEZ ...) @@ -446,12 +446,12 @@ // Generic move (Move [s] {t} dst src mem) && s > 0 && s <= 3*8*moveSize(t.Alignment(), config) && logLargeCopy(v, s) => - (LoweredMove [makeValAndOff(int32(s),int32(t.Alignment()))] dst src mem) + (LoweredMove [s] {t.Alignment()} dst src mem) // Generic move uses a loop (Move [s] {t} dst src mem) && s > 3*8*moveSize(t.Alignment(), config) && logLargeCopy(v, s) => - (LoweredMoveLoop [makeValAndOff(int32(s),int32(t.Alignment()))] dst src mem) + (LoweredMoveLoop [s] {t.Alignment()} dst src mem) // Boolean ops; 0=false, 1=true (AndB ...) => (AND ...) diff --git a/src/cmd/compile/internal/ssa/_gen/RISCV64Ops.go b/src/cmd/compile/internal/ssa/_gen/RISCV64Ops.go index 56979bd968cde3..9aaa83533f2962 100644 --- a/src/cmd/compile/internal/ssa/_gen/RISCV64Ops.go +++ b/src/cmd/compile/internal/ssa/_gen/RISCV64Ops.go @@ -293,15 +293,15 @@ func init() { // general unrolled zeroing // arg0 = address of memory to zero // arg1 = mem - // auxint = element size and type alignment + // auxint = size + // aux = alignment (as an int64) // returns mem // mov ZERO, (OFFSET)(Rarg0) { name: "LoweredZero", - aux: "SymValAndOff", + aux: "SizeAndAlign", typ: "Mem", argLength: 2, - symEffect: "Write", faultOnNilArg0: true, addrSinkArg0: true, reg: regInfo{ @@ -310,15 +310,15 @@ func init() { }, // general unaligned zeroing // arg0 = address of memory to zero (clobber) - // arg2 = mem - // auxint = element size and type alignment + // arg1 = mem + // auxint = size + // aux = alignment (as an int64) // returns mem { name: "LoweredZeroLoop", - aux: "SymValAndOff", + aux: "SizeAndAlign", typ: "Mem", argLength: 2, - symEffect: "Write", needIntTemp: true, faultOnNilArg0: true, addrSinkArg0: true, @@ -332,14 +332,14 @@ func init() { // arg0 = address of dst memory (clobber) // arg1 = address of src memory (clobber) // arg2 = mem - // auxint = size and type alignment + // auxint = size + // aux = alignment (as an int64) // returns mem // mov (offset)(Rarg1), TMP // mov TMP, (offset)(Rarg0) { name: "LoweredMove", - aux: "SymValAndOff", - symEffect: "Write", + aux: "SizeAndAlign", argLength: 3, reg: regInfo{ inputs: []regMask{gpMask.minus(regNamed["X5"]), gpMask.minus(regNamed["X5"])}, @@ -354,8 +354,9 @@ func init() { // general unaligned move // arg0 = address of dst memory (clobber) // arg1 = address of src memory (clobber) - // arg3 = mem - // auxint = alignment + // arg2 = mem + // auxint = size + // aux = alignment (as an int64) // returns mem // ADD $sz, X6 //loop: @@ -367,9 +368,8 @@ func init() { // BNE X6, Rarg1, loop { name: "LoweredMoveLoop", - aux: "SymValAndOff", + aux: "SizeAndAlign", argLength: 3, - symEffect: "Write", reg: regInfo{ inputs: []regMask{gpMask.minus(r5toR6), gpMask.minus(r5toR6)}, clobbers: r5toR6, diff --git a/src/cmd/compile/internal/ssa/_gen/generic.rules b/src/cmd/compile/internal/ssa/_gen/generic.rules index 80e7c8e9b6ead3..6c7f1767448e9e 100644 --- a/src/cmd/compile/internal/ssa/_gen/generic.rules +++ b/src/cmd/compile/internal/ssa/_gen/generic.rules @@ -2237,10 +2237,7 @@ (RotateLeft8 x (Sub(64|32|16|8) (Const(64|32|16|8) [c]) y)) && c&7 == 0 => (RotateLeft8 x (Neg(64|32|16|8) y)) // Ensure we don't do Const64 rotates in a 32-bit system. -(RotateLeft64 x (Const64 [c])) && config.PtrSize == 4 => (RotateLeft64 x (Const32 [int32(c)])) -(RotateLeft32 x (Const64 [c])) && config.PtrSize == 4 => (RotateLeft32 x (Const32 [int32(c)])) -(RotateLeft16 x (Const64 [c])) && config.PtrSize == 4 => (RotateLeft16 x (Const32 [int32(c)])) -(RotateLeft8 x (Const64 [c])) && config.PtrSize == 4 => (RotateLeft8 x (Const32 [int32(c)])) +(RotateLeft(64|32|16|8) x (Const64 [c])) && config.PtrSize == 4 => (RotateLeft(64|32|16|8) x (Const32 [int32(c)])) // Rotating by c, then by d, is the same as rotating by c+d. // We're trading a rotate for an add, which seems generally a good choice. It is especially good when c and d are constants. diff --git a/src/cmd/compile/internal/ssa/_gen/main.go b/src/cmd/compile/internal/ssa/_gen/main.go index c2c1a04ad00883..6ca71b7e00c45d 100644 --- a/src/cmd/compile/internal/ssa/_gen/main.go +++ b/src/cmd/compile/internal/ssa/_gen/main.go @@ -377,13 +377,13 @@ func genOp() { } if v.faultOnNilArg0 { fmt.Fprintln(w, "faultOnNilArg0: true,") - if v.aux != "Sym" && v.aux != "SymOff" && v.aux != "SymValAndOff" && v.aux != "Int64" && v.aux != "Int32" && v.aux != "" { + if v.aux != "Sym" && v.aux != "SymOff" && v.aux != "SymValAndOff" && v.aux != "Int64" && v.aux != "Int32" && v.aux != "SizeAndAlign" && v.aux != "" { log.Fatalf("faultOnNilArg0 with aux %s not allowed", v.aux) } } if v.faultOnNilArg1 { fmt.Fprintln(w, "faultOnNilArg1: true,") - if v.aux != "Sym" && v.aux != "SymOff" && v.aux != "SymValAndOff" && v.aux != "Int64" && v.aux != "Int32" && v.aux != "" { + if v.aux != "Sym" && v.aux != "SymOff" && v.aux != "SymValAndOff" && v.aux != "Int64" && v.aux != "Int32" && v.aux != "SizeAndAlign" && v.aux != "" { log.Fatalf("faultOnNilArg1 with aux %s not allowed", v.aux) } } diff --git a/src/cmd/compile/internal/ssa/_gen/rulegen.go b/src/cmd/compile/internal/ssa/_gen/rulegen.go index 6cabb0006c980e..d9e96b61f58537 100644 --- a/src/cmd/compile/internal/ssa/_gen/rulegen.go +++ b/src/cmd/compile/internal/ssa/_gen/rulegen.go @@ -1460,7 +1460,7 @@ func opHasAuxInt(op opData) bool { switch op.aux { case "Bool", "Int8", "Int16", "Int32", "Int64", "Int128", "UInt8", "Float32", "Float64", "SymOff", "CallOff", "SymValAndOff", "TypSize", "ARM64BitField", "FlagConstant", "CCop", - "PanicBoundsC", "PanicBoundsCC", "ARM64ConditionalParams": + "PanicBoundsC", "PanicBoundsCC", "ARM64ConditionalParams", "SizeAndAlign": return true } return false @@ -1469,7 +1469,7 @@ func opHasAuxInt(op opData) bool { func opHasAux(op opData) bool { switch op.aux { case "String", "Sym", "SymOff", "Call", "CallOff", "SymValAndOff", "Typ", "TypSize", - "S390XCCMask", "S390XRotateParams", "PanicBoundsC", "PanicBoundsCC": + "S390XCCMask", "S390XRotateParams", "PanicBoundsC", "PanicBoundsCC", "SizeAndAlign": return true } return false @@ -1828,6 +1828,8 @@ func (op opData) auxType() string { return "PanicBoundsC" case "PanicBoundsCC": return "PanicBoundsCC" + case "SizeAndAlign": + return "int64" default: return "invalid" } @@ -1872,6 +1874,8 @@ func (op opData) auxIntType() string { return "arm64ConditionalParams" case "PanicBoundsC", "PanicBoundsCC": return "int64" + case "SizeAndAlign": + return "int64" default: return "invalid" } diff --git a/src/cmd/compile/internal/ssa/check.go b/src/cmd/compile/internal/ssa/check.go index 9396f8dccbe087..9f08a88586142f 100644 --- a/src/cmd/compile/internal/ssa/check.go +++ b/src/cmd/compile/internal/ssa/check.go @@ -219,8 +219,14 @@ func checkFunc(f *Func) { case auxPanicBoundsC, auxPanicBoundsCC: canHaveAux = true canHaveAuxInt = true + case auxSizeAndAlign: + if _, ok := v.Aux.(int64Aux); !ok { + f.Fatalf("value %v has Aux type %T, want int64Aux", v, v.Aux) + } + canHaveAux = true + canHaveAuxInt = true default: - f.Fatalf("unknown aux type for %s", v.Op) + f.Fatalf("unknown aux type %T for %s", opcodeTable[v.Op].auxType, v.Op) } if !canHaveAux && v.Aux != nil { f.Fatalf("value %s has an Aux value %v but shouldn't", v.LongString(), v.Aux) diff --git a/src/cmd/compile/internal/ssa/nilcheck.go b/src/cmd/compile/internal/ssa/nilcheck.go index 7f88b14086d539..4db37450af69b2 100644 --- a/src/cmd/compile/internal/ssa/nilcheck.go +++ b/src/cmd/compile/internal/ssa/nilcheck.go @@ -302,7 +302,7 @@ func nilcheckelim2(f *Func) { case auxInt64: // ARM uses this auxType for duffcopy/duffzero/alignment info. // It does not affect the effective address. - case auxNone: + case auxNone, auxSizeAndAlign: // offset is zero. default: v.Fatalf("can't handle aux %s (type %d) yet\n", v.auxString(), int(opcodeTable[v.Op].auxType)) diff --git a/src/cmd/compile/internal/ssa/op.go b/src/cmd/compile/internal/ssa/op.go index 85fdb08075c1f2..0c4e902aece0e6 100644 --- a/src/cmd/compile/internal/ssa/op.go +++ b/src/cmd/compile/internal/ssa/op.go @@ -384,6 +384,7 @@ const ( auxS390XCCMask // aux is a s390x 4-bit condition code mask auxS390XCCMaskInt8 // aux is a s390x 4-bit condition code mask, auxInt is an int8 immediate auxS390XCCMaskUint8 // aux is a s390x 4-bit condition code mask, auxInt is a uint8 immediate + auxSizeAndAlign // auxInt is an int64 size, aux is an int64 alignment ) // A SymEffect describes the effect that an SSA Value has on the variable diff --git a/src/cmd/compile/internal/ssa/opGen.go b/src/cmd/compile/internal/ssa/opGen.go index 86ede5910c5bdd..7aa8af2e78e5b9 100644 --- a/src/cmd/compile/internal/ssa/opGen.go +++ b/src/cmd/compile/internal/ssa/opGen.go @@ -96044,11 +96044,10 @@ var opcodeTable = [...]opInfo{ }, { name: "LoweredZero", - auxType: auxSymValAndOff, + auxType: auxSizeAndAlign, argLen: 2, faultOnNilArg0: true, addrSinkArg0: true, - symEffect: SymWrite, reg: regInfo{ inputs: []inputInfo{ {0, regMask{v1: 1006632944, v2: 0}}, // X5 X6 X7 X8 X9 X10 X11 X12 X13 X14 X15 X16 X17 X18 X19 X20 X21 X22 X23 X24 X25 X26 X28 X29 X30 @@ -96057,12 +96056,11 @@ var opcodeTable = [...]opInfo{ }, { name: "LoweredZeroLoop", - auxType: auxSymValAndOff, + auxType: auxSizeAndAlign, argLen: 2, needIntTemp: true, faultOnNilArg0: true, addrSinkArg0: true, - symEffect: SymWrite, reg: regInfo{ inputs: []inputInfo{ {0, regMask{v1: 1006632944, v2: 0}}, // X5 X6 X7 X8 X9 X10 X11 X12 X13 X14 X15 X16 X17 X18 X19 X20 X21 X22 X23 X24 X25 X26 X28 X29 X30 @@ -96072,13 +96070,12 @@ var opcodeTable = [...]opInfo{ }, { name: "LoweredMove", - auxType: auxSymValAndOff, + auxType: auxSizeAndAlign, argLen: 3, faultOnNilArg0: true, faultOnNilArg1: true, addrSinkArg0: true, addrSinkArg1: true, - symEffect: SymWrite, reg: regInfo{ inputs: []inputInfo{ {0, regMask{v1: 1006632928, v2: 0}}, // X6 X7 X8 X9 X10 X11 X12 X13 X14 X15 X16 X17 X18 X19 X20 X21 X22 X23 X24 X25 X26 X28 X29 X30 @@ -96089,13 +96086,12 @@ var opcodeTable = [...]opInfo{ }, { name: "LoweredMoveLoop", - auxType: auxSymValAndOff, + auxType: auxSizeAndAlign, argLen: 3, faultOnNilArg0: true, faultOnNilArg1: true, addrSinkArg0: true, addrSinkArg1: true, - symEffect: SymWrite, reg: regInfo{ inputs: []inputInfo{ {0, regMask{v1: 1006632896, v2: 0}}, // X7 X8 X9 X10 X11 X12 X13 X14 X15 X16 X17 X18 X19 X20 X21 X22 X23 X24 X25 X26 X28 X29 X30 diff --git a/src/cmd/compile/internal/ssa/rewrite.go b/src/cmd/compile/internal/ssa/rewrite.go index c1fbc51dc3e98d..78a576df514a18 100644 --- a/src/cmd/compile/internal/ssa/rewrite.go +++ b/src/cmd/compile/internal/ssa/rewrite.go @@ -794,6 +794,14 @@ func (stringAux) CanBeAnSSAAux() {} func auxToString(i Aux) string { return string(i.(stringAux)) } + +type int64Aux int64 + +func (int64Aux) CanBeAnSSAAux() {} + +func int64ToAux(v int64) Aux { + return int64Aux(v) +} func auxToSym(i Aux) Sym { // TODO: kind of a hack - allows nil interface through s, _ := i.(Sym) diff --git a/src/cmd/compile/internal/ssa/rewriteRISCV64.go b/src/cmd/compile/internal/ssa/rewriteRISCV64.go index e516259cc09d0c..de8a2dcd41ab4e 100644 --- a/src/cmd/compile/internal/ssa/rewriteRISCV64.go +++ b/src/cmd/compile/internal/ssa/rewriteRISCV64.go @@ -3102,7 +3102,7 @@ func rewriteValueRISCV64_OpMove(v *Value) bool { } // match: (Move [s] {t} dst src mem) // cond: s > 0 && s <= 3*8*moveSize(t.Alignment(), config) && logLargeCopy(v, s) - // result: (LoweredMove [makeValAndOff(int32(s),int32(t.Alignment()))] dst src mem) + // result: (LoweredMove [s] {t.Alignment()} dst src mem) for { s := auxIntToInt64(v.AuxInt) t := auxToType(v.Aux) @@ -3113,13 +3113,14 @@ func rewriteValueRISCV64_OpMove(v *Value) bool { break } v.reset(OpRISCV64LoweredMove) - v.AuxInt = valAndOffToAuxInt(makeValAndOff(int32(s), int32(t.Alignment()))) + v.AuxInt = int64ToAuxInt(s) + v.Aux = int64ToAux(t.Alignment()) v.AddArg3(dst, src, mem) return true } // match: (Move [s] {t} dst src mem) // cond: s > 3*8*moveSize(t.Alignment(), config) && logLargeCopy(v, s) - // result: (LoweredMoveLoop [makeValAndOff(int32(s),int32(t.Alignment()))] dst src mem) + // result: (LoweredMoveLoop [s] {t.Alignment()} dst src mem) for { s := auxIntToInt64(v.AuxInt) t := auxToType(v.Aux) @@ -3130,7 +3131,8 @@ func rewriteValueRISCV64_OpMove(v *Value) bool { break } v.reset(OpRISCV64LoweredMoveLoop) - v.AuxInt = valAndOffToAuxInt(makeValAndOff(int32(s), int32(t.Alignment()))) + v.AuxInt = int64ToAuxInt(s) + v.Aux = int64ToAux(t.Alignment()) v.AddArg3(dst, src, mem) return true } @@ -10768,7 +10770,7 @@ func rewriteValueRISCV64_OpZero(v *Value) bool { } // match: (Zero [s] {t} ptr mem) // cond: s <= 24*moveSize(t.Alignment(), config) - // result: (LoweredZero [makeValAndOff(int32(s),int32(t.Alignment()))] ptr mem) + // result: (LoweredZero [s] {t.Alignment()} ptr mem) for { s := auxIntToInt64(v.AuxInt) t := auxToType(v.Aux) @@ -10778,13 +10780,14 @@ func rewriteValueRISCV64_OpZero(v *Value) bool { break } v.reset(OpRISCV64LoweredZero) - v.AuxInt = valAndOffToAuxInt(makeValAndOff(int32(s), int32(t.Alignment()))) + v.AuxInt = int64ToAuxInt(s) + v.Aux = int64ToAux(t.Alignment()) v.AddArg2(ptr, mem) return true } // match: (Zero [s] {t} ptr mem) // cond: s > 24*moveSize(t.Alignment(), config) - // result: (LoweredZeroLoop [makeValAndOff(int32(s),int32(t.Alignment()))] ptr mem) + // result: (LoweredZeroLoop [s] {t.Alignment()} ptr mem) for { s := auxIntToInt64(v.AuxInt) t := auxToType(v.Aux) @@ -10794,7 +10797,8 @@ func rewriteValueRISCV64_OpZero(v *Value) bool { break } v.reset(OpRISCV64LoweredZeroLoop) - v.AuxInt = valAndOffToAuxInt(makeValAndOff(int32(s), int32(t.Alignment()))) + v.AuxInt = int64ToAuxInt(s) + v.Aux = int64ToAux(t.Alignment()) v.AddArg2(ptr, mem) return true } diff --git a/src/cmd/compile/internal/ssa/rewritegeneric.go b/src/cmd/compile/internal/ssa/rewritegeneric.go index a5e0478a40acb7..83adb5397f8437 100644 --- a/src/cmd/compile/internal/ssa/rewritegeneric.go +++ b/src/cmd/compile/internal/ssa/rewritegeneric.go @@ -27987,6 +27987,7 @@ func rewriteValuegeneric_OpRotateLeft16(v *Value) bool { v_0 := v.Args[0] b := v.Block config := b.Func.Config + typ := &b.Func.Config.Types // match: (RotateLeft16 x (Const16 [c])) // cond: c%16 == 0 // result: x @@ -28430,21 +28431,20 @@ func rewriteValuegeneric_OpRotateLeft16(v *Value) bool { v.AddArg2(x, v0) return true } - // match: (RotateLeft16 x (Const64 [c])) + // match: (RotateLeft16 x (Const64 [c])) // cond: config.PtrSize == 4 - // result: (RotateLeft16 x (Const32 [int32(c)])) + // result: (RotateLeft16 x (Const32 [int32(c)])) for { x := v_0 if v_1.Op != OpConst64 { break } - t := v_1.Type c := auxIntToInt64(v_1.AuxInt) if !(config.PtrSize == 4) { break } v.reset(OpRotateLeft16) - v0 := b.NewValue0(v.Pos, OpConst32, t) + v0 := b.NewValue0(v.Pos, OpConst32, typ.UInt32) v0.AuxInt = int32ToAuxInt(int32(c)) v.AddArg2(x, v0) return true @@ -28532,6 +28532,7 @@ func rewriteValuegeneric_OpRotateLeft32(v *Value) bool { v_0 := v.Args[0] b := v.Block config := b.Func.Config + typ := &b.Func.Config.Types // match: (RotateLeft32 x (Const32 [c])) // cond: c%32 == 0 // result: x @@ -28975,21 +28976,20 @@ func rewriteValuegeneric_OpRotateLeft32(v *Value) bool { v.AddArg2(x, v0) return true } - // match: (RotateLeft32 x (Const64 [c])) + // match: (RotateLeft32 x (Const64 [c])) // cond: config.PtrSize == 4 - // result: (RotateLeft32 x (Const32 [int32(c)])) + // result: (RotateLeft32 x (Const32 [int32(c)])) for { x := v_0 if v_1.Op != OpConst64 { break } - t := v_1.Type c := auxIntToInt64(v_1.AuxInt) if !(config.PtrSize == 4) { break } v.reset(OpRotateLeft32) - v0 := b.NewValue0(v.Pos, OpConst32, t) + v0 := b.NewValue0(v.Pos, OpConst32, typ.UInt32) v0.AuxInt = int32ToAuxInt(int32(c)) v.AddArg2(x, v0) return true @@ -29077,6 +29077,7 @@ func rewriteValuegeneric_OpRotateLeft64(v *Value) bool { v_0 := v.Args[0] b := v.Block config := b.Func.Config + typ := &b.Func.Config.Types // match: (RotateLeft64 x (Const64 [c])) // cond: c%64 == 0 // result: x @@ -29520,21 +29521,20 @@ func rewriteValuegeneric_OpRotateLeft64(v *Value) bool { v.AddArg2(x, v0) return true } - // match: (RotateLeft64 x (Const64 [c])) + // match: (RotateLeft64 x (Const64 [c])) // cond: config.PtrSize == 4 - // result: (RotateLeft64 x (Const32 [int32(c)])) + // result: (RotateLeft64 x (Const32 [int32(c)])) for { x := v_0 if v_1.Op != OpConst64 { break } - t := v_1.Type c := auxIntToInt64(v_1.AuxInt) if !(config.PtrSize == 4) { break } v.reset(OpRotateLeft64) - v0 := b.NewValue0(v.Pos, OpConst32, t) + v0 := b.NewValue0(v.Pos, OpConst32, typ.UInt32) v0.AuxInt = int32ToAuxInt(int32(c)) v.AddArg2(x, v0) return true @@ -29622,6 +29622,7 @@ func rewriteValuegeneric_OpRotateLeft8(v *Value) bool { v_0 := v.Args[0] b := v.Block config := b.Func.Config + typ := &b.Func.Config.Types // match: (RotateLeft8 x (Const8 [c])) // cond: c%8 == 0 // result: x @@ -30065,21 +30066,20 @@ func rewriteValuegeneric_OpRotateLeft8(v *Value) bool { v.AddArg2(x, v0) return true } - // match: (RotateLeft8 x (Const64 [c])) + // match: (RotateLeft8 x (Const64 [c])) // cond: config.PtrSize == 4 - // result: (RotateLeft8 x (Const32 [int32(c)])) + // result: (RotateLeft8 x (Const32 [int32(c)])) for { x := v_0 if v_1.Op != OpConst64 { break } - t := v_1.Type c := auxIntToInt64(v_1.AuxInt) if !(config.PtrSize == 4) { break } v.reset(OpRotateLeft8) - v0 := b.NewValue0(v.Pos, OpConst32, t) + v0 := b.NewValue0(v.Pos, OpConst32, typ.UInt32) v0.AuxInt = int32ToAuxInt(int32(c)) v.AddArg2(x, v0) return true diff --git a/src/cmd/compile/internal/ssa/value.go b/src/cmd/compile/internal/ssa/value.go index 279099d4785d88..bd11e526fea6df 100644 --- a/src/cmd/compile/internal/ssa/value.go +++ b/src/cmd/compile/internal/ssa/value.go @@ -152,6 +152,10 @@ func (v *Value) AuxArm64ConditionalParams() arm64ConditionalParams { return auxIntToArm64ConditionalParams(v.AuxInt) } +func (v *Value) AuxSizeAndAlign() (int64, int64) { + return v.AuxInt, int64(v.Aux.(int64Aux)) +} + // long form print. v# = opcode [aux] args [: reg] (names) func (v *Value) LongString() string { if v == nil { @@ -250,6 +254,8 @@ func (v *Value) auxString() string { return fmt.Sprintf(" {%v}", v.Aux) case auxFlagConstant: return fmt.Sprintf("[%s]", flagConstant(v.AuxInt)) + case auxSizeAndAlign: + return fmt.Sprintf(" [size=%d] {align=%d}", v.AuxInt, v.Aux) case auxNone: return "" default: diff --git a/src/cmd/compile/internal/ssagen/ssa.go b/src/cmd/compile/internal/ssagen/ssa.go index 6b215a833df011..7052dc2dc8b5ce 100644 --- a/src/cmd/compile/internal/ssagen/ssa.go +++ b/src/cmd/compile/internal/ssagen/ssa.go @@ -5114,6 +5114,16 @@ func (s *state) call(n *ir.CallExpr, k callKind, returnResultAddr bool, deferExt callArgs = append(callArgs, s.putArg(n, t.Param(i).Type)) } + // In -race mode, we need to call racefuncexit before a tail call. + // A tail call reuses our frame and returns directly to our caller, + // so this is the last chance we get to tell the race detector that + // this function is done. Note: this has to happen after the + // arguments are evaluated, otherwise races in the argument + // expressions would be attributed to the caller instead. + if k == callTail && s.instrumentEnterExit { + s.rtcall(ir.Syms.Racefuncexit, true, nil) + } + callArgs = append(callArgs, s.mem()) // call target diff --git a/src/cmd/compile/internal/test/issue_81478_test.go b/src/cmd/compile/internal/test/issue_81478_test.go new file mode 100644 index 00000000000000..e3af6eb40bd5e0 --- /dev/null +++ b/src/cmd/compile/internal/test/issue_81478_test.go @@ -0,0 +1,98 @@ +// Copyright 2026 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +package test + +import ( + "internal/platform" + "internal/testenv" + "os" + "path/filepath" + "runtime" + "strings" + "testing" +) + +// Test that compiler-generated wrappers that end in a tail call still tell the +// race detector that the wrapper has returned. This is a regression test for +// #81478, where the missing racefuncexit call made every wrapper call leak an +// entry in the race detector's shadow stack, which in turn made later +// synchronization operations dramatically slower. +func TestIssue81478(t *testing.T) { + if !platform.RaceDetectorSupported(runtime.GOOS, runtime.GOARCH) { + t.Skipf("race detector not supported on %s/%s", runtime.GOOS, runtime.GOARCH) + } + testenv.MustHaveGoBuild(t) + + dir := t.TempDir() + src := filepath.Join(dir, "x.go") + if err := os.WriteFile(src, []byte(issue81478src), 0644); err != nil { + t.Fatalf("could not write file: %v", err) + } + + cmd := testenv.Command(t, testenv.GoToolPath(t), "tool", "compile", "-race", "-p=main", "-S", "-o", filepath.Join(dir, "x.o"), src) + out, err := cmd.CombinedOutput() + if err != nil { + t.Fatalf("compile failed: %v\n%s", err, out) + } + + // W.F tail calls through an interface, S.G tail calls a static function. + // Whether the tail call is actually emitted is architecture-dependent, but + // either way the enter/exit instrumentation must be balanced. + for _, wrapper := range []string{"main.(*W).F", "main.(*S).G"} { + body, ok := funcBody(string(out), wrapper) + if !ok { + t.Errorf("no assembly found for %s", wrapper) + continue + } + enter := strings.Count(body, "runtime.racefuncenter") + exit := strings.Count(body, "runtime.racefuncexit") + if enter == 0 { + t.Errorf("%s: not instrumented for the race detector\n%s", wrapper, body) + } else if enter != exit { + t.Errorf("%s: unbalanced race instrumentation: %d racefuncenter, %d racefuncexit\n%s", wrapper, enter, exit, body) + } + } +} + +// funcBody returns the part of the -S output that belongs to the symbol fn: +// its header line plus the indented instruction and relocation lines below it. +func funcBody(out, fn string) (string, bool) { + lines := strings.Split(out, "\n") + for i, line := range lines { + if !strings.HasPrefix(line, fn+" STEXT") { + continue + } + end := i + 1 + for end < len(lines) && strings.HasPrefix(lines[end], "\t") { + end++ + } + return strings.Join(lines[i:end], "\n"), true + } + return "", false +} + +var issue81478src = ` +package main + +type I interface{ F() } + +type T struct{ n int } + +func (t *T) F() { t.n++ } + +//go:noinline +func (t *T) G() { t.n++ } + +// W.F is a wrapper around an embedded interface method. +type W struct{ I } + +// S.G is a wrapper around an embedded pointer's method. +type S struct{ *T } + +var _ I = &W{} +var _ = (&S{}).G + +func main() {} +` diff --git a/src/cmd/compile/internal/types2/named_test.go b/src/cmd/compile/internal/types2/named_test.go index 76093060994807..e241c83b9de2bf 100644 --- a/src/cmd/compile/internal/types2/named_test.go +++ b/src/cmd/compile/internal/types2/named_test.go @@ -123,7 +123,7 @@ package p type T struct{} func (T) a() {} -func (T) c() {} +func (T) c[X any](x X) {} func (T) b() {} ` // should get the same method order each time diff --git a/src/cmd/compile/internal/walk/stmt.go b/src/cmd/compile/internal/walk/stmt.go index 2c01fd10f124e5..b9999cd9ea4732 100644 --- a/src/cmd/compile/internal/walk/stmt.go +++ b/src/cmd/compile/internal/walk/stmt.go @@ -7,6 +7,7 @@ package walk import ( "cmd/compile/internal/base" "cmd/compile/internal/ir" + "cmd/compile/internal/reflectdata" ) // The result of walkStmt MUST be assigned back to n, e.g. @@ -138,6 +139,16 @@ func walkStmt(n ir.Node) ir.Node { case ir.OTAILCALL: n := n.(*ir.TailCallStmt) + // Since go.dev/cl/751465, the compiler emits tail calls for wrappers + // for embedded interfaces. But a tail call never reaches walkCall, so + // the interface calls are not marked as used, causing the linker to + // drop the callee. See issues #81089 and #81340. + // TODO: Should we just call walkCall here? + if n.Call.Op() == ir.OCALLINTER { + usemethod(n.Call) + reflectdata.MarkUsedIfaceMethod(n.Call) + } + var init ir.Nodes n.Call.Fun = walkExpr(n.Call.Fun, &init) diff --git a/src/cmd/cover/cover.go b/src/cmd/cover/cover.go index bc24c52f138a9b..82c419483d0f69 100644 --- a/src/cmd/cover/cover.go +++ b/src/cmd/cover/cover.go @@ -387,17 +387,31 @@ func insideStatement(pos token.Pos, stmts []ast.Stmt) bool { // mergeRangesWithinStatements merges consecutive ranges when a later range's // start position falls strictly inside a statement. This prevents counter // insertion inside multi-line statements such as const (...) blocks. -func mergeRangesWithinStatements(ranges []Range, stmts []ast.Stmt) []Range { - if len(ranges) <= 1 { - return ranges - } - merged := []Range{ranges[0]} - for _, r := range ranges[1:] { - if insideStatement(r.pos, stmts) { - // Extend previous range to cover this one. - merged[len(merged)-1].end = r.end +type rangeWithStatements struct { + Range + numStmt int +} + +func mergeRangesWithinStatements(ranges []Range, stmts []ast.Stmt) []rangeWithStatements { + merged := make([]rangeWithStatements, 0, len(ranges)) + for _, r := range ranges { + // Statements are sorted by source position, so use binary search to + // find the statements whose positions fall within this range. + first, _ := slices.BinarySearchFunc(stmts, r.pos, func(s ast.Stmt, p token.Pos) int { + return cmp.Compare(s.Pos(), p) + }) + last, _ := slices.BinarySearchFunc(stmts, r.end, func(s ast.Stmt, p token.Pos) int { + return cmp.Compare(s.Pos(), p) + }) + numStmt := last - first + + if len(merged) > 0 && insideStatement(r.pos, stmts) { + // Extend the previous range to cover this one. + last := &merged[len(merged)-1] + last.end = r.end + last.numStmt += numStmt } else { - merged = append(merged, r) + merged = append(merged, rangeWithStatements{Range: r, numStmt: numStmt}) } } return merged @@ -952,7 +966,7 @@ func (f *File) addCounters(pos, insertPos, blockEnd token.Pos, list []ast.Stmt, if i == 0 { insertOffset = f.offset(insertPos) } - f.edit.Insert(insertOffset, f.newCounter(r.pos, r.end, last)+";") + f.edit.Insert(insertOffset, f.newCounter(r.pos, r.end, r.numStmt)+";") } } list = list[last:] diff --git a/src/cmd/cover/cover_test.go b/src/cmd/cover/cover_test.go index b86ebd0d149a0b..1991f140c79957 100644 --- a/src/cmd/cover/cover_test.go +++ b/src/cmd/cover/cover_test.go @@ -864,3 +864,53 @@ func main() { got := coverRanges(t, src) compareRanges(t, src, got, want) } + +// TestStatementCountsAfterCommentSplit verifies that splitting a basic block +// at blank or comment-only lines preserves the statement count of the block. +func TestStatementCountsAfterCommentSplit(t *testing.T) { + testenv.MustHaveGoBuild(t) + + src := []byte(`package main + +func main() { + a := 1 + b := 2 + + c := 3 + d := 4 + + if a == 0 { + return + } + + println(a + b + c + d) +}`) + tmpdir := t.TempDir() + srcPath := filepath.Join(tmpdir, "test.go") + if err := os.WriteFile(srcPath, src, 0666); err != nil { + t.Fatal(err) + } + cmd := testenv.Command(t, testcover(t), "-mode=set", srcPath) + out, err := cmd.Output() + if err != nil { + t.Fatalf("cover failed: %v\nOutput: %s", err, out) + } + + re := regexp.MustCompile(`(?s)NumStmt: \[\d+\]uint16\{(.*?)\n\t\}`) + m := re.FindSubmatch(out) + if m == nil { + t.Fatalf("NumStmt array not found in output:\n%s", out) + } + var got []int + for _, entry := range regexp.MustCompile(`(?m)^\s*(\d+),`).FindAllSubmatch(m[1], -1) { + n, err := strconv.Atoi(string(entry[1])) + if err != nil { + t.Fatal(err) + } + got = append(got, n) + } + want := []int{2, 2, 1, 1, 1} + if !slices.Equal(got, want) { + t.Errorf("NumStmt = %v, want %v", got, want) + } +} diff --git a/src/cmd/go.mod b/src/cmd/go.mod index 523528eae2e7c8..068d8d22476f1a 100644 --- a/src/cmd/go.mod +++ b/src/cmd/go.mod @@ -11,7 +11,7 @@ require ( golang.org/x/sys v0.45.0 golang.org/x/telemetry v0.0.0-20260519152614-eab6ae52b5e2 golang.org/x/term v0.43.0 - golang.org/x/tools v0.45.1-0.20260826175739-e1f45aa8aed5 + golang.org/x/tools v0.45.1-0.20260917185718-3dd9077b8066 ) require ( diff --git a/src/cmd/go.sum b/src/cmd/go.sum index 12edd38b08e27f..a5248cabd280a0 100644 --- a/src/cmd/go.sum +++ b/src/cmd/go.sum @@ -22,7 +22,7 @@ golang.org/x/term v0.43.0 h1:S4RLU2sB31O/NCl+zFN9Aru9A/Cq2aqKpTZJ6B+DwT4= golang.org/x/term v0.43.0/go.mod h1:lrhlHNdQJHO+1qVYiHfFKVuVioJIheAc3fBSMFYEIsk= golang.org/x/text v0.37.0 h1:Cqjiwd9eSg8e0QAkyCaQTNHFIIzWtidPahFWR83rTrc= golang.org/x/text v0.37.0/go.mod h1:a5sjxXGs9hsn/AJVwuElvCAo9v8QYLzvavO5z2PiM38= -golang.org/x/tools v0.45.1-0.20260826175739-e1f45aa8aed5 h1:vnaehVTejSNXKBvdx1mZMwNt7N3hxwIMC9CMJbf/iLI= -golang.org/x/tools v0.45.1-0.20260826175739-e1f45aa8aed5/go.mod h1:LuUGqqaXcXMEFEruIVJVm5mgDD8vww/z/SR1gQ4uE/0= +golang.org/x/tools v0.45.1-0.20260917185718-3dd9077b8066 h1:xPYC80nKQHF/E3LNhvTNYVga54nQm+ziyokKY2lud1c= +golang.org/x/tools v0.45.1-0.20260917185718-3dd9077b8066/go.mod h1:LuUGqqaXcXMEFEruIVJVm5mgDD8vww/z/SR1gQ4uE/0= rsc.io/markdown v0.0.0-20240306144322-0bf8f97ee8ef h1:mqLYrXCXYEZOop9/Dbo6RPX11539nwiCNBb1icVPmw8= rsc.io/markdown v0.0.0-20240306144322-0bf8f97ee8ef/go.mod h1:8xcPgWmwlZONN1D9bjxtHEjrUtSEa3fakVF8iaewYKQ= diff --git a/src/cmd/go/internal/fips140/fips140.go b/src/cmd/go/internal/fips140/fips140.go index 64e8bded3235d6..f5d3c2c758459e 100644 --- a/src/cmd/go/internal/fips140/fips140.go +++ b/src/cmd/go/internal/fips140/fips140.go @@ -2,7 +2,7 @@ // Use of this source code is governed by a BSD-style // license that can be found in the LICENSE file. -// Package fips implements support for the GOFIPS140 build setting. +// Package fips140 implements support for the GOFIPS140 build setting. // // The GOFIPS140 build setting controls two aspects of the build: // @@ -230,48 +230,62 @@ func initDir() { file := filepath.Join(cfg.GOROOT, "lib/fips140", v+".zip") ctx := context.Background() - // The FIPS 140-3 Security Policy require checking the SHA-256 hash of the - // zip file. Verify it once against fips140.sum before unpacking it. - if _, err := modfetch.DownloadDir(ctx, mod); err != nil { - sumfile := filepath.Join(cfg.GOROOT, "lib/fips140/fips140.sum") - if err := verifyZipSum(file, sumfile); err != nil { - base.Fatalf("go: verifying GOFIPS140=%v: %v", v, err) - } + sumfile := filepath.Join(cfg.GOROOT, "lib/fips140/fips140.sum") + _, ziphash, err := lookupZipSum(sumfile, filepath.Base(file)) + if err != nil { + base.Fatalf("go: verifying GOFIPS140=%v: %v", v, err) } - zdir, err := modfetch.NewFetcher().Unzip(ctx, mod, file) + // fips140.sum records both the SHA-256 hash of the zip file and its + // go.sum-style module zip hash. Unzip uses the cached copy in the + // module cache only if the cache records that module zip hash for it, + // and otherwise discards the copy and unpacks the snapshot again. + // + // The FIPS 140-3 Security Policy requires checking the SHA-256 hash + // of the zip file. Unzip calls verify exactly when it is about to + // unpack the zip file, so the hash is checked once per unpacking + // (whatever the reason for it) rather than on every go command. + verify := func() error { return verifyZipSum(file, sumfile) } + zdir, err := modfetch.NewFetcher().Unzip(ctx, mod, file, ziphash, verify) if err != nil { base.Fatalf("go: unpacking GOFIPS140=%v: %v", v, err) } dir = filepath.Join(zdir, "fips140") } -// verifyZipSum checks that the SHA-256 hash of zipfile matches the entry -// for its base name in sumfile, which is expected to be in the format of -// GOROOT/lib/fips140/fips140.sum: "NAME SHA256HEX" lines, with "#" comments. -func verifyZipSum(zipfile, sumfile string) error { +// lookupZipSum returns the SHA-256 hash and module zip hash recorded for +// name in sumfile, which is expected to be in the format of +// GOROOT/lib/fips140/fips140.sum: "NAME SHA256HEX H1HASH" lines, with "#" +// comments. H1HASH is the module zip hash in the format used by go.sum. +func lookupZipSum(sumfile, name string) (sha256hex, ziphash string, err error) { sums, err := os.ReadFile(sumfile) if err != nil { - return err + return "", "", err } - name := filepath.Base(zipfile) - var want string for line := range strings.SplitSeq(string(sums), "\n") { line = strings.TrimSpace(line) if line == "" || strings.HasPrefix(line, "#") { continue } - n, h, ok := strings.Cut(line, " ") - if !ok { + f := strings.Fields(line) + if len(f) != 3 || f[0] != name { continue } - if n == name { - want = strings.TrimSpace(h) - break + if !strings.HasPrefix(f[2], "h1:") { + return "", "", fmt.Errorf("malformed module zip hash %q for %s in %s", f[2], name, sumfile) } + return f[1], f[2], nil } - if want == "" { - return fmt.Errorf("no SHA-256 hash for %s in %s", name, sumfile) + return "", "", fmt.Errorf("no entry for %s in %s", name, sumfile) +} + +// verifyZipSum checks that the SHA-256 hash of zipfile matches the entry +// for its base name in sumfile. +func verifyZipSum(zipfile, sumfile string) error { + name := filepath.Base(zipfile) + want, _, err := lookupZipSum(sumfile, name) + if err != nil { + return err } f, err := os.Open(zipfile) if err != nil { diff --git a/src/cmd/go/internal/fips140/fips_test.go b/src/cmd/go/internal/fips140/fips_test.go index 8f4a669eef5b9b..bb9ccb70c6925d 100644 --- a/src/cmd/go/internal/fips140/fips_test.go +++ b/src/cmd/go/internal/fips140/fips_test.go @@ -15,6 +15,8 @@ import ( "slices" "strings" "testing" + + "golang.org/x/mod/sumdb/dirhash" ) var update = flag.Bool("update", false, "update GOROOT/lib/fips140/fips140.sum") @@ -33,8 +35,8 @@ func TestSums(t *testing.T) { t.Fatal(err) } - format := func(name string, sum [32]byte) string { - return fmt.Sprintf("%s %x\n", name, sum[:]) + format := func(name string, sum [32]byte, ziphash string) string { + return fmt.Sprintf("%s %x %s\n", name, sum[:], ziphash) } want := make(map[string]string) @@ -43,8 +45,12 @@ func TestSums(t *testing.T) { if err != nil { t.Fatal(err) } + ziphash, err := dirhash.HashZip(zip, dirhash.DefaultHash) + if err != nil { + t.Fatal(err) + } name := filepath.Base(zip) - want[name] = format(name, sha256.Sum256(data)) + want[name] = format(name, sha256.Sum256(data), ziphash) } // Process diff, deleting or correcting stale lines. @@ -117,8 +123,13 @@ func TestVerifyZipSum(t *testing.T) { } } + const ( + zeroSum = "0000000000000000000000000000000000000000000000000000000000000000" + ziphash = "h1:AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=" + ) + // Matching hash with comments and a second entry passes. - write(fmt.Sprintf("# comment\n\nv1.2.3.zip %x\nother.zip 0000000000000000000000000000000000000000000000000000000000000000\n", sum[:])) + write(fmt.Sprintf("# comment\n\nv1.2.3.zip %x %s\nother.zip %s h1:other=\n", sum[:], ziphash, zeroSum)) if err := verifyZipSum(zipfile, sumfile); err != nil { t.Errorf("unexpected error: %v", err) } @@ -130,7 +141,7 @@ func TestVerifyZipSum(t *testing.T) { } // Wrong hash fails. - write("v1.2.3.zip 0000000000000000000000000000000000000000000000000000000000000000\n") + write(fmt.Sprintf("v1.2.3.zip %s %s\n", zeroSum, ziphash)) if err := verifyZipSum(zipfile, sumfile); err == nil { t.Errorf("expected error when hash does not match") } diff --git a/src/cmd/go/internal/modfetch/cache.go b/src/cmd/go/internal/modfetch/cache.go index 2e6e7a5048034b..7b728906d34899 100644 --- a/src/cmd/go/internal/modfetch/cache.go +++ b/src/cmd/go/internal/modfetch/cache.go @@ -115,15 +115,10 @@ func DownloadDir(ctx context.Context, m module.Version) (string, error) { return dir, err } - // Special case: ziphash is not required for the golang.org/fips140 module, - // because it is unpacked from a file in GOROOT, not downloaded. - // We've already checked that it's not a partial unpacking, so we're happy. - if m.Path == "golang.org/fips140" { - return dir, nil - } - - // Check if a .ziphash file exists. It should be created before the - // zip is extracted, but if it was deleted (by another program?), we need + // Check if a .ziphash file exists. For downloaded modules it is created + // before the zip is extracted; for GOFIPS140 snapshots, which Fetcher.Unzip + // unpacks from GOROOT, it is created only after extraction completes. + // Either way, if it is missing (deleted by another program?), we need // to re-calculate it. Note that checkMod will repopulate the ziphash // file if it doesn't exist, but if the module is excluded by checks // through GONOSUMDB or GOPRIVATE, that check and repopulation won't happen. diff --git a/src/cmd/go/internal/modfetch/fetch.go b/src/cmd/go/internal/modfetch/fetch.go index 28cc87921ad162..74cbbe06468c8e 100644 --- a/src/cmd/go/internal/modfetch/fetch.go +++ b/src/cmd/go/internal/modfetch/fetch.go @@ -7,6 +7,15 @@ package modfetch import ( "archive/zip" "bytes" + "cmd/go/internal/base" + "cmd/go/internal/cfg" + "cmd/go/internal/fsys" + "cmd/go/internal/gover" + "cmd/go/internal/lockedfile" + "cmd/go/internal/str" + "cmd/go/internal/trace" + "cmd/internal/par" + "cmd/internal/robustio" "context" "crypto/sha256" "encoding/base64" @@ -20,22 +29,15 @@ import ( "strings" "sync" - "cmd/go/internal/base" - "cmd/go/internal/cfg" - "cmd/go/internal/fsys" - "cmd/go/internal/gover" - "cmd/go/internal/lockedfile" - "cmd/go/internal/str" - "cmd/go/internal/trace" - "cmd/internal/par" - "cmd/internal/robustio" - "golang.org/x/mod/module" "golang.org/x/mod/sumdb/dirhash" modzip "golang.org/x/mod/zip" ) -var ErrToolchain = errors.New("internal error: invalid operation on toolchain module") +var ( + ErrToolchain = errors.New("internal error: invalid operation on toolchain module") + ErrFIPS140 = errors.New("golang.org/fips140 is bundled with the Go distribution and cannot be downloaded") +) // Download downloads the specific module version to the // local download cache and returns the name of the directory @@ -44,6 +46,9 @@ func (f *Fetcher) Download(ctx context.Context, mod module.Version) (dir string, if gover.IsToolchain(mod.Path) { return "", ErrToolchain } + if mod.Path == "golang.org/fips140" { + return "", ErrFIPS140 + } if err := checkCacheDir(ctx); err != nil { base.Fatal(err) } @@ -70,10 +75,29 @@ func (f *Fetcher) Download(ctx context.Context, mod module.Version) (dir string, }) } -// Unzip is like Download but is given the explicit zip file to use, -// rather than downloading it. This is used for the GOFIPS140 zip files, -// which ship in the Go distribution itself. -func (f *Fetcher) Unzip(ctx context.Context, mod module.Version, zipfile string) (dir string, err error) { +// Unzip is like Download but for GOFIPS140 zip files which ship with +// the Go distribution itself. +// +// Unzip performs a check like the go.sum check for downloaded modules: +// A cached copy of mod is used only if the module cache records ziphash +// as its module zip hash. As with go.sum, this rejects a cached copy +// that was populated from some other source; it does not detect +// modifications made to the module cache directory itself. +// +// Otherwise, any existing copy is discarded, zipfile is unpacked +// again, and ziphash is stamped upon success. +// +// If verify is non-nil, it is called before zipfile is unpacked. +// Unzip does not call verify when it returns a cached copy. +// +// If verify returns an error, the module cache is left untouched and +// Unzip returns that error. +// +// ziphash must be non-empty. +func (f *Fetcher) Unzip(ctx context.Context, mod module.Version, zipfile, ziphash string, verify func() error) (dir string, err error) { + if ziphash == "" { + return "", module.VersionError(mod, errors.New("internal error: Unzip called with empty module zip hash")) + } if err := checkCacheDir(ctx); err != nil { base.Fatal(err) } @@ -85,15 +109,47 @@ func (f *Fetcher) Unzip(ctx context.Context, mod module.Version, zipfile string) dir, err = DownloadDir(ctx, mod) if err == nil { // The directory has already been completely extracted (no .partial file exists). - return dir, nil + ok, err := haveZipHash(ctx, mod, ziphash) + if err != nil { + return "", err + } + if ok { + return dir, nil + } } else if dir == "" || !errors.Is(err, fs.ErrNotExist) { return "", err } - return unzip(ctx, mod, zipfile) + return unzip(ctx, mod, zipfile, ziphash, verify) }) } +func haveZipHash(ctx context.Context, mod module.Version, ziphash string) (bool, error) { + path, err := CachePath(ctx, mod, "ziphash") + if err != nil { + return false, err + } + data, err := lockedfile.Read(path) + if errors.Is(err, fs.ErrNotExist) { + return false, nil + } + if err != nil { + return false, err + } + return strings.TrimSpace(string(data)) == ziphash, nil +} + +func writeZipHash(ctx context.Context, mod module.Version, ziphash string) error { + path, err := CachePath(ctx, mod, "ziphash") + if err != nil { + return err + } + if err := os.MkdirAll(filepath.Dir(path), 0o777); err != nil { + return err + } + return lockedfile.Write(path, strings.NewReader(ziphash), 0o666) +} + func (f *Fetcher) download(ctx context.Context, mod module.Version) (dir string, err error) { ctx, span := trace.StartSpan(ctx, "modfetch.download "+mod.String()) defer span.Done() @@ -114,10 +170,16 @@ func (f *Fetcher) download(ctx context.Context, mod module.Version) (dir string, return "", err } - return unzip(ctx, mod, zipfile) + return unzip(ctx, mod, zipfile, "", nil) } -func unzip(ctx context.Context, mod module.Version, zipfile string) (dir string, err error) { +// unzip extracts zipfile into the module cache directory for mod, +// unless a complete copy is already there (and, if ziphash is non-empty, +// the module cache records ziphash for it). +// +// If verify is non-nil, it is called once the decision to extract has +// been made, before any existing copy is discarded. +func unzip(ctx context.Context, mod module.Version, zipfile, ziphash string, verify func() error) (dir string, err error) { unlock, err := lockVersion(ctx, mod) if err != nil { return "", err @@ -130,10 +192,29 @@ func unzip(ctx context.Context, mod module.Version, zipfile string) (dir string, // Check whether the directory was populated while we were waiting on the lock. dir, dirErr := DownloadDir(ctx, mod) if dirErr == nil { - return dir, nil + if ziphash == "" { + return dir, nil + } + ok, err := haveZipHash(ctx, mod, ziphash) + if err != nil { + return "", err + } + if ok { + return dir, nil + } + dirErr = &DownloadDirPartialError{dir, errors.New("ziphash file does not match")} } _, dirExists := dirErr.(*DownloadDirPartialError) + // We are going to extract zipfile. Verify it first, while the module + // cache is still untouched, so that a bad zip file never costs us a + // cached copy and a good one is verified at most once per extraction. + if verify != nil { + if err := verify(); err != nil { + return "", err + } + } + // Clean up any remaining temporary directories created by old versions // (before 1.16), as well as partially extracted directories (indicated by // DownloadDirPartialError, usually because of a .partial file). This is only @@ -151,6 +232,19 @@ func unzip(ctx context.Context, mod module.Version, zipfile string) (dir string, return "", err } } + // Remove any recorded module zip hash before extracting, and record the + // new one only after the .partial file is removed (below). That way a + // .ziphash file for the module exists only beside a completely extracted + // directory, no matter where an unpack is interrupted. + if ziphash != "" { + hashPath, err := CachePath(ctx, mod, "ziphash") + if err != nil { + return "", err + } + if err := robustio.RemoveAll(hashPath); err != nil { + return "", err + } + } partialPath, err := CachePath(ctx, mod, "partial") if err != nil { @@ -187,6 +281,13 @@ func unzip(ctx context.Context, mod module.Version, zipfile string) (dir string, if err := os.Remove(partialPath); err != nil { return "", err } + // The directory is complete: record the module zip hash of the zip + // file it was extracted from. See the removal above. + if ziphash != "" { + if err := writeZipHash(ctx, mod, ziphash); err != nil { + return "", err + } + } if !cfg.ModCacheRW { makeDirsReadOnly(dir) @@ -199,6 +300,10 @@ var downloadZipCache par.ErrCache[module.Version, string] // DownloadZip downloads the specific module version to the // local zip cache and returns the name of the zip file. func (f *Fetcher) DownloadZip(ctx context.Context, mod module.Version) (zipfile string, err error) { + if mod.Path == "golang.org/fips140" { + return "", ErrFIPS140 + } + // The par.Cache here avoids duplicate work. return downloadZipCache.Do(mod, func() (string, error) { zipfile, err := CachePath(ctx, mod, "zip") @@ -776,11 +881,10 @@ func checkModSum(f *Fetcher, mod module.Version, h string) error { } f.mu.Unlock() - if done { + if done && mod.Path != "golang.org/toolchain" { return nil } - // Not listed, so we want to add them. // Consult checksum database if appropriate. if useSumDB(mod) { // Calls base.Fatalf if mismatch detected. @@ -789,6 +893,10 @@ func checkModSum(f *Fetcher, mod module.Version, h string) error { } } + if done { + return nil + } + // Add mod+h to go.sum, if it hasn't appeared already. if inited { f.mu.Lock() diff --git a/src/cmd/go/internal/modfetch/fetch_test.go b/src/cmd/go/internal/modfetch/fetch_test.go new file mode 100644 index 00000000000000..4a941510e5b19e --- /dev/null +++ b/src/cmd/go/internal/modfetch/fetch_test.go @@ -0,0 +1,101 @@ +// Copyright 2026 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +//go:build !js && !wasip1 + +package modfetch + +import ( + "errors" + "io/fs" + "os" + "path/filepath" + "testing" + + "golang.org/x/mod/module" + "golang.org/x/mod/sumdb/dirhash" + modzip "golang.org/x/mod/zip" +) + +func TestUnzipVerify(t *testing.T) { + var ( + ctx = t.Context() + src = t.TempDir() + mod = module.Version{Path: "example.com/unzip", Version: "v1.0.0"} + ) + if err := os.WriteFile(filepath.Join(src, "go.mod"), []byte("module example.com/unzip\n"), 0o666); err != nil { + t.Fatal(err) + } + + zipfile := filepath.Join(t.TempDir(), "v1.0.0.zip") + zf, err := os.Create(zipfile) + if err != nil { + t.Fatal(err) + } + if err := modzip.CreateFromDir(zf, mod, src); err != nil { + t.Fatal(err) + } + if err := zf.Close(); err != nil { + t.Fatal(err) + } + ziphash, err := dirhash.HashZip(zipfile, dirhash.DefaultHash) + if err != nil { + t.Fatal(err) + } + + unzip := func(verify func() error) (string, error) { + dir, err := NewFetcher().Unzip(ctx, mod, zipfile, ziphash, verify) + if dir != "" { + t.Cleanup(func() { RemoveAll(dir) }) + } + return dir, err + } + + var ( + errVerify = errors.New("verify failed") + verifyFnFails = func() error { return errVerify } + verifyFnPasses = func() error { return nil } + ) + if _, err := unzip(verifyFnFails); !errors.Is(err, errVerify) { + t.Fatalf("unzip with empty cache: got %v, want %v", err, errVerify) + } + if _, err := DownloadDir(ctx, mod); !errors.Is(err, fs.ErrNotExist) { + t.Fatalf("DownloadDir after failed verify: got %v, want %v", err, fs.ErrNotExist) + } + + if _, err := unzip(verifyFnPasses); err != nil { + t.Fatal(err) + } + if _, err := DownloadDir(ctx, mod); err != nil { + t.Fatal(err) + } + if ok, err := haveZipHash(ctx, mod, ziphash); err != nil || !ok { + t.Fatalf("haveZipHash = (%v, %v), want (true, nil)", ok, err) + } + + if _, err := unzip(verifyFnFails); err != nil { + t.Fatalf("unzip with cached copy: got %v, want nil", err) + } + + if err := writeZipHash(ctx, mod, "h1:AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA="); err != nil { + t.Fatal(err) + } + if _, err := unzip(verifyFnFails); !errors.Is(err, errVerify) { + t.Fatalf("unzip with mismatched ziphash: got %v, want %v", err, errVerify) + } + if _, err := DownloadDir(ctx, mod); err != nil { + t.Fatal(err) + } + + if _, err := unzip(verifyFnPasses); err != nil { + t.Fatal(err) + } + if _, err := DownloadDir(ctx, mod); err != nil { + t.Fatal(err) + } + + if ok, err := haveZipHash(ctx, mod, ziphash); err != nil || !ok { + t.Errorf("haveZipHash after re-unpack: got (%v, %v), want (true, nil)", ok, err) + } +} diff --git a/src/cmd/go/internal/modfetch/sumdb.go b/src/cmd/go/internal/modfetch/sumdb.go index ea7d561d7b9e0b..6fb435b010d67c 100644 --- a/src/cmd/go/internal/modfetch/sumdb.go +++ b/src/cmd/go/internal/modfetch/sumdb.go @@ -35,9 +35,8 @@ import ( func useSumDB(mod module.Version) bool { if mod.Path == "golang.org/toolchain" { must := true - // Downloaded toolchains cannot be listed in go.sum, - // so we require checksum database lookups even if - // GOSUMDB=off or GONOSUMDB matches the pattern. + // Toolchain downloads must be verified against the checksum database, + // even if GOSUMDB=off or GONOSUMDB matches the pattern. // If GOSUMDB=off, then the eventual lookup will fail // with a good error message. diff --git a/src/cmd/go/internal/work/exec.go b/src/cmd/go/internal/work/exec.go index 8dcf324529f5df..3ed01c0fbd49b8 100644 --- a/src/cmd/go/internal/work/exec.go +++ b/src/cmd/go/internal/work/exec.go @@ -1349,6 +1349,7 @@ func (b *Builder) vet(ctx context.Context, a *Action) error { h := cache.NewHash("vet " + a.Package.ImportPath) fmt.Fprintf(h, "vet %q\n", b.toolID("vet")) + fmt.Fprintf(h, "vetxonly %v\n", vcfg.VetxOnly) vetFlags := VetFlags diff --git a/src/cmd/go/testdata/script/fipssnap_modcache.txt b/src/cmd/go/testdata/script/fipssnap_modcache.txt new file mode 100644 index 00000000000000..ca43ec9a7d6ce2 --- /dev/null +++ b/src/cmd/go/testdata/script/fipssnap_modcache.txt @@ -0,0 +1,61 @@ +# The module cache copy of a FIPS snapshot is used only if the module +# cache records the module zip hash from GOROOT/lib/fips140/fips140.sum +# for it, the way downloaded modules are checked against go.sum. +# Otherwise it is discarded and the snapshot is unpacked again. +# +# This detects a cache entry that was not unpacked from the bundled +# snapshot (or whose hash record is missing). Like go.sum, it does not +# protect the module cache from being modified in place: the replaced +# source file below stands in for stale contents, not an attacker. + +env snap=v1.26.0 +env GOFIPS140=$snap +env GOMODCACHE=$WORK/modcache +env GOFLAGS=-modcacherw + +# Go+BoringCrypto conflicts with GOFIPS140. +[GOEXPERIMENT:boringcrypto] skip + +env ziphash=$GOMODCACHE/cache/download/golang.org/fips140/@v/$snap.ziphash +env srcfile=$GOMODCACHE/golang.org/fips140@$snap/fips140/$snap/sha256/sha256.go + +# unpacking the snapshot records its module zip hash in the module cache +go list -f '{{.DefaultGODEBUG}}' +stdout fips140=on +exists $ziphash +exists $srcfile +grep '^h1:dtoPX1ALGGp4rMLzyh6oqIkYRXnXxRkpPWu56l5DFpM=$' $ziphash +cp $ziphash good.ziphash + +# a recorded hash that does not match fips140.sum +# discards the cached copy and unpacks the snapshot again +cp bad.ziphash $ziphash +rm $srcfile +cp stale/sha256.go $srcfile +go list -f '{{.DefaultGODEBUG}}' +stdout fips140=on +exists $srcfile +! grep stale $srcfile +cmp $ziphash good.ziphash + +# so does a missing hash +rm $ziphash +rm $srcfile +cp stale/sha256.go $srcfile +go list -f '{{.DefaultGODEBUG}}' +stdout fips140=on +exists $srcfile +! grep stale $srcfile +cmp $ziphash good.ziphash + +-- go.mod -- +module m +-- x.go -- +package main +import _ "crypto/sha256" +func main() { +} +-- bad.ziphash -- +h1:AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA= +-- stale/sha256.go -- +package sha256 // stale diff --git a/src/cmd/go/testdata/script/gotoolchain_switch_gosum.txt b/src/cmd/go/testdata/script/gotoolchain_switch_gosum.txt new file mode 100644 index 00000000000000..aaac9a9b816f1b --- /dev/null +++ b/src/cmd/go/testdata/script/gotoolchain_switch_gosum.txt @@ -0,0 +1,40 @@ +# A toolchain switch triggered by a dependency must verify the downloaded +# toolchain against the checksum database, even when the main module's go.sum +# already lists a matching entry for golang.org/toolchain. + +[!exec:/bin/sh] skip 'the fake proxy serves shell scripts instead of binaries' +env TESTGO_VERSION=go1.21.0 +env GOTOOLCHAIN=local +env sumdb=$GOSUMDB +env proxy=$GOPROXY +env dbname=localhost.localdev/sumdb + +# Record the toolchain in go.sum, then drop it from go.mod so that only +# the go.sum line remains, as it would in an attacker-supplied repository. +go get golang.org/toolchain@v0.0.1-go1.999testmod.$GOOS-$GOARCH +go mod edit -droprequire golang.org/toolchain +grep '^golang.org/toolchain v0.0.1-go1.999testmod.[a-z0-9\-]* h1:' go.sum +go mod edit -require rsc.io/future@v1.0.0 + +# Point at a checksum database that disagrees with go.sum and the download. +# GONOSUMDB keeps rsc.io/future out of the way; it does not apply to the toolchain. +# Clear cached lookups and the cached tree head so the server is consulted. +go clean -modcache +rm $GOPATH/pkg/sumdb/$dbname/latest +env GOTOOLCHAIN=auto +env GONOSUMDB=rsc.io +env GOSUMDB=$sumdb' '$proxy/sumdb-wrong +! go get . +stderr 'switching to go1.999testmod' +stderr 'golang.org/toolchain@v0.0.1-go1.999testmod.[a-z0-9\-]*: verifying (module|go.mod): checksum mismatch' +stderr 'localhost.localdev/sumdb: h1:wrong' +stderr 'SECURITY ERROR' + +-- go.mod -- +module example + +go 1.21 +-- example.go -- +package example + +import _ "rsc.io/future" diff --git a/src/cmd/go/testdata/script/mod_download_toolchain_gosum.txt b/src/cmd/go/testdata/script/mod_download_toolchain_gosum.txt new file mode 100644 index 00000000000000..9954fb263f7e7c --- /dev/null +++ b/src/cmd/go/testdata/script/mod_download_toolchain_gosum.txt @@ -0,0 +1,34 @@ +# A matching line in go.sum must not bypass checksum database verification +# for a downloaded toolchain, since useSumDB requires the checksum database +# for golang.org/toolchain even when GOSUMDB=off. + +env GOTOOLCHAIN=local +env sumdb=$GOSUMDB +env proxy=$GOPROXY +env dbname=localhost.localdev/sumdb + +go get golang.org/toolchain@v0.0.1-go1.999testmod.$GOOS-$GOARCH +grep '^golang.org/toolchain v0.0.1-go1.999testmod.[a-z0-9\-]* h1:' go.sum +grep '^golang.org/toolchain v0.0.1-go1.999testmod.[a-z0-9\-]*/go.mod h1:' go.sum + +# With the checksum database disabled, the matching go.sum entry +# must not be accepted on its own. +env GOSUMDB=off +! go mod download golang.org/toolchain +stderr 'checksum database disabled by GOSUMDB=off' + +# With a checksum database that disagrees with go.sum, the download +# must be rejected even though go.sum matches the downloaded bits. +# Clear cached lookups and the cached tree head so the server is consulted. +go clean -modcache +rm $GOPATH/pkg/sumdb/$dbname/latest +env GOSUMDB=$sumdb' '$proxy/sumdb-wrong +! go mod download golang.org/toolchain +stderr 'verifying (module|go.mod): checksum mismatch' +stderr 'localhost.localdev/sumdb: h1:wrong' +stderr 'SECURITY ERROR' + +-- go.mod -- +module example.com/m + +go 1.21 diff --git a/src/cmd/go/testdata/script/test_chatty_ascii.txt b/src/cmd/go/testdata/script/test_chatty_ascii.txt new file mode 100644 index 00000000000000..6d2570f626443d --- /dev/null +++ b/src/cmd/go/testdata/script/test_chatty_ascii.txt @@ -0,0 +1,24 @@ +# Make sure the ESC character isn't doubled. +! go test -v + +stdout ' x_test.go:11: \x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\n \x0b\x0c\x0d\x0e\x0f\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f\n' +stdout ' x_test.go:12: \x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\n \x0b\x0c\x0d\x0e\x0f\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f\n' +! stdout '\x1b\x1b' + +-- go.mod -- +module p + +-- x_test.go -- +package p + +import "testing" + +func Test(t *testing.T) { + var s string + for i := rune(0); i < ' '; i++ { + s += string(i) + } + + t.Log(s) + t.Error(s) +} diff --git a/src/cmd/go/testdata/script/vet_cache.txt b/src/cmd/go/testdata/script/vet_cache.txt index 624df5573240c1..3871352d91b670 100644 --- a/src/cmd/go/testdata/script/vet_cache.txt +++ b/src/cmd/go/testdata/script/vet_cache.txt @@ -15,6 +15,14 @@ stderr 'fmt.Sprint call has possible Printf formatting directive' ! go vet example.com/a stderr 'fmt.Sprint call has possible Printf formatting directive' +# Vetting example.com/a only as a dependency of example.com/b +# records facts but reports no diagnostics. A later +# 'go vet example.com/a' must not reuse that result. +env GOCACHE=$WORK/gocache2 +go vet example.com/b +! go vet example.com/a +stderr 'fmt.Sprint call has possible Printf formatting directive' + -- go.mod -- module example.com @@ -24,3 +32,8 @@ package a import "fmt" var _ = fmt.Sprint("%s") // oops! + +-- b/b.go -- +package b + +import _ "example.com/a" diff --git a/src/cmd/internal/obj/ppc64/asm_test.go b/src/cmd/internal/obj/ppc64/asm_test.go index 9f1acf4b62e2ee..8b6d22e69243be 100644 --- a/src/cmd/internal/obj/ppc64/asm_test.go +++ b/src/cmd/internal/obj/ppc64/asm_test.go @@ -13,9 +13,12 @@ import ( "os" "path/filepath" "regexp" + "strconv" "strings" "testing" + "internal/abi" + "cmd/internal/obj" "cmd/internal/objabi" ) @@ -556,3 +559,109 @@ func TestOptabReinit(t *testing.T) { t.Errorf("rerunning buildop changes optab size from %d to %d", optabLen, reinitOptabLen) } } + +// A tail call is lowered to "MOVD Rx, CTR; BR (CTR)". runtime.asyncPreempt +// does not preserve CTR, and its resume sequence leaves CTR holding the resume +// PC, so a goroutine preempted anywhere between the load of CTR and the branch +// would resume by branching to the wrong place (for a preemption at the branch +// itself, to that very instruction, spinning there forever). Check that the +// whole sequence is marked as an unsafe point. See go.dev/issue/78576. +const tailCallSrc = ` +// 4 = NOSPLIT, 512 = NOFRAME (see textflag.h). +TEXT ·leafNoFrame(SB),4|512,$0-0 + MOVD $0, R3 + RET (R3) + +TEXT ·leafFrame(SB),4,$8-0 + MOVD $0, R3 + RET (R3) + +TEXT ·nonLeaf(SB),4,$8-0 + CALL ·leafNoFrame(SB) + MOVD $0, R3 + RET (R3) + +// A tail call that is not the last instruction in the function, +// to check that the unsafe point ends at the branch and does not +// swallow the rest of the function. +TEXT ·leafNoFrameCond(SB),4|512,$0-0 + CMP R3, $0 + BEQ skip + RET (R3) +skip: + MOVD $0, R3 + RET +` + +func TestTailCallUnsafePoint(t *testing.T) { + testenv.MustHaveGoBuild(t) + + dir := t.TempDir() + tmpfile := filepath.Join(dir, "x.s") + if err := os.WriteFile(tmpfile, []byte(tailCallSrc), 0644); err != nil { + t.Fatalf("can't write output: %v\n", err) + } + + for _, goarch := range []string{"ppc64", "ppc64le"} { + cmd := testenv.Command(t, testenv.GoToolPath(t), "tool", "asm", "-o", filepath.Join(dir, "x.o"), "-S", tmpfile) + cmd.Env = append(os.Environ(), "GOARCH="+goarch, "GOOS=linux") + out, err := cmd.CombinedOutput() + if err != nil { + t.Fatalf("GOARCH=%s: assembly failed: %v, output:\n%s", goarch, err, out) + } + + // Walk the -S output tracking the current PCDATA_UnsafePoint value. + // Every branch through CTR must be an unsafe point, and the unsafe + // point must end at the branch: the next instruction, if any, must be + // a safe point again, and so must the end of the function. + branches := 0 + sym := "" + unsafePoint := int64(abi.UnsafePointSafe) + atBranch := false + endFunc := func() { + if sym != "" && unsafePoint != abi.UnsafePointSafe { + t.Errorf("GOARCH=%s: %s: unsafe point %d at end of function, want %d", + goarch, sym, unsafePoint, abi.UnsafePointSafe) + } + } + for _, line := range strings.Split(string(out), "\n") { + f := strings.Fields(line) + // Lines look like: + // sym STEXT ... + // 0x0014 00020 (x.s:14)PCDATA$0,$-2 + // Other lines (hex dumps, relocations) have no (file:line) field. + switch { + case len(f) >= 2 && f[1] == "STEXT": + endFunc() + sym, unsafePoint, atBranch = f[0], abi.UnsafePointSafe, false + case len(f) < 4 || !strings.HasPrefix(f[2], "("): + // Not an instruction. + case len(f) >= 6 && f[3] == "PCDATA" && f[4] == fmt.Sprintf("$%d,", abi.PCDATA_UnsafePoint): + v, err := strconv.ParseInt(strings.TrimPrefix(f[5], "$"), 10, 64) + if err != nil { + t.Fatalf("GOARCH=%s: can't parse %q: %v", goarch, line, err) + } + unsafePoint = v + case f[3] == "PCDATA" || f[3] == "FUNCDATA" || f[3] == "TEXT": + // Not a real instruction. + case len(f) >= 5 && f[3] == "JMP" && f[4] == "CTR": + branches++ + atBranch = true + if unsafePoint != abi.UnsafePointUnsafe { + t.Errorf("GOARCH=%s: %s\n\tbranch through CTR has unsafe point %d, want %d", + goarch, strings.TrimSpace(line), unsafePoint, abi.UnsafePointUnsafe) + } + case atBranch: + atBranch = false + if unsafePoint != abi.UnsafePointSafe { + t.Errorf("GOARCH=%s: %s\n\tinstruction after branch through CTR has unsafe point %d, want %d", + goarch, strings.TrimSpace(line), unsafePoint, abi.UnsafePointSafe) + } + } + } + endFunc() + if want := strings.Count(tailCallSrc, "RET\t(R3)"); branches != want { + t.Errorf("GOARCH=%s: found %d branches through CTR, want %d; output:\n%s", goarch, branches, want, out) + } + } +} diff --git a/src/cmd/internal/obj/ppc64/obj9.go b/src/cmd/internal/obj/ppc64/obj9.go index 323420fefa4917..fc5d05b4fe0cdb 100644 --- a/src/cmd/internal/obj/ppc64/obj9.go +++ b/src/cmd/internal/obj/ppc64/obj9.go @@ -966,6 +966,7 @@ func preprocess(ctxt *obj.Link, cursym *obj.LSym, newprog obj.ProgAlloc) { } retTarget, retReg := p.To.Sym, p.To.Reg + finish := func(last *obj.Prog) {} if retReg == obj.REG_NONE { retReg = REG_LR } else { @@ -980,6 +981,15 @@ func preprocess(ctxt *obj.Link, cursym *obj.LSym, newprog obj.ProgAlloc) { p.To.Reg = REG_CTR retReg = REG_CTR p.Link = x + // Everything from here through the BR (CTR) must be an + // unsafe point. runtime.asyncPreempt does not preserve CTR, + // and its resume sequence leaves CTR holding the resume PC, + // so a goroutine preempted at the BR (CTR) would resume by + // branching to that very instruction and spin there forever. + c.ctxt.StartUnsafePoint(p, c.newprog) + finish = func(last *obj.Prog) { + c.ctxt.EndUnsafePoint(last, c.newprog, -1) + } p = x } @@ -995,6 +1005,7 @@ func preprocess(ctxt *obj.Link, cursym *obj.LSym, newprog obj.ProgAlloc) { p.To.Sym = retTarget } p.Mark |= BRANCH + finish(p) break } @@ -1020,6 +1031,7 @@ func preprocess(ctxt *obj.Link, cursym *obj.LSym, newprog obj.ProgAlloc) { q.Link = p.Link p.Link = q + finish(q) break } @@ -1089,6 +1101,8 @@ func preprocess(ctxt *obj.Link, cursym *obj.LSym, newprog obj.ProgAlloc) { q1.Link = q.Link prev.Link = q1 + finish(q1) + case AADD: if p.To.Type == obj.TYPE_REG && p.To.Reg == REGSP && p.From.Type == obj.TYPE_CONST { p.Spadj = int32(-p.From.Offset) diff --git a/src/cmd/link/internal/ld/lib.go b/src/cmd/link/internal/ld/lib.go index 9cbc12919f1d00..b9848b232ce4ec 100644 --- a/src/cmd/link/internal/ld/lib.go +++ b/src/cmd/link/internal/ld/lib.go @@ -1479,6 +1479,14 @@ func (ctxt *Link) hostlink() { // resolving a lazy binding. See issue 38824. // Force eager resolution to work around. argv = append(argv, "-Wl,-flat_namespace", "-Wl,-bind_at_load") + if combineDwarf && linkerFlagSupported(ctxt.Arch, argv[0], "", "-Wl,-no_fixup_chains") { + // As of macOS 27, the dynamic linker checks the number of segments + // in recorded in LC_DYLD_CHAINED_FIXUPS and reject the shared object + // with a mismatch. When combining DWARF, we add a segment but + // currently don't fix up the recorded number. Pass -no_fixup_chains to + // the C linker as a workaround. See issue 81793. + argv = append(argv, "-Wl,-no_fixup_chains") + } } if !combineDwarf { argv = append(argv, "-Wl,-S") // suppress STAB (symbolic debugging) symbols diff --git a/src/cmd/vendor/golang.org/x/tools/go/analysis/passes/modernize/errorsastype.go b/src/cmd/vendor/golang.org/x/tools/go/analysis/passes/modernize/errorsastype.go index f52f202e22def9..0e3f17fbcc07a1 100644 --- a/src/cmd/vendor/golang.org/x/tools/go/analysis/passes/modernize/errorsastype.go +++ b/src/cmd/vendor/golang.org/x/tools/go/analysis/passes/modernize/errorsastype.go @@ -32,7 +32,7 @@ var ErrorsAsTypeAnalyzer = &analysis.Analyzer{ Run: errorsastype, } -// errorsastype offers a fix to replace error.As with the newer +// errorsastype offers a fix to replace errors.As with the newer // errors.AsType[T] following this pattern: // // var myerr *MyErr @@ -228,6 +228,11 @@ func canUseErrorsAsType(info *types.Info, index *typeindex.Index, curCall inspec len(curDecl.Node().(*ast.GenDecl).Specs) != 1 { return // not a simple "var v T" decl } + // AsType requires that its type argument implements error. + // Reject if v does not implement error. + if !types.AssignableTo(v.Type(), errorType) { + return + } // Have: // var v *MyErr diff --git a/src/cmd/vendor/golang.org/x/tools/go/analysis/passes/modernize/modernize.go b/src/cmd/vendor/golang.org/x/tools/go/analysis/passes/modernize/modernize.go index d3870e93dbc98c..12700b24230446 100644 --- a/src/cmd/vendor/golang.org/x/tools/go/analysis/passes/modernize/modernize.go +++ b/src/cmd/vendor/golang.org/x/tools/go/analysis/passes/modernize/modernize.go @@ -136,6 +136,7 @@ var ( builtinTrue = types.Universe.Lookup("true") byteSliceType = types.NewSlice(types.Typ[types.Byte]) omitemptyRegex = regexp.MustCompile(`(?:^json| json):"[^"]*(,omitempty)(?:"|,[^"]*")\s?`) + errorType = types.Universe.Lookup("error").Type() ) // lookup returns the symbol denoted by name at the position of the cursor. diff --git a/src/cmd/vendor/modules.txt b/src/cmd/vendor/modules.txt index 27220d946fe80e..b35ee5fc5d5841 100644 --- a/src/cmd/vendor/modules.txt +++ b/src/cmd/vendor/modules.txt @@ -73,7 +73,7 @@ golang.org/x/text/internal/tag golang.org/x/text/language golang.org/x/text/transform golang.org/x/text/unicode/norm -# golang.org/x/tools v0.45.1-0.20260826175739-e1f45aa8aed5 +# golang.org/x/tools v0.45.1-0.20260917185718-3dd9077b8066 ## explicit; go 1.25.0 golang.org/x/tools/cmd/bisect golang.org/x/tools/cover diff --git a/src/compress/flate/deflate.go b/src/compress/flate/deflate.go index fcd9d8c8d57ecd..f7e58e8dfc2f5c 100644 --- a/src/compress/flate/deflate.go +++ b/src/compress/flate/deflate.go @@ -250,6 +250,7 @@ func (d *compressor) fillWindow(b []byte) { // Update window information. d.windowEnd += n s.index = n + d.blockStart = d.windowEnd } // findMatch finds the longest match starting at pos in the hash chain starting diff --git a/src/compress/flate/deflate_test.go b/src/compress/flate/deflate_test.go index 47ef8f0da087d3..073c63a5eef550 100644 --- a/src/compress/flate/deflate_test.go +++ b/src/compress/flate/deflate_test.go @@ -494,6 +494,40 @@ func TestWriterDict(t *testing.T) { } } +// TestNonCompressedBlockDoesntLeakDict checks that the dictionary isn't sent when +// sending a non-compressed block. See https://go.dev/issue/80538 +func TestNonCompressedBlockDoesntLeakDict(t *testing.T) { + data := make([]byte, 763) + rand.New(rand.NewSource(42)).Read(data) + dict := []byte("0123456789abcdefghij") + for l := range BestCompression + 1 { + t.Run(fmt.Sprintf("level=%d", l), func(t *testing.T) { + var b bytes.Buffer + w, err := NewWriterDict(&b, l, dict) + if err != nil { + t.Fatalf("NewWriterDict: %v", err) + } + if _, err := w.Write(data); err != nil { + t.Fatalf("Write: %v", err) + } + if err := w.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + got, err := io.ReadAll(NewReaderDict(&b, dict)) + if err != nil { + t.Fatalf("NewReaderDict: %v", err) + } + if !bytes.Equal(got, data) { + t.Errorf("round trip mismatch: got %d bytes, want %d (dictionary emitted: %v)", + len(got), len(data), bytes.HasPrefix(got, dict)) + } + if b.Len() != 0 { + t.Errorf("compressed stream not fully consumed: %d bytes left", b.Len()) + } + }) + } +} + // See https://golang.org/issue/2508 func TestRegression2508(t *testing.T) { if testing.Short() { diff --git a/src/crypto/tls/ech.go b/src/crypto/tls/ech.go index 91aea409196e77..75a5677b11e3c4 100644 --- a/src/crypto/tls/ech.go +++ b/src/crypto/tls/ech.go @@ -313,6 +313,7 @@ func decodeInnerClientHello(outer *clientHelloMsg, encoded []byte) (*clientHello recon.AddBytes(compressionMethods) }) recon.AddUint16LengthPrefixed(func(recon *cryptobyte.Builder) { + var outerExtensionsSeen bool for !extensions.Empty() { var extension uint16 var extData cryptobyte.String @@ -322,18 +323,30 @@ func decodeInnerClientHello(outer *clientHelloMsg, encoded []byte) (*clientHello return } if extension == extensionECHOuterExtensions { - if !extData.ReadUint8LengthPrefixed(&extData) { + if outerExtensionsSeen { + recon.SetError(errors.New("tls: invalid outer extensions")) + return + } + outerExtensionsSeen = true + var outerExtensions cryptobyte.String + if !extData.ReadUint8LengthPrefixed(&outerExtensions) || !extData.Empty() || + outerExtensions.Empty() { recon.SetError(errors.New("tls: invalid inner client hello")) return } + // OuterExtensions reconstruction per RFC 9849, Appendix A. + // i scans the outer extensions in order and never rewinds, + // so a referenced type that is out of order or duplicated + // cannot be found again and is rejected. var i int - for !extData.Empty() { + for !outerExtensions.Empty() { var extType uint16 - if !extData.ReadUint16(&extType) { + if !outerExtensions.ReadUint16(&extType) { recon.SetError(errors.New("tls: invalid inner client hello")) return } - if extType == extensionEncryptedClientHello { + if extType == extensionEncryptedClientHello || + extType == extensionECHOuterExtensions { recon.SetError(errors.New("tls: invalid outer extensions")) return } @@ -350,6 +363,7 @@ func decodeInnerClientHello(outer *clientHelloMsg, encoded []byte) (*clientHello recon.AddUint16LengthPrefixed(func(recon *cryptobyte.Builder) { recon.AddBytes(rawOuterExts[i].data) }) + i++ } } else { recon.AddUint16(extension) diff --git a/src/crypto/tls/ech_test.go b/src/crypto/tls/ech_test.go index 5cf5035fceccf9..07ecc2cafda0d9 100644 --- a/src/crypto/tls/ech_test.go +++ b/src/crypto/tls/ech_test.go @@ -9,6 +9,8 @@ import ( "encoding/hex" "strings" "testing" + + "golang.org/x/crypto/cryptobyte" ) func TestDecodeECHConfigLists(t *testing.T) { @@ -116,3 +118,130 @@ func TestECHPadding(t *testing.T) { } }) } + +func TestDecodeInnerClientHelloOuterExtensions(t *testing.T) { + outer := &clientHelloMsg{ + vers: VersionTLS12, + random: make([]byte, 32), + cipherSuites: []uint16{TLS_AES_128_GCM_SHA256}, + compressionMethods: []uint8{compressionNone}, + ocspStapling: true, + supportedCurves: []CurveID{CurveP256}, + encryptedClientHello: []byte{byte(innerECHExt)}, + } + outer.original = mustMarshal(t, outer) + + encodeInner := func(buildExts func(*cryptobyte.Builder)) []byte { + var b cryptobyte.Builder + b.AddUint16(VersionTLS12) + b.AddBytes(make([]byte, 32)) + b.AddUint8(0) + b.AddUint16LengthPrefixed(func(b *cryptobyte.Builder) { + b.AddUint16(TLS_AES_128_GCM_SHA256) + }) + b.AddUint8LengthPrefixed(func(b *cryptobyte.Builder) { + b.AddUint8(compressionNone) + }) + b.AddUint16LengthPrefixed(func(b *cryptobyte.Builder) { + buildExts(b) + b.AddUint16(extensionEncryptedClientHello) + b.AddUint16LengthPrefixed(func(b *cryptobyte.Builder) { + b.AddUint8(uint8(innerECHExt)) + }) + b.AddUint16(extensionSupportedVersions) + b.AddUint16LengthPrefixed(func(b *cryptobyte.Builder) { + b.AddUint8LengthPrefixed(func(b *cryptobyte.Builder) { + b.AddUint16(VersionTLS13) + }) + }) + }) + return b.BytesOrPanic() + } + + outerExts := func(b *cryptobyte.Builder, extTypes ...uint16) { + b.AddUint16(extensionECHOuterExtensions) + b.AddUint16LengthPrefixed(func(b *cryptobyte.Builder) { + b.AddUint8LengthPrefixed(func(b *cryptobyte.Builder) { + for _, extType := range extTypes { + b.AddUint16(extType) + } + }) + }) + } + + for _, tc := range []struct { + name string + buildExts func(*cryptobyte.Builder) + wantErr string + }{ + { + name: "valid order", + buildExts: func(b *cryptobyte.Builder) { + outerExts(b, extensionStatusRequest, extensionSupportedCurves) + }, + }, + { + name: "duplicate reference", + buildExts: func(b *cryptobyte.Builder) { + outerExts(b, extensionStatusRequest, extensionStatusRequest) + }, + wantErr: "tls: invalid outer extensions", + }, + { + name: "references encrypted_client_hello", + buildExts: func(b *cryptobyte.Builder) { + outerExts(b, extensionEncryptedClientHello) + }, + wantErr: "tls: invalid outer extensions", + }, + { + name: "references ech_outer_extensions", + buildExts: func(b *cryptobyte.Builder) { + outerExts(b, extensionECHOuterExtensions) + }, + wantErr: "tls: invalid outer extensions", + }, + { + name: "multiple ech_outer_extensions", + buildExts: func(b *cryptobyte.Builder) { + outerExts(b, extensionStatusRequest) + outerExts(b, extensionSupportedCurves) + }, + wantErr: "tls: invalid outer extensions", + }, + { + name: "odd-length reference list", + buildExts: func(b *cryptobyte.Builder) { + b.AddUint16(extensionECHOuterExtensions) + b.AddUint16LengthPrefixed(func(b *cryptobyte.Builder) { + b.AddUint8LengthPrefixed(func(b *cryptobyte.Builder) { + b.AddUint8(0) + }) + }) + }, + wantErr: "tls: invalid inner client hello", + }, + { + name: "empty reference list", + buildExts: func(b *cryptobyte.Builder) { + b.AddUint16(extensionECHOuterExtensions) + b.AddUint16LengthPrefixed(func(b *cryptobyte.Builder) { + b.AddUint8(0) + }) + }, + wantErr: "tls: invalid inner client hello", + }, + } { + t.Run(tc.name, func(t *testing.T) { + encoded := encodeInner(tc.buildExts) + _, err := decodeInnerClientHello(outer, encoded) + if tc.wantErr == "" { + if err != nil { + t.Fatalf("decodeInnerClientHello returned %v, want nil", err) + } + } else if err == nil || err.Error() != tc.wantErr { + t.Fatalf("decodeInnerClientHello returned %v, want %q", err, tc.wantErr) + } + }) + } +} diff --git a/src/encoding/json/decode_test.go b/src/encoding/json/decode_test.go index c8d078a72d3a6b..066f4b1bbf3ce5 100644 --- a/src/encoding/json/decode_test.go +++ b/src/encoding/json/decode_test.go @@ -436,6 +436,7 @@ var unmarshalTests = []struct { {CaseName: Name(""), in: `true`, ptr: new(bool), out: true}, {CaseName: Name(""), in: `1`, ptr: new(int), out: 1}, {CaseName: Name(""), in: `1.2`, ptr: new(float64), out: 1.2}, + {CaseName: Name(""), in: `1e1000`, ptr: new(float64), out: float64(0), err: &UnmarshalTypeError{Value: "number 1e1000", Type: reflect.TypeFor[float64](), Offset: len64("1e1000")}}, {CaseName: Name(""), in: `-5`, ptr: new(int16), out: int16(-5)}, {CaseName: Name(""), in: `2`, ptr: new(Number), out: Number("2"), useNumber: true}, {CaseName: Name(""), in: `2`, ptr: new(Number), out: Number("2")}, @@ -2496,9 +2497,10 @@ func TestPrefilled(t *testing.T) { // Values here change, cannot reuse table across runs. tests := []struct { CaseName - in string - ptr any - out any + in string + ptr any + out any + wantErr error }{{ CaseName: Name(""), in: `{"X": 1, "Y": 2}`, @@ -2529,17 +2531,29 @@ func TestPrefilled(t *testing.T) { in: `[3]`, ptr: &[...]int{1, 2}, out: &[...]int{3, 0}, + }, { + CaseName: Name(""), + in: `1e1000`, + ptr: addr(float64(math.Pi)), + out: addr(float64(math.Pi)), + wantErr: &UnmarshalTypeError{Value: "number 1e1000", Type: reflect.TypeFor[float64](), Offset: len64("1e1000")}, + }, { + CaseName: Name(""), + in: `1e1000`, + ptr: addr(any(float64(0))), + out: addr(any(float64(0))), + wantErr: &UnmarshalTypeError{Value: "number 1e1000", Type: reflect.TypeFor[float64](), Offset: len64("1e1000") + 1}, }} for _, tt := range tests { t.Run(tt.Name, func(t *testing.T) { ptrstr := fmt.Sprintf("%v", tt.ptr) err := Unmarshal([]byte(tt.in), tt.ptr) // tt.ptr edited here - if err != nil { - t.Errorf("%s: Unmarshal error: %v", tt.Where, err) - } if !reflect.DeepEqual(tt.ptr, tt.out) { t.Errorf("%s: Unmarshal(%#q, %T):\n\tgot: %v\n\twant: %v", tt.Where, tt.in, ptrstr, tt.ptr, tt.out) } + if !reflect.DeepEqual(err, tt.wantErr) { + t.Errorf("%s: Unmarshal(%#q, %T) error:\n\tgot: %v\n\twant: %v", tt.Where, tt.in, ptrstr, err, tt.wantErr) + } }) } } diff --git a/src/encoding/json/encode_test.go b/src/encoding/json/encode_test.go index 87074eabd46399..f81cdd9582fcdd 100644 --- a/src/encoding/json/encode_test.go +++ b/src/encoding/json/encode_test.go @@ -429,6 +429,43 @@ func TestUnsupportedValues(t *testing.T) { } } +// Issue 81176: UnsupportedValueError.Value should hold the NaN or ±Inf value. +func TestUnsupportedValueErrorValue(t *testing.T) { + type NamedFloat float64 + tests := []struct { + CaseName + in any + want any + }{ + {Name(""), NamedFloat(math.NaN()), NamedFloat(math.NaN())}, + {Name(""), math.Inf(-1), math.Inf(-1)}, + {Name(""), NamedFloat(math.Inf(1)), NamedFloat(math.Inf(1))}, + {Name(""), map[string]float64{"x": math.Inf(1)}, math.Inf(1)}, + {Name(""), []NamedFloat{NamedFloat(math.NaN())}, NamedFloat(math.NaN())}, + {Name(""), struct{ F float64 }{math.Inf(-1)}, math.Inf(-1)}, + } + for _, tt := range tests { + t.Run(tt.Name, func(t *testing.T) { + _, err := Marshal(tt.in) + uve, ok := err.(*UnsupportedValueError) + if !ok { + t.Fatalf("%s: Marshal error:\n\tgot: %T\n\twant: %T", tt.Where, err, new(UnsupportedValueError)) + } + got := uve.Value + want := reflect.ValueOf(tt.want) + if got.Type() != want.Type() { + t.Fatalf("%s: UnsupportedValueError.Value.Type = %v, want %v", tt.Where, got.Type(), want.Type()) + } + equalFloat := func(x, y float64) bool { + return x == y || math.IsNaN(x) == math.IsNaN(y) + } + if !equalFloat(got.Float(), want.Float()) { + t.Fatalf("%s: UnsupportedValueError.Value.Float = %v, want %v", tt.Where, got.Float(), want.Float()) + } + }) + } +} + // Issue 43207 func TestMarshalTextFloatMap(t *testing.T) { m := map[textfloat]string{ @@ -1111,6 +1148,30 @@ func TestNilMarshalerTextMapKey(t *testing.T) { } } +// textMarshalerString is a string kind that implements encoding.TextMarshaler. +type textMarshalerString string + +func (s textMarshalerString) MarshalText() ([]byte, error) { + return []byte("X_" + string(s)), nil +} + +func (s textMarshalerString) AppendText(b []byte) ([]byte, error) { + return append(b, ("X_" + string(s))...), nil +} + +// Issue 81355: string-kind map keys are used directly even if the key type +// implements encoding.TextMarshaler. MarshalText is still called for values. +func TestStringKindTextMarshalerMapKey(t *testing.T) { + got, err := Marshal(map[textMarshalerString]textMarshalerString{"foo": "bar"}) + if err != nil { + t.Fatalf("Marshal error: %v", err) + } + const want = `{"foo":"X_bar"}` + if string(got) != want { + t.Errorf("Marshal:\n\tgot: %s\n\twant: %s", got, want) + } +} + var re = regexp.MustCompile // syntactic checks on form of marshaled floating point numbers. diff --git a/src/encoding/json/internal/internal.go b/src/encoding/json/internal/internal.go index 456676ea2529aa..1f57724b390023 100644 --- a/src/encoding/json/internal/internal.go +++ b/src/encoding/json/internal/internal.go @@ -24,6 +24,16 @@ var ( ErrNilInterface = errors.New("cannot derive concrete type for nil interface with finite type set") ) +// ValueError wraps the error with a Go value relevant to the error. +// This is only used for v1 compatibility purposes. +type ValueError struct { + Val any + Err error +} + +func (ve *ValueError) Error() string { return ve.Err.Error() } +func (ve *ValueError) Unwrap() error { return ve.Err } + var ( // TransformMarshalError converts a v2 error into a v1 error. // It is called only at the top-level of a Marshal function. diff --git a/src/encoding/json/v2/arshal_any.go b/src/encoding/json/v2/arshal_any.go index 63ac6c44d83479..2367502b6d47c8 100644 --- a/src/encoding/json/v2/arshal_any.go +++ b/src/encoding/json/v2/arshal_any.go @@ -93,7 +93,7 @@ func unmarshalValueAny(dec *jsontext.Decoder, uo *jsonopts.Struct) (any, error) } fv, err := strconv.ParseFloat(string(val), 64) if err != nil { - return fv, newUnmarshalErrorAfterWithValue(dec, float64Type, errors.Unwrap(err)) + return 0.0, newUnmarshalErrorAfterWithValue(dec, float64Type, errors.Unwrap(err)) } return fv, nil default: diff --git a/src/encoding/json/v2/arshal_default.go b/src/encoding/json/v2/arshal_default.go index 77afaa40d89b11..9ba764aa750288 100644 --- a/src/encoding/json/v2/arshal_default.go +++ b/src/encoding/json/v2/arshal_default.go @@ -672,6 +672,9 @@ func makeFloatArshaler(t reflect.Type) *arshaler { if math.IsNaN(fv) || math.IsInf(fv, 0) { if !allowNonFinite { err := fmt.Errorf("unsupported value: %v", fv) + if mo.Flags.Get(jsonflags.ReportErrorsWithLegacySemantics) { + err = &internal.ValueError{Val: va.Interface(), Err: err} + } return newMarshalErrorBefore(enc, t, err) } return enc.WriteToken(jsontext.Float(fv)) @@ -760,10 +763,10 @@ func makeFloatArshaler(t reflect.Type) *arshaler { break } fv, err := strconv.ParseFloat(string(val), bits) - va.SetFloat(fv) if err != nil { return newUnmarshalErrorAfterWithValue(dec, t, errors.Unwrap(err)) } + va.SetFloat(fv) return nil } return newUnmarshalErrorAfter(dec, t, nil) diff --git a/src/encoding/json/v2/arshal_methods.go b/src/encoding/json/v2/arshal_methods.go index 337a56aa708e0c..14fe953359877f 100644 --- a/src/encoding/json/v2/arshal_methods.go +++ b/src/encoding/json/v2/arshal_methods.go @@ -178,7 +178,9 @@ func makeMethodArshaler(fncs *arshaler, t reflect.Type) *arshaler { prevMarshal := fncs.marshal fncs.marshal = func(enc *jsontext.Encoder, va addressableValue, mo *jsonopts.Struct) error { if mo.Flags.Get(jsonflags.CallMethodsWithLegacySemantics) && - (needAddr && va.forcedAddr) { + ((needAddr && va.forcedAddr) || + (export.Encoder(enc).Tokens.Last.NeedObjectName()) && t.Kind() == reflect.String) { + // Do not call MarshalText on unaddressable values and map keys of string kind. return prevMarshal(enc, va, mo) } marshaler, _ := reflect.TypeAssert[encoding.TextMarshaler](va.Addr()) @@ -204,7 +206,9 @@ func makeMethodArshaler(fncs *arshaler, t reflect.Type) *arshaler { prevMarshal := fncs.marshal fncs.marshal = func(enc *jsontext.Encoder, va addressableValue, mo *jsonopts.Struct) (err error) { if mo.Flags.Get(jsonflags.CallMethodsWithLegacySemantics) && - (needAddr && va.forcedAddr) { + ((needAddr && va.forcedAddr) || + (export.Encoder(enc).Tokens.Last.NeedObjectName()) && t.Kind() == reflect.String) { + // Do not call AppendText on unaddressable values and map keys of string kind. return prevMarshal(enc, va, mo) } appender, _ := reflect.TypeAssert[encoding.TextAppender](va.Addr()) @@ -228,6 +232,7 @@ func makeMethodArshaler(fncs *arshaler, t reflect.Type) *arshaler { fncs.marshal = func(enc *jsontext.Encoder, va addressableValue, mo *jsonopts.Struct) error { if mo.Flags.Get(jsonflags.CallMethodsWithLegacySemantics) && ((needAddr && va.forcedAddr) || export.Encoder(enc).Tokens.Last.NeedObjectName()) { + // Do not call MarshalJSON on unaddressable values and map keys. return prevMarshal(enc, va, mo) } marshaler, _ := reflect.TypeAssert[Marshaler](va.Addr()) @@ -259,6 +264,7 @@ func makeMethodArshaler(fncs *arshaler, t reflect.Type) *arshaler { fncs.marshal = func(enc *jsontext.Encoder, va addressableValue, mo *jsonopts.Struct) error { if mo.Flags.Get(jsonflags.CallMethodsWithLegacySemantics) && ((needAddr && va.forcedAddr) || export.Encoder(enc).Tokens.Last.NeedObjectName()) { + // Do not call MarshalJSONTo on unaddressable values and map keys. return prevMarshal(enc, va, mo) } xe := export.Encoder(enc) @@ -330,6 +336,7 @@ func makeMethodArshaler(fncs *arshaler, t reflect.Type) *arshaler { fncs.unmarshal = func(dec *jsontext.Decoder, va addressableValue, uo *jsonopts.Struct) error { if uo.Flags.Get(jsonflags.CallMethodsWithLegacySemantics) && export.Decoder(dec).Tokens.Last.NeedObjectName() { + // Do not call UnmarshalJSON on map keys. return prevUnmarshal(dec, va, uo) } val, err := dec.ReadValue() @@ -355,6 +362,7 @@ func makeMethodArshaler(fncs *arshaler, t reflect.Type) *arshaler { fncs.unmarshal = func(dec *jsontext.Decoder, va addressableValue, uo *jsonopts.Struct) error { if uo.Flags.Get(jsonflags.CallMethodsWithLegacySemantics) && export.Decoder(dec).Tokens.Last.NeedObjectName() { + // Do not call UnmarshalJSONFrom on map keys. return prevUnmarshal(dec, va, uo) } xd := export.Decoder(dec) diff --git a/src/encoding/json/v2/arshal_test.go b/src/encoding/json/v2/arshal_test.go index 554ee654786363..d90184db547e82 100644 --- a/src/encoding/json/v2/arshal_test.go +++ b/src/encoding/json/v2/arshal_test.go @@ -5541,7 +5541,7 @@ func TestUnmarshal(t *testing.T) { name: jsontest.Name("Floats/Float32/Overflow"), inBuf: `-1e1000`, inVal: addr(float32(32.32)), - want: addr(float32(math.Inf(-1))), + want: addr(float32(32.32)), wantErr: EU(strconv.ErrRange).withVal(`-1e1000`).withType('0', T[float32]()), }, { name: jsontest.Name("Floats/Float64/Pi"), @@ -5557,13 +5557,25 @@ func TestUnmarshal(t *testing.T) { name: jsontest.Name("Floats/Float64/Overflow"), inBuf: `-1e1000`, inVal: addr(float64(64.64)), - want: addr(float64(math.Inf(-1))), + want: addr(float64(64.64)), wantErr: EU(strconv.ErrRange).withVal(`-1e1000`).withType('0', T[float64]()), }, { name: jsontest.Name("Floats/Any/Overflow"), inBuf: `1e1000`, inVal: new(any), - want: addr(any(float64(math.Inf(+1)))), + want: addr(any(float64(0))), + wantErr: EU(strconv.ErrRange).withVal(`1e1000`).withType('0', T[float64]()), + }, { + name: jsontest.Name("Floats/Any/ExistingFloat32/Overflow"), + inBuf: `1e1000`, + inVal: addr(any(float32(32.32))), + want: addr(any(float32(32.32))), + wantErr: EU(strconv.ErrRange).withVal(`1e1000`).withType('0', T[float32]()), + }, { + name: jsontest.Name("Floats/Any/ExistingFloat64/Overflow"), + inBuf: `1e1000`, + inVal: addr(any(float64(64.64))), + want: addr(any(float64(64.64))), wantErr: EU(strconv.ErrRange).withVal(`1e1000`).withType('0', T[float64]()), }, { name: jsontest.Name("Floats/Named"), diff --git a/src/encoding/json/v2_decode_test.go b/src/encoding/json/v2_decode_test.go index a904b3c1718b3d..6218440d7d2939 100644 --- a/src/encoding/json/v2_decode_test.go +++ b/src/encoding/json/v2_decode_test.go @@ -436,6 +436,7 @@ var unmarshalTests = []struct { {CaseName: Name(""), in: `true`, ptr: new(bool), out: true}, {CaseName: Name(""), in: `1`, ptr: new(int), out: 1}, {CaseName: Name(""), in: `1.2`, ptr: new(float64), out: 1.2}, + {CaseName: Name(""), in: `1e1000`, ptr: new(float64), out: float64(0), err: &UnmarshalTypeError{Value: "number 1e1000", Type: reflect.TypeFor[float64](), Offset: len64("1e1000")}}, {CaseName: Name(""), in: `-5`, ptr: new(int16), out: int16(-5)}, {CaseName: Name(""), in: `2`, ptr: new(Number), out: Number("2"), useNumber: true}, {CaseName: Name(""), in: `2`, ptr: new(Number), out: Number("2")}, @@ -2525,9 +2526,10 @@ func TestPrefilled(t *testing.T) { // Values here change, cannot reuse table across runs. tests := []struct { CaseName - in string - ptr any - out any + in string + ptr any + out any + wantErr error }{{ CaseName: Name(""), in: `{"X": 1, "Y": 2}`, @@ -2558,17 +2560,29 @@ func TestPrefilled(t *testing.T) { in: `[3]`, ptr: &[...]int{1, 2}, out: &[...]int{3, 0}, + }, { + CaseName: Name(""), + in: `1e1000`, + ptr: addr(float64(math.Pi)), + out: addr(float64(math.Pi)), + wantErr: &UnmarshalTypeError{Value: "number 1e1000", Type: reflect.TypeFor[float64](), Offset: len64("1e1000")}, + }, { + CaseName: Name(""), + in: `1e1000`, + ptr: addr(any(float64(0))), + out: addr(any(float64(0))), + wantErr: &UnmarshalTypeError{Value: "number 1e1000", Type: reflect.TypeFor[float64](), Offset: len64("1e1000")}, }} for _, tt := range tests { t.Run(tt.Name, func(t *testing.T) { ptrstr := fmt.Sprintf("%v", tt.ptr) err := Unmarshal([]byte(tt.in), tt.ptr) // tt.ptr edited here - if err != nil { - t.Errorf("%s: Unmarshal error: %v", tt.Where, err) - } if !reflect.DeepEqual(tt.ptr, tt.out) { t.Errorf("%s: Unmarshal(%#q, %T):\n\tgot: %v\n\twant: %v", tt.Where, tt.in, ptrstr, tt.ptr, tt.out) } + if !reflect.DeepEqual(err, tt.wantErr) { + t.Errorf("%s: Unmarshal(%#q, %T) error:\n\tgot: %v\n\twant: %v", tt.Where, tt.in, ptrstr, err, tt.wantErr) + } }) } } diff --git a/src/encoding/json/v2_encode_test.go b/src/encoding/json/v2_encode_test.go index 11c521864916e0..0e938bceefd703 100644 --- a/src/encoding/json/v2_encode_test.go +++ b/src/encoding/json/v2_encode_test.go @@ -434,6 +434,43 @@ func TestUnsupportedValues(t *testing.T) { } } +// Issue 81176: UnsupportedValueError.Value should hold the NaN or ±Inf value. +func TestUnsupportedValueErrorValue(t *testing.T) { + type NamedFloat float64 + tests := []struct { + CaseName + in any + want any + }{ + {Name(""), NamedFloat(math.NaN()), NamedFloat(math.NaN())}, + {Name(""), math.Inf(-1), math.Inf(-1)}, + {Name(""), NamedFloat(math.Inf(1)), NamedFloat(math.Inf(1))}, + {Name(""), map[string]float64{"x": math.Inf(1)}, math.Inf(1)}, + {Name(""), []NamedFloat{NamedFloat(math.NaN())}, NamedFloat(math.NaN())}, + {Name(""), struct{ F float64 }{math.Inf(-1)}, math.Inf(-1)}, + } + for _, tt := range tests { + t.Run(tt.Name, func(t *testing.T) { + _, err := Marshal(tt.in) + uve, ok := err.(*UnsupportedValueError) + if !ok { + t.Fatalf("%s: Marshal error:\n\tgot: %T\n\twant: %T", tt.Where, err, new(UnsupportedValueError)) + } + got := uve.Value + want := reflect.ValueOf(tt.want) + if got.Type() != want.Type() { + t.Fatalf("%s: UnsupportedValueError.Value.Type = %v, want %v", tt.Where, got.Type(), want.Type()) + } + equalFloat := func(x, y float64) bool { + return x == y || math.IsNaN(x) == math.IsNaN(y) + } + if !equalFloat(got.Float(), want.Float()) { + t.Fatalf("%s: UnsupportedValueError.Value.Float = %v, want %v", tt.Where, got.Float(), want.Float()) + } + }) + } +} + // Issue 43207 func TestMarshalTextFloatMap(t *testing.T) { m := map[textfloat]string{ @@ -1116,6 +1153,30 @@ func TestNilMarshalerTextMapKey(t *testing.T) { } } +// textMarshalerString is a string kind that implements encoding.TextMarshaler. +type textMarshalerString string + +func (s textMarshalerString) MarshalText() ([]byte, error) { + return []byte("X_" + string(s)), nil +} + +func (s textMarshalerString) AppendText(b []byte) ([]byte, error) { + return append(b, ("X_" + string(s))...), nil +} + +// Issue 81355: string-kind map keys are used directly even if the key type +// implements encoding.TextMarshaler. MarshalText is still called for values. +func TestStringKindTextMarshalerMapKey(t *testing.T) { + got, err := Marshal(map[textMarshalerString]textMarshalerString{"foo": "bar"}) + if err != nil { + t.Fatalf("Marshal error: %v", err) + } + const want = `{"foo":"X_bar"}` + if string(got) != want { + t.Errorf("Marshal:\n\tgot: %s\n\twant: %s", got, want) + } +} + var re = regexp.MustCompile // syntactic checks on form of marshaled floating point numbers. diff --git a/src/encoding/json/v2_inject.go b/src/encoding/json/v2_inject.go index 3dd5e080a77cd5..8e0ed8c97bc27d 100644 --- a/src/encoding/json/v2_inject.go +++ b/src/encoding/json/v2_inject.go @@ -7,6 +7,7 @@ package json import ( + "errors" "fmt" "reflect" "strconv" @@ -44,14 +45,16 @@ func transformMarshalError(root any, err error) error { } else { // Historically, this was only reported for NaN or ±Inf values // and cycles detected in the value. - // The Value field used to be populated with the reflect.Value, - // but this is no longer supported. + var v reflect.Value + if err, ok := errors.AsType[*internal.ValueError](err.Err); ok { + v = reflect.ValueOf(err.Val) + } errStr := err.Err.Error() if err.Err == internal.ErrCycle && err.GoType != nil { errStr += " via " + err.GoType.String() } errStr = strings.TrimPrefix(errStr, "unsupported value: ") - return &UnsupportedValueError{Str: errStr} + return &UnsupportedValueError{Value: v, Str: errStr} } } else if ok { return (*UnsupportedValueError)(nil) diff --git a/src/go/internal/gcimporter/ureader.go b/src/go/internal/gcimporter/ureader.go index 5b3051c3ceea94..1e9f5c03cc77c9 100644 --- a/src/go/internal/gcimporter/ureader.go +++ b/src/go/internal/gcimporter/ureader.go @@ -5,6 +5,7 @@ package gcimporter import ( + "cmp" "go/token" "go/types" "internal/pkgbits" @@ -561,11 +562,26 @@ func (pr *pkgReader) objIdx(idx pkgbits.Index) (*types.Package, string) { named.SetUnderlying(underlying) - for i, n := 0, r.Len(); i < n; i++ { - named.AddMethod(r.method()) - } - if r.Version().Has(pkgbits.GenericMethods) { + // V4 (go1.27.0) emitted all non-generic methods + // before all generic ones, discarding source + // order: a bug (go.dev/issue/81188). + // V5 (go1.27.x) fixes it by emitting an explicit + // index along with each method. + type indexedMethod struct { + index int // (or -1 in V4) + fn *types.Func + } + + var methods []indexedMethod + + // ordinary methods + for range r.Len() { + idx, m := r.method() + methods = append(methods, indexedMethod{idx, m}) + } + + // generic methods for range r.Len() { // Careful: objIdx is used to read in package-scoped declarations, which // methods are not. Instead, decode it here. This makes it easier to @@ -580,11 +596,31 @@ func (pr *pkgReader) objIdx(idx pkgbits.Index) (*types.Package, string) { pkg, name := r.selector() rtparams := r.typeParamNames(true) recv := r.param(types.RecvVar) + methodIdx := -1 + if r.Version().Has(pkgbits.PreserveMethodOrder) { + methodIdx = r.Len() + } tparams := r.typeParamNames(false) sig := r.signature(recv, rtparams, tparams) pr.retireReader(r) - named.AddMethod(types.NewFunc(pos, pkg, name, sig)) + methods = append(methods, indexedMethod{methodIdx, types.NewFunc(pos, pkg, name, sig)}) + } + + if r.Version().Has(pkgbits.PreserveMethodOrder) { + slices.SortFunc(methods, func(a, b indexedMethod) int { + return cmp.Compare(a.index, b.index) + }) + } + + for _, m := range methods { + named.AddMethod(m.fn) + } + + } else { + for range r.Len() { + _, m := r.method() + named.AddMethod(m) } } @@ -696,8 +732,12 @@ func (r *reader) typeParamNames(isGenMeth bool) []*types.TypeParam { return tparams } -func (r *reader) method() *types.Func { +func (r *reader) method() (int, *types.Func) { r.Sync(pkgbits.SyncMethod) + idx := -1 + if r.Version().Has(pkgbits.PreserveMethodOrder) { + idx = r.Len() + } pos := r.pos() pkg, name := r.selector() @@ -705,7 +745,7 @@ func (r *reader) method() *types.Func { sig := r.signature(r.param(types.RecvVar), rparams, nil) _ = r.pos() // TODO(mdempsky): Remove; this is a hacker for linker.go. - return types.NewFunc(pos, pkg, name, sig) + return idx, types.NewFunc(pos, pkg, name, sig) } func (r *reader) qualifiedIdent() (*types.Package, string) { return r.ident(pkgbits.SyncSym) } diff --git a/src/go/types/named_test.go b/src/go/types/named_test.go index effeeb6f5c4299..34584d47caaa43 100644 --- a/src/go/types/named_test.go +++ b/src/go/types/named_test.go @@ -138,7 +138,7 @@ package p type T struct{} func (T) a() {} -func (T) c() {} +func (T) c[X any](x X) {} func (T) b() {} ` // should get the same method order each time diff --git a/src/html/template/escape_test.go b/src/html/template/escape_test.go index 3b26f7815fee37..d00a587d215c72 100644 --- a/src/html/template/escape_test.go +++ b/src/html/template/escape_test.go @@ -1884,6 +1884,14 @@ func TestEscapeText(t *testing.T) { "")) + var buf bytes.Buffer + if err := tmpl.Execute(&buf, `x/.exec(alert(1))}`); err != nil { + t.Fatalf("Execute: %v", err) + } + want := "" + if got := buf.String(); got != want { + t.Errorf("got: %s\nwant: %s", got, want) + } +} + func TestCVE202656858(t *testing.T) { tests := []struct { name string @@ -2327,3 +2347,49 @@ func TestCVE202656858(t *testing.T) { }) } } + +func TestIssue81823(t *testing.T) { + tests := []struct { + name string + tmpl string + input string + want string + }{ + { + name: "yield", + tmpl: ``, + input: `/;alert(1)//`, + want: ``, + }, + { + name: "property yield", + tmpl: "", + input: `1;pwned=1;0`, + want: "", + }, + { + name: "private yield", + tmpl: "", + input: `1;pwned=1;0`, + want: "", + }, + { + name: "property in", + tmpl: "", + input: `1;pwned=1;0`, + want: "", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + tmpl := Must(New("test").Parse(tt.tmpl)) + var buf strings.Builder + if err := tmpl.Execute(&buf, tt.input); err != nil { + t.Fatalf("Execute: %v", err) + } + if got := buf.String(); got != tt.want { + t.Errorf("got: %s\nwant: %s", got, tt.want) + } + }) + } +} diff --git a/src/html/template/js.go b/src/html/template/js.go index e2db30f966ed7d..354414f5dadc89 100644 --- a/src/html/template/js.go +++ b/src/html/template/js.go @@ -96,7 +96,10 @@ func nextJSCtx(s []byte, preceding jsCtx) jsCtx { for j > 0 && isJSIdentPart(rune(s[j-1])) { j-- } - if regexpPrecederKeywords[string(s[j:])] { + // An IdentifierName after a property-access dot is a property name, + // which precedes a div op. + if regexpPrecederKeywords[string(s[j:])] && + !bytes.HasSuffix(bytes.TrimRight(s[:j], jsWhitespace), []byte(".")) { return jsCtxRegexp } } @@ -106,8 +109,9 @@ func nextJSCtx(s []byte, preceding jsCtx) jsCtx { return jsCtxDivOp } -// regexpPrecederKeywords is a set of reserved JS keywords that can precede a -// regular expression in JS source. +// regexpPrecederKeywords is a set of JS keywords that can precede a regular +// expression in JS source. It deliberately treats the context-sensitive +// keyword yield as a keyword. var regexpPrecederKeywords = map[string]bool{ "break": true, "case": true, @@ -123,6 +127,7 @@ var regexpPrecederKeywords = map[string]bool{ "try": true, "typeof": true, "void": true, + "yield": true, } var jsonMarshalType = reflect.TypeFor[json.Marshaler]() diff --git a/src/html/template/js_test.go b/src/html/template/js_test.go index 015d97e6b50119..fbcfeebda05c7a 100644 --- a/src/html/template/js_test.go +++ b/src/html/template/js_test.go @@ -65,6 +65,12 @@ func TestNextJsCtx(t *testing.T) { {jsCtxRegexp, "return\t"}, {jsCtxRegexp, "return\n"}, {jsCtxRegexp, "return\u2028"}, + {jsCtxRegexp, "yield"}, + // A keyword after property access is a property name. + {jsCtxDivOp, "x.yield"}, + {jsCtxDivOp, "x?.yield"}, + {jsCtxDivOp, "x.\nyield"}, + {jsCtxDivOp, "x.in"}, // Identifiers can be divided and cannot validly be preceded by // a regular expressions. Semicolon insertion cannot happen // between an identifier and a regular expression on a new line diff --git a/src/html/template/transition.go b/src/html/template/transition.go index d9d4f63beba807..a49ae700b2c59a 100644 --- a/src/html/template/transition.go +++ b/src/html/template/transition.go @@ -331,7 +331,14 @@ func tJS(c context, s []byte) (context, int) { case '#': if i+1 < len(s) && s[i+1] == '!' { c.state, i = stateJSLineCmt, i+1 + break + } + // A private identifier such as #yield is never a keyword, and + // precedes a div op. + for i+1 < len(s) && isJSIdentPart(rune(s[i+1])) { + i++ } + c.jsCtx = jsCtxDivOp case '{': // We only care about tracking brace depth if we are inside of a // template literal. @@ -382,7 +389,7 @@ func tJSTmpl(c context, s []byte) (context, int) { case '$': if len(s) >= i+2 && s[i+1] == '{' { c.jsBraceDepth = append(c.jsBraceDepth, 0) - c.state = stateJS + c.state, c.jsCtx = stateJS, jsCtxRegexp return c, i + 2 } case '`': diff --git a/src/internal/godebugs/table.go b/src/internal/godebugs/table.go index a8c63cbc4388c8..abc6934b208056 100644 --- a/src/internal/godebugs/table.go +++ b/src/internal/godebugs/table.go @@ -47,6 +47,7 @@ var All = []Info{ {Name: "httplaxcontentlength", Package: "net/http", Changed: 22, Old: "1"}, {Name: "httpmuxgo121", Package: "net/http", Changed: 22, Old: "1"}, {Name: "httpservecontentkeepheaders", Package: "net/http", Changed: 23, Old: "1"}, + {Name: "httpservecontentmaxranges", Package: "net/http", Changed: 26, Old: "1"}, {Name: "installgoroot", Package: "go/build"}, {Name: "jstmpllitinterp", Package: "html/template", Opaque: true}, // bug #66217: remove Opaque //{Name: "multipartfiles", Package: "mime/multipart"}, diff --git a/src/internal/pkgbits/pkgbits_test.go b/src/internal/pkgbits/pkgbits_test.go index b6e421e7387600..8488fea00d263b 100644 --- a/src/internal/pkgbits/pkgbits_test.go +++ b/src/internal/pkgbits/pkgbits_test.go @@ -16,6 +16,8 @@ func TestRoundTrip(t *testing.T) { pkgbits.V1, pkgbits.V2, pkgbits.V3, + pkgbits.V4, + pkgbits.V5, } { pw := pkgbits.NewPkgEncoder(version, -1) w := pw.NewEncoder(pkgbits.SectionMeta, pkgbits.SyncPublic) @@ -40,6 +42,8 @@ var ( _ [1]bool = [pkgbits.V1]bool{} _ [2]bool = [pkgbits.V2]bool{} _ [3]bool = [pkgbits.V3]bool{} + _ [4]bool = [pkgbits.V4]bool{} + _ [5]bool = [pkgbits.V5]bool{} ) func TestVersions(t *testing.T) { @@ -60,6 +64,8 @@ func TestVersions(t *testing.T) { {pkgbits.V1, pkgbits.DerivedInfoNeeded}, {pkgbits.V2, pkgbits.AliasTypeParamNames}, {pkgbits.V3, pkgbits.CompactCompLiterals}, + {pkgbits.V4, pkgbits.GenericMethods}, + {pkgbits.V5, pkgbits.PreserveMethodOrder}, } { if !c.v.Has(c.f) { t.Errorf("Expected version %v to have field %v", c.v, c.f) @@ -77,6 +83,15 @@ func TestVersions(t *testing.T) { {pkgbits.V0, pkgbits.CompactCompLiterals}, {pkgbits.V1, pkgbits.CompactCompLiterals}, {pkgbits.V2, pkgbits.CompactCompLiterals}, + {pkgbits.V0, pkgbits.GenericMethods}, + {pkgbits.V1, pkgbits.GenericMethods}, + {pkgbits.V2, pkgbits.GenericMethods}, + {pkgbits.V3, pkgbits.GenericMethods}, + {pkgbits.V0, pkgbits.PreserveMethodOrder}, + {pkgbits.V1, pkgbits.PreserveMethodOrder}, + {pkgbits.V2, pkgbits.PreserveMethodOrder}, + {pkgbits.V3, pkgbits.PreserveMethodOrder}, + {pkgbits.V4, pkgbits.PreserveMethodOrder}, } { if c.v.Has(c.f) { t.Errorf("Expected version %v to not have field %v", c.v, c.f) diff --git a/src/internal/pkgbits/version.go b/src/internal/pkgbits/version.go index 6c424ee0c01de4..684b471753da61 100644 --- a/src/internal/pkgbits/version.go +++ b/src/internal/pkgbits/version.go @@ -1,4 +1,4 @@ -// Copyright 2021 The Go Authors. All rights reserved. +// Copyright 2024 The Go Authors. All rights reserved. // Use of this source code is governed by a BSD-style // license that can be found in the LICENSE file. @@ -37,6 +37,10 @@ const ( // V4: encodes generic methods as standalone function objects V4 + // V5: encodes the index of methods to preserve relative order + // of nongeneric and generic methods (go.dev/issue/81188). + V5 + numVersions = iota ) @@ -76,6 +80,10 @@ const ( // Generic methods may appear as standalone function objects. GenericMethods + // Method index is encoded to preserve relative order of + // nongeneric and generic methods. + PreserveMethodOrder + numFields = iota ) @@ -85,6 +93,7 @@ var introduced = [numFields]Version{ AliasTypeParamNames: V2, CompactCompLiterals: V3, GenericMethods: V4, + PreserveMethodOrder: V5, } // removed is the version a field was removed in or 0 for fields diff --git a/src/internal/poll/fd_windows.go b/src/internal/poll/fd_windows.go index 269bf4f9b628dd..85ee10e92ed1b5 100644 --- a/src/internal/poll/fd_windows.go +++ b/src/internal/poll/fd_windows.go @@ -377,6 +377,10 @@ type FD struct { // message based socket connection. ZeroReadIsEOF bool + // KeepFileCompletionModes prevents Init from changing the file object's + // completion notification modes. + KeepFileCompletionModes bool + // Whether the handle is owned by os.File. isFile bool @@ -462,6 +466,22 @@ func (fd *FD) Init(net string, pollable bool) error { // behavior below, as it requires an extra syscall. fd.waitOnSuccess = true + if fd.KeepFileCompletionModes { + // Query the existing skip-success mode so we don't wait for a + // suppressed completion or skip waiting for an expected one. + var info windows.FILE_IO_COMPLETION_NOTIFICATION_INFORMATION + if err := windows.NtQueryInformationFile(fd.Sysfd, &windows.IO_STATUS_BLOCK{}, + unsafe.Pointer(&info), uint32(unsafe.Sizeof(info)), windows.FileIoCompletionNotificationInformation); err != nil { + // Without knowing the modes, neither waiting for a completion + // packet on success nor skipping it is safe. Leave the handle + // unassociated and use explicit events for pending I/O instead. + // Inline success needs no wait, and deadlines are unavailable. + fd.waitOnSuccess = false + return nil + } + fd.waitOnSuccess = info.Flags&syscall.FILE_SKIP_COMPLETION_PORT_ON_SUCCESS == 0 + } + // It is safe to add overlapped handles that also perform I/O // outside of the runtime poller. The runtime poller will ignore // I/O completion notifications not initiated by us. @@ -471,16 +491,18 @@ func (fd *FD) Init(net string, pollable bool) error { } fd.associated = true - // FILE_SKIP_SET_EVENT_ON_HANDLE is always safe to use. We don't use that feature - // and it adds some overhead to the Windows I/O manager. - // See https://devblogs.microsoft.com/oldnewthing/20200221-00/?p=103466. - modes := uint8(syscall.FILE_SKIP_SET_EVENT_ON_HANDLE) - if canSkipCompletionPortOnSuccess(fd.Sysfd, fd.kind == kindNet) { - modes |= syscall.FILE_SKIP_COMPLETION_PORT_ON_SUCCESS - } - if syscall.SetFileCompletionNotificationModes(fd.Sysfd, modes) == nil { - if modes&syscall.FILE_SKIP_COMPLETION_PORT_ON_SUCCESS != 0 { - fd.waitOnSuccess = false + if !fd.KeepFileCompletionModes { + // FILE_SKIP_SET_EVENT_ON_HANDLE is always safe to use. We don't use that feature + // and it adds some overhead to the Windows I/O manager. + // See https://devblogs.microsoft.com/oldnewthing/20200221-00/?p=103466. + modes := uint8(syscall.FILE_SKIP_SET_EVENT_ON_HANDLE) + if canSkipCompletionPortOnSuccess(fd.Sysfd, fd.kind == kindNet) { + modes |= syscall.FILE_SKIP_COMPLETION_PORT_ON_SUCCESS + } + if syscall.SetFileCompletionNotificationModes(fd.Sysfd, modes) == nil { + if modes&syscall.FILE_SKIP_COMPLETION_PORT_ON_SUCCESS != 0 { + fd.waitOnSuccess = false + } } } return nil diff --git a/src/internal/runtime/syscall/linux/defs_linux_arm.go b/src/internal/runtime/syscall/linux/defs_linux_arm.go index cef556d5f6f986..94a5eee1712904 100644 --- a/src/internal/runtime/syscall/linux/defs_linux_arm.go +++ b/src/internal/runtime/syscall/linux/defs_linux_arm.go @@ -25,7 +25,7 @@ const ( ) type EpollEvent struct { - Events uint32 - _pad uint32 - Data [8]byte // to match amd64 + Events uint32 + pad_cgo_0 uint32 + Data uint64 } diff --git a/src/internal/runtime/syscall/linux/defs_linux_arm64.go b/src/internal/runtime/syscall/linux/defs_linux_arm64.go index eabddbac1bc063..b87eb357a34644 100644 --- a/src/internal/runtime/syscall/linux/defs_linux_arm64.go +++ b/src/internal/runtime/syscall/linux/defs_linux_arm64.go @@ -25,7 +25,7 @@ const ( ) type EpollEvent struct { - Events uint32 - _pad uint32 - Data [8]byte // to match amd64 + Events uint32 + pad_cgo_0 uint32 + Data uint64 } diff --git a/src/internal/runtime/syscall/linux/defs_linux_loong64.go b/src/internal/runtime/syscall/linux/defs_linux_loong64.go index 08e5d49b83c9bd..b87eb357a34644 100644 --- a/src/internal/runtime/syscall/linux/defs_linux_loong64.go +++ b/src/internal/runtime/syscall/linux/defs_linux_loong64.go @@ -26,6 +26,6 @@ const ( type EpollEvent struct { Events uint32 - pad_cgo_0 [4]byte - Data [8]byte // unaligned uintptr + pad_cgo_0 uint32 + Data uint64 } diff --git a/src/internal/runtime/syscall/linux/defs_linux_mips64x.go b/src/internal/runtime/syscall/linux/defs_linux_mips64x.go index b5794e5002af5e..195fc679f33a7b 100644 --- a/src/internal/runtime/syscall/linux/defs_linux_mips64x.go +++ b/src/internal/runtime/syscall/linux/defs_linux_mips64x.go @@ -28,6 +28,6 @@ const ( type EpollEvent struct { Events uint32 - pad_cgo_0 [4]byte - Data [8]byte // unaligned uintptr + pad_cgo_0 uint32 + Data uint64 } diff --git a/src/internal/runtime/syscall/linux/defs_linux_mipsx.go b/src/internal/runtime/syscall/linux/defs_linux_mipsx.go index 1fb4d919d1a318..824eb94829b657 100644 --- a/src/internal/runtime/syscall/linux/defs_linux_mipsx.go +++ b/src/internal/runtime/syscall/linux/defs_linux_mipsx.go @@ -28,6 +28,6 @@ const ( type EpollEvent struct { Events uint32 - pad_cgo_0 [4]byte + pad_cgo_0 uint32 Data uint64 } diff --git a/src/internal/runtime/syscall/linux/defs_linux_ppc64x.go b/src/internal/runtime/syscall/linux/defs_linux_ppc64x.go index ee93ad345b810f..63a71377ce7975 100644 --- a/src/internal/runtime/syscall/linux/defs_linux_ppc64x.go +++ b/src/internal/runtime/syscall/linux/defs_linux_ppc64x.go @@ -28,6 +28,6 @@ const ( type EpollEvent struct { Events uint32 - pad_cgo_0 [4]byte - Data [8]byte // unaligned uintptr + pad_cgo_0 uint32 + Data uint64 } diff --git a/src/internal/runtime/syscall/linux/defs_linux_riscv64.go b/src/internal/runtime/syscall/linux/defs_linux_riscv64.go index 08e5d49b83c9bd..b87eb357a34644 100644 --- a/src/internal/runtime/syscall/linux/defs_linux_riscv64.go +++ b/src/internal/runtime/syscall/linux/defs_linux_riscv64.go @@ -26,6 +26,6 @@ const ( type EpollEvent struct { Events uint32 - pad_cgo_0 [4]byte - Data [8]byte // unaligned uintptr + pad_cgo_0 uint32 + Data uint64 } diff --git a/src/internal/runtime/syscall/linux/defs_linux_s390x.go b/src/internal/runtime/syscall/linux/defs_linux_s390x.go index da11c704081abc..edea09e9619871 100644 --- a/src/internal/runtime/syscall/linux/defs_linux_s390x.go +++ b/src/internal/runtime/syscall/linux/defs_linux_s390x.go @@ -26,6 +26,6 @@ const ( type EpollEvent struct { Events uint32 - pad_cgo_0 [4]byte - Data [8]byte // unaligned uintptr + pad_cgo_0 uint32 + Data uint64 } diff --git a/src/internal/syscall/windows/at_windows.go b/src/internal/syscall/windows/at_windows.go index f2f95b871186cc..f7687159776795 100644 --- a/src/internal/syscall/windows/at_windows.go +++ b/src/internal/syscall/windows/at_windows.go @@ -6,6 +6,7 @@ package windows import ( "internal/oserror" + "internal/stringslite" "runtime" "structs" "syscall" @@ -141,19 +142,27 @@ func Openat(dirfd syscall.Handle, name string, flag uint64, perm uint32) (_ sysc } var h syscall.Handle - err := NtCreateFile( - &h, - SYNCHRONIZE|access, - objAttrs, - &IO_STATUS_BLOCK{}, - nil, - fileAttrs, - FILE_SHARE_READ|FILE_SHARE_WRITE|FILE_SHARE_DELETE, - disposition, - FILE_OPEN_FOR_BACKUP_INTENT|options, - nil, - 0, - ) + var err error + if TestOpenatFallback && flag&O_NOFOLLOW_ANY != 0 { + err = STATUS_INVALID_PARAMETER + } else { + err = NtCreateFile( + &h, + SYNCHRONIZE|access, + objAttrs, + &IO_STATUS_BLOCK{}, + nil, + fileAttrs, + FILE_SHARE_READ|FILE_SHARE_WRITE|FILE_SHARE_DELETE, + disposition, + FILE_OPEN_FOR_BACKUP_INTENT|options, + nil, + 0, + ) + } + if err == STATUS_INVALID_PARAMETER && flag&O_NOFOLLOW_ANY != 0 { + h, err = openatFallback(name, SYNCHRONIZE|access, *objAttrs, fileAttrs, disposition, FILE_OPEN_FOR_BACKUP_INTENT|options) + } if err != nil { return h, ntCreateFileError(err, flag) } @@ -177,11 +186,78 @@ func Openat(dirfd syscall.Handle, name string, flag uint64, perm uint32) (_ sysc return h, nil } +// TestOpenatFallback should only be used for testing purposes. +// When set, Openat simulates a system that does not support OBJ_DONT_REPARSE. +var TestOpenatFallback bool + +// openatFallback implements O_NOFOLLOW_ANY for a single path component on +// Windows versions that do not support OBJ_DONT_REPARSE (including Windows 10 +// build 10240). See go.dev/issue/78131. +func openatFallback(name string, access uint32, attrs OBJECT_ATTRIBUTES, fileAttrs, disposition, options uint32) (syscall.Handle, error) { + // FILE_OPEN_REPARSE_POINT only prevents following the final component. + // All current production callers that request O_NOFOLLOW_ANY pass a + // single component: os.Root splits paths (including symlink targets) + // before opening each component, and RemoveAll passes a basename or a + // directory entry name relative to an open parent directory. Thus this + // fallback also supports multi-component paths passed to those APIs. + // Reject other paths here rather than weaken O_NOFOLLOW_ANY: this is + // not a general replacement for OBJ_DONT_REPARSE. + if attrs.RootDirectory == 0 || name == ".." || + stringslite.IndexByte(name, '\\') >= 0 || stringslite.IndexByte(name, '/') >= 0 || stringslite.IndexByte(name, ':') >= 0 { + return syscall.InvalidHandle, STATUS_INVALID_PARAMETER + } + attrs.Attributes &^= OBJ_DONT_REPARSE + var h syscall.Handle + err := NtCreateFile( + &h, access, &attrs, &IO_STATUS_BLOCK{}, nil, fileAttrs, + FILE_SHARE_READ|FILE_SHARE_WRITE|FILE_SHARE_DELETE, disposition, + (options|FILE_OPEN_REPARSE_POINT)&^FILE_DELETE_ON_CLOSE, nil, 0, + ) + if err != nil { + return syscall.InvalidHandle, err + } + // Inspect the handle, not the path, before truncation or delete-on-close. + // Skip this check if opening the reparse point itself was requested, + // or if O_CREAT|O_EXCL requires creating a new file without following links. + if options&FILE_OPEN_REPARSE_POINT == 0 { + var info syscall.ByHandleFileInformation + err = syscall.GetFileInformationByHandle(h, &info) + if err == nil && info.FileAttributes&syscall.FILE_ATTRIBUTE_REPARSE_POINT != 0 { + err = STATUS_REPARSE_POINT_ENCOUNTERED + } + if err != nil { + syscall.CloseHandle(h) + return syscall.InvalidHandle, err + } + } + if options&FILE_DELETE_ON_CLOSE == 0 { + return h, nil + } + defer syscall.CloseHandle(h) + + // Only enable deletion after checking for a reparse point. An empty + // name relative to h reopens the same file, even if it has been renamed + // or its directory entry has been replaced since the first open. + // The file already exists, including when created with O_EXCL. + attrs.RootDirectory = h + attrs.ObjectName = &NTUnicodeString{} + var dh syscall.Handle + err = NtOpenFile( + &dh, access, &attrs, &IO_STATUS_BLOCK{}, + FILE_SHARE_READ|FILE_SHARE_WRITE|FILE_SHARE_DELETE, + options|FILE_OPEN_REPARSE_POINT, + ) + if err != nil { + return syscall.InvalidHandle, err + } + return dh, nil +} + // ntCreateFileError maps error returns from NTCreateFile to user-visible errors. func ntCreateFileError(err error, flag uint64) error { s, ok := err.(NTStatus) if !ok { - // Shouldn't really be possible, NtCreateFile always returns NTStatus. + // The Openat fallback can also return Win32 errors. return err } switch s { @@ -222,7 +298,7 @@ func Mkdirat(dirfd syscall.Handle, name string, mode uint32) error { syscall.FILE_ATTRIBUTE_NORMAL, syscall.FILE_SHARE_READ|syscall.FILE_SHARE_WRITE|syscall.FILE_SHARE_DELETE, FILE_CREATE, - FILE_DIRECTORY_FILE, + FILE_DIRECTORY_FILE|FILE_OPEN_REPARSE_POINT, nil, 0, ) diff --git a/src/internal/syscall/windows/types_windows.go b/src/internal/syscall/windows/types_windows.go index 55ab0a44b990f0..7ab9032ac15ea7 100644 --- a/src/internal/syscall/windows/types_windows.go +++ b/src/internal/syscall/windows/types_windows.go @@ -272,7 +272,14 @@ type FILE_LINK_INFORMATION struct { FileName [syscall.MAX_PATH]uint16 } -const FileReplaceCompletionInformation = 61 +const ( + FileIoCompletionNotificationInformation = 41 + FileReplaceCompletionInformation = 61 +) + +type FILE_IO_COMPLETION_NOTIFICATION_INFORMATION struct { + Flags uint32 +} // https://learn.microsoft.com/en-us/windows-hardware/drivers/ddi/ntifs/ns-ntifs-_file_completion_information type FILE_COMPLETION_INFORMATION struct { diff --git a/src/mime/multipart/formdata.go b/src/mime/multipart/formdata.go index d0e0151a6fbe4a..f60fbb99954f37 100644 --- a/src/mime/multipart/formdata.go +++ b/src/mime/multipart/formdata.go @@ -93,17 +93,16 @@ func (r *Reader) readForm(maxMemory int64) (_ *Form, err error) { // The relationship between these parameters, as well as the overly-large and // unconfigurable 10 MB added on to maxMemory, is unfortunate but difficult to change // within the constraints of the API as documented. + if maxMemory < 0 { + maxMemory = 0 + } maxFileMemoryBytes := maxMemory if maxFileMemoryBytes == math.MaxInt64 { maxFileMemoryBytes-- } maxMemoryBytes := maxMemory + int64(10<<20) if maxMemoryBytes <= 0 { - if maxMemory < 0 { - maxMemoryBytes = 0 - } else { - maxMemoryBytes = math.MaxInt64 - } + maxMemoryBytes = math.MaxInt64 } var copyBuf []byte for { diff --git a/src/net/http/fs.go b/src/net/http/fs.go index 240a8dc675ddc5..3efd9eb7c1f75b 100644 --- a/src/net/http/fs.go +++ b/src/net/http/fs.go @@ -12,6 +12,7 @@ import ( "internal/godebug" "io" "io/fs" + "math" "mime" "mime/multipart" "net/http/internal" @@ -1017,6 +1018,25 @@ func (r httpRange) mimeHeader(contentType string, size int64) textproto.MIMEHead } } +// GODEBUG=httpservecontentmaxranges= controls the maximum number of ranges that will +// be processed in a Range header. Setting httpservecontentmaxranges=0 disables the limit. +var httpservecontentmaxranges = godebug.New("httpservecontentmaxranges") + +const defaultMaxContentRanges = 200 + +func maxContentRanges() int { + maxRanges := defaultMaxContentRanges + if v := httpservecontentmaxranges.Value(); v != "" { + if n, err := strconv.Atoi(v); err == nil && n >= 0 { + maxRanges = n + if maxRanges == 0 { + maxRanges = math.MaxInt + } + } + } + return maxRanges +} + // parseRange parses a Range header string as per RFC 7233. // errNoOverlap is returned if none of the ranges overlap. func parseRange(s string, size int64) ([]httpRange, error) { @@ -1029,6 +1049,14 @@ func parseRange(s string, size int64) ([]httpRange, error) { } var ranges []httpRange noOverlap := false + numRanges := strings.Count(s[len(b):], ",") + 1 + maxRanges := maxContentRanges() + if (numRanges > maxRanges) != (numRanges > defaultMaxContentRanges) { + httpservecontentmaxranges.IncNonDefault() + } + if numRanges > maxRanges { + return nil, nil // ignore header with too many ranges + } for ra := range strings.SplitSeq(s[len(b):], ",") { ra = textproto.TrimString(ra) if ra == "" { diff --git a/src/net/http/http1_server_test.go b/src/net/http/http1_server_test.go index 17c31a877c08f0..d6ca39c997f4bd 100644 --- a/src/net/http/http1_server_test.go +++ b/src/net/http/http1_server_test.go @@ -6,14 +6,17 @@ package http_test import ( "bufio" + "bytes" "errors" "internal/nettest" - "internal/synctest" "io" "net/http" "net/http/httptest" + "slices" "strings" + "sync" "testing" + "testing/synctest" ) // An http1ServerTest tests an HTTP/1 server using a fake network. @@ -81,6 +84,18 @@ func (tc *http1TestConn) writeMessage(lines ...string) { } } +// readRequest reads a request from the connection (not including the request body). +func (tc *http1TestConn) readRequest() *http.Request { + t := tc.t + t.Helper() + synctest.Wait() + req, err := http.ReadRequest(tc.bufr) + if err != nil { + t.Fatalf("ReadRequest: %v", err) + } + return req +} + // readResponse reads a response from the connection (not including the response body). func (tc *http1TestConn) readResponse() *http.Response { t := tc.t @@ -93,6 +108,60 @@ func (tc *http1TestConn) readResponse() *http.Response { return resp } +func (tc *http1TestConn) wantResponse(wantStart string, wantHeaders http.Header) { + t := tc.t + t.Helper() + synctest.Wait() + gotStart, err := tc.bufr.ReadString('\n') + if err != nil { + t.Fatalf("read from conn: %q, %v; want start line %q", gotStart, err, wantStart) + } + if got, want := gotStart, wantStart+"\r\n"; got != want { + t.Fatalf("read start line:\n%q\nwant:\n%q", got, want) + } + gotHeaders := make(http.Header) + for { + line, err := tc.bufr.ReadString('\n') + if err != nil { + t.Fatalf("read from conn: %v (want header)", err) + } + line, ok := strings.CutSuffix(line, "\r\n") + if !ok { + t.Fatalf("header line has no CRLF suffix: %q", line) + } + if line == "" { + break + } + k, v, ok := strings.Cut(line, ": ") + if !ok { + t.Fatalf("invalid header line: %q", line) + } + gotHeaders[k] = append(gotHeaders[k], v) + } + for k, wantv := range wantHeaders { + gotv := gotHeaders[k] + if !slices.Equal(gotv, wantv) { + t.Errorf("header %v = %q, want %q", k, gotv, wantv) + } + } + if t.Failed() { + t.FailNow() + } +} + +// wantBytes asserts that the given bytes can be read from the connection. +func (tc *http1TestConn) wantBytes(want []byte) { + t := tc.t + t.Helper() + synctest.Wait() + got := make([]byte, len(want)) + n, err := io.ReadFull(tc.bufr, got) + got = got[:n] + if err != nil || !bytes.Equal(want, got) { + t.Fatalf("want bytes %q, got %q and error %v", want, got, err) + } +} + // wantIdle asserts that the connection is not closed and has no pending data to read. func (tc *http1TestConn) wantIdle() { t := tc.t @@ -112,3 +181,97 @@ func (tc *http1TestConn) wantClosed() { t.Fatalf("read from conn: %q; expect conn to be closed", got) } } + +type testHandler struct { + t *testing.T + mu sync.Mutex + calls []*testHandlerCall + closed bool +} + +func newTestHandler(t *testing.T) *testHandler { + h := &testHandler{t: t} + t.Cleanup(func() { + // testHandler.Close should be called before the server shuts down. + // Catch the case where we forgot to do this. + if !h.closed { + t.Errorf("testHandler.Close not called") + } + }) + return h +} + +func (h *testHandler) Close() { + h.t.Helper() + synctest.Wait() + h.mu.Lock() + defer h.mu.Unlock() + if len(h.calls) > 0 { + h.t.Errorf("test finished with %v handler calls unhandled", len(h.calls)) + } + for _, call := range h.calls { + call.exit() + } + h.calls = nil + h.closed = true +} + +func (h *testHandler) ServeHTTP(w http.ResponseWriter, req *http.Request) { + call := &testHandlerCall{ + w: w, + req: req, + ch: make(chan func()), + } + h.mu.Lock() + if h.closed { + h.t.Errorf("test handler called after close") + } + h.calls = append(h.calls, call) + h.mu.Unlock() + for f := range call.ch { + f() + } +} + +func (h *testHandler) nextCall() *testHandlerCall { + h.t.Helper() + synctest.Wait() + h.mu.Lock() + defer h.mu.Unlock() + if len(h.calls) == 0 { + h.t.Fatal("expected server handler call, got none") + } + call := h.calls[0] + h.calls = h.calls[1:] + h.t.Cleanup(call.exit) + return call +} + +// testHandlerCall is a call to the server handler's ServeHTTP method. +type testHandlerCall struct { + w http.ResponseWriter + req *http.Request + closeOnce sync.Once + ch chan func() +} + +// do executes f in the handler's goroutine. +func (call *testHandlerCall) do(f func(http.ResponseWriter, *http.Request)) { + donec := make(chan struct{}) + call.ch <- func() { + defer close(donec) + f(call.w, call.req) + } + <-donec +} + +// exit causes the handler to return. +func (call *testHandlerCall) exit() { + call.closeOnce.Do(func() { + close(call.ch) + }) +} + +func joinCRLF(s ...string) string { + return strings.Join(s, "\r\n") +} diff --git a/src/net/http/http1_transport_test.go b/src/net/http/http1_transport_test.go new file mode 100644 index 00000000000000..7394b93c723f9f --- /dev/null +++ b/src/net/http/http1_transport_test.go @@ -0,0 +1,256 @@ +// Copyright 2026 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +package http_test + +import ( + "bufio" + "context" + "errors" + "internal/nettest" + "net" + "net/http" + "slices" + "sync" + "testing" + "testing/synctest" +) + +// TestHTTP1TransportTest is an example of using http1TransportTest. +func TestHTTP1TransportTest(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + tt := newHTTP1TransportTest(t) + + // tt.roundTrip immediately returns a testRoundTrip, + // which we can use to examine the state of the RoundTrip call. + sentReq, _ := http.NewRequest("GET", "http://example.tld/request/path", nil) + rt := tt.roundTrip(sentReq) + if rt.done() { + t.Fatalf("RoundTrip unexpectedly returned before reading response") + } + + // Expect that the Transport dials a new connection. + // dial.connect provides it with a connection, and gives us the other half. + dial := tt.wantDial("tcp", "example.tld:80") + conn := dial.connect() + + // Read the request written by the Transport. + req := conn.readRequest() + if got, want := req.URL.Path, sentReq.URL.Path; got != want { + t.Fatalf("read request path %q, want %q", got, want) + } + + // Respond, finishing the request. + conn.writeMessage( + "HTTP/1.1 200 OK", + "Content-Length: 0", + "", + ) + rt.wantStatus(200) + }) +} + +// An http1TransportTest tests an HTTP/1 transport using a fake network. +// It must be used in a synctest bubble. +type http1TransportTest struct { + t *testing.T + tr *http.Transport + + dialsMu sync.Mutex + dials []*http1TestDial +} + +func newHTTP1TransportTest(t *testing.T) *http1TransportTest { + tt := &http1TransportTest{ + t: t, + tr: &http.Transport{}, + } + tt.tr.DialContext = (*http1TransportTestDialer)(tt).dialContext + return tt +} + +func (tt *http1TransportTest) roundTrip(req *http.Request) *testRoundTrip { + return newTestRoundTrip(tt.t, tt.tr, req) +} + +func newTestRoundTrip(t *testing.T, roundTripper http.RoundTripper, req *http.Request) *testRoundTrip { + ctx, cancel := context.WithCancel(req.Context()) + req = req.WithContext(ctx) + rt := &testRoundTrip{ + t: t, + donec: make(chan struct{}), + cancel: cancel, + } + go func() { + defer close(rt.donec) + rt.resp, rt.respErr = roundTripper.RoundTrip(req) + }() + synctest.Wait() + + t.Cleanup(func() { + if !rt.done() { + return + } + res, _ := rt.result() + if res != nil { + res.Body.Close() + } + }) + + return rt +} + +func (tt *http1TransportTest) newClientConn(scheme, address string) (*http.ClientConn, *http1TestConn) { + t := tt.t + t.Helper() + + var ( + clientConn *http.ClientConn + err = errors.New("still running") + ) + go func() { + clientConn, err = tt.tr.NewClientConn(t.Context(), scheme, address) + }() + synctest.Wait() + netConn := tt.wantDial("tcp", address).connect() + synctest.Wait() + if err != nil { + t.Fatalf("NewClientConn: %v (want success)", err) + } + t.Cleanup(func() { + netConn.conn.Close() + clientConn.Close() + }) + return clientConn, netConn +} + +func (tt *http1TransportTest) wantDial(network, address string) *http1TestDial { + tt.t.Helper() + synctest.Wait() + tt.dialsMu.Lock() + defer tt.dialsMu.Unlock() + for i, dial := range tt.dials { + if dial.network == network && dial.address == address { + tt.dials = slices.Delete(tt.dials, i, i+1) + return dial + } + } + if len(tt.dials) == 0 { + tt.t.Fatalf("want dial for %q, %q; got none", network, address) + } else { + tt.t.Fatalf("want dial for %q, %q; got %q, %q", network, address, tt.dials[0].network, tt.dials[0].address) + } + return nil +} + +type connOrError struct { + conn net.Conn + err error +} + +type http1TestDial struct { + t *testing.T + network string + address string + resultc chan connOrError +} + +func (dial *http1TestDial) connect() *http1TestConn { + cliConn, srvConn := nettest.NewConnPair() + dial.t.Cleanup(func() { + srvConn.Close() + }) + dial.resultc <- connOrError{conn: cliConn} + srvConn.SetReadError(errWouldBlock) // effectively make reads non-blocking + return &http1TestConn{ + t: dial.t, + conn: srvConn, + bufr: bufio.NewReader(srvConn), + } +} + +type http1TransportTestDialer http1TransportTest + +func (tt *http1TransportTestDialer) dialContext(ctx context.Context, network, address string) (net.Conn, error) { + dial := &http1TestDial{ + t: tt.t, + network: network, + address: address, + resultc: make(chan connOrError, 1), + } + tt.dialsMu.Lock() + tt.dials = append(tt.dials, dial) + tt.dialsMu.Unlock() + select { + case res := <-dial.resultc: + return res.conn, res.err + case <-tt.t.Context().Done(): + return nil, errors.New("test ended") + } +} + +// testRoundTrip manages a RoundTrip in progress. +type testRoundTrip struct { + t *testing.T + resp *http.Response + respErr error + donec chan struct{} + cancel context.CancelFunc +} + +// done reports whether RoundTrip has returned. +func (rt *testRoundTrip) done() bool { + synctest.Wait() + select { + case <-rt.donec: + return true + default: + return false + } +} + +// result returns the result of the RoundTrip. +func (rt *testRoundTrip) result() (*http.Response, error) { + t := rt.t + t.Helper() + synctest.Wait() + select { + case <-rt.donec: + default: + t.Fatalf("RoundTrip is not done; want it to be") + } + return rt.resp, rt.respErr +} + +// response returns the response of a successful RoundTrip. +// If the RoundTrip unexpectedly failed, it calls t.Fatal. +func (rt *testRoundTrip) response() *http.Response { + t := rt.t + t.Helper() + resp, err := rt.result() + if err != nil { + t.Fatalf("RoundTrip returned unexpected error: %v", rt.respErr) + } + if resp == nil { + t.Fatalf("RoundTrip returned nil *Response and nil error") + } + return resp +} + +// err returns the (possibly nil) error result of RoundTrip. +func (rt *testRoundTrip) err() error { + t := rt.t + t.Helper() + _, err := rt.result() + return err +} + +// wantStatus indicates the expected response StatusCode. +func (rt *testRoundTrip) wantStatus(want int) { + t := rt.t + t.Helper() + if got := rt.response().StatusCode; got != want { + t.Fatalf("got response status %v, want %v", got, want) + } +} diff --git a/src/net/http/httputil/reverseproxy.go b/src/net/http/httputil/reverseproxy.go index ba9b507cd1587c..cc5bb7b4afa021 100644 --- a/src/net/http/httputil/reverseproxy.go +++ b/src/net/http/httputil/reverseproxy.go @@ -518,6 +518,21 @@ func (p *ReverseProxy) ServeHTTP(rw http.ResponseWriter, req *http.Request) { } } + if outreq.Method == "CONNECT" { + // We cannot handle CONNECT requests. + // (Perhaps we should just send a 405 and never call ErrorHandler? + // More consistent to always call ErrorHandler or ModifyResponse, + // so we do that for now.) + err := errors.New("client sent unsupported CONNECT request") + if p.ErrorHandler != nil { + p.ErrorHandler(rw, outreq, err) + } else { + p.logf("http: proxy error: %v", err) + rw.WriteHeader(http.StatusMethodNotAllowed) + } + return + } + if _, ok := outreq.Header["User-Agent"]; !ok { // If the outbound request doesn't have a User-Agent header set, // don't send the default Go HTTP client User-Agent. diff --git a/src/net/http/httputil/reverseproxy_test.go b/src/net/http/httputil/reverseproxy_test.go index aa26d900cee8f2..02f9c72981539f 100644 --- a/src/net/http/httputil/reverseproxy_test.go +++ b/src/net/http/httputil/reverseproxy_test.go @@ -2215,3 +2215,25 @@ func (rc *testReadWriteCloser) Close() error { } return nil } + +func TestReverseProxyCONNECT(t *testing.T) { + proxy := &ReverseProxy{ + Rewrite: func(r *ProxyRequest) { + backendURL, _ := url.Parse("http://backend.tld/") + r.SetURL(backendURL) + }, + } + frontend := httptest.NewServer(proxy) + defer frontend.Close() + + req, _ := http.NewRequest("CONNECT", frontend.URL, strings.NewReader("body")) + resp, err := frontend.Client().Do(req) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + + if got, want := resp.StatusCode, http.StatusMethodNotAllowed; got != want { + t.Errorf("on response to CONNECT: got status %v, want %v", got, want) + } +} diff --git a/src/net/http/internal/http2/flow.go b/src/net/http/internal/http2/flow.go index b7dbd186957ef7..5066a39912b38d 100644 --- a/src/net/http/internal/http2/flow.go +++ b/src/net/http/internal/http2/flow.go @@ -10,6 +10,9 @@ package http2 // flow control window update. const inflowMinRefresh = 4 << 10 +// maxFlowWindow is the maximum size of a flow control window. +const maxFlowWindow = (1 << 31) - 1 + // inflow accounts for an inbound flow control window. // It tracks both the latest window sent to the peer (used for enforcement) // and the accumulated unsent window. @@ -74,47 +77,99 @@ func takeInflows(f1, f2 *inflow, n uint32) bool { return true } -// outflow is the outbound flow control window's size. +// connOutflow is connection-level outbound flow control. +type connOutflow struct { + initial int32 // SETTINGS_INITIAL_WINDOW_SIZE, changes with settings updates + n int32 // connection-level flow control window + flowErr bool // set when a flow control error is encountered +} + +func (f *connOutflow) init() { + f.initial = initialWindowSize // initial stream window size + f.n = initialWindowSize // current connection window size +} + +func (f *connOutflow) changeInitialWindowSize(size int64) bool { + if size > maxFlowWindow { + f.flowErr = true + return false + } + f.initial = int32(size) + return true +} + +func (f *connOutflow) add(n int32) bool { + sum := int64(f.n) + int64(n) + if sum > maxFlowWindow { + f.flowErr = true + return false + } + f.n += n + return true +} + +// outflow is the stream-level outbound flow control window's size. type outflow struct { _ incomparable - // n is the number of DATA bytes we're allowed to send. - // An outflow is kept both on a conn and a per-stream. - n int32 + // delta is the difference between the stream's flow control window and + // the connection's initial window size (conn.initial). + // + // Another view is that delta is the number of flow control bytes provided to this + // stream in WINDOW_UPDATE frames, less the number of bytes sent on the stream. + delta int32 // conn points to the shared connection-level outflow that is - // shared by all streams on that conn. It is nil for the outflow - // that's on the conn directly. - conn *outflow + // shared by all streams on that conn. + conn *connOutflow } -func (f *outflow) setConnFlow(cf *outflow) { f.conn = cf } - -func (f *outflow) available() int32 { - n := f.n - if f.conn != nil && f.conn.n < n { - n = f.conn.n +func (f *outflow) available() (int32, bool) { + if f.conn == nil { + return maxFlowWindow, true // only happens in tests } - return n + if f.conn.flowErr { + // Block all sending once any stream observes a flow control error. + return 0, false + } + n := int64(f.conn.initial) + int64(f.delta) + if n > maxFlowWindow { + f.conn.flowErr = true + return 0, false + } + return min(int32(n), f.conn.n), true } func (f *outflow) take(n int32) { - if n > f.available() { - panic("internal error: took too much") + if f.conn == nil { + return // only happens in tests } - f.n -= n - if f.conn != nil { - f.conn.n -= n + avail, _ := f.available() + if n > avail { + panic("internal error: took too much") } + f.delta -= n + f.conn.n -= n } // add adds n bytes (positive or negative) to the flow control window. // It returns false if the sum would exceed 2^31-1. func (f *outflow) add(n int32) bool { - sum := f.n + n - if (sum > n) == (f.n > 0) { - f.n = sum - return true + if f.conn == nil { + return true // only happens in tests + } + avail := int64(f.conn.initial) + int64(f.delta) + if avail > maxFlowWindow { + // An earlier change to the initial window pushed this stream over the limit. + // This is a connection-level flow control error. + f.conn.flowErr = true + return false } - return false + if avail+int64(n) > maxFlowWindow { + // This update would push the stream over the limit. + // This is a stream-level flow control error. + return false + } + f.delta += n + return true } diff --git a/src/net/http/internal/http2/flow_test.go b/src/net/http/internal/http2/flow_test.go index cae4f38c0c4260..b4b0f47584a05c 100644 --- a/src/net/http/internal/http2/flow_test.go +++ b/src/net/http/internal/http2/flow_test.go @@ -60,43 +60,42 @@ func TestTakeInflows(t *testing.T) { func TestOutFlow(t *testing.T) { var st outflow - var conn outflow + var conn connOutflow + st.conn = &conn st.add(3) conn.add(2) - if got, want := st.available(), int32(3); got != want { - t.Errorf("available = %d; want %d", got, want) - } - st.setConnFlow(&conn) - if got, want := st.available(), int32(2); got != want { - t.Errorf("after parent setup, available = %d; want %d", got, want) + if got, ok := st.available(); !ok || got != 2 { + t.Errorf("available = %d, %v; want 2, true", got, ok) } st.take(2) - if got, want := conn.available(), int32(0); got != want { + if got, want := conn.n, int32(0); got != want { t.Errorf("after taking 2, conn = %d; want %d", got, want) } - if got, want := st.available(), int32(0); got != want { - t.Errorf("after taking 2, stream = %d; want %d", got, want) + if got, ok := st.available(); !ok || got != 0 { + t.Errorf("after taking 2, stream = %d, %v; want 0, true", got, ok) } } func TestOutFlowAdd(t *testing.T) { var f outflow + f.conn = &connOutflow{} + f.conn.add(1<<31 - 1) if !f.add(1) { t.Fatal("failed to add 1") } if !f.add(-1) { t.Fatal("failed to add -1") } - if got, want := f.available(), int32(0); got != want { - t.Fatalf("size = %d; want %d", got, want) + if got, ok := f.available(); !ok || got != 0 { + t.Fatalf("size = %d, %v; want 0, true", got, ok) } if !f.add(1<<31 - 1) { t.Fatal("failed to add 2^31-1") } - if got, want := f.available(), int32(1<<31-1); got != want { - t.Fatalf("size = %d; want %d", got, want) + if got, ok := f.available(); !ok || got != 1<<31-1 { + t.Fatalf("size = %d, %v; want %d, true", got, ok, 1<<31-1) } if f.add(1) { t.Fatal("adding 1 to max shouldn't be allowed") @@ -105,6 +104,8 @@ func TestOutFlowAdd(t *testing.T) { func TestOutFlowAddOverflow(t *testing.T) { var f outflow + f.conn = &connOutflow{} + f.conn.add(1<<31 - 1) if !f.add(0) { t.Fatal("failed to add 0") } @@ -126,14 +127,14 @@ func TestOutFlowAddOverflow(t *testing.T) { if !f.add(-3) { t.Fatal("failed to add -3") } - if got, want := f.available(), int32(-2); got != want { - t.Fatalf("size = %d; want %d", got, want) + if got, ok := f.available(); !ok || got != -2 { + t.Fatalf("size = %d, %v; want -2, true", got, ok) } if !f.add(1<<31 - 1) { t.Fatal("failed to add 2^31-1") } - if got, want := f.available(), int32(1+-3+(1<<31-1)); got != want { - t.Fatalf("size = %d; want %d", got, want) + if got, ok := f.available(); !ok || got != 1+-3+(1<<31-1) { + t.Fatalf("size = %d, %v; want %d, true", got, ok, 1+-3+(1<<31-1)) } } diff --git a/src/net/http/internal/http2/frame.go b/src/net/http/internal/http2/frame.go index 5567c2ae97a164..f30fa57cea1e00 100644 --- a/src/net/http/internal/http2/frame.go +++ b/src/net/http/internal/http2/frame.go @@ -1846,9 +1846,13 @@ func (fr *Framer) readMetaFrame(hf *HeadersFrame) (Frame, error) { mh := &MetaHeadersFrame{ HeadersFrame: hf, } - var remainSize = fr.maxHeaderListSize() + type headerBudget struct { + remainSize uint32 + count int + } + headers := headerBudget{remainSize: fr.maxHeaderListSize()} + trailers := headerBudget{remainSize: fr.maxHeaderListSize()} var sawRegular bool - var headerCount int var invalid error // pseudo header field errors hdec := fr.ReadMetaHeaders @@ -1858,13 +1862,6 @@ func (fr *Framer) readMetaFrame(hf *HeadersFrame) (Frame, error) { if VerboseLogs && fr.logReads { fr.debugReadLoggerf("http2: decoded hpack field %+v", hf) } - headerCount++ - if limit := fr.maxHeaderValueCount(); limit > 0 && headerCount > limit { - hdec.SetEmitEnabled(false) - mh.Truncated = true - remainSize = 0 - return - } if !httpguts.ValidHeaderFieldValue(hf.Value) { // Don't include the value in the error, because it may be sensitive. invalid = headerFieldValueError(hf.Name) @@ -1886,14 +1883,31 @@ func (fr *Framer) readMetaFrame(hf *HeadersFrame) (Frame, error) { return } - size := hf.Size() - if size > remainSize { + var budget *headerBudget + var size uint32 + if hf.Name == "trailer" { + budget = &trailers + fieldCount := strings.Count(hf.Value, ",") + 1 + budget.count += fieldCount + // Rather than actually constructing hpack.HeaderField for each + // trailer field and cumulatively adding its Size, just do the math + // manually to avoid unnecessary work. This does make it so + // whitespaces after comma are counted against the budget, but that + // should be innocuous. + size = uint32(len(hf.Value)-fieldCount+1) + uint32(fieldCount)*hpack.HeaderField{}.Size() + } else { + budget = &headers + budget.count++ + size = hf.Size() + } + if countLimit := fr.maxHeaderValueCount(); (countLimit > 0 && budget.count > countLimit) || size > budget.remainSize { hdec.SetEmitEnabled(false) mh.Truncated = true - remainSize = 0 + headers.remainSize = 0 + trailers.remainSize = 0 return } - remainSize -= size + budget.remainSize -= size mh.Fields = append(mh.Fields, hf) }) @@ -1909,10 +1923,10 @@ func (fr *Framer) readMetaFrame(hf *HeadersFrame) (Frame, error) { // skip parsing the fragment and close the connection. // // "Too much" is either any CONTINUATION frame after we've already - // exceeded the max header list size (in which case remainSize is 0), - // or a frame whose encoded size is more than twice the remaining - // header list bytes we're willing to accept. - if int64(len(frag)) > int64(2*remainSize) { + // exceeded the max header list size (if so, both budgets are 0), or a + // frame whose encoded size is more than twice the remaining header + // list bytes we're willing to accept. + if int64(len(frag)) > 2*int64(headers.remainSize+trailers.remainSize) { if VerboseLogs { log.Printf("http2: header list too large") } diff --git a/src/net/http/internal/http2/frame_test.go b/src/net/http/internal/http2/frame_test.go index 2d3318be9695c4..70dcdce3a3ae2c 100644 --- a/src/net/http/internal/http2/frame_test.go +++ b/src/net/http/internal/http2/frame_test.go @@ -1130,6 +1130,15 @@ func TestMetaFrameHeader(t *testing.T) { oneKBString := strings.Repeat("a", 1<<10) + // A Trailer declaration of 20 fields: ~250 bytes on the wire, but ~900 + // bytes by our accounting since we add 32 bytes for each entry in + // Request.Trailer. + var trailerNames []string + for i := range 20 { + trailerNames = append(trailerNames, fmt.Sprintf("x-trailer-%d", i)) + } + trailerDecl := strings.Join(trailerNames, ",") + tests := [...]struct { name string w func(*Framer) @@ -1267,6 +1276,29 @@ func TestMetaFrameHeader(t *testing.T) { want: streamError(1, ErrCodeProtocol), wantErrReason: `invalid header field value for "key"`, }, + 13: { + name: "trailer_declaration_okay", + w: func(f *Framer) { + write(f, encodeHeaderRaw(t, ":method", "GET", ":path", "/", "trailer", trailerDecl)) + }, + maxHeaderListSize: 1024, + want: want(FlagHeadersEndHeaders, 193, + ":method", "GET", + ":path", "/", + "trailer", trailerDecl, + ), + }, + 14: { + name: "trailer_declaration_truncated", + w: func(f *Framer) { + write(f, encodeHeaderRaw(t, ":method", "GET", ":path", "/", "trailer", trailerDecl)) + }, + maxHeaderListSize: 512, + want: truncated(want(FlagHeadersEndHeaders, 193, + ":method", "GET", + ":path", "/", + )), + }, } for i, tt := range tests { buf := new(bytes.Buffer) diff --git a/src/net/http/internal/http2/server.go b/src/net/http/internal/http2/server.go index f3dc9d537e1b51..fe35030c7e88bc 100644 --- a/src/net/http/internal/http2/server.go +++ b/src/net/http/internal/http2/server.go @@ -291,7 +291,6 @@ func (s *Server) serveConn(c net.Conn, opts *ServeConnOpts, newf func(*serverCon doneServing: make(chan struct{}), clientMaxStreams: math.MaxUint32, // Section 6.5.2: "Initially, there is no limit to this value" advMaxStreams: uint32(conf.MaxConcurrentStreams), - initialStreamSendWindowSize: initialWindowSize, initialStreamRecvWindowSize: int32(conf.MaxReceiveBufferPerStream), maxFrameSize: initialMaxFrameSize, pingTimeout: conf.PingTimeout, @@ -334,7 +333,7 @@ func (s *Server) serveConn(c net.Conn, opts *ServeConnOpts, newf func(*serverCon // These start at the RFC-specified defaults. If there is a higher // configured value for inflow, that will be updated when we send a // WINDOW_UPDATE shortly after sending SETTINGS. - sc.flow.add(initialWindowSize) + sc.flow.init() sc.inflow.init(initialWindowSize) sc.hpackEncoder = hpack.NewEncoder(&sc.headerWriteBuf) sc.hpackEncoder.SetMaxDynamicTableSizeLimit(uint32(conf.MaxEncoderHeaderTableSize)) @@ -476,7 +475,7 @@ type serverConn struct { wroteFrameCh chan frameWriteResult // from writeFrameAsync -> serve, tickles more frame writes bodyReadCh chan bodyReadMsg // from handlers -> serve serveMsgCh chan any // misc messages & code to send to / run on the serve loop - flow outflow // conn-wide (not stream-specific) outbound flow control + flow connOutflow // conn-wide (not stream-specific) outbound flow control inflow inflow // conn-wide inbound flow control tlsState *tls.ConnectionState // shared by all handlers, like net/http remoteAddrStr string @@ -508,6 +507,9 @@ type serverConn struct { needToSendSettingsAck bool unackedSettings int // how many SETTINGS have we sent without ACKs? pendingDecoderTableSize uint32 // if non-zero, HPACK decoder table size to apply on SETTINGS ack + pendingEncoderTableSize bool // peer changed SETTINGS_HEADER_TABLE_SIZE; apply to hpackEncoder before the next frame write + encoderTableSizeMin uint32 // smallest SETTINGS_HEADER_TABLE_SIZE since the last apply + encoderTableSize uint32 // latest SETTINGS_HEADER_TABLE_SIZE queuedControlFrames int // control frames in the writeSched queue clientMaxStreams uint32 // SETTINGS_MAX_CONCURRENT_STREAMS from client (our PUSH_PROMISE limit) advMaxStreams uint32 // our SETTINGS_MAX_CONCURRENT_STREAMS advertised the client @@ -518,7 +520,6 @@ type serverConn struct { maxPushPromiseID uint32 // ID of the last push promise (even), or 0 if there have been no pushes streams map[uint32]*stream unstartedHandlers []unstartedHandler - initialStreamSendWindowSize int32 initialStreamRecvWindowSize int32 maxFrameSize int32 peerMaxHeaderListSize uint32 // zero means unknown (default) @@ -591,7 +592,8 @@ type stream struct { // immutable: sc *serverConn id uint32 - body *pipe // non-nil if expecting DATA frames + body *pipe // non-nil if expecting DATA frames + reqBody *requestBody cw closeWaiter // closed wait stream transitions to closed state ctx context.Context cancelCtx func() @@ -1441,6 +1443,16 @@ func (sc *serverConn) startFrameWrite(wr FrameWriteRequest) { sc.writingFrame = true sc.needsFrameFlush = true + if sc.pendingEncoderTableSize { + // hpackEncoder may be in use by writeFrameAsync, so SETTINGS + // changes to it are deferred until no frame is being written. + // Replaying the smallest size before the latest one keeps the + // encoder's view identical to having applied every change + // (RFC 7541, Section 4.2). + sc.pendingEncoderTableSize = false + sc.hpackEncoder.SetMaxDynamicTableSize(sc.encoderTableSizeMin) + sc.hpackEncoder.SetMaxDynamicTableSize(sc.encoderTableSize) + } if wr.write.staysWithinBuffer(sc.bw.Available()) { sc.writingFrameAsync = false err := wr.write.writeFrame(sc) @@ -1540,6 +1552,11 @@ func (sc *serverConn) scheduleFrameWrite() { } sc.inFrameScheduleLoop = true for !sc.writingFrameAsync { + if sc.flow.flowErr && (!sc.inGoAway || sc.goAwayCode == ErrCodeNo) { + sc.inGoAway = true + sc.needToSendGoAway = true + sc.goAwayCode = ErrCodeFlowControl + } if sc.needToSendGoAway { sc.needToSendGoAway = false sc.startFrameWrite(FrameWriteRequest{ @@ -1563,6 +1580,9 @@ func (sc *serverConn) scheduleFrameWrite() { sc.startFrameWrite(wr) continue } + if sc.flow.flowErr { + continue + } } if sc.needsFrameFlush { sc.startFrameWrite(FrameWriteRequest{write: flushFrameWriter{}}) @@ -1795,6 +1815,10 @@ func (sc *serverConn) processWindowUpdate(f *WindowUpdateFrame) error { return nil } if !st.flow.add(int32(f.Increment)) { + if st.flow.conn.flowErr { + // This is a lazily-detected connection-level flow control error. + return sc.countError("bad_flow", ConnectionError(ErrCodeFlowControl)) + } return sc.countError("bad_flow", streamError(f.StreamID, ErrCodeFlowControl)) } default: // connection-level flow control @@ -1854,10 +1878,6 @@ func (sc *serverConn) closeStream(st *stream, err error) { } } if p := st.body; p != nil { - // Return any buffered unread bytes worth of conn-level flow control. - // See golang.org/issue/16481 - sc.sendWindowUpdate(nil, p.Len()) - p.CloseWithError(err) } if e, ok := err.(StreamError); ok { @@ -1922,7 +1942,12 @@ func (sc *serverConn) processSetting(s Setting) error { } switch s.ID { case SettingHeaderTableSize: - sc.hpackEncoder.SetMaxDynamicTableSize(s.Val) + // Applied by startFrameWrite; see comment there. + if !sc.pendingEncoderTableSize || s.Val < sc.encoderTableSizeMin { + sc.encoderTableSizeMin = s.Val + } + sc.encoderTableSize = s.Val + sc.pendingEncoderTableSize = true case SettingEnablePush: sc.pushEnabled = s.Val != 0 case SettingMaxConcurrentStreams: @@ -1953,28 +1978,14 @@ func (sc *serverConn) processSetting(s Setting) error { func (sc *serverConn) processSettingInitialWindowSize(val uint32) error { sc.serveG.check() - // Note: val already validated to be within range by - // processSetting's Valid call. - - // "A SETTINGS frame can alter the initial flow control window - // size for all current streams. When the value of - // SETTINGS_INITIAL_WINDOW_SIZE changes, a receiver MUST - // adjust the size of all stream flow control windows that it - // maintains by the difference between the new value and the - // old value." - old := sc.initialStreamSendWindowSize - sc.initialStreamSendWindowSize = int32(val) - growth := int32(val) - old // may be negative - for _, st := range sc.streams { - if !st.flow.add(growth) { - // 6.9.2 Initial Flow Control Window Size - // "An endpoint MUST treat a change to - // SETTINGS_INITIAL_WINDOW_SIZE that causes any flow - // control window to exceed the maximum size as a - // connection error (Section 5.4.1) of type - // FLOW_CONTROL_ERROR." - return sc.countError("setting_win_size", ConnectionError(ErrCodeFlowControl)) - } + if !sc.flow.changeInitialWindowSize(int64(val)) { + // 6.9.2 Initial Flow Control Window Size + // "An endpoint MUST treat a change to + // SETTINGS_INITIAL_WINDOW_SIZE that causes any flow + // control window to exceed the maximum size as a + // connection error (Section 5.4.1) of type + // FLOW_CONTROL_ERROR." + return sc.countError("setting_win_size", ConnectionError(ErrCodeFlowControl)) } return nil } @@ -2247,7 +2258,7 @@ func (sc *serverConn) processHeaders(f *MetaHeadersFrame) error { if st.reqTrailer != nil { st.trailer = make(Header) } - st.body = req.Body.(*requestBody).pipe // may be nil + st.body = st.reqBody.pipe // may be nil st.declBodyBytes = req.ContentLength handler := sc.handler.ServeHTTP @@ -2262,7 +2273,7 @@ func (sc *serverConn) processHeaders(f *MetaHeadersFrame) error { st.readDeadline = time.AfterFunc(sc.hs.ReadTimeout(), st.onReadTimeout) } - return sc.scheduleHandler(id, rw, req, handler) + return sc.scheduleHandler(st, rw, req, handler) } func (sc *serverConn) upgradeRequest(req *ServerRequest) { @@ -2374,7 +2385,6 @@ func (sc *serverConn) newStream(id, pusherID uint32, state streamState, priority } st.cw.Init() st.flow.conn = &sc.flow // link to conn-level counter - st.flow.add(sc.initialStreamSendWindowSize) st.inflow.init(sc.initialStreamRecvWindowSize) if writeTimeout := sc.hs.WriteTimeout(); writeTimeout > 0 { st.writeDeadline = time.AfterFunc(writeTimeout, st.onWriteTimeout) @@ -2468,7 +2478,7 @@ func (sc *serverConn) newWriterAndRequest(st *stream, f *MetaHeadersFrame) (*res if _, ok := rp.Header["Content-Length"]; !ok { req.ContentLength = -1 } - req.Body.(*requestBody).pipe = &pipe{ + st.reqBody.pipe = &pipe{ b: &dataBuffer{expected: req.ContentLength}, } } @@ -2478,17 +2488,12 @@ func (sc *serverConn) newWriterAndRequest(st *stream, f *MetaHeadersFrame) (*res func (sc *serverConn) newWriterAndRequestNoBody(st *stream, rp httpcommon.ServerRequestParam) (*responseWriter, *ServerRequest, error) { sc.serveG.check() - var tlsState *tls.ConnectionState // nil if not scheme https - if rp.Scheme == "https" { - tlsState = sc.tlsState - } - res := httpcommon.NewServerRequest(rp) if res.InvalidReason != "" { return nil, nil, sc.countError(res.InvalidReason, streamError(st.id, ErrCodeProtocol)) } - body := &requestBody{ + st.reqBody = &requestBody{ conn: sc, stream: st, needsContinue: res.NeedsContinue, @@ -2504,9 +2509,9 @@ func (sc *serverConn) newWriterAndRequestNoBody(st *stream, rp httpcommon.Server Proto: "HTTP/2.0", ProtoMajor: 2, ProtoMinor: 0, - TLS: tlsState, + TLS: sc.tlsState, Host: rp.Authority, - Body: body, + Body: st.reqBody, Trailer: res.Trailer, } return rw, &rw.rws.req, nil @@ -2525,11 +2530,12 @@ type unstartedHandler struct { rw *responseWriter req *ServerRequest handler func(*ResponseWriter, *ServerRequest) + body *pipe } // scheduleHandler starts a handler goroutine, // or schedules one to start as soon as an existing handler finishes. -func (sc *serverConn) scheduleHandler(streamID uint32, rw *responseWriter, req *ServerRequest, handler func(*ResponseWriter, *ServerRequest)) error { +func (sc *serverConn) scheduleHandler(st *stream, rw *responseWriter, req *ServerRequest, handler func(*ResponseWriter, *ServerRequest)) error { sc.serveG.check() maxHandlers := sc.advMaxStreams if sc.curHandlers < maxHandlers { @@ -2541,10 +2547,11 @@ func (sc *serverConn) scheduleHandler(streamID uint32, rw *responseWriter, req * return sc.countError("too_many_early_resets", ConnectionError(ErrCodeEnhanceYourCalm)) } sc.unstartedHandlers = append(sc.unstartedHandlers, unstartedHandler{ - streamID: streamID, + streamID: st.id, rw: rw, req: req, handler: handler, + body: st.body, }) return nil } @@ -2558,6 +2565,10 @@ func (sc *serverConn) handlerDone() { u := sc.unstartedHandlers[i] if sc.streams[u.streamID] == nil { // This stream was reset before its goroutine had a chance to start. + if u.body != nil { + u.body.BreakWithError(errClosedBody) + sc.sendWindowUpdate(nil, u.body.Len()) + } continue } if sc.curHandlers >= maxHandlers { @@ -2582,6 +2593,12 @@ func (sc *serverConn) runHandler(rw *responseWriter, req *ServerRequest, handler if req.MultipartForm != nil { req.MultipartForm.RemoveAll() } + if b := rw.rws.stream.reqBody; b != nil { + // Closing the body refunds flow control credit for any unconsumed data. + // (reqBody is nil for Upgrade: h2c requests, but those do not use flow + // control for the request body.) + b.Close() + } if didPanic { e := recover() sc.writeFrameFromHandler(FrameWriteRequest{ @@ -2679,7 +2696,7 @@ func (sc *serverConn) noteBodyReadFromHandler(st *stream, n int, err error) { func (sc *serverConn) noteBodyRead(st *stream, n int) { sc.serveG.check() sc.sendWindowUpdate(nil, n) // conn-level - if st.state != stateHalfClosedRemote && st.state != stateClosed { + if st != nil && st.state != stateHalfClosedRemote && st.state != stateClosed { // Don't send this WINDOW_UPDATE if the stream is closed // remotely. sc.sendWindowUpdate(st, n) @@ -2727,6 +2744,9 @@ func (b *requestBody) Close() error { b.closeOnce.Do(func() { if b.pipe != nil { b.pipe.BreakWithError(errClosedBody) + if unread := b.pipe.Len(); unread > 0 { + b.conn.noteBodyReadFromHandler(nil, unread, errClosedBody) + } } }) return nil @@ -2744,9 +2764,6 @@ func (b *requestBody) Read(p []byte) (n int, err error) { if err == io.EOF { b.sawEOF = true } - if b.conn == nil { - return - } b.conn.noteBodyReadFromHandler(b.stream, n, err) return } diff --git a/src/net/http/internal/http2/server_test.go b/src/net/http/internal/http2/server_test.go index 9358d3bb33601a..d3e8b76d191cfb 100644 --- a/src/net/http/internal/http2/server_test.go +++ b/src/net/http/internal/http2/server_test.go @@ -610,6 +610,53 @@ func testServer(t *testing.T) { <-gotReq } +func TestServer_Request_TLS(t *testing.T) { + for _, unencrypted := range []bool{false, true} { + for _, scheme := range []string{"https", "http", ""} { + name := scheme + if scheme == "" { + name = "CONNECT" + } + t.Run(fmt.Sprintf("unencrypted=%v/%s", unencrypted, name), func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + gotTLS := make(chan *tls.ConnectionState, 1) + st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) { + gotTLS <- r.TLS + }, func(s *http.Server) { + s.Protocols = new(http.Protocols) + s.Protocols.SetHTTP2(!unencrypted) + s.Protocols.SetUnencryptedHTTP2(unencrypted) + }) + st.greet() + headers := []string{":method", "CONNECT", ":authority", "example.com:443"} + if scheme != "" { + headers = []string{":method", "GET", ":authority", "example.com", ":scheme", scheme, ":path", "/"} + } + st.writeHeaders(HeadersFrameParam{ + StreamID: 1, + BlockFragment: st.encodeHeaderRaw(headers...), + EndStream: true, + EndHeaders: true, + }) + state := <-gotTLS + if unencrypted { + if state != nil { + t.Fatalf("Request.TLS = %v; want nil for an unencrypted connection", state) + } + } else { + if state == nil { + t.Fatal("Request.TLS = nil; want TLS connection state") + } + if !state.HandshakeComplete || state.NegotiatedProtocol != "h2" { + t.Errorf("Request.TLS = %+v; want completed HTTP/2 TLS handshake", state) + } + } + }) + }) + } + } +} + func TestServer_Request_Get(t *testing.T) { synctest.Test(t, testServer_Request_Get) } func testServer_Request_Get(t *testing.T) { testServerRequest(t, func(st *serverTester) { @@ -1395,6 +1442,9 @@ func testServer_Send_RstStream_After_Bogus_WindowUpdate(t *testing.T) { t.Fatal(err) } st.wantRSTStream(1, ErrCodeFlowControl) + // Connection is still alive, even if the stream has been reset. + st.writePing(false, [8]byte{}) + st.wantFrameType(FramePing) } // testServerPostUnblock sends a hanging POST with unsent data to handler, @@ -2929,6 +2979,64 @@ func testServer_MaxEncoderHeaderTableSize(t *testing.T) { } } +// TestServer_HeaderTableSizeDuringWrite tests that a +// SETTINGS_HEADER_TABLE_SIZE change from the client is not applied to the +// server's HPACK encoder while a frame write, which may be using the encoder +// on another goroutine, is in progress. +func TestServer_HeaderTableSizeDuringWrite(t *testing.T) { + synctest.Test(t, testServer_HeaderTableSizeDuringWrite) +} +func testServer_HeaderTableSizeDuringWrite(t *testing.T) { + st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {}) + defer st.Close() + st.greet() + enc := st.sc.TestHPACKEncoder() + + // Leave the response HEADERS write for stream 1 in progress while the + // client changes SETTINGS_HEADER_TABLE_SIZE twice. + st.cc.(*tls.Conn).NetConn().(*synctestNetConn).SetReadBufferSize(0) + st.bodylessReq1() + synctest.Wait() + st.writeSettings(Setting{SettingHeaderTableSize, 0}) + st.writeSettings(Setting{SettingHeaderTableSize, 2048}) + synctest.Wait() + // Table size should not be changed immediately. To avoid concurrent use of + // the encoder, we apply the size change right before we use the encoder to + // write frames. + if got, want := enc.MaxDynamicTableSize(), uint32(InitialHeaderTableSize); got != want { + t.Errorf("during frame write: encoder header table size = %d, want %d", got, want) + } + st.cc.(*tls.Conn).NetConn().(*synctestNetConn).SetReadBufferSize(math.MaxInt) + st.wantHeaders(wantHeader{streamID: 1, endStream: true}) + st.wantSettingsAck() + + st.writeHeaders(HeadersFrameParam{ + StreamID: 3, + BlockFragment: st.encodeHeader(), + EndStream: true, + EndHeaders: true, + }) + synctest.Wait() + hf := readFrame[*HeadersFrame](t, st) + if hf.StreamID != 3 { + t.Fatalf("got HEADERS for stream %d, want stream 3", hf.StreamID) + } + // The response must signal both the smallest (0) and the final (2048) + // table size (RFC 7541, Section 4.2), just like an encoder that saw both + // size changes directly. + var want bytes.Buffer + wantEnc := hpack.NewEncoder(&want) + wantEnc.SetMaxDynamicTableSize(0) + wantEnc.SetMaxDynamicTableSize(2048) + wantEnc.WriteField(hpack.HeaderField{Name: ":status", Value: "200"}) + if got := hf.HeaderBlockFragment(); !bytes.HasPrefix(got, want.Bytes()) { + t.Errorf("stream 3 header block = %x, want prefix %x", got, want.Bytes()) + } + if got, want := enc.MaxDynamicTableSize(), uint32(2048); got != want { + t.Errorf("after frame write: encoder header table size = %d, want %d", got, want) + } +} + // Issue 12843 func TestServerDoS_MaxHeaderListSize(t *testing.T) { synctest.Test(t, testServerDoS_MaxHeaderListSize) } func testServerDoS_MaxHeaderListSize(t *testing.T) { @@ -3671,6 +3779,166 @@ func testServerReturnsStreamAndConnFlowControlOnBodyClose(t *testing.T) { }) } +func TestServerResetStreamUnreadBody(t *testing.T) { + synctestSubtest(t, "read_after_reset", func(t *testing.T) { + st := newServerTester(t, nil) + defer st.Close() + + st.greet() + st.writeHeaders(HeadersFrameParam{ + StreamID: 1, + BlockFragment: st.encodeHeader(":method", "POST"), + EndHeaders: true, + }) + call := st.nextHandlerCall() + + const size = InflowMinRefresh + st.writeData(1, false, make([]byte, size)) + st.writeRSTStream(1, ErrCodeCancel) + st.sync() + + // The handler reads the body after the stream has been reset. + call.do(func(w http.ResponseWriter, req *http.Request) { + io.ReadAll(req.Body) + }) + st.wantWindowUpdate(0, size) + + // Handler exits; no second refund. + call.exit() + st.wantIdle() + }) + + synctestSubtest(t, "exit_without_reading", func(t *testing.T) { + st := newServerTester(t, nil) + defer st.Close() + + st.greet() + st.writeHeaders(HeadersFrameParam{ + StreamID: 1, + BlockFragment: st.encodeHeader(":method", "POST"), + EndHeaders: true, + }) + call := st.nextHandlerCall() + + const size = InflowMinRefresh + st.writeData(1, false, make([]byte, size)) + st.writeRSTStream(1, ErrCodeCancel) + st.sync() + + // The handler exits without reading the body. Flow control is returned on exit. + call.exit() + st.wantWindowUpdate(0, size) + st.wantIdle() + }) + + synctestSubtest(t, "read_after_handler_exit", func(t *testing.T) { + st := newServerTester(t, nil) + defer st.Close() + + st.greet() + st.writeHeaders(HeadersFrameParam{ + StreamID: 1, + BlockFragment: st.encodeHeader(":method", "POST"), + EndHeaders: true, + }) + call := st.nextHandlerCall() + + const size = InflowMinRefresh + st.writeData(1, false, make([]byte, size)) + synctest.Wait() + + // The handler exits without reading the body. + call.exit() + + // Flow control is returned when the handler exits. + st.wantUnorderedFrames( + func(f *WindowUpdateFrame) bool { + return f.StreamID == 0 && f.Increment == size + }, + func(f *HeadersFrame) bool { + return f.StreamID == 1 && f.StreamEnded() + }, + func(f *RSTStreamFrame) bool { + return f.StreamID == 1 && f.ErrCode == ErrCodeNo + }, + ) + + // Reading the body after the handler exits must not cause a double refund. + io.ReadAll(call.req.Body) + st.wantIdle() + }) + + synctestSubtest(t, "unstarted_handler", func(t *testing.T) { + st := newServerTester(t, nil, func(h2 *http.HTTP2Config) { + h2.MaxConcurrentStreams = 1 + }) + defer st.Close() + + st.greet() + + // Stream 1 uses the single handler slot. + st.writeHeaders(HeadersFrameParam{ + StreamID: 1, + BlockFragment: st.encodeHeader(), + EndStream: true, + EndHeaders: true, + }) + call := st.nextHandlerCall() + + // Reset stream 1 so the client can open another stream without + // exceeding the concurrent stream limit. + st.writeRSTStream(1, ErrCodeCancel) + + // Stream 3 is queued in unstartedHandlers because the handler for + // stream 1 is still executing. + st.writeHeaders(HeadersFrameParam{ + StreamID: 3, + BlockFragment: st.encodeHeader(":method", "POST"), + EndHeaders: true, + }) + + const size = InflowMinRefresh + st.writeData(3, false, make([]byte, size)) + st.writeRSTStream(3, ErrCodeCancel) + st.wantIdle() + + // Stream 1 handler exits. Stream 3 is removed from unstartedHandlers + // and its flow control is returned. + call.exit() + st.wantWindowUpdate(0, size) + st.wantIdle() + }) + + // Same thing as exit_without_reading, but some imaginary middleware + // replaced Request.Body first. + synctestSubtest(t, "middleware_replaces_body", func(t *testing.T) { + st := newServerTester(t, nil) + defer st.Close() + + st.greet() + st.writeHeaders(HeadersFrameParam{ + StreamID: 1, + BlockFragment: st.encodeHeader(":method", "POST"), + EndHeaders: true, + }) + call := st.nextHandlerCall() + + call.do(func(w http.ResponseWriter, r *http.Request) { + // Wrap r.Body like some middleware might. + r.Body = http.MaxBytesReader(w, r.Body, 1000) + }) + + const size = InflowMinRefresh + st.writeData(1, false, make([]byte, size)) + st.writeRSTStream(1, ErrCodeCancel) + st.sync() + + call.exit() + st.wantWindowUpdate(0, size) + st.wantIdle() + }) +} + func TestServerIdleTimeout(t *testing.T) { synctest.Test(t, testServerIdleTimeout) } func testServerIdleTimeout(t *testing.T) { if testing.Short() { @@ -4376,9 +4644,21 @@ func testServerWindowUpdateOnBodyClose(t *testing.T) { } } + st.wantHeaders(wantHeader{ + streamID: 1, + endStream: true, + }) + // Writing data after the stream is reset immediately returns flow control credit. st.writeData(1, false, content[windowSize/2:]) - st.wantWindowUpdate(0, windowSize/2) + st.wantUnorderedFrames( + func(f *WindowUpdateFrame) bool { + return f.StreamID == 0 && f.Increment == windowSize/2 + }, + func(f *RSTStreamFrame) bool { + return f.StreamID == 1 && f.ErrCode == ErrCodeStreamClosed + }, + ) } func TestNoErrorLoggedOnPostAfterGOAWAY(t *testing.T) { @@ -5409,6 +5689,9 @@ func testServerSettingsFlowControlUpdateBeyondLimit(t *testing.T) { EndStream: false, // data coming EndHeaders: true, }) + call := st.nextHandlerCall() + http.NewResponseController(call.w).Flush() + st.wantFrameType(FrameHeaders) // Give this stream some additional flow control. const windowIncrease = 1000 @@ -5419,6 +5702,12 @@ func testServerSettingsFlowControlUpdateBeyondLimit(t *testing.T) { const maxWindowSize = (1 << 31) - 1 // RFC 9113, 6.9.1 const maxInitialWindowSize = maxWindowSize - windowIncrease st.writeSettings(Setting{SettingInitialWindowSize, maxInitialWindowSize + 1}) + st.wantSettingsAck() + + // We detect this condition lazily. Write something to the stream so we notice. + call.w.Write([]byte("hello")) + http.NewResponseController(call.w).Flush() + st.wantGoAway(1, ErrCodeFlowControl) } @@ -5437,6 +5726,9 @@ func testServerSettingsFlowControlUpdateWithinLimit(t *testing.T) { EndStream: false, // data coming EndHeaders: true, }) + call := st.nextHandlerCall() + http.NewResponseController(call.w).Flush() + st.wantFrameType(FrameHeaders) // Give this stream some additional flow control. const windowIncrease = 1000 @@ -5448,6 +5740,10 @@ func testServerSettingsFlowControlUpdateWithinLimit(t *testing.T) { const maxInitialWindowSize = maxWindowSize - windowIncrease st.writeSettings(Setting{SettingInitialWindowSize, maxInitialWindowSize}) st.wantSettingsAck() + + call.w.Write([]byte("hello")) + http.NewResponseController(call.w).Flush() + st.wantFrameType(FrameData) st.wantIdle() } diff --git a/src/net/http/internal/http2/transport.go b/src/net/http/internal/http2/transport.go index 9d63ed26c85cfc..d1e19583ba097a 100644 --- a/src/net/http/internal/http2/transport.go +++ b/src/net/http/internal/http2/transport.go @@ -144,11 +144,11 @@ type ClientConn struct { idleTimeout time.Duration // or 0 for never idleTimer *time.Timer - mu sync.Mutex // guards following - cond *sync.Cond // hold mu; broadcast on flow/closed changes - flow outflow // our conn-level flow control quota (cs.outflow is per stream) - inflow inflow // peer's conn-level flow control - doNotReuse bool // whether conn is marked to not be reused for any future requests + mu sync.Mutex // guards following + cond *sync.Cond // hold mu; broadcast on flow/closed changes + flow connOutflow // our conn-level flow control quota (cs.outflow is per stream) + inflow inflow // peer's conn-level flow control + doNotReuse bool // whether conn is marked to not be reused for any future requests closing bool closed bool closedOnIdle bool // true if conn was closed for idleness @@ -170,7 +170,6 @@ type ClientConn struct { maxConcurrentStreams uint32 peerMaxHeaderListSize uint64 peerMaxHeaderTableSize uint32 - initialWindowSize uint32 initialStreamRecvWindowSize int32 readIdleTimeout time.Duration pingTimeout time.Duration @@ -644,7 +643,6 @@ func (t *Transport) newClientConn(c net.Conn, singleUse bool, internalStateHook readerDone: make(chan struct{}), nextStreamID: 1, maxFrameSize: 16 << 10, // spec default - initialWindowSize: 65535, // spec default initialStreamRecvWindowSize: int32(conf.MaxReceiveBufferPerStream), maxConcurrentStreams: initialMaxConcurrentStreams, // "infinite", per spec. Use a smaller value until we have received server settings. strictMaxConcurrentStreams: conf.StrictMaxConcurrentRequests, @@ -669,7 +667,7 @@ func (t *Transport) newClientConn(c net.Conn, singleUse bool, internalStateHook } cc.cond = sync.NewCond(&cc.mu) - cc.flow.add(int32(initialWindowSize)) + cc.flow.init() // TODO: adjust this writer size to account for frame size + // MTU + crypto/tls record padding. @@ -1064,6 +1062,33 @@ func (cc *ClientConn) closeForError(err error) { cc.closeConn() } +func (cc *ClientConn) goAwayAndClose(code ErrCode) { + cc.mu.Lock() + closed := cc.closed + cc.closing = true + cc.closed = true + cc.mu.Unlock() + if closed { + return + } + if f := cc.fr.countError; f != nil { + f(fmt.Sprintf("conn_close_error_%s", code.stringToken())) + } + done := make(chan struct{}) + go func() { + defer close(done) + cc.wmu.Lock() + cc.fr.WriteGoAway(0, code, nil) + cc.bw.Flush() + cc.wmu.Unlock() + }() + select { + case <-done: + case <-time.After(250 * time.Millisecond): + } + cc.closeForError(fmt.Errorf("http2: closing connection with %v", code)) +} + // Close closes the client connection immediately. // // In-flight requests are interrupted. For a graceful shutdown, use Shutdown instead. @@ -1583,7 +1608,6 @@ func (cs *clientStream) cleanupWriteRequest(err error) { if bodyClosed != nil { <-bodyClosed } - if err != nil && cs.sentEndStream { // If the connection is closed immediately after the response is read, // we may be aborted before finishing up here. If the stream was closed @@ -1596,7 +1620,9 @@ func (cs *clientStream) cleanupWriteRequest(err error) { } if err != nil { cs.abortStream(err) // possibly redundant, but harmless - if cs.sentHeaders { + if ce, ok := err.(ConnectionError); ok { + cc.goAwayAndClose(ErrCode(ce)) + } else if cs.sentHeaders { if se, ok := err.(StreamError); ok { if se.Cause != errFromPeer { cc.writeStreamReset(cs.ID, se.Code, false, err) @@ -1922,8 +1948,12 @@ func (cs *clientStream) awaitFlowControl(maxBytes int) (taken int32, err error) return 0, errRequestCanceled default: } - if a := cs.flow.available(); a > 0 { - take := a + avail, ok := cs.flow.available() + if !ok { + return 0, ConnectionError(ErrCodeFlowControl) + } + if avail > 0 { + take := avail if int(take) > maxBytes { take = int32(maxBytes) // can't truncate int; take is int32 @@ -1984,8 +2014,7 @@ type resAndError struct { // requires cc.mu be held. func (cc *ClientConn) addStreamLocked(cs *clientStream) { - cs.flow.add(int32(cc.initialWindowSize)) - cs.flow.setConnFlow(&cc.flow) + cs.flow.conn = &cc.flow cs.inflow.init(cc.initialStreamRecvWindowSize) cs.ID = cc.nextStreamID cc.nextStreamID += 2 @@ -2390,18 +2419,42 @@ func (rl *clientConnReadLoop) handleResponse(cs *clientStream, f *MetaHeadersFra return nil, nil } + // Delete various headers that might mess up framing for HTTP/1. This is + // not a problem for HTTP/2, but someone might use HTTP/2 transport as a + // reverse proxy which forwards the response to an HTTP/1 client. Our + // HTTP/1 client transport will properly reject improper headers such as + // multiple conflicting Content-Length headers, but other implementations + // might not. + // TODO: just reject such responses? We deleted them for compatibility + // since this was done in a security fix (go.dev/issue/81115). However, + // rejecting them seems entirely reasonable and relatively safe. + + // Connection-specific header fields must not appear in an HTTP/2 message, + // and any message containing them is malformed. RFC 9113, Section 8.2.2. + for _, k := range connHeaders { + delete(res.Header, k) + } res.ContentLength = -1 - if clens := res.Header["Content-Length"]; len(clens) == 1 { - if cl, err := strconv.ParseUint(clens[0], 10, 63); err == nil { - res.ContentLength = int64(cl) + if clens, ok := res.Header["Content-Length"]; ok { + // Repeated Content-Length values may be collapsed into one only if + // they are identical per RFC 9110 Section 8.6. + // No need to trim whitespace, HTTP/2 header values must not have + // extraneous whitespace per RFC 9113 Section 8.2.1. + // Non-canonical headers are already rejected by our framer at this + // point. + conflicting := slices.ContainsFunc(clens[1:], func(clen string) bool { + return clen != clens[0] + }) + cl, err := strconv.ParseUint(clens[0], 10, 63) + if conflicting || err != nil { + delete(res.Header, "Content-Length") } else { - // TODO: care? unlike http/1, it won't mess up our framing, so it's - // more safe smuggling-wise to ignore. + res.Header["Content-Length"] = clens[:1] + res.ContentLength = int64(cl) } - } else if len(clens) > 1 { - // TODO: care? unlike http/1, it won't mess up our framing, so it's - // more safe smuggling-wise to ignore. - } else if f.StreamEnded() && !cs.isHead { + } + + if res.ContentLength < 0 && f.StreamEnded() && !cs.isHead { res.ContentLength = 0 } @@ -2724,6 +2777,10 @@ const ( func (rl *clientConnReadLoop) streamByID(id uint32, headerOrData bool) *clientStream { rl.cc.mu.Lock() defer rl.cc.mu.Unlock() + return rl.streamByIDLocked(id, headerOrData) +} + +func (rl *clientConnReadLoop) streamByIDLocked(id uint32, headerOrData bool) *clientStream { if headerOrData { // Work around an unfortunate gRPC behavior. // See comment on ClientConn.rstStreamPingsBlocked for details. @@ -2842,18 +2899,10 @@ func (rl *clientConnReadLoop) processSettingsNoWrite(f *SettingsFrame) error { case SettingMaxHeaderListSize: cc.peerMaxHeaderListSize = uint64(s.Val) case SettingInitialWindowSize: - // Adjust flow control of currently-open - // frames by the difference of the old initial - // window size and this one. - delta := int32(s.Val) - int32(cc.initialWindowSize) - for _, cs := range cc.streams { - if !cs.flow.add(delta) { - return ConnectionError(ErrCodeFlowControl) - } + if !cc.flow.changeInitialWindowSize(int64(s.Val)) { + return ConnectionError(ErrCodeFlowControl) } cc.cond.Broadcast() - - cc.initialWindowSize = s.Val case SettingHeaderTableSize: cc.henc.SetMaxDynamicTableSize(s.Val) cc.peerMaxHeaderTableSize = s.Val @@ -2895,29 +2944,30 @@ func (rl *clientConnReadLoop) processSettingsNoWrite(f *SettingsFrame) error { func (rl *clientConnReadLoop) processWindowUpdate(f *WindowUpdateFrame) error { cc := rl.cc - cs := rl.streamByID(f.StreamID, notHeaderOrDataFrame) - if f.StreamID != 0 && cs == nil { - return nil - } - cc.mu.Lock() defer cc.mu.Unlock() - - fl := &cc.flow - if cs != nil { - fl = &cs.flow - } - if !fl.add(int32(f.Increment)) { - // For stream, the sender sends RST_STREAM with an error code of FLOW_CONTROL_ERROR - if cs != nil { + if f.StreamID == 0 { + if !cc.flow.add(int32(f.Increment)) { + return ConnectionError(ErrCodeFlowControl) + } + } else { + cs := rl.streamByIDLocked(f.StreamID, notHeaderOrDataFrame) + if cs == nil { + return nil + } + if !cs.flow.add(int32(f.Increment)) { + if cs.flow.conn.flowErr { + // This is a lazily-detected connection-level flow control error. + return ConnectionError(ErrCodeFlowControl) + } + // For stream, the sender sends RST_STREAM with + // an error code of FLOW_CONTROL_ERROR. rl.endStreamErrorLocked(cs, StreamError{ StreamID: f.StreamID, Code: ErrCodeFlowControl, }) return nil } - - return ConnectionError(ErrCodeFlowControl) } cc.cond.Broadcast() return nil diff --git a/src/net/http/internal/http2/transport_test.go b/src/net/http/internal/http2/transport_test.go index 6cbd8cb676a431..8b1bba1acb1aab 100644 --- a/src/net/http/internal/http2/transport_test.go +++ b/src/net/http/internal/http2/transport_test.go @@ -17,6 +17,7 @@ import ( "fmt" "io" "log" + "maps" "math/rand" "net" "net/http" @@ -27,6 +28,7 @@ import ( "os" "reflect" "runtime" + "slices" "sort" "strconv" "strings" @@ -1748,7 +1750,8 @@ func testTransportSettingsFlowControlUpdateBeyondLimit(t *testing.T) { tc := newTestClientConn(t) tc.greet() - req, _ := http.NewRequest("GET", "https://dummy.tld/", nil) + body := tc.newRequestBody() + req, _ := http.NewRequest("GET", "https://dummy.tld/", body) rt := tc.roundTrip(req) tc.wantFrameType(FrameHeaders) @@ -1761,6 +1764,12 @@ func testTransportSettingsFlowControlUpdateBeyondLimit(t *testing.T) { const maxWindowSize = (1 << 31) - 1 // RFC 9113, 6.9.1 const maxInitialWindowSize = maxWindowSize - windowIncrease tc.writeSettings(Setting{SettingInitialWindowSize, maxInitialWindowSize + 1}) + tc.wantSettingsAck() + + // We detect this condition lazily. Write something to the stream so we notice. + body.writeBytes(1) + body.closeWithError(io.EOF) + tc.wantGoAway(0, ErrCodeFlowControl) } @@ -1773,7 +1782,8 @@ func testTransportSettingsFlowControlUpdateWithinLimit(t *testing.T) { tc := newTestClientConn(t) tc.greet() - req, _ := http.NewRequest("GET", "https://dummy.tld/", nil) + body := tc.newRequestBody() + req, _ := http.NewRequest("GET", "https://dummy.tld/", body) rt := tc.roundTrip(req) tc.wantFrameType(FrameHeaders) @@ -1787,6 +1797,15 @@ func testTransportSettingsFlowControlUpdateWithinLimit(t *testing.T) { const maxInitialWindowSize = maxWindowSize - windowIncrease tc.writeSettings(Setting{SettingInitialWindowSize, maxInitialWindowSize}) tc.wantSettingsAck() + + body.writeBytes(1) + body.closeWithError(io.EOF) + tc.wantData(wantData{ + streamID: 1, + endStream: true, + size: 1, + multiple: true, + }) tc.wantIdle() } @@ -2010,6 +2029,161 @@ func TestTransportRejectsContentLengthWithSign(t *testing.T) { } } +// TestTransportResponseContentLength checks that a Content-Length we cannot +// validate is dropped from the response rather than passed on to the caller, +// who may be forwarding it to an HTTP/1 endpoint. +func TestTransportResponseContentLength(t *testing.T) { + tests := []struct { + name string + clValues []string + wantLen int // -1 means the header is expected to be dropped. + }{ + { + name: "single value", + clValues: []string{"3"}, + wantLen: 3, + }, + { + name: "identical duplicate values", + clValues: []string{"3", "3", "3"}, + wantLen: 3, + }, + { + name: "different duplicate values", + clValues: []string{"3", "1", "3"}, + wantLen: -1, + }, + { + name: "extraneous whitespace", + clValues: []string{" 3"}, + wantLen: -1, + }, + { + name: "identical duplicate values with extraneous whitespace", + clValues: []string{"3", "3", " 3"}, + wantLen: -1, + }, + { + name: "plus sign", + clValues: []string{"+3"}, + wantLen: -1, + }, + { + name: "non-numeric", + clValues: []string{"abc"}, + wantLen: -1, + }, + { + name: "empty value", + clValues: []string{""}, + wantLen: -1, + }, + { + name: "no header", + wantLen: -1, + }, + } + for _, tt := range tests { + synctestSubtest(t, tt.name, func(t *testing.T) { + tc := newTestClientConn(t) + tc.greet() + + req, _ := http.NewRequest("GET", "https://dummy.tld/", nil) + rt := tc.roundTrip(req) + + headers := []string{":status", "200"} + for _, val := range tt.clValues { + headers = append(headers, "content-length", val) + } + tc.wantFrameType(FrameHeaders) + tc.writeHeaders(HeadersFrameParam{ + StreamID: rt.streamID(), + EndHeaders: true, + BlockFragment: tc.makeHeaderBlockFragment(headers...), + }) + body := slices.Repeat([]byte("a"), max(tt.wantLen, 1)) + tc.writeData(rt.streamID(), true, body) + + res := rt.response() + rt.wantBody(body) + + if res.ContentLength != int64(tt.wantLen) { + t.Errorf("got ContentLength = %d, want %d", res.ContentLength, tt.wantLen) + } + var wantHeader []string + if tt.wantLen >= 0 { + wantHeader = []string{strconv.FormatInt(int64(tt.wantLen), 10)} + } + if got := res.Header["Content-Length"]; !slices.Equal(got, wantHeader) { + t.Errorf("got Header[%q] = %q, want %q", "Content-Length", got, wantHeader) + } + }) + } +} + +// TestTransportResponseConnHeaders checks that connection-related headers, +// which are not valid in HTTP/2 and which an HTTP/1 endpoint may use for +// framing, are dropped from the response. +func TestTransportResponseConnHeaders(t *testing.T) { + tests := []struct { + name string + fields []string + wantHeader http.Header + }{ + { + name: "unaffected header", + fields: []string{"content-type", "text/plain"}, + wantHeader: http.Header{"Content-Type": {"text/plain"}}, + }, + { + name: "transfer-encoding", + fields: []string{"transfer-encoding", "chunked"}, + }, + { + name: "transfer-encoding alongside content-length", + fields: []string{"content-length", "-1", "transfer-encoding", "chunked"}, + }, + { + name: "connection and keep-alive", + fields: []string{"connection", "keep-alive", "keep-alive", "timeout=5"}, + }, + { + name: "proxy-connection", + fields: []string{"proxy-connection", "keep-alive"}, + }, + { + name: "upgrade", + fields: []string{"upgrade", "websocket"}, + }, + } + for _, tt := range tests { + synctestSubtest(t, tt.name, func(t *testing.T) { + tc := newTestClientConn(t) + tc.greet() + + req, _ := http.NewRequest("GET", "https://dummy.tld/", nil) + rt := tc.roundTrip(req) + + headers := []string{":status", "200"} + headers = append(headers, tt.fields...) + tc.wantFrameType(FrameHeaders) + tc.writeHeaders(HeadersFrameParam{ + StreamID: rt.streamID(), + EndHeaders: true, + BlockFragment: tc.makeHeaderBlockFragment(headers...), + }) + tc.writeData(rt.streamID(), true, []byte("body")) + + res := rt.response() + rt.wantBody([]byte("body")) + + if !maps.EqualFunc(res.Header, tt.wantHeader, slices.Equal) { + t.Errorf("got Header = %q, want %q", res.Header, tt.wantHeader) + } + }) + } +} + // golang.org/issue/14048 // golang.org/issue/64766 func TestTransportFailsOnInvalidHeadersAndTrailers(t *testing.T) { diff --git a/src/net/http/internal/http2/writesched.go b/src/net/http/internal/http2/writesched.go index 883c4ddc0d07e3..b4a45368f0c3d0 100644 --- a/src/net/http/internal/http2/writesched.go +++ b/src/net/http/internal/http2/writesched.go @@ -115,7 +115,11 @@ func (wr FrameWriteRequest) Consume(n int32) (FrameWriteRequest, FrameWriteReque } // Might need to split after applying limits. - allowed := min(n, wr.stream.flow.available()) + avail, ok := wr.stream.flow.available() + if !ok { + return empty, empty, 0 + } + allowed := min(n, avail) if wr.stream.sc.maxFrameSize < allowed { allowed = wr.stream.sc.maxFrameSize } diff --git a/src/net/http/internal/http2/writesched_test.go b/src/net/http/internal/http2/writesched_test.go index 6dbd7f0adc7203..0eb156a3e7823a 100644 --- a/src/net/http/internal/http2/writesched_test.go +++ b/src/net/http/internal/http2/writesched_test.go @@ -76,6 +76,7 @@ func TestFrameWriteRequestWithData(t *testing.T) { id: 1, sc: &serverConn{maxFrameSize: 16}, } + st.flow.conn = &connOutflow{} const size = 32 wr := FrameWriteRequest{&writeData{st.id, make([]byte, size), true}, st, make(chan error)} if got, want := wr.DataSize(), size; got != want { @@ -113,6 +114,8 @@ func TestFrameWriteRequestData(t *testing.T) { id: 1, sc: &serverConn{maxFrameSize: 16}, } + st.flow.conn = &connOutflow{} + st.flow.conn.add(maxFlowWindow) // conn-level flow is large const size = 32 wr := FrameWriteRequest{&writeData{st.id, make([]byte, size), true}, st, make(chan error)} if got, want := wr.DataSize(), size; got != want { diff --git a/src/net/http/range_test.go b/src/net/http/range_test.go index 114987ed2c6984..6cfef250d52c13 100644 --- a/src/net/http/range_test.go +++ b/src/net/http/range_test.go @@ -5,6 +5,8 @@ package http import ( + "runtime/metrics" + "strings" "testing" ) @@ -77,3 +79,100 @@ func TestParseRange(t *testing.T) { } } } + +func TestParseRangeLimit(t *testing.T) { + for _, tc := range []struct { + name string + godebug string + numRanges int + want int + wantNonDef bool + }{ + { + name: "default limit not exceeded", + numRanges: defaultMaxContentRanges, + want: defaultMaxContentRanges, + }, + { + name: "default limit exceeded", + numRanges: defaultMaxContentRanges + 1, + want: 0, + }, + { + name: "small limit exceeded", + godebug: "httpservecontentmaxranges=10", + numRanges: 11, + want: 0, + wantNonDef: true, + }, + { + name: "small limit not exceeded", + godebug: "httpservecontentmaxranges=10", + numRanges: 10, + want: 10, + }, + { + name: "disabled limit", + godebug: "httpservecontentmaxranges=0", + numRanges: 2 * defaultMaxContentRanges, + want: 2 * defaultMaxContentRanges, + wantNonDef: true, + }, + { + name: "large limit exceeded", + godebug: "httpservecontentmaxranges=300", + numRanges: 301, + want: 0, + }, + { + name: "large limit not exceeded", + godebug: "httpservecontentmaxranges=300", + numRanges: 300, + want: 300, + wantNonDef: true, + }, + { + name: "large limit not exceeded by small input", + godebug: "httpservecontentmaxranges=300", + numRanges: 10, + want: 10, + }, + { + name: "invalid limit negative", + godebug: "httpservecontentmaxranges=-5", + numRanges: defaultMaxContentRanges + 1, + want: 0, + }, + { + name: "invalid limit non-numeric", + godebug: "httpservecontentmaxranges=abc", + numRanges: defaultMaxContentRanges + 1, + want: 0, + }, + } { + t.Run(tc.name, func(t *testing.T) { + t.Setenv("GODEBUG", tc.godebug) + var m [1]metrics.Sample + m[0].Name = "/godebug/non-default-behavior/httpservecontentmaxranges:events" + metrics.Read(m[:]) + before := m[0].Value.Uint64() + + rangeHeader := "bytes=" + strings.Repeat("0-0,", tc.numRanges-1) + "0-0" + ranges, err := parseRange(rangeHeader, 10_000_000) + if err != nil { + t.Fatalf("parseRange(%q): %v", rangeHeader, err) + } + if got := len(ranges); got != tc.want { + t.Errorf("len(ranges) = %v, want %v", got, tc.want) + } + + metrics.Read(m[:]) + after := m[0].Value.Uint64() + if tc.wantNonDef && after <= before { + t.Errorf("metric did not increment: before=%d, after=%d", before, after) + } else if !tc.wantNonDef && after != before { + t.Errorf("metric unexpectedly incremented: before=%d, after=%d", before, after) + } + }) + } +} diff --git a/src/net/http/serve_test.go b/src/net/http/serve_test.go index 45b861969aeba7..f9ceb6ca0303c5 100644 --- a/src/net/http/serve_test.go +++ b/src/net/http/serve_test.go @@ -3511,6 +3511,24 @@ func testRequestHeaderValueCountLimit(t *testing.T, mode testMode) { }, wantStatus: 431, }, + { + // Comma separated Trailer values are counted as multiple, because + // each value becomes its own field / a key in Request.Trailer. + // This is different from TestRequestTrailerHeaderValueCountLimit + // which tests the actual sending of the trailer, this just tests + // the Trailer header declaration. + name: "comma separated trailer values count as multiple", + limit: 15, + setup: func(req *Request) { + req.Body = NoBody + req.TransferEncoding = []string{"chunked"} + req.Trailer = make(Header) + for i := range 16 { + req.Trailer[fmt.Sprintf("X-Trailer-%d", i)] = nil + } + }, + wantStatus: 431, + }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { @@ -8317,3 +8335,98 @@ func TestServerIdleKeepAliveNonstandardTimeoutError(t *testing.T) { }) }) } + +func TestServerCONNECTSuccess(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const code = 200 + body := []byte("body") + srv := &Server{ + Handler: HandlerFunc(func(w ResponseWriter, req *Request) { + w.WriteHeader(code) + w.Write(body) + }), + } + l := fakeNetListen() + defer l.Close() + go srv.Serve(l) + + conn := l.connect() + defer conn.Close() + io.WriteString(conn, "CONNECT backend.example.tld:80 HTTP/1.1\r\nHost: example.tld\r\n\r\n") + + bufr := bufio.NewReader(conn) + synctest.Wait() + resp, err := ReadResponse(bufr, nil) + if err != nil { + t.Fatal(err) + } + if resp.StatusCode != code { + t.Errorf("response status = %v, want %v", resp.StatusCode, code) + } + if resp.ContentLength != -1 { + t.Errorf("Content-Length: %v; want absent", resp.ContentLength) + } + if len(resp.TransferEncoding) > 0 { + t.Errorf("Transfer-Encoding: %q; want absent", resp.TransferEncoding) + } + for _, h := range []string{"Content-Length", "Transfer-Encoding"} { + if got, ok := resp.Header[h]; ok { + t.Errorf("response header %q = %q; want absent", h, got) + } + } + got := make([]byte, len(body)) + if _, err := io.ReadFull(bufr, got); err != nil || !bytes.Equal(got, body) { + t.Fatalf("want bytes %q, got %q (err %v)", body, got, err) + } + if !conn.IsClosedByPeer() { + t.Errorf("connection not closed by peer") + } + if got, err := bufr.Peek(32); len(got) != 0 || err != io.EOF { + t.Errorf("read from conn: %q, %v; expect conn to be closed", got, err) + } + }) +} + +func TestServerCONNECTFailure(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const code = 409 + body := []byte("body") + srv := &Server{ + Handler: HandlerFunc(func(w ResponseWriter, req *Request) { + w.Header().Set("Content-Length", strconv.Itoa(len(body))) + w.WriteHeader(code) + w.Write(body) + }), + } + l := fakeNetListen() + defer l.Close() + go srv.Serve(l) + + conn := l.connect() + defer conn.Close() + io.WriteString(conn, "CONNECT backend.example.tld:80 HTTP/1.1\r\nHost: example.tld\r\n\r\n") + + bufr := bufio.NewReader(conn) + synctest.Wait() + resp, err := ReadResponse(bufr, nil) + if err != nil { + t.Fatal(err) + } + if resp.StatusCode != code { + t.Errorf("response status = %v, want %v", resp.StatusCode, code) + } + if !resp.Close { + t.Errorf("Connection: close not set; want it to be") + } + got := make([]byte, len(body)) + if _, err := io.ReadFull(bufr, got); err != nil || !bytes.Equal(got, body) { + t.Fatalf("want bytes %q, got %q (err %v)", body, got, err) + } + if !conn.IsClosedByPeer() { + t.Errorf("connection not closed by peer") + } + if got, err := bufr.Peek(32); len(got) != 0 || err != io.EOF { + t.Errorf("read from conn: %q, %v; expect conn to be closed", got, err) + } + }) +} diff --git a/src/net/http/server.go b/src/net/http/server.go index 3c73abf2aeb14a..5baa45c3201f95 100644 --- a/src/net/http/server.go +++ b/src/net/http/server.go @@ -1395,7 +1395,11 @@ func (cw *chunkWriter) writeHeader(p []byte) { w := cw.res keepAlivesEnabled := w.conn.server.doKeepAlives() - isHEAD := w.req.Method == "HEAD" + + // Consult w.conn.lastMethod instead of w.req.Method, + // just in case a middleware layer modified w.req. + isHEAD := w.conn.lastMethod == "HEAD" + isCONNECT := w.conn.lastMethod == "CONNECT" // header is written out to w.conn.buf below. Depending on the // state of the handler, we either own the map or not. If we @@ -1507,6 +1511,11 @@ func (cw *chunkWriter) writeHeader(p []byte) { w.closeAfterReply = true } + if isCONNECT { + // Don't reuse a connection after a CONNECT, even if we reject it. + w.closeAfterReply = true + } + // We do this by default because there are a number of clients that // send a full request before starting to read the response, and they // can deadlock if we start writing the response with unconsumed body @@ -1581,7 +1590,13 @@ func (cw *chunkWriter) writeHeader(p []byte) { hasCL = false } - if w.req.Method == "HEAD" || !bodyAllowedForStatus(code) || code == StatusNoContent { + isSuccessfulCONNECT := isCONNECT && code >= 200 && code < 300 + if isSuccessfulCONNECT { + // Tunnel established, connection is no longer HTTP. + delHeader("Transfer-Encoding") + delHeader("Content-Length") + setHeader.contentLength = nil + } else if isHEAD || !bodyAllowedForStatus(code) || code == StatusNoContent { // Response has no body. delHeader("Transfer-Encoding") } else if hasCL { @@ -1631,8 +1646,13 @@ func (cw *chunkWriter) writeHeader(p []byte) { !isProtocolSwitchResponse(w.status, header) if delConnectionHeader { delHeader("Connection") - if w.req.ProtoAtLeast(1, 1) { + // Don't set Connection: close on a 2xx CONNECT response, + // even though we will close the connection if the handler doesn't hijack it. + // If the handler does hijack the connection, the Connection: close is confusing. + if w.req.ProtoAtLeast(1, 1) && !isSuccessfulCONNECT { setHeader.connection = "close" + } else { + setHeader.connection = "" } } @@ -3235,9 +3255,12 @@ type Server struct { // MaxHeaderValueCount controls the maximum number of header // values that the server is willing to parse from a request. // If zero, DefaultMaxHeaderValueCount is used. - // Note that comma-separated values in a single header line are - // counted once, while values sent as multiple header lines are - // counted multiple times. + // Comma-separated values in a single header line are counted + // once, while values sent as multiple header lines are + // counted multiple times. An exception to this is the Trailer + // header, whose comma-separated values are counted separately, + // as each of them is expected to be received later as an + // individual trailer header line. MaxHeaderValueCount int // TLSNextProto optionally specifies a function to take over diff --git a/src/net/http/transfer.go b/src/net/http/transfer.go index 019ba7d2fdaadb..997b9a00404364 100644 --- a/src/net/http/transfer.go +++ b/src/net/http/transfer.go @@ -550,7 +550,7 @@ func readTransfer(msg any, r *bufio.Reader, maxTrailerHeaders int64) (err error) } // Trailer - t.Trailer, err = fixTrailer(t.Header, t.Chunked) + t.Trailer, err = fixTrailer(t.Header, t.Chunked, maxTrailerHeaders) if err != nil { return err } @@ -772,7 +772,7 @@ func shouldClose(major, minor int, header Header, removeCloseHeader bool) bool { } // Parse the trailer header. -func fixTrailer(header Header, chunked bool) (Header, error) { +func fixTrailer(header Header, chunked bool, maxHeaders int64) (Header, error) { vv, ok := header["Trailer"] if !ok { return nil, nil @@ -788,6 +788,12 @@ func fixTrailer(header Header, chunked bool) (Header, error) { return nil, nil } header.Del("Trailer") + for _, v := range vv { + maxHeaders -= int64(strings.Count(v, ",") + 1) + if maxHeaders < 0 { + return nil, errTooLarge + } + } trailer := make(Header) var err error @@ -824,11 +830,12 @@ type body struct { doEarlyClose bool // whether Close should stop early maxTrailerHeaders int64 // how many trailer header values are allowed - mu sync.Mutex // guards following, and calls to Read and Close - sawEOF bool - closed bool - earlyClose bool // Close called and we didn't read to the end of src - onHitEOF func() // if non-nil, func to call when EOF is Read + mu sync.Mutex // guards following, and calls to Read and Close + sawEOF bool + closed bool + earlyClose bool // Close called and we didn't read to the end of src + dropTrailer bool // if true, do not populate hdr.Trailer + onHitEOF func() // if non-nil, func to call when EOF is Read } // ErrBodyReadAfterClose is returned when reading a [Request] or [Response] @@ -953,6 +960,13 @@ func (b *body) readTrailer() error { } return err } + // When we are automatically draining a response body, let the trailer + // still be parsed above (so connection can be reused). However, do not + // actually populate b.hdr.Trailer. Doing so is racy as we do not own b.hdr + // anymore when automatic draining occurs. + if b.dropTrailer { + return nil + } switch rr := b.hdr.(type) { case *Request: mergeSetHeader(&rr.Trailer, Header(hdr)) @@ -970,6 +984,12 @@ func mergeSetHeader(dst *Header, src Header) { maps.Copy(*dst, src) } +func (b *body) discardTrailer() { + b.mu.Lock() + defer b.mu.Unlock() + b.dropTrailer = true +} + // unreadDataSizeLocked returns the number of bytes of unread input. // It returns -1 if unknown. // b.mu must be held. diff --git a/src/net/http/transport.go b/src/net/http/transport.go index 90db79b4882c0f..4104d3632d3ef9 100644 --- a/src/net/http/transport.go +++ b/src/net/http/transport.go @@ -2430,10 +2430,16 @@ const maxPostCloseReadBytes = 256 << 10 // has been closed. const maxPostCloseReadTime = 50 * time.Millisecond -func maybeDrainBody(body io.Reader) bool { +func maybeDrainBody(r io.Reader) bool { drainedCh := make(chan bool, 1) go func() { - if _, err := io.CopyN(io.Discard, body, maxPostCloseReadBytes+1); err == io.EOF { + // When we drain the body and (hopefully) reach EOF, we might + // potentially need to deal with trailers. Make sure they are discarded + // so the connection can actually be reused. + if b, ok := r.(*body); ok { + b.discardTrailer() + } + if _, err := io.CopyN(io.Discard, r, maxPostCloseReadBytes+1); err == io.EOF { drainedCh <- true } else { drainedCh <- false @@ -2447,6 +2453,10 @@ func maybeDrainBody(body io.Reader) bool { } } +// errClosedEarly is an internal-only error used to indicate that a response body +// was closed early prior to EOF. +var errClosedEarly = errors.New("net/http: response body closed early") + func (pc *persistConn) readLoop() { closeErr := errReadLoopExiting // default value, if not changed below defer func() { @@ -2526,12 +2536,17 @@ func (pc *persistConn) readLoop() { pc.mu.Unlock() bodyWritable := resp.bodyIsWritable() + isConnect := rc.treq.Request.Method == "CONNECT" hasBody := rc.treq.Request.Method != "HEAD" && resp.ContentLength != 0 - if resp.Close || rc.treq.Request.Close || resp.StatusCode <= 199 || bodyWritable { + if resp.Close || rc.treq.Request.Close || resp.StatusCode <= 199 || bodyWritable || isConnect { // Don't do keep-alive on error if either party requested a close // or we get an unexpected informational (1xx) response. // StatusCode 100 is already handled above. + // + // Don't do keep-alive after sending a CONNECT request. + // Only a 2xx response converts the connection into a tunnel, + // but for safety we'll drop the connection even after getting a non-2xx. alive = false } @@ -2565,19 +2580,17 @@ func (pc *persistConn) readLoop() { continue } - waitForBodyRead := make(chan bool, 2) + waitForBodyRead := make(chan error, 1) body := &bodyEOFSignal{ body: resp.Body, earlyCloseFn: func() error { - waitForBodyRead <- false + waitForBodyRead <- errClosedEarly <-eofc // will be closed by deferred call at the end of the function return nil - }, fn: func(err error) error { - isEOF := err == io.EOF - waitForBodyRead <- isEOF - if isEOF { + waitForBodyRead <- err + if err == io.EOF { <-eofc // see comment above eofc declaration } else if err != nil { if cerr := pc.canceled(); cerr != nil { @@ -2607,19 +2620,29 @@ func (pc *persistConn) readLoop() { // the bufio.Reader, wait for the caller goroutine to finish // reading the response body. (or for cancellation or death) select { - case bodyEOF := <-waitForBodyRead: - tryDrain := !bodyEOF && resp.ContentLength <= maxPostCloseReadBytes - if tryDrain { - eofc <- struct{}{} - bodyEOF = maybeDrainBody(body.body) + case err := <-waitForBodyRead: + tryPutIdle := func() { + alive = alive && + !pc.sawEOF && + pc.wroteRequest() && + tryPutIdleConn(rc.treq) } - alive = alive && - bodyEOF && - !pc.sawEOF && - pc.wroteRequest() && - tryPutIdleConn(rc.treq) - if !tryDrain && bodyEOF { + switch err { + case io.EOF: + tryPutIdle() + eofc <- struct{}{} + case errClosedEarly: + // Read resp before signaling eofc: the send lets the caller's + // Close return, and resp belongs to the caller after that. + tryDrain := alive && resp.ContentLength <= maxPostCloseReadBytes eofc <- struct{}{} + if tryDrain && maybeDrainBody(body.body) { + tryPutIdle() + } else { + alive = false + } + default: + alive = false } case <-rc.treq.ctx.Done(): alive = false @@ -3215,16 +3238,16 @@ func canonicalAddr(url *url.URL) string { // once, right before its final (error-producing) Read or Close call // returns. fn should return the new error to return from Read or Close. // -// If earlyCloseFn is non-nil and Close is called before io.EOF is -// seen, earlyCloseFn is called instead of fn, and its return value is +// If earlyCloseFn is non-nil and Close is called before any final error from +// Read is seen, earlyCloseFn is called instead of fn, and its return value is // the return value from Close. type bodyEOFSignal struct { body io.ReadCloser mu sync.Mutex // guards following 4 fields closed bool // whether Close has been called rerr error // sticky Read error - fn func(error) error // err will be nil on Read io.EOF - earlyCloseFn func() error // optional alt Close func used if io.EOF not seen + fn func(error) error // called on final body.Read non-nil error (or body.Close if earlyCloseFn is not run) + earlyCloseFn func() error // called if body.Close is called before body.Read ever returns a non-nil error } var errReadOnClosedResBody = errors.New("http: read on closed response body") @@ -3260,8 +3283,16 @@ func (es *bodyEOFSignal) Close() error { return nil } es.closed = true - if es.earlyCloseFn != nil && es.rerr != io.EOF { - return es.earlyCloseFn() + if es.earlyCloseFn != nil && es.rerr == nil { + earlyCloseFn := es.earlyCloseFn + es.earlyCloseFn = nil + es.fn = nil + return earlyCloseFn() + } + if es.rerr != nil && es.rerr != io.EOF { + // Read already returned this error and readLoop gave up the + // connection. Draining would only read the same error again. + return nil } err := es.body.Close() return es.condfn(err) @@ -3272,9 +3303,10 @@ func (es *bodyEOFSignal) condfn(err error) error { if es.fn == nil { return err } - err = es.fn(err) + fn := es.fn es.fn = nil - return err + es.earlyCloseFn = nil + return fn(err) } // gzipReader wraps a response body so it can lazily diff --git a/src/net/http/transport_test.go b/src/net/http/transport_test.go index c70b5855184b2d..e44d1cb9784d2b 100644 --- a/src/net/http/transport_test.go +++ b/src/net/http/transport_test.go @@ -6496,6 +6496,40 @@ func testTransportCONNECTBidi(t *testing.T, mode testMode) { } } +func TestTransportCONNECTRejected(t *testing.T) { + runSynctest(t, testTransportCONNECTRejected, []testMode{http1Mode}) +} +func testTransportCONNECTRejected(t *testing.T, mode testMode) { + tt := newHTTP1TransportTest(t) + + sentReq := &Request{ + Method: "CONNECT", + URL: &url.URL{ + Scheme: "http", + Opaque: "backend.example.tld:80", + Host: "proxy.example.tld", + }, + Host: "proxy.example.tld", + Header: make(Header), + } + rt := tt.roundTrip(sentReq) + + dial := tt.wantDial("tcp", "proxy.example.tld:80") + conn := dial.connect() + recvReq := conn.readRequest() + if got, want := recvReq.URL.Path, sentReq.URL.Path; got != want { + t.Fatalf("read request path %q, want %q", got, want) + } + + conn.writeMessage( + "HTTP/1.1 405 We Have No Connections Today", + "Content-Length: 0", + "", + ) + rt.wantStatus(405) + conn.wantClosed() +} + func TestTransportRequestReplayable(t *testing.T) { someBody := io.NopCloser(strings.NewReader("")) tests := []struct { @@ -7494,6 +7528,199 @@ func TestTransportReqCancelerCleanupOnRequestBodyWriteError(t *testing.T) { }) } +func TestTransportResponseBodyDrainReadAndClose(t *testing.T) { + tests := []struct { + name string + read bool + closeEarly bool + closeAfterRead bool + serverTruncate bool + wantReadErr error + wantCloseErr error + wantReuse bool + }{ + // go.dev/issue/81404. + { + name: "concurrent early close and read to eof", + read: true, + closeEarly: true, + wantReadErr: io.EOF, + wantReuse: true, + }, + { + name: "concurrent early close and unexpected read error", + read: true, + closeEarly: true, + serverTruncate: true, + wantReadErr: io.ErrUnexpectedEOF, + wantReuse: false, + }, + { + name: "unexpected read error without any close", + read: true, + serverTruncate: true, + wantReadErr: io.ErrUnexpectedEOF, + wantReuse: false, + }, + { + name: "unexpected read error followed by close", + read: true, + closeAfterRead: true, + serverTruncate: true, + wantReadErr: io.ErrUnexpectedEOF, + // golang.org/issue/81511: Close does not repeat the Read error. + wantCloseErr: nil, + wantReuse: false, + }, + { + name: "early close without any read", + read: false, + closeEarly: true, + wantReuse: true, + }, + { + name: "read all then close", + read: true, + closeAfterRead: true, + wantReadErr: io.EOF, + wantReuse: true, + }, + } + + synctestSubtest := func(t *testing.T, name string, f func(*testing.T)) { + t.Helper() + t.Run(name, func(t *testing.T) { + t.Helper() + synctest.Test(t, f) + }) + } + for _, tc := range tests { + synctestSubtest(t, tc.name, func(t *testing.T) { + tt := newHTTP1TransportTest(t) + req, _ := NewRequest("GET", "http://example.tld/", nil) + rt := tt.roundTrip(req) + conn := tt.wantDial("tcp", "example.tld:80").connect() + conn.readRequest() + conn.writeMessage( + "HTTP/1.1 200 OK", + "Transfer-Encoding: chunked", + "", + ) + res := rt.response() + + readErr := make(chan error, 1) + if tc.read { + go func() { + var buf [1]byte + _, err := res.Body.Read(buf[:]) + readErr <- err + }() + // Wait for the Read to block on chunk data. + synctest.Wait() + } + + if tc.closeEarly { + if err := res.Body.Close(); !errors.Is(err, tc.wantCloseErr) { + t.Fatalf("Close = %v, want %v", err, tc.wantCloseErr) + } + } + + if tc.serverTruncate { + conn.conn.CloseWrite() + } else { + conn.writeMessage( + "0", + "", + ) + } + + if tc.read { + synctest.Wait() + if err := <-readErr; !errors.Is(err, tc.wantReadErr) { + t.Fatalf("Read = %v, want %v", err, tc.wantReadErr) + } + } + + if tc.closeAfterRead { + if err := res.Body.Close(); !errors.Is(err, tc.wantCloseErr) { + t.Fatalf("Close = %v, want %v", err, tc.wantCloseErr) + } + } + + if !tc.wantReuse { + conn.wantClosed() + } else { + synctest.Wait() + req, _ = NewRequest("GET", "http://example.tld/next", nil) + rt = tt.roundTrip(req) + if got := conn.readRequest().URL.Path; got != "/next" { + t.Fatalf("request path = %q, want /next", got) + } + conn.writeMessage( + "HTTP/1.1 200 OK", + "Content-Length: 0", + "", + ) + rt.wantStatus(200) + } + }) + } +} + +// When a response body is automatically drained, it should read the trailer to +// make sure the connection is reusable. Additionally, it should not populate +// the Response.Trailer: by the time Close returns, Response belongs to the +// caller and cannot be written to anymore by the draining goroutine. +func TestTransportResponseBodyDrainDropsTrailers(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + tt := newHTTP1TransportTest(t) + req, _ := NewRequest("GET", "http://example.tld/", nil) + rt := tt.roundTrip(req) + conn := tt.wantDial("tcp", "example.tld:80").connect() + conn.readRequest() + conn.writeMessage( + "HTTP/1.1 200 OK", + "Transfer-Encoding: chunked", + "Trailer: X-Test", + "", + "5", + "hello", + ) + res := rt.response() + + if err := res.Body.Close(); err != nil { + t.Fatalf("Close = %v, want nil", err) + } + conn.writeMessage( + "0", + "X-Test: v", + "", + ) + synctest.Wait() + + // Trailer must not be populated, to avoid racy write, and to be + // consistent with how trailers are normally empty when one does not + // read to EOF. + if got := res.Trailer.Get("X-Test"); got != "" { + t.Error("drained response body unexpectedly has populated trailers") + } + + // The trailer must still have been consumed from the connection, or + // the connection would not be reusable. + req, _ = NewRequest("GET", "http://example.tld/next", nil) + rt = tt.roundTrip(req) + if got := conn.readRequest().URL.Path; got != "/next" { + t.Fatalf("request path = %q, want /next", got) + } + conn.writeMessage( + "HTTP/1.1 200 OK", + "Content-Length: 0", + "", + ) + rt.wantStatus(200) + }) +} + func TestValidateClientRequestTrailers(t *testing.T) { run(t, testValidateClientRequestTrailers) } @@ -7858,3 +8085,22 @@ func testIssue61474(t *testing.T, mode testMode) { }) } } + +// After Body.Close returns, the Response belongs to the caller. readLoop +// used to read resp.ContentLength after letting Close return; run with -race. +func TestTransportResponseWriteAfterEarlyClose(t *testing.T) { + run(t, testTransportResponseWriteAfterEarlyClose, []testMode{http1Mode}) +} +func testTransportResponseWriteAfterEarlyClose(t *testing.T, mode testMode) { + cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) { + io.WriteString(w, "hello") + })) + res, err := cst.c.Get(cst.ts.URL) + if err != nil { + t.Fatal(err) + } + if err := res.Body.Close(); err != nil { + t.Fatal(err) + } + res.ContentLength = 0 +} diff --git a/src/net/textproto/reader.go b/src/net/textproto/reader.go index b9ec4654db5883..53c65806a9ee4f 100644 --- a/src/net/textproto/reader.go +++ b/src/net/textproto/reader.go @@ -180,7 +180,7 @@ func (r *Reader) readContinuedLineSlice(lim int64, validateFirstLine func([]byte } line, err := r.readLineSlice(lim - int64(len(r.buf))) if err != nil { - break + return nil, err } r.buf = append(r.buf, trim(line)...) } @@ -521,6 +521,10 @@ func (r *Reader) ReadMIMEHeader() (MIMEHeader, error) { // readMIMEHeader is a version of ReadMIMEHeader which takes a limit on the header size. // It is called by the mime/multipart and net/http package. func readMIMEHeader(r *Reader, maxMemory, maxHeaders int64) (MIMEHeader, error) { + if maxMemory < 0 { + return nil, errMessageTooLarge + } + // Avoid lots of small slice allocations later by allocating one // large one ahead of time which we'll cut up into smaller // slices. If this isn't big enough later, we allocate small ones. @@ -535,12 +539,6 @@ func readMIMEHeader(r *Reader, maxMemory, maxHeaders int64) (MIMEHeader, error) m := make(MIMEHeader, hint) - // Account for 400 bytes of overhead for the MIMEHeader, plus 200 bytes per entry. - // Benchmarking map creation as of go1.20, a one-entry MIMEHeader is 416 bytes and large - // MIMEHeaders average about 200 bytes per entry. - maxMemory -= 400 - const mapEntryOverhead = 200 - // The first line cannot start with a leading space. if buf, err := r.R.Peek(1); err == nil && (buf[0] == ' ' || buf[0] == '\t') { const errorLimit = 80 // arbitrary limit on how much of the line we'll quote @@ -582,6 +580,10 @@ func readMIMEHeader(r *Reader, maxMemory, maxHeaders int64) (MIMEHeader, error) vv := m[key] if vv == nil { + // Account for per-entry overhead. + // 200 bytes is based on benchmarks circa Go 1.20. + const mapEntryOverhead = 200 + maxMemory -= int64(len(key)) maxMemory -= mapEntryOverhead } diff --git a/src/net/textproto/reader_test.go b/src/net/textproto/reader_test.go index 3b2d003bcf2f07..57db4bd82f9d72 100644 --- a/src/net/textproto/reader_test.go +++ b/src/net/textproto/reader_test.go @@ -344,6 +344,31 @@ func TestReadMIMEHeaderAllocations(t *testing.T) { } } +type testReader struct { + read func([]byte) (int, error) +} + +func (r testReader) Read(p []byte) (n int, err error) { return r.read(p) } + +func TestReadMIMEHeaderShortLimitLongLine(t *testing.T) { + // Small limit, long first line. + r := NewReader(bufio.NewReader( + io.MultiReader( + strings.NewReader("K:"), + testReader{ + read: func(p []byte) (n int, err error) { + for i := range p { + p[i] = ' ' + } + return len(p), nil + }, + }))) + _, err := readMIMEHeader(r, 1, 1) + if err != errMessageTooLarge { + t.Fatalf("readMIMEHeader = %v, want errMessageTooLarge", err) + } +} + type readResponseTest struct { in string inCode int diff --git a/src/os/file_windows.go b/src/os/file_windows.go index 8f0827a23debdf..a7d82e691054fd 100644 --- a/src/os/file_windows.go +++ b/src/os/file_windows.go @@ -100,11 +100,15 @@ func newFile(h syscall.Handle, name string, kind newFileKind, nonBlocking bool) panic("newFile with unknown kind") } + // Completion notification modes are shared by all handles to the file + // object. Preserve them for handles passed to NewFile, since other users + // of the file object may rely on those modes. See go.dev/issue/80979. f := &File{&file{ pfd: poll.FD{ - Sysfd: h, - IsStream: true, - ZeroReadIsEOF: true, + Sysfd: h, + IsStream: true, + ZeroReadIsEOF: true, + KeepFileCompletionModes: kind == kindNewFile, }, name: name, }} diff --git a/src/os/os_windows_test.go b/src/os/os_windows_test.go index 540364cf420cdb..d5952b0c8aa325 100644 --- a/src/os/os_windows_test.go +++ b/src/os/os_windows_test.go @@ -2341,6 +2341,70 @@ func TestOpenFileTruncateNamedPipe(t *testing.T) { f.Close() } +func TestFileKeepsCompletionNotificationModes(t *testing.T) { + // NewFile must preserve completion notification modes and perform I/O + // correctly with each combination. See go.dev/issue/80979. + t.Parallel() + for _, tt := range []struct { + name string + modes uint8 + }{ + {"none", 0}, + {"skipSuccess", syscall.FILE_SKIP_COMPLETION_PORT_ON_SUCCESS}, + {"skipEvent", syscall.FILE_SKIP_SET_EVENT_ON_HANDLE}, + {"both", syscall.FILE_SKIP_COMPLETION_PORT_ON_SUCCESS | syscall.FILE_SKIP_SET_EVENT_ON_HANDLE}, + } { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + name := filepath.Join(t.TempDir(), "file") + namep, err := syscall.UTF16PtrFromString(name) + if err != nil { + t.Fatal(err) + } + h, err := syscall.CreateFile(namep, syscall.GENERIC_READ|syscall.GENERIC_WRITE, + 0, nil, syscall.CREATE_ALWAYS, syscall.FILE_FLAG_OVERLAPPED, 0) + if err != nil { + t.Fatal(err) + } + if err := syscall.SetFileCompletionNotificationModes(h, tt.modes); err != nil { + syscall.CloseHandle(h) + t.Fatal(err) + } + f := os.NewFile(uintptr(h), name) + if f == nil { + syscall.CloseHandle(h) + t.Fatal("NewFile returned nil") + } + defer f.Close() + + // Query h directly: calling f.Fd would disassociate it from the poller. + var info windows.FILE_IO_COMPLETION_NOTIFICATION_INFORMATION + if err := windows.NtQueryInformationFile(h, &windows.IO_STATUS_BLOCK{}, + unsafe.Pointer(&info), uint32(unsafe.Sizeof(info)), windows.FileIoCompletionNotificationInformation); err != nil { + t.Fatal(err) + } + if info.Flags != uint32(tt.modes) { + t.Fatalf("completion modes = %#x; want %#x", info.Flags, tt.modes) + } + // Check that NewFile has initialized the runtime poller. + if err := f.SetDeadline(time.Time{}); err != nil { + t.Fatal(err) + } + const want = "hello" + if n, err := f.Write([]byte(want)); err != nil || n != len(want) { + t.Fatalf("Write = %d, %v; want %d, nil", n, err, len(want)) + } + buf := make([]byte, len(want)) + if n, err := f.ReadAt(buf, 0); err != nil || n != len(want) { + t.Fatalf("ReadAt = %d, %v; want %d, nil", n, err, len(want)) + } + if string(buf) != want { + t.Fatalf("ReadAt returned %q; want %q", buf, want) + } + }) + } +} + func TestNewFileStdinBlocked(t *testing.T) { // See https://go.dev/issue/75949. t.Parallel() diff --git a/src/os/root_openat.go b/src/os/root_openat.go index 62e6be665a9233..ee786176ce1fb6 100644 --- a/src/os/root_openat.go +++ b/src/os/root_openat.go @@ -153,13 +153,13 @@ func rootMkdirAll(r *Root, fullname string, perm FileMode) error { openLastComponentFunc := func(parent sysfdType, name string, endsInSlash bool) (struct{}, error) { err := mkdirat(parent, name, perm) if err == syscall.EEXIST { - mode, e := modeAt(parent, name) + fi, e := lstatat(parent, name) if e == nil { - if mode.IsDir() { + if fi.Mode().IsDir() { // The target of MkdirAll is an existing directory. err = nil - } else if mode&ModeSymlink != 0 { - // The target of MkdirAll is a symlink. + } else if isLink(fi) { + // The target of MkdirAll is a symlink or junction. // For consistency with os.MkdirAll, // succeed if the link resolves to a directory. // We don't return errSymlink here, because we don't @@ -420,7 +420,7 @@ Loop: case err != nil: return case fi.Mode().Type() == fs.ModeDir: - case fi.Mode().Type() == fs.ModeSymlink: + case isLink(fi): if runtime.GOOS != "windows" || flags&doInRootAlwaysResolveTerminalSlash != 0 { err = checkSymlink(dirfd, parts[i], syscall.ENOTDIR) } else { diff --git a/src/os/root_test.go b/src/os/root_test.go index 8d20ba0b8ed04e..353218a09533a4 100644 --- a/src/os/root_test.go +++ b/src/os/root_test.go @@ -2179,7 +2179,7 @@ func runRootMultiTest2(t *testing.T, f func(*testing.T, *rootMultiTest) (string, if desc.ref.template != "BASE" { return false } - if desc.kind == testFileSymlink && desc.target.ref.template != "BASE" { + if desc.kind.isLink() && desc.target.ref.template != "BASE" { return false } return true @@ -2383,11 +2383,12 @@ func newRootTest(t *testing.T, source, target testFileDesc, inRoot bool) *rootMu type testFileKind int const ( - testFileUnused = testFileKind(iota) - testFileAbsent // file does not exist - testFileFile // regular file - testFileDir // directory - testFileSymlink // symlink + testFileUnused = testFileKind(iota) + testFileAbsent // file does not exist + testFileFile // regular file + testFileDir // directory + testFileSymlink // symlink + testFileJunction // Windows directory junction testFileMax // testFileError represents a path which fails during resolution, @@ -2407,6 +2408,8 @@ func (kind testFileKind) String() string { return "dir" case testFileSymlink: return "symlink" + case testFileJunction: + return "junction" case testFileError: return "error" default: @@ -2414,6 +2417,10 @@ func (kind testFileKind) String() string { } } +func (kind testFileKind) isLink() bool { + return kind == testFileSymlink || kind == testFileJunction +} + // testFileRef is a kind of reference to a file. // // Many path names can refer to the same file: f, ./f, /abs/path/to/f, somedir/../f, etc. @@ -2480,12 +2487,14 @@ func (ref testFileRef) hasSlashSuffix() bool { type testFileDesc struct { kind testFileKind ref testFileRef - target *testFileDesc // symlink target, nil when kind is not testFileSymlink + target *testFileDesc // symlink or junction target, nil when !kind.isLink() } var rootComprehensive = flag.Bool("root_comprehensive", false, "run many more os.Root test variations (slow, uncertain value)") +var createJunction func(t *testing.T, link, target string) + // allTestFileDescs returns an iterator over all the testFileDescs we use in tests. func allTestFileDescs() iter.Seq[testFileDesc] { // A testFileDesc contains a reference type ("name", "d/../name", "../r/name", etc.) and @@ -2494,7 +2503,7 @@ func allTestFileDescs() iter.Seq[testFileDesc] { // When the kind is symlink, the desc contains a reference type and file kind for // the link target as well. We only exercise one level of symlink (although we // could do more), so this means a testFileDesc effectively contains four axes of - // variation: ref, kind, symlink ref, symlink kind. + // variation: ref, kind, link ref, link kind. // // For example: // @@ -2508,9 +2517,10 @@ func allTestFileDescs() iter.Seq[testFileDesc] { // but this is quite a few tests and gets quite slow. So by default we exclude // some variations. We test: // - // - every reference to every kind, except symlink - // - direct and direct/ references to a symlink to every reference to a file - // - a direct reference to a symlink to a direct reference to every kind (except file) + // - direct and direct/ references to a symlink or junction + // to every reference to a file + // - a direct reference to a symlink or junction + // to a direct reference to every kind (except file) // // The full set of variations may be enabled with the -comprehensive_root_tests flag. @@ -2518,7 +2528,7 @@ func allTestFileDescs() iter.Seq[testFileDesc] { // Every type of reference to every type of file, except symlink. for _, ref := range testFileRefs { for kind := range testFileMax { - if kind == testFileUnused || kind == testFileSymlink { + if kind == testFileUnused || kind.isLink() { continue } desc := testFileDesc{ @@ -2536,27 +2546,32 @@ func allTestFileDescs() iter.Seq[testFileDesc] { if !*rootComprehensive { refs = testFileLimitedRefs } - for _, ref := range refs { - for linkKind := range testFileMax { - if linkKind == testFileUnused || linkKind == testFileSymlink { - continue - } + for _, kind := range []testFileKind{testFileSymlink, testFileJunction} { + if kind == testFileJunction && runtime.GOOS != "windows" { + continue + } + for _, ref := range refs { + for linkKind := range testFileMax { + if linkKind == testFileUnused || linkKind.isLink() { + continue + } - linkRefs := testFileRefs - if !*rootComprehensive && linkKind != testFileFile && linkKind != testFileDir { - linkRefs = testFileLimitedRefs - } - for _, linkRef := range linkRefs { - desc := testFileDesc{ - kind: testFileSymlink, - ref: ref, - target: &testFileDesc{ - kind: linkKind, - ref: linkRef, - }, + linkRefs := testFileRefs + if !*rootComprehensive && linkKind != testFileFile && linkKind != testFileDir { + linkRefs = testFileLimitedRefs } - if !yield(desc) { - return + for _, linkRef := range linkRefs { + desc := testFileDesc{ + kind: kind, + ref: ref, + target: &testFileDesc{ + kind: linkKind, + ref: linkRef, + }, + } + if !yield(desc) { + return + } } } } @@ -2577,7 +2592,7 @@ func allTestFileDescs() iter.Seq[testFileDesc] { // So, open "file1/", where file1 is a symlink to "DIR/../file2", where file2 is a directory. func (desc testFileDesc) String() string { s := desc.ref.name + strings.ToUpper(desc.kind.String()[:1]) - if desc.kind == testFileSymlink { + if desc.kind.isLink() { s += desc.target.String() } return s @@ -2592,6 +2607,9 @@ func (desc testFileDesc) escapes() bool { if desc.kind == testFileSymlink { return desc.target.escapes() } + if desc.kind == testFileJunction { + return true + } return false } @@ -2605,7 +2623,7 @@ func (desc testFileDesc) lescapes() bool { // On Windows, a trailing slash does not cause symlink resolution. return false } - if desc.ref.hasSlashSuffix() && desc.kind == testFileSymlink { + if desc.ref.hasSlashSuffix() && desc.kind.isLink() { return desc.target.escapes() } return false @@ -2613,7 +2631,7 @@ func (desc testFileDesc) lescapes() bool { // finalKind reports the kind of the file after following all symlinks. func (desc testFileDesc) finalKind() testFileKind { - if desc.kind == testFileSymlink { + if desc.kind.isLink() { return desc.target.finalKind() } return desc.kind @@ -2622,11 +2640,11 @@ func (desc testFileDesc) finalKind() testFileKind { func (desc testFileDesc) lfinalKind() testFileKind { switch runtime.GOOS { case "windows": - if desc.ref.hasSlashSuffix() && desc.kind == testFileSymlink && desc.target.kind != testFileDir { + if desc.ref.hasSlashSuffix() && desc.kind.isLink() && desc.target.kind != testFileDir { return testFileError } default: - if desc.ref.hasSlashSuffix() && desc.kind == testFileSymlink { + if desc.ref.hasSlashSuffix() && desc.kind.isLink() { return desc.target.finalKind() } } @@ -2653,6 +2671,11 @@ func (desc testFileDesc) isError() bool { return true } return isError(*desc.target, hasSuffix) + case testFileJunction: + if desc.target.kind == testFileFile || (hasSuffix && desc.target.kind != testFileDir) { + return true + } + return isError(*desc.target, hasSuffix) default: return hasSuffix } @@ -2661,12 +2684,17 @@ func (desc testFileDesc) isError() bool { } func (desc testFileDesc) isSymlinkToDir() bool { - if desc.kind != testFileSymlink { + if !desc.kind.isLink() { return false } if desc.ref.escapes { return false } + if desc.kind == testFileJunction { + // Windows junctions are always directory links, regardless of what the + // target of the junction might be. + return true + } if desc.finalKind() == testFileDir { return true } @@ -2681,7 +2709,7 @@ func (desc testFileDesc) anySlashSuffix() bool { if len(name) > 0 && os.IsPathSeparator(name[len(name)-1]) { return true } - if desc.kind == testFileSymlink { + if desc.kind.isLink() { return desc.target.anySlashSuffix() } return false @@ -2736,6 +2764,14 @@ func (desc testFileDesc) create(t *testing.T, dir, base, token string) (fi os.Fi if err := os.Symlink(linktarget, path); err != nil { t.Fatal(err) } + case testFileJunction: + // Directory junction. We create a target named "s_"+base. + if runtime.GOOS != "windows" { + t.Skip("junctions not supported on " + runtime.GOOS) + } + linktarget := desc.target.ref.path(dir, "s_"+base) + fi = desc.target.create(t, dir, "s_"+base, token) + createJunction(t, path, linktarget) default: t.Fatalf("can't create file of kind: %v", desc.kind) } @@ -2796,7 +2832,7 @@ func dirTreeContents(t *testing.T, dir string) (contents []string) { switch d.Type() { case fs.ModeDir: ent += "/" - case fs.ModeSymlink: + case fs.ModeSymlink, fs.ModeIrregular: target, err := root.Readlink(path) if err != nil { t.Fatal(err) @@ -2863,6 +2899,8 @@ func TestRootMultiOpen(t *testing.T) { got := test.describeFile(t, f) switch { + case test.target.isError(): + test.wantError(t, gotErr, errAny) case test.root != nil && test.target.escapes(): // The operation escapes the root. test.wantError(t, gotErr, os.ErrPathEscapes) @@ -3030,7 +3068,7 @@ func TestRootMultiLink(t *testing.T) { test.wantError(t, gotErr, os.ErrPathEscapes) case test.source.lfinalKind() == testFileAbsent: test.wantError(t, gotErr, errAny) - case test.source.kind == testFileSymlink: + case test.source.kind.isLink(): // os.Link(old, new) may or may not deference old when it is a symlink. // POSIX says that link(2) should deference the source, but implementations // are inconsistent. @@ -3079,6 +3117,11 @@ func TestRootMultiLstat(t *testing.T) { if got, want := gotStat.Mode().Type(), fs.ModeSymlink; got != want { test.errorf(t, "got mode %v, want %v", got, want) } + case finalKind == testFileJunction: + test.wantError(t, gotErr, nil) + if got, want := gotStat.Mode().Type(), fs.ModeIrregular; got != want { + test.errorf(t, "got mode %v, want %v", got, want) + } case gotErr != nil: default: if !os.SameFile(gotStat, test.targetInfo) { @@ -3108,7 +3151,7 @@ func TestRootMultiMkdir(t *testing.T) { case test.root != nil && test.target.ref.escapes: // "mkdir ../target", or equivalent escaping path. test.wantError(t, gotErr, os.ErrPathEscapes) - case test.target.slashSuffix() && test.target.kind == testFileSymlink: + case test.target.slashSuffix() && test.target.kind.isLink(): // "mkdir symlink/", inconsistent behavior across platforms // as to whether this follows the symlink or not. // @@ -3173,7 +3216,7 @@ func testRootMultiMkdirAll(t *testing.T, test *rootMultiTest, targetPath string) // "mkdir ../target", or equivalent escaping path. test.wantError(t, gotErr, errAny) return "", errSkipRootConsistencyCheck - case test.root != nil && test.target.kind == testFileSymlink && test.target.target.kind == testFileAbsent && targetPath != test.targetPath: + case test.root != nil && test.target.kind.isLink() && test.target.target.kind == testFileAbsent && targetPath != test.targetPath: // A minor inconsistency between Root.MkdirAll and os.MkdirAll: // When an intermediate component of the tree being constructed is a // dangling symlink, Root.MkdirAll will follow the symlink and create @@ -3207,7 +3250,7 @@ func TestRootMultiRename(t *testing.T) { gotErr := rename(test.sourcePath, test.targetPath) if runtime.GOOS == "windows" && - (test.source.finalKind() != test.target.finalKind() || test.source.kind == testFileSymlink || test.target.kind == testFileSymlink) { + (test.source.finalKind() != test.target.finalKind() || test.source.kind.isLink() || test.target.kind.isLink()) { // os.Rename on Windows is implemented using MoveFileEx, // while Root.Rename is implemented using NtSetInformationFileEx // with an explicit request for POSIX semantics. @@ -3272,6 +3315,8 @@ func TestRootMultiReadFile(t *testing.T) { } switch { + case test.target.isError(): + test.wantError(t, gotErr, errAny) case test.root != nil && test.target.escapes(): test.wantError(t, gotErr, os.ErrPathEscapes) case test.target.finalKind() == testFileAbsent: @@ -3430,7 +3475,7 @@ func TestRootMultiReadlink(t *testing.T) { switch { case test.root != nil && test.target.lescapes(): test.wantError(t, gotErr, os.ErrPathEscapes) - case test.target.kind != testFileSymlink: + case !test.target.kind.isLink(): test.wantError(t, gotErr, errAny) case test.target.anySlashSuffix(): default: @@ -3483,6 +3528,8 @@ func TestRootMultiOpenFile(t *testing.T) { got := test.describeFile(t, f) switch { + case test.target.isError(): + test.wantError(t, gotErr, errAny) case test.root != nil && test.target.escapes(): test.wantError(t, gotErr, os.ErrPathEscapes) case test.target.finalKind() == testFileAbsent: diff --git a/src/os/root_unix.go b/src/os/root_unix.go index e88058715387b3..36db082e4f36d8 100644 --- a/src/os/root_unix.go +++ b/src/os/root_unix.go @@ -306,6 +306,10 @@ func readlinkat(fd int, name string) (string, error) { } } +func isLink(fi FileInfo) bool { + return fi.Mode()&ModeSymlink != 0 +} + // isDirectoryLink always returns false, because Unix systems don't have separate // symlink types for files and directories. // (See the Windows version of this function for more details.) diff --git a/src/os/root_windows.go b/src/os/root_windows.go index e03e5d4c2cc73b..a4dc73ac7f4625 100644 --- a/src/os/root_windows.go +++ b/src/os/root_windows.go @@ -425,8 +425,14 @@ func lstatat(parent syscall.Handle, name string) (FileInfo, error) { return fi, nil } -// isDirectoryLink reports whether fi (assumed to be a symlink) is a directory link. -// Windows symlinks come in two flavors: file and directory. This function distinguishes +// isLink reports whether fi is a symlink or other surrogate reparse point (such as a junction). +func isLink(fi FileInfo) bool { + fs, ok := fi.(*fileStat) + return ok && fs.isReparseTagNameSurrogate() +} + +// isDirectoryLink reports whether fi (assumed to be a link) is a directory link. +// Windows links come in two flavors: file and directory. This function distinguishes // between the two. func isDirectoryLink(fi FileInfo) bool { fs, ok := fi.(*fileStat) diff --git a/src/os/root_windows_test.go b/src/os/root_windows_test.go index ea604d18b19590..d15a3f79ef0015 100644 --- a/src/os/root_windows_test.go +++ b/src/os/root_windows_test.go @@ -19,6 +19,46 @@ import ( "unsafe" ) +func init() { + createJunction = func(t *testing.T, link, target string) { + if !filepath.IsAbs(target) { + target = filepath.Dir(link) + `\` + target + } + target, err := syscall.FullPath(target) + if err != nil { + t.Fatal(err) + } + var rd reparseData + rd.addSubstituteName(`\??\` + target) + rd.addPrintName(target) + if err := createMountPoint(link, &rd); err != nil { + t.Fatal(err) + } + } +} + +func TestRootOpenatFallback(t *testing.T) { + windows.TestOpenatFallback = true + t.Cleanup(func() { windows.TestOpenatFallback = false }) + + // Exercise the existing path traversal and symlink confinement cases + // with OBJ_DONT_REPARSE unavailable, as on Windows 10 build 10240. + t.Run("OpenFile", TestRootOpen_File) + t.Run("OpenDirectory", TestRootOpen_Directory) + t.Run("Create", TestRootCreate) + t.Run("Stat", TestRootStat) + t.Run("Lstat", TestRootLstat) + t.Run("RemoveAll", TestRootRemoveAll) + t.Run("RemoveAllNoRoot", TestRemoveAll) + t.Run("DeleteOnClose", testRootOpenFileDeleteOnClose) + t.Run("LegacyDelete", func(t *testing.T) { + windows.TestDeleteatFallback = true + t.Cleanup(func() { windows.TestDeleteatFallback = false }) + t.Run("RemoveAll", TestRootRemoveAll) + t.Run("RemoveAllNoRoot", TestRemoveAll) + }) +} + // Verify that Root.Open rejects Windows reserved names. func TestRootWindowsDeviceNames(t *testing.T) { r, err := os.OpenRoot(t.TempDir()) @@ -327,6 +367,10 @@ func TestRootOpenFileFlags(t *testing.T) { func TestRootOpenFileDeleteOnClose(t *testing.T) { t.Parallel() + testRootOpenFileDeleteOnClose(t) +} + +func testRootOpenFileDeleteOnClose(t *testing.T) { dir := t.TempDir() root, err := os.OpenRoot(dir) if err != nil { diff --git a/src/runtime/iface.go b/src/runtime/iface.go index 6385d1c9d055b1..9f1914372e7194 100644 --- a/src/runtime/iface.go +++ b/src/runtime/iface.go @@ -174,11 +174,11 @@ func (t *itabTableType) add(m *itab) { for i := uintptr(1); ; i++ { p := (**itab)(add(unsafe.Pointer(&t.entries), h*goarch.PtrSize)) m2 := *p - if m2 == m { - // A given itab may be used in more than one module - // and thanks to the way global symbol resolution works, the - // pointed-to itab may already have been inserted into the - // global 'hash'. + if m2 != nil && m2.Inter == m.Inter && m2.Type == m.Type { + // A plugin has its own copy of the itabs that the main program + // also has. Type switches and type assertions compare against + // the itab that is already in the table, so don't add a second + // itab for the same interface/type pair. return } if m2 == nil { diff --git a/src/runtime/metrics/doc.go b/src/runtime/metrics/doc.go index d9013d21932f12..57a3f644226530 100644 --- a/src/runtime/metrics/doc.go +++ b/src/runtime/metrics/doc.go @@ -334,6 +334,11 @@ Below is the full list of supported metrics, ordered lexicographically. by the net/http package due to a non-default GODEBUG=httpservecontentkeepheaders=... setting. + /godebug/non-default-behavior/httpservecontentmaxranges:events + The number of non-default behaviors executed + by the net/http package due to a non-default + GODEBUG=httpservecontentmaxranges=... setting. + /godebug/non-default-behavior/installgoroot:events The number of non-default behaviors executed by the go/build package due to a non-default GODEBUG=installgoroot=... setting. diff --git a/src/runtime/os_windows.go b/src/runtime/os_windows.go index 73e2b213f3f030..d89270bc4462cc 100644 --- a/src/runtime/os_windows.go +++ b/src/runtime/os_windows.go @@ -754,7 +754,13 @@ func semacreate(mp *m) { // //go:nowritebarrierrec func newosproc(mp *m) { - thandle, err := createThread(0, unsafe.Pointer(abi.FuncPCABI0(tstart_stdcall)), unsafe.Pointer(mp)) + // LockOSThread can reach newosproc on a goroutine stack when starting + // the template thread. createThread does not switch stacks itself. + var thandle uintptr + var err uint32 + systemstack(func() { + thandle, err = createThread(0, unsafe.Pointer(abi.FuncPCABI0(tstart_stdcall)), unsafe.Pointer(mp)) + }) if thandle == 0 { if atomic.Load(&exiting) != 0 { // CreateThread may fail if called diff --git a/src/runtime/os_workdir_ios_arm64.go b/src/runtime/os_workdir_ios_arm64.go index 40acc2f07888f3..f6e6d16cf6e718 100644 --- a/src/runtime/os_workdir_ios_arm64.go +++ b/src/runtime/os_workdir_ios_arm64.go @@ -19,6 +19,10 @@ func initWorkingDir() { writeErrStr("runtime/cgo: no main bundle\n") return } + if !bundleHasInfoPlist(bundle) { + // Not an app bundle; keep the inherited working directory. + return + } url := cfBundleCopyBundleURL(bundle) if url == 0 { // No app bundle URL found. @@ -62,3 +66,33 @@ func initWorkingDir() { writeErrStr(") failed\n") } } + +// bundleHasInfoPlist reports whether bundle contains an Info.plist. The main +// bundle of an executable that is not inside an app bundle is the directory +// holding the executable, which has no Info.plist. It can also happen on +// Corellium virtual devices. +func bundleHasInfoPlist(bundle uintptr) bool { + const ( + infoName = "Info\x00" + infoType = "plist\x00" + ) + name := cfStringCreateWithCString(0, unsafe.StringData(infoName), _kCFStringEncodingUTF8) + if name == 0 { + writeErrStr("runtime/cgo: cannot create Info.plist strings\n") + return false + } + typ := cfStringCreateWithCString(0, unsafe.StringData(infoType), _kCFStringEncodingUTF8) + if typ == 0 { + cfRelease(name) + writeErrStr("runtime/cgo: cannot create Info.plist strings\n") + return false + } + url := cfBundleCopyResourceURL(bundle, name, typ, 0) + cfRelease(name) + cfRelease(typ) + if url == 0 { + return false + } + cfRelease(url) + return true +} diff --git a/src/runtime/sys_ios_arm64.go b/src/runtime/sys_ios_arm64.go index 34dccf3207be4e..15417bf1f6046c 100644 --- a/src/runtime/sys_ios_arm64.go +++ b/src/runtime/sys_ios_arm64.go @@ -40,6 +40,25 @@ func cfBundleCopyBundleURL(bundle uintptr) uintptr { } func cfBundleCopyBundleURL_trampoline() +//go:nosplit +func cfBundleCopyResourceURL(bundle, resourceName, resourceType, subDirName uintptr) uintptr { + args := struct { + bundle uintptr + resourceName uintptr + resourceType uintptr + subDirName uintptr + ret uintptr + }{ + bundle: bundle, + resourceName: resourceName, + resourceType: resourceType, + subDirName: subDirName, + } + libcCall(unsafe.Pointer(abi.FuncPCABI0(cfBundleCopyResourceURL_trampoline)), unsafe.Pointer(&args)) + return args.ret +} +func cfBundleCopyResourceURL_trampoline() + //go:nosplit func cfURLGetFileSystemRepresentation(url uintptr, resolveAgainstBase bool, path *byte, pathLen uintptr) bool { args := struct { @@ -124,6 +143,7 @@ func cfRelease_trampoline() //go:cgo_import_dynamic libc_CFBundleGetMainBundle CFBundleGetMainBundle "/System/Library/Frameworks/CoreFoundation.framework/Versions/A/CoreFoundation" //go:cgo_import_dynamic libc_CFBundleCopyBundleURL CFBundleCopyBundleURL "/System/Library/Frameworks/CoreFoundation.framework/Versions/A/CoreFoundation" +//go:cgo_import_dynamic libc_CFBundleCopyResourceURL CFBundleCopyResourceURL "/System/Library/Frameworks/CoreFoundation.framework/Versions/A/CoreFoundation" //go:cgo_import_dynamic libc_CFURLGetFileSystemRepresentation CFURLGetFileSystemRepresentation "/System/Library/Frameworks/CoreFoundation.framework/Versions/A/CoreFoundation" //go:cgo_import_dynamic libc_CFStringCreateWithCString CFStringCreateWithCString "/System/Library/Frameworks/CoreFoundation.framework/Versions/A/CoreFoundation" //go:cgo_import_dynamic libc_CFBundleGetValueForInfoDictionaryKey CFBundleGetValueForInfoDictionaryKey "/System/Library/Frameworks/CoreFoundation.framework/Versions/A/CoreFoundation" diff --git a/src/runtime/sys_ios_arm64.s b/src/runtime/sys_ios_arm64.s index 72c17811c49ba8..6a253074744088 100644 --- a/src/runtime/sys_ios_arm64.s +++ b/src/runtime/sys_ios_arm64.s @@ -25,6 +25,16 @@ TEXT runtime·cfBundleCopyBundleURL_trampoline(SB),NOSPLIT,$0 MOVD R0, 8(R19) RET +TEXT runtime·cfBundleCopyResourceURL_trampoline(SB),NOSPLIT,$0 + MOVD R0, R19 + MOVD 8(R0), R1 // arg 2 resourceName + MOVD 16(R0), R2 // arg 3 resourceType + MOVD 24(R0), R3 // arg 4 subDirName + MOVD 0(R0), R0 // arg 1 bundle + BL libc_CFBundleCopyResourceURL(SB) + MOVD R0, 32(R19) + RET + TEXT runtime·cfURLGetFileSystemRepresentation_trampoline(SB),NOSPLIT,$0 MOVD R0, R19 MOVD 8(R0), R1 // arg 2 resolveAgainstBase diff --git a/src/runtime/unsafepoint_test.go b/src/runtime/unsafepoint_test.go index 79f0171854191f..235eaab645f553 100644 --- a/src/runtime/unsafepoint_test.go +++ b/src/runtime/unsafepoint_test.go @@ -5,14 +5,17 @@ package runtime_test import ( + "internal/abi" "internal/testenv" "os" "os/exec" "reflect" + "regexp" "runtime" "strconv" "strings" "testing" + "unsafe" ) // This is the function we'll be testing. @@ -120,3 +123,83 @@ func TestUnsafePoint(t *testing.T) { // write barrier proper into adjacent instructions (in both directions). // Hopefully we can clean up the latter at some point. } + +// tailCallOuter embeds an interface, so the compiler generates a wrapper for +// the promoted method M. On ppc64 that wrapper ends in a tail call: +// +// MOVD Rx, CTR +// BR (CTR) +// +// runtime.asyncPreempt does not preserve CTR, and its resume sequence leaves +// CTR holding the resume PC, so a goroutine preempted at that branch would +// resume by branching to that very instruction and spin there forever. The +// branch must therefore be marked as an unsafe point. See go.dev/issue/78576. +type tailCallInner interface{ M() int } + +type tailCallImpl struct{} + +func (tailCallImpl) M() int { return 42 } + +type tailCallOuter struct{ tailCallInner } + +var tailCallValue tailCallInner = tailCallOuter{tailCallImpl{}} + +func TestUnsafePointTailCall(t *testing.T) { + switch runtime.GOARCH { + case "ppc64", "ppc64le": + default: + t.Skipf("test not enabled for %s", runtime.GOARCH) + } + testenv.MustHaveExec(t) + + if got := tailCallValue.M(); got != 42 { + t.Fatalf("tailCallValue.M() = %d, want 42", got) + } + + // The itab's first method slot holds the generated wrapper for + // tailCallOuter.M, which is the function containing the tail call. + iface := (*struct { + tab *abi.ITab + data unsafe.Pointer + })(unsafe.Pointer(&tailCallValue)) + f := runtime.FuncForPC(iface.tab.Fun[0]) + if f == nil { + t.Fatal("no func for the tailCallOuter.M wrapper") + } + + // See TestUnsafePoint for why objdump works here. + cmd := exec.Command(testenv.GoToolPath(t), "tool", "objdump", "-s", "^"+regexp.QuoteMeta(f.Name())+"$", os.Args[0]) + out, err := cmd.CombinedOutput() + if err != nil { + t.Fatalf("can't objdump %v:\n%s", err, out) + } + + // Walk the disassembly and check the branch through CTR. As in + // TestUnsafePoint, only offsets from the function entry are meaningful. + var entry uint64 + branches := 0 + for _, line := range strings.Split(string(out), "\n")[1:] { + parts := strings.Fields(strings.TrimSpace(line)) + if len(parts) < 4 || !strings.HasPrefix(parts[0], ":") { + continue + } + pc, err := strconv.ParseUint(parts[1][2:], 16, 64) + if err != nil { + t.Fatalf("can't parse pc %s: %v", parts[1], err) + } + if entry == 0 { + entry = pc + } + t.Logf("%s", strings.TrimSpace(line)) + if parts[3] != "BR" || parts[4] != "(CTR)" { + continue + } + branches++ + if !runtime.UnsafePoint(f.Entry() + uintptr(pc-entry)) { + t.Errorf("%s\n\tbranch through CTR must be marked unsafe, but isn't", strings.TrimSpace(line)) + } + } + if branches != 1 { + t.Errorf("found %d branches through CTR in %s, want 1; output:\n%s", branches, f.Name(), out) + } +} diff --git a/src/testing/sub_test.go b/src/testing/sub_test.go index cad14bea9bfd00..4fa26c1f4395f0 100644 --- a/src/testing/sub_test.go +++ b/src/testing/sub_test.go @@ -195,12 +195,18 @@ func TestTRun(t *T) { === RUN chatty with recursion === RUN chatty with recursion/#00 === RUN chatty with recursion/#00/#00 +=== RUN chatty with recursion/#00/#01 + sub_test.go:NNN: ^V^O^N^[ --- PASS: chatty with recursion (N.NNs) --- PASS: chatty with recursion/#00 (N.NNs) - --- PASS: chatty with recursion/#00/#00 (N.NNs)`, + --- PASS: chatty with recursion/#00/#00 (N.NNs) + --- PASS: chatty with recursion/#00/#01 (N.NNs)`, f: func(t *T) { t.Run("", func(t *T) { t.Run("", func(t *T) {}) + t.Run("", func(t *T) { + t.Log(string(markFraming) + string(markErrBegin) + string(markErrEnd) + string(markEscape)) + }) }) }, }, { diff --git a/src/testing/testing.go b/src/testing/testing.go index 832d9e597fb4a2..c7431e5c3523a4 100644 --- a/src/testing/testing.go +++ b/src/testing/testing.go @@ -1260,7 +1260,9 @@ func (o *outputWriter) writeLine(b []byte, errBegin, errEnd bool) { } // Escape the framing marker. - b = escapeMarkers(b) + if o.c.chatty.json { + b = escapeMarkers(b) + } // If this is the start of an error, add ^O to the start of the output. var strErrBegin, strErrEnd string diff --git a/test/fixedbugs/issue81089.go b/test/fixedbugs/issue81089.go new file mode 100644 index 00000000000000..347a5884b4ef0b --- /dev/null +++ b/test/fixedbugs/issue81089.go @@ -0,0 +1,28 @@ +// run + +// Copyright 2026 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +package main + +type Conn interface{ Hello() } +type Pool[T Conn] struct{} + +func (p *Pool[T]) Hello(conn T) { conn.Hello() } + +type PoolConn struct{ Conn } +type CustomConn struct{} + +func (p CustomConn) Hello() { called = true } + +func NewPool[T Conn]() *Pool[T] { return &Pool[T]{} } + +var called bool + +func main() { + NewPool[*PoolConn]().Hello(&PoolConn{Conn: CustomConn{}}) + if !called { + panic("the embedded interface method Hello was not called") + } +} diff --git a/test/fixedbugs/issue81240.go b/test/fixedbugs/issue81240.go new file mode 100644 index 00000000000000..a98e89048d258b --- /dev/null +++ b/test/fixedbugs/issue81240.go @@ -0,0 +1,17 @@ +// compile + +// Copyright 2026 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +package p + +const W = 32 << (^uintptr(0) >> 63) // 32 or 64 + +type T struct { + a [1<<(W-30) - 1]byte +} + +func f(x, y *T) { + *x = *y +} diff --git a/test/fixedbugs/issue81242.go b/test/fixedbugs/issue81242.go new file mode 100644 index 00000000000000..c9cacce29bccfa --- /dev/null +++ b/test/fixedbugs/issue81242.go @@ -0,0 +1,17 @@ +// compile + +// Copyright 2026 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +package p + +const W = 32 << (^uintptr(0) >> 63) // 32 or 64 + +type T struct { + a [1<<(W-30) - 1]byte +} + +func f(t *T) { + *t = T{} +} diff --git a/test/fixedbugs/issue81612.go b/test/fixedbugs/issue81612.go new file mode 100644 index 00000000000000..6d80d56a634a10 --- /dev/null +++ b/test/fixedbugs/issue81612.go @@ -0,0 +1,125 @@ +// run + +// Copyright 2026 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +// The stack-allocated slice backing store optimization must not kick in +// when the address of a field of a slice element escapes. &s[i].f points +// into s's backing store just like &s[i] does, so the backing store has +// to be moved to the heap. + +package main + +type def struct { + ID int64 +} + +//go:noinline +func resolve() *int64 { + defs := []def{{77}} + var matches []def + for _, d := range defs { + matches = append(matches, d) + } + switch len(matches) { + case 1: + return &matches[0].ID + default: + candidates := make([]int64, 0, len(matches)) + for _, m := range matches { + candidates = append(candidates, m.ID) + } + return &candidates[0] + } +} + +var sink int64 +var escaped []def + +// resolveNoRange reaches the same bug without ranging over the slice: +// here the exclusive->nonexclusive transition is the assignment to +// escaped, which is on a different path than the &matches[0].ID. +// +//go:noinline +func resolveNoRange(n int) *int64 { + var matches []def + for i := range n { + matches = append(matches, def{int64(77 + i)}) + } + if len(matches) == 1 { + return &matches[0].ID + } + escaped = matches + return &sink +} + +type elem struct{ a [4]int64 } + +var escapedElems []elem + +// resolveViaSliceArr reaches the same bug through a slice of an array +// field: &t[0].a[:][0] points into t's backing store just as &t[0].a[0] +// does, because the OSLICEARR shares storage with the array. +// +//go:noinline +func resolveViaSliceArr(n int) *int64 { + var t []elem + for i := range n { + t = append(t, elem{a: [4]int64{int64(77 + i)}}) + } + if len(t) == 1 { + return &t[0].a[:][0] + } + escapedElems = t + return &sink +} + +// clobber overwrites the stack frames below main's. +// +//go:noinline +func clobber(n int) { + var buf [256]int64 + for i := range buf { + buf[i] = 0xdeadbeef + } + if n > 0 { + clobber(n - 1) + } + sink = buf[n&255] +} + +func main() { + p := resolve() + if *p != 77 { + println("resolve: before clobber:", *p) + panic("wrong value") + } + clobber(4) + if *p != 77 { + println("resolve: after clobber:", *p) + panic("value destroyed by unrelated stack traffic") + } + + p = resolveNoRange(1) + if *p != 77 { + println("resolveNoRange: before clobber:", *p) + panic("wrong value") + } + clobber(4) + if *p != 77 { + println("resolveNoRange: after clobber:", *p) + panic("value destroyed by unrelated stack traffic") + } + + p = resolveViaSliceArr(1) + if *p != 77 { + println("resolveViaSliceArr: before clobber:", *p) + panic("wrong value") + } + clobber(4) + if *p != 77 { + println("resolveViaSliceArr: after clobber:", *p) + panic("value destroyed by unrelated stack traffic") + } +} diff --git a/test/fixedbugs/issue81687.go b/test/fixedbugs/issue81687.go new file mode 100644 index 00000000000000..610fec95990c69 --- /dev/null +++ b/test/fixedbugs/issue81687.go @@ -0,0 +1,23 @@ +// compile + +// Copyright 2026 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +package p + +// The rotate amounts are only rewritten on 32-bit platforms, +// so this only failed to compile there. + +func rotate(x uint32, n int) uint32 { + if n < 0 { + s := uint(-n) + return x<>(32-s) + } + s := uint(n) + return x>>s | x<<(32-s) +} + +func f(x uint32) uint32 { + return rotate(rotate(x, 7), -7) +}