diff --git a/src/Lean/Elab/Do.lean b/src/Lean/Elab/Do.lean index 68cc516707..81f0996c81 100644 --- a/src/Lean/Elab/Do.lean +++ b/src/Lean/Elab/Do.lean @@ -935,15 +935,18 @@ partial def doSeqToCode : List Syntax → M CodeBlock forCodeBlock ← withNewVars newVars $ withFor (doSeqToCode forElems); ⟨isForInMap, uvars, forInBody⟩ ← toForInTerm x forCodeBlock; uvarsTuple ← liftMacroM $ mkTuple ref (uvars.map (mkIdentFrom ref)); - auxDo ← if isForInMap then do + if isForInMap then do forInTerm ← `($(xs).forInMap $uvarsTuple fun $x $uvarsTuple => $forInBody); - `(do let r ← $forInTerm; $uvarsTuple:term := r.2; return ensureExpectedType! "type mismatch, 'for'" r.1) - else do { + auxDo ← `(do let r ← $forInTerm; $uvarsTuple:term := r.2; return ensureExpectedType! "type mismatch, 'for'" r.1); + doSeqToCode (getDoSeqElems (getDoSeq auxDo) ++ doElems) + else do forInTerm ← `($(xs).forIn $uvarsTuple fun $x $uvarsTuple => $forInBody); - `(do let r ← $forInTerm; $uvarsTuple:term := r) - }; - let doElemsNew := getDoSeqElems (getDoSeq auxDo); - doSeqToCode (doElemsNew ++ doElems) + if doElems.isEmpty then do + auxDo ← `(do let r ← $forInTerm; $uvarsTuple:term := r; return ensureExpectedType! "type mismatch, 'for'" PUnit.unit); + doSeqToCode (getDoSeqElems (getDoSeq auxDo) ++ doElems) + else do + auxDo ← `(do let r ← $forInTerm; $uvarsTuple:term := r); + doSeqToCode (getDoSeqElems (getDoSeq auxDo) ++ doElems) else if k == `Lean.Parser.Term.doMatch then throwError "WIP" else if k == `Lean.Parser.Term.doTry then diff --git a/tests/lean/doNotation1.lean b/tests/lean/doNotation1.lean index 88fe03cb10..f7fe0b32d5 100644 --- a/tests/lean/doNotation1.lean +++ b/tests/lean/doNotation1.lean @@ -23,3 +23,7 @@ Vector.cons 1 v def f5 (xs : List Nat) : List Bool := do for x in xs do x := true -- invalid reassigned + +def f6 (xs : List Nat) : IO (List Nat) := do +for x in xs do + IO.println x diff --git a/tests/lean/doNotation1.lean.expected.out b/tests/lean/doNotation1.lean.expected.out index 812eab849f..1c85dc0525 100644 --- a/tests/lean/doNotation1.lean.expected.out +++ b/tests/lean/doNotation1.lean.expected.out @@ -12,3 +12,7 @@ doNotation1.lean:24:0: error: type mismatch, 'for' has type List Nat but is expected to have type List Bool +doNotation1.lean:28:0: error: type mismatch, 'for' has type + PUnit +but is expected to have type + List Nat