You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用Cats的State Monad实现栈安全的递归while循环?

使用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的实现太复杂了,怀疑写法有问题,因此有三个问题:

  1. 是否可以简化该实现?
  2. 是否可以使用State[S, Unit]替代State[Unit, A],让状态保存在正确的位置?
  3. 不使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.13 16:10:30