如何用Cats的State Monad实现栈安全的递归while循环?
问题背景
我想用Cats库的State Monad实现一个while循环函数,最初的实现代码如下:
def whileLoopState[S](cond: S => Boolean)(block: S => S): State[S, Unit] = State { state => if (cond(state)) { val nextState = block(state) whileLoopState(cond)(block).run(nextState).value } else { (state, ()) } }
这个实现存在栈不安全的问题,因为递归调用不在尾位置,执行下面的代码会触发栈溢出:
whileLoopState[Int](s => s > 0) { s => println(s) s - 1 }.run(10000).value
后来我用Cats Monad的tailRecM方法(所有Monad实例都实现了这个栈安全的递归方法),实现了另一个版本:
type WhileLoopState[A] = State[Unit, A] def whileLoopStateTailRec[S](cond: S => Boolean)(block: S => S)(initialState: S): WhileLoopState[S] = Monad[WhileLoopState] .tailRecM(initialState) { newState => State { _ => if (cond(newState)) { val nextState = block(newState) ((), Left(nextState)) } else { ((), Right(newState)) } } }
这个版本可以正常运行:
whileLoopStateTailRec[Int](s => s > 0) { s => println(s) s - 1 } (10000).run().value
但我觉得这个whileLoopStateTailRec的实现太复杂了,怀疑写法有问题,因此有三个问题:
- 是否可以简化该实现?
- 是否可以使用
State[S, Unit]替代State[Unit, A],让状态保存在正确的位置? - 不使用
tailRecM的情况下,能否基于State Monad实现栈安全的递归函数?
问题解答
1. 可以大幅简化实现
你的tailRecM版本确实绕了弯路——完全没必要把State的状态类型设为Unit,直接针对State[S, Unit]使用tailRecM即可。简化后的实现如下:
import cats.Monad import cats.data.State import cats.implicits._ def whileLoopState[S](cond: S => Boolean)(block: S => S): State[S, Unit] = Monad[State[S, *]].tailRecM(()) { _ => State.get[S].flatMap { s => if (cond(s)) State.modify(block).as(Left(())) else Right(()).pure[State[S, *]] } }
核心逻辑是:借助tailRecM的循环机制,每次从State中取出当前状态,判断条件——满足则执行block修改状态,返回Left(())继续循环;不满足则返回Right(())终止循环,最终返回Unit类型的结果。
2. 完全可以用State[S, Unit]替代
你的需求本质就是维护类型为S的状态,最终返回Unit,所以State[S, Unit]才是最贴合需求的类型。上面的简化实现已经直接使用该类型,状态完全保存在正确位置,调用方式和你最初的版本一致:
whileLoopState[Int](_ > 0) { s => println(s) s - 1 }.run(10000).value
这样既符合State Monad的常规用法,又保持了API的一致性。
3. 不使用tailRecM也能实现栈安全
有两种常见方式:
方式一:在State的run函数中直接写尾递归
State本质上是S => (S, A)的函数,直接在这个函数内部实现尾递归循环,天然栈安全,代码也最简洁:
import cats.data.State def whileLoopState[S](cond: S => Boolean)(block: S => S): State[S, Unit] = State { s => @scala.annotation.tailrec def loop(current: S): (S, Unit) = if (cond(current)) loop(block(current)) else (current, ()) loop(s) }
用@scala.annotation.tailrec注解确保编译器识别尾递归,避免栈溢出。
方式二:使用Trampoline包装递归
如果需要更通用的栈安全递归(比如复杂的monadic递归),可以用Cats的Trampoline数据类型包装递归步骤:
import cats.data.State import cats.free.Trampoline def whileLoopState[S](cond: S => Boolean)(block: S => S): State[S, Unit] = State { s => def loop(current: S): Trampoline[(S, Unit)] = if (cond(current)) Trampoline.defer(loop(block(current))) else Trampoline.done((current, ())) loop(s).run }
Trampoline.defer会将递归调用惰性化,避免栈帧累积,最终通过.run执行整个蹦床流程。
内容的提问来源于stack exchange,提问作者jcz

