blob: e812be4624a7ed5bb860b0aa676eaa8c3535b9be [file]
// 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 ssa
import (
"slices"
"strings"
"cmd/compile/internal/ssa/block"
"cmd/compile/internal/ssa/ssaop"
)
// KnownBits does constant folding across bitfields
func KnownBits(f *Func) {
kb := &knownBitsState{
entries: f.Cache.AllocKnownBitsEntriesSlice(f.NumValues()),
seenValues: f.Cache.allocBitset(f.NumValues()),
reachableBlocks: f.Cache.allocBitset(f.NumBlocks()),
}
defer f.Cache.FreeKnownBitsEntriesSlice(kb.entries)
defer f.Cache.freeBitset(kb.seenValues)
defer f.Cache.freeBitset(kb.reachableBlocks)
clear(kb.seenValues)
clear(kb.entries)
clear(kb.reachableBlocks)
blocks := f.Postorder()
for _, b := range blocks {
kb.reachableBlocks.Set(uint32(b.ID))
}
for _, b := range slices.Backward(blocks) {
for _, v := range b.Values {
if v.Uses == 0 && f.Pass.Debug == 0 {
continue
}
if !(v.Type.IsInteger() || v.Type.IsBoolean()) {
continue
}
switch v.Op {
case ssaop.OpConst64, ssaop.OpConst32, ssaop.OpConst16, ssaop.OpConst8, ssaop.OpConstBool:
continue
}
val, k := kb.fold(v)
if k != -1 {
continue
}
if f.Pass.Debug > 0 {
var pval any = val
if v.Type.IsBoolean() {
pval = val != 0
}
f.Warnl(v.Pos, "known value of %v (%v): %v", v, v.Op, pval)
}
var c *Value
switch v.Type.Size() {
case 1:
if v.Type.IsBoolean() {
c = f.ConstBool(v.Type, val != 0)
break
}
c = f.ConstInt8(v.Type, int8(val))
case 2:
c = f.ConstInt16(v.Type, int16(val))
case 4:
c = f.ConstInt32(v.Type, int32(val))
case 8:
c = f.ConstInt64(v.Type, val)
default:
panic("unreachable; unknown integer size")
}
v.CopyOf(c)
}
}
}
// allPossibleValues iterates over all values that could exist.
// It scales exponentially with the number of unknown bits,
// the exact number of iterations will be uint128(1)<<bits.OnesCount64(^known)
// thus be careful with what values you pass to it.
func allPossibleValues(value, known int64) func(yield func(v int64) bool) {
unknown := ^known
return func(yield func(v int64) bool) {
// This finds the next valid value for the variable bits.
// It is equivalent to (s|known + 1) & unknown.
// The s|known step creates blocks of 1s in all the known bits.
// +1 finds the next possible value, the blocks of 1s set in the previous step allows it to skip over blocks of known bits.
// & unknown clears garbage generated by the blocks of ones and overflow.
//
// You can transform (s|known + 1) & unknown into (s - unknown) & unknown through:
// (s + known + 1) & unknown: s | known → s + known (since s & known == 0)
// (s + ^unknown + 1) & unknown: known → ^unknown (definition of unknown)
// (s + -unknown) & unknown: ^unknown + 1 → -unknown (two's complement negation)
// (s - unknown) & unknown: s + -unknown → s - unknown (arithmetic)
for s := int64(0); ; s = (s - unknown) & unknown {
// fixed bits | current variable bits gives the current iteration
if !yield(value | s) {
return
}
if s == unknown {
break
}
}
}
}
type knownBitsEntry struct {
// Two invariants:
// 1. unknown bits are always set to 0 inside value
// 2. all values are sign-extended to int64 (inspired by RISC-V's xlen=64)
// This means let's say you know an 8 bits value is 0b10??????,
// known = int64(int8(0b11000000))
// value = int64(int8(0b10000000))
// 3. booleans are stored as 1 byte values who are either 0 or 1.
known, value int64
}
type knownBitsState struct {
entries []knownBitsEntry // indexed by Value.ID
seenValues bitset // indexed by Value.ID (at the bit level)
reachableBlocks bitset // indexed by Block.ID (at the bit level)
}
func (kb *knownBitsState) fold(v *Value) (value, known int64) {
if kb.seenValues.Test(uint32(v.ID)) {
return kb.entries[v.ID].value, kb.entries[v.ID].known
}
defer func() {
// maintain the invariants:
// 3. booleans are stored as 1 byte values who are either 0 or 1.
if v.Type.IsBoolean() {
value &= 1
known |= ^1
}
// 2. all values are sign-extended to int64 (inspired by RISC-V's xlen=64)
switch v.Type.Size() {
case 1:
value = int64(int8(value))
known = int64(int8(known))
case 2:
value = int64(int16(value))
known = int64(int16(known))
case 4:
value = int64(int32(value))
known = int64(int32(known))
case 8:
default:
panic("unreachable; unknown integer size")
}
// 1. unknown bits are always set to 0 inside value
value &= known
kb.entries[v.ID].known = known
kb.entries[v.ID].value = value
if v.Block.Func.Pass.Debug > 1 {
v.Block.Func.Warnl(v.Pos, "known bits state %v: %v", v, kb.entries[v.ID])
}
}()
kb.seenValues.Set(uint32(v.ID)) // set seen early to give up on loops
switch v.Op {
// TODO: rotates, ...
case ssaop.OpConst64, ssaop.OpConst32, ssaop.OpConst16, ssaop.OpConst8, ssaop.OpConstBool:
return v.AuxInt, -1
case ssaop.OpAnd64, ssaop.OpAnd32, ssaop.OpAnd16, ssaop.OpAnd8, ssaop.OpAndB:
x, xk := kb.fold(v.Args[0])
y, yk := kb.fold(v.Args[1])
onesInBoth := x & y
zerosInX := ^x & xk
zerosInY := ^y & yk
return x & y, onesInBoth | zerosInX | zerosInY
case ssaop.OpOr64, ssaop.OpOr32, ssaop.OpOr16, ssaop.OpOr8, ssaop.OpOrB:
x, xk := kb.fold(v.Args[0])
y, yk := kb.fold(v.Args[1])
zerosInBoth := ^x & ^y & (xk & yk)
onesInX := x
onesInY := y
return x | y, onesInX | onesInY | zerosInBoth
case ssaop.OpXor64, ssaop.OpXor32, ssaop.OpXor16, ssaop.OpXor8:
x, xk := kb.fold(v.Args[0])
y, yk := kb.fold(v.Args[1])
return x ^ y, xk & yk
case ssaop.OpCom64, ssaop.OpCom32, ssaop.OpCom16, ssaop.OpCom8, ssaop.OpNot:
x, xk := kb.fold(v.Args[0])
return ^x, xk
case ssaop.OpPhi:
set := false
for i, arg := range v.Args {
if !kb.isLiveInEdge(v.Block, uint(i)) {
continue
}
a, k := kb.fold(arg)
if !set {
value, known = a, k
set = true
} else {
known &^= value ^ a
known &= k
}
if known == 0 {
break
}
}
return value, known
case ssaop.OpCopy, ssaop.OpCvtBoolToUint8,
ssaop.OpSignExt8to16, ssaop.OpSignExt8to32, ssaop.OpSignExt8to64, ssaop.OpSignExt16to32, ssaop.OpSignExt16to64, ssaop.OpSignExt32to64,
// The defer block handles maintaining the sign-extension invariant using v.Type.Size()
// thus we can just pass Truncs as-is.
ssaop.OpTrunc64to32, ssaop.OpTrunc64to16, ssaop.OpTrunc64to8, ssaop.OpTrunc32to16, ssaop.OpTrunc32to8, ssaop.OpTrunc16to8:
return kb.fold(v.Args[0])
case ssaop.OpEq64, ssaop.OpEq32, ssaop.OpEq16, ssaop.OpEq8, ssaop.OpEqB:
x, xk := kb.fold(v.Args[0])
y, yk := kb.fold(v.Args[1])
differentBits := x ^ y
if differentBits&xk&yk != 0 {
return 0, -1
}
if xk == -1 && yk == -1 {
return BoolToAuxInt(x == y), -1
}
return 0, -1 << 1
case ssaop.OpNeq64, ssaop.OpNeq32, ssaop.OpNeq16, ssaop.OpNeq8, ssaop.OpNeqB:
x, xk := kb.fold(v.Args[0])
y, yk := kb.fold(v.Args[1])
differentBits := x ^ y
if differentBits&xk&yk != 0 {
return 1, -1
}
if xk == -1 && yk == -1 {
return BoolToAuxInt(x != y), -1
}
return 0, -1 << 1
case ssaop.OpZeroExt8to16, ssaop.OpZeroExt8to32, ssaop.OpZeroExt8to64, ssaop.OpZeroExt16to32, ssaop.OpZeroExt16to64, ssaop.OpZeroExt32to64:
x, k := kb.fold(v.Args[0])
srcSize := v.Args[0].Type.Size() * 8
mask := int64(1<<srcSize - 1)
return x & mask, k | ^mask
case ssaop.OpLsh8x8, ssaop.OpLsh16x8, ssaop.OpLsh32x8, ssaop.OpLsh64x8,
ssaop.OpLsh8x16, ssaop.OpLsh16x16, ssaop.OpLsh32x16, ssaop.OpLsh64x16,
ssaop.OpLsh8x32, ssaop.OpLsh16x32, ssaop.OpLsh32x32, ssaop.OpLsh64x32,
ssaop.OpLsh8x64, ssaop.OpLsh16x64, ssaop.OpLsh32x64, ssaop.OpLsh64x64:
return kb.computeKnownBitsForShift(v, func(x, xk, xSize, shift int64) (value, known int64) {
return x << shift, xk<<shift | (1<<shift - 1)
})
case ssaop.OpRsh8Ux8, ssaop.OpRsh16Ux8, ssaop.OpRsh32Ux8, ssaop.OpRsh64Ux8,
ssaop.OpRsh8Ux16, ssaop.OpRsh16Ux16, ssaop.OpRsh32Ux16, ssaop.OpRsh64Ux16,
ssaop.OpRsh8Ux32, ssaop.OpRsh16Ux32, ssaop.OpRsh32Ux32, ssaop.OpRsh64Ux32,
ssaop.OpRsh8Ux64, ssaop.OpRsh16Ux64, ssaop.OpRsh32Ux64, ssaop.OpRsh64Ux64:
return kb.computeKnownBitsForShift(v, func(x, xk, xSize, shift int64) (value, known int64) {
x &= (1<<xSize - 1)
xk |= -1 << xSize
return int64(uint64(x) >> shift), int64(uint64(xk)>>shift | (^uint64(0) << (64 - shift)))
})
case ssaop.OpRsh8x8, ssaop.OpRsh16x8, ssaop.OpRsh32x8, ssaop.OpRsh64x8,
ssaop.OpRsh8x16, ssaop.OpRsh16x16, ssaop.OpRsh32x16, ssaop.OpRsh64x16,
ssaop.OpRsh8x32, ssaop.OpRsh16x32, ssaop.OpRsh32x32, ssaop.OpRsh64x32,
ssaop.OpRsh8x64, ssaop.OpRsh16x64, ssaop.OpRsh32x64, ssaop.OpRsh64x64:
return kb.computeKnownBitsForShift(v, func(x, xk, xSize, shift int64) (value, known int64) {
return x >> shift, xk >> shift
})
default:
return 0, 0
}
}
func (kbe knownBitsEntry) String() string {
lut := []rune{ // indexed by knownBit<<1 | valueBit
0b00: '?',
0b01: '¿', // violates invariant 1
0b10: '0',
0b11: '1',
}
var sb strings.Builder
sb.Grow(64)
for i := 63; i >= 0; i-- {
bits := (kbe.known>>i&1)<<1 | (kbe.value >> i & 1)
sb.WriteRune(lut[bits])
}
return sb.String()
}
func (kb *knownBitsState) isLiveInEdge(b *Block, index uint) bool {
inEdge := b.Preds[index]
return kb.isLiveOutEdge(inEdge.B, uint(inEdge.I))
}
func (kb *knownBitsState) isLiveOutEdge(b *Block, index uint) bool {
if !kb.reachableBlocks.Test(uint32(b.ID)) {
return false
}
switch b.Kind {
case block.BlockFirst:
return index == 0
case block.BlockPlain, block.BlockIf, block.BlockDefer, block.BlockRet, block.BlockRetJmp, block.BlockExit, block.BlockJumpTable:
return true
default:
panic("unreachable; unknown block kind")
}
}
// computeKnownBitsForShift computes the known bits for a shift operation.
// Considering the following piece of code x = x << uint8(i)
// The algorithm is based on two observations:
//
// 1. computing a shift of a lattice by a constant (i) is easy:
// value, known = x<<i, xk<<i|(1<<i-1)
// each point in the lattice is shifted by the constant, all new shifted in bits are known zeros.
//
// 2. x = uint8(x) << i is equivalent to
//
// switch i {
// case 0: x0 = x << 0
// case 1: x1 = x << 1
// case 2: x2 = x << 2
// case 3: x3 = x << 3
// case 4: x4 = x << 4
// case 5: x5 = x << 5
// case 6: x6 = x << 6
// case 7: x7 = x << 7
// default: xd = x << 8
// }
// x = phi(x0, x1, x2, x3, x4, x5, x6, x7, xd)
//
// The algorithm below then models the phi in the equivalence above using same intersection algorithm phi uses.
// We also leverage known bits of the shift amount to remove "branches" in the switch that are proved to be impossible.
func (kb *knownBitsState) computeKnownBitsForShift(v *Value, doShiftByAConst func(x, xk, xSize, shift int64) (value, known int64)) (value, known int64) {
xSize := v.Args[0].Type.Size() * 8
x, xk := kb.fold(v.Args[0])
y, yk := kb.fold(v.Args[1])
if uint64(y) >= uint64(xSize) {
return doShiftByAConst(x, xk, xSize, 64)
}
set := false
if v.AuxInt == 0 && uint64(^yk) >= uint64(xSize) {
// this implement the default case of the equivalent switch above.
// if the shift isn't bounded and there are unknown bits above the shift size we might completely stomp all bits.
value, known = doShiftByAConst(x, xk, xSize, 64)
set = true
}
yk |= ^(xSize - 1)
for i := range allPossibleValues(y, yk) {
a, k := doShiftByAConst(x, xk, xSize, i)
if !set {
value, known = a, k
set = true
} else {
known &^= value ^ a
known &= k
}
if known == 0 {
break
}
}
return value & known, known
}