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₂) /-- Fold a source evaluation using its actual whole-program execution prefix. -/ noncomputable def foldStmt_eval (prog : Program) {s : Stmt} {ρ₀ ρ₁ : Env} (h : EvalStmt ρ₀ s ρ₁) : (e : Embed s.cfg prog.cfg) → (pre : Traceₗ prog.cfg prog.initialState (e.f (Stmt.cfg_sufficient h).entry) [] ρ₀) → EvalStmt ρ₀ (foldStmt prog (result prog) s e) ρ₁ := by induction h with | basic ρ₀ ρ₁ bs hbs => intro e pre have hr : Reaches e.singletonIndex ρ₀ ρ₁ := ⟨pre, by rw [Program.code, e.nodes_singletonIndex]; exact .some hbs⟩ cases hbs with | noop => exact .basic _ _ _ (.noop _) | assign x expr v hev => exact .basic _ _ _ (.assign _ _ _ _ (foldExpr_eval (ConstAnalysis.analyze_correct_at prog hr).1 hev)) | andThen ρ₀ ρ₁ ρ₂ s₁ s₂ h₁ h₂ ih₁ ih₂ => intro e pre exact .andThen _ _ _ _ _ (ih₁ ((Embed.sequenceLeft s₁.cfg s₂.cfg).trans e) pre) (ih₂ ((Embed.sequenceRight s₁.cfg s₂.cfg).trans e) (Path.append pre (Path.embed e ((Stmt.cfg_sufficient h₁).beforeRight (Stmt.cfg_sufficient h₂))))) | ifTrue ρ₀ ρ₁ cond z s₁ s₂ hc hz h ih => intro e pre exact .ifTrue _ _ _ _ _ _ hc hz (ih ((Embed.overlayLeft s₁.cfg s₂.cfg).trans e) pre) | ifFalse ρ₀ ρ₁ cond s₁ s₂ hc h ih => intro e pre exact .ifFalse _ _ _ _ _ hc (ih ((Embed.overlayRight s₁.cfg s₂.cfg).trans e) pre) | whileTrue ρ₀ ρ₁ ρ₂ cond z body hc hz hb hr ihb ihr => intro e pre exact .whileTrue _ _ _ _ _ _ hc hz (ihb ((Embed.loop body.cfg).trans e) (Path.append pre (Path.embed e (Stmt.cfg_sufficient hb).beforeBody))) (ihr e (Path.append pre (Path.embed e ((Stmt.cfg_sufficient hb).loop.beforeRest (Stmt.cfg_sufficient hr))))) | whileFalse ρ cond body hc => intro e pre exact .whileFalse _ _ _ hc /-- Constant folding preserves every terminating source evaluation. -/ noncomputable def foldProgram_eval (prog : Program) {ρ : Env} (h : EvalStmt [] prog.rootStmt ρ) : EvalStmt [] (foldProgram prog) ρ := foldStmt_eval prog h prog.rootEmbed (Stmt.cfg_sufficient h).beforeRoot end ConstantTransform end Spa