import Spa.Language.Base import Spa.Language.Program import Spa.Lattice.FiniteMap import Spa.Analysis.Constant import Spa.Analysis.Forward /-! # Constant folding Rewrites each assignment's right-hand side to a literal wherever the constant analysis (`Spa/Analysis/Constant.lean`) pins its value down. The traversal recurses over the plain `Stmt`, threading a `GGraph.Embed` of the current subtree's CFG into the program's (`Program.rootEmbed`, then one `Embed.trans` per descent). At an assignment, `Embed.singletonIndex` gives its CFG state, and the facts to fold with are `Forward.joinForKey` at that state — the join over predecessors, i.e. the values *entering* the node, which is what the right-hand side reads. (`Forward.variablesAt` would be the values *leaving* it, which already include this assignment's own effect.) -/ namespace Spa namespace ConstantTransform open GGraph Forward variable (prog : Program) abbrev result := Forward.result ConstLattice prog /-- Replace an expression by a literal when the analysis pins its value down, recursing into its subexpressions otherwise. Whole subexpressions are tried first, so `(x + 1) - x` folds outright when `x` is known, rather than only in its leaves. -/ def foldExpr (vs : VariableValues ConstLattice prog) : Expr → Expr | .num n => .num n | .var k => match ConstAnalysis.eval prog (.var k) vs with | .mk z => .num z | _ => .var k | .add a b => match ConstAnalysis.eval prog (.add a b) vs with | .mk z => .num z | _ => .add (foldExpr vs a) (foldExpr vs b) | .sub a b => match ConstAnalysis.eval prog (.sub a b) vs with | .mk z => .num z | _ => .sub (foldExpr vs a) (foldExpr vs b) /-- Constant-fold every assignment in a statement. `sv` is the analysis result, taken as a parameter rather than read from `result` at each node: it is a fixpoint computation, so recomputing it per assignment would make folding quadratic in the analysis. Guards of `ifElse`/`whileLoop` are deliberately left alone. `Stmt.cfg` gives a conditional's guard no node at all (`ifElse` overlays the two branches), so there is no state whose entry facts describe where it is evaluated. A `whileLoop`'s guard does have a candidate — the loop header `GGraph.loopIn` — but tying the guard's evaluation environment to that node needs a lemma that does not exist yet, so folding it here would be an unproven soundness claim. -/ def foldStmt (sv : StateVariables ConstLattice prog) : (s : Stmt) → Embed s.cfg prog.cfg → Stmt | .basic .noop, _ => .basic .noop | .basic (.assign k v), e => .basic (.assign k (foldExpr prog (joinForKey e.singletonIndex sv) v)) | .andThen s₁ s₂, e => .andThen (foldStmt sv s₁ ((Embed.sequenceLeft s₁.cfg s₂.cfg).trans e)) (foldStmt sv s₂ ((Embed.sequenceRight s₁.cfg s₂.cfg).trans e)) | .ifElse cond s₁ s₂, e => .ifElse cond (foldStmt sv s₁ ((Embed.overlayLeft s₁.cfg s₂.cfg).trans e)) (foldStmt sv s₂ ((Embed.overlayRight s₁.cfg s₂.cfg).trans e)) | .whileLoop cond body, e => .whileLoop cond (foldStmt sv body ((Embed.loop body.cfg).trans e)) /-- Constant-fold a whole program, running the analysis once. -/ def foldProgram : Stmt := foldStmt prog (result prog) prog.rootStmt prog.rootEmbed /-! ## Correctness Folding preserves meaning *provided the facts folded with actually hold of the environment folded in*. That proviso is the whole content: `foldExpr` is sound against any `vs` that over-approximates `ρ`, and it is the analysis engine's job (`Forward.analyze_correct_at`) to supply such a `vs` at each program point. -/ variable {prog} /-- If the analysis pins an expression to a constant and its facts hold of `ρ`, then the expression really does evaluate to that constant. This is `ValidExprEvaluator` specialised to the `.mk` case, where `interpConst` says exactly `v = .int z`. -/ lemma eq_int_of_eval_mk {vs : VariableValues ConstLattice prog} {ρ : Env} {e : Expr} {v : Value} {z : ℤ} (hev : EvalExpr ρ e v) (hvs : ⟦vs⟧ ρ) (hz : ConstAnalysis.eval prog e vs = .mk z) : v = .int z := by have h := ValidExprEvaluator.valid (L := ConstLattice) (prog := prog) hev hvs rw [show ExprEvaluator.eval e vs = ConstAnalysis.eval prog e vs from rfl, hz] at h exact h /-- **Expression folding is meaning-preserving.** Whenever `vs` over-approximates `ρ`, the folded expression evaluates in `ρ` to whatever the original did. -/ theorem foldExpr_eval {vs : VariableValues ConstLattice prog} {ρ : Env} (hvs : ⟦vs⟧ ρ) : ∀ {e : Expr} {v : Value}, EvalExpr ρ e v → EvalExpr ρ (foldExpr prog vs e) v := by intro e induction e with | num n => intro v hev; simpa [foldExpr] using hev | var k => intro v hev simp only [foldExpr] split · case h_1 z hz => rw [eq_int_of_eval_mk hev hvs hz]; exact EvalExpr.num ρ z · exact hev | add a b iha ihb => intro v hev simp only [foldExpr] split · case h_1 z hz => rw [eq_int_of_eval_mk hev hvs hz]; exact EvalExpr.num ρ z · cases hev with | add _ _ z₁ z₂ h₁ h₂ => exact EvalExpr.add ρ _ _ z₁ z₂ (iha h₁) (ihb h₂) | sub a b iha ihb => intro v hev simp only [foldExpr] split · case h_1 z hz => rw [eq_int_of_eval_mk hev hvs hz]; exact EvalExpr.num ρ z · cases hev with | sub _ _ z₁ z₂ h₁ h₂ => exact EvalExpr.sub ρ _ _ z₁ z₂ (iha h₁) (ihb h₂) end ConstantTransform end Spa