|
1 | 1 | module |
2 | 2 |
|
3 | 3 | public import Test.Common |
| 4 | +meta import Test.Common |
| 5 | +public import Std.Tactic.Do |
4 | 6 |
|
5 | 7 | set_option linter.style.header false |
| 8 | +set_option linter.hashCommand false |
6 | 9 |
|
7 | 10 | /-! |
8 | | -# `rdo`: `for` loops over a single collection |
| 11 | +# `rdo`: `for` and `while` loops |
9 | 12 |
|
10 | 13 | `rdo` has its own `for … rdo …` parser, expander and elaborator, mirroring core's but emitting |
11 | | -`MeasurableSpaceForIn.forIn`. Instances exist for `List`, `Array` and `Vector`. |
| 14 | +`MeasurableSpaceForIn.forIn`. Instances exist for `List`, `Array` and `Vector`, and for `Lean.Loop`, |
| 15 | +which `while … rdo` loops over, at the core monads. |
12 | 16 |
|
13 | 17 | A loop can sit under another construct, including another loop: the enclosing one learns what the |
14 | 18 | loop does to the control flow from the `ControlInfo` handler of `rdoFor`, which is that of core's |
@@ -235,6 +239,100 @@ noncomputable def countPairsOfHeads (n : ℕ) : Measure ℕ := rdo |
235 | 239 | c := c + 1 |
236 | 240 | return c |
237 | 241 |
|
| 242 | +/-! ## `while` loops |
| 243 | +
|
| 244 | +`while c rdo body` is a loop over `Lean.Loop`, as in core. At a core monad it is core's loop, which |
| 245 | +the kernel cannot unfold, so these programs are checked with `#guard` rather than `rfl`, and proved |
| 246 | +through `mvcgen`. There is no instance at `Measure` yet. |
| 247 | +-/ |
| 248 | + |
| 249 | +/-- A `while` loop, counting down from `n`. -/ |
| 250 | +def countdown (n : ℕ) : IdM ℕ := rdo |
| 251 | + let mut i := n |
| 252 | + let mut steps := 0 |
| 253 | + while 0 < i rdo |
| 254 | + i := i - 1 |
| 255 | + steps := steps + 1 |
| 256 | + return steps |
| 257 | + |
| 258 | +#guard IdM.run (countdown 5) = 5 |
| 259 | + |
| 260 | +#guard IdM.run (countdown 0) = 0 |
| 261 | + |
| 262 | +open Std.Do in |
| 263 | +set_option mvcgen.warning false in |
| 264 | +theorem countdown_eq (n : ℕ) : IdM.run (countdown n) = n := by |
| 265 | + generalize h : IdM.run (countdown n) = r |
| 266 | + apply Id.of_wp_run_eq h |
| 267 | + simp only [countdown, MeasurableSpaceForIn.forIn, MeasurableSpaceBind.mBind, |
| 268 | + MeasurableSpacePure.mPure] |
| 269 | + dsimp only [IdM, Monad.toMeasurableSpaceMonad] |
| 270 | + mvcgen invariants |
| 271 | + · fun st => ⟨st.1⟩ |
| 272 | + · ⇓ c => match c with |
| 273 | + | .inl st => ⌜st.1 + st.2 = n⌝ |
| 274 | + | .inr st => ⌜st.2 = n⌝ |
| 275 | + all_goals simp_all <;> omega |
| 276 | + |
| 277 | +/-- `break` out of a `while` loop. -/ |
| 278 | +def halveUntilOdd (n : ℕ) : IdM ℕ := rdo |
| 279 | + let mut k := n |
| 280 | + while 0 < k rdo |
| 281 | + if k % 2 = 1 then |
| 282 | + break |
| 283 | + k := k / 2 |
| 284 | + return k |
| 285 | + |
| 286 | +#guard IdM.run (halveUntilOdd 24) = 3 |
| 287 | + |
| 288 | +#guard IdM.run (halveUntilOdd 0) = 0 |
| 289 | + |
| 290 | +/-- An early `return` out of a `while` loop. -/ |
| 291 | +def firstSquareAbove (n : ℕ) : IdM ℕ := rdo |
| 292 | + let mut k := 0 |
| 293 | + while true rdo |
| 294 | + if k * k > n then |
| 295 | + return k |
| 296 | + k := k + 1 |
| 297 | + return 0 |
| 298 | + |
| 299 | +#guard IdM.run (firstSquareAbove 10) = 4 |
| 300 | + |
| 301 | +/-- `while let`, consuming a list one element at a time. -/ |
| 302 | +def sumByPopping (xs : List ℕ) : IdM ℕ := rdo |
| 303 | + let mut rest := xs |
| 304 | + let mut s := 0 |
| 305 | + while let x :: xs' := rest rdo |
| 306 | + s := s + x |
| 307 | + rest := xs' |
| 308 | + return s |
| 309 | + |
| 310 | +#guard IdM.run (sumByPopping [1, 2, 3]) = 6 |
| 311 | + |
| 312 | +/-- `while h : c`, which hands the body a proof of the condition. -/ |
| 313 | +def countdownWithProof (n : ℕ) : IdM ℕ := rdo |
| 314 | + let mut i := n |
| 315 | + let mut steps := 0 |
| 316 | + while h : 0 < i rdo |
| 317 | + have : i - 1 < i := Nat.sub_lt h Nat.one_pos |
| 318 | + i := i - 1 |
| 319 | + steps := steps + 1 |
| 320 | + return steps |
| 321 | + |
| 322 | +#guard IdM.run (countdownWithProof 4) = 4 |
| 323 | + |
| 324 | +/-- A `while` loop nested inside a `for` loop. -/ |
| 325 | +def sumOfLogs (xs : List ℕ) : IdM ℕ := rdo |
| 326 | + let mut s := 0 |
| 327 | + for x in xs rdo |
| 328 | + let mut k := x |
| 329 | + while 1 < k rdo |
| 330 | + k := k / 2 |
| 331 | + s := s + 1 |
| 332 | + return s |
| 333 | + |
| 334 | +#guard IdM.run (sumOfLogs [1, 2, 8]) = 4 |
| 335 | + |
238 | 336 | end Test.Loops |
239 | 337 |
|
240 | 338 | end |
0 commit comments