Scala缓存线程安全问题复现与疑问:为何单线程写入会有问题?
同事指出下面这个Scala Cache对象存在线程安全问题:
import scala.collection.mutable object Cache { private val data = mutable.Map[Int, String]() def get(key: Int): Option[String] = data.get(key) def set(key: Int, value: String): Unit = data(key) = value }
我写了测试代码复现问题:
def expensiveMethod(): String = { Thread.sleep(1) Thread.currentThread.getName } for (i <- 1 to 3) { new Thread(() => { val cachedValue = Cache.get(1) if (cachedValue.nonEmpty) { println(s"${Thread.currentThread.getName}: cached value = ${cachedValue}") } else { val newValue = expensiveMethod() Cache.set(1, newValue) val finalValue = Cache.get(1) println(s"${Thread.currentThread.getName}: final value = ${finalValue}") } }).start() }
输出结果:
// Thread-2: final value = Some(Thread-2) // Thread-3: final value = Some(Thread-2) // Thread-1: final value = Some(Thread-2)
你的困惑点拆解与解答
1. 为什么“多个线程竞争写入仅一个成功”是问题?
你可能误解了缓存的核心作用:缓存的目标是避免重复执行昂贵计算,而不是纠结最终写入的值是谁。从输出能看到,三个线程都走到了else分支,也就是都执行了expensiveMethod()——这才是最严重的问题!本来应该只有第一个线程执行耗时计算,后面的线程直接读缓存结果,现在三个线程都做了重复的无用功,完全违背了缓存的初衷。
2. 核心问题是不是“写入后读取前被其他线程修改”?
不是,那只是竞态条件的一个表现而已。核心问题是检查-计算-设置(Check-Then-Act)的操作不是原子的:
- 线程A执行
get(1),发现值为空 - 线程A开始执行
expensiveMethod(),这时候线程B、C也同时执行了get(1),同样发现值为空,也启动了昂贵计算 - 等线程A计算完写入值,线程B、C的计算也快完成了,它们会覆盖掉之前的值,而整个过程中重复计算的问题已经发生了。
为什么你的尝试没完全解决问题?
1. 给set方法加synchronized锁
给set加锁只能保证写入操作本身是原子的,不会导致mutable.Map的内部结构损坏,但get→判断为空→执行计算的过程还是在锁外面,多个线程依然能同时进入else分支,重复执行昂贵计算,只是写入的时候不会出现Map的并发修改异常,所以你看到的是“异常现象减少”,但核心问题没解决。
2. 用concurrent.TrieMap替代mutable.Map
TrieMap是线程安全的Map,它的单个get、put操作都是原子的,但同样,组合的Check-Then-Act操作不是原子的。多个线程还是能同时通过get(1)的检查,进入计算流程,只是因为TrieMap的操作速度更快,线程之间的时间窗口更小,重复计算的概率降低了,但依然存在竞态条件,需要多次测试才会出现。
正确的修复方案
要解决这个问题,必须把检查-计算-设置的整个流程变成原子操作,下面是两种常见的修复方式:
方式一:用synchronized包裹整个缓存逻辑
修改Cache对象,把get和set的逻辑合并,用synchronized保证原子性:
import scala.collection.mutable object Cache { private val data = mutable.Map[Int, String]() def getOrCompute(key: Int, compute: => String): String = data.synchronized { data.getOrElseUpdate(key, compute) } }
测试代码改成:
def expensiveMethod(): String = { Thread.sleep(1) Thread.currentThread.getName } for (i <- 1 to 3) { new Thread(() => { val value = Cache.getOrCompute(1, expensiveMethod()) println(s"${Thread.currentThread.getName}: value = ${value}") }).start() }
这样整个“检查是否存在→不存在则计算→写入”的流程在同一个锁里,只有第一个线程会执行expensiveMethod(),后面的线程直接读取缓存值。
方式二:用TrieMap的原子操作
TrieMap提供了原子性的computeIfAbsent方法(Scala 2.13+支持),可以直接保证只有当key不存在时才执行计算逻辑:
import scala.collection.concurrent.TrieMap object Cache { private val data = TrieMap[Int, String]() def getOrCompute(key: Int, compute: => String): String = { data.computeIfAbsent(key, _ => compute) } }
多个线程同时调用时,只会有一个线程执行昂贵计算,其他线程会阻塞等待计算结果,完全避免了竞态条件。
内容的提问来源于stack exchange,提问作者Big McLargeHuge

