diff --git a/UnitTest.lean b/UnitTest.lean index cd19e9b8b2..5bade00ca0 100644 --- a/UnitTest.lean +++ b/UnitTest.lean @@ -14,6 +14,7 @@ import UnitTest.DataFlowFramework.Dominance import UnitTest.DataFlowFramework.DeadCodeAnalysis import UnitTest.DataFlowFramework.EntryState import UnitTest.DataFlowFramework.ModArithRangeAnalysis +import UnitTest.DataFlowFramework.KnownBitsAnalysis import UnitTest.ConstantValue import UnitTest.Evaluate import UnitTest.Interp diff --git a/UnitTest/DataFlowFramework/EntryState.lean b/UnitTest/DataFlowFramework/EntryState.lean index 79941330a0..5061b76d94 100644 --- a/UnitTest/DataFlowFramework/EntryState.lean +++ b/UnitTest/DataFlowFramework/EntryState.lean @@ -47,19 +47,19 @@ private def transfer (irCtx : WfIRContext OpCode) : Array TestDomain := Array.replicate (op.getNumResults! irCtx.raw) ⊥ -/-- Sparse test analysis configured with the type-sensitive entry-state hook. -/ +/-- Sparse test analysis configured with the type sensitive entry state hook. -/ private def customEntryStateAnalysis : DataFlowAnalysis := SparseForwardDataFlowAnalysis.new .test .test transfer (entryState := entryState) -/-- Sparse test analysis using the framework's default top entry state. -/ -private def defaultEntryStateAnalysis : DataFlowAnalysis := - SparseForwardDataFlowAnalysis.new .test .test transfer +/-- Sparse test analysis configured with a constant top entry state. -/ +private def topEntryStateAnalysis : DataFlowAnalysis := + SparseForwardDataFlowAnalysis.new .test .test transfer (entryState := fun _ _ => ⊤) /-- Read the test lattice element attached to an SSA value. -/ private def getElement (value : ValuePtr) (dfCtx : DataFlowContext) : TestDomain := SparseFact.getElement .test value dfCtx -/-- Compare one named SSA value's state with the expected test-domain value. -/ +/-- Compare one named SSA value's state with the expected test domain value. -/ private def checkValue (name : String) (expected : TestDomain) @@ -85,10 +85,10 @@ private def checkNoFact | some _ => return #[s!"{name}: expected no stored fact for bottom"] /-- -Input shared by the custom and default entry-state checks. It exercises both places -where entry-state facts participate in sparse propagation: +Input shared by the type sensitive and constant top entry state checks. It exercises both places +where entry state facts participate in sparse propagation: -* `entryArg` and `forwardedArg` are entry-block arguments whose states cannot yet +* `entryArg` and `forwardedArg` are entry block arguments whose states cannot yet come from call sites. * `fallbackArg` verifies that the state of `forwardedArg` propagates through a valid `cf.br` to a non-entry block argument. @@ -104,7 +104,7 @@ private def testInput := r#""builtin.module"() ({ }) : () -> () }) : () -> ()"# -/-- Verify that an analysis can override the default with a type-sensitive entry state. -/ +/-- Verify that an analysis can use a type sensitive entry state. -/ private def testCustomEntryState : String := runWithAnalyses testInput #[customEntryStateAnalysis] fun top dfCtx ctx => match recoverNames top ctx testInput with @@ -114,9 +114,9 @@ private def testCustomEntryState : String := checkValue "fallbackArg" (.value 16) recovered dfCtx ++ checkNoFact "implicitBottom" recovered dfCtx -/-- Verify that omitting the entry-state hook conservatively assigns top. -/ -private def testDefaultEntryState : String := - runWithAnalyses testInput #[defaultEntryStateAnalysis] fun top dfCtx ctx => +/-- Verify that a constant top entry state conservatively assigns top. -/ +private def testTopEntryState : String := + runWithAnalyses testInput #[topEntryStateAnalysis] fun top dfCtx ctx => match recoverNames top ctx testInput with | .error err => #[err] | .ok recovered => @@ -134,6 +134,6 @@ info: "ok" info: "ok" -/ #guard_msgs in -#eval! testDefaultEntryState +#eval! testTopEntryState end EntryStateTest diff --git a/UnitTest/DataFlowFramework/KnownBitsAnalysis.lean b/UnitTest/DataFlowFramework/KnownBitsAnalysis.lean new file mode 100644 index 0000000000..f0e658cc11 --- /dev/null +++ b/UnitTest/DataFlowFramework/KnownBitsAnalysis.lean @@ -0,0 +1,201 @@ +import UnitTest.DataFlowFramework.Helpers + +import Veir.Analysis.DataFlow.KnownBitsAnalysis + +open Veir + +namespace KnownBitsDataflow + +/-- Expected masks for one named SSA value. -/ +private structure ExpectedKnownBits where + name : String + bitwidth : Nat + zero : Nat + one : Nat + +private def knownBitsToString : KnownBitsLattice → String + | .bottom => "bottom" + | .known bits => s!"i{bits.bitwidth}(zero={bits.zero.toNat}, one={bits.one.toNat})" + +private def compareKnownBits + (dfCtx : DataFlowContext) + (recovered : RecoveredNames) + (expected : Array ExpectedKnownBits) : MismatchReport := Id.run do + let mut report := #[] + for e in expected do + let some value := recovered.values[e.name]? + | report := report.push s!"known bits {e.name}: missing SSA value" + continue + let observed : KnownBitsLattice := SparseFact.getElement .knownBits value dfCtx + let expectedZero := BitVec.ofNat e.bitwidth e.zero + let expectedOne := BitVec.ofNat e.bitwidth e.one + let isMatch := match observed with + | .known bits => + bits.bitwidth == e.bitwidth && + bits.zero.toNat == expectedZero.toNat && + bits.one.toNat == expectedOne.toNat + | .bottom => false + if !isMatch then + report := report.push <| + s!"known bits {e.name}: expected i{e.bitwidth}" ++ + s!"(zero={expectedZero.toNat}, one={expectedOne.toNat}), " ++ + s!"observed {knownBitsToString observed}" + report + +private def run (mlir : String) (expected : Array ExpectedKnownBits) : String := + runWithAnalyses mlir #[Veir.KnownBitsAnalysis] fun top dfCtx irCtx => + match recoverNames top irCtx mlir with + | .error err => #[err] + | .ok recovered => compareKnownBits dfCtx recovered expected + +/-- Arith constants and bitwise operations preserve partial known-bit information. -/ +def runArithKnownBitsExample : String := + let mlir := r#""builtin.module"() ({ +^bb0: + "func.func"() <{function_type = (i8) -> (), sym_name = "known_bits_arith"}> ({ + ^entry(%x : i8): + %c240 = "arith.constant"() <{value = 240 : i8}> : () -> i8 + %c3 = "arith.constant"() <{value = 3 : i8}> : () -> i8 + %c5 = "arith.constant"() <{value = 5 : i8}> : () -> i8 + %sum = "arith.addi"(%c3, %c5) : (i8, i8) -> i8 + %anded = "arith.andi"(%x, %c240) : (i8, i8) -> i8 + %ored = "arith.ori"(%anded, %c3) : (i8, i8) -> i8 + %xored = "arith.xori"(%ored, %c5) : (i8, i8) -> i8 + "func.return"() : () -> () + }) : () -> () +}) : () -> ()"# + let expected := + #[ { name := "x", bitwidth := 8, zero := 0, one := 0 } + , { name := "c240", bitwidth := 8, zero := 15, one := 240 } + , { name := "c3", bitwidth := 8, zero := 252, one := 3 } + , { name := "c5", bitwidth := 8, zero := 250, one := 5 } + , { name := "sum", bitwidth := 8, zero := 247, one := 8 } + , { name := "anded", bitwidth := 8, zero := 15, one := 0 } + , { name := "ored", bitwidth := 8, zero := 12, one := 3 } + , { name := "xored", bitwidth := 8, zero := 9, one := 6 } + ] + run mlir expected + +/-- LLVM spellings and variadic Comb operations use the same transfer functions. -/ +def runLLVMAndCombKnownBitsExample : String := + let mlir := r#""builtin.module"() ({ +^bb0: + "func.func"() <{function_type = (i8) -> (), sym_name = "known_bits_dialects"}> ({ + ^entry(%x : i8): + %lc240 = "llvm.mlir.constant"() <{value = 240 : i8}> : () -> i8 + %lc3 = "llvm.mlir.constant"() <{value = 3 : i8}> : () -> i8 + %lc5 = "llvm.mlir.constant"() <{value = 5 : i8}> : () -> i8 + %land = "llvm.and"(%x, %lc240) : (i8, i8) -> i8 + %lor = "llvm.or"(%land, %lc3) : (i8, i8) -> i8 + %lxor = "llvm.xor"(%lor, %lc5) : (i8, i8) -> i8 + %hc240 = "hw.constant"() <{value = 240 : i8}> : () -> i8 + %hc15 = "hw.constant"() <{value = 15 : i8}> : () -> i8 + %hc3 = "hw.constant"() <{value = 3 : i8}> : () -> i8 + %cand = "comb.and"(%hc240, %hc15, %hc3) : (i8, i8, i8) -> i8 + %cor = "comb.or"(%hc240, %hc15, %hc3) : (i8, i8, i8) -> i8 + %cxor = "comb.xor"(%hc240, %hc15, %hc3) : (i8, i8, i8) -> i8 + "func.return"() : () -> () + }) : () -> () +}) : () -> ()"# + let expected := + #[ { name := "land", bitwidth := 8, zero := 15, one := 0 } + , { name := "lor", bitwidth := 8, zero := 12, one := 3 } + , { name := "lxor", bitwidth := 8, zero := 9, one := 6 } + , { name := "cand", bitwidth := 8, zero := 255, one := 0 } + , { name := "cor", bitwidth := 8, zero := 0, one := 255 } + , { name := "cxor", bitwidth := 8, zero := 3, one := 252 } + ] + run mlir expected + +/-- Arithmetic, shifts, casts, comparisons, and flags use LLVM-style known-bit transfer rules. -/ +def runLLVMStyleTransfersExample : String := + let mlir := r#""builtin.module"() ({ +^bb0: + "func.func"() <{function_type = (i8) -> (), sym_name = "known_bits_transfers"}> ({ + ^entry(%x : i8): + %c1 = "arith.constant"() <{value = 1 : i8}> : () -> i8 + %c3 = "arith.constant"() <{value = 3 : i8}> : () -> i8 + %c12 = "arith.constant"() <{value = 12 : i8}> : () -> i8 + %c16 = "arith.constant"() <{value = 16 : i8}> : () -> i8 + %c127 = "arith.constant"() <{value = 127 : i8}> : () -> i8 + %c128 = "arith.constant"() <{value = 128 : i8}> : () -> i8 + %low = "arith.andi"(%x, %c127) : (i8, i8) -> i8 + %shl = "arith.shli"(%x, %c3) : (i8, i8) -> i8 + %lshr = "arith.shrui"(%x, %c3) : (i8, i8) -> i8 + %mul = "arith.muli"(%x, %c12) : (i8, i8) -> i8 + %udiv = "arith.divui"(%x, %c16) : (i8, i8) -> i8 + %urem = "arith.remui"(%x, %c16) : (i8, i8) -> i8 + %self_add = "arith.addi"(%x, %x) : (i8, i8) -> i8 + %nuw = "arith.addi"(%x, %c128) <{overflowFlags = #arith.overflow}> : (i8, i8) -> i8 + %nsw = "arith.addi"(%low, %c1) <{overflowFlags = #arith.overflow}> : (i8, i8) -> i8 + %wide = "arith.extui"(%low) : (i8) -> i16 + %ult = "arith.cmpi"(%low, %c128) <{predicate = 6 : i64}> : (i8, i8) -> i1 + %pop = "llvm.intr.ctpop"(%low) : (i8) -> i8 + %clz = "llvm.intr.ctlz"(%low) <{is_zero_poison = 0 : i1}> : (i8) -> i8 + %reverse = "llvm.intr.bitreverse"(%low) : (i8) -> i8 + %fshl = "llvm.intr.fshl"(%x, %low, %c3) : (i8, i8, i8) -> i8 + %saturating = "llvm.intr.uadd.sat"(%x, %c128) : (i8, i8) -> i8 + %bswap = "llvm.intr.bswap"(%wide) : (i16) -> i16 + "func.return"() : () -> () + }) : () -> () +}) : () -> ()"# + let expected := + #[ { name := "low", bitwidth := 8, zero := 128, one := 0 } + , { name := "shl", bitwidth := 8, zero := 7, one := 0 } + , { name := "lshr", bitwidth := 8, zero := 224, one := 0 } + , { name := "mul", bitwidth := 8, zero := 3, one := 0 } + , { name := "udiv", bitwidth := 8, zero := 240, one := 0 } + , { name := "urem", bitwidth := 8, zero := 240, one := 0 } + , { name := "self_add", bitwidth := 8, zero := 1, one := 0 } + , { name := "nuw", bitwidth := 8, zero := 0, one := 128 } + , { name := "nsw", bitwidth := 8, zero := 128, one := 0 } + , { name := "wide", bitwidth := 16, zero := 65408, one := 0 } + , { name := "ult", bitwidth := 1, zero := 0, one := 1 } + , { name := "pop", bitwidth := 8, zero := 248, one := 0 } + , { name := "clz", bitwidth := 8, zero := 240, one := 0 } + , { name := "reverse", bitwidth := 8, zero := 1, one := 0 } + , { name := "fshl", bitwidth := 8, zero := 4, one := 0 } + , { name := "saturating", bitwidth := 8, zero := 0, one := 128 } + , { name := "bswap", bitwidth := 16, zero := 33023, one := 0 } + ] + run mlir expected + +/-- Joining exact values retains only the bits on which both values agree. -/ +def testKnownBitsJoin : String := + let joined := + KnownBitsLattice.join + (.constant 8 165) + (.constant 8 167) + match joined with + | .known bits => + if bits.bitwidth == 8 && bits.zero.toNat == 88 && bits.one.toNat == 165 then + "ok" + else + s!"unexpected join: {knownBitsToString joined}" + | .bottom => s!"unexpected join: {knownBitsToString joined}" + +/-- +info: "ok" +-/ +#guard_msgs in +#eval! runArithKnownBitsExample + +/-- +info: "ok" +-/ +#guard_msgs in +#eval! runLLVMAndCombKnownBitsExample + +/-- +info: "ok" +-/ +#guard_msgs in +#eval! runLLVMStyleTransfersExample + +/-- +info: "ok" +-/ +#guard_msgs in +#eval! testKnownBitsJoin + +end KnownBitsDataflow diff --git a/Veir/Analysis.lean b/Veir/Analysis.lean index 6c1f3fe8e1..ff22206ecb 100644 --- a/Veir/Analysis.lean +++ b/Veir/Analysis.lean @@ -7,3 +7,4 @@ public import Veir.Analysis.DataFlow.Printer public import Veir.Analysis.DataFlow.ModArithRangeAnalysis public import Veir.Analysis.DataFlow.SparseForwardDataFlowAnalysis public import Veir.Analysis.DataFlow.SparseConstantPropagationAnalysis +public import Veir.Analysis.DataFlow.KnownBitsAnalysis diff --git a/Veir/Analysis/DataFlow/Domains/KnownBitsDomain.lean b/Veir/Analysis/DataFlow/Domains/KnownBitsDomain.lean new file mode 100644 index 0000000000..0926dbda40 --- /dev/null +++ b/Veir/Analysis/DataFlow/Domains/KnownBitsDomain.lean @@ -0,0 +1,862 @@ +module + +public import Veir.Analysis.DataFlow.Domains.AbstractDomain +public import Veir.Interpreter.RuntimeValue +import Veir.Meta.Tactic.BVDecide + +public section + +namespace Veir + +/-! +# Known bits domain + +This file defines the abstract value used by known bits analysis. Known bits store +two masks: `zero` marks bits known to be zero and `one` marks bits known to be one. +Bits absent from both masks are unknown, and the masks are disjoint by construction. +The transfer operations follow LLVM's `KnownBits` implementation: + and +. +-/ + +/-- Two masks describing the known zero and known one bits of a fixed width integer. -/ +structure KnownBits where + bitwidth : Nat + zero : BitVec bitwidth + one : BitVec bitwidth + /-- No bit can be known to be both zero and one. -/ + disjoint : zero &&& one = 0 +deriving DecidableEq, Repr + +namespace KnownBits + +/-- Construct known bits while conservatively discarding any contradictory mask bits. -/ +def ofMasks {bitwidth : Nat} (zero one : BitVec bitwidth) : KnownBits := + let consistent := ~~~(zero &&& one) + { bitwidth + zero := zero &&& consistent + one := one &&& consistent + disjoint := by + ext i hi + simp [consistent] + grind } + +/-- No bits are known for an integer of the given width. -/ +def unknown (bitwidth : Nat) : KnownBits := + ofMasks (bitwidth := bitwidth) 0 0 + +/-- Every bit of a concrete integer is known. -/ +def constant (bitwidth : Nat) (value : Int) : KnownBits := + let bits := BitVec.ofInt bitwidth value + ofMasks (~~~bits) bits + +/-- Whether every bit has a known value. -/ +def isConstant (bits : KnownBits) : Bool := + bits.zero = ~~~bits.one + +/-- Render known bits from most to least significant using LLVM's `0`/`1`/`?` notation. -/ +def toPattern (bits : KnownBits) : String := + String.ofList <| (List.range bits.bitwidth).reverse.map fun i => + if bits.zero.getLsbD i then '0' + else if bits.one.getLsbD i then '1' + else '?' + +instance : ToString KnownBits where + toString := toPattern + +/-- The smallest unsigned value represented by these masks. -/ +def unsignedMin (bits : KnownBits) : BitVec bits.bitwidth := + bits.one + +/-- The largest unsigned value represented by these masks. -/ +def unsignedMax (bits : KnownBits) : BitVec bits.bitwidth := + ~~~bits.zero + +/-- The smallest signed value represented by these masks. -/ +def signedMin (bits : KnownBits) : BitVec bits.bitwidth := + if bits.bitwidth = 0 || bits.zero.msb then bits.one + else bits.one ||| BitVec.ofNat bits.bitwidth (2 ^ (bits.bitwidth - 1)) + +/-- The largest signed value represented by these masks. -/ +def signedMax (bits : KnownBits) : BitVec bits.bitwidth := + let max := ~~~bits.zero + if bits.bitwidth = 0 || bits.one.msb then max + else max &&& ~~~(BitVec.ofNat bits.bitwidth (2 ^ (bits.bitwidth - 1))) + +/-- A mask containing the lowest `count` bits. -/ +def lowMask (bitwidth count : Nat) : BitVec bitwidth := + BitVec.ofNat bitwidth (2 ^ (min count bitwidth) - 1) + +/-- A mask containing the highest `count` bits. -/ +def highMask (bitwidth count : Nat) : BitVec bitwidth := + ~~~lowMask bitwidth (bitwidth - min count bitwidth) + +/-- The number of consecutive one bits at the least-significant end. -/ +def countTrailingOnes {bitwidth : Nat} (value : BitVec bitwidth) : Nat := + (~~~value).ctz.toNat + +/-- The number of consecutive one bits at the most-significant end. -/ +def countLeadingOnes {bitwidth : Nat} (value : BitVec bitwidth) : Nat := + (~~~value).clz.toNat + +def countMinTrailingZeros (bits : KnownBits) : Nat := countTrailingOnes bits.zero +def countMaxTrailingZeros (bits : KnownBits) : Nat := bits.one.ctz.toNat +def countMinLeadingZeros (bits : KnownBits) : Nat := countLeadingOnes bits.zero +def countMaxLeadingZeros (bits : KnownBits) : Nat := bits.one.clz.toNat +def countMinLeadingOnes (bits : KnownBits) : Nat := countLeadingOnes bits.one +def countMaxLeadingOnes (bits : KnownBits) : Nat := bits.zero.clz.toNat + +/-- Keep only facts true of both possible results. -/ +def intersect (lhs rhs : KnownBits) : Option KnownBits := + if h : lhs.bitwidth = rhs.bitwidth then + some (ofMasks (lhs.zero &&& rhs.zero.cast h.symm) (lhs.one &&& rhs.one.cast h.symm)) + else + none + +/-- Known bits produced by bitwise AND. -/ +def bitwiseAnd? (lhs rhs : KnownBits) : Option KnownBits := + if h : lhs.bitwidth = rhs.bitwidth then + let rhsZero := rhs.zero.cast h.symm + let rhsOne := rhs.one.cast h.symm + some (ofMasks (lhs.zero ||| rhsZero) (lhs.one &&& rhsOne)) + else + none + +/-- Known bits produced by bitwise OR. -/ +def bitwiseOr? (lhs rhs : KnownBits) : Option KnownBits := + if h : lhs.bitwidth = rhs.bitwidth then + let rhsZero := rhs.zero.cast h.symm + let rhsOne := rhs.one.cast h.symm + some (ofMasks (lhs.zero &&& rhsZero) (lhs.one ||| rhsOne)) + else + none + +/-- Known bits produced by bitwise XOR. -/ +def bitwiseXor? (lhs rhs : KnownBits) : Option KnownBits := + if h : lhs.bitwidth = rhs.bitwidth then + let rhsZero := rhs.zero.cast h.symm + let rhsOne := rhs.one.cast h.symm + some (ofMasks + ((lhs.zero &&& rhsZero) ||| (lhs.one &&& rhsOne)) + ((lhs.zero &&& rhsOne) ||| (lhs.one &&& rhsZero))) + else + none + +/-- Combine independent facts about the same value, returning unknown on contradiction. -/ +def refineWith (lhs rhs : KnownBits) : Option KnownBits := + if h : lhs.bitwidth = rhs.bitwidth then + let zero := lhs.zero ||| rhs.zero.cast h.symm + let one := lhs.one ||| rhs.one.cast h.symm + if zero &&& one = 0 then some (ofMasks zero one) else some (unknown lhs.bitwidth) + else + none + +/-- Known common prefix of every value in the inclusive unsigned interval `[lower, upper]`. -/ +def fromUnsignedInterval {bitwidth : Nat} + (lower upper : BitVec bitwidth) : KnownBits := + if upper.ult lower then + unknown bitwidth + else + let differing := lower ^^^ upper + let unknownBits := bitwidth - differing.clz.toNat + let known := ~~~lowMask bitwidth unknownBits + ofMasks (~~~lower &&& known) (lower &&& known) + +/-- Retain the low `newWidth` bits, matching LLVM's `KnownBits::trunc`. -/ +def trunc (bits : KnownBits) (newWidth : Nat) : KnownBits := + ofMasks (bits.zero.setWidth newWidth) (bits.one.setWidth newWidth) + +/-- Extend with unknown high bits, matching LLVM's `KnownBits::anyext`. -/ +def anyext (bits : KnownBits) (newWidth : Nat) : KnownBits := + trunc bits newWidth + +/-- Zero-extend known bits, marking every new high bit as zero. -/ +def zext (bits : KnownBits) (newWidth : Nat) : KnownBits := + let added := newWidth - bits.bitwidth + ofMasks (bits.zero.setWidth newWidth ||| highMask newWidth added) (bits.one.setWidth newWidth) + +/-- Sign-extend known bits. Unknown sign bits produce unknown extension bits. -/ +def sext (bits : KnownBits) (newWidth : Nat) : KnownBits := + ofMasks (bits.zero.signExtend newWidth) (bits.one.signExtend newWidth) + +/-- Extract `width` bits beginning at `lowBit`. -/ +def extract (bits : KnownBits) (lowBit width : Nat) : KnownBits := + ofMasks (bits.zero.extractLsb' lowBit width) (bits.one.extractLsb' lowBit width) + +/-- Concatenate two known-bit values. -/ +def concat (high low : KnownBits) : KnownBits := + ofMasks (high.zero ++ low.zero) (high.one ++ low.one) + +/-- Reverse the order of the bits. -/ +def reverse (bits : KnownBits) : KnownBits := + ofMasks bits.zero.reverse bits.one.reverse + +/-- Compute the base known bits for `lhs + rhs + carry`. -/ +def addCarry? (lhs rhs : KnownBits) (carryZero carryOne : Bool) : Option KnownBits := + if h : lhs.bitwidth = rhs.bitwidth then + let rhsZero := rhs.zero.cast h.symm + let rhsOne := rhs.one.cast h.symm + let possibleZero := (~~~lhs.zero) + (~~~rhsZero) + BitVec.ofNat lhs.bitwidth (!carryZero).toNat + let possibleOne := lhs.one + rhsOne + BitVec.ofNat lhs.bitwidth carryOne.toNat + let carryKnownZero := ~~~(possibleZero ^^^ lhs.zero ^^^ rhsZero) + let carryKnownOne := possibleOne ^^^ lhs.one ^^^ rhsOne + let known := (lhs.zero ||| lhs.one) &&& (rhsZero ||| rhsOne) &&& + (carryKnownZero ||| carryKnownOne) + some (ofMasks (~~~possibleZero &&& known) (possibleOne &&& known)) + else + none + +/-- A mask of `count` high bits immediately below the sign bit. -/ +private def highBitsBelowSign (bitwidth count : Nat) : BitVec bitwidth := + highMask bitwidth (count + 1) &&& ~~~highMask bitwidth 1 + +/-- LLVM's shared transfer for addition and subtraction with no-wrap flags. -/ +private def addSub? (add : Bool) (lhs rhs : KnownBits) (nsw nuw : Bool) : Option KnownBits := + if h : lhs.bitwidth = rhs.bitwidth then + let rhs := ofMasks (rhs.zero.cast h.symm) (rhs.one.cast h.symm) + let width := lhs.bitwidth + if width = 0 then + some (unknown 0) + else + let base := + if lhs.zero = 0 && lhs.one = 0 || rhs.zero = 0 && rhs.one = 0 then + unknown width + else if add then + (addCarry? lhs rhs true false).getD (unknown width) + else + (addCarry? lhs (ofMasks rhs.one rhs.zero) false true).getD (unknown width) + let (zero, one) := Id.run do + let mut zero := base.zero.setWidth width + let mut one := base.one.setWidth width + if nuw then + if add then + let maximum := 2 ^ width - 1 + let minimum := BitVec.ofNat width + (min maximum (lhs.unsignedMin.toNat + rhs.unsignedMin.toNat)) + if nsw then + let count := countLeadingOnes (minimum.setWidth (width - 1)) + one := one ||| highBitsBelowSign width count + one := one ||| highMask width (countLeadingOnes minimum) + else + let maximum := BitVec.ofNat width + (lhs.unsignedMax.toNat - rhs.unsignedMin.toNat) + if nsw then + let count := (maximum.setWidth (width - 1)).clz.toNat + zero := zero ||| highBitsBelowSign width count + zero := zero ||| highMask width maximum.clz.toNat + if nsw then + let signedLower := -(2 ^ (width - 1) : Int) + let signedUpper := (2 ^ (width - 1) : Int) - 1 + let rawMinimum := if add then lhs.signedMin.toInt + rhs.signedMin.toInt + else lhs.signedMin.toInt - rhs.signedMax.toInt + let rawMaximum := if add then lhs.signedMax.toInt + rhs.signedMax.toInt + else lhs.signedMax.toInt - rhs.signedMin.toInt + let minimum := BitVec.ofInt width (max signedLower (min signedUpper rawMinimum)) + let maximum := BitVec.ofInt width (max signedLower (min signedUpper rawMaximum)) + if !minimum.msb then + let count := countLeadingOnes (minimum.setWidth (width - 1)) + one := one ||| highBitsBelowSign width count + zero := zero ||| highMask width 1 + if maximum.msb then + let count := (maximum.setWidth (width - 1)).clz.toNat + zero := zero ||| highBitsBelowSign width count + one := one ||| highMask width 1 + return (zero, one) + if zero &&& one ≠ 0 then some (constant width 0) + else some (ofMasks zero one) + else + none + +/-- Known bits for addition, including information supplied by `nsw` and `nuw`. -/ +def add? (lhs rhs : KnownBits) (nsw nuw : Bool := false) : Option KnownBits := + addSub? true lhs rhs nsw nuw + +/-- Known bits for subtraction, including information supplied by `nsw` and `nuw`. -/ +def sub? (lhs rhs : KnownBits) (nsw nuw : Bool := false) : Option KnownBits := + addSub? false lhs rhs nsw nuw + +/-- LLVM-style known bits for multiplication. -/ +def mul? (lhs rhs : KnownBits) : Option KnownBits := + if h : lhs.bitwidth = rhs.bitwidth then + let rhsZero := rhs.zero.cast h.symm + let rhsOne := rhs.one.cast h.symm + let width := lhs.bitwidth + let maxProduct := lhs.unsignedMax.toNat * (~~~rhsZero).toNat + let leadingZeros := if maxProduct < 2 ^ width + then (BitVec.ofNat width maxProduct).clz.toNat else 0 + let lhsKnownLow := countTrailingOnes (lhs.zero ||| lhs.one) + let rhsKnownLow := countTrailingOnes (rhsZero ||| rhsOne) + let lhsTrailingZeros := countMinTrailingZeros lhs + let rhsTrailingZeros := countTrailingOnes rhsZero + let trailingZeros := lhsTrailingZeros + rhsTrailingZeros + let smallestKnown := min (lhsKnownLow - lhsTrailingZeros) + (rhsKnownLow - rhsTrailingZeros) + let resultKnownLow := min (smallestKnown + trailingZeros) width + let bottomKnown := lhs.one * rhsOne + let low := lowMask width resultKnownLow + some (ofMasks (highMask width leadingZeros ||| (~~~bottomKnown &&& low)) (bottomKnown &&& low)) + else + none + +/-- Whether a concrete unsigned value satisfies the known-bit masks. -/ +def containsNat (bits : KnownBits) (value : Nat) : Bool := + let value := BitVec.ofNat bits.bitwidth value + value &&& bits.zero = 0 && value &&& bits.one = bits.one + +/-- LLVM's upper bound for a possibly out-of-range shift amount. -/ +private def maxShiftAmount (rhs : KnownBits) (bitwidth : Nat) : Nat := + if bitwidth = 0 then 0 + else if bitwidth &&& (bitwidth - 1) = 0 then + let extractedWidth := min (Nat.log2 bitwidth) rhs.bitwidth + rhs.unsignedMax.toNat % (2 ^ extractedWidth) + else + min rhs.unsignedMax.toNat (bitwidth - 1) + +/-- Intersect the facts produced by LLVM's feasible shift-amount range. -/ +private def forEachShiftAmount + (lhs rhs : KnownBits) + (minimum maximum : Nat) + (transfer : Nat → Option KnownBits) : KnownBits := Id.run do + let mut result : Option KnownBits := none + for amount in List.range (maximum + 1) do + if minimum ≤ amount && rhs.containsNat amount then + if let some shifted := transfer amount then + result := match result with + | none => some shifted + | some current => current.intersect shifted + -- LLVM uses zero as the non-conflicting representative when every result is poison. + return result.getD (constant lhs.bitwidth 0) + +/-- Known bits for a left shift, including `nsw` and `nuw` constraints. -/ +def shl (lhs rhs : KnownBits) (nsw nuw : Bool := false) : KnownBits := + let minimum := min rhs.unsignedMin.toNat lhs.bitwidth + let initialMaximum := maxShiftAmount rhs lhs.bitwidth + let maximum := Id.run do + let mut maximum := initialMaximum + if nuw && nsw then + let count := lhs.countMaxLeadingZeros + if count ≠ 0 then maximum := min maximum (count - 1) + if nuw then maximum := min maximum lhs.countMaxLeadingZeros + if nsw then + let count := max lhs.countMaxLeadingZeros lhs.countMaxLeadingOnes + if count ≠ 0 then maximum := min maximum (count - 1) + return maximum + forEachShiftAmount lhs rhs minimum maximum fun amount => + let shiftedZero := (lhs.zero <<< amount) ||| lowMask lhs.bitwidth amount + let shiftedOne := lhs.one <<< amount + if nsw then + let shiftedOutZero := + nuw && amount ≠ 0 || lhs.zero &&& highMask lhs.bitwidth amount ≠ 0 + let shiftedOutOne := lhs.one &&& highMask lhs.bitwidth amount ≠ 0 + let zero := if shiftedOutZero then shiftedZero ||| highMask lhs.bitwidth 1 + else shiftedZero + let one := if !shiftedOutZero && shiftedOutOne then + shiftedOne ||| highMask lhs.bitwidth 1 else shiftedOne + if zero &&& one ≠ 0 then none else some (ofMasks zero one) + else + some (ofMasks shiftedZero shiftedOne) + +/-- Known bits for a logical right shift. -/ +def lshr (lhs rhs : KnownBits) (exact : Bool := false) : KnownBits := + let minimum := min rhs.unsignedMin.toNat lhs.bitwidth + let maximum := maxShiftAmount rhs lhs.bitwidth + let maximum := if exact then min maximum lhs.countMaxTrailingZeros else maximum + forEachShiftAmount lhs rhs minimum maximum fun amount => + some <| ofMasks + ((lhs.zero >>> amount) ||| highMask lhs.bitwidth amount) + (lhs.one >>> amount) + +/-- Known bits for an arithmetic right shift. -/ +def ashr (lhs rhs : KnownBits) (exact : Bool := false) : KnownBits := + let minimum := min rhs.unsignedMin.toNat lhs.bitwidth + let maximum := maxShiftAmount rhs lhs.bitwidth + let maximum := if exact then min maximum lhs.countMaxTrailingZeros else maximum + forEachShiftAmount lhs rhs minimum maximum fun amount => + some <| ofMasks (lhs.zero.sshiftRight amount) (lhs.one.sshiftRight amount) + +/-- Add trailing-bit facts implied by an exact division. -/ +private def refineExactDivision + (result lhs rhs : KnownBits) (exact : Bool) : KnownBits := + if !exact || lhs.bitwidth = 0 then + result + else + let minimum : Int := lhs.countMinTrailingZeros - rhs.countMaxTrailingZeros + let maximum : Int := lhs.countMaxTrailingZeros - rhs.countMinTrailingZeros + if maximum < 0 then + constant lhs.bitwidth 0 + else + let (zero, one) := Id.run do + let mut zero := result.zero.setWidth lhs.bitwidth + let mut one := result.one.setWidth lhs.bitwidth + if lhs.one.getLsbD 0 then one := one ||| BitVec.ofNat lhs.bitwidth 1 + if 0 ≤ minimum then + let trailing := minimum.toNat + zero := zero ||| lowMask lhs.bitwidth trailing + if minimum = maximum then + one := one ||| BitVec.ofNat lhs.bitwidth (2 ^ trailing) + return (zero, one) + if zero &&& one ≠ 0 then constant lhs.bitwidth 0 else ofMasks zero one + +/-- Known bits for unsigned division, following LLVM's upper-zero-bit estimate. -/ +def udiv? (lhs rhs : KnownBits) (exact : Bool := false) : Option KnownBits := + if h : lhs.bitwidth = rhs.bitwidth then + let rhs := ofMasks (rhs.zero.cast h.symm) (rhs.one.cast h.symm) + if lhs.isConstant && lhs.one = 0 || rhs.isConstant && rhs.one = 0 then + some (constant lhs.bitwidth 0) + else + let maximumResult := if rhs.unsignedMin = 0 then lhs.unsignedMax + else lhs.unsignedMax.udiv rhs.unsignedMin + let result := ofMasks (highMask lhs.bitwidth maximumResult.clz.toNat) 0 + some (refineExactDivision result lhs rhs exact) + else + none + +/-- Known bits for signed division. The non-negative case has full unsigned precision. -/ +def sdiv? (lhs rhs : KnownBits) (exact : Bool := false) : Option KnownBits := + if h : lhs.bitwidth = rhs.bitwidth then + let rhs := ofMasks (rhs.zero.cast h.symm) (rhs.one.cast h.symm) + if lhs.isConstant && lhs.one = 0 || rhs.isConstant && rhs.one = 0 then + some (constant lhs.bitwidth 0) + else if lhs.zero.msb && rhs.zero.msb then + udiv? lhs rhs exact + else + let estimate : Option (BitVec lhs.bitwidth) := + if lhs.one.msb && rhs.one.msb then + let numerator := lhs.signedMin + let denominator := rhs.signedMax + let intMin := BitVec.ofNat lhs.bitwidth (2 ^ (lhs.bitwidth - 1)) + let signedMax := BitVec.ofNat lhs.bitwidth (2 ^ (lhs.bitwidth - 1) - 1) + if numerator == intMin && denominator == ~~~(0 : BitVec lhs.bitwidth) then + some signedMax + else + some (numerator.sdiv denominator) + else if lhs.one.msb && rhs.zero.msb && + (exact || (-lhs.signedMax).toNat ≥ rhs.signedMax.toNat) then + let denominator := rhs.signedMin + some (if denominator = 0 then lhs.signedMin else lhs.signedMin.sdiv denominator) + else if lhs.zero.msb && lhs.one ≠ 0 && rhs.one.msb && + (exact || lhs.signedMin.toNat ≥ (-rhs.signedMin).toNat) then + some (lhs.signedMax.sdiv rhs.signedMax) + else + none + let result := match estimate with + | some value => + if value.msb then ofMasks 0 (highMask lhs.bitwidth (countLeadingOnes value)) + else ofMasks (highMask lhs.bitwidth value.clz.toNat) 0 + | none => unknown lhs.bitwidth + some (refineExactDivision result lhs rhs exact) + else + none + +/-- Preserve low dividend bits when the divisor is known to be even. -/ +private def remLowBits (lhs rhs : KnownBits) : KnownBits := + if rhs.unsignedMax ≠ 0 && rhs.zero.getLsbD 0 then + let mask := lowMask lhs.bitwidth rhs.countMinTrailingZeros + ofMasks (lhs.zero &&& mask) (lhs.one &&& mask) + else + unknown lhs.bitwidth + +/-- Known bits for unsigned remainder. -/ +def urem? (lhs rhs : KnownBits) : Option KnownBits := + if h : lhs.bitwidth = rhs.bitwidth then + let rhs := ofMasks (rhs.zero.cast h.symm) (rhs.one.cast h.symm) + let low := remLowBits lhs rhs + if rhs.isConstant && rhs.one.toNat ≠ 0 && rhs.one.toNat &&& (rhs.one.toNat - 1) = 0 then + let high := ~~~(BitVec.ofNat lhs.bitwidth (rhs.one.toNat - 1)) + some (ofMasks (low.zero.setWidth lhs.bitwidth ||| high) + (low.one.setWidth lhs.bitwidth)) + else + let leaders := max lhs.countMinLeadingZeros rhs.countMinLeadingZeros + some (ofMasks (low.zero.setWidth lhs.bitwidth ||| highMask lhs.bitwidth leaders) + (low.one.setWidth lhs.bitwidth)) + else + none + +/-- Known bits for signed remainder, including the exact power-of-two case. -/ +def srem? (lhs rhs : KnownBits) : Option KnownBits := + if h : lhs.bitwidth = rhs.bitwidth then + let rhs := ofMasks (rhs.zero.cast h.symm) (rhs.one.cast h.symm) + let lowFacts := remLowBits lhs rhs + if rhs.isConstant && rhs.one.toNat ≠ 0 && rhs.one.toNat &&& (rhs.one.toNat - 1) = 0 then + let low := BitVec.ofNat lhs.bitwidth (rhs.one.toNat - 1) + if lhs.zero.msb || low &&& lhs.zero = low then + some (ofMasks (lowFacts.zero.setWidth lhs.bitwidth ||| ~~~low) + (lowFacts.one.setWidth lhs.bitwidth)) + else if lhs.one.msb && low &&& lhs.one ≠ 0 then + some (ofMasks (lowFacts.zero.setWidth lhs.bitwidth) + (lowFacts.one.setWidth lhs.bitwidth ||| ~~~low)) + else + some lowFacts + else + let rhsSignBits := if rhs.zero.msb then rhs.countMinLeadingZeros + else if rhs.one.msb then rhs.countMinLeadingOnes else 1 + if lhs.one.msb && lowFacts.one ≠ 0 then + some (ofMasks (lowFacts.zero.setWidth lhs.bitwidth) + (lowFacts.one.setWidth lhs.bitwidth ||| highMask lhs.bitwidth + (max lhs.countMinLeadingOnes rhsSignBits))) + else if lhs.zero.msb then + some (ofMasks + (lowFacts.zero.setWidth lhs.bitwidth ||| highMask lhs.bitwidth + (max lhs.countMinLeadingZeros rhsSignBits)) + (lowFacts.one.setWidth lhs.bitwidth)) + else + some lowFacts + else + none + +/-- Known result of an integer comparison, when it can be decided from the masks. -/ +def compare? (predicate : Data.LLVM.IntPred) (lhs rhs : KnownBits) : Option Bool := + if h : lhs.bitwidth = rhs.bitwidth then + let rhsZero := rhs.zero.cast h.symm + let rhsOne := rhs.one.cast h.symm + let rhs := ofMasks rhsZero rhsOne + match predicate with + | .eq => + if lhs.isConstant && rhs.isConstant then some (lhs.one == rhsOne) + else if lhs.one &&& rhsZero ≠ 0 || rhsOne &&& lhs.zero ≠ 0 then some false + else none + | .ne => + if lhs.isConstant && rhs.isConstant then some (lhs.one != rhsOne) + else if lhs.one &&& rhsZero ≠ 0 || rhsOne &&& lhs.zero ≠ 0 then some true + else none + | .ugt => + if lhs.unsignedMax.ule rhs.unsignedMin then some false + else if rhs.unsignedMax.ult lhs.unsignedMin then some true + else none + | .uge => + if rhs.unsignedMax.ult lhs.unsignedMin then some true + else if lhs.unsignedMax.ult rhs.unsignedMin then some false + else none + | .ult => + if rhs.unsignedMax.ule lhs.unsignedMin then some false + else if lhs.unsignedMax.ult rhs.unsignedMin then some true + else none + | .ule => + if lhs.unsignedMax.ule rhs.unsignedMin then some true + else if rhs.unsignedMax.ult lhs.unsignedMin then some false + else none + | .sgt => + if lhs.signedMax.sle rhs.signedMin then some false + else if rhs.signedMax.slt lhs.signedMin then some true + else none + | .sge => + if rhs.signedMax.slt lhs.signedMin then some true + else if lhs.signedMax.slt rhs.signedMin then some false + else none + | .slt => + if rhs.signedMax.sle lhs.signedMin then some false + else if lhs.signedMax.slt rhs.signedMin then some true + else none + | .sle => + if lhs.signedMax.sle rhs.signedMin then some true + else if rhs.signedMax.slt lhs.signedMin then some false + else none + else + none + +/-- Known bits for a comparison result. -/ +def compare (predicate : Data.LLVM.IntPred) (lhs rhs : KnownBits) : KnownBits := + match compare? predicate lhs rhs with + | some value => constant 1 (if value then 1 else 0) + | none => unknown 1 + +/-- Known bits for unsigned maximum. -/ +private def makeGE (bits : KnownBits) (value : BitVec bits.bitwidth) : KnownBits := + let leading := countLeadingOnes (bits.zero ||| value) + let forcedOnes := value &&& highMask bits.bitwidth leading + ofMasks bits.zero (bits.one ||| forcedOnes) + +private def complementValue (bits : KnownBits) : KnownBits := + ofMasks bits.one bits.zero + +private def flipSignBit (bits : KnownBits) : KnownBits := + let sign := highMask bits.bitwidth 1 + ofMasks + ((bits.zero &&& ~~~sign) ||| (bits.one &&& sign)) + ((bits.one &&& ~~~sign) ||| (bits.zero &&& sign)) + +def umax? (lhs rhs : KnownBits) : Option KnownBits := + if h : lhs.bitwidth = rhs.bitwidth then + let rhs := ofMasks (rhs.zero.cast h.symm) (rhs.one.cast h.symm) + if rhs.unsignedMax.toNat ≤ lhs.unsignedMin.toNat then some lhs + else if lhs.unsignedMax.toNat ≤ rhs.unsignedMin.toNat then some rhs + else (lhs.makeGE rhs.unsignedMin).intersect (rhs.makeGE lhs.unsignedMin) + else + none + +/-- Known bits for unsigned minimum. -/ +def umin? (lhs rhs : KnownBits) : Option KnownBits := + if lhs.bitwidth ≠ rhs.bitwidth then none else do + let maximum ← lhs.complementValue.umax? rhs.complementValue + return maximum.complementValue + +/-- Known bits for signed maximum. -/ +def smax? (lhs rhs : KnownBits) : Option KnownBits := + if lhs.bitwidth ≠ rhs.bitwidth then none else do + let maximum ← lhs.flipSignBit.umax? rhs.flipSignBit + return maximum.flipSignBit + +/-- Known bits for signed minimum. -/ +def smin? (lhs rhs : KnownBits) : Option KnownBits := + if lhs.bitwidth ≠ rhs.bitwidth then none else do + let minimum ← lhs.flipSignBit.umin? rhs.flipSignBit + return minimum.flipSignBit + +/-- Known bits for population count from the attainable result interval. -/ +def ctpop (bits : KnownBits) : KnownBits := + let lower := bits.one.cpop.toNat + let upper := bits.bitwidth - bits.zero.cpop.toNat + fromUnsignedInterval (BitVec.ofNat bits.bitwidth lower) (BitVec.ofNat bits.bitwidth upper) + +/-- Known bits for count-leading-zeros. -/ +def ctlz (bits : KnownBits) (zeroIsPoison : Bool := false) : KnownBits := + let lower := bits.countMinLeadingZeros + let upper := if zeroIsPoison && bits.one = 0 then bits.bitwidth - 1 + else bits.countMaxLeadingZeros + fromUnsignedInterval (BitVec.ofNat bits.bitwidth lower) (BitVec.ofNat bits.bitwidth upper) + +/-- Known bits for count-trailing-zeros. -/ +def cttz (bits : KnownBits) (zeroIsPoison : Bool := false) : KnownBits := + let lower := bits.countMinTrailingZeros + let upper := if zeroIsPoison && bits.one = 0 then bits.bitwidth - 1 + else bits.countMaxTrailingZeros + fromUnsignedInterval (BitVec.ofNat bits.bitwidth lower) (BitVec.ofNat bits.bitwidth upper) + +/-- High half of an unsigned multiplication. -/ +def mulhu? (lhs rhs : KnownBits) : Option KnownBits := do + if lhs.bitwidth ≠ rhs.bitwidth then none else + let width := lhs.bitwidth + let product ← (lhs.zext (2 * width)).mul? (rhs.zext (2 * width)) + return product.extract width width + +/-- High half of a signed multiplication. -/ +def mulhs? (lhs rhs : KnownBits) : Option KnownBits := do + if lhs.bitwidth ≠ rhs.bitwidth then none else + let width := lhs.bitwidth + let product ← (lhs.sext (2 * width)).mul? (rhs.sext (2 * width)) + return product.extract width width + +/-- Unsigned saturating addition. -/ +def uaddSat? (lhs rhs : KnownBits) : Option KnownBits := + if lhs.bitwidth ≠ rhs.bitwidth then none else + let limit := 2 ^ lhs.bitwidth - 1 + let lower := min limit (lhs.unsignedMin.toNat + rhs.unsignedMin.toNat) + let upper := min limit (lhs.unsignedMax.toNat + rhs.unsignedMax.toNat) + some (fromUnsignedInterval (BitVec.ofNat lhs.bitwidth lower) (BitVec.ofNat lhs.bitwidth upper)) + +/-- Unsigned saturating subtraction. -/ +def usubSat? (lhs rhs : KnownBits) : Option KnownBits := + if lhs.bitwidth ≠ rhs.bitwidth then none else + let lower := lhs.unsignedMin.toNat - rhs.unsignedMax.toNat + let upper := lhs.unsignedMax.toNat - rhs.unsignedMin.toNat + some (fromUnsignedInterval (BitVec.ofNat lhs.bitwidth lower) (BitVec.ofNat lhs.bitwidth upper)) + +private def fromSignedInterval (bitwidth : Nat) (lower upper : Int) : KnownBits := + let lower := BitVec.ofInt bitwidth lower + let upper := BitVec.ofInt bitwidth upper + if lower.msb = upper.msb then fromUnsignedInterval lower upper else unknown bitwidth + +/-- Signed saturating addition. -/ +def saddSat? (lhs rhs : KnownBits) : Option KnownBits := + if lhs.bitwidth ≠ rhs.bitwidth then none else + if lhs.bitwidth = 0 then some (unknown 0) else + let lowerLimit := -(2 ^ (lhs.bitwidth - 1) : Int) + let upperLimit := (2 ^ (lhs.bitwidth - 1) : Int) - 1 + let lower := max lowerLimit (lhs.signedMin.toInt + rhs.signedMin.toInt) + let upper := min upperLimit (lhs.signedMax.toInt + rhs.signedMax.toInt) + some (fromSignedInterval lhs.bitwidth lower upper) + +/-- Signed saturating subtraction. -/ +def ssubSat? (lhs rhs : KnownBits) : Option KnownBits := + if lhs.bitwidth ≠ rhs.bitwidth then none else + if lhs.bitwidth = 0 then some (unknown 0) else + let lowerLimit := -(2 ^ (lhs.bitwidth - 1) : Int) + let upperLimit := (2 ^ (lhs.bitwidth - 1) : Int) - 1 + let lower := max lowerLimit (lhs.signedMin.toInt - rhs.signedMax.toInt) + let upper := min upperLimit (lhs.signedMax.toInt - rhs.signedMin.toInt) + some (fromSignedInterval lhs.bitwidth lower upper) + +/-- Absolute value, preserving LLVM's useful sign and trailing-zero facts. -/ +def abs (bits : KnownBits) (intMinIsPoison : Bool := false) : KnownBits := + if bits.bitwidth = 0 || bits.zero.msb then + bits + else if bits.one.msb then + (constant bits.bitwidth 0).sub? bits intMinIsPoison false |>.getD (unknown bits.bitwidth) + else + let minTrailing := bits.countMinTrailingZeros + let maxTrailing := bits.countMaxTrailingZeros + let sign := highMask bits.bitwidth 1 + let signZero := if intMinIsPoison || bits.one ≠ 0 && bits.one ≠ sign then sign else 0 + let lowestOne := if minTrailing = maxTrailing && minTrailing < bits.bitwidth + then BitVec.ofNat bits.bitwidth (2 ^ minTrailing) else 0 + ofMasks (lowMask bits.bitwidth minTrailing ||| signZero) lowestOne + +private def byteSwapMask {bitwidth : Nat} (value : BitVec bitwidth) : BitVec bitwidth := + if bitwidth % 8 ≠ 0 then value else + let bytes := bitwidth / 8 + let swapped := (List.range bitwidth).foldl (init := 0) fun result source => + if value.getLsbD source then + let target := (bytes - 1 - source / 8) * 8 + source % 8 + result ||| 2 ^ target + else + result + BitVec.ofNat bitwidth swapped + +/-- Reverse the byte order when the bitwidth is byte-aligned. -/ +def byteSwap (bits : KnownBits) : KnownBits := + ofMasks (byteSwapMask bits.zero) (byteSwapMask bits.one) + +/-- Repeat an input bit pattern to a requested result width. -/ +def replicate (bits : KnownBits) (resultWidth : Nat) : KnownBits := + if bits.bitwidth = 0 || resultWidth % bits.bitwidth ≠ 0 then + unknown resultWidth + else + let repetitions := resultWidth / bits.bitwidth + let repeatMask (value : BitVec bits.bitwidth) := + let repeated := (List.range repetitions).foldl (init := 0) fun result index => + result ||| value.toNat <<< (index * bits.bitwidth) + BitVec.ofNat resultWidth repeated + ofMasks (repeatMask bits.zero) (repeatMask bits.one) + +/-- Known bits for a funnel shift left by a concrete amount. -/ +def fshl? (lhs rhs : KnownBits) (amount : Nat) : Option KnownBits := + if h : lhs.bitwidth = rhs.bitwidth then + if lhs.bitwidth = 0 then some (unknown 0) else + let rhsZero := rhs.zero.cast h.symm + let rhsOne := rhs.one.cast h.symm + let amount := amount % lhs.bitwidth + if amount = 0 then some lhs else + some (ofMasks + ((lhs.zero <<< amount) ||| (rhsZero >>> (lhs.bitwidth - amount))) + ((lhs.one <<< amount) ||| (rhsOne >>> (lhs.bitwidth - amount)))) + else + none + +/-- Known bits for a funnel shift right by a concrete amount. -/ +def fshr? (lhs rhs : KnownBits) (amount : Nat) : Option KnownBits := + if h : lhs.bitwidth = rhs.bitwidth then + if lhs.bitwidth = 0 then some (unknown 0) else + let rhsZero := rhs.zero.cast h.symm + let rhsOne := rhs.one.cast h.symm + let amount := amount % lhs.bitwidth + if amount = 0 then some (ofMasks rhsZero rhsOne) else + some (ofMasks + ((rhsZero >>> amount) ||| (lhs.zero <<< (lhs.bitwidth - amount))) + ((rhsOne >>> amount) ||| (lhs.one <<< (lhs.bitwidth - amount)))) + else + none + +/-- Keep only facts known on both incoming control flow paths. -/ +def join? (lhs rhs : KnownBits) : Option KnownBits := + lhs.intersect rhs + +end KnownBits + +/-- +Sparse lattice for known bits. `bottom` is an uninitialized sparse value and `known` +contains width aware masks. A `known` value with two zero masks is the unique representation +of an integer for which no bits are known. +-/ +inductive KnownBitsLattice where + | bottom + | known (bits : KnownBits) +deriving DecidableEq, Repr + +namespace KnownBitsLattice + +instance : Bot KnownBitsLattice where + bot := .bottom + +/-- No bit facts are known, but the integer width is known. -/ +def unknown (bitwidth : Nat) : KnownBitsLattice := + .known (KnownBits.unknown bitwidth) + +/-- An exact fixed width integer value. -/ +def constant (bitwidth : Nat) (value : Int) : KnownBitsLattice := + .known (KnownBits.constant bitwidth value) + +instance : ToString KnownBitsLattice where + toString + | .bottom => "bottom" + | .known bits => bits.toPattern + +/-- The concrete runtime values represented by a known bits lattice element. -/ +@[expose] def γ : KnownBitsLattice → Set RuntimeValue + | .bottom => ⊥ + | .known bits => fun concrete => + match concrete with + | .int bitwidth (.val value) => + ∃ h : bitwidth = bits.bitwidth, + let value := value.cast h + value &&& bits.zero = 0 ∧ value &&& bits.one = bits.one + | _ => False + +@[simp] theorem not_mem_γ_bottom (value : RuntimeValue) : value ∉ γ .bottom := fun h => h.elim + +/-- Normalize membership in a known-bits value to masks at the concrete value's width. -/ +theorem mem_γ_known_masks_iff + {bits : KnownBits} + {bitwidth : Nat} + {value : BitVec bitwidth} : + RuntimeValue.int bitwidth (.val value) ∈ γ (.known bits) ↔ + ∃ (zero one : BitVec bitwidth) (disjoint : zero &&& one = 0), + bits = ⟨bitwidth, zero, one, disjoint⟩ ∧ + value &&& zero = 0 ∧ value &&& one = one := by + constructor + · rcases bits with ⟨bitsWidth, zero, one, disjoint⟩ + rintro ⟨hwidth, hzero, hone⟩ + change bitwidth = bitsWidth at hwidth + subst bitsWidth + simp at hzero hone + exact ⟨zero, one, disjoint, rfl, hzero, hone⟩ + · rintro ⟨zero, one, disjoint, rfl, hzero, hone⟩ + exact ⟨rfl, hzero, hone⟩ + +/-- Characterize known-bits membership as facts about each concrete bit. -/ +theorem mem_γ_known_iff + {bits : KnownBits} + {bitwidth : Nat} + {value : BitVec bitwidth} : + RuntimeValue.int bitwidth (.val value) ∈ γ (.known bits) ↔ + ∃ (zero one : BitVec bitwidth) (disjoint : zero &&& one = 0), + bits = ⟨bitwidth, zero, one, disjoint⟩ ∧ + (∀ i (hi : i < bitwidth), zero[i] = true → value[i] = false) ∧ + (∀ i (hi : i < bitwidth), one[i] = true → value[i] = true) := by + rw [mem_γ_known_masks_iff] + constructor + · rintro ⟨zero, one, disjoint, hbits, hzero, hone⟩ + refine ⟨zero, one, disjoint, hbits, ?_, ?_⟩ + · intro i hi hzeroTrue + have hzeroBit := congrArg (fun value => value[i]) hzero + simp at hzeroBit + veir_bv_decide + · intro i hi honeTrue + have honeBit := congrArg (fun value => value[i]) hone + simp at honeBit + veir_bv_decide + · rintro ⟨zero, one, disjoint, hbits, hzero, hone⟩ + refine ⟨zero, one, disjoint, hbits, ?_, ?_⟩ + · ext i hi + have hzeroBit := hzero i hi + simp at hzeroBit ⊢ + veir_bv_decide + · ext i hi + have honeBit := hone i hi + simp at honeBit ⊢ + veir_bv_decide + +/-- Join facts arriving along different control flow paths. -/ +def join : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice + | .bottom, rhs => rhs + | lhs, .bottom => lhs + | .known lhs, .known rhs => + match lhs.join? rhs with + | some bits => .known bits + | none => .unknown lhs.bitwidth + +instance : Join KnownBitsLattice where + join := join + +end KnownBitsLattice + +end Veir diff --git a/Veir/Analysis/DataFlow/Facts.lean b/Veir/Analysis/DataFlow/Facts.lean index 266c5840b2..2f01936a60 100644 --- a/Veir/Analysis/DataFlow/Facts.lean +++ b/Veir/Analysis/DataFlow/Facts.lean @@ -2,6 +2,7 @@ module public import Veir.GlobalOpInfo public import Veir.Analysis.DataFlow.Domains.IntegerRangeDomain +public import Veir.Analysis.DataFlow.Domains.KnownBitsDomain public import Veir.Analysis.DataFlow.Domains.LivenessDomain public import Veir.Rewriter.InsertPoint public import Veir.Analysis.DataFlow.Domains.ConstantDomain @@ -75,6 +76,7 @@ inductive AnalysisKind where | sparseConstantPropagation | integerRange | modArithRange + | knownBits deriving BEq, Hashable, Repr, DecidableEq /-- @@ -89,6 +91,7 @@ inductive FactKind where | sparseConstant | integerRange | modArithRange + | knownBits deriving BEq, ReflBEq, LawfulBEq, Hashable, Repr, DecidableEq /-- @@ -167,6 +170,7 @@ The fact specific data stored for each fact kind. | .sparseConstant => SparsePayload AbstractConstant (Option OpCode) | .integerRange => SparsePayload IntegerRangeLattice Unit | .modArithRange => SparsePayload IntegerRangeLattice Unit + | .knownBits => SparsePayload KnownBitsLattice /-- A dataflow fact stored by the framework. diff --git a/Veir/Analysis/DataFlow/KnownBitsAnalysis.lean b/Veir/Analysis/DataFlow/KnownBitsAnalysis.lean new file mode 100644 index 0000000000..a7c5a43daa --- /dev/null +++ b/Veir/Analysis/DataFlow/KnownBitsAnalysis.lean @@ -0,0 +1,746 @@ +module + +public import Veir.Analysis.DataFlow.Domains.KnownBitsDomain +public import Veir.Analysis.DataFlow.SparseForwardDataFlowAnalysis + +import Veir.Interfaces.FoldInterfaces + +public section + +namespace Veir + +/-! +# Known bits analysis + +This sparse forward analysis tracks fixed width integer bits that are provably zero +or one. Its transfer functions follow LLVM's `KnownBits` algorithms for arithmetic, +bitwise operations, shifts, casts, comparisons, division, remainder, and structural +bit operations across the Arith, Comb, and LLVM dialects. +-/ + +namespace KnownBitsAnalysis + +instance : SparseFactSpec .knownBits KnownBitsLattice where + Metadata := Unit + payloadEq := rfl + +/-- Return the width of the first integer result, if one exists. -/ +private def resultWidth? (resultTypes : Array TypeAttr) : Option Nat := do + let resultType ← resultTypes[0]? + let .integerType intType := resultType.val | none + return intType.bitwidth + +/-- Lift a unary known bits operation to the sparse lattice, preserving bottom. -/ +private def liftUnary + (operation : KnownBits → KnownBits) : KnownBitsLattice → KnownBitsLattice + | .bottom => .bottom + | .known operand => .known (operation operand) + +/-- +Lift a partial binary known bits operation to the sparse lattice. A width mismatch +produces an unknown value at the left operand's width. +-/ +private def liftBinary + (operation : KnownBits → KnownBits → Option KnownBits) : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice + | .bottom, _ | _, .bottom => .bottom + | .known lhs, .known rhs => + (operation lhs rhs).map (.known ·) |>.getD (.unknown lhs.bitwidth) + +/-- Lift a total binary known bits operation to the sparse lattice, preserving bottom. -/ +private def liftBinaryTotal + (operation : KnownBits → KnownBits → KnownBits) : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice + | .bottom, _ | _, .bottom => .bottom + | .known lhs, .known rhs => .known (operation lhs rhs) + +/-- Apply a unary lattice operation, returning `fallback` when the operand count is not one. -/ +private def applyUnary + (fallback operands : Array KnownBitsLattice) + (operation : KnownBitsLattice → KnownBitsLattice) : Array KnownBitsLattice := + match operands.toList with + | [operand] => #[operation operand] + | _ => fallback + +/-- Apply a binary lattice operation, returning `fallback` when the operand count is not two. -/ +private def applyBinary + (fallback operands : Array KnownBitsLattice) + (operation : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice) : + Array KnownBitsLattice := + match operands.toList with + | [lhs, rhs] => #[operation lhs rhs] + | _ => fallback + +/-- Apply a ternary lattice operation, returning `fallback` when the operand count is not three. -/ +private def applyTernary + (fallback operands : Array KnownBitsLattice) + (operation : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice → KnownBitsLattice) : + Array KnownBitsLattice := + match operands.toList with + | [first, second, third] => #[operation first second third] + | _ => fallback + +/-- Fold a binary lattice operation over the operands, returning `fallback` when they are empty. -/ +private def applyVariadic + (fallback operands : Array KnownBitsLattice) + (operation : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice) : + Array KnownBitsLattice := + match operands.toList with + | [] => fallback + | first :: rest => #[rest.foldl operation first] + +/-- +Apply a binary lattice operation that produces two results, returning `fallback` +when the operand count is not two. +-/ +private def applyBinaryPair + (fallback operands : Array KnownBitsLattice) + (operation : KnownBitsLattice → KnownBitsLattice → + KnownBitsLattice × KnownBitsLattice) : Array KnownBitsLattice := + match operands.toList with + | [lhs, rhs] => + let (first, second) := operation lhs rhs + #[first, second] + | _ => fallback + +/-- Whether a binary operation uses the same SSA value for both operands. -/ +private def hasSameBinaryOperands (op : OperationPtr) (irCtx : WfIRContext OpCode) : Bool := + match (op.getOperands! irCtx.raw).toList with + | [lhs, rhs] => lhs == rhs + | _ => false + +/-- Determine the known overflow bit for unsigned addition. -/ +private def unsignedAddOverflow? (lhs rhs : KnownBits) : Option KnownBits := + if lhs.bitwidth ≠ rhs.bitwidth then none else + let limit := 2 ^ lhs.bitwidth + let minSum := lhs.unsignedMin.toNat + rhs.unsignedMin.toNat + let maxSum := lhs.unsignedMax.toNat + rhs.unsignedMax.toNat + if limit ≤ minSum then some (KnownBits.constant 1 1) + else if maxSum < limit then some (KnownBits.constant 1 0) + else some (KnownBits.unknown 1) + +/-- Convert a fully known lattice value into a concrete runtime value for folding. -/ +private def exactRuntimeValue? : KnownBitsLattice → Option RuntimeValue + | .known bits => + if bits.isConstant then + some (.int bits.bitwidth (.val bits.one)) + else + none + | .bottom => none + +/-- Convert one generic fold result back into the known bits lattice. -/ +private def knownBitsOfFoldResult + (operands : Array KnownBitsLattice) + (resultType : TypeAttr) + (result : FoldDecision) : KnownBitsLattice := + match resultType.val, result with + | .integerType intType, .useOperand index => + operands[index]?.getD (.unknown intType.bitwidth) + | .integerType intType, .useConstant (.int bitwidth (.val value)) => + if h : bitwidth = intType.bitwidth then + let value := value.cast h + .known (KnownBits.ofMasks (~~~value) value) + else + .unknown intType.bitwidth + | .integerType intType, .useConstant _ => .unknown intType.bitwidth + | _, _ => ⊥ + +/-- Try the operation's generic fold hook, providing concrete values for exact operands. -/ +private def foldOperation? + (op : OperationPtr) + (operands : Array KnownBitsLattice) + (resultTypes : Array TypeAttr) + (irCtx : WfIRContext OpCode) : Option (Array KnownBitsLattice) := do + let opType := op.getOpType! irCtx.raw + let exactOperands := operands.map exactRuntimeValue? + let results ← opType.foldsTo + (op.getProperties! irCtx.raw opType) resultTypes exactOperands + return (results.zip resultTypes).map fun (result, resultType) => + knownBitsOfFoldResult operands resultType result + +/-- Zero extend known bits, first refining the sign bit when `nneg` is set. -/ +private def zeroExtend (resultWidth : Nat) (nneg : Bool) (bits : KnownBits) : KnownBits := + let bits := if nneg then + (bits.refineWith (KnownBits.ofMasks (KnownBits.highMask bits.bitwidth 1) 0)).getD bits + else bits + bits.zext resultWidth + +/-- Select one lattice value for a constant condition, or intersect both possible values. -/ +private def selectLattice + (condition trueValue falseValue : KnownBitsLattice) : KnownBitsLattice := + match condition, trueValue, falseValue with + | .bottom, _, _ | _, .bottom, _ | _, _, .bottom => ⊥ + | .known condition, .known trueValue, .known falseValue => + if condition.isConstant then + .known (if condition.one ≠ 0 then trueValue else falseValue) + else + .known <| (trueValue.intersect falseValue).getD (KnownBits.unknown trueValue.bitwidth) + +/-- Apply a funnel shift when its shift amount is exactly known. -/ +private def funnelShiftLattice + (operation : KnownBits → KnownBits → Nat → Option KnownBits) + (lhs rhs amount : KnownBitsLattice) : KnownBitsLattice := + match lhs, rhs, amount with + | .bottom, _, _ | _, .bottom, _ | _, _, .bottom => ⊥ + | .known lhs, .known rhs, .known amount => + if amount.isConstant then + (operation lhs rhs amount.one.toNat).map (.known ·) |>.getD (.unknown lhs.bitwidth) + else + .unknown lhs.bitwidth + +private def transferArithConstant (bitwidth : Nat) (value : Int) : KnownBitsLattice := + .constant bitwidth value + +private def transferLLVMConstant (bitwidth : Nat) (value : Int) : KnownBitsLattice := + .constant bitwidth value + +private def transferHWConstant (bitwidth : Nat) (value : Int) : KnownBitsLattice := + .constant bitwidth value + +private def transferArithAndI : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.bitwiseAnd? + +private def transferLLVMAnd : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.bitwiseAnd? + +private def transferCombAnd : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.bitwiseAnd? + +private def transferArithOrI : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.bitwiseOr? + +private def transferLLVMOr : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.bitwiseOr? + +private def transferCombOr : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.bitwiseOr? + +private def transferArithXorI : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.bitwiseXor? + +private def transferLLVMXor : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.bitwiseXor? + +private def transferCombXor : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.bitwiseXor? + +private def transferArithAddI (selfAdd nsw nuw : Bool) : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + fun lhs rhs => + if selfAdd then + match lhs, rhs with + | .bottom, _ | _, .bottom => ⊥ + | .known lhs, .known _ => + .known (lhs.shl (KnownBits.constant 8 1) nsw nuw) + else + liftBinary (fun lhs rhs => lhs.add? rhs nsw nuw) lhs rhs + +private def transferLLVMAdd (selfAdd nsw nuw : Bool) : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + fun lhs rhs => + if selfAdd then + match lhs, rhs with + | .bottom, _ | _, .bottom => ⊥ + | .known lhs, .known _ => + .known (lhs.shl (KnownBits.constant 8 1) nsw nuw) + else + liftBinary (fun lhs rhs => lhs.add? rhs nsw nuw) lhs rhs + +private def transferCombAdd : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.add? + +private def transferArithAddUIExtended + (lhs rhs : KnownBitsLattice) : KnownBitsLattice × KnownBitsLattice := + match lhs, rhs with + | .bottom, _ | _, .bottom => (⊥, ⊥) + | .known lhs, .known rhs => + let sum := (lhs.add? rhs).map (.known ·) |>.getD (.unknown lhs.bitwidth) + let overflow := unsignedAddOverflow? lhs rhs |>.map (.known ·) |>.getD (.unknown 1) + (sum, overflow) + +private def transferArithSubI (nsw nuw : Bool) : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary fun lhs rhs => lhs.sub? rhs nsw nuw + +private def transferLLVMSub (nsw nuw : Bool) : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary fun lhs rhs => lhs.sub? rhs nsw nuw + +private def transferCombSub : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.sub? + +private def transferArithSubUIExtended + (lhs rhs : KnownBitsLattice) : KnownBitsLattice × KnownBitsLattice := + match lhs, rhs with + | .bottom, _ | _, .bottom => (⊥, ⊥) + | .known lhs, .known rhs => + let difference := (lhs.sub? rhs).map (.known ·) |>.getD (.unknown lhs.bitwidth) + (difference, .known (lhs.compare .ult rhs)) + +private def transferArithMulI : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.mul? + +private def transferLLVMMul : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.mul? + +private def transferCombMul : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.mul? + +private def transferArithMulUIExtended + (lhs rhs : KnownBitsLattice) : KnownBitsLattice × KnownBitsLattice := + match lhs, rhs with + | .bottom, _ | _, .bottom => (⊥, ⊥) + | .known lhs, .known rhs => + let low := (lhs.mul? rhs).map (.known ·) |>.getD (.unknown lhs.bitwidth) + let high := (lhs.mulhu? rhs).map (.known ·) |>.getD (.unknown lhs.bitwidth) + (low, high) + +private def transferArithMulSIExtended + (lhs rhs : KnownBitsLattice) : KnownBitsLattice × KnownBitsLattice := + match lhs, rhs with + | .bottom, _ | _, .bottom => (⊥, ⊥) + | .known lhs, .known rhs => + let low := (lhs.mul? rhs).map (.known ·) |>.getD (.unknown lhs.bitwidth) + let high := (lhs.mulhs? rhs).map (.known ·) |>.getD (.unknown lhs.bitwidth) + (low, high) + +private def transferArithShLI (nsw nuw : Bool) : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinaryTotal fun lhs rhs => lhs.shl rhs nsw nuw + +private def transferLLVMShL (nsw nuw : Bool) : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinaryTotal fun lhs rhs => lhs.shl rhs nsw nuw + +private def transferCombShL : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinaryTotal KnownBits.shl + +private def transferArithShrUI (exact : Bool) : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinaryTotal fun lhs rhs => lhs.lshr rhs exact + +private def transferLLVMLShr (exact : Bool) : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinaryTotal fun lhs rhs => lhs.lshr rhs exact + +private def transferCombShrU : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinaryTotal KnownBits.lshr + +private def transferArithShrSI (exact : Bool) : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinaryTotal fun lhs rhs => lhs.ashr rhs exact + +private def transferLLVMAShr (exact : Bool) : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinaryTotal fun lhs rhs => lhs.ashr rhs exact + +private def transferCombShrS : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinaryTotal KnownBits.ashr + +private def transferArithDivUI (exact : Bool) : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary fun lhs rhs => lhs.udiv? rhs exact + +private def transferLLVMUDiv (exact : Bool) : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary fun lhs rhs => lhs.udiv? rhs exact + +private def transferCombDivU : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.udiv? + +private def transferArithDivSI (exact : Bool) : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary fun lhs rhs => lhs.sdiv? rhs exact + +private def transferLLVMSDiv (exact : Bool) : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary fun lhs rhs => lhs.sdiv? rhs exact + +private def transferCombDivS : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.sdiv? + +private def transferArithRemUI : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.urem? + +private def transferLLVMURem : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.urem? + +private def transferCombModU : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.urem? + +private def transferArithRemSI : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.srem? + +private def transferLLVMSRem : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.srem? + +private def transferCombModS : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.srem? + +private def transferArithExtUI (resultWidth : Nat) (nneg : Bool) : + KnownBitsLattice → KnownBitsLattice := + liftUnary (zeroExtend resultWidth nneg) + +private def transferLLVMZExt (resultWidth : Nat) (nneg : Bool) : + KnownBitsLattice → KnownBitsLattice := + liftUnary (zeroExtend resultWidth nneg) + +private def transferArithExtSI (resultWidth : Nat) : KnownBitsLattice → KnownBitsLattice := + liftUnary (KnownBits.sext · resultWidth) + +private def transferLLVMSExt (resultWidth : Nat) : KnownBitsLattice → KnownBitsLattice := + liftUnary (KnownBits.sext · resultWidth) + +private def transferArithTruncI (resultWidth : Nat) : KnownBitsLattice → KnownBitsLattice := + liftUnary (KnownBits.trunc · resultWidth) + +private def transferLLVMTrunc (resultWidth : Nat) : KnownBitsLattice → KnownBitsLattice := + liftUnary (KnownBits.trunc · resultWidth) + +private def transferArithCmpI (predicate : Data.LLVM.IntPred) : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinaryTotal fun lhs rhs => lhs.compare predicate rhs + +private def transferLLVMICmp (predicate : Data.LLVM.IntPred) : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinaryTotal fun lhs rhs => lhs.compare predicate rhs + +private def transferCombICmp (predicate : Data.LLVM.IntPred) : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinaryTotal fun lhs rhs => lhs.compare predicate rhs + +private def transferArithSelect : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + selectLattice + +private def transferLLVMSelect : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + selectLattice + +private def transferCombMux : + KnownBitsLattice → KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + selectLattice + +private def transferArithMaxUI : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.umax? + +private def transferLLVMUMax : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.umax? + +private def transferArithMinUI : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.umin? + +private def transferLLVMUMin : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.umin? + +private def transferArithMaxSI : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.smax? + +private def transferLLVMSMax : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.smax? + +private def transferArithMinSI : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.smin? + +private def transferLLVMSMin : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.smin? + +private def transferCombConcat : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinaryTotal KnownBits.concat + +private def transferCombExtract (lowBit resultWidth : Nat) : + KnownBitsLattice → KnownBitsLattice := + liftUnary fun bits => bits.extract lowBit resultWidth + +private def transferCombReverse : KnownBitsLattice → KnownBitsLattice := + liftUnary KnownBits.reverse + +private def transferLLVMBitReverse : KnownBitsLattice → KnownBitsLattice := + liftUnary KnownBits.reverse + +private def transferCombReplicate (resultWidth : Nat) : KnownBitsLattice → KnownBitsLattice := + liftUnary (KnownBits.replicate · resultWidth) + +private def transferLLVMByteSwap : KnownBitsLattice → KnownBitsLattice := + liftUnary KnownBits.byteSwap + +private def transferLLVMFShL + (lhs rhs amount : KnownBitsLattice) : KnownBitsLattice := + funnelShiftLattice KnownBits.fshl? lhs rhs amount + +private def transferLLVMFShR + (lhs rhs amount : KnownBitsLattice) : KnownBitsLattice := + funnelShiftLattice KnownBits.fshr? lhs rhs amount + +private def transferLLVMCountPopulation : KnownBitsLattice → KnownBitsLattice := + liftUnary KnownBits.ctpop + +private def transferLLVMCountLeadingZeros (isZeroPoison : Bool) : + KnownBitsLattice → KnownBitsLattice := + liftUnary (KnownBits.ctlz · isZeroPoison) + +private def transferLLVMCountTrailingZeros (isZeroPoison : Bool) : + KnownBitsLattice → KnownBitsLattice := + liftUnary (KnownBits.cttz · isZeroPoison) + +private def transferLLVMAbs (isIntMinPoison : Bool) : + KnownBitsLattice → KnownBitsLattice := + liftUnary (KnownBits.abs · isIntMinPoison) + +private def transferLLVMUAddSat : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.uaddSat? + +private def transferLLVMUSubSat : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.usubSat? + +private def transferLLVMSAddSat : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.saddSat? + +private def transferLLVMSSubSat : KnownBitsLattice → KnownBitsLattice → KnownBitsLattice := + liftBinary KnownBits.ssubSat? + +/-- +Infer known bits for one operation. Bottom operands cause the transfer to wait for +more information; unsupported integer results receive a width-aware unknown value. +-/ +def transfer + (op : OperationPtr) + (operands : Array KnownBitsLattice) + (irCtx : WfIRContext OpCode) : Array KnownBitsLattice := + let numResults := op.getNumResults! irCtx.raw + let resultTypes := op.getResultTypes! irCtx.raw + let pessimisticUpdates := resultTypes.map fun resultType => + match resultType.val with + | .integerType intType => .unknown intType.bitwidth + | _ => ⊥ + + if op.getNumRegions! irCtx.raw ≠ 0 then + pessimisticUpdates + else if operands.any (· = ⊥) then + Array.replicate numResults ⊥ + else + let opType := op.getOpType! irCtx.raw + let resultWidth := (resultWidth? resultTypes).getD 0 + match foldOperation? op operands resultTypes irCtx with + | some results => results + | none => + match opType with + | OpCode.arith Arith.constant => + let props := op.getProperties! irCtx.raw (OpCode.arith Arith.constant) + #[transferArithConstant props.value.type.bitwidth props.value.value] + | OpCode.llvm Llvm.mlir__constant => + let props := op.getProperties! irCtx.raw (OpCode.llvm Llvm.mlir__constant) + match props.value with + | .integer attr => #[transferLLVMConstant attr.type.bitwidth attr.value] + | _ => pessimisticUpdates + | OpCode.hw HW.constant => + let props := op.getProperties! irCtx.raw (OpCode.hw HW.constant) + #[transferHWConstant props.value.type.bitwidth props.value.value] + | OpCode.arith Arith.andi => + applyBinary pessimisticUpdates operands transferArithAndI + | OpCode.llvm Llvm.and => + applyBinary pessimisticUpdates operands transferLLVMAnd + | OpCode.comb Comb.and => + applyVariadic pessimisticUpdates operands transferCombAnd + | OpCode.arith Arith.ori => + applyBinary pessimisticUpdates operands transferArithOrI + | OpCode.llvm Llvm.or => + applyBinary pessimisticUpdates operands transferLLVMOr + | OpCode.comb Comb.or => + applyVariadic pessimisticUpdates operands transferCombOr + | OpCode.arith Arith.xori => + applyBinary pessimisticUpdates operands transferArithXorI + | OpCode.llvm Llvm.xor => + applyBinary pessimisticUpdates operands transferLLVMXor + | OpCode.comb Comb.xor => + applyVariadic pessimisticUpdates operands transferCombXor + | OpCode.arith Arith.addi => + let props := op.getProperties! irCtx.raw (OpCode.arith Arith.addi) + applyBinary pessimisticUpdates operands <| + transferArithAddI (hasSameBinaryOperands op irCtx) props.attr.nsw props.attr.nuw + | OpCode.llvm Llvm.add => + let props := op.getProperties! irCtx.raw (OpCode.llvm Llvm.add) + applyBinary pessimisticUpdates operands <| + transferLLVMAdd (hasSameBinaryOperands op irCtx) props.nsw props.nuw + | OpCode.comb Comb.add => + applyVariadic pessimisticUpdates operands transferCombAdd + | OpCode.arith Arith.addui_extended => + applyBinaryPair pessimisticUpdates operands transferArithAddUIExtended + | OpCode.arith Arith.subi => + let props := op.getProperties! irCtx.raw (OpCode.arith Arith.subi) + applyBinary pessimisticUpdates operands <| + transferArithSubI props.attr.nsw props.attr.nuw + | OpCode.llvm Llvm.sub => + let props := op.getProperties! irCtx.raw (OpCode.llvm Llvm.sub) + applyBinary pessimisticUpdates operands <| transferLLVMSub props.nsw props.nuw + | OpCode.comb Comb.sub => + applyBinary pessimisticUpdates operands transferCombSub + | OpCode.arith Arith.subui_extended => + applyBinaryPair pessimisticUpdates operands transferArithSubUIExtended + | OpCode.arith Arith.muli => + applyBinary pessimisticUpdates operands transferArithMulI + | OpCode.llvm Llvm.mul => + applyBinary pessimisticUpdates operands transferLLVMMul + | OpCode.comb Comb.mul => + applyVariadic pessimisticUpdates operands transferCombMul + | OpCode.arith Arith.mului_extended => + applyBinaryPair pessimisticUpdates operands transferArithMulUIExtended + | OpCode.arith Arith.mulsi_extended => + applyBinaryPair pessimisticUpdates operands transferArithMulSIExtended + | OpCode.arith Arith.shli => + let props := op.getProperties! irCtx.raw (OpCode.arith Arith.shli) + applyBinary pessimisticUpdates operands <| + transferArithShLI props.attr.nsw props.attr.nuw + | OpCode.llvm Llvm.shl => + let props := op.getProperties! irCtx.raw (OpCode.llvm Llvm.shl) + applyBinary pessimisticUpdates operands <| transferLLVMShL props.nsw props.nuw + | OpCode.comb Comb.shl => + applyBinary pessimisticUpdates operands transferCombShL + | OpCode.arith Arith.shrui => + let props := op.getProperties! irCtx.raw (OpCode.arith Arith.shrui) + applyBinary pessimisticUpdates operands <| transferArithShrUI props.exact + | OpCode.llvm Llvm.lshr => + let props := op.getProperties! irCtx.raw (OpCode.llvm Llvm.lshr) + applyBinary pessimisticUpdates operands <| transferLLVMLShr props.exact + | OpCode.comb Comb.shru => + applyBinary pessimisticUpdates operands transferCombShrU + | OpCode.arith Arith.shrsi => + let props := op.getProperties! irCtx.raw (OpCode.arith Arith.shrsi) + applyBinary pessimisticUpdates operands <| transferArithShrSI props.exact + | OpCode.llvm Llvm.ashr => + let props := op.getProperties! irCtx.raw (OpCode.llvm Llvm.ashr) + applyBinary pessimisticUpdates operands <| transferLLVMAShr props.exact + | OpCode.comb Comb.shrs => + applyBinary pessimisticUpdates operands transferCombShrS + | OpCode.arith Arith.divui => + let props := op.getProperties! irCtx.raw (OpCode.arith Arith.divui) + applyBinary pessimisticUpdates operands <| transferArithDivUI props.exact + | OpCode.llvm Llvm.udiv => + let props := op.getProperties! irCtx.raw (OpCode.llvm Llvm.udiv) + applyBinary pessimisticUpdates operands <| transferLLVMUDiv props.exact + | OpCode.comb Comb.divu => + applyBinary pessimisticUpdates operands transferCombDivU + | OpCode.arith Arith.divsi => + let props := op.getProperties! irCtx.raw (OpCode.arith Arith.divsi) + applyBinary pessimisticUpdates operands <| transferArithDivSI props.exact + | OpCode.llvm Llvm.sdiv => + let props := op.getProperties! irCtx.raw (OpCode.llvm Llvm.sdiv) + applyBinary pessimisticUpdates operands <| transferLLVMSDiv props.exact + | OpCode.comb Comb.divs => + applyBinary pessimisticUpdates operands transferCombDivS + | OpCode.arith Arith.remui => + applyBinary pessimisticUpdates operands transferArithRemUI + | OpCode.llvm Llvm.urem => + applyBinary pessimisticUpdates operands transferLLVMURem + | OpCode.comb Comb.modu => + applyBinary pessimisticUpdates operands transferCombModU + | OpCode.arith Arith.remsi => + applyBinary pessimisticUpdates operands transferArithRemSI + | OpCode.llvm Llvm.srem => + applyBinary pessimisticUpdates operands transferLLVMSRem + | OpCode.comb Comb.mods => + applyBinary pessimisticUpdates operands transferCombModS + | OpCode.arith Arith.extui => + let props := op.getProperties! irCtx.raw (OpCode.arith Arith.extui) + applyUnary pessimisticUpdates operands <| transferArithExtUI resultWidth props.nneg + | OpCode.llvm Llvm.zext => + let props := op.getProperties! irCtx.raw (OpCode.llvm Llvm.zext) + applyUnary pessimisticUpdates operands <| transferLLVMZExt resultWidth props.nneg + | OpCode.arith Arith.extsi => + applyUnary pessimisticUpdates operands <| transferArithExtSI resultWidth + | OpCode.llvm Llvm.sext => + applyUnary pessimisticUpdates operands <| transferLLVMSExt resultWidth + | OpCode.arith Arith.trunci => + applyUnary pessimisticUpdates operands <| transferArithTruncI resultWidth + | OpCode.llvm Llvm.trunc => + applyUnary pessimisticUpdates operands <| transferLLVMTrunc resultWidth + | OpCode.arith Arith.cmpi => + let props := op.getProperties! irCtx.raw (OpCode.arith Arith.cmpi) + applyBinary pessimisticUpdates operands <| transferArithCmpI props.predicate + | OpCode.llvm Llvm.icmp => + let props := op.getProperties! irCtx.raw (OpCode.llvm Llvm.icmp) + applyBinary pessimisticUpdates operands <| transferLLVMICmp props.predicate + | OpCode.comb Comb.icmp => + let props := op.getProperties! irCtx.raw (OpCode.comb Comb.icmp) + match Data.LLVM.IntPred.fromNat props.predicate.value.toNat with + | some predicate => + applyBinary pessimisticUpdates operands <| transferCombICmp predicate + | none => pessimisticUpdates + | OpCode.arith Arith.select => + applyTernary pessimisticUpdates operands transferArithSelect + | OpCode.llvm Llvm.select => + applyTernary pessimisticUpdates operands transferLLVMSelect + | OpCode.comb Comb.mux => + applyTernary pessimisticUpdates operands transferCombMux + | OpCode.arith Arith.maxui => + applyBinary pessimisticUpdates operands transferArithMaxUI + | OpCode.llvm Llvm.intr__umax => + applyBinary pessimisticUpdates operands transferLLVMUMax + | OpCode.arith Arith.minui => + applyBinary pessimisticUpdates operands transferArithMinUI + | OpCode.llvm Llvm.intr__umin => + applyBinary pessimisticUpdates operands transferLLVMUMin + | OpCode.arith Arith.maxsi => + applyBinary pessimisticUpdates operands transferArithMaxSI + | OpCode.llvm Llvm.intr__smax => + applyBinary pessimisticUpdates operands transferLLVMSMax + | OpCode.arith Arith.minsi => + applyBinary pessimisticUpdates operands transferArithMinSI + | OpCode.llvm Llvm.intr__smin => + applyBinary pessimisticUpdates operands transferLLVMSMin + | OpCode.comb Comb.concat => + applyVariadic pessimisticUpdates operands transferCombConcat + | OpCode.comb Comb.extract => + let props := op.getProperties! irCtx.raw (OpCode.comb Comb.extract) + applyUnary pessimisticUpdates operands <| + transferCombExtract props.lowBit.value.toNat resultWidth + | OpCode.comb Comb.reverse => + applyUnary pessimisticUpdates operands transferCombReverse + | OpCode.llvm Llvm.intr__bitreverse => + applyUnary pessimisticUpdates operands transferLLVMBitReverse + | OpCode.comb Comb.replicate => + applyUnary pessimisticUpdates operands <| transferCombReplicate resultWidth + | OpCode.llvm Llvm.intr__bswap => + applyUnary pessimisticUpdates operands transferLLVMByteSwap + | OpCode.llvm Llvm.intr__fshl => + applyTernary pessimisticUpdates operands transferLLVMFShL + | OpCode.llvm Llvm.intr__fshr => + applyTernary pessimisticUpdates operands transferLLVMFShR + | OpCode.llvm Llvm.intr__ctpop => + applyUnary pessimisticUpdates operands transferLLVMCountPopulation + | OpCode.llvm Llvm.intr__ctlz => + let props := op.getProperties! irCtx.raw (OpCode.llvm Llvm.intr__ctlz) + applyUnary pessimisticUpdates operands <| + transferLLVMCountLeadingZeros props.is_zero_poison + | OpCode.llvm Llvm.intr__cttz => + let props := op.getProperties! irCtx.raw (OpCode.llvm Llvm.intr__cttz) + applyUnary pessimisticUpdates operands <| + transferLLVMCountTrailingZeros props.is_zero_poison + | OpCode.llvm Llvm.intr__abs => + let props := op.getProperties! irCtx.raw (OpCode.llvm Llvm.intr__abs) + applyUnary pessimisticUpdates operands <| + transferLLVMAbs props.is_int_min_poison + | OpCode.llvm Llvm.intr__uadd__sat => + applyBinary pessimisticUpdates operands transferLLVMUAddSat + | OpCode.llvm Llvm.intr__usub__sat => + applyBinary pessimisticUpdates operands transferLLVMUSubSat + | OpCode.llvm Llvm.intr__sadd__sat => + applyBinary pessimisticUpdates operands transferLLVMSAddSat + | OpCode.llvm Llvm.intr__ssub__sat => + applyBinary pessimisticUpdates operands transferLLVMSSubSat + | _ => pessimisticUpdates + +end KnownBitsAnalysis + +/-- Sparse forward known bits analysis for fixed-width integer SSA values. -/ +def KnownBitsAnalysis : DataFlowAnalysis := + SparseForwardDataFlowAnalysis.new + .knownBits + .knownBits + KnownBitsAnalysis.transfer + (entryState := fun value irCtx => + match (value.getType! irCtx.raw).val with + | .integerType intType => .unknown intType.bitwidth + | _ => ⊥) + +end Veir diff --git a/Veir/Analysis/DataFlow/SparseConstantPropagationAnalysis.lean b/Veir/Analysis/DataFlow/SparseConstantPropagationAnalysis.lean index 464288b9e7..5b56f9cd78 100644 --- a/Veir/Analysis/DataFlow/SparseConstantPropagationAnalysis.lean +++ b/Veir/Analysis/DataFlow/SparseConstantPropagationAnalysis.lean @@ -71,7 +71,8 @@ def SparseConstantPropagationAnalysis : DataFlowAnalysis := { SparseForwardDataFlowAnalysis.new .sparseConstant SparseConstantPropagation.kind - SparseConstantPropagation.transfer with + SparseConstantPropagation.transfer + (entryState := fun _ _ => ⊤) with printer? := some printer } end Veir diff --git a/Veir/Analysis/DataFlow/SparseForwardDataFlowAnalysis.lean b/Veir/Analysis/DataFlow/SparseForwardDataFlowAnalysis.lean index 3687dcb984..9017a6da74 100644 --- a/Veir/Analysis/DataFlow/SparseForwardDataFlowAnalysis.lean +++ b/Veir/Analysis/DataFlow/SparseForwardDataFlowAnalysis.lean @@ -360,18 +360,14 @@ private def visit Build a sparse forward analysis over one abstract value domain. Sparse facts default to `⊥`. Whenever control flow loses precision, the framework -conservatively joins the entry state into the affected values. The entry state defaults -to `⊤`; analyses only need to override it when they have a more precise analysis-specific -state. +conservatively joins the analysis specific entry state into the affected values. -/ def new (kind : FactKind) [SparseFactSpec kind Domain] - [Top Domain] (analysisKind : AnalysisKind) (transfer : TransferFn Domain) - (entryState : EntryStateFn Domain := fun _ _ => ⊤) - : DataFlowAnalysis := + (entryState : EntryStateFn Domain) : DataFlowAnalysis := { kind := analysisKind init := init kind analysisKind entryState transfer visit := visit kind analysisKind entryState transfer }